news 2026/10/10 20:23:21

基于Python的手写数字识别系统:从MNIST到卷积网络实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
基于Python的手写数字识别系统:从MNIST到卷积网络实战

简介:这份资源面向Python初学者、机器学习入门者以及需要完成课程设计的学生,提供一套完整的手写数字识别系统实现方案。核心思路是先用Windows画图软件绘制28×28像素、黑底白字的数字图像作为输入,再交由训练好的多元线性回归模型完成0~9的十分类识别,帮助读者理解从图像预处理到模型推理的完整链路。压缩包共16个文件,约251KB,包含2个Python脚本分别负责训练与测试,2个CSV文件存放权重与标签数据,9个BMP样本图像用于验证,另附设计报告Word文档、README说明及LICENSE协议,结构紧凑、便于直接运行与二次修改。目前已有4327人学习下载,说明该方案在同类课程设计中认可度较高。读者可据此获得可复现的训练与测试代码、现成数据集、完整设计报告以及模型权重文件,既能快速跑通手写数字识别流程,也能对照报告梳理多分类问题的建模思路与实现细节。

1. 手写数字识别系统:从一张 28×28 灰度图到能跑起来的 Python 工程

很多人第一次接触图像分类,都是从 MNIST 手写数字识别开始的。标题里这个「基于 Python 实现的手写数字识别系统」,说白了就是一套能把手写数字图片喂进去、吐出 0 到 9 分类结果的完整代码工程,通常包含数据加载、模型定义、训练、评估和推理几个模块。它解决的核心问题是:让你在没有 GPU 集群、没有海量标注数据的条件下,用一台普通笔记本就能把「图像分类」这条链路从头到尾跑通一遍。适合谁?刚学完 Python 基础、装过 numpy 和 cv2、想找一个能真正跑出准确率的小项目练手的人;也适合需要快速验证某个模型改动是否有效的老手,拿它当基准测试集。别小看这个数据集,它虽然简单,但把数据预处理、模型结构、训练循环、超参调节这些坑全暴露出来了,跑通它,后面换更复杂的数据集心里就有底了。

2. 先搞清楚数据长什么样:MNIST 的加载与预处理

2.1 为什么 MNIST 是 28×28 灰度图而不是彩色图

MNIST 里的每张图都是 28 像素宽、28 像素高,单通道灰度值范围 0 到 255,0 代表纯黑,255 代表纯白。手写数字识别不需要颜色信息,笔画形状才是关键特征,所以用灰度图既省计算量又够用。数据集分四份:训练集 60000 张、训练标签 60000 个、测试集 10000 张、测试标签 10000 个。常见做法是用torchvision.datasets.MNIST或tensorflow.keras.datasets.mnist直接下载,但国内网络下载可能卡住,我一般会提前把四个压缩包放到本地目录,用download=False加载。这里有个容易翻车的点:原始图片是 PIL 格式,像素值 0 到 255 的整数,直接送进网络会导致梯度爆炸,必须先归一化到 0 到 1 之间,再减均值除标准差。均值 0.1307、标准差 0.3081 是 MNIST 训练集统计出来的经验值,用这两个数做标准化,模型收敛会稳很多。

2.2 用 PyTorch 加载 MNIST 的最小可跑代码

import torch from torchvision import datasets, transforms from torch.utils.data import DataLoader # 定义预处理:转张量 + 标准化 transform = transforms.Compose([ transforms.ToTensor(), # 把 PIL 图转成 [0,1] 的 FloatTensor,形状 [1,28,28] transforms.Normalize((0.1307,), (0.3081,)) # 减均值除标准差 ]) # 加载训练集和测试集,download=False 表示用本地已下载好的数据 train_set = datasets.MNIST(root='./data', train=True, download=False, transform=transform) test_set = datasets.MNIST(root='./data', train=False, download=False, transform=transform) # 批大小设为 64,训练集打乱顺序,测试集不用打乱 train_loader = DataLoader(train_set, batch_size=64, shuffle=True, num_workers=2) test_loader = DataLoader(test_set, batch_size=1000, shuffle=False, num_workers=2) # 检查一下一个 batch 的形状 images, labels = next(iter(train_loader)) print(images.shape) # 期望输出 torch.Size([64, 1, 28, 28]) print(labels.shape) # 期望输出 torch.Size([64])

