Skip to content

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