ARTICLE DETAIL

资讯详情

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

PyTorch实战:手写数字识别与Tkinter手写板GUI全流程

PyTorch实战:手写数字识别与Tkinter手写板GUI全流程 简介这份资源面向希望入门深度学习与计算机视觉的 Python 开发者、学生及算法爱好者提供一套可直接运行的手写数字识别完整方案解决从模型训练到交互式识别演示的落地问题。压缩包共 5 个文件以 4 个 py 脚本和 1 个 pth 模型为主整体约 1.53MB其中 py 文件分别承担网络结构定义、训练流程、工具函数与 GUI 主程序等职责pth 为已训练好的权重文件。资源基于 Pytorch 构建网络包含卷积层与全连接层训练代码可自行调整参数重新训练同时附带训练了 140 个 epoch 的模型无需等待即可直接加载使用配合 PyQt5 实现的 GUI 界面用户能在手写板上书写数字并实时获得识别结果。目前已有 2380 人学习下载适合作为课程设计、毕业项目或深度学习练手案例帮助读者理解卷积神经网络原理、掌握模型训练与推理流程并体验从算法到可视化交互的完整闭环。1. 从一张手写涂鸦到识别结果这套 PyTorch 方案到底解决什么问题很多人做 MNIST 手写数字识别训练完模型准确率能到 99%但真让他拿鼠标在屏幕上写个数字模型却认不出来。问题不在模型在于训练数据和推理数据之间隔了一条鸿沟MNIST 是 28×28 的灰度图笔画居中、粗细均匀、背景纯黑而你用鼠标画出来的线条又粗又歪位置偏、背景白、笔画带抗锯齿。这套「Python 手写数字识别带手写板 GUI 界面 PyTorch 代码 含训练模型」的方案核心就是把这套链路打通——从 PyTorch 训练一个 CNN到用 Tkinter 搭一个能画板输入的手写板 GUI再到把画板上的笔迹预处理成模型能吃的张量最后实时输出识别结果。适合刚入门 PyTorch、想做一个能拿得出手的完整小项目的人也适合已经会训模型但没做过 GUI 推理闭环的开发者。下面按「训练 → 预处理 → GUI → 联调 → 避坑」的顺序拆开讲。2. 用 PyTorch 训练 MNIST 模型网络结构、超参和保存方式2.1 为什么选 CNN 而不是全连接MNIST 虽然简单但全连接网络对像素位置极其敏感——你把数字往右挪两个像素全连接的输出就可能全错。卷积核自带平移不变性两层卷积就能把局部笔画特征横、竖、弧提取出来再经过池化降维最后接全连接分类。我一般用两个卷积块Conv → ReLU → MaxPool重复两次然后展平接两层全连接。这个结构参数量在 100 万左右CPU 上训练 5 个 epoch 大概两三分钟GPU 上几十秒。不要一上来就上 ResNetMNIST 用不着反而容易过拟合。2.2 训练脚本与关键参数import torch import torch.nn as nn import torch.optim as optim from torchvision import datasets, transforms from torch.utils.data import DataLoader # 数据预处理转张量 归一化到 [-1, 1] transform transforms.Compose([ transforms.ToTensor(), # 像素值从 [0,255] 变 [0,1] transforms.Normalize((0.1307,), (0.3081,)) # MNIST 全局均值和标准差 ]) train_set datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) test_set datasets.MNIST(root./data, trainFalse, downloadTrue, transformtransform) train_loader DataLoader(train_set, batch_size64, shuffleTrue) test_loader DataLoader(test_set, batch_size1000, shuffleFalse) class Net(nn.Module): def __init__(self): super().__init__() self.conv1 nn.Sequential( nn.Conv2d(1, 32, kernel_size3, padding1), # 1 通道进32 通道出 nn.ReLU(), nn.MaxPool2d(2) # 28x28 - 14x14 ) self.conv2 nn.Sequential( nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(2) # 14x14 - 7x7 ) self.fc nn.Sequential( nn.Flatten(), nn.Linear(64 * 7 * 7, 128), nn.ReLU(), nn.Dropout(0.25), # 防过拟合 nn.Linear(128, 10) # 10 个数字类别 ) def forward(self, x): x self.conv1(x) x self.conv2(x) return self.fc(x) device torch.device(cuda if torch.cuda.is_available() else cpu) model Net().to(device) optimizer optim.Adam(model.parameters(), lr1e-3) criterion nn.CrossEntropyLoss() for epoch in range(5): 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) pred model(data).argmax(dim1) correct pred.eq(target).sum().item() print(fEpoch {epoch1}, Test Acc: {correct / len(test_set):.4f}) # 保存模型只存 state_dict加载时先建结构再 load torch.save(model.state_dict(), mnist_cnn.pth)逻辑说明Normalize用的 (0.1307, 0.3081) 是 MNIST 训练集的全局统计值不是随便填的换成别的数据集要重新算。Dropout(0.25)放在全连接层之间卷积层后面不加因为卷积本身参数共享已经有一定正则效果。保存用state_dict()而不是整个模型对象这样加载时不依赖原始类定义路径换台机器也能跑。参数说明batch_size64是 MNIST 上的经验值显存小可以降到 32但太小会让 BatchNorm如果加了不稳定。lr1e-3配 Adam 是安全起点如果 loss 震荡就降到 5e-4。epoch5足够收敛到 99% 以上再多容易过拟合——你会看到测试准确率不升反降。2.3 模型保存与加载的两种方式对比方式代码优点缺点只存 state_dicttorch.save(model.state_dict(), x.pth)文件小、跨设备兼容加载前必须重建模型结构存整个模型torch.save(model, x.pth)加载一行搞定依赖类定义路径换目录就报错我一般用第一种配合一个load_model()函数把结构定义和加载写在一起GUI 启动时调用。3. 手写板 GUI 怎么做Tkinter 画板 实时预处理管线3.1 为什么选 Tkinter 而不是 PyQtPyQt 功能强、界面好看但安装包大、授权复杂PyQt5 是 GPL商用要注意。Tkinter 是 Python 标准库自带不用额外装代码量少做一个画板加按钮加标签的界面足够了。这个项目的 GUI 需求就三样一块能画画的 Canvas、一个识别按钮、一个显示结果的 Label。Tkinter 完全够用而且打包成 exe 时体积小很多。3.2 画板事件绑定与笔迹采集import tkinter as tk from PIL import Image, ImageDraw class DrawBoard: def __init__(self, root): self.root root self.canvas tk.Canvas(root, width280, height280, bgblack) self.canvas.pack() self.image Image.new(L, (280, 280), 0) # 灰度图背景黑 self.draw ImageDraw.Draw(self.image) self.last_x, self.last_y None, None # 绑定鼠标事件 self.canvas.bind(Button-1, self.start_draw) self.canvas.bind(B1-Motion, self.paint) self.canvas.bind(ButtonRelease-1, self.reset) def start_draw(self, event): self.last_x, self.last_y event.x, event.y def paint(self, event): if self.last_x and self.last_y: # 在 Canvas 上画白线 self.canvas.create_line(self.last_x, self.last_y, event.x, event.y, fillwhite, width12, capstyletk.ROUND) # 同步画到 PIL 图像上供后续预处理 self.draw.line([self.last_x, self.last_y, event.x, event.y], fill255, width12) self.last_x, self.last_y event.x, event.y def reset(self, event): self.last_x, self.last_y None, None def clear(self): self.canvas.delete(all) self.image Image.new(L, (280, 280), 0) self.draw ImageDraw.Draw(self.image)逻辑说明Canvas 负责显示PIL Image 负责存数据。两边同步画保证你看到的和模型拿到的是一致的。width12是笔刷粗细这个值很关键——太细模型认不出MNIST 笔画相对粗太粗会糊成一团。capstyletk.ROUND让线条拐弯处圆滑接近真实手写。参数说明Canvas 尺寸设 280×280是 MNIST 的 28×28 放大 10 倍方便鼠标操作。背景设黑色、笔迹白色和 MNIST 配色一致省得后面反色。如果你习惯白底黑字那预处理时要加一步反色否则模型输入分布对不上。3.3 从画板到模型输入预处理四步走画板上是 280×280 的图模型要的是 28×28 且居中的图。中间差四步裁剪、缩放、居中、归一化。import numpy as np from PIL import Image def preprocess(pil_image): # 1. 找到笔迹的边界框 bbox pil_image.getbbox() # 返回非零区域的 (left, upper, right, lower) if bbox is None: return None # 空画板 cropped pil_image.crop(bbox) # 2. 等比缩放到 20x20 以内保持长宽比 w, h cropped.size scale 20.0 / max(w, h) new_w, new_h int(w * scale), int(h * scale) resized cropped.resize((new_w, new_h), Image.LANCZOS) # 3. 粘贴到 28x28 黑底中心 canvas Image.new(L, (28, 28), 0) paste_x (28 - new_w) // 2 paste_y (28 - new_h) // 2 canvas.paste(resized, (paste_x, paste_y)) # 4. 转 numpy 再转 tensor归一化 arr np.array(canvas, dtypenp.float32) / 255.0 arr (arr - 0.1307) / 0.3081 # 和训练时一致的归一化 tensor torch.from_numpy(arr).unsqueeze(0).unsqueeze(0) # [1,1,28,28] return tensor逻辑说明getbbox()找到笔迹的最小外接矩形这一步去掉了画板上的空白区域让数字不管写在哪个角落都能被正确识别。缩放到 20×20 而不是直接 28×28是 MNIST 官方预处理的做法——留出 4 像素边距让数字居中。LANCZOS插值比默认的 NEAREST 更平滑减少锯齿对识别的干扰。参数说明20.0这个缩放目标别改改成 28 会让数字顶满边框模型没见过这种样本。归一化的均值和标准差必须和训练时完全一致差一点准确率就掉。unsqueeze两次是因为模型输入要求[batch, channel, H, W]单张图要补上 batch 和 channel 维度。4. 把 GUI 和模型接起来推理线程、实时识别与打包4.1 推理函数与按钮绑定import torch from model import Net # 假设网络定义在 model.py class App: def __init__(self, root): self.root root self.board DrawBoard(root) self.model Net() self.model.load_state_dict(torch.load(mnist_cnn.pth, map_locationcpu)) self.model.eval() btn_frame tk.Frame(root) btn_frame.pack() tk.Button(btn_frame, text识别, commandself.predict).pack(sidetk.LEFT) tk.Button(btn_frame, text清空, commandself.board.clear).pack(sidetk.LEFT) self.result_label tk.Label(root, text结果, font(Arial, 24)) self.result_label.pack() def predict(self): tensor preprocess(self.board.image) if tensor is None: self.result_label.config(text请先写一个数字) return with torch.no_grad(): output self.model(tensor) pred output.argmax(dim1).item() prob torch.softmax(output, dim1).max().item() self.result_label.config(textf结果{pred} 置信度{prob:.2%})逻辑说明map_locationcpu保证在有 GPU 的机器上训练、在没 GPU 的机器上推理时不会报错。torch.no_grad()关掉梯度计算省内存也快一点。softmax后的最大值当置信度低于 60% 时你可以选择不显示结果或者提示「请写清楚一点」。参数说明load_state_dict要求模型结构和保存时完全一致改过网络层数或名字就会报 key 不匹配。如果报错先 print 一下model.state_dict().keys()和保存文件的 keys 对比。4.2 用 PyInstaller 打包成 exepip install pyinstaller pyinstaller --onefile --windowed --add-data mnist_cnn.pth;. main.py--onefile打成单个 exe--windowed去掉黑色控制台窗口--add-data把模型文件一起打进去。注意 Windows 上用分号分隔源和目标Linux/Mac 上用冒号。打包后模型路径要用sys._MEIPASS处理否则 exe 找不到 pth 文件。import sys, os def resource_path(relative): if hasattr(sys, _MEIPASS): return os.path.join(sys._MEIPASS, relative) return os.path.join(os.path.abspath(.), relative) self.model.load_state_dict(torch.load(resource_path(mnist_cnn.pth), map_locationcpu))4.3 实时识别要不要每画一笔就推理有人想做成「边写边识别」每落一笔就调一次模型。技术上可行但体验不好——你写「5」的第一笔是横模型可能识别成「7」或「1」结果标签疯狂跳。我一般做成按钮触发写完点一下识别。如果非要做实时加一个 300ms 的防抖定时器停止画线 300ms 后再推理这样结果稳定得多。5. 避坑与排查手写板识别翻车的五个典型场景5.1 现象训练准确率 99%画板上写啥都识别成同一个数字原因预处理和训练的输入分布不一致。最常见的是画板背景白色、笔迹黑色而 MNIST 是黑底白字。模型学到的是「白色像素是笔画」你给它白底它把整个背景当笔画。解决要么画板设黑底白笔要么在预处理里加arr 255 - arr反色。检查方法把预处理后的 28×28 图用plt.imshow显示出来和 MNIST 样本对比肉眼一看就知道对不对。5.2 现象数字写在画板角落就识别错写中间就正常原因没有做居中处理。MNIST 所有数字都在 28×28 的中心区域模型对位置敏感。你写在角落缩放后数字偏到一边模型没见过这种分布。解决getbbox()裁剪 等比缩放 粘贴到中心三步缺一不可。裁剪后如果长宽比差异大比如「1」很窄缩放时保持长宽比不要强行拉成正方形。5.3 现象加载模型时报RuntimeError: Error(s) in loading state_dict原因保存和加载时的网络结构不一致。常见于改了层名、加了层、或者保存时用了DataParallelkey 会多一个module.前缀。解决先 print 两边的 keys 对比。如果是module.前缀问题加载时用model.load_state_dict({k.replace(module., ): v for k, v in state_dict.items()})。如果是结构改了老老实实重建模型。5.4 现象打包成 exe 后运行报FileNotFoundError: mnist_cnn.pth原因PyInstaller 打包后文件被解压到临时目录相对路径找不到。解决用上面resource_path()函数通过sys._MEIPASS定位。另外--add-data的参数在 Windows 和 Linux 上分隔符不同写错了文件根本不会被打进去。5.5 现象识别置信度很低结果在几个数字之间跳原因笔刷太细或太粗。太细时缩放后笔画断断续续太粗时数字糊成一团。MNIST 的笔画宽度在 28×28 上大概 2-3 像素对应 280×280 画板上 20-30 像素。解决笔刷宽度设 12-15 比较合适。另外画的时候尽量写大一点占满画板 2/3 以上给getbbox足够的裁剪空间。6. 进阶技巧用数据增强和置信度过滤把识别率再拉一截训练时的数据增强能显著提升画板识别的鲁棒性。MNIST 自带的数字都是居中且规整的而手写板输入有随机偏移和旋转。在transform里加RandomAffinetransform_train transforms.Compose([ transforms.RandomAffine(degrees10, translate(0.1, 0.1), scale(0.9, 1.1)), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ])degrees10允许 ±10 度旋转translate(0.1, 0.1)允许 10% 的平移scale允许 90%-110% 缩放。这三个参数模拟了手写时的自然抖动。注意增强只加在训练集上测试集和推理时不要加否则准确率评估会偏低。另一个技巧是置信度过滤。推理时如果prob 0.6不直接显示数字而是提示「请写大一点或写清楚一点」。这比硬给一个错误结果体验好。你还可以把 softmax 输出的 top-3 都显示出来让用户自己判断。增强参数推荐值作用degrees10模拟手写倾斜translate0.1模拟位置偏移scale0.9-1.1模拟书写大小差异我自己的习惯是每次改完预处理或增强先把画板上的图存成 png和 MNIST 的样本拼在一起看肉眼确认分布一致了再跑训练。这个「肉眼对齐」的笨办法帮我省了无数次调参的来回。另外模型文件别只存一个训练时每个 epoch 存一个最后挑测试集上最好的那个用别拿最后一个——最后一个往往已经过拟合了。希望帮到你。本文还有配套的精品资源点击获取
返回列表