三亩地 三亩地SAN MU DI · CODE DIARY
ARTICLE DETAIL

日记详情

真实记录编程学习的某一天,欢迎挑你感兴趣的翻一翻。

JavaScript深度学习:从入门到实战

JavaScript深度学习:从入门到实战

1. JavaScript与深度学习的奇妙结合

当大多数人听到"深度学习"这个词时,脑海中浮现的往往是Python、TensorFlow或PyTorch这些传统工具。但今天我要告诉你一个可能让你惊讶的事实:JavaScript也能成为深度学习的强大工具。作为一名在Web开发和机器学习交叉领域工作多年的工程师,我见证了JavaScript生态系统中深度学习工具的惊人成长。

你可能会有疑问:为什么要在JavaScript中做深度学习?答案很简单——因为Web无处不在。想象一下,直接在浏览器中运行图像识别模型而不需要服务器支持,或者在移动设备上离线执行自然语言处理任务。这些场景正是JavaScript深度学习的独特优势所在。

我清楚地记得第一次用TensorFlow.js完成一个MNIST手写数字识别项目时的震撼——不需要任何复杂的后端部署,模型直接在浏览器中训练和推理,整个过程流畅得令人难以置信。从那时起,我就迷上了用JavaScript探索深度学习的可能性。

2. JavaScript深度学习生态全景

2.1 核心工具库介绍

JavaScript深度学习生态系统已经相当丰富,以下是最主流的几个工具:

  1. TensorFlow.js:Google推出的JavaScript版TensorFlow,支持在浏览器和Node.js环境中运行

    • 提供与Python版相似的高级API
    • 支持WebGL加速
    • 可以直接加载预训练的Python模型
  2. Brain.js:专注于神经网络的轻量级库

    • 简单易用的API
    • 支持多种网络类型(前馈、RNN、LSTM等)
    • 非常适合快速原型开发
  3. ML5.js:基于TensorFlow.js的高级封装

    • 提供预训练模型(图像分类、姿态检测等)
    • 设计理念是让机器学习对创意编码者更友好
  4. ONNX.js:支持运行ONNX格式模型的运行时

    • 可以与其他框架训练的模型互操作
    • 支持WebAssembly加速

2.2 浏览器与Node.js环境对比

选择运行环境是JavaScript深度学习项目的第一个关键决策:

特性浏览器环境Node.js环境
计算性能依赖WebGL,中等可调用原生扩展,高性能
模型大小限制受内存限制较大可处理更大模型
部署便利性无需服务器,直接运行需要服务器环境
离线能力完全离线依赖服务器
隐私性数据完全在客户端处理数据需要发送到服务器
预训练模型支持部分支持支持更广泛的模型格式
开发调试便利性可直接使用浏览器开发者工具需要额外工具

提示:对于需要快速原型验证或注重隐私保护的项目,浏览器环境是更好的选择;而对于需要处理大型模型或复杂计算的任务,Node.js环境更合适。

3. 实战:构建你的第一个JavaScript深度学习项目

3.1 环境准备与基础配置

让我们从最基础的开始——搭建开发环境。与Python生态不同,JavaScript深度学习项目通常不需要复杂的虚拟环境管理。

基础工具安装:

# 初始化项目 npm init -y # 安装TensorFlow.js核心库 npm install @tensorflow/tfjs # 如果需要Node.js后端支持 npm install @tensorflow/tfjs-node # 对于有GPU的机器(仅限Node.js环境) npm install @tensorflow/tfjs-node-gpu

HTML基础模板(浏览器环境):

<!DOCTYPE html> <html> <head> <title>我的第一个JS深度学习项目</title> <script src="https://cdn.jsdelivr.net/npm/@tensorflow/tfjs@latest"></script> </head> <body> <script src="index.js"></script> </body> </html>

3.2 线性回归实战

让我们从一个简单的线性回归问题开始,这能帮助你理解JavaScript深度学习的基本工作流程。

完整代码示例:

