从第一性原理重新学习 LLM 与强化学习
写作说明:本文的想法和初始草稿材料来自我本人;ChatGPT 协助了结构整理、文字编辑与翻译。
过去我已经使用过很多语言模型和强化学习背后的概念。有些是在博士阶段接触到的,有些则已经变成今天做 AI 时几乎每天都会遇到的词:Softmax、Cross-Entropy、Attention、KL Divergence、Value Function、Policy Gradient。
这些公式我见过,也会用。
但随着大语言模型逐渐变成我日常工作中几乎无处不在的工具,我开始越来越明显地感觉到一件事:认识一个数学公式、知道怎么使用它,和真正理解它为什么长成这样,并不是一回事。
比如 Softmax:
\[p_i = \frac{e^{z_i}}{\sum_j e^{z_j}}.\]在二分类的情况下,这个结构尤其容易看出来。如果
\[\Delta = z_1-z_2,\]那么第一类的 probability 可以写成
\[p_1 = \frac{e^{\Delta}}{e^{\Delta}+1}.\]拖动下面的 slider,可以直接看到:当两个 logits 之间的相对差值变化时,两个类别的 probability 如何随之变化。
这个交互图通过 Marimo 直接在浏览器中运行。查看源 notebook。
这个公式我已经见过无数次。
但为什么要用 exponential?
为什么偏偏要使用自然指数的底数 e?
为什么不能把这些 score 直接除以它们的总和?
为什么给所有 logits 都加上同一个常数,最后的 probability 却完全不会改变?
一旦开始这样追问,很多原本已经非常熟悉的公式突然又变得陌生起来。
比如,当我们写:
\[P(x_t\mid x_{<t})\]这里的 probability 到底是从哪里来的?
数据集给我们的首先是 observations,而不是一张写好了 probability 的表。
再比如 Maximum Likelihood:为什么一件事情被我们观察到了,就意味着模型应该提高它的 probability?
又比如 Policy Gradient:
\[\nabla_\theta \log \pi_\theta(a\mid s)\]action 明明是离散地 sample 出来的,我们为什么还能通过这样的式子去训练一个 policy?
这些问题都非常基础。
某种意义上,它们甚至比我之前做过的研究问题基础得多。
也正因如此,它们让我有一点不舒服。
我已经对这些工具足够熟悉,熟悉到可以看到这些公式以后直接继续往下走,而不再停下来问它们为什么成立。
所以我决定停下来。
我所说的“第一性原理”是什么
这个项目的目的并不是收集一份“做 LLM 必须掌握的公式清单”。
我也不想因为用了 “first principles” 这个名字,就假装所有成功的机器学习方法都能够从几个数学公理中唯一地推导出来。
机器学习中有很多东西,本来就是选择。
一个 probability identity、一个 modeling assumption、一种 parameterization、一个 optimization convention 和一个 engineering decision,它们的地位并不相同。
理解一个公式,其中的一部分就是弄清楚:它究竟属于哪一种东西。
所以每遇到一个数学结构,我想重新回到几个问题:
- 我们到底在解决什么问题?
- 我们真正观察到的 quantity 是什么?
- 哪些地方需要 assumption?
- 哪一步是 mathematical identity?
- 哪一步是 modeling choice?
- 为什么这种选择特别有用?
- 如果换一种做法,会失去什么?
我心里有一个很简单的检验方法:
如果我把公式忘掉了,我还能不能根据问题本身,把背后的 reasoning 重新构造出来?
Softmax 就是一个很好的例子。
记住
\[p_i = \frac{e^{z_i}}{\sum_j e^{z_j}}\]并不难。
但当它进一步变成:
\[\frac{p_i}{p_j}=e^{z_i-z_j},\]或者:
\[z_i-z_j=\log\frac{p_i}{p_j},\]理解就发生了一点变化。
logits 不再只是一些等待被 normalization 的神秘数字。
它们之间的差,直接描述了 probability ratio 的 log。
这正是我希望在这些文章里保留下来的东西:不仅是最后得到的公式,还有为什么我当时会觉得这个公式值得重新想一遍。
用解释来检验理解
写这个系列还有另外一个原因。
我希望训练自己解释技术问题的能力。
一个概念只停留在脑子里的时候,很容易产生“我已经懂了”的感觉。但真正写出来以后,标准会突然提高。
如果我无法解释为什么这一步能从前一步得到,如果我不得不用一句“反正就是这么做的”来跳过一个 modeling choice,那么那个理解上的空缺就会暴露出来。
所以对我来说,解释并不是学习结束以后才做的事情。
解释本身就是检验我是否真正学会的一部分。
慢慢地,我形成了一个大概这样的学习过程:
\[\text{困惑} \rightarrow \text{推导} \rightarrow \text{直觉} \rightarrow \text{实现} \rightarrow \text{解释}.\]当然,真实过程并不会永远这么整齐。
有时候写代码会暴露出一个原本没有意识到的误解。
有时候一个很小的数值例子,比完整的推导更有用。
有时候一个听起来很有道理的解释,一旦试着精确地写出来,就会发现其实并没有真的说清楚。
这些失败本身也是学习的一部分。
所以我不太想在最后的文章里把它们全部清洗掉,再重写成一条像教科书一样、所有事情从一开始就显得顺理成章的路线。
“为什么 Softmax 要用 exponential?”这个问题值得留下,很大程度上正是因为我真的问过这个问题。
用 LLM 来学习 LLM
这个项目还有一个有点自指的地方:
我的很多学习过程,本身就是通过和 ChatGPT 不断对话完成的。
我把它当成 tutor,也把它当成一个 thinking partner。
我会提出问题,质疑它的解释,尝试自己的理解,和它一起推导,然后过一段时间又回到同一个问题,直到我能够用自己的方式把它说清楚。
但这并没有取消“自己理解”这一步。
某种意义上,它反而让这条界限变得更重要。
LLM 可以非常快地生成一段听起来很合理的解释。
但“听起来合理”不是我希望使用的标准。
我仍然需要检查 algebra,确认 assumption,用简单例子测试 intuition,并最终判断:这个解释对我来说到底是不是真的成立。
所以这些对话也会成为这个系列的一部分 source material。
我希望尽可能保留真实发生过的问题、误解和理解上的变化,而不是在事情结束以后,重新虚构一条更漂亮、更顺畅的学习历史。
用一个语言模型去重新理解语言模型背后的数学,本身有一种有点奇怪的循环。
我也很好奇,这种循环什么时候能够帮助理解,什么时候又会变成一种让我们绕过理解的捷径。
从公式走到代码
与此同时,我也在维护一个叫做 llm-from-first-principles 的 GitHub repository。
它和这个博客做的是两件互补的事情。
写文章是在问:
我能不能把它解释清楚?
而实现代码是在问:
我能不能在不把机制藏起来的情况下把它做出来?
repository 是我亲手拆解和实现这些组件的地方。
博客则是我尝试重新构造这些组件背后的 reasoning 的地方。
Softmax 也是一个很小、但很能说明问题的例子:为什么我不想停在公式本身。
纸面上,我们写:
\[p_i = \frac{e^{z_i}}{\sum_j e^{z_j}}.\]从数学上说,这个定义没有问题。但如果在程序里机械地照着这个式子计算,数值上可能并不稳定:当某个 logit 是很大的正数时,exponential 可能 overflow。
稳定的实现会先找到:
\[m = \max_j z_j,\]然后把所有 logits 同时平移:
\[z_i' = z_i-m.\]最后再计算:
\[p_i = \frac{e^{z_i-m}}{\sum_j e^{z_j-m}}.\]这会得到完全相同的 probabilities,因为 Softmax 对所有 logits 同时加上或减去同一个常数是不变的:
\[\operatorname{softmax}(z) = \operatorname{softmax}(z-c).\]这是我在 PyTorch 中手动实现的版本:
import torch
def softmax_stable(x: torch.Tensor) -> torch.Tensor:
max_x = x.max(
dim=-1, keepdim=True
).values # .values => get values, here [0] the same.
shifted_x = x - max_x
exp_x = torch.exp(shifted_x)
row_sums = exp_x.sum(dim=-1, keepdim=True)
return exp_x / row_sums
减去最大值并不是重新定义了 Softmax,而是在数值上更安全地计算同一个数学函数。
理解一个数学公式,并不等于理解它在代码里真正如何成立。
公式写在纸面上的样子,和它要在程序里可靠运行所需要的细节之间,通常还有一层:numerical stability、tensor shapes、masking、broadcasting、memory layout、gradient flow,或者其他被紧凑数学记号省略掉的约束。
这正是我最初想开这个 repository 的原因之一。
我希望这两个过程最终能够形成一个循环:
\[\text{问题} \rightarrow \text{推导} \rightarrow \text{实现} \rightarrow \text{解释} \rightarrow \text{再次提问}.\]我并不期待这个过程最后能够给出一套完整的理论,解释“为什么大语言模型会成功”。
这个问题远远大于任何一组数学公式,而且其中还有很多事情我们并没有真正理解。
我的目标要小一些,但对我来说也更实际:
让这套机器里越来越少的部分看起来像魔法。
我想从 probability 开始。
我们训练语言模型时不断谈论 token 和 sequence 的 probability,但一开始真正拿到手里的,却只是 observations。
那么,这个 probability 究竟是从哪里来的?
这就是第一篇文章想要开始回答的问题。