《基于 TensorFlow.js 的机器学习实战:JavaScript 线性回归从原理到落地》
一、机器学习与线性回归:基础概念
1.1 什么是机器学习?
机器学习是人工智能的一个分支,它让计算机能够通过数据学习规律,而非依赖硬编码规则。简单来说,机器学习系统会从历史数据中 “归纳” 出模式,再用这些模式预测未知结果。例如,通过历史房价数据学习 “面积与房价的关系”,进而预测新房屋的价格。
1.2 什么是线性回归?
线性回归是机器学习中最基础的监督学习算法(输入数据包含标签 / 目标值),专门用于解决连续值预测问题(如价格、温度、销量等)。其核心思想是:假设特征与目标值之间存在线性关系,通过数据拟合出最优的线性模型。
从数学角度看,对于单一特征(一元线性回归),模型可表示为:y=wx+b
其中:
- x 是输入特征(如房屋面积);
- y 是预测的目标值(如房价);
- w 是权重(直线斜率,反映x对y的影响程度);
- b 是偏置(直线截距,x=0时的基准值)。
1.3 线性回归的核心目标
线性回归的目标是找到最优的w和b,使模型预测值ypred=wx+b尽可能接近真实值ytrue。衡量 “接近程度” 的指标称为损失函数,最常用的是均方误差(MSE):
MSE=n1∑i=1n(ytrue,i−ypred,i)2
(n为样本数量,MSE越小,模型拟合效果越好)
为了最小化 MSE,我们会使用梯度下降算法:通过不断计算损失函数对w和b的偏导数(梯度),逐步调整w和b,直到损失收敛到最小值。
二、JavaScript 实现线性回归:基于 TensorFlow.js
TensorFlow.js 是 Google 推出的 JavaScript 机器学习库,支持在浏览器或 Node.js 中运行机器学习模型。我们将用它实现一个完整的一元线性回归案例,步骤包括:数据生成、模型构建、训练、预测与可视化。
2.1 环境准备
无需复杂配置,通过 CDN 直接引入 TensorFlow.js:
<!-- 引入TensorFlow.js核心库 -->
<script src="https://cdn.jsdelivr.net/npm/@tensorflow/tfjs@4.14.0/dist/tf.min.js"></script>
<!-- 用于可视化的Canvas -->
<canvas id="regressionPlot" width="800" height="500"></canvas>
2.2 步骤 1:生成模拟数据
为了验证模型,我们先生成一组符合线性关系的数据。假设真实规律为y=1.5x+2,并添加随机噪声(模拟现实中的数据波动):
/**
* 生成模拟数据集
* @returns {Object} 包含特征x和目标y的张量
*/
function generateDataset() {
// 生成100个x值(范围0-10)
const x = tf.tensor2d(
Array.from({ length: 100 }, () => Math.random() * 10), // 随机数0-10
[100, 1] // 形状:100个样本,每个样本1个特征
);
// 生成y值:y = 1.5x + 2 + 噪声(噪声范围-1到1)
const y = x.mul(1.5) // 1.5x
.add(2) // +2
.add(tf.randomNormal([100, 1], 0, 1)); // 加噪声(均值0,标准差1)
return { x, y };
}
// 生成并获取数据
const { x, y } = generateDataset();
代码说明:
tf.tensor2d用于创建二维张量(机器学习中数据的基本格式),randomNormal生成符合正态分布的噪声,让数据更贴近真实场景。
2.3 步骤 2:构建线性回归模型
用 TensorFlow.js 的sequential(序列式)模型构建线性回归器。线性回归本质是一个 “单层神经网络”,仅包含一个全连接层(dense):
/**
* 构建线性回归模型
* @returns {tf.Sequential} 编译好的模型
*/
function buildModel() {
// 初始化序列模型
const model = tf.sequential();
// 添加全连接层(实现y = wx + b的计算)
model.add(tf.layers.dense({
units: 1, // 输出维度:1个预测值
inputShape: [1], // 输入维度:1个特征
// 权重和偏置的初始值(可选,这里用默认值)
kernelInitializer: 'randomNormal', // w的初始值
biasInitializer: 'zeros' // b的初始值
}));
// 编译模型:配置优化器和损失函数
model.compile({
optimizer: tf.train.sgd(0.01), // 随机梯度下降优化器,学习率0.01
loss: 'meanSquaredError' // 损失函数:均方误差
});
return model;
}
// 构建模型
const model = buildModel();
关键参数解释:
units: 1:输出 1 个值(预测的 y);inputShape: [1]:输入 1 个特征(x);- 优化器
sgd(0.01):通过梯度下降更新w和b,学习率 0.01 控制更新步长(步长过大会震荡,过小会训练慢);- 损失函数
meanSquaredError:即 MSE,用于衡量预测误差。
2.4 步骤 3:训练模型
用生成的数据集训练模型,通过多次迭代(epochs)优化w和b:
/**
* 训练模型
* @param {tf.Sequential} model 模型实例
* @param {tf.Tensor} x 特征数据
* @param {tf.Tensor} y 目标数据
*/
async function trainModel(model, x, y) {
console.log('开始训练...');
// 训练配置
const trainingConfig = {
epochs: 200, // 迭代次数:200次
callbacks: {
// 每20次迭代打印一次损失
onEpochEnd: (epoch, logs) => {
if ((epoch + 1) % 20 === 0) {
console.log(`第${epoch + 1}轮,损失:${logs.loss.toFixed(4)}`);
}
}
}
};
// 开始训练(fit方法返回Promise,需用async/await处理)
const history = await model.fit(x, y, trainingConfig);
console.log('训练完成!最终损失:', history.history.loss[history.history.loss.length - 1].toFixed(4));
}
// 启动训练
trainModel(model, x, y).then(() => {
// 训练完成后执行预测和可视化
predictAndVisualize(model);
});
训练逻辑:
model.fit会自动用梯度下降调整w和b,每次迭代后计算损失。随着训练进行,损失会逐渐降低(理想情况下接近噪声的方差)。
2.5 步骤 4:预测与可视化
训练完成后,用模型预测新数据,并通过 Canvas 绘制原始数据和拟合直线,直观展示效果:
/**
* 预测并可视化结果
* @param {tf.Sequential} model 训练好的模型
*/
function predictAndVisualize(model) {
// 生成测试数据(x从0到10,更密集的点)
const xTest = tf.tensor2d(
Array.from({ length: 50 }, (_, i) => i * 0.2), // 0, 0.2, 0.4, ..., 10
[50, 1]
);
// 用模型预测y值
const yPred = model.predict(xTest);
// 将张量转为普通数组(用于绘制)
const xOriginal = x.arraySync().flat(); // 原始x
const yOriginal = y.arraySync().flat(); // 原始y
const xTestArr = xTest.arraySync().flat(); // 测试x
const yPredArr = yPred.arraySync().flat(); // 预测y
// 获取Canvas上下文
const canvas = document.getElementById('regressionPlot');
const ctx = canvas.getContext('2d');
const margin = 50; // 边距
const width = canvas.width - 2 * margin;
const height = canvas.height - 2 * margin;
// 清除画布
ctx.clearRect(0, 0, canvas.width, canvas.height);
// 绘制原始数据点(蓝色)
xOriginal.forEach((xi, i) => {
// 将数据坐标映射到Canvas像素坐标
const xPx = margin + (xi / 10) * width; // x范围0-10,映射到0-width
const yPx = canvas.height - margin - (yOriginal[i] / 20) * height; // y约0-20,映射到0-height
ctx.fillStyle = 'blue';
ctx.beginPath();
ctx.arc(xPx, yPx, 4, 0, 2 * Math.PI);
ctx.fill();
});
// 绘制拟合直线(红色)
ctx.beginPath();
ctx.strokeStyle = 'red';
ctx.lineWidth = 3;
xTestArr.forEach((xi, i) => {
const xPx = margin + (xi / 10) * width;
const yPx = canvas.height - margin - (yPredArr[i] / 20) * height;
if (i === 0) {
ctx.moveTo(xPx, yPx); // 起点
} else {
ctx.lineTo(xPx, yPx); // 连线
}
});
ctx.stroke();
// 输出学习到的w和b(接近真实值1.5和2)
const [wTensor, bTensor] = model.layers[0].getWeights();
const w = wTensor.arraySync()[0][0].toFixed(2);
const b = bTensor.arraySync()[0].toFixed(2);
console.log(`模型学习到的关系:y = ${w}x + ${b}`);
}
可视化逻辑:通过坐标映射将数据值(x:0-10,y:0-20)转换为 Canvas 像素位置,蓝色点为原始数据,红色线为模型拟合结果。
三、运行结果与解读
-
控制台输出:训练过程中会打印每 20 轮的损失(如从初始的 50 + 逐渐降至 1 左右),最终输出学习到的w和b(例如
y = 1.48x + 2.05),接近真实值 1.5 和 2。 -
Canvas 可视化:蓝色点均匀分布在红色直线周围,说明模型成功拟合了数据的线性规律。
四、常见问题与解决方案
4.1 TensorFlow.js 加载失败
- 现象:控制台报错
tf is not defined。 - 原因:CDN 链接失效或网络问题。
- 解决方案:
- 更换 CDN(如
https://cdn.jsdelivr.net/npm/@tensorflow/tfjs); - 本地下载
tf.min.js,通过相对路径引入。
- 更换 CDN(如
4.2 模型训练不收敛(损失不下降)
- 现象:损失始终很高(如 100+)或波动剧烈。
- 原因:学习率不合理(过大或过小)。
- 解决方案:
- 调整学习率:若损失波动大,减小学习率(如 0.001);若损失下降慢,增大学习率(如 0.05);
- 增加迭代次数(
epochs):复杂数据可能需要更多训练轮次。
4.3 可视化异常(点 / 线超出画布)
- 现象:数据点或直线跑到 Canvas 外面。
- 原因:坐标映射逻辑错误(数据范围与画布范围不匹配)。
- 解决方案:
- 打印
xOriginal和yOriginal的最大值,确认数据范围; - 调整映射公式中的分母(如 y 的范围若为 0-30,分母应改为 30)。
- 打印
4.4 浏览器兼容性问题
- 现象:在旧浏览器(如 IE)中无法运行。
- 原因:TensorFlow.js 依赖现代 JavaScript 特性(如 ES6、WebGL)。
- 解决方案:使用 Chrome、Edge、Firefox 等现代浏览器;若需兼容旧环境,可通过 Babel 转译代码。
五、总结与扩展
线性回归是机器学习的 “入门钥匙”,本文通过 JavaScript 实现了完整流程,核心要点包括:
- 线性回归通过y=wx+b拟合特征与目标的线性关系;
- 用均方误差衡量误差,用梯度下降优化参数;
- TensorFlow.js 简化了模型构建与训练的复杂度。
在此基础上,你可以进一步探索:
- 多元线性回归:处理多个特征(如用面积 + 房龄预测房价);
- 多项式回归:通过特征转换(如x2)拟合非线性关系;
- 正则化:解决过拟合问题(如 L1/L2 正则化)。
更多推荐


所有评论(0)