// 生成合成数据 function generateData(numPoints, coeff, sigma = 0.04) { return tf.tidy(() => { const [a, b, c, d] = [ tf.scalar(coeff.a), tf.scalar(coeff.b), tf.scalar(coeff.c), tf.scalar(coeff.d) ]; const xs = tf.randomUniform([numPoints], -1, 1); // 三次多项式:y = a*x^3 + b*x^2 + c*x + d const ys = a.mul(xs.pow(tf.scalar(3))) .add(b.mul(xs.square())) .add(c.mul(xs)) .add(d) .add(tf.randomNormal([numPoints], 0, sigma)); return {xs, ys}; }); } // 定义模型 function createModel() { const model = tf.sequential(); model.add(tf.layers.dense({ units: 1, inputShape: [1], activation: 'linear' })); model.compile({ optimizer: tf.train.sgd(0.1), loss: 'meanSquaredError' }); return model; } // 训练模型 async function trainModel(model, xs, ys) { const history = await model.fit(xs, ys, { epochs: 100, callbacks: { onEpochEnd: (epoch, logs) => { if (epoch % 10 === 0) { console.log(`Epoch ${epoch}: loss = ${logs.loss}`); } } } }); return history; } // 主函数 async function run() { const coeff = {a: -0.8, b: -0.2, c: 0.9, d: 0.5}; const trainingData = generateData(100, coeff); const model = createModel(); await trainModel(model, trainingData.xs, trainingData.ys); // 测试模型 const testXs = tf.tensor1d([-0.5, 0, 0.5]); const preds = model.predict(testXs); preds.print(); } run();

代码解析:

  1. generateData函数创建了一个带有噪声的三次多项式数据集。这里使用了tf.tidy来确保中间张量被正确清理。

  2. createModel定义了一个简单的线性回归模型,虽然我们的数据是非线性的,但线性模型仍然可以学习到一定的趋势。

  3. trainModel函数展示了基本的训练流程,包括回调函数的使用来监控训练进度。

  4. 注意所有操作都是异步的,因为TensorFlow.js的许多操作返回Promise。

注意:在浏览器环境中运行大量计算时,可能会阻塞UI线程。对于复杂任务,考虑使用Web Worker或将计算拆分为小块。

3.3 可视化训练过程

可视化是理解模型行为的关键。在浏览器环境中,我们可以轻松使用Chart.js等库来可视化训练过程。

增强版训练函数:

