要不是昨天有人在群里发了一张 torchvision 下载 MNIST 数据集报 404 的截图,我差点忘了自己当年入坑 PyTorch 时也被这一下卡了整整一个晚上。想学 CNN 却连最经典的手写数字分类数据集都拉不下来,这几乎是每个初学者都会经历的一道“开胃败仗”。
这篇东西不是官方文档的复述,而是我实际从零跑通 PyTorch + MNIST + CNN 全流程的记录。包括环境怎么搭、数据加载器怎么写、模型结构为什么这样堆、训练循环里哪些细节不处理必踩坑,以及最后如何把模型保存下来、导出 ONNX、用它识别真实手写图片。我会把自己踩过的 404、CUDA 不可用、黑白翻转这些坑一一说清楚,把你最可能卡住的位置提前点上灯。
1. 为什么MNIST是CNN入门路上绕不开的那个数据集
不少人会问:都 2024 年了,MNIST 这种 28×28 的老古董还有什么好学的?我的看法是,正因为简单,它才是最好的“试金石”。
MNIST 由 6 万张训练图片和 1 万张测试图片组成,每张都是 28×28 的灰度图,内容是 0 到 9 的手写数字。它的特点非常鲜明:单通道、小尺寸、类别均衡、背景干净。这意味着你不需要多贵重的 GPU,普通 CPU 都能在几分钟内把模型训练到 99% 以上。对初学者来说,这种“付出有回报、调试有反馈”的节奏感太重要了。
从另一个角度讲,MNIST 又足够“像样”。手写数字存在笔画粗细不均、位置偏移、形变、断笔等现象,这些特征和真实图像识别任务面临的问题是同一类问题,只是程度更轻。你在 MNIST 上学会的卷积、池化、数据增强、过拟合分析、训练评估流程,换到 CIFAR-10、ImageNet 时依然成立,只是数据和模型规模变大而已。
还有个很容易被忽略的价值:MNIST 的错误可视化非常直观。训练集上预测错了,把图片画出来,肉眼看一眼就知道是模型犯糊涂还是标注本身就有问题。这种“即时纠错”的体验对建立直觉极有帮助,不是每个数据集都能给你。
所以别嫌它“太简单”。我见过一些新同学一上手就啃 Transformer、Diffusion,结果连梯度都不回传,愁眉苦脸一整天。先把 MNIST 上的 CNN 流程吃透,后面的复杂网络才谈得上有章法。
2. 环境搭建踩坑实录:版本匹配、CUDA、WSL与MNIST下载404
环境搭建是最劝退新手的一步。这里啰嗦几句,我踩过的坑几乎都是在这里踩的。
2.1 安装策略:conda 还是 pip,日常开发怎么选
我的习惯是先用 Anaconda 建一个独立环境,避免把系统的 Python 环境搅乱。操作非常简单:
conda create -n mnist python=3.10 conda activate mnist接着安装 PyTorch。官网首页会给出对应 CUDA 版本的安装命令,这个必须自己去查,因为 PyTorch 版本、CUDA 版本、Python 版本三者之间有兼容关系。
选 pip 还是 conda?我个人的建议是:能用 conda 就用 conda,尤其在 Windows 上。conda 对 MKL、OpenMP 这类底层依赖的处理更省心,pip 偶尔会遇到 PyTorch 装好了但 import 时动态链接库报错的情况。
这里最容易被忽略的是 CPU 版和 GPU 版的区别。如果电脑没有 NVIDIA 显卡,直接装 CPU 版就够了;如果有显卡,务必看清楚安装命令里是 cu118、cu121 还是 cu124 这种后缀,它代表配套的 CUDA 版本。不要盲目装最新的 CUDA,PyTorch 官方的支持列表是你该相信的,不是显卡驱动面板里显示的数字。
安装完一个习惯性自检:
import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else "CPU")如果torch.cuda.is_available()返回 False,是大多数环境问题集中爆发的信号。原因可能包括:装成了 CPU 版、NVIDIA 驱动版本太老、系统 CUDA 版本与 PyTorch 不匹配、在 WSL 里没启用 GPU 支持等。
2.2 在 WSL/Linux 环境里配置 GPU 支持
现在很多同学在 WSL 里跑实验。WSL 2 本身支持 CUDA,前提是 Windows 侧已经装好 NVIDIA 显卡驱动,并且 WSL 里的 PyTorch 必须装 CUDA 版。
我遇到的一个典型情况是:在 WSL 里nvidia-smi能正常显示,但 PyTorch 就是检测不到 CUDA。原因往往是torch.cuda.is_available()检测的是 PyTorch 编译时链接的 CUDA runtime,而不是nvidia-smi里的驱动版本。解决办法是在 WSL 的 conda/pip 环境里重新安装对应 CUDA 版本的 PyTorch,而不是依赖系统级 CUDA。
还有个容易踩的坑:Windows 和 WSL 的文件系统互访。如果把项目放在/mnt/c/...路径下,训练时读取小文件可能很慢。MNIST 这种小数据集还好,如果换到大数据集,建议把代码和数据都放在 WSL 的 ext4 目录里,速度和稳定性会好很多。
2.3 torchvision 下载 MNIST 报 404?这里有一个手工落地的方案
很多同学是在这一步崩溃的:执行datasets.MNIST(root='./data', download=True)后,终端出现HTTP Error 404: Not Found,数据集下载失败。
这个问题的根源通常是下载源访问不顺畅。MNIST 数据集文件本身是公开的、体积也就几 MB 到十几 MB,但下载源在某些网络环境下并不稳定。我不建议反复重试,更靠谱的是手工下载这 4 个文件:
train-images-idx3-ubyte.gztrain-labels-idx1-ubyte.gzt10k-images-idx3-ubyte.gzt10k-labels-idx1-ubyte.gz
下载完成后,把它们放到./data/MNIST/raw/目录下,注意目录层级必须匹配。然后把代码里download=True改成download=False,再运行一次。torchvision 检测到本地已经有原始文件,就会自动解压并生成格式化的数据。
这个手工方案的原理其实很简单:datasets.MNIST的第一步是从远端拉取原始 gzip 文件,如果本地 raw 目录已经有这些文件,它就不管网络了。从教育角度讲,手工下载文件也让你更清楚数据从哪来,必要时还能编写脚本自行校验。
另外补充一点:如果运行环境本身访问下载源比较慢,即使不报 404,也可能卡很久。这时候可以设置一个超时时间,仔细观察报错信息,不要盲目等。
2.4 训练时遇到“CUDA out of memory”怎么办
MNIST 数据集很玩具,基本不会 OOM,除非你把 batch size 调到几千。但很多人用同一个 torch 环境跑大网络时就会撞上。养成好习惯:数据、模型、损失函数全部to(device),batch size 不要盲目翻倍。如果真的遇到显存不足,先把batch_size减半,或者减少 DataLoader 的num_workers,这招在 MNIST 阶段能覆盖大多数情况。
3. 数据管线:Dataset、DataLoader和归一化里的小心计
模型写得再漂亮,数据管线没做好也白搭。MNIST 的数据读取用 torchvision 几行就能搞定,但背后的几个细节值得掰开揉碎讲。
先看一段常用写法:
from torch.utils.data import DataLoader from torchvision import datasets, transforms transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_set = datasets.MNIST(root='./data', train=True, transform=transform, download=False) test_set = datasets.MNIST(root='./data', train=False, transform=transform, download=False) train_loader = DataLoader(train_set, batch_size=128, shuffle=True, num_workers=4, pin_memory=True) test_loader = DataLoader(test_set, batch_size=256, shuffle=False, num_workers=4, pin_memory=True)3.1 ToTensor 不只是“转张量”这么简单
transforms.ToTensor()完成两件事:把 PIL Image 或 numpy 数组(H×W×C)转成 PyTorch 张量,并在最后增加一个通道维度(C×H×W);同时把像素值从 0~255 线性缩放到 0~1。
很多第一次写代码的人会自己手动做:
x = torch.tensor(np.array(img)).float()结果训练直接崩一脸。原因很典型:通道顺序不对,或者 dtype 不是浮点型、范围没有归一化。ToTensor帮你把这些琐碎细节一次性处理干净,能少踩很多坑。
MNIST 原始图片是单通道灰度图,所以ToTensor()之后张量形状是[1, 28, 28],第一个 1 就是通道维。如果后续模型第一层Conv2d的in_channels=1,正好对上。
3.2 归一化为什么非做不可
灰度图转成 0~1 范围已经不错了,但还不够。Normalize((0.1307,), (0.3081,))里的两个数字是 MNIST 全体训练集像素的均值和标准差。
为什么要额外归一化?因为神经网络对输入的尺度很敏感。输入特征的值域如果不一致,早期梯度更新的方向就会有很大波动,训练要么不稳定,要么收敛很慢。把数据变成“零均值、单位方差”之后,每个像素的特征分布更接近,优化器迭代起来也快得多。这两组数字并不是拍脑袋定的,而是统计出来的,所以直接沿用就好。
这个参数等模型训练完、部署新图片时还要再一次用到。很多人模型训练没问题,但预测单张图时忘了套同样的归一化,结果准确率莫名其妙掉到 20% 以下,这点后面还会再提。
3.3 数据增强:MNIST 上简单做就行,别上头
有人刚学到数据增强,就恨不得把旋转、平移、缩放、加噪声全部堆上去。但 MNIST 做增强要克制。数字 6 和 9 只要稍加旋转就容易混淆,人眼都容易认错;过度旋转反而让测试集性能下降。
如果确实想加,我推荐两种轻量方式:
RandomAffine加一个很小的角度范围,比如正负 10 度,外加 10% 左右的平移;RandomErasing随机擦除一小块,模拟笔画不完整的情况。
在 MNIST 上做增强的问题在于,测试集和训练集分布差异不大,增强带来的泛化收益有限,反而增加了训练时间。我的经验是:先跑一次不加增强的基线,精确率大约 99% 左右,再决定是否加增强,否则你很难判断提升究竟是增强带来的还是代码修改带来的。
3.4 DataLoader 的 tiny 细节
batch_size在没有显卡的情况下不宜太大,否则计算慢。CPU 上我建议batch_size=64或128;GPU 上256也毫无压力。
shuffle只对训练集有意义,测试集不要 shuffle,否则你后续打印一张图对一张标签时,顺序全乱掉,排查问题会很难受。
num_workers在 Linux 上可以设成 4 或者 8,但在 Windows 上经常会报BrokenPipeError或者DataLoader worker (pid...) is killed。这是 Windows 多进程的经典问题,省事方案是设成num_workers=0。虽然数据读取慢一点,但能避免一个巨坑。我自己的经验是:MNIST 数据集太小,num_workers=0的额外开销完全在可接受范围,没必要为了炫技给自己找麻烦。
pin_memory=True在 GPU 训练时可以把数据固定到页锁定内存里,减少主机到 GPU 的拷贝耗时。但它必须配合device='cuda'才有意义,纯 CPU 训练时开着不会有坏处,也不用指望有多大提升。
4. 搭建CNN模型:从LeNet思路到现代小改动
MNIST 的 CNN 网络结构并不复杂,经典 LeNet-5 已经能做到很好的效果。不过 LeNet 是上世纪 90 年代的设计,很多写法在现代 PyTorch 里可以做点小改进,比如加上 BatchNorm 和 Dropout。
我常用的一个结构是这个:
import torch import torch.nn as nn class MNISTCNN(nn.Module): def __init__(self): super().__init__() self.features = nn.Sequential( nn.Conv2d(1, 32, kernel_size=3, padding=1), nn.BatchNorm2d(32), nn.ReLU(inplace=True), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size=3, padding=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True), nn.MaxPool2d(2), ) self.classifier = nn.Sequential( nn.Flatten(), nn.Linear(64 * 7 * 7, 128), nn.ReLU(inplace=True), nn.Dropout(0.2), nn.Linear(128, 10), ) def forward(self, x): return self.classifier(self.features(x))4.1 每一层为什么这么放
先看输入输出尺寸的变化。输入是[1, 28, 28]。第一次卷积,kernel_size=3, padding=1,会让尺寸保持不变,仍然是28×28,因为(28 - 3 + 2*1) / 1 + 1 = 28。接着MaxPool2d(2)把宽高减半,变成14×14。第二次卷积同理保持14×14,再池化变成7×7。此时通道数是 64,所以展平后是64*7*7 = 3136。
用 padding=1 的好处是让特征图尺寸成倍数缩小,计算形状很直观。如果不加 padding,第一层卷积后变成26×26,池化后13×13,第二次卷积后11×11,池化后5×5,展平尺寸是64*5*5 = 1600。两者都能用,但我建议初学者用 padding=1 的写法,减少形状推算时的出错概率。
两个卷积块就够了吗?对 MNIST 来说,两个卷积层已经能捕捉到笔画、边缘、闭环结构等关键特征。继续堆更多卷积层,精度的提升非常有限,还会让训练更慢、过拟合风险更高。训练 CNN 不是层数越多越好,而是够用就好。
4.2 BatchNorm 与 Dropout 的位置奥妙
BatchNorm2d放在卷积之后、ReLU 之前,已经成为一种主流做法。它的作用是把当前 batch 的分布拉回均值为 0、方差为 1 的状态,缓解网络内部协变量偏移问题。
但 BatchNorm 在训练和推理时的行为是不一样的。训练时它统计当前 batch 的均值和方差;推理时它使用训练过程中累计得到的全局均值和方差。这就是为什么后面训练循环里model.train()和model.eval()必须成对切换,否则会出现很诡异的现象:训练时 loss 正常下降,测试时结果一塌糊涂。
Dropout(0.2)放在全连接层之前的含义是:每次前向传播有 20% 的神经元输出被随机置零,迫使网络不依赖个别神经元,提高泛化能力。Dropout 同样只在训练时生效,在eval()模式下它会自动关闭。这两个层倒不是说对 MNIST 提升巨大,但它们教会你的知识点会沿用到所有后续网络里。
4.3 为什么不急着上残差网络
现在很多人一上来就 ResNet,这没必要。ResNet 处理的是深层的梯度退化问题,它需要足够深的网络才能体现优势。在 MNIST 上强行套 ResNet,反而网络过重,训练慢,而且要处理输入尺寸、padding、stride 等问题,对新手极不友好。
我的建议是:先把上面这个轻量 CNN 从零写明白,弄懂 forward 里每一步的形状变化、每个模块的作用,再逐步引入更重的结构。真正的高手不是只会调包,而是对组合结构有直觉。
5. 训练循环:CrossEntropy、Adam与99%背后的工程细节
模型定义好,接下来是训练。MNIST 的训练循环说简单也简单,说难也有几个细节会卡住人。
完整代码参考这段:
import torch import torch.optim as optim device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = MNISTCNN().to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=1e-3) for epoch in range(5): model.train() total_loss = 0.0 correct = 0 for x, y in train_loader: x, y = x.to(device), y.to(device) optimizer.zero_grad() out = model(x) loss = criterion(out, y) loss.backward() optimizer.step() total_loss += loss.item() * x.size(0) pred = out.argmax(dim=1) correct += (pred == y).sum().item() train_acc = correct / len(train_loader.dataset) print(f"epoch {epoch+1}, train loss: {total_loss / len(train_loader.dataset):.4f}, train acc: {train_acc:.4f}") model.eval() test_correct = 0 with torch.no_grad(): for x, y in test_loader: x, y = x.to(device), y.to(device) out = model(x) pred = out.argmax(dim=1) test_correct += (pred == y).sum().item() test_acc = test_correct / len(test_loader.dataset) print(f"test acc: {test_acc:.4f}")5.1 损失函数为什么是 CrossEntropyLoss 而不是 MSELoss
分类任务的标准选择是交叉熵损失,对应 torch 里的nn.CrossEntropyLoss。
这里有个细节必须说清楚:CrossEntropyLoss内部已经包含 softmax 操作。也就是说,它接收的是模型输出的原始 logits,而不是经过 softmax 的概率值。很多人第一次手写模型时,会在nn.Linear之后再加一个nn.Softmax,然后把自己坑了。原因是:logits 在进入 CE loss 前还要再过一次 log_softmax,你手动的 softmax 反而让梯度在数值上变得不稳定,而且和网络最后的概率标签对不上。
把nn.CrossEntropyLoss理解成“log_softmax + NLLLoss”的合并体就够了。模型最后一层不用 softmax,它在训练时由 loss 函数内部处理,在推理时你用out.argmax(dim=1)拿到预测标签即可。
5.2 优化器选 Adam 还是 SGD
初学者我推荐 Adam,学习率设1e-3,基本不需要太精细的调参就能快速收敛。原因在于 Adam 自带自适应学习率,对梯度尺度不敏感,比较“皮实”。
但我也得坦白说,真正调参经验更丰富之后,SGD with Momentum 往往能刷出比 Adam 更高的测试精度。SGD 对学习率更加敏感,需要更精细的 schedule,但它收敛到的最小值通常更平缓,泛化性更好。MNIST 任务上这个差异很小,没必要为了 0.1 个百分点折腾太久。
经验值供参考:MNIST 上 Adam + lr=1e-3,通常 3 个 epoch 训练集就有 99.5% 以上,5 个 epoch 测试集能达到 99.2% 到 99.5%。再往高刷的边际成本远大于收益。
5.3 model.train() 与 model.eval() 的切换不能省略
很多初学者写完循环不切model.train()和model.eval(),之前 4.2 我讲过 BatchNorm 和 Dropout 的存在,它们对 train 和 eval 两种状态反应完全不同。
model.train()是在告诉模型:现在是训练阶段,BatchNorm 使用当前 batch 统计量,Dropout 随机失活。model.eval()则进入推理模式:BatchNorm 使用训练期间保存的全局统计量,Dropout 关闭,同时nn.Flatten之外的任何可能影响结果的结构都要稳定下来。
如果把这段切换忘了,测试时模型内部某些层的行为仍然保持训练模式,输出结果会出现波动,准确率看起来忽上忽下。这是新手最容易忽略但又最影响判断的坑。
5.4no_grad的作用
测试阶段用with torch.no_grad():包裹起来。它的意思是告诉 PyTorch 不需要计算梯度。
如果你不加这层保护,测试前向传播仍然会构建计算图、暂存中间变量,不但占内存,还浪费时间。在 MNIST 这种小任务上影响不明显,但换成大一点的模型,比如后续做 CIFAR-10 时,这个习惯必须提前养成。
5.5 训练过程的实时监控不是可选项
我强烈建议每个 epoch 打印训练集 loss、训练集准确率和测试集准确率。别只盯测试准确率,因为训练集 loss 的变化能告诉你模型是否还在收敛,是否已经开始过拟合。
一个经典信号是:训练集准确率不断逼近 100%,但测试集准确率停滞甚至下降。这是过拟合的警报。这时候可以做的事包括增加 Dropout、加入轻量数据增强、减少训练 epoch,而不是继续盲刷训练轮数。
6. “准确率挺高”是好事,但MNIST会骗人:过拟合与数据分布
很多人看到测试集 99.4% 就觉得自己“搞定了”。实际上 MNIST 太简单,准确率非常容易被刷上去,也正因为这样,它非常适合用来分析过拟合和分布漂移。
6.1 训练几个 epoch 就过拟合了
MNIST 的 6 万张训练图片相比模型参数量并不算多。当模型训练到第 4、5 个 epoch 时,训练集的准确率往往会达到 99.9% 以上,而测试集却停留在 99.3% 左右。这条“缝”就是过拟合。
我习惯的做法是这样的:把训练集每个类的准确率单独算一遍,再看测试集的混淆矩阵,重点关注到底哪些数字互相混淆。比如常见的 4 和 9、3 和 8、7 和 2,这种混淆在笔画结构上很合理,如果连人眼都觉得难分,那模型的错误就有可解释性。
6.2 测试集不能拿来反复调参
这里有一条我在学习和带人时反复强调的规则:测试集只能你最终用一次,用来报告效果,不能作为日常调参的指标。因为一旦你根据测试集性能反复修改超参数、模型结构,测试集的信息就“污染”了模型选择过程,最终报告的数字会失真。
正确做法是把训练集再切出一块验证集,或者用交叉验证,调参只看验证集。MNIST 官方给了测试集,但很多人不自觉地把测试集当成验证集用。虽然 MNIST 太简单,这个做法坏处不大,但坏习惯一旦养成,后面做大项目时一定会吃亏。
6.3 一张真实手写图片就把模型打回原形
MNIST 图片的背景是黑色、数字是白色。现实里拿手机拍一张白纸黑字的手写数字,导入模型前如果只是简单地 resize 到 28×28、转成灰度图,大概率识别效果惨不忍睹。
原因有两层:
- 颜色极性是反的。MNIST 训练数据大多是黑底白字,白底黑字输入相当于把图像颜色反转了,模型看到的特征分布彻底变了。
- 数字不一定居中,笔画粗细也相差很大。MNIST 的预处理把数字大致居中,但现实图片随手一拍的偏移和缩放,网络没见过就很容易误判。
解决方案是在预处理里多做几步:把图片转成灰度、根据阈值做二值化、裁掉过于冗余的边缘、缩放到 20×20 再放到 28×28 画布中央。这一步说白了是尽量把真实图片处理成和 MNIST 训练分布一致的样子。很多人在模型上折腾半天,结果问题出在预处理上,这是最尴尬的翻车现场。
6.4 保存错例比保存准确率数字更有用
训练完后我会专门做一个步骤:把测试集里预测错误的那批(image, label, pred)三元组保存成图片,按类别分开存放。这样我能直观看到模型在哪些数字上犯迷糊。
这种做法比贴一个“test acc: 0.9930”有用得多。因为准确率只是一个标量,错例却是可解释的证据。哪怕你没时间画混淆矩阵,抽几十张错例图看一遍,也能快速判断问题是数据预处理、模型容量还是训练策略导致的。这个习惯从 MNIST 阶段养成,会让你后面处理真实项目时受益很多。
7. 保存模型、导出ONNX以及拿真实手写图片推理
到这里模型已经训练好了,但博文不能止步于训练,后面还有几张“实战底牌”。
7.1 正确保存和加载方式
PyTorch 保存模型有几种写法,我建议尽量保存state_dict,也就是模型权重字典,而不是整个模型对象:
torch.save(model.state_dict(), 'mnist_cnn.pth')加载的时候需要先定义同样的网络结构,再加载权重:
model = MNISTCNN() model.load_state_dict(torch.load('mnist_cnn.pth', map_location='cpu')) model.eval()用map_location='cpu'的好处是:模型如果在 GPU 上训练,加载到没有 GPU 的机器上也不会报设备不匹配的错误。保存整个 model 对象虽然省事,但一旦网络代码变化或者依赖版本升级,反序列化很可能出问题,实际部署时非常脆弱。
7.2 导出为 ONNX 的流程
导出 ONNX 是为了把模型从 PyTorch 生态中解脱出来,方便部署到移动端、服务端或者嵌入到推理引擎里。MNIST 这个例子导出 ONNX 非常简单:
dummy_input = torch.randn(1, 1, 28, 28) torch.onnx.export( model, dummy_input, 'mnist_cnn.onnx', input_names=['input'], output_names=['output'], dynamic_axes={'input': {0: 'batch_size'}, 'output': {0: 'batch_size'}} )dynamic_axes设置成动态 batch 后,推理时可以一次传入多张图片,不至于每次只能处理单张。导出完成后可以用onnxruntime验证一遍输入输出的形状,确保和预期一致。
还有一个细节:ONNX 导出时模型要在eval()模式。如果在train()模式下导出,BatchNorm 和 Dropout 的行为仍然停留在训练状态,导出后的模型在推理端结果会不稳定。
7.3 用真实手写图片做推理
这里我用一个非常简化的单张图片推理示例,前提是图片已经过灰度化、二值化、裁剪缩放处理,最终得到 28×28 灰度张量:
import cv2 import numpy as np import torch from torchvision import transforms def load_and_preprocess(img_path): img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) img = cv2.resize(img, (28, 28), interpolation=cv2.INTER_AREA) # 如果检测到背景偏黑、前景偏白,可以先保持原样; # 如果背景是白色,需要反转像素值,让数字变成白色 if np.mean(img) > 127: img = 255 - img img = img.astype(np.float32) / 255.0 img = (img - 0.1307) / 0.3081 img = torch.from_numpy(img).unsqueeze(0).unsqueeze(0) # [1,1,28,28] return imgnp.mean(img) > 127这一步是我的常用判断:如果整图平均亮度很高,说明大概率是白底黑字,直接翻转像素值,让数字变成高亮像素,而不是黑底白字。这个自动判断不是万无一失,但对大多数真实图片足够有效。
推理部分就是常规操作:
with torch.no_grad(): pred = model(img).argmax(dim=1).item() print(pred)这里最容易出问题的地方是张量形状。灰度图的形状是[H, W],加两个维度后变成[1, 1, H, W],如果写成img.unsqueeze(0)少了通道维,模型会直接报维度错误,或者in_channels=1与输入通道数 3 对不上。还有同学习惯用 PIL 的Image.open,却忘了convert('L'),把三通道 RGB 图直接喂进单通道模型,也会翻车。
7.4 把整个推理流程封装成一个类,后续省力
等模型稳定之后,我就不再写裸函数了,而是把它封装成一个小类,方便在 API 服务和命令行工具里复用。核心是三个方法:preprocess、predict、predict_batch。类的好处是初始化时加载模型一次,推理时不重复加载,服务响应速度快很多。
封装的时候注意一个细节:模型加载后必须调用model.eval(),同时保持 forward 只做前向传播,不在类内部改任何参数。这样你在调试和部署时,行为和结果都可复现,不会出现一个接口今天识别成功明天识别结果不同的奇怪现象。
我在实际使用中的体会是,MNIST 这个项目虽然小,但它几乎覆盖了一个深度学习完整流程里所有的关键节点:环境配置、数据处理、模型设计、训练评估、导出部署、真实场景适配。把这套流程跑顺了,后面的路会走得踏实很多。最后再分享一个小技巧:训练完成后,把测试集里预测错误的图片全部保存到一个单独目录,每隔一段时间拿出来看看。你会发现,模型真正学到的和你以为它学到的,经常是两个完全不同的东西。