| """ |
| yk_diffusion: a from-scratch hybrid language model. |
| |
| A single Transformer runs in two modes, selected by a learned mode embedding: |
| - AR mode (mode=0): causal attention -> standard next-token autoregressive LM ("normal") |
| - DIFF mode (mode=1): bidirectional attention -> masked discrete diffusion denoising |
| |
| A time embedding conditions the diffusion mask ratio (MDLM / LLaDA-style absorbing-state training). |
| The same weights serve both behaviours; a <MODE> signal picks which one at inference time. |
| """ |
|
|
| import math |
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from torch.utils.checkpoint import checkpoint |
|
|
|
|
| |
| |
| |
|
|
| class RMSNorm(nn.Module): |
| def __init__(self, d, eps=1e-6): |
| super().__init__() |
| self.eps = eps |
| self.weight = nn.Parameter(torch.ones(d)) |
|
|
| def forward(self, x): |
| var = x.pow(2).mean(-1, keepdim=True) |
| x = x * torch.rsqrt(var + self.eps) |
| return x * self.weight |
|
|
|
|
| class RotaryEmbedding(nn.Module): |
| def __init__(self, head_dim, max_len=8192, base=10000.0, |
| rope_scale=1.0): |
| super().__init__() |
| |
| |
| |
| |
| base = base * (rope_scale ** (head_dim / (head_dim - 2))) |
| inv = 1.0 / (base ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)) |
| self.register_buffer("inv_freq", inv) |
|
|
| def forward(self, seq_len, device): |
| t = torch.arange(seq_len, device=device, dtype=self.inv_freq.dtype) |
| freqs = torch.outer(t, self.inv_freq) |
| emb = torch.cat([freqs, freqs], dim=-1) |
| return emb.cos(), emb.sin() |
|
|
|
|
| def rotate_half(x): |
| x1, x2 = x.chunk(2, dim=-1) |
| return torch.cat((-x2, x1), dim=-1) |
|
|
|
|
| def apply_rope(x, cos, sin): |
| |
| return x * cos + rotate_half(x) * sin |
|
|
|
|
| class Attention(nn.Module): |
| def __init__(self, d, n_heads): |
| super().__init__() |
| assert d % n_heads == 0 |
| self.d = d |
| self.n_heads = n_heads |
| self.hd = d // n_heads |
| self.scale = self.hd ** -0.5 |
| self.qkv = nn.Linear(d, 3 * d, bias=False) |
| self.proj = nn.Linear(d, d, bias=False) |
|
|
| def forward(self, x, cos, sin, attn_mask=None): |
| B, L, _ = x.shape |
| qkv = self.qkv(x).reshape(B, L, 3, self.n_heads, self.hd) |
| qkv = qkv.permute(2, 0, 3, 1, 4) |
| q, k, v = qkv[0], qkv[1], qkv[2] |
| q = apply_rope(q, cos, sin) |
| k = apply_rope(k, cos, sin) |
| scores = (q @ k.transpose(-2, -1)) * self.scale |
| if attn_mask is not None: |
| scores = scores + attn_mask |
| attn = scores.softmax(dim=-1) |
| out = (attn @ v).transpose(1, 2).reshape(B, L, self.d) |
| return self.proj(out) |
|
|
|
|
| class MLP(nn.Module): |
| def __init__(self, d, d_ff): |
| super().__init__() |
| self.w1 = nn.Linear(d, d_ff, bias=False) |
| self.w3 = nn.Linear(d, d_ff, bias=False) |
| self.w2 = nn.Linear(d_ff, d, bias=False) |
|
|
| def forward(self, x): |
| return self.w2(F.silu(self.w1(x)) * self.w3(x)) |
|
|
|
|
| class Block(nn.Module): |
| def __init__(self, d, n_heads, d_ff): |
| super().__init__() |
| self.ln1 = RMSNorm(d) |
| self.attn = Attention(d, n_heads) |
| self.ln2 = RMSNorm(d) |
| self.mlp = MLP(d, d_ff) |
|
|
| def forward(self, x, cos, sin, attn_mask): |
| x = x + self.attn(self.ln1(x), cos, sin, attn_mask) |
| x = x + self.mlp(self.ln2(x)) |
| return x |
|
|
|
|
| |
| |
| |
|
|
| class YKDiff(nn.Module): |
| def __init__(self, cfg): |
| super().__init__() |
| self.cfg = cfg |
| d = cfg["d_model"] |
| self.vocab = cfg["vocab_size"] |
| self.max_len = cfg["max_len"] |
|
|
| self.tok_emb = nn.Embedding(self.vocab, d) |
| self.mode_emb = nn.Embedding(2, d) |
| self.time_emb = nn.Linear(1, d, bias=False) |
| self.rope = RotaryEmbedding(d // cfg["n_heads"], max_len=self.max_len, |
| rope_scale=cfg.get("rope_scale", 1.0)) |
|
|
| self.blocks = nn.ModuleList([ |
| Block(d, cfg["n_heads"], cfg["d_ff"]) for _ in range(cfg["n_layers"]) |
| ]) |
| self.norm = RMSNorm(d) |
| self.lm_head = nn.Linear(d, self.vocab, bias=False) |
| with torch.no_grad(): |
| self.lm_head.weight.copy_(self.tok_emb.weight) |
|
|
| self._causal = None |
|
|
| @property |
| def device(self): |
| return next(self.parameters()).device |
|
|
| def _causal_mask(self, L, device): |
| if self._causal is None or self._causal.shape[-1] < L: |
| m = torch.full((L, L), float("-inf"), device=device) |
| m = torch.triu(m, diagonal=1) |
| self._causal = m |
| return self._causal[:L, :L] |
|
|
| def forward(self, idx, mode, t=None, attn_mask=None): |
| """ |
| idx: [B, L] long token ids |
| mode: [B] long (0 AR, 1 DIFF) |
| t: [B] float|None diffusion mask ratio (None -> 0) |
| attn_mask: [L, L]|None explicit mask; if None, AR uses causal, DIFF uses none |
| """ |
| B, L = idx.shape |
| x = self.tok_emb(idx) |
| x = x + self.mode_emb(mode).unsqueeze(1) |
| if t is None: |
| t = torch.zeros(B, device=idx.device) |
| x = x + self.time_emb(t.unsqueeze(-1)).unsqueeze(1) |
|
|
| cos, sin = self.rope(L, idx.device) |
| cos = cos.unsqueeze(0).unsqueeze(0) |
| sin = sin.unsqueeze(0).unsqueeze(0) |
|
|
| if attn_mask is None: |
| |
| attn_mask = self._causal_mask(L, idx.device) if mode[0].item() == 0 else None |
| else: |
| attn_mask = attn_mask.to(idx.device) |
|
|
| |
| |
| training = self.training and torch.is_grad_enabled() |
| for blk in self.blocks: |
| if training: |
| x = checkpoint(blk, x, cos, sin, attn_mask, |
| use_reentrant=False) |
| else: |
| x = blk(x, cos, sin, attn_mask) |
| x = self.norm(x) |
| return self.lm_head(x) |
|
|