从零手写大模型 · 训练篇 01 - Loss Function(损失函数)

一颗消失的星星

1801 年,意大利天文学家皮亚齐发现了一颗新天体——谷神星。但他只观测到了一小段弧线的数据,谷神星就运行到太阳背后,从视野里消失了。

问题来了:凭这几个残缺、带误差的观测点,能不能算出它的完整轨道,预测它什么时候会从太阳背后重新出现?

24 岁的高斯给出了答案:能。他用的方法,后来被称为最小二乘法。这一集,我们就把这套 200 多年前的方法讲透,再看看它跟大模型训练之间,到底是什么关系。

高斯面对的问题

假设望远镜在几个不同时刻,观测到谷神星出现在几个位置(简化成一维,方便手算)。真实轨道是一条曲线,但每次观测都会有误差——仪器精度、大气扰动,总会让记录下来的位置跟真实位置差一点。

用最简单的情形说明这个思路:假设我们观测到 3 个数据点(x 代表时间,y 代表观测到的位置):

x = 1, y = 2.1
x = 2, y = 3.9
x = 3, y = 6.2

我们猜测这些点背后藏着一条直线关系:y = wx + b。问题是:w 和 b 该取多少,才能让这条直线"最贴合"这三个观测点?

为什么是"平方"而不是直接求差

对每一个观测点,我们画的这条线预测出来的值,和真实观测到的值之间,都有一个误差:

误差 = 预测值 - 真实值 = (w·x + b) - y

如果直接把 3 个点的误差加起来判断"贴合程度好不好",会出问题:比如误差分别是 +0.5、-0.5、0,加起来是 0,看起来"完美贴合",但实际上第一、第二个点都偏了 0.5,并不完美。正负误差会互相抵消,掩盖真实的偏差。

高斯的处理方式是:把每个误差平方,这样正负号消失,同时大的误差会被放得更大(平方是非线性的,误差越大惩罚越重),这样更能真实反映"整体偏差有多大":

损失 = (预测值 - 真实值)² 加总所有点
     = [(w·1+b) - 2.1]² + [(w·2+b) - 3.9]² + [(w·3+b) - 6.2]²

这个"预测值与真实值之差的平方和",就是最小二乘法(Least Squares)名字的由来——目标是找到一组 w、b,让这个平方和最小。

怎么找到那个"最小"点:求导

高斯是怎么算出让损失最小的 w 和 b 的?他用的是微积分里最基础的一个事实:一个函数在最低点(极小值)的地方,切线是水平的,也就是导数为 0。

想象一下:如果把"损失"画成一座山谷的形状(w、b 是横纵坐标,损失是高度),山谷最低的那个点,一定是"往哪个方向挪一点点,高度都不会再降低"的地方——这正是导数(某个方向上的变化率)等于 0 的地方。

把损失函数分别对 w 和 b 求导,让导数都等于 0,就能解出这组最优的 w、b。用上面那 3 个点实际算一遍(过程略去,只看结果):

w ≈ 2.05
b ≈ -0.03

也就是说,最贴合这 3 个观测点的直线大概是 y ≈ 2.05x - 0.03。高斯当年处理的不是简单的直线,而是复杂得多的椭圆轨道方程,但思路完全一样:先写出"预测值与观测值之差的平方和"作为衡量误差的标准,再用求导找到误差最小的那组参数。他用这套方法算出的位置,后来天文学家真的在那个位置重新找到了谷神星。

把这套流程抽象出来:三步走

高斯这套方法,拆出来看只有三步:

  1. 假设一组参数(w、b),用它做出预测
  2. 写一个数字,衡量预测和真实值差多少(这就是"损失")
  3. 调整参数,让这个数字变小(用求导找到最优点)

这三步,就是"损失函数"这个概念最原始的样子——用一个数字量化"错得有多离谱",再想办法让这个数字变小。

这跟大模型训练是什么关系

我们手写的 GPT 模型里,也有大量的参数(推理篇 EP13 加载的 GPT-2 权重,光是每一层里 Q/K/V 的投影矩阵,就是成千上万个数字)。训练这些参数,走的其实是完全相同的三步:

  1. 模型现在的这一组参数,对"猫吃"后面该接哪个词,给出一个预测
  2. 需要一个数字,衡量这个预测和真实答案(比如"鱼")差多少——这就是模型的损失
  3. 想办法调整这些参数,让这个损失变小

高斯只有 2 个参数(w、b),可以直接求导解方程算出精确答案。但我们的模型有成千上万个参数,没办法这样直接解方程——这个"参数太多,该怎么求导"的问题,正是下一集要讲的核心内容:反向传播。但不管参数是 2 个还是几百万个,"写一个数字衡量误差,再想办法把这个数字变小"这个根本思路,是完全一样的。

也就是说,你现在理解了最小二乘法在干什么,就已经理解了"损失函数"这个概念的本质——大模型训练里那个损失函数(叫 Cross-Entropy),只是换了一种方式去"衡量误差",但它存在的目的、它在整个训练流程里扮演的角色,跟高斯这里的"差的平方和",没有任何区别。

下集预告

这一集我们理解了损失函数如何衡量预测误差。下一集是训练篇 EP02「反向传播」:从 loss 出发,用链式法则算出每一层参数的梯度,为后续的参数更新做好准备。

← 返回训练篇目录