| 613 | return torch.concat(logits_all, 1) |
| 614 | |
| 615 | def forward_train(self, x, targets=None): |
| 616 | bs = x.shape[0] |
| 617 | |
| 618 | if not self.ds: |
| 619 | visual_f = x + self.vis_pos_embed |
| 620 | elif self.pos2d: |
| 621 | visual_f = x + self.vis_pos_embed[:, :, :x.shape[2], :x.shape[3]] |
| 622 | else: |
| 623 | visual_f = x |
| 624 | max_len_curr = targets[3].max() |
| 625 | subs = targets[1][:, :max_len_curr, :] # b, n, subs_l |
| 626 | mask_next = torch.where(subs == self.bos_next, float('-inf'), |
| 627 | 0) # b, n, subs_l |
| 628 | prompt_next_embed = self.prompt_next_embed.tile( |
| 629 | [bs, max_len_curr, 1, 1]) |
| 630 | prompt_char_next = torch.concat([ |
| 631 | prompt_next_embed[:, :, :1, :], |
| 632 | prompt_next_embed[:, :, 1:, :] + self.char_embed(subs) |
| 633 | ], 2) # b, n, subs_l, dim |
| 634 | next = self.next_token.tile([bs, max_len_curr, 1, 1]) |
| 635 | |
| 636 | max_len_curr_pre = targets[6].max() |
| 637 | subs = targets[4][:, :max_len_curr_pre, :] # b, n, subs_l |
| 638 | mask_pre = torch.where(subs == self.bos_pre, float('-inf'), |
| 639 | 0) # b, n, subs_l |
| 640 | prompt_pre_embed = self.prompt_pre_embed.tile( |
| 641 | [bs, max_len_curr_pre, 1, 1]) |
| 642 | prompt_char_pre = torch.concat([ |
| 643 | prompt_pre_embed[:, :, :1, :], |
| 644 | prompt_pre_embed[:, :, 1:, :] + self.char_embed(subs) |
| 645 | ], 2) # b, n, sub_l, dim |
| 646 | pre = self.pre_token.tile([bs, max_len_curr_pre, 1, 1]) # b, n, 1, dim |
| 647 | |
| 648 | prompt_char = torch.concat([prompt_char_next, prompt_char_pre], 1) |
| 649 | next_pre = torch.concat([next, pre], 1) |
| 650 | |
| 651 | mask_pad = torch.zeros([bs * (max_len_curr + max_len_curr_pre), 1], |
| 652 | dtype=torch.float32, |
| 653 | device=x.device) |
| 654 | mask = torch.concat([mask_next, mask_pre], 1).flatten(0, 1) |
| 655 | mask = torch.concat([mask_pad, mask], 1) |
| 656 | next_pre = next_pre.flatten(0, 1) |
| 657 | prompt_char = prompt_char.flatten(0, 1) |
| 658 | for layer in self.cmff_decoder: |
| 659 | next_pre = layer(next_pre, prompt_char, visual_f, |
| 660 | mask.unsqueeze(1)) |
| 661 | answer1_pred = self.ques1_head(self.norm_pred(next_pre)) |
| 662 | logits = answer1_pred[:, :max_len_curr] |
| 663 | |
| 664 | label = torch.concat( |
| 665 | [targets[2][:, :max_len_curr], targets[5][:, :max_len_curr_pre]], |
| 666 | 1) |
| 667 | loss1 = F.cross_entropy(answer1_pred.flatten(0, 1), |
| 668 | label.flatten(0, 1), |
| 669 | ignore_index=self.ignore_index, |
| 670 | reduction='mean') |
| 671 | loss = {'loss': loss1} |
| 672 | return [loss, logits] |