这段代码里ToTensor()做了两件事:把像素值从 0 到 255 缩放到 0 到 1,同时把维度从 HWC 转成 CHW,因为 PyTorch 卷积层要求通道维在前。Normalize的两个参数是均值和标准差,注意写法是元组,单通道就写一个值。batch_size设 64 是常见起点,显存不够就降到 32,想训练更快可以升到 128,但太大可能泛化变差。num_workers在 Windows 上有时会报错,设成 0 最稳,Linux 下可以设 2 或 4 加速数据读取。如果本地没有数据,把download改成True让它自己下,但记得检查./data/MNIST/raw目录下有没有四个文件:train-images-idx3-ubyte、train-labels-idx1-ubyte、t10k-images-idx3-ubyte、t10k-labels-idx1-ubyte,缺一个都会报错。

2.3 数据增强要不要做:手写数字场景的取舍

很多人一上来就加随机旋转、随机裁剪,结果准确率反而掉了。手写数字的笔画方向是有意义的,6 和 9 旋转一下就分不清了,所以旋转角度要控制在正负 10 度以内。常见做法是只加RandomAffine(degrees=10, translate=(0.1, 0.1)),让数字在图中轻微平移和旋转,模拟不同人的书写习惯。但如果你用的是全连接网络而不是卷积网络,数据增强带来的收益很小,因为全连接层对位置变化很敏感,增强反而增加学习难度。我一般会先不加增强跑一个基线,看测试集准确率能不能到 97% 以上,如果能,再考虑加增强冲 99%。基线都跑不到 97%,说明模型结构或训练参数有问题,加增强是治标不治本。

3. 模型选型:从全连接到卷积,到底用哪个

3.1 全连接网络为什么也能到 97% 但不够稳

最简单的做法是把 28×28 的图拉平成一个 784 维向量,接两层全连接,第一层 512 个神经元加 ReLU,第二层 10 个神经元输出 logits。这种结构在 MNIST 上训练几十轮也能到 97% 左右,但有两个硬伤:一是参数量大,784×512 就是 40 万个权重,容易过拟合;二是对平移敏感,数字往左挪两个像素,全连接层的权重就对不上了。我试过把测试集里的图整体右移 3 个像素,全连接网络准确率直接从 97% 掉到 82%,卷积网络只掉了不到 1%。所以如果你的系统要处理真实拍照的手写数字,位置不可能每次都居中,全连接网络就是个坑。

3.2 一个够用又不过时的卷积网络结构

下面这个结构是我在多个项目里验证过的,参数量不到 10 万,MNIST 测试集准确率稳定在 99.2% 以上:

import torch.nn as nn import torch.nn.functional as F class Net(nn.Module): def __init__(self): super(Net, self).__init__() # 第一个卷积层:输入1通道,输出32通道,卷积核3x3 self.conv1 = nn.Conv2d(1, 32, 3, padding=1) # 第二个卷积层:输入32通道,输出64通道,卷积核3x3 self.conv2 = nn.Conv2d(32, 64, 3, padding=1) # 最大池化层:窗口2x2,步长2 self.pool = nn.MaxPool2d(2, 2) # 全连接层:经过两次池化后,特征图大小是 7x7,通道64 self.fc1 = nn.Linear(64 * 7 * 7, 128) self.fc2 = nn.Linear(128, 10) # Dropout 层,防止过拟合 self.dropout = nn.Dropout(0.25) def forward(self, x): # 第一层:卷积 -> ReLU -> 池化,输出 [batch, 32, 14, 14] x = self.pool(F.relu(self.conv1(x))) # 第二层:卷积 -> ReLU -> 池化,输出 [batch, 64, 7, 7] x = self.pool(F.relu(self.conv2(x))) # 拉平 x = x.view(-1, 64 * 7 * 7) # 全连接 + Dropout + ReLU x = self.dropout(F.relu(self.fc1(x))) # 输出层,不加 softmax,因为 CrossEntropyLoss 内部会做 x = self.fc2(x) return x model = Net() print(sum(p.numel() for p in model.parameters())) # 打印参数量

