Container module with an encoder, a recurrent or transformer module, and a decoder.
| 105 | return self.dropout(x) |
| 106 | |
| 107 | class TransformerModel(nn.Transformer): |
| 108 | """Container module with an encoder, a recurrent or transformer module, and a decoder.""" |
| 109 | |
| 110 | def __init__(self, ntoken, ninp, nhead, nhid, nlayers, dropout=0.5): |
| 111 | super(TransformerModel, self).__init__(d_model=ninp, nhead=nhead, dim_feedforward=nhid, num_encoder_layers=nlayers) |
| 112 | self.model_type = 'Transformer' |
| 113 | self.src_mask = None |
| 114 | self.pos_encoder = PositionalEncoding(ninp, dropout) |
| 115 | |
| 116 | self.input_emb = nn.Embedding(ntoken, ninp) |
| 117 | self.ninp = ninp |
| 118 | self.decoder = nn.Linear(ninp, ntoken) |
| 119 | |
| 120 | self.init_weights() |
| 121 | |
| 122 | def _generate_square_subsequent_mask(self, sz): |
| 123 | return torch.log(torch.tril(torch.ones(sz,sz))) |
| 124 | |
| 125 | def init_weights(self): |
| 126 | initrange = 0.1 |
| 127 | nn.init.uniform_(self.input_emb.weight, -initrange, initrange) |
| 128 | nn.init.zeros_(self.decoder.bias) |
| 129 | nn.init.uniform_(self.decoder.weight, -initrange, initrange) |
| 130 | |
| 131 | def forward(self, src, has_mask=True): |
| 132 | if has_mask: |
| 133 | device = src.device |
| 134 | if self.src_mask is None or self.src_mask.size(0) != len(src): |
| 135 | mask = self._generate_square_subsequent_mask(len(src)).to(device) |
| 136 | self.src_mask = mask |
| 137 | else: |
| 138 | self.src_mask = None |
| 139 | |
| 140 | src = self.input_emb(src) * math.sqrt(self.ninp) |
| 141 | src = self.pos_encoder(src) |
| 142 | output = self.encoder(src, mask=self.src_mask) |
| 143 | output = self.decoder(output) |
| 144 | return F.log_softmax(output, dim=-1) |