PyTorch 神经网络实战
理论看十遍不如训练一遍。本教程用 MNIST 手写数字数据集,训练一个卷积神经网络(CNN)。
一、加载数据
from torchvision import datasets, transforms
train = datasets.MNIST(root='./data', train=True, transform=transforms.ToTensor(), download=True)
loader = torch.utils.data.DataLoader(train, batch_size=64, shuffle=True)
二、定义网络
卷积层提取特征,池化降维,全连接层做分类:
import torch.nn as nn
class CNN(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(
nn.Conv2d(1, 32, 3), nn.ReLU(), nn.MaxPool2d(2),
nn.Conv2d(32, 64, 3), nn.ReLU(), nn.MaxPool2d(2),
nn.Flatten(), nn.Linear(64 * 5 * 5, 10)
)
def forward(self, x):
return self.net(x)
三、训练循环
model = CNN()
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
loss_fn = nn.CrossEntropyLoss()
for imgs, labels in loader:
preds = model(imgs)
loss = loss_fn(preds, labels)
opt.zero_grad(); loss.backward(); opt.step()
四、训练技巧
- 用
model.train()/model.eval()切换模式 - 记录 loss 曲线,欠拟合加层、过拟合加 dropout 或数据增强
- 显存不足先降 batch_size
跑完 2-3 个 epoch 准确率就能超过 97%。这些代码稍加改造,就是后面训练 Transformer 的骨架。