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 # weight tying # ---- proper init for tied embedding (fixes huge initial loss) ---- 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) # keep padding row at zero 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)