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