机器学习模型训练完只是第一步真正难的是让它在用户的手机、平板、笔记本上跑起来——不依赖服务器、不传数据、断网也能用。TensorFlow.js 就是干这个的把模型直接塞进浏览器用 JavaScript 调用 GPU 做推理。我最近拿它做了几个端侧推理的小项目从模型转换到 Web Worker 多线程调度踩了一圈坑这篇就把整个链路拆开讲清楚包括 WebGPU 后端怎么开、内存怎么管、主线程怎么不卡。适合已经会点前端、想把手里的模型搬到浏览器里跑的人也适合做机器学习但不想折腾服务端部署的开发者。1. 为什么要把推理放到用户设备上1.1 端侧推理解决的三个现实问题先说清楚动机不然很容易做成为了用而用。把模型放到浏览器里跑最直接的好处是数据不出设备。比如做一个本地图片分类或者文本情感分析的工具用户的照片、输入的文字全程在浏览器内存里处理压根不经过网络请求隐私这块天然就干净。这一点在涉及个人内容的场景里特别关键你不需要向用户解释我们不会上传你的数据因为代码本身就证明了这一点。第二个好处是延迟。走服务端推理一次请求要经历序列化、网络传输、排队、反序列化哪怕服务器就在同城往返也得几十毫秒起步遇到网络抖动直接上百毫秒。端侧推理省掉了整个网络环节输入到输出就是一次前向传播的时间。对于交互式的场景——比如实时滤镜、手写识别、姿态检测——这个差距是能明显感知到的。第三个好处是成本。服务端推理是要烧 GPU 的用户量一上来推理成本线性增长。端侧推理把算力成本转移到了用户设备上你的服务器只需要托管静态资源一个 CDN 就能扛住。当然这不是说端侧能完全替代服务端大模型、需要频繁更新的模型还是得放服务端但那些轻量级的、固定的推理任务放端侧是划算的。1.2 什么样的模型适合搬到浏览器不是所有模型都适合端侧。我的判断标准有三条模型体积、计算量、更新频率。模型体积直接决定加载时间。一个 50MB 的模型用户首次打开页面要下载半天体验直接崩掉。我的经验是经过量化压缩后控制在 5MB 以内的模型端侧体验才比较舒服。超过 20MB 的除非你能接受较长的首次加载否则还是老老实实放服务端。计算量决定了推理速度。浏览器能拿到的算力有限尤其是移动端。一个在服务器 GPU 上跑 10ms 的模型在手机浏览器上可能要 200ms 甚至更久。所以端侧模型通常是小而专的——专门做一件事参数量控制在百万级到千万级。更新频率也很重要。如果你的模型每周都要重新训练、更新权重那每次更新都要让用户重新下载模型文件这个成本不低。反过来那些训练一次就能稳定用很久的模型比如通用的图像分类、人脸检测、关键词识别就非常适合端侧。1.3 TensorFlow.js 在端侧推理里的位置TensorFlow.js 不是一个独立的机器学习框架它是 TensorFlow 生态在 JavaScript 侧的延伸。它的核心价值在于统一的 API 抽象同一套代码底层可以跑在 WebGL、WebGPU、WASM 甚至纯 CPU 上你不需要为每个后端写不同的实现。它主要分两块能力一是推理加载已经训练好的模型TF SavedModel、Keras、TF Hub 或者 TF.js 自己的格式在浏览器里做前向传播二是训练可以在浏览器里做迁移学习、微调甚至从零训练小模型。实际项目里用得最多的是推理训练更多是锦上添花。和 ONNX Runtime Web 相比TensorFlow.js 的生态更完整模型转换工具链更成熟文档也更全。ONNX Runtime Web 的优势在于对 ONNX 格式的支持更原生如果你手里的模型本来就是 ONNX 的用它可能更省事。但如果你的模型来自 TensorFlow/Keras 体系TensorFlow.js 是更自然的选择。2. 模型转换从训练产物到浏览器可加载格式2.1 转换工具链的选择与安装TensorFlow.js 不能直接加载.h5或者 SavedModel需要先转成它自己的格式。转换工具是tensorflowjs_converter通过 pip 安装pip install tensorflowjs装完之后会得到一个命令行工具。这里有个坑tensorflowjs的版本要和你的 TensorFlow 版本大致匹配否则转换时可能报一些莫名其妙的算子不支持错误。我一般会先确认环境里的 TF 版本再装对应版本的转换器。转换的目标格式有两种Layers 格式和Graph 格式。Layers 格式对应 Keras 模型转换后是一个model.json加一组.bin权重文件Graph 格式对应 TF 的 SavedModel结构类似但加载方式不同。现在新项目基本都用 Layers 格式因为 Keras 是主流而且 Layers 格式对自定义层的支持更好。2.2 转换命令与关键参数假设你有一个 Keras 模型model.h5转换命令是tensorflowjs_converter \ --input_formatkeras \ --output_formattfjs_layers_model \ --weight_shard_size_bytes4194304 \ model.h5 \ ./web_model这里几个参数值得说。--weight_shard_size_bytes控制权重文件的分片大小默认是 4MB 左右。为什么要分片因为浏览器加载大文件时单个文件太大会导致加载卡顿分片之后可以并行下载而且配合 HTTP/2 的多路复用效果更好。我一般保持默认的 4MB除非模型特别小那就没必要分片。还有一个参数是--quantize_float16或者--quantize_uint8用来做量化。量化能把模型体积压到原来的 1/2 到 1/4代价是精度会掉一点。实测下来float16 量化对大多数分类模型的影响很小精度掉个 0.5% 以内但体积直接减半非常划算。uint8 量化更激进适合对精度不那么敏感的场景。tensorflowjs_converter \ --input_formatkeras \ --output_formattfjs_layers_model \ --quantize_float16 \ model.h5 \ ./web_model转换完成后web_model目录里会有model.json和若干.bin文件。model.json里描述了模型的拓扑结构和权重文件的位置.bin就是实际的权重数据。2.3 转换后必须验证的三件事转换完不要直接扔到前端用先在 Node 环境里验证一遍。TensorFlow.js 有 Node 版本可以加载模型跑一次推理确认输出正常。第一件事是检查输入输出张量的形状。用model.inputs和model.outputs打印出来确认和你训练时一致。我遇到过一次转换后输入形状从[null, 224, 224, 3]变成了[null, 224, 224]原因是模型定义时用了Flatten之后又接了全连接转换器把通道维度搞丢了。这种问题在浏览器里报错很隐晦提前在 Node 里发现能省很多时间。第二件事是对比转换前后的输出。拿同一份输入数据分别在原始 Keras 模型和转换后的 TF.js 模型上跑一遍看输出的差异。如果差异在 1e-5 量级说明转换没问题如果差异很大那多半是某个算子转换时出了问题。第三件事是确认模型体积。转换后的总体积所有.bin加起来要符合预期。如果量化后体积没怎么变可能是量化参数没生效需要检查命令。3. 浏览器端加载与推理的完整链路3.1 后端选择WebGL、WebGPU 还是 WASMTensorFlow.js 支持多种后端选哪个直接决定推理速度。我整理了一个对比后端适用场景优势局限WebGL兼容性优先支持面最广老设备也能跑算子覆盖不全部分模型跑不了WebGPU性能优先速度快支持新算子需要较新的浏览器WASMCPU 兜底兼容性最好算子全速度慢适合小模型CPU调试用无需额外依赖最慢仅用于验证实际项目里我的策略是优先 WebGPU回退 WebGL最后 WASM。TensorFlow.js 提供了自动选择后端的机制但自动选择不一定最优我一般会手动检测import * as tf from tensorflow/tfjs; async function selectBackend() { if (await tf.setBackend(webgpu)) { await tf.ready(); return webgpu; } if (await tf.setBackend(webgl)) { await tf.ready(); return webgl; } await tf.setBackend(wasm); await tf.ready(); return wasm; }WebGPU 的开启需要浏览器支持目前主流的新版本浏览器都已经默认开启。它的优势在于能更充分地利用 GPU 的并行计算能力尤其是卷积类操作速度比 WebGL 快不少。但要注意WebGPU 后端对算子的支持还在完善中如果你的模型里有比较冷门的算子可能会回退到 CPU 执行反而更慢。3.2 模型加载的时机与进度反馈模型加载是异步的而且可能比较慢所以加载时机和进度反馈都要设计好。加载时机上不要在页面一打开就加载模型。用户可能只是来看看根本没打算用推理功能。我的做法是懒加载等用户触发了需要推理的操作比如点击开始识别按钮再开始加载模型。这样首屏体验不受影响。进度反馈上tf.loadLayersModel支持onProgress回调const model await tf.loadLayersModel(/models/web_model/model.json, { onProgress: (fraction) { console.log(加载进度: ${(fraction * 100).toFixed(1)}%); updateProgressBar(fraction); } });这个回调会随着权重文件的分片下载不断触发可以拿来做进度条。注意fraction是 0 到 1 之间的小数不是百分比要自己乘 100。还有一个细节模型加载完成后第一次推理会特别慢。因为第一次推理时TensorFlow.js 要把权重上传到 GPU、编译着色器WebGL 后端或者初始化计算管线WebGPU 后端。这个开销可能有好几百毫秒。解决办法是加载完成后立刻跑一次预热推理用全零或者随机的输入跑一遍把编译开销提前消化掉。用户感知到的第一次真实推理就快了。// 预热 const warmupInput tf.zeros([1, 224, 224, 3]); const warmupOutput model.predict(warmupInput); warmupOutput.dispose(); warmupInput.dispose();3.3 输入数据的预处理与张量构造浏览器里的输入通常是ImageData、HTMLImageElement、HTMLVideoElement或者HTMLCanvasElement需要转成张量。TensorFlow.js 提供了tf.browser.fromPixels来做这件事const tensor tf.browser.fromPixels(imageElement); // 形状是 [height, width, 3]值域 0-255但模型通常需要特定的输入形状和值域。比如一个 224x224 的模型需要先 resize再归一化const resized tf.image.resizeBilinear(tensor, [224, 224]); const normalized resized.div(255.0); // 归一化到 0-1 const batched normalized.expandDims(0); // 增加 batch 维度这里每一步都会创建新的张量旧的张量必须手动释放否则内存会持续增长。TensorFlow.js 不像 Python 有垃圾回收能自动管理 GPU 内存它需要显式调用dispose()。我一般用tf.tidy来包裹const input tf.tidy(() { const tensor tf.browser.fromPixels(imageElement); const resized tf.image.resizeBilinear(tensor, [224, 224]); const normalized resized.div(255.0); return normalized.expandDims(0); });tf.tidy会自动释放内部创建的所有中间张量只保留返回值。这个习惯一定要养成不然跑几十次推理之后页面就卡死了。3.4 推理执行与结果解析推理本身很简单一行代码const output model.predict(input);但output是一个张量需要转成 JavaScript 能用的数据。如果是分类任务通常要取 argmax 和对应的概率const probabilities await output.data(); const maxIndex probabilities.indexOf(Math.max(...probabilities)); const confidence probabilities[maxIndex];output.data()是异步的因为它要把 GPU 上的数据读回 CPU。这个读回操作是有开销的频繁调用会拖慢速度。如果只是要 argmax可以用output.argMax(1).data()直接拿到索引减少数据传输量。用完的input和output都要dispose()。如果放在tf.tidy里input会被自动释放但output是返回值需要手动释放。4. 让主线程不卡Web Worker 与内存管理4.1 为什么推理必须放到 Web WorkerJavaScript 是单线程的推理是计算密集型任务如果放在主线程页面会直接卡住——按钮点不动、动画停摆、滚动卡顿。用户会以为页面崩了。Web Worker 是浏览器提供的多线程机制可以把推理放到后台线程执行主线程继续负责 UI 渲染和交互。TensorFlow.js 在 Worker 里能正常工作而且 WebGL 和 WebGPU 后端在 Worker 里也能用OffscreenCanvas 的支持下。Worker 的通信通过postMessage和onmessage完成。主线程把图像数据传给 WorkerWorker 推理完把结果传回来。注意传图像数据时用Transferable Objects可以避免拷贝开销// 主线程 const imageData ctx.getImageData(0, 0, width, height); worker.postMessage({ type: predict, data: imageData }, [imageData.data.buffer]);第二个参数是 transfer 列表把imageData.data.buffer转移过去主线程这边就失效了但省掉了一次内存拷贝。对于大图像这个优化很明显。4.2 Worker 里的模型加载与消息协议Worker 里加载模型和主线程一样但要注意 Worker 里没有 DOM不能用document相关 API。模型加载完成后给主线程发个消息通知// worker.js importScripts(https://cdn.jsdelivr.net/npm/tensorflow/tfjs4.x/dist/tf.min.js); let model null; self.onmessage async (event) { const { type, data } event.data; if (type load) { model await tf.loadLayersModel(data.modelUrl); // 预热 const warmup tf.zeros([1, 224, 224, 3]); model.predict(warmup).dispose(); warmup.dispose(); self.postMessage({ type: loaded }); } if (type predict) { const input tf.tidy(() { const tensor tf.browser.fromPixels(data); const resized tf.image.resizeBilinear(tensor, [224, 224]); return resized.div(255.0).expandDims(0); }); const output model.predict(input); const probabilities await output.data(); input.dispose(); output.dispose(); self.postMessage({ type: result, data: Array.from(probabilities) }); } };消息协议要设计得清晰用type字段区分不同的消息类型。我一般会定义load、predict、result、error这几种类型方便扩展。4.3 内存泄漏的排查与张量生命周期管理TensorFlow.js 的内存泄漏是最容易踩的坑。表现是页面跑一段时间后越来越卡最后崩溃。原因是张量没释放GPU 内存持续增长。排查方法是定期打印内存使用情况console.log(tf.memory()); // { numTensors: 42, numDataBuffers: 42, numBytes: 12345678, ... }numTensors是当前存活的张量数量。如果这个数字随着推理次数不断增加说明有泄漏。正常情况下每次推理结束后numTensors应该回到一个稳定的基线。管理张量生命周期的原则有三条能 tidy 就 tidy、返回值手动 dispose、循环里尤其注意。循环里最容易漏因为每次迭代都创建张量如果不在迭代末尾释放很快就爆了。for (const image of images) { const result tf.tidy(() { const input preprocess(image); return model.predict(input); }); // result 是 tidy 的返回值需要手动释放 const data await result.data(); result.dispose(); // 处理 data }还有一个隐蔽的坑tf.browser.fromPixels在传入HTMLVideoElement时如果视频在播放每次调用都会创建新张量。做实时视频推理时一定要确保每帧的张量都被释放。5. 实测中的性能数据与优化取舍5.1 不同后端的推理耗时对比我拿一个 MobileNetV2 的变体约 3.5MBfloat16 量化后 1.8MB在几台设备上做了测试输入 224x224batch size 为 1结果如下设备后端首次推理稳定推理桌面 ChromeWebGPU320ms12ms桌面 ChromeWebGL480ms25ms桌面 ChromeWASM900ms85ms中端手机WebGPU650ms45ms中端手机WebGL820ms78ms中端手机WASM1500ms220ms数据很直观WebGPU 在稳定推理阶段优势明显桌面端能到 12ms手机端 45ms都算流畅。WebGL 慢一倍左右但兼容性好。WASM 只适合兜底或者模型特别小的情况。首次推理的耗时主要是编译和初始化开销WebGPU 的首次开销比 WebGL 还大一些但稳定后反超。所以预热这一步不能省。5.2 模型量化对精度和速度的实际影响量化是端侧推理的常规操作但影响要实测。我拿一个图像分类模型做了对比量化方式体积精度Top-1推理耗时无量化14MB76.2%25msfloat167MB76.0%24msuint83.5MB74.8%22msfloat16 几乎无损体积减半速度基本不变是性价比最高的选择。uint8 体积更小但精度掉了 1.4 个百分点速度提升也不明显。所以我的默认策略是float16 量化只有在体积实在压不下来时才考虑 uint8。需要注意的是量化对不同类型的模型影响不一样。分类模型对量化比较鲁棒但检测模型、分割模型对量化更敏感尤其是回归输出的部分。做检测任务时我一般不做 uint8 量化float16 是底线。5.3 批处理与流式推理的取舍端侧推理通常是单张输入因为用户一次只处理一张图或者一帧视频。但有些场景可以批处理比如一次上传多张图片做分类。批处理的好处是能摊薄每次推理的固定开销吞吐量更高。但坏处是内存占用成倍增长而且延迟变高——要等所有输入都准备好才能开始推理。我的经验是交互式场景用单张后台批处理场景用 batch。如果做批处理batch size 控制在 4 到 8 之间比较稳妥再大就容易触发内存问题。流式推理比如视频逐帧处理要注意丢帧策略。如果推理速度跟不上视频帧率不能傻等要主动丢帧。我的做法是维护一个是否正在推理的标志如果上一帧还没处理完就跳过当前帧let isProcessing false; async function processFrame(videoElement) { if (isProcessing) return; isProcessing true; try { const result await runInference(videoElement); renderResult(result); } finally { isProcessing false; } }这样能保证 UI 不堆积任务虽然会丢帧但整体流畅度更好。6. 几个容易翻车的细节6.1 跨域与模型文件托管模型文件通过fetch加载所以必须配置正确的 CORS 头。如果模型放在 CDN 上CDN 要允许你的域名跨域访问。这个坑很常见本地开发时用localhost没问题一部署到线上就报 CORS 错误。解决办法是在模型文件的响应头里加上Access-Control-Allow-Origin。如果用对象存储一般都有配置项如果用 Nginx加一行add_header就行。还有一个细节是模型文件的 MIME 类型。.bin文件最好返回application/octet-streammodel.json返回application/json。有些服务器对未知扩展名会返回text/plain虽然大多数情况下不影响但个别环境会出问题。6.2 移动端浏览器的内存限制移动端浏览器对内存的限制比桌面严格得多。iOS Safari 尤其敏感单个标签页的内存超过一定阈值大概几百 MB就会被系统杀掉。所以移动端的模型体积和中间张量都要控制。我的做法是移动端只加载量化后的模型推理时用tf.tidy严格管理内存避免同时存在多个大张量。如果模型实在太大考虑分片加载或者用更小的模型架构。另外移动端的 WebGL 有纹理尺寸限制通常是 4096x4096 或者 8192x8192。如果中间张量的某个维度超过这个限制会直接报错。做高分辨率图像处理时要注意必要时先降采样。6.3 模型版本更新与缓存策略模型文件是静态资源浏览器会缓存。但模型更新后如果缓存没失效用户用的还是旧模型。解决办法是给模型文件加版本号或者哈希比如model.json?v2或者model-abc123.json。这样更新时 URL 变了浏览器会重新下载。model.json里引用了权重文件的路径如果权重文件也带哈希转换时要确保路径正确。TensorFlow.js 的转换器默认用相对路径部署时保持目录结构一致就行。还有一个策略是用 Service Worker 做离线缓存。把模型文件缓存到本地用户第二次打开时直接从缓存读取加载速度飞快。但要注意缓存的更新策略模型更新时要能触发重新下载。6.4 推理结果的置信度处理端侧模型的精度通常不如服务端大模型所以置信度阈值要设得合理。分类任务里如果最高概率低于某个阈值比如 0.6应该返回无法确定而不是硬给一个结果。这个阈值需要根据实际场景调。安防类的场景宁可漏报不可误报阈值设高推荐类的场景可以宽松一些。我一般会先在验证集上跑一遍看不同阈值下的准确率和召回率再决定。还有一个技巧是多帧投票。视频推理时连续几帧的结果做投票能显著降低误判。比如连续 5 帧里有 3 帧以上是同一类别才认为结果有效。这个在姿态检测、手势识别里特别有用。7. 端侧推理的边界与后续扩展TensorFlow.js 能做的事比很多人想象的多但边界也很清楚。它适合中小型模型、固定任务、对隐私和延迟敏感的场景。大语言模型、需要频繁更新的推荐模型、超大规模视觉模型还是得靠服务端。如果要把端侧推理做得更完整有几个方向可以扩展。一是模型热更新通过版本化 URL 和 Service Worker 实现无感更新二是多模型协同比如一个轻量模型做初筛命中后再调用稍大的模型做精细判断三是联邦学习的思路端侧做推理的同时收集梯度定期回传聚合但这个涉及的东西比较多落地要谨慎。我在实际项目里最大的体会是端侧推理的瓶颈往往不在模型本身而在工程细节。内存管理、线程调度、加载策略、缓存更新这些看起来不起眼的地方才是决定用户体验的关键。模型转换和推理调用可能半天就搞定了但把这些细节打磨好花的时间是前者的好几倍。所以如果你打算做端侧推理心理预期要放对——算法只是入场券工程才是主战场。