☰
PyTorch CNN手写数字识别全流程:从数据整理到网页交互实战
2026/10/1 3:50:53 网站建设 项目流程

简介:这份资源面向希望入门深度学习与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负责把图片路径和标签写成训练用的 txt,02深度学习模型训练.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.html

dataset下面每个数字文件夹就是一个类别,文件夹名就是标签。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", encoding="utf-8") as f: f.write("\n".join(all_samples[:split_idx])) with open(output_val, "w", encoding="utf-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", encoding="utf-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_classes=10): super().__init__() self.features = nn.Sequential( nn.Conv2d(1, 32, kernel_size=3, padding=1), # 输入1通道灰度图 nn.ReLU(), nn.MaxPool2d(2), # 28x28 -> 14x14 nn.Conv2d(32, 64, kernel_size=3, padding=1), 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, padding=1)表示输入 1 通道、输出 32 通道、卷积核 3x3、边缘补 1 圈保持尺寸不变。两次MaxPool2d(2)把 28x28 降到 7x7,最后 flatten 成 6477=3136 维向量送进全连接。如果你的图片不是 28x28,比如是 64x64,那全连接层的输入维度要跟着改,否则会报维度不匹配。训练时的超参一般设 batch_size=32 或 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(f"Epoch {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_location="cpu")) 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()), dtype=torch.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(host="127.0.0.1", port=4399)

这段代码的关键点:map_location="cpu"保证在没 GPU 的机器上也能加载模型;convert("L")把彩色图转灰度,因为训练用的是单通道;resize((28, 28))必须和训练输入尺寸一致,不一致会报维度错误。归一化那步最容易翻车,训练时如果用了transforms.Normalize(mean=[0.5], std=[0.5]),推理时也要做同样的变换,否则识别结果会乱跳。我见过有人训练时归一化了但推理时忘了,模型把 8 认成 3,查了半天以为是模型没训好。

5.2 浏览器端上传图片的格式要求

前端index.html里一般用<input type="file" accept="image/*">让用户选图,然后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,或者直接把脚本里的port=4399改成 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 之间,太低没意义,太高会拒掉很多正确样本。

从那以后我每次跑完训练都会先看混淆矩阵再决定要不要调参,光看准确率数字太容易自我感觉良好了。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询