简介:这份资源面向希望入门深度学习与Web交互的开发者,提供一套基于PyTorch的手写数字识别完整实践方案,涵盖从数据处理、CNN模型训练到网页端部署的全流程。包内共131个文件,以124张jpg图片构成分类数据集,另含3个Python脚本、3个txt说明文本和1个html页面,压缩包约3.88MB,结构紧凑便于快速上手。已有95人学习下载。读者可依次运行数据集文本生成、模型训练与HTML服务脚本,训练过程会输出每个epoch的验证集损失与准确率日志,并保存本地模型;随后通过本地URL在浏览器中打开交互页面,直观体验手写数字识别效果。资源同时附带环境安装说明,适合具备Python基础、想打通CNN训练与Web部署链路的学习者参考实践。
1. 网页版 CNN 手写数字识别:从 zip 包到浏览器里跑通全流程
拿到一个名为「web网页html版通过cnn训练手写数字识别-含图片数据集.zip」的压缩包时,多数人的第一反应是:里面到底是训练脚本还是推理页面?能不能不装 Python 环境,直接在浏览器里画个数字就出结果?这个标题其实指向一条很具体的落地链路——用 CNN 做 MNIST 手写数字识别,把训练好的模型搬到 HTML 页面上,让用户用鼠标写一个数字,前端完成推理并返回预测结果,同时包里附带一份图片数据集供训练和验证使用。
它解决的核心痛点是「演示门槛」:传统做法要装 CUDA、配 PyTorch、跑 Flask 服务,给非技术同事看效果时经常卡在环境上。而 web 网页 html 版把推理环节放到浏览器,打开一个 .html 文件就能交互,适合教学演示、课程作业、内部技术分享。适合谁?一是刚接触 CNN 卷积神经网络、想找一个完整闭环练手的开发者;二是需要快速做原型验证的产品或教研人员;三是手里已有 MNIST 手写数字识别图片数据集,想把它用起来而不是只跑一遍官方示例的人。下面按「数据怎么组织 → 模型怎么训 → 怎么导出到前端 → 页面怎么接 → 坑在哪」的顺序拆开讲。
2. 图片数据集怎么组织:MNIST 目录结构与预处理参数
2.1 压缩包里数据集常见的两种形态
标题里明确写了「含图片数据集」,这意味着它不是直接给你 IDX 格式的二进制文件,而是按类别分文件夹的图片。常见做法是dataset/train/0/到dataset/train/9/十个子目录,每个目录下是若干张 28×28 的灰度 PNG 或 JPG;测试集同理放在dataset/test/下。这种结构的好处是肉眼可查、方便增删样本,坏处是读取比 IDX 慢,需要自己写 Dataset 类。
先确认目录层级,别急着写训练代码:
# 查看压缩包解压后的目录结构,确认图片是按类别分文件夹 find dataset -maxdepth 2 -type d | sort # 统计每个类别下的图片数量,检查是否严重不均衡 for d in dataset/train/*/; do echo -n "$d "; ls "$d" | wc -l; done第一段命令列出两级目录,确认train和test下确实是 0-9 十个文件夹;第二段统计每类样本数。MNIST 原始训练集每类约 6000 张,如果某个类别只有几百张,训练时就要考虑加权采样或数据增强,否则模型会偏向样本多的类。
2.2 预处理必须对齐训练与推理
图片数据集最容易翻车的地方是「训练时一套预处理,推理时另一套」。CNN 对手写数字的输入要求通常是:灰度单通道、尺寸 28×28、像素值归一化到 [0,1] 或标准化到均值 0.1307、标准差 0.3081。训练脚本里这样写:
import torch from torchvision import transforms from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder # 训练集预处理:转灰度、缩放到28、转张量、标准化 train_tf = transforms.Compose([ transforms.Grayscale(num_output_channels=1), transforms.Resize((28, 28)), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) # 测试集只做相同变换,不做随机增强 test_tf = transforms.Compose([ transforms.Grayscale(num_output_channels=1), transforms.Resize((28, 28)), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_ds = ImageFolder('dataset/train', transform=train_tf) test_ds = ImageFolder('dataset/test', transform=test_tf) train_loader = DataLoader(train_ds, batch_size=64, shuffle=True, num_workers=2) test_loader = DataLoader(test_ds, batch_size=256, shuffle=False, num_workers=2)Grayscale(1)保证三通道图片被压成单通道,避免前端传灰度图、训练用 RGB 导致通道数不匹配;Resize((28,28))统一尺寸;Normalize的两个数值是 MNIST 全局统计量,训练和推理必须一致,否则预测结果会整体偏移。batch_size=64是显存和收敛速度的折中,num_workers=2在 Windows 上如果报错就改成 0。
提示:如果压缩包里的图片已经是 28×28 灰度图,
Resize可以保留但不会改变尺寸;如果图片是白底黑字而 MNIST 是黑底白字,需要在预处理里加反色,否则模型学到的特征完全相反。
3. 用 CNN 训练手写数字识别:网络结构与关键超参
3.1 一个够用又不臃肿的 CNN 结构
MNIST 手写数字识别不需要 ResNet 这种量级,两层卷积加两层全连接就能到 99% 以上。结构设计上,卷积层负责提取笔画边缘和局部形状,池化层降维,全连接层做分类。下面这个结构在 CPU 上几分钟就能训完一轮:
import torch.nn as nn import torch.nn.functional as F class DigitCNN(nn.Module): def __init__(self): super().__init__() # 第一层卷积:1通道输入,16个3x3卷积核 self.conv1 = nn.Conv2d(1, 16, kernel_size=3, padding=1) self.conv2 = nn.Conv2d(16, 32, kernel_size=3, padding=1) self.pool = nn.MaxPool2d(2, 2) self.fc1 = nn.Linear(32 * 7 * 7, 128) self.fc2 = nn.Linear(128, 10) self.dropout = nn.Dropout(0.25) 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, 32 * 7 * 7) # 展平 x = F.relu(self.fc1(x)) x = self.dropout(x) x = self.fc2(x) return xpadding=1保证 3×3 卷积后尺寸不变,两次 2×2 池化把 28×28 降到 7×7,所以全连接输入是 32×7×7。Dropout(0.25)放在全连接层之间,抑制过拟合。如果数据集里样本较少,可以把卷积核数量减半,避免参数过多。
3.2 训练循环与必调参数
训练脚本的核心是损失函数、优化器和学习率。分类任务用交叉熵,优化器用 Adam 起步,学习率 1e-3:
import torch from torch import optim device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = DigitCNN().to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=1e-3) for epoch in range(10): model.train() for imgs, labels in train_loader: imgs, labels = imgs.to(device), labels.to(device) optimizer.zero_grad() out = model(imgs) loss = criterion(out, labels) loss.backward() optimizer.step() # 每轮结束在测试集上评估 model.eval() correct = total = 0 with torch.no_grad(): for imgs, labels in test_loader: imgs, labels = imgs.to(device), labels.to(device) pred = model(imgs).argmax(dim=1) correct += (pred == labels).sum().item() total += labels.size(0) print(f'epoch {epoch+1}, acc {correct/total:.4f}')lr=1e-3是 Adam 的常用起点,如果 loss 震荡就降到 1e-4;epoch=10对 MNIST 足够,验证准确率通常在 98.5% 到 99.2% 之间。每轮评估时切到model.eval()并关闭梯度,否则 Dropout 和 BatchNorm 行为不一致,准确率会偏低。
注意:如果训练准确率很高但测试准确率明显低,优先检查测试集预处理是否和训练集完全一致,而不是急着加正则。
4. 从 PyTorch 模型到 HTML 可用的推理文件
4.1 导出 ONNX 而不是直接存 state_dict
HTML 页面要跑推理,不能依赖 PyTorch 运行时。常见做法是导出 ONNX,再用 onnxruntime-web 在浏览器加载。导出时固定输入形状为[1,1,28,28],动态 batch 在网页端用不到:
import torch model.eval() dummy = torch.randn(1, 1, 28, 28) torch.onnx.export( model, dummy, 'digit_cnn.onnx', input_names=['input'], output_names=['output'], opset_version=11, dynamic_axes=None )opset_version=11兼容性较好,onnxruntime-web 支持稳定。dynamic_axes=None表示输入尺寸固定,网页端每次只推理一张图,不需要动态轴。导出后可以用onnx.checker验证文件完整性,避免前端加载时报模型解析错误。
4.2 前端页面加载模型的最小结构
HTML 页面需要引入 onnxruntime-web 的脚本,然后在 canvas 上监听鼠标事件,把绘制结果转成 28×28 灰度数组。核心逻辑如下:
<!DOCTYPE html> <html lang="zh-cn"> <head> <meta charset="utf-8"> <title>手写数字识别</title> <script src="https://cdn.jsdelivr.net/npm/onnxruntime-web/dist/ort.min.js"></script> </head> <body> <canvas id="pad" width="280" height="280" style="border:1px solid #ccc"></canvas> <button id="predict">识别</button> <div id="result"></div> <script> // 初始化画布为黑底白字,与MNIST一致 const canvas = document.getElementById('pad'); const ctx = canvas.getContext('2d'); ctx.fillStyle = '#000'; ctx.fillRect(0, 0, 280, 280); ctx.strokeStyle = '#fff'; ctx.lineWidth = 18; ctx.lineCap = 'round'; let drawing = false; canvas.addEventListener('mousedown', e => { drawing = true; ctx.beginPath(); ctx.moveTo(e.offsetX, e.offsetY); }); canvas.addEventListener('mousemove', e => { if (drawing) { ctx.lineTo(e.offsetX, e.offsetY); ctx.stroke(); } }); canvas.addEventListener('mouseup', () => drawing = false); async function predict() { // 缩小到28x28并取灰度 const small = document.createElement('canvas'); small.width = 28; small.height = 28; const sctx = small.getContext('2d'); sctx.drawImage(canvas, 0, 0, 28, 28); const imgData = sctx.getImageData(0, 0, 28, 28).data; const input = new Float32Array(28 * 28); for (let i = 0; i < 28 * 28; i++) { // 取红色通道,归一化到[0,1],再按训练均值方差标准化 input[i] = (imgData[i * 4] / 255 - 0.1307) / 0.3081; } const tensor = new ort.Tensor('float32', input, [1, 1, 28, 28]); const session = await ort.InferenceSession.create('digit_cnn.onnx'); const out = await session.run({ input: tensor }); const logits = out.output.data; let best = 0; for (let i = 1; i < 10; i++) if (logits[i] > logits[best]) best = i; document.getElementById('result').innerText = '预测:' + best; } document.getElementById('predict').onclick = predict; </script> </body> </html>画布初始化成黑底白字,和 MNIST 训练数据一致;drawImage把 280×280 缩到 28×28,相当于前端做了一次 Resize;归一化和标准化公式必须和训练脚本里的Normalize完全对应,否则输入分布偏移,预测会乱。ort.InferenceSession.create每次点击都创建会话效率低,实际项目里应该在页面加载时创建一次并缓存。
5. 避坑与排查:网页版 CNN 识别最常见的 5 个翻车点
5.1 现象:页面能画但预测永远是同一个数字
原因通常是输入张量的形状或通道顺序不对。ONNX 模型期望[1,1,28,28],如果前端传成[1,28,28,1]或[784],onnxruntime 可能不报错但输出恒定。解决:在new ort.Tensor时打印 shape,确认是[1,1,28,28];同时检查input数组长度是否为 784。
5.2 现象:训练准确率 99%,网页上识别率很低
这是血泪经验里最常见的一条。训练时用了Normalize((0.1307,),(0.3081,)),前端只做了/255没做标准化,输入分布差了一个量级。解决:前端归一化公式写成(pixel/255 - 0.1307) / 0.3081,两个常数从训练脚本里抄,不要凭记忆写。
5.3 现象:onnxruntime-web 加载模型报 404 或跨域
原因是用file://直接打开 HTML 时,浏览器禁止加载同目录下的 .onnx 文件。解决:用python -m http.server 8000在目录下起一个本地静态服务,通过http://localhost:8000/页面.html访问;或者把模型转成 base64 内嵌,但文件会变大。
5.4 现象:画布上写的数字偏上或偏左,识别不准
MNIST 的数字是居中且经过尺寸归一化的,而用户在 canvas 上随手写的位置不固定。解决:在缩小到 28×28 之前,先计算笔画的包围盒,把它裁剪出来再等比缩放到 20×20,最后放到 28×28 画布中央,模拟 MNIST 的居中效果。
5.5 现象:第一次点击识别很慢,后面正常
原因是每次点击都InferenceSession.create,重复加载模型。解决:把 session 创建放在页面初始化阶段,用全局变量保存,点击时只调session.run。如果模型较大,可以在页面上加一个加载状态提示。
6. 把识别准确率再往上推:数据增强与前端预处理的联动技巧
训练侧还有一点余量可以挖。如果压缩包里的图片数据集样本量比原始 MNIST 少,或者包含一些拍摄的手写数字,直接训练容易过拟合。我一般会加轻量数据增强:随机旋转 ±10 度、随机平移 10% 以内、随机缩放 0.9 到 1.1。注意增强只加在训练集,测试集和前端推理保持原始变换。旋转角度不要超过 15 度,否则 6 和 9 会互相混淆,这是实际调参时踩过的坑。
前端侧有一个容易被忽略的技巧:把用户绘制过程做一次「笔画居中 + 尺寸归一化」。具体做法是遍历 28×28 灰度数组,找到非零像素的最小/最大行列,裁剪后缩放到 20×20,再粘贴到 28×28 中心。这样即使用户写在画布角落,输入分布也更接近训练数据。实测这个改动能让网页端识别率提升几个百分点,尤其是对写得偏小的数字。
验证方法上,不要只看单张。准备 20 张测试集图片,用脚本批量走一遍前端相同的预处理和 ONNX 推理,统计准确率,和 PyTorch 直接推理的结果对比。如果两者差距超过 1%,说明前端预处理和训练预处理有偏差,回去逐项核对归一化、通道顺序和尺寸。
| 检查项 | 训练侧 | 前端侧 | 不一致的后果 |
|---|---|---|---|
| 通道数 | 1 | 1 | 形状报错或输出恒定 |
| 尺寸 | 28×28 | 28×28 | 全连接层维度不匹配 |
| 归一化 | /255 | /255 | 输入范围差 255 倍 |
| 标准化 | (x-0.1307)/0.3081 | 同左 | 分布偏移,预测乱 |
| 颜色 | 黑底白字 | 黑底白字 | 特征相反,准确率骤降 |
最后说一个习惯:每次改完前端预处理,我一定先拿训练集里的一张图,用前端同一套代码跑一遍,看预测是否和训练时一致。这一步能挡住大部分「模型没问题、页面有问题」的玄学故障。希望帮到你。
本文还有配套的精品资源,点击获取