尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Python手写数字识别实战:从MNIST数据到CNN模型训练

Python手写数字识别实战:从MNIST数据到CNN模型训练 简介面向课程设计与机器学习入门者这套基于 Python 的手写数字识别系统覆盖从模型训练到识别测试的完整流程。项目将 09 识别视为多分类问题采用多元线性回归模型包含训练脚本、测试脚本、独热编码标签与权重数据并附有 28×28 黑白数字图像可立即验证识别效果。资源共 16 个文件以 Word 设计报告、Python 源码、CSV 数据、BMP 图片和说明文档为主压缩包仅 251KB结构紧凑便于对照学习与二次开发。目前已有 4327 人学习下载。设计报告梳理了整体思路与实验步骤源码及样例数据则支持完整复现能帮助读者快速搭建手写识别演示也为课程设计文档撰写和答辩准备提供参考。1. 手写数字识别不是玩具从MNIST到能跑的Python工程化系统把一张带着手写数字的图片丢给程序让它在一秒内告诉你这是几这就是基于Python实现的手写数字识别系统。它经常被当作教科书里的MNIST入门项目但真要在自己机器上跑通、把测试集准确率稳定推到99%以上需要把数据、模型、训练、验证这条链路上的每个细节都串起来。这篇笔记会从图像分类原理讲到CNN训练再列出几个我实际踩过的坑。适合两种人刚学Python想做一个完整项目练手的以及要用Python快速实现表单数字自动录入的从业者。2. 从像素到数字手写识别系统背后的图像分类逻辑2.1 为什么手写数字识别是所有图像分类的Hello World手写数字识别在技术本质上是单字符图像分类。一张MNIST手写数字灰度图是28乘以28的像素矩阵每个像素取值范围0到255。把矩阵展平就得到784维的数值特征向量。传统机器学习路线会直接把784个像素当作特征扔给SVM、随机森林或k近邻算法。k近邻在不做任何调参时就能拿到97%左右的准确率但每预测一张图都要和上万张训练图算一次距离越到后面越慢工程落地不划算。到了深度学习时代这个任务变成了卷积神经网络的主场。卷积核能自动学习局部笔画、边缘、拐角这类低层特征再逐层组合成完整的数字结构。一个只有两层卷积的轻量CNN在MNIST上就能稳定达到99.2%以上稍微加一点归一化和Dropout就能到99.5%。这个准确率天花板很有参考意义如果某份代码在MNIST上连99%都到不了多半不是模型结构的问题而是数据预处理或训练超参出了问题。手写数字识别还常被误认为等同于OCR。实际上它只解决单字符分类不负责检测文字区域也不处理连续文本切分。但它是OCR流水线里最核心的识别子模块先把图像里的数字区域切出来再交给这套模型做分类就能完成验证码识别、票据数字校验、表单自动录入等真实场景。正在写这个源码包或者说你想照着实现这套系统时可以从这个定位倒推需要哪些模块。2.2 MNIST数据集的真实结构四份IDX文件与标签含义MNIST源自NIST手写样本库。常用版本包含60000张训练图和10000张测试图每张都是28乘以28的灰度图。原始数据不是PNG或JPG而是IDX二进制格式。理解这个格式很重要因为很多从网上下载的教程会先转成图片再让新手用文件路径读取结果训练代码和真实数据格式脱节部署时还得重新写一套加载逻辑。MNIST原始文件按用途分为四份训练图片、训练标签、测试图片、测试标签。图片文件里的每条记录是784个无符号字节按行优先排列成28乘28矩阵。标签文件里每个样本是一个uint8数字范围0到9。文件头部用大端序存储魔数和各维度尺寸。实际写解析函数时因为外部库已经处理了这些细节你通常不会直接碰二进制但一旦遇到下载损坏、文件不完整就需要回到这个格式去排查。这份数据集的设计很巧妙训练集和测试集来源不同写字的群体不完全一致天然自带一点分布差异所以测试集准确率才能代表模型的泛化能力。很多小白喜欢在训练集上反复调参把loss压得很低看训练集准确率接近百分百就以为完了结果测试集一测立刻打回原形。后面我会专门讲如何用测试集判断模型是不是真的学会了。2.3 选PyTorch而不选sklearn和TensorFlow的理由手写数字识别有非常多的实现路线。学机器学习时sklearn里一个SVM加上像素化特征就能跑起来调一下C和gamma也能达到98%但是传统算法的上限很快碰到而且特征工程要手工做换一张不同风格的图就崩。TensorFlow在生产部署方面生态完善但API变化大新手在环境配置和版本匹配上容易卡几个小时。PyTorch的直观之处在于动态图网络结构在运行时就是Python对象打印模型、打印中间张量、断点调试都顺理成章所以中小型视觉项目我一般首选PyTorch。另一个常被提到的选择是JAX社区活跃但在Windows上的支持不如PyTorch相关教程也更偏向研究论文复现。至于纯手写反向传播的网络只适合用来理解原理不适合作为系统交付。考虑到这个项目的目标是把手写数字识别跑通并能够继续扩展PyTorch是综合成本最低的方案。如果你只是交一份作业sklearn更快但你要把这个方向往深做最好从PyTorch开始。3. 环境准备与数据加载让MNIST在你的机器上跑起来3.1 Python环境配置与依赖安装拿到源码包后第一步永远是独立环境。用conda创建一个专门的环境避免把系统Python搞乱。Python 3.8到3.11都比较稳妥我建议用3.10因为主流库对它的兼容时间最长。创建并激活环境的命令如下conda create -n mnist python3.10 -y conda activate mnist pip install torch torchvision matplotlib numpy如果你的网络下载慢可以临时换国内镜像源pip install torch torchvision matplotlib numpy -i https://pypi.tuna.tsinghua.edu.cn/simple安装完成后验证一下PyTorch是否能正常导入python -c import torch; print(torch.__version__)如果输出版本号说明基础环境没问题。这里有一个常见的坑在VSCode里运行代码前需要先选择Python解释器让它指向mnist环境的Python而不是系统默认的全局Python。很多人在终端里装好了库但VSCode右下角还指着别处的解释器一运行就报ModuleNotFoundError其实不是代码问题是环境选错了。3.2 下载MNIST数据集TorchVision接口与离线文件放在哪用TorchVision自带的数据集接口加载MNIST非常省事。这里需要一次性定义好预处理流程把输入转换成张量并做归一化。MNIST训练集的全图均值和标准差大概在0.1307和0.3081这两个数值是社区反复验证过的经验值直接用就行。from torchvision import datasets, transforms transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_data datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) test_data datasets.MNIST(root./data, trainFalse, downloadTrue, transformtransform)这段代码里root是数据存放根目录trainTrue表示加载训练集downloadTrue表示如果本地没有文件就自动下载。transform参数会被作用到每一张图上ToTensor负责把像素从0到255压缩到0到1区间并把28乘以28的二维数组改成1乘以28乘以28的三维张量Normalize再按均值0.1307和标准差0.3081做标准化。标准化不是可选项它对训练收敛速度影响很大具体我会在避坑章节展开。实际运行中downloadTrue并不总是顺利。原始文件托管在Lab的开源数据页面某些网络环境下直连容易失败或卡住。常见处理方法是在另一台能访问的机器上把四个压缩包下载好然后放进当前机器上项目的data/MNIST/raw目录文件名保持原始命名再重新运行上面的代码。TorchVision检查到文件已经存在会跳过下载直接解压和读取。如果版本较新的TorchVision还有SHA256校验文件名或内容不对会立刻报错这时只用删除raw目录里的损坏文件重新放一份即可。3.3 不依赖TorchVision自己写IDX解析函数如果你不打算使用TorchVision或者想彻底搞懂数据到底长什么样可以自己写一个IDX加载函数。我以前排查一个奇怪的数据错位问题时就是靠这段代码把所有文件读出来和官方校验值对了一遍import struct import numpy as np def load_idx_images(path): with open(path, rb) as f: magic, num, rows, cols struct.unpack(IIII, f.read(16)) data np.frombuffer(f.read(), dtypenp.uint8) data data.reshape(num, rows, cols) return data def load_idx_labels(path): with open(path, rb) as f: magic, num struct.unpack(II, f.read(8)) labels np.frombuffer(f.read(), dtypenp.uint8) return labels这里的struct.unpack使用了大端序格式字符串IIII表示四个无符号整数。图片文件的魔数通常是2051标签文件的魔数通常是2049。如果你直接把PNG后缀的文件循环读进来解析就会失败。自己写加载逻辑的意义在于遇到数据损坏或形状不对时能亲手确认文件数量、图像尺寸而不是把问题掩盖在高层接口里。3.4 可视化一个批次先看数据再谈训练写完数据加载后我习惯先跑一个可视化确认数据和标签是配对的再做任何训练。数字识别任务的训练数据一旦标签错位模型几乎不可能收敛而这种错误光看loss曲线很难发现。下面这段代码从DataLoader里拿一个批次画出前10张图import matplotlib.pyplot as plt images, labels next(iter(train_loader)) fig, axes plt.subplots(2, 5, figsize(8, 4)) for i, ax in enumerate(axes.flat): ax.imshow(images[i].squeeze(), cmapgray) ax.set_title(flabel: {labels[i].item()}) ax.axis(off) plt.show()squeeze把1乘以28乘以28张量里的通道维度去掉变成28乘以28的灰度矩阵。cmapgray保证用灰度显示。如果发现图像颜色反了或者数字边缘有异常白框说明前面的预处理和实际数据不一致尽早排查比训练半天后再回看数据要省事得多。4. 训练一个CNN手写数字识别模型从网络结构到超参4.1 设计一个轻量CNN为什么28乘28的图不需要ResNetMNIST图像尺寸只有28乘以28内容又是简单笔画结构不需要搬出ResNet或VGG那种几十层的大网络。网络太深反而会在小数据集上过拟合训练也慢。我常用的轻量结构是两层卷积加两层全连接。第一层卷积从1个通道扩到32个通道提取边缘和笔画第二层卷积从32个通道扩到64个通道把局部特征组合成更抽象的模式。每个卷积层后面加BatchNorm让每层输入分布稳定再接MaxPooling把特征图尺寸减半最后展平送进全连接层。import torch import torch.nn as nn class Net(nn.Module): def __init__(self): super().__init__() self.conv1 nn.Conv2d(1, 32, 3, padding1) self.bn1 nn.BatchNorm2d(32) self.conv2 nn.Conv2d(32, 64, 3, padding1) self.bn2 nn.BatchNorm2d(64) self.pool nn.MaxPool2d(2) self.fc1 nn.Linear(64 * 7 * 7, 128) self.drop nn.Dropout(0.25) self.fc2 nn.Linear(128, 10) def forward(self, x): x self.pool(torch.relu(self.bn1(self.conv1(x)))) x self.pool(torch.relu(self.bn2(self.conv2(x)))) x x.view(x.size(0), -1) x self.drop(torch.relu(self.fc1(x))) return self.fc2(x)这里conv1的padding1保持卷积后大小不变。输入1乘以28乘以28经过第一次池化后变成32乘以14乘以14经过第二次池化后变成64乘以7乘以7。所以全连接层的输入维度是64乘以7乘以7。BatchNorm放在激活函数之前是常见做法。Dropout放在第一个全连接层之后比例0.25用来减轻过拟合。最后一层fc2输出10维向量每个值对应数字0到9的未归一化分数。4.2 数据加载器与训练循环PyTorch标准流程有了网络之后需要把数据集封装成DataLoader。DataLoader会自动按批次组合数据并在训练时打乱顺序。训练集必须打乱否则每个epoch内样本顺序固定会影响梯度更新质量。测试集不需要打乱因为我们只关心最终准确率。from torch.utils.data import DataLoader import torch.optim as optim batch_size 128 train_loader DataLoader(train_data, batch_sizebatch_size, shuffleTrue) test_loader DataLoader(test_data, batch_sizebatch_size, shuffleFalse) device torch.device(cuda if torch.cuda.is_available() else cpu) model Net().to(device) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-3)batch_size选择128这是MNIST训练中比较稳的数值。如果显存很小可以用64但不要低于32否则每个batch的梯度方差太大损失曲线会反复震荡。Adam优化器对新手友好学习率1e-3是默认值通常不需要单独调整。CrossEntropyLoss在PyTorch内部已经包含了softmax所以模型前向输出不需要手动softmax损失函数的输入是原始logits和整数标签。训练循环最核心的步骤有三个清零梯度、计算损失、反向传播更新参数。每个epoch结束时打印平均损失观察它是否持续下降。下面的代码是标准模板for epoch in range(10): model.train() running_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() running_loss loss.item() avg_loss running_loss / len(train_loader) print(fepoch {epoch 1}, avg_loss {avg_loss:.4f})model.train()这一行不能省它把BatchNorm和Dropout切到训练模式。如果不写Dropout不会生效模型效果会变差。每个batch的images输入形状是128乘以1乘以28乘以28转换到device后确保所有计算都在同一设备上。损失下降缓慢时先看数据是否归一化正确再考虑学习率设置不要一上来就换大模型。4.3 测试集评估与保存模型准确率算对才算完训练完成后需要用测试集做一次完整评估。评估阶段必须调用model.eval()关闭Dropout让BatchNorm使用训练阶段得到的均值和方差。然后包在torch.no_grad()里减少显存和内存开销。预测结果取10维输出的最大值索引就是模型认为的数字。model.eval() correct total 0 with torch.no_grad(): for images, labels in test_loader: images, labels images.to(device), labels.to(device) outputs model(images) preds outputs.argmax(dim1) correct (preds labels).sum().item() total labels.size(0) print(ftest accuracy: {correct / total:.4f}) torch.save(model.state_dict(), mnist_cnn.pth)保存state_dict而不是整个模型是最佳实践。state_dict只包含参数和缓冲区体积小且不受PyTorch小版本API变动影响。加载时先创建相同结构的Net实例再调用load_state_dict。如果加载时报unexpected key或missing key说明网络结构和保存时的结构对不上检查构造函数是否修改过。4.4 收敛细节输入范围、通道顺序与设备切换手写数字识别里最容易丢分的地方不是网络结构而是数据形状。TorchVision的ToTensor会把原始0到255的像素值缩放到0到1并自动把28乘以28的二维数组变成1乘以28乘以28的张量。如果你自己用numpy读取图像一定要手动除以255并reshape成(1,28,28)再转tensor。否则数值范围差了两个数量级loss下降曲线像一个平台怎么也上不去。设备切换同样要小心。当代码同时支持CPU和GPU时正确的做法是先把标签和图像都传到device而不是只传模型。用cpu训练整个epoch大约几分钟可以用GPU但batch_size过小反而更慢。训练时如果打了一堆warning说GPU利用率低大可不必担心MNIST这种小任务本来就是CPU友好型的。5. 避坑/常见问题/排查手写数字识别从90%到99%的关键障碍5.1 下载MNIST一直失败不是代码问题是网络问题现象执行datasets.MNIST时卡在Download进度条或直接抛出URLError、HTTPError。原因MNIST原始文件托管在某些海外站点直连速度不稳定在部分网络环境下基本下不动。TorchVision的下载代码本身没有做重试和断点续传。解决不用死磕自动下载。找一台能正常访问的机器把四个.gz文件下载下来文件名保持不变放到项目目录下的data/MNIST/raw文件夹。重新运行下载代码TorchVision检测到文件存在就会跳过下载。如果文件损坏它会重新下载或直接报校验错误此时删除坏文件重新放一份。以后做项目遇到数据下载类报错时第一反应应该是哪些第三方库给我们封装了网络下载优先用离线文件替代。5.2 训练损失降不下去测试准确率卡在94%现象loss下降到0.3左右开始震荡测试集准确率在94%到95%之间怎么调参都上不了99%。原因最常见的图省事写法是只做ToTensor不做Normalize。原始MNIST图像均值为0.1307、标准差0.3081如果输入直接是0到1的像素值模型输入分布和训练时预期分布不一致会拖慢收敛。第二个原因是网络里没有BatchNorm或者在全连接层前少了一层Dropout。解决在transform里加入Normalize((0.1307,), (0.3081,))。加完之后准确率通常能直接提升3到4个百分点。注意Normalize的均值标准差是一个一维tuple因为灰度图只有一个通道所以每个参数只写一个数值。如果是三通道图片需要三个数值。这个改动是所有优化里性价比最高的比更换模型有效多了。5.3 测试集有99%单张真实图片却预测错现象模型在MNIST测试集上跑到99.2%但拿自己写的数字或网图测试时经常识别错误甚至完全不像数字的图也在乱报。原因这是训练域和推理域不一致。MNIST训练集里的图全部是黑底白字的28乘以28灰度图数字笔画居中。课堂演示用的真实图片往往是彩色、白底黑字、扫描件还有边框和噪声分类器没见过这些分布。解决写一个预处理函数在送入模型前把任意输入图片统一成MNIST风格。先转灰度图再判断是否需要反色把白底黑字变成黑底白字然后裁剪出数字周围的空边等比缩放到20乘以20大小粘贴到28乘以28画布中央最后做归一化。这段处理是工程落地的关键模型本身反而不用改动。很多网上找的python手写数字识别代码只教训练不教这个预处理导致一换图片就翻车。5.4 GPU检测通过但训练速度反而比CPU慢现象torch.cuda.is_available()返回True每个epoch用时比论文里写的CPU训练时间还长。GPU显存占用很低但损耗明显。原因MNIST图像小模型也小单个batch的处理时间非常短。GPU启动和kernel调度的开销远大于计算收益尤其batch_size只有64时GPU几乎一直在等待数据从内存搬进显存。解决把batch_size调成128或256同时可以在DataLoader里设置num_workers2或4提升数据读取效率。Windows上num_workers建议用0否则可能报DataLoader worker的worker错误。如果一个epoch仍然要几分钟直接强制用CPU训练这种小任务CPU的压力不大训练结果完全一致没必要追求把GPU跑满。5.5 加载模型state_dict时提示missing key现象保存了state_dict下次启动程序后加载报错提示Missing key(s) in state_dict: conv1.weight。原因保存模型权重时Net类定义和加载时的Net类结构不一致。常见原因包括加载脚本里忘了把网络完整定义出来、中间改过层名或卷积核数量。解决把模型定义放到一个单独的model.py文件里训练脚本和预测脚本都import同一个类。避免在Jupyter里训练时临时改网络结构再另写一个脚本加载。同理加载时要先实例化model Net()再model.load_state_dict(torch.load(mnist_cnn.pth, map_locationdevice))。如果保存的是整个modeltorch.load拿到的是Net对象直接能调用但跨环境的兼容性不如state_dict稳定出现类型错误时优先检查是不是搞混了两种保存方式。6. 把模型变成能用的工具写一个识别单张图片的predict脚本训练模型只是万里长征一半真正让系统可用的是接收任意一张图片输出0到9结果的predict脚本。核心思路是和训练时用同一个网络结构加载权重后对输入做同样的预处理。from PIL import Image, ImageOps def preprocess(img_path): img Image.open(img_path).convert(L) img ImageOps.invert(img) # 白底黑字变黑底白字 bbox img.point(lambda p: p 128).getbbox() if bbox: img img.crop(bbox) img.thumbnail((20, 20)) canvas Image.new(L, (28, 28), 0) canvas.paste(img, ((28 - img.width) // 2, (28 - img.height) // 2)) img_t torch.from_numpy(np.array(canvas, dtypenp.float32) / 255.0) img_t img_t.unsqueeze(0).unsqueeze(0) img_t (img_t - 0.1307) / 0.3081 return img_t这段预处理的关键在于先找数字的实际边界再缩放。如果一开始就把整张图缩放到28乘以28周围边框和空白比例会跟着变形模型看到的是被压扁的数字。裁剪后再等比缩放到20乘以20并粘贴到28乘以28画布中央正好接近MNIST原始训练样本的数字占比。预测时用torch.no_grad()把预处理后的张量unsqueeze加上batch维forward输出后argmax就是最终结果。我自己的验证习惯是训练完先做三件事第一打印测试集混淆矩阵重点关注4和9、7和2这类易混组合第二挑出模型置信度最低的10个样本看它们的长相判断是标注问题还是书写风格特殊第三拿一张手机拍的、带一点歪斜的数字图测试经过上面这个预处理后如果准确率依然有明显下降再考虑增加随机旋转的数据增强。这套验证做完系统能不能上线心里就有数了。MNIST上的99%只是一个相对简单的起点真实场景里的手写数字还会遇到不同字体、不同粗细、倾斜、重叠等复杂情况。这套基于Python和PyTorch的代码是一个很好的底座你可以在它上面继续接图像预处理、数据增强、多模型集成。我最初就是靠这个项目学会了怎么从零搭一个视觉识别流程过程中不少时间花在看配置文件和处理数据集上回头想想都很值得。希望帮到你。本文还有配套的精品资源点击获取
返回列表