Args: enc_inputs: [batch_size, src_len] dec_inputs: [batch_size, tgt_len]
(self, enc_inputs, dec_inputs)
| 59 | self.soft_cross_entropy = layer.SoftMaxCrossEntropy() |
| 60 | |
| 61 | def forward(self, enc_inputs, dec_inputs): |
| 62 | """ |
| 63 | Args: |
| 64 | enc_inputs: [batch_size, src_len] |
| 65 | dec_inputs: [batch_size, tgt_len] |
| 66 | |
| 67 | """ |
| 68 | # enc_outputs: [batch_size, src_len, d_model], |
| 69 | # enc_self_attns: [n_layers, batch_size, n_heads, src_len, src_len] |
| 70 | enc_outputs, enc_self_attns = self.encoder(enc_inputs) |
| 71 | |
| 72 | # dec_outputs: [batch_size, tgt_len, d_model] |
| 73 | # dec_self_attns: [n_layers, batch_size, n_heads, tgt_len, tgt_len] |
| 74 | # dec_enc_attn: [n_layers, batch_size, tgt_len, src_len] |
| 75 | dec_outputs, dec_self_attns, dec_enc_attns = self.decoder(dec_inputs, enc_inputs, enc_outputs) |
| 76 | |
| 77 | # dec_logits: [batch_size, tgt_len, tgt_vocab_size] |
| 78 | dec_logits = self.linear3d(dec_outputs) |
| 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) |