ARTICLE DETAIL

资讯详情

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

PyTorch CNN手写数字识别:从模型训练到网页实时部署全流程

PyTorch CNN手写数字识别:从模型训练到网页实时部署全流程 简介面向需要快速上手图像识别网页应用的开发者这份压缩包提供了一个基于PyTorch CNN的手写数字识别完整方案。压缩包共131个文件以124张JPG图片构成数据集主体另有3个Python脚本、3个TXT文本和1个HTML页面整体大小仅3.88MB结构紧凑其中TXT文本用于记录环境依赖及数据标签Python脚本分别承担数据集文本生成、模型训练、本地网页服务HTML则提供前端交互界面。目前已有96人学习下载。使用脉络非常清晰先运行01脚本读取各分类图片并生成对应的标签文本再运行02脚本基于该文本完成CNN训练保存模型与日志最后启动03脚本即可在浏览器中打开本地URL进行手写数字实时识别验证。无论是课程设计、毕业设计还是PyTorch入门实践这套资源都能帮助理解从数据准备、模型训练到网页部署的完整流程且模型文件可留存便于后续调整参数或扩展类别。1. 网页版手写数字识别为什么值得跑一遍从PyTorch CNN到浏览器的完整闭环如果你刚学完深度学习正想看一个模型怎么被外部程序真正调用这个web网页式的项目应该是同类型里路径最短的一个不需要React不需要云服务器只用PyTorch训练好CNN卷积神经网络再用一个本地html网页作为推理入口就能打开浏览器手写数字并实时看到识别结果。资源里连图片数据集都带好了不用去下MNIST断网环境下也能把训练和推理整条链路走通。适合两类人一是刚学完CNN基础、想让“训练好的模型”跑在一个真实web页面上验证效果的学生二是需要快速搭一个演示demo给产品、给老师、给需求方看效果的工程师。全文我会按“数据集处理 → 模型训练 → web联调 → 踩坑记录 → 上线前验证”的顺序把这套资源里三个脚本的实际用法和边界讲清楚。2. 先把图片数据集看明白01脚本如何把文件夹“翻译”成训练txt2.1 目录结构与文件命名规则解压这个资源包之后目录里会有数据集文件夹、三个.py脚本、index.html和requirement.txt。很多新手拿到手第一件事是直接打开02深度学习模型训练.py运行结果一上来就报错找不到文件——因为02脚本依赖的是01脚本生成的txt文件而01脚本依赖的是数据集文件夹的目录结构。所以第一步不是看模型代码而是先把数据集的摆放方式搞清楚。数据集的常见组织方式是根目录下按数字0到9各建一个子文件夹比如数据集/0/、数据集/1/每个子文件夹里放对应类别的jpg图片。文件名本身无所谓像004.jpg、048.jpg、04_flip.jpg这种命名基本不影响训练因为标签是从“所在文件夹名”读出来的。看到04_flip.jpg这类文件说明数据已经做过水平翻转增强属于扩充样本的常规操作。我建议拿到资源后先手动浏览一遍每个类别文件夹确认三件事图片是不是统一jpg格式、每个类别的样本数是否大致均衡、图片尺寸是不是一致。如果尺寸不统一后面训练脚本里一定要有resize逻辑兜底。目录/文件作用数据集/0~9/存放各数字类别的jpg图片标签由文件夹名决定01数据集文本生成制作.py扫描数据集目录生成 train.txt 和 val.txt02深度学习模型训练.py读取txt中的路径和标签训练CNN并保存模型03html_server.py启动Flask服务加载模型提供HTTP推理接口index.html浏览器端手写画布页面负责采集笔迹并展示识别结果requirement.txtPyTorch、Flask等依赖清单2.2 01脚本运行时发生了什么01脚本干的事情本质上就是“打标签”把每个图片文件的绝对路径和它对应的数字类别写进txt的一行里类别按所在文件夹名从0开始逐个编号。下面这类逻辑是这个脚本最常见的实现方式。如果你的资源里的代码略有出入核心逻辑通常也就是这段的变体import os dataset_root 数据集 classes sorted(os.listdir(dataset_root)) # 按文件夹名排序得到类别列表 train_lines, val_lines [], [] for cls_id, cls_name in enumerate(classes): cls_dir os.path.join(dataset_root, cls_name) imgs [f for f in os.listdir(cls_dir) if f.lower().endswith(.jpg)] split int(len(imgs) * 0.8) # 前80%做训练集后20%做验证集 for i, img in enumerate(imgs): line f{os.path.join(cls_dir, img)} {cls_id}\n if i split: train_lines.append(line) else: val_lines.append(line) with open(train.txt, w, encodingutf-8) as f: f.writelines(train_lines) with open(val.txt, w, encodingutf-8) as f: f.writelines(val_lines) print(ftrain samples: {len(train_lines)}, val samples: {len(val_lines)})这段代码里的cls_id是核心它直接由文件夹名的排序位置决定而不是从图片内容推断。这样做的好处是标签和文件夹天然对应一眼能看懂坏处是一旦你新增了类别文件夹或者改动了文件夹名排序规则标签含义就可能整体位移。split按列表顺序直接切开这是一种省事的划分方式。更严谨的做法是先random.shuffle(imgs)再切分避免同一类图片按文件名排序后训练集和验证集分布不均。我在实际项目里还会顺手加一句print(classes)确认类别顺序是[0,1,2,...]而不是[1,0]。2.3 生成结果怎么核对脚本跑完后立刻验证txt内容这一步很多教程都不提但恰恰是性价比最高的排查手段。打开命令行切换到脚本所在目录执行type train.txt | more如果看到的行是数据集\0\005.jpg 0这样的格式说明路径和标签列都正常。这里有个隐蔽问题值得注意Windows下os.path.join生成的是反斜杠分隔路径比如数据集\0\005.jpg。PyTorch在Windows上读取没问题但如果后面你想把这份txt拷到Linux服务器或者Google Colab上训练反斜杠路径大概率报文件不存在。我在Windows上跑此类项目时习惯在01脚本末尾加两行替换逻辑with open(train.txt, w, encodingutf-8) as f: f.writelines(line.replace(\\, /) for line in train_lines) with open(val.txt, w, encodingutf-8) as f: f.writelines(line.replace(\\, /) for line in val_lines)顺手把分隔符统一成正斜杠省得以后踩跨平台的坑。标签列也要抽一眼如果一个文件夹里有几十张图片这些图片的标签值应该完全相同不可能出现同一文件夹里一会儿是0一会儿是1。我自己见过最快翻车的场景就是有人把某个类别图片放错了文件夹比如把“7”的图片放进了“1”的目录里肉眼看不出来训练出来的模型识别准确率始终在70%左右上不去查了半天数据才发现是标签错了。3. CNN训练主流程02脚本的网络结构、超参数与日志解读3.1 网络结构一个够用且好复现的LeNet变体02脚本的核心是训练一个CNN卷积神经网络。手写数字识别这个任务MNIST级别的数据集用LeNet这个量级的网络就足够了不需要搬ResNet或者VGG过来——数据量不大模型太大反而容易过拟合训练速度也慢。这个资源里最可能出现的结构是两层卷积加两层全连接这也是绝大多数入门级CNN手写数字项目的标配。模型定义部分通常是这样的import torch.nn as nn class DigitCNN(nn.Module): def __init__(self, num_classes10): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 32, kernel_size3, padding1), # 32个卷积核 nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(64 * 7 * 7, 128), # 输入按28x28图片计算 nn.ReLU(inplaceTrue), nn.Dropout(0.3), nn.Linear(128, num_classes), ) def forward(self, x): return self.classifier(self.features(x))这个结构的参数含义需要说清楚nn.Conv2d(1, 32, ...)中的1是输入通道数因为手写数字图片是单通道灰度图32是第一层卷积输出通道数越大代表提取的特征越多但计算量也越大。原始图片按28x28灰度处理的话经过两次MaxPool2d(2)尺寸从28x28变成7x7所以全连接层输入维度是64 * 7 * 7。Dropout(0.3)是训练时随机丢弃30%的神经元防止模型过度依赖某些节点导致过拟合推理阶段它会自动失效。如果你的资源里图片原始分辨率不是28x28全连接层的输入维度要根据(原尺寸 / 4)重新计算这是改结构时最容易报错的地方。3.2 训练循环与超参数设置训练脚本的常规流程是用torch.utils.data.DataLoader读取上一章生成的train.txt和val.txt完成图片解码、灰度转化、尺寸缩放和归一化之后喂给模型。图片文件夹里的jpg是RGB三通道的但CNN的第一个卷积层只接收单通道所以必须先转成灰度图。这一步偷懒不做运行时就会报通道数不匹配的错。一般会看到这样的加载代码from torch.utils.data import Dataset, DataLoader from PIL import Image import torchvision.transforms as T transform T.Compose([ T.Grayscale(num_output_channels1), T.Resize((28, 28)), T.ToTensor(), T.Normalize((0.5,), (0.5,)) ]) class DigitsDataset(Dataset): def __init__(self, txt_path, transformNone): self.samples [] self.transform transform with open(txt_path, r, encodingutf-8) as f: for line in f: path, label line.strip().split() self.samples.append((path, int(label))) def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label self.samples[idx] img Image.open(path).convert(L) if self.transform: img self.transform(img) return img, label这段代码里的Normalize((0.5,), (0.5,))表示把像素值从[0,1]归一化到[-1,1]具体公式是(x - 0.5) / 0.5。这个均值0.5和标准差0.5是手写数字这类灰度图的常见选择不需要像ImageNet那样用三通道的均值和方差。Image.open(path).convert(L)是先以原始格式打开图片再手动转灰度双保险。这里有个细节.convert(L)转出来的灰度值范围是0到255ToTensor()会自动缩放到[0,1]顺序不能颠倒。超参数方面这类项目最常见的是学习率0.001、batch size 128、训练50到100个epoch。我的实践经验是数据集图片总量不大时epoch设50已经能让验证集准确率稳定在98%以上。如果你发现50轮结束时准确率还在爬升没有平台期才需要加到100轮。优化器用Adam学习率0.001基本不用调比SGD省心。训练循环里每个epoch结束后记录验证集的loss和accuracy输出到本地log文件这正是资源简介里说的“log日志保存本地里面记录了每个epoch的验证集损失值和准确率”。3.3 模型保存与日志解读训练完通常会把model.state_dict()存成.pth文件同时把每个epoch的loss和accuracy追加写入一个train.log。日志的读法有讲究如果print出来的训练loss一直在降而验证集准确率到了某个epoch后不再上升甚至掉头向下说明开始过拟合了这时候不是继续加训练轮数而是减小学习率或者加大Dropout。如果loss从头到尾纹丝不动、准确率一直卡在10%左右那基本是数据或标签的问题不是模型的问题。10%这个数值对应随机乱猜的水平因为一共10个数字。另外一个常见习惯是把效果最好的epoch对应的模型单独保存一份而不是等到训练全部跑完才保存最后一次的结果。写法很简单best_acc 0.0 for epoch in range(epochs): train_acc train_one_epoch(model, train_loader) val_acc, val_loss evaluate(model, val_loader) if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_model.pth) print(fepoch {epoch}: train_acc{train_acc:.4f}, val_acc{val_acc:.4f}, val_loss{val_loss:.4f})torch.save只存state_dict()而不存整个model对象原因是state_dict只包含张量参数体积小、跨PyTorch版本兼容性好加载时只需要先实例化DigitCNN()再load_state_dict。这个习惯在换电脑、换Python环境时能省掉很多莫名其妙的版本兼容问题。我见过有人直接torch.save(model)结果换台机器加载失败报数据类型不匹配原因就是保存了整个对象结构PyTorch版本一升级字段就对不上了。4. Web端到端联调03脚本与index.html怎么把识别结果送到浏览器4.1 Flask服务器与路由03脚本的角色是把训练好的模型包装成一个本地HTTP服务。它用Flask起一个web服务器默认监听127.0.0.1:4399。打开浏览器访问这个地址时Flask返回index.html页面页面里用户手写一个数字并点击识别浏览器把图片数据POST回服务器服务器跑一次模型前向推理把预测结果以JSON格式返回给前端展示。整个交互不依赖任何外部服务全在本机完成。03脚本的典型骨架是这样的from flask import Flask, request, jsonify, render_template import torch import torchvision.transforms as T app Flask(__name__) model DigitCNN() model.load_state_dict(torch.load(best_model.pth, map_locationcpu)) model.eval() app.route(/) def index(): return render_template(index.html) app.route(/predict, methods[POST]) def predict(): data request.get_json() img_data data[image] # 前端传过来的base64编码图片 img decode_base64_to_tensor(img_data) with torch.no_grad(): output model(img) pred output.argmax(dim1).item() prob torch.softmax(output, dim1).max().item() return jsonify({prediction: pred, probability: prob}) if __name__ __main__: app.run(host127.0.0.1, port4399, debugFalse)这中间有几个必须注意的参数map_locationcpu是必须的——如果你本机有GPU但没装CUDA版PyTorch或者训练时用了GPU、测试时换了台纯CPU机器不加这个参数直接load会报错说没有GPU设备。model.eval()千万别漏它会把Dropout和BatchNorm切到推理模式漏掉的话同一张图每次预测结果可能都不一样。torch.no_grad()把梯度计算关掉推理时省内存也提速。为什么端口选4399没有特殊规定Flask默认是5000这个项目大概率是作者在写demo时顺手指定的避开常见端口免得和本机别的服务冲突。如果你访问时发现4399被别的程序占了把port4399改成一个没被占用的端口即可。4.2 index.html 画布与提交逻辑index.html是这个资源里web项目的入口页面它在浏览器里渲染一个手写画布。典型的实现是放一个canvas元素监听鼠标或触摸事件来画线识别按钮按下时把canvas内容导出成图片发送给后端。这里需要理解一个关键逻辑canvas本身是网页上的矢量绘图区要交给CNN识别必须先把它变成一张和训练数据同分布的图片。最省事的做法是调用canvas.toDataURL(image/png)拿到base64编码的PNG字符串POST给Flask后端解码。前端核心提交代码大概是canvas idboard width280 height280/canvas button onclickpredict()识别/button div idresult/div script function predict() { const canvas document.getElementById(board); const dataURL canvas.toDataURL(image/png); fetch(/predict, { method: POST, headers: { Content-Type: application/json }, body: JSON.stringify({ image: dataURL }) }).then(r r.json()).then(res { document.getElementById(result).innerText 预测结果: ${res.prediction}, 置信度: ${(res.probability * 100).toFixed(1)}%; }); } /script这里canvas设成280x280是为了手写好用因为28x28太小直接用鼠标在那么小的区域里写数字根本写不开。280x280写完后后端在预处理时会缩放到28x28。这个“前端大画布、后端小输入”的设计是这类web demo的通用做法用户写起来舒服模型吃到嘴里还是标准尺寸。有个细节容易被忽略toDataURL返回的字符串形如data:image/png;base64,xxxxxx后端要先按逗号切分出xxxxxx部分再做base64解码不要把data:image/png;base64,这个前缀一起喂给解码器否则图片解出来是坏的。4.3 一次完整的识别请求链路理解了前后端各自的任务整条链路就清晰了用户在canvas上写数字点击识别浏览器把280x280的canvas内容导出为PNG的base64字符串POST出去Flask收到后用base64解码还原成图片转灰度缩放到28x28归一化到[-1,1]变成[1, 1, 28, 28]的tensor第一维是batch size喂给模型模型前向传播得到10个类别的分数argmax取最大分数对应的下标作为预测数字最后把数字和置信度拼成JSON返回页面显示。链路里的每一个环节都依赖上一个环节的输出格式正确任何一步的格式没对上结果就是报错或者识别结果莫名其妙。我在给这类项目做联调时习惯先不用浏览器直接用Python的requests库模拟一次POST请求确认后端逻辑通顺再打开网页。这样排查问题时能把“前端问题”和“后端问题”快速切分如果是requests能通而浏览器不行问题基本出在canvas导出或CORS如果requests都不通问题出在Flask路由或模型加载。资源里的03脚本跑起来后终端会打印一个本地地址复制到浏览器打开即可不用手动敲。如果你确实要手动输入注意是http://127.0.0.1:4399不是localhost:4399——两者大多数时候等价但当本机hosts文件动过手脚时结果可能不一样。5. 五个坑排查笔记从URL打不开到识别率突然翻车5.1 URL打不开服务没起来还是端口被占用现象浏览器里输入http://127.0.0.1:4399一直转圈或提示无法访问但03脚本明明已经运行了。原因最常见的是端口被占用。4399这个端口在有些电脑上会被游戏相关进程或其他开发服务占掉Flask启动时报错Address already in use但很多人不看终端窗口的报错以为脚本跑起来了。还有一种情况是03脚本启动时加载模型失败代码抛异常退出了Flask服务根本没起来。解决先看终端有没有输出Running on http://127.0.0.1:4399这一行。接着执行netstat -ano | findstr 4399查看端口占用情况如果被别的进程占用就改掉03脚本里的port4399比如改成port8080。我在给这个web项目调试时每次改端口后都会顺手清空浏览器缓存再刷新避免浏览器把旧的失败连接缓存下来。5.2 准确率卡在10%先怀疑txt标签错位再怀疑数据预处理现象02脚本训练了50个epoch训练损失降到了0.1以下但验证准确率一直徘徊在10%到20%和乱猜差不多。原因这个症状几乎可以锁定是数据标签和路径错位。最常见的根因是01脚本在生成txt时某个类别的图片路径写到了错误标签下面或者图片转换灰度时出了问题。另一个高频元凶是数据集目录下混入了非数字文件夹比如一个名为.DS_Store的文件或者误放进去的说明文档导致cls_id整体位移。解决回到第2章用type train.txt | more抽查各个类别段的标签列是否连续正确。再检查数据集根目录下有没有隐藏文件或非图片文件有就删掉或移走。我处理过最离谱的一次是资源里有一个README.txt被放在了数据集/根目录01脚本把它当成一个类别结果所有真实数字的标签全部错位准确率必然崩。检查顺序是先看txt内容再看目录结构最后怀疑代码。5.3 浏览器手写识别差但训练和验证集准确率都正常现象训练日志显示验证集准确率98%但在网页canvas上手写数字识别结果经常是错的而且错的毫无规律。原因这是web推理项目最常见的翻车点——前端图片和训练图片的分布不一致。训练数据是白底黑字、数字居中且尺寸规范浏览器canvas里的手写笔画可能很细、位置靠边、甚至因为鼠标手抖导致数字断成两截。CNN对输入分布非常敏感训练时没见过的笔画粗细和位置偏移都会让它胡乱预测。解决先看后端有没有做正确的归一化比如canvas导出的是RGB彩色PNG而后端处理只接了灰度通道。其次在canvas尺寸不变的情况下把笔刷线宽调大、建议用户尽量把数字写在画布中心区域。最有效的一招是我在项目里强制在缩放前做一步“按内容边界裁剪”找到图片中最小的非背景像素包围盒裁掉四周空白再缩放相当于自动居中。这个30行不到的预处理函数能直接拉升一截浏览器端的真实手感。如果仍然不理想就把训练数据里的图片做随机平移几个像素、随机缩放粗细模拟手写偏差。5.4 显示CUDA out of memory或运行到一半卡死现象02脚本训练时报CUDA out of memory或者CPU环境下一运行就卡得动不了每轮epoch要等很久。原因显存溢出通常是batch size设得太大了。128张28x28的图片本来占不了多少显存但如果你用的是2G显存的旧显卡且代码里同时加载了模型、优化器、中间特征图就可能爆。卡死则大概率是Windows下DataLoader的num_workers设成了大于0的值多进程数据加载在Windows的PyTorch里经常出现诡异卡顿。解决显存溢出就把batch size从128调到32或16模型变小一点或者加一句torch.cuda.empty_cache()。卡死就把DataLoader(..., num_workers0)。我一般调试期一律num_workers0只有确认数据集读取正确后才调大这是Windows上跑PyTorch的血泪经验。如果你根本没有NVIDIA显卡但代码里写了model.cuda()也会报错这种情况把.cuda()调用全删掉纯CPU训练28x28的小数据集只是慢一点完全跑得动。5.5 地址输入没问题但页面白屏index.html没被渲染出来现象Flask启动成功地址也是对的但浏览器打开后一片空白控制台报404或500。原因Flask默认从templates目录找模板文件。如果index.html直接和03html_server.py放在同一层目录而代码里写的是render_template(index.html)Flask会去templates/index.html找找不到就404。还有一种情况是index.html里引用了外部的css或js文件路径写成了绝对路径而Flask默认静态文件目录是static/文件没放对位置导致页面渲染不出来。解决先确认项目目录结构——如果index.html在根目录而不是templates子目录下就把03脚本里的渲染方式改成send_file(index.html)。我拿到这套资源后第一件事就是看目录结构再决定渲染方式这个顺序错了会浪费很多时间。另外提醒一点03脚本启动流程里如果Flask服务是线程模式它会默认处理并发请求不需要额外配置别自己加多线程代码画蛇添足。6. 上线前多做一步用“单张图片批量回测”把模型边界摸清6.1 写一个轻量的predict.py逐张验证训练完模型、web服务也能跑之后很多人会直接在浏览器里随便写几个数字就宣布“成功了”。但浏览器手写样本太少根本测不出模型边界。我的习惯是先把模型离线验证一遍再做web联调。写一个单独的脚本把所有验证集图片逐张喂给模型统计每个类别的准确率顺便把预测错的图打印出来看一眼。import torch from PIL import Image import torchvision.transforms as T transform T.Compose([ T.Grayscale(num_output_channels1), T.Resize((28, 28)), T.ToTensor(), T.Normalize((0.5,), (0.5,)) ]) model DigitCNN() model.load_state_dict(torch.load(best_model.pth, map_locationcpu)) model.eval() class_correct [0] * 10 class_total [0] * 10 with open(val.txt, r, encodingutf-8) as f: for line in f: path, label line.strip().split() img transform(Image.open(path).convert(L)).unsqueeze(0) with torch.no_grad(): pred model(img).argmax(dim1).item() label int(label) class_total[label] 1 if pred label: class_correct[label] 1 else: print(fmisclassified: {path}, true{label}, pred{pred}) for i in range(10): if class_total[i] 0: acc class_correct[i] / class_total[i] print(fdigit {i}: {class_correct[i]}/{class_total[i]} {acc:.2%})这个脚本的价值在最后那几行输出它能告诉你模型到底在哪些数字上容易犯错。如果打印出来的错误集中在“7和1”“0和6”这类外形接近的数字对说明模型学习到的特征分布正常只是相似数字样本不够多如果某个类别准确率特别低多半是该类训练样本太少或者数据集里该类图片本身拍得不清晰。看过这份错误清单后我才能判断问题出在数据、模型还是前后端预处理而不是盲目调参。6.2 画布数字和训练样本的差异意识即使离线批量回测全部类别都能到98%以上浏览器里识别的效果仍然可能打折这不是模型的问题而是“训练分布”和“使用分布”不一致。训练数据里的数字是印刷体或规整手写体居中、粗细均匀、笔画完整浏览器里用户用鼠标写出来的数字往往潦草、笔画细、位置偏。CNN对这类偏移极其敏感所以在资源基础上想要进一步提升真实体验最对症的做法是在训练时加一点“随机扰动增强”对训练图片随机平移、随机旋转几度、随机缩放笔画粗细。这种增强用torchvision.transforms.RandomAffine一两行代码就能实现它带来的真实识别率提升往往比换更大的网络结构更明显。从那以后我每次拿到这类“训练web演示”项目都会强制走一遍“离线批量回测 → 看错误清单 → 再做浏览器测试”的流程。直接打开浏览器乱画几个数字表面上跑通了但模型的短板在哪、真实手感如何心里完全没底。先离线测出边界再上web页面验证交互前后端问题也能更快定位。希望帮到你。本文还有配套的精品资源点击获取
返回列表