padding=1是为了让卷积后特征图大小不变,28×28 进,28×28 出,再经过 2×2 池化变成 14×14。第二次卷积后 14×14 池化成 7×7,所以全连接层输入是 64×7×7。Dropout(0.25)放在全连接层之前,训练时随机丢弃 25% 的神经元,测试时自动关闭。注意输出层不要加 softmax,因为nn.CrossEntropyLoss内部已经包含了 log_softmax 和 NLLLoss,再加一次会导致数值不稳定。参数量打印出来大概是 42 万左右,比全连接网络小一个数量级,但准确率更高。

3.3 训练循环里必须盯住的三个参数

import torch.optim as optim device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = Net().to(device) optimizer = optim.Adam(model.parameters(), lr=0.001) criterion = nn.CrossEntropyLoss() for epoch in range(10): model.train() running_loss = 0.0 for i, (inputs, labels) in enumerate(train_loader): inputs, labels = inputs.to(device), labels.to(device) optimizer.zero_grad() # 清空上一轮梯度 outputs = model(inputs) # 前向传播 loss = criterion(outputs, labels) # 计算损失 loss.backward() # 反向传播 optimizer.step() # 更新参数 running_loss += loss.item() print(f"Epoch {epoch+1}, Loss: {running_loss/len(train_loader):.4f}")

学习率lr=0.001是 Adam 的常用起点,太大容易震荡,太小收敛慢。批大小 64 配合这个学习率,一般 5 到 10 轮就能收敛。optimizer.zero_grad()必须放在前向传播之前,否则梯度会累加,这是新手最常见的翻车点之一。损失函数用CrossEntropyLoss,它期望的输入是未归一化的 logits,标签是 0 到 9 的整数,不要自己转成 one-hot。如果训练损失降到 0.01 以下但测试准确率卡在 98% 上不去,大概率是过拟合了,可以加大 Dropout 比例或者加 L2 正则化。

4. 避坑与排查:手写数字识别系统最常见的五个翻车现场

4.1 现象:训练损失正常下降,测试准确率一直 10% 左右

原因:标签和输出对不上。常见情况是用了CrossEntropyLoss但输出层加了 softmax,或者标签被错误地转成了 one-hot 而损失函数期望整数标签。另一个可能是数据加载时shuffle设成了False,且数据集本身按类别排序,导致每个 batch 全是同一个数字。

解决:检查输出层有没有多余的 softmax,检查标签形状是不是[batch]而不是[batch, 10]。把shuffle改成True,打印一个 batch 的标签看看是不是均匀分布。

4.2 现象:训练到一半 loss 突然变成 nan

原因:学习率太大导致梯度爆炸,或者输入数据没有归一化,像素值 0 到 255 直接进网络。也有可能是log(0)的问题,但CrossEntropyLoss内部做了数值稳定处理,概率极低。

解决:先把学习率降到 0.0001 试一轮,如果 loss 正常下降再慢慢调回去。确认Normalize那一步有没有写错,均值和标准差是不是写反了。加梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5)也能兜底。

4.3 现象:Windows 上num_workers大于 0 时报BrokenPipeError

原因:Windows 的 DataLoader 多进程实现和 Linux 不同,子进程会重新导入主模块,如果代码没有放在if __name__ == '__main__':保护块里,就会无限递归创建进程。

解决:把训练代码包进if __name__ == '__main__':,或者直接把num_workers设成 0。设成 0 训练速度会慢一些,但最稳。Linux 下没这个问题,可以放心用 2 到 4。

4.4 现象:测试集准确率比训练集低 3% 以上

原因:过拟合。模型把训练集的噪声也学进去了,泛化能力差。MNIST 只有 6 万张训练图,如果模型参数量太大或者训练轮数太多,很容易过拟合。

解决:加 Dropout,加 L2 正则化(在优化器里设weight_decay=1e-4),或者减少训练轮数。也可以做数据增强,让每轮看到的图略有不同。我一般会监控测试集准确率,连续 3 轮不提升就停,别硬跑 50 轮。

