说实话,刚看到“TensorFlow.js”这个项目名的时候,我第一反应是回想起两年前接的一个需求:客户要在网页里做一个快递单 OCR 识别,但数据不能出内网,也没有云服务器预算,更离谱的是使用场景是展厅大屏,观众点开页面就要能识别。传统的“Python 训练模型 + 后端 API + 前端展示”这条路直接被堵死。最后把目光放在了 TensorFlow.js 上——把机器学习模型直接塞进浏览器,训练和推理全部在本地完成。前后端零服务,数据不出页面,刷新即重置。这个项目最终顺利交付,我也因此把“在浏览器里跑机器学习”这件事从实验性质彻底变成了我的常规技术方案之一。
如果你也是前端开发者、刚接触机器学习的程序员,或者单纯想给网页加一点“智能”却不想折腾 Python 环境,这篇文章很适合照着做。它不是什么学术大课,而是一个完整项目的拆解复盘:从选型思路、底层原理,到可复现的实操代码、调优技巧和踩坑实录,全部按我实际做过的路径记录下来。
1. 项目思路拆解:为什么偏偏是浏览器里的机器学习
1.1 需求倒逼选型:不是炫技,是交付要求
很多人在聊 TensorFlow.js 时容易陷入“为了在浏览器跑 AI 而在浏览器跑 AI”的误区。我做选型时从来不看技术新不新,只看交付条件。当时客户提的三个硬指标是这样的:
- 数据不能出内网,所有图片和识别结果必须留在浏览器本地;
- 不能部署额外服务,产品是一个 U 盘拷贝就能跑的离线网页包;
- 识别过程要可交互,识别结果要实时反馈到页面上。
这三条直接排除了 Flask/FastAPI 后端方案,也排除了 Node.js 中间层方案。你总不能给展厅大屏配一台 GPU 服务器,更不能让客户的快递单图片先飞到某个云厂商服务器再回来。TensorFlow.js 几乎是唯一的技术选项:模型文件打包进静态资源,前端直接加载,推理全部走浏览器本地 WebGL,图片数据从头到尾不经过任何服务器。换个角度看,它等价于把传统 AI 产品的“模型服务层”整个下沉到了客户端。
后来我复盘这种模式,觉得它本质上和“把数据库从服务器搬到 SQLite”是同一类思路:当你的用户量大不起来、隐私要求又高、部署环境又不可控时,把重计算放在终端反而是最稳的架构。TensorFlow.js 做的就是把模型的推理甚至训练能力全部打包成一个 JS 包,前端开发者不需要懂 Python、不需要懂 CUDA,也不需要考虑高并发,只要懂 JavaScript 就能让页面“有智能”。
1.2 浏览器做机器学习的边界在哪里
当然,“搬进浏览器”不等于“所有机器学习都能在浏览器里跑”。我在项目完成后给团队做内部分享时,第一张 PPT 就写着:推理为主、小规模训练为辅、原型验证为补充。这是我对浏览器端机器学习的基本定位。
浏览器能做好推理这件事。像 MobileNet、EfficientNet 这类为端侧设计的模型,参数在 400 万到 700 万这个量级,经过转换后在浏览器里单次推理基本在几十到一百多毫秒。如果模型设计得当,甚至能做到实时视频流的人像分割,这个我在后面实操部分会展示。
浏览器也能做轻量训练和迁移学习。像 MNIST、鸢尾花这类小数据集,或者给预训练模型只微调最后几层,WebGL 加速下跑起来并不算太慢。但要说在浏览器里从零训练一个 ResNet50,或者拿几万张图片做完整训练,那纯属自找麻烦。核心瓶颈有两个:一是数据加载和预处理都在前端受内存限制,二是 GPU 纹理大小有上限,大模型训练会导致内存爆炸。
所以做项目规划时,我的默认分工是这样的:重训练用 Python 在本地或服务器完成,训练好之后转成 tfjs 模型部署到前端;浏览器端只承担推理和对少量数据的在线微调。这个分工在我做过的几乎所有 TensorFlow.js 项目里都成立,也是最省时间的路线。
1.3 TensorFlow.js 生态全景:不只是“一个库”
很多人以为 TensorFlow.js 就是 npm 里一个叫@tensorflow/tfjs的包,其实它是一整套生态。了解这套生态,你在做技术选型时才能知道边界在哪。我列一下实际项目里用得上的几个模块:
@tensorflow/tfjs:核心库,包含 Layers API(类似 Keras)、底层 Ops API、以及 WebGL/WASM/WebGPU 后端管理;@tensorflow/tfjs-converter:负责把 Python 训练好的模型(Keras H5、SavedModel、TF Hub 模型)转成浏览器可加载的格式;@tensorflow/tfjs-data:提供数据管道,比如tf.data.generator,适合边读数据边训练的场景;@tensorflow/tfjs-vis:可视化训练过程的损失曲线和准确率,调试时极其好用;@tensorflow/tfjs-react-native:React Native 环境下的适配层,移动端场景会用到。
转换器是我特别想提醒重点关注的。它决定了你的模型能否从“Python 世界”无损地搬到“浏览器世界”。转换器把 Keras 的 H5 文件或 TensorFlow SavedModel 里的网络结构和权重导出成两个产物:一个model.json描述网络结构和权重分片信息,多个二进制.bin文件存权重。前端通过tf.loadLayersModel('.../model.json')就能恢复整个模型,依赖关系特别清晰,部署起来也不用处理乱七八糟的动态链接库。
2. 核心机制拆解:浏览器里跑 ML 的那套底牌
2.1 三块后端:WebGL、WASM 和 WebGPU 怎么选
先说结论:你写模型代码时不需要关心后端,但遇到性能问题、兼容性问题时必须懂后端。TensorFlow.js 的架构有点像数据库的存储引擎:上层是统一的张量操作接口,下层可切换不同的执行后端。目前主流是三个:
WebGL 是默认后端,也是最初让 TensorFlow.js 能在浏览器里跑深度学习的关键。它的原理是把张量映射成 GPU 纹理,把矩阵乘法、卷积这类算子编译成着色器程序(GLSL),再交给 GPU 并行执行。生活化一点理解:WebGL 相当于你在浏览器里租了一辆跑车,但不能自由改装,只能用指定的“着色器语言”这条赛道去跑,速度确实快,但你有多少赛道宽度(纹理大小)和油箱容量(显存)都被限制死了。
WASM 是 CPU 上运行的后备方案。它在浏览器里跑编译好的 C++ 代码,配合 SIMD 指令也能获得不错的性能。遇到设备不支持 WebGL、或者 WebGL 纹理限制把你的模型卡住时,WASM 兜底能力非常强,兼容性远好于 WebGL。代价是速度慢于 GPU,但比纯 JavaScript 的 CPU 实现快得多。
WebGPU 是近两年的新选择,可以把它理解成“新一代 WebGL 赛道”,能更高效地暴露 GPU 能力,训练性能比 WebGL 更强。不过兼容性目前还不算全面,我在正式项目里基本不会主动开启 WebGPU,只有用户明确在 Chrome 系最新版本上才考虑。
代码层面切换后端非常直接:
import * as tf from '@tensorflow/tfjs'; await tf.setBackend('webgl'); console.log(tf.getBackend()); // 'webgl'注意:
setBackend返回 Promise。有些人在setBackend('wasm')后立刻执行模型加载,结果后端还没就绪就报错。正确写法是await tf.setBackend('wasm')之后再加载模型。
2.2 张量与内存管理:入门最容易翻车的坑
在浏览器里写 TensorFlow.js,你会直接和“张量”(Tensor)打交道。可以把它想象成多维数组,但它实际驻留在 GPU 显存或 WASM 线性内存里。这里有个非常反直觉的地方:就算 JavaScript 有垃圾回收,浏览器的 GC 也没法自动回收 GPU 纹理内存。结果就是你不用的张量必须手动释放,否则内存会一路涨到浏览器崩溃。
我见过太多新手在循环里写predict,然后内存暴涨把 Tab 页卡死。因为每次predict都会返回新张量,中间计算结果也没释放。解法是学会两件事:tf.tidy和tf.dispose。
const result = tf.tidy(() => { const a = tf.tensor2d([[1, 2], [3, 4]]); const b = tf.tensor2d([[5, 6], [7, 8]]); return a.matMul(b); }); // 在这里 a 和 b 会被自动回收,result 作为返回值保留 result.dispose();tf.tidy有点像沙盒:它执行回调函数,函数内创建的、且没有被返回出去的张量都会在函数结束后自动释放。返回出来的张量需要你手动dispose。奖品是你没法在tidy外面使用那些被回收的中间张量,一旦用了就会报 Already disposed 错误。
如果你不确定当前内存状态,开发时打开控制台直接执行:
console.log(tf.memory());这个命令会返回类似{ numTensors: 128, numBytes: 51200, numBytesInGPU: 45056 }的结果。我每次排查内存泄漏都靠它:跑一段重复预测,然后盯住numBytesInGPU是不是只增不减。增速明显且永不复原,那就是有张量没释放。
另一个新手常踩的坑是dataSync。它有同步阻塞的特点,会直接卡死主线程,尤其在张量很大时。我当时识别快递单图片,一张高清图转成张量后执行dataSync(),整个页面直接白屏。后来改成await tensor.data()异步获取数据,UI 才保住流畅。
2.3 模型从哪来:三种途径按需选
用 TensorFlow.js 做项目的第一个决策就是模型来源。我总结下来有三种路径,各自适用场景完全不同。
第一种是纯前端从零训练。适合数据量小、模型结构简单的原型项目,比如手写数字识别、文本情感分类。优点是环境零依赖,缺点是只能跑小规模任务。
第二种是加载转换后的预训练模型。这是真实项目里最高频的做法。你用 Python 的 TensorFlow/Keras 训练好一个模型,然后通过tensorflowjs_converter转成 tfjs 格式,再放到前端加载。它兼顾了 Python 生态的成熟度和浏览器部署的便利性。我快递单项目里用的文本检测模型就是这么干的,Python 训练一周,浏览器推理只需 60 毫秒。
第三种是迁移学习。前端加载一个 MobileNet 这类通用特征提取器,冻结前面的卷积层,只训练最后几层全连接层来适配你的新类别。这在浏览器里也能实时训练,是我做个性化手势识别时最常用的一招,用户拍几张照片就能让模型学会新姿态。
模型转换命令通常长这样:
pip install tensorflowjs tensorflowjs_converter \ --input_format=keras \ --output_format=tfjs_model \ mnist_cnn.h5 \ ./tfjs_export转换成功后tfjs_export目录下会出现model.json和权重分片文件。前端部署时要把整个目录当静态资源上传到服务器或打包进页面,取model.json的完整 URL 传给tf.loadLayersModel就能加载。
这里有个隐藏痛点:tensorflowjs_converter并不能保证所有算子都能转换成功。如果你的模型里用了 TensorFlow 自定义算子或者某些不常用算子,转换阶段会直接报 Unsupported Ops。我踩过一次坑,当时一个模型里用了自定义的损失函数,Keras 里能跑,转换时直接失败。解决办法要么换用 tfjs 支持的算子重写,要么在模型里去掉自定义逻辑,这一步必须在训练阶段就考虑好。
3. 实操实录:用浏览器训练并识别手写数字
3.1 先搭一个能跑的最小项目骨架
为了把整条链路说清楚,我用“手写数字识别”做例子。这是机器学习界 Hello World,但它麻雀虽小五脏俱全:有数据输入、模型构建、训练、推理、展示,完整覆盖了 TensorFlow.js 项目的核心环节。
先搭页面骨架。我用最传统的单页结构,没有框架依赖,任何人复制下来都能跑:
<!DOCTYPE html> <html lang="zh-CN"> <head> <meta charset="UTF-8"> <title>浏览器手写数字识别</title> <script src="https://cdn.jsdelivr.net/npm/@tensorflow/tfjs@4/dist/tf.min.js"></script> <script src="https://cdn.jsdelivr.net/npm/@tensorflow/tfjs-vis@1/dist/tfvis.min.js"></script> </head> <body> <h3>在表格里手写一个数字,右侧返回识别结果</h3> <canvas id="draw" width="280" height="280" style="border:1px solid #aaa;background:#000;"></canvas> <button id="trainBtn">开始训练</button> <button id="predictBtn">识别</button> <div id="output">尚未识别</div> <script src="./main.js"></script> </body> </html>这里有个小细节:canvas 画布我设为 280x280,但模型输入是 28x28。实际推理时要把 canvas 的像素内容缩放 10 倍,从 280 降到 28。为什么不直接画 28x28?因为用户手写需要足够的画布空间,直接画到 28x28 的 canvas 上笔迹会粗糙得没法用。所以推理阶段做缩放是必然的。
3.2 方案 A:在浏览器里直接训练一个简化版卷积网络
既然是全浏览器方案,不妨直接在浏览器里训练一把。我用的 MNIST 数据集不是完整 6 万张,而是随机抽了 5000 张作训练集。为什么不用完整数据集?因为 6 万张 28x28 图像转换成浮点张量后占用内存接近 188MB,浏览器直接吃不住。5000 张约 15MB,配合 batchsize=64,既有代表性又不会卡死页面。
模型结构是经典 LeNet 的极简版:
const model = tf.sequential(); model.add(tf.layers.conv2d({ inputShape: [28, 28, 1], filters: 8, kernelSize: 3, activation: 'relu' })); model.add(tf.layers.maxPooling2d({ poolSize: [2, 2] })); model.add(tf.layers.conv2d({ filters: 16, kernelSize: 3, activation: 'relu' })); model.add(tf.layers.maxPooling2d({ poolSize: [2, 2] })); model.add(tf.layers.flatten()); model.add(tf.layers.dense({ units: 10, activation: 'softmax' })); model.compile({ optimizer: 'adam', loss: 'categoricalCrossentropy', metrics: ['accuracy'] });可以算一下这个模型参数量,非常轻量:第一层卷积 3x3x1x8 再加 8 个偏置,是 80 个参数;第二层卷积 3x3x8x16 再加 16 个偏置,是 1168 个参数;最后全连接是 7x7x16 个输入到 10 个输出,约 7850 个参数。总共不到一万个参数,在浏览器里训练压力很小。
接下来是数据准备。把 MNIST 的二进制文件放到本地静态目录,解析成两个 Float32Array:xs存像素值,ys存标签。然后抽取子集、归一化、做 one-hot 编码。
const X_SUBSET = 5000; const xs = tf.tensor4d(trainXs.slice(0, X_SUBSET), [X_SUBSET, 28, 28, 1]).div(255); const labels = tf.tensor1d(trainYs.slice(0, X_SUBSET), 'int32'); const ys = tf.oneHot(labels, 10).toFloat(); labels.dispose();为什么用.div(255)?因为图像原始像素是 0-255 整数,模型训练时用 0-1 范围浮点更稳定。如果不归一化,输入数值范围过大,梯度更新会剧烈震荡,模型很难收敛。很多第一次训练的人不归一化,盯着 loss 曲线看了半天死活不降,八成就是这个原因。
训练调用比想象的简单:
await model.fit(xs, ys, { batchSize: 64, epochs: 3, validationSplit: 0.1, callbacks: { onEpochEnd: (epoch, logs) => { console.log(`Epoch ${epoch + 1}: loss=${logs.loss.toFixed(4)}, acc=${logs.acc.toFixed(4)}`); } } });batchSize=64是平衡点。取值太小,GPU 利用率不足;取值太大,显存压力上升且训练波动变大。epochs=3是因为浏览器训练成本摆在那,3 轮之后准确率就能到 90% 上下,再往下收益边际递减,没必要为了刷精度在页面里干等。
训练结束后记得把xs和ys这两个张量dispose掉,否则它们会一直占着 GPU 纹理不释放:
xs.dispose(); ys.dispose();3.3 方案 B:加载 Python 训练好的模型做推理
真实项目里更多是方案 B:用 Python 训练,浏览器只承担推理。我前面提到用tensorflowjs_converter转换已经训练好的模型。假设你已经把 Keras 版的mnist_cnn.h5转成了tfjs_export目录,部署到静态资源路径/models/mnist/tfjs_export/下面,前端加载非常简单:
const model = await tf.loadLayersModel('/models/mnist/tfjs_export/model.json');注意model.json内部用了相对路径指向.bin权重文件。转换器默认生成group1-shard1of1.bin这种命名,前端加载后它会自动按相对路径去 fetch 权重分片。部署时千万不要单独把model.json挪走而漏了.bin,否则会卡在“找不到权重文件”的报错上。
推理阶段,要从 canvas 里拿到用户手写的像素,转换为模型输入格式:
const canvas = document.getElementById('draw'); const ctx = canvas.getContext('2d'); async function predictDigit() { // 将 canvas 内容转成 28x28 单通道张量 const tensor = tf.browser.fromPixels(canvas, 1) .resizeNearestNeighbor([28, 28]) .toFloat() .div(255) .expandDims(0); // 推理 + 找最大概率类别 const pred = model.predict(tensor); const result = pred.argMax(-1).dataSync()[0]; document.getElementById('output').innerText = `识别结果:${result}`; tensor.dispose(); pred.dispose(); }这段代码有几个关键点需要展开解释,不然你很容易写出一个“看起来对但结果全错”的版本。
第一,tf.browser.fromPixels(canvas, 1)把 canvas 像素读成维度为[height, width, 1]的张量。第二个参数1表示只取一个颜色通道。MNIST 本来就是灰度图,而黑色背景上的白色笔迹取红绿蓝哪个通道都行,但取单通道运算量小很多。
第二,resizeNearestNeighbor是最近邻缩放。280x280 到 28x28 缩小了 10 倍,选最近邻算法速度快,且有二进制分类场景下边缘保留比较好。如果希望识别结果更平滑,可以改成resizeBilinear,但速度慢一点。以实战经验说,手写数字识别用最近邻完全够用。
第三,expandDims(0)的目的是把[28, 28, 1]变成[1, 28, 28, 1],增加一个 batch 维度。模型的inputShape是[28, 28, 1],但predict方法要求传入带 batch 维度的四维张量。这可能是整个推理流程里最容易漏的一步,漏了之后会报跟 input shape 相关的错误。
第四,pred.argMax(-1)是在最后一维里找最大值索引。pred是 10 个概率组成的张量,argMax(-1)返回概率最大的那个数字,再dataSync()[0]把结果从张量里取出来。到这里,60 行代码就把一个可用的小型识别 Demo 跑通了。
3.4 这个方案还能怎么优化
方案 A 和方案 B 只是起点。我实际交付时还做了两件增强:第一是把模型加入 IndexedDB 缓存,第二次打开页面直接秒加载;第二是给识别结果做一个置信度显示,把pred.max().dataSync()[0]取出来显示百分比。置信度低于 0.7 时提示“请重新书写”,这样用户体验比傻乎乎返回一个错误结果要友好得多。
还有一个小细节:在方案 A 中训练结束后可以一键保存模型。调用model.save('indexeddb://handwritten-model'),模型就存到浏览器本地了。下次打开页面用tf.loadLayersModel('indexeddb://handwritten-model')直接加载,不用再重新训练。注意页面跨域名时 IndexedDB 不通用,所以这个方案只适合单机部署的 Demo 场景。
4. 性能优化与排查实录:卡顿、加载慢、内存泄漏三板斧
4.1 模型加载慢:拆分、压缩、离线缓存
模型加载慢是最常见的性能投诉。有一次模型文件接近 100MB,客户在普通办公网下打开页面转了十几秒才出结果,体验稀碎。那次我把问题拆成三个层面解决。
第一层是优化模型体积。先看能否用更小的模型结构,其次用量化压缩。tensorflowjs_converter支持--quantization_bytes=1或--quantization_bytes=2,意思是用 1 个字节或 2 个字节表示原来 4 字节的浮点权重。量化到 1 字节后模型体积能缩到原本的四分之一左右,代价是精度轻微下降,但这个下降在分类任务上通常可接受。命令示例:
tensorflowjs_converter \ --input_format=keras \ --output_format=tfjs_model \ --quantization_bytes=1 \ mnist_cnn.h5 \ ./tfjs_quantized第二层是 HTTP 层面打开缓存。确认静态服务器给model.json和.bin文件配置了Cache-Control: max-age=31536000。模型文件通常不会频繁变动,缓存一个月完全没问题,这意味着第二次访问时模型直接走浏览器缓存,加载接近瞬时。
第三层是用 IndexedDB 主动缓存。TensorFlow.js 自带indexeddb://这个 IOHandler,可以让模型在首次下载后被持久化到本地:
// 首次:从 HTTP 加载并复制到 IndexedDB await tf.loadLayersModel('/models/model.json'); await tf.io.copyModel('/models/model.json', 'indexeddb://my-model'); // 后续:直接从 IndexedDB 读取 const model = await tf.loadLayersModel('indexeddb://my-model');这个做法对离线网页包特别有效:首次打开花点时间下载模型,之后用户断网都能继续识别。
如果发现页面加载模型时控制台报 404,先别急着怀疑代码,去检查
model.json里的weightsManifest路径是否和实际部署目录结构一致。这方面出问题的概率比代码本身高得多。
4.2 训练时页面卡成幻灯片:把主线程让出来
在浏览器里训练模型时页面卡顿是另一类高频问题。尤其是训练循环里要做数据读取、张量创建、模型计算,这些活儿全部挤在主线程里跑,用户的页面自然掉帧。我曾经在一个 3 万条数据的小数据集上训练,没加任何保护,结果页面标签页直接变成“未响应”白色卡片,客户当场截图发过来我以为代码死循环了。
解决办法是训练大循环里主动让出主线程。TensorFlow.js 提供了一个后端相关的工具方法tf.nextFrame(),返回一个 Promise,可以把当前执行往后推一帧,让浏览器有机会去处理渲染和交互。
一个典型的训练循环可以写成:
const BATCH_SIZE = 64; for (let epoch = 0; epoch < 3; epoch++) { for (let i = 0; i < trainXs.length; i += BATCH_SIZE) { const xsBatch = tf.tensor4d(trainXs.slice(i, i + BATCH_SIZE), [BATCH_SIZE, 28, 28, 1]); const ysBatch = tf.oneHot(tf.tensor1d(trainYs.slice(i, i + BATCH_SIZE), 'int32'), 10); const result = model.trainOnBatch(xsBatch, ysBatch); xsBatch.dispose(); ysBatch.dispose(); await tf.nextFrame(); } }trainOnBatch一次只训练一个 batch,不会像model.fit一样把整个数据集一次性吞进去。配合每训练一个 batch 就await tf.nextFrame(),页面虽然不能完全保持 60 帧,但至少不会卡到假死。
如果计算量进一步加大,就上 Web Worker。TensorFlow.js 官方支持在 Worker 里加载和运行模型,把训练的大头计算放到后台线程,主线程只负责 UI。成本是通信复杂度上升,需要处理 Worker 线程里的模型实例管理和错误传播。我的建议是:小模型先试nextFrame(),巨大模型再考虑 Worker,不要一上来就把架构搞复杂。
4.3 内存泄漏:tf.memory()当场抓现行
内存泄漏这个问题我在第 2.2 节已经提到了基础机制,这里讲一个具体的排查现场。有一次我做实时手势识别,相当于摄像头每一帧都做一次predict,跑了几分钟后页面先是越来越卡,然后直接崩溃。打开 DevTools 看 GPU 进程的内存占用,一路飙升到 1.5GB。
当时我的代码长这样:
function detectGesture(frame) { const tensor = preprocess(frame); const pred = model.predict(tensor); const result = pred.argMax(-1).dataSync()[0]; return result; }看起来没问题,但实际上每一帧都产生了至少三个张量:tensor、pred、pred.argMax()的结果。它们使用完之后全都没有被释放。视频流一秒钟 15 帧,运行 10 分钟就是 9000 帧,每帧漏 3 个张量,不崩才怪。
修复后的版本:
function detectGesture(frame) { return tf.tidy(() => { const tensor = preprocess(frame); const pred = model.predict(tensor); return pred.argMax(-1).dataSync()[0]; }); }tf.tidy包裹了整个函数,函数里创建的中间张量会自动清理,只保留返回的数值,内存不再增长。这里有个容易误解的点:tf.tidy回调里return的如果是张量,它不会被自动回收,而是作为返回值传递出来,需要在外部再次dispose。我在实际中见过有人把返回值也留在tidy里,然后试图在外面dataSync(),结果拿到一个已经被释放的张量,直接报错。
以后凡是看到“越跑越慢、最后崩溃”的 TensorFlow.js 页面,第一个排查动作永远是打开控制台跑tf.memory(),记录numTensors和numBytesInGPU两个值,每执行 100 次操作后再看一次。数值只增不减,基本就是泄漏实锤了。
5. 跨浏览器兼容、移动端适配与常见报错速查
5.1 不同浏览器和设备的真实差异
TensorFlow.js 给我们的跨浏览器体验确实比想象中平滑,但“平滑”不等于“没有差异”。我实际测试过 Chrome、Edge、Safari 以及各类国产 WebView,结论是:Chrome 系的 WebGL 实现最好,推理性能和稳定性都排第一;Edge 换到 Chromium 内核后和 Chrome 基本一致;Safari 的问题时而出现,主要集中在 iOS 设备上的纹理大小限制和内存不足;普通 WebView 则看内核版本,Android 上有些 WebView 内核太老,WebGL 2 支持不完整。
这里要特别提醒一件事:如果你用 HBuilderX 这类工具的内置浏览器做调试,一定要确认内核版本。内置浏览器不等同于用户手机里的浏览器,WebGL 功能可能不完整。一旦出现tf.getBackend()返回'cpu'而代码里默认是'webgl'的情况,页面可能运行得很怪。我在移动端项目调试时固定流程是:先console.log(tf.getBackend()),再console.log(tf.memory()),确认后端正常才开始跑。
iOS Safari 下遇到 WebGL 纹理上限是家常便饭。28x28 的 MNIST 当然没问题,但一旦输入图片分辨率超过 2048x2048,或者模型中间层输出在纹理里装不下,就可能崩。遇到这种情况我的策略是:给页面加一个后端切换开关,在设置面板里强制指定tf.setBackend('wasm')。WASM 后端跑在 CPU 上,不依赖 GPU 纹理,兼容性一下就好了,虽然速度稍慢但胜在稳定。
5.2 报错速查表
这几年我收集了不少报错现场,整理成一张速查表。遇到问题时可以直接对照着排查:
| 报错信息 | 原因 | 解决方案 |
|---|---|---|
Error: Cannot find a valid backend | 后端没有加载或 WebGL 不可用 | 确认tf.min.js已引入;检查浏览器是否开启硬件加速;考虑加载 WASM 后端 |
Cannot read properties of undefined (reading 'rank') | model.json加载失败或权重路径错误 | 检查模型文件路径、model.json里weightsManifest指向的.bin文件是否存在 |
Input tensors to model must have the same dtype | 输入张量类型与模型预期不一致 | 检查是否有toFloat()或cast('float32'),确保输入为 float32 |
The shape of dict provides unexpected shape | predict传入的 input shape 不对 | 确认带 batch 维度[batch, h, w, c],必要时用expandDims |
Tensor is disposed | 张量被释放后还在使用 | 检查tidy作用域,不要在tidy外使用内部张量 |
WebGL context lost | 浏览器 GPU 进程崩溃或上下文重置 | 监听webglcontextlost事件,重置模型实例,必要时重新加载页面 |
Cannot read property 'dataSync' of null | predict结果为空或模型没有正确加载 | 先确认模型 load 完成再执行预测,await model不要漏 |
5.3 哪些场景最适合 TensorFlow.js
做完手写数字识别这个项目后,我最大的收获是知道了“什么东西不该放浏览器里跑”。适合用 TensorFlow.js 的场景有几个共同特点:数据隐私敏感、实时性要求高、离线可用、模型规模可控。
数据隐私敏感最好理解,图片不出浏览器,天然符合合规要求。实时性要求高,是因为省掉了网络请求,本地 GPU 推理几十毫秒就能出结果,视频帧率都能跟上。离线可用则是静态资源天然具备的特性,模型打包进网页后,断网也能用。模型规模可控具体说就是模型权重文件不超过几十 MB,再大的模型就不建议前端硬扛了。
我在给客户做内部表格识别时,还留了一条“逃生通道”:页面里放了个隐藏参数,允许强制切换 WASM 后端。这在碰上了 WebGL 纹理限制或桌面浏览器某个实习生改了 GPU 设置导致 WebGL 不可用的场景下,成了救命的保险栓。每次有人问我 TensorFlow.js 真能用吗,我都会说:推理能用,训练要克制。凡是给客户交付用 TensorFlow.js 的项目,我都会加一个“后端选择”的隐藏开关,默认 WebGL,遇到兼容问题切成 WASM。这个习惯救过我至少三次场,你也可以直接抄。