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

Method forward

detrsmpl/models/utils/transformer.py:158–184  ·  view source on GitHub ↗

Forward function for `TransformerDecoder`. Args: query (Tensor): Input query with shape `(num_query, bs, embed_dims)`. Returns: Tensor: Results with shape [1, num_query, bs, embed_dims] when return_intermediate is `False`, oth

(self, query, *args, **kwargs)

Source from the content-addressed store, hash-verified

156 self.post_norm = None
157
158 def forward(self, query, *args, **kwargs):
159 """Forward function for `TransformerDecoder`.
160
161 Args:
162 query (Tensor): Input query with shape
163 `(num_query, bs, embed_dims)`.
164
165 Returns:
166 Tensor: Results with shape [1, num_query, bs, embed_dims] when
167 return_intermediate is `False`, otherwise it has shape
168 [num_layers, num_query, bs, embed_dims].
169 """
170 if not self.return_intermediate:
171 x = super().forward(query, *args, **kwargs)
172 if self.post_norm:
173 x = self.post_norm(x)[None]
174 return x
175
176 intermediate = []
177 for layer in self.layers:
178 query = layer(query, *args, **kwargs)
179 if self.return_intermediate:
180 if self.post_norm is not None:
181 intermediate.append(self.post_norm(query))
182 else:
183 intermediate.append(query)
184 return torch.stack(intermediate)
185
186
187@TRANSFORMER_LAYER_SEQUENCE.register_module()

Callers

nothing calls this directly

Calls 1

forwardMethod · 0.45

Tested by

no test coverage detected