TensorFlow.js 实战
模型加载/训练/迁移、WebGL 后端加速、浏览器端实时推理应用
TensorFlow.js 概述
TensorFlow.js 是 TensorFlow 的 JavaScript 版本,由 Google 维护的开源机器学习库。它让开发者可以在浏览器和 Node.js 环境中运行机器学习模型,无需安装 Python 或配置 GPU 驱动。
TensorFlow.js 支持三种计算后端:
| 后端 | 计算设备 | 特点 | 适用场景 |
|---|---|---|---|
| CPU(tfjs-backend-cpu) | CPU | 兼容性最好,所有浏览器均支持 | 模型调试、小规模推理、无 WebGL 环境 |
| WebGL(tfjs-backend-webgl) | GPU(通过 WebGL) | 性能最佳,利用 GPU 并行计算 | 实时推理、训练中型模型、图像处理 |
| WebAssembly(tfjs-backend-wasm) | CPU(通过 WASM) | 比 CPU 后端快 2-10 倍,需加载 .wasm 文件 | 不支持 WebGL 的现代浏览器、移动端 |
核心 API 包括:tf.tensor() 创建张量、tf.sequential() 构建顺序模型、tf.loadLayersModel() 加载模型、model.fit() 训练模型、model.predict() 推理预测。
模型加载
TensorFlow.js 支持加载多种格式的预训练模型,可以在浏览器端直接进行推理。
加载 Keras / TF.js 格式模型
使用 tf.loadLayersModel() 加载由 TensorFlow/Keras 训练后转换的模型(model.json + weight 文件):
// 从 URL 加载 MobileNet 进行图像分类
async function loadMobileNet() {
const model = await tf.loadLayersModel(
'https://storage.googleapis.com/tfjs-models/tfjs/mobilenet_v1_0.25_224/model.json'
);
return model;
}加载 GraphModel 格式模型
使用 tf.loadGraphModel() 加载 TensorFlow SavedModel 或 Frozen Model 转换的模型,适合包含复杂计算图的模型:
// 加载 GraphModel 格式的模型
async function loadGraphModel() {
const model = await tf.loadGraphModel(
'https://storage.googleapis.com/tfjs-models/tfjs/mobilenet_v2_1.0_224/model.json'
);
return model;
}加载模型并进行图像分类
以下是一个完整的 MobileNet 图像分类示例:
// 加载 MobileNet 并对图像进行分类
async function classifyImage(imageElement) {
// 加载模型
const model = await tf.loadLayersModel(
'https://storage.googleapis.com/tfjs-models/tfjs/mobilenet_v1_0.25_224/model.json'
);
// 加载 ImageNet 标签
const labels = await fetch(
'https://storage.googleapis.com/tfjs-models/tfjs/mobilenet_v1_0.25_224/labels.txt'
).then(res => res.text());
const labelList = labels.split('\n').filter(l => l.trim() !== '');
// 预处理图像
const tensor = tf.browser.fromPixels(imageElement)
.resizeNearestNeighbor([224, 224])
.toFloat()
.sub([123, 117, 104]) // 减去均值
.expandDims(0);
// 推理
const predictions = model.predict(tensor);
const topK = tf.topk(predictions, 3);
const topKIndices = topK.indices.dataSync();
const topKScores = topK.values.dataSync();
const result = [];
for (let i = 0; i < 3; i++) {
result.push({
label: labelList[topKIndices[i]],
score: topKScores[i]
});
}
// 清理内存
tensor.dispose();
predictions.dispose();
topK.indices.dispose();
topK.values.dispose();
return result;
}模型训练
TensorFlow.js 支持在浏览器中训练模型,利用客户端的计算资源进行训练,无需服务器。这对于隐私敏感数据(如图片、用户行为数据)特别有价值。
数据管理
使用 tf.data.Dataset 管理训练数据,支持从数组、CSV 文件等来源创建数据集:
// 创建数据集
const xs = tf.tensor2d([1, 2, 3, 4, 5, 6, 7, 8, 9, 10], [10, 1]);
const ys = tf.tensor2d([3, 5, 7, 9, 11, 13, 15, 17, 19, 21], [10, 1]);
const dataset = tf.data.array([...Array(10).keys()].map(i => ({
xs: [i + 1],
ys: [2 * (i + 1) + 1]
}))).batch(2);浏览器中训练线性回归模型
以下示例在浏览器中训练一个简单的线性回归模型(拟合 y = 2x + 1):
// 生成训练数据
const xs = tf.tensor2d([1, 2, 3, 4, 5, 6, 7, 8, 9, 10], [10, 1]);
const ys = tf.tensor2d([3, 5, 7, 9, 11, 13, 15, 17, 19, 21], [10, 1]);
// 定义模型
const model = tf.sequential();
model.add(tf.layers.dense({ units: 1, inputShape: [1] }));
// 编译模型
model.compile({
optimizer: tf.train.sgd(0.01),
loss: 'meanSquaredError'
});
// 训练模型
async function trainModel() {
const history = await model.fit(xs, ys, {
batchSize: 2,
epochs: 100,
shuffle: true,
callbacks: {
onEpochEnd: (epoch, logs) => {
console.log(`Epoch ${epoch}: loss = ${logs.loss.toFixed(4)}`);
}
}
});
// 预测
const output = model.predict(tf.tensor2d([11], [1, 1]));
console.log(`预测 x=11 时 y = ${output.dataSync()[0].toFixed(2)}`);
// 预期输出约为 23(2 * 11 + 1 = 23)
// 清理内存
output.dispose();
}训练循环说明
model.fit() 是浏览器端训练的核心方法,支持以下关键参数:
batchSize:每批训练的样本数,影响内存占用和收敛速度epochs:训练轮数shuffle:是否在每个 epoch 前打乱数据validationSplit:从训练集中划分验证集的比例callbacks:训练回调,支持onEpochBegin、onEpochEnd、onBatchEnd等
迁移学习
迁移学习是在浏览器中实现高效训练的常用技术。通过冻结预训练模型的底层,只训练新增的顶层,可以在少量数据和较短训练时间内获得良好的分类效果。
冻结预训练模型层
// 加载预训练的 MobileNet,不包含顶层分类层
const mobilenet = await tf.loadLayersModel(
'https://storage.googleapis.com/tfjs-models/tfjs/mobilenet_v1_0.25_224/model.json'
);
// 冻结所有层
for (const layer of mobilenet.layers) {
layer.trainable = false;
}
// 获取 mobilenet 的输出层(去掉原始分类层)
const truncatedModel = tf.model({
inputs: mobilenet.inputs,
outputs: mobilenet.getLayer('reshape_2').output
});
// 构建迁移学习模型
const model = tf.sequential();
model.add(truncatedModel);
model.add(tf.layers.flatten());
model.add(tf.layers.dense({
units: 128,
activation: 'relu'
}));
model.add(tf.layers.dropout({ rate: 0.5 }));
model.add(tf.layers.dense({
units: 5, // 自定义类别数
activation: 'softmax'
}));
// 编译模型
model.compile({
optimizer: tf.train.adam(0.001),
loss: 'categoricalCrossentropy',
metrics: ['accuracy']
});使用 tf.model() 构建迁移学习模型
通过 tf.model() 可以更灵活地构建迁移学习模型,支持多输入或多输出场景:
// 使用 tf.model() 构建迁移学习模型
async function buildTransferModel() {
const baseModel = await tf.loadLayersModel(
'https://storage.googleapis.com/tfjs-models/tfjs/mobilenet_v1_0.25_224/model.json'
);
// 冻结底层
for (const layer of baseModel.layers) {
layer.trainable = false;
}
// 截断到指定层
const truncatedOutput = baseModel.getLayer('conv_pw_13_relu').output;
// 添加自定义层
const flatten = tf.layers.flatten().apply(truncatedOutput);
const dense1 = tf.layers.dense({ units: 128, activation: 'relu' }).apply(flatten);
const dropout = tf.layers.dropout({ rate: 0.5 }).apply(dense1);
const output = tf.layers.dense({ units: 10, activation: 'softmax' }).apply(dropout);
// 构建新模型
const model = tf.model({
inputs: baseModel.inputs,
outputs: output
});
model.compile({
optimizer: tf.train.adam(0.0001),
loss: 'categoricalCrossentropy',
metrics: ['accuracy']
});
return model;
}WebGL 后端加速
WebGL 后端是 TensorFlow.js 在浏览器中实现高性能计算的关键,它将张量操作映射为 WebGL 纹理上的着色器程序。
配置 WebGL 后端
// 设置 WebGL 后端(默认即为 WebGL)
await tf.setBackend('webgl');
console.log('当前后端:', tf.getBackend());
// 检查是否支持 WebGL
const isWebGLSupported = tf.ENV.features.WEBGL_VERSION >= 1;
console.log('WebGL 支持:', isWebGLSupported);WebGL 后端 vs CPU 后端性能对比
| 操作类型 | CPU 后端 | WebGL 后端 | 加速比 |
|---|---|---|---|
| 矩阵乘法(1024x1024) | ~120ms | ~8ms | 约 15x |
| 卷积(224x224x3, 32 filters) | ~450ms | ~25ms | 约 18x |
| MobileNet 推理 | ~350ms | ~55ms | 约 6x |
| 批量归一化 | ~80ms | ~5ms | 约 16x |
内存管理
在浏览器中使用 GPU 内存需要特别注意内存管理,避免 WebGL 纹理内存泄漏导致浏览器卡顿或崩溃。
// 使用 tf.tidy() 自动清理中间张量
function computeWithTidy(input) {
return tf.tidy(() => {
const a = tf.mul(input, 2);
const b = tf.add(a, 1);
const c = tf.relu(b);
// a 和 b 会在 tidy 结束时自动释放
return c;
});
}
// 手动使用 tf.dispose() 清理
function computeWithDispose(input) {
const a = tf.mul(input, 2);
const b = tf.add(a, 1);
const result = tf.relu(b);
// 手动释放不再使用的中间张量
a.dispose();
b.dispose();
return result;
}
// 查看内存使用情况
console.log(tf.memory());
// 输出: { numTensors: ..., numBytes: ..., numDataBuffers: ..., unreliable: false }浏览器端实时推理应用
以下示例展示如何使用摄像头进行实时图像分类,结合 MobileNet 迁移学习模型和 requestAnimationFrame 实现流畅的实时推理。
// 摄像头实时分类应用
class RealTimeClassifier {
constructor(videoElement, canvasElement) {
this.video = videoElement;
this.canvas = canvasElement;
this.ctx = canvasElement.getContext('2d');
this.model = null;
this.isRunning = false;
this.classNames = ['class_0', 'class_1', 'class_2', 'class_3', 'class_4'];
}
// 初始化:加载模型和摄像头
async init() {
// 加载预训练模型
this.model = await tf.loadLayersModel(
'/models/my_custom_model/model.json'
);
// 启动摄像头
const stream = await navigator.mediaDevices.getUserMedia({
video: { width: 640, height: 480, facingMode: 'environment' }
});
this.video.srcObject = stream;
await this.video.play();
// 开始推理循环
this.isRunning = true;
this.loop();
}
// 推理循环
loop() {
if (!this.isRunning) return;
// 绘制视频帧到 canvas
this.ctx.drawImage(this.video, 0, 0, 224, 224);
// 预处理图像
const tensor = tf.tidy(() => {
return tf.browser.fromPixels(this.canvas)
.toFloat()
.sub([123, 117, 104])
.div(255)
.expandDims(0);
});
// 推理
const predictions = this.model.predict(tensor);
const probabilities = predictions.dataSync();
// 获取最高概率类别
const maxIndex = probabilities.indexOf(Math.max(...probabilities));
const confidence = probabilities[maxIndex];
// 在 canvas 上显示结果
this.ctx.fillStyle = 'white';
this.ctx.font = '24px sans-serif';
this.ctx.fillText(
`分类: ${this.classNames[maxIndex]} (${(confidence * 100).toFixed(1)}%)`,
10, 40
);
// 清理内存
tensor.dispose();
predictions.dispose();
// 继续下一帧
requestAnimationFrame(() => this.loop());
}
// 停止推理
stop() {
this.isRunning = false;
if (this.video.srcObject) {
this.video.srcObject.getTracks().forEach(track => track.stop());
}
}
}使用方式
const video = document.getElementById('videoElement');
const canvas = document.getElementById('canvasElement');
const classifier = new RealTimeClassifier(video, canvas);
classifier.init();总结
TensorFlow.js 将机器学习的强大能力带入浏览器端,使得开发者可以构建无需服务器、保护用户隐私的智能 Web 应用。通过 WebGL 后端加速,浏览器端的深度学习推理性能已经接近原生水平;结合迁移学习技术,可以在浏览器中用少量数据快速训练出高精度的自定义模型。合理使用 tf.tidy() 和 tf.dispose() 进行内存管理,是构建稳定高效的浏览器端 ML 应用的关键。