MCPcopy Create free account
hub / github.com/Topdu/OpenOCR / forward_train

Method forward_train

openrec/modeling/decoders/smtr_decoder.py:615–672  ·  view source on GitHub ↗
(self, x, targets=None)

Source from the content-addressed store, hash-verified

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]

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected