import math def generate_src_mask(sz, device): return torch.triu(torch.full((sz, sz), True, device=device), diagonal=1) # class PositionalEmbedding(nn.Module): def __init__(self, sequence_length, embed_dim): super().__init__() self.sqrt_embed_dim = math.sqrt(embed_dim) self.pos_embed = nn.Parameter(torch.empty((1, sequence_length, embed_dim))) nn.init.uniform_(self.pos_embed, -1.0, 1.0) # def forward(self, x): return x * self.sqrt_embed_dim + self.pos_embed[:,:x.size(1)] # # class WordPredictionTransformerModel(nn.Module): def __init__(self, sequence_length, num_embed, embed_dim, pad_idx, num_heads, num_layers, output_dim, dropout, norm_first, activation): super().__init__() self.vocab_size = num_embed self.sequence_length = sequence_length self.embed_dim = embed_dim self.sqrt_embed_dim = math.sqrt(embed_dim) self.embed = nn.Sequential( nn.Embedding(num_embed, embed_dim, pad_idx), PositionalEmbedding(sequence_length, embed_dim), nn.LayerNorm(embed_dim), nn.Dropout(p=0.1), ) encoder_layer = nn.TransformerEncoderLayer( d_model=embed_dim, nhead=num_heads, dropout=dropout, batch_first=True, norm_first=norm_first, activation=activation, ) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_layers) self.fc = nn.Sequential( nn.Linear(embed_dim, embed_dim * 4), nn.LayerNorm(embed_dim * 4), nn.LeakyReLU(), nn.Dropout(p=dropout), nn.Linear(embed_dim * 4, output_dim), ) # def forward(self, x): src_attention_mask = generate_src_mask(x.size(1), x.device) x = self.embed(x) x = self.encoder(x, is_causal=True, mask=src_attention_mask) x = self.fc(x) x = x.permute(0, 2, 1) return x # #