ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

MNIST手写数字识别实战:从PyTorch数据加载到CNN模型部署的完整流程

MNIST手写数字识别实战:从PyTorch数据加载到CNN模型部署的完整流程 简介面向计算机视觉初学者与机器学习课程设计场景这份资源以经典MNIST手写数字识别为切入点完整覆盖数据读取、归一化预处理、模型构建、参数调整、训练测试与准确率评估的闭环流程适用于期末大作业或项目实战。压缩包共9个文件包含4个gz格式的MNIST标准数据训练/测试图片及标签、2个py源码文件卷积网络实现与初始化模块、1份txt说明文档以及build、swo辅助文件整体仅11.07MB下载后免去额外找数据的麻烦。源码已经严格调试运行即可看到识别效果通过调整网络结构、学习率或批次大小可进一步理解模型调优对精度的影响。目前已有64人学习下载对希望快速上手深度学习、夯实图像分类原理并完成课程设计任务的学生是份高效实用的参考。1. 手写数字识别为什么都拿 MNIST 开刀一个 10 分类问题背后的完整工程链手写数字识别跟 MNIST 这两个词在视觉领域基本是绑定出现的。任何一本深度学习入门书、任何一节 AI 课的第一份作业几乎都是同一个题目用 Python 读入 MNIST 手写数字数据集训练一个模型把 0 到 9 的灰度图认出来。数据只有 28×28 像素、70000 张图一张普通显卡都用不满但它把数据加载、张量变换、网络设计、训练评估、推理部署这条完整链路全部串起来了。这篇文章就是照着这条链路写的先讲数据怎么完整拿到手再讲模型怎么选最后给出可复现的完整代码和踩坑记录。适合刚跑通 Python 基础、想用一个项目把 PyTorch 流程走完的人也适合需要快速搭一个图像分类基准实验的工程师。2. 数据与环境准备完整代码跑起来之前先把 MNIST 数据老老实实拿到手MNIST 数据本身不复杂但国内网络环境下torchvision默认下载源经常 404这一关先把很多人卡住了。这里给出两条路一条是让 PyTorch 自动下载另一条是手动下载四个.gz文件再离线加载。两条路的代码我都会给全实际项目里我一般直接走离线那条省得每台机器都要跟网络搏斗。2.1 用 torchvision 一行代码加载 MNIST 数据集常见做法是直接用torchvision.datasets.MNIST它内部封装了下载、解压、读取、缓存整套逻辑。先看最小可用的写法import torch from torchvision import datasets, transforms transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_data datasets.MNIST( root./data, trainTrue, transformtransform, downloadTrue ) test_data datasets.MNIST( root./data, trainFalse, transformtransform, downloadTrue ) print(f训练集大小: {len(train_data)}) print(f测试集大小: {len(test_data)}) print(f单张图片形状: {train_data[0][0].shape})这段代码的逻辑是指定根目录./data声明训练集还是测试集传入预处理流水线transform然后让downloadTrue自动补全数据。输出应该是训练集 60000 张、测试集 10000 张、单张图片形状torch.Size([1, 28, 28])。这里有两个参数需要认真理解。第一是transformToTensor()会把原始 PIL 图片从 0 到 255 的 uint8 变成 0 到 1 的 float32 张量同时把维度从 28×28 变成 1×28×28补上通道维。第二是Normalize((0.1307,), (0.3081,))这两个数字是 MNIST 全量数据的均值和方法预计算值作用是把像素分布拉成接近标准正态模型收敛会明显更快。很多新手不写 Normalize 直接训练Loss 也能降但同样的 epoch 数精度会差 1 到 2 个百分点这就是数据预处理带来的实打实差距。跑完代码去看./data目录会发现里面多了一个MNIST/raw文件夹四个.gz压缩包和四个解压后的.ubyte文件都在里面。这就是完整数据的本体后面离线加载就是利用这个目录结构。2.2 torchvision 下载 MNIST 报 404手动下载与离线加载热词里出现“torchvision 下载 mnist 会 404”这不是个例。官方源放在国外服务器上国内直连经常返回 404 或者超时。我一般绕开数据集类的自动下载自己先把四个压缩包抓下来再让downloadFalse走离线加载。完整代码如下import gzip import os import urllib.request BASE_URL https://ossci-datasets.s3.amazonaws.com/mnist RAW_DIR ./data/MNIST/raw files { train-images-idx3-ubyte.gz: None, train-labels-idx1-ubyte.gz: None, t10k-images-idx3-ubyte.gz: None, t10k-labels-idx1-ubyte.gz: None, } os.makedirs(RAW_DIR, exist_okTrue) for fname in files.keys(): dest os.path.join(RAW_DIR, fname) if os.path.exists(dest): print(f{fname} 已存在跳过) continue url f{BASE_URL}/{fname} print(f正在下载 {fname} ...) urllib.request.urlretrieve(url, dest) print(四个 gz 文件就位)这段代码就是把官网公开的四个压缩包逐个下载到MNIST/raw目录下。BASE_URL指向 MNIST 数据的镜像托管地址urllib.request.urlretrieve是 Python 自带的下载函数不依赖 wget 和 curl。如果这个镜像也访问不了还可以把BASE_URL换成能访问的 MNIST 官方源文件名保持不变即可。下载完成后有一个关键动作检查压缩包大小是否完整。训练图像包约 9.9MB训练标签约 0.03MB这是网上能查到的公开信息。很多 404 之后手动下载的包其实是 HTML 错误页PyTorch 读取时会直接崩或者解压报错所以在加载前先看一眼文件大小是很便宜的排查方式。确认无误后离线加载只需把download改为Falsetrain_data datasets.MNIST( root./data, trainTrue, transformtransform, downloadFalse )这里的奥妙在于torchvision.datasets.MNIST的构造函数只有在downloadTrue并且原始文件缺失时才发起网络请求数据文件已经存在时完全走本地读取天然支持断网环境。2.3 DataLoader 参数设置batch size、shuffle 与 num_workers数据准备好之后要用 DataLoader 包一层否则训练时只能一张图一张图地喂效率极低。我用的是下面这组参数from torch.utils.data import DataLoader batch_size 64 train_loader DataLoader( train_data, batch_sizebatch_size, shuffleTrue, num_workers2, pin_memoryTrue ) test_loader DataLoader( test_data, batch_sizebatch_size, shuffleFalse, num_workers2, pin_memoryTrue )shuffleTrue只给训练集每个 epoch 打乱样本顺序防止模型学到样本顺序的假规律测试集不需要打乱。num_workers2表示用两个子进程预取数据能掩盖磁盘读取延迟Windows 系统下如果报错先把它设回 0。pin_memoryTrue是给 CUDA 训练加速用的把数据锁页从 CPU 拷贝到 GPU 时能省一点时间CPU 训练设了也无害。DataLoader 产出的每个 batch 是四维张量形状为[64, 1, 28, 28]对应 batch 大小、通道数、高、宽。标签是形状为[64]的长整型张量。理解这个形状很重要后面定义模型时第一层输入维度就要跟它对齐。3. 模型选型与网络设计识别手写数字用 MLP 还是 CNN参数怎么定数据拿到手之后下一个问题是模型选什么。手写数字识别有两个主流路线全连接网络MLP和卷积网络CNN。很多人上来就上 ResNet其实没必要。MNIST 是 28×28 单通道小图LeNet-5 结构的 CNN 就能跑到 99% 附近MLP 认真调也能到 97%。我建议是先用 MLP 把训练流程跑通再切到 CNN 提精度两步都写在下面。3.1 全连接网络方案先把训练链路跑通全连接网络思路最直白把 28×28 的图像拉平成 784 维向量过两层线性层加 ReLU 激活最后输出 10 个数字的得分。定义如下import torch.nn as nn class MLP(nn.Module): def __init__(self): super().__init__() self.net nn.Sequential( nn.Flatten(), nn.Linear(784, 256), nn.ReLU(), nn.Linear(256, 128), nn.ReLU(), nn.Linear(128, 10) ) def forward(self, x): return self.net(x)nn.Flatten()把[64, 1, 28, 28]变成[64, 784]然后进入两个隐藏层。中间维度选了 256 和 128这是容量和速度的平衡点再大训练变慢但精度提升很小再小欠拟合明显。最后一层输出 10 个值对应 0 到 9 这 10 个类别的原始得分后面接交叉熵损失函数会自动做 softmax 归一化不需要手动加。MLP 的作用是验证整条数据加载、损失计算、反向传播链路是否正常。第一次跑如果 Loss 从 2.3 左右稳步下降说明前面代码都没问题。如果这一步就翻车问题多半不在模型而在数据没进去排查范围被一下子缩小了。3.2 用 CNN 提精度LeNet-5 结构拆解与维度推演想把手写数字识别精度推到 99% 附近还是得上卷积。经典 LeNet-5 结构专为 MNIST 设计我做了一点现代化修改换用 ReLU、加 Dropout。代码见下import torch.nn as nn class LeNet5(nn.Module): def __init__(self): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 6, kernel_size5, padding2), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(6, 16, kernel_size5), nn.ReLU(), nn.MaxPool2d(2) ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(16 * 5 * 5, 120), nn.ReLU(), nn.Dropout(0.25), nn.Linear(120, 84), nn.ReLU(), nn.Linear(84, 10) ) def forward(self, x): x self.features(x) x self.classifier(x) return x维度推演是理解卷积网络的关键。输入是[64, 1, 28, 28]。第一层卷积用 5×5 卷积核、padding2输出高宽保持 28×28通道数从 1 变 6经过 2×2 最大池化后变成 6×14×14。第二层卷积没有 padding5×5 卷积核会把 14×14 变成 10×10通道数从 6 变 16再池化变成 16×5×5。所以全连接层输入维度是 16×5×5400这就是nn.Linear(400, 120)的来历。自己改网络结构时最容易在这步出错改动卷积核大小或池化参数后全连接层的输入维度必然变但 PyTorch 不会提前告诉你是错的要等前向传播跑到那一层才报维度不匹配。我的经验是每次改完结构先拿一个随机张量过一遍模型确认输出形状再进训练循环。3.3 激活函数与 Dropout为什么这样配整个网络里 ReLU 和 Dropout 的搭配是有讲究的。ReLU 解决了深层网络的梯度消失问题比早期 LeNet 用的 tanh 收敛快得多但 ReLU 有个毛病是神经元可能“死掉”——一旦输出为负梯度就是 0再也激活不回来。在小数据集上这个风险偏低但 Dropout 的存在能把过拟合压住间接降低这种风险。Dropout 只在训练时生效它随机把 25% 的神经元输出置零强迫网络不依赖单个特征。在验证和推理阶段必须把它关闭否则预测结果会抖动这就是坑章节里要重点讲的model.eval()的用途。一个经验值层数不深的小网络Dropout 放 0.25 到 0.5 之间即可放太大模型欠拟合Loss 在训练集上都降不下去这时候不要怀疑学习率先怀疑 Dropout 是不是太狠了。4. 训练与验证的完整代码损失函数、优化器、batch size 怎么配合模型定义好了接下来是训练主循环。这一章给出的代码就是可以直接复制运行的最小完整实现包含训练、验证、模型保存三件事。我会把每个关键参数讲清楚并说明改参数后最可能出现的副作用。4.1 训练主循环完整可运行代码下面的代码在 CPU 上几分钟就能完成 5 个 epochGPU 上更快。我用的是交叉熵损失加 Adam 优化器import torch import torch.nn as nn from torch.optim import Adam device torch.device(cuda if torch.cuda.is_available() else cpu) model LeNet5().to(device) criterion nn.CrossEntropyLoss() optimizer Adam(model.parameters(), lr0.001) epochs 5 for epoch in range(epochs): model.train() train_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() train_loss loss.item() * images.size(0) avg_train_loss train_loss / len(train_loader.dataset) model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in test_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() accuracy correct / total * 100 print(fEpoch {epoch 1}/{epochs} | 训练损失: {avg_train_loss:.4f} | 测试准确率: {accuracy:.2f}%) torch.save(model.state_dict(), mnist_cnn.pth)这段代码的核心流程是每个 epoch 先进入model.train()模式做训练遍历训练集的一个个 batch清零梯度、前向计算、算损失、反向传播、更新参数一个 epoch 结束后切换到model.eval()模式在测试集上统计预测正确率最后把所有参数保存到mnist_cnn.pth文件里。几个需要注意的细节。optimizer.zero_grad()必须在loss.backward()之前调用否则梯度会在旧梯度上累加Loss 看起来在降但方向是错的这是新手最常踩的隐性坑。loss.item()取的是 Python 标量不会保留计算图用来做日志打印images.size(0)是当前 batch 的样本数乘起来折算成该 batch 的总损失。验证阶段用torch.no_grad()包住明确告诉 PyTorch 不需要计算梯度省显存也加速。4.2 训练参数怎么调学习率、batch size、epoch 的配合参数之间不是孤立的我按重要程度排一下。第一是学习率Adam 默认 0.001 对 MNIST 这种小数据集是最稳的起点。调大 10 倍到 0.01Loss 可能前几个 batch 还降得猛后面就在 0.3 附近震荡调小 10 倍到 0.0001收敛肉眼可见地慢5 个 epoch 根本不够用。第二是 batch size64 是我在 CPU 和 GPU 上都能稳定跑的折中值。调到 256 能更充分用满显存但梯度估计更平滑、收敛更快同时每个 epoch 的更新次数变少需要更多 epoch 才能达到同等精度。第三是 epoch 数5 个是基线越多精度越高但收益递减10 到 15 个已经能把 LeNet-5 推到接近上限。这些参数的交互关系很像玄学但有个定心丸MNIST 够小参数差一点也能出结果只是精度和速度的交换比不同。我一般先用小规模试跑定位明显问题再往大调而不是一上来就追求最好的参数组合。4.3 验证与保存只看 Loss 会骗人准确率才是硬指标训练损失下降只能说明模型在训练集上拟合得越来越好不能代表泛化能力。验证集准确率才是模型能不能用的硬指标。上面代码里测试准确率如果是第一次跑到 98% 以上说明整条链路是健康的如果训练损失降到了 0.05 以下但测试准确率只有 90% 出头就是过拟合了优先调高 Dropout 或减小模型容量。保存模型用torch.save(model.state_dict(), mnist_cnn.pth)只保存参数不保存结构。这样做的好处是模型文件小、跨版本兼容性好加载时需要先实例化模型再load_state_dict。后面推理章节我会把加载和预测的代码一起给全这里先不展开。5. 复现途中常见的 5 个坑从数据 404 到模型不收敛的排查记录这一章不写理论全部是实操中容易踩进去的坑。每条按照“现象、原因、解决”的顺序展开你可以直接拿来做排查清单。5.1 torchvision 下载 MNIST 一直 404代码卡死现象datasets.MNIST(..., downloadTrue)运行时抛 HTTP 404 错误或者长时间卡在下载步骤。服务器返回的往往是一段 XML 错误信息不是正常的 gz 压缩包。原因默认下载源对部分网络环境不稳定官方源迁移过资源路径老的下载 URL 已经失效同时下载时没有超时机制一旦连接挂在半路就死等。解决放弃自动下载改用第 2 章的离线手动下载方案先把四个 gz 文件用迅雷或浏览器准备好放进MNIST/raw目录再以downloadFalse加载。这个方案完全绕开网络问题而且一次下载终身复用。5.2 训练 Loss 在 0.3 附近降不下去准确率只有 90%现象Loss 前几个 epoch 从 2.3 降到了 0.3之后再也降不动测试准确率卡在 90% 上下。原因先检查数据有没有归一化。如果不做Normalize((0.1307,), (0.3081,))像素值在 0 到 1 之间但均值 0.5、方差大会让网络参数更新路径绕来绕去。另一种可能是学习率偏大Adam 虽然自适应但 0.01 起步在这类任务上仍然容易震荡。解决确认transform里包含 Normalize并把学习率压回 0.001。这两个都排除后再把 Dropout 从 0.25 提到 0.5 看看是否过度自信。我遇到过几次所谓的不收敛最后都不是玄学就是归一化忘了写。5.3 验证集准确率忽高忽低同一个模型每次预测结果不同现象打印测试准确率时每次跑的结果不一样甚至同一张图预测两次得到不同类别。原因模型处于训练模式。Dropout层在model.train()模式下是开启的每过一次前向传播都有随机神经元被丢弃输出自然抖动model.eval()会统一关闭 Dropout 和 BN 的训练行为。解决在验证和推理前必须调用model.eval()并搭配torch.no_grad()。这两件套缺一不可前者管层行为后者管梯度计算。5.4 DataLoader 的 num_workers 设置后直接报错退出现象在 Windows 上把num_workers设为 2 或更大一运行就报 RuntimeError提示与多进程启动方式相关。原因Windows 下 DataLoader 多进程使用spawn方式启动子进程需要if __name__ __main__:保护入口普通脚本里直接写循环体就会炸。解决训练代码放在if __name__ __main__:代码块内或者直接把num_workers设成 0。小数据集上训练本身不慢数据预取带来的收益没那么大零 workers 最省心。5.5 CUDA out of memory但模型明明很小现象LeNet-5 这么小的网络也报显存不足或者跑到中间某个 epoch 突然崩掉。原因最常见的是验证阶段忘了torch.no_grad()每个测试 batch 都建了计算图显存越积越多另一种是把batch_size调到 512 甚至 1024MNIST 图虽然小但数据加载的中间副本同样吃显存。解决验证代码段加上torch.no_grad()把 batch size 调回 64。如果还崩检查是否有其他程序占着显存。CPU 训练完全不存在这个问题最多慢一点。6. 把训练好的模型用起来可视化、自定义图片推理与模型导出训练完成只是开始模型真正要能在实际场景里被调用才算落地。这一章给出三件套预测置信度可视化、用自己的手写图片验证模型、把模型导出成可部署格式。6.1 在测试集上打印每个类别的预测置信度手写数字识别在业务里往往不是要一个“它是几”的硬结论而是要知道模型有多少把握。下面这段代码直接从测试集取一张图输出 10 个类别的置信度import torch import torch.nn.functional as F model LeNet5() model.load_state_dict(torch.load(mnist_cnn.pth, map_locationcpu)) model.eval() image, label test_data[0] with torch.no_grad(): output model(image.unsqueeze(0)) probs F.softmax(output, dim1) for i, prob in enumerate(probs.squeeze().tolist()): print(f数字 {i}: {prob:.2%}) print(f真实标签: {label})image.unsqueeze(0)是把形状从[1, 28, 28]变成[1, 1, 28, 28]补上 batch 维这是模型输入的要求。F.softmax把原始得分转成概率分布10 个值加起来恰好等于 1。实际项目里如果最高置信度都不到 60%说明这张图本身质量可疑业务上应该返回“不确定”而不是硬认。6.2 用自己的手写图片推理预处理是关键拿自己拍的或者画的手写数字让模型认是验证模型泛化能力最有说服力的方式。难点不在推理本身在预处理环节。手机拍的照片尺寸大、背景杂而模型只认识 28×28 的黑底白字灰度图所以预处理要完成缩放、去背景、反色三步。完整代码如下from PIL import Image def preprocess_image(image_path): img Image.open(image_path).convert(L) # 灰度 img img.resize((28, 28), Image.Resampling.LANCZOS) # 缩放到 28x28 import numpy as np arr np.array(img, dtypenp.float32) # 白字黑底MNIST 是黑底白字若图片是白底黑字则反转 if arr.mean() 128: arr 255.0 - arr arr arr / 255.0 arr (arr - 0.1307) / 0.3081 tensor torch.from_numpy(arr).unsqueeze(0).unsqueeze(0) return tensor这段代码有个关键判断计算整张图的平均像素值如果偏亮白底黑字就做反色变成 MNIST 模型期望的黑底白字。均值是否超过 128 是一个人工规则对大多数手写图够用。推理时调用preprocess_image后直接把张量喂给模型输出取argmax。6.3 把模型导出为 TorchScript脱离 Python 也能部署如果想把模型接到 C 服务或者移动端PyTorch 模型不能直接跨语言使用TorchScript 是官方推荐的中间格式。导出代码只有三行scripted torch.jit.script(model.cpu()) scripted.save(mnist_cnn.pt)导出后在 Python 侧验证一次确保逻辑没有变化loaded torch.jit.load(mnist_cnn.pt) with torch.no_grad(): result loaded(image.unsqueeze(0)) print(TorchScript 模型输出:, torch.argmax(result, dim1).item())我习惯把 TorchScript 文件和训练好的.pth参数文件一起归档版本号写在文件名里比如mnist_cnn_v3.pth和mnist_cnn_v3.pt。这样线上跑挂了、告警了随手就能找到上一次可用的版本回滚这是我在项目里养成的后悔药习惯几次救过大命。整个流程走到这里从数据下载、模型训练、坑位排查到推理部署就闭环了。拿 MNIST 练手写的这套思路换成 CIFAR-10、猫狗分类或者更业务的分类任务流程骨架完全一样差别只在数据预处理和网络结构上。这个方向值得投入花一个周末把链路跑通后面再做任何视觉分类项目都会顺手很多。希望帮到你。本文还有配套的精品资源点击获取
返回列表