第 13 课:训练循环 —— 让模型真正开始学习

2026-09-14 干徒
RustLLM训练

第 13 课:训练循环 —— 让模型真正开始学习

代码位置:src/train.rs 演示入口:src/main.rs

1. 本课要搞懂的问题

  1. 训练一个模型,代码上到底要做哪几步?顺序为什么不能乱?
  2. loss.backward() 到底做了什么?为什么调用一次,所有参数的梯度就都有了?
  3. 梯度爆炸是什么?clip_grad_norm 是怎么防止它的?
  4. 为什么每步训练结束都要"清零梯度"?不清零会怎样?

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(&params, 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(&params, 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 课实现的 Module trait,把模型里所有可训练张量(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_sizeblock_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(); }
    }
}

拆开看有三步:

  1. 初始化:把 loss 自己的梯度设为 1.0。因为我们要算的是 ∂loss/∂θ,链式法则的起点就是 ∂loss/∂loss = 1
  2. 拓扑排序:前向时每个中间结果都记住了自己的 parents(谁算出了我)和 backward(如何把梯度传回给我的输入)。DFS 从 loss 出发往下走到所有叶子节点,得到"先依赖、后被依赖"的拓扑序。
  3. 逆序传播:按拓扑序的反向(从 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. 动手练习

  1. clip_grad_norm(&params, 1.0) 这一行注释掉再跑 cargo run,观察 loss 是否更容易出现剧烈波动(体会它到底在防什么)。
  2. max_norm 从 1.0 改成 0.01 再训练,对比收敛速度(阈值太小会"拖慢"训练)。
  3. 试着删掉 opt.zero_grad(),观察 loss 曲线会发生什么,并解释原因(提示:梯度在跨步累加)。
  4. 在循环里打印第 5 步更新前后某个参数 params[0].data()[0] 的变化,验证"更新确实发生在 step 之后"。

12. 本课总结

  • 训练循环六步,顺序固定:采样 → 前向 → 损失 → 反向 → 裁剪 → 更新 → 清零(裁剪是防爆炸的保险,可视为第 3.5 步)
  • loss.backward():置初值 → DFS 拓扑排序 → 逆序沿链式法则累加梯度,一次调用搞定所有参数
  • clip_grad_norm:总梯度范数超阈值就等比缩放,只限大小、不改方向
  • opt.zero_grad():跨步清零,防止梯度污染
  • 下一步(第 14 课):搞清楚循环第一步 sample_batch 返回的 x/y 到底是怎么从语料里构造出来的
Rust 大语言模型 学习指南共 22 章
1从零用 Rust 实现大语言模型 —— 学习计划2第 1 课:张量 Tensor —— 一切的基础3第 2 课:自动微分 Autograd —— 让模型学会「自我修正」4第 3 课:张量运算扩展 —— 广播、归约、softmax、批量矩阵乘法5第 4 课:模块化重构 —— 项目结构分层与 Rc<RefCell> 架构6第 5 课:线性层与激活函数 —— 神经网络的「积木」7第 6 课:损失函数与优化器 —— 让模型知道「错在哪、怎么改」8第 7 课:第一个 MLP —— 教会神经网络算 XOR9第 8 课:BPE 分词器 —— 让模型「读懂」文字10第 9 课:注意力机制 —— 让 token 互相「看」11第 10 课:多头注意力 —— 让模型「多角度」看世界12第 11 课:位置编码与归一化 —— 让序列带上「位置感」13第 12 课:完整 GPT 模型 —— 把积木拼成能预测下一个词的模型14第 13 课:训练循环 —— 让模型真正开始学习本篇15第 14 课:数据加载 —— 文本如何变成训练样本16第 15 课:推理与采样 —— 让模型「创造性」地生成17第 16 课:训练小 GPT —— 看 loss 从 1.46 一路降到 0.1618第 17 课:AdamW 优化器 —— 给梯度下降装上「惯性」和「自适应步长」19第 18 课:KV Cache —— 让逐 token 生成不再重复计算20第 19 课:RoPE 旋转位置编码 —— 把「相对位置」揉进注意力21第 20 课:学习率调度与收尾 —— warmup、cosine decay 与全项目总结22第 21 课:GPU 加速训练与推理