news 2026/10/1 11:33:17

网页版手写数字识别:从MNIST数据集到CNN推理的完整落地指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
网页版手写数字识别:从MNIST数据集到CNN推理的完整落地指南

简介:这份资源面向希望入门深度学习与Web交互的开发者,提供一套基于PyTorch的手写数字识别完整项目,涵盖从数据处理到网页端展示的全流程。包内共131个文件,以124张jpg图片构成分类数据集,另含3个Python脚本、3个txt说明与1个html页面,压缩包约3.88MB,结构轻量、便于本地运行。项目依次通过数据集文本生成、CNN模型训练与HTML服务启动三个脚本串联,训练过程会输出每个epoch的验证集损失与准确率日志,并保存本地模型,最终生成可交互的网页URL,让读者直观体验模型推理效果。已有95人学习,适合作为课程设计、毕业项目或CNN实战练手素材,帮助理解图像分类数据组织、模型训练与前后端联调的关键环节。

1. 网页版手写数字识别:从 MNIST 图片数据集到 CNN 推理的完整落地

很多人第一次接触深度学习,都是从手写数字识别开始的。但真正把它做成一个「打开浏览器就能用」的网页版工具,中间要跨过的坑远比想象中多:图片数据集怎么组织、CNN 模型怎么训练、训练好的权重怎么塞进 HTML 页面、用户手写的数字怎么预处理成模型能吃的张量。这套「web 网页 html 版通过 CNN 训练手写数字识别 + 图片数据集」的方案,解决的正是这条从数据到网页推理的完整链路。它适合两类人:一是想找一个能跑通、能改、能展示的深度学习入门项目的前端或全栈工程师;二是已经会写 PyTorch 或 TensorFlow 训练脚本,但不知道怎么把模型搬到浏览器里给非技术用户用的算法同学。核心思路不复杂——用 Python 训练一个轻量 CNN,导出成 ONNX 或 TF.js 格式,再用一个纯 HTML 页面加载模型、接收 canvas 手写输入、实时输出识别结果。整套东西不需要后端服务器,双击 HTML 文件就能跑。

2. 图片数据集怎么组织:MNIST 的目录结构与预处理流水线

2.1 为什么不能直接把 PNG 丢给训练脚本

MNIST 原始格式是 IDX 二进制文件,不是常见的图片文件夹结构。很多网上下载的「手写数字图片数据集.zip」解压后是一堆 28x28 的 PNG,按 0-9 分文件夹存放。这种结构对人友好,但对训练脚本来说需要额外写 Dataset 类去遍历目录、读取图片、转灰度、归一化。更关键的是,如果数据集里混入了非 28x28 的图片,或者灰度值范围不是 0-255,训练时会出现 loss 不下降或者准确率卡在 10% 的玄学现象。我一般会先写一个数据审计脚本,把所有图片的尺寸、通道数、像素值范围统计一遍,确认没有脏数据再开始训练。

2.2 用 torchvision 构建可复现的数据加载器

假设数据集已经按train/0/、train/1/……test/0/、test/1/的目录结构组织好了,下面这段代码可以直接抄作业:

import torch from torchvision import datasets, transforms from torch.utils.data import DataLoader # 训练集变换:转灰度、转张量、归一化到 [-1, 1] train_transform = transforms.Compose([ transforms.Grayscale(num_output_channels=1), # 强制单通道 transforms.Resize((28, 28)), # 统一尺寸 transforms.ToTensor(), # 像素值从 0-255 转到 0-1 transforms.Normalize((0.5,), (0.5,)) # 再转到 [-1, 1] ]) # 测试集用同样的变换,保证分布一致 test_transform = transforms.Compose([ transforms.Grayscale(num_output_channels=1), transforms.Resize((28, 28)), transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) train_dataset = datasets.ImageFolder(root='./data/train', transform=train_transform) test_dataset = datasets.ImageFolder(root='./data/test', transform=test_transform) train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=2) test_loader = DataLoader(test_dataset, batch_size=64, shuffle=False, num_workers=2) print(f"训练集样本数: {len(train_dataset)}, 类别: {train_dataset.classes}")

这段代码的关键参数有三个:Grayscale(num_output_channels=1)确保输入是单通道,因为 MNIST 是灰度图,如果数据集里混了 RGB 图片不转灰度,后面卷积层的in_channels就对不上;Resize((28, 28))是硬性要求,CNN 的全连接层输入维度是按 28x28 算的,尺寸不对直接报维度错误;Normalize((0.5,), (0.5,))把像素值从 [0,1] 映射到 [-1,1],这个操作对收敛速度影响很大,不归一化的话训练到 5 个 epoch 准确率可能还在 60% 徘徊。

