PyTorch手写数字识别系统:MNIST三种CNN模型对比与GUI部署实战(附完整代码)
PyTorch手写数字识别系统:MNIST三种CNN模型对比与GUI部署实战(附完整代码)
本文基于 PyTorch 框架,使用 MNIST 数据集训练了 SimpleCNN、DeepCNN、MLP 三种神经网络模型,测试准确率分别达到 99.27%、99.52%、98.20%。系统包含完整的训练、评估、GUI 交互三个模块,支持手写输入和图片上传两种识别方式,并实现了多模型概率集成推理。
数据集:https://pan.quark.cn/s/655076a96f84#/list/share
代码:https://pan.quark.cn/s/158eb54aa740
一、什么是手写数字识别
手写数字识别(Handwritten Digit Recognition)是利用计算机视觉和深度学习技术,将手写的阿拉伯数字(0-9)图像自动分类为对应数字类别的任务。它是最早成功应用卷积神经网络(CNN)的实际问题之一,也是深度学习入门的标准教学项目。
MNIST 数据集由 Yann LeCun 等人于 1998 年发布,包含 60,000 张训练图像和 10,000 张测试图像,每张图像为 28×28 像素的灰度手写数字。截至 2024 年,MNIST 上最优模型的测试错误率已降至 0.21% 以下(即准确率超过 99.79%),是目前基准最完善、研究最深入的图像分类数据集之一。
本系统实现了三种不同复杂度的神经网络,在 MNIST 测试集上的表现如下:
| 模型 | 网络结构 | 测试准确率 | 特点 |
|---|---|---|---|
| DeepCNN | 3层卷积 + BatchNorm + 2层全连接 | 99.52% | 精度最高,推荐使用 |
| SimpleCNN | 2层卷积 + 2层全连接 | 99.27% | 结构简洁,性价比高 |
| MLP | 3层全连接感知机 | 98.20% | 基线模型,验证CNN优势 |
| 综合模型 | 三模型概率平均集成 | 99.40%+ | 降低单一模型偏差 |
二、技术栈与环境配置
系统基于以下技术栈构建:
| 组件 | 版本要求 | 用途 |
|---|---|---|
| Python | 3.8+ | 运行环境 |
| PyTorch | ≥ 2.0.0 | 深度学习框架 |
| torchvision | ≥ 0.15.0 | 数据集加载与图像变换 |
| Pillow | ≥ 9.0.0 | 图像处理(GUI预处理) |
| NumPy | ≥ 1.21.0 | 数组运算 |
| Matplotlib | ≥ 3.5.0 | 训练曲线与评估图表 |
安装依赖:
pip install -r requirements.txt
如需 GPU 加速训练,参考 PyTorch 官网 安装对应 CUDA 版本的 PyTorch。系统支持自动检测 CUDA,无 GPU 时回退到 CPU 运行。
三、项目结构
手写数字识别系统/
├── model.py # 三种模型定义(SimpleCNN / DeepCNN / MLP)
├── train.py # 训练脚本(自动下载MNIST并训练所有模型)
├── test.py # 测试评估脚本(准确率、类别分析、可视化)
├── gui.py # 图形界面(手写/上传图片实时识别)
├── draw_report_figures.py # 混淆矩阵、PR曲线等高级评估图表
├── requirements.txt # 依赖清单
├── data/ # MNIST数据集(首次运行自动下载)
├── models/ # 训练好的模型权重
│ ├── SimpleCNN.pth
│ ├── DeepCNN.pth
│ └── MLP.pth
└── img/ # 训练过程生成的可视化图表
四、三种神经网络模型设计
4.1 模型架构对比
三种模型代表了从简单到复杂的递进关系,适合理解不同网络结构对识别精度的影响。
| 维度 | SimpleCNN | DeepCNN | MLP |
|---|---|---|---|
| 卷积层数 | 2 | 3 | 0 |
| 全连接层数 | 2 | 2 | 3 |
| BatchNorm | 无 | 有 | 无 |
| Dropout | 0.25 + 0.5 | 0.25 + 0.5 | 0.3 |
| 特征提取 | 卷积+池化 | 卷积+BN+池化 | 无(直接展平) |
| 适用场景 | 轻量部署 | 高精度识别 | 教学基线 |
4.2 SimpleCNN:2层卷积网络
SimpleCNN 采用经典的 LeNet 风格架构,两个卷积层提取局部特征,最大池化降维,两个全连接层完成分类。
class SimpleCNN(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(1, 32, 3, 1) # 1→32通道, 3×3卷积
self.conv2 = nn.Conv2d(32, 64, 3, 1) # 32→64通道
self.dropout1 = nn.Dropout(0.25)
self.dropout2 = nn.Dropout(0.5)
self.fc1 = nn.Linear(9216, 128) # 64×12×12=9216 → 128
self.fc2 = nn.Linear(128, 10) # 128 → 10类
def forward(self, x):
x = F.relu(self.conv1(x))
x = F.relu(self.conv2(x))
x = F.max_pool2d(x, 2)
x = self.dropout1(x)
x = torch.flatten(x, 1)
x = F.relu(self.fc1(x))
x = self.dropout2(x)
x = self.fc2(x)
return F.log_softmax(x, dim=1)
4.3 DeepCNN:3层卷积 + BatchNorm
DeepCNN 在 SimpleCNN 基础上增加了第三层卷积和 BatchNorm 归一化层,梯度传播更稳定,训练收敛更快,最终准确率最高。
class DeepCNN(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(1, 32, 3, 1, padding=1)
self.bn1 = nn.BatchNorm2d(32)
self.conv2 = nn.Conv2d(32, 64, 3, 1, padding=1)
self.bn2 = nn.BatchNorm2d(64)
self.conv3 = nn.Conv2d(64, 128, 3, 1, padding=1)
self.bn3 = nn.BatchNorm2d(128)
self.dropout1 = nn.Dropout(0.25)
self.dropout2 = nn.Dropout(0.5)
self.fc1 = nn.Linear(128 * 3 * 3, 256)
self.fc2 = nn.Linear(256, 10)
def forward(self, x):
x = F.relu(self.bn1(self.conv1(x)))
x = F.max_pool2d(x, 2)
x = F.relu(self.bn2(self.conv2(x)))
x = F.max_pool2d(x, 2)
x = F.relu(self.bn3(self.conv3(x)))
x = F.max_pool2d(x, 2)
x = self.dropout1(x)
x = torch.flatten(x, 1)
x = F.relu(self.fc1(x))
x = self.dropout2(x)
x = self.fc2(x)
return F.log_softmax(x, dim=1)
4.4 MLP:3层全连接感知机
MLP 不使用卷积操作,直接将 28×28=784 像素展平输入三层全连接网络。作为基线模型,它验证了卷积结构在图像任务中的必要性。
class MLP(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(784, 256)
self.fc2 = nn.Linear(256, 128)
self.fc3 = nn.Linear(128, 10)
self.dropout = nn.Dropout(0.3)
def forward(self, x):
x = torch.flatten(x, 1)
x = F.relu(self.fc1(x))
x = self.dropout(x)
x = F.relu(self.fc2(x))
x = self.dropout(x)
x = self.fc3(x)
return F.log_softmax(x, dim=1)
4.5 模型注册表
系统使用统一的模型注册表管理三种架构,训练和测试脚本通过遍历该字典实现批量处理:
MODELS = {
"SimpleCNN": SimpleCNN,
"DeepCNN": DeepCNN,
"MLP": MLP,
}
五、训练流程
5.1 超参数配置
| 参数 | 值 | 说明 |
|---|---|---|
| BATCH_SIZE | 64 | 批次大小 |
| EPOCHS | 15 | 训练轮数 |
| LEARNING_RATE | 0.001 | Adam学习率 |
| 优化器 | Adam | 自适应矩估计 |
| 损失函数 | CrossEntropyLoss | 交叉熵 |
| 数据标准化 | 均值0.1307, 标准差0.3081 | MNIST统计量 |
5.2 数据预处理
MNIST 图像在输入网络前经过标准化处理,使用 MNIST 数据集的全局均值(0.1307)和标准差(0.3081)进行归一化,加速收敛:
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))
])
5.3 训练核心循环
每个模型训练 15 个 Epoch,每个 Epoch 内完成前向传播、反向传播、参数更新,并在测试集上评估:
def train_one_model(model_name, model_class):
model = model_class().to(DEVICE)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE)
for epoch in range(1, EPOCHS + 1):
# 训练阶段
model.train()
for data, target in train_loader:
data, target = data.to(DEVICE), target.to(DEVICE)
optimizer.zero_grad()
output = model(data)
loss = criterion(output, target)
loss.backward()
optimizer.step()
# 测试阶段
model.eval()
with torch.no_grad():
for data, target in test_loader:
data, target = data.to(DEVICE), target.to(DEVICE)
output = model(data)
# 计算准确率...
torch.save(model.state_dict(), f"models/{model_name}.pth")
训练完成后,系统自动生成训练曲线对比图和准确率柱状图,直观展示三个模型的收敛过程:

从训练曲线可以看出,DeepCNN 收敛最快(2-3 个 Epoch 即接近最优),SimpleCNN 次之(4-5 个 Epoch),MLP 收敛最慢且最终精度最低。这验证了更深的卷积网络在特征提取上的优势,以及 BatchNorm 对训练稳定性的提升作用。
六、模型评估
6.1 准确率对比
三个模型在 MNIST 测试集(10,000 张图像)上的最终准确率:

DeepCNN 以 99.52% 的准确率领先,比 SimpleCNN 高 0.25 个百分点,比 MLP 高 1.32 个百分点。CNN 类模型均突破 99%,而 MLP 停留在 98% 左右,说明卷积操作对图像局部特征的提取能力远优于全连接网络。
6.2 各类别准确率分析
系统对每个数字(0-9)单独统计识别准确率,生成类别准确率柱状图。以 DeepCNN 为例:

DeepCNN 在所有数字类别上的准确率均超过 99%。其中数字 1 的识别率最高(接近 100%),因为其形态简单、笔画特征明确;数字 5 和 8 的识别率相对略低,因为手写体中这两个数字的形态变异较大,容易与 3、6、9 等数字混淆。
6.3 混淆矩阵
系统通过 draw_report_figures.py 生成混淆矩阵,展示每个数字被误分类的情况:

混淆矩阵对角线表示正确分类的数量,非对角线表示误分类。DeepCNN 的混淆矩阵中,绝大多数误分类集中在 3↔5、4↔9、7↔1 等形态相似的数字对之间,这与人类视觉感知中的混淆模式一致。
6.4 Precision-Recall 曲线
系统还计算了宏平均 Precision-Recall 曲线和置信度阈值分析,帮助确定最优置信度阈值:

七、图像预处理算法
GUI 中用户上传的图片尺寸、背景、光照各不相同,直接输入模型会导致识别失败。系统实现了一套完整的预处理流程,将任意图片转换为 MNIST 标准 28×28 格式。
7.1 预处理流程
| 步骤 | 操作 | 作用 |
|---|---|---|
| 1 | 灰度转换 | 去除颜色信息,保留亮度 |
| 2 | 高斯平滑 | 消除纸张纹理和噪点 |
| 3 | 边缘背景估计 | 判断前景/背景颜色 |
| 4 | 自动反色 | 白底黑字→黑底白字 |
| 5 | Otsu阈值分割 | 自适应二值化 |
| 6 | 边界框裁剪 | 提取数字区域 |
| 7 | 等比缩放 | 缩放至20×20,保留4像素边距 |
| 8 | 灰度重心居中 | 对齐到MNIST标准位置 |
| 9 | 标准化归一化 | 均值0.1307,标准差0.3081 |
7.2 Otsu自动阈值算法
Otsu 算法通过最大化类间方差自动确定最优阈值,比固定阈值更适合手机拍照、阴影和不同纸张背景。核心实现:
# 遍历所有可能阈值,找到使类间方差最大的阈值
for value in range(256):
background_weight += histogram[value]
if background_weight == 0:
continue
foreground_weight = total - background_weight
if foreground_weight == 0:
break
background_sum += value * histogram[value]
mean_background = background_sum / background_weight
mean_foreground = (total_sum - background_sum) / foreground_weight
score = background_weight * foreground_weight * (
mean_background - mean_foreground
) ** 2
if score > best_score:
best_score = score
threshold = value
7.3 灰度重心居中
MNIST 数据集中的数字以图像重心为中心。系统计算数字像素的灰度重心,并将其平移到 28×28 图像的中心位置(13.5, 13.5),减少手写偏移对识别的影响:
mass = canvas.sum()
if mass > 0:
yy, xx = np.indices(canvas.shape)
center_x = (xx * canvas).sum() / mass
center_y = (yy * canvas).sum() / mass
shift_x = round(13.5 - center_x)
shift_y = round(13.5 - center_y)
# 执行平移...
八、GUI交互界面
8.1 界面设计
GUI 基于 Tkinter 构建,采用深色科技风格设计,左右分栏布局:

