| import math, torch |
| import torch.nn as nn |
|
|
|
|
| class PositionalEncoding(nn.Module): |
| def __init__(self, d_model, max_len=512): |
| 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 TranslitModel(nn.Module): |
| def __init__(self, vocab_size, d_model=384, nhead=6, |
| num_layers=4, dim_ff=1536, dropout=0.1, max_len=512): |
| super().__init__() |
| self.d_model = d_model |
|
|
| self.embed = nn.Embedding(vocab_size, 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=num_layers, |
| num_decoder_layers=num_layers, |
| dim_feedforward=dim_ff, |
| dropout=dropout, |
| batch_first=True, |
| norm_first=True, |
| ) |
|
|
| self.out = nn.Linear(d_model, vocab_size) |
| self.out.weight = self.embed.weight |
|
|
| |
| nn.init.normal_(self.embed.weight, mean=0.0, 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): |
| src_pad = src == 0 |
| tgt_pad = 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=src_pad, |
| tgt_key_padding_mask=tgt_pad, |
| memory_key_padding_mask=src_pad, |
| ) |
| return self.out(h) |