简介:这份资源面向深度学习入门者与需要快速验证手写数字识别效果的开发者,提供基于MNIST数据集训练前馈神经网络的完整方案,解决从零搭建模型时环境配置繁琐、训练耗时的问题。压缩包共5个文件,包含3个Python脚本与2个H5模型文件,整体约1.39MB,脚本覆盖数据加载、网络定义与训练流程,H5文件分别保存模型参数与完整模型结构,便于直接加载推理或继续微调。资源已有6118人学习下载,说明其在实际练习与项目起步阶段具备较高参考价值。拿到后可直接运行脚本复现训练过程,也可跳过训练环节,用现成模型完成手写数字图片的预测验证,同时对照代码理解前馈神经网络的基本结构与训练要点,适合作为课程作业、实验报告或入门练手的轻量级素材。
1. 从一次“模型训练完却不敢上线”说起:MNIST 手写数字识别到底该怎么落地
很多人第一次跑 MNIST 手写数字识别,都会经历同一个瞬间:训练脚本跑完,终端打印出 99% 的准确率,心里一阵激动,然后……就没有然后了。模型文件躺在checkpoints/里,既不知道怎么在别的机器上复现,也不知道怎么把它塞进一个真实的小工具里。更尴尬的是,换台机器重新torchvision.datasets.MNIST(download=True),直接给你甩一个 404,或者卡在下载进度条上不动,这就是热搜里那个“torchvision下载mnist会404”的真实来源。
这篇笔记要解决的就是这条链路:用 MNIST 数据集训练一个手写数字识别模型,把完整代码写清楚,把训练好的模型文件怎么保存、怎么加载、怎么验证讲透,让你拿到代码就能跑,跑完就能用。它适合两类人:一类是刚入门深度学习、想找一个能完整走通“数据→训练→保存→推理”闭环的从业者;另一类是手头有个小需求(比如票据数字识别、表单数字录入),想先用 MNIST 练手验证方案可行性的工程师。MNIST 本身很简单,但“简单数据集 + 完整落地链路”恰恰是很多人缺的那一课。
2. 先把数据和网络这两件事定下来:MNIST 加载与模型选型的取舍
2.1 MNIST 数据集的结构与三种加载方式
MNIST 一共 70000 张 28×28 的灰度图,其中 60000 张训练、10000 张测试,10 个类别对应数字 0 到 9。它的原始格式是 IDX(一种二进制格式),不是常见的图片文件夹结构,所以你不能直接拿ImageFolder去读。常见做法有三种,我一般按场景选:
第一种,直接用torchvision.datasets.MNIST,最省事,适合快速验证。第二种,提前把 IDX 转成 PNG 或 numpy 数组,适合需要自己做数据增强、或者训练框架不是 PyTorch 的场景。第三种,用sklearn.datasets.fetch_openml('mnist_784'),适合只做传统机器学习(比如 SVM、KNN)的对比实验。
先看最常用的 torchvision 方式,这里有个关键点:download=True触发的下载地址在某些网络环境下会失败,所以生产环境我一般提前把四个压缩包(train-images-idx3-ubyte.gz、train-labels-idx1-ubyte.gz、t10k-images-idx3-ubyte.gz、t10k-labels-idx1-ubyte.gz)放到./data/MNIST/raw/目录下,再让 torchvision 去读,避免每次训练都依赖网络。
import torch from torchvision import datasets, transforms from torch.utils.data import DataLoader # 关键:先定义 transform,ToTensor 会把 0-255 的像素归一化到 0-1 transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) # MNIST 全局均值和标准差 ]) # root 指向本地目录,download=False 表示只用本地已存在的文件 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=2) test_loader = DataLoader(test_set, batch_size=256, shuffle=False, num_workers=2) print(len(train_set), len(test_set)) # 60000 10000这段代码里有两个参数值得说清楚。Normalize((0.1307,), (0.3081,))里的两个数是 MNIST 训练集的全局均值和标准差,用它们归一化能让输入分布更接近标准正态,收敛更稳;如果你不做归一化,模型也能训,但前期 loss 下降会明显更慢。num_workers=2在 Windows 上如果报错,直接改成 0,这是血泪经验,别硬扛。
2.2 模型选型:为什么我推荐先上一个小 CNN
MNIST 上能用的模型很多,从逻辑回归到 ResNet 都能跑。但选型要看目标:如果你是要一个能快速复现、参数量小、CPU 也能推理的模型,一个小型卷积网络(CNN)是最优解。全连接网络在 MNIST 上也能到 97% 左右,但它对平移敏感,泛化到你自己手写的数字时掉点明显;CNN 的卷积核天然有平移不变性,实测在真实手写场景下更稳。
我常用的结构是两层卷积 + 两层全连接,参数量约 120 万,训练 5 个 epoch 就能到 99% 以上。下面给出完整定义:
import torch.nn as nn import torch.nn.functional as F class SmallCNN(nn.Module): def __init__(self): super().__init__() # 输入 1x28x28 self.conv1 = nn.Conv2d(1, 32, kernel_size=3, padding=1) # -> 32x28x28 self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1) # -> 64x28x28 self.pool = nn.MaxPool2d(2) # 每次减半 self.fc1 = nn.Linear(64 * 7 * 7, 128) self.fc2 = nn.Linear(128, 10) self.dropout = nn.Dropout(0.25) def forward(self, x): x = self.pool(F.relu(self.conv1(x))) # 32x14x14 x = self.pool(F.relu(self.conv2(x))) # 64x7x7 x = x.view(x.size(0), -1) # 展平 x = F.relu(self.fc1(x)) x = self.dropout(x) return self.fc2(x)padding=1保证卷积后尺寸不变,这样两次池化后正好是 7×7,全连接层的输入维度64*7*7就是这么来的。Dropout(0.25)放在全连接之后,是为了抑制过拟合,MNIST 数据量不大,不加 dropout 训练集准确率会明显高于测试集。如果你把卷积核改成 5×5,那padding要改成 2,否则尺寸对不上,这是新手最容易翻车的地方。
3. 训练脚本怎么写:从 loss 曲线到模型文件落盘
3.1 训练循环与三个必调参数
训练循环本身不复杂,但有几个参数直接决定你能不能复现出 99%。我把它们列成表,方便你对照调整:
| 参数 | 推荐值 | 作用与调整建议 |
|---|---|---|
| 学习率 lr | 1e-3 | Adam 的默认值,太大 loss 震荡,太小收敛慢 |
| batch_size | 128 | 太小梯度噪声大,太大显存吃紧且泛化略差 |
| epoch | 5~8 | MNIST 上 5 轮足够,再多容易过拟合 |
| 优化器 | Adam | 比 SGD 收敛快,适合快速验证 |
| 损失函数 | CrossEntropyLoss | 多分类标准选择,内部含 softmax |
下面是完整训练代码,包含每轮在测试集上的评估,以及最优模型保存逻辑:
import torch from torch import optim device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = SmallCNN().to(device) optimizer = optim.Adam(model.parameters(), lr=1e-3) criterion = nn.CrossEntropyLoss() best_acc = 0.0 for epoch in range(1, 6): model.train() for imgs, labels in train_loader: imgs, labels = imgs.to(device), labels.to(device) optimizer.zero_grad() out = model(imgs) loss = criterion(out, labels) loss.backward() optimizer.step() # 每轮结束做一次评估 model.eval() correct = total = 0 with torch.no_grad(): for imgs, labels in test_loader: imgs, labels = imgs.to(device), labels.to(device) pred = model(imgs).argmax(dim=1) correct += (pred == labels).sum().item() total += labels.size(0) acc = correct / total print(f'epoch {epoch}, test acc = {acc:.4f}') # 只保存效果最好的那一版,避免最后一轮过拟合反而变差 if acc > best_acc: best_acc = acc torch.save(model.state_dict(), 'mnist_cnn_best.pth') print('best acc:', best_acc)这里有个细节:torch.save(model.state_dict(), ...)保存的是参数字典,不是整个模型对象。这样做的好处是加载时不依赖原来的类定义路径,只要你有SmallCNN这个类就能恢复;坏处是你必须保留模型定义代码。如果你想要“一个文件走天下”,可以用torch.save(model, ...)保存整个对象,但跨版本加载容易出兼容问题,我一般不用。
3.2 模型文件怎么存、怎么读、怎么验证
训练完你会得到一个mnist_cnn_best.pth,通常几百 KB 到 1 MB 出头。加载它只需要三步:重建模型结构、加载参数、切到 eval 模式。
# 加载模型 model = SmallCNN().to(device) model.load_state_dict(torch.load('mnist_cnn_best.pth', map_location=device)) model.eval() # 用测试集里的一张图验证 img, label = test_set[0] with torch.no_grad(): logits = model(img.unsqueeze(0).to(device)) # 加 batch 维度 pred = logits.argmax(dim=1).item() print('真实标签:', label, '预测:', pred)map_location=device是为了在只有 CPU 的机器上也能加载 GPU 训出来的权重,不加的话会报找不到 CUDA 设备的错。img.unsqueeze(0)是因为单张图没有 batch 维度,模型 forward 里x.size(0)会取错,这是推理阶段最常见的翻车点之一。
提示:如果你要把模型交给别人用,建议同时给出模型定义代码和加载示例,否则对方拿到
.pth也不知道怎么还原结构。
4. 避坑与排查:MNIST 训练里最容易踩的五个坑
4.1 下载 404 或卡住不动
现象:执行datasets.MNIST(download=True)时报 HTTP 404,或者进度条长时间停在 0%。原因:torchvision 默认的下载源在某些网络环境下不可达,或者本地raw目录里存在不完整的临时文件。解决:手动把四个 gz 文件放到./data/MNIST/raw/,并确认文件名完全一致;如果之前下过一半,把raw目录清空重来。这一步做完,download=False就能稳定读取。
4.2 训练准确率高但测试准确率上不去
现象:训练集准确率 99.9%,测试集只有 97%。原因:模型过拟合,或者归一化参数用错。解决:先确认Normalize用的是 MNIST 的均值和标准差,而不是 ImageNet 的;再检查是否加了 dropout;如果还不行,把 epoch 从 10 降到 5,MNIST 不需要训太久。
4.3 推理时维度报错
现象:RuntimeError: Expected 4D input (got 3D input)。原因:单张图没有 batch 维度。解决:推理前用img.unsqueeze(0)补一维,或者用DataLoader包一层。这个错误几乎每个新手都会遇到一次,记住就好。
4.4 保存的模型换台机器加载失败
现象:RuntimeError: Error(s) in loading state_dict。原因:保存和加载时模型结构不一致,比如卷积核数量改了、全连接层维度改了。解决:加载前先打印model.state_dict().keys()和保存时的 keys 对比,确保结构完全一致。如果只是想做推理,建议保存时连模型定义一起打包。
4.5 CPU 推理速度慢
现象:单张图推理要几百毫秒。原因:模型没切到 eval 模式,或者没加torch.no_grad()。解决:推理前调用model.eval(),并用with torch.no_grad():包住前向过程,速度能提升数倍。如果还嫌慢,可以把模型转成 ONNX 或 TorchScript,这是进阶做法,后面会提。
5. 进阶技巧:把 MNIST 模型变成能真正用起来的小工具
训练和保存只是第一步,真正让这个方案有价值的是“能推理”。我一般会做两件事:一是把模型导出成 TorchScript,摆脱对 Python 类定义的依赖;二是写一个最小的推理脚本,接收一张 28×28 的灰度图,输出预测数字。
先看 TorchScript 导出:
# 导出为 TorchScript,推理时不需要原始类定义 model.eval() example = torch.randn(1, 1, 28, 28).to(device) traced = torch.jit.trace(model, example) traced.save('mnist_cnn_scripted.pt') # 加载并推理 loaded = torch.jit.load('mnist_cnn_scripted.pt') with torch.no_grad(): out = loaded(example) print(out.argmax(dim=1).item())torch.jit.trace会记录一次前向的计算图,所以example的 shape 必须和真实输入一致。导出后的.pt文件可以直接在 C++ 里加载,也可以被其他语言通过 LibTorch 调用,这是把模型交给非 Python 环境的标准做法。
再给一个验证方法:拿你自己手写的数字拍照,用 OpenCV 做灰度化、二值化、缩放到 28×28,再送进模型。这一步能直接暴露模型在真实数据上的短板——MNIST 的测试集太干净了,真实手写数字的笔画粗细、倾斜角度都不同,准确率通常会掉几个点。我的习惯是每次改完模型,都拿自己写的 10 个数字测一遍,记录哪些数字容易错(通常是 4 和 9、3 和 5),再决定要不要加数据增强。
最后说一个我自己的教训:早期我总想着把准确率刷到 99.9% 再上线,结果发现真实场景里那 0.1% 的提升毫无意义,反而是一个能稳定加载、推理速度可控的模型更有价值。MNIST 手写数字识别这个方向,值不值得做?如果你是想走通深度学习落地链路,它非常值得,因为成本低、反馈快;但如果你指望它直接解决复杂的票据识别,那还需要在数据增强和模型结构上继续投入。希望帮到你。
本文还有配套的精品资源,点击获取