2.3 数据集划分的坑:别让测试集泄漏进训练集

常见做法是 6 万张训练、1 万张测试。但如果你是从网上下载的「图片数据集.zip」,很可能作者已经把训练集和测试集混在一起了。我踩过一次坑:解压后发现所有图片都在一个文件夹里,按文件名前缀分 train 和 test,结果前缀规则不统一,导致部分测试图片混进了训练集,模型在测试集上准确率 99.2%,但实际部署到网页上用户手写的数字识别率只有 70% 左右。后来写了个脚本按文件名哈希重新划分,才把这个问题解决。建议在训练前先跑一遍:

import os, hashlib all_images = [] for root, _, files in os.walk('./data/all'): for f in files: if f.endswith('.png') or f.endswith('.jpg'): all_images.append(os.path.join(root, f)) # 按文件名哈希划分,保证可复现 train_files, test_files = [], [] for path in all_images: h = int(hashlib.md5(os.path.basename(path).encode()).hexdigest(), 16) if h % 10 < 8: train_files.append(path) else: test_files.append(path) print(f"重新划分后 训练: {len(train_files)}, 测试: {len(test_files)}")

3. CNN 模型怎么搭:从 LeNet 到轻量级网页推理的取舍

3.1 网页端推理对模型结构的硬约束

在服务器上跑 CNN,你可以堆到几十层,参数量上亿也无所谓。但要把模型塞进 HTML 页面,让浏览器用 WebAssembly 或 WebGL 跑推理,模型大小最好控制在 5MB 以内,参数量控制在 100 万以下。LeNet-5 是经典选择:两个卷积层、两个池化层、三个全连接层,参数量约 6 万,导出成 ONNX 后不到 1MB。但 LeNet 的准确率在 MNIST 上大概 98.5% 左右,如果数据集质量一般,可能掉到 97%。我一般会在 LeNet 基础上加一层 BatchNorm 和 Dropout,准确率能拉到 99% 以上,参数量只增加几千。

3.2 可直接复现的 CNN 训练脚本

下面这个模型结构是我在多个网页版手写数字识别项目里反复用过的,平衡了准确率和模型体积:

import torch.nn as nn import torch.nn.functional as F class HandwritingCNN(nn.Module): def __init__(self): super(HandwritingCNN, self).__init__() # 第一个卷积块:1通道输入,16个3x3卷积核 self.conv1 = nn.Conv2d(1, 16, kernel_size=3, padding=1) self.bn1 = nn.BatchNorm2d(16) # 第二个卷积块:16通道输入,32个3x3卷积核 self.conv2 = nn.Conv2d(16, 32, kernel_size=3, padding=1) self.bn2 = nn.BatchNorm2d(32) # 池化层:2x2最大池化 self.pool = nn.MaxPool2d(2, 2) # 全连接层:经过两次池化后,28x28 -> 14x14 -> 7x7 self.fc1 = nn.Linear(32 * 7 * 7, 128) self.dropout = nn.Dropout(0.3) self.fc2 = nn.Linear(128, 10) # 10个数字类别 def forward(self, x): # 第一层:卷积 -> BN -> ReLU -> 池化 x = self.pool(F.relu(self.bn1(self.conv1(x)))) # 第二层:卷积 -> BN -> ReLU -> 池化 x = self.pool(F.relu(self.bn2(self.conv2(x)))) # 展平 x = x.view(-1, 32 * 7 * 7) # 全连接 -> Dropout -> 输出 x = F.relu(self.fc1(x)) x = self.dropout(x) x = self.fc2(x) return x model = HandwritingCNN() total_params = sum(p.numel() for p in model.parameters()) print(f"模型总参数量: {total_params}") # 约 42 万,导出 ONNX 后约 1.7MB

这个结构里,padding=1保证卷积后尺寸不变,两次MaxPool2d(2,2)把 28x28 降到 7x7,全连接层输入维度就是32*7*7=1568。Dropout(0.3)是防止过拟合的关键,如果训练集只有几千张图片,不加 Dropout 训练准确率能到 100% 但测试集只有 95%。BatchNorm 层在推理时可以融合进卷积层,导出 ONNX 时用torch.onnx.export会自动处理,不会增加推理耗时。

3.3 训练循环与关键超参数

