从零实现一次完整训练循环

把 Batch、Forward、Loss、Backward 和 Optimizer Step 真正连起来

目录

  前七课分别看过数据、模型、显存、多卡和混合精度。它们最后都会汇到同一段代码里:训练循环。

  训练循环并不神秘。模型先根据一批输入做预测,用 Loss 衡量答案有多错,再通过 Backward 找出每个参数应该往哪边调整,最后由 Optimizer 真正更新参数。下一批数据到来后,重复这一过程。

一次训练迭代就是:取数据 → 清梯度 → Forward → 算 Loss → Backward → 更新参数。

一次完整训练循环

从一批 Token 到一次参数更新,点击查看原始 SVG

1、训练循环到底在循环什么

  一个 Epoch 表示完整看过一遍训练数据。一个 Step 表示处理一个 Batch,并尝试更新一次参数。

训练过程
├─ Epoch 1
│  ├─ Batch 1 → Step 1
│  ├─ Batch 2 → Step 2
│  └─ ...
├─ Epoch 2
└─ ...

  大模型的数据量很大,实际训练往往更关心 Step,而不是一定要跑完多少个 Epoch。

2、准备输入和正确答案

  自回归语言模型的输入 x 和标签 y 来自同一段 Token,只错开一个位置:

原始序列:[我, 喜欢, 学习, 大模型, 。]
输入 x :[我, 喜欢, 学习, 大模型]
标签 y :[喜欢, 学习, 大模型, 。]

  如果 Batch Size 是 B,序列长度是 T,那么 x 和 y 的形状通常都是 [B, T]。

3、Forward:得到预测

  把 x 送进模型,得到每个位置对整个词表的预测分数 logits:

\[x:[B,T]\quad\longrightarrow\quad logits:[B,T,V]\]

  这里的 $V$ 是词表大小。Forward 只负责“根据当前参数给出预测”,还没有修改模型。

4、Loss:衡量错得多远

  语言模型通常使用交叉熵 Loss。它会比较每个位置的 logits 和正确 Token ID:

loss = F.cross_entropy(
    logits.reshape(-1, vocab_size),
    y.reshape(-1),
)

  reshape 把 [B, T, V] 展平成 [B×T, V],标签则从 [B, T] 变成 [B×T]。只是换了形状,没有改变对应关系。

  Loss 越小,说明模型给正确 Token 的概率越高。但一次 Batch 的 Loss 下降,不代表模型已经学会泛化,所以还要定期看验证集。

5、Backward:计算梯度

  loss.backward() 会沿计算图反向计算梯度:

\[\frac{\partial L}{\partial \theta}\]

  它回答的是:“如果稍微改变参数 $\theta$,Loss 会朝哪个方向变化?”

  Backward 只计算并累积梯度,并不直接修改参数。PyTorch 默认把新梯度加到旧梯度上,因此每次普通更新前要先执行:

optimizer.zero_grad(set_to_none=True)

6、Optimizer Step:更新参数

  optimizer.step() 才是真正修改参数的动作。以最简单的梯度下降为例:

\[\theta_{new}=\theta_{old}-\eta\nabla_\theta L\]

  其中 $\eta$ 是学习率。AdamW 还会保存动量等优化器状态,但核心目的相同:根据梯度更新参数,让下一次预测更接近正确答案。

  实际训练中常在更新前做梯度裁剪,避免偶发的大梯度让训练突然失控:

torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

7、一段完整的 PyTorch 代码

  下面这段代码省略了模型定义,但保留了一次训练真正需要的顺序。假设 train_loader 每次返回形状为 [B, T] 的 x 和 y。

import torch
import torch.nn.functional as F

device = "cuda"
model = model.to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)

model.train()

for step, (x, y) in enumerate(train_loader, start=1):
    x = x.to(device, non_blocking=True)
    y = y.to(device, non_blocking=True)

    # 1. 清掉上一步留下的梯度
    optimizer.zero_grad(set_to_none=True)

    # 2. Forward:大规模矩阵计算使用 BF16
    with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
        logits = model(x)                    # [B, T, V]
        loss = F.cross_entropy(
            logits.reshape(-1, logits.size(-1)),
            y.reshape(-1),
        )

    # 3. Backward:把梯度写入 parameter.grad
    loss.backward()

    # 4. 先裁剪,再更新参数
    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
    optimizer.step()

    if step % 100 == 0:
        print(f"step={step}, loss={loss.item():.4f}")

  使用 BF16 时通常不需要 GradScaler。如果硬件只适合 FP16,可以使用 torch.amp.GradScaler 对 Loss 做动态缩放;顺序会变成 scale → backward → unscale → clip → step → update。

训练不只是这五行

  真实项目还会在循环外或固定 Step 插入一些工作:

工作 放置位置 作用
学习率调度 optimizer.step() 后 控制不同阶段的更新幅度
验证集评估 每隔若干 Step 检查模型是否真的变好
保存 Checkpoint 评估后或定期 保存模型、优化器和当前 Step
梯度累积 多个 Micro Batch 之后 小显存模拟更大的 Batch
日志监控 每隔若干 Step 观察 Loss、学习率、吞吐和显存

  评估时要切换到 model.eval() 和 torch.no_grad(),完成后再切回 model.train()。保存 Checkpoint 时不仅要保存模型参数,也要保存优化器状态和训练进度,否则无法真正从中断位置继续。

最容易写错的 5 个地方

  1. 忘记清梯度,导致不同 Step 的梯度意外累积。
  2. logits 与 y 没有正确错位,模型在学习复制当前 Token。
  3. 只保存模型参数,没有保存优化器和 Step。
  4. 验证后忘记切回 model.train()。
  5. 使用 FP16 时先裁剪缩放后的梯度,而不是先 unscale_。

记住这 5 件事

  1. Forward 产生预测,Loss 衡量预测与答案的差距。
  2. Backward 计算梯度,但不会更新参数。
  3. Optimizer Step 才真正修改模型参数。
  4. PyTorch 默认累积梯度,普通训练每一步都要先清梯度。
  5. 训练循环之外,还要做好评估、日志和 Checkpoint。

第八课复习总图

大模型第八课复习总图

第八课复习总图,点击查看原始 SVG

  到这里,我们已经把第一阶段的知识串成了一条完整链路。下一课开始拆开模型内部:一个 Transformer Block 里,到底有哪些组件?