Forward function for `Transformer`. Args: x (Tensor): Input query with shape [bs, c, h, w] where c = embed_dims. mask (Tensor): The key_padding_mask used for encoder and decoder, with shape [bs, h, w]. query_embed (Tensor):
(self, x, mask, query_embed, pos_embed)
| 306 | self._is_init = True |
| 307 | |
| 308 | def forward(self, x, mask, query_embed, pos_embed): |
| 309 | """Forward function for `Transformer`. |
| 310 | |
| 311 | Args: |
| 312 | x (Tensor): Input query with shape [bs, c, h, w] where |
| 313 | c = embed_dims. |
| 314 | mask (Tensor): The key_padding_mask used for encoder and decoder, |
| 315 | with shape [bs, h, w]. |
| 316 | query_embed (Tensor): The query embedding for decoder, with shape |
| 317 | [num_query, c]. |
| 318 | pos_embed (Tensor): The positional encoding for encoder and |
| 319 | decoder, with the same shape as `x`. |
| 320 | |
| 321 | Returns: |
| 322 | tuple[Tensor]: results of decoder containing the following tensor. |
| 323 | |
| 324 | - out_dec: Output from decoder. If return_intermediate_dec \ |
| 325 | is True output has shape [num_dec_layers, bs, |
| 326 | num_query, embed_dims], else has shape [1, bs, \ |
| 327 | num_query, embed_dims]. |
| 328 | - memory: Output results from encoder, with shape \ |
| 329 | [bs, embed_dims, h, w]. |
| 330 | """ |
| 331 | bs, c, h, w = x.shape |
| 332 | # use `view` instead of `flatten` for dynamically exporting to ONNX |
| 333 | x = x.view(bs, c, -1).permute(2, 0, 1) # [bs, c, h, w] -> [h*w, bs, c] |
| 334 | pos_embed = pos_embed.view(bs, c, -1).permute(2, 0, 1) |
| 335 | query_embed = query_embed.unsqueeze(1).repeat( |
| 336 | 1, bs, 1) # [num_query, dim] -> [num_query, bs, dim] |
| 337 | mask = mask.view(bs, -1) # [bs, h, w] -> [bs, h*w] |
| 338 | memory = self.encoder(query=x, |
| 339 | key=None, |
| 340 | value=None, |
| 341 | query_pos=pos_embed, |
| 342 | query_key_padding_mask=mask) |
| 343 | target = torch.zeros_like(query_embed) |
| 344 | # out_dec: [num_layers, num_query, bs, dim] |
| 345 | out_dec = self.decoder(query=target, |
| 346 | key=memory, |
| 347 | value=memory, |
| 348 | key_pos=pos_embed, |
| 349 | query_pos=query_embed, |
| 350 | key_padding_mask=mask) |
| 351 | out_dec = out_dec.transpose(1, 2) |
| 352 | memory = memory.permute(1, 2, 0).reshape(bs, c, h, w) |
| 353 | return out_dec, memory |
| 354 | |
| 355 | |
| 356 | @TRANSFORMER.register_module() |