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