1. 为什么要把模型搬到用户设备上跑
1.1 从一次线上事故说起
去年我负责的一个图像分类小工具上线,用户上传图片到服务端,后端用 Python 加载模型做推理,再把结果返回前端。功能本身不复杂,但上线第三天就出事了:一张 4K 分辨率的照片,从上传到拿到结果整整等了 6 秒多,用户以为页面卡死,直接关掉走人。更麻烦的是,那天流量稍微一冲,后端推理队列就排满了,CPU 打满,连带其他接口一起变慢。
那次之后我认真复盘了一遍,问题的根子不在模型精度,也不在服务器配置,而在于架构选错了。图片分类这种任务,模型本身只有几 MB,推理耗时在本地设备上也就几十毫秒,我却非要绕一圈:图片编码上传、网络传输、服务端排队、推理、结果回传。这一圈下来,延迟被放大了几十倍,还白白消耗了带宽和服务器算力。
后来我把整个推理链路搬到了浏览器里,用TensorFlow.js重写了一遍。同样的图片,本地推理 80 毫秒出结果,服务器只负责托管静态资源,压力几乎归零。这就是端侧推理最朴素的价值:数据不出设备、延迟极低、服务器成本可控。
1.2 端侧推理到底解决了什么问题
很多人一听到"机器学习"就默认要跑在服务器上,其实这是个思维惯性。把模型放到用户设备上跑,至少能解决四类问题。
第一类是隐私敏感。用户的照片、语音、输入的文字,如果全部上传到服务器,就涉及数据合规和用户信任问题。端侧推理让原始数据压根不离开设备,只把结果留在本地,这在医疗、金融、个人助理类应用里是刚需。
第二类是延迟敏感。实时手势识别、视频滤镜、输入法联想这类场景,用户对延迟的容忍度是毫秒级的。走一趟网络往返,光 RTT 就几十毫秒起步,体验直接崩掉。
第三类是成本敏感。推理是要烧算力的,用户量一大,GPU 服务器账单非常可观。把推理分摊到千万台用户设备上,等于让用户帮你免费算,边际成本趋近于零。
第四类是离线可用。网络不稳定或者干脆没网的环境下,服务端推理直接歇菜,端侧推理照样能跑。
1.3 TensorFlow.js 在端侧推理里的位置
端侧推理不是只有一条路。移动端可以用 TFLite,桌面端可以用 ONNX Runtime,浏览器里则主要是TensorFlow.js。它最大的优势是跨平台、零安装、开箱即用——只要用户有一个现代浏览器,你的模型就能跑起来,不需要用户下载 App,不需要考虑 iOS 还是 Android。
TensorFlow.js 支持三种运行后端:纯 CPU 的 JavaScript 后端、基于 WebGL 的 GPU 后端,以及新一代的WebGPU后端。它还支持Web Worker,可以把推理放到后台线程,避免阻塞主线程导致页面卡顿。这几个关键词——TensorFlow.js、端侧推理、WebGPU、Web Worker——基本就是浏览器端机器学习的全部核心拼图。
这篇文章我会把这套东西从头到尾讲一遍:模型怎么来、怎么加载、怎么选后端、怎么用 Worker 隔离、怎么调优。适合已经会一点前端、想入门端侧机器学习的人,也适合做后端想了解前端推理能力的同学。
2. 整体方案设计与技术选型思路
2.1 三种模型来源,怎么选
在浏览器里跑模型,第一步是搞清楚模型从哪来。TensorFlow.js 支持三种模型来源,各有适用场景。
第一种是直接用 JavaScript 定义模型。用tf.sequential()或者tf.model()一层层搭,适合结构简单、层数少的模型,比如几层的全连接网络。优点是灵活、可调试,缺点是复杂模型手写太痛苦,而且没法复用 Python 那边训练好的权重。
第二种是转换已有的 Keras / TensorFlow 模型。这是生产环境最常用的方式。你在 Python 里用 Keras 训练好模型,保存成 SavedModel 或者 HDF5,然后用tensorflowjs_converter工具转成 TensorFlow.js 能识别的格式。转换后会得到一个model.json加一组.bin权重分片文件。
第三种是直接用现成的预训练模型。TensorFlow.js 官方维护了一批开箱即用的模型,比如 MobileNet、PoseNet、COCO-SSD、FaceMesh 等,通过 npm 包直接引入,几行代码就能用。
我的建议是:新手先用官方预训练模型跑通流程,理解推理链路;有自定义需求再走转换路线;只有极简模型才考虑纯 JS 手写。
2.2 后端选择:CPU、WebGL 还是 WebGPU
这是端侧推理里最影响性能的一个决策。TensorFlow.js 的"后端"指的是底层用什么算力执行张量运算。
| 后端 | 算力来源 | 适用场景 | 主要限制 |
|---|---|---|---|
| CPU (js) | 纯 JavaScript | 小模型、兼容性兜底 | 速度慢,大模型卡顿明显 |
| WebGL | 显卡 GPU | 主流选择,兼容性好 | 部分算子不支持,精度偶有偏差 |
| WebGPU | 新一代 GPU API | 新浏览器,性能最佳 | 浏览器支持还在铺开 |
选后端的逻辑很简单:能用 WebGPU 就用 WebGPU,不支持就退到 WebGL,再不行才用 CPU。TensorFlow.js 提供了tf.setBackend()和tf.ready()来做这件事,后面实操部分我会给出完整的探测代码。
这里有个容易被忽略的点:WebGL 后端在移动端和桌面端的表现差异很大。桌面独显跑 WebGL 很稳,但一些老手机的 GPU 驱动对浮点纹理支持不完整,可能出现结果异常。所以上线前一定要在目标机型上实测,不能只在开发机上跑通就完事。
2.3 为什么要用 Web Worker
JavaScript 是单线程的,主线程既要处理 DOM 渲染,又要响应用户交互。如果你在主线程里跑一个 200 毫秒的推理,页面就会卡顿 200 毫秒,用户滑动、点击全部延迟。模型越大,卡顿越明显。
Web Worker的作用就是开一个后台线程,把推理任务丢进去,主线程继续负责 UI。推理完成后通过postMessage把结果传回来。这样即使推理要跑几百毫秒,页面依然丝滑。
不过 Worker 也有代价:线程间通信需要序列化数据,传大张量会有开销。所以我的经验是:推理本身耗时超过 50 毫秒的,就值得上 Worker;几十毫秒以内的小模型,主线程直接跑反而更省事。
2.4 一个完整的端侧推理架构长什么样
把上面的选择拼起来,一个生产级的端侧推理架构大致是这样:
- 模型文件托管在 CDN,首次加载后由浏览器缓存
- 页面启动时探测后端能力,优先 WebGPU,回退 WebGL
- 推理逻辑封装在 Web Worker 里,主线程只负责输入采集和结果展示
- 输入数据(图片、摄像头帧)在主线程预处理后传给 Worker
- Worker 内完成张量转换、推理、后处理,把结构化结果回传
这套架构的好处是职责清晰、可维护、性能可控。下面几节我会把每一块拆开讲透。
3. 核心细节解析与实操要点
3.1 模型转换:从 Python 到浏览器
假设你在 Python 里训练了一个图像分类模型,保存为saved_model目录。转换命令大致如下:
tensorflowjs_converter \ --input_format=tf_saved_model \ --output_format=tfjs_graph_model \ --signature_name=serving_default \ --saved_model_tags=serve \ ./saved_model \ ./web_model转换完成后,web_model目录里会有model.json和若干group1-shard1ofN.bin文件。model.json描述网络结构,.bin是权重分片。
这里有几个坑我必须提醒:
注意:转换工具对 TensorFlow 版本很敏感。转换脚本的版本最好和训练时用的 TF 版本对齐,否则可能出现算子不支持或者权重对不上的问题。我踩过一次,训练用 TF 2.8,转换环境是 TF 2.15,结果某个自定义层直接报错,折腾了半天才发现是版本问题。
另一个要点是权重量化。默认转换出来是 float32,模型体积大。可以用--quantize_float16把权重压成 16 位浮点,体积直接减半,精度损失通常很小。如果对精度要求不高,还能用--quantize_uint8压到 8 位,体积再降,但精度损失要实测评估。
3.2 模型加载与预热
模型加载是端侧推理里最容易被低估的环节。一个 10MB 的模型,在慢网络下加载可能要好几秒。所以加载策略要讲究。
import * as tf from '@tensorflow/tfjs'; async function loadModel() { // 优先尝试 WebGPU,失败回退 WebGL try { await tf.setBackend('webgpu'); await tf.ready(); console.log('使用 WebGPU 后端'); } catch (e) { await tf.setBackend('webgl'); await tf.ready(); console.log('回退到 WebGL 后端'); } const model = await tf.loadGraphModel('/models/web_model/model.json'); return model; }加载完之后,强烈建议做一次预热推理。第一次推理往往包含算子编译、显存分配等一次性开销,耗时可能是后续推理的好几倍。用一个全零张量先跑一遍,把这条路径"热"起来,用户真正用的时候就不会遇到第一次特别慢的情况。
// 预热:用真实输入形状的零张量跑一次 const warmupInput = tf.zeros([1, 224, 224, 3]); const warmupResult = model.predict(warmupInput); warmupResult.dispose(); warmupInput.dispose();3.3 张量内存管理:别让内存泄漏拖垮页面
TensorFlow.js 的张量是手动管理内存的,这一点和 Python 的自动垃圾回收完全不同。你每创建一个张量,就占用一块显存或内存,不主动dispose()就不会释放。在长时间运行的应用里,这会导致内存持续增长,最后页面崩溃。
我见过最常见的泄漏场景是:在摄像头循环里每帧都创建张量,但从不释放。跑几分钟内存就爆了。
// 错误示范:每帧都泄漏 function processFrame(imageElement) { const tensor = tf.browser.fromPixels(imageElement); const result = model.predict(tensor); return result; // tensor 和中间张量都没释放 } // 正确做法:用 tf.tidy 自动清理 function processFrame(imageElement) { return tf.tidy(() => { const tensor = tf.browser.fromPixels(imageElement); const result = model.predict(tensor); return result; // tidy 会释放除返回值外的所有张量 }); }tf.tidy()是官方推荐的清理方式,它会自动释放回调函数内创建的所有张量,只保留返回值。但要注意,返回值本身还是需要调用方手动释放。另外,异步函数里不能用tidy,得手动dispose。
提示:调试内存问题时,可以定期打印
tf.memory().numTensors,如果这个数字持续增长不回落,基本就是泄漏了。
3.4 输入预处理:别在细节上翻车
模型对输入的要求很严格,形状、数值范围、通道顺序错一个,结果就完全不对。常见的预处理包括:
- 尺寸缩放:模型要求 224x224,你的图片是 1920x1080,得先缩放
- 归一化:很多模型要求输入在 [0,1] 或 [-1,1],而
fromPixels出来是 [0,255] - 通道顺序:TensorFlow 用 RGB,某些模型可能要求 BGR
function preprocess(imageElement) { return tf.tidy(() => { let tensor = tf.browser.fromPixels(imageElement); // 缩放到模型输入尺寸 tensor = tf.image.resizeBilinear(tensor, [224, 224]); // 归一化到 [0,1] tensor = tensor.toFloat().div(255.0); // 增加 batch 维度 tensor = tensor.expandDims(0); return tensor; }); }这里有个性能细节:tf.image.resizeBilinear在 GPU 上跑很快,但如果你的图片本来就接近目标尺寸,直接用fromPixels的第二个参数指定尺寸可能更省事。另外,如果输入源是canvas或video,fromPixels能直接读取,不用先转成 ImageData。
4. 完整实操流程与核心环节实现
4.1 项目初始化与依赖安装
我用 Vite 搭一个最小可运行的项目,这样打包快、开发体验好。
npm create vite@latest tfjs-demo -- --template vanilla cd tfjs-demo npm install @tensorflow/tfjs @tensorflow/tfjs-backend-webgpu npm install注意@tensorflow/tfjs默认包含 CPU 和 WebGL 后端,WebGPU 后端需要单独装@tensorflow/tfjs-backend-webgpu。装完之后在入口文件里引入,后端才会注册。
4.2 后端探测与初始化
后端探测要放在所有推理之前,而且要考虑异步就绪的问题。
import * as tf from '@tensorflow/tfjs'; import '@tensorflow/tfjs-backend-webgpu'; async function initBackend() { const backends = ['webgpu', 'webgl', 'cpu']; for (const name of backends) { try { const ok = await tf.setBackend(name); if (ok) { await tf.ready(); console.log(`后端就绪: ${tf.getBackend()}`); return tf.getBackend(); } } catch (e) { console.warn(`${name} 不可用,尝试下一个`); } } throw new Error('没有可用的后端'); }这段代码的关键是await tf.ready()。setBackend只是切换,真正初始化完成要等ready。很多人漏了这一步,结果第一次推理报错说后端没准备好。
4.3 用 Web Worker 隔离推理
Worker 的写法分两部分:主线程侧和 Worker 侧。
主线程侧:
const worker = new Worker(new URL('./inference.worker.js', import.meta.url), { type: 'module' }); worker.onmessage = (e) => { const { type, payload } = e.data; if (type === 'result') { renderResult(payload); } }; // 发送图片数据(ImageData 是可转移对象,性能好) function requestInference(imageData) { worker.postMessage({ type: 'infer', payload: imageData }, [imageData.data.buffer]); }Worker 侧:
import * as tf from '@tensorflow/tfjs'; import '@tensorflow/tfjs-backend-webgpu'; let model = null; async function ensureModel() { if (model) return model; await tf.setBackend('webgl'); await tf.ready(); model = await tf.loadGraphModel('/models/web_model/model.json'); // 预热 tf.tidy(() => model.predict(tf.zeros([1, 224, 224, 3]))); return model; } self.onmessage = async (e) => { const { type, payload } = e.data; if (type !== 'infer') return; const m = await ensureModel(); const result = tf.tidy(() => { let input = tf.browser.fromPixels(payload); input = tf.image.resizeBilinear(input, [224, 224]).toFloat().div(255); input = input.expandDims(0); return m.predict(input); }); const data = await result.data(); result.dispose(); self.postMessage({ type: 'result', payload: Array.from(data) }); };这里有个重要细节:Worker 里不能直接用document和canvas,但tf.browser.fromPixels可以接受ImageData对象。所以主线程要把 canvas 转成ImageData再传过去。用postMessage的第二个参数做转移(transfer),可以避免拷贝开销。
注意:WebGPU 后端在 Worker 里的支持情况因浏览器而异,实测 Chrome 较新版本可以,但 Safari 和部分 Firefox 版本还不稳定。生产环境建议 Worker 里先用 WebGL,主线程再根据能力决定是否用 WebGPU。
4.4 摄像头实时推理的完整链路
把摄像头、Worker、渲染串起来,就是一个实时推理应用。
const video = document.getElementById('video'); const canvas = document.getElementById('canvas'); const ctx = canvas.getContext('2d'); async function startCamera() { const stream = await navigator.mediaDevices.getUserMedia({ video: { width: 640, height: 480 } }); video.srcObject = stream; await video.play(); loop(); } let running = false; function loop() { if (!running) return; if (video.readyState === video.HAVE_ENOUGH_DATA) { canvas.width = video.videoWidth; canvas.height = video.videoHeight; ctx.drawImage(video, 0, 0); const imageData = ctx.getImageData(0, 0, canvas.width, canvas.height); requestInference(imageData); } requestAnimationFrame(loop); }这里我用requestAnimationFrame驱动循环,而不是setInterval。原因是rAF和浏览器渲染节奏同步,不会在页面不可见时浪费算力。但要注意,如果推理比帧率慢,会积压任务。实际项目里我会加一个"上一帧没处理完就跳过"的标志位,避免队列堆积。
4.5 性能实测数据
我在一台 2021 款 MacBook Pro(M1 Pro)和一台中端安卓机上测了同一个 MobileNet 模型(输入 224x224),数据如下:
| 环境 | 后端 | 单次推理耗时 | 首帧耗时 |
|---|---|---|---|
| M1 Pro / Chrome | WebGPU | 12ms | 180ms |
| M1 Pro / Chrome | WebGL | 18ms | 220ms |
| M1 Pro / Chrome | CPU | 95ms | 400ms |
| 中端安卓 / Chrome | WebGL | 45ms | 650ms |
| 中端安卓 / Chrome | CPU | 320ms | 1200ms |
从数据能看出几个规律:WebGPU 比 WebGL 快 30% 左右,CPU 后端慢一个数量级;首帧耗时远高于稳态,预热非常必要;移动端整体比桌面慢 2-3 倍,但 WebGL 依然可用。
5. 常见问题与排查技巧实录
5.1 推理结果和 Python 对不上
这是最高频的问题。模型在 Python 里跑得好好的,搬到浏览器结果就偏了。排查顺序建议这样:
第一,检查预处理是否一致。Python 里用的归一化方式、通道顺序、resize 算法,浏览器里必须一模一样。我遇到过 Python 用cv2.resize的双线性插值,浏览器用resizeBilinear,理论上一样,但边界处理有细微差异,导致结果有零点几个百分点的偏差。
第二,检查后端精度。WebGL 后端默认用 float32,但某些算子会降精度。可以试试强制 CPU 后端跑一遍,如果 CPU 结果对、WebGL 不对,那就是后端精度问题。
第三,检查模型转换。用tfjs_graph_model还是tfjs_layers_model,输出节点名字对不对,这些都会影响结果。
5.2 页面卡顿、掉帧
如果推理在主线程跑,卡顿几乎是必然的。解决方案就是上 Worker。但上了 Worker 还卡,通常是这两个原因:
一是数据传输开销大。每帧传一个 640x480 的 ImageData,就是 1.2MB,60 帧就是 70MB/s 的拷贝量。用 transfer 转移 buffer 可以避免拷贝,但转移后主线程的 buffer 就失效了,得重新创建。
二是Worker 里也在做重活。如果预处理(resize、归一化)也在 Worker 里做,那 Worker 线程本身就很忙。可以考虑把预处理放主线程用 GPU 加速,Worker 只做推理。
5.3 内存持续增长
前面提过,核心原因是张量没释放。排查方法:
setInterval(() => { const info = tf.memory(); console.log(`张量数: ${info.numTensors}, 字节: ${info.numBytes}`); }, 5000);如果numTensors一直涨,就去找哪里创建了张量没释放。常见位置:fromPixels、predict的返回值、中间计算张量。用tf.tidy包起来基本能解决 90% 的问题。
5.4 常见问题速查表
| 现象 | 可能原因 | 解决方向 |
|---|---|---|
| 首次推理特别慢 | 未预热,算子编译开销 | 加载后跑一次零张量预热 |
| 结果精度偏差 | 预处理不一致 / 后端精度 | 对齐预处理,试 CPU 后端对比 |
| 页面卡顿 | 主线程推理 | 迁移到 Web Worker |
| 内存持续增长 | 张量未释放 | 用 tf.tidy 或手动 dispose |
| WebGPU 报错 | 浏览器不支持 | 回退 WebGL |
| Worker 里报错 | 后端未注册 | Worker 内单独 import 后端包 |
| 移动端结果异常 | GPU 驱动兼容性 | 回退 CPU 或换 WebGL 参数 |
5.5 几个我踩过的坑
第一个坑是模型文件没配 CORS。模型托管在 CDN 上,如果没开跨域头,loadGraphModel会直接失败,而且报错信息很含糊。解决办法是在 CDN 配置里加上Access-Control-Allow-Origin。
第二个坑是在 Worker 里用了主线程的 tf 实例。Worker 是独立上下文,必须自己 import 一遍 TensorFlow.js,不能共享主线程的实例。这个错误很隐蔽,因为不报错,只是行为异常。
第三个坑是忽略了tf.ready()。切换后端后不 await ready,直接推理,会拿到未初始化的后端,结果不可预测。这个坑我在两个项目里都踩过,现在养成了切换后端必 await 的习惯。
第四个坑是模型版本和代码不匹配。模型更新了但前端缓存了旧版本,导致输入输出对不上。解决办法是给模型 URL 加版本号或者 hash,强制刷新缓存。
6. 性能调优与进阶方向
6.1 模型层面的优化
端侧推理的性能,一半靠模型本身。几个有效的优化手段:
量化是最直接的。float16 量化体积减半,精度损失通常小于 1%;int8 量化体积再减半,但需要校准数据集,精度损失要评估。对于分类任务,int8 往往够用;对于检测和分割,建议谨慎。
剪枝是去掉模型中不重要的权重,让模型变稀疏。TensorFlow.js 对稀疏模型的支持有限,实际收益不如量化明显。
换更小的骨干网络。MobileNetV3 比 MobileNetV2 又快又准,EfficientNet-Lite 系列也是端侧友好。如果精度允许,直接换小模型比任何优化都有效。
6.2 运行时层面的优化
批处理在端侧通常不适用,因为实时场景一次就一张图。但如果你的场景是批量处理(比如一次上传多张图),把 batch size 调大能提升 GPU 利用率。
输入尺寸是性能的平方级影响因素。224x224 换成 320x320,计算量增加一倍多。如果精度允许,尽量用小输入。
复用张量能减少分配开销。对于固定形状的输入,可以预分配一个张量,每次往里写数据,而不是每帧新建。
6.3 WebGPU 的现状与预期
WebGPU 是浏览器 GPU 计算的新标准,相比 WebGL 有几个本质优势:支持计算着色器、支持更灵活的内存布局、没有 WebGL 的那些历史包袱。实测下来,同样的模型 WebGPU 比 WebGL 快 20%-40%,模型越大差距越明显。
目前的限制主要是浏览器支持。Chrome 从 113 版本开始默认开启,Edge 跟进,Firefox 和 Safari 还在推进中。所以生产环境必须做好回退,不能只依赖 WebGPU。
6.4 端侧推理还能怎么扩展
跑通基础推理之后,有几个方向可以继续深入。
一是端侧训练。TensorFlow.js 支持在浏览器里做迁移学习,用用户的数据微调模型。比如一个手写识别应用,可以让用户教它认自己的字迹,数据不出设备。这个方向隐私友好,但训练开销大,要控制好规模。
二是模型集成。把多个小模型串起来,各司其职。比如先用人脸检测定位,再用关键点模型提取特征,最后用分类模型判断表情。这种流水线在端侧完全可行。
三是和 WebAssembly 结合。TensorFlow.js 的 WASM 后端在某些场景下比 WebGL 更稳,尤其是需要高精度计算的时候。可以把它作为 WebGL 之外的另一个回退选项。
四是离线优先的 PWA。把模型和代码都缓存到 Service Worker 里,用户装一次之后完全离线可用。这对网络不稳定的场景特别有价值。
我在实际项目里的体会是,端侧推理不是"能不能跑"的问题,而是"怎么跑得好"的问题。模型选型、后端选择、Worker 隔离、内存管理,每一环都影响最终体验。把这几点做扎实,浏览器里的机器学习完全能达到生产可用的水平。最后再分享一个小技巧:如果你的模型加载慢,可以在页面空闲时用requestIdleCallback提前预加载,等用户真正触发功能时,模型早就准备好了,感知延迟几乎为零。