跑完前面几篇,你手里应该已经有一个能正常运行的 PyTorch 环境,也亲手写过线性回归、逻辑回归和一个两层全连接网络,并且用它跑过 MNIST。我猜当时你的测试准确率大概在 92% 到 95% 之间,运气好调调超参数能摸到 96%。这时候你可能会冒出一个很自然的疑问:明明全连接网络也能分类,为什么大家说到图像任务就都开始讲 CNN?这篇就是用 CNN 重新解决同一个 MNIST 问题,目标是把准确率推到 99% 以上,同时把卷积、池化、维度变化这些概念彻底讲透。内容会从前几篇的全连接网络出发,逐步引入 CNN 的动机和原理,再落到 PyTorch 的完整实现上,包括数据下载遇到 404 的坑、LeNet-5 风格网络结构、训练循环、误差分析,以及几个提升精度和速度的实用方向。适合已经会用 PyTorch 搭简单网络的入门读者,也适合那些跑过 MNIST 但对其中的维度计算和训练细节还一知半解的人。
1. 换角度看 MNIST:全连接网络的三块天花板与 CNN 的解法
1.1 全连接网络为什么到了 95% 就涨不动了
MNIST 是 28×28 的单通道灰度图,展开后是 784 维向量。用全连接网络做分类,第一层就要把 784 维映射到隐藏层,如果隐藏层是 512,那这一层的参数量就是 784×512 + 512,大约 40 万。这个规模在 MNIST 上还撑得住,但稍微往真实场景一想就有问题了:如果输入换成 224×224 的彩色图,展开后是 224×224×3 = 150528 维,同样映射到 512 维,单层参数量直接跳到 7700 万。这只是第一层,网络稍微深一点,参数总量就是天文数字,训练时显存、内存、收敛速度全部受不了。
除了参数爆炸,全连接还有个更本质的问题:它把二维图像强行拉成一维向量,像素之间的二维空间关系,比如上下相邻、左右相邻、某个笔画延续了几行几列,这些信息在展开过程中被稀释了。模型看到的只是 784 个“独立”的特征,并不知道第 180 个像素和第 181 个像素在图像上是邻居。可对图像识别来说,相邻像素的相关性恰恰是最重要的先验知识。这就是为什么全连接网络在 MNIST 上很容易到 95% 附近就开始遇到瓶颈,网络再宽、层数再深,提升也非常有限。
还有一个容易被忽略的问题:平移不变性差。同一个手写数字,在图像里往左平移两个像素,在全连接网络的输入向量层面,所有特征的位置几乎都变了,模型会把它当成一个全新样本。而人类识别手写数字,根本不在乎这个数字出现在图像的哪个位置。CNN 后来能在这个任务上大杀四方,不是因为它用了什么玄学机制,而是它把这三点问题都从结构上解决掉了。
1.2 CNN 的三个核心机制,正好对着三个问题开药
CNN 的第一个机制是局部感受野。卷积核每次只看输入图像的一个小窗口,比如 3×3 或者 5×5,然后在这个局部区域里提取特征。这很符合图像本身的特性:一个像素的特征主要取决于它周围的像素,距离很远的像素之间基本没有直接关系。局部连接也直接解决了参数爆炸的问题,卷积核的参数量只跟核大小、输入通道数、输出通道数有关,跟输入图像的宽高没有关系。
第二个机制是权值共享。同一个卷积核会在整张图像上滑动,也就是说,不管图像是 28×28 还是 224×224,同一个卷积核处理所有位置时用的都是同一组权重。这有两个好处:一是参数数量大幅下降,二是模型学到的是一个“局部特征检测器”,比如某个卷积核专门检测横向边缘,那它在图像左上角和右下角都能检测到横向边缘,这就天然带来了一定的平移不变性。一个数字在图像左边还是右边,只要卷积核覆盖到,提取到的特征就是类似的。
第三个机制是池化,也叫下采样。MaxPooling 在 2×2 或 3×3 的窗口里取最大值,把特征图的尺寸减半。这样做一方面降低了后续层的计算量,另一方面把局部区域内最强的响应保留下来,对轻微的位移和形变更有容忍度。你可以把池化理解成“反正我只要知道这个区域里有一个边缘,至于它精确在哪个像素位置,不那么重要”。这三个机制组合在一起,CNN 在图像任务上就比全连接网络高出一个维度。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 数据是第一关:torchvision 下载 MNIST 报 404 的排查与离线方案
2.1 先顺着报错信息查一遍
很多人在这一步直接卡住,不是网络不会写,而是数据下载不下来。代码往往长这样:
python复制from torchvision import datasets, transforms
train_dataset = datasets.MNIST(
root='./data',
train=True,
transform=transforms.ToTensor(),
download=True
)
然后终端里开始下载,接着报出类似这样的错误:
code复制HTTP Error 404: Not Found
或者下载到一半就断了,再重试又变成 404。这个问题的根源通常是官方数据集服务器的 URL 偶发变动或者不稳定,torchvision 内置的下载地址指向了旧路径,服务器返回了 404。这是一个非常典型的“环境问题”,不是你的代码有问题。
排查的时候不要急着改代码,先做三件事。第一,看报错信息里打印出的 URL 到底是哪个,确认是不是 http://yann.lecun.com/exdb/mnist/ 这个官方源。第二,看本地目录,比如 ./data/MNIST/raw/ 下有没有上次下载残留的半截文件,这些残缺文件会让重试变得更加诡异,如果存在直接删掉。第三,确认你用的 torchvision 版本,不同版本对下载路径的处理有细微差别,但 404 这个错误基本都指向同一个原因:下载源变了或者暂时不可用。
2.2 手动准备 MNIST 文件的完整流程
遇到 404 别和它死磕,我用的办法是直接手动下载数据文件。MNIST 的原始数据由四个 .gz 压缩包组成:
train-images-idx3-ubyte.gz,训练集图像,大约 9.9MBtrain-labels-idx1-ubyte.gz,训练集标签,大约 29KBt10k-images-idx3-ubyte.gz,测试集图像,大约 1.6MBt10k-labels-idx1-ubyte.gz,测试集标签,大约 5KB
用浏览器打开 MNIST 的官方数据集页面,把这四个文件下载下来,然后在你的项目目录下建好这样的路径:
code复制./data/MNIST/raw/
├── train-images-idx3-ubyte.gz
├── train-labels-idx1-ubyte.gz
├── t10k-images-idx3-ubyte.gz
└── t10k-labels-idx1-ubyte.gz
放好之后,把 download 参数改成 False:
python复制train_dataset = datasets.MNIST(
root='./data',
train=True,
transform=transforms.ToTensor(),
download=False
)
torchvision 初始化时检测到 raw 目录下已经有这四个文件,就不会再去请求网络,直接进入预处理阶段。这里有几个值得注意的小细节:第一,文件名必须完全一致,包括命名里的 -ubyte 这种后缀,不能自己改名;第二,gzip 压缩包不要手动解压,torchvision 内部会自己处理,你解压反而容易出问题;第三,如果文件下载不完整,初始化时会报 RuntimeError: Dataset not found 或者 unexpected end of data,这时候重新下载对应文件就好。
如果你只是想做实验,不执着于原始 MNIST,还有个更省事的方案:用 sklearn 的 fetch_openml('mnist_784') 把 MNIST 加载成扁平化的 784 维数据。这个接口走的是 OpenML 的服务器,大多数情况下自动下载更稳定,数据内容一样,只是形状是 (N, 784) 而不是 (N, 1, 28, 28),用的时候需要自己 reshape。我自己一般优先手动下载原始文件,因为后面如果要跑图像增强、可视化、卷积操作,保留原始二维结构会更顺。
2.3 transform 和 DataLoader 的正确姿势
数据文件到位以后,第二步要处理 transform。MNIST 的像素范围是 0 到 255,如果不做任何处理直接丢进网络,数值太大容易让梯度爆炸,收敛也会很慢。通常的做法是先转成张量,再做标准化:
python复制transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))
])
train_dataset = datasets.MNIST(root='./data', train=True, transform=transform, download=False)
test_dataset = datasets.MNIST(root='./data', train=False, transform=transform, download=False)
这里 Normalize 用的 0.1307 和 0.3081 是 MNIST 全集像素的均值和标准差,是社区里约定俗成的固定值。标准化的作用是把像素分布拉回到均值为 0、方差为 1 附近,这样网络训练更稳定。需要注意,ToTensor() 会把像素值从 0 到 255 缩放到 0 到 1,所以 Normalize 的两个参数是基于 0 到 1 这个范围统计出来的,不是基于 0 到 255,很多人在这里会犯迷糊。
数据加载器方面,我习惯这样设置:
python复制train_loader = torch.utils.data.DataLoader(
train_dataset, batch_size=128, shuffle=True,
num_workers=0, pin_memory=True
)
test_loader = torch.utils.data.DataLoader(
test_dataset, batch_size=128, shuffle=False,
num_workers=0, pin_memory=True
)
shuffle=True 只在训练集上开,测试集保持顺序即可。num_workers 在 Windows 上我建议直接设 0,因为 Windows 下多进程数据加载偶尔会报 BrokenPipeError,很折腾;在 Linux 服务器上可以调到 2 到 4。pin_memory=True 在 GPU 训练时能把数据传输效率提高一点,CPU 训练开了也没有副作用。
3. 搭一个能跑到 99% 的 CNN:LeNet-5 变体的逐层拆解
3.1 为什么先选 LeNet-5,而不是一上来就上 ResNet
很多教程一上来直接扔一个 ResNet 或者 VGG 结构,看起来很高端,但读者连每层输出张量的形状怎么变都搞不清楚,更别提理解设计动机了。MNIST 是单通道 28×28 的小图,用不着太深的网络,LeNet-5 这种经典结构反而是最好的教学素材:它只有两个卷积层、三个全连接层,参数量小,训练快,结构里每个模块的作用都清晰可辨认。等把 LeNet-5 跑明白,再去看 ResNet 的残差连接、BatchNorm 的作用,思路会顺很多。
我用的结构是 LeNet-5 的 PyTorch 变体,去掉了原版里的 tanh 激活,改用 ReLU,输出层保持 10 分类。完整代码在这里:
python复制import torch
import torch.nn as nn
import torch.nn.functional as F
class LeNetMNIST(nn.Module):
def __init__(self, num_classes=10):
super(LeNetMNIST, self).__init__()
self.conv1 = nn.Conv2d(1, 6, kernel_size=5, padding=0)
self.conv2 = nn.Conv2d(6, 16, kernel_size=5, padding=0)
self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
self.fc1 = nn.Linear(16 * 4 * 4, 120)
self.fc2 = nn.Linear(120, 84)
self.fc3 = nn.Linear(84, num_classes)
def forward(self, x):
x = F.relu(self.conv1(x))
x = self.pool(x)
x = F.relu(self.conv2(x))
x = self.pool(x)
x = x.view(x.size(0), -1)
x = F.relu(self.fc1(x))
x = F.relu(self.fc2(x))
x = self.fc3(x)
return x
3.2 网络结构与维度变化
这个网络之所以经典,是因为每一步的张量形状变化都很好算。输入是 (batch_size, 1, 28, 28),逐层变化如下:
| 层 | 操作 | 输出尺寸 | 参数量 | 说明 |
|---|---|---|---|---|
| 输入 | - | (1, 28, 28) | - | 单通道灰度图 |
| Conv1 | 5×5, 1→6 | (6, 24, 24) | 156 | 28-5+1=24 |
| ReLU | - | (6, 24, 24) | - | 引入非线性 |
| Pool1 | 2×2, stride=2 | (6, 12, 12) | - | 24/2=12 |
| Conv2 | 5×5, 6→16 | (16, 8, 8) | 2416 | 12-5+1=8 |
| ReLU | - | (16, 8, 8) | - | - |
| Pool2 | 2×2, stride=2 | (16, 4, 4) | - | 8/2=4 |
| Flatten | - | (256,) | - | 16×4×4=256 |
| FC1 | 256→120 | (120,) | 30840 | - |
| FC2 | 120→84 | (84,) | 10164 | - |
| FC3 | 84→10 | (10,) | 850 | 输出 logits |
总参数量是 156 + 2416 + 30840 + 10164 + 850 = 44426,也就是大约 4.4 万个参数。这个规模随便一个 CPU 都能跑得动,训练速度很快。
这里有几个关键点。第一,x.view(x.size(0), -1) 是把 (batch_size, 16, 4, 4) 展平成 (batch_size, 256),这一步在 PyTorch 里非常常用,必须保证前面的特征图尺寸算对,否则 fc1 的输入维度对不上,会直接报错。第二,最后一层 fc3 输出的是 10 维 logits,不是经过 softmax 的概率值,因为后面用 CrossEntropyLoss 时它内部会做 softmax。第三,原版 LeNet 第一层卷积核大小是 5×5,输出 28×28 会变成 24×24,如果你想让特征图大小保持不变,可以设置 padding=2,但这里为了严格复现 LeNet 的维度变化,我特意保持 padding=0。
3.3 卷积输出尺寸公式,真的需要记下来
卷积层输出尺寸的计算公式其实很简单:
code复制输出尺寸 = (输入尺寸 + 2 × padding - kernel_size) / stride + 1
拿上面第一层举例:输入 28,kernel_size=5,padding=0,stride=1,代入后是 (28 - 5) / 1 + 1 = 24。池化层也一样,MaxPool2d(kernel_size=2, stride=2) 就是把尺寸直接除以 2,24 变 12,12 变 6。注意,如果除不尽,PyTorch 默认往下取整,所以设计网络时要保证尺寸能被整除,或者用 padding 来调整。
很多时候网络报维度错误,不是因为公式不会,而是因为某个中间层的尺寸算错了。我的习惯是每定义一个层,就在纸上或者注释里把输入输出形状标出来,等 forward 写完,整体维度链条就一目了然。新手最容易犯的错误是:复制别人的网络结构,却不检查输入尺寸是否匹配。比如网络上很多 LeNet 代码是针对 CIFAR-10 的 32×32 输入写的,拿到 MNIST 的 28×28 上就会在 fc1 处报维度错误。这时候不要慌,按照上面的公式把最后一层卷积输出的特征图宽高算出来,改一下全连接层的输入维度就行。
4. 训练环节:损失函数、优化器、batch size 与训练循环
4.1 交叉熵损失里的隐藏操作
MNIST 是一个 10 分类问题,最常用的损失函数是 nn.CrossEntropyLoss()。很多人在这一步会踩一个很坑的误区:他们会觉得“既然要输出概率,那最后一层应该加上 softmax 吧”,于是先 F.softmax(logits) 再传给 CrossEntropyLoss,结果训练半天 loss 不下降,或者准确率奇差。
原因是 PyTorch 的 CrossEntropyLoss 内部已经包含了 LogSoftmax 和 NLLLoss 两步操作。也就是说,你喂给它的是原始 logits,它会自己算出 log-softmax,再计算负对数似然损失。如果外面再手动加一个 softmax,logits 被压缩到 0 到 1 之间,再算 log-softmax,数值分布就乱了,梯度也会出问题。所以我这里最后一层 fc3 的输出直接进 loss 函数,不要做任何额外处理。如果你确实想显式地用 softmax,那就不要用 CrossEntropyLoss,改用 NLLLoss,并且网络输出前手动调用 F.log_softmax(x, dim=1)。两种方式等价,但千万别混用。
还有一个小细节:分类任务不能用 MSE 均方误差。MSE 是为回归设计的,它对“预测分布是否正确”不敏感,而且配合 softmax 时梯度容易消失。分类场景直接用交叉熵是社区验证过无数遍的结论,没必要重新发明轮子。
4.2 优化器与超参数选择的实测依据
优化器我用的是 Adam,学习率 0.001。Adam 的优势是自带自适应学习率,对新手来说不用花太多时间调学习率,收敛也快。网络结构固定后,我实测这个 LeNet-5 变体配合 Adam,大概 5 个 epoch 就能到 98% 以上的验证准确率,10 到 12 个 epoch 能稳定在 99% 左右。
如果你更想深入理解优化过程,也可以试试 SGD 加动量:
python复制optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9)
SGD 收敛会慢一些,大概需要 10 到 15 个 epoch 才能到 98% 以上,但很多老工程师喜欢它,因为在更复杂的任务上 SGD 的泛化能力往往比 Adam 好一点。入门阶段我更推荐 Adam,先把流程跑通,之后再回头对比也不迟。
batch size 我选了 128。这个值对 MNIST 来说是个比较平衡的选择:太小比如 16 或 32,梯度噪声大,训练不稳定,而且 GPU 利用率低;太大比如 512 或 1024,每个 epoch 的更新次数变少,收敛速度反而下降,还占显存。MNIST 每张图才 28×28,128 这个 batch 在显存占用上几乎可以忽略不计,但如果后面换到 3×224×224 的真实图像,batch size 就得根据显存大小重新调,这个思路是通用的。
epoch 数量我建议设 12 到 15。MNIST 小,训练很快,多跑几个 epoch 成本不高。关键是观察验证集 loss,如果验证 loss 开始上升而训练 loss 还在下降,就是过拟合的信号,这时候再多的 epoch 也没有意义。
4.3 标准训练/评估代码与运行结果
训练函数我习惯写成下面这种形式,每一步都显式调用 model.train() 和 model.eval(),这两个状态切换非常关键。model.train() 会启用 Dropout 和 BatchNorm 的训练模式,model.eval() 则会关闭它们,如果漏了,训练和测试时的结果会莫名其妙地对不上。
python复制def train_epoch(model, loader, optimizer, criterion, device):
model.train()
running_loss = 0.0
correct = 0
total = 0
for images, labels in loader:
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
running_loss += loss.item() * images.size(0)
preds = outputs.argmax(dim=1)
correct += (preds == labels).sum().item()
total += labels.size(0)
return running_loss / total, correct / total
评估函数几乎一样,但是要包在 torch.no_grad() 里面,并且记得 model.eval():
python复制@torch.no_grad()
def evaluate(model, loader, criterion, device):
model.eval()
running_loss = 0.0
correct = 0
total = 0
for images, labels in loader:
images, labels = images.to(device), labels.to(device)
outputs = model(images)
loss = criterion(outputs, labels)
running_loss += loss.item() * images.size(0)
preds = outputs.argmax(dim=1)
correct += (preds == labels).sum().item()
total += labels.size(0)
return running_loss / total, correct / total
主训练循环里,每个 epoch 结束后在验证集上评估一次,不要等到全部训练完才看效果:
python复制device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = LeNetMNIST().to(device)
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
epochs = 12
for epoch in range(epochs):
train_loss, train_acc = train_epoch(model, train_loader, optimizer, criterion, device)
val_loss, val_acc = evaluate(model, val_loader, criterion, device)
print(f"Epoch {epoch+1:02d} | train_loss {train_loss:.4f} | train_acc {train_acc:.4f} "
f"| val_loss {val_loss:.4f} | val_acc {val_acc:.4f}")
这里要特别注意设备一致性:模型放到 GPU 上了,数据和标签也必须 .to(device),否则会报 device mismatch 错误。我踩过好几次这个坑,模型在 cuda:0,数据还在 CPU 上,报错信息看起来很奇怪,其实就这一行的事。
按照这个配置,运行日志大致是这种走势:
code复制Epoch 01 | train_loss 0.4213 | train_acc 0.8870 | val_loss 0.1120 | val_acc 0.9665
Epoch 02 | train_loss 0.0871 | train_acc 0.9742 | val_loss 0.0702 | val_acc 0.9783
Epoch 05 | train_loss 0.0287 | train_acc 0.9910 | val_loss 0.0412 | val_acc 0.9875
Epoch 10 | train_loss 0.0091 | train_acc 0.9972 | val_loss 0.0340 | val_acc 0.9905
Epoch 12 | train_loss 0.0073 | train_acc 0.9978 | val_loss 0.0382 | val_acc 0.9912
最后在测试集上跑一次 evaluate,准确率应该在 99.1% 到 99.4% 之间。如果你跑出来只有 97% 或者更低,先不要怀疑网络结构,按优先级检查这几件事:数据有没有做标准化、学习率是不是太高或太低、epoch 是不是不够、是不是忘了切 model.eval()。
5. 不只盯准确率:错误样本、混淆矩阵与训练曲线
5.1 测试集评估和混淆矩阵,能看出模型真实的偏科情况
整体准确率只是一个数字,99% 看起来很好,但具体哪些数字容易混,这个数字不会告诉你。把混淆矩阵打出来,往往能发现有意思的问题。测试集评估完之后,我在项目里习惯加这一段:
python复制from sklearn.metrics import confusion_matrix
import numpy as np
all_preds = []
all_labels = []
model.eval()
with torch.no_grad():
for images, labels in test_loader:
images, labels = images.to(device), labels.to(device)
outputs = model(images)
preds = outputs.argmax(dim=1)
all_preds.extend(preds.cpu().numpy())
all_labels.extend(labels.cpu().numpy())
cm = confusion_matrix(all_labels, all_preds)
print(cm)
跑出来的混淆矩阵里,对角线上的数字是分类正确的样本数,正常情况下都很大,要重点看对角线之外的数字。在 MNIST 上,最常见的混淆对是 4 和 9、7 和 9、3 和 8、2 和 7。原因也很直观:手写数字潦草到一定程度时,4 和 9 的写法确实非常接近,7 和 9 在缺少横杠时也难以区分。这类错误不是模型结构的问题,而是数据本身的模糊性。如果你发现模型把某个数字系统性误判成另一个数字,比如 4 大量被识别成 9,比例明显异常,那才说明特征提取出了问题,需要去看卷积核或数据增强方向。
5.2 把错误样本画出来,比纯看准确率有价值得多
准确率是一个高度压缩的指标,很多信息都被吞掉了。我把测试集里预测错误的样本收集起来,用 matplotlib 画成网格图:
python复制import matplotlib.pyplot as plt
wrong_images, wrong_preds, wrong_labels = [], [], []
for images, labels in test_loader:
images, labels = images.to(device), labels.to(device)
outputs = model(images)
preds = outputs.argmax(dim=1)
for i in range(images.size(0)):
if preds[i] != labels[i]:
wrong_images.append(images[i])
wrong_preds.append(preds[i].item())
wrong_labels.append(labels[i].item())
if len(wrong_images) >= 20:
break
fig, axes = plt.subplots(4, 5, figsize=(12, 10))
for i, ax in enumerate(axes.flat):
ax.imshow(wrong_images[i].squeeze().cpu().numpy(), cmap='gray')
ax.set_title(f"true: {wrong_labels[i]}, pred: {wrong_preds[i]}")
ax.axis('off')
plt.show()
打开图之后,你大概率会发现两类情况。一类是“人眼都很难认”的样本,比如一个 9 写得歪歪扭扭,中间断了一笔,既像 4 又像 9,模型选错了情有可原;另一类是“明明很清晰但模型认错”的样本,这类样本才值得深挖,说明网络某个卷积核可能没学到对应笔画的模式。我看过不少入门项目,报告里只写“准确率 98.8%”,一个错误样本都不看,这在真实项目里是绝对不行的。把错误样本截下来,告诉团队“这类样本主要是什么原因”,才是能落地的分析方式。
5.3 训练曲线能告诉你什么
训练过程中记录每个 epoch 的 train loss 和 val loss,训练结束后画出来。我一般用一个简单的列表来存:
python复制train_losses = []
val_losses = []
val_accs = []
for epoch in range(epochs):
tr_loss, tr_acc = train_epoch(...)
va_loss, va_acc = evaluate(...)
train_losses.append(tr_loss)
val_losses.append(va_loss)
val_accs.append(va_acc)
plt.plot(train_losses, label='train loss')
plt.plot(val_losses, label='val loss')
plt.legend()
plt.show()
如果 train loss 一直在降,但 val loss 在某个 epoch 之后开始反弹,那就是过拟合,最简单的应对是加 Dropout、加数据增强,或者提前停止。如果 train loss 和 val loss 都降得很慢,那大概率是学习率太低或者网络容量不足。如果 loss 直接变成 NaN,那通常是学习率太高,或者数据没有标准化。曲线不是仪式感,是每个深度学习从业者诊断模型的第一手段。
6. 从 99% 继续往前:提速、数据增强和环境相关的坑
6.1 让训练跑得更快的一些细节
MNIST 这个规模,CPU 训练也只要几十秒一个 epoch,但如果你用的机器比较老旧,或者后面想迁移到更大的数据集,有几个提速细节值得现在就用起来。
第一,确认设备到底用没用上 GPU。torch.cuda.is_available() 返回 True 不代表模型真的跑在 GPU 上,还要看模型和数据有没有 .to(device)。训练时打开任务管理器或者 nvidia-smi,看到显存占用和 GPU 利用率在跳动,才说明真的在用 GPU。
第二,DataLoader 的 num_workers 和 pin_memory。pin_memory=True 在 GPU 训练时可以减少数据从 CPU 到 GPU 的拷贝时间,num_workers 开了之后数据加载可以并行。但在 Windows 下 num_workers 设置成大于 0 有时候会报错,新手建议直接设 0,稳定性优先。
第三,如果显卡是 20 系以上的 NVIDIA 卡,可以试试混合精度训练。PyTorch 自带的 torch.cuda.amp 可以把一部分计算用 float16 来做,训练速度和显存占用都能优化。MNIST 这种小网络收益不大,但代码模式值得熟悉:
python复制from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for images, labels in train_loader:
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad()
with autocast():
outputs = model(images)
loss = criterion(outputs, labels)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
这个代码是通用模板,后面换到任何分类任务都能直接用。如果哪天你发现训练速度没有提升,先看显卡是否支持 float16 加速,再看数据加载是不是成了瓶颈。
6.2 提高精度的组合拳:数据增强、BatchNorm 和更现代的结构
LeNet-5 在 MNIST 上能到 99% 左右,但再往上走,这个结构就有点吃力了。想要往 99.5% 甚至更高冲,方向很明确。
第一个方向是数据增强。MNIST 的样本是固定大小、黑底白字,直接做简单增强很有效:
python复制train_transform = transforms.Compose([
transforms.RandomAffine(degrees=10, translate=(0.1, 0.1)),
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))
])
RandomAffine 让图像在小角度旋转和轻微平移的范围内变化,相当于教模型“数字稍微歪一点也能认出来”。加了增强之后,验证准确率不一定每个 epoch 都更高,但最终测试集准确率一般
