目录
前七课分别看过数据、模型、显存、多卡和混合精度。它们最后都会汇到同一段代码里:训练循环。
训练循环并不神秘。模型先根据一批输入做预测,用 Loss 衡量答案有多错,再通过 Backward 找出每个参数应该往哪边调整,最后由 Optimizer 真正更新参数。下一批数据到来后,重复这一过程。
一次训练迭代就是:取数据 → 清梯度 → Forward → 算 Loss → Backward → 更新参数。
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:
这里的 $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() 会沿计算图反向计算梯度:
它回答的是:“如果稍微改变参数 $\theta$,Loss 会朝哪个方向变化?”
Backward 只计算并累积梯度,并不直接修改参数。PyTorch 默认把新梯度加到旧梯度上,因此每次普通更新前要先执行:
optimizer.zero_grad(set_to_none=True)
6、Optimizer Step:更新参数
optimizer.step() 才是真正修改参数的动作。以最简单的梯度下降为例:
其中 $\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 个地方
- 忘记清梯度,导致不同 Step 的梯度意外累积。
logits与y没有正确错位,模型在学习复制当前 Token。- 只保存模型参数,没有保存优化器和 Step。
- 验证后忘记切回
model.train()。 - 使用 FP16 时先裁剪缩放后的梯度,而不是先
unscale_。
记住这 5 件事
- Forward 产生预测,Loss 衡量预测与答案的差距。
- Backward 计算梯度,但不会更新参数。
- Optimizer Step 才真正修改模型参数。
- PyTorch 默认累积梯度,普通训练每一步都要先清梯度。
- 训练循环之外,还要做好评估、日志和 Checkpoint。
第八课复习总图
到这里,我们已经把第一阶段的知识串成了一条完整链路。下一课开始拆开模型内部:一个 Transformer Block 里,到底有哪些组件?