Skip to content

最小可运行示例:20 行学会完整训练流程

目标:用不到 20 行核心代码,让模型从数据里"学出"一个我们故意藏起来的规律 y = 2x + 1。跑通它,PyTorch 训练的完整骨架就都在你脑子里了——实战项目只是把这个骨架放大。

问题设定

我们生成一批带噪声的样本点,模型不知道真实公式,只能通过数据逼近它:

python
import torch
import torch.nn as nn

torch.manual_seed(42)                                    # 固定随机数,结果可复现

x = torch.rand(256, 1) * 10                              # 256 个 0~10 之间的输入
y = x * 2 + 1 + torch.randn(256, 1) * 0.5                # y = 2x + 1 + 噪声

model = nn.Linear(1, 1)                                  # 最简模型:y = wx + b
loss_fn = nn.MSELoss()                                   # 回归任务用均方误差
opt = torch.optim.SGD(model.parameters(), lr=0.05)       # 随机梯度下降

for epoch in range(200):                                 # 训练 200 轮
    pred = model(x)                                      # ① 前向传播
    loss = loss_fn(pred, y)                              # ② 计算损失
    opt.zero_grad()                                      # ③ 梯度清零
    loss.backward()                                      # ④ 反向传播求梯度
    opt.step()                                           # ⑤ 更新参数
    if (epoch + 1) % 50 == 0:
        print(f"epoch {epoch+1:3d}  loss={loss.item():.4f}  "
              f"w={model.weight.item():.3f}  b={model.bias.item():.3f}")

保存为 hello.py,在激活了 PyTorch 环境的终端执行:

powershell
python hello.py

预期输出

text
epoch  50  loss=0.4321  w=1.703  b=1.382
epoch 100  loss=0.2765  w=1.845  b=1.192
epoch 150  loss=0.2387  w=1.912  b=1.103
epoch 200  loss=0.2298  w=1.945  b=1.061

(具体数字因硬件/版本略有差异,趋势一致即可。)

看懂输出:loss 一路下降,w 从随机值逐渐逼近 2,b 逼近 1 —— 模型凭数据"发现"了隐藏公式。loss 最终停在 0.23 左右而不是 0,因为数据本身带噪声(方差 0.5²=0.25),这是正常的、也是正确的——模型不该去拟合噪声

逐行拆解:五步训练循环

①~⑤ 这五行就是所有 PyTorch 训练代码的心脏,无论模型多大、任务多复杂,循环体内永远是这五步:

步骤代码在做什么
① 前向传播pred = model(x)数据流过网络,得到预测
② 计算损失loss = loss_fn(pred, y)量化预测与真实值的差距
③ 梯度清零opt.zero_grad()清掉上一轮遗留的梯度(PyTorch 默认累加!)
④ 反向传播loss.backward()Autograd 沿计算图求出每个参数的梯度
⑤ 参数更新opt.step()优化器按 w -= lr * grad 的思路微调参数

顺序提示zero_grad()backward() 之前还是 step() 之后都可以,只要在下一轮求梯度前清零即可。社区常见两种写法,本教程统一用"先清零"风格。

你可能会遇到

  • loss 变成 NaN:多半是学习率太大,梯度爆炸。把 lr=0.05 改成 0.005 再试。
  • loss 纹丝不动:检查学习率是否太小,或 opt 是否传了 model.parameters()
  • 输出 w≈1.9 不是精确的 2:正常。有限数据 + 噪声 + 训练轮数有限,逼近即可,不必苛求相等。
  • Windows 下多进程相关报错:本示例没有用到,实战章节 num_workers > 0 时再注意。

小结

20 行代码覆盖了:数据构造、模型定义、损失函数、优化器、五步训练循环。这份骨架会原样出现在下一章实战项目里——只不过数据换成真正的图片集、模型换成多层网络、循环加上验证环节。

下一步 → 实战项目:手写数字识别