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

Method forward_train

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

Source from the content-addressed store, hash-verified

110 nn.init.zeros_(m.bias)
111
112 def forward_train(self, memory, data=None):
113 labels, reflect_ids, noisy_batch, masked_indices, p_mask, length = data
114 p_mask = p_mask[:, None].repeat(1, labels.shape[1])
115 noisy_data_length = length + 1
116 noisy_data_length = noisy_data_length[:,
117 None].repeat(1, labels.shape[1])
118
119 tgts = self.embedding(noisy_batch)
120 tgts = self.positional_encoding(tgts) + self.pos_embed
121
122 for decoder_layer in self.decoder:
123 tgts = decoder_layer(tgts, memory, self_mask=None)
124 logits = self.tgt_word_prj(tgts)
125 token_loss = F.cross_entropy(
126 logits[masked_indices],
127 labels[masked_indices],
128 reduction='none',
129 ignore_index=self.ignore_index) / p_mask[masked_indices]
130 loss = torch.sum(
131 token_loss / noisy_data_length[masked_indices]) / labels.shape[0]
132
133 if reflect_ids is not None:
134 reflect_tgts = self.embedding(reflect_ids)
135 reflect_tgts = self.positional_encoding(
136 reflect_tgts) + self.pos_embed
137 for decoder_layer in self.decoder:
138 reflect_tgts = decoder_layer(reflect_tgts,
139 memory,
140 self_mask=None)
141 reflect_logits = self.tgt_word_prj(reflect_tgts)
142 reflect_loss = F.cross_entropy(reflect_logits.flatten(0, 1),
143 labels.flatten(0, 1),
144 reduction='mean',
145 ignore_index=self.ignore_index)
146 loss = self.rec_loss_weight * loss + self.reflect_loss_weight * reflect_loss
147
148 return loss
149
150 def forward_train_all(self, memory, data=None):
151

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected