coderofpears's picture
Upload model.py with huggingface_hub
4f9392a verified
Raw
History Blame Contribute Delete
6.89 kB
"""
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
# ----------------------------------------------------------------------------
# Building blocks
# ----------------------------------------------------------------------------
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__()
# NTK-aware scaling: stretch the base so the trained (short) context
# extends to longer inference windows without retraining. With
# rope_scale = context_ratio (e.g. 128k/8k = 16) the model trained at
# 8k still positions tokens correctly at 128k.
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) # [L, head_dim/2]
emb = torch.cat([freqs, freqs], dim=-1) # [L, head_dim]
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):
# x: [B, h, L, hd]; cos/sin: [1, 1, L, hd]
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) # [3, B, h, L, hd]
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
# ----------------------------------------------------------------------------
# The model
# ----------------------------------------------------------------------------
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) # 0 = AR, 1 = DIFF
self.time_emb = nn.Linear(1, d, bias=False) # diffusion mask ratio conditioning
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) # weight tying
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) # [1,1,L,hd]
sin = sin.unsqueeze(0).unsqueeze(0)
if attn_mask is None:
# AR default causal; DIFF default bidirectional (None)
attn_mask = self._causal_mask(L, idx.device) if mode[0].item() == 0 else None
else:
attn_mask = attn_mask.to(idx.device)
# gradient checkpointing: trade ~20% compute for ~4x less
# activation memory, so a big batch fits on 16 GB.
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) # [B, L, vocab]