| 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 | |