简介一份基于TensorFlow.js框架的人工智能实验项目专为前端开发者设计用于在浏览器中完成深度学习模型的训练与推理演示。项目覆盖神经网络算法与计算机视觉通过交互式界面直观呈现图像识别等任务背后的实现流程也涉及数据准备、网络结构选择、参数调整与性能评估实验代码可帮助理解权重初始化、损失计算与优化器等概念。压缩包共120个文件大小2.69MB代码以JS脚本、JSON配置与TSX组件为主另有HTML交互页面、PNG/GIF可视化素材以及Markdown、Word与文本说明文档结构清晰便于查阅。当前已有73人学习下载。资源内含可直接运行的训练示例、cartpole效果动图及附赠文档新手可依照说明完成从训练到推理的完整流程有经验者也能参考其中的模型结构与浏览器端部署思路快速将AI能力接入自己的前端项目。1. 很多人拿到 TensorFlow.js 实验项目压缩包第一反应是翻错地方很多人拿到「基于 TensorFlow.js 的 JavaScript 人工智能实验项目」这类压缩包第一反应是翻神经网络代码结果每行 JavaScript 都看得懂却不知道张量在做什么。这个标题指向的路径很明确用纯前端技术栈在浏览器里完成深度学习模型的训练与推理并把过程做成可交互演示。它适合的读者是前端开发者不是算法工程师——不用搭 Python 环境、不用理解 CUDA浏览器支持 WebGL 就能训练卷积神经网络再在 Canvas 上画数字让模型实时识别。这类项目把训练、推理、可视化压缩在同一条技术栈里损失曲线怎么变化、卷积核学到了什么都能用 JavaScript 函数直接渲染到页面上哪怕只是拿它当人工智能大作业的起点也够你把整条链路走通。2. TensorFlow.js 的运行机制与前端项目初始化2.1 为什么选 TensorFlow.js模型在哪里训练推理就在哪里发生标题把 TensorFlow.js 和 JavaScript 并列出现说明这个实验项目刻意选了「从前到后只用前端技术」的路线。这类项目解压后通常是一个纯静态站点index.html、js/、models/、data/ 四个部分没有后端服务。常见做法是安装两个 npm 包tensorflow/tfjs 负责张量运算与模型生命周期tensorflow/tfjs-vis 负责把损失值、准确率、混淆矩阵渲染成页面侧边栏的交互图表。选型理由可以拆成三条运行时只依赖 WebGL模型参数的更新和预测都在用户打开的页面里发生不需要 Node.js 服务也就没有跨域和接口鉴权的问题适合做教学演示。训练完的模型是 tf.LayersModel 对象save() 之后得到 model.json 加权重分片文件可以直接扔进静态站点部署推理端不用再装任何深度学习环境。对前端开发者来说深度学习环境配置这一步被完全省掉浏览器本身就是运行时断网也能复现实验。如果只做交互演示、不打算引入打包器用 script 标签方式最快。在 HBuilder 里新建一个静态页面配置好 html、css、javascript 之后把两行 CDN 地址贴进 bodyscript srchttps://cdn.jsdelivr.net/npm/tensorflow/tfjs/dist/tf.min.js/script script srchttps://cdn.jsdelivr.net/npm/tensorflow/tfjs-vis/dist/tf-vis.min.js/script注意 tfjs 的 API 大量使用 ES6 的 Promise 与 async/await调试时先确认浏览器版本别太旧。如果你习惯用 import 方式加载本地模块直接双击 index.html 大概率会报Failed to load module script这类 MIME 错误因为它要求资源必须通过 HTTP 提供用npx serve起一个本地静态服务器最省心。2.2 WebGL、CPU 与 WASM 后端初始化顺序和内存基线tfjs 在浏览器里默认优先选 WebGL 后端但不同设备差异很大。一个稳妥的初始化函数长这样async function initBackend() { if (tf.getBackend() webgl) return; try { await tf.setBackend(webgl); // 请求 WebGL 后端 await tf.ready(); // 等待后端真正可用 } catch (e) { await tf.setBackend(cpu); // 失败回退 CPU console.warn(WebGL 初始化失败已回退到 CPU 后端, e); } const mem tf.memory(); console.log(后端: ${tf.getBackend()}, 张量: ${mem.numTensors}, 字节: ${mem.numBytes}); }setBackend 返回 Promise要在 await 之后调用 tf.ready() 确认后端真正可用再开始创建模型。tf.memory() 输出的 numTensors 是排查内存泄漏的关键基线页面加载完成后记一个数字跑完一次训练或推理再记一次增量大于 0 说明有张量没释放。常见泄漏点是循环里反复调用数据预处理函数代码里我一般把所有临时张量包进 tf.tidy 或手动 dispose。三个后端的取舍区别整理成一张表初始化时按设备能力选择后端启用方式优势限制WebGLtf.setBackend(webgl)GPU 并行训练推理快显存有限批量一大就溢出CPUtf.setBackend(cpu)无显存限制调试直观计算慢大模型容易卡顿WASM额外引入 tfjs-backend-wasm比 CPU 快可多线程要加载 wasm 文件移动端兼容一般对于走人工智能学习路线的前端开发者我建议先强制用 CPU 后端把训练逻辑跑通确认数值和收敛方向没问题再切到 WebGL 对比训练耗时通常能看到一个数量级的差距。切换后端必须在创建模型之前完成否则已分配的张量和编译好的 shader 都要作废重建。3. 用 TensorFlow.js 训练一个卷积神经网络从数据到损失曲线3.1 数据加载与预处理784 个像素怎么变成张量动手深度学习的第一步不是调参而是搞清楚数据怎么进入计算图。以最常见的 MNIST 手写数字为例data/mnist.json 里每张图是 784 个 0 到 255 的像素值labels 是对应数字。加载函数只需要做三件事转成四维张量、归一化、把标签变成 one-hot。async function loadMNIST() { const res await fetch(./data/mnist.json); const raw await res.json(); // 形状 [样本数, 高, 宽, 通道数]并把像素压到 0~1 const xs tf.tensor4d(raw.images, [raw.images.length, 28, 28, 1]) .div(255); // 标签转 one-hot10 对应最后一个 softmax 层的输出维度 const ys tf.oneHot(tf.tensor1d(raw.labels, int32), 10); return { xs, ys }; }tensor4d 的四个维度分别对应 [样本数, 高, 宽, 通道数]MNIST 是灰度图所以通道是 1train 集 6 万张 28×28 的 float32 数据约占 188 MB浏览器里能装下但你已经能感知到 WebGL 显存的天花板。div(255) 把像素压到 0~1 区间卷积层权重初始化通常假设输入在这个尺度跳过去会导致前几层激活值过大收敛明显变慢。tf.oneHot 的第二个参数 10 是类别数和最后一层 softmax 的 units 必须一致这是训练时最常出现的 shape 不匹配来源。3.2 用 sequential 搭一个 CNN参数规模怎么算出来的计算机视觉这类图结构数据神经网络算法里首选卷积网络卷积核的权值共享让参数数量比全连接低一个量级局部感受野则让模型对笔画位置偏移更鲁棒。tfjs 的 layers API 和 Keras 几乎一一对应function createModel() { const model tf.sequential(); model.add(tf.layers.conv2d({ inputShape: [28, 28, 1], // 第一层必须写完整三分量 filters: 32, kernelSize: 3, activation: relu })); model.add(tf.layers.maxPooling2d({ poolSize: 2, strides: 2 })); model.add(tf.layers.conv2d({ filters: 64, kernelSize: 3, activation: relu })); model.add(tf.layers.maxPooling2d({ poolSize: 2, strides: 2 })); model.add(tf.layers.flatten()); model.add(tf.layers.dropout({ rate: 0.2 })); // 抑制过拟合 model.add(tf.layers.dense({ units: 128, activation: relu })); model.add(tf.layers.dense({ units: 10, activation: softmax })); return model; }第一层 inputShape 必须写完整的三分量后续层通过上一层的输出自动推断。第一层卷积 32 个 3×3 卷积核只在输入上扫一遍参数是 3×3×1×32 加 32 个偏置共 320 个经过一次池化后特征图变成 14×14第二层 64 个卷积核的参数变成 3×3×32×64 加 64共 18496 个。第二次池化后是 7×7×64flatten 成 3136 维向量全连接层参数约 40 万占整个模型绝大部分。整个网络约 42 万参数在浏览器里训练没有压力console 里运行 model.summary() 就能核对每一层的数据。dropout 只加在全连接层之前rate 0.2 表示训练时随机丢弃 20% 的神经元tfjs 在 predict 阶段会自动关闭 dropout不需要你手动切换。卷积层的激活函数用 relu 而不是 sigmoid主要是为了规避梯度消失深层网络里这个选择比激活函数本身的非线性更重要。3.3 训练循环与可视化epochs、batchSize 和验证集怎么定训练之前先 compile 指定优化器、损失函数和评估指标然后用 tfvis 的 fitCallbacks 把训练过程画出来这是标题里「交互式演示」的核心部分model.compile({ optimizer: rmsprop, loss: categoricalCrossentropy, metrics: [accuracy] }); const container { name: MNIST 训练, tab: 训练曲线 }; const callbacks tfvis.show.fitCallbacks( container, [loss, val_loss, acc, val_acc], // 四类指标画成曲线 { height: 260, callbacks: [onEpochEnd] } ); await model.fit(xs, ys, { epochs: 5, // 完整遍历训练集 5 遍 batchSize: 128, // 每次梯度更新用 128 个样本 validationSplit: 0.2, // 切 20% 做验证 shuffle: true, callbacks });epoch 指完整遍历训练集一遍epochs: 5 就是让网络把 4.8 万张训练图反复看 5 次。batchSize 决定每次梯度更新用多少样本128 在 WebGL 显存里比较安全低端机型降到 64 更稳。validationSplit 会从训练集末尾切 20% 作为验证集注意提前设置 shuffle 再切否则验证集可能全是同一类数字。fitCallbacks 把四个指标画成四条曲线val_loss 在第三个 epoch 开始回升而 loss 还在降就是过拟合的典型信号优先加大 dropout 或减少 epochs。训练参数的经验值整理如下实验项目里按这个表起步再微调参数推荐值作用与坑epochs5 ~ 10太少欠拟合太多过拟合batchSize64 ~ 128太大容易 WebGL 内存溢出validationSplit0.15 ~ 0.2太小验证曲线波动大shuffletrue不设或设晚验证集分布会偏提示fit 返回的 history 里有每一轮的 loss 和 acc即使不用 tfvis也可以用 canvas 自己把 history.history.loss 画成折线图这和标题里的可视化展示是同一件事的两条实现路径。3.4 训练阶段最常见的两类报错第一类是 WebGL 内存溢出报错关键字一般是OUT_OF_MEMORY或INVALID_VALUE。常见做法是优先把 batchSize 从 128 降到 64同时检查数据预处理代码里有没有把中间张量泄漏在循环外把整段预处理包进 tf.tidy 是最省事的修法。第二类是 shape 不匹配报错会给出期望维度。典型场景是 one-hot 类别数写成 9 或者 12和最后一层 softmax 的 units 不一致另外注意把 validationSplit 切出来的标签也检查一遍维度fit 报错信息里的 hidden state 提示一般在 console 里能直接定位到是哪一行 createModel 出的问题。4. 交互式推理与计算机视觉演示的实现4.1 Canvas 手写输入到张量一条必须写对的转换管线训练只是前半程把模型的推理能力真正 harness 进浏览器交互才是标题里「交互式演示」的重头戏。手写数字识别最常见的做法是在 Canvas 上用手指或鼠标写字松手后把画布内容转成模型能吃的张量。这中间每一小步出错推理结果都会变得离谱function canvasToTensor(canvas) { const imageData canvas.getContext(2d).getImageData(0, 0, 280, 280); return tf.tidy(() { return tf.browser.fromPixels(imageData, 1) // RGBA 转灰度 .resizeBilinear([28, 28]) // 双线性缩放笔画不断裂 .div(255) // 归一化 .expandDims(0); // 加批次维 [1,28,28,1] }); }转换管线分四步getImageData 拿到整个画布的 RGBA 像素fromPixels 的第二个参数 numChannels1 表示转成单通道灰度这一步会把 RGBA 按亮度公式合掉resizeBilinear 缩到 28×28这里选双线性而不是最近邻是为了让缩小后的笔画边缘平滑避免出现断笔最后 expandDims 在索引 0 的位置插入一维把形状从 [28, 28, 1] 变成 [1, 28, 28, 1]因为 predict 永远接收带批次维的输入。整段代码包在 tf.tidy 里中间产生的临时张量会在函数返回时自动释放。提示fromPixels 的输入可以是 HTMLCanvasElement、ImageData 或 video 元素摄像头实时分析也是同一条管线把 canvas 换成视频帧即可。4.2 推理代码与置信度可视化top-5 是怎么来的predict 返回的 logits 经过最后一层 softmax 之后已经是十个类别的概率分布。把概率降序排序取前五个画成柱状图用户就能直观看到模型在「纠结」哪些数字const input canvasToTensor(canvas); const logits model.predict(input, { batchSize: 1 }); // 固定批次省一次 shader 编译 const probs Array.from(await logits.data()); const topK [...probs.entries()] .sort((a, b) b[1] - a[1]) .slice(0, 5); // 取概率最高的 5 个类别 tfvis.render.barchart( { name: Top-5 预测结果, tab: 推理 }, topK.map(([label, prob]) ({ index: label, value: prob })) ); input.dispose(); logits.dispose();predict 的第二个参数传 { batchSize: 1 }是让 WebGL 后端不必再做动态批次推导省一次 shader 重编译对交互场景首帧延迟影响很明显。logits.data() 返回 Promiseawait 之后得到 Float32Arrayspread 成普通数组再排序。最后两行 dispose 不能省尤其是画布连续触发识别的场景每个事件都泄漏两个张量几次之后页面就开始卡。置信度分布的另一个用途是检查模型缺陷如果模型经常把 9 认成 4并且这两个类别的概率都接近 0.5说明训练数据里这两个数字的形态太接近而不是模型坏了。4.3 迁移学习让实验项目具备真正的计算机视觉能力如果你想演示的不是 MNIST 而是真实摄像头场景——比如手势分类、猫狗识别这类任务——从头训练一个卷积网络在浏览器里代价太高。常见做法是加载一个已经在 ImageNet 上预训练好的 MobileNet冻结主干只在最后一层换成自己的分类头这就是迁移学习const baseModel await tf.loadLayersModel(./models/mobilenet/model.json); // 取倒数第二层的输出作为特征向量最后一层是面向千类的 softmax const featureLayer baseModel.layers[baseModel.layers.length - 2].output; const featureExtractor tf.model({ inputs: baseModel.inputs, outputs: featureLayer }); featureExtractor.trainable false; // 冻结主干权重 const head tf.sequential(); head.add(tf.layers.dense({ inputShape: [featureExtractor.outputs[0].shape[1]], // 特征维度 units: 4, // 换成自己的类别数 activation: softmax })); const model tf.model({ inputs: featureExtractor.inputs, outputs: head.apply(featureExtractor.outputs[0]) }); model.compile({ optimizer: tf.train.adam(0.0001), // 迁移学习用较小学习率 loss: categoricalCrossentropy, metrics: [accuracy] });取倒数第二层的输出作为特征向量是因为预训练模型的最后一层是面向 ImageNet 千分类的 softmax直接截断后只保留前面几百维的通用视觉特征。featureExtractor.trainable false 让主干权重在训练中不再更新梯度只回流到新加的 head。学习率用 0.0001 而不是默认的 0.001因为底座特征已经很好了学习率太大会在新分类头上反复震荡。从头训练和迁移学习的选择可以用这张表对比维度从头训练迁移学习数据量需求每类数千张起步每类几十张可用训练时间分钟到小时级几秒到几十秒前端可行性受 WebGL 显存限制主干冻结显存占用小适用场景数据形态特殊通用视觉特征够用的任务5. 模型导出、体积压缩与推理性能的三个实测技巧5.1 用 tensorflowjs_converter 把训练好的模型转成网页格式实验中如果先在 Node.js 或 Python 里训好了模型要导到前端项目常见做法是用官方转换器pip install tensorflowjs tensorflowjs_converter \ --input_format keras \ --output_format tfjs_layers_model \ --quantization_bytes 1 \ ./mnist_model.h5 ./web_modelquantization_bytes 1 表示把权重量化到 8 bit模型文件体积大约缩小到原来的四分之一精度损失通常不到 0.5%浏览器的内存占用也会同步下降。转换完成后 web_model 目录里是 model.json 加若干 .bin 分片浏览器端用 tf.loadLayersModel(./web_model/model.json) 就能加载不需要任何后端。5.2 用 IndexedDB 缓存模型避免每次刷新都走网络模型体积压缩完之后还有一次加载成本。把第一次下载的模型缓存到 IndexedDB后续刷新直接读本地async function loadModelWithCache(remotePath) { const cacheKey indexeddb://cached-web-model; // 换模型时改这个键 try { return await tf.loadLayersModel(cacheKey); } catch (e) { const model await tf.loadLayersModel(remotePath); await model.save(cacheKey); return model; } }第一次访问走网络save 会把模型权重和结构写进 IndexedDB第二次 loadLayersModel 遇到 indexeddb:// 前缀就直接读缓存。注意缓存键要带版本号模型更新后改成新的键否则一直加载旧权重。5.3 一条命令量化推理延迟与内存泄漏交互式演示项目验收时我一般会在控制台跑一段压测脚本同时拿到两个关键指标const before tf.memory().numTensors; const input canvasToTensor(canvas); const results []; for (let i 0; i 20; i) { const t0 performance.now(); const out model.predict(input, { batchSize: 1 }); out.dataSync(); // 强制同步确保计算完成 results.push(performance.now() - t0); out.dispose(); } input.dispose(); const avg results.reduce((a, b) a b, 0) / results.length; const leaked tf.memory().numTensors - before; // 应等于 0 console.log(平均推理延迟: ${avg.toFixed(2)}ms, 泄漏张量: ${leaked});延迟取的是连续 20 次预测的平均值第一次 predict 因为要编译 shader 会明显偏慢所以循环前先 predict 一次做 warmup 更准确。泄漏张量如果是 0说明推理路径上每个张量都释放干净了大于 0 就回到前面提到的 tf.tidy 和 dispose 检查。平均延迟超过 50ms 时优先检查画布绘制分辨率是不是太高把绘制区从 1920 降到 280 之后getImageData 和 resize 的开销会立刻降下来。本文还有配套的精品资源点击获取