03-训练核心 核心 预计 40 分钟 ★ 最小路径

实战:手写 MLP 训练循环

本知识点把模块 01~03 的全部零件(张量、autograd、MLP、交叉熵、优化器、正则化、数据管道)组装成一个 60 行左右、可直接运行的完整 PyTorch 训练循环,在 MNIST 上达到 98% 测试准确率。

一句话定义

本知识点把模块 01~03 的全部零件(张量、autograd、MLP、交叉熵、优化器、正则化、数据管道)组装成一个 60 行左右、可直接运行的完整 PyTorch 训练循环,在 MNIST 上达到 98% 测试准确率。

直观类比

前面的知识点是零件与图纸,本节是第一次整车下线:装好、点火、看仪表(损失曲线)、调螺丝(超参),跑通之后再谈改装(kp-021 换 CNN 引擎)。

为什么重要

「能看懂」与「能写出来」之间隔着本知识点。训练循环是深度学习工程的原子单位:kp-021 换成 CNN、kp-022 换成 RNN,骨架都不变——变的只是模型与数据两行。此后 kp-027 的调试方法论、kp-028 的实验管理都以「有一个能跑的基线循环」为前提。

前置知识

  • kp-004(PyTorch 张量与 autograd)、kp-007(训练在更新什么)、kp-009/kp-010(损失与优化器)、kp-013/kp-015(正则化与数据管道)。
  • 环境:CPU 即可,MNIST 每轮 1~2 分钟。

核心概念

  • 五件套节奏:zero_grad → forward → loss → backward → step,每个迭代固定五步。
  • 模型定义:nn.Sequential 或继承 nn.Module 两种写法,参数自动注册。
  • 双模式循环:外层 epoch、内层 batch;训练循环与评估循环分离,评估用 no_grad。
  • 指标(metrics):损失之外记录准确率,训练/验证双线记录(kp-013 双曲线)。
  • 检查点(checkpoint):保存最优验证模型的状态字典,早停恢复的载体。

原理与机制

循环把 kp-001 机制链逐行翻译:DataLoader 吐出 [B,784] 批 → 模型前向(kp-005 层变换)→ CrossEntropyLoss(kp-009,注意直接吃 logits)→ backward(kp-007/003 反向传播)→ optimizer.step(kp-010 更新)→ zero_grad(kp-003 的累加器清零)。评估循环的三条纪律:model.eval() 切换 BN/Dropout 行为(kp-014)、torch.no_grad() 关闭图记录(kp-004)、指标按样本数加权平均而非按批平均(尾批样本少会扭曲均值)。超参数据 kp-011 的扫描结论取一组安全默认:Adam lr=1e−3、batch 128、3~5 epoch 即到 98%。保存最优模型的判据是验证准确率而非训练损失(kp-013 纪律)。跑通后的标准动作是把本轮配置与结果写进实验记录(kp-028),然后做一组单变量实验(如换 SGD)体验「控制变量」的工作流。

公式或模型

本节不适用——理由:本节是既有公式(kp-005/009/010 各式)的代码化整合,所有数学已在对应知识点给出,重复罗列不产生新信息。

图示

for epoch:
  ┌─ 训练:model.train()
  │   for x,y in train_loader:
  │      zero_grad → ŷ=f(x) → loss=CE(ŷ,y) → loss.backward() → opt.step()
  ├─ 评估:model.eval() + no_grad()
  │   for x,y in val_loader: 统计 val_loss / val_acc(按样本加权)
  └─ 若 val_acc 创新高 → 保存 checkpoint

实例或案例

完整可运行代码(PyTorch 2.x,MNIST 自动下载):

import torch, torch.nn as nn
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

tf = transforms.Compose([transforms.ToTensor(),
    transforms.Normalize((0.1307,), (0.3081,))])
tr = datasets.MNIST('./data', train=True,  download=True, transform=tf)
va = datasets.MNIST('./data', train=False, download=True, transform=tf)
tr_loader = DataLoader(tr, batch_size=128, shuffle=True)
va_loader = DataLoader(va, batch_size=512)

model = nn.Sequential(
    nn.Flatten(),
    nn.Linear(784, 512), nn.ReLU(), nn.Dropout(0.2),
    nn.Linear(512, 256), nn.ReLU(), nn.Dropout(0.2),
    nn.Linear(256, 10))
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
lossf = nn.CrossEntropyLoss()

best = 0.0
for epoch in range(5):
    model.train()                                   # ① 训练模式
    for x, y in tr_loader:
        opt.zero_grad()                             # ② 清梯度
        loss = lossf(model(x), y)                   # ③ 前向+损失
        loss.backward()                             # ④ 反向
        opt.step()                                  # ⑤ 更新
    model.eval()                                    # 评估模式
    correct = total = vloss = 0
    with torch.no_grad():
        for x, y in va_loader:
            logits = model(x)
            vloss += lossf(logits, y).item() * y.size(0)
            correct += (logits.argmax(1) == y).sum().item()
            total += y.size(0)
    acc = correct / total
    print(f'epoch {epoch}: val_loss={vloss/total:.4f} val_acc={acc:.4f}')
    if acc > best:
        best = acc
        torch.save(model.state_dict(), 'best.pt')   # 保存最优
print('best val acc:', best)                        # 预期 ≈ 0.98

排错演练建议:故意注释掉 zero_grad 观察损失震荡;故意漏掉 model.eval() 观察 Dropout 污染评估——两个「疫苗实验」能把 kp-003 与 kp-014 的坑一次性打进肌肉记忆。

常见误区

  • 损失按批平均再求总平均:尾批样本少导致权重失真,应按「损失×批样本数」累积。
  • 评估循环不关 no_grad:显存持续增长且速度慢数倍(kp-004 误区在项目里的显影)。
  • 只记录训练指标:没有验证双线就无从判断过拟合(kp-013 的心电图没有第二根线)。
  • 随机种子里外不全:DataLoader 的 shuffle 也有自己的生成器,完整复现需 seed_worker(kp-028 展开完整清单)。

与其他知识点的关系

本节是模块 01~03 的总装线与模块 04/05 的模板:kp-021 把 model 一行换成 CNN、增强管道换成 kp-015 的版本;kp-027 的调试方法论以本循环的损失曲线为观察对象;kp-028 的记录动作直接挂在本循环的 epoch 边界。

自测题

  1. 写出训练五件套的顺序,并解释哪一步对应 kp-003 的哪个机制。

答案要点:zero_grad→forward→loss→backward→step;backward 即自动微分反向遍历计算图,step 消费 .grad。

  1. 为什么评估循环必须 model.eval() + no_grad() 双开关?

答案要点:eval 切换 BN/Dropout 推理行为;no_grad 关闭图记录省显存提速,两者职责不同缺一不可。

  1. 验证准确率的正确平均方式是什么?

答案要点:按样本数加权(Σ批正确数/Σ批样本数),不能对各批准确率简单平均。

延伸阅读

  • PyTorch 官方教程 «Optimization» 与 «Training a Classifier»(与本节代码互为参照)。
  • 李沐《动手学深度学习》4.5 节「模型选择、欠拟合和过拟合」到 5 章(softmax 回归简洁版)。

学习状态

状态保存在浏览器本地,用于首页与路径页的进度统计。