| 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) |
|
|