外观
llm/model.py
配套源码,运行方法见同目录 README。返回实验总览。
python
"""Tiny decoder-only Transformer with explicit NumPy backward passes.
Pre-RMSNorm, RoPE, SwiGLU, causal multi-head attention, untied output head.
Float64 keeps finite-difference checks interpretable. This is CPU teaching
code, not an optimized training framework or an official assignment answer.
"""
from dataclasses import dataclass, asdict
import json
import numpy as np
@dataclass
class Config:
vocab_size: int
dim: int = 32
heads: int = 4
hidden: int = 64
layers: int = 1
context: int = 24
seed: int = 7
def __post_init__(self):
if min(self.vocab_size, self.dim, self.heads, self.hidden,
self.layers, self.context) < 1:
raise ValueError("model dimensions must be positive")
if self.dim % self.heads or (self.dim // self.heads) % 2:
raise ValueError("dim/heads must be an even integer for RoPE")
def softmax(x):
z = np.exp(x - x.max(axis=-1, keepdims=True))
return z / z.sum(axis=-1, keepdims=True)
def cross_entropy(logits, targets):
"""Mean token NLL and derivative with respect to logits."""
if logits.shape[:-1] != targets.shape:
raise ValueError("targets must match all non-vocabulary axes")
shifted = logits - logits.max(axis=-1, keepdims=True)
log_z = np.log(np.exp(shifted).sum(axis=-1))
selected = np.take_along_axis(shifted, targets[..., None], axis=-1)[..., 0]
loss = float((log_z - selected).mean())
grad = softmax(logits)
flat = grad.reshape(-1, grad.shape[-1])
flat[np.arange(targets.size), targets.reshape(-1)] -= 1.0
grad /= targets.size
return loss, grad
def rmsnorm(x, gain, eps=1e-5):
inv = 1.0 / np.sqrt(np.mean(x * x, axis=-1, keepdims=True) + eps)
return x * inv * gain, (x, gain, inv)
def rmsnorm_backward(dy, cache):
x, gain, inv = cache
dgain = np.sum(dy * x * inv, axis=tuple(range(dy.ndim - 1)))
u = dy * gain
dx = inv * u - x * inv**3 * np.mean(u * x, axis=-1, keepdims=True)
return dx, dgain
def rope(x, inverse=False):
"""x: (batch, heads, time, head_dim); rotation of adjacent pairs."""
t, d = x.shape[-2:]
angle = np.arange(t)[:, None] * 10000.0 ** (-np.arange(0, d, 2) / d)
if inverse:
angle = -angle
c, s = np.cos(angle), np.sin(angle)
even, odd = x[..., 0::2], x[..., 1::2]
out = np.empty_like(x)
out[..., 0::2] = even * c - odd * s
out[..., 1::2] = even * s + odd * c
return out
def causal_attention(q, k, v):
"""Dense reference attention; returns output and backward cache."""
scores = q @ k.swapaxes(-1, -2) / np.sqrt(q.shape[-1])
t = q.shape[-2]
scores = np.where(np.tril(np.ones((t, t), dtype=bool)), scores, -np.inf)
p = softmax(scores)
return p @ v, (q, k, v, p)
def attention_backward(dout, cache):
q, k, v, p = cache
dp = dout @ v.swapaxes(-1, -2)
ds = p * (dp - (dp * p).sum(axis=-1, keepdims=True))
scale = np.sqrt(q.shape[-1])
dq = ds @ k / scale
dk = ds.swapaxes(-1, -2) @ q / scale
dv = p.swapaxes(-1, -2) @ dout
return dq, dk, dv
class TinyTransformer:
def __init__(self, config):
self.config = config
rng = np.random.default_rng(config.seed)
d, f = config.dim, config.hidden
self.params = {
"embedding": rng.normal(0, .1, (config.vocab_size, d)),
"final_gain": np.ones(d),
"head": rng.normal(0, 1 / np.sqrt(d), (d, config.vocab_size)),
}
for layer in range(config.layers):
p = f"block{layer}."
self.params[p + "gain1"] = np.ones(d)
self.params[p + "gain2"] = np.ones(d)
for name in ("wq", "wk", "wv", "wo"):
self.params[p + name] = rng.normal(0, 1 / np.sqrt(d), (d, d))
self.params[p + "wo"] /= np.sqrt(2 * config.layers)
for name in ("w1", "w3"):
self.params[p + name] = rng.normal(0, 1 / np.sqrt(d), (d, f))
self.params[p + "w2"] = rng.normal(0, 1 / np.sqrt(f * 2 * config.layers), (f, d))
def split_heads(self, x):
b, t, d = x.shape
return x.reshape(b, t, self.config.heads, d // self.config.heads).transpose(0, 2, 1, 3)
@staticmethod
def join_heads(x):
b, h, t, dh = x.shape
return x.transpose(0, 2, 1, 3).reshape(b, t, h * dh)
def forward(self, ids):
ids = np.asarray(ids, dtype=np.int64)
if ids.ndim != 2 or not 0 < ids.shape[1] <= self.config.context:
raise ValueError("ids must have shape (batch, time), with 0 < time <= context")
if np.any(ids < 0) or np.any(ids >= self.config.vocab_size):
raise ValueError("token id out of range")
x = self.params["embedding"][ids]
caches = []
for layer in range(self.config.layers):
p = f"block{layer}."
z, norm1 = rmsnorm(x, self.params[p + "gain1"])
q = rope(self.split_heads(z @ self.params[p + "wq"]))
k = rope(self.split_heads(z @ self.params[p + "wk"]))
v = self.split_heads(z @ self.params[p + "wv"])
att, ac = causal_attention(q, k, v)
joined = self.join_heads(att)
r = x + joined @ self.params[p + "wo"]
n, norm2 = rmsnorm(r, self.params[p + "gain2"])
u, gate = n @ self.params[p + "w1"], n @ self.params[p + "w3"]
sigmoid = .5 * (1 + np.tanh(u / 2))
silu = u * sigmoid
hidden = silu * gate
x = r + hidden @ self.params[p + "w2"]
caches.append((p, z, norm1, ac, joined, n, norm2, u, gate, sigmoid, silu, hidden))
final, normf = rmsnorm(x, self.params["final_gain"])
logits = final @ self.params["head"]
return logits, (ids, caches, final, normf)
@staticmethod
def weight_grad(x, dy):
return x.reshape(-1, x.shape[-1]).T @ dy.reshape(-1, dy.shape[-1])
def backward(self, dlogits, cache):
ids, caches, final, normf = cache
g = {name: np.zeros_like(value) for name, value in self.params.items()}
g["head"] = self.weight_grad(final, dlogits)
dx, g["final_gain"] = rmsnorm_backward(dlogits @ self.params["head"].T, normf)
for p, z, norm1, ac, joined, n, norm2, u, gate, sigmoid, silu, hidden in reversed(caches):
g[p + "w2"] = self.weight_grad(hidden, dx)
dhidden = dx @ self.params[p + "w2"].T
du = dhidden * gate * (sigmoid + u * sigmoid * (1 - sigmoid))
dgate = dhidden * silu
g[p + "w1"] = self.weight_grad(n, du)
g[p + "w3"] = self.weight_grad(n, dgate)
dn = du @ self.params[p + "w1"].T + dgate @ self.params[p + "w3"].T
drnorm, g[p + "gain2"] = rmsnorm_backward(dn, norm2)
dr = dx + drnorm # residual branch
g[p + "wo"] = self.weight_grad(joined, dr)
datt = self.split_heads(dr @ self.params[p + "wo"].T)
dq, dk, dv = attention_backward(datt, ac)
dq, dk, dv = self.join_heads(rope(dq, True)), self.join_heads(rope(dk, True)), self.join_heads(dv)
dz = np.zeros_like(z)
for name, grad in (("wq", dq), ("wk", dk), ("wv", dv)):
g[p + name] = self.weight_grad(z, grad)
dz += grad @ self.params[p + name].T
dxnorm, g[p + "gain1"] = rmsnorm_backward(dz, norm1)
dx = dr + dxnorm # second residual branch
# Repeated token IDs must accumulate; fancy-index += is not sufficient.
np.add.at(g["embedding"], ids, dx)
return g
def loss(self, x, y, backward=False):
logits, cache = self.forward(x)
value, dlogits = cross_entropy(logits, y)
return (value, self.backward(dlogits, cache)) if backward else value
def save(self, path, tokenizer):
metadata = json.dumps({"config": asdict(self.config), "tokenizer": tokenizer.to_dict()})
np.savez_compressed(path, metadata=np.array(metadata), **self.params)
@classmethod
def load(cls, path):
from bpe import tokenizer_from_dict
with np.load(path, allow_pickle=False) as saved:
meta = json.loads(str(saved["metadata"]))
model = cls(Config(**meta["config"]))
for name in model.params:
model.params[name][...] = saved[name]
return model, tokenizer_from_dict(meta["tokenizer"])
class AdamW:
def __init__(self, params, lr=.003, betas=(.9, .999), eps=1e-8, decay=.01):
self.params, self.lr, self.betas, self.eps, self.decay = params, lr, betas, eps, decay
self.m = {k: np.zeros_like(v) for k, v in params.items()}
self.v = {k: np.zeros_like(v) for k, v in params.items()}
self.step_number = 0
def step(self, grads, max_norm=1.0):
norm = np.sqrt(sum(float(np.sum(g * g)) for g in grads.values()))
if not np.isfinite(norm):
raise FloatingPointError("non-finite gradient")
scale = min(1.0, max_norm / (norm + 1e-12))
self.step_number += 1
b1, b2 = self.betas
for name, value in self.params.items():
grad = grads[name] * scale
self.m[name] = b1 * self.m[name] + (1 - b1) * grad
self.v[name] = b2 * self.v[name] + (1 - b2) * grad**2
mhat = self.m[name] / (1 - b1**self.step_number)
vhat = self.v[name] / (1 - b2**self.step_number)
# Only matrices decay. Normalization gains are excluded.
if value.ndim > 1:
value *= 1 - self.lr * self.decay
value -= self.lr * mhat / (np.sqrt(vhat) + self.eps)
return float(norm)
def make_batch(tokens, batch_size, context, rng):
tokens = np.asarray(tokens, dtype=np.int64)
if len(tokens) <= context:
raise ValueError("need at least context + 1 tokens")
starts = rng.integers(0, len(tokens) - context, size=batch_size)
windows = tokens[starts[:, None] + np.arange(context + 1)[None, :]]
return windows[:, :-1], windows[:, 1:]
def generate(model, prompt_ids, count, temperature=.8, top_k=10, seed=9):
if not prompt_ids:
raise ValueError("prompt must contain at least one token")
if temperature < 0 or top_k < 0 or count < 0:
raise ValueError("temperature, top_k, count must be nonnegative")
rng = np.random.default_rng(seed)
ids = list(prompt_ids)
for _ in range(count):
# Deliberately recompute the full window: no KV cache in this CPU lab.
logits = model.forward(np.array([ids[-model.config.context:]]))[0][0, -1].copy()
if temperature == 0:
next_id = int(np.argmax(logits))
else:
logits /= temperature
if 0 < top_k < len(logits):
keep = np.argsort(logits, kind="stable")[-top_k:]
logits[np.setdiff1d(np.arange(len(logits)), keep)] = -np.inf
next_id = int(rng.choice(len(logits), p=softmax(logits)))
ids.append(next_id)
return ids