
机器学习模型训练完只是走完了上半场真正决定用户体验的往往是推理这一段跑在哪里、跑得怎么样。把模型塞进浏览器让它在用户的设备上直接完成计算这件事在几年前还像是玩具如今随着 WebGPU 的普及和 TensorFlow.js 的持续迭代已经能撑起不少真实业务场景了。我最近用 TensorFlow.js 做了一个端侧推理的小项目从模型转换、算子兼容、线程调度到内存回收踩了一圈坑也摸清了一些门道。这篇就把整个实战过程拆开讲包括为什么选端侧、WebGPU 和 WebGL 后端怎么选、Web Worker 怎么配合、模型转换时哪些算子会翻车、性能怎么量化。不管你是刚接触机器学习想找个能跑起来的落地场景还是已经会训练模型但不知道怎么搬到前端这篇应该都能给你一些能直接抄的配置和思路。1. 为什么要把推理放到用户设备上1.1 端侧推理解决的三个真实痛点先说清楚动机不然很容易陷入为了用新技术而用的陷阱。把推理放到用户设备上最直接的好处是数据不出端。用户的照片、输入文本、传感器数据全在本地处理不上传服务器这在隐私敏感的场景里是硬需求。比如做一个本地的人脸打码工具、一个离线的文本分类器用户根本不愿意把原始数据传出去端侧推理就是唯一解。第二个痛点是延迟。一次网络往返哪怕服务器就在同城加上排队、序列化、反序列化轻松就是几十到几百毫秒。而端侧推理省掉了这一整段用户点一下按钮结果几乎立刻出来。对于交互式应用比如实时滤镜、手写识别、游戏里的 AI 对手这个差距是体验级的。第三个是成本。推理放在服务器上每一个用户每一次调用都在烧 GPU 和带宽。用户量一上来账单非常可观。把计算分摊到用户设备上服务器只需要负责分发模型文件边际成本几乎为零。当然代价是首次加载模型有流量开销这个后面会讲怎么优化。提示端侧推理不是万能药。模型特别大、需要频繁更新、或者用户设备性能参差不齐的场景仍然要考虑服务端推理或者混合方案。选型前先想清楚你的瓶颈到底在哪。1.2 TensorFlow.js 在端侧生态里的位置端侧推理的方案不止一种ONNX Runtime Web、WebLLM、原生 WebAssembly 手写算子都有人用。TensorFlow.js 的优势在于生态完整从 Python 侧的 Keras/TF 模型转换工具到浏览器里的推理 API再到可视化调试工具是一条龙的。你训练好的模型用tensorflowjs_converter一转就能用不用自己写算子绑定。它的运行时后端也做得比较成熟支持 WebGL、WebGPU、WASM 以及纯 CPU 回退。这意味着同一份代码在高性能设备上走 WebGPU在老设备上自动降级到 WebGL 或 WASM兼容性不用你操心。对于我这种不想在设备适配上花太多精力的人来说这一点非常省心。另外它的 API 设计对前端开发者友好tf.loadLayersModel、model.predict这些调用方式和 Keras 很像学习成本低。如果你本来就会一点 Python 侧的机器学习迁移过来几乎没有门槛。1.3 什么样的项目适合端侧推理不是所有项目都适合。我总结了几条判断标准你可以对照自己的场景模型体积可控量化后能压到几 MB 到几十 MB 之间首次加载不会让用户等太久。推理频率高或交互性强比如每帧都要跑、每次输入都要跑放服务端延迟受不了。数据隐私敏感原始数据不适合离开用户设备。离线可用是加分项用户在地铁、飞机上也能用。模型更新不频繁不需要每天热更新权重。反过来如果模型是几百 MB 的大语言模型、需要频繁 A/B 测试不同版本、或者用户设备普遍很弱那端侧就不是好选择。我见过有人硬把一个大模型塞进浏览器结果加载要一分钟用户体验极差这就属于选型失误。2. 环境搭建与后端选型WebGPU、WebGL 还是 WASM2.1 三种后端的能力边界对比TensorFlow.js 的运行时后端决定了算子跑在什么硬件上直接关系到性能。选错后端可能慢十倍。我把三种主要后端的特性整理成表方便你对照后端硬件适用场景主要限制WebGPUGPU现代 API大矩阵运算、卷积、Transformer浏览器支持还在铺开部分算子缺失WebGLGPU老 API兼容性最好的 GPU 方案精度受限部分算子用 CPU 回退WASMCPU多线程小模型、GPU 不可用时的兜底大模型慢受内存限制CPUCPU单线程调试、极简场景最慢仅用于验证WebGPU 是这几年的重点它比 WebGL 更接近现代图形 API支持计算着色器做矩阵乘法这类密集计算效率高很多。实测下来同一个卷积模型WebGPU 比 WebGL 快 2 到 5 倍不等具体看模型结构。但它的坑在于浏览器覆盖率虽然主流浏览器的新版本都在推但用户群里总有一部分人用不了所以必须有回退方案。WebGL 是当前兼容性最好的 GPU 后端几乎所有支持 WebGL 2.0 的浏览器都能跑。它的短板是精度和算子覆盖某些操作会因为浮点精度问题走 CPU 回退反而变慢。WASM 则是纯 CPU 方案配合多线程和 SIMD 能跑出不错的速度适合小模型或者 GPU 完全不可用的环境。2.2 后端自动选择的代码实现我的做法是写一个优先级探测函数按 WebGPU、WebGL、WASM 的顺序尝试哪个能用用哪个。TensorFlow.js 提供了tf.setBackend和tf.ready配合tf.getBackend可以确认当前实际生效的后端。import * as tf from tensorflow/tfjs; import tensorflow/tfjs-backend-webgpu; import tensorflow/tfjs-backend-wasm; async function initBackend() { const candidates [webgpu, webgl, wasm, cpu]; for (const name of candidates) { 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()。设置后端是异步的如果你不等待就调用model.predict可能拿到错误的结果或者直接报错。我一开始就踩过这个坑页面加载后立刻推理结果第一次总是失败第二次才正常排查半天才发现是没等ready。注意WASM 后端需要额外配置.wasm文件的路径通常用setWasmPaths指定 CDN 或本地目录。如果路径不对WASM 会静默回退到 CPU性能断崖式下跌而且不容易发现。2.3 后端切换带来的精度差异这里有个容易被忽略的问题不同后端的浮点精度不完全一致。WebGL 在某些设备上用的是 16 位浮点纹理WebGPU 和 WASM 通常是 32 位。这意味着同一个模型在不同后端上输出的数值可能有微小差异。对于分类任务这个差异通常不影响最终结果因为 argmax 之后类别还是一样的。但对于回归任务比如预测一个具体数值差异就可能被放大。我的经验是如果业务对数值精度敏感要么统一强制用某个后端要么在训练时就考虑到量化误差让模型对精度不敏感。实测中我还遇到过一个情况某个自定义算子在 WebGL 上走 CPU 回退导致整个推理链路被拖慢。排查方法是打开 TensorFlow.js 的调试日志看哪些算子被回退了。这个后面在性能优化章节会详细讲。3. 模型转换从 Python 到浏览器的完整链路3.1 转换前的模型瘦身策略模型转换不是简单地把文件格式换一下转换前该做的瘦身一定要做否则转出来的文件大得吓人。我一般按这个顺序处理剪枝去掉对输出贡献小的权重减小模型规模。量化把 32 位浮点权重压成 16 位甚至 8 位整数体积能降一半到四分之三。算子融合把连续的卷积、批归一化、激活函数合并减少推理时的算子调用次数。量化是性价比最高的一步。TensorFlow.js 支持训练后量化转换时加参数就行。但要注意量化会带来精度损失尤其是 8 位整数量化对某些模型影响明显。我的做法是先量化然后在验证集上跑一遍看精度掉多少能接受就用不能接受就退回 16 位浮点。tensorflowjs_converter \ --input_formatkeras \ --output_formattfjs_graph_model \ --quantize_uint8 \ model.h5 \ ./web_model上面这条命令把 Keras 模型转成 TensorFlow.js 的 GraphModel 格式并做 8 位量化。--quantize_uint8是体积最小的选项如果精度不够可以换成--quantize_float16。3.2 GraphModel 与 LayersModel 的选择TensorFlow.js 支持两种模型格式LayersModel 和 GraphModel。前者对应 Keras 的层式模型后者对应 TensorFlow 的 SavedModel 计算图。选择哪个取决于你的模型来源和需求。LayersModel 的好处是可以在浏览器里继续训练API 也更直观。缺点是转换时对自定义层支持有限某些 Keras 自定义层转不过去。GraphModel 支持更广的算子推理性能通常也更好但不能在浏览器里训练只能推理。我的建议是纯推理场景优先用 GraphModel尤其是从 TensorFlow SavedModel 转过来的。如果你需要在浏览器里做迁移学习或者微调那就用 LayersModel。我这次的项目是纯推理所以选了 GraphModel转换过程顺畅很多。3.3 转换后必做的算子兼容性检查转换成功不代表能跑。有些算子在 Python 侧有在 TensorFlow.js 里没有实现或者只在特定后端有。转换工具会给出警告但很多人不看日志直接上线结果运行时才报错。我的做法是转换后立刻写一个最小测试脚本用随机输入跑一遍model.predict看有没有报错。同时打开tf.enableDebugMode()它会打印每个算子的执行情况包括哪些走了 CPU 回退。tf.enableDebugMode(); const model await tf.loadGraphModel(web_model/model.json); const dummy tf.randomNormal([1, 224, 224, 3]); const out model.predict(dummy); out.print();如果某个算子不支持控制台会明确告诉你算子名字。这时候有两条路一是换一个等价的算子组合二是自己用 TensorFlow.js 的底层 API 实现一个。前者更省事后者工作量大但灵活。我遇到过一个自定义激活函数不支持最后用几个基础算子拼出来替代效果基本一致。4. Web Worker 与主线程的协作模式4.1 为什么推理必须离开主线程JavaScript 是单线程的主线程负责渲染、事件响应、脚本执行。如果你在主线程里跑推理尤其是大模型页面会直接卡死用户点什么都没反应。哪怕推理只要几百毫秒用户也能明显感觉到卡顿。Web Worker 是浏览器提供的后台线程可以独立执行脚本不阻塞主线程。把推理放进 Worker主线程只管 UI 和通信页面始终保持流畅。这是端侧推理的标准做法几乎没有例外。但 Worker 也有代价它和主线程之间只能通过消息传递通信数据要序列化。对于大张量序列化和反序列化的开销不能忽略。所以怎么传数据、传什么格式是有讲究的。4.2 张量数据的传递与零拷贝默认情况下postMessage传的是结构化克隆大数组会被完整复制一份开销很大。对于图像数据这种几十 MB 的输入复制一次就是几十毫秒。解决办法是用Transferable Objects把 ArrayBuffer 的所有权转移过去避免复制。// 主线程 const imageData ctx.getImageData(0, 0, w, h); const buffer imageData.data.buffer; worker.postMessage({ type: infer, buffer, width: w, height: h }, [buffer]); // Worker 内 self.onmessage async (e) { const { buffer, width, height } e.data; const data new Uint8ClampedArray(buffer); const tensor tf.tensor3d(data, [height, width, 3], int32); const result model.predict(tensor); const output await result.data(); self.postMessage({ type: result, output }); };注意postMessage的第二个参数[buffer]它声明这个 buffer 是转移而非复制。转移之后主线程里的原 buffer 会被清空不能再访问。这个细节要小心如果你转移后还想用原数据就得先复制一份。提示转移 ArrayBuffer 是零拷贝的但张量本身在 Worker 里还是要创建。如果输入是图像建议在 Worker 里直接从 buffer 构造张量避免在主线程先转成张量再传。4.3 多 Worker 并行与任务队列单个 Worker 只能串行处理任务。如果推理请求很密集比如视频逐帧处理一个 Worker 会成为瓶颈。这时候可以起多个 Worker组成一个任务队列谁空闲谁接活。我的实现是一个简单的 Worker 池启动 N 个 WorkerN 通常取navigator.hardwareConcurrency的一半左右留出余量给主线程维护一个任务队列每个 Worker 完成一个任务后从队列取下一个。class WorkerPool { constructor(size, workerUrl) { this.workers []; this.idle []; this.queue []; for (let i 0; i size; i) { const w new Worker(workerUrl); w.onmessage (e) this.onDone(w, e.data); this.workers.push(w); this.idle.push(w); } } run(task) { return new Promise((resolve) { const job { task, resolve }; const w this.idle.pop(); if (w) this.dispatch(w, job); else this.queue.push(job); }); } dispatch(w, job) { w._job job; w.postMessage(job.task); } onDone(w, data) { w._job.resolve(data); const next this.queue.shift(); if (next) this.dispatch(w, next); else this.idle.push(w); } }这个池子要注意 Worker 数量不能太多。每个 Worker 都会加载一份模型内存占用是叠加的。我试过开 8 个 Worker结果内存直接爆了页面崩溃。后来改成 2 到 4 个稳定很多。具体数量要看你模型大小和设备内存。5. 性能优化的几个关键抓手5.1 用 tf.tidy 和 dispose 管住内存TensorFlow.js 的张量是手动管理内存的不像 JavaScript 的普通对象有垃圾回收。你每创建一个张量它就占用一块显存或内存不主动释放就会泄漏。跑几百次推理之后页面就会因为内存耗尽而崩溃。tf.tidy是官方推荐的方案它包裹一个函数函数执行完自动释放函数内创建的所有中间张量只保留返回值。const result tf.tidy(() { const input tf.tensor3d(data, [h, w, 3]); const normalized input.div(255).sub(0.5).mul(2); return model.predict(normalized); }); // 这里 normalized 和 input 已经被释放result 还在要注意model.predict返回的张量不会被tidy释放因为它是返回值。用完这个 result 之后记得手动result.dispose()。我一开始以为tidy会管所有东西结果发现输出张量一直累积跑久了还是崩。判断有没有泄漏可以用tf.memory()看当前张量数量和字节数。如果每次推理后这个数字都往上涨那就是有泄漏。5.2 批处理与输入尺寸的权衡批处理能提高 GPU 利用率一次推理多个样本比逐个推理快。但端侧场景下用户通常一次只处理一个输入批处理用不上。不过如果你的场景是批量处理比如一次上传多张图片那就可以攒一批一起推理。输入尺寸是另一个关键。模型训练时的输入尺寸是固定的推理时如果传更大的图要么被缩放要么报错。缩放会损失精度但能保证速度。我的做法是在 Worker 里先把输入 resize 到模型期望的尺寸用 canvas 的drawImage做高质量缩放再转成张量。尺寸的选择要平衡精度和速度。224x224 是很多视觉模型的标配推理快精度够用。如果你需要更高精度比如医学图像那可能要用 512x512 甚至更大但推理时间会成倍增加。实测下来输入尺寸翻倍推理时间大约翻三到四倍因为计算量是尺寸的平方关系。5.3 首次加载与模型缓存模型文件通常有几 MB 到几十 MB首次加载要下载用户要等。优化手段有两个一是用 CDN 加速分发二是用浏览器缓存避免重复下载。TensorFlow.js 的loadGraphModel支持从 IndexedDB 缓存加载。你可以用tf.io的模型保存和加载 API把模型存到 IndexedDB下次直接从本地读省掉网络请求。async function loadModelWithCache(url) { const cacheKey indexeddb://my-model; try { return await tf.loadGraphModel(cacheKey); } catch { const model await tf.loadGraphModel(url); await model.save(cacheKey); return model; } }这个方案要注意缓存失效。模型更新后旧缓存还在用户拿到的还是老模型。解决办法是在缓存 key 里带上版本号比如indexeddb://my-model-v2版本变了就重新下载。注意IndexedDB 有存储配额限制不同浏览器不一样。模型太大可能存不下要做好降级处理存不下就直接走网络加载。6. 踩坑实录那些让我排查半天的诡异问题6.1 第一次推理总是失败或结果异常这个问题我遇到过两次表现是页面加载后第一次predict报错或者输出全是 NaN第二次就正常了。原因有两个可能一是没等tf.ready()后端还没初始化完二是模型权重还没加载完就调用了推理。排查方法是加日志确认tf.ready()的 Promise 已经 resolve模型加载的 Promise 也 resolve 了再触发推理。我后来把初始化和推理做成两个明确的阶段UI 上显示模型加载中加载完才允许用户操作问题就没了。6.2 WebGL 后端下的精度陷阱有一次做回归任务预测一个连续值在 WebGPU 上结果正常切到 WebGL 后数值偏差很大。查了半天发现是 WebGL 的浮点纹理精度问题某些设备只支持 16 位浮点累加误差被放大。解决办法是给模型加一个后处理把输出限制在合理范围内或者干脆在检测到 WebGL 后端时切换到 WASM。虽然 WASM 慢一点但精度稳定。这个取舍要看业务精度优先就牺牲速度速度优先就接受误差。6.3 Worker 里加载模型失败的隐蔽原因在 Worker 里加载模型路径问题和主线程不一样。主线程里相对路径是相对于页面 URLWorker 里是相对于 Worker 脚本的 URL。我一开始用相对路径主线程能加载Worker 里就 404。解决办法是用绝对路径或者用import.meta.url动态计算。另外Worker 里加载 WASM 后端时setWasmPaths的路径也要重新设置因为 Worker 的环境和主线程隔离。这个坑很隐蔽报错信息也不明确容易卡住。6.4 内存泄漏的渐进式排查内存泄漏是最难查的因为它是渐进的跑一会儿才崩。我的排查流程是用tf.memory()打印每次推理前后的张量数量看是否持续增长。如果增长检查所有predict的返回值有没有dispose。检查tidy里有没有创建了不该保留的张量。检查事件监听、定时器有没有在组件卸载时清理。有一次发现是requestAnimationFrame循环里每帧都创建张量但没释放跑几分钟就崩了。把张量创建移出循环或者每帧用tidy包裹问题解决。7. 端侧推理的边界与后续演进方向7.1 当前方案的性能天花板说实话端侧推理的性能天花板还是明显的。用户设备的 GPU 和服务器没法比一个在服务器上几十毫秒的模型在手机上可能要几百毫秒甚至更久。所以端侧适合的是中小模型 高频调用的组合大模型还是得靠服务端。另外端侧的算力是共享的用户可能同时开着视频、游戏你的推理只能分到一部分资源。要做好性能波动的心理准备UI 上给用户明确的反馈别让用户以为卡死了。7.2 模型更新与版本管理端侧模型更新是个麻烦事。模型缓存在用户设备上你发了新版本用户不一定能及时拿到。我的做法是版本号写进缓存 key同时在应用启动时异步检查新版本有更新就后台下载下次启动生效。这样用户不会因为下载模型而等待。版本管理还要考虑兼容性。新模型可能依赖新的算子老浏览器不支持。所以要么保证新模型向后兼容要么在检测到不支持时回退到旧模型。7.3 什么情况下该考虑混合方案纯端侧不是唯一选择。有些场景适合混合简单输入走端侧复杂输入走服务端或者端侧先做粗筛服务端做精算。这样既保证了响应速度又保证了精度。判断标准是如果端侧能覆盖 80% 的常见情况剩下 20% 的疑难杂症交给服务端整体体验和成本都是最优的。我下一个项目就打算这么做端侧跑一个轻量模型做实时反馈用户确认后再调服务端跑大模型出最终结果。这套东西跑下来我最大的体会是端侧推理的难点不在模型本身而在工程细节。模型转换、后端选型、线程调度、内存管理每一环都有坑但每一环也都有成熟的解法。把这几块理顺了TensorFlow.js 跑在用户设备上这件事其实比想象中稳。