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)