ARTICLE DETAIL

资讯详情

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

PyTorch实战:MNIST手写数字识别CNN源码解析与踩坑指南

PyTorch实战:MNIST手写数字识别CNN源码解析与踩坑指南 简介基于Python深度学习实现MNIST手写数据集识别的完整源码与数据包面向计算机、电子信息工程、数学等专业学生尤其适合课程设计、期末大作业或毕业设计阶段作为参考资料。压缩包共18个文件包含5个Python源码、5个pyc编译缓存、3个JSON配置、2个图像数据集、2个标签文件及1个pkl模型其中pyc可提升加载速度JSON用于配置环境与调试参数idx文件为标准MNIST图像/标签数据pkl为可直接使用的序列化模型。整体大小19.77MB目录结构清晰。源码提供了从数据加载、网络结构定义、层级运算到激活与损失函数、训练与识别流程的完整闭环读者可清晰掌握卷积神经网络的实现细节并借此扩展其他图像识别任务。目前已有511人学习下载建议具备一定Python基础、能自行调试与扩展代码的学习者使用。1. 这套源码到底在解决什么问题Python 深度学习的第一个可跑通闭环如果你刚装好 Python准备试水深度学习但又不想一上来就看几百页理论那这个以 mnist 手写数据集识别为目标的源码包恰好是门槛最低的一条路径。它讲的是用深度学习里最经典的卷积网络CNN让程序认出 09 的手写数字数据是公开的 MNIST模型代码和数据集都被打包成一份可直接运行的项目。对很多人来说第一次训练出 99% 左右的测试准确率比任何教程都更能建立信心。这篇笔记会沿着“环境准备 → 数据落地 → 模型训练 → 评估推理”的顺序把每一步的参数选择和踩坑点拆开讲清楚让你拿到类似源码包时能独立跑通也能动手改网络结构做实验。2. 先把数据拿到手MNIST 下载、预处理与 DataLoader 的三个细节解压这类以 rar 形式分发的深度学习入门工程后你大概率会看到模型定义脚本、训练脚本、数据目录和一个说明文件。按照这类项目的常见组织方式数据一般不会打包在压缩包里而是首次运行时自动下载到本地也就是说跑通代码的第一步其实是把 MNIST 下载流程搞定再谈模型。在开始写数据相关代码之前确认 Python 环境是值得花两分钟做的事。命令行里先跑一下python --version能正常输出版本号说明解释器没问题接着安装依赖pip install torch torchvision numpytorch 是深度学习框架torchvision 负责提供 MNIST 数据集和常见图像变换numpy 用于数值处理。如果你打算用 GPU 训练需要去 PyTorch 官网选择对应的 CUDA 版本安装命令如果只是跑通流程CPU 版本就够了MNIST 图片只有 28×28 大小CPU 训练十轮也只需要几分钟。2.1 为什么用 PyTorch 而不是其他框架很多看过《深度学习入门》这类书就是网上大家常说的“深度学习鱼书”的读者最初是用 NumPy 从零手写了两层神经网络。手写实现能帮你理解反向传播原理但真正做实验时我建议直接用 PyTorch原因很实际torchvision 内置了 MNIST 数据集动态计算图让调试更直观而且社区里的源码和教程大多数都用它。如果你之前看过李沐老师的《动手学深度学习》会发现 PyTorch 版的代码风格和这份题目里的源码非常接近迁移成本几乎为零。另外一个现实原因是MNIST 的 TensorFlow 1.x 老教程里很多接口已经废弃照抄会踩一堆兼容性坑。PyTorch 的接口稳定性好很多同一套代码从两年前到现在基本还能直接跑。对于刚进入深度学习领域的新手来说少一个环境层面的变量就多一分跑通的概率。2.2 数据下载与本地缓存torchvision 的 404 坑和离线兜底torchvision 里加载 MNIST 的标准写法是下面的样子这段代码也是大多数源码的开头部分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 )代码逻辑不复杂root指定数据存放目录trainTrue加载训练集downloadTrue表示本地没有时自动下载。关键在transform里的两步操作ToTensor()把 PIL 图像转成张量并把像素值从 0255 缩放到 01Normalize用均值 0.1307 和标准差 0.3081 做标准化这两个值是官方统计好的 MNIST 全局统计量直接拿来用就行不需要自己算。实际执行时很多人的第一个拦路虎不是网络结构而是下载报错。最近一两年 torchvision 下载 MNIST 会 404 的情况越来越常见原因并不是代码写错了而是 MNIST 官方服务器限制了脚本直接抓取文件。提示如果你看到URLError: urlopen error [Errno 404] Not Found或Downloading ... failed基本可以确定是被 MNIST 官方源拦截了。我一般会这样兜底用浏览器手动访问 MNIST 官方页面下载四个文件分别是train-images-idx3-ubyte.gz、train-labels-idx1-ubyte.gz、t10k-images-idx3-ubyte.gz、t10k-labels-idx1-ubyte.gz然后放进本地目录./data/MNIST/raw/下面再把download参数改成False代码就能直接使用了。train_data datasets.MNIST( root./data, trainTrue, transformtransform, downloadFalse # 文件已放在 ./data/MNIST/raw/ ) test_data datasets.MNIST( root./data, trainFalse, transformtransform, downloadFalse )torchvision 的逻辑是启动时先检查raw目录下是否存在这四个 gz 文件存在就直接解压处理不会再发起网络请求。所以手动放文件的路径必须准确放错目录仍然会报错。经过去重验证./data/MNIST/raw/是它的固定检查路径不要自作主张改结构。2.3 DataLoader 参数的配置batch_size、shuffle 和 num_workers拿到 Dataset 对象之后还需要用 DataLoader 把数据包装成可迭代批次这一步直接决定了训练时的 GPU 利用率和数据读取速度from torch.utils.data import DataLoader train_loader DataLoader( train_data, batch_size64, shuffleTrue, num_workers2, pin_memoryTrue ) test_loader DataLoader( test_data, batch_size256, shuffleFalse, num_workers2, pin_memoryTrue )这里每个参数都值得说清楚batch_size64表示每批取 64 张图片太小梯度更新频繁但震荡大太大则显存占用高MNIST 这种小图用 64 或 128 比较平衡shuffleTrue只在训练集开启作用是打乱样本顺序避免模型学到数据集排列顺序里的假规律测试集不需要打乱num_workers是并行读取数据的进程数Linux 下设为 2 或 4 能明显提升数据加载速度但 Windows 下多进程会引发一些兼容问题后面避坑章节会展开说。参数推荐值作用batch_size64每批样本数影响梯度稳定性和显存占用shuffle训练 True / 测试 False打乱样本顺序打破数据排列偏置num_workersLinux 24 / Windows 0数据预取进程数Windows 多进程容易报错pin_memoryTrue开启锁页内存加速 GPU 传输pin_memoryTrue在 GPU 训练时能减少主机到显存的拷贝时间但如果用 CPU 训练这个参数没有收益保持默认也无妨。3. 模型与训练用简单 CNN 在十分钟内逼近 99% 准确率数据就绪之后进入核心部分定义网络结构并跑训练循环。MNIST 识别任务非常适合作为深度学习 CNN 的入门项目因为图像尺寸小、类别清晰一个两层卷积网络就已经有足够表达能力不需要搬出 ResNet 这种重型结构。3.1 网络结构设计两层卷积加全连接参数为什么这样定我一般会先定义一个简洁的卷积神经网络结构如下import torch import torch.nn as nn class MnistCNN(nn.Module): def __init__(self): super().__init__() self.conv_block nn.Sequential( nn.Conv2d(1, 32, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(2), ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(64 * 7 * 7, 128), nn.ReLU(), nn.Dropout(0.5), nn.Linear(128, 10), ) def forward(self, x): return self.classifier(self.conv_block(x))网络设计的逻辑要从前向传播的维度变化来理解。输入是 1×28×28 的灰度图第一个卷积层把通道数从 1 变成 32kernel_size3加padding1保证图像尺寸不变再经过MaxPool2d(2)把 28×28 缩小到 14×14第二个卷积层把 32 通道扩到 64 通道再次池化后变成 7×7。所以最终展平送入全连接层的维度是64 * 7 * 7 3136这一步的计算必须和池化结果严格对应否则训练一启动就会报维度不匹配的错误。全连接部分最后输出 10 个值对应 09 十个类别的得分。这里不需要手动加 Softmax因为后面使用的交叉熵损失函数内部已经包含了 Softmax 计算。Dropout(0.5) 放在全连接层之间用来随机丢弃一半神经元是防止过拟合的经典手段。3.2 训练循环与损失函数交叉熵加 Adam 的搭配逻辑模型定义好之后训练循环的写法有固定套路几乎所有 PyTorch 源码都逃不出这个框架device torch.device(cuda if torch.cuda.is_available() else cpu) model MnistCNN().to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-3) epochs 10 for epoch in range(epochs): model.train() running_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() logits model(images) loss criterion(logits, labels) loss.backward() optimizer.step() running_loss loss.item() * images.size(0) epoch_loss running_loss / len(train_loader.dataset) print(fEpoch {epoch 1}/{epochs}, Loss: {epoch_loss:.4f})几个关键操作不能省optimizer.zero_grad()必须在每次反向传播之前清空上一轮梯度否则 PyTorch 会默认累加梯度导致参数更新方向错误model.train()的作用是打开 Dropout让模型每个批次随机丢弃神经元这一步很容易漏掉漏掉后训练和评估的表现都会异常。损失函数选择CrossEntropyLoss它把模型输出的原始 logits 直接和整数标签做比较内部完成 Softmax 和负对数似然的计算。优化器选 Adam 是因为它对学习率不敏感默认的 1e-3 在 MNIST 上表现稳定省去了手动调节动量等参数的过程。对刚进入深度学习领域的新手来说这能减少很多未知变量。3.3 学习率、epoch 和 batch_size 的联动配置训练参数之间是相互影响的不能单独调一个。这里给一组我常用的起始配置并解释怎么根据现象调整学习率初始 1e-3如果训练损失震荡不下降优先降到 3e-4epoch10 轮足够MNIST 数据量不大跑太多轮容易过拟合batch_size64 起步显存充足可以试 128但不建议直接上 512小批量梯度噪声反而可能让泛化更好。如果训练到第 3 轮时 loss 还在 1.0 以上一般不是参数问题而是数据预处理错了常见的是忘记归一化导致像素值范围不对。反过来如果 loss 降得很快但测试准确率上不去那就是过拟合的信号需要加强 Dropout 或做数据增强这些在避坑章节会具体展开。4. 评估与落地准确率、混淆矩阵与模型保存训练结束后我们需要回答一个更实际的问题这个模型到底靠不靠谱只看训练集 loss 没有意义必须用模型没见过的测试集来评估。MNIST 的标准做法是拿官方划分好的 10000 张测试图片做验证这也是试卷上的“闭卷考试”。4.1 测试集评估切换 eval 模式并计算准确率评估代码和训练循环略有不同差别在于不需要计算梯度也不需要更新参数def evaluate(model, loader, device): model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) logits model(images) preds torch.argmax(logits, dim1) correct (preds labels).sum().item() total labels.size(0) return correct / total test_acc evaluate(model, test_loader, device) print(f测试集准确率: {test_acc:.4f})model.eval()的作用是关闭 Dropout使所有神经元都参与推理with torch.no_grad()关闭自动求导省去中间变量存储显存占用会大幅下降。torch.argmax(logits, dim1)按行取最大值所在的索引也就是模型预测的数字类别。在 MNIST 上用两层 CNN 训练 10 轮测试集准确率通常在 0.98 到 0.99 以上具体会因随机种子和参数略有波动。如果你的结果明显低于 0.95建议回头检查数据归一化是否生效以及训练时是否真的打开了model.train()。4.2 混淆矩阵看哪些数字容易被认错准确率只是一个总分想知道模型具体在哪些数字上犯糊涂就需要混淆矩阵。这一步可以借助 sklearn 快速实现from sklearn.metrics import confusion_matrix all_preds [] all_labels [] model.eval() with torch.no_grad(): for images, labels in test_loader: images images.to(device) logits model(images) preds torch.argmax(logits, dim1).cpu() all_preds.extend(preds.numpy()) all_labels.extend(labels.numpy()) cm confusion_matrix(all_labels, all_preds) print(cm)输出是一个 10×10 的矩阵行代表真实标签列代表预测标签对角线上的数字是正确分类的样本数。你会发现 4 和 9、3 和 8、7 和 9 这些成对数字的混淆频率明显偏高因为它们的形状确实存在相似笔迹。如果某个非对角线位置数值特别大可以针对性增加该类别的训练样本或做数据增强。4.3 模型保存与加载state_dict 和完整模型怎么选训练结束后面临一个实际工程问题模型要怎么带走PyTorch 里有两种常见做法我只推荐其中一种torch.save(model.state_dict(), mnist_cnn.pth)保存而不是保存整个模型对象。这样做的原因是state_dict只存参数不存网络结构文件小、可读性强加载时必须先手动构建出相同的网络结构然后才能把参数灌进去model MnistCNN().to(device) model.load_state_dict(torch.load(mnist_cnn.pth, map_locationdevice)) model.eval()torch.load中的map_location参数很实用如果你在 GPU 上训练换到一台没有显卡的机器推理时加上map_locationcpu就能避免设备不匹配的报错。加载后记得调用model.eval()这和数据加载一样都是新手最容易漏掉的细节。4.4 用自己的手写图片做推理验证测试集准确率高只能说明模型在标准数据上表现好更有意思的验证是拿自己画出来的数字试试。用手机拍一张或者画图软件写一个数字保存成图片然后用下面这段代码推理from PIL import Image import torchvision.transforms as T def predict_image(model, image_path, device): img Image.open(image_path).convert(L) img img.resize((28, 28)) transform T.Compose([ T.ToTensor(), T.Normalize((0.1307,), (0.3081,)) ]) tensor transform(img).unsqueeze(0).to(device) model.eval() with torch.no_grad(): logits model(tensor) pred torch.argmax(logits, dim1).item() return pred这里有一个很隐蔽的坑MNIST 训练集是黑底白字而你自己在白色画板上写的字是白底黑字颜色正好相反直接推理结果大概率是错的。我一般会在推理前先做一次颜色反转把白底黑字变成黑底白字import numpy as np array np.array(img, dtypenp.float32) array 255.0 - array img Image.fromarray(array.astype(uint8))做完反转再走ToTensor和Normalize结果会可靠得多。这一步能直观感受到模型泛化能力的边界。如果你的手写数字书写风格偏草书识别错误也不意外MNIST 本身是标准手写体数据集对连笔字的容忍度有限。5. 常见问题避坑记录从 404 到过拟合一次说清代码能跑通是第一步跑通之后能稳定复现才是真本事。这里整理五个我在复现 MNIST 识别源码时真实遇到过的坑按“现象 → 原因 → 解决”的方式记录希望你能少走几段弯路。5.1 torchvision 下载 MNIST 报 404现象执行训练脚本时控制台输出URLError: urlopen error [Errno 404]然后程序退出./data目录下只有空的文件夹。原因MNIST 官方服务器对脚本下载做了限制torchvision 内置下载逻辑用的是 urllib 请求被服务器拒绝后无法拿到数据文件。解决用浏览器打开 MNIST 官方数据页面手动下载四个 gz 文件放到./data/MNIST/raw/目录下然后把所有downloadTrue改成downloadFalse。放好文件后可以先用一条 Python 命令验证from torchvision import datasets datasets.MNIST(root./data, trainTrue, downloadFalse)不报错就说明识别成功。另外如果你的网络环境访问国外站点不稳定也可以找国内镜像站下载这四个文件注意文件名字和原始官方文件名保持一致。5.2 全连接层维度不匹配现象训练代码一跑到第一个 batch 就报错错误信息里类似size mismatch for linear.linear1.weight: [128 x 3136] vs [128 x 12544]。原因nn.Linear的输入维度写错了。假设你从两层卷积改成三层卷积或者调大了池化步幅特征图尺寸不再是 7×7展平后的长度就跟着变了。维度不匹配的本质是没有按实际输出尺寸来定义全连接层。解决在写Linear之前先用一段测试代码打印特征图尺寸dummy torch.zeros(1, 1, 28, 28) with torch.no_grad(): x model.conv_block(dummy) print(x.shape) # 比如 torch.Size([1, 64, 7, 7])拿到[batch, channel, h, w]之后Linear的输入就是channel * h * w。改网络时每次都要重新验证这一步不能靠心算。5.3 忘记切换 model.eval()测试集结果忽高忽低现象准确率评估代码写了但每次跑出来的结果都不一样而且明显偏低比如只有 79% 左右训练集准确率却接近 100%。原因评估前没有调用model.eval()。这会导致模型处在训练模式Dropout 仍然随机丢弃神经元推理结果自然不稳定。解决在评估函数开头加一行model.eval()确认关闭 Dropout。同理训练循环里要记得调用model.train()否则 Dropout 永远不生效。两个模式是成对出现的漏掉一个就会出现测试结果飘忽或者训练不收敛的怪问题。顺带一提这个坑在很多开源源码里都能找到不是少数人才犯的错。5.4 过拟合训练集 98%测试集 91%现象训练 loss 一路降到 0.02 附近测试准确率却停留在 91% 左右怎么调 epoch 都上不去。原因模型记忆了训练集纹理而不是学到数字的通用特征典型的过拟合。MNIST 虽然简单但训练轮数过多、dropout 强度不够时同样会出现这个现象。解决最有效的三个手段一是检查Dropout(0.5)是否真的存在于全连接层之间二是把训练轮数从 10 降到 7 或 8 看测试集变化三是做数据增强在训练集上加入小幅随机旋转和平移。对 MNIST 来说数据增强效果很明显随机旋转 5 度左右测试准确率往往能往上走 0.5 到 1 个百分点。5.5 Windows 下 DataLoader 报 BrokenPipeError现象代码在 Linux 上正常换到 Windows 上运行程序在第一个 epoch 快结束时报错提示BrokenPipeError或DataLoader worker (pid xxx) exited unexpectedly。原因Windows 上num_workers 0时DataLoader 多进程模型和主进程之间存在资源竞争尤其是代码直接写在脚本顶层而不是封装在if __name__ __main__里时很容易触发。解决两个办法最简单的是把num_workers改成 0效率损失在 MNIST 这种小数据集上可以忽略另一个是把训练入口用if __name__ __main__:包起来让多进程安全启动。我一般两种一起上因为换机器后你无法预知环境是否继续触发。6. 进阶技巧与最终验收数据增强、学习率衰减和手写实测跑通基础流程之后想进一步逼近 99% 以上的准确率可以加入数据增强和学习率调度这两招也符合深度学习训练的标准做法值得在源码基础上长期保留。数据增强不是盲目堆操作。MNIST 手写数字识别里语义敏感度很高水平翻转和随机裁剪这类常规增强手段反而会破坏数字语义比如把 6 翻转成 9。我常用的增强只保留轻微旋转和平移transform_train transforms.Compose([ transforms.RandomRotation(5), transforms.RandomAffine(degrees0, translate(0.05, 0.05)), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ])旋转 5 度能模拟书写歪斜平移 0.05 倍宽高能模拟落笔位置偏移这些都能提升模型鲁棒性。注意训练集和测试集的 transform 要分开定义测试集仍然只做ToTensor和Normalize不能加入任何随机扰动。学习率衰减方面我习惯把 Adam 的 1e-3 初始学习率和ReduceLROnPlateau配合使用当验证损失连续一个 epoch 不下降时学习率自动乘以 0.5scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemin, patience1, factor0.5 ) # 每个 epoch 结束后用验证集 loss 更新调度器 scheduler.step(val_loss)最后是一份适用于这个源码包的验收清单跑完照此检查第一测试集准确率不低于 0.97这是底线第二混淆矩阵对角线明显突出没有哪一对数字的混淆异常严重第三自己手写 3 到 5 张数字图片颜色反转后推理正确率让人满意。我自己第二次复现时忘了把model.train()和model.eval()区分开测试准确率一直卡在 79%折腾了一个多小时才发现是 Dropout 的锅。这种教训经历过一遍就很难再忘希望帮到你避坑。本文还有配套的精品资源点击获取
返回列表