Skip to content

实战第 3 步:训练循环(并故意踩第 2 个坑)

数据通了、模型通了,现在把最小示例里的五步循环放大成真正的工程循环:多 epoch、批次迭代、设备迁移、每轮验证。然后我们故意删掉一行代码,看看"训练不收敛"这种最折磨新手的症状长什么样。

完整的 train.py

替换 train.py 的主逻辑(数据部分不变):

python
# train.py —— 完整训练脚本
import time
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from model import DigitMLP

BATCH_SIZE = 64
EPOCHS = 5
LR = 1e-3

transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.1307,), (0.3081,)),
])


def evaluate(model, loader, loss_fn, device):
    """在验证/测试数据上计算平均损失与准确率(不更新参数)"""
    model.eval()                          # 切评估模式:关 Dropout
    total_loss, correct, seen = 0.0, 0, 0
    with torch.no_grad():                 # 评估不需要建计算图
        for images, labels in loader:
            images, labels = images.to(device), labels.to(device)
            logits = model(images)
            loss = loss_fn(logits, labels)
            total_loss += loss.item() * labels.size(0)
            correct += (logits.argmax(dim=1) == labels).sum().item()
            seen += labels.size(0)
    return total_loss / seen, correct / seen


def main():
    device = torch.device("cuda" if torch.cuda.is_available()
                          else "mps" if torch.backends.mps.is_available()
                          else "cpu")
    print(f"设备: {device}")

    train_set = datasets.MNIST("./data", train=True,  download=True, transform=transform)
    test_set  = datasets.MNIST("./data", train=False, download=True, transform=transform)
    train_loader = DataLoader(train_set, batch_size=BATCH_SIZE, shuffle=True)
    test_loader  = DataLoader(test_set,  batch_size=256)

    model = DigitMLP().to(device)         # 模型整体搬到 GPU/MPS
    loss_fn = nn.CrossEntropyLoss()
    opt = torch.optim.AdamW(model.parameters(), lr=LR)
    sched = torch.optim.lr_scheduler.StepLR(opt, step_size=3, gamma=0.5)  # 第3轮起学习率减半

    best_acc = 0.0
    for epoch in range(1, EPOCHS + 1):
        model.train()                     # 每个 epoch 开始切回训练模式
        t0, total_loss, n = time.time(), 0.0, 0
        for images, labels in train_loader:
            images, labels = images.to(device), labels.to(device)  # 数据与模型同设备
            logits = model(images)        # ① 前向
            loss = loss_fn(logits, labels)   # ② 损失
            opt.zero_grad()               # ③ 梯度清零(下一节要"案发现场",先别删它)
            loss.backward()               # ④ 反向求梯度
            opt.step()                    # ⑤ 更新参数
            total_loss += loss.item() * labels.size(0)
            n += labels.size(0)
        sched.step()                      # 学习率调度按 epoch 步进

        val_loss, val_acc = evaluate(model, test_loader, loss_fn, device)
        print(f"epoch {epoch}: train_loss={total_loss/n:.4f}  "
              f"val_loss={val_loss:.4f}  val_acc={val_acc*100:.2f}%  "
              f"({time.time()-t0:.1f}s)")

        if val_acc > best_acc:            # 只保存历史最优权重
            best_acc = val_acc
            torch.save(model.state_dict(), "weights/best.pt")
            print(f"  ↳ 保存最优模型 acc={best_acc*100:.2f}%")


if __name__ == "__main__":
    main()

终端运行(CPU 上约 1~2 分钟/epoch,GPU 上几秒):

powershell
python train.py

预期输出(数字会有波动):

text
设备: cpu
epoch 1: train_loss=0.3072  val_loss=0.1701  val_acc=95.02%  (38.2s)
  ↳ 保存最优模型 acc=95.02%
epoch 2: train_loss=0.1421  val_loss=0.1265  val_acc=96.28%  (37.6s)
  ↳ 保存最优模型 acc=96.28%
epoch 3: train_loss=0.1010  val_loss=0.1061  val_acc=96.86%  (37.9s)
  ↳ 保存最优模型 acc=96.86%
epoch 4: train_loss=0.0816  val_loss=0.1033  val_acc=97.01%  (38.1s)
  ↳ 保存最优模型 acc=97.01%
epoch 5: train_loss=0.0695  val_loss=0.1010  val_acc=97.08%  (38.4s)
  ↳ 保存最优模型 acc=97.08%

准确率 ≥ 97%,达标。在正式庆祝前——

踩坑时刻:删掉 opt.zero_grad() 会怎样?

这是真实世界发生率最高的"训练玄学问题"根源。把第 ③ 行注释掉:

python
            # opt.zero_grad()   # 手滑删掉了这行……

重新 python train.py(终端)。观察现象:

text
epoch 1: train_loss=0.6203  val_loss=0.5800  val_acc=82.15%
epoch 2: train_loss=0.5900  val_loss=0.6000  val_acc=80.30%
epoch 3: train_loss=0.6100  val_loss=0.5700  val_acc=83.50%
...(loss 在 0.6 附近反复横跳,准确率卡在 80% 上不去)

程序不报错、loss 不降、acc 卡住——比报错更可怕的是"看似在训练"。这就是心智模型篇预言过的:PyTorch 的梯度默认累加。每一轮的 .grad 里混着前面所有轮次的梯度,越积越大,更新方向被历史梯度污染,参数在原地打摆。

正确的调试思路(比答案重要)

遇到"不收敛"不要无脑换模型,按序排查:

  1. 过拟合一个 batch 试试:取 1 个 batch 反复训练几十步,loss 应该趋近 0。连这个都做不到,说明训练管线(而非模型容量)有问题。
  2. 打印梯度for p in model.parameters(): print(p.grad.norm())。梯度范数持续暴涨 → 累加没清 or 学习率过大;梯度恒为 0 → 参数没注册进优化器 or 图断了(比如 detach 用错)。
  3. 回看五步循环是否齐全zero_grad → backward → step 一个都不能少。

本例属于第 2 条的"梯度暴涨",恢复 opt.zero_grad() 立刻满血复活。

版本提示:PyTorch 2.x 推荐写 opt.zero_grad(set_to_none=True),比置零更快、更省显存(直接丢弃 grad 张量而非填 0),新代码可无脑使用。

训练完成,手里已经有 weights/best.pt 了。最后一步 → 第 4 步:评估与导出