第 13 课:训练循环 —— 让模型真正开始学习
第 13 课:训练循环 —— 让模型真正开始学习
代码位置:src/train.rs 演示入口:src/main.rs
1. 本课要搞懂的问题
- 训练一个模型,代码上到底要做哪几步?顺序为什么不能乱?
loss.backward()到底做了什么?为什么调用一次,所有参数的梯度就都有了?- 梯度爆炸是什么?
clip_grad_norm是怎么防止它的? - 为什么每步训练结束都要"清零梯度"?不清零会怎样?
2. 训练循环总览:六步走
训练 = 把"前向 → 算损失 → 反向 → 更新"这个动作反复执行成千上万次。
train_gpt 里每一次循环(一个 step)严格按下面的顺序执行:
| 步骤 | 动作 | 代码 | 干什么 |
|---|---|---|---|
| 1 | 采样 batch | loader.sample_batch(rng) |
从语料里随机取一批 (x, y) 训练对 |
| 2 | 前向 + 损失 | model.forward(...) → cross_entropy_loss(...) |
让模型预测一遍,算出"错得有多离谱" |
| 3 | 反向 | loss.backward() |
沿计算图把损失对每个参数的偏导(梯度)算出来 |
| 4 | 梯度裁剪 | clip_grad_norm(¶ms, 1.0) |
总梯度范数超阈值就等比缩小,防梯度爆炸 |
| 5 | 更新参数 | opt.step() |
按梯度方向调整每个参数,让损失变小 |
| 6 | 清零梯度 | opt.zero_grad() |
把梯度归零,为下一步重新计算做准备 |
顺序不能乱:前向必须发生在反向之前(没有前向就没有计算图);更新必须发生在反向之后(没有梯度就没东西可更新);清零必须放在最后(否则下一步的梯度会叠加在旧梯度上)。
对应代码(src/train.rs 中的 train_gpt,省略了日志打印):
for step in 0..steps {
// 1. 采样 batch
let (x, y) = loader.sample_batch(rng);
let b = batch_size;
let t = block_size;
// 2. 前向 + 损失
let logits = model.forward(&x, b, t, None);
let loss = cross_entropy_loss(&logits, &y);
// 3. 反向
loss.backward();
// 4. 梯度裁剪
clip_grad_norm(¶ms, 1.0);
// 5. 更新参数(设置当前学习率)
opt.lr = scheduler.lr();
opt.step();
// 6. 清零梯度
opt.zero_grad();
scheduler.step();
}
下面把这六步逐一拆开讲。
3. 准备阶段:参数、优化器、学习率调度器
进入循环前,train_gpt 先做了三件准备工作:
let params = model.parameters(); // 收集模型所有参数
let mut opt = AdamW::new(max_lr, params.clone(), 0.01); // 优化器(第 17 课详述)
let mut scheduler = LRScheduler::new(warmup_steps, steps, max_lr, max_lr * 0.1); // 学习率调度(第 20 课详述)
parameters():通过第 5 课实现的Moduletrait,把模型里所有可训练张量(embedding、每层 Linear、LayerNorm 的权重与偏置)收进一个Vec<Tensor>。优化器更新和梯度裁剪都靠这份清单。AdamW:负责"怎么更新参数",内部为每个参数维护动量 m 和二阶动量 v(第 17 课专门讲)。LRScheduler:决定每一步用多大的学习率——前期 warmup 从 0 线性爬升,后期 cosine 衰减。本课只需知道它输出一个学习率、赋值给opt.lr即可。
4. 第一步:采样 batch
let (x, y) = loader.sample_batch(rng);
x 是一批"输入 token",y 是"标准答案"(下一个 token)。它们怎么来的、batch_size 和 block_size 是什么意思,是第 14 课的主题,本课先把它当黑盒:每次调用就得到一批新的 (x, y)。
5. 第二步:前向 + 损失
let logits = model.forward(&x, b, t, None); // [B*T, vocab_size]
let loss = cross_entropy_loss(&logits, &y); // 标量
model.forward让数据走一遍整个 Transformer,输出logits:每个位置"预测下一个 token"的分数(未归一化),形状[B*T, vocab_size]。cross_entropy_loss(第 6 课已实现)比较预测与标准答案 y:
loss = -mean( log softmax(logits)[i, y[i]] )
模型给正确 token 的概率越高,loss 越小;loss 为 0 意味着预测完全正确。训练的目标就是把这个数压到最低。
6. 第三步:反向 —— loss.backward() 做了什么
这是自动微分(第 2 课)的核心。看 src/autograd.rs 里的实现:
pub fn backward(&self) {
assert_eq!(self.rank(), 0, "backward() 只支持标量(0 维)输出");
// 1. 先把输出自身的梯度设为 1(d loss / d loss = 1)
self.grad.borrow_mut()[0] = 1.0;
// 2. 迭代式 DFS 拓扑排序(避免递归栈溢出,见第 4 课 6.2 节)
// 用显式栈模拟递归,得到"先子节点后父节点"的拓扑序
let mut stack: Vec<(Tensor, usize)> = Vec::new();
let mut visited: HashSet<usize> = HashSet::new();
// ... 三色标记法遍历 ...
// 3. 把拓扑序反过来:从 loss 开始,一层一层往回传梯度
for t in order.iter().rev() {
if let Some(b) = &t.backward { b(); }
}
}
拆开看有三步:
- 初始化:把 loss 自己的梯度设为 1.0。因为我们要算的是
∂loss/∂θ,链式法则的起点就是∂loss/∂loss = 1。 - 拓扑排序:前向时每个中间结果都记住了自己的
parents(谁算出了我)和backward(如何把梯度传回给我的输入)。DFS 从 loss 出发往下走到所有叶子节点,得到"先依赖、后被依赖"的拓扑序。 - 逆序传播:按拓扑序的反向(从 loss 往参数方向)依次调用每个节点的
backward闭包。每个节点把自己的梯度g按链式法则乘上局部导数,累加到它的输入(包括参数)的梯度上。
loss ──► softmax ──► matmul ──► ... ──► Linear(W) ──► embedding
g=1 累积到输入 累积到输入 累积到 W 累积到表
为什么是"累加"而不是"赋值"?因为一个参数会被计算图中很多地方共用(比如同一个权重矩阵被整批样本共享),每个分支的贡献都要加在一起,这才是真正的偏导数。
所以调用一次 loss.backward(),模型里所有参数(params 里每个张量)的 .grad() 就都拿到了 ∂loss/∂θ。
7. 第四步:梯度裁剪 clip_grad_norm
问题:梯度爆炸。 深层网络 + 长序列训练时,链式法则连乘会导致梯度呈指数级放大,一步更新就可能把参数"推飞",loss 直接变成 NaN/无穷大。
对策: 在更新参数之前,先检查所有参数梯度的总范数;如果超过阈值 max_norm,就整体等比缩放,让方向不变、大小受控。
代码(src/train.rs):
pub fn clip_grad_norm(params: &[Tensor], max_norm: f32) {
// 1. 算总范数:所有梯度元素的平方和,再开根号
let mut total = 0.0f32;
for p in params {
let g = p.grad();
for &v in &g { total += v * v; }
}
let norm = total.sqrt();
// 2. 超过阈值才裁剪:整体乘 scale = max_norm / norm
if norm > max_norm {
let scale = max_norm / norm;
for p in params {
let g = p.grad();
let scaled: Vec<f32> = g.iter().map(|&v| v * scale).collect();
p.grad_set(scaled); // 覆盖写回(tensor.rs 里专为裁剪提供的 API)
}
}
}
对应公式:
norm = sqrt( Σ gᵢ² )
gᵢ' = gᵢ · min(1, max_norm / norm)
| 情况 | 行为 |
|---|---|
norm ≤ max_norm |
梯度正常,什么都不做 |
norm > max_norm |
所有梯度乘 max_norm / norm,总范数被压回 max_norm |
注意
max_norm这里取 1.0,训练 LLM 时这是很常见的取值。裁剪不改变梯度方向,只是限制步长上限——这比粗暴地调低学习率更精准。
8. 第五步:更新参数 opt.step()
opt.lr = scheduler.lr(); // 这一步用多大的学习率,由调度器决定
opt.step(); // 按梯度更新所有参数
本课用的是 AdamW(第 17 课详述),核心更新规则(简写):
m ← β₁·m + (1-β₁)·g 一阶动量(记住梯度方向)
v ← β₂·v + (1-β₂)·g² 二阶动量(感知坡度陡缓)
θ ← θ - lr·m̂/(√v̂ + ε) - lr·wd·θ
直觉:沿着梯度的反方向迈一小步,步子大小由学习率和"历史上梯度的大小"共同决定。与最朴素的 SGD(θ = θ - lr·g,第 6 课)相比,AdamW 对每个参数自适应地调整步长,训练更稳。
9. 第六步:清零梯度 opt.zero_grad()
opt.zero_grad(); // 内部就是对每个参数调用 p.zero_grad()
为什么必须清零?回顾第 6 节:反向传播是累加梯度。如果不清零,下一步前向+反向得到的梯度会叠加上一步的旧梯度,参数更新方向就被污染了,loss 会来回震荡甚至发散。
正确的时间线(注意第 3 步的累加发生在"同一步内、多次贡献之间",第 6 步的清零发生在"不同步之间"):
step 0: 梯度 = 0 → 前向 → 反向(累加) → 裁剪 → 更新 → 清零 → 梯度 = 0
step 1: 梯度 = 0 → 前向 → 反向(累加) → 裁剪 → 更新 → 清零 → 梯度 = 0
...
10. 看训练日志:loss 在下降
train_gpt 末尾每隔 eval_every 步打印一行(main.rs 里每 100 步一次):
step 0 | lr 0.00006 | loss 3.4658
step 100 | lr 0.00138 | loss 1.8701
step 200 | lr 0.00228 | loss 1.4210
...
step 599 | lr 0.00030 | loss 0.8142
判读方法:
- loss 总体下降 → 模型在学习,六步循环工作正常;
- loss 出现 NaN/剧烈抖动 → 大概率梯度爆炸,检查
clip_grad_norm的阈值或调低max_lr; - loss 降得很慢/不降 → 学习率太小,或数据/模型配置有问题(下一课会讲数据怎么来)。
11. 动手练习
- 把
clip_grad_norm(¶ms, 1.0)这一行注释掉再跑cargo run,观察 loss 是否更容易出现剧烈波动(体会它到底在防什么)。 - 把
max_norm从 1.0 改成 0.01 再训练,对比收敛速度(阈值太小会"拖慢"训练)。 - 试着删掉
opt.zero_grad(),观察 loss 曲线会发生什么,并解释原因(提示:梯度在跨步累加)。 - 在循环里打印第 5 步更新前后某个参数
params[0].data()[0]的变化,验证"更新确实发生在 step 之后"。
12. 本课总结
- 训练循环六步,顺序固定:采样 → 前向 → 损失 → 反向 → 裁剪 → 更新 → 清零(裁剪是防爆炸的保险,可视为第 3.5 步)
loss.backward():置初值 → DFS 拓扑排序 → 逆序沿链式法则累加梯度,一次调用搞定所有参数clip_grad_norm:总梯度范数超阈值就等比缩放,只限大小、不改方向opt.zero_grad():跨步清零,防止梯度污染- 下一步(第 14 课):搞清楚循环第一步
sample_batch返回的 x/y 到底是怎么从语料里构造出来的