MCPcopy Create free account
hub / github.com/apache/singa / Transformer

Class Transformer

examples/trans/model.py:29–102  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

27
28
29class 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)

Callers 1

runFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected