(
self,
tgt: torch.Tensor,
memory: torch.Tensor,
tgt_mask: Optional[Tensor] = None,
tgt_padding_mask: Optional[Tensor] = None,
tgt_query: Optional[Tensor] = None,
tgt_query_mask: Optional[Tensor] = None,
pos_query: torch.Tensor = None,
)
| 242 | return param_names |
| 243 | |
| 244 | def decode( |
| 245 | self, |
| 246 | tgt: torch.Tensor, |
| 247 | memory: torch.Tensor, |
| 248 | tgt_mask: Optional[Tensor] = None, |
| 249 | tgt_padding_mask: Optional[Tensor] = None, |
| 250 | tgt_query: Optional[Tensor] = None, |
| 251 | tgt_query_mask: Optional[Tensor] = None, |
| 252 | pos_query: torch.Tensor = None, |
| 253 | ): |
| 254 | N, L = tgt.shape |
| 255 | # <bos> stands for the null context. We only supply position information for characters after <bos>. |
| 256 | null_ctx = self.text_embed(tgt[:, :1]) |
| 257 | |
| 258 | if tgt_query is None: |
| 259 | tgt_query = pos_query[:, :L] |
| 260 | tgt_emb = pos_query[:, :L - 1] + self.text_embed(tgt[:, 1:]) |
| 261 | tgt_emb = self.dropout(torch.cat([null_ctx, tgt_emb], dim=1)) |
| 262 | |
| 263 | tgt_query = self.dropout(tgt_query) |
| 264 | return self.decoder(tgt_query, tgt_emb, memory, tgt_query_mask, |
| 265 | tgt_mask, tgt_padding_mask) |
| 266 | |
| 267 | def forward(self, x, data=None, pos_query=None): |
| 268 | if self.training: |
no outgoing calls
no test coverage detected