PyTorch CNN手写数字识别全流程:从数据整理到网页交互实战

📅 发布时间:2026/10/1 3:50:57
PyTorch CNN手写数字识别全流程:从数据整理到网页交互实战
简介这份资源面向希望入门深度学习与Web交互的开发者提供一套基于PyTorch的手写数字识别完整实践方案涵盖从数据处理、CNN模型训练到网页端部署的全流程。压缩包共131个文件以124张jpg图片构成分类数据集另含3个txt说明与日志、3个Python脚本及1个html页面整体约3.88MB结构紧凑便于快速上手。已有95人学习下载。资源按编号依次提供数据集文本生成、模型训练与HTML服务脚本训练过程会输出每个epoch的验证集损失与准确率日志并保存本地模型启动服务后可通过本机浏览器访问交互页面直观体验识别效果。适合作为CNN图像分类的练手项目帮助读者理解数据组织、训练评估与前后端联调的完整链路。1. 从一堆散图到网页可交互这个 CNN 手写数字识别包到底能跑通什么如果你手头有一批按类别分文件夹存放的手写数字图片想快速验证「数据整理 → CNN 训练 → 网页端实时识别」这条链路而不是从零去搭 Flask 或 FastAPI 的接口那这个包值得拆开看看。它把三件事串成了一条线01数据集文本生成制作.py负责把图片路径和标签写成训练用的 txt02深度学习模型训练.py用 PyTorch 跑 CNN 并把模型和日志落到本地03html_server.py起一个本地 HTTP 服务浏览器打开http://127.0.0.1:4399就能上传图片看识别结果。技术栈是 Python PyTorch 原生 HTML 前端没有额外的前端框架依赖适合刚接触 CNN 分类、想找一个能跑通全流程的练手项目也适合需要快速给非技术同事演示识别效果的场景。数据集里已经带了005.jpg、048.jpg、031.jpg这类按数字命名的样本还有04_flip.jpg这种翻转增强图说明作者在数据层面已经考虑过简单增广。下面按实际拆包顺序把环境、数据、训练、网页交互和踩坑点逐个讲透。2. 环境配置与依赖锁定requirements.txt 里没写全的坑2.1 为什么不能直接 pip install -r requirements.txt拿到包之后第一反应通常是找requirements.txt然后一把梭。但这个包的依赖文件只列了核心库没有锁版本号也没有区分 CPU 和 GPU 环境。PyTorch 的安装命令跟 CUDA 版本强绑定如果你直接pip install torch默认拉的是 CPU 版训练速度会慢到让你怀疑人生。常见做法是先确认本机显卡驱动支持的 CUDA 版本再去 PyTorch 官网拿对应的安装命令。我一般会先跑nvidia-smi看右上角的 CUDA Version然后去 pytorch.org 选对应版本复制命令。如果机器没有 N 卡那就老老实实用 CPU 版把 batch size 调小一点也能跑。# 先看显卡和 CUDA 版本没有 N 卡就跳过这步 nvidia-smi # 有 N 卡且 CUDA 12.1 的情况用官方源装 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121 # 没有 N 卡或不想折腾直接装 CPU 版 pip install torch torchvision # 其余依赖单独补不要迷信 requirements.txt 的版本 pip install opencv-python pillow numpy flask上面这段命令的逻辑是先探测硬件再决定装哪个变体的 PyTorch。参数上唯一要盯的是--index-url后面的 CUDA 版本号写错了会装成 CPU 版或者直接报找不到包。opencv-python和pillow是图片读取和预处理用的flask是03html_server.py起服务的基础。注意不要同时装opencv-python和opencv-python-headless两者冲突会导致cv2.imshow报错虽然这个项目用不到显示窗口但混装后 import 会出玄学问题。2.2 虚拟环境与目录结构确认强烈建议用 conda 或 venv 建独立环境因为这个包对 PyTorch 版本有隐性要求跟你系统里其他项目的 torch 版本很可能打架。建好环境后把下载的 zip 解压到一个纯英文路径下路径里不要有中文和空格否则01数据集文本生成制作.py在写 txt 时可能因为编码问题把路径写乱。解压后你应该看到类似这样的结构project/ ├── 01数据集文本生成制作.py ├── 02深度学习模型训练.py ├── 03html_server.py ├── requirements.txt ├── dataset/ │ ├── 0/ │ │ ├── 005.jpg │ │ └── 048.jpg │ ├── 1/ │ │ └── 031.jpg │ └── ... └── index.htmldataset下面每个数字文件夹就是一个类别文件夹名就是标签。index.html是前端页面03html_server.py会把它渲染出来。确认结构没问题再往下走不然后面训练时找不到图会报FileNotFoundError回头查路径很浪费时间。3. 数据集文本生成从文件夹到 txt 的转换逻辑与参数3.1 01 脚本到底干了什么01数据集文本生成制作.py的核心任务就一件事遍历dataset下每个类别文件夹把每张图片的路径和对应标签写成一个 txt 文件通常还会按比例拆成训练集和验证集。这个设计的好处是训练脚本不用关心图片怎么存的只读 txt 就行换数据集时只要重新跑一遍这个脚本。常见做法是用os.walk或pathlib遍历然后用random.shuffle打乱后按 8:2 或 7:3 切分。下面是我拆包后还原出来的核心逻辑你可以对照原脚本看import os import random dataset_dir dataset output_train train.txt output_val val.txt val_ratio 0.2 # 验证集比例20% all_samples [] # 遍历每个类别文件夹文件夹名就是标签 for label_name in os.listdir(dataset_dir): class_dir os.path.join(dataset_dir, label_name) if not os.path.isdir(class_dir): continue for img_name in os.listdir(class_dir): if img_name.lower().endswith((.jpg, .png, .jpeg)): img_path os.path.join(class_dir, img_name) all_samples.append(f{img_path}\t{label_name}) random.shuffle(all_samples) # 打乱顺序避免按类别聚集 split_idx int(len(all_samples) * (1 - val_ratio)) with open(output_train, w, encodingutf-8) as f: f.write(\n.join(all_samples[:split_idx])) with open(output_val, w, encodingutf-8) as f: f.write(\n.join(all_samples[split_idx:])) print(f总样本 {len(all_samples)}训练集 {split_idx}验证集 {len(all_samples)-split_idx})这段代码的关键参数是val_ratio设 0.2 意味着 80% 训练、20% 验证。如果你的数据集本身很小比如每个数字只有几十张那验证集可能只有几张图评估结果波动会很大这时候可以调到 0.1 或者用交叉验证。另一个要注意的是random.shuffle之前没有设随机种子每次跑生成的 txt 都不一样想复现实验就加一行random.seed(42)。txt 的格式是「路径 tab 标签」训练脚本按 tab 切分所以路径里不能有 tab 字符Windows 路径里的反斜杠在 Python 字符串里也要注意转义。3.2 标签映射与类别不均衡的处理这个包默认用文件夹名当标签也就是0到9这十个字符串。训练时 PyTorch 的CrossEntropyLoss需要标签是 0 到 9 的整数所以02脚本里会有一个label_map把字符串转成索引。如果你自己加了一个10文件夹想识别两位数那标签映射就要改成动态生成不能写死。常见做法是先把所有类别名排序然后{name: idx for idx, name in enumerate(sorted(classes))}。另外如果某些数字的样本特别少比如8只有 20 张而1有 200 张训练时模型会偏向多数类验证集准确率看着高但实际对8的识别很差。解决办法是在01脚本里做欠采样或过采样或者训练时给CrossEntropyLoss传weight参数。我一般会先跑一遍统计看看每个类别的数量差三倍以上就要处理。from collections import Counter labels [line.split(\t)[1].strip() for line in open(train.txt, encodingutf-8)] count Counter(labels) print(count) # 看每个数字有多少张差太多就要做均衡这段统计代码不复杂但很实用跑完你心里就有数了。如果发现不均衡最简单的做法是在01脚本里对少数类重复采样让每个类别的样本数接近虽然会引入过拟合风险但比模型完全学不会少数类要好。4. CNN 模型训练网络结构、超参与日志解读4.1 02 脚本里的 CNN 长什么样02深度学习模型训练.py是核心它读train.txt和val.txt用 PyTorch 定义 CNN跑若干个 epoch最后保存模型权重和日志。手写数字识别的 CNN 通常不会太深两层卷积加两层全连接就够了输入是 28x28 灰度图。下面是我根据常见实现还原的结构你对照原脚本看层数和通道数是否一致import torch import torch.nn as nn class SimpleCNN(nn.Module): def __init__(self, num_classes10): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 32, kernel_size3, padding1), # 输入1通道灰度图 nn.ReLU(), nn.MaxPool2d(2), # 28x28 - 14x14 nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(2), # 14x14 - 7x7 ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(64 * 7 * 7, 128), nn.ReLU(), nn.Linear(128, num_classes), ) def forward(self, x): x self.features(x) x self.classifier(x) return x这个结构的参数含义Conv2d(1, 32, 3, padding1)表示输入 1 通道、输出 32 通道、卷积核 3x3、边缘补 1 圈保持尺寸不变。两次MaxPool2d(2)把 28x28 降到 7x7最后 flatten 成 64773136 维向量送进全连接。如果你的图片不是 28x28比如是 64x64那全连接层的输入维度要跟着改否则会报维度不匹配。训练时的超参一般设 batch_size32 或 64学习率 0.001优化器用 Adam损失函数用CrossEntropyLoss。epoch 数看数据集大小几千张图跑 10 到 20 个 epoch 就收敛了。4.2 训练日志里该盯哪几个数训练完成后本地会生成 log 文件里面记录了每个 epoch 的验证集损失和准确率。很多人只看准确率觉得到 99% 就万事大吉但损失值不降反升往往是过拟合的信号。我一般会同时看训练损失和验证损失如果训练损失一直降而验证损失先降后升说明模型开始死记硬背训练集这时候要么加 dropout要么早停。日志里如果出现准确率在某个 epoch 后剧烈波动比如从 98% 掉到 85%常见原因是学习率太大可以试试调小到 0.0001 或者加学习率衰减。另外如果验证准确率一直卡在 10% 左右那基本是标签映射错了模型在瞎猜回去检查label_map和 txt 里的标签是否对得上。# 训练循环里记录日志的常见写法 for epoch in range(epochs): model.train() for imgs, labels in train_loader: # ... 前向、反向、优化 ... pass model.eval() val_loss, correct, total 0, 0, 0 with torch.no_grad(): for imgs, labels in val_loader: outputs model(imgs) loss criterion(outputs, labels) val_loss loss.item() correct (outputs.argmax(1) labels).sum().item() total labels.size(0) acc correct / total print(fEpoch {epoch}: val_loss{val_loss:.4f}, val_acc{acc:.4f}) # 把上面这行同时写进 log 文件这段代码里model.eval()和torch.no_grad()必须成对出现否则验证时会更新梯度还占用显存。argmax(1)取的是第二个维度的最大值索引也就是预测类别。日志建议同时写文件和打印到控制台方便你边跑边看不用等跑完再翻文件。5. 网页交互与本地服务03 脚本起服务后排错指南5.1 03html_server.py 的请求链路03html_server.py干的事是起一个 HTTP 服务把index.html返回给浏览器同时提供一个接收图片并返回识别结果的接口。前端页面上通常有一个文件选择框和一个「识别」按钮用户选图后通过fetch或表单提交把图片传到后端后端用训练好的模型推理再把结果返回给页面显示。这个链路里最容易出问题的是端口占用和跨域。端口 4399 如果被其他程序占了服务起不来报Address already in use解决办法是改脚本里的端口号或者把占用进程杀掉。跨域问题在本地开发时一般不会遇到因为前后端同源但如果你把index.html单独用文件方式打开而不是通过服务访问那fetch就会因为file://协议被浏览器拦截。from flask import Flask, request, jsonify, render_template import torch from PIL import Image import io app Flask(__name__) model SimpleCNN() model.load_state_dict(torch.load(model.pth, map_locationcpu)) model.eval() app.route(/) def index(): return render_template(index.html) app.route(/predict, methods[POST]) def predict(): file request.files[image] img Image.open(io.BytesIO(file.read())).convert(L).resize((28, 28)) # 转成 tensor 并归一化具体归一化参数要和训练时一致 tensor torch.tensor(list(img.getdata()), dtypetorch.float32).view(1, 1, 28, 28) / 255.0 with torch.no_grad(): output model(tensor) pred output.argmax(1).item() return jsonify({digit: pred}) if __name__ __main__: app.run(host127.0.0.1, port4399)这段代码的关键点map_locationcpu保证在没 GPU 的机器上也能加载模型convert(L)把彩色图转灰度因为训练用的是单通道resize((28, 28))必须和训练输入尺寸一致不一致会报维度错误。归一化那步最容易翻车训练时如果用了transforms.Normalize(mean[0.5], std[0.5])推理时也要做同样的变换否则识别结果会乱跳。我见过有人训练时归一化了但推理时忘了模型把 8 认成 3查了半天以为是模型没训好。5.2 浏览器端上传图片的格式要求前端index.html里一般用input typefile acceptimage/*让用户选图然后FormData打包发送。这里要注意的是用户选的图可能是任意尺寸和格式后端必须做兼容处理。如果用户传了一张 4000x3000 的彩色照片后端直接 resize 到 28x28 会丢失大量信息识别率会很低。常见做法是在前端先用 canvas 把图片缩放到 28x28 再上传或者后端加一步自适应二值化把背景和笔画分离。这个包默认假设用户上传的是类似数据集里的手写数字图背景干净、笔画清晰如果你拿一张复杂背景的照片去测识别不准是正常的不是模型的问题。提示服务起来后如果浏览器打不开http://127.0.0.1:4399先检查终端有没有报错再确认端口没被占用最后看防火墙有没有拦本地回环。6. 避坑与常见问题排查6.1 训练时 loss 不降反升现象跑了几十个 epoch训练损失一直在 2.3 左右震荡准确率跟随机猜差不多。原因通常是学习率设太大或者数据标签对不上。先检查 txt 里的标签是不是 0 到 9 的字符串再看label_map有没有把0映射成 0。如果标签没问题把学习率从 0.01 降到 0.001 或 0.0001 再试。另一个隐蔽原因是图片读取时通道数不对比如用cv2.imread读出来是 BGR 三通道但模型第一层是Conv2d(1, ...)维度不匹配会直接报错而不是 loss 不降所以这个可能性较小。6.2 验证集准确率很高但网页识别全错现象日志里 val_acc 到 99%但网页上传图片识别结果乱七八糟。原因几乎可以肯定是推理时的预处理和训练时不一致。训练时可能用了ToTensor()自动归一化到 [0,1]推理时如果忘了除 255输入值域变成 [0,255]模型直接懵了。解决办法是把训练时的 transform 管道原样复制到推理代码里或者手动做同样的除 255 和归一化。另外检查 resize 的插值方式训练用Image.BILINEAR推理也要用同一种用NEAREST会引入锯齿导致识别偏差。6.3 03 脚本启动报端口被占用现象运行03html_server.py后终端报OSError: [Errno 98] Address already in use。原因是 4399 端口被其他进程占了可能是上次没关干净的服务也可能是别的软件。解决办法是在终端跑lsof -i:4399找到 PID 然后kill -9或者直接把脚本里的port4399改成 4400 或其他空闲端口。改端口后浏览器地址也要跟着改别忘了。6.4 数据集里图片格式不统一导致读取失败现象01脚本跑一半报UnidentifiedImageError或cannot identify image file。原因是数据集里混了非图片文件或者有些 jpg 其实是 webp 改了后缀。解决办法是在遍历时加 try-except 跳过坏图或者用PIL.Image.open验证后再写入 txt。我一般会在01脚本里加一行Image.open(img_path).verify()做校验坏图直接打印路径跳过不中断整个流程。6.5 模型保存后加载报 key 不匹配现象02脚本保存的模型在03脚本里load_state_dict时报Missing key(s)或Unexpected key(s)。原因是保存时用了torch.save(model, model.pth)保存整个模型对象而加载时用了load_state_dict两者格式不兼容。正确做法是保存时用torch.save(model.state_dict(), model.pth)加载时先实例化模型再load_state_dict。如果已经保存错了可以用torch.load加载整个对象然后取.state_dict()补救。7. 进阶技巧用混淆矩阵定位模型到底把哪个数字认错了跑通全流程之后光看一个总准确率是不够的你根本不知道模型在哪些数字上容易翻车。我习惯在验证阶段加一个混淆矩阵把每个类别的预测分布打出来。具体做法是收集所有验证集的预测结果和真实标签用 sklearn 的confusion_matrix算一下然后打印成表格。这样一眼就能看出是 6 和 8 混了还是 1 和 7 不分。下面是我常用的代码片段from sklearn.metrics import confusion_matrix import numpy as np all_preds, all_labels [], [] model.eval() with torch.no_grad(): for imgs, labels in val_loader: outputs model(imgs) preds outputs.argmax(1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) cm confusion_matrix(all_labels, all_preds) print(混淆矩阵行真实列预测) for i, row in enumerate(cm): print(f{i}: {row})这段代码跑完你会得到 10x10 的矩阵对角线是正确识别的数量非对角线就是错分。如果发现 6 被大量认成 8那说明模型对这两个数字的区分特征学得不够可以考虑在数据增强里加一点旋转或弹性形变让模型见过更多 6 的变体。如果某个数字整行都是 0那说明验证集里根本没有这个类别的样本回去检查01脚本的切分逻辑是不是把某个类全分到训练集了。另一个实用技巧是给推理加一个置信度阈值。网页端返回结果时如果模型对最高分的置信度低于 0.6就提示「不确定请重新上传」而不是硬给一个错误答案。这个阈值怎么定跑一遍验证集看正确识别的样本里最低置信度是多少取那个值往下浮一点就行。我一般会设 0.5 到 0.7 之间太低没意义太高会拒掉很多正确样本。从那以后我每次跑完训练都会先看混淆矩阵再决定要不要调参光看准确率数字太容易自我感觉良好了。希望帮到你。本文还有配套的精品资源点击获取