Skip to content

llm/train.py ​

配套源码,运行方法见同目录 README。返回实验总览。

python
"""Usage from repository root: python3 labs/llm/train.py --steps 300."""
import argparse
import hashlib
import json
from pathlib import Path
import platform
import time
import numpy as np
from bpe import ByteBPE, CharacterTokenizer
from model import Config, TinyTransformer, AdamW, make_batch, generate


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--steps", type=int, default=300)
    parser.add_argument("--seed", type=int, default=7)
    parser.add_argument("--batch-size", type=int, default=8)
    parser.add_argument("--context", type=int, default=24)
    parser.add_argument("--dim", type=int, default=32)
    parser.add_argument("--lr", type=float, default=.003)
    parser.add_argument("--tokenizer", choices=["char", "bpe"], default="char")
    parser.add_argument("--merges", type=int, default=32)
    parser.add_argument("--corpus", type=Path, default=Path(__file__).with_name("corpus.txt"))
    parser.add_argument("--out", type=Path, default=Path(__file__).with_name("runs"))
    args = parser.parse_args()
    if args.steps < 1 or args.batch_size < 1 or args.lr <= 0:
        parser.error("steps, batch-size and lr must be positive")
    raw = args.corpus.read_text(encoding="utf-8")
    documents = raw.splitlines(keepends=True)
    cutoff = int(.8 * len(documents))
    train_docs, val_docs = documents[:cutoff], documents[cutoff:]
    tokenizer = (CharacterTokenizer.train(train_docs) if args.tokenizer == "char"
                 else ByteBPE.train(train_docs, args.merges))
    train_ids = tokenizer.encode("".join(train_docs))
    val_ids = tokenizer.encode("".join(val_docs))
    config = Config(len(tokenizer.vocab), dim=args.dim, hidden=2 * args.dim,
                    context=args.context, seed=args.seed)
    model, rng = TinyTransformer(config), np.random.default_rng(args.seed)
    optimizer = AdamW(model.params, lr=args.lr)
    # Held-out evaluation batches are fixed; neither the tokenizer nor SGD sees them.
    eval_rng = np.random.default_rng(101)
    train_eval = make_batch(train_ids, 32, args.context, eval_rng)
    val_eval = make_batch(val_ids, 32, args.context, eval_rng)
    history = [{"step": 0, "train_nll": model.loss(*train_eval), "val_nll": model.loss(*val_eval)}]
    print(json.dumps(history[-1]), flush=True)
    start = time.perf_counter()
    for step in range(1, args.steps + 1):
        x, y = make_batch(train_ids, args.batch_size, args.context, rng)
        loss, grads = model.loss(x, y, backward=True)
        # Short linear warmup and cosine decay, with a 10% floor.
        warmup = min(20, max(1, args.steps // 10))
        if step <= warmup:
            optimizer.lr = args.lr * step / warmup
        else:
            progress = (step - warmup) / max(1, args.steps - warmup)
            optimizer.lr = args.lr * (.1 + .9 * .5 * (1 + np.cos(np.pi * progress)))
        grad_norm = optimizer.step(grads)
        if step % 50 == 0 or step == args.steps:
            record = {"step": step, "train_nll": model.loss(*train_eval),
                      "val_nll": model.loss(*val_eval), "batch_nll": loss,
                      "gradient_norm_before_clip": grad_norm, "lr": optimizer.lr}
            history.append(record)
            print(json.dumps(record), flush=True)
    elapsed = time.perf_counter() - start
    prompt = "the fox "
    sample = tokenizer.decode(generate(model, tokenizer.encode(prompt), 120, seed=args.seed))
    args.out.mkdir(parents=True, exist_ok=True)
    model.save(args.out / "model.npz", tokenizer)
    metadata = {"python": platform.python_version(), "numpy": np.__version__,
                "device": "CPU / NumPy float64", "seed": args.seed,
                "corpus_sha256": hashlib.sha256(raw.encode()).hexdigest(),
                "train_documents": len(train_docs), "validation_documents": len(val_docs),
                "train_tokens": len(train_ids), "validation_tokens": len(val_ids),
                "parameters": sum(p.size for p in model.params.values()),
                "elapsed_seconds": elapsed, "history": history, "sample": sample,
                "note": "Original tiny synthetic corpus; not a quality benchmark. Checkpoint is for inference, not exact optimizer resume."}
    (args.out / "metrics.json").write_text(json.dumps(metadata, indent=2), encoding="utf-8")
    print("sample:", sample)
    print("saved:", args.out.resolve())


if __name__ == "__main__":
    main()