跳到主要内容

图解机器学习:线性回归模型

前言

最近有些好奇,大模型到底是如何训练的。每天都在跟 ChatGPT、Claude 一起工作,作为一个计算机行业的工作者,难免对它们的实现产生好奇心。

大模型与传统软件的区别

如今的大模型和我们过去认知中的软件并不是相同的东西。

过去开发传统软件时,我们会预先定义一系列规则。用户输入数据后,软件按照这些既定规则进行处理,并输出相应的答案。

传统软件通过预先定义的规则处理数据并输出答案

而大模型的工作方式恰好相反。训练时,我们向模型提供大量输入及其对应答案,让模型从数据中学习两者之间的规律。传统软件是“先定义规则,再根据输入得到答案”;大模型则是“先提供输入和答案,再从中归纳规则”。

大模型通过数据和答案训练学习出规则

线性回归模型

如今的大模型本质上都是预测型模型,而线性回归作为最简单的预测模型,虽然体量微不足道,却已经浓缩了大模型的核心:模型 → 损失函数 → 优化

所谓训练模型,本质上就是不断优化、让模型能力逐步升级。今天,我们就借线性回归这个最简单的模型,学习线性回归模型的三种求解方法:

  • 穷举法:把参数的候选取值一个个代入试,谁让损失最小就用谁。
  • 最小二乘法:直接对损失函数求导、令导数为 0,一步解出最优参数。
  • 梯度下降法:沿损失函数梯度反方向一步步迭代,逼近最低点。

举个例子:小明在电脑城打工,我们不知道他每组装一台电脑具体能赚多少钱,只知道他每天组装了多少台电脑、当天拿到了多少工资。

电脑数量 x工资 y
550
10100
20200
15150
…………
xy

这时,聪明的你肯定已经想到设一个一元一次方程,只要找出 x 和 y 的关系,就能知道装一台电脑赚多少钱。

损失函数

损失函数就是衡量“模型预测”的记分器,分数越小越好。

还是沿用上面的例子,我们需要知道小明装一台机器赚多少钱。假设装一台机器赚 w 元,此时我们的方程就是 y = w * x。因为我们并不知道 w 到底是多少,所以就需要猜 w 到底是多少,然后代入 x,计算预测结果与实际结果之间有多少差别。

  • 假设 w = 10:预测 50、100、200、150,和真实工资一分不差 → 没有任何损失。
  • 假设 w = 5:预测 25、50、100、75,每个都差一截 → 损失很大。
  • 假设 w = 15:预测 75、150、300、225,每个都超一截 → 损失也很大。

注意,损失函数和 w 不是一回事。损失函数是我们自己定义的“记分器”,而 w 才是真正要找的参数。

我们可以发现,在不同的 w 下损失各不相同。我们需要算出误差的平均值,所以就需要用到下面这个公式:

L(w, b) = (1/n) · Σ (yᵢ − ŷᵢ)²

给不懂数学的小伙伴解释一下这段“鬼画符”的含义:

  • L(w, b):损失函数,输入参数 wb,输出一个“差多少”的分数。
  • n:数据点个数。
  • Σ:求和,把每个点算一遍再加起来。
  • yᵢ:第 i 个点的真实值。
  • ŷᵢ:第 i 个点的预测值,ŷᵢ = w * xᵢ + b。

由上面的例子可以得出,损失越小,直线越贴合;损失越大,直线越离谱。于是问题就统一转换为寻找让 L(w, b) 最小的参数。

如果把 b 固定,那么 L(w, b) 就是一个开口向上的抛物线(梦回高中数学),最低点就是最优解。

穷举法

穷举法,就是把 w 可能的范围切成很多份,得到一串候选值,然后一个一个代入计算损失,挑最小的那个。

  • 优点:直观、几乎不需要数学;给定候选网格内一定能找到最小值。
  • 缺点:计算量爆炸。w 试 100 个、b 试 100 个,就是 100 × 100 = 1 万个组合,参数一多,维度相乘,直接算不动了。

穷举法适合建立直觉,不适合实战。

最小二乘法(解析解)

既然最低点处曲线是“平的”,也就是导数为 0,那就直接求导、令导数等于 0,把 wb 解出来,一步到位。

对损失函数分别求偏导并令其为 0:

∂L/∂b = (2/n) Σ (w·xᵢ + b − yᵢ) = 0      … ①
∂L/∂w = (2/n) Σ (w·xᵢ + b − yᵢ)·xᵢ = 0 … ②

先解 ①,求 b:

Σ (w·xᵢ + b − yᵢ) = 0
n·b = Σ yᵢ − w·Σ xᵢ
b = ȳ − w·x̄ … ③

其中 ȳ 是均值。

把 ③ 代入 ②,求 w:

Σ (w·xᵢ + ȳ − w·x̄ − yᵢ)·xᵢ = 0
Σ [ w·(xᵢ − x̄) − (yᵢ − ȳ) ]·xᵢ = 0
w·Σ (xᵢ − x̄)·xᵢ = Σ (yᵢ − ȳ)·xᵢ

利用均值性质 Σ (xᵢ − x̄) = 0,可得 Σ (xᵢ − x̄)·xᵢ = Σ (xᵢ − x̄)²;同理,右边 Σ (yᵢ − ȳ)·xᵢ = Σ (yᵢ − ȳ)(xᵢ − x̄)。于是:

w = Σ (xᵢ − x̄)(yᵢ − ȳ) / Σ (xᵢ − x̄)²    … ④

③④ 就是一元线性回归的闭式解。写成多元矩阵形式更简洁:

θ = (XᵀX)⁻¹ Xᵀ y

优点:一步算出、结果精确,不用迭代、不用选超参数。

缺点:要求 XᵀX 可逆(不可逆时要用伪逆或加正则项);特征很多时矩阵求逆代价高(约 O(d³)),内存也吃紧。

梯度下降法(逐步逼近)

最小二乘法“一步到位”虽爽,但特征很多、或损失函数复杂到解不出解析解时就不灵了。于是换个思路:不指望一步到位,而是沿梯度反方向,一步步“滑”到最低点

先算损失对参数的偏导:

∂L/∂w = (2/n) Σ (w·xᵢ + b − yᵢ)·xᵢ
∂L/∂b = (2/n) Σ (w·xᵢ + b − yᵢ)

梯度方向是“上升最快”的方向,取反方向就是“下降最快”,反复迭代:

w ← w − α · ∂L/∂w
b ← b − α · ∂L/∂b

其中 α学习率,控制每步迈多大。

优点:能处理海量数据和高维参数;几乎适用于任何可求导的损失函数;配合小批量(mini-batch)可扩展到深度学习。

缺点:要调学习率 α——太大时震荡甚至发散,太小时收敛太慢;可能陷入局部最优(对线性回归这种凸损失则不会);结果是近似解。

三种方法对比

方法求解方式精度适用场景
穷举法暴力遍历候选值网格内精确、网格外受限教学演示、小参数空间
最小二乘法求导 = 0 得解析解精确特征不多、XᵀX 可逆
梯度下降法沿梯度反方向迭代近似(可控制)大规模、高维、复杂损失

小结

三者的底层逻辑一致:把“拟合数据”转化为“最小化损失函数”,再去求这个最小值

  • 穷举法:最笨、最直观,帮你“看见”问题。
  • 最小二乘法:数学上一步到位,适合小规模精确求解。
  • 梯度下降法:工程上最通用,是大规模机器学习和深度学习的基础。

搞懂这三条路,线性回归乃至大部分机器学习的“训练”本质,就都通了。

💬 评论区