左侧输入区域:
- 280×280 像素画布,支持鼠标手写或上传图片
- 四个操作按钮:上传图片、手写输入、清除、识别
右侧结果区域:
- 预测数字大号显示(72px Consolas 字体)
- 置信度百分比
- 0-9 各数字概率分布柱状图
8.2 模型切换与集成推理
GUI 支持在三个单模型和综合模型之间切换。综合模型将三个模型的预测概率取平均,降低单一模型的偏差:
if self.ensemble_models:
model_probs = [
torch.exp(model(input_tensor)) for model in self.ensemble_models
]
probs_tensor = torch.stack(model_probs).mean(dim=0)
8.3 手写输入实现
手写模式通过监听鼠标事件,在 PIL ImageDraw 和 Tkinter Canvas 上同步绘制,确保手写图像可被模型读取:
def _on_drag(self, event):
if self.draw_mode and self.drawing and self.last_pos:
# Tkinter画布显示
self.canvas.create_line(self.last_pos[0], self.last_pos[1],
event.x, event.y, fill="white",
width=DRAW_WIDTH, capstyle="round")
# PIL图像同步绘制(供模型读取)
draw = ImageDraw.Draw(self.draw_image)
draw.line([self.last_pos[0], self.last_pos[1], event.x, event.y],
fill=255, width=DRAW_WIDTH)
九、快速上手
第一步:训练模型
python train.py
程序自动下载 MNIST 数据集(约 11MB),依次训练三个模型,权重保存至 models/ 目录,训练曲线和对比图保存至 img/ 目录。
第二步:测试评估
python test.py
对已训练模型在测试集上评估,输出每个模型的总体准确率和各类别准确率,生成可视化图表。
第三步:启动GUI
python gui.py
启动图形界面,支持手写输入或上传图片进行实时识别。
第四步:高级评估(可选)
python draw_report_figures.py
生成混淆矩阵、PR曲线、置信度阈值分析等高级评估图表。
十、FAQ
MNIST数据集是什么?
MNIST(Modified National Institute of Standards and Technology)是深度学习领域最常用的手写数字图像数据集,包含 60,000 张训练图像和 10,000 张测试图像,每张为 28×28 像素灰度图。由 Yann LeCun 等人从美国国家标准与技术研究院的原始数据集修改而来,广泛用于图像分类算法的基准测试。
为什么CNN比MLP在图像任务上表现更好?
CNN 通过卷积核在图像上滑动提取局部特征,具有平移不变性和参数共享特性,能用更少的参数学到有效的空间特征。MLP 将图像展平为一维向量,丢失了空间结构信息,且参数量随输入尺寸急剧增长。在本项目中,CNN 模型准确率均超过 99%,而 MLP 仅为 98.20%。
BatchNorm为什么能提升训练效果?
Batch Normalization 对每个 mini-batch 的特征进行标准化,使输入分布保持稳定,缓解内部协变量偏移问题。它能加速训练收敛、允许使用更大的学习率,并起到一定的正则化效果。本项目中 DeepCNN 加入了 BatchNorm,收敛速度比 SimpleCNN 快 2-3 个 Epoch。
GUI中手写识别效果不好怎么办?
将数字写在画布中央,笔画清晰粗细适中,避免过于潦草或偏离中心。系统会自动进行灰度重心居中,但严重偏离画布中心的书写仍可能影响识别。上传手机拍摄的手写图片时,确保光线均匀、背景对比度足够。
如何只训练某一个模型?
编辑 train.py 文件末尾的循环,将 MODELS.items() 替换为指定模型,例如:
for model_name, model_class in {"DeepCNN": DeepCNN}.items():
all_results[model_name] = train_one_model(model_name, model_class)
训练时下载数据集失败怎么办?
MNIST 数据集约 11MB。如遇网络问题,可手动下载后放入 data/MNIST/raw/ 目录。需要四个文件:train-images-idx3-ubyte、train-labels-idx1-ubyte、t10k-images-idx3-ubyte、t10k-labels-idx1-ubyte。
十一、总结
本系统实现了从数据准备、模型训练、性能评估到交互应用的完整流程。三种模型的对比清晰展示了网络深度、BatchNorm、卷积操作对识别精度的影响:DeepCNN 凭借三层卷积和 BatchNorm 达到 99.52% 的最高准确率,SimpleCNN 以更简洁的结构达到 99.27%,MLP 作为基线停留在 98.20%。
GUI 中的图像预处理流程(Otsu阈值、自动反色、灰度重心居中)解决了实际使用中图片格式多样的问题,多模型集成推理进一步提升了识别鲁棒性。整个系统代码结构清晰,适合作为 PyTorch 深度学习入门的完整实践项目。
关键词:PyTorch手写数字识别、MNIST数据集、CNN卷积神经网络、SimpleCNN、DeepCNN、BatchNorm、图像分类、深度学习入门、Python机器学习、GUI数字识别
更多推荐

所有评论(0)