ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

网页版手写数字识别:用TensorFlow.js在浏览器中训练CNN

网页版手写数字识别:用TensorFlow.js在浏览器中训练CNN 简介面向深度学习初学者的手写数字识别PyTorch项目基于CNN模型并附带可直接运行的网页交互界面。压缩包内含图片数据集、环境依赖清单与三个Python脚本依次执行即可完成从数据准备、模型训练到浏览器端识别演示的完整流程适合课程设计或毕设快速复现。资源共131个文件包括124张JPG样本图片、3个TXT路径标签文件、3个PY脚本和1个HTML页面整体仅3.88MB轻量易部署。已有96人学习下载。训练中会自动保存模型并记录每个epoch的验证集损失与准确率日志方便分析收敛曲线网页端通过本地URLhttp://127.0.0.1:4399打开可直接测试手写数字识别效果。整个项目步骤紧凑代码注释清晰从生成标签到模型部署均有对应脚本对希望掌握CNN训练流程、数据集组织方式以及模型与网页交互部署的读者具有较高参考价值。1. 网页版手写数字识别是什么为什么要折腾这个 zip拿到这个 zip 时很多人第一反应是真假手写数字识别这种经典问题不装个 Python 跑 Keras网页上能完成 CNN 训练吗实际用过 TensorFlow.js 就会知道它把深度学习搬到了浏览器而这份 web 网页 html 版的手写数字识别项目正是用浏览器内联的 JavaScript 训练卷积神经网络连图片数据集都一并打好包放在 data 目录下。你不再需要配置 Python 环境、不需要 GPU甚至不需要把数据上传到服务器。解压后双击 index.html从训练、验证到画板识别整个过程都在网页里完成。对课程设计、AI 演示、给前端团队科普 CNN 的人来说这是最轻量的一条入门路径。接下来我会把文件结构、CNN 如何搭建、图片数据集怎么加载与预处理、训练参数怎么调以及网页端训练独有的几个坑全部拆开照着展开就能复现不用猜。2. 拆开 zip 看原理CNN 在浏览器里怎么一步步认出数字2.1 解压后的目录里藏了什么文件结构核对先别急着双击 index.html。凡是带“图片数据集”的网页项目目录结构多半是下面这个样子handwritten_cnn_web/ ├── index.html ├── css/ │ └── style.css ├── js/ │ ├── tf.min.js │ ├── data.js │ ├── model.js │ └── ui.js ├── data/ │ ├── train/ │ │ ├── 0/0001.png, 0002.png, ... │ │ ├── 1/... │ │ └── 9/... │ └── test/ │ ├── 0/... │ └── 9/... └── README.mdindex.html 的 head 里就是meta charsetutf-8这类标准声明没有后端模板所有交互靠原生 JavaScript 完成。各文件作用可以核对下面这张表目录/文件作用index.html页面入口放 canvas 画板、训练按钮、损失曲线展示位css/style.css页面样式控制画板大小和按钮布局js/tf.min.jsTensorFlow.js 库浏览器里的深度学习引擎js/data.js图片路径清单manifest、加载图片的工具函数js/model.jsCNN 模型定义、训练流程、模型保存与加载js/ui.js前端交互画布绘制、识别结果展示、日志打印data/train 与 data/test按标签分好文件夹的图片数据集命名规则是“标签/文件名.png”为什么要强调“图片数据集”很多在线 demo 直接加载 MNIST 二进制文件要在前端解析 idx 格式调试起来很痛苦。把数据做成 PNG 图片按目录存至少在浏览器 Network 面板里能直观看到哪张图 404也方便替换成自己手写的样本。打开项目以后第一步不是看模型代码而是核对js/data.js里的 manifest 路径和 data 目录是否一一对应。路径写错一个字符训练根本跑不起来。2.2 CNN 为什么能稳定识别数字卷积、池化、全连接各管什么事一张 28×28 的手写数字图展开成一个向量就是 784 维但直接接全连接层会把像素的左右相邻关系打散。CNN 的优势在于用滑动卷积核保留空间结构。第一层卷积通常学习的是笔画边缘横、竖、斜、拐角第二层卷积会把边缘组合成小结构比如圆弧、交叉点到了后面的层才逐渐对应到“0”的封闭环、“8”的两个窟窿这类整体形状。池化层不是单纯压缩图片它让模型对笔画粗细和轻微偏移更鲁棒。最大池化在一个小窗口里只取响应最强烈的值相当于告诉网络“这里有没有一个强特征”而不是精确记住它在哪个像素位置。这样参数骤减过拟合风险也降下来。最后必须靠全连接层来分类。卷积层输出的是特征图全连接层把特征图拉平成一个长向量再通过 softmax 输出十个类别的概率。对网页版来说参数少是硬指标。28×28 的灰度图用两个卷积层加一个全连接层就能达到 98% 以上如果把每层 filter 数量从 32 降到 8、16训练速度可以快将近一倍适合现场演示。2.3 浏览器内训练还是模型导入两条路线的取舍同为“web html 版”实际有两种不同做法。第一种是把 Keras 训练好的模型通过tensorflowjs_converter导出成 JSON 加权重文件网页只加载模型做推理。代码量小、执行快但看不到“训练”过程。第二种是直接用 TensorFlow.js 在浏览器里当场训练用户点“开始训练”页面实时打印 loss 和准确率课程设计和演示效果更好缺点是迭代一轮比较慢。这份标题既然叫“通过 cnn 训练手写数字识别”我默认走第二种。把训练集控制在几千张以内一个 epoch 也就是几秒到几十秒。如果纯粹为了识别手写数字先把 batchSize 和 epochs 调小流程通了再加大体验会顺很多。两条路线的对比是这样的对比项浏览器内训练Python 预训练 网页推理环境依赖浏览器即可需要 Python 转换一次训练过程可见可见、可交互不可见训练速度慢受浏览器线程限制快可用 GPU模型体积可动态调整固定可能较大适合场景教学、演示、算法展示生产环境、已有模型浏览器内训练最大的隐藏成本是内存而不是 CPU 速度。数据一次性堆进张量WebGL 后端会把中间结果存成纹理显存占用比你以为的大得多。所以后面所有的代码和参数调整都围绕“在有限资源里把流程跑通”这个目标来设计。3. 图片数据集准备与加载从原始 MNIST 到浏览器能吃的张量3.1 数据集怎么来从 MNIST 批量导出成 PNG或者自建数字样本MNIST 原始包是四个二进制文件画图工具打不开也没法直接在 index.html 里用img引用。为了塞进网页项目常见做法是先把图像批量导出成 28×28 的 PNG再按标签分文件夹。如果你拿到的 zip 里已经整理好了图片可以直接跳到 3.2如果想自己重新生成一份样本或者加入自己手写的数字图片我一般会保留这样一个 Python 脚本在项目里import numpy as np import struct from pathlib import Path from PIL import Image def extract_and_save(image_path, label_path, output_dir, sample_num2000, seed0): # 读取 MNIST 二进制图像文件 with open(image_path, rb) as f: magic, num, rows, cols struct.unpack(IIII, f.read(16)) images np.frombuffer(f.read(rows * cols * num), dtypenp.uint8).reshape(num, rows, cols) # 读取标签文件 with open(label_path, rb) as f: magic, num struct.unpack(II, f.read(8)) labels np.frombuffer(f.read(num), dtypenp.uint8) # 随机抽一部分避免训练集过大 rng np.random.RandomState(seed) idx np.arange(len(labels)) rng.shuffle(idx) idx idx[:sample_num] for index in idx: label int(labels[index]) out_dir Path(output_dir) / str(label) out_dir.mkdir(parentsTrue, exist_okTrue) img Image.fromarray(images[index], modeL) img.save(out_dir / f{index}.png) if __name__ __main__: extract_and_save(train-images-idx3-ubyte, train-labels-idx1-ubyte, data/train, sample_num2000)逻辑说明先按 MNIST 标准的大端格式读出 16 个字节头拿到图像数量、行数、列数再把后续的二进制按uint8解析成[num, rows, cols]的数组。随机抽取 2000 张按标签写到data/train/0、data/train/1这样的目录文件名里带上原始索引保证不会覆盖重复。参数说明sample_num控制总样本数数字越小训练越快但准确率会下降seed固定随机顺序方便复现。如果你拿到的 zip 里已经有完整图片目录这段脚本就当备用工具不需要每次重跑。3.2 目录和 manifest浏览器不能遍历目录那就用数组把路径写出来浏览器出于安全策略不允许 JavaScript 直接读取本地目录列表所以所有图片路径必须在代码里明确列出来。项目里常见做法是在js/data.js里维护一个 manifest 数组const TRAIN_MANIFEST [ { path: data/train/0/0002.png, label: 0 }, { path: data/train/0/0003.png, label: 0 }, // 继续把所有训练图片列进来 { path: data/train/9/0998.png, label: 9 } ]; const TEST_MANIFEST [ { path: data/test/0/0001.png, label: 0 }, // 测试集同理 ];如果图片有两三千张手写 manifest 太容易漏。我一般会用一个小 Python 脚本直接生成这个数组避免自己折腾from pathlib import Path def gen_manifest(data_dir, js_name): entries [] for label_dir in sorted(Path(data_dir).iterdir()): if not label_dir.is_dir(): continue for img_path in sorted(label_dir.iterdir())[:500]: entries.append( f {{ path: {data_dir}/{label_dir.name}/{img_path.name}, label: {label_dir.name} }} ) js_content const TRAIN_MANIFEST [\n ,\n.join(entries) \n]; Path(js_name).write_text(js_content, encodingutf-8) gen_manifest(data/train, data.js)这段脚本把data/train下每个标签目录里的图片列出来输出成 JS 数组。注意路径分隔符统一用 Unix 格式/在 Windows 上用Path生成的\会让浏览器解析失败这是数据加载阶段最容易踩的坑。如果同时有测试集就把data_dir换成data/test函数名改成gen_test_manifest再输出一份 TEST_MANIFEST 即可。3.3 从 PNG 到张量缩放、灰度归一化和 one-hot 标签图片加载的关键是把浏览器 Image 对象读成张量。常见做法是用Image.decode()等图片解码完成后再用tf.browser.fromPixels取像素async function loadImages(manifest) { const xs []; const ys []; for (const item of manifest) { const img new Image(); img.src item.path; // 等图片真正解码完避免拿到空白像素 await img.decode(); const tensor tf.browser.fromPixels(img, 1) .resizeBilinear([28, 28]) .toFloat() .div(tf.scalar(255)); xs.push(tensor); const labelTensor tf.oneHot( tf.tensor1d([item.label], int32), 10 ).squeeze(); ys.push(labelTensor); } return { xs: tf.stack(xs), ys: tf.stack(ys) }; }逻辑说明tf.browser.fromPixels(img, 1)把图片读成单通道灰度第二个参数传1表示通道数如果图片本身是彩色也可以传3保留 RGB但手写数字任务用灰度更稳。resizeBilinear把任意尺寸的图片压成 28×28避免不同来源图片尺寸不一致导致模型输入崩溃。.div(tf.scalar(255))把像素从 0~255 归一到 0~1这一步不做训练时很容易出现 NaN loss。标签用tf.oneHot变成 10 维向量squeeze()确保每一行 label 的 shape 是[10]而不是[1, 10]否则训练时维度对不上。内存方面如果一次性加载三千张 28×28 图片浏览器会创建大量张量对象老手机容易崩。解决办法有三个方向一是控制训练集总量二是用完的中间张量及时dispose()三是用tf.data异步流。对课程设计来说控制图片数量最省事。我一般会先把训练集压到每类 200 张左右跑通整个流程后再考虑加大数据量。4. 在网页里用 TensorFlow.js 实现 CNN模型结构、训练参数和进度显示4.1 模型结构用 sequential 定义 LeNet 风格的 CNN接下来在js/model.js里定义 CNN。输入是 28×28×1 的张量依次放卷积层、池化层、再卷积、再池化、全连接层function createModel() { const model tf.sequential(); model.add(tf.layers.conv2d({ inputShape: [28, 28, 1], filters: 8, kernelSize: 3, padding: same, activation: relu })); model.add(tf.layers.maxPooling2d({ poolSize: 2 })); model.add(tf.layers.conv2d({ filters: 16, kernelSize: 3, padding: same, activation: relu })); model.add(tf.layers.maxPooling2d({ poolSize: 2 })); model.add(tf.layers.flatten()); model.add(tf.layers.dense({ units: 128, activation: relu })); model.add(tf.layers.dropout({ rate: 0.2 })); model.add(tf.layers.dense({ units: 10, activation: softmax })); return model; }参数说明filters指卷积核数量第一层用 8 个、第二层用 16 个已经够用。如果你用的是 5000 张以上的图片数据集可以改成 16 和 32准确率会再涨一点但训练时间明显变长。kernelSize: 3表示 3×3 卷积核是 28×28 小图最常用的配置5×5 感受野更大但参数更多。padding: same保证卷积后特征图尺寸不变边缘信息不会过早丢失。dropout({ rate: 0.2 })在全连接层后随机丢掉 20% 的神经元防止小数据集过拟合。如果训练集只有 1000 张rate 可以提高到 0.5。最后一层用 softmax 输出十个数字的概率和我们在 3.3 里生成的 one-hot 标签正好对应。4.2 编译和训练batchSize、epochs、validationSplit 分别怎么取值模型定义好之后先编译再调用fitmodel.compile({ optimizer: adam, loss: categoricalCrossentropy, metrics: [accuracy] }); const history await model.fit(trainData.xs, trainData.ys, { batchSize: 64, epochs: 10, validationSplit: 0.2, shuffle: true, callbacks: { onEpochEnd: (epoch, logs) { console.log(Epoch ${epoch 1} -- loss: ${logs.loss.toFixed(4)}, acc: ${logs.acc.toFixed(4)}); updateUI(epoch, logs); } } });optimizer: adam是网页端最省心的选择默认学习率 0.001大多数情况下不用调。如果 loss 不降或发散可以先检查数据是否归一化而不是急着改学习率。categoricalCrossentropy是多分类的标准损失配合 one-hot 标签使用。batchSize直接影响内存占用。64 是速度和稳定性的折中32 更稳但每个 epoch 更慢128 收敛可能出现抖动。网页演示时训练集 3000 张、batchSize 64、epochs 10 是可以接受的起点。validationSplit: 0.2表示把全部样本的 20% 留出来做验证每个 epoch 结束时会额外算一次验证准确率方便观察是否过拟合。shuffle: true让每个 epoch 的数据顺序重新打乱避免同一类图片扎堆导致梯度更新方向偏。训练过程中logs.val_acc是验证准确率字段注意别和 Python Keras 里的val_accuracy搞混。TensorFlow.js 回调日志里用的字段名是val_acc和浏览器端 API 一致。4.3 让训练过程可感知显示损失曲线和验证准确率callback 里除了打印日志还可以把logs.loss和logs.val_acc画成折线。网页项目里最朴素的做法是拿一个 canvas 或 SVG 画图不需要引入图表库。把每个 epoch 的数值推入数组每次drawChart重新渲染。这样在演示时观众能亲眼看到 loss 从 1.2 掉到 0.05比干等进度条有说服力得多。如果回调里拿不到val_acc说明validationSplit没被正确传入或者字段名拼错。另一个常见问题是训练时页面加了个按钮点击后model.fit异步执行但按钮状态没有锁定用户再次点击会同时启动两个训练任务导致浏览器卡死。我习惯在训练开始时禁用按钮训练结束再恢复这个小细节能避免很多演示现场翻车。5. 网页训练手写数字识别常见踩坑记录跨域、内存、NaN 和速度5.1 打开 index.html 后控制台全是 404本地文件跨域拦路现象双击 index.html页面能打开但图片和 TensorFlow.js 权重一个都加载不进来Network 面板里全是Failed to load resource。原因浏览器不允许网页脚本通过fetch()或Image读取本地文件系统。file://协议下目录访问和跨域限制更严格这和 CNN 无关是 Web 安全机制。解决在项目根目录起一个本地 web 服务器。最简单的方式是打开终端执行python -m http.server 8000然后浏览器访问http://localhost:8000。如果不方便用 Python也可以用 VS Code 的 Live Server 插件或npx serve .。改完以后刷新图片加载路径就正常了。5.2 训练过程中页面越来越卡最后标签页直接崩溃现象点“开始训练”前几秒很流畅到第三个 epoch 浏览器就开始掉帧CPU 风扇狂转最后页面弹出崩溃提示。原因TensorFlow.js 默认优先走 WebGL 后端如果模型、数据、中间激活值反复创建张量没有释放内存会持续累积。更常见的是用tf.stack一次性把几千张图片堆叠虽然 28×28 图片本身不大但训练时梯度、WebGL 纹理来回复制显存占用会快速膨胀。解决限制训练集规模首轮先用 1000 张跑通把batchSize控制在 32~64在加载数据循环里对临时创建的中间张量调用dispose()。如果项目已经上了自定义数据集另一个降内存峰值的办法是使用 TensorFlow.js 的tf.dataAPI 流式生成批次而不是一次性加载全部图片。5.3 训练没几轮 loss 突然变成 NaN现象loss 从 0.5 一路正常下降到第三轮开始出现NaN准确率直接掉到 10% 附近。原因最常见的是图片没有归一化。有些示例从 canvas 取像素后直接参与训练像素值还在 0~255Adam 优化器遇到大数值或除零就会产生 NaN。另一个原因是 one-hot 标签没有正确 squeeze导致 label 维度多了 1损失函数计算时越界。这里需要补充NaN 常常意味着数据预处理有 bug不是模型架构问题。解决在fromPixels之后加.div(tf.scalar(255))检查 labelTensor 的 shape 是不是[10]。如果数据已经归一化再尝试把 Adam 的学习率从默认 0.001 调到 0.0005。还有一个小细节tf.oneHot的 depth 参数要写 10而不是类别数加 1搞错了会在训练不崩的情况下识别全错。5.4 验证准确率卡在 90% 附近怎么都上不去现象训练集准确率能到 99%但验证准确率一直在 90% 上下浮动再增加 epochs 也没有改善。原因数据来源太单一或者模型容量不足。还有一个隐蔽原因训练集和测试集没有彻底 shuffle如果同一张图片因为随机种子被同时切到训练和验证验证准确率会虚高反过来如果验证集里恰好集中了难分的手写样本90% 也可能只是偶然而已。解决先确认shuffle: true然后把评估改成独立的 test 目录来计算准确率而不是依赖validationSplit。模型侧把第一层卷积 filters 从 8 提到 16再在 dropout 后加一个dense({units: 64})通常能跨过 95%。另外数字 1 和 7、4 和 9 在 28×28 分辨率下确实容易混淆验证集难样本偏多时90% 并不一定是异常。5.5 浏览器训练速度慢到让人想放弃CPU 与 WebGL 取舍现象同样一个 CNN 在 Python 里一个 epoch 可能是 1 秒在网页里要等 30 秒。原因TensorFlow.js 可能没有启用 WebGL 后端或者浏览器硬件加速被关闭。卷积运算在 GPU 后端能快一个数量级但在老集成显卡上WebGL 反而可能比 CPU 更慢。解决代码开头调用tf.setBackend(webgl)再用tf.getBackend()确认后端生效。如果机器配置很低把 epochs 从 20 降到 5训练集从 5000 张降到 1500 张演示效果也还够。另一个变通思路是网页里同时内置一个预训练好的模型训练按钮变成“微调”把训练时间压缩到几秒钟既保留交互感又避免现场等待。6. 网页识别与调优最后一步保存模型、画布预测和一批小技巧6.1 把训练好的模型保存到浏览器本地model.save(indexeddb://handwriting-cnn)可以把权重和网络结构存进浏览器 IndexedDB。下次打开页面后直接tf.loadLayersModel(indexeddb://handwriting-cnn)就能恢复模型不用重新训练。如果要在另一台电脑上用可以改成model.save(downloads://handwriting-cnn)浏览器会下载一个 JSON 和一堆权重文件放到项目里再加载。6.2 用 Canvas 画板做实时识别const imgTensor tf.browser.fromPixels(canvas, 1) .resizeBilinear([28, 28]) .toFloat() .div(tf.scalar(255)) .expandDims(0); const pred model.predict(imgTensor).argMax(1).dataSync()[0];这里 canvas 建议设成 280×280鼠标或手指画的线条要粗一些这样缩放到 28×28 时笔画不至于断掉。如果识别结果总是偏可以检查图片是否有白底黑字的翻转问题fromPixels拿到的是 RGBA如果要反色需要先对像素做tf.tensor2d运算再归一化。补充一个调优技巧训练前把图片里数字的质心移到画布中心能让准确率提升 1%~2%。具体做法是先算出非零像素的质心再做平移。网页端可以直接在fromPixels之后用tf.image的裁剪和填充来实现代价是训练时间略增但对真实手写输入的稳定性提升很明显。最后说句血泪经验我第一次拿类似结构做演示时没检查图片路径当着客户的面在控制台刷出一排 404。后来我会在页面加载时先跑一次 manifest 校验确保前几张图能被Image正常解码再点亮“开始训练”按钮。哪怕数据再小也要先保证一张图能完整地从文件路径变成张量再谈扩大训练集。这个顺序掌握住后面的路就好走多了。希望这些踩坑经历能帮到你。本文还有配套的精品资源点击获取
返回列表