PyTorch 是当前深度学习和大模型研究的主流框架。它的核心是一套固定套路:用张量表示数据、用 nn.Module 搭网络、用循环完成”前向-反向-更新”。这篇教程用一个手写数字分类任务,带你把这套流程完整跑通。

训练的核心循环

无论多复杂的模型,训练都是这个循环在重复:

前置准备

pip install torch torchvision

有 GPU 的话,安装对应 CUDA 版本的 PyTorch 会快很多。

步骤一:准备数据

用 torchvision 加载 MNIST 手写数字数据集,并用 DataLoader 分批:

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

transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,)),
])
train_set = datasets.MNIST(root="./data", train=True, download=True, transform=transform)
test_set = datasets.MNIST(root="./data", train=False, download=True, transform=transform)

train_loader = DataLoader(train_set, batch_size=64, shuffle=True)
test_loader = DataLoader(test_set, batch_size=1000)

步骤二:定义网络

继承 nn.Module,在 __init__ 里搭层,在 forward 里定义数据流:

import torch.nn as nn
import torch.nn.functional as F

class Net(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(28 * 28, 128)
self.fc2 = nn.Linear(128, 64)
self.fc3 = nn.Linear(64, 10) # 10 个数字类别

def forward(self, x):
x = x.view(-1, 28 * 28) # 把 28x28 图像展平
x = F.relu(self.fc1(x))
x = F.relu(self.fc2(x))
return self.fc3(x)

model = Net()

步骤三:定义损失函数和优化器

import torch.optim as optim

criterion = nn.CrossEntropyLoss() # 分类任务用交叉熵
optimizer = optim.Adam(model.parameters(), lr=1e-3)

步骤四:编写训练循环

这是最关键的一步,几个动作缺一不可:

def train(epoch):
model.train()
for images, labels in train_loader:
optimizer.zero_grad() # 1. 清空上一轮梯度
outputs = model(images) # 2. 前向传播
loss = criterion(outputs, labels) # 3. 计算损失
loss.backward() # 4. 反向传播
optimizer.step() # 5. 更新参数
print(f"Epoch {epoch} 完成,loss={loss.item():.4f}")

for epoch in range(3):
train(epoch)

zero_grad 一定要记得,否则梯度会不断累加导致训练出错。

步骤五:评估模型

评估时关闭梯度计算,既省内存又加速:

def evaluate():
model.eval()
correct = 0
with torch.no_grad():
for images, labels in test_loader:
preds = model(images).argmax(dim=1)
correct += (preds == labels).sum().item()
print(f"测试集准确率:{correct / len(test_set) * 100:.2f}%")

evaluate()

步骤六:保存与加载模型

torch.save(model.state_dict(), "mnist_net.pth")   # 保存参数

loaded = Net() # 之后加载
# weights_only=True 只反序列化权重,避免加载不受信任文件时执行恶意代码
loaded.load_state_dict(torch.load("mnist_net.pth", weights_only=True))
loaded.eval()

常见问题

  • loss 不下降:检查学习率、数据是否归一化、标签是否对应
  • 显存溢出(CUDA out of memory):减小 batch_size
  • 训练准确率高、测试低:过拟合,加 Dropout 或正则
  • 结果不可复现:设置 torch.manual_seed(0)

小结

PyTorch 训练神经网络的骨架非常固定:DataLoader 供数据、nn.Module 定网络、损失加优化器定目标、循环里”清梯度→前向→算损失→反向→更新”。把这个套路刻进肌肉记忆,之后无论 CNN、RNN 还是 Transformer,都只是替换网络结构而已。