| 79 | return dec_logits, enc_self_attns, dec_self_attns, dec_enc_attns |
| 80 | |
| 81 | def train_one_batch(self, enc_inputs, dec_inputs, dec_outputs, pad): |
| 82 | out, _, _, _ = self.forward(enc_inputs, dec_inputs) |
| 83 | shape = out.shape[-1] |
| 84 | out = autograd.reshape(out, [-1, shape]) |
| 85 | |
| 86 | out_np = tensor.to_numpy(out) |
| 87 | preds_np = np.argmax(out_np, -1) |
| 88 | |
| 89 | dec_outputs_np = tensor.to_numpy(dec_outputs) |
| 90 | dec_outputs_np = dec_outputs_np.reshape(-1) |
| 91 | |
| 92 | y_label_mask = dec_outputs_np != pad |
| 93 | correct = preds_np == dec_outputs_np |
| 94 | acc = np.sum(y_label_mask * correct) / np.sum(y_label_mask) |
| 95 | dec_outputs = tensor.from_numpy(dec_outputs_np) |
| 96 | |
| 97 | loss = self.soft_cross_entropy(out, dec_outputs) |
| 98 | self.opt(loss) |
| 99 | return out, loss, acc |
| 100 | |
| 101 | def set_optimizer(self, opt): |
| 102 | self.opt = opt |