ARTICLE DETAIL

资讯详情

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

PyTorch+PyQt5手写数字识别GUI实战:从训练到交互识别

PyTorch+PyQt5手写数字识别GUI实战:从训练到交互识别 简介这份资源面向希望入门深度学习与计算机视觉的 Python 开发者以及需要完成课程设计或练手项目的学生提供了一套可直接运行的手写数字识别完整方案。核心基于 PyTorch 搭建卷积神经网络包含卷积层与全连接层训练脚本可自行调整参数重新训练同时附带已训练 140 个 epoch 的 pth 模型无需等待即可加载推理。界面部分采用 PyQt5 实现手写板 GUI用户可在窗口内直接书写数字并实时获得识别结果适合作为神经网络落地交互的参考范例。压缩包共 5 个文件以 4 个 py 脚本和 1 个 pth 模型为主脚本分别承担网络结构定义、训练流程、工具函数与 GUI 主程序整体约 1.53MB结构精简、依赖清晰。目前已有 2380 人学习下载读者可借此理解从数据加载、模型训练到界面集成的完整链路并在此基础上替换数据集或调整网络层数快速迁移到其他图像分类任务中。1. 从一张手写数字截图说起这套 PyTorch PyQt5 资源到底能跑出什么很多人第一次接触深度学习都是从 MNIST 手写数字识别开始的。但真正让人卡住的往往不是模型本身而是「训练完了怎么用」。你跑完train.py终端打印出一串准确率然后呢想验证一下效果还得写个脚本读图片、转张量、跑推理折腾半天。这套资源解决的就是这个断层它把 PyTorch 训练的卷积网络和 PyQt5 手写板 GUI 拼在一起你打开界面直接用鼠标写一个数字点识别结果就出来了。资源包里包含train.py、net.py、model.pth、main.py、utils.py和mnist_gui目录。其中model.pth是已经训练了 140 个 epoch 的权重文件意味着你不需要先跑训练就能直接体验识别效果。网络结构是卷积层加全连接层的经典组合训练代码完整可复现。适合两类人一是刚学完 PyTorch 基础想找个完整项目练手的二是需要快速搭一个手写识别演示原型给非技术同事看的。下面从环境搭建到模型结构再到 GUI 交互把这份资源拆开讲清楚。2. 环境搭建与依赖安装把 PyTorch 和 PyQt5 装进同一个 Python 环境2.1 为什么建议用 conda 而不是 pip 裸装PyTorch 和 PyQt5 对 Python 版本和底层库的依赖比较敏感尤其是 Windows 上直接用系统 Python 装 PyTorch经常会遇到 DLL 加载失败或者 numpy 版本冲突。常见做法是用 conda 建一个独立环境把 Python 版本锁在 3.8 到 3.10 之间这三个版本对 PyTorch 和 PyQt5 的兼容性最稳。# 创建独立环境Python 版本建议 3.9 conda create -n mnist_gui python3.9 -y conda activate mnist_gui # 安装 PyTorchCPU 版本就够用这个项目不需要 GPU # 如果你有 NVIDIA 显卡且装了 CUDA可以去 PyTorch 官网复制对应命令 pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu # 安装 PyQt5 和图像处理库 pip install PyQt5 opencv-python pillow numpy这里有几个参数需要说明。--index-url指定的是 PyTorch 官方 CPU 版仓库下载速度比默认源快而且不会把 CUDA 相关的几个 G 的包拖下来。opencv-python和pillow在 GUI 里处理手写板画布转张量时会用到numpy是 PyTorch 的隐式依赖显式装上避免版本问题。提示如果你之前装过 PyTorch 但版本混乱建议先pip uninstall torch torchvision再重新装不要在一个已经装了 TensorFlow 的环境里硬塞 PyTorchnumpy 版本冲突会让你排查到怀疑人生。2.2 验证环境是否就绪装完之后别急着跑main.py先做两步验证。第一步确认 PyTorch 能正常导入并且版本符合预期import torch import torchvision import PyQt5.QtCore print(PyTorch:, torch.__version__) print(TorchVision:, torchvision.__version__) print(PyQt5 QtCore:, PyQt5.QtCore.QT_VERSION_STR) # 检查是否能创建张量并做一次简单卷积 x torch.randn(1, 1, 28, 28) conv torch.nn.Conv2d(1, 32, kernel_size3, padding1) out conv(x) print(Conv output shape:, out.shape) # 应该是 torch.Size([1, 32, 28, 28])这段代码做了三件事打印三个核心库的版本号确认导入路径没有冲突创建一个模拟的 MNIST 输入张量形状是[batch1, channel1, height28, width28]用一个卷积层跑一次前向传播输出通道数 32、空间尺寸不变说明 PyTorch 的计算图能正常执行。如果这一步报错后面 GUI 里的识别一定跑不通先在这里解决。第二步确认 PyQt5 能弹出窗口。有些 Linux 服务器没有图形界面PyQt5 装了也弹不出来这种情况需要在本地机器上跑或者用 X11 转发。Windows 和 macOS 一般不会有这个问题。3. 网络结构与训练脚本拆解卷积层怎么堆、全连接层怎么接3.1 net.py 里的网络定义逻辑net.py定义了一个继承自nn.Module的类名字可能是Net或ConvNet具体以你拿到的代码为准。结构上分两块特征提取部分用卷积层加池化层堆叠分类部分用全连接层输出 10 个类别的概率。典型的写法是这样import torch 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, kernel_size3, padding1) # 第二层卷积32 通道进64 通道出 self.conv2 nn.Conv2d(32, 64, kernel_size3, padding1) # 池化层2x2 窗口步长 2每次把空间尺寸减半 self.pool nn.MaxPool2d(2, 2) # 全连接层经过两次池化后 28x28 变成 7x764 通道展平后是 64*7*7 self.fc1 nn.Linear(64 * 7 * 7, 128) self.fc2 nn.Linear(128, 10) # 10 个数字类别 def forward(self, x): x self.pool(F.relu(self.conv1(x))) # 28x28 - 14x14 x self.pool(F.relu(self.conv2(x))) # 14x14 - 7x7 x x.view(-1, 64 * 7 * 7) # 展平 x F.relu(self.fc1(x)) x self.fc2(x) return x关键参数解释一下。padding1配合kernel_size3保证卷积后空间尺寸不变这样两次池化后刚好从 28 降到 7全连接层的输入维度64*7*7就是这么来的。如果你改了卷积核大小或者池化方式这个维度必须跟着改否则view那一步会报形状不匹配。fc1的 128 是隐藏层宽度可以调大但会增加参数量和训练时间MNIST 这种简单任务 128 足够。3.2 train.py 的训练流程与超参数train.py负责加载 MNIST 数据集、定义损失函数和优化器、跑训练循环。核心代码结构如下import torch import torch.nn as nn import torch.optim as optim from torchvision import datasets, transforms from torch.utils.data import DataLoader from net import Net # 数据预处理转张量 归一化 transform transforms.Compose([ transforms.ToTensor(), # 像素值从 0-255 转到 0-1 transforms.Normalize((0.1307,), (0.3081,)) # MNIST 全局均值和标准差 ]) # 加载训练集和测试集 train_dataset datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) test_dataset datasets.MNIST(root./data, trainFalse, downloadTrue, transformtransform) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue) test_loader DataLoader(test_dataset, batch_size1000, shuffleFalse) device torch.device(cuda if torch.cuda.is_available() else cpu) model Net().to(device) optimizer optim.Adam(model.parameters(), lr0.001) criterion nn.CrossEntropyLoss() for epoch in range(1, 141): # 资源里的模型训练了 140 个 epoch model.train() for batch_idx, (data, target) in enumerate(train_loader): data, target data.to(device), target.to(device) optimizer.zero_grad() output model(data) loss criterion(output, target) loss.backward() optimizer.step() # 每个 epoch 结束后在测试集上评估 model.eval() correct 0 with torch.no_grad(): for data, target in test_loader: data, target data.to(device), target.to(device) output model(data) pred output.argmax(dim1) correct pred.eq(target).sum().item() acc correct / len(test_dataset) print(fEpoch {epoch}, Test Accuracy: {acc:.4f}) # 保存模型权重 torch.save(model.state_dict(), model.pth)Normalize里的(0.1307,)和(0.3081,)是 MNIST 数据集的全局均值和标准差这两个数字是固定的不要改成别的。batch_size64是常见选择显存不够就降到 32CPU 训练的话 64 也能跑。Adam优化器学习率0.001对这个网络规模比较合适用 SGD 的话需要调到 0.01 左右并加动量。140 个 epoch 在 CPU 上大概要跑一两个小时GPU 上几分钟就完事。如果你不想等直接用资源里的model.pth就行。注意torch.save(model.state_dict(), model.pth)保存的是参数字典加载的时候需要先实例化Net()再load_state_dict。如果保存的是整个模型对象加载时依赖的类定义路径必须一致换目录就容易翻车。这份资源用的是 state_dict 方式更稳妥。4. GUI 手写板与推理管线从鼠标轨迹到 28x28 张量4.1 main.py 里的 PyQt5 界面布局main.py用 PyQt5 搭了一个窗口核心控件是一块画布区域和一个识别按钮。画布通常用QWidget或QLabel自定义鼠标事件来实现记录鼠标按下和移动时的坐标点用QPainter画线。识别按钮触发后把画布内容截取出来缩放到 28x28转成张量送进模型。import sys import numpy as np from PyQt5.QtWidgets import QApplication, QMainWindow, QWidget, QPushButton, QVBoxLayout, QLabel from PyQt5.QtGui import QPainter, QPen, QImage, QPixmap from PyQt5.QtCore import Qt, QPoint import torch from net import Net class DrawBoard(QWidget): def __init__(self): super().__init__() self.setFixedSize(280, 280) self.image QImage(280, 280, QImage.Format_RGB32) self.image.fill(Qt.black) # 黑底白字和 MNIST 风格一致 self.last_point QPoint() self.drawing False def mousePressEvent(self, event): if event.button() Qt.LeftButton: self.drawing True self.last_point event.pos() def mouseMoveEvent(self, event): if self.drawing and (event.buttons() Qt.LeftButton): painter QPainter(self.image) pen QPen(Qt.white, 20, Qt.SolidLine, Qt.RoundCap, Qt.RoundJoin) painter.setPen(pen) painter.drawLine(self.last_point, event.pos()) painter.end() self.last_point event.pos() self.update() def mouseReleaseEvent(self, event): if event.button() Qt.LeftButton: self.drawing False def paintEvent(self, event): canvas QPainter(self) canvas.drawImage(self.rect(), self.image, self.image.rect()) def clear(self): self.image.fill(Qt.black) self.update() def get_array(self): # 把 QImage 转成 numpy 数组再缩放到 28x28 ptr self.image.constBits() ptr.setsize(self.image.byteCount()) arr np.array(ptr).reshape(280, 280, 4) # RGBA gray arr[:, :, 0] # 取红色通道即可因为是黑白图 # 用 OpenCV 或 PIL 缩放到 28x28 import cv2 resized cv2.resize(gray.astype(np.uint8), (28, 28), interpolationcv2.INTER_AREA) return resized画笔宽度设成 20 是因为 280x280 的画布缩到 28x28 是 10 倍缩小20 像素宽的线缩完大概 2 像素和 MNIST 里数字的笔画粗细接近。如果设太细缩小后笔画会断设太粗数字会糊成一团。Qt.RoundCap和Qt.RoundJoin让线条端点圆滑避免出现尖锐拐角影响识别。4.2 推理时的张量预处理与模型加载拿到 28x28 的 numpy 数组后不能直接送进模型需要做和训练时一致的归一化def predict(model, img_array): # img_array 是 28x28 的 uint8值域 0-255 tensor torch.from_numpy(img_array).float() / 255.0 # 转到 0-1 tensor (tensor - 0.1307) / 0.3081 # 和训练时一致的归一化 tensor tensor.unsqueeze(0).unsqueeze(0) # 变成 [1, 1, 28, 28] model.eval() with torch.no_grad(): output model(tensor) pred output.argmax(dim1).item() prob torch.softmax(output, dim1).max().item() return pred, probunsqueeze(0)两次分别在第 0 维加 batch 维、第 1 维加通道维这是 PyTorch 卷积层要求的输入格式[N, C, H, W]。归一化用的均值和标准差必须和训练时完全一致否则模型看到的输入分布变了准确率会掉得很难看。torch.no_grad()关闭梯度计算推理时省内存也快。加载模型权重的代码在main.py启动时执行model Net() model.load_state_dict(torch.load(model.pth, map_locationtorch.device(cpu))) model.eval()map_locationcpu是防止在 GPU 上训练的模型加载到 CPU 机器时报错。如果你确定只在 CPU 上跑这个参数加上更保险。5. 避坑与排查识别不准、界面卡死、模型加载失败怎么查5.1 手写数字识别总是不对现象写了个很清楚的「3」识别成「8」或「5」。原因最常见的是画布背景和 MNIST 数据集不一致。MNIST 是黑底白字如果你 GUI 里用了白底黑字模型看到的输入是反的识别率会暴跌。另一个原因是画笔太粗或太细缩放后笔画形态和训练数据差异大。解决确认画布fill(Qt.black)画笔颜色Qt.white。画笔宽度在 15 到 25 之间调写完数字占画布中央 60% 到 80% 面积比较合适。如果还是不准可以在get_array里加一步二值化把灰度值大于 50 的置为 255其余置 0让输入更接近 MNIST 的纯黑白风格。5.2 点击识别按钮后界面卡住现象按下识别按钮窗口无响应几秒钟然后才出结果。原因推理代码跑在主线程里PyTorch 加载模型和第一次前向传播需要初始化计算图耗时较长。如果每次点击都重新加载模型卡顿会更明显。解决把模型加载放在__init__里只执行一次识别时只跑前向传播。如果还是卡可以用QThread把推理放到子线程主线程只负责更新界面。不过 MNIST 这个网络很小CPU 推理一次也就几十毫秒通常不需要上线程检查一下是不是每次点击都在torch.load。5.3 model.pth 加载报错 Unexpected key(s) in state_dict现象load_state_dict抛出RuntimeError: Unexpected key(s) in state_dict: conv1.weight, ...。原因保存模型时的网络结构和加载时的网络结构不一致。比如训练时用了conv1、conv2、fc1、fc2四个命名加载时的Net类里改成了别的名字或者层数对不上。解决打开net.py确认类名和层命名确保和训练时用的完全一致。如果资源里的model.pth和net.py是配套的直接跑不应该报这个错。如果你自己改了网络结构又加载旧权重那只能重新训练。5.4 PyQt5 报 Could not load the Qt platform plugin windows现象运行main.py时直接崩溃提示找不到 Qt 平台插件。原因PyQt5 安装不完整或者环境变量QT_QPA_PLATFORM_PLUGIN_PATH指向了错误的路径。常见于在 conda 环境里混用了 pip 和 conda 安装的 PyQt5。解决先pip uninstall PyQt5 PyQt5-tools然后pip install PyQt5重新装。如果还不行在代码最开头加import os; os.environ[QT_QPA_PLATFORM_PLUGIN_PATH] 让 Qt 自己找插件路径。Windows 上还可以试试用conda install pyqt走 conda 渠道装。5.5 训练时 loss 不下降或变成 nan现象跑train.py前几个 batch loss 正常后面突然变成 nan。原因学习率太大Adam 的lr0.001一般不会但如果你改成了 0.01 或更高梯度爆炸就会出 nan。另一个可能是输入数据没有归一化像素值 0-255 直接送进网络第一层卷积输出值过大。解决确认transforms.Normalize那一步没被注释掉。学习率调回 0.001如果还出 nan在loss.backward()后面加torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5)做梯度裁剪。6. 进阶玩法用自己写的数字微调模型并验证泛化资源里的model.pth是在标准 MNIST 上训了 140 个 epoch测试集准确率通常能到 99% 以上。但你自己在 GUI 上写的数字风格和 MNIST 不完全一样偶尔识别错很正常。一个实用的进阶做法是把你手写的数字保存下来打上标签微调最后几层让模型适应你的笔迹。具体操作分三步。第一步在 GUI 里加一个「保存样本」按钮把当前画布转成 28x28 数组后存成 png文件名用「标签_序号.png」格式比如3_001.png。第二步写一个微调脚本加载model.pth把最后全连接层的参数解冻用你自己收集的几百张样本跑几个 epoch学习率调到 0.0001避免把预训练特征冲掉。# 微调脚本核心逻辑 model Net() model.load_state_dict(torch.load(model.pth)) # 只训练 fc1 和 fc2冻结卷积层 for param in model.conv1.parameters(): param.requires_grad False for param in model.conv2.parameters(): param.requires_grad False optimizer optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr0.0001) # 后续训练循环和 train.py 一致但数据集换成你自己的手写样本第三步验证微调效果。把微调后的权重存成model_finetuned.pth在 GUI 里加载新权重再写同样的数字看识别结果是否改善。注意保留原始model.pth作为回退万一微调把模型带偏了还能切回来。我自己的习惯是每次改完网络结构或换了数据集先把model.pth备份一份命名带上日期和准确率比如model_20250101_acc995.pth。这样后面试各种微调方案时随时能回到一个已知可用的版本不用重新训 140 个 epoch。从那以后我每次动训练脚本之前都强制走一遍备份流程省下了不少后悔药。希望帮到你。本文还有配套的精品资源点击获取
返回列表