一句话总结:训练循环的核心确实只有前向、loss、反向、更新四步;真正容易出错的是四步周围的数据边界、评估模式、学习率、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 是 。cl100k_base 的完整词表很大,所以随机基线在 11 左右并不奇怪;但这不意味着所有“随机初始化模型”都必然精确得到这个数,初始化和 logits 尺度也会影响分布。
Perplexity 常写成 。它方便比较相同 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_dict | AdamW 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 怎样验证这个脚本?
至少做四类检查:
- Batch 边界:x/y 恰好错位一位,最后一个合法窗口不越界;
- 真正更新:一次
optimizer.step()后至少一个参数发生变化; - 评估无副作用:
estimate_loss()前后恢复原来的 train/eval mode,且不留下梯度; - 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 进行生成。