PyTorch 是当前深度学习和大模型研究的主流框架。它的核心是一套固定套路:用张量表示数据、用 nn.Module 搭网络、用循环完成”前向-反向-更新”。这篇教程用一个手写数字分类任务,带你把这套流程完整跑通。
训练的核心循环
无论多复杂的模型,训练都是这个循环在重复:
graph LR
A[前向传播 算预测] --> B[计算损失]
B --> C[反向传播 算梯度]
C --> D[优化器更新参数]
D --> E{完成所有轮次}
E -->|否| A
E -->|是| F[训练结束]
前置准备
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)
def forward(self, x): x = x.view(-1, 28 * 28) 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() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() 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()
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,都只是替换网络结构而已。