尧图精选

从MNIST训练到PySide6 GUI:手写数字识别全流程实战

🕒 发布时间:2026/10/1 4:51:23 📁 来源:尧图网络
简介基于Python卷积神经网络实现MNIST手写数字数据集识别并配套GUI界面的完整工程包适合计算机、电子信息工程、数学等专业学生用于课程设计、期末大作业或毕业设计参考。工程采用CNN模型完成手写数字分类并提供图形界面便于直观演示和交互操作覆盖数据加载、模型构建、训练推理与界面集成等关键环节可帮助读者快速理解深度学习项目的基本架构。压缩包共22个文件大小仅3.41MB包含3个Python源码文件模型训练、识别逻辑、GUI界面、10张数字样本图片、5个工程配置文件、1个权重文件、1个图标及markdown说明文档目录结构清晰便于按模块阅读和二次开发。说明文档对项目背景、代码结构和运行方式做了梳理适合有一定Python和深度学习基础的人群作为参考资料自助学习。目前已有614人学习下载适合需要快速上手CNN图像识别项目、实现MNIST分类GUI演示或在此基础上扩展功能的开发者。1. 为什么说这个项目是大多数人的第一个深度学习落地项目手写数字识别在深度学习里几乎等同于编程里的 Hello World。MNIST 数据集里的每一张图片都是 28x28 的灰度数字没有任何背景干扰任务目标单一把一张图正确分类成 0 到 9 中的一个数字。这样一个任务用卷积神经网络跑起来十分钟训练就能到 99% 左右的准确率不需要 GPU不需要分布式一台普通笔记本就能完成。但真正让很多人在这个项目上卡住的不是模型本身而是两件事一是数据集下载和预处理的各种幺蛾子二是把训练好的模型接到 GUI 界面上时暴露出的工程问题。很多教程只讲到模型准确率 99% 就结束了但你要交作业、做毕设或者演示给同事看一个能鼠标点几下就出结果的界面才是真正的分水岭。这个标题里的项目本质上就是一条完整的最小可落地链路数据准备 → CNN 训练 → 模型保存 → GUI 加载模型做推理。适合的人群很明确正在入门深度学习的开发者、需要交课程设计/毕业设计的学生、想把算法演示做成小工具的工程师。下文会按这条链路把每一步拆开包括可复现的代码、参数选择和踩坑记录。2. 先把 MNIST 拿捏住数据集下载、目录结构与预处理2.1 离线数据集是最靠谱的方案别在下载上耗时间很多人的第一个坑就出在数据集下载上。如果你用 PyTorch 的torchvision.datasets.MNIST它会默认尝试从官方的 yann.lecun.com 拉数据这个地址在国内经常连不上热词里那条“torchvision下载mnist会404”说的就是这件事。我一般不会去动downloadTrue这个参数而是直接先把四个文件下好放进项目目录里的data/mnist/文件夹。mkdir -p data/mnist cd data/mnist wget https://ossci-datasets.s3.amazonaws.com/mnist/train-images-idx3-ubyte.gz wget https://ossci-datasets.s3.amazonaws.com/mnist/train-labels-idx1-ubyte.gz wget https://ossci-datasets.s3.amazonaws.com/mnist/t10k-images-idx3-ubyte.gz wget https://ossci-datasets.s3.amazonaws.com/mnist/t10k-labels-idx1-ubyte.gz gzip -d *.gz这段命令把 MNIST 常见的四个文件下载到本地并解压。.gz文件是压缩过的二进制解压后得到的是 IDX 格式的原始文件PyTorch 的MNIST类可以直接读取不需要你手动解析。如果你不熟悉命令行也可以在浏览器里下载这四个文件解压后放进同样的目录结构里。注意文件名必须保持原样PyTorch 的加载器是按文件名匹配的。这个方案的好处是整个训练过程完全离线不会因为网络波动让你反复翻车而且换一台机器跑的时候直接把data目录拷过去就行。2.2 transform 里的两个参数决定模型能学到什么数据加载代码里最关键的不是路径而是transform。下面是我最常用的一套配置既能保证模型收敛稳定又不会过度增加计算量。from torchvision import datasets, transforms from torch.utils.data import DataLoader transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_dataset datasets.MNIST( root./data, trainTrue, transformtransform, downloadFalse ) test_dataset datasets.MNIST( root./data, trainFalse, transformtransform, downloadFalse ) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue, num_workers2) test_loader DataLoader(test_dataset, batch_size128, shuffleFalse, num_workers2)这里有两件值得说明的事。第一ToTensor()会把 PIL 图像或者 numpy 数组转成 0 到 1 之间的浮点张量并把通道维度提到最前面从 HxW 变成 CxHxW。第二Normalize((0.1307,), (0.3081,))用的是 MNIST 数据集的全局均值和标准差这两个值是公开的统计量。归一化之后数据大致落在 -1 到 1 之间梯度下降会更稳定收敛速度明显比不归一化快。关于num_workers在 Windows 上如果设置大于 0 有时会报BrokenPipeError那可以改成 0。在 Linux 或 macOS 上保留 2 到 4 通常没问题。shuffleTrue只在训练集上开启测试集不需要打乱因为评估时顺序无关紧要。注意downloadFalse的前提是data/mnist/下已经有解压好的四个文件。如果路径或文件名对不上加载时会报RuntimeError: Dataset not found。这时候优先检查目录结构别急着改代码。3. 把 CNN 搭到 99%网络结构、训练参数与模型保存3.1 卷积神经网络结构图两层卷积加全连接足够很多人第一次接触卷积神经网络容易被各种经典结构吓到VGG、ResNet、GoogLeNet每个听起来都很厉害。但针对 MNIST 这种 28x28 的小图一个两卷积层加全连接的小网络就足够跑出 99% 以上的准确率。结构简单的好处是训练快、容易调试而且 GUI 推理时延时很低。更深的网络在这个任务上属于杀鸡用牛刀收益非常有限。下面这个模型是 MNIST 任务的经典配置用 PyTorch 实现。import torch.nn as nn class CNN(nn.Module): def __init__(self): super(CNN, self).__init__() self.conv1 nn.Conv2d(1, 32, kernel_size3, padding1) self.conv2 nn.Conv2d(32, 64, kernel_size3, padding1) self.pool nn.MaxPool2d(2, 2) self.fc1 nn.Linear(64 * 7 * 7, 128) self.dropout nn.Dropout(0.25) self.fc2 nn.Linear(128, 10) def forward(self, x): x self.pool(torch.relu(self.conv1(x))) x self.pool(torch.relu(self.conv2(x))) x x.view(x.size(0), -1) x torch.relu(self.fc1(x)) x self.dropout(x) x self.fc2(x) return x逐一说明参数的含义。conv1的输入通道是 1因为 MNIST 是灰度图输出通道 32 意味着提取 32 种不同的特征kernel_size3是 3x3 卷积核是当前实践中最常用的尺寸小卷积核参数量少且能堆叠出大感受野。padding1保持卷积后特征图尺寸不变28x28 输入经过 conv1 后仍然是 28x28。MaxPool2d(2, 2)把尺寸减半从 28x28 变为 14x14第二次卷积池化后再减半到 7x7所以此时特征张量维度是 64x7x7这就是fc1输入维度 64*7*7 的由来。Dropout(0.25)是训练时随机丢弃 25% 的神经元防止过拟合。这个比例在 MNIST 这个小数据集上不是必须的但它能让你在增大网络宽度时不那么担心过拟合。fc2输出 10 个节点对应 10 个数字类别。需要注意的是这里最后一层没有接 softmax因为 PyTorch 的CrossEntropyLoss内部已经包含了 softmax 计算你在 forward 里再加 softmax 反而会导致结果错误。3.2 训练一个能用的模型优化器、损失函数和收敛判断模型定义好后训练代码的套路是固定的但有几个参数值得认真对待。import torch import torch.nn as nn from torch import optim model CNN() device torch.device(cuda if torch.cuda.is_available() else cpu) model.to(device) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr0.001) 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() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() # 每个 epoch 结束后在测试集上评估一次 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.data, 1) total labels.size(0) correct (predicted labels).sum().item() accuracy 100.0 * correct / total print(fEpoch {epoch1}/{epochs}, Loss: {running_loss/len(train_loader):.4f}, Accuracy: {accuracy:.2f}%) torch.save(model.state_dict(), mnist_cnn.pth)Adam优化器的初始学习率0.001是经验值。对于 MNIST 这种任务这个值几乎不需要调收敛速度快且稳定。如果你用 SGD学习率需要调大到 0.01 左右而且收敛过程会更曲折。CrossEntropyLoss是分类任务的标准选择它会同时计算 softmax 和交叉熵梯度传导效率高。训练 10 个 epoch 通常已经足够。在我的经验里第 3 个 epoch 准确率就能到 97% 以上第 6 到第 8 个 epoch 会稳定在 99% 左右再往后提升很有限。如果你发现训练集准确率很高但测试集一直上不去可以把Dropout比例从 0.25 提到 0.4或者减小全连接层的宽度到 64。最后一行torch.save(model.state_dict(), mnist_cnn.pth)只保存了模型参数没有保存整个模型结构。这样做的优点是文件更小、跨版本兼容性更好代价是加载时必须先实例化模型再载入参数。GUI 界面里要复用的就是这套流程。注意如果你要保存整个模型可以用torch.save(model, mnist_cnn_full.pth)但官方并不推荐这种方式因为依赖源文件路径代码改动后容易加载失败。统一用state_dict更干净。4. 给 CNN 加一张脸用 PySide6 搭建可交互的 GUI 界面4.1 GUI 工具选型为什么 PySide6 比 Tkinter 更值得写标题里明确要求带 GUI那摆在面前的问题就是用哪套方案。Tkinter 是 Python 自带的不需要额外安装对纯新手友好但画出来的界面风格比较老旧坐标布局调起来费劲。PySide6Qt 的 Python 绑定界面现代化布局用QVBoxLayout这类东西自动管理缩放不变形按钮和画布组件的交互响应也更好。做课程设计的话PySide6 的观感明显更专业。安装方式没什么玄学pip install PySide6如果下载慢用国内镜像站pip install PySide6 -i https://pypi.tuna.tsinghua.edu.cn/simplePySide6 安装包比较大有两百多 MB耐心等待即可。安装完成后GUI 程序的核心逻辑是界面上有一个画图区域用户用鼠标写一个数字点击“识别”按钮后程序把画布内容缩放成 28x28、转成模型需要的张量格式送入 CNN 前向推理推理结果实时显示在标签上。4.2 主窗口与手写画布的实现细节下面是一份可以直接跑通的最小 GUI 代码核心组件是一个继承自QWidget的画布和两个按钮。这份代码省略了训练部分默认你已经通过第 3 章的流程生成了mnist_cnn.pth。import sys import torch import torch.nn as nn from PySide6.QtWidgets import QApplication, QWidget, QVBoxLayout, QHBoxLayout, QPushButton, QLabel from PySide6.QtGui import QPainter, QPen, QImage, QColor from PySide6.QtCore import Qt, QPoint class CNN(nn.Module): # 与训练时完全一致的网络结构 def __init__(self): super(CNN, self).__init__() self.conv1 nn.Conv2d(1, 32, kernel_size3, padding1) self.conv2 nn.Conv2d(32, 64, kernel_size3, padding1) self.pool nn.MaxPool2d(2, 2) self.fc1 nn.Linear(64 * 7 * 7, 128) self.dropout nn.Dropout(0.25) self.fc2 nn.Linear(128, 10) def forward(self, x): x self.pool(torch.relu(self.conv1(x))) x self.pool(torch.relu(self.conv2(x))) x x.view(x.size(0), -1) x torch.relu(self.fc1(x)) x self.dropout(x) x self.fc2(x) return x class PaintBoard(QWidget): def __init__(self): super().__init__() self.setFixedSize(280, 280) self.image QImage(280, 280, QImage.Format_RGB32) self.image.fill(Qt.white) self.last_pos None def mousePressEvent(self, event): if event.button() Qt.LeftButton: self.last_pos event.position().toPoint() def mouseMoveEvent(self, event): if self.last_pos: painter QPainter(self.image) pen QPen(QColor(Qt.black), 15, Qt.SolidLine) painter.setPen(pen) painter.drawLine(self.last_pos, event.position().toPoint()) self.last_pos event.position().toPoint() self.update() def mouseReleaseEvent(self, event): self.last_pos None def paintEvent(self, event): painter QPainter(self) painter.drawImage(0, 0, self.image) def clear(self): self.image.fill(Qt.white) self.update()这段代码里有几个容易出错的地方。setFixedSize(280, 280)让画布区域固定为 280x280 像素这个尺寸不是随便定的MNIST 原始图片是 28x28280 是 28 的 10 倍做缩放时可以直接除以 10省去很多坐标换算的麻烦。画笔宽度15对应原始图片里的约 1.5 像素粗这个粗细是经过实际测试的——太细的话模型容易识别失败太粗的话数字宽度失真15 是一个比较稳的值。mouseMoveEvent里每次都要重新创建QPainter并重用上一个鼠标位置last_pos来画线这样才能保证快速移动鼠标时笔画是连续的。如果你在mousePressEvent里只画一个点快速书写时会出现断线识别准确率会明显下降。4.3 把画布内容喂给模型缩放、张量转换与推理画布里的图像是 280x280 的 RGB 图而模型要求的是 1x1x28x28 的浮点张量中间要经过一个转换函数。import numpy as np def image_to_tensor(board: PaintBoard): # 缩小到 28x28 small_img board.image.scaled(28, 28, Qt.IgnoreAspectRatio, Qt.SmoothTransformation) # 转为 numpy 数组并取灰度 ptr small_img.bits() ptr.setsize(28 * 28 * 4) arr np.array(ptr).reshape(28, 28, 4).copy() gray arr[:, :, 0].astype(np.float32) # 取红色通道黑白图下等价于灰度值 # 归一化到 0~1再反转颜色白底变为黑底 gray 1.0 - gray / 255.0 # 转换为 PyTorch 张量并加 batch 和 channel 维度 tensor torch.from_numpy(gray).unsqueeze(0).unsqueeze(0) return tensorscaled之后的图像用bits()获取原始像素指针转成 numpy 数组后取红色通道。因为画布是黑白图RGB 三个通道值相同取任意一个都行。gray 1.0 - gray / 255.0这一步是关键画布是白色背景、黑色笔画归一化后白色是 1、黑色是 0而 MNIST 数据集恰好相反——黑色背景、白色数字训练时白像素接近 1。如果不做颜色反转模型会把笔画当成背景识别结果几乎每次都是错的。有了张量之后推理就很简单了def predict(tensor): model.eval() with torch.no_grad(): output model(tensor) pred torch.argmax(output, dim1).item() confidence torch.softmax(output, dim1).max().item() return pred, confidencemodel.eval()会关闭 Dropout让推理结果稳定。torch.no_grad()关闭梯度计算省内存且推理更快。torch.argmax取概率最大的类别作为预测结果。torch.softmax把 logits 转成 0 到 1 之间的概率值方便在界面上显示“置信度”。主窗口的逻辑就是把上面这些组件串起来class MainWindow(QWidget): def __init__(self): super().__init__() self.setWindowTitle(MNIST 手写数字识别) self.board PaintBoard() self.result_label QLabel(请在上面写一个数字然后点击识别) btn_predict QPushButton(识别) btn_clear QPushButton(清空) btn_predict.clicked.connect(self.on_predict) btn_clear.clicked.connect(self.board.clear) layout QVBoxLayout() layout.addWidget(self.board) layout.addWidget(self.result_label) btn_layout QHBoxLayout() btn_layout.addWidget(btn_predict) btn_layout.addWidget(btn_clear) layout.addLayout(btn_layout) self.setLayout(layout) # 加载训练好的模型 self.model CNN() self.model.load_state_dict(torch.load(mnist_cnn.pth, map_locationcpu)) self.model.eval() def on_predict(self): tensor image_to_tensor(self.board) pred, conf predict(self.model, tensor) self.result_label.setText(f识别结果{pred}置信度{conf:.2%}) if __name__ __main__: app QApplication(sys.argv) window MainWindow() window.show() sys.exit(app.exec())clicked.connect把按钮点击事件绑定到对应的方法上。torch.load里的map_locationcpu是为了保证在只有 CPU 的机器上也能正常加载模型。完整跑起来之后你会发现一个很有意思的现象用鼠标写数字的识别率往往低于测试集的 99%因为鼠标笔画粗细、位置、数字大小都和训练集里的标准化图像有差异。这属于正常现象下面这段专门讲怎么把这个差异降到最低。5. 从跑通到稳定5 个新手最容易翻车的环节5.1 画布写数字识别不准不是模型的问题现象模型测试集准确率 99%但在 GUI 上手写 0 到 9错三四个很正常。原因测试集里的数字是经过居中、大小归一化的标准图而你在画布上随手写的数字可能偏小、偏左、笔画过粗或过细。CNN 对位置和尺寸有一定的容忍度但容忍度有限。另外鼠标书写体验本来就比笔差写出来的形状和手写体差异较大。解决一个立竿见影的办法是把画布改成 560x560然后在转换时先对图像做轮廓检测找到数字的外接矩形裁剪后等比缩放到 28x28。这样等效于帮模型做了一次“注意力对齐”识别率能提升一大截。如果不想引入 OpenCV也可以用简单的像素遍历找非白色区域的外接框代码量不大但效果明显。def crop_digit(image_np): # image_np 是 280x280x4 的 numpy 数组白底黑字 gray image_np[:, :, 0] rows np.any(gray 128, axis1) cols np.any(gray 128, axis0) if not rows.any(): return None rmin, rmax np.where(rows)[0][[0, -1]] cmin, cmax np.where(cols)[0][[0, -1]] # 向外扩展 5 像素避免笔画被切掉 rmin max(0, rmin - 5); rmax min(gray.shape[0], rmax 6) cmin max(0, cmin - 5); cmax min(gray.shape[1], cmax 6) return gray[rmin:rmax, cmin:cmax], (rmin, rmax, cmin, cmax)这个函数先找出所有非白像素的行和列范围然后裁剪出数字所在的最小矩形。扩展 5 像素是为了防止数字的边缘笔画比如 1 的底部、7 的横杠末端被硬切掉。得到裁剪图后再用cv2.resize或者 numpy 插值缩放到 28x28。这一步实现了类似 MNIST 原始数据集的居中效果比直接整图缩放科学得多。5.2 Windows 下模型加载失败报错内容却隐晦现象torch.load(mnist_cnn.pth)在 PyCharm 里直接运行没问题打包成 exe 或者换目录运行就报错。原因PyTorch 的torch.load默认依赖原训练脚本里的模块路径。如果你的模型类定义在model.py里打包时没有把这个模块一起包含进去加载就会失败。另一个常见原因是路径问题——GUI 程序用相对路径加载模型但工作目录和 GUI 脚本所在目录不一致。解决统一用state_dict格式保存模型参数并在 GUI 脚本里重新定义一次完整的模型类像我第 4 章里做的那样。路径上用Path(__file__).parent拼出模型文件的绝对路径而不是依赖os.getcwd()。from pathlib import Path model_path Path(__file__).parent / mnist_cnn.pth self.model.load_state_dict(torch.load(model_path, map_locationcpu))5.3 CPU 机器推理卡顿现象点击“识别”按钮后界面卡住一两秒才出结果。原因如果模型在 GPU 上训练时保存的state_dict参数带 CUDA 标记CPU 机器加载时map_locationcpu已经能解决。真正的卡顿往往是image_to_tensor里用bits()取指针后没有.copy()导致 numpy 数组和 QImage 内存共享后续操作触发 Qt 事件循环阻塞。解决bits()之后一定要.copy()断开共享引用。另外把模型加载和推理放到初始化时预热一次比如加载后立即跑一个全零张量把第一次推理的耗时提前消耗掉。我在自己的项目里测过预热后单次推理稳定在 5 到 10 毫秒完全感知不到卡顿。5.4 训练时 loss 不下降准确率一直在 10% 附近现象loss在初始值附近抖动准确率约等于瞎猜。原因最常见的两个一是归一化参数写错比如把均值标准差写成(0.5, 0.5)导致输入数据分布被破坏二是没有调用optimizer.zero_grad()梯度一直在累加。解决先检查transformMNIST 的标准值是官方统计的(0.1307,)和(0.3081,)。然后在每个 batch 的optimizer.step()之前确认有optimizer.zero_grad()。如果这两处都对把学习率从0.001降到0.0003再试。MNIST 上这些玄学问题基本都能通过这三个检查解决。5.5 GUI 界面中文乱码现象标题栏和按钮上的中文变成方块。原因Qt 在部分中文字体缺失的 Linux 机器上会 fallback 失败Windows 上很少出现这个问题。解决在QApplication创建后设置默认字体from PySide6.QtGui import QFont app QApplication(sys.argv) app.setFont(QFont(Microsoft YaHei, 10))如果是 Linux 服务器上跑 GUI需要先确认系统装了中文字体比如fonts-noto-cjk。这个坑在课程演示现场遇到会比较尴尬提前处理掉能省不少事。6. 让 99% 的准确率再往前走一步可视化与边界场景验证模型训练好了GUI 也能识别了接下来值得做两件事一件是让自己更信服模型的判断依据另一件是找到 GUI 的识别边界在哪。第一件事是可视化卷积核和特征图。加载训练好的模型后把第一层卷积核直接画出来你会看到它们像一些方向边缘检测器把一张输入数字的中间层特征图打印出来能看到网络在逐层抽象笔画的局部结构。很多博客里贴的“卷积神经网络结构图”就是这么来的。实现上只需要一行model.conv1.weight.data就能拿到权重torchvision.utils.make_grid可以把它们排成网格图保存成图片文件。这些图放进课程设计报告或答辩 PPT 里比一张准确率曲线更有说服力。第二件事是系统性地测一下 GUI 的边界。我会用这样的方法每个数字写 10 遍记录识别错误集中在哪个数字上。按我的经验最容易混的是 4 和 9、7 和 1。4 和 9 混淆的根源往往是你写 4 时最后一笔收得过于急促顶部开口很小放大后像 9 的小圆圈7 和 1 混淆则是因为很多人的手写 7 不带横杠。解决的方式不是改模型而是在 GUI 里加一条使用提示“请按标准手写体书写7 尽量带横杠、4 尽量开口。”这种交互层面的引导比任何算法优化都直接有效。如果你想把这件事做得更严谨一点可以在 GUI 里接入数字样本采集功能每识别一个手写数字把裁剪后的 28x28 图片和真实标签保存到本地文件夹积累几百条后增量微调模型。这是一个很自然的模型迭代闭环也是从“跑通”到“能用”的过渡标志。我在自己的项目里体验过从 99% 到 99.5% 的提升靠的往往不是换更大的网络而是这些数据和交互层面的打磨。这个项目做到这里横跨了数据、模型、界面、部署四条线虽然每个环节都不深但它把深度学习从“Jupyter Notebook 里跑一个 cell”带到了“别人能上手用的工具软件”。这也是一开始值得在这个方向上花时间的原因。希望这份经验能帮你在同样的路上少折腾几个晚上。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联 返回资讯列表 →