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)完整流程:

  1. 数据预处理:将图像转换为 28x28 的灰度张量。
  2. 模型加载:加载预训练的卷积神经网络。
  3. 实时预测:监听画布绘制事件并输出分类结果。

代码片段:

const canvas = document.getElementById('draw-canvas');
const ctx = canvas.getContext('2d');
// 监听绘制事件并调用 model.predict()

注意事项

  • 模型大小限制:浏览器缓存限制建议模型小于 50MB。
  • 跨域问题:若模型托管在 CDN,需配置 CORS 头部。
  • 移动端兼容性:iOS 的 WebGL 性能可能较低,需测试降级方案。
Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