一个 PyTorch 模型训练的完整流程

学 PyTorch 时,单个概念看懂不难,真正容易乱的是完整训练流程。

这篇文章先不追求复杂模型,只把一条最基础的训练主线串起来。

第一步:准备数据

训练模型前,先要把数据整理成模型能吃的形式。

通常会经历几步:

  • 读取原始数据

  • 做必要的清洗和预处理

  • 转成 Tensor

  • 封装成 Dataset

  • 用 DataLoader 批量加载

如果数据这一步没处理好,后面模型再复杂也很难救回来。

第二步:定义模型

PyTorch 里通常会继承 nn.Module 定义模型:

复制代码
from torch import nn
​
class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.linear = nn.Linear(10, 2)
​
    def forward(self, x):
        return self.linear(x)

这里 __init__ 定义模型有哪些层,forward 定义数据怎么流过这些层。

第三步:选择损失函数

损失函数负责衡量模型预测错了多少。

分类任务常见:

复制代码
loss_fn = nn.CrossEntropyLoss()

回归任务常见:

复制代码
loss_fn = nn.MSELoss()

损失函数要和任务类型匹配,这一点很重要。

第四步:选择优化器

优化器负责根据梯度更新参数。

常见写法:

复制代码
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)

这里 model.parameters() 告诉优化器要更新哪些参数,lr 是学习率。

第五步:训练循环

最核心的训练循环通常长这样:

复制代码
for x, y in train_loader:
    pred = model(x)
    loss = loss_fn(pred, y)
​
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

这几行非常重要。

可以按顺序理解:

  1. 前向传播,得到预测

  2. 计算 loss

  3. 清空旧梯度

  4. 反向传播,计算新梯度

  5. 优化器更新参数

这就是 PyTorch 训练模型的核心骨架。

第六步:验证模型

训练时还需要在验证集上观察效果。

验证阶段通常不需要计算梯度,所以会写:

复制代码
model.eval()
with torch.no_grad():
    for x, y in val_loader:
        pred = model(x)

这样可以减少显存占用,也避免误更新模型。

小结

一个 PyTorch 训练流程可以压缩成这样:

复制代码
数据 -> 模型 -> loss -> backward -> optimizer.step -> 验证

刚开始不要急着堆复杂结构。

先把这条主线真正跑通,后面再换模型、调参数、加可视化,都会轻松很多。

技术图:把关键链路画清楚

可运行实验:跑通一个最小训练与验证闭环

训练循环真正需要关注的是数据流和状态变化:训练阶段计算梯度,验证阶段关闭梯度并只统计指标。

复制代码
import torch
from torch import nn
​
torch.manual_seed(0)
x = torch.arange(0, 6, dtype=torch.float32).unsqueeze(1)
y = 2 * x + 1
model = nn.Linear(1, 1)
optimizer = torch.optim.SGD(model.parameters(), lr=0.05)
loss_fn = nn.MSELoss()
for epoch in range(101):
    optimizer.zero_grad()
    loss = loss_fn(model(x), y)
    loss.backward()
    optimizer.step()
    if epoch in (0, 50, 100): print(f"epoch={epoch} loss={loss.item():.6f}")
with torch.no_grad(): print(f"x=7 prediction={model(torch.tensor([[7.0]])).item():.3f}")

运行结果:

复制代码
epoch=0 loss=41.809498
epoch=50 loss=0.000142
epoch=100 loss=0.000007
x=7 prediction=14.996

Loss 持续下降,模型最终接近真实关系 y=2x+1,所以输入 7 时预测接近 15。验证与推理用 torch.no_grad() 避免无意义的计算图。

常见误区

  1. 验证时只写 model.eval() 就会关闭梯度。二者职责不同,通常还要配合 torch.no_grad()

  2. 只保存模型对象最方便。更稳妥的做法是保存 state_dict、配置和预处理信息。

动手练习

增加一个验证集,每 10 个 epoch 记录训练与验证 loss,并画出两条曲线。


本文首发于「去你想去的地方」: 一个 PyTorch 模型训练的完整流程 | 去你想去的地方

完整学习路线、视频版和后续更新请访问原文。

相关推荐
小宋10216 分钟前
大模型成本怎么控制:Token统计、缓存与模型路由实战
人工智能·缓存
阿里云大数据AI技术7 分钟前
DataWorks Data Agent 实战课堂(七):数据治理 Agent——AI 驱动的自动化治理
人工智能·agent
天一生水water7 分钟前
基于重构的无监督/单类时间序列异常检测
人工智能·深度学习·重构·transformer
知识分享小能手17 分钟前
深度学习学习教程,从入门到精通,数值计算 — 知识点详解(4)
人工智能·深度学习·学习
三声三视17 分钟前
审计清单函数名写成 def?tri-checklist 的 diff 解析在 Python 改名场景下悄悄翻车
人工智能·ai·skillhub·tri-checklist·tri-skills
嘟嘟嘟952743 分钟前
Kubernetes 引入 KYAML:更安全的 YAML 子集
人工智能·架构·开源
智购科技自动售货机厂家44 分钟前
2026自动售货机设备清洁效果自动验证:从图像比对到评分算法的工程实践~YH
人工智能·算法·计算机视觉
财迅通Ai1 小时前
光智科技全景透视:一家被“光学元件”标签遮蔽的稀散金属材料平台
大数据·人工智能·科技·光智科技
小刘在重生~1 小时前
Java常用类|String类详解 + BigDecimal精准计算(含面试题)
java·开发语言·python
AI技术新视界1 小时前
失控的认知外包:大语言模型如何像病毒般入侵人类思维与社会生态
人工智能·llm·认知科学