尧图精选

TensorFlow.js教程:浏览器端机器学习与图像分类实战

🕒 发布时间:2026/10/1 6:01:49 📁 来源:尧图网络
1. 项目定位TensorFlow.js到底帮我们解决了什么问题先说一个很实际的场景。上个月朋友接了个需求要在 H5 页面里做实时的人体关键点检测用户随便传一张照片浏览器直接标出肩膀、手肘、膝盖的位置而且照片不能上传到服务器。他一开始想走传统后端方案Python 起一个推理服务前端把图片传上去算完再拿结果。但聊完需求就发现这条路走不通数据必须留在用户本地光这一条就把服务端推理解法堵死了。后来换成 TensorFlow.js事情一下子简单了模型随前端资源一起下发推理在浏览器里完成图片从头到尾不出用户设备。这就是 TensorFlow.js 最核心的价值把机器学习运行环境从服务端搬到浏览器让推理和训练都能跑在用户的设备上。如果你是做前端的想在页面里加点智能化能力又不想碰 Python 和服务器或者你虽然有现成的 Python 训练好的模型但不想为了一个功能专门维护推理服务那这个库就是你的主场。这个库不是“玩具级”的存在。Google 官方维护底层能调用 WebGL 跑 GPU 加速也能退到 WebAssembly 用 CPU 硬算。它可以加载别人训练好的模型做推理也可以把 Python 里 TensorFlow 训练的模型转换后拿到浏览器用甚至能从头在浏览器里训练模型。往下读之前先记住一件事TensorFlow.js 不是一个全新的机器学习框架而是把 TensorFlow 的能力平移到了 JavaScript 生态里让你用写前端的方式去写机器学习代码。1.1 它和传统机器学习方案的本质区别传统机器学习流程大家应该很熟数据准备、模型训练、模型部署、接口调用。部署这一步通常发生在服务器上客户端发请求服务器跑推理。这样做没什么问题但有几个场景天然吃亏。第一是隐私敏感场景。医疗影像、身份信息、企业内部数据这类数据一旦上传服务器就面临合规压力很多项目干脆因为这个砍掉了智能功能。第二是实时交互场景。视频流、摄像头画面、频繁的鼠标操作每帧都传到服务器推理网络延迟先不说服务器费用也扛不住。第三是弱网和离线场景。移动端用户在地铁、电梯、地下车库网络说断就断推理服务必须连续可用这在传统方案里是全链路的高可用工程问题。TensorFlow.js 把推理搬到浏览器后这三个问题同时被绕开了。数据不出端带宽消耗变成了一次性的模型下载实时性取决于本机 GPU 和 CPU 的能力。前端领域经常讨论的“边缘计算”落到用户端其实就是这么回事把算力放到离用户最近的地方。1.2 什么样的项目最适合用它不是所有机器学习功能都适合在浏览器里跑。我自己的经验是适合和不适合的界限还挺清晰的。适合的图像分类、人脸检测、姿态估计、物体检测、语音命令识别、简单的文本情感分析这些模型尺寸在几 MB 到几十 MB 之间推理延迟百毫秒级用户体验可接受。交互式的项目尤其合适比如摄像头前实时人脸关键点追踪、手势控制、拍照后立刻给图片打标签这类功能如果走服务端延迟会彻底毁掉体验。不适合的超大规模模型的推理比如几十 GB 的推荐模型、大语言模型浏览器根本装不下高频的批量数据处理比如要处理一万张图片的OCR任务浏览器端跑会慢得让人崩溃这类活儿应该丢给服务端批处理。判断标准就两条模型打包后最大不能超过用户愿意等的时间上限推理延迟必须能撑得起交互。拿捏好这两条基本不会选错方案。2. 核心机制浏览器里到底怎么跑机器学习很多人第一次接触 TensorFlow.js 都会有个疑问JavaScript 不是用来写网页的吗它怎么承担得起矩阵运算和梯度下降这种重活儿要回答这个问题得看看这个库在底层到底做了什么。2.1 张量和自动微分用前端的方式理解机器学习TensorFlow.js 里的基本操作单位是张量你可以直接把它理解成一个“多维数组”。标量是 0 维张量向量是 1 维张量矩阵是 2 维张量图片是 3 维张量一批图片就是 4 维张量。前端开发者天天跟数据打交道这个抽象其实很好接受。库把 TensorFlow 的核心操作几乎全部搬了过来add、matMul、conv2d、maxPool 这些底层运算都在自动求梯度的能力也在。你在浏览器里可以用tf.train.sgd这类优化器配合自定义损失函数从头训练一个线性回归或者小型神经网络。虽然实践中直接在浏览器里训练模型的场景不算多但理解这套机制能帮你更好地理解推理过程中每一步在干嘛。浏览器端做推理的直观感觉是你不需要理解反向传播的数学细节只要会用model.predict()就够了。但如果你想调优模型或者排查问题知道张量的形状对齐、数据类型、归一化方式这些基础知识能少踩很多坑。后面第四章我会专门讲前端预处理里最容易出错的地方。2.2 后端加速WebGL 和 WASM 在背后做了什么TensorFlow.js 的架构设计里有一个“后端”概念术语叫 Backend。最常用的有三个WebGL 后端、WebAssembly 后端、纯 JavaScript 后端。WebGL 后端是默认的主力。它在页面内部创建一个不可见的 WebGL 渲染上下文把张量数据上传成 GPU 纹理然后通过着色器程序执行矩阵运算。通俗点说GPU 在浏览器里的“本职工作”是绘制画面但 TensorFlow.js 借了它的算力来干数学运算的活儿。优点是快图像模型的推理速度能接近原生体验缺点是它在某些低端设备上兼容性不稳定而且 GPU 纹理内存需要格外小心管理用完了不释放就会越积越多一会儿我会专门讲这个问题。WebAssembly 后端相当于用 CPU 硬算。它把运算逻辑编译成 WASM 字节码在浏览器里接近原生速度执行兼容性更好在 GPU 不可用时可以自动降级。纯 JS 后端基本只剩兜底作用性能差但任何环境都能跑。这个多后端机制是我认为 TensorFlow.js 设计得最聪明的地方。你不用关心用户机器上有没有 GPU、浏览器版本老不老库会自己选择最优路径。但也因为这种动态性实际使用中会出现一些诡异现象同一段代码在 Chrome 上流畅到某个版本的 Safari 上就慢得离谱很可能就是后端选择差异导致的。2.3 三条路线预训练模型、模型导入和浏览器端训练TensorFlow.js 的使用方式有三条路线你可以根据手头资源来选择。第一条是用官方维护的预训练模型。tensorflow-models 仓库下有一批开箱即用的模型比如 MobileNet 图像分类、Coco-SSD 目标检测、PoseNet 姿态估计、FaceDetector 人脸检测。用这些模型的体验和调第三方 API 很像一个load()方法下载模型一个classify()或detect()方法出结果。如果你只是想快速验证一个想法这是最快的一条路。第二条是把 Python 里训练好的模型转换后拿过来用。训练还是用 TensorFlow Python 版做训练完导出模型再用官方提供的转换工具把它转成 TensorFlow.js 的格式然后放在前端加载。这条路线适合有正式训练流程的团队把训练和部署解耦我在第五章会给出具体的转换命令和踩坑过程。第三条是直接在浏览器里训练模型。用tf.layers搭神经网络用model.fit训练。这条路线上手门槛低不用装 Python 环境但说实话性能上限不高只适合教学演示和轻量场景真要做生产级训练还是得回到 Python 生态。3. 环境准备不装 Python 的工程化配置TensorFlow.js 的开发环境比传统机器学习轻太多了。不需要 Anaconda不需要 CUDA不需要 GPU 驱动一个 Node.js 环境加一个浏览器就够。下面是完整的初始化过程前端熟手可以直接跳过去。3.1 搭建 npm 项目并安装依赖打开终端创建项目目录并初始化mkdir tfjs-browser-demo cd tfjs-browser-demo npm init -y核心依赖就两个tensorflow/tfjs是核心库tensorflow-models/mobilenet是图像分类模型。另外我建议装一个打包工具用 Vite 起步最简单开发体验和部署都很顺。npm install tensorflow/tfjs tensorflow-models/mobilenet npm install --save-dev vite安装完成后在 package.json 里配置脚本{ scripts: { dev: vite, build: vite build, preview: vite preview } }这里有个选择需要解释一下。为什么用 Vite 而不是直接开一个静态服务器两个原因。第一TensorFlow.js 官方模型文件需要通过网络加载直接双击 HTML 文件用 file:// 协议打开浏览器会因为跨域限制拦截模型请求必须通过本地 HTTP 服务访问。第二Vite 是 ESM 友好的import语句用起来很自然省掉一堆构建配置。Vite 默认的 5173 端口足够用了开发服务器起起来是这样的npm run dev3.2 一个最小可运行的页面结构在项目根目录建index.html内容不需要复杂一个文件选择器加一个 canvas 就够了!DOCTYPE html html langzh-CN head meta charsetUTF-8 meta nameviewport contentwidthdevice-width, initial-scale1.0 titletfjs 图像分类 Demo/title style body { font-family: system-ui, sans-serif; max-width: 720px; margin: 0 auto; padding: 24px; } img { max-width: 100%; margin-top: 16px; } #results { margin-top: 16px; line-height: 1.8; } /style /head body h1浏览器里的图像分类/h1 input typefile acceptimage/* iduploader / br/ img idimage alt待分类图片 styledisplay:none; / div idresults选一张图片试试/div script typemodule src./src/main.js/script /body /html这个页面把用户选中的图片展示出来并交给模型分类。实际项目中你还会遇到图片压缩、canvas 绘制、多模型并行这类需求但先把这条最小链路跑通后面都好说。3.3 依赖体积与版本锁定的问题这里必须提醒一句TensorFlow.js 核心库的体积并不小。生产构建后仅核心库压缩后大概在 600KB 左右MobileNet 模型文件又是 5MB 左右。对移动端项目来说这是需要认真对待的资源开销首屏不能无脑加载建议按需引入模型用动态 import 让模型加载和首屏渲染并行。另外TensorFlow.js 的版本更新节奏比较快API 偶有调整生产项目一定要锁定精确版本号建议直接写死不要用^这种范围版本。我见过不止一个项目因为自动升级了依赖导致推理结果突然不对排查半天发现是模型输出格式变了。锁定版本是前端工程的常识但在机器学习库上体现得更明显。4. 实操 Demo让 MobileNet 在浏览器里做图像分类环境就绪开始写核心代码。这一章我会走完整个流程页面交互、模型加载、数据预处理、推理、结果展示每一步都讲清楚“为什么这么做”。4.1 加载模型与初始化新建src/main.js先写模型加载逻辑import * as tf from tensorflow/tfjs; import * as mobilenet from tensorflow-models/mobilenet; const model await mobilenet.load({ version: 2, alpha: 1.0 });await mobilenet.load()会在浏览器后台下载模型权重。参数里的version: 2表示用 MobileNetV2结构更轻、速度更快alpha: 1.0是模型宽度系数数值越小模型越小、精度越低。如果你的项目对体积敏感可以把 alpha 降到 0.5 或 0.25体积和精度都会明显变化。我实际测试过alpha 从 1.0 降到 0.5加载体积能减少大约 60%分类准确率大概掉两三个百分点大多数普通图片分类场景完全够用。模型加载完成后给文件选择器绑定事件。这里有一个很常见的坑用户在文件框里选了图片不能直接拿File对象做模型输入模型要的不是文件对象也不是 Blob而是一个 HTMLImageElement 或者 canvas。所以需要先把文件转成可显示的图片元素const uploader document.getElementById(uploader); const imageEl document.getElementById(image); const resultsEl document.getElementById(results); uploader.addEventListener(change, (event) { const file event.target.files[0]; if (!file) return; const url URL.createObjectURL(file); imageEl.src url; imageEl.style.display block; imageEl.onload () { classifyImage(); URL.revokeObjectURL(url); }; });这里URL.createObjectURL和revokeObjectURL的配对使用是前端基础但很多人会忘掉 revoke导致大图反复选择时浏览器内存缓慢上涨。注意这个手法后面排查内存问题时你会感激它。4.2 数据处理从像素到张量的关键一步这是全篇文章最关键的一步响应一下“机器学习中的数据处理”这个核心问题。图片从 DOM 进入模型之前必须被转成特定形状和取值范围的张量。不同的模型有不同的要求MobileNet V2 的要求是图片缩放到 224x224 分辨率像素值从 0 到 255 归一化到 -1 到 1 区间张量形状为[1, 224, 224, 3]。代码看起来是这样async function classifyImage() { const tensor tf.browser.fromPixels(imageEl) .resizeNearestNeighbor([224, 224]) .toFloat() .sub(127.5) .div(127.5) .expandDims(0); const predictions await model.classify(tensor); tensor.dispose(); renderResults(predictions); }一步步解释。tf.browser.fromPixels(imageEl)是浏览器环境专属的 API能把图片元素或 canvas 直接转成[height, width, 3]的张量第三维是 RGB 三个通道。.resizeNearestNeighbor([224, 224])做尺寸缩放模型要求什么尺寸就缩到什么尺寸不能多也不能少。.toFloat()把整数像素转成浮点数因为后续的减法和除法需要浮点参与。.sub(127.5).div(127.5)是归一化MobileNet 的输入范围是 -1 到 1像素值 255 减去 127.5 再除以 127.5 正好变成 10 变成 -1 附近。最后的.expandDims(0)在第一维增加一个 1让形状从[224, 224, 3]变成[1, 224, 224, 3]这个 1 表示“batch”即一次输入一张图片。数据处理的本质就两件事对齐形状、对齐数值范围。这两个条件任何一个不对模型输出的结果就会莫名其妙。前端项目里最容易出问题的就是数值范围有人写 RGB 转灰度花了半天最后发现模型要求的其实是 RGB 三通道根本不用转灰度。4.3 结果解析与展示模型输出的形状是[1, 1000]表示 1000 个类别的置信度分数。直接用model.classify的话库已经帮我们做了 softmax 和后处理返回的是置信度排名前 N 的结果数组。渲染部分的代码很直白function renderResults(predictions) { resultsEl.innerHTML h2识别结果/h2; predictions.slice(0, 3).forEach((p, i) { const div document.createElement(div); div.textContent ${i 1}. ${p.className}${(p.probability * 100).toFixed(1)}%; resultsEl.appendChild(div); }); }到这里一个能在浏览器里做图像分类的完整 Demo 就跑通了。从用户选图到看到结果整个流程全部在浏览器内完成没有任何一行代码把图片发到服务器。5. 模型迁移把 Python 里训练好的模型搬进浏览器官方预训练模型只覆盖了常见任务。真实项目的模型往往是用自己的数据训练出来的这时候需要走模型转换这条路线。核心思路是训练阶段仍然在 Python 里做部署阶段把模型转成 TensorFlow.js 格式。5.1 转换工具安装与命令转换工具是 Python 包需要装一下pip install tensorflowjs假设你在 Python 里训练了一个 Keras 模型保存成了my_model.h5转换命令如下tensorflowjs_converter \ --input_formatkeras \ --output_formattfjs_layers_model \ my_model.h5 \ ./web_model如果是 TensorFlow SavedModel 格式则用tensorflowjs_converter \ --input_formattf_saved_model \ --output_formattfjs_graph_model \ --output_node_namesprediction \ ./saved_model \ ./web_model转换完成后web_model目录里会生成两类文件一个model.json描述网络结构和权重文件的索引若干个.bin分片文件存放实际权重。浏览器加载时读model.json再按索引加载各个分片。分片是自动产生的一个 20MB 的模型会被切成多个几 MB 的小文件这是为了方便浏览器并行加载和缓存。5.2 前端怎么加载转换后的模型转换后的模型在前端加载方式很简单const model await tf.loadLayersModel(/models/model.json); const tensor tf.browser.fromPixels(imageEl) .resizeBilinear([224, 224]) .toFloat() .div(127.5) .sub(1); const result model.predict(tensor.expandDims(0));注意看预处理部分这里我用的是.div(127.5).sub(1)结果同样是 -1 到 1只是计算顺序和 MobileNet 官方示例相反。处理结果完全一样但有人会把这当成两种不同的归一化方式其实只是数学顺序不同。这提醒我们一件事接手别人的模型时一定要搞清楚训练时的预处理和推理时的预处理必须完全一致这是机器学习部署里最常见的连环坑。5.3 量化压缩让模型从 20MB 瘦到 5MB模型体积是浏览器端推理的实际瓶颈。TensorFlow.js 转换工具提供了权重量化功能可以在模型大小和精度之间做取舍。命令加几个参数tensorflowjs_converter \ --input_formatkeras \ --output_formattfjs_layers_model \ --quantization_bytes1 \ my_model.h5 \ ./web_model_quantized--quantization_bytes1表示用 1 个字节存浮点权重原来 4 字节的 float32 直接变成 1 字节的 uint8理论体积缩小到四分之一。代价是精度轻微下降。实际项目中图像分类、目标检测这类任务用 1 字节量化通常感知不到差别但有些对数值敏感的模型可能掉点明显建议上线前做一次全面评估。我通常的流程是先不量化部署一版测量模型体积和推理精度再量化一版对比两次结果差距在业务可接受范围内就上量化版。6. 性能优化与内存管理别让页面卡成幻灯片浏览器端推理最大的问题不是“能不能跑”而是“跑得有多流畅”。我自己优化过几个项目把实操经验浓缩在这章。6.1 推理延迟分布在哪一次完整的浏览器端推理包含三段时间模型加载时间、数据预处理时间、核心推理时间。模型加载是一次性的但体积大时慢得让人焦虑20MB 模型在弱网下可能要等十几秒。这时要做两件事加载时给出进度反馈不能让用户对着白屏干等提前把模型预加载并且缓存好第二次访问时浏览器磁盘缓存直接命中。数据预处理往往比想象中慢。tf.browser.fromPixels操作涉及像素数据的 CPU 到 GPU 拷贝大图比如 4000x3000 像素的高清照片转张量就非常耗时。我建议在进入模型前先对图片做压缩具体说用 canvas 先把图片画成模型需要的尺寸比如 224x224再做后续操作比直接对原图操作快得多。这也是 6.2 里内存优化的前提因为在相同尺寸下更小的输入意味着更少的计算量。6.2 内存泄漏与 dispose 的正确姿势回看 4.2 的代码有一行被很多人忽略tensor.dispose()。TensorFlow.js 管理的内存和普通 JavaScript 对象不同它底层大量使用 WebGL 纹理和 GPU 显存JavaScript 的垃圾回收器管不到这些资源。每次fromPixels、resize、div都会产生新的张量如果不显式释放GPU 显存就会持续增长页面表现得越来越卡最后直接崩溃。更规范的做法是用tf.tidy包装整个推理过程const predictions tf.tidy(() { const tensor tf.browser.fromPixels(imageEl) .resizeNearestNeighbor([224, 224]) .toFloat() .sub(127.5) .div(127.5) .expandDims(0); return model.predict(tensor); });tf.tidy会在回调函数执行完后自动释放内部创建的所有中间张量只保留返回值。这个 API 是前端 ML 的救星用好了基本不会遇到内存问题。注意一个细节返回值本身不会自动释放如果后续还要用用完还是要手动 dispose。对内存的直观感受是浏览器标签页打开开发者工具切到内存面板连续做推理操作观察 GPU 内存曲线。如果曲线只升不降那就是有张量泄漏了。我排查过的项目里几乎都是缺了 dispose 或者没用 tidy。6.3 后端降级与低端设备策略并非所有用户的浏览器都支持 WebGL。部分低端安卓设备的 WebGL 实现有 bug或者用户关闭了硬件加速。TensorFlow.js 在检测到 WebGL 不可用时会自动降级到 WASM但如果你在业务代码里显式指定了后端降级逻辑就会被绕过。建议在应用启动时做一个后端检测await tf.setBackend(webgl);如果设置失败再回退到 wasmtry { await tf.setBackend(webgl); } catch (e) { await tf.setBackend(wasm); }低端设备上还要考虑模型选择策略。同一类功能可以准备大模型和小模型两个版本根据设备性能动态选择。比如检测设备内存小于 2GB 时用 alpha 0.25 的轻量模型否则用 alpha 1.0 的标准模型。这个策略不需要多复杂的判断逻辑换来的是大量低端用户的可用性提升。7. 常见问题与排查技巧实录最后把我在实际项目中碰到的问题整理成速查表每个都附上排查思路。这些坑大概率你也会踩到早看早避。现象常见原因排查方法控制台报 404 模型加载失败模型路径错误或静态资源目录没配置对确认 model.json 的 URL 从服务器访问正常图片选了没反应File 对象没有转成图片元素在 onload 回调里加断点确认图片是否加载完成推理结果全部是垃圾值预处理和训练时不匹配归一化范围或通道顺序错误对照训练脚本里的预处理代码逐行核对页面越用越卡张量没有释放GPU 显存泄漏用 tf.engine().numTensors() 监控张量数量是否持续增长移动端推理特别慢WebGL 不可用或模型太大检查 tf.getBackend() 返回的后端类型考虑模型量化模型加载频繁超时模型文件过大多次加载用 HTTP 缓存并做模型预加载7.1 一个被反复问到的 CORS 问题很多新手在本地把 HTML 文件直接拖进浏览器然后发现模型加载失败控制台报跨域错误。这其实不是 TensorFlow.js 的坑而是浏览器安全策略的正常表现。file://协议下请求任何外部资源都会被拦截模型文件放在本地也逃不掉。解决方案就是用本地 HTTP 服务也就是前面我用 Vite 的原因。不要在这一步折腾什么代理、压缩工具先把开发服务器起对后面全流程都会顺很多。7.2 三个最容易忽略的隐藏坑第一个移动端图片方向。手机拍照的图片默认可能带 EXIF 旋转信息直接fromPixels读出来可能是横着的。处理办法是用 canvas 配合createImageBitmap做一次标准渲染让浏览器自动处理方向再进模型。第二个批量推理和单张推理的差异。如果一次要识别多张图片建议把它们堆成一个 batch用tf.stack组合成[n, 224, 224, 3]的张量一次 predict 完成。批量推理比循环单张推理快很多在 GPU 上原因是矩阵乘法本身利用了并行计算。第三个模型的分类标签映射表。MobileNet 这类预训练模型自带标签映射你的自定义模型没有。推理结果只是类别编号必须在代码里维护一份 id 到名称的映射表。很多人训练模型时觉得 LabelMap 无所谓到了前端才发现没有映射表根本没法展示结果。最后再多说一句关于模型安全的话。前端模型文件对用户是可见的任何人都可能通过开发者工具把模型下载走。如果你的模型本身是核心商业资产部署在浏览器端需要权衡泄漏风险。这一点在产品设计早期就应该想清楚而不是等技术选型做完了再纠结。我个人做这个项目的最大感受是TensorFlow.js 让我把“机器学习能力”和“前端交互体验”彻底融在了一起。以前做智能功能前后端之间的网络请求、数据格式对齐、服务运维全部要管现在模型在浏览器里焦点回到了模型本身和用户体验本身。如果你手上正好有“必须跑在浏览器里的机器学习”需求我的建议是先跑通一个最小 Demo把模型体积和后端加速这两件事摸清楚然后再往工程化方向扩展。后面模型推理速度还可以往 WebGPU 方向继续挖掘模型体积和首次下载体验也还有很大优化空间这套技术栈的演进速度其实比大多数人想象中要快。
上一篇/下一篇内容由系统自动关联 返回资讯列表 →