4.5 现象:推理时单张图片预测结果乱跳

原因:推理时没有把模型切换到eval()模式,Dropout 和 BatchNorm 还在按训练模式工作。另外,单张图片没有做 batch 维度扩展,形状是[1,28,28]而不是[1,1,28,28],卷积层会报错或算出奇怪结果。

解决:推理前调用model.eval(),并用torch.no_grad()包住前向传播。图片预处理要和训练时完全一致,包括归一化的均值和标准差。单张图用unsqueeze(0)加一个 batch 维度。

5. 把准确率从 99% 推到 99.5% 的两个技巧

第一个技巧是学习率预热和衰减。前 2 轮用 0.0001 的小学习率预热,让模型先稳定下来,然后升到 0.001 跑 5 轮,最后 3 轮再降到 0.0001 精细调整。这个策略在 MNIST 上能稳定提升 0.2 到 0.3 个百分点。实现方式是用torch.optim.lr_scheduler.StepLR,每 3 轮把学习率乘以 0.5,或者用CosineAnnealingLR让学习率按余弦曲线平滑下降。我一般会打印每轮的学习率,确认调度器真的生效了,血泪经验是有人忘了调用scheduler.step(),结果学习率一直没变。

第二个技巧是模型集成。训练 3 到 5 个结构相同但初始化不同的模型,推理时把它们的输出 logits 平均一下,再取 argmax。单模型 99.2%,三个模型集成后能到 99.5% 左右。代价是推理时间翻三倍,如果系统对延迟不敏感,这个投入很值。集成时注意每个模型都要用eval()模式,并且用torch.no_grad()包住,否则显存会爆。下面是一个简单的集成推理代码:

def ensemble_predict(models, image_tensor): # image_tensor 形状 [1,1,28,28] logits_sum = 0 with torch.no_grad(): for model in models: model.eval() logits_sum += model(image_tensor) # 平均后取最大值的索引 return torch.argmax(logits_sum / len(models), dim=1).item()

这个函数接收一个图片张量和一组模型,返回预测的数字。注意logits_sum初始化为 0 后直接加张量,PyTorch 会自动处理类型。如果模型在 GPU 上,图片张量也要.to(device)。集成虽然简单,但有个坑:如果某个模型训练失败,准确率只有 90%,它会拖累整体表现,所以集成前先单独评估每个模型,低于 99% 的直接扔掉。

最后说个我自己的习惯:每次改完模型结构或超参,先跑一个 3 轮的快速实验,看损失下降趋势对不对,趋势对了再跑完整 10 轮。这样能省下大量等待时间,也避免在错误方向上浪费算力。手写数字识别系统虽然小,但把它跑稳、跑透,后面遇到更复杂的图像分类任务,心里就有谱了。希望帮到你。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/10 20:15:57

GitHub高Star开源远程控制工具实测与选型指南

我在很多场合都被人问过同一个问题:我需要一个能远程控制电脑的工具,到底装哪个好?先不说商业软件的授权和价格,光是GitHub上一堆开源项目就够让人挑花眼的。GitHub本身提供了一个特别方便的筛选维度——按Star数量排序。Star虽然…

作者头像 李华
网站建设 2026/10/10 20:14:08

轴承转子齿轮系统非线性动力学MATLAB仿真与故障特征分析

做旋转机械动力学仿真的人,迟早会撞上一整套连环问题:轴承转子系统怎么建模、齿轮传动的时变刚度怎么处理、裂纹故障怎么引入、非线性振动算出来之后怎么判断它是周期解还是混沌。这几个问题单独拎出来每一个都有大量文献,但真正落到MATLAB里…

作者头像 李华
网站建设 2026/10/10 20:11:03

Agent平台线上超时故障复盘:一次工具调用拖垮整个系统

开发 Agent Platform,踩了一次真实的线上超时故障下午两点半,手机连着震了七八次,全是告警群的消息。打开监控面板看到可用性从 99.99% 直线跌到 90% 附近,第一反应是模型供应商又出问题了——毕竟 Agent 平台对外的体验几乎完全绑…

作者头像 李华