import torch.optim as optim device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = HandwritingCNN().to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=0.001) for epoch in range(15): model.train() running_loss = 0.0 for images, labels in train_loader: images, labels = images.to(device), labels.to(device) optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() # 每个 epoch 结束后在测试集上评估 model.eval() correct, total = 0, 0 with torch.no_grad(): for images, labels in test_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) _, predicted = torch.max(outputs, 1) total += labels.size(0) correct += (predicted == labels).sum().item() acc = 100 * correct / total print(f"Epoch {epoch+1}, Loss: {running_loss/len(train_loader):.4f}, Test Acc: {acc:.2f}%")

学习率设 0.001 配合 Adam 优化器,15 个 epoch 基本能收敛到 99% 以上。如果 loss 震荡厉害,把学习率降到 0.0005;如果收敛太慢,加到 0.002 但不要超过 0.005,否则容易跳过最优解。batch_size=64是显存和训练速度的平衡点,显存小于 4GB 的话改成 32。

4. 模型导出与网页集成:ONNX 转 TF.js 的完整链路

4.1 为什么选 ONNX 作为中间格式

PyTorch 训练出来的.pth文件不能直接在浏览器里跑。常见做法是先把 PyTorch 模型导出成 ONNX,再用onnx-tf转成 TensorFlow SavedModel,最后用tensorflowjs_converter转成 TF.js 能加载的model.json+权重分片文件。这条链路虽然长,但每一步都有成熟的命令行工具,比直接用 PyTorch 的torch.jit导出然后找 JS 运行时靠谱得多。ONNX 的好处是算子标准统一,导出时如果遇到不支持的算子会直接报错,不会等到浏览器里才翻车。

4.2 导出 ONNX 并验证

import torch.onnx # 切换到推理模式,Dropout 和 BatchNorm 行为会改变 model.eval() # 构造一个假输入,维度必须和实际推理时一致 dummy_input = torch.randn(1, 1, 28, 28).to(device) torch.onnx.export( model, dummy_input, "handwriting_cnn.onnx", input_names=["input"], output_names=["output"], dynamic_axes={"input": {0: "batch_size"}, "output": {0: "batch_size"}}, opset_version=11 ) print("ONNX 导出完成")

opset_version=11是兼容性最好的选择,TF.js 对 11 的支持最稳定。dynamic_axes把 batch 维度设为动态,这样网页端可以一次识别一张图,也可以批量识别。导出后建议用onnxruntime跑一遍验证:

import onnxruntime as ort import numpy as np sess = ort.InferenceSession("handwriting_cnn.onnx") test_input = np.random.randn(1, 1, 28, 28).astype(np.float32) result = sess.run(None, {"input": test_input}) print(f"ONNX 推理输出形状: {result[0].shape}") # 应该是 (1, 10)

4.3 转成 TF.js 并在 HTML 里加载

# 安装转换工具 pip install onnx-tf tensorflowjs # ONNX 转 TensorFlow SavedModel onnx-tf convert -i handwriting_cnn.onnx -o saved_model # SavedModel 转 TF.js tensorflowjs_converter --input_format=tf_saved_model \ --output_format=tfjs_graph_model \ saved_model \ web_model

转换完成后web_model文件夹里会有model.json和若干.bin权重文件。HTML 页面里加载模型的代码:

<!DOCTYPE html> <html lang="zh-cn"> <head> <meta charset="utf-8"> <title>手写数字识别</title> <script src="https://cdn.jsdelivr.net/npm/@tensorflow/tfjs@4.0.0/dist/tf.min.js"></script> </head> <body> <canvas id="canvas" width="280" height="280" style="border:1px solid #ccc;"></canvas> <button id="predict">识别</button> <div id="result">等待输入...</div> <script> let model; // 加载 TF.js 模型 async function loadModel() { model = await tf.loadGraphModel('web_model/model.json'); console.log('模型加载完成'); } loadModel(); // 识别按钮点击事件 document.getElementById('predict').onclick = async () => { const canvas = document.getElementById('canvas'); // 从 canvas 获取像素数据,缩放到 28x28 const tensor = tf.browser.fromPixels(canvas, 1) .resizeNearestNeighbor([28, 28]) .toFloat() .div(255.0) // 归一化到 [0,1] .sub(0.5) // 再减 0.5 .div(0.5) // 再除 0.5,等价于 [-1,1] 归一化 .expandDims(0); // 增加 batch 维度 const prediction = model.predict(tensor); const scores = await prediction.data(); const digit = scores.indexOf(Math.max(...scores)); document.getElementById('result').innerText = `识别结果: ${digit}`; }; </script> </body> </html>

这段 HTML 里最关键的是预处理要和训练时完全一致:div(255.0)把像素从 0-255 转到 0-1,sub(0.5).div(0.5)再转到 [-1,1],和训练时的Normalize((0.5,), (0.5,))对应。如果这里少了一步或者顺序错了,识别率会断崖式下跌。canvas 的 280x280 是 28 的 10 倍,方便用户手写,推理前用resizeNearestNeighbor缩到 28x28。

