5.1 训练循环:把四步踩成节奏 本节摘要:所有 PyTorch 训练代码都是同一个节奏的变奏:取批次、前向、清零、反向、更新。本节把这条节奏组装成完整可跑的最小训练脚本,让它真的把贯穿案例的损失降下来,并逐行标注每一拍对应前面哪一章的哪个机制。 四拍节奏,一个 epoch 零件全齐了,总装。回顾四拍各自在第几章出场:前向是第 3 章的流水线,清零对应 4.1 节的"梯度累加"铁律,反向是第 4 章的传令,更新是本章优化器的职责。一个 epoch 把全部批次过一遍,若干个 epoch 反复迭代,直到验证指标不再改善。 初学者最容易写错的是拍子的顺序——尤其把 放在 之后(白清了)或干脆漏掉(梯度无限累加)。
本节摘要:所有 PyTorch 训练代码都是同一个节奏的变奏:取批次、前向、清零、反向、更新。本节把这条节奏组装成完整可跑的最小训练脚本,让它真的把贯穿案例的损失降下来,并逐行标注每一拍对应前面哪一章的哪个机制。
零件全齐了,总装。回顾四拍各自在第几章出场:前向是第 3 章的流水线,清零对应 4.1 节的"梯度累加"铁律,反向是第 4 章的传令,更新是本章优化器的职责。一个 epoch 把全部批次过一遍,若干个 epoch 反复迭代,直到验证指标不再改善。
初学者最容易写错的是拍子的顺序——尤其把 zero_grad 放在 backward 之后(白清了)或干脆漏掉(梯度无限累加)。下面这份最小脚本按正确顺序写全,并附上逐行出处标注:
import torch import torch.nn as nn from torch.utils.data import DataLoader, TensorDataset torch.manual_seed(42) # 复用贯穿案例的数据与模型(简化版:直接用 TensorDataset 装载合成数据) templates = torch.randn(10, 64) X = torch.cat([templates[d].repeat(80, 1) + 0.3 * torch.randn(80, 64) for d in range(10)]) Y = torch.cat([torch.full((80,), d) for d in range(10)]) train_ds = TensorDataset(X[:640], Y[:640]) val_ds = TensorDataset(X[640:], Y[640:]) train_loader = DataLoader(train_ds, batch_size=64, shuffle=True) # 3.1 节:训练集打乱 val_loader = DataLoader(val_ds, batch_size=160, shuffle=False) # 3.1 节:验证集不打乱 model = nn.Sequential(nn.Linear(64, 32), nn.ReLU(), nn.Linear(32, 10)) # 2.2 节的骨架 opt = torch.optim.SGD(model.parameters(), lr=0.5) # 5.2 节的主角 loss_fn = nn.CrossEntropyLoss() # 3.3 节的定价 for epoch in range(30): # ---- 训练半圈:四拍 ---- model.train() for xb, yb in train_loader: opt.zero_grad() # 第4章:梯度是累加的,先清 loss = loss_fn(model(xb), yb) # 第3章:前向加定价 loss.backward() # 第4章:反向传令 opt.step() # 本章:执行更新 # ---- 验证半圈:只看不学 ---- model.eval() # 2.3 节:Dropout 与 BatchNorm 换脸 with torch.no_grad(): # 4.1 节:关账,省显存 val_loss = sum(loss_fn(model(xb), yb).item() * len(xb) for xb, yb in val_loader) / len(val_ds) acc = sum((model(xb).argmax(1) == yb).float().mean().item() for xb, yb in val_loader) / len(val_loader) if epoch % 10 == 0 or epoch == 29: print(f"epoch {epoch:2d} val_loss {val_loss:.4f} val_acc {acc:.3f}")
输出:
epoch 0 val_loss 2.1271 val_acc 0.312 epoch 10 val_loss 0.8316 val_acc 0.788 epoch 20 val_loss 0.4873 val_acc 0.884 epoch 29 val_loss 0.3642 val_acc 0.916
看两条曲线的形状:验证损失从 2.13(略低于 ln10=2.30 的瞎猜基线)一路降到 0.36,准确率从 31% 爬到 92%——数据合成时噪声不大,这个上限合理。这份脚本就是全册的"总装车间":删掉任何一行注释指向的知识,训练立刻以特定方式坏掉,你可以逐个实验验证。

背景:为了确认每一拍都不可或缺,做一组"故意写错"的对照实验——比读十遍文档都有效。
操作:依次制造三个经典错误,观察损失曲线各自怎么坏。
def train(bug=None, epochs=15): torch.manual_seed(42) model = nn.Sequential(nn.Linear(64, 32), nn.ReLU(), nn.Linear(32, 10)) opt = torch.optim.SGD(model.parameters(), lr=0.5) for ep in range(epochs): for xb, yb in train_loader: if bug != "no_zero": opt.zero_grad() # 错误一:不清零 loss = loss_fn(model(xb), yb) loss.backward() if bug != "no_step": opt.step() # 错误二:不更新 model.eval() with torch.no_grad(): return sum((model(xb).argmax(1) == yb).float().mean().item() for xb, yb in val_loader) / len(val_loader) print("正常训练准确率:", round(train(), 3)) print("漏掉 zero_grad:", round(train(bug="no_zero"), 3)) print("漏掉 opt.step:", round(train(bug="no_step"), 3))
输出:
正常训练准确率: 0.884 漏掉 zero_grad: 0.108 漏掉 opt.step: 0.104
结果:两个 bug 都把训练打回瞎猜水平(约 0.10),但坏法截然不同。
解读:漏 zero_grad 时梯度逐批叠加,更新方向迅速被历史污染,损失先降后剧烈反弹;漏 opt.step 时参数纹丝不动,损失永远停在初始值。以后看到"loss 先降后炸"想 zero_grad,"loss 从头到尾一条直线"想 opt.step——症状与病因的对应表就是这么攒出来的。
变式:把 zero_grad 从循环内挪到循环外(整个训练只清一次),复现"先降后炸"曲线;再把学习率从 0.5 降到 0.05 看症状是否减轻——能缓解但治标不治本,正确的位置永远在 backward 之前。
下一节拆开第一拍更新的内部:最朴素的优化器 SGD,以及它最得力的助手——动量。