File size: 5,058 Bytes
12496fc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 | """Small reference decoder: GQA, RoPE, RMSNorm, SwiGLU and causal SDPA.
Dense correctness baseline. No unimplemented MoE/distributed claims.
"""
from dataclasses import asdict, dataclass
import torch
from torch import nn
from torch.nn import functional as F
@dataclass(frozen=True)
class ModelConfig:
vocab_size: int = 259
hidden_size: int = 128
layers: int = 4
heads: int = 4
kv_heads: int = 2
intermediate_size: int = 384
max_context: int = 256
rope_theta: float = 10000.0
def __post_init__(self):
for key, value in asdict(self).items():
if value <= 0:
raise ValueError(f"{key} must be positive")
if self.hidden_size % self.heads or self.heads % self.kv_heads:
raise ValueError("Hidden/head and query/KV counts must divide evenly")
if (self.hidden_size // self.heads) % 2:
raise ValueError("RoPE requires even head dimension")
class RMSNorm(nn.Module):
def __init__(self, dim):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
y = x.float() * torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + 1e-6)
return y.to(x.dtype) * self.weight
def rope(x, theta):
length, dim = x.shape[-2:]
freq = 1.0 / (theta ** (torch.arange(0, dim, 2, device=x.device).float() / dim))
angles = torch.outer(torch.arange(length, device=x.device), freq)
cos, sin = angles.cos().to(x.dtype), angles.sin().to(x.dtype)
a, b = x[..., 0::2], x[..., 1::2]
return torch.stack((a * cos - b * sin, a * sin + b * cos), -1).flatten(-2)
class Attention(nn.Module):
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
d = cfg.hidden_size // cfg.heads
self.q = nn.Linear(cfg.hidden_size, cfg.heads * d, bias=False)
self.k = nn.Linear(cfg.hidden_size, cfg.kv_heads * d, bias=False)
self.v = nn.Linear(cfg.hidden_size, cfg.kv_heads * d, bias=False)
self.o = nn.Linear(cfg.hidden_size, cfg.hidden_size, bias=False)
def forward(self, x):
b, t, _ = x.shape
c = self.cfg
d = c.hidden_size // c.heads
q = rope(self.q(x).view(b, t, c.heads, d).transpose(1, 2), c.rope_theta)
k = rope(self.k(x).view(b, t, c.kv_heads, d).transpose(1, 2), c.rope_theta)
v = self.v(x).view(b, t, c.kv_heads, d).transpose(1, 2)
k = k.repeat_interleave(c.heads // c.kv_heads, dim=1)
v = v.repeat_interleave(c.heads // c.kv_heads, dim=1)
y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
return self.o(y.transpose(1, 2).contiguous().view(b, t, c.hidden_size))
class Block(nn.Module):
def __init__(self, c):
super().__init__()
self.norm1, self.norm2 = RMSNorm(c.hidden_size), RMSNorm(c.hidden_size)
self.attn = Attention(c)
self.gate = nn.Linear(c.hidden_size, c.intermediate_size, bias=False)
self.up = nn.Linear(c.hidden_size, c.intermediate_size, bias=False)
self.down = nn.Linear(c.intermediate_size, c.hidden_size, bias=False)
def forward(self, x):
x = x + self.attn(self.norm1(x))
h = self.norm2(x)
return x + self.down(F.silu(self.gate(h)) * self.up(h))
class NexoraLM(nn.Module):
def __init__(self, config: ModelConfig):
super().__init__()
self.config = config
self.embedding = nn.Embedding(config.vocab_size, config.hidden_size)
self.blocks = nn.ModuleList(Block(config) for _ in range(config.layers))
self.norm = RMSNorm(config.hidden_size)
self.apply(self._init)
@staticmethod
def _init(module):
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, std=0.02)
def forward(self, ids, labels=None):
if ids.ndim != 2 or not 0 < ids.shape[1] <= self.config.max_context:
raise ValueError("Expected nonempty batch x sequence within configured context")
x = self.embedding(ids)
for block in self.blocks:
x = block(x)
logits = F.linear(self.norm(x), self.embedding.weight)
loss = None if labels is None else F.cross_entropy(
logits.reshape(-1, self.config.vocab_size), labels.reshape(-1), ignore_index=-100)
return logits, loss
@torch.no_grad()
def generate(self, ids, max_new_tokens=64, temperature=0.0, eos_id=258):
if max_new_tokens < 0 or temperature < 0:
raise ValueError("Invalid generation limits")
self.eval()
for _ in range(max_new_tokens):
logits, _ = self(ids[:, -self.config.max_context:])
scores = logits[:, -1]
nxt = scores.argmax(-1, keepdim=True) if temperature == 0 else torch.multinomial(
(scores / temperature).softmax(-1), 1)
ids = torch.cat((ids, nxt), 1)
if (nxt == eos_id).all():
break
return ids
def parameter_count(self):
return sum(p.numel() for p in self.parameters())
|