ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

MNIST手写数字识别实战:从模型训练到保存加载的完整落地链路

MNIST手写数字识别实战:从模型训练到保存加载的完整落地链路 简介这份资源面向深度学习入门者与需要快速验证手写数字识别效果的开发者提供基于MNIST数据集训练前馈神经网络的完整方案解决从零搭建模型时环境配置繁琐、训练耗时的问题。压缩包共5个文件包含3个Python脚本与2个H5模型文件整体约1.39MB脚本覆盖数据加载、网络定义与训练流程H5文件分别保存模型参数与完整模型结构便于直接加载推理或继续微调。资源已有6118人学习下载说明其在实际练习与项目起步阶段具备较高参考价值。拿到后可直接运行脚本复现训练过程也可跳过训练环节用现成模型完成手写数字图片的预测验证同时对照代码理解前馈神经网络的基本结构与训练要点适合作为课程作业、实验报告或入门练手的轻量级素材。1. 从一次“模型训练完却不敢上线”说起MNIST 手写数字识别到底该怎么落地很多人第一次跑 MNIST 手写数字识别都会经历同一个瞬间训练脚本跑完终端打印出 99% 的准确率心里一阵激动然后……就没有然后了。模型文件躺在checkpoints/里既不知道怎么在别的机器上复现也不知道怎么把它塞进一个真实的小工具里。更尴尬的是换台机器重新torchvision.datasets.MNIST(downloadTrue)直接给你甩一个 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 方式这里有个关键点downloadTrue触发的下载地址在某些网络环境下会失败所以生产环境我一般提前把四个压缩包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 # 关键先定义 transformToTensor 会把 0-255 的像素归一化到 0-1 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) # MNIST 全局均值和标准差 ]) # root 指向本地目录downloadFalse 表示只用本地已存在的文件 train_set datasets.MNIST(root./data, trainTrue, transformtransform, downloadFalse) test_set datasets.MNIST(root./data, trainFalse, transformtransform, downloadFalse) train_loader DataLoader(train_set, batch_size128, shuffleTrue, num_workers2) test_loader DataLoader(test_set, batch_size256, shuffleFalse, num_workers2) print(len(train_set), len(test_set)) # 60000 10000这段代码里有两个参数值得说清楚。Normalize((0.1307,), (0.3081,))里的两个数是 MNIST 训练集的全局均值和标准差用它们归一化能让输入分布更接近标准正态收敛更稳如果你不做归一化模型也能训但前期 loss 下降会明显更慢。num_workers2在 Windows 上如果报错直接改成 0这是血泪经验别硬扛。2.2 模型选型为什么我推荐先上一个小 CNNMNIST 上能用的模型很多从逻辑回归到 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_size3, padding1) # - 32x28x28 self.conv2 nn.Conv2d(32, 64, kernel_size3, padding1) # - 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)padding1保证卷积后尺寸不变这样两次池化后正好是 7×7全连接层的输入维度64*7*7就是这么来的。Dropout(0.25)放在全连接之后是为了抑制过拟合MNIST 数据量不大不加 dropout 训练集准确率会明显高于测试集。如果你把卷积核改成 5×5那padding要改成 2否则尺寸对不上这是新手最容易翻车的地方。3. 训练脚本怎么写从 loss 曲线到模型文件落盘3.1 训练循环与三个必调参数训练循环本身不复杂但有几个参数直接决定你能不能复现出 99%。我把它们列成表方便你对照调整参数推荐值作用与调整建议学习率 lr1e-3Adam 的默认值太大 loss 震荡太小收敛慢batch_size128太小梯度噪声大太大显存吃紧且泛化略差epoch5~8MNIST 上 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(), lr1e-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(dim1) correct (pred labels).sum().item() total labels.size(0) acc correct / total print(fepoch {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_locationdevice)) model.eval() # 用测试集里的一张图验证 img, label test_set[0] with torch.no_grad(): logits model(img.unsqueeze(0).to(device)) # 加 batch 维度 pred logits.argmax(dim1).item() print(真实标签:, label, 预测:, pred)map_locationdevice是为了在只有 CPU 的机器上也能加载 GPU 训出来的权重不加的话会报找不到 CUDA 设备的错。img.unsqueeze(0)是因为单张图没有 batch 维度模型 forward 里x.size(0)会取错这是推理阶段最常见的翻车点之一。提示如果你要把模型交给别人用建议同时给出模型定义代码和加载示例否则对方拿到.pth也不知道怎么还原结构。4. 避坑与排查MNIST 训练里最容易踩的五个坑4.1 下载 404 或卡住不动现象执行datasets.MNIST(downloadTrue)时报 HTTP 404或者进度条长时间停在 0%。原因torchvision 默认的下载源在某些网络环境下不可达或者本地raw目录里存在不完整的临时文件。解决手动把四个 gz 文件放到./data/MNIST/raw/并确认文件名完全一致如果之前下过一半把raw目录清空重来。这一步做完downloadFalse就能稳定读取。4.2 训练准确率高但测试准确率上不去现象训练集准确率 99.9%测试集只有 97%。原因模型过拟合或者归一化参数用错。解决先确认Normalize用的是 MNIST 的均值和标准差而不是 ImageNet 的再检查是否加了 dropout如果还不行把 epoch 从 10 降到 5MNIST 不需要训太久。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(dim1).item())torch.jit.trace会记录一次前向的计算图所以example的 shape 必须和真实输入一致。导出后的.pt文件可以直接在 C 里加载也可以被其他语言通过 LibTorch 调用这是把模型交给非 Python 环境的标准做法。再给一个验证方法拿你自己手写的数字拍照用 OpenCV 做灰度化、二值化、缩放到 28×28再送进模型。这一步能直接暴露模型在真实数据上的短板——MNIST 的测试集太干净了真实手写数字的笔画粗细、倾斜角度都不同准确率通常会掉几个点。我的习惯是每次改完模型都拿自己写的 10 个数字测一遍记录哪些数字容易错通常是 4 和 9、3 和 5再决定要不要加数据增强。最后说一个我自己的教训早期我总想着把准确率刷到 99.9% 再上线结果发现真实场景里那 0.1% 的提升毫无意义反而是一个能稳定加载、推理速度可控的模型更有价值。MNIST 手写数字识别这个方向值不值得做如果你是想走通深度学习落地链路它非常值得因为成本低、反馈快但如果你指望它直接解决复杂的票据识别那还需要在数据增强和模型结构上继续投入。希望帮到你。本文还有配套的精品资源点击获取
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进