async function trainModelWithVisualization(model, xs, ys) { // 准备画布 const ctx = document.getElementById('chart').getContext('2d'); const chart = new Chart(ctx, { type: 'scatter', data: { datasets: [ { label: '原始数据', data: Array.from(ys.dataSync()).map((y, i) => ({ x: xs.dataSync()[i], y: y })), backgroundColor: 'rgba(75, 192, 192, 0.6)' }, { label: '模型预测', data: [], backgroundColor: 'rgba(255, 99, 132, 0.6)', showLine: true } ] }, options: { scales: { x: { type: 'linear', position: 'center' } } } }); // 自定义回调 const callbacks = { onEpochEnd: async (epoch, logs) => { if (epoch % 5 === 0) { const testXs = tf.linspace(-1, 1, 100); const preds = model.predict(testXs.reshape([100, 1])); chart.data.datasets[1].data = Array.from(testXs.dataSync()) .map((x, i) => ({ x: x, y: preds.dataSync()[i] })); chart.update(); await tf.nextFrame(); // 让浏览器有机会渲染 } } }; await model.fit(xs, ys, { epochs: 200, batchSize: 32, callbacks: callbacks }); }

这个增强版训练函数会在训练过程中动态更新图表,让你直观地看到模型如何逐步拟合数据。这种即时反馈对于调试模型和超参数非常有价值。

4. 进阶主题:图像分类实战

4.1 使用预训练模型

TensorFlow.js提供了多种预训练模型,让我们可以快速实现强大的功能而不必从头训练。

MobileNet图像分类示例:

async function loadAndPredict() { // 加载MobileNet模型 const model = await tf.loadLayersModel( 'https://storage.googleapis.com/tfjs-models/tfjs/mobilenet_v1_0.25_224/model.json' ); // 加载图像 const image = document.getElementById('my-image'); const tensor = tf.browser.fromPixels(image) .resizeNearestNeighbor([224, 224]) .toFloat() .expandDims(); // 预处理(ImageNet标准化) const mean = tf.tensor1d([0.485, 0.456, 0.406]); const std = tf.tensor1d([0.229, 0.224, 0.225]); const preprocessed = tensor.div(255.0).sub(mean).div(std); // 预测 const predictions = model.predict(preprocessed); const top5 = Array.from(predictions.dataSync()) .map((p, i) => ({ probability: p, className: IMAGENET_CLASSES[i] })) .sort((a, b) => b.probability - a.probability) .slice(0, 5); console.log(top5); }

关键点说明:

  1. 我们直接从Google的服务器加载了MobileNet模型,这是一个在ImageNet数据集上预训练的卷积神经网络。

  2. 图像预处理步骤非常重要,必须与模型训练时使用的预处理方式一致(这里是ImageNet的标准归一化)。

  3. tf.browser.fromPixels是TensorFlow.js提供的专门用于处理浏览器图像元素的API。

  4. 预测结果包含了1000个ImageNet类别的概率,我们只取前5个最高概率的结果。

4.2 迁移学习实战

预训练模型很强大,但要让它们解决我们的特定问题,通常需要进行迁移学习。

迁移学习代码示例:

async function transferLearning() { // 加载MobileNet的基础部分(去掉顶层) const baseModel = await tf.loadLayersModel( 'https://storage.googleapis.com/tfjs-models/tfjs/mobilenet_v1_0.25_224/model.json' ); // 截断模型,保留特征提取层 const layer = baseModel.getLayer('conv_pw_13_relu'); const truncatedModel = tf.model({ inputs: baseModel.inputs, outputs: layer.output }); // 冻结基础模型权重 truncatedModel.trainable = false; // 添加新的可训练层 const newModel = tf.sequential(); newModel.add(truncatedModel); newModel.add(tf.layers.globalAveragePooling2d()); newModel.add(tf.layers.dense({ units: 10, activation: 'softmax' })); // 编译模型 newModel.compile({ optimizer: tf.train.adam(0.0001), loss: 'categoricalCrossentropy', metrics: ['accuracy'] }); // 准备自定义数据(这里需要你自己的数据集) // const {trainXs, trainYs, testXs, testYs} = prepareYourData(); // 训练新层 // await newModel.fit(trainXs, trainYs, { // epochs: 20, // validationData: [testXs, testYs] // }); }

迁移学习的关键步骤:

  1. 加载预训练模型并截断,保留特征提取部分。

  2. 冻结基础模型的权重,防止它们在训练过程中被修改。

  3. 添加新的可训练层来适应你的特定任务。

  4. 使用较小的学习率训练,因为特征已经相对较好,只需要微调。

实战技巧:对于小型数据集,可以尝试不同的数据增强技术(旋转、翻转、颜色调整等)来提高模型泛化能力。TensorFlow.js提供了tf.image命名空间下的多种图像处理操作。

5. 性能优化与生产部署

5.1 性能优化技巧

JavaScript深度学习面临的最大挑战之一是性能。以下是一些经过验证的优化技巧:

  1. 内存管理

    • 使用tf.tidy()自动清理中间张量
    • 手动调用tensor.dispose()释放不再需要的张量
    • 监控内存使用:tf.memory()可以打印当前内存状态
  2. 批量处理

    • 尽量使用批量操作而不是循环处理单个样本
    • 例如,使用model.predictOnBatch()而不是循环调用model.predict()
  3. WebGL优化

    • 确保张量形状在连续操作中保持一致,减少WebGL着色器重新编译
    • 对于小型张量,CPU可能比WebGL更快(可以使用tf.setBackend('cpu')测试)
  4. 模型量化

    • 使用16位浮点或8位整数权重减小模型大小
    • TensorFlow.js支持从量化的TensorFlow Lite模型转换

5.2 生产部署策略

将JavaScript深度学习模型部署到生产环境需要考虑几个关键因素:

部署选项对比:

部署方式优点缺点适用场景
纯客户端无服务器成本,隐私友好受限于设备性能,模型大小受限轻量级应用,注重隐私
客户端+服务器平衡计算负载,支持更大模型需要服务器基础设施大多数生产应用
WebAssembly性能接近原生,支持复杂模型加载时间较长,兼容性问题性能关键型应用
边缘计算低延迟,减少网络传输部署复杂度高IoT、实时处理应用
CDN缓存模型快速加载,减轻服务器压力模型更新有延迟静态模型,不频繁更新

模型优化建议:

  1. 模型剪枝:移除对输出影响较小的神经元,减小模型大小

    // 使用TensorFlow.js的模型转换工具进行剪枝 const prunedModel = await tf.graphModel.convertToPruned(model, pruningConfig);
  2. 模型量化:降低权重精度

    // 转换为16位浮点 const quantizedModel = await model.save(tf.io.withSaveHandler(async (artifacts) => { return { modelTopology: artifacts.modelTopology, weightSpecs: artifacts.weightSpecs, weightData: new Uint16Array(artifacts.weightData.buffer) }; }));
  3. 模型分割:将大模型分成多个部分,按需加载

    // 先加载模型骨架 const modelPart1 = await tf.loadLayersModel('model_part1.json'); // 用户交互后再加载剩余部分 button.onclick = async () => { const modelPart2 = await tf.loadLayersModel('model_part2.json'); // 组合模型... };
  4. 渐进式加载:先加载轻量级模型,后台下载完整模型

    // 先加载精简版 const lightModel = await tf.loadLayersModel('light_model.json'); // 后台下载完整模型 const fullModelPromise = tf.loadLayersModel('full_model.json') .then(model => { console.log('完整模型已就绪'); return model; }); // 当需要更高精度时 const fullModel = await fullModelPromise;

监控与维护:

  1. 实现模型性能监控:

    // 记录推理时间 const start = performance.now(); const results = await model.predict(input); const duration = performance.now() - start; // 发送到分析服务 fetch('/analytics', { method: 'POST', body: JSON.stringify({ model: 'image-classifier', version: '1.2', inferenceTime: duration, device: navigator.userAgent }) });
  2. 建立模型版本控制:

    // 检查模型版本 async function checkModelVersion() { const response = await fetch('/model-version'); const { latestVersion } = await response.json(); if (localStorage.getItem('modelVersion') !== latestVersion) { console.log('新模型版本可用'); // 触发模型更新流程 } }
  3. 实现回退机制:

    async function loadModelWithFallback() { try { return await tf.loadLayersModel('https://example.com/new-model.json'); } catch (error) { console.error('加载新模型失败,回退到旧版本', error); return await tf.loadLayersModel('https://example.com/fallback-model.json'); } }

6. 常见问题与调试技巧

6.1 典型错误与解决方案

在JavaScript深度学习开发中,有几个常见陷阱需要注意:

  1. 内存泄漏

    • 症状:浏览器标签内存使用量持续增长,最终崩溃
    • 原因:未正确释放张量内存
    • 解决方案
      // 错误方式:循环创建未释放的张量 for (let i = 0; i < 1000; i++) { const tensor = tf.tensor([i]); // 内存泄漏! } // 正确方式:使用tf.tidy自动清理 for (let i = 0; i < 1000; i++) { tf.tidy(() => { const tensor = tf.tensor([i]); // 自动清理 }); }
  2. WebGL上下文丢失

    • 症状:模型突然停止工作,控制台报WebGL错误
    • 原因:浏览器回收了WebGL资源(如切换标签页后返回)
    • 解决方案
      // 监听上下文丢失事件 const canvas = document.getElementById('webgl-canvas'); canvas.addEventListener('webglcontextlost', (e) => { e.preventDefault(); console.log('WebGL上下文丢失,需要重新初始化模型'); // 重新加载模型 });
  3. 模型加载失败

    • 症状:模型文件加载超时或返回404
    • 原因:路径错误、CORS限制或网络问题
    • 解决方案
      // 使用try-catch处理加载错误 try { const model = await tf.loadLayersModel('model.json'); } catch (error) { console.error('模型加载失败:', error); // 回退到CDN或本地缓存 const model = await tf.loadLayersModel('fallback-model.json'); }
  4. NaN损失值

    • 症状:训练过程中损失值突然变成NaN
    • 原因:学习率过高、数据未归一化或数值不稳定
    • 解决方案
      • 降低学习率(如从0.01降到0.001)
      • 检查输入数据是否已正确归一化(通常缩放到0-1或-1到1)
      • 添加梯度裁剪:
        model.compile({ optimizer: tf.train.adam(0.001), loss: 'meanSquaredError', clipValue: 0.5 // 裁剪梯度 });

6.2 调试工具与技术

有效的调试可以节省大量开发时间:

  1. TensorFlow.js自带工具

    // 启用调试模式 tf.enableDebugMode(); // 检查张量值(注意:会阻塞执行) const tensor = tf.tensor([1, 2, 3]); tensor.print(); // 打印张量内容 // 内存状态 console.log(tf.memory());
  2. 浏览器性能分析

    • 使用Chrome DevTools的Performance面板记录模型训练过程
    • 检查主要耗时操作和内存分配情况
    • 注意WebGL调用和着色器编译时间
  3. 模型可视化

    // 在浏览器中显示模型结构 function visualizeModel(model) { const surface = { name: '模型结构', tab: '模型' }; tfvis.show.modelSummary(surface, model); // 显示层详情 model.layers.forEach((layer, i) => { const layerSurface = { name: `层 ${i}: ${layer.name}`, tab: '层详情' }; tfvis.show.layer(layerSurface, layer); }); }
  4. 训练过程监控

    // 使用tfvis可视化训练指标 const metrics = ['loss', 'val_loss', 'acc', 'val_acc']; const container = { name: '训练指标', tab: '训练' }; const callbacks = tfvis.show.fitCallbacks(container, metrics); await model.fit(data, labels, { epochs: 100, validationSplit: 0.2, callbacks: callbacks });
  5. 自定义日志回调

    class CustomCallback extends tf.Callback { onEpochBegin(epoch, logs) { console.log(`Epoch ${epoch} 开始`); } onBatchEnd(batch, logs) { if (batch % 10 === 0) { console.log(`批次 ${batch}: 损失 = ${logs.loss.toFixed(4)}`); } } onEpochEnd(epoch, logs) { console.log(`Epoch ${epoch} 结束: 损失 = ${logs.loss.toFixed(4)}, 准确率 = ${logs.acc.toFixed(4)}`); } } await model.fit(data, labels, { epochs: 10, callbacks: new CustomCallback() });

6.3 跨浏览器兼容性问题

不同浏览器对WebGL的支持程度不同,可能导致性能差异甚至功能失效:

  1. WebGL特性检测

    function checkWebGLSupport() { const canvas = document.createElement('canvas'); const gl = canvas.getContext('webgl') || canvas.getContext('experimental-webgl'); if (!gl) { console.error('WebGL不支持!将回退到CPU'); tf.setBackend('cpu'); return false; } // 检查关键扩展 const extensions = [ 'OES_texture_float', 'WEBGL_draw_buffers', 'OES_element_index_uint' ]; extensions.forEach(ext => { if (!gl.getExtension(ext)) { console.warn(`扩展 ${ext} 不支持,某些功能可能受限`); } }); return true; }
  2. 浏览器特定问题

    • Safari:对WebGL 2.0支持有限,可能需要polyfill
    • iOS:有严格的内存限制,大模型容易崩溃
    • Firefox:WebGL性能通常较好,但可能有不同的精度行为
  3. 优雅降级策略

    async function loadModelWithFallback() { try { // 先尝试WebGL await tf.setBackend('webgl'); return await tf.loadLayersModel('complex-model.json'); } catch (webglError) { console.warn('WebGL失败:', webglError); try { // 回退到WASM await tf.setBackend('wasm'); console.log('使用WASM后端'); return await tf.loadLayersModel('simplified-model.json'); } catch (wasmError) { console.warn('WASM失败:', wasmError); // 最后尝试CPU await tf.setBackend('cpu'); console.log('使用CPU后端'); return await tf.loadLayersModel('lightweight-model.json'); } } }
  4. 性能基准测试

    async function benchmarkModel(model, input, iterations = 100) { // 预热 await model.predict(input).data(); // 测试推理时间 const start = performance.now(); for (let i = 0; i < iterations; i++) { await model.predict(input).data(); } const duration = performance.now() - start; console.log(`平均推理时间: ${(duration / iterations).toFixed(2)}ms`); return duration / iterations; } // 在不同后端上运行基准测试 async function compareBackends() { const model = await loadModel(); const input = tf.randomNormal(model.inputs[0].shape); const backends = ['webgl', 'wasm', 'cpu']; const results = {}; for (const backend of backends) { await tf.setBackend(backend); results[backend] = await benchmarkModel(model, input); } console.table(results); }

7. 前沿探索与未来方向

7.1 WebGPU与下一代加速

WebGPU是即将到来的新一代图形API,有望大幅提升浏览器中的计算性能:

// 检测WebGPU支持 async function checkWebGPU() { if (!('gpu' in navigator)) { console.log('WebGPU不支持'); return false; } try { const adapter = await navigator.gpu.requestAdapter(); const device = await adapter.requestDevice(); console.log('WebGPU可用:', adapter, device); return true; } catch (error) { console.error('WebGPU初始化失败:', error); return false; } } // TensorFlow.js未来可能会支持WebGPU后端 // tf.setBackend('webgpu');

WebGPU相比WebGL的优势:

  • 更低的开销,更高效的并行计算
  • 更好的显存管理
  • 支持计算着色器(compute shaders)
  • 更现代的API设计

7.2 模型压缩与量化新技术

最新的模型压缩技术可以让深度学习模型更适应浏览器环境:

  1. 知识蒸馏:训练小型"学生"模型模仿大型"教师"模型的行为

    // 伪代码示例 async function distillModel(teacher, student, data) { const temperature = 2.0; // 软化概率分布 const lambda = 0.5; // 平衡系数 student.compile({ optimizer: 'adam', loss: (yTrue, yPred) => { const teacherPred = teacher.predict(yTrue); const softTargets = tf.softmax(teacherPred.div(temperature)); const studentSoft = tf.softmax(yPred.div(temperature)); const kld = tf.losses.klDivergence(softTargets, studentSoft); const originalLoss = tf.losses.softmaxCrossEntropy(yTrue, yPred); return originalLoss.mul(1 - lambda).add(kld.mul(lambda * temperature * temperature)); } }); await student.fit(data.x, data.y, { epochs: 10 }); }
  2. 结构化剪枝:移除整个神经元或通道而不仅仅是单个权重

    // 使用TensorFlow.js模型优化工具包 import * as tfmo from '@tensorflow/tfjs-model-optimization'; const pruningParams = { pruningSchedule: tfmo.pruning.PolynomialDecay( 0.5, // 初始稀疏率 0.9, // 最终稀疏率 1000, // 开始步 2000, // 结束步 2.0 // 幂次 ) }; const model = tf.sequential(); model.add(tfmo.pruning.prune( tf.layers.dense({ units: 100, activation: 'relu' }), pruningParams ));
  3. 量化感知训练:在训练过程中模拟量化效果,提高最终量化模型的精度

    // 伪代码示例 import * as tfquant from '@tensorflow/tfjs-quantization'; const originalModel = await tf.loadLayersModel('original-model.json'); const quantizedModel = tfquant.quantizeModel(originalModel, { // 量化配置 weightBits: 8, activationBits: 8, mode: 'quantization-aware-training' }); // 继续训练量化模型 await quantizedModel.fit(trainData, trainLabels, { epochs: 5, callbacks: { onEpochEnd: (epoch, logs) => { console.log(`量化感知训练 Epoch ${epoch}: loss = ${logs.loss}`); } } });

7.3 联邦学习与隐私保护

JavaScript的广泛分布特性使其成为联邦学习的理想平台:

// 简化的联邦学习客户端代码 class FederatedClient { constructor(modelUrl) { this.model = null; this.localData = []; } async initialize() { this.model = await tf.loadLayersModel(modelUrl); } async trainOnLocalData() { // 在本地数据上训练 await this.model.fit(this.localData.x, this.localData.y, { epochs: 1, batchSize: 32 }); // 提取权重更新 const weights = this.model.getWeights(); const update = weights.map(w => w.clone()); return update; } async applyGlobalUpdate(globalWeights) { // 应用服务器聚合后的全局权重 this.model.setWeights(globalWeights); } } // 服务器端伪代码 async function aggregateUpdates(clientUpdates) { // 平均所有客户端的更新 const averagedWeights = []; const numClients = clientUpdates.length; for (let i = 0; i < clientUpdates[0].length; i++) { let sum = tf.zeros(clientUpdates[0][i].shape); for (const update of clientUpdates) { sum = sum.add(update[i]); } averagedWeights.push(sum.div(tf.scalar(numClients))); } return averagedWeights; }

联邦学习的优势:

  • 数据保留在客户端,保护隐私
  • 利用分布式设备计算资源
  • 可以持续改进模型而无需集中收集数据

7.4 与Web技术的深度集成

JavaScript深度学习正与各种Web技术深度融合:

  1. WebAssembly SIMD:单指令多数据加速

    // 检测SIMD支持 const simdSupported = (() => { try { return new WebAssembly.Instance( new WebAssembly.Module(new Uint8Array([0,97,115,109,1,0,0,0,1,5,1,96,0,1,123,3,2,1,0,10,10,1,8,0,65,0,253,15,253,98,11])) ).exports.f() === 42; } catch { return false; } })(); console.log('SIMD支持:', simdSupported);
  2. Web Workers并行计算

    // 主线程 const worker = new Worker('tf-worker.js'); worker.onmessage = (e) => { const { prediction } = e.data; console.log('收到预测结果:', prediction); }; worker.postMessage({ imageData: canvas.toDataURL() }); // tf-worker.js importScripts('https://cdn.jsdelivr.net/npm/@tensorflow/tfjs@latest'); let model; async function loadModel() { model = await tf.loadLayersModel('model.json'); } self.onmessage = async (e) => { if (!model) await loadModel(); const { imageData } = e.data; const tensor = preprocess(imageData); const prediction = model.predict(tensor); const results = await prediction.data(); self.postMessage({ prediction: results }); };
  3. WebRTC实时视频处理

    // 从摄像头获取视频流 async function setupCamera() { const stream = await navigator.mediaDevices.getUserMedia({ video: true }); const video = document.getElementById('video'); video.srcObject = stream; return new Promise((resolve) => { video.onloadedmetadata = () => { resolve(video); }; }); } // 实时处理视频帧 async function processVideoFrames(model, video) { const canvas = document.createElement('canvas'); const ctx = canvas.getContext('2d'); canvas.width = video.videoWidth; canvas.height = video.videoHeight; async function processFrame() { ctx.drawImage(video, 0, 0); const imageData = ctx.getImageData(0, 0, canvas.width, canvas.height); const tensor = tf.browser.fromPixels(imageData) .resizeNearestNeighbor([224, 224]) .toFloat(); const prediction = model.predict(tensor.expandDims()); const results = await prediction.data(); // 显示结果 displayResults(results); // 下一帧 requestAnimationFrame(processFrame); } processFrame(); }
  4. Web Components封装

    // 定义自定义元素 class ImageClassifier extends HTMLElement { constructor() { super(); this.attachShadow({ mode: 'open' }); this.model = null; } async connectedCallback() { this.shadowRoot.innerHTML = ` <style> :host { display: block; } #results { margin-top: 10px; } </style> <input type="file" id="file-input" accept="image/*"> <div id="results"></div> `; this.model = await tf.loadLayersModel(this.getAttribute('model-url')); this.shadowRoot.getElementById('file-input') .addEventListener('change', this.handleImage.bind(this)); } async handleImage(event) { const file = event.target.files[0]; const img = await createImageBitmap(file); const tensor = tf.browser.fromPixels(img) .resizeNearestNeighbor([224, 224]) .toFloat(); const prediction = this.model.predict(tensor.expandDims()); const results = await prediction.data(); this.shadowRoot.getElementById('results').textContent = `预测结果: ${JSON.stringify(results)}`; }
← 返回列表