MCPcopy Create free account
hub / github.com/MotrixLab/AiOS / forward

Method forward

detrsmpl/models/utils/transformer.py:308–353  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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()

Callers 2

forwardMethod · 0.45
forwardMethod · 0.45

Calls

no outgoing calls

Tested by

no test coverage detected