import math from typing import Optional, Tuple import torch import torch.nn as nn import torch.nn.functional as F VOCAB_SIZE = 32_768 MAX_CONTEXT = 2_048 TRAIN_SEQ_LEN = 512 HIDDEN = 896 LAYERS = 20 HEADS = 14 KV_HEADS = 7 INTERMEDIATE = 2_816 ROPE_THETA = 10_000.0 RMS_EPS = 1e-6 def masked_cross_entropy(logits, labels): valid = labels.ne(-100) if int(valid.sum().item()) <= 0: raise RuntimeError('masked_cross_entropy received zero supervised targets') return F.cross_entropy(logits.float().reshape(-1, VOCAB_SIZE), labels.reshape(-1), ignore_index=-100) class RMSNorm(nn.Module): def __init__(self, d=HIDDEN, eps=RMS_EPS): super().__init__() self.weight = nn.Parameter(torch.ones(d)) self.eps = eps def forward(self, x): return x * torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + self.eps).to(x.dtype) * self.weight def rotate_half(x): half = x.shape[-1] // 2 return torch.cat((-x[..., half:], x[..., :half]), dim=-1) class Rotary(nn.Module): def __init__(self, head_dim, max_seq=MAX_CONTEXT, theta=ROPE_THETA): super().__init__() inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2).float() / head_dim)) t = torch.arange(max_seq, dtype=torch.float32) freqs = torch.outer(t, inv_freq) emb = torch.cat((freqs, freqs), dim=-1) self.register_buffer('cos', emb.cos()[None, :, :], persistent=False) self.register_buffer('sin', emb.sin()[None, :, :], persistent=False) def forward(self, q, k): n = q.shape[-2] cos = self.cos[:, :n].to(q.device, q.dtype) sin = self.sin[:, :n].to(q.device, q.dtype) return q * cos + rotate_half(q) * sin, k * cos + rotate_half(k) * sin class M31Attention(nn.Module): def __init__(self): super().__init__() assert HIDDEN % HEADS == 0 assert HEADS % KV_HEADS == 0 self.head_dim = HIDDEN // HEADS self.q = nn.Linear(HIDDEN, HIDDEN, bias=False) self.k = nn.Linear(HIDDEN, KV_HEADS * self.head_dim, bias=False) self.v = nn.Linear(HIDDEN, KV_HEADS * self.head_dim, bias=False) self.o = nn.Linear(HIDDEN, HIDDEN, bias=False) self.rope = Rotary(self.head_dim, MAX_CONTEXT, ROPE_THETA) def forward(self, x): b, t, _ = x.shape q = self.q(x).view(b, t, HEADS, self.head_dim).transpose(1, 2) k = self.k(x).view(b, t, KV_HEADS, self.head_dim).transpose(1, 2) v = self.v(x).view(b, t, KV_HEADS, self.head_dim).transpose(1, 2) repeat = HEADS // KV_HEADS k = k.repeat_interleave(repeat, dim=1) v = v.repeat_interleave(repeat, dim=1) q, k = self.rope(q, k) y = F.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True) return self.o(y.transpose(1, 2).contiguous().view(b, t, HIDDEN)) class M31Block(nn.Module): def __init__(self): super().__init__() self.n1 = RMSNorm(HIDDEN, RMS_EPS) self.attn = M31Attention() self.n2 = RMSNorm(HIDDEN, RMS_EPS) self.gate = nn.Linear(HIDDEN, INTERMEDIATE, bias=False) self.up = nn.Linear(HIDDEN, INTERMEDIATE, bias=False) self.down = nn.Linear(INTERMEDIATE, HIDDEN, bias=False) def forward(self, x): x = x + self.attn(self.n1(x)) h = self.n2(x) h = F.silu(self.gate(h)) * self.up(h) return x + self.down(h) class M31Model(nn.Module): def __init__(self): super().__init__() self.embed = nn.Embedding(VOCAB_SIZE, HIDDEN) self.blocks = nn.ModuleList([M31Block() for _ in range(LAYERS)]) self.norm = RMSNorm(HIDDEN, RMS_EPS) self.apply(self._init_weights) nn.init.normal_(self.embed.weight, mean=0.0, std=0.02) self.num_parameters = sum(p.numel() for p in self.parameters()) if self.num_parameters >= 250_000_000: raise RuntimeError('Hard parameter ceiling violated.') @staticmethod def _init_weights(m): if isinstance(m, nn.Linear): nn.init.normal_(m.weight, mean=0.0, std=0.02 / math.sqrt(2 * LAYERS)) elif isinstance(m, nn.Embedding): pass def forward(self, input_ids, labels: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: x = self.embed(input_ids) for block in self.blocks: x = block(x) x = self.norm(x) logits = F.linear(x, self.embed.weight) loss = masked_cross_entropy(logits, labels) if labels is not None else None return logits, loss