| 27 | |
| 28 | |
| 29 | class Transformer(model.Model): |
| 30 | def __init__(self, src_n_token, tgt_n_token, d_model=512, n_head=8, dim_feedforward=2048, n_layers=6): |
| 31 | """ |
| 32 | Transformer model |
| 33 | Args: |
| 34 | src_n_token: the size of source vocab |
| 35 | tgt_n_token: the size of target vocab |
| 36 | d_model: the number of expected features in the encoder/decoder inputs (default=512) |
| 37 | n_head: the number of heads in the multi head attention models (default=8) |
| 38 | dim_feedforward: the dimension of the feedforward network model (default=2048) |
| 39 | n_layers: the number of sub-en(de)coder-layers in the en(de)coder (default=6) |
| 40 | """ |
| 41 | super(Transformer, self).__init__() |
| 42 | |
| 43 | self.opt = None |
| 44 | self.src_n_token = src_n_token |
| 45 | self.tgt_n_token = tgt_n_token |
| 46 | self.d_model = d_model |
| 47 | self.n_head = n_head |
| 48 | self.dim_feedforward = dim_feedforward |
| 49 | self.n_layers = n_layers |
| 50 | |
| 51 | # encoder / decoder / linear |
| 52 | self.encoder = TransformerEncoder(src_n_token=src_n_token, d_model=d_model, n_head=n_head, |
| 53 | dim_feedforward=dim_feedforward, n_layers=n_layers) |
| 54 | self.decoder = TransformerDecoder(tgt_n_token=tgt_n_token, d_model=d_model, n_head=n_head, |
| 55 | dim_feedforward=dim_feedforward, n_layers=n_layers) |
| 56 | |
| 57 | self.linear3d = Linear3D(in_features=d_model, out_features=tgt_n_token, bias=False) |
| 58 | |
| 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) |
| 83 | shape = out.shape[-1] |
| 84 | out = autograd.reshape(out, [-1, shape]) |
| 85 | |
| 86 | out_np = tensor.to_numpy(out) |