1. JavaScript与深度学习的奇妙结合
当大多数人听到"深度学习"这个词时,脑海中浮现的往往是Python、TensorFlow或PyTorch这些传统工具。但今天我要告诉你一个可能让你惊讶的事实:JavaScript也能成为深度学习的强大工具。作为一名在Web开发和机器学习交叉领域工作多年的工程师,我见证了JavaScript生态系统中深度学习工具的惊人成长。
你可能会有疑问:为什么要在JavaScript中做深度学习?答案很简单——因为Web无处不在。想象一下,直接在浏览器中运行图像识别模型而不需要服务器支持,或者在移动设备上离线执行自然语言处理任务。这些场景正是JavaScript深度学习的独特优势所在。
我清楚地记得第一次用TensorFlow.js完成一个MNIST手写数字识别项目时的震撼——不需要任何复杂的后端部署,模型直接在浏览器中训练和推理,整个过程流畅得令人难以置信。从那时起,我就迷上了用JavaScript探索深度学习的可能性。
2. JavaScript深度学习生态全景
2.1 核心工具库介绍
JavaScript深度学习生态系统已经相当丰富,以下是最主流的几个工具:
TensorFlow.js:Google推出的JavaScript版TensorFlow,支持在浏览器和Node.js环境中运行
- 提供与Python版相似的高级API
- 支持WebGL加速
- 可以直接加载预训练的Python模型
Brain.js:专注于神经网络的轻量级库
- 简单易用的API
- 支持多种网络类型(前馈、RNN、LSTM等)
- 非常适合快速原型开发
ML5.js:基于TensorFlow.js的高级封装
- 提供预训练模型(图像分类、姿态检测等)
- 设计理念是让机器学习对创意编码者更友好
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-gpuHTML基础模板(浏览器环境):
<!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();代码解析:
generateData函数创建了一个带有噪声的三次多项式数据集。这里使用了tf.tidy来确保中间张量被正确清理。createModel定义了一个简单的线性回归模型,虽然我们的数据是非线性的,但线性模型仍然可以学习到一定的趋势。trainModel函数展示了基本的训练流程,包括回调函数的使用来监控训练进度。注意所有操作都是异步的,因为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); }关键点说明:
我们直接从Google的服务器加载了MobileNet模型,这是一个在ImageNet数据集上预训练的卷积神经网络。
图像预处理步骤非常重要,必须与模型训练时使用的预处理方式一致(这里是ImageNet的标准归一化)。
tf.browser.fromPixels是TensorFlow.js提供的专门用于处理浏览器图像元素的API。预测结果包含了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] // }); }迁移学习的关键步骤:
加载预训练模型并截断,保留特征提取部分。
冻结基础模型的权重,防止它们在训练过程中被修改。
添加新的可训练层来适应你的特定任务。
使用较小的学习率训练,因为特征已经相对较好,只需要微调。
实战技巧:对于小型数据集,可以尝试不同的数据增强技术(旋转、翻转、颜色调整等)来提高模型泛化能力。TensorFlow.js提供了
tf.image命名空间下的多种图像处理操作。
5. 性能优化与生产部署
5.1 性能优化技巧
JavaScript深度学习面临的最大挑战之一是性能。以下是一些经过验证的优化技巧:
内存管理:
- 使用
tf.tidy()自动清理中间张量 - 手动调用
tensor.dispose()释放不再需要的张量 - 监控内存使用:
tf.memory()可以打印当前内存状态
- 使用
批量处理:
- 尽量使用批量操作而不是循环处理单个样本
- 例如,使用
model.predictOnBatch()而不是循环调用model.predict()
WebGL优化:
- 确保张量形状在连续操作中保持一致,减少WebGL着色器重新编译
- 对于小型张量,CPU可能比WebGL更快(可以使用
tf.setBackend('cpu')测试)
模型量化:
- 使用16位浮点或8位整数权重减小模型大小
- TensorFlow.js支持从量化的TensorFlow Lite模型转换
5.2 生产部署策略
将JavaScript深度学习模型部署到生产环境需要考虑几个关键因素:
部署选项对比:
| 部署方式 | 优点 | 缺点 | 适用场景 |
|---|---|---|---|
| 纯客户端 | 无服务器成本,隐私友好 | 受限于设备性能,模型大小受限 | 轻量级应用,注重隐私 |
| 客户端+服务器 | 平衡计算负载,支持更大模型 | 需要服务器基础设施 | 大多数生产应用 |
| WebAssembly | 性能接近原生,支持复杂模型 | 加载时间较长,兼容性问题 | 性能关键型应用 |
| 边缘计算 | 低延迟,减少网络传输 | 部署复杂度高 | IoT、实时处理应用 |
| CDN缓存模型 | 快速加载,减轻服务器压力 | 模型更新有延迟 | 静态模型,不频繁更新 |
模型优化建议:
模型剪枝:移除对输出影响较小的神经元,减小模型大小
// 使用TensorFlow.js的模型转换工具进行剪枝 const prunedModel = await tf.graphModel.convertToPruned(model, pruningConfig);模型量化:降低权重精度
// 转换为16位浮点 const quantizedModel = await model.save(tf.io.withSaveHandler(async (artifacts) => { return { modelTopology: artifacts.modelTopology, weightSpecs: artifacts.weightSpecs, weightData: new Uint16Array(artifacts.weightData.buffer) }; }));模型分割:将大模型分成多个部分,按需加载
// 先加载模型骨架 const modelPart1 = await tf.loadLayersModel('model_part1.json'); // 用户交互后再加载剩余部分 button.onclick = async () => { const modelPart2 = await tf.loadLayersModel('model_part2.json'); // 组合模型... };渐进式加载:先加载轻量级模型,后台下载完整模型
// 先加载精简版 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;
监控与维护:
实现模型性能监控:
// 记录推理时间 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 }) });建立模型版本控制:
// 检查模型版本 async function checkModelVersion() { const response = await fetch('/model-version'); const { latestVersion } = await response.json(); if (localStorage.getItem('modelVersion') !== latestVersion) { console.log('新模型版本可用'); // 触发模型更新流程 } }实现回退机制:
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深度学习开发中,有几个常见陷阱需要注意:
内存泄漏:
- 症状:浏览器标签内存使用量持续增长,最终崩溃
- 原因:未正确释放张量内存
- 解决方案:
// 错误方式:循环创建未释放的张量 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]); // 自动清理 }); }
WebGL上下文丢失:
- 症状:模型突然停止工作,控制台报WebGL错误
- 原因:浏览器回收了WebGL资源(如切换标签页后返回)
- 解决方案:
// 监听上下文丢失事件 const canvas = document.getElementById('webgl-canvas'); canvas.addEventListener('webglcontextlost', (e) => { e.preventDefault(); console.log('WebGL上下文丢失,需要重新初始化模型'); // 重新加载模型 });
模型加载失败:
- 症状:模型文件加载超时或返回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'); }
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 调试工具与技术
有效的调试可以节省大量开发时间:
TensorFlow.js自带工具:
// 启用调试模式 tf.enableDebugMode(); // 检查张量值(注意:会阻塞执行) const tensor = tf.tensor([1, 2, 3]); tensor.print(); // 打印张量内容 // 内存状态 console.log(tf.memory());浏览器性能分析:
- 使用Chrome DevTools的Performance面板记录模型训练过程
- 检查主要耗时操作和内存分配情况
- 注意WebGL调用和着色器编译时间
模型可视化:
// 在浏览器中显示模型结构 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); }); }训练过程监控:
// 使用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 });自定义日志回调:
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的支持程度不同,可能导致性能差异甚至功能失效:
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; }浏览器特定问题:
- Safari:对WebGL 2.0支持有限,可能需要polyfill
- iOS:有严格的内存限制,大模型容易崩溃
- Firefox:WebGL性能通常较好,但可能有不同的精度行为
优雅降级策略:
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'); } } }性能基准测试:
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 模型压缩与量化新技术
最新的模型压缩技术可以让深度学习模型更适应浏览器环境:
知识蒸馏:训练小型"学生"模型模仿大型"教师"模型的行为
// 伪代码示例 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 }); }结构化剪枝:移除整个神经元或通道而不仅仅是单个权重
// 使用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 ));量化感知训练:在训练过程中模拟量化效果,提高最终量化模型的精度
// 伪代码示例 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技术深度融合:
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);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 }); };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(); }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)}`; }