MCPcopy Create free account
hub / github.com/pytorch/examples / TransformerModel

Class TransformerModel

word_language_model/model.py:107–144  ·  view source on GitHub ↗

Container module with an encoder, a recurrent or transformer module, and a decoder.

Source from the content-addressed store, hash-verified

105 return self.dropout(x)
106
107class 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)

Callers 1

main.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected