尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

MNIST手写数字识别实战:PyTorch完整训练与推理指南

发布时间:2026/9/24 23:20:23

资讯中心
01
ARTICLE

MNIST手写数字识别实战:PyTorch完整训练与推理指南

MNIST手写数字识别实战:PyTorch完整训练与推理指南
简介这套基于Python与MNIST数据集的手写数字识别项目源码及完整数据是面向计算机相关专业学生、适合用于课程设计、期末大作业或人工智能入门实战的经典实践资源。手写数字识别作为计算机视觉领域的入门项目资源包内提供了从数据预处理归一化、中心化、模型构建逻辑回归、SVM或卷积神经网络到训练与测试评估的完整实现代码均经过严格调试下载后即可运行减少了环境配置和数据准备方面的繁琐工作。压缩包共9个文件大小约11.07MB主要包含2个Python源码文件、4个GZ格式的MNIST标准数据集文件以及说明文档与构建辅助文件文件结构清晰便于按需查阅。目前已有64人学习/下载。通过实际运行和修改该项目学习者能够直观理解图像分类的基本流程掌握参数调整、准确率评估等关键技能同时项目内附说明文档有助于快速上手既适合新手模仿学习也可作为课程设计的直接参考具有较高的实用价值。1. MNIST 手写数字识别为什么我劝你先从这份数据开始上手写数字识别绝大多数人卡住的第一件事不是“模型不会写”而是“数据没拿到、环境没跑通”。MNIST 是计算机视觉里最经典的一份入门数据6万张训练图片、1万张测试图片每张都是 28×28 的灰度手写数字内容就是 0 到 9。数据量小单张图只有 784 个像素用 Python 在普通笔记本上几分钟就能训练一轮特别适合把“数据下载、预处理、训练、评估、推理”整条链路完整走通。网上一搜“Python MNIST 手写数字识别完整代码”结果常常只给一段模型代码数据获取和预处理全跳过了复制下来根本跑不动。这篇文章按我实际落地经验来写从环境安装、数据集下载开始到写出能直接跑的完整训练与推理代码再落到新手最容易翻车的几个坑上。适合刚跑通 Python 环境、想完整复现第一个视觉模型的人也适合需要快速搭一个数字识别 demo 的工程师。2. 搭环境与拿数据跑通手写数字识别要过的两道坎2.1 用 Python 3.10 PyTorch 搭好基础环境版本选型与安装命令做手写数字识别深度学习框架我一般直接选 PyTorch原因很直接torchvision 里内置了 MNIST 数据集类不用自己写下载和解析逻辑接口也比其他框架更贴近“写论文代码”的习惯遇到问题搜索量最大、答案最多。Python 版本建议 3.10 或 3.11太老的版本装新版本 PyTorch 会有兼容问题太新的版本某些算子库可能还没跟上。如果你连 Python 环境都还没准备好先把基础解释器装好在 VSCode 里选对解释器并确认python命令能跑通再往下走。用 conda 建一个独立环境是个好习惯装坏了删掉重来就是“后悔药”conda create -n mnist python3.10 -y conda activate mnist pip install torch torchvision python -c import torch, torchvision; print(torch.__version__, torchvision.__version__)最后一行如果能正常输出两个版本号说明安装成功。注意torch和torchvision是配套发布的用pip install torch torchvision一起装时 pip 会自动选匹配版本如果你分别安装导致版本对不上最常见的报错是AttributeError: module torchvision has no attribute datasets。大版本匹配即可PyTorch 2.x 配对应版本的 torchvision下面代码都能跑。2.2 torchvision 下载 MNIST 报 404/超时手动落盘与文件校验第一次跑datasets.MNIST(root./data, downloadTrue)时很多人会遇到下载卡住、报 404 或超时的问题。这个数据文件托管在公开网络上网络波动、连接超时、服务端临时不可达都可能中断下载而 torchvision 自带的下载逻辑没有断点续传中断后残留的半截文件会导致下次运行继续报错。报 404 只是最显眼的一种更烦人的是“下载到一半不动了”和“下一次运行还从中间开始继续失败”。我的处理思路很简单不等它的下载器手动把四个文件下载好放到约定目录下然后让 torchvision 跳过下载、直接使用本地文件。MNIST 在 torchvision 里的预期文件路径是data/MNIST/raw/四个文件是文件名大小参考内容train-images-idx3-ubyte.gz约 9.5 MB6 万张训练图片28×28 灰度train-labels-idx1-ubyte.gz约 28 KB6 万个训练标签t10k-images-idx3-ubyte.gz约 1.6 MB1 万张测试图片t10k-labels-idx1-ubyte.gz约 4.5 KB1 万个测试标签下载完成后我习惯先跑一个校验脚本核对文件字节数是否完整。MNIST 这四个文件的字节数是公开且长期稳定的如果大小不对基本就是下了一半或者文件损坏import os raw_dir data/MNIST/raw expected { train-images-idx3-ubyte.gz: 9912422, train-labels-idx1-ubyte.gz: 28881, t10k-images-idx3-ubyte.gz: 1648877, t10k-labels-idx1-ubyte.gz: 4542, } for name, size in expected.items(): path os.path.join(raw_dir, name) real os.path.getsize(path) if os.path.exists(path) else 0 status ok if real size else bad/missing print(f{name}: {real} bytes - {status})这段脚本的作用是把四个文件和字典里的大小逐项比对任何一项不是ok就直接定位到具体文件不要等训练时才发现数据读取异常。torchvision 的 MNIST 类在检测到raw目录下已有同名文件时会跳过下载直接解压所以手动放置文件是最稳的做法。注意文件名必须与列表完全一致包括中间的连字符和.gz后缀改动一个字符torchvision 都会把它当成新文件再尝试下载一遍。3. 数据预处理与 DataLoader把 28×28 的像素变成模型能吃的张量3.1 ToTensor 与 Normalize灰度图归一化与通道维度的来龙去脉MNIST 原始数据是 0 到 255 的灰度像素值PIL 读出来是 (28, 28) 的二维数组。模型不能直接吃这种格式必须先转成张量并归一化。torchvision 提供了一套标准预处理方案我直接抄作业from torchvision import transforms transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ])ToTensor做两件事把 PIL 图像或 numpy 数组从 (H, W) 变成 (C, H, W)这里就是 (1, 28, 28)因为灰度图只有一个通道同时把像素值从 0-255 缩放到 0-1。Normalize再按公式(x - mean) / std把像素分布调整为均值为 0、方差为 1其中 0.1307 和 0.3081 是 MNIST 全量数据的像素均值与标准差属于公开的统计先验。为什么非得归一化神经网络对输入尺度非常敏感。输入范围如果是 0-255第一层加权求和后的数值会偏大激活函数容易进入饱和区梯度更新不稳定归一到均值为 0、方差为 1 之后各层输出的分布更接近模型初始化时的假设收敛路径会顺很多。这点在很多教程里被一笔带过但实际训练时你会发现不做归一化 loss 下降明显变慢甚至震荡。3.2 DataLoader 参数batch_size、shuffle 与 num_workers 怎么设数据转换做好后用DataLoader把数据集包一层按批次取数据。我常用的配置如下from torch.utils.data import DataLoader from torchvision import datasets train_dataset datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) test_dataset datasets.MNIST(root./data, trainFalse, downloadTrue, transformtransform) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue, num_workers2) test_loader DataLoader(test_dataset, batch_size256, shuffleFalse, num_workers2)三个参数值得细说参数我的设置理由batch_size训练 64测试 25664 是 CPU 和 GPU 上都很稳的起点测试不需要反传梯度可以开大一点加快评估shuffle训练 True测试 False训练集不打乱的话连续多个 batch 可能是同一个数字梯度更新会产生偏向测试集只是统计结果顺序无所谓num_workersWindows 用 0 或 2子进程加载数据能提速但 Windows 下开多了容易报BrokenPipeError后面避坑章详细讲downloadTrue在这里是安全的因为第 2.2 节已经把文件放进raw目录它会直接使用本地文件不会真的触发网络下载。第一次运行我建议先打印一批数据的形状确认一下images应该是(64, 1, 28, 28)labels是(64,)不对的话检查预处理。3.3 用 matplotlib 看一眼数据验证数据管线的第一张“后悔药”数据管线搭好后先别急着训练。用 matplotlib 可视化一批样本能立刻发现预处理错误、标签错位、图片方向等问题比训练完才发现结果离谱要省事得多import matplotlib.pyplot as plt dataiter iter(train_loader) images, labels next(dataiter) fig, axes plt.subplots(2, 5, figsize(6, 3)) for i, ax in enumerate(axes.flat): ax.imshow(images[i][0], cmapgray) ax.set_title(flabel{labels[i].item()}) ax.axis(off) plt.show()这里有一个新手常踩的坑Normalize 之后的像素值已经不再是 0-255 范围了imshow默认按当前数据范围映射颜色图片可能显示成全黑或对比度极低。这通常不代表数据坏了只是可视化时需要关掉归一化或者直接用cmapgray并设置vmin-0.5, vmax0.5之类的范围。我在这一步翻过车看到全黑图后以为数据预处理写错了排查了半天才发现只是显示问题。确认图片内容清晰、标题和数字对得上再进训练环节。4. 从零写网络与训练循环手写数字识别最小可跑完整代码4.1 两层全连接网络与 ReLU为什么不用 CNN 也能到 97%很多人一上来就搜“CNN 手写数字识别”但在这个数据集上一个简单的全连接网络已经能到 97% 左右的测试准确率。先把全连接版跑通理解训练循环的每个环节再换成卷积网络去对比效果才是性价比最高的路径。网络定义如下import torch import torch.nn as nn class Net(nn.Module): def __init__(self): super().__init__() self.fc1 nn.Linear(28 * 28, 128) self.fc2 nn.Linear(128, 64) self.fc3 nn.Linear(64, 10) def forward(self, x): x x.view(-1, 28 * 28) x torch.relu(self.fc1(x)) x torch.relu(self.fc2(x)) return self.fc3(x)输入是 784 维的像素向量先映射到 128 维再压到 64 维最后输出 10 个数对应 0 到 9 的类别得分。中间维度 128、64 没有严格公式是按问题规模拍的经验值MNIST 是简单任务不需要很大的隐藏层如果你想试试网络容量对结果的影响把 128 改成 256 或 512 对比一下即可。relu的作用是引入非线性如果不加激活函数多层线性层叠在一起等效于一层表达能力会大打折扣。输出层不接 softmax因为后面用的交叉熵损失函数内部已经做了 softmax 计算。4.2 训练循环交叉熵损失、SGD 与损失下降节奏网络定义好后训练循环是整个项目的核心。下面的代码可以直接复制运行每轮打印平均损失和测试集准确率import torch.nn as nn from torch.optim import SGD model Net() criterion nn.CrossEntropyLoss() optimizer SGD(model.parameters(), lr0.01, momentum0.9) epochs 8 for epoch in range(epochs): total_loss 0.0 for images, labels in train_loader: optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) avg_loss total_loss / len(train_loader.dataset) print(fepoch {epoch 1}/{epochs}, avg_loss{avg_loss:.4f}) model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in test_loader: outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() print(fepoch {epoch 1}/{epochs}, test_acc{correct / total:.4f}) model.train()几个关键点拆开讲。CrossEntropyLoss是多分类任务的标准损失它内部把输出层的 logits 做了 softmax 再算交叉熵所以模型最后一层不需要手动加 softmax。SGD带momentum0.9是经典组合在小数据集上比不带动量的普通 SGD 收敛更稳也比 Adam 更容易体现出“调学习率”的效果适合学习。optimizer.zero_grad()必须每步调用否则 PyTorch 会累加梯度导致更新方向错误。model.eval()和torch.no_grad()是评估时的标准操作前者把 dropout、batchnorm 切换到推理模式后者告诉 autograd 不要记录梯度省内存也加速。lr0.01是经验起点。我习惯先把 loss 曲线跑出来看下降节奏正常情况第 1 轮平均损失在 0.2 到 0.4 之间之后稳步下降8 轮之后测试准确率在 97% 到 98% 之间。如果第 1 轮损失就跌破 0.1说明学习率可能偏大如果 loss 完全不动先检查数据标签是否对齐再检查学习率是不是太小。4.3 保存与加载模型训练一次以后直接复用训练完把模型权重保存下来后面写推理脚本就不用重新训练了torch.save(model.state_dict(), mnist_fc.pt) # 需要加载时 model Net() model.load_state_dict(torch.load(mnist_fc.pt, map_locationcpu)) model.eval()我只保存state_dict而不是整个模型对象原因是权重字典不依赖模型类定义的源码位置跨文件、跨机器加载都能用如果直接torch.save(model, ...)换台机器或改动类名后很容易反序列化失败。map_locationcpu是为了防止在只有 CPU 的机器上加载 GPU 训练的权重时报显存错误。加载后记得调eval()否则权重里如果带了 dropout 之类的模块推理结果会有随机性。5. 评估指标与推理演示准确率、混淆矩阵和预测一张新图5.1 用 accuracy 与混淆矩阵读结果别只盯一个数字训练结束后的准确率只是整体概况想让模型改进有方向我一般会再画一张混淆矩阵看看谁和谁总被搞混import numpy as np model.eval() confusion np.zeros((10, 10), dtypeint) with torch.no_grad(): for images, labels in test_loader: outputs model(images) preds torch.argmax(outputs, dim1) for t, p in zip(labels.numpy(), preds.numpy()): confusion[t][p] 1 print(混淆矩阵行为真实标签列为预测标签) print(confusion)对角线上的数字是被正确分类的样本数越大越好。非对角线上如果某个值明显偏大意味着模型系统性犯同类错误。MNIST 上最常见的混淆对是 4 和 9、3 和 8、7 和 2因为它们在笔画结构上确实接近。看到这类结果不用急这是数据本身的干扰不是代码 bug。混淆矩阵还能帮你判断“准确率差不多但错误分布不同”的两个模型哪个更适合你的业务如果你的应用场景里 4 被认成 9 的代价特别高那混淆矩阵比总准确率更能指导选型。5.2 用训练好的模型识别一张本机图片从 PIL 到预测结果的完整推理代码训练完模型不是终点能识别一张自己写的数字图才算真正闭环。推理代码的关键是预处理必须和训练时完全一致否则模型看到的分布一变准确率直线下降from PIL import Image, ImageOps import torchvision.transforms as transforms def predict_image(path, model): img Image.open(path).convert(L) img img.resize((28, 28)) # 如果图片是白底黑字需要取反成黑底白字和 MNIST 训练数据方向一致 img ImageOps.invert(img) transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) tensor transform(img).unsqueeze(0) model.eval() with torch.no_grad(): outputs model(tensor) pred torch.argmax(outputs, dim1).item() probs torch.softmax(outputs, dim1).squeeze().tolist() return pred, probs pred, probs predict_image(my_digit.png, model) print(f预测结果: {pred}) print(f各类别概率: {[round(p, 4) for p in probs]})这段代码里有三个坑。第一真实图片往往是白底黑字而 MNIST 训练数据是黑底白字白色像素是数字不取反的话模型看到的颜色极性反了预测基本是错的。第二resize默认用双线性插值如果你训练时用的缩放方式不同输入分布会有细微偏差推理时最好固定用同一种插值算法。第三输出层是 logitsargmax取的是最大得分的类别softmax后的概率值可以作为“模型置信度”参考——如果最高概率还不到 0.5说明这张图大概率不在训练数据分布里结果不可信。注意如果图片本身就是黑底白字取反这步会弄巧成拙建议先显示原图确认背景色再决定是否调用ImageOps.invert。6. 手写数字识别的 5 个常见坑与排查手册6.1 torchvision 下载 MNIST 失败不止 404 这一种表现现象运行datasets.MNIST(downloadTrue)时要么直接报 HTTP 404要么卡在下载进度不动要么第二次运行时报“文件已存在但解压失败”。原因torchvision 的下载器不支持断点续传网络波动中断后会残留一个半截文件下次运行它认为文件已存在跳过下载但解压时才发现文件损坏。404 错误是这个环节最常见的报错但“卡住不动”和“解压失败”同样高频。解决按第 2.2 节的做法手动下载四个.gz文件到data/MNIST/raw/用校验脚本确认字节数后把downloadTrue留空或者直接删掉。如果网络环境不稳定手动下载是成功率最高的方案不要反复重试自带的下载器。6.2 loss 不降或飙升学习率、初始化与数据乱序现象训练几轮后平均损失不降反升或者 loss 剧烈震荡测试准确率一直在 10% 附近徘徊。原因最常见的是学习率太大SGD 带动量后在最优解附近来回冲其次是训练时shuffleFalse同一个 batch 里全是同一个数字梯度方向单一导致震荡再就是权重初始化异常但 PyTorch 默认初始化一般不会出问题。解决把lr从 0.01 降到 0.001 再跑一轮同时确认train_loader里shuffleTrue。如果还不行给全局设固定随机种子复现训练过程排除随机性干扰torch.manual_seed(42)6.3 训练集 99% 测试集只有 85%过拟合信号与应对现象训练集准确率已经 99%测试集却卡在 85% 上下gap 一直拉大。原因模型容量超过任务需求把训练集里的噪声和细节也背下来了。MNIST 本身不容易过拟合但如果你把网络换成大几百维的隐藏层或者训练轮数拉到 50 以上照样会出现这个 gap。另一个隐性原因是反复拿测试集调参调着调着测试集信息泄漏进超参选择里。解决先用验证集调参、测试集只做最终评估过拟合明显时给模型加 dropoutself.dropout nn.Dropout(0.2) # forward 里在 relu 之后加一行 # x self.dropout(x)Dropout(0.2)表示训练时随机丢弃 20% 的隐藏层神经元迫使网络不依赖单一通路。另外把epochs从 8 提到 20观察测试准确率到 10 轮左右是否开始下降这就是早停信号。6.4 预测自己写的数字总出错预处理不一致是最隐蔽的坑现象测试集准确率 97%拿自己手写的数字一预测就错而且错得毫无规律。原因基本可以断定是推理链路的预处理和训练时不一致。常见的有三种尺寸缩放方式不同笔画粗细差异导致归一化特征偏移白底黑字没取反。MNIST 原图是 28×28 小图字迹占图片面积较大你用手机拍一张白纸上的小字再缩到 28×28笔画会很细像素分布和训练数据差异很大。解决把推理预处理固定为和训练完全相同的管线包括resize插值方式、灰度转换、取反、归一化参数。笔画的粗细问题可以在缩放前先对图片做形态学膨胀或加粗处理或者把画布做大、字写粗一点再缩放效果通常立竿见影。6.5 num_workers 设得过大Windows 上 DataLoader 卡死现象Windows 下num_workers4或更大程序运行到 DataLoader 取数据时卡住无输出或者直接报BrokenPipeError。原因Windows 的multiprocessing默认走spawn方式子进程会重新执行主模块导入逻辑和 Linux 的fork行为不一样worker 一多就容易卡死或报错。解决Windows 上把num_workers设为 0 或 2这是最省事的方案同时把主逻辑包在if __name__ __main__:里这能规避大部分 spawn 机制下的执行顺序问题。如果数据加载确实是性能瓶颈再考虑换机器或者换平台而不是硬调 worker 数。7. 把模型从 demo 用到能演示整合推理脚本与下一步扩展7.1 一份“完整代码”应该怎么组织训练入口与推理入口分离很多人拿到“完整代码”后直接跑跑通就结束了等到想改数据、换模型时才发现所有逻辑搅在一起改一处动全身。我建议至少分成两个文件train.py负责训练、评估、保存权重predict.py负责加载权重、预处理图片、输出预测。公共部分比如网络结构和预处理变换抽到一个utils.py里两边 importmnist/ ├── data/MNIST/raw/ # 数据文件 ├── train.py # 训练与评估 ├── predict.py # 单张图片推理 └── utils.py # 网络定义与预处理这份组织方式的好处是训练脚本只关注 loss 曲线和权重保存推理脚本只关注图片输入和输出结果后续把全连接网络升级成 CNN只需要改utils.py里的网络类两个入口文件都不动。我习惯给推理脚本加一个--image参数让它能从命令行直接接收图片路径方便在终端里反复测不同图片。7.2 从数字识别走向更多CNN、导出与摄像头扩展全连接版跑通后下一步有三个自然的扩展方向。第一把网络换成简单的 CNN在fc1之前加两层Conv2d MaxPool2d输入改为 (1, 28, 28) 的四维张量测试准确率能到 99% 左右这也是热门搜索里“cnn minist 手写数字识别”的常规做法。第二用torch.jit.trace或torch.onnx.export把模型导出成通用格式脱离 Python 环境部署到服务端或者移动端这是从“自己演示”走向“别人能调用”的必要一步。第三用 OpenCV 的VideoCapture读取摄像头画面把画面预处理好后逐帧传给predict_image就能做一个实时手写数字识别的小工具体验感完全不同。我第一次跑通这个项目时最深的教训不是模型选得不够先进而是数据预处理没有对齐模型换了好几个准确率始终上不去。后来把每一步的输入输出都打印出来核对才定位到是取反和缩放的问题。这篇从数据到推理的完整路径希望帮你在手写数字识别这条路上少走一段弯路。本文还有配套的精品资源点击获取
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

◈

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

◐

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

▲

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。