ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

浏览器端CNN手写数字识别:从模型训练到部署全流程

浏览器端CNN手写数字识别:从模型训练到部署全流程 简介基于PyTorch框架的手写数字识别网页演示项目借助卷积神经网络识别多类别手写数字并将训练好的模型封装为本地网页交互服务覆盖数据准备、模型训练到网页部署的完整流程无需额外收集数据适合深度学习初学者快速上手。压缩包内共131个文件主体为124张手写数字图片并包含3个文本文件环境依赖与标签路径、3个Python脚本分别用于数据集文本生成、模型训练、启动网页服务和1个网页文件整体仅3.88MB目录清晰。目前已有94人学习下载。使用时可依次运行三个脚本01读取各文件夹图片生成训练/验证标签02训练卷积神经网络并保存模型同时输出每个迭代周期的验证集损失与准确率日志方便监控收敛和调整参数03启动本地服务在浏览器打开 http://127.0.0.1:4399 即可在线体验手写数字识别。整个流程逻辑完整适合课程设计、毕业设计或技术展示是一套能快速跑通的迷你实践项目。1. 解压这个zip之前先想清楚它在解决什么问题解压这个zip之前先想清楚一件事它把完整的手写数字识别链路拆成了两半——Python负责用CNN训练模型HTML页面负责在浏览器里加载模型并实时识别你手写的数字。也就是说最终交付物不是一条命令行而是一个能打开就用的网页画个数字、点一下识别浏览器在本地跑完卷积计算把预测结果和置信度返回给你。整个推理过程不需要服务器也不依赖GPU。对做毕设、Web前端可视化和想快速看到CNN效果的人来说这是成本最低的一条技术路径。我按平时动手的顺序看数据、训模型、导出到网页、排错、进阶把每个环节拆开讲透新手能照做熟手能避坑。2. CNN识别手写数字为什么这个任务天生适合卷积网络2.1 从784个像素到10个类别卷积在做什么手写数字识别本质上是把一张28×28的灰度图映射到09这十个类别。如果用全连接网络硬做输入层784个节点第一层放128个神经元光这一层就有约10万个可训练参数。更麻烦的是全连接网络对位置极度敏感同样的一个数字整体往左平移两个像素激活值完全变样网络可能就从“认成3”变成“认成8”。而卷积神经网络用3×3或5×5的卷积核在整张图上滑动同一个卷积核在不同位置提取同一类局部特征参数规模大幅缩减同时也天然具备平移不变性。对这类zip里的项目来说训练数据通常是28×28灰度图模型结构也基本沿着LeNet-5的路线走两层卷积加池化提取特征再接两层全连接做分类。第一层卷积学到的是横线、竖线、斜边这类底层边缘第二层卷积把这些边缘组合成弧线、圆圈、交叉点。手写数字的笔画粗细因人而异最大池化层取窗口内的最强响应相当于对局部做了模糊对齐这正是3和8、4和9这类笔画位置相近的数字能被区分开的关键。2.2 这个项目里CNN结构怎么定层数、卷积核与Dropout位置常见的手写数字CNN结构非常固定从MNIST时代的实践沉淀下来基本就是下面这个形态。我一般会直接把这个结构作为基线再根据数据集规模微调from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Conv2D, MaxPooling2D, Flatten, Dense, Dropout model Sequential([ Conv2D(32, (3, 3), activationrelu, input_shape(28, 28, 1)), MaxPooling2D((2, 2)), Conv2D(64, (3, 3), activationrelu), MaxPooling2D((2, 2)), Flatten(), Dense(128, activationrelu), Dropout(0.5), Dense(10, activationsoftmax) ])第一层用32个3×3卷积核第二层翻倍到64个这个“通道数逐层翻倍”是卷积网络最常用的经验规律因为底层边缘特征种类有限而组合特征数量会越来越多。两个MaxPooling层都取2×2窗口把28×28逐步压到7×7既减小了计算量又让后续全连接层拿到的特征图对笔画偏移更鲁棒。Dropout放在倒数第二层Dense(128)之后、输出层之前因为卷积层靠权重共享天然有抗过拟合能力而Flatten之后的全连接层占了模型85%以上的参数是过拟合的重灾区。激活函数全部用ReLU输出层用Softmax得到十个类别的概率分布这个配置基本不需要改动。真正需要调的是epochs和Dropout比例如果zip里的数据只有两三千张Dropout可以从0.5提到0.6epochs适当加多如果数据量到了MNIST级别0.5的Dropout配合40个epoch就能稳定跑到99%以上的验证集准确率。2.3 用Keras还是PyTorch由部署路径倒推拿到zip先翻目录看到.h5训练好的模型文件加一个.html页面基本可以断定这是Keras加TensorFlow.js的链路。这条链路最顺模型用tf.keras训练保存为.h5一条tensorflowjs_converter命令就能导出浏览器可加载的model.json和权重bin文件前端用tensorflow.js调用model.predict直接完成推理。PyTorch也可以做Web部署但要走ONNX Runtime Web中间多了一道模型转换和算子的适配工作除非zip里放的是.pt权重否则没必要绕这个弯。选型这件事经常被新手忽略但它直接决定后面的排错难度。我见过不少人在PyTorch里训练精度很高结果导出到浏览器时发现某个算子不被ONNX Runtime Web支持回头改网络结构白白浪费一下午。反过来如果一开始就知道最终产物是网页直接用Keras是最省心的选择。HTML页面本身不挑框架但模型从训练到浏览器这段路的顺畅程度Keras加TF.js目前仍是最优解。3. 图片数据集把zip里的散图整理成CNN吃得到的矩阵3.1 先扫描再训练看清数据集形态比调参重要“含图片数据集”听起来简单但zip解压后的数据形态常见有两种一种是按train/0、train/1这样的标签目录存放PNG或JPG图片另一种是单个.npz文件里面直接就是x_train、y_train、x_test、y_test四个数组。先跑一段扫描脚本把数据集的结构、数量、图片尺寸和通道数摸清楚比直接开训练更能避免后面浪费时间。import os from PIL import Image root dataset # zip解压后的数据集根目录 for label in sorted(os.listdir(root)): label_dir os.path.join(root, label) if not os.path.isdir(label_dir): continue files [f for f in os.listdir(label_dir) if f.lower().endswith((.png, .jpg, .jpeg))] first_img Image.open(os.path.join(label_dir, files[0])) print(flabel{label} count{len(files)} fsize{first_img.size} mode{first_img.mode})这段脚本逐个标签目录统计文件数并读取第一张图的尺寸和颜色模式。看输出就能判断两件事一是各类别样本是否均衡如果某个数字只有几十张训练时这个类的准确率大概率上不去二是图片尺寸和通道是否统一打印出size28×28和modeL说明数据已经处理过直接用如果出现modeRGB甚至RGBA后面必须统一转灰度否则模型输入维度对不上训练直接报错。3.2 目录结构的散图统一转成numpy训练集如果数据集是目录结构需要先把所有图片读进来统一尺寸、转灰度、归一化组装成numpy数组。这一步的代码几乎是复制粘贴就能用import os import numpy as np from PIL import Image from sklearn.model_selection import train_test_split def load_images_from_dir(root, size(28, 28)): xs, ys [], [] for label in sorted(os.listdir(root)): label_dir os.path.join(root, label) if not os.path.isdir(label_dir): continue for name in os.listdir(label_dir): if not name.lower().endswith((.png, .jpg, .jpeg)): continue img Image.open(os.path.join(label_dir, name)) img img.convert(L).resize(size, Image.Resampling.BICUBIC) xs.append(np.asarray(img, dtypenp.uint8)) ys.append(int(label)) return np.stack(xs), np.array(ys) x_all, y_all load_images_from_dir(dataset) print(x_all.shape, y_all.shape) # 期望得到 (N, 28, 28) 和 (N,) x_train, x_val, y_train, y_val train_test_split( x_all, y_all, test_size0.2, random_state42, stratifyy_all ) np.savez_compressed(dataset.npz, x_trainx_train, y_trainy_train, x_testx_val, y_testy_val)几个关键点。convert(L)把RGB和RGBA统一成8位灰度这样不管原图是彩色还是带透明通道都不会报错。resize用BICUBIC插值而不是默认的最近邻缩放后的数字边缘更平滑特征更接近MNIST的训练分布。train_test_split里加stratifyy_all保证每个数字在训练集和验证集里的占比一致避免某个数字在验证集里恰好只有一两个样本导致准确率剧烈抖动。random_state固定为42这样每次切分的数据分布一致模型对比才有意义。最后用np.savez_compressed保存成压缩npz后面训练脚本直接load不用再回扫图片。3.3 数据增强做不做、做到什么程度数据增强在这个项目里的价值取决于样本量。如果zip里自带的数据有两三万张不做增强也能训得不错如果只有几千张轻度增强能把验证集准确率拉高1到3个百分点。手写数字增强要克制旋转角度超过15度就会把3和8、7和9这种结构相近的数字搅浑反而引入噪声。from tensorflow.keras.preprocessing.image import ImageDataGenerator datagen ImageDataGenerator( rotation_range10, # 旋转角度范围不要超过15 width_shift_range0.1, # 水平平移10% height_shift_range0.1, # 垂直平移10% zoom_range0.1, # 缩放10% fill_modenearest # 空位用最近像素填充 ) x_train_4d x_train[..., np.newaxis] # (N, 28, 28) - (N, 28, 28, 1) datagen.fit(x_train_4d) model.fit(datagen.flow(x_train_4d, y_train_cat, batch_size128), epochs40, validation_data(x_val[..., np.newaxis], y_val_cat))fill_modenearest值得单独说。旋转和平移后图片边缘会出现空像素如果用zero填充黑色空位会被模型误认为是笔画干扰训练用nearest填充空位颜色跟边缘像素一致干扰小很多。还有一个容易踩的坑ImageDataGenerator的flow要求输入是四维数组x_train必须补上通道维变成(N, 28, 28, 1)再传进去少这一维会直接报维度不匹配的错误。4. 把CNN塞进Web训练、导出与canvas手写板推理4.1 训练脚本精度过98%再谈部署训练这一步的目标只有一个得到一个验证集准确率足够高的模型。手写数字识别在MNIST级别数据上做到99%是小菜一碟但如果zip里数据只有几千张98%以上也够用。先把训练脚本跑通import numpy as np from tensorflow.keras.utils import to_categorical from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Conv2D, MaxPooling2D, Flatten, Dense, Dropout from tensorflow.keras.optimizers import Adam data np.load(dataset.npz) x_train, y_train data[x_train], data[y_train] x_val, y_val data[x_test], data[y_test] x_train x_train.astype(float32) / 255.0 x_val x_val.astype(float32) / 255.0 if x_train.ndim 3: x_train x_train[..., np.newaxis] x_val x_val[..., np.newaxis] y_train to_categorical(y_train, 10) y_val_cat to_categorical(y_val, 10) model Sequential([ Conv2D(32, (3, 3), activationrelu, input_shape(28, 28, 1)), MaxPooling2D((2, 2)), Conv2D(64, (3, 3), activationrelu), MaxPooling2D((2, 2)), Flatten(), Dense(128, activationrelu), Dropout(0.5), Dense(10, activationsoftmax) ]) model.compile(optimizerAdam(learning_rate1e-3), losscategorical_crossentropy, metrics[accuracy]) model.fit(x_train, y_train, batch_size128, epochs40, validation_data(x_val, y_val_cat), verbose2) model.save(model.h5)三个参数值得注意。Adam优化器用默认学习率1e-3就够手写数字任务不需要精细调学习率如果发现训练loss波动大降到1e-4重跑一次。batch_size128在内存和收敛速度之间比较平衡显存不够或CPU训练就改64。epochs这里写了40但如果数据量小模型会在20个epoch左右就收敛后面只是震荡建议训练时盯着验证集准确率连续5个epoch不提升就可以停了不用死等40轮跑完。4.2 用tensorflowjs_converter导出浏览器可加载的模型训练好的model.h5不能直接被网页加载要先用TensorFlow.js的转换工具导出一份浏览器可识别的格式。这一步本身不复杂但环境准备和参数选择有讲究pip install tensorflowjs tensorflowjs_converter \ --input_formatkeras \ --output_formattfjs_layers_model \ --quantization_bytes1 \ model.h5 web_model转换完成后web_model目录下会出现一个model.json和若干个group1-shard1of1.bin文件。--quantization_bytes1表示把32位浮点权重量化到8位模型体积缩到原来的四分之一。对于手写数字识别这种简单任务量化后的精度损失通常小于0.1%肉眼完全无感但网页加载速度会明显变快。如果拿到的zip里原本就带网页和模型这一步往往是跳过的但自己从头训时别漏了。转换之后有个容易忽略的坑model.json里记录的是相对路径bin文件必须和model.json放在同一个目录下不能只拷贝model.json进网页目录否则加载时404。4.3 HTML页面拿canvas当输入拿tfjs做推理这是整个web项目里最核心的页面。用canvas做手写板把280×280的画布内容缩放到28×28归一化后喂给模型做预测。完整页面代码如下!doctype html html langzh-cn head meta charsetutf-8 meta nameviewport contentwidthdevice-width, initial-scale1 titleCNN 手写数字识别浏览器版/title script srchttps://cdn.jsdelivr.net/npm/tensorflow/tfjs4/script style body { font-family: sans-serif; display: flex; flex-direction: column; align-items: center; gap: 12px; padding: 24px; } canvas#board { border: 2px solid #333; background: #fff; touch-action: none; } canvas#thumb { display: none; } .row { display: flex; gap: 12px; } button { font-size: 16px; padding: 6px 18px; } /style /head body h3画一个数字点击识别/h3 canvas idboard width280 height280/canvas canvas idthumb width28 height28/canvas div classrow button idrecognize识别/button button idclear清空/button /div div idresult等候输入.../div script let model null; const board document.getElementById(board); const ctx board.getContext(2d); const thumb document.getElementById(thumb); const tctx thumb.getContext(2d); ctx.fillStyle #fff; ctx.fillRect(0, 0, 280, 280); ctx.strokeStyle #000; ctx.lineWidth 14; ctx.lineCap round; let drawing false; board.addEventListener(pointerdown, e { drawing true; ctx.beginPath(); ctx.moveTo(e.offsetX, e.offsetY); board.setPointerCapture(e.pointerId); }); board.addEventListener(pointermove, e { if (!drawing) return; ctx.lineTo(e.offsetX, e.offsetY); ctx.stroke(); }); board.addEventListener(pointerup, () drawing false); tf.loadLayersModel(web_model/model.json).then(m { model m; document.getElementById(result).textContent 模型已加载可以开始写数字; }); document.getElementById(recognize).onclick () { if (!model) return; tctx.clearRect(0, 0, 28, 28); tctx.drawImage(board, 0, 0, 28, 28); const img tctx.getImageData(0, 0, 28, 28); const input new Float32Array(28 * 28); for (let i 0; i 28 * 28; i) { input[i] img.data[i * 4] / 255.0; } const tensor tf.tensor4d(input, [1, 28, 28, 1]); const pred model.predict(tensor); const probs Array.from(pred.dataSync()); const digit probs.indexOf(Math.max(...probs)); const top3 probs.map((p, i) [i, p]) .sort((a, b) b[1] - a[1]) .slice(0, 3) .map(([i, p]) i ( p.toFixed(2) )) .join( ); document.getElementById(result).textContent 识别结果 digit 置信度 probs[digit].toFixed(3) top3: top3; tensor.dispose(); pred.dispose(); }; document.getElementById(clear).onclick () { ctx.fillStyle #fff; ctx.fillRect(0, 0, 280, 280); ctx.strokeStyle #000; ctx.lineWidth 14; ctx.lineCap round; document.getElementById(result).textContent 等候输入...; }; /script /body /html预处理这段是精度的核心。tctx.drawImage(board, 0, 0, 28, 28)把280×280的画布整体缩放到28×28缩放过程等于做了一次像素混合相当于简化版的抗锯齿。取像素时只取R通道因为画布是白底黑笔迹RGB三个通道值相等取一个就够。除以255把像素从0255的整数映射到01的浮点数这和训练时的归一化完全一致。容易遗漏的是input数组的构造顺序getImageData返回的data是一维数组每个像素占四字节RGBA所以第i个像素的R通道在data[i*4]不能写成data[i]。预测完tensor.dispose()和pred.dispose()必须调用。TF.js在Web环境下内存管理比较敏感每次predict都会创建新的tensor不释放的话连续识别几十次后页面就会明显卡顿最后变成黑色页面崩溃。4.4 模型自检把测试集跑一遍前端链路很多人训练阶段准确率99%部署到网页就变成70%原因几乎都出在图片预处理和训练数据没对齐。与其来回猜不如直接在网页里写一段自检代码把测试集前50张图渲到canvas上走和手写识别完全相同的预处理逻辑和真实标签做对比async function selfCheck(xTest, yTest, n 50) { let hit 0; for (let i 0; i n; i) { const canvas document.createElement(canvas); canvas.width canvas.height 28; const c canvas.getContext(2d); const img new ImageData(new Uint8ClampedArray(xTest[i]), 28, 28); c.putImageData(img, 0, 0); const pixel c.getImageData(0, 0, 28, 28); const input new Float32Array(28 * 28); for (let j 0; j 784; j) input[j] pixel.data[j * 4] / 255.0; const pred model.predict(tf.tensor4d(input, [1, 28, 28, 1])).dataSync(); const digit pred.indexOf(Math.max(...pred)); if (digit yTest[i]) hit; } console.log(前50张自检准确率:, hit / n); }这段代码的作用是把测试集图片从numpy数组还原成canvas像素走一遍和手写板完全相同的缩放、归一化、预测流程。如果这里自检准确率和Python端训练时验证集准确率偏差超过1%就说明前端预处理链路里有环节不一致需要优先排查画布缩放和归一化逻辑而不是怀疑模型本身。5. 避坑网页版手写识别最容易翻车的五个地方5.1 画布预处理和训练数据不一致预测结果像随机数现象网页上无论写什么数字识别结果总是在某两个类之间反复横跳或者概率集中在一个类上写3和写8的结果几乎一样像是玄学。原因前端画布是白底黑字训练数据如果恰好是反色格式比如黑底白字或者训练时图片经过了居中裁剪而前端没有喂给模型的像素分布就和训练时不一致模型的输出自然不可信。解决把前端预处理固定成一套标准流程缩放28×28、取灰度、除以255归一化、reshape成(1, 28, 28, 1)并用console打印几组关键数字背景像素值应该是255笔迹像素值接近0。如果发现相反说明画布颜色和训练数据反了在输入模型前加一句input[i] 1 - input[i]反转即可。5.2 model.json加载报404或Failed to fetch现象双击HTML文件打开页面控制台报Failed to fetch model.json或跨域错误模型一直加载不出来。原因浏览器出于安全策略在file://协议下限制fetch请求本地直接双击打开的HTML页面不能加载同目录的模型文件。解决不要双击打开在项目目录起一个本地静态服务器用python -m http.server 8080启动浏览器访问http://localhost:8080。同时确认web_model目录下的model.json和group1-shard1of1.bin确实在同一级目录且引用的路径大小写一致。用TF.js的CDN加载参考loadLayersModel时用的路径是web_model/model.json以HTML所在的页面为相对路径基准不要写成绝对路径/User/xxx这种。5.3 前端识别精度远低于Python训练精度差距在1%以上现象Python训练脚本里验证集准确率99%浏览器里同样一张图片自检只有80%左右。原因归一化翻车。最常见的是忘记除以255直接把0255的像素喂给模型或者reshape少了通道维把(N, 28, 28)当成(N, 28, 28, 1)用的函数实际上一维数组顺序不对。第二个常见原因是画布缩放用了drawImage默认的平滑参数而训练时用的是BICUBIC插值边缘像素分布有差异。解决在自检脚本里把输入矩阵打印出来肉眼对比训练集的灰度分布。确认背景是255、笔画是0、数值范围在01之间。如果是插值方式的问题drawImage可以加第三个参数设置imageSmoothingQuality比如ctx.imageSmoothingQuality high或者干脆忽略这0.5%的误差只要整体趋势正常就先往下走。5.4 连续预测几次后页面越来越卡最后崩溃现象第一次点识别正常连续操作十几次后页面开始掉帧最终直接卡死。原因是TF.js的内存泄漏。每次model.predict都会创建tensor如果predict之后不手动释放这些tensor会一直留在内存里垃圾回收机制在Web环境下不会及时回收。解决在predict结束后立即调用tensor.dispose()和pred.dispose()或者在预测代码外包一层tf.tidy(() { ... })TF.js会在函数执行完自动清理内部的中间tensor。另一个习惯是model只加载一次每次识别都复用同一个model实例不要重复loadLayersModel。5.5 手机上canvas画不出笔迹页面跟着手指滚动现象手机浏览器里打开页面手指在画布上滑动时画面不出现笔迹或者笔迹断断续续页面还在上下滚动。原因触摸事件和鼠标事件是两套体系直接监听mousedown、mousemove在触屏设备上不会触发。另外canvas没有设置touch-action: none时浏览器会把在画布上的滑动当作页面滚动事件被默认行为吞掉。解决统一用Pointer Eventspointerdown、pointermove、pointerup同时覆盖鼠标和触摸。同时给canvas加CSS touch-action: none并在pointerdown事件里调用board.setPointerCapture(e.pointerId)这样手指移出画布范围后事件仍然持续触发笔迹不会断。这条在手机演示时几乎是必踩的坑提前写进代码能省很多麻烦。6. 不止是演示浏览器端微调与双模型兜底真要把这个网页版手写数字识别当工具用而不是只当一个演示玩具有两个升级方向值得做浏览器端微调和双模型兜底。6.1 在浏览器里做一次个性化适配TF.js的模型同样支持trainable控制。把前面的卷积层全部冻结只训练最后一个Dense(128)和输出层用户自己画8个样本并标注真实数字就能在浏览器里做一次快速微调model.layers.forEach((layer, idx) { if (idx model.layers.length - 2) layer.trainable false; }); model.compile({ optimizer: rmsprop, loss: categoricalCrossentropy, metrics: [accuracy] }); await model.fit(xs, ys, { epochs: 30, batchSize: 4 });冻结卷积层是为了防止小样本把已经提取好的底层特征破坏掉只训练分类头让模型往个人笔迹方向偏移一点点。rmsprop学习率默认0.001微调时不用改。我实测过画8个样本做一轮微调个人笔迹的整体误识别率能从2%压到0.5%以下。微调后的权重可以用model.save(indexeddb://custom)存进浏览器IndexedDB下次打开页面自动加载就实现了“越用越准”。6.2 双模型兜底与犹豫状态有些数字对比如4和9、3和8个人笔迹下实在写得太像单模型输出会长期处于低置信度状态。思路是页面上同时加载主模型和副模型主模型是MNIST预训练副模型是本zip数据训练的。推理时如果主模型top1概率低于0.7就用两个模型的softmax输出加权融合主模型0.6、副模型0.4融合后仍然低于阈值直接提示“这个字我不确定”不强行给答案。验证方法很简单拿测试集固定抽20张图在页面离线跑一遍和前端完全一致的预处理链路对比准确率再用手写同一数字10次看预测分布方差。如果方差很大多半是预处理没对齐训练分布而不是模型本身的问题。我现在跑任何带canvas输入的前端模型第一件事永远是打印归一化之后的28×28矩阵确认底色、笔迹方向、缩放方式都和训练数据一致再做精度对比——这是连续三次“画板看着正常、识别结果全错”之后养成的习惯。希望帮到你。本文还有配套的精品资源点击获取
返回列表