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

Method forward

openrec/modeling/decoders/bus_decoder.py:78–133  ·  view source on GitHub ↗

Args: tokens: (N, T, C) where T is length, N is batch size and C is classes number lengths: (N,)

(self, img_feat, data=None)

Source from the content-addressed store, hash-verified

76 return logits
77
78 def forward(self, img_feat, data=None):
79 """
80 Args:
81 tokens: (N, T, C) where T is length, N is batch size and C is classes number
82 lengths: (N,)
83 """
84 img_feat = img_feat + self.v_embeding
85 B, L, C = img_feat.shape
86
87 # --------------------------------------------------------------------------
88 # decoder procedure
89 T = self.max_length
90 zeros = img_feat.new_zeros((B, T, C))
91 zeros_len = img_feat.new_zeros(B)
92 query = self.pos_encoder(zeros)
93
94 # 1. vision decode
95 v_embed = torch.cat((img_feat, self.l_mask.repeat(B, T, 1)),
96 dim=1) # v
97 padding_mask = _get_mask(
98 self.max_length + zeros_len,
99 self.max_length) # 对tokens长度以外的padding # B, maxlen maxlen
100 v_mask = torch.zeros((1, 1, self.max_length, L),
101 device=img_feat.device).tile([B, 1, 1,
102 1]) # maxlen L
103 mask = torch.cat((v_mask, padding_mask), 3)
104 v_logits = self.forward_decoder(query, v_embed, mask=mask)
105
106 # 2. language decode
107 if self.training and self.pretraining:
108 tgt = torch.where(data[0] == self.ignore_index, 0, data[0])
109 tokens = F.one_hot(tgt, num_classes=self.out_channels)
110 tokens = tokens.float()
111 lengths = data[-1]
112 else:
113 tokens = torch.softmax(v_logits, dim=-1)
114 lengths = _get_length(v_logits)
115 tokens = tokens.detach()
116 token_embed = self.proj(tokens) # (N, T, E)
117 token_embed = self.token_encoder(token_embed) # (T, N, E)
118 token_embed = token_embed + self.l_embeding
119
120 padding_mask = _get_mask(lengths,
121 self.max_length) # 对tokens长度以外的padding
122 mask = torch.cat((v_mask, padding_mask), 3)
123 l_embed = torch.cat((self.v_mask.repeat(B, L, 1), token_embed), dim=1)
124 l_logits = self.forward_decoder(query, l_embed, mask=mask)
125
126 # 3. vision language decode
127 vl_embed = torch.cat((img_feat, token_embed), dim=1)
128 vl_logits = self.forward_decoder(query, vl_embed, mask=mask)
129
130 if self.training:
131 return {'align': [vl_logits], 'lang': l_logits, 'vision': v_logits}
132 else:
133 return F.softmax(vl_logits, -1)

Callers

nothing calls this directly

Calls 3

forward_decoderMethod · 0.95
_get_maskFunction · 0.85
_get_lengthFunction · 0.85

Tested by

no test coverage detected