ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

TensorFlow.js 浏览器端线性回归实战:从零到模型预测的完整指南

TensorFlow.js 浏览器端线性回归实战:从零到模型预测的完整指南 最近有人问我前端工程师想入门机器学习应该从哪里下手。我给的答复一直是同一个用 TensorFlow.js 在浏览器里跑一个线性回归模型。这事儿听起来简单但麻雀虽小五脏俱全数据生成、模型定义、训练、预测、评估全流程都有而且全部在浏览器里完成不需要装 Python 环境不需要 GPU 服务器一个 Chrome 标签页就够了。我自己带新人做前端智能化方向时第一课就是这个。这篇文章就完整记录一下这套代码示例的落地过程包括每一步为什么要这么写、参数该怎么调、以及我在实测里踩过的坑。1. 整体设计为什么是线性回归为什么在浏览器里跑1.1 浏览器端机器学习的现实价值TensorFlow.js 最吸引人的点不是它能跑多复杂的模型而是它把机器学习的门槛降到了非常低的位置。用户不需要安装任何东西打开网页就能体验模型训练和推理数据不用上传服务器隐私性好响应也快。训练好的模型可以直接在网页里做实时预测比如根据用户输入预估价格、预判流失概率、分析传感器数据等等。在实际项目中浏览器端训练不会用来跑大模型更多是两类场景一类是轻量级个性化模型比如根据用户的历史行为做本地化预测另一类是学习与演示像这篇例子里做的线性回归本质上是一个完整的教学闭环。因为用户零安装、零配置你发一个链接过去对方在手机浏览器里也能看到训练过程和结果这比让用户去装 Anaconda 友好太多了。1.2 线性回归作为第一个实验的合理性我见过不少人一上来就想用 TensorFlow.js 做图片分类、做目标检测结果被复杂的网络结构、数据增强、迁移学习这些东西直接劝退。线性回归作为入门实验好处是它足够简单线性层就是神经网络最基础的计算单元理解了它后续再学 Dense 层堆叠、多层感知机就非常顺。具体来说线性回归要完成的任务是拟合一个形如 y kx b 的关系。数据只有一维特征模型只需要一个全连接层损失函数是均方误差优化器是随机梯度下降。这套组合在 TensorFlow.js 里代码量很小但机器学习里最核心的那几个概念——数据张量化、模型编译、迭代训练、损失下降、模型预测——全部覆盖到了。把这一步跑通你就拥有了排错和阅读更复杂模型代码的基础能力。2. 环境准备把 TensorFlow.js 跑起来2.1 两种引入方式CDN 与 npmTensorFlow.js 的引入方式和普通前端库没有区别CDN 适合写演示页、CodePen、本地单文件测试npm 适合正经的前端工程项目。CDN 方式直接在 HTML 里加script srchttps://cdn.jsdelivr.net/npm/tensorflow/tfjs4.20.0/dist/tf.min.js/script这样打开页面后浏览器全局就会挂一个 tf 对象后面的代码直接用 tf.tensor、tf.sequential 这些 API 就行了。优点是零构建步骤双击 HTML 文件就能跑。缺点是全局变量污染、版本升级不便、离线环境没法用。npm 方式适合工程化项目npm install tensorflow/tfjs然后在你需要的地方引入import * as tf from tensorflow/tfjs;两者本质上是同一套 API只是模块化与打包方式不同。我个人的习惯是演示和教学用 CDN因为省事学生拿到文件就能跑实际项目用 npm配合 Vite 或 Webpack 做构建还能利用 tree-shaking 减小体积。有一点需要注意浏览器版默认用的是 WebGL 后端如果你的机器不支持 WebGLTensorFlow.js 会退回 CPU 后端后面我会专门讲这个问题。2.2 浏览器兼容性与性能基线这套东西在 Chrome、Edge、Firefox 和 Safari 上都能跑但在不同浏览器上的表现有差异。我最常用的环境是 Chrome因为它的 WebGL 实现最成熟训练速度相对也最快。浏览器WebGL 后端训练表现备注Chrome稳定最快开发调试首选DevTools 可查看 WebGL 状态Edge稳定与 Chrome 接近同为 Chromium 内核Firefox基本可用略慢某些版本会出现 WebGL 上下文丢失Safari可用不稳定长时间训练容易触达内存限制对于线性回归这种极小模型CPU 后端其实也完全够用训练时间也就是几百毫秒到几秒的事。真正需要在意的是老设备的 WebGL 兼容问题建议在生产环境里做后端自动检测await tf.setBackend(webgl);如果报错就回退到 CPUconst result await tf.setBackend(cpu);3. 数据生成先造一套能复现的训练集3.1 合成数据与噪声设计训练机器学习模型的第一步是拿到数据。真实场景里你可能要采集几十上百条样本但作为代码示例最方便的做法是人工合成一套带噪声的数据。我们定义一个真实的关系y 2x 1。然后在这个关系的基础上加一些随机噪声让数据看起来不那么完美。为什么要加噪声因为如果你给模型喂的数据是绝对精确的直线那么模型学到的只是机械的公式你无法观察到损失曲线的实际行为。加一点噪声模型才需要真正“学”。function generateData(numPoints) { const xs []; const ys []; for (let i 0; i numPoints; i) { const x (Math.random() * 2 - 1); // x 范围 [-1, 1] const noise (Math.random() - 0.5) * 0.3; const y 2 * x 1 noise; xs.push(x); ys.push(y); } return { xs, ys }; }这里把 x 的范围限制在 [-1, 1] 是刻意的。如果不限制范围x 散布到几万几十万模型的输入尺度会变得很大训练时梯度容易不稳定。后面模型要预测的 x 值也最好别超出这个范围太远这是训练和预测数据分布一致性原则非常重要。如果你看了实际数据会发现它们大体排成一条斜线但每个点都偏离了真实直线一点这就是噪声的效果。噪声的幅度我控制在 ±0.15相对于 y 的真实变化范围-1 到 3来说信噪比还是不错的模型能比较容易地学到 2 和 1 这两个参数。3.2 张量转换从普通数组到 TensorTensorFlow.js 的核心数据结构是 Tensor跟 JavaScript 里的普通数组不是一回事。Tensor 可以理解为一个多维数组外加一套自动微分系统所有梯度计算都建立在 Tensor 上。所以数据生成出来之后必须转换为 Tensor 才能参与后续的运算。const xsTensor tf.tensor2d(xs, [xs.length, 1]); const ysTensor tf.tensor2d(ys, [ys.length, 1]);这里 tensor2d 的第二个参数是 shape[numPoints, 1] 表示总共 numPoints 行、每行 1 列。为什么一定要是二维因为 TensorFlow.js 的 Dense 层默认接收的输入是 [batchSize, inputFeatures]也就是一批样本每个样本有 inputFeatures 个特征。线性回归只有一个输入特征所以每个样本可以表示成一个长度为 1 的数组batch 维度由样本数量承担。转换成 Tensor 之后数据不再是我们熟悉的数组了你没法直接用 console.log 把里面的值打出来要这样取const values await xsTensor.array(); // 转为普通数组 // 或者 const values xsTensor.dataSync(); // 同步取出 Float32Array这里我特别想提醒Tensor 是占用 WebGL 显存的如果你的代码里创建了大量 Tensor 没有释放显存就会越占越多最终导致页面崩溃。处理方式是调用 .dispose() 方法释放或者用 tf.tidy 自动管理临时 Tensor。这个后面在训练部分细说。3.3 训练集与测试集的拆分原则模型在训练时见过的数据叫训练集没见过的数据叫测试集。我在这套示例里把 200 个样本中的 170 个用于训练30 个用于测试。测试集的作用不是参与参数更新而是在训练结束后检查模型泛化能力——如果模型在训练集上误差很小、在测试集上误差很大那就是过拟合了说明模型只是死记硬背了训练样本没有学会背后的规律。const trainX xsTensor.slice([0, 0], [trainCount, 1]); const trainY ysTensor.slice([0, 0], [trainCount, 1]); const testX xsTensor.slice([trainCount, 0], [testCount, 1]); const testY ysTensor.slice([trainCount, 0], [testCount, 1]);slice 方法第一个参数是起始位置 [startRow, startCol]第二个参数是切出的尺寸 [rows, cols]。用前 170 个样本做训练后 30 个做测试这种顺序切分在数据本身就是随机生成的前提下没问题。真实项目里数据如果存在时间顺序或类别分布不均就要用随机打乱的方式切分不能简单截取。4. 模型定义用几行代码看懂 Dense 层4.1 Sequential 模型与 Dense 层参数TensorFlow.js 提供了两种模型定义方式Sequential顺序模型和 Functional函数式模型。线性回归用 Sequential 就足够了它像搭积木一样一层一层往上加前一层的输出自动作为后一层的输入。const model tf.sequential(); model.add(tf.layers.dense({ inputShape: [1], units: 1, useBias: true }));逐行解释tf.sequential() 创建一个顺序模型容器tf.layers.dense 是全连接层在线性回归这里就相当于一个线性变换层。units: 1 表示这一层有 1 个神经元inputShape: [1] 表示每个样本输入是 1 个特征useBias: true 表示要学习偏置项 b。所以模型要学的参数就两个输入到输出的权重 k以及偏置 b。模型内部的数学表达式就是 output k * input b这正是我们生成数据用的关系式。这个模型里可训练参数的数量是 2我们可以通过下面的代码确认console.log(model.summary());4.2 优化器与损失函数的选择逻辑模型定义好之后还需要指定优化器和损失函数这一步在 TensorFlow.js 里叫做编译compilemodel.compile({ optimizer: tf.train.sgd(0.1), loss: tf.losses.meanSquaredError });优化器是模型学习的方式。sgd 是随机梯度下降Stochastic Gradient Descent它的核心思路是每次随机选一小批样本计算梯度然后朝着让损失减小的方向更新参数。0.1 是学习率也就是每次参数更新的步长。学习率的取值非常有讲究。太大模型可能来回震荡永远不收敛太小训练几百轮损失还是降不下来。0.1 对于这个简单线性模型来说是个比较合适的起点。很多教程会把学习率藏起来不讲但它是调参时第一个该动的旋钮。我把 0.01、0.1、0.3 三个值都测过0.01 需要更多轮次才能收敛0.3 损失曲线已经出现明显震荡0.1 是稳定且快速的选择。损失函数是预测值与真实值的差距度量。均方误差MSE是回归问题最常用的损失函数它先计算每个样本预测值和真实值差的平方然后取平均。平方的作用是放大大误差、抵消正负误差方向上的抵消这样模型会优先优化误差较大的样本。理论上线性回归还可以用平均绝对误差MAE但 MAE 在误差为零处不可导对梯度下降不友好所以 MSE 更常用。4.3 编译时发生了什么很多初学者会以为 compile 只是做个配置登记其实它内部做了大量初始化工作为每个可训练参数创建对应的变量Variable初始化梯度记录器构建优化器的内部状态。理论上不调用 compile 也能继续训练但模型不知道用什么方式更新参数也不知道怎么算误差训练过程就是无头苍蝇。所以 compile 这一步是必需的顺序也不能乱必须在 model 定义完成后、调用 fit 之前。还有一个容易忽视的细节compile 里指定 optimizer 时既可以传优化器的字符串名字也可以传优化器实例。建议尽量传实例因为这样你可以控制学习率const rmsprop tf.train.rmsprop(0.01); const adam tf.train.adam(0.01);对线性回归这种凸优化问题SGD 已经足够好了不需要动用 Adam 这类自适应优化器。但如果你把这个例子扩展成多层神经网络把优化器换成 adam 往往能让训练更稳。5. 训练循环从 fit 到实时监控5.1 fit 参数逐个拆解训练是整套流程里最关键的一步。TensorFlow.js 的 fit 方法几乎把训练循环都封装好了我们只需要传入训练数据和一些超参数。async function trainModel() { const history await model.fit(trainX, trainY, { epochs: 100, batchSize: 32, shuffle: true, callbacks: { onEpochEnd: (epoch, logs) { console.log(Epoch ${epoch}: loss ${logs.loss.toFixed(4)}); } } }); }epochs 是全量数据被完整遍历的次数。100 轮对这个模型来说已经很多了通常 30 到 50 轮就能看到损失收敛到很小的值。batchSize 是指每一次参数更新使用多少个样本。这里 32 表示每次从训练集里随机取 32 个样本算平均梯度后更新一次参数。170 个训练样本一个 epoch 就会分成约 6 个 batch每次 batch 更新一次参数所以 100 个 epoch 实际上发生了约 600 次参数更新。shuffle: true 表示每个 epoch 开始前都会打乱样本顺序。这能避免模型学到样本顺序里的偶然规律比如连续几个样本恰好都偏向高值模型就会被带偏。对梯度下降来说样本顺序确实是会影响的小批量梯度下降比全批次梯度下降多了随机性合理的打乱能让梯度方向更接近真实的方向。5.2 训练过程监控与损失期望fit 返回一个 history 对象里面记录每一轮结束时的损失值。如果你观察训练过程中的 loss会看到一个快速下降然后逐渐平台化的曲线。刚启动时模型参数是随机初始化的预测结果和真实值差距很大loss 可能在 2 到 4 之间。随着训练推进loss 会一路下降。第一个 epoch 结束我实测的 loss 通常在 1 到 2 之间到第 20 轮loss 会在 0.1 以下到第 60 轮以后基本稳定在 0.01 上下。那 0.01 是个什么概念呢MSE 是误差平方的平均0.01 意味着平均每个样本的预测误差大约是 0.1。考虑到我们加的噪声幅度就是 ±0.15这个误差水平已经非常接近理论最优了因为即使是完美模型也无法消除数据本身的噪声。在训练过程中实时监控是很重要的。我做了一个简单的页面把每一轮的 loss 绘制成折线图放在网页侧边栏这样训练是否正常一眼就能看出来。如果 loss 不降反升大概率是学习率太大或者数据归一化出了问题。5.3 浏览器训练中的异步与内存管理TensorFlow.js 的 fit 是异步的返回一个 Promise。这意味着训练不会阻塞浏览器主线程页面在训练过程中是不会卡死的。但也带来了一个开发习惯问题如果你不 await 这个 Promise代码会继续往下执行此时模型可能还没训练好你直接调用 predict 就会得到不靠谱的结果。await model.fit(...); const prediction model.predict(testX);从代码执行的先后顺序来说必须先等 fit 完成才能做预测。这就是 async/await 存在的意义。另外一个非常关键的问题是内存管理。训练过程中会创建大量中间 Tensor包括每轮的梯度、每个 batch 的特征张量等。TensorFlow.js 提供了两个手段来管理内存tf.tidy(() { const input tf.tensor2d([0.5], [1, 1]); const output model.predict(input); // 在这个回调中创建的 Tensor 会被自动清理 return output; });在这个示例中fit 函数内部会自动管理大部分中间 Tensor但你自己创建的数据张量、预测时要手动注意。一个最常见的错误是循环里反复创建 Tensor 而不做 dispose这会让你的网页在长时间训练或多次训练后越来越卡最终浏览器标签页崩溃。我在写这个示例的时候训练完成后会立刻释放原始数据的 TensortrainX.dispose(); trainY.dispose(); testX.dispose(); testY.dispose();如果你的模型还需要继续用测试集做评估那评估完再释放也不迟。理解了这个逻辑你在浏览器里跑再多次训练都不会怕内存泄漏。6. 预测与评估模型到底学会了没有6.1 predict 方法的输入输出格式训练完成后模型的使用就比较简单了。predict 是同步方法输入一个 Tensor输出一个预测 Tensor。const inputTensor tf.tensor2d([0.5], [1, 1]); const outputTensor model.predict(inputTensor); const result outputTensor.dataSync(); console.log(x0.5 时预测 y ${result[0]});这里有两处要注意。第一是输入数据的 shape我们定义模型时 inputShape 是 [1]单个样本就是一个长度为 1 的数组即使只有一个样本也要包一层 batch 维度变成 [1, 1]也就是 1 行 1 列。第二是 dataSync() 会阻塞主线程去同步取回数据在训练过程里不建议用但在预测这种一次性操作里问题不大。如果模型训练得足够好输入 0.5预测结果应该接近真实关系 y 2*0.5 1 2再加上噪声的波动应该在 2 附近。我实测多次预测值通常在 1.95 到 2.05 之间浮动非常接近。6.2 用测试集评估模型性能单点预测只能说明模型不是完全没学到东西更严谨的做法是在整个测试集上计算预测误差。const predictions model.predict(testX); const mse tf.metrics.meanSquaredError(testY, predictions); const mseValue mse.dataSync()[0]; console.log(测试集 MSE${mseValue.toFixed(4)});注意 tf.metrics.meanSquaredError 这几个 API 的写法第一个参数是真实标签第二个参数是预测结果顺序不能搞反。结果应该和训练结束时的 loss 很接近因为测试集和训练集的数据分布相同。我实际跑的时候测试集 MSE 大约在 0.011 到 0.018 之间。这个数字比训练集的略大一点属于正常现象毕竟测试集是模型没见过的新样本表现总会稍差一点。如果测试集 MSE 明显大于训练集比如大出几倍那就要警惕过拟合了但因为线性回归本身的容量很小几乎不会出现严重过拟合。6.3 可视化把拟合效果画出来训练好的模型用文字描述太抽象我习惯把原始数据点和模型拟合出的直线画在一起一眼就能看出效果好坏。纯前端可以用 Canvas 手绘也可以直接叠在图表里。这里分享一个简单的思路先生成一组均匀分布的 x 值然后输入模型得到预测 y最后把点画成一条线。const xsForLine []; for (let i -1; i 1; i 0.05) { xsForLine.push(i); } const lineTensor tf.tensor2d(xsForLine, [xsForLine.length, 1]); const ysPred model.predict(lineTensor); const ysData await ysPred.array();拿到这些预测点之后在 Canvas 里把训练数据点绘制为散点把预测点绘制为折线。你会发现无论噪声怎么分布拟合出的直线斜率都在 2 附近、截距都在 1 附近这就是模型学到的规律。我自己还有一个额外的验证方式直接把模型学到的权重打出来。const weights model.getLayer(null, 0).getWeights(); const k weights[0].dataSync()[0]; const b weights[1].dataSync()[0]; console.log(学到的参数k${k.toFixed(3)}, b${b.toFixed(3)});getWeights 返回一个数组第一个元素是权重矩阵第二个元素是偏置向量。我实测多次k 通常在 1.95 到 2.05 之间b 通常在 0.93 到 1.06 之间。看到那两个熟悉的数字出现在眼前你会真切地感到机器学习不是黑魔法它就是通过数据把未知参数逐步逼近真实值的过程。7. 常见问题与排查实录7.1 浏览器差异导致的训练环境问题在浏览器里做深度学习最大的不确定因素就是浏览器本身。我的默认开发浏览器是 Chrome调试时用 DevTools 的 Performance 面板查看训练过程对性能的影响。Chrome 的 WebGL 后端一般不会有问题但遇到“Error: WebGL is not supported”之类的报错时禁用 WebGL 会导致显存不可用TensorFlow.js 会自动回退到 CPU 后端。最典型的坑出现在 Safari 上。长时间训练或多次创建模型后Safari 偶尔会报内存压力错误。原因是 Safari 的 WebGL 显存管理策略和 Chrome 不同有些已经释放的 Tensor 没能及时回收。如果你需要通过网页把训练过程分享给一大群人建议在代码里对后端做一层降级判断同时限制训练轮次别在移动端 Safari 上跑 500 个 epoch。我还遇到过一个比较极端的例子某些浏览器的无痕模式默认禁用了 WebGL。用户在无痕窗口打开你的训练页面直接报错。你可以在页面加载时加一段提示告诉用户如果训练失败请检查浏览器设置或换用普通模式。7.2 训练不收敛时该怎么办训练不收敛是初学者最容易碰到的瓶颈现象是 loss 在训练若干轮之后依然没有明显下降或者下降到一个平台就再也下不去了。排查顺序我一般是这样先看数据范围是否合理再看学习率是否合适最后看模型结构是否有问题。数据范围的问题最隐蔽。如果你生成数据时 x 的范围是 [0, 1000]而 y 的范围是 [1, 2000]这两个数字的量级不一致梯度更新时某些参数会比其他参数更新得快得多训练就会不稳定。解决办法是把数据归一化到 [-1, 1] 或 [0, 1] 区间。我在这篇文章里一开始就把 x 限制在 [-1, 1]就是为了避免这类问题。学习率的问题最好排查。如果 loss 一路飙升或剧烈震荡说明学习率太大如果 loss 下降非常缓慢说明学习率太小。我会用对数尺度的学习率去尝试0.1 不行就 0.01再不行就 0.001很快能找到合适的量级。7.3 多次训练导致的内存与状态问题这个示例做出来之后交互面板上我加了一个“重新训练”按钮方便测试不同超参数。结果发现一个问题点击按钮多次之后页面越来越卡偶尔还会报“Cannot read properties of undefined”的错误。排查下来的结论是问题不在 fit 本身而在我创建的新模型和数据 Tensor 没有清理。旧模型仍然占用显存旧数据的 Tensor 也没有释放。修复方式很简单每次重新训练前model.dispose(); trainX.dispose(); trainY.dispose();如果你为了让模型可以反复预测保留模型不释放那也要注意另一个问题每次 fit 之前旧的优化器状态可能还在。虽然 TensorFlow.js 的 fit 会在新训练中重置优化器状态但为了保险我建议每次重新训练时直接创建一个新模型而不是反复调用旧的 fit。这能确保训练状态完全干净避免之前训练的影响残留到新一轮。7.4 开发者工具里的那些隐藏报错有时候训练页面没有报错但你在 DevTools 的 Console 里会看到 TensorFlow.js 打出的各种警告。最常见的是 tf.tensor() 被用在非 tidy 环境下创建张量它会提示你使用 tf.tidy 包裹或者手动 dispose。还有一类是 “WARNING: The WebGL backend is WEBGL_DISABLED”这说明浏览器不支持 WebGL 或已被禁用。这些警告不一定会导致你的页面崩溃但它们是潜在风险的信号。我处理这些警告的原则是不在生产代码里保留任何 tf.tensor 创建的临时数据全部用 tf.tidy 包裹不在非必要情况下执行 dataSync()尽量用 async 的 data()。8. 写在最后的实操心得整个示例从数据生成到模型预测加起来不到 150 行代码但把机器学习的核心流程完完整整走了一遍。我带过不少同事从零开始接触 TensorFlow.js这个例子从来没有失手过。它最大的价值在于给了你一个“快速反馈”的学习环境——调整学习率、更换优化器、修改数据噪声立刻就能从损失曲线和预测结果里看到变化。这种即时反馈是传统 Python 机器学习教学里很难给到的体验。最后分享一个小技巧。训练完成后我习惯去浏览器控制台里跑几条命令比如用模型预测几个极端位置的数值比如 x -0.99 和 x 0.99看两端预测是否稳定。如果两端的预测都符合线性规律说明模型没有在数据边界处跑偏。这个技巧虽然简单但在验证模型稳健性时特别实用也是我每次演示这个示例时必做的一个环节。
返回列表