(self, vis_feat, tgt)
| 274 | return res |
| 275 | |
| 276 | def forward_train(self, vis_feat, tgt): |
| 277 | sem_feat, sem_mask = self.semantic_branch(tgt) |
| 278 | pos_feat = self.positional_branch(sem_feat) |
| 279 | output = self.mdcdp( |
| 280 | sem_feat, |
| 281 | vis_feat, |
| 282 | pos_feat, |
| 283 | tgt_mask=sem_mask, |
| 284 | memory_mask=None, |
| 285 | ) |
| 286 | |
| 287 | logit = self.tgt_word_prj(output) |
| 288 | return logit |
| 289 | |
| 290 | def forward_test(self, vis_feat): |
| 291 | bs = vis_feat.size(0) |