5. 避坑与排查:网页版手写数字识别最常见的 5 个翻车现场

5.1 现象:网页上识别结果永远是同一个数字

原因通常是 canvas 背景色和笔迹颜色反了。MNIST 数据集是黑底白字,但网页 canvas 默认是白底黑字。如果训练时用的是黑底白字,推理时没有做颜色反转,模型看到的输入分布完全不对,输出就会坍缩到一个固定类别。解决办法是在预处理时加一步1 - pixel/255做反转,或者训练时就把数据集反色成白底黑字。

5.2 现象:模型加载成功但 predict 报维度错误

TF.js 的loadGraphModel加载后,输入张量的形状必须和导出时一致。如果导出 ONNX 时 dummy_input 是(1, 1, 28, 28),网页端传进去的也必须是四维张量。常见错误是忘了expandDims(0),传了个(28, 28)进去,报错信息通常是Expected input shape [1,1,28,28] but got [28,28]。排查方法是在model.predict之前打印tensor.shape确认。

5.3 现象:训练准确率 99% 但网页上手写数字识别率不到 80%

这是最典型的「数据集分布和真实输入不匹配」。MNIST 的数字是居中、大小统一、笔画粗细一致的,但用户在 canvas 上写的数字可能偏左、偏小、笔画很细。解决办法有两个:一是在训练时加数据增强,随机平移、缩放、旋转;二是在网页端预处理时做居中裁剪,把用户写的数字从 canvas 里抠出来,缩放到 20x20 再放到 28x28 画布中央,模拟 MNIST 的构图。

5.4 现象:ONNX 转 TF 时报不支持的算子

PyTorch 的某些算子(比如AdaptiveAvgPool2d)在 ONNX opset 11 里没有对应实现,转换时会报Unsupported operator。解决办法是把模型里的自适应池化换成固定尺寸的MaxPool2d或AvgPool2d,因为输入尺寸固定是 28x28,不需要自适应。如果已经用了,改模型结构重新训练比找算子映射表快得多。

5.5 现象:网页打开后模型加载极慢或卡死

TF.js 的模型文件如果超过 10MB,在移动端浏览器上加载会非常慢。检查web_model文件夹里.bin文件的总大小,如果超过 5MB,说明模型参数量太大。回到训练脚本,把全连接层的 128 个神经元降到 64,或者把第二个卷积层的 32 通道降到 16,重新导出。另一个原因是权重分片太多,tensorflowjs_converter默认按 4MB 分片,可以在命令里加--weight_shard_size_bytes=1000000控制分片大小。

6. 进阶技巧:用 Canvas 预处理把网页识别率再拉高 5 个百分点

前面提到用户在 canvas 上写的数字和 MNIST 分布不匹配,最有效的补救措施是在推理前做一次「居中裁剪 + 尺寸归一化」。具体做法是:从 canvas 拿到像素数据后,先扫描所有非背景像素的边界框,把数字区域抠出来,缩放到 20x20,再放到 28x28 的黑色画布正中央。这样处理后的输入和 MNIST 的构图几乎一致,实测识别率能从 82% 提升到 94% 左右。

