Skip to content

进阶与最佳实践

实战项目让你"能跑",这一篇让你"跑得快、跑得稳、跑得可信"。内容分三块:性能、工程规范、常见反模式。

性能:从 CPU 分钟级到 GPU 秒级

1. 混合精度训练(AMP)—— 免费的午餐

GPU(图灵架构及以上)用 float16/bfloat16 做矩阵乘,速度约 2 倍、显存约省一半,精度几乎无损:

python
from torch.amp import autocast, GradScaler

scaler = GradScaler()                      # float16 需要损失缩放防下溢
for images, labels in train_loader:
    opt.zero_grad(set_to_none=True)
    with autocast("cuda", dtype=torch.float16):   # 这段内的运算自动降精度
        loss = loss_fn(model(images), labels)
    scaler.scale(loss).backward()          # 缩放后的 loss 反传
    scaler.step(opt)
    scaler.update()

A100/H100 等支持 bf16 的卡可直接 dtype=torch.bfloat16 且不需要 GradScaler。注意:AMP 只在 CUDA 上有意义,CPU/MPS 训练不需要。

2. torch.compile —— 一行代码换加速

PyTorch 2.x 招牌功能,把模型编译成优化过的内核(2.7+ 一行 torch.compile 更省心):

python
model = torch.compile(DigitMLP().to(device))   # 首次运行有编译开销,之后提速

小模型收益有限(可能反而慢在编译上),中等以上模型常见 10~40% 训练提速。报错频繁的环境(部分自定义算子)可以先不启用。

3. 数据管道别拖 GPU 后腿

nvidia-smi(终端)显示 GPU 利用率忽高忽低 → 大概率数据加载是瓶颈:

  • DataLoader(..., num_workers=4, pin_memory=True):多进程预取 + 锁页内存加速 CPU→GPU 拷贝
  • 把昂贵的预处理挪出 __getitem__,或改用 GPU 端变换(kornia 等库)

4. 显存不够的自救顺序

  1. 减小 batch_size(最直接)
  2. 梯度累积:batch 开小,但攒 4 个 micro-batch 的梯度再 step(),等效大 batch:
python
# 累积 4 个 micro-batch 的梯度再更新一次(等效 batch_size×4)
for i, (images, labels) in enumerate(train_loader):
    loss = loss_fn(model(images.to(device)), labels.to(device)) / 4  # 除以累积步数
    loss.backward()
    if (i + 1) % 4 == 0:
        opt.step()
        opt.zero_grad(set_to_none=True)
  1. 开启 AMP;4. 冻结暂时不训的层(requires_grad_(False));5. 用 torch.no_grad() 包住一切评估代码(忘记它是最常见的显存泄漏源)。

工程规范:demo 与生产项目的差异

维度Demo 做法生产做法
随机性不管种子torch.manual_seed + DataLoader worker_init_fn,结果可复现
实验跟踪printTensorBoard / Weights & Biases 记录 loss、acc、超参、GPU 占用
配置写死在代码里argparse / YAML 管理超参,一次实验一个配置快照
Checkpoint只存最终模型定期存档(每 N epoch),进程崩溃可续训
数据划分训练时反复看测试集train/val/test 严格分离,测试集只在验收时用一次
环境"我电脑上能跑"requirements.txt 锁版本 + 容器化(镜像内 CUDA 版本一致)
监控训完不管上线后监控推理延迟、数据漂移、准确率衰减

续训 checkpoint 的标准姿势(保存与加载要对称):

python
# 保存
torch.save({"epoch": epoch, "model": model.state_dict(),
            "opt": opt.state_dict(), "sched": sched.state_dict()}, "ckpt.pt")

# 加载并继续
# ⚠️ weights_only=False 会执行 pickle 反序列化,仅限加载【自己生成的】checkpoint 文件!
# 来路不明的模型文件一律不要加载(等价于运行陌生人的代码)
ck = torch.load("ckpt.pt", map_location=device, weights_only=False)
model.load_state_dict(ck["model"]); opt.load_state_dict(ck["opt"])
sched.load_state_dict(ck["sched"]); start_epoch = ck["epoch"] + 1

常见反模式及替代方案

反模式 1:在训练循环里做 Python 级同步

python
acc = (pred == labels).float().mean().item()   # .item() 每次强制 GPU→CPU 同步

每 batch 一两次无所谓,在细粒度循环里狂用会显著拖慢训练。替代:累加张量,epoch 结束再 .item()

反模式 2:模型、数据、优化器三者的设备/时序关系搞错

model.to(device) 是原地搬参数并返回自身,写不写 model = ... 效果相同;真正容易出事的是顺序——先 model.to(device) 再构建优化器、再喂数据,且每个 batch 记得 images.to(device)。三者任一不同屏,就会撞上 device 类报错。

反模式 3:评估/推理时忘记 eval() + no_grad() 成对出现

只用 no_grad() 不切 eval():Dropout 继续随机丢,指标抖动。只切 eval() 不用 no_grad():白白占显存。两个都要。

反模式 4:学习率一调就 10 倍暴改

调参先做学习率扫描:固定其他超参,跑 [3e-3, 1e-3, 3e-4] 三档各 1 epoch 看 loss 下降速度,选最快档。比玄学加减高效得多。

反模式 5:train 指标很好,直接宣布胜利

永远同时看 train 和 val 两条曲线:train 好 val 差 = 过拟合(加数据/加 Dropout/早停);两边都差 = 欠拟合(模型太小/训练不足/学习率不当)。只看一个数字的决策都是瞎猜。

反模式 6:把 requires_grad=False 的 buffer 当参数改

running_mean 这类统计量是 buffer 不是 parameter,冻结模型、复制权重时注意 state_dict() 包含 buffer,而 parameters() 不包含。

小结

  • 性能三板斧:AMP、torch.compile、喂饱 GPU(workers + pin_memory);显存不够就梯度累积。
  • 生产的核心关键词是可复现可恢复:种子、配置快照、对称的 checkpoint。
  • 反模式大多源于"把训练循环当普通 Python 循环"——记住循环里每个操作背后都有设备同步、显存和计算图。

下一步:遇到具体报错时去排错与 FAQ速查。