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

Method forward_train_all

openrec/modeling/decoders/mdiff_decoder.py:150–203  ·  view source on GitHub ↗
(self, memory, data=None)

Source from the content-addressed store, hash-verified

148 return loss
149
150 def forward_train_all(self, memory, data=None):
151
152 labels, reflect_ids_all, noisy_batch_all, masked_indices_all, p_mask_all, length = data
153 bs, L = labels.shape
154 tgts = self.embedding(noisy_batch_all.flatten(0, 1))
155 tgts = self.positional_encoding(tgts) + self.pos_embed
156 tgts = tgts.reshape(bs, self.sample_k, L, -1)
157
158 for decoder_layer in self.decoder:
159 tgts = decoder_layer(tgts,
160 memory,
161 self_mask=None,
162 sample_k=self.sample_k)
163 logits_all = self.tgt_word_prj(tgts) # bs, sample_k, L, c_num
164
165 reflect_tgts = self.embedding(reflect_ids_all.flatten(0, 1))
166 reflect_tgts = self.positional_encoding(reflect_tgts) + self.pos_embed
167 reflect_tgts = reflect_tgts.reshape(bs, self.sample_k, L, -1)
168
169 for decoder_layer in self.decoder:
170 reflect_tgts = decoder_layer(reflect_tgts,
171 memory,
172 self_mask=None,
173 sample_k=self.sample_k)
174 reflect_logits_all = self.tgt_word_prj(reflect_tgts)
175
176 loss = []
177 for i in range(self.sample_k):
178 p_mask = p_mask_all[:, i]
179 masked_indices = masked_indices_all[:, i]
180 logits = logits_all[:, i]
181
182 p_mask = p_mask[:, None].repeat(1, labels.shape[1])
183 noisy_data_length = length + 1
184 noisy_data_length = noisy_data_length[:, None].repeat(
185 1, labels.shape[1])
186 token_loss = F.cross_entropy(
187 logits[masked_indices],
188 labels[masked_indices],
189 reduction='none',
190 ignore_index=self.ignore_index) / p_mask[masked_indices]
191 denoise_loss_i = torch.sum(
192 token_loss /
193 noisy_data_length[masked_indices]) / labels.shape[0]
194
195 reflect_logits = reflect_logits_all[:, i]
196 reflect_loss_i = F.cross_entropy(reflect_logits.flatten(0, 1),
197 labels.flatten(0, 1),
198 reduction='mean',
199 ignore_index=self.ignore_index)
200 loss_i = self.rec_loss_weight * denoise_loss_i + self.reflect_loss_weight * reflect_loss_i
201 loss.append(loss_i)
202
203 return sum(loss) / len(loss)
204
205 def forward(self, src, data=None):
206 """Take in and process masked source/target sequences.

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected