一句话总结:训练循环的核心确实只有前向、loss、反向、更新四步;真正容易出错的是四步周围的数据边界、评估模式、学习率、checkpoint 与恢复状态。

从语料切分、批次采样到训练、评估和 checkpoint 的完整训练闭环

📦 2024 年原始演示仓库github.com/waylandzhang/Transformer-from-scratch。历史 demo 把数据、模型、训练和生成放在一个文件;本章代码与第 18 章审校后的独立 model.py 配套。


19.1 模型到底在学什么?

刚初始化的模型不是“一无所知”:它已经有架构、位置编码和一组随机参数,只是还没有从这份语料中学到下一个 token 的规律。

继续用原来的“农夫山泉”例子。Tokenizer 先把文字变成 token IDs,然后从一段连续序列构造错位一位的输入和目标:

连续 Token: [, , , , , , , , , , ...]

输入 x:      [, , , , , , , ]
目标 y:      [, , , , , , , ]

在一个 forward 中,位置 0 学着预测“夫”,位置 1 根据“农、夫”学着预测“山”,后面依次类推。Causal mask 保证每个位置不能偷看目标右侧的 token。

四个动作是:

Forward  得到 logits  loss
Backward  得到各参数梯度
Optimizer step  修改参数
重复,但定期在 validation split 上停下来检查

训练 loss 下降只说明模型越来越适合训练样本。Validation loss 才帮助我们判断这种进步有没有延伸到未参与更新的数据。


19.2 数据:先切分,再随机抽窗口

19.2.1 Tokenizer 的完整词表

tokenizer_name = "cl100k_base"
tokenizer = tiktoken.get_encoding(tokenizer_name)
token_ids = tokenizer.encode(text, disallowed_special=())
tokens = torch.tensor(token_ids, dtype=torch.long)  # 留在 CPU

model_config = ModelConfig(
    vocab_size=tokenizer.n_vocab,
    context_length=128,
    d_model=80,
    n_layers=6,
    n_heads=4,
    dropout=0.1,
)

这里用 tokenizer.n_vocab,不能用“语料中出现过的最大 ID + 1”。后者会让模型词表随样本碰巧出现的 token 改变。

cl100k_base 只是本书 demo 选择的一个现成 tokenizer,不是“GPT-3 tokenizer”。训练自己的模型时,tokenizer 本身也是模型契约的一部分,checkpoint 必须记录它的名字或版本。

整份 token 数据先留在 CPU,只把当前 batch 搬到计算设备。小语料全部放 GPU 也许能跑,但这个习惯无法扩展到真正的大数据集。

19.2.2 连续切分,避免窗口跨界

def split_tokens(tokens, train_fraction, context_length):
    split = int(len(tokens) * train_fraction)
    train_data = tokens[:split]
    valid_data = tokens[split:]
    if min(len(train_data), len(valid_data)) <= context_length:
        raise ValueError("each split needs more than context_length tokens")
    return train_data, valid_data

对一条连续语料,先切分再采样,能保证训练窗口不会伸进 validation 区域。若语料由许多文档组成,更好的做法通常是先按文档分组、去重、再切分;不要把所有 token 打散后随机切,因为相邻段落或重复文档会泄漏到两边。

19.2.3 用向量索引构造 x 与 y

def get_batch(data, batch_size, context_length, device, generator):
    max_start = len(data) - context_length
    starts = torch.randint(
        0, max_start, (batch_size,), generator=generator
    )
    offsets = torch.arange(context_length)
    x = data[starts[:, None] + offsets]
    y = data[starts[:, None] + offsets + 1]
    return x.to(device), y.to(device)

torch.randint 的上界不包含在内。这里最后一个合法起点是 len(data)-context_length-1;它的 y 刚好取到数据最后一个 token,不会越界。


19.3 Loss 的数值该怎样读?

第 18 章的 Model.forward(idx, targets) 已经调用:

F.cross_entropy(
    logits.reshape(-1, vocab_size),
    targets.reshape(-1),
)

输入是 raw logits,不要先手动 Softmax。对每个 token,交叉熵等于正确 target 的负对数概率,再对 batch 与位置求平均。

如果模型对 V 个 token 真的是均匀分布,loss 是 lnV\ln Vcl100k_base 的完整词表很大,所以随机基线在 11 左右并不奇怪;但这不意味着所有“随机初始化模型”都必然精确得到这个数,初始化和 logits 尺度也会影响分布。

Perplexity 常写成 exp(loss)\exp(\text{loss})。它方便比较相同 tokenizer、相同数据处理、相同评估边界下的 run;换了 tokenizer 后,不能只拿一个 perplexity 数字横向宣布谁更好。

旧版列出的 loss 从 10.847 下降到 2.798,没有对应的可重现实验日志,因此本版不再把它当成真实结果。你的 loss 应该由当前数据、代码和 seed 实际跑出来。


19.4 评估要稳定,也要恢复原模式

@torch.inference_mode()
def estimate_loss(model, train_data, valid_data, config, device):
    was_training = model.training
    model.eval()
    try:
        result = {}
        eval_generator = torch.Generator().manual_seed(config.seed + 1)
        for name, data in (("train", train_data), ("valid", valid_data)):
            losses = []
            for _ in range(config.eval_batches):
                x, y = get_batch(
                    data,
                    config.batch_size,
                    model.config.context_length,
                    device,
                    eval_generator,
                )
                _, loss = model(x, y)
                losses.append(loss.detach().cpu())
            result[name] = torch.stack(losses).mean().item()
        return result
    finally:
        model.train(was_training)

这里有四个细节:

  • torch.inference_mode() 不记录梯度;
  • model.eval() 关闭 Dropout;这个模型没有 BatchNorm,没必要拿 BatchNorm 表格分散注意力;
  • 每次评估都从固定 generator seed 开始,因此不同 step 比较的是同一组随机窗口;
  • finally 恢复调用前的模式,而不是武断地总切回 training。

评估若只抽少量 batch,仍只是 validation loss 的 Monte Carlo 估计,会有方差。正式报告要扩大覆盖或遍历固定 validation 集。


19.5 Optimizer、调度与裁剪必须连起来

19.5.1 AdamW 参数组

def configure_optimizer(model, config):
    decay, no_decay = [], []
    for _, parameter in model.named_parameters():
        if not parameter.requires_grad:
            continue
        (decay if parameter.ndim >= 2 else no_decay).append(parameter)

    groups = [
        {"params": decay, "weight_decay": config.weight_decay},
        {"params": no_decay, "weight_decay": 0.0},
    ]
    return torch.optim.AdamW(
        groups,
        lr=config.peak_lr,
        betas=(0.9, 0.95),
    )

这是一条常见启发式:对矩阵状权重做 decay,对 bias 与 Norm 的一维参数不做。AdamW 的 weight decay 是与 loss gradient 解耦的收缩,不是简单的“L2 正则化”;Adam 也不是为每个参数自动找到一个最佳学习率。第 17 章已经解释过两者边界。

19.5.2 Warmup + cosine

def lr_at_step(step, total_steps, warmup_steps, peak_lr, min_lr):
    if step <= warmup_steps:
        return peak_lr * step / warmup_steps
    progress = (step - warmup_steps) / (total_steps - warmup_steps)
    cosine = 0.5 * (1.0 + math.cos(math.pi * progress))
    return min_lr + (peak_lr - min_lr) * cosine

训练 step 从 1 开始,warmup 最后一步到 peak_lr,总训练最后一步到 min_lr。这个边界与第 17 章保持一致。

19.5.3 一次参数更新

lr = lr_at_step(step, total_steps, warmup_steps, peak_lr, min_lr)
for group in optimizer.param_groups:
    group["lr"] = lr

optimizer.zero_grad(set_to_none=True)
_, loss = model(x, y)
loss.backward()
grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()

这里的 clip_grad_norm_ 返回裁剪前的总梯度范数,适合记录和排查尖峰。它不是“loss 大就自动修好”的开关;数据错误和数值错误仍要单独诊断。


19.6 Checkpoint 不只是模型权重

只保存 model.state_dict() 可以做推理,却不能忠实地继续训练。AdamW 的一阶矩、二阶矩、当前 step、随机采样器状态都丢了。

本章保存:

内容为什么需要
model_state_dict当前模型参数
optimizer_state_dictAdamW moments 与参数组状态
step下一步从哪里继续
model_config / train_config重建模型并审计配方
tokenizer_name保持 token IDs 语义一致
corpus_sha256防止恢复时悄悄换了语料
CPU / CUDA / MPS RNG state尽量延续 Dropout 与采样状态
batch generator state延续训练窗口的随机序列

保存时先写同目录临时文件,再 replace,避免进程中断留下半个 checkpoint。加载时要核对模型配置、训练配置、tokenizer 和语料指纹;本地 checkpoint 使用 weights_only=True,不可信来源的 checkpoint 仍不应该随便执行或加载。

即使保存了这些状态,不同硬件、kernel、PyTorch 版本和非确定性算子仍可能让 GPU run 不能 bit-for-bit 重现。Checkpoint 保证的是完整训练状态,而不是跨所有环境的数学同一条轨迹。


19.7 完整 train.py

import argparse
import hashlib
import math
from dataclasses import asdict, dataclass
from pathlib import Path

import tiktoken
import torch

from model import Model, ModelConfig


@dataclass
class TrainConfig:
    batch_size: int = 8
    total_steps: int = 500
    eval_interval: int = 50
    eval_batches: int = 10
    peak_lr: float = 3e-4
    min_lr: float = 3e-5
    warmup_steps: int = 50
    weight_decay: float = 0.1
    grad_clip: float = 1.0
    seed: int = 1337

    def __post_init__(self):
        if min(self.batch_size, self.total_steps, self.eval_interval,
               self.eval_batches, self.warmup_steps) <= 0:
            raise ValueError("step and batch settings must be positive")
        if self.warmup_steps >= self.total_steps:
            raise ValueError("warmup_steps must be smaller than total_steps")
        if not 0.0 <= self.min_lr <= self.peak_lr:
            raise ValueError("learning rates must satisfy 0 <= min_lr <= peak_lr")
        if self.weight_decay < 0 or self.grad_clip <= 0:
            raise ValueError("weight_decay and grad_clip are invalid")


def select_device():
    if torch.cuda.is_available():
        return torch.device("cuda")
    if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
        return torch.device("mps")
    return torch.device("cpu")


def split_tokens(tokens, train_fraction, context_length):
    if not 0.0 < train_fraction < 1.0:
        raise ValueError("train_fraction must be between 0 and 1")
    split = int(len(tokens) * train_fraction)
    train_data = tokens[:split]
    valid_data = tokens[split:]
    if min(len(train_data), len(valid_data)) <= context_length:
        raise ValueError("each split needs more than context_length tokens")
    return train_data, valid_data


def get_batch(data, batch_size, context_length, device, generator):
    max_start = len(data) - context_length
    if max_start <= 0:
        raise ValueError("data is too short for context_length")
    starts = torch.randint(
        0, max_start, (batch_size,), generator=generator
    )
    offsets = torch.arange(context_length)
    x = data[starts[:, None] + offsets]
    y = data[starts[:, None] + offsets + 1]
    return x.to(device), y.to(device)


@torch.inference_mode()
def estimate_loss(model, train_data, valid_data, config, device):
    was_training = model.training
    model.eval()
    try:
        result = {}
        eval_generator = torch.Generator().manual_seed(config.seed + 1)
        for name, data in (("train", train_data), ("valid", valid_data)):
            losses = []
            for _ in range(config.eval_batches):
                x, y = get_batch(
                    data,
                    config.batch_size,
                    model.config.context_length,
                    device,
                    eval_generator,
                )
                _, loss = model(x, y)
                losses.append(loss.detach().cpu())
            result[name] = torch.stack(losses).mean().item()
        return result
    finally:
        model.train(was_training)


def configure_optimizer(model, config):
    decay, no_decay = [], []
    for _, parameter in model.named_parameters():
        if not parameter.requires_grad:
            continue
        (decay if parameter.ndim >= 2 else no_decay).append(parameter)
    groups = [
        {"params": decay, "weight_decay": config.weight_decay},
        {"params": no_decay, "weight_decay": 0.0},
    ]
    return torch.optim.AdamW(
        groups,
        lr=config.peak_lr,
        betas=(0.9, 0.95),
    )


def lr_at_step(step, total_steps, warmup_steps, peak_lr, min_lr):
    if not 1 <= step <= total_steps:
        raise ValueError("step must be inside the training range")
    if not 0 < warmup_steps < total_steps:
        raise ValueError("warmup_steps must be inside the training range")
    if step <= warmup_steps:
        return peak_lr * step / warmup_steps
    progress = (step - warmup_steps) / (total_steps - warmup_steps)
    cosine = 0.5 * (1.0 + math.cos(math.pi * progress))
    return min_lr + (peak_lr - min_lr) * cosine


def save_checkpoint(path, model, optimizer, step, model_config,
                    train_config, tokenizer_name, corpus_sha256,
                    train_generator):
    path.parent.mkdir(parents=True, exist_ok=True)
    temporary = path.with_suffix(path.suffix + ".tmp")
    cuda_rng = torch.cuda.get_rng_state_all() if torch.cuda.is_available() else None
    mps_rng = (
        torch.mps.get_rng_state()
        if hasattr(torch, "mps") and torch.backends.mps.is_available()
        else None
    )
    torch.save({
        "model_state_dict": model.state_dict(),
        "optimizer_state_dict": optimizer.state_dict(),
        "step": step,
        "model_config": asdict(model_config),
        "train_config": asdict(train_config),
        "tokenizer_name": tokenizer_name,
        "corpus_sha256": corpus_sha256,
        "torch_rng_state": torch.get_rng_state(),
        "cuda_rng_state": cuda_rng,
        "mps_rng_state": mps_rng,
        "train_generator_state": train_generator.get_state(),
    }, temporary)
    temporary.replace(path)


def load_checkpoint(path, model, optimizer, model_config,
                    train_config, tokenizer_name, corpus_sha256,
                    train_generator, device):
    checkpoint = torch.load(path, map_location="cpu", weights_only=True)
    if checkpoint["model_config"] != asdict(model_config):
        raise ValueError("checkpoint model_config does not match")
    if checkpoint["train_config"] != asdict(train_config):
        raise ValueError("checkpoint train_config does not match")
    if checkpoint["tokenizer_name"] != tokenizer_name:
        raise ValueError("checkpoint tokenizer does not match")
    if checkpoint["corpus_sha256"] != corpus_sha256:
        raise ValueError("checkpoint corpus does not match")
    model.load_state_dict(checkpoint["model_state_dict"])
    optimizer.load_state_dict(checkpoint["optimizer_state_dict"])
    torch.set_rng_state(checkpoint["torch_rng_state"])
    train_generator.set_state(checkpoint["train_generator_state"])
    if device.type == "cuda" and checkpoint["cuda_rng_state"] is not None:
        torch.cuda.set_rng_state_all(checkpoint["cuda_rng_state"])
    if device.type == "mps" and checkpoint["mps_rng_state"] is not None:
        torch.mps.set_rng_state(checkpoint["mps_rng_state"])
    return int(checkpoint["step"])


def parse_args():
    parser = argparse.ArgumentParser()
    parser.add_argument("--data", type=Path, required=True)
    parser.add_argument("--output", type=Path,
                        default=Path("model/checkpoint.pt"))
    parser.add_argument("--resume", type=Path)
    return parser.parse_args()


def main():
    args = parse_args()
    train_config = TrainConfig()
    tokenizer_name = "cl100k_base"
    tokenizer = tiktoken.get_encoding(tokenizer_name)

    text = args.data.read_text(encoding="utf-8")
    corpus_sha256 = hashlib.sha256(text.encode("utf-8")).hexdigest()
    token_ids = tokenizer.encode(text, disallowed_special=())
    tokens = torch.tensor(token_ids, dtype=torch.long)  # keep corpus on CPU

    model_config = ModelConfig(
        vocab_size=tokenizer.n_vocab,
        context_length=128,
        d_model=80,
        n_layers=6,
        n_heads=4,
        dropout=0.1,
    )
    train_data, valid_data = split_tokens(
        tokens, train_fraction=0.9,
        context_length=model_config.context_length,
    )

    device = select_device()
    torch.manual_seed(train_config.seed)
    if device.type == "cuda":
        torch.cuda.manual_seed_all(train_config.seed)

    model = Model(model_config).to(device)
    optimizer = configure_optimizer(model, train_config)
    train_generator = torch.Generator().manual_seed(train_config.seed)
    start_step = 0

    if args.resume is not None:
        start_step = load_checkpoint(
            args.resume, model, optimizer, model_config,
            train_config, tokenizer_name, corpus_sha256,
            train_generator, device,
        )

    model.train()
    if start_step == 0:
        metrics = estimate_loss(
            model, train_data, valid_data, train_config, device
        )
        print(f"step=0 train={metrics['train']:.4f} "
              f"valid={metrics['valid']:.4f}")

    for step in range(start_step + 1, train_config.total_steps + 1):
        lr = lr_at_step(
            step,
            train_config.total_steps,
            train_config.warmup_steps,
            train_config.peak_lr,
            train_config.min_lr,
        )
        for group in optimizer.param_groups:
            group["lr"] = lr

        x, y = get_batch(
            train_data,
            train_config.batch_size,
            model_config.context_length,
            device,
            train_generator,
        )
        optimizer.zero_grad(set_to_none=True)
        _, loss = model(x, y)
        loss.backward()
        grad_norm = torch.nn.utils.clip_grad_norm_(
            model.parameters(), train_config.grad_clip
        )
        optimizer.step()

        should_evaluate = (
            step % train_config.eval_interval == 0
            or step == train_config.total_steps
        )
        if should_evaluate:
            metrics = estimate_loss(
                model, train_data, valid_data, train_config, device
            )
            print(
                f"step={step} lr={lr:.2e} grad={grad_norm.item():.3f} "
                f"train={metrics['train']:.4f} valid={metrics['valid']:.4f}"
            )
            save_checkpoint(
                args.output, model, optimizer, step,
                model_config, train_config, tokenizer_name, corpus_sha256,
                train_generator,
            )


