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

发布时间:2026/10/1 11:33:18
网页版手写数字识别:从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_channels1), # 强制单通道 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_channels1), transforms.Resize((28, 28)), transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) train_dataset datasets.ImageFolder(root./data/train, transformtrain_transform) test_dataset datasets.ImageFolder(root./data/test, transformtest_transform) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue, num_workers2) test_loader DataLoader(test_dataset, batch_size64, shuffleFalse, num_workers2) print(f训练集样本数: {len(train_dataset)}, 类别: {train_dataset.classes})这段代码的关键参数有三个Grayscale(num_output_channels1)确保输入是单通道因为 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_size3, padding1) self.bn1 nn.BatchNorm2d(16) # 第二个卷积块16通道输入32个3x3卷积核 self.conv2 nn.Conv2d(16, 32, kernel_size3, padding1) 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这个结构里padding1保证卷积后尺寸不变两次MaxPool2d(2,2)把 28x28 降到 7x7全连接层输入维度就是32*7*71568。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(), lr0.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(fEpoch {epoch1}, 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_size64是显存和训练速度的平衡点显存小于 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_version11 ) print(ONNX 导出完成)opset_version11是兼容性最好的选择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(fONNX 推理输出形状: {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_formattf_saved_model \ --output_formattfjs_graph_model \ saved_model \ web_model转换完成后web_model文件夹里会有model.json和若干.bin权重文件。HTML 页面里加载模型的代码!DOCTYPE html html langzh-cn head meta charsetutf-8 title手写数字识别/title script srchttps://cdn.jsdelivr.net/npm/tensorflow/tfjs4.0.0/dist/tf.min.js/script /head body canvas idcanvas width280 height280 styleborder:1px solid #ccc;/canvas button idpredict识别/button div idresult等待输入.../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-1sub(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_bytes1000000控制分片大小。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[idx1] data[idx2]) / 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 张就重新训练一轮。网页版手写数字识别这个项目模型结构其实不是瓶颈真正的功夫都在数据预处理和输入对齐上。希望帮到你。本文还有配套的精品资源点击获取