TensorFlow.js实战:浏览器端机器学习模型部署
·
TensorFlow.js 简介
TensorFlow.js 是一个基于 JavaScript 的库,用于在浏览器和 Node.js 中训练和部署机器学习模型。支持直接运行预训练模型或迁移学习,无需后端服务器。
环境准备
安装 TensorFlow.js 库:
npm install @tensorflow/tfjs
或通过 CDN 引入:
<script src="https://cdn.jsdelivr.net/npm/@tensorflow/tfjs@latest"></script>
加载预训练模型
使用 tf.loadLayersModel 加载已转换的 TensorFlow 模型(需转换为 JSON 格式):
async function loadModel() {
const model = await tf.loadLayersModel('model/model.json');
return model;
}
浏览器端模型推理
输入数据需转换为 Tensor 格式,调用 model.predict 执行推理:
const input = tf.tensor2d([[...inputArray]], [1, inputShape]);
const output = model.predict(input);
const result = output.dataSync(); // 获取预测结果
模型训练与迁移学习
通过 tf.model 定义新模型结构,结合预训练层进行迁移学习:
const baseModel = await tf.loadLayersModel('base-model.json');
baseModel.trainable = false; // 冻结预训练层
const newModel = tf.sequential({
layers: [
baseModel,
tf.layers.dense({ units: 10, activation: 'softmax' })
]
});
newModel.compile({ optimizer: 'adam', loss: 'categoricalCrossentropy' });
性能优化技巧
- WebGL 加速:默认启用 GPU 加速,可通过
tf.backend()检查后端类型。 - 内存管理:手动释放 Tensor 内存以避免泄漏:
tf.dispose([tensor1, tensor2]); - 量化模型:使用
tfjs-converter将模型转换为 8 位或 16 位以减少体积。
实际应用示例
手写数字识别(MNIST)完整流程:
- 数据预处理:将图像转换为 28x28 的灰度张量。
- 模型加载:加载预训练的卷积神经网络。
- 实时预测:监听画布绘制事件并输出分类结果。
代码片段:
const canvas = document.getElementById('draw-canvas');
const ctx = canvas.getContext('2d');
// 监听绘制事件并调用 model.predict()
注意事项
- 模型大小限制:浏览器缓存限制建议模型小于 50MB。
- 跨域问题:若模型托管在 CDN,需配置 CORS 头部。
- 移动端兼容性:iOS 的 WebGL 性能可能较低,需测试降级方案。
更多推荐


所有评论(0)