if __name__ == "__main__":
    main()

19.8 怎样验证这个脚本?

至少做四类检查:

  1. Batch 边界:x/y 恰好错位一位,最后一个合法窗口不越界;
  2. 真正更新:一次 optimizer.step() 后至少一个参数发生变化;
  3. 评估无副作用estimate_loss() 前后恢复原来的 train/eval mode,且不留下梯度;
  4. Checkpoint 往返:保存后重新构建 model/optimizer,参数、optimizer moments、step 和 batch generator state 都能恢复,配置或语料变化则应该拒绝恢复。

短跑时要同时看:

  • loss 是不是有限数;
  • gradient norm 有没有突然爆炸;
  • train 与 valid 是否都在相同评估窗口上取得进展;
  • resume 后的第一步是不是接着原 step,而不是从 0 偷偷重来。

不要为了让示例看起来漂亮而编一条平滑下降曲线。真实小模型会抖动,少量 validation batch 也会有抽样误差。


19.9 本章总结

  • 下一 token 训练来自同一连续序列的 x 与右移一位的 y
  • 先切 train/valid,再在各自内部采样窗口,避免跨边界泄漏;
  • 语料留在 CPU,batch 才搬到 device;词表大小来自 tokenizer 契约;
  • 评估要关闭 Dropout、关闭梯度、使用可比较窗口并恢复原模式;
  • AdamW 参数组、warmup/cosine、gradient clipping 要在同一个 step 中正确排序;
  • 能继续训练的 checkpoint 必须包含 optimizer、step、配置、语料指纹和随机状态,不只是模型权重。

本章交付物

  • 能解释 x/y 的一位错位和 causal mask 的分工
  • 能写出无越界、无 train/valid 跨界的 batch sampler
  • 能区分随机均匀 loss 基线、validation loss 和 perplexity 的适用边界
  • 能写出含 schedule、gradient clipping 与稳定评估的训练循环
  • 能保存并恢复完整训练状态

下一章预告

训练脚本会留下一个包含模型配置、tokenizer 名称和权重的 checkpoint。下一章写 inference.py:安全加载它、恢复完全相同的模型契约,再把 prompt 变成 token IDs 进行生成。

引用本文 / Cite
Zhang, Wayland (2026). 第 19 章:手写 Train.py - 训练循环. In Transformer 架构:从直觉到实现. https://waylandz.com/llm-transformer-book/%E7%AC%AC19%E7%AB%A0-%E6%89%8B%E5%86%99Train.py-%E8%AE%AD%E7%BB%83%E5%BE%AA%E7%8E%AF/
@incollection{zhang2026transformer_19_-_Train_py-,
  author = {Zhang, Wayland},
  title = {第 19 章:手写 Train.py - 训练循环},
  booktitle = {Transformer 架构:从直觉到实现},
  year = {2026},
  url = {https://waylandz.com/llm-transformer-book/%E7%AC%AC19%E7%AB%A0-%E6%89%8B%E5%86%99Train.py-%E8%AE%AD%E7%BB%83%E5%BE%AA%E7%8E%AF/}
}