| import math, torch |
| import torch.nn as nn |
|
|
|
|
| class PositionalEncoding(nn.Module): |
| def __init__(self, d_model, max_len=256): |
| super().__init__() |
| pe = torch.zeros(max_len, d_model) |
| pos = torch.arange(max_len).unsqueeze(1).float() |
| div = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) |
| pe[:, 0::2] = torch.sin(pos * div) |
| pe[:, 1::2] = torch.cos(pos * div) |
| self.register_buffer("pe", pe.unsqueeze(0)) |
|
|
| def forward(self, x): |
| return x + self.pe[:, :x.size(1)] |
|
|
|
|
| class MTModel(nn.Module): |
| def __init__(self, vocab, d_model=512, nhead=8, layers=6, dim_ff=2048, |
| dropout=0.1, max_len=256): |
| super().__init__() |
| self.d_model = d_model |
| self.embed = nn.Embedding(vocab, d_model, padding_idx=0) |
| self.pos = PositionalEncoding(d_model, max_len) |
| self.transformer = nn.Transformer( |
| d_model=d_model, nhead=nhead, |
| num_encoder_layers=layers, num_decoder_layers=layers, |
| dim_feedforward=dim_ff, dropout=dropout, |
| batch_first=True, norm_first=True) |
| self.out = nn.Linear(d_model, vocab) |
| self.out.weight = self.embed.weight |
| nn.init.normal_(self.embed.weight, std=d_model ** -0.5) |
| nn.init.zeros_(self.out.bias) |
| with torch.no_grad(): |
| self.embed.weight[0].fill_(0) |
|
|
| def forward(self, src, tgt_in): |
| sp, tp = src == 0, tgt_in == 0 |
| causal = nn.Transformer.generate_square_subsequent_mask(tgt_in.size(1), device=src.device) |
| s = self.pos(self.embed(src) * math.sqrt(self.d_model)) |
| t = self.pos(self.embed(tgt_in) * math.sqrt(self.d_model)) |
| h = self.transformer(s, t, tgt_mask=causal, |
| src_key_padding_mask=sp, tgt_key_padding_mask=tp, |
| memory_key_padding_mask=sp) |
| return self.out(h) |
|
|