from __future__ import annotations import math import torch from torch import nn def sinusoidal_time_embedding(t: torch.Tensor, dimension: int) -> torch.Tensor: half = dimension // 2 frequencies = torch.exp( -math.log(10_000) * torch.arange(half, device=t.device, dtype=torch.float32) / max(half - 1, 1) ) angles = t.float().unsqueeze(1) * frequencies.unsqueeze(0) embedding = torch.cat([torch.sin(angles), torch.cos(angles)], dim=1) if dimension % 2: embedding = torch.nn.functional.pad(embedding, (0, 1)) return embedding class MultiMDMTransformer(nn.Module): def __init__( self, vocab_size: int, clean_vocab_size: int, seq_len: int, num_masks: int, d_model: int = 128, nhead: int = 4, num_layers: int = 2, dim_feedforward: int = 256, dropout: float = 0.1, ): super().__init__() self.vocab_size = vocab_size self.clean_vocab_size = clean_vocab_size self.seq_len = seq_len self.num_masks = num_masks self.d_model = d_model self.token_embedding = nn.Embedding(vocab_size, d_model) self.position_embedding = nn.Embedding(seq_len, d_model) self.time_mlp = nn.Sequential( nn.Linear(d_model, d_model * 2), nn.SiLU(), nn.Linear(d_model * 2, d_model), ) layer = nn.TransformerEncoderLayer( d_model=d_model, nhead=nhead, dim_feedforward=dim_feedforward, dropout=dropout, activation="gelu", batch_first=True, norm_first=True, ) self.transformer = nn.TransformerEncoder(layer, num_layers=num_layers) self.norm = nn.LayerNorm(d_model) self.output_head = nn.Linear(d_model, clean_vocab_size) self.mask_class_head = nn.Linear(d_model, num_masks) def forward( self, input_ids: torch.Tensor, t: torch.Tensor, attention_mask: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: batch, length = input_ids.shape if length > self.seq_len: raise ValueError(f"input length {length} exceeds seq_len {self.seq_len}") positions = torch.arange(length, device=input_ids.device) hidden = self.token_embedding(input_ids) hidden = hidden + self.position_embedding(positions).unsqueeze(0) time_hidden = self.time_mlp( sinusoidal_time_embedding(t, self.d_model) ).unsqueeze(1) hidden = hidden + time_hidden padding_mask = None if attention_mask is None else ~attention_mask.bool() hidden = self.transformer(hidden, src_key_padding_mask=padding_mask) hidden = self.norm(hidden) return self.output_head(hidden), self.mask_class_head(hidden)