function preprocessCanvas(canvas) { const ctx = canvas.getContext('2d'); const imageData = ctx.getImageData(0, 0, canvas.width, canvas.height); const data = imageData.data; // 找到非背景像素的边界框 let minX = canvas.width, minY = canvas.height, maxX = 0, maxY = 0; for (let y = 0; y < canvas.height; y++) { for (let x = 0; x < canvas.width; x++) { const idx = (y * canvas.width + x) * 4; // 假设笔迹是深色,背景是浅色,亮度小于 128 视为笔迹 const brightness = (data[idx] + data[idx+1] + data[idx+2]) / 3; if (brightness < 128) { minX = Math.min(minX, x); minY = Math.min(minY, y); maxX = Math.max(maxX, x); maxY = Math.max(maxY, y); } } } // 抠出数字区域并缩放到 20x20 const digitWidth = maxX - minX + 1; const digitHeight = maxY - minY + 1; const tempCanvas = document.createElement('canvas'); tempCanvas.width = 20; tempCanvas.height = 20; const tempCtx = tempCanvas.getContext('2d'); tempCtx.drawImage(canvas, minX, minY, digitWidth, digitHeight, 0, 0, 20, 20); // 放到 28x28 画布中央 const finalCanvas = document.createElement('canvas'); finalCanvas.width = 28; finalCanvas.height = 28; const finalCtx = finalCanvas.getContext('2d'); finalCtx.fillStyle = '#000'; finalCtx.fillRect(0, 0, 28, 28); finalCtx.drawImage(tempCanvas, 4, 4); // 居中偏移 4 像素 return finalCanvas; }

这段代码的核心逻辑是「先找边界,再缩放,最后居中」。brightness < 128是判断笔迹的阈值,如果用户用的是浅色笔迹,需要把阈值调高或者加一个颜色反转。drawImage的九参数版本可以把源 canvas 的指定区域画到目标 canvas 的指定位置,这里把数字区域缩放到 20x20 后,再画到 28x28 画布的 (4,4) 位置,正好居中。实测这个预处理步骤对潦草手写的识别率提升最明显,尤其是数字写得很小或者偏在角落的情况。

另一个技巧是「多模型投票」:训练两个结构略有差异的 CNN(比如一个用 3x3 卷积核,一个用 5x5),推理时两个模型都跑一遍,取平均概率最高的类别。代价是模型体积翻倍,但如果对准确率要求高、不在乎多加载 1MB 权重,这个方案能把识别率再拉高 1-2 个百分点。我一般会在项目里保留一个「高精度模式」开关,默认用单模型,用户手动开启后才加载第二个模型。

最后说一个我自己的习惯:每次改完预处理逻辑,不要只测自己写的数字,找几个同事用鼠标、触控板、手指分别写一遍,把识别错误的样本截图存下来,攒够 50 张就重新训练一轮。网页版手写数字识别这个项目,模型结构其实不是瓶颈,真正的功夫都在数据预处理和输入对齐上。希望帮到你。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/1 11:31:13

数据库第二次作业全攻略:从库表设计到死锁排查

先交代一点背景。数据库第二次作业&#xff0c;放在很多计算机相关专业的培养方案里&#xff0c;正好是从“会写SQL”过渡到“能把数据库用在真实系统里”的那道坎。第一次作业往往是建表、插入、简单查询&#xff0c;第二次作业就开始上强度了&#xff1a;外键约束、索引优化、…

作者头像 李华
网站建设 2026/10/1 11:31:04

Java程序员转战大模型应用团队:一个月真实感受与收藏必备学习资料

作者分享了从Java开发转向大模型应用团队的第一个月的真实体验。文章指出&#xff0c;虽然大模型应用开发不像网上说的那么“高大上”&#xff0c;但确实比传统业务开发更有意思。作者发现&#xff0c;转行并非等于从零开始&#xff0c;技术栈的快速更新和业务问题的解决更能体…

作者头像 李华
网站建设 2026/10/1 11:30:12

MySQL索引之魂:B+树如何用三层结构解决磁盘IO与查询性能难题

1. 从一条慢查询开始&#xff1a;为什么索引结构会成为数据库的命门做后端开发这些年&#xff0c;我见过太多“SQL优化三板斧”式的操作——加索引、改查询、跑EXPLAIN&#xff0c;好像只要把索引列加上就万事大吉。直到有一次&#xff0c;线上一个订单表到了千万级&#xff0c…

作者头像 李华
网站建设 2026/10/1 11:29:40

C++20 Concepts入门:用约束告别模板报错地狱

这套C的模板从入门到放弃&#xff0c;就卡在报错上。每次递归展开几十层&#xff0c;错误信息动辄几百行&#xff0c;看一眼就头大。C20的Concepts甩掉了这口最大的锅——它把对模板参数的约束直接提升成了语言一等公民&#xff0c;让编译器能明确告诉你“你要的int版本不存在&…

作者头像 李华
网站建设 2026/10/1 11:29:37

MySQL视图、存储过程与触发器:边界、代价与避坑指南

如果你经常跟MySQL打交道&#xff0c;一定绕不开视图、存储过程和触发器这三样东西。它们能把复杂的SQL拆成清晰的功能块&#xff0c;也能在你没想到的角落变成性能黑洞。我上一份工作维护的订单库里&#xff0c;几十个视图、七八个存储过程、外加一堆触发器&#xff0c;改一个…

作者头像 李华
网站建设 2026/10/1 11:28:07

VMware Workstation Pro免费授权、下载安装与汉化问题全解

VMware Workstation Pro 现在的下载和授权&#xff0c;确实和两三年前完全不一样了。我在搜安装包的时候也发现&#xff0c;搜索引擎前排一堆第三方下载站&#xff0c;动不动就带个“高速下载器”&#xff0c;点了之后全家桶安排得明明白白。再加上很多人第一次知道 Workstatio…

作者头像 李华