加速器系统课程 / 主线 / 专题四 / 反向与优化器

反向与优化器

Completing the Bill: Backward, Recompute, and the Optimizer
专题一算完前向,结论是「装不下」。
—— 可那只是半张账单。而剩下那半张里最大的一块, 既不是权重也不是激活。

这一讲把训练一步的账补全:反向要付的三倍算力、 那些不能算完就扔的中间激活、以及每参数 16 字节的优化器状态补完之后,ZeRO 那三级分法就不用背了 —— 它是从这张账单里长出来的。

前置 专题一 部分小节需要 专题三 后续 专题五 · 并行策略 读法 按「谁最大」读,不按流程读

课程作者 Chris Yang·Google Cloud AI Infra 架构师

第 零 节

前向只是半张账单

专题一跟着一个 token 走完了前向,最后算出一句话:装不进任何一块卡。 ⛔ 可那只是账单的一小半

真正训练一步,显存里是同时压着四样东西的:权重梯度优化器状态,外加一大堆不能算完就扔的中间激活这一讲把这张账单补全。

⭐⭐⭐ 而补完之后会发现一件多数人想不到的事: 最大的那一块,既不是权重,也不是激活 —— 是优化器状态。

⭐ 这一讲按「谁最大」的顺序讲,而不是按「训练流程」的顺序。

流程的顺序是:前向 → 反向 → 更新。 但那个顺序会让最大的那一块最后才出场,而它恰恰是决定一切的那一块。

⭐⭐ 所以这一讲的每一节只回答同一个问题: 这一项有多大,能不能省,省它要拿什么去换。

第 一 节

反向要付什么 —— 三倍算力,外加一堆扔不掉的中间结果

反向传播不是「再跑一遍」。它要付两样东西: 大约两倍于前向的算力,以及 —— 更要命的 ——  一路累加、不能提前扔掉的中间激活

1.1 先把难点问对 —— 难的不是求导,是求三千亿个导数

⭐ 求导是高中的事。 这里难的是:要求快三千亿个导数,而且每一步都要重求一遍

不过「求导是高中的事」这句话会劝退一批人 ——  所以先花一格,把这一讲要用到的三个词说完。 没有极限、没有公式,三句话,一张图。

导数、偏导数、链式法则 —— 其实是同一个词:兑换率 这一格不讲极限、不讲公式 —— 只讲这三个词到底在说什么事 导数:动一格,变几格 偏导数:其余按住 链式法则:一路乘 导数问的不是「现在是多少」,是「你动一格,它动几格」 把参数想成一个旋钮,loss 是它右边那个读数 一个参数 (三千亿个之一) 往上推一点点 +0.01 这个参数 1 格 loss 3 格 两把尺一格都是 0.01 那这个旋钮的导数就是 0.03 ÷ 0.01 = 3 读作:你动一格,它动三格 不是「读数是 5.00」 —— 那是,这是兑换率 顺带记住它的单位loss 每参数 —— 后面讲「那一步该迈多大」那一节,整节都是被这个单位逼出来的。 Ⓑ 那「偏」字是什么意思 —— 就三个字:按住不动 一台三千亿个旋钮的调音台,一次只拧一个,其余全按住 按住 按住 按住 按住 只动这一个 按住 按住 按住 按住 …… 其余三千亿个,全按住 偏导数 = 其余全按住时,这一个旋钮的兑换率。 而把三千亿个旋钮各自的那个数排成一列 —— 那一列就叫「梯度」。 Ⓒ 可旋钮不直接连到读数 —— 中间隔着好几级 每一级有自己的兑换率,总兑换率就是一路乘起来 推子 × 2.0 中间量甲 × 0.5 中间量乙 × 3.0 难听程度 总兑换率 = 一路乘起来 2.0 × 0.5 × 3.0 = 3 这就是换汇 人民币 → 港币 → 美元,每步一个汇率 总汇率当然是乘出来的 —— 链式法则就这一件事 而「这一串数从哪一头开始乘」 —— 结果完全一样,代价差一万年。那是紧接着 `fig-reverse` 那一格的事。 所以这一讲后面所有的东西,都只用到这三句话 ① 导数 = 你动一格,它动几格(兑换率)。② 偏导数 = 其余全按住时的那个兑换率。③ 链式法则 = 中间每一级的兑换率,一路乘起来。—— 没有极限,没有 ε,没有要背的公 式。 唯一一个真的要小心的点:兑换率不是「值」。读数是 5.00 跟「动一格变三格」是两件完全不同的事 —— 本讲后面每次说「梯度大」,说的都是后者
⭐⭐⭐ 整讲的入口那一格 —— 导数、偏导数、链式法则,其实是同一个词:兑换率
⭐⭐ 导数问的不是「现在是多少」,是「你动一格,它动几格」。Ⓑ「偏」字的全部含义就三个字 —— 按住不动;而三千亿个旋钮各自那个数排成一列,那一列就叫梯度
Ⓒ 链式法则就是换汇:人民币→港币→美元,每步一个汇率,总汇率当然是乘出来的。—— 而「从哪一头开始乘」代价差一万年,那是 1.5 的事。
出处与口径

⚠️ Ⓐ 的 5.00 → 5.03、Ⓒ 的 ×2 / ×0.5 / ×3 都是编出来的示意数,唯一的作用是让「相除」和「相乘」这两件事看得见。⭐ 脚本里 assert 了两条:总兑换率必须真的是三个乘积,而且链条里要有一级是缩小的 —— 不然「乘起来」会被读成「越乘越大」,而那正是梯度消失/爆炸那一节要讲的反面。

⭐ 这一格刻意不碰极限:严格地说导数是「动的那一点点趋于 0 时的极限」,而这张图画的是差商。⛔ 对本讲够用 —— 而且 §1.1 那个「笨办法」用的**正是差商**,所以这里不严格反而接得更顺。

⭐⭐ 「兑换率」这个说法不是为了好听,它自带两个钩子,两个都是本讲自己的:一串数相乘,从哪头开始乘代价差一万年(§1.5);而单位对不上所以必须再乘一个折算系数(§3.3)。

先看看笨办法长什么样 —— 这一步不能跳过, 不知道笨办法有多笨,就不会觉得反向传播有多神。

笨办法(有限差分): 我想知道某个参数对 loss 有多大影响,就把它动一丁点, 整个网络重跑一遍前向,看 loss 变了多少。

那三千亿个参数,就是三千亿次前向。 一次前向按一秒算,跑完要将近一万年 —— 而这还只是一步

⭐⭐⭐ 所以反向传播真正解决的问题,不是「怎么求导」, 是怎么把三千亿次前向压成一次

1.2 种子:第一个梯度长什么样

起点很朴素 —— loss 是一个数。 (这一点后面是全部关键,先记着。)

那第一个梯度是什么?拿最常见的交叉熵来说:网络最后吐出一个概率分布 —— 下一个字是「的」的概率 0.3、是「了」的概率 0.2…… 而正确答案是一个只有一格是 1、其余全是 0 的东西。

⭐⭐ 第一个梯度 = 你猜的,减去正确答案。

猜高了的地方是正数,该高没高的地方是负数。 就这么简单。这个差,就是整条反向链的种子

📌 这是 softmax + 交叉熵这一对的经典结果 ——  它们凑在一起,导数会漂亮地约化成「预测减真值」。 不是所有 loss 都这么好看,但这一对是今天的默认组合。

⚠️ 一个说准了才不会误导的细节: 那个「减」不是对概率求导,是对 logits 求导 —— 也就是 softmax 之前那一层的输出。
⭐ 这个区别有实际后果:正因为约掉的是 softmax 那一步, 这个梯度才不会在 softmax 饱和的地方消失 —— softmax 和交叉熵总是成对出现,一半的原因就在这儿。

梯度是什么 —— 一屋子人各提各的要求,最后取个平均 这一格不讲公式,只讲这件事在干嘛 ① 一条样本的诉求 ② 诉求怎么往回递 ③ 众口难调,取平均 Ⓐ 一条样本看完输出,只会提一个要求 正确答案是「的」,而网络给「的」只打了 0.3 差 0.70 现在 0.30 差 0.20 现在 0.20 差 0.15 现在 0.15 差 0.10 现在 0.10 灰柱 = 网络现在给的 | 虚线 = 正确答案 | 箭头长度 = 要推多少 它的要求就一句话 「的」推上去,其他推下去 而且推多少,看差多远 这就是本讲说的那颗种子 「预测 − 真值」 —— 换成人话就是这一句 注意它只是「希望」—— 没人能直接改输出 Ⓑ 可输出改不了 —— 想让一个数变大,只有三条路 而第三条走不通,于是它变成了对上一层的新要求 ① 改偏置 直接给它加一点 最省事,但能调的余地小 ② 改权重 把连过来的线加粗 上游越亮的那根线,加粗越划算 ③ 让上一层更亮 要求前一层把该亮的点亮起来 可上一层也改不了 —— 于是这条要求往回递一层 第 ② 条那句「上游越亮越划算」,正是反向为什么必须用到前向存下来的值 —— 要知道改哪根线回报最大, 得先知道那根线的上游当时有多亮,而那个亮度是前向算出来的。激活扔不掉,根子就在这一句上。 而第 ③ 条「要求往回递一层」,就是反向传播这个名字的全部含义。 Ⓒ 但不能只听一个人的 只听那一条样本的,网络会学会把什么都答成「的」 样本 1 「的」推上去 样本 2 「了」推上去 样本 3 「在」推上去 …… 各提各的 全都满足? 做不到 于是把所有人的诉求加起来取平均 —— 那个平均,就是梯度。 而多卡训练里那一道「跨卡汇总」,干的就是这个平均 —— 每张卡先收自己那批人的意见,再凑到一起。 所以 batch 不是「为了跑得快」才有的 —— 它首先是「别只听一个人的」。 整个反向传播,用人话讲完就是这三步 一条样本提要求 → 要求往回递一层又一层 → 所有样本的要求取平均。后面那些张量、链式法则、矩阵乘,都是这三句话的算法实现,不是另外一件事。 而这一讲真正关心的账,全挂在第二步上:要求往回递的时候,必须回头看前向留下的东西 —— 这就是激活扔不掉、显存下不来的全部原因。
⭐⭐⭐ 全讲最口语的一张 —— 「梯度」这个词唬人,本意却只是一屋子人各提各的要求,最后取个平均
⭐⭐ Ⓑ 那句「上游越亮的那根线,加粗越划算」,就是 §1.6「反向必须用到前向的值」的人话版 —— 激活扔不掉,根子在这一句上。
Ⓒ 只听一条样本的,网络会学会把什么都答成同一个字。所以要平均 —— 而多卡训练里那道「跨卡汇总」,干的就是这个平均。
出处与口径

📌 「推一下(nudge),推多少跟差多远成正比」「想让一个神经元更亮有三条路」「改权重要按上游亮度成比例,回报最大」「众口难调,只能取平均;只听那张 2 的,网络会把所有图都判成 2」四个装置,取自 3Blue1Brown《What is backpropagation really doing?》官方讲义 —— 已逐条核过原文,图是我们自己重画的,例子换成了本讲一直在用的「下一个字」。

⚠️ 他在「按上游亮度成比例」那里顺带提了赫布理论(neurons that fire together wire together),但原文自己就说这个类比并不严格(未训练的网络并没有在「想」那个答案),所以只记在这儿,不画上图。

⭐ Ⓑ 与 §1.6、Ⓒ 与 §1.8 的那两处接头,原文都没有,是本讲自己的合题 —— 3B1B 讲的是「反向传播在干嘛」,本讲要的是「它为什么这么费显存」。

1.3 往回传的不是一个数 —— 这里是最多人卡住的地方

很多人脑子里的画面是「一个梯度值一层一层传下去」。不是这样。

往回传的是一整个张量,形状跟这一层的输出一模一样它的含义是:「loss 对我这一层输出的每一个位置,各有多敏感」 —— 这一层输出有多少个数,它就有多少个数。

⭐⭐ 每一层拿到这个上游传来的敏感度之后,只做两件事:

  • ① 算出自己那块权重的梯度 ——  用这个敏感度,乘上前向时存下来的输入
  • ② 算出该继续往下游传的敏感度 ——  用这个敏感度,乘上自己的权重

⭐ 这两件事各是一次矩阵乘。前向一次、反向两次 —— 三倍算力就是从这儿来的,不是估的,是数出来的。

三倍算力是数出来的 —— 前向一次矩阵乘,反向两次 这一格不讲链式法则,只数乘法做了几次 前向 1 次 反向 2 次 合计 3 次 Ⓐ 把一层摊开 —— 总共只有三次矩阵乘,一次向前、两次向后 只数矩阵乘:norm / 激活 / 偏置在这笔账里可以忽略 ① 前向 算这一层的输出 输入 X × 权重 W 输出 Y → 交给下一层 ② 反向 · 权重梯度 算「我这块权重该怎么改」 输入 Xᵀ × 上游敏感度 dY 权重梯度 dW → 交给优化器 这一块不是新算的 —— 是 ① 里那个输入被存下来了 ③ 反向 · 传给下游 算「前一层该收到什么」 上游敏感度 dY × 权重 Wᵀ 新敏感度 dX → 交给前一层 三块积的面积一样大 —— 所以 1 + 2 = 3 不是个比喻,是数出来的。 而 ② 那块虚线的,就是激活扔不掉的全部原因 Ⓑ 全图的钥匙:一条线进来,分成两支 不是「反向比较慢」这种含糊说法 —— 是两支,而且一支都省不掉 上游传来的敏感度 (就这一个东西) × 前向存下来的输入 一次矩阵乘 这块权重的梯度 交给优化器 到此为止 这是这一步真正要的东西 · 到这儿这一支就不往前了 × 这一层的权重 一次矩阵乘 新的敏感度 喂给前一层 接着往左走 链条靠它 · 少了它,再往前就断了 两支乘的东西不一样(一支乘输入、一支乘权重)—— 所以合并不了;一支断了链条就断,一支没了这一层就白算。一个都省不掉,这就是那个 2。 Ⓒ 于是这笔账就封口了 推理 只有前向 一遍 = 1× 训练 前向 + 反向 一遍 一遍 一遍 = 3× 训练 + 全量重算 一遍 一遍 一遍 一遍 = 4× 「一遍」= 一次走完整个网络的矩阵乘量。推理只买一遍,训练要买三遍 这张图顺带回答了另外两个常见疑问 「为什么训练比推理贵这么多」 —— 算力上就是这个 3 倍; 但真正拉开差距的不是它,而是显存里那一整条从头挂到尾的激活(下一张图)。 「反向能不能只算一次」 —— 能,如果你不打算继续往前传(比如只微调最后一层)。 那种情况下第 ③ 次确实可以省掉 —— 冻结层为什么便宜,原因就在这儿。
⭐⭐ Ⓑ 是这张图的钥匙 —— 反向不是「比较慢」,是它要回答两个不同的问题:我这块权重该怎么改、上游该收到什么。两个问题,两次乘法。
Ⓐ 第②步那句「用前向存下来的输入」是整个专题的枢纽 —— 激活扔不掉、以及下一节那笔交易,全挂在这一句上。
出处与口径

⚠️ 「3 倍」是矩阵乘口径的常用近似:只数 matmul,忽略 norm / 激活 / 偏置 / 通信。真实 step 里这些占比不大,但不是零

⛔ 图上第 ② 步那句「用前向存下来的输入」是整个专题的枢纽 —— 激活扔不掉、以及下一节那笔重算交易,全都挂在这一句上

跟着数走一遍 —— 三个节点,前向一遍,反向一遍 前面讲的全是道理。道理点头很容易,跟着数走一遍才知道自己卡在哪 黑字:前向的值 红字:反向的梯度 Ⓐ 式子是 f = (a + b) × c 同一条线上两个数:上面黑的是前向算出来的,下面红的是反向传回来的 a = 2 梯度 4 b = -3 梯度 4 c = 4 梯度 -1 × f = -4 梯度 1 q = a + b -1 ↤ 4 前向从左往右走一遍,反向沿同一批线从右往左走一遍 —— 红字就是「这个东西变一点,f 变多少 Ⓑ 常见的门只有三种脾气 —— 而这三种脾气就是三种连线的形状 不用背公式 —— 线越粗梯度越大,虚线表示负的,三张图的骨架完全一样,只看红线怎么走 + 加法门 像一个分发器 a b 上游 4 4 4 两条一样粗 —— 原样各拿一份 × 乘法门 像一个交换器 a b × 上游 1 4 -1 两条交叉 —— 各自拿对方的前向值 max 门 像一个路由器 a b max 上游 4 4 ✕ 0 一条断了 —— 全给赢的那个 Ⓒ 而中间那个门,正是这一整讲的账单的源头 别跳过这一格 —— 后面所有关于「激活」的话都从这儿来 × 反向要乘它 所以扔不掉 就这一个 每一层 每一步 ……于是变成这么一片 每一个小方块都是同样一件事 —— 某一层某一步存下来的那个前向值 所以激活不是「框架顺手缓存的东西」 —— 是反向的数学要求它在场 所以「重算」那个决定,在这张小电路上也说得通 把 q 扔掉,反向要用的时候临时把 a + b 再加一遍 —— 一次加法,换一个数的存储空间。这就是重算,整套逻辑一个字都不用改。 而这也解释了为什么加法便宜、矩阵乘贵:重算一个加法门几乎不要钱,重算一个矩阵乘要把那一大坨乘法再做一遍。
⭐⭐⭐ 这是整讲唯一一处「小到能用眼睛跟着数走」的地方 —— 同一条线上两个数:黑的是前向,红的是反向。
⭐⭐ Ⓑ 那三个门不用背:加法是分发器、乘法是交换器、max 是路由器。记住脾气就能手推任何一张电路。
Ⓒ 别跳过 —— 乘法门反向要用前向的值,而这就是整讲那张激活账单的源头:激活不是框架顺手缓存的东西,是反向的数学要求它在场。
出处与口径

📌 「电路图 + 三个门(分发 / 交换 / 路由)」这个讲法取自 CS231n 的反向传播讲义 —— 图是我们自己重画的,数也是自己挑的

⭐ 图上每个数都可以自己验:脚本里带 assert —— a 和 b 的梯度都是 4、c 的梯度是 −1,对不上就不让构建。

1.4 ⭐ 为什么偏导数一出现,这件事就变容易了

⭐⭐⭐ 因为偏导数是局部的。

「偏」这个字,就是全部的关窍。

普通的导数问的是:这个东西变了,结果变多少。
偏导数问的是:其它全部按住不动,只动这一个,结果变多少。

⭐ 而正因为其它都按住了,这个问题就缩到了一个算子身上 —— 它不需要知道外面的世界长什么样。这就是「局部」的意思。

每一个算子只需要知道「我自己是怎么把输入变成输出的」, 就能写出自己的反向规则。 它完全不需要知道前面是什么、后面是什么、整个网络长什么样。

⭐⭐ 这就是为什么自动微分能做成一个通用库

矩阵乘写一个 backward,softmax 写一个 backward,加法写一个 backward —— 然后框架只干一件事:按顺序倒着把它们串一遍。

⭐ 没有任何人需要手推整个网络的导数。

反过来想才知道这有多可怕: 如果你非要写出「loss 对第一层某个权重」的解析式, 那个式子要穿过 61 层展开 —— 项数是天文数字,写不出来。

⭐⭐ 链式法则真正的意思是:你永远不用把它展开。 你只要把 61 个局部的小导数,按顺序乘起来。

—— 那「乘起来」之后会怎样? 这一问有一个非常出名的答案,而它是链式法则最直接的后果

一路乘下去会怎样 —— 梯度消失,和梯度爆炸 上一格说「把 61 个小导数按顺序乘起来」—— 这一格说乘起来之后发生了什么 每层 ×0.8:消失 ×1.0:刚好 ×1.2:爆炸 Ⓐ 同一个 60 层的网络,只改「每层的兑换率」—— 三条线都是真乘出来的 横轴左边是靠输入的层,右边是靠输出的层;纵轴是那一层拿到的梯度有多大 1 0 第 1 层(最靠输入) 第 60 层(最靠输出) 冲出画面 这条贴在轴上 —— 不是画漏了,是真的小到画不出来 Ⓑ 换一根对数轴才装得下 —— 两头差了 10 个数量级 而那两个兑换率只差两成 10^-6 10^-4 10^-2 10^0 10^2 10^4 每层 ×0.8 ≈ 1.9 × 10⁻⁶ 每层 ×1.0 = 1 每层 ×1.2 ≈ 4.7 万 相差 10 个数量级 所以被指数放大的不是梯度,是「每层偏离 1 多少」那个微小的偏差 —— 偏一点点,乘六十次就回不来了。 Ⓒ 那「每层的兑换率」从哪来 —— 其中一截是激活函数的导数 这两条是真画的函数曲线,不是示意 导数 输入 → 1.0 0.25 ReLU 的导数 正半轴恒等于 1 sigmoid 的导数,峰值就这么高 sigmoid 的导数最大是多少 σ(1 − σ) 在 σ = ½ 处最大 0.25 光这一下,每层就先乘了个 ≤ ¼ 的数 权重那一乘还在外面 —— 所以不能说 「一定衰减四倍」,只能说它先天往下压 这就是为什么 ReLU 一换上来,深网络忽然就训得动了 —— 它把每层那个「先天往下压」的系数,从 ¼ 变回了 1。 这一格只回答「会怎样」和「为什么」—— 治法在本讲后面,各有各的位置 要记的只有一句:链式法则是连乘,而连乘是指数的。所以真正要盯的不是「梯度大不大」,是「每层那个兑换率离 1 有多远」 —— 离一点点,乘几十层就回不来了。 而两头的症状完全不同:消失是静悄悄的(loss 就是不降,前面几层跟没训一样,没有任何报错);爆炸是吵闹的(loss 直接飞掉或者变 NaN)。 所以爆炸好查,消失难查。
⭐⭐⭐ 链式法则是连乘,而连乘是指数的。同一个 60 层网络,每层兑换率 0.8 还是 1.2 —— 只差两成,最靠输入那层拿到的梯度差了 10 个数量级。
⭐⭐ Ⓐ 那条贴在轴上的红线不是画漏了,是真的小到画不出来;Ⓑ 得换一根对数轴才装得下这两头。
Ⓒ 是这一格唯一能当场算清楚的那截原因:sigmoid 的导数最大只有 0.25(σ(1−σ) 的闭式最大值),光这一下每层就先乘了个 ≤ ¼ 的数 —— 而 ReLU 正半轴恒为 1。深网络忽然训得动,就是这么来的。
出处与口径

⭐ Ⓐ 三条曲线、Ⓑ 那两个端点值、Ⓒ 那个 0.25,全是脚本算的,没有一个是抄来的:0.8 的 59 次方 ≈ 1.9e-06、1.2 的 59 次方 ≈ 46956,两者相差 10 个数量级;0.25 是 σ(1−σ) 的闭式最大值,脚本里还用一万个采样点复核了一遍。

⚠️ Ⓒ 只挑了激活函数这一截 —— 每层真正的兑换率还要乘上权重矩阵那一下(以及归一化那一下)。选它是因为它是这一格里唯一能当场算清楚的,⛔ 不是因为它是唯一的原因。

📌 出处与年份(完整的故事在讲义里):这件事 1991 年就被正式指出了 —— Hochreiter 的硕士论文《Untersuchungen zu dynamischen neuronalen Netzen》(慕尼黑工业大学)。⭐ 它是德文写的,从没在英文期刊上发表过。三年后 Bengio、Simard、Frasconi 在 IEEE Trans. Neural Networks 5(2):157–166 发表《Learning long-term dependencies with gradient descent is difficult》—— 标题就是结论

⚠️ 这条的归属至今仍有争议,本讲只说「1991 年那篇论文正式指出」,不卷进优先权之争。⭐ 而有意思的是结局:2001 年 Hochreiter、Bengio、Frasconi、Schmidhuber 四个人合写了一章《Gradient flow in recurrent nets》—— 当年分头发现同一件事的两拨人,十年后坐到一起把它总结了。

⚠️ 还有一条口径要说清:这个问题最早是在循环网络上发现的,不是在我们这一格画的前馈网络上。⭐ 但它们是同一件事 —— 循环网络按时间展开之后,就是一个「层数 = 序列长度」的超深网络。那篇综述里报的实验跨度是 1000 步,换算过来就是一千层。

⭐⭐⭐ 对数轴让你看得见十个数量级,也让你感觉不到十个数量级有多大。 两条轨共用一根横轴、一根扫描线 —— 所以它们永远在说同一层
上轨是对数刻度:三条线才刚刚张开,看着「还差不多」。
⭐ 下轨是线性刻度,也就是这些数真实的样子第 8 层蓝线就顶出了画框,第 18 层红线已经贴在地上看不见了 —— 而同一层,上轨那把扇子才张开三成
⛔ 每层只差两成的兑换率,十几层就已经没法放进同一张图里了。 (14 秒无声循环,Manim 渲染,脚本在 tools/manim/。 出框层与消失层都是脚本当场算的,并有断言钉住。)

1.5 ⭐⭐ 最深的一条:为什么是从后往前

数学上两个方向都成立,算出来的结果一模一样⭐ 区别只有一个 —— 你得把整条链走多少遍。

为什么是从后往前 —— 因为那一头只有一个数 这里的关键不是链式法则怎么乘(乘法从哪头开始都一样)—— 是你得重复几遍 正向:插在输入端 反向:插在输出端 Ⓐ 正向模式 —— 种子插在左边,有几个参数就要插几次 每插一次,就要把整条链从头走到尾一遍 嵌入 第 1 层 第 2 层 最后一层 loss 第 1 遍 第 2 遍 第 3 遍 3000 亿 每一遍走完,只拿到一个参数的梯度 —— 因为你这一遍只扰动了它一个 所以它的代价跟参数个数成正比。三千亿个参数,就是三千亿遍整条链。 画成从最左边插只是为了整齐 —— 参数分布在每一层,种子就插在它所在的那一层。不变的是「一个参数一遍」。 Ⓑ 反向模式 —— 种子插在右边,而右边只有一个数 所以只插一次,而且回来的路上每个参数的梯度顺手就到手了 嵌入 第 1 层 第 2 层 最后一层 loss 种子 = 1 只有这一个 梯度到手 梯度到手 梯度到手 梯度到手 梯度到手 走到最左边的时候,三千亿个梯度已经全部在手里了 —— 而你只走了一遍 一路上传的始终是一个东西(那个种子沿途被改写),所以那些巨大的中间雅可比从来不用真的算出来 Ⓒ 换个看法:反向那一遍,其实就是同一张网倒着走 不是「另一套算法」—— 同样的形状,只有三处换了 W W Wᵀ Wᵀ 前向 反向 同一张网 只有三处换了 这个常数是 σ′(z) 前向那一遍算好的 ① 每条边的箭头掉了个头 ② 同一个 W,转置过来 —— 没有新参数 ③ 节点从「套非线性」换成「乘一个常数」 落点在第 ③ 处:那个常数在前向就定下来了 —— 所以前向算出来的东西必须留在场上,反向才有东西可乘。 Ⓓ 所以规则只有一条,而且它跟神经网络没关系 别把它记成「反向传播永远更好」—— 换个形状它就反过来 从窄的那一头起步。 三千亿个参数 → 1 个 loss 输入多、输出少 → 从后往前(就是反向传播) 这就是我们的情况 10 个参数 → 100 万维输出 输入少、输出多 → 从前往后 这时候反向传播反而亏 判据:看到任何一个「要求一大堆偏导」的问题,先数输入多还是输出多 它决定的不是「算得对不对」—— 两个方向算出来的结果一模一样,决定的是你要把整条链走多少遍 而这正是推理里没有的那一半:推理只有一遍前向,压根不存在「往回走」这个动作,也就不存在「沿途把梯度收下来」这件事。
⭐⭐⭐ 两排只差一件事:种子插在哪一头 —— 而这一件事,就把代价从「三千亿遍」压到「一遍」。
⭐⭐ Ⓑ 里每个算子底下那个 ✓ 是关键:回来的路上顺手就收下了 —— 不是走完再统一算。
⭐⭐⭐ Ⓒ 是另一个看法:反向那一遍就是同一张网倒着走 —— 同样的形状,只有三处换了。第三处最要紧:每个节点从「套一个非线性」换成「乘一个常数」,而那个常数在前向就定死了。
Ⓓ 的反例一定要看:「反向传播更好」是条件成立的,条件就是输入多输出少。
出处与口径

⭐ 这一格是结构,不是数据 —— 图上唯一的量是「三千亿」,那是模型规模,前面已经立过。

📌 「正向模式 / 反向模式」是自动微分的标准术语;反向传播是反向模式用在神经网络上的那个特例。

⭐⭐⭐ 两条链同时开跑 —— 看谁先跑完。 上面是正向模式:小球一次只能为一个参数跑一趟,跑完才点亮那一个, 于是八个参数就得跑八趟(右边堆了八块)。
下面是反向模式:小球从 loss 出发往回走一趟,沿途所有参数一起点亮 —— 一趟就完事,右边永远只有一块。
⭐ 注意它早早就停在那儿了,而上面那条还在吭哧 —— 那份等待就是三千亿倍代价的样子。 (9 秒无声循环,Manim 渲染,脚本在 tools/manim/。)

⭐⭐⭐ 所以反向传播成立的全部理由就一句话: 参数有几千亿个,而 loss 只有一个。

多输入、单输出。 ⛔ 如果 loss 不是一个数、而是一百万维的输出,这套就不划算了 —— 那时候反倒该用从前往后。

⭐⭐⭐ 所以 1.1 那个问题的答案是:那三千亿次前向, 压成的是一次反向

这条判据可以直接迁移: 看到任何一个「求一大堆偏导」的问题,先数一下输入多、还是输出多 —— 它决定了你该从哪头开始。

1.6 闭环:反向必须用到前向的值 —— 这就是激活扔不掉的根本原因

回头看 1.3 那两件事里的第①件: 算自己权重的梯度,要用「前向时存下来的输入」。

⛔⛔ 这就是激活必须留着的根本原因 ——  不是谁设计得不好,是反向的数学本身要求它在场

前向每算出一个中间结果,都得一直挂在那儿, 等反向走回来的时候用。整条网络走完,它们全都还在。

⭐⭐ 这一节的两笔账到此都立住了: 算力 (前向 1 + 反向 2); 显存里多出一整条从头挂到尾的激活 而下一节要做的,就是拿第一样去换第二样。

激活是一座山 —— 前向一路堆,山顶在 loss 那一刻 横轴是时间,不是层号 · 权重和优化器状态没画 —— 它们是水平的,会把山形压扁 前向 · 堆 峰值 反向 · 拆 Ⓐ 每前进一层,就多挂一份中间结果 —— 而且一份都不能提前扔 为什么不能扔:反向算权重梯度时要用前向那一刻的输入 显存里的激活 前向:第 1 层 → 第 61 层 反向:第 61 层 → 第 1 层 每过一层,多挂一份 每走回一层,释放一份 峰值在这一刻 前向刚算完、反向还没开始 Ⓑ 这座山有多高 —— 一条 128K 序列,V3 那个规模 自己按算子推的估算,当量级看,别当准数 不开重算 4.15 TiB 一整条全挂着 光这一项就已经装不下 开了全量重算 106.75 GiB 每层只留入口那一份 约 40 倍的差距 两个框的高度是按真实比例画的 —— 右边那条薄片就是重算之后剩下的厚度 下一节整节都在讲这两栏之间那个箭头 这张图顺带把两件事一起讲了 「为什么训练比推理贵」:算力只贵 3 倍,可推理根本没有这座山 —— 它算完一层就把中间结果扔了,只留 KV cache。真正拉开差距的是显存,不是算力。 「峰值出现在哪一刻」:就是山顶那一竖 —— 前向刚结束、反向还没开始。 到第五节会把权重和优化器状态那两条水平带叠上来,山形不变,只是整体抬高。
⭐⭐⭐ 横轴是时间,不是层号 —— 换成层号,「什么时候最挤」这个问题就提不出来了。
山顶那一竖同时回答了第五节要问的「峰值在哪一刻」:前向刚算完、反向还没开始。
出处与口径

⚠️ 台阶画了 12 级只是为了看得清 —— 真实是 61 层,而且每层内部还有若干个中间张量,山坡比图上细密得多

⛔ 山形画成直上直下是简化:真实曲线会因为 MoE 派发、attention 那几个大中间量而有凸起,但「顶点在前向末尾」这个结论不受影响

⚠️ 4.15 TiB 与 106.75 GiB 两个数是自己按算子推的(输入:V3 的 config + 官方参考实现的 MLA 前向),没有第三方背书

📌 图 Ⓑ 那两张卡片的读法: 它们是同一条序列的两种配置,不是两个模型。

⚠️ 两个数都是自己按算子推的(输入是 V3 的 config + 官方参考实现的 MLA 前向),没有第三方背书⭐ 但就算差一倍,结论也不变:这个规模上,装不下。

1.7 把那张原始账单摊开 —— 一层里到底挂了些什么

⭐ 这一小节可以整段跳过。 它是给要自己动手算的人看的 —— 上面那两个数从哪来,全在这儿。

基准单位先立住:一份 hidden 宽的张量 = 131,072 × 7,168 × 2 B1.75 GiB。 下面所有数都是它的倍数。(序列 131,072、batch 1、bf16。)

MLA 子层留什么

留下的张量宽度大小
norm 的输入(= 残差入口,重算模式下唯一留的那份7,1681.75 GiB
norm 的输出7,1681.75 GiB
Q 降维结果(layernorm 前后各一份)1,5360.75 GiB
KV 降维结果(含 RoPE 那 64 维)576 / 5120.27 GiB
Q 展开后128 头 × 192 = 24,5766.00 GiB
K、V 解压后128 头 × 256 = 32,7688.00 GiB
attention 输出128 头 × 128 = 16,3844.00 GiB
logsumexp(fp32)1280.06 GiB
小计约 22.6 GiB

MoE 子层留什么

留下的张量份数 × 宽度大小
norm 的输入 / 输出2 × 7,1683.50 GiB
路由分数2560.06 GiB
派发出去的激活9 份 × 7,16815.75 GiB
gate / up / SwiGLU 乘积3 × 9 份 × 2,04813.50 GiB
专家输出(合并前)9 份 × 7,16815.75 GiB
小计约 48.6 GiB

⭐⭐⭐ 「9 份」是这张表里最值得停一下的地方。

9 = 被激活的 8 个路由专家 + 1 个共享专家。 每个 token 的那份激活被复制了九遍 —— 进去一次、出来一次, 光这两项就 31.5 GiB

⛔ 而它在专题一那张参数量的表上完全看不出来—— MoE 在参数账上很划算(只激活一小部分), 可在激活账上,它要付九份的复制费。

⭐ 合起来:一层 MoE 块约 71 GiB,一层 dense 块约 40 GiB; 58 层 MoE + 3 层 dense ≈ 4,245 GiB ≈ 4.15 TiB

⚠️ 这张表的三条边界,讲的时候必须说清楚。

  • 它是按算子逐项推的,不是从哪份报告抄的。 不同框架的算子融合程度、是否顺手重算便宜算子、 MoE 派发是否真的物化九份副本,差别都很大 —— 当量级看,别当准数。
  • 注意力分数矩阵没算在里面。 专题一算过它在 128K 上是 4 TiB/层。 FlashAttention 压根不把它写进显存,所以它不出现在这张表上 —— ⛔ 换成不带 flash 的朴素实现,这张表整个作废。
  • 那个 106.75 GiB 是「已经做过一次交易之后」的账。 它假设每层只留入口那一份 —— 而那恰恰就是开了重算之后的样子。 ⭐ 所以先看原始账单、再看下一节,顺序不能颠倒

1.8 ⭐ 梯度算出来之后,还没完

⭐ 1.3 只讲到梯度算出来那一刻就停了。 可它离「被用掉」还隔着好几道 ——  而这几道每一道都是推理里不存在的。

梯度的一生 —— 从算出来到被扔掉,中间要过五道 看三条泳道:做什么能不能藏推理里有没有 能藏进计算里 藏不了 推理里都没有 Ⓐ 五道依次走完,而其中只有一道能跟计算叠在一起 也只有一道会把所有卡钉在一起等 算出来 反向走到哪层,哪层的梯度就 出来 从后往前陆续出来,不是最后一起出 推理里:没有 跨卡汇总 每张卡看的数据不同,得取平 能藏 前面几层还在算,后面的梯度 已经能传了 推理里:没有 裁剪 按全局范数,超了整体缩回去 藏不了 要等所有卡所有层都到齐,强制一次同步 推理里:没有 累积 几个 micro-batc h 先加起来(可选) 它改的是时间线,不是总量 推理里:没有 用掉就扔 交给优化器,这一份就没用了 正因为用完就扔,它才敢用 bf16 存 推理里:没有 第三条泳道是这张图的落点:这五道,推理里一道都没有 —— 它们整条都是训练才有的东西。 而这也是为什么「训练比推理贵」不只是「多跑一遍」 —— 多出来的是一整条流水线 只有第 ③ 道会把所有卡钉在一起等 —— 而它传的其实只是一个数 「全局范数」是跨全部参数的一个标量:得先把每一块的平方和都算出来、汇总成一个数、开根号,才知道要不要缩、缩多少。通信量可以忽略,代价是那一次同步。 所以按全局范数裁剪天然是「一步一次」。 但别读成「逐层裁剪不存在」—— 它存在,只是保的不是同一个量:逐层保的是每层各自的范数,全局保的是这一步总共迈多大。
⭐⭐ 看三条泳道,别只看第一条 —— 第一条是「做什么」,第二条是「能不能跟计算叠在一起」,第三条才是落点。
⭐⭐⭐ 第三条泳道全是紫的 —— 这一格是现场那条界(推理里没有的那些东西)最直观的一次落地。
而五道里只有裁剪那一道会把所有卡钉在一起等,偏偏它传的只是一个数。
出处与口径

⚠️ 图上不含任何量 —— 通信到底多大、能藏掉多少,是 fig-batch 和专题五的事。

⭐ 第 ⑤ 道那句「用完就扔」是本讲另一处的伏笔:正因为它一次性,它才敢用 bf16 存 —— 而一旦要做累积,它就变成累加量,得升回 fp32。

⭐⭐ 第 ② 步那笔通信,有个很反直觉的性质

纯数据并行下,每一步要把整份梯度在所有卡之间汇总一遍。 量级很好算:参数量 × 每参数字节数

⭐⭐⭐ 关键在这儿:这笔通信量跟 batch 完全无关—— 参数有多少就传多少,你喂一条序列和喂一千条,传的是同样多的字节

⭐⭐ 而计算量是随 batch 线性涨的。所以:

batch 越大,这笔通信被摊得越薄。

这一下就解释了两件原本看着不相干的事:

  • 为什么大模型都在爬 batch。 GPT-3 从 32K token 一路爬到 3.2M,PaLM 从 1M 翻到 4M, V3 从 3072 条序列爬到 15360 条(4K 上下文,折合 12.6M → 62.9M token) —— 它们全都在往上爬,而且都写着「越到后面越大」。
  • 为什么梯度累积除了省显存还顺带便宜。 累积 K 个 micro-batch 才更新一次,就是把通信次数除以了 K —— 它省的不只是激活。
「batch 开大就更快」 —— 这句话有个前提,而它通常不写 这一格画的不是结论,是结论的适用区间 左段:没吃满 右段:吃满了 通信:两段都无关 Ⓐ 横轴 batch、纵轴一次更新要多久 —— 这条线是折的,不是直的 大多数讲法只画了左半段 batch → 一次更新的时间 并行度在这儿吃满 几乎是平的 batch 翻倍,时间几乎不变 正比于 batch batch 翻倍,时间也翻倍 卡还没喂饱 卡已经满了 左段 batch 开大是白赚的(步数变少,每步没变贵);右段就不白赚了(步数变少,每步同比变贵)。 Ⓑ 同一句话,在两段里一句成立、一句不成立 而教程里那句「大 batch 更快」,说的几乎都是左段 左段 并行度没吃满 「跑完固定的数据量,大 batch 更快」 成立 —— 步数少了,而每步没变贵 小模型、单卡、教学例子,多半在这儿 右段 并行度吃满了 同一句话 不成立 —— 每步变贵的倍数,正好抵掉步数少的倍数 几百亿参数、几百张卡的训练,基本在这儿 Ⓒ 但通信那一笔不分左右段 —— 它跟 batch 完全无关 所以「batch 越大摊得越薄」这句话,对通信永远成立 每步要汇总的梯度 = 参数量 × 每参数字节数 喂一条序列和喂一千条,传的一样多 所以它只会被摊薄 batch 翻倍 → 每 token 分摊的通信减半 这一条在左段右段都成立 这一格真正要留下的,不是「大 batch 快不快」,是那个「看情况」怎么看 计算和通信在这件事上行为不一样:计算分两段(没吃满时白赚,吃满后不赚),通信不分段(永远摊得越薄越好)。—— 所以「batch 该开多大」不是一个数,是看你现在卡在 哪一栏 判据(给自己的):借别人的讲法时,连同它的前提一起借。 只搬结论、把前提留在原处,是本讲已经踩过的坑。
⭐⭐⭐ 先看那条线是折的 —— 整张图的全部信息就在那个拐点上。
看图的顺序建议:先 Ⓐ 找到拐点,再 Ⓑ 对照两边各自成立什么,最后 Ⓒ 看通信那条为什么不长这样
⚠️ 这张图一个实测数字都没有,只画形状 —— 拐点落在哪,取决于你的模型、卡和并行配置,得自己量。
出处与口径

📌 「有平行运算时,大小 batch 跑一次的时间差别不大;而一个 epoch 反过来」出自李宏毅2021 年《类神经网络训练不起来怎么办(二):批次与动量》。⛔ 图是我们自己重画的。

⚠️ 图上不含任何实测数字 —— 只画两段的形状。拐点落在哪,取决于模型、卡、并行配置,得自己在目标配置上量

⭐ 那条「能藏」的泳道,值得单独说一句

⭐⭐ 「算」和「传」能叠在一起,是训练里最重要的一类优化 —— 而推理侧没有对应物:它没有反向, 也就没有这条可以边算边传的长尾巴。

⭐ 裁剪那一道为什么藏不了、以及「逐层裁剪」为什么是另一件事, 图下面的落点带写了。这一节只留一句: 「这一步总共迈多大」才是跟 3.3 对话的那个量。

📌 顺带把 ④ 说清楚:梯度累积到底省什么

⭐ 一句话:梯度累积改的是时间线,不是总量

这一节的边界要说清楚: 上面只讲了「这一步存在、它有多大、能不能藏起来」。 至于在不同的切法下它究竟是 all-reduce 还是别的形态、量各自差多少 —— 那是专题五整讲的事。

⚠️ 通信量那个「参数量 × 每参数字节」是量级口径: 真实的环形汇总要来回搬,实际搬运量大约是它的两倍; 而 MoE 的专家部分走的又是另一套。当量级看。

第 二 节

第一个真正的「决策」—— 拿算力换显存,换多少算划算

这是全课第一次出现真正的取舍不是「有没有更好的办法」,而是两样东西只能选一样,你选哪个。

2.1 重算是什么 —— 一句话:不存过程,只存存档点

前向每算出一个中间结果就留着,是因为上一节那条 —— 反向要用。重算说:我不留了。

反向要用的时候,我从这一层的入口再往前跑一遍,现算出来。 ⭐ 就像游戏存档 —— 不存整个过程,只存一个存档点, 要用的时候从那儿重打一遍。

2.2 全量重算:这笔交易到底划不划算

⭐ 先把量级说清楚 —— 这才是要记住的东西。

⭐⭐⭐ 这个兑换比例,夸张到不像是个「权衡」

所以在大模型训练里,它默认就是开着的 ——  值得讨论的从来不是开不开,而是开到哪一档(图 Ⓑ 那三档)。

全量重算 —— 用三分之一的算力,换掉四十分之三十九的显存 两边都归一到「原来 = 100%」 · 否则一边 TiB 一边 TFLOPs,两根条没法比 付出:算力 省下:显存 Ⓐ 这笔交易的两边 —— 形状不对称到不像是个权衡 条长代表「相对原来剩多少 / 涨多少」 付出 算力 原来:3×(前向 1 + 反向 2) 现在:4×(多跑一遍前向) +33% 多出来的那一小截,就是「再跑一遍前向」 省下 显存 原来:4.15 TiB 激活 现在:106.75 GiB —— 只剩 2.5% 这根条短到几乎看不见 —— 那正是这张图要说的事 Ⓑ 所以值得讨论的从来不是「开不开」 而是开到哪一档 不开 全留 只在显存宽裕时才合理 选择性 按比值挑着扔 约 1.9% 算力换约 77% 显存 —— 性价比最高 全量 每层只留入口那一份 33% 算力换 97% —— 显存实在不够时 显存本来就宽裕的时候(小模型、短序列),重算就是纯亏 Ⓒ 换一个轴:隔几层留一个存档点 注意这跟 Ⓑ 不是同一件事 —— Ⓑ 是一层之内留哪些,这里是隔几层留一个 全存 每一层都留 💾 💾 💾 💾 💾 💾 💾 💾 💾 💾 💾 💾 💾 💾 💾 💾 存 16 份 + 段内 1 层 = 17 隔 4 层留一个 L = 16,而 √16 = 4 💾 💾 💾 💾 存 4 份 + 段内 4 层 = 8 只存开头那一个 重算时整条 16 层都得在场 💾 存 1 份 + 段内 16 层 = 17 同时在场 一个方块 = 一层。💾 就是留下来的存档点 两头一样高 —— 存得越少,并不是越省。少存一个存档点,就要多扛一段重算时的中间结果。
⭐⭐ 要看的是两根条的不对称 —— 付出那边只多出一小截,省下那边短到几乎看不见。
⚠️ 两根条都归一化了,所以可以直接并排看 —— 单位不同的量并排放,读者第一反应是去比长度,而那个比较没有意义。
⭐⭐⭐ Ⓒ 换了另一个轴:Ⓑ 问的是「一层之内留哪些」,Ⓒ 问的是「隔几层留一个」。看那三个「同时在场」—— 17 / 8 / 17,两头一样高。
所以最优点在中间,而它落在 √L 上 —— 这也是 ZeRO 论文说「把激活降到大约总量的平方根」的由来。
出处与口径

⚠️ 4.15 TiB / 106.75 GiB 是自己按算子推的估算(V3 的 config + 官方参考实现的 MLA 前向),当量级看;就算差一倍,「不对称」这个结论也不变

⭐ 「3× → 4×」是矩阵乘口径:重算等于把前向那一遍再买一次,所以多出来的正好是 1/3 —— ZeRO 论文原话也是这个数(33% re-computation overhead)

⭐ Ⓒ 那三个「同时在场」由 L/k + k 当场算出,脚本内 assert 最优点落在 √L、且两头等高 —— 这一格的全部意思就是那个等高

⛔ Ⓑ 里那两个百分比都按「一个 step」做分母(前向 + 反向 = 3 遍)。早先「选择性」那一档写的是 6%,那是拿一遍前向当分母的旧口径—— 并排摆着不能比,已统一。

📌 Ⓑ 里「选择性」那一档的 1.9% / 77%,推导过程与名次表见下一张图与本节正文;DeepSeek-V3 报告(arXiv 2412.19437)明写他们重算全部 RMSNorm 与 MLA 上投影 —— 正是那一档

2.3 ⭐ 选择性重算:判据只有一个数

对每一个中间张量问同一个问题把它扔掉、反向时重算回来,每省一个字节要付多少次浮点运算?

比值越小越该扔。就这一个数,整张排序表自己就出来了。

⭐⭐ 对线性层,这个数有闭式解,而且漂亮得出乎意料:

一个 [S,k] × [k,n] 的矩阵乘 ——  重算代价 2·S·k·n FLOPs,产出张量 S·n·2 字节(bf16)。 一除,每字节代价 = k。

⭐⭐⭐ 重算一个线性层的输出,每字节要付的 FLOPs 恰好等于它的输入宽度Sn 全约掉了 ——  跟序列长度无关,跟输出多宽也无关。

所以「哪个线性层最该重算」这个问题,答案是看谁的输入最窄 —— 不用算,扫一眼 config 就知道。

2.4 但 attention 不服从这条 —— 而且它必然会跟线性层交叉

attention 不是线性层:代价随 S² 涨,产出只随 S 涨。 一除,每字节代价随序列长度线性上升

为什么同一条判据会给出相反的答案 —— 两条斜率不同的线,必然相交 纵轴:扔掉它、再算回来 —— 一个字节的代价 · 双对数坐标 —— 常数在这里是水平线,正比于 S 的是直线 便宜 → 该扔 贵 → 该留 attention Ⓐ 横轴是序列长度 —— 整张图只有这一个变量在动 线性层那几条是水平的:它们的代价跟序列长度无关 1K 4K 16K 64K 256K 序列长度(token) 1K 10K 100K FLOPs/字节 K/V 解压(输入宽 512) Q 展开(输入宽 1,536) 专家输出(输入宽 2,048) gate / up / 路由(输入宽 7,168) V3 的 attention = 1.25 × 序列长度 GPT-3 的 attention = 1.00 × 序列长度 GPT-3 最宽的线性层(12,288) GPT-3 自己的交点 12,288 交点 S ≈ 5,734 左边 attention 最便宜,右边最贵 GPT-3 2,048 → 2,048 FLOPs/字节 V3 131,072 → 163,840 FLOPs/字节 同一条判据,结论翻转 —— 而两个模型的线都不一样 2022 年那篇(arXiv 2205.05198)说:挑「占显存不少、但重算起来不贵」的扔,它选中了 attention。 看橙色那一组:GPT-3 在 2,048 处只要 2,048,而它自己的交点在 12,288 —— 远在左边,attention 确实最便宜。论文没错。 再看紫色那一组:V3 在 131,072 处是 163,840,而它自己的交点在 5,734 —— 远在右边,attention 成了最该留的那一个。 判据一个字没改 —— 而且注意:两个模型的斜 率和对照宽度都不一样,翻转靠的是「斜率不同必然相交」这件事本身,不靠任何一组具体数值。
⭐⭐⭐ 全专题最值钱的一张 —— 「两条斜率不同的线必然相交」这句话,图讲一秒,字讲一段。
双对数不是为了好看:只有在这个坐标上,「常数」才是水平线、「正比于 S」才是直线。
两根竖线就是那场翻转的全部原因 —— 判据一个字没改,变的只有序列长度。
出处与口径

⭐ 线性层那条闭式解:一个 [S,k]×[k,n] 的矩阵乘,每字节代价 = k(重算 2·S·k·n FLOPs ÷ 产出 S·n·2 字节,S 和 n 全约掉)—— 所以它在图上必然是水平线

⚠️ 两条斜率分别是 V3 的 1.25(qk 192 / v 128)与 GPT-3 的 1.00(qk = v = 128);口径都是 causal 折半。⛔ 换个模型就得重画它自己那两条线(qk 192 / v 128)与 causal 折半的口径;换个模型斜率会变,但「它是斜的」不会变 —— 这才是要记的东西

⛔ 交点 5,734 不是工程阈值,它还取决于你拿哪条线性层做对照(图上四条给的交点各不相同)

📌 2022 那篇的收益数字(GPT-3 省 70% 付 2.7%)用的是未折半的 Megatron 口径;换成本图统一的 causal 口径,GPT-3 的 attention 占比是 1.37%

⭐⭐⭐ 图上那两条斜率不同的线,就是这一节的全部内容 —— 一条平的(线性层,代价 = 输入宽度),一条在爬的(attention)。 斜率不同,就必然相交。

⭐⭐⭐ 一条是常数(输入宽度),一条在往上爬 —— 那它们必然相交

对 V3:1.25·S = 7168 → S ≈ 5,734

  • 短于 ~5.7K —— attention 比大多数线性层还便宜, 是个好的重算对象
  • 长于 ~5.7K —— 它一路变成最贵的那个,而且线性地越来越贵

⚠️ 5,734 是个粗略交叉点,不是工程阈值。 它依赖 causal 折半的口径、依赖 V3 的头维度、依赖你拿哪个线性层做对照。 ⭐ 要记的是那句「两条线斜率不同所以必然相交」,不是这个数。

把 V3 在 128K 上排一遍 —— 一层 MoE 块,按「每 GiB 要付多少 TFLOP」升序:

张量省显存重算代价每 GiB 付结论
MoE 派发(复制成 9 份)15.75 GiB~00白捡
RMSNorm 输出(每层 2 个)3.50 GiB0.01 TFLOP~0白捡
SwiGLU 乘积 9 份(逐元素)4.50 GiB0.01 TFLOP~0白捡
K/V 解压(输入宽 512)8.00 GiB4.40 TFLOP0.55划算
Q 展开(输入宽 1,536)6.00 GiB9.90 TFLOP1.65划算
专家输出 9 份(输入宽 2,048)15.75 GiB34.63 TFLOP2.20划算
gate / up / 路由 / 降维(输入宽 7,168)9.14 GiB70.35 TFLOP7.70边际
attention 输出4.00 GiB703.69 TFLOP175.92绝不

⭐ 这张表不用背,它是上面那条闭式解直接排出来的 ——  前六行的「每 GiB 付」就是各自的输入宽度换了个单位, 你拿 config 自己也能排一遍。 ⛔ 顺带自曝:这张表早先有一行是手填的, 「省显存 9.14 × 每 GiB 7.70」算出来是 70 TFLOP,那一栏却写着 39.08。 而这张表的卖点正是「你自己也能排一遍」 ——  照着排的人会正好撞上它。现在整张表由闭式解生成,手改不了。

⭐⭐ 回答那个直觉问题:「有没有占显存很小、算力却很大的东西?」 —— 有,就是 attention,而且极端。 它比排在前一档的贵 23 倍

⛔ 但结论是反过来的:正因为如此,它是最该留着的那一个,不是最该重算的。 判据只有一条比值 —— 比值高的留,比值低的扔。

「小于 3」这个 3 是哪来的?—— 我定的,不是算出来的。

它在表上的位置很清楚:专家输出那一行是 2.20,下一行 gate/up 是 7.70 —— 3 就落在这两行中间那个空当里,放 2.5 或者 5 结果完全一样。

⭐⭐⭐ 而真正定这条线的东西不在这张表上它是你现在的兑换率 —— 一 GiB 显存对你值多少 TFLOP
显存快爆了就往右挪(多扔几项,多付点算力),算力是瓶颈就往左挪。

⛔ 所以别把 3 记成阈值 ——  要记的是「表是按比值排的,你只需要决定在哪儿切一刀」。 ⭐ 排序是客观的,切在哪是你的配置说了算。

⭐⭐ 照这条线切(把比值小于 3 的全收下来):一层省 53.5 GiB,付 48.9 TFLOP。 换算到一个 step(前向 + 反向 = 3 遍 ≈ 2,571 TFLOP): 只多付 1.9% 的算力,换掉约 77% 的激活显存

对比全量重算 —— 同一个 step 口径下是「多付 33%,换掉 97%」。 ⭐ 选择性重算的性价比高一个数量级,这就是它值得单独配的原因。

⚠️ 这两个百分比早先分母不一样(一个除一遍前向、一个除整个 step), 并排比是不能比的 —— 而这一讲自己在 2.5 就写着「跨文献比这类比例前, 先确认对方折没折半」。已统一到 step 口径。

⭐⭐ 一条第一方证据:这张排序表不是我们自己推着玩的。

DeepSeek-V3 的报告里写得很直白 —— 他们 「重算全部 RMSNorm 运算和 MLA 上投影」, 理由是「开销很小,却显著降低了存激活的显存需求」。

⭐ 这正好落在上面那张表最便宜的那几行里。 判据是我们自己推的,但推出来的名单跟人家实际配的对上了 —— 这比多算一遍更有说服力。

📌 出处:arXiv 2412.19437 §3.2.3。

2.5 ⭐⭐ 同一条判据,2022 年给出的是相反的答案

「选择性重算」这个概念出自 Korthikanti 等,Reducing Activation Recomputation in Large Transformer Models,arXiv 2205.05198(2022-05)。

它的判据跟 2.3 一模一样 —— 论文原话是:挑那些 「占显存不少、但重算起来不贵」的部分。

⛔ 但它选出来要重算的,恰恰就是 attention ——  QKᵀ、softmax、softmax dropout、以及 attention over V。 收益是:GPT-3 省 70% 激活显存只付 2.7% 算力; MT-NLG 省 65%1.6%

⭐⭐⭐ 为什么结论会反过来?因为它的序列长度是 2,048。

⛔ 但要用 GPT-3 自己的斜率算。它是标准 MHA,qk = v = 128, 所以斜率是 1.00,不是 V3 那个 1.25 ——  代进去是 1.00 × 2048 = 2,048。 而 GPT-3 自己最宽的线性层是 12,288(V3 是 7,168)。

2,048 远小于 12,288 —— 它落在自己交点的左边很远处。 在 2K 上,attention 真的是最划算的重算对象。论文没错。

⭐⭐⭐ 判据一个字没改,结论却翻转了。但有一句要说准:

变的不只是序列长度。两个模型的斜率不同(1.00 vs 1.25)、 对照的线性层也不同(12,288 vs 7,168)—— 两条线都得重画。

⭐⭐ 换句话说:能迁移的是那个形状,不是那些数。 —— 换个模型,你要重画两条线,而不是重记一组数字。

⛔ 顺带自曝一处:这张图早先把 V3 的 1.25 直接套在 GPT-3 上, 标成了 2,560。—— 而这一讲从头到尾在讲「别拿 A 的常数套 B」。 判据:同一张图里出现两个模型时,每个模型的每一条线都要用它自己的常数。

⚠️ 一处口径要说明白,不然数字对不上:论文那个 2.7% 用的是 Megatron 传统公式,没有为 causal 折半。换成本课统一的 causal 口径, GPT-3 的 attention 占比是 1.37%。两个都对,只是约定不同 ——  跨文献比这类比例前,先确认对方折没折半。

顺带一句:这件事今天已经不用你操心了。 FlashAttention 天生就不把分数矩阵写进显存、反向时现算 —— 等于把论文那条建议内建成了默认行为

这里我原来写「所以现代框架的候选名单里 attention 根本不出现」——  去翻了源码,完全说反了。

Megatron-LM 的选择性重算开关 --recompute-modules, 候选项里第一个就是 core_attn而且它是默认值recompute_modules = ["core_attn"])。

⭐⭐⭐ 而这个默认值,正是 2.6 那条判据的活标本它是 2,048 那个年代选出来的最优解:在那个长度上它对, 在 128K 上按我们这张表,它恰恰是最不该重算的那一个 —— 但它还在那里当默认。

更值得看的是名单上的其余几项layernormmoe_actmla_up_proj —— 对照 2.3 那张表:RMSNorm 输出、SwiGLU 乘积、K/V 解压与 Q 展开。 正好是最便宜的那几行,一个不多一个不少。 我们那条闭式解排出来的名单,和框架实际提供的选项对上了。

📌 megatron/core/transformer/transformer_config.pyrecompute_modules 字段(2026-09 主干)。 ⛔ 判据:说「现代框架都不这么干了」之前,去 grep 一下那个框架。 这类断言听起来像常识,而它正好是最容易过期的一类。

2.5b ⛔ 重算还有一笔不在账上的代价:它可能算出不一样的东西

⭐ 前面算的是两本能算的账。可重算还有第三样代价, 它不在任何一本账上 —— 而且它不报错。

重算的前提是:把同一段前向再跑一遍,跑出来的要一模一样。 —— 有三种情况会让它不一样。

① 随机数 —— 这一条框架替你挡住了。

前向里有 dropout。如果重算时抽到的是另一套掩码, 那反向用的就不是前向那张网络了 —— 梯度直接是错的,而且没有任何报错。

⭐⭐ PyTorch 的做法很干脆:把前向那一刻的随机数状态存下来, 重算前先恢复。 —— torch.utils.checkpointpreserve_rng_state默认就是 True

⭐ 而这一条正好解释了 6.4 那个清单里的第三项: checkpoint 里除了权重和优化器状态,还得存随机数状态 ——  原因是同一个:随机数是训练结果的一部分,不是「运行时的临时东西」。

② 副作用 —— 这一条没人替你挡

如果那段前向除了返回结果,还顺手改了别的东西 —— 更新了一个滑动统计量、累加了一个计数器、写了一行日志 —— 重算会把它再改一遍。

⛔ 框架不会替你回滚这些,因为它根本不知道你改了什么。 —— 判据很简单:被重算的那一段,最好是个纯函数。

⚠️ ③ 逐比特不一致 —— 通常无害,但它有个具体的受害者。

同一段计算跑两遍,结果可能差在最后几位 —— 归约的顺序不同、挑的 kernel 不同,都会这样。

⭐ 对训练本身,这个量级的差别基本没影响。 ⛔ 但它让「逐比特可复现」变难 ——  而那正是 6.7 里 PaLM 做到的那件事, 也是 6.2 那个「回滚 + 跳数据」的救火办法所依赖的前提。

⭐⭐ 一句话:重算把「算力 ↔ 显存」这笔交易谈成了, 但它同时悄悄引入了一个「结果要可重现」的要求—— 前两本账会告诉你划不划算,这一条不会:它只在出事的时候现身。

2.6 ⚠️ 本章落点:收益不能照抄别人的

⛔⛔ 同一个模型、同一个开关,换一个规模, 收益可能从正的变成负的

原因不神秘:重算改变的是计算与访存的配比, 而这个配比在不同并行配置、不同芯片数下本来就不同。

⭐⭐⭐ 由此推出一条通用规则(它不只对 remat 成立): 凡是会改变数据分片形状的参数,都不能跨规模照抄。 小规模上验过的结论,到目标规模上必须重验。

2.5 那个例子说明这条规则比听起来更狠序列长度也在这个名单里,而且它翻转的不是收益的大小, 是收益的正负号 —— 在 2K 上最该重算的那一项, 到 128K 上变成最不该动的那一项。

📌 另一个值得知道的量级 —— 而这个数可以直接算出来:

有用的活是 3 份(前向 1 + 反向 2),全量重算再加 1 份 —— 总共执行了 4 份。所以重算占已执行算力的四分之一

这里我原来写的是「三分之一」 ——  那是把「比原来多付三分之一」(1 ÷ 3)和「占执行总量的多少」(1 ÷ 4) 搞混了。同一个 1 份,分母不同。 📌 顺带:那个「多三分之一」不是我们推的, ZeRO 论文原话就是「33% re-computation overhead」(arXiv 1910.02054 §3.2)。

这顺带解释了一个常见困惑 —— 为什么 MFU 看起来那么低: 分母里有一大块被重算吃了,它做了功,但不算进「有效算力」。

第 三 节

优化器 —— 最大的一块显存,和最难调的那个数

前面两节都在跟激活较劲。可把账摊开一看 ——  那个从头到尾一言不发的角色,才是最大的一块。

3.1 一个参数到底要占多少字节

先花两句话把前提说清楚,不然下面那张账单看不懂。

训练的时候,同一个参数在显存里同时存着两份
① 一份省地方的(bf16,2 字节)—— 前向和反向都拿它算;
② 一份精确的(fp32,4 字节)—— 专门用来记账,每一步的更新加在它身上

这套做法就叫「混合精度」。 —— 为什么非得存两份,下一小节(3.2)整节都在回答

混合精度训练下,每一个参数身上挂着五样东西 ——  而其中只有第一样是「模型本身」。

优化器谱系 —— 每一代往每个参数身上挂了几份状态 这张图不按算法怎么算排,按账单排 · 横轴只表示先后,不是线性年份 0 份 1 份 2 份 往回走 Ⓐ 三十年只在加同一样东西 —— 直到 2024 年有人把它拿掉一份 「份数」指的是每个参数要额外存几个跟它一样大的张量 0 份 1 份 2 份 SGD 只按梯度走一步 + 动量 1964 记住上一步往哪走 AdaGrad 2011 每个参数一个学习率 RMSProp 2012 把累加换成滑动平均 Adam 2014 动量 + 自适应,两份都要 AdamW 2017 修 weight decay 那个 bug Muon 2024 去掉 v,改为正交化 横轴只表示先后,不是线性年份 Muon 是往回走的那一步 三十年一路加,2024 年 有人把二阶矩那一份整个拿掉了 AdamW 不是新算法 —— 它是在修一个 bug 论文摘要原话:L2 正则和 weight decay 对普通 SGD 是等价的,但对 Adam 这种自适应方法不等价。 因为 L2 那一项是混在梯度里进去的,再被二阶矩的分母一除 —— 梯度 大的参数,它的 weight decay 被稀释掉了 所以 AdamW 的修法是把 weight decay 从梯度里解耦出来,更新参数时单独减一下。那个 W 就是 decoupled 的意思。 Ⓑ 所以「每参数 16 字节」是哪五样 —— 其中 12 字节是 fp32 的那三份 混合精度训练的常规配置 bf16 权重 2 B 前向反向都用它 bf16 梯度 2 B 反向算出来的 fp32 主权重 4 B 真身,用来累加 fp32 动量 m 4 B Adam 的第一份 fp32 二阶矩 v 4 B Muon 去掉的就是它 = 16 B 2 B 推理也要 就是你下载到的那份权重 14 B 训练才要 这一讲讲的,全是这一段 而这一讲那句「最大的一块是优化器状态」,落到数上就是这么来的: 权重只占 2 B,优化器那边(主权重 + m + v)占 12 B —— 六倍 而换成 Muon:二阶矩那一份没了 —— 16 B 变成 12 B,少四分之一。 但它只管二维参数 —— embedding 和输出头仍然走 AdamW,所以整模型省不到四分之一。
⭐⭐ Ⓐ 的形状是这张图的全部 —— 三十年一路往上加,2024 年有人往回走了一步
而 Ⓑ 把「最大的一块是优化器状态」翻译成了一根尺子:权重只占 2 B,优化器那边占 12 B。
出处与口径

Adam:Kingma & Ba,arXiv 1412.6980(ICLR 2015);AdamW:Loshchilov & Hutter,arXiv 1711.05101(ICLR 2019)—— Ⓐ 那条带子里的说法是它摘要的转述

Muon:Keller Jordan 等 2024,作者 writeup kellerjordan.github.io/posts/muon。「只用于 2D 参数」「标量 / 向量 / 输入输出层仍用 AdamW」「per-step wallclock 比 AdamW 慢」三条都是原文

动量:Polyak 1964(heavy ball)/Nesterov 1983;AdaGrad:Duchi, Hazan, Singer, JMLR 2011(无 arXiv);RMSProp:Hinton 2012 Coursera 第 6 讲 —— 从未正式发表,这一条本身值得一提

另有三条省状态的岔路本图没画:Adafactor(arXiv 1804.04235,把 v 分解成一行加一列)、Lion(arXiv 2302.06675,只用梯度符号、单动量)、8-bit Adam(状态量化,不改算法)

本图不含任何收敛速度的对照 —— 本课没有这些优化器的对照实测,画曲线就是编。图上只有结构与各自论文的定位

⭐⭐ 这一节只要记住一句:权重 2 字节,优化器那边 12 字节。

⭐ 所以「显存里最大的一块是优化器状态」不是一个修辞 ——  它就是 2 比 12 这个比。

而且这 12 字节全是 fp32。下一小节说为什么它们非 fp32 不可。

⛔⛔ 顺手拆掉一个最顽固的误解:「用了混合精度,显存不就省一半吗?」

⭐⭐⭐ 常驻这一块,它一个字节都没省 —— 两边算一遍,答案都是 16。

权重梯度动量 m二阶矩 v主权重合计
纯 fp32444416 B
混合精度2244416 B

看那一列:权重和梯度各省了 2 字节, 可换来的是多出一份 4 字节的主权重 —— 2 + 2 正好被 4 吃掉。

⭐⭐ 那混合精度到底换来了什么?—— 两样,都不是常驻显存。

  • 速度。矩阵乘走 bf16,用的是芯片上那套专门的单元 —— 这才是它的主要目的。
  • 激活。激活那一大块是 bf16 存的,这里确实省了一半 —— 而按 5.4 那个小例子,激活经常比常驻还大。

判据:说「省显存」之前,先问省的是哪一块 常驻和激活是两本账,一句「省一半」把它们混在一起, 结论就会在你最需要它的时候是错的(比如估一张卡装不装得下)。

3.2 哪些量必须高精度 —— 判据只有一句:看它累不累加

图上那五格,为什么有三格是 fp32、两格是 bf16? ⭐ 不用一格一格记 —— 根子上只有一个画面。

精度问题,根子上只有一个画面 —— 大数吃小数 这一格里每一个数都是脚本当场算的 —— 你也可以自己跑一遍 bf16:范围够,精细不够 fp16:精细够,范围不够 fp32:都够,但四倍大 Ⓐ 一个小数在机器里分成两段存 —— 指数管「能多大」,尾数管「能多细」 看清楚:bf16 和 fp32 的指数段一样长 fp32 32 位 指数 8 位 尾数 23 位 bf16 16 位 指数 8 位 尾数 7 位 指数跟 fp32 一样长 → 范围一样,只是变粗 fp16 16 位 指数 5 位 尾数 10 位 指数短了 3 位 → 小的梯度直接掉到 0 bf16 和 fp16 都是 16 位,只是分法不同 —— 所以 fp16 要 loss scaling,不是因为「16 位不够用」,是因为它把位数分给了尾数。 Ⓑ 于是「加了等于没加」 —— 这不是比喻,是真的一动不动 权重 1.0,每步加 0.0003(训练后期的典型量级) 1.0000000 1.0078125 1.0156250 1.0234375 这中间什么都没有 一格 = 0.0078125 要加的量 0.0003 只有一格的 3.8% —— 四舍五入直接舍回原地 要连加 26 次 才够跨过一格 可它跨不过去 bf16 里连加 1,000 次 1.0 → 1.0 一步都没动 —— 每次都被舍回去了 同样 1,000 次,改用 fp32 1.0 → 1.3 该加的都加上了 所以「主权重留一份 fp32」不是保险起见 —— 不留,训练到后期就真的停在原地了。 Ⓒ 那谁必须 fp32?判据一句话:老的贡献会不会永远不走 不是「累不累加」—— 那个说法太粗,被 DeepSeek-V3 一条实测打中过 10⁰ 10¹ 10² 10³ 10⁴ 10⁵ 这一步的贡献,到现在还剩多少 一半 主权重 必须 fp32 v:半衰期 693 步 m:6.6 步 梯度/激活:下一步就没了 只有那条一直平着的必须 fp32 —— 它装的是十万步前那一点点增量,而那一点现在还在里面。其余三条都会被忘掉,忘得掉的就存得粗。 所以 12 个 fp32 字节装的正好是「不走的」,2 个 bf16 字节装的正好是「会走的」 —— 这条线不是拍出来的,是这条曲线画出来的。 梯度一做累积就又变成「要留一阵子」,于是升回 fp32。 而这条规则可以被打破 —— 只要把「每次都朝同一边抹零」去掉 真正致命的不是「加的量小」,是「每次都朝同一个方向抹零」。换成随机舍入(按余数大小决定进位概率),那么每步的期望增量正好等于真实增量 —— 丢掉的那部分,多步之 后会被找回来。(这个恒等式脚本里验了。) 但这不是推荐做法:它要额外的硬件/框架支持,而 fp32 主权重便宜又省心。放在这儿只是因为 —— 知道一条规则「为什么成立」,才知道它什么时候可以不成立。
⭐⭐⭐ Ⓑ 是这一节的全部 —— 权重 1.0、每步加 0.0003,在 bf16 里连加一千次还是 1.0。不是慢,是一步都没动。换 fp32 同样一千次:1.0 → 1.3。
⭐⭐ 这些数不是引来的,是脚本当场算的 —— 里面有一个真的 bf16 舍入函数,三条 assert 盯着。你可以自己跑一遍。
Ⓐ 顺带解掉一个常见误解:fp16 要 loss scaling 跟「16 位不够」无关—— 它跟 bf16 一样是 16 位,只是把 3 位从指数挪给了尾数,于是范围小了,小梯度直接掉到 0。
出处与口径

📌 Ⓑ 里每一个数都是脚本用一个真的 bf16 舍入函数算出来的(取 float32 的高 16 位,round-to-nearest-even,只用标准库 struct)。三条 assert 盯着:不到半格必须被舍回、跨一整格必须加得上、连加 1,000 次必须仍然等于 1.0。⭐ 判据:一张图要让人相信某个算术结果,就让脚本真的算那个算术。

📌 Ⓒ 那条「不是累不累加,是老的贡献会不会永远不走」是被一条第一方反证逼准的:DeepSeek-V3 技术报告(arXiv 2412.19437 §3.3.3)明说 AdamW 的一阶矩二阶矩用 bf16「未观察到性能下降」,而主权重、以及用于累积的梯度仍保 fp32。⛔ 判据(元级):一条判据被反例打中的时候,先别扔它 —— 多半是它的措辞比它的机制粗。

⭐ Ⓐ 顺带解掉 6.6 那条:fp16 需要 loss scaling 跟「16 位不够」无关 —— 它跟 bf16 一样是 16 位,只是把 3 位从指数挪给了尾数,于是范围小了三个二进制数量级,小梯度直接下溢到 0。

⭐⭐⭐ 所以那 12 个 fp32 字节,装的正好是「老的贡献不走」的那几个; 两个 bf16 字节,装的正好是「用完就扔」的。 —— 这不是巧合,这是同一条判据画出来的线。

❓ 那梯度能不能是 fp32?—— 能,而且有些实现就是。

最常见的场景是梯度累积:几个 micro-batch 的梯度要相加。 一直在 bf16 里加,小的会被大的吃掉 —— 累几次就掉精度了。 所以不少实现专门开一个 fp32 的 accumulator

那样这一栏就是 4 字节,总账从 16 变 18

📌 我们图上那个 16 是哪来的:它是 ZeRO 论文的经典口径 —— 2 字节权重 + 2 字节梯度 + 12 字节 fp32 那三份,原文写作 2Ψ + 2Ψ + 12Ψ = 16ΨarXiv 1910.02054⭐ 所以这个数有出处,不是我们自己定的;但它是一种配置,不是定律 —— 报这个数的时候要把口径一起报。

⚠️ 一处小订正:那两个 2 字节,论文里是 fp16,不是 bf16 —— 2019 年那会儿大家用的是 fp16。 今天普遍换成了 bf16,字节数一样,所以这张账一个字不用改; 但两者不是一回事(见 6.6 「bf16 要不要 loss scaling」那条)。

⚠️ 一条第一方反证,把上面那句话逼得更准了。

DeepSeek-V3 的报告里明说:AdamW 的一阶矩和二阶矩都用 bf16, 「没有可观察到的性能退化」;而主权重和梯度仍然保 fp32 ——  梯度保 fp32 的理由,恰恰就是「要用于累积」。

⭐⭐ 方向没错,措辞太粗。真正的分界不是「累不累加」,是 会不会无限累加

  • 主权重 —— 跨十万步不断加进微小增量,老的贡献一直在里面。 这是真的无限累加,必须 fp32。
  • 动量 m / 二阶矩 v —— 它们是滑动平均, 每步都乘一个小于 1 的系数再掺新值 —— 老的贡献会被衰减掉。 ⭐ 滑动平均不是累加,所以 bf16 扛得住。
  • 梯度 —— 本来是一次性的;一旦要做梯度累积,它就变成累加量, 于是又得升回 fp32。

⭐ 收紧后的判据:看它是不是「老的贡献永远不走」。 —— 永远不走的必须 fp32;会被衰减掉的、用完就扔的,低精度就够。

⚠️ 「必须 fp32」这句话要留个口子 —— 有人真的把它省掉了。

主权重之所以要 fp32,是因为四舍五入会把小更新整个抹掉。 ⭐ 而如果把它换成随机舍入 —— 按余数的大小决定进位的概率 —— 那么被抹掉的部分,在多步之后期望上会被找回来

⭐ 那么判据其实比「累不累加」还要细一层: 真正致命的不是累加,是累加 + 每次都朝同一个方向抹零 —— 把这个偏差去掉,bf16 主权重就是可行的。

⛔ 这不是推荐做法 —— 它要额外的支持,而 fp32 主权重便宜又省心。 放在这里是因为:知道一条规则「为什么成立」,才知道它什么时候可以不成立。

📌 出处:DeepSeek-V3 技术报告,arXiv 2412.19437 §3.3.3(原文分节是「低精度存储与通信」,不是 §3.2.3 那个重算小节)。 ⭐ 原文这句话几乎是逐条印证上面那张清单「用 BF16 而非 FP32 追踪 AdamW 的一阶和二阶矩,未观察到性能下降; 但主权重、以及用于 batch 累积的梯度,仍保留 FP32。」 —— 连「梯度因为要累积所以升回 fp32」这一条都写在里面。顺带说明这张 16 字节的表不是定律 ——  V3 这套配下来,每参数的账跟经典口径已经不是一个数了。

⭐ 这条判据出了这一讲还能用:FP8 训练为什么难, 问的就是同一句 —— 精度往下压的时候,哪些量能压、哪些必须留住, 看的还是它累不累加。那是另一课的事。

3.3 ⭐ 梯度是怎么变成更新的 —— 中间那个折算系数

梯度下降为什么可以 —— 一切的本源,先把它讲透 前面讲了梯度怎么算出来,这一格讲为什么顺着它走就能变好 一维:一条线 二维:一片山地 三千亿维:不画了 Ⓐ 先看一维 —— 脑子里放一个球,从山上滚下来 规则只有一句:斜率为正就往左挪,为负就往右挪 某个参数 → loss 起点 A 起点 B 落在这儿 落在这儿 球一步比一步挪得少 没人让它慢下来 —— 是坡自己变平了 A 落进了更浅的那个谷 只因为它从左边出发 —— 跟谁更优无关 自动刹车 步长 ∝ 斜率 越接近谷底,坡越平 坡越平,步子越小 所以它不会在谷底 来回冲过头 —— 这一条是免费送的,不用额外做什么 Ⓑ 加到两个参数 —— 「斜率」这个词就不够用了 一个数说不清方向,得用一个向量 一维:问「斜率是正是负」 一个数就够了 二维:问「往哪个方向走,降得最快」 一个数说不清 —— 它得是个方向 这个方向就叫梯度 梯度指的是上坡最快的方向 所以下坡就取它的相反数 而它的长度还顺带告诉你:这个坡有多陡 顺带记一句,后面会兑现:「最快」是相对于你怎么量「这一步迈了多大」说的 —— 换一把尺,最快的方向就跟着变(`fig-muon` Ⓒ)。 Ⓒ 到了三千亿个参数 —— 别再想「山」了,换个读法 三千亿维的山画不出来,但那一列数是看得懂的 所有参数 排成一列 0.31 −1.24 0.07 三千亿个数 负梯度 也排成一列 +0.002 −0.910 +0.004 一一对应 每一项告诉你两件事 正负:这个参数该往上推还是往下推 相对大小:哪一项改起来更要紧 第二件才是这一格的重点 —— 它不只说往哪走,还说该先动谁 所以「为什么可以」这个问题,答案分成能保证的和不能保证的两半 能保证的:每一步都沿着局部下降最快的方向走,而且步长跟坡度成正比 —— 只要步子别迈得太离谱,loss 就一路在降 不能保证的:你落在哪个谷取决于从哪儿出发,它不保证那是最低的一个。—— 梯度下降从来没承诺过最优,它只承诺每一步都在变好
⭐⭐⭐ 这是全讲的本源那一格 —— 前面讲了梯度怎么算出来,这里讲为什么顺着它走就能变好
⭐⭐ Ⓐ 那两串球是真跑了一遍梯度下降算出来的,不是摆上去的:步子一步比一步小 —— 没人让它慢,是坡自己变平了
而 A 那一串落进了更浅的谷 —— 只因为它从左边出发。看图时先比两个落点的高低,再回头看起点。
出处与口径

📌 Ⓑ 末尾那句「最快是相对于你怎么量一步」取自苏剑林《为什么我们偏爱各向同性?基于最速下降的理解》—— 原话是「梯度反方向是损失下降最快的方向,但这结论是有前提的,最关键的前提是它选取的度量是欧氏范数,如果换一个范数,那么最速方向也就变了」。⭐ 这句前提几乎所有教程都略过,而略过它,Muon 就只能被当成「又一个新优化器」。

📌 「球滚下山」「步长 ∝ 斜率所以不会冲过头」「落在哪个谷取决于起点」「高维时看那一列数的正负与相对大小」四个讲法,取自 3Blue1Brown《Gradient descent, how neural networks learn》官方讲义 —— 已逐条核过原文,图是我们自己重画的

⭐ Ⓒ 那一整列负梯度,正是反向传播一遍算出来的那一份 —— 所以这一讲的第一节和第三节,接头就在这儿。

⭐⭐⭐ 一维 → 二维 → 不再画空间。 第一幕:同一条曲线、同一个规则,两个球落进了不同的谷
第二幕:升到二维,换成等高线俯视 —— 形状还看得见。
⭐ 第三幕是关键:三千亿维想象不出来,所以干脆不画空间了 —— 把参数画成一列条,颜色是符号、长度是大小, 它们一根根缩短,就是在下降。
(15 秒无声循环,Manim 渲染。骨架取自 3Blue1Brown 神经网络系列的 叙事结构(源码, CC BY-NC-SA 4.0)—— 只借思路,画面与数据全部自算。 脚本在 tools/manim/。)

⛔ 那张图末尾留了一句不太好听的话:落在哪个谷,取决于你从哪儿出发。 一维的图上这件事看着很吓人 —— 满眼都是坑,随便掉哪个都出不来。 下面这张就是来拆这个画面的。

导数为零 ≠ 走不动了 —— 上一格那个「掉进坑里出不来」,是一维骗你的 这一格回答:为什么参数越多,梯度下降反而越走得通 真谷底:走不了 鞍点:还能走 实测:没遇到过真谷底 Ⓐ 坡度为零的地方不止一种 —— 看两个方向就能分开 每一格画的是同一个点,只是沿两个不同方向切一刀看剖面 真谷底(局部最小) 方向 1 往上 ↑ 方向 2 往上 ↑ 两个方向都往上 往哪走都变差 —— 真的卡住了 山顶(局部最大) 方向 1 往下 ↓ 方向 2 往下 ↓ 两个方向都往下 一推就走 —— 训练里基本遇不到 鞍点(垭口) 方向 1 往上 ↑ 方向 2 往下 ↓ 一个往上,一个往下 顺着往下那个方向,接着走 「垭口」就是鞍点:沿着翻山的那条路,它是整条路的最高点;横过来沿山脊走,它又是最低点。 中文管这种地形叫鞍部 Ⓑ 要当真谷底,得所有方向都朝上 —— 而只要有一个朝下,就还能走 一维图上只有两个方向可选 一维图上 2 个方向 (图上画了 2 个) 两个都朝上很容易凑齐 —— 所以那种图上坑一个接一个 一个小模型 几万个方向 (图上画了 40 个) 四十个全朝上已经很苛刻了 —— 何况几万个 我们在谈的规模 几千亿个方向 画到这儿就到头了 —— 真实的数再乘十亿倍,这一行根本画不下 只要有这一个 只要漏掉一个,它就只是个垭口 —— 接着走 这是定义层面的话,不是概率论证 —— 各个方向的弯曲方向并不彼此独立,所以这里不给「罕见到什么程度」配任何数字。 Ⓒ 而这不只是推理 —— 量过 把训练跑到梯度为零,再看这个点有多少方向是朝上的 0 1 横轴:朝上的方向占比(minimum ratio) 实际训练收敛到的点,都落在这一段 = 1 所有方向都朝上 = 真正的局部最小值 实测一次都没到过 换句话说:训练停下来的那些地方,总还剩一批方向是朝下的 —— 没走,不是因为没路,是因为坡太平了 这根轴是示意,不是实测散点 —— 只画「观测都停在 1 之前」这一个事实,原图见李宏毅投影片。 所以「梯度下降为什么可以」,最深的那半个答案是:参数越多,它越不容易真卡住 好消息:上一格那个「掉进浅坑出不来」的画面,是一维特有的。维度一多,坡度为零的地方绝大多数只是垭口 —— 总还剩一个朝下的方向。 坏消息:卡不住,不等于走得快。垭口附近坡极平、梯度极小,朴素梯度下降(步长 ∝ 斜率)会在那儿磨很久 —— 后面几格讲的那些优化器,一多半是在治这个。
⭐⭐⭐ 这张图拆掉上一张留下的那个阴影 —— 「掉进浅坑出不来」是一维特有的画面
⭐⭐ Ⓐ 记住垭口这一个词就够了:沿着翻山的路它是最高点,横过来沿山脊走它又是最低点 —— 坡度为零,却照样有路可走。
Ⓑ 要当真谷底,所有方向都得朝上;几千亿个参数就是几千亿个方向,漏一个它就只是垭口。Ⓒ 实测也是这样:训练停下来的那些点,一个真谷底都没遇到过
但别读成「所以随便训」—— 卡不住不等于走得快,垭口附近坡太平,后面那些优化器一多半是在治这个。
出处与口径

📌 「critical point 分 local minima 与 saddle point」「判据是 Hessian 特征值全正与否」「沿负特征值方向可继续下降,但实践中很少这么做」「实测 never reach a real local minima」四条,取自李宏毅2021《类神经网络训练不起来怎么办(一):局部最小值与鞍点》投影片 —— 已逐条核过原文,图是我们自己重画的。

⭐ 他课上还引了《三体Ⅲ·死神永生》里那个能在高维取物的魔法师作比,投影片原话是「从三维空间看它是封死的,在更高维度里并不是」。⛔ 小说情节我们没核,所以只转述这一句,不展开 —— 图上的主比喻用我们自己的垭口

⚠️ Ⓑ 刻意不给概率:「每个方向朝上朝下各半,所以 d 维全朝上的概率是 2 的 −d 次方」这句话听起来很顺,但特征值的符号并不独立,那个数是编的。只讲定义层面必然成立的部分。

⭐ 原片收尾那句也在本讲别处出现:更小的 batch 和动量都有助于逃离临界点 —— 噪声和惯性,正好是 `fig-batch` 与 Adam 那两格在讲的东西。

⛔ 可这张图是切了两刀给你看剖面 ——  而「一个方向上翘、另一个方向下沉」本来是个三维的形状。 ⭐ 下面让它转起来。

⭐⭐⭐ 同一个鞍点,转起来看。 上面那张图只能切两刀给你看剖面;而它真正的形状是这样的。 红线沿一个方向往上翘(走过去 loss 变大), 绿线沿另一个方向往下沉 —— 两条在中间那个点交叉。
⭐ 站在那个黑点上,脚下确实是平的 —— 可绿色那条路一直都在。 (8 秒无声循环,用 Manim 渲染 ——  就是 3Blue1Brown 那套工具的社区版。脚本在 tools/manim/。)

⛔ 上面那张图欠了半句:卡不住,可垭口附近坡太平,它会在那儿磨很久。 —— 下面这张是那半句的第一个答案。

动量 —— 同一个起点,只多给它一点惯性 上一格的落点是「卡不住,但走得慢」—— 这一格是那句话的第一个答案 没动量:停在浅坑 有动量:冲过去 β:记多久 Ⓐ 同一条 loss、同一个起点 —— 两条轨迹都是真滚出来的 唯一的差别:右边那条记得上一步怎么动的 浅坑 深谷 停在这儿 滚到这儿 同一个起点 灰:没有动量 滚进浅坑就停了 —— 而且是真停 那一点的梯度已经归零 绿:加了动量 带着上一步的速度冲了过去 还冲过了头,再荡回来 但说准一点:它冲得过去,是因为那个坑够浅动量让你更不容易被小坑绊住,它没承诺过能找到最低的那个谷。 Ⓑ 它凭什么冲得过去 —— 看它刚越过坑底那一刻,两股力在对着干 箭头的长度和方向都按真实数值画 —— 取自上面那条绿轨迹刚越过浅坑的那一步 从这一点出发 ← 往深谷那边 往回退 → 梯度这一项 − 学习率 × 梯度 0.092 它在把它往回拉 惯性这一项 β × 上一步的移动 1.304 它要继续往前 净移动 两项加起来 1.213 惯性赢了 13 倍 看清楚这一刻在发生什么:梯度是反对它继续走的 —— 它刚越过坑底,坡正拽着它回去。是惯性把它带出去的。 Ⓒ 那 β 是什么 —— 「上一步」其实是「过去所有步」,只是越老越轻 每往前追一步就再乘一个 β,所以它是一排按 β 衰减的柱子 这一步 1 步前 9 步前 17 步前 有效窗口 ≈ 10 步 窗口 = 1 ÷ (1 − β) —— β = 0.9 就是 10 步左右。β 不是玄学,它是一个「记多久」的旋钮。 Ⓓ 最后一件容易漏的:你降学习率的时候,其实顺手把摩擦力拧大了 条的长度 = 摩擦系数 λ —— 那才是物理量,β 只是它和步长凑出来的一个数 原来 η = 0.9 β = 0.85 λ = 0.158 记忆窗口 7 步 只把 η 降到 1/10 η = 0.09 β = 0.85(没动) λ = 0.500 摩擦悄悄涨了 3.16 倍 记忆窗口 7 步 同时把 β 提到 0.953 β ← 1−(1−β)√r λ = 0.158 摩擦纹丝不动 记忆窗口 21 步 摩擦系数本该是个常数 —— 可 β 不动、只动 η,它就跟着变了。调完之后记忆窗口从 7 步拉到 21 步(正好 1/√r 倍)。 动量补上了上一格欠的那半句 —— 但它补的是「慢」,不是「最优」 上一格说「卡不住,但走得慢」,动量正是治那个「慢」的第一招:在坡平的地方,梯度这一项快没了,可上一步的移动还在 —— 于是它继续滑。 别把它讲成「能找到全局最优」。图上那条能出来,是因为那个坑够浅;换个更深的坑它一样出不来。它降低的是「起点决定一切」的程度,不是消除它。
⭐⭐⭐ 同一条 loss、同一个起点,只多给它一点惯性。灰的那条滚进半路那个浅坑就停了(梯度真的归零,动不了);绿的那条冲过去,落进了后面更深的谷。两条都是真滚出来的。
⭐⭐ Ⓑ 是机关所在,而且比预想的更狠:在它刚越过坑底那一刻,梯度其实是在把它往回拉 —— 是惯性以十几倍的优势把它带出去的。箭头长度和方向都按真实数值画。
但别读成「能找到全局最优」—— 它冲得过去是因为那个坑够浅。它降低的是「起点决定一切」的程度,不是消除它。
出处与口径

⭐ Ⓐ 两条轨迹、Ⓑ 三个箭头的长度、Ⓒ 那排柱子,全是脚本算的。loss 是自己造的一维函数 L(x) = 0.015x² + 0.45cos(1.1x),造它的唯一要求是「路上正好一个浅坑、一个深谷」—— 坑在哪、多深,是让脚本自己找出来的,不是我标上去的。四条 assert 盯着:没动量那条必须停在浅坑且梯度归零、有动量那条必须落进深谷、两者落点要拉开。

⭐ Ⓑ 特意挑的是绿轨迹正穿过浅坑的那一步,并 assert 了「这一步的梯度近乎为 0,而移动量仍然很大」 —— 这一条不成立的话,这张图就没什么可讲的了

📌 这个主意有多老:1964 年。Polyak《Some methods of speeding up the convergence of iteration methods》—— 比反向传播那篇 Nature 还早 22 年。⭐ 他给这个方法起的名字是「小重球法」(the method of a small heavy sphere)—— 我们上面画的那个球,名字就是从这儿来的。

⭐⭐⭐ 而 Ⓑ 那个「梯度往回拉、惯性把它带走」的画面,原文逐字说过:「The motion proceeds not in the direction of the force (i.e. antigradient) because of the presence of inertia」—— 运动不沿着力(也就是负梯度)的方向走,因为有惯性。他还说那一项会让它「沿着谷底走」。

⭐⭐ 最有意思的一条:Polyak 在 1964 年给的经验取值是 ρ = 0.8 – 0.99,而今天 Adam 的默认 β₁ = 0.9 —— 六十年过去,还在这个区间里。(脚本里拿这个区间 assert 了图上用的两个 β。)⭐ 他连调参顺序都写了:先把 ρ 设成 0 调好学习率,等收敛慢下来再把动量加上;并报告实测「多数情况下比梯度法快,最多十倍」。

⛔ 有一条我没能核实,所以不写:常有人说 1986 年那篇反向传播的 Nature 论文里就带了动量项。这一轮没拿到原文(链接 404),所以本讲不提这一条 —— 等拿到原文再说。

📌 讲法取自李宏毅《类神经网络训练不起来怎么办(一)》:他把动量写成「Movement = 上一步的移动 − 当前的梯度」(不只看梯度,还看上一步怎么动的),并在收尾写明 「smaller batch size and momentum help escape critical points」。⛔ 图是我们自己重画的,曲线和数都是自己跑的。

⭐⭐⭐ 同一条曲线、同一个起点,差别只在「上一步的速度留不留」。 灰球没有动量,滚进第一个浅坑就停住了; 蓝球带着攒下来的速度直接冲过去,落进了更深的谷。
⭐ 注意它冲过了头才荡回来 —— 那不是 bug,是动量真实的副作用, 也正是下一段要治的东西。
⛔ 两颗球跑的是同一份代码,唯一的差别是 β 从 0 变成 0.9。 (23 秒,Manim 渲染,带播放控件、不循环。 地形与这两条轨迹是脚本当场跑出来的:「先遇到的坑更浅」「灰球停在浅坑」 「蓝球落进深谷」「速度峰值高出近三倍」各有一条断言钉住。 脚本在 tools/manim/。)

⭐ 这三张图把方向说完了 —— 顺着负梯度走 loss 就在降,维度一多基本卡不死,平地上还能靠惯性滑过去。 这一节只剩一个问题:那一步,到底该迈多大?

最朴素的那一版只有一行:新权重 = 老权重 − 学习率 × 梯度⭐ 但你仔细想 —— 这一行凭什么成立?

⭐⭐ 两边的单位对不上。

梯度的含义是「这个参数变一点,loss 变多少」——  单位是 loss ÷ 参数。 可你要的是「参数该挪多少」—— 单位是参数本身

⛔ 所以中间必须乘一个东西把它折过来 —— 那就是学习率。 反推一下它的单位:参数² ÷ loss

⭐⭐⭐ 所以学习率不是一个纯数字它跟你的 loss 有多大、你的权重有多大,全绑在一起 —— 这就是为什么换个模型、换个 batch,学习率就得重调。 它压根不是个通用常数。

这个量纲问题一暴露,SGD 的根本毛病也就出来了

全网络共用一个学习率,可不同参数的梯度尺度能差好几个数量级。 —— 下面两条轨迹用的是同一个谷、同一套规则,只改了学习率。

一个学习率伺候不了所有参数 —— 而中间那个「折中」也不存在 同一个谷、同一套规则,只改学习率 —— 两条轨迹都是真跑出来的 大 5%:炸 小 5%:慢 中间:也不行 Ⓐ 这个谷一个方向陡、一个方向平 —— 陡峭程度差 25 倍 竖着弹的是陡的那个参数,横着爬的是平的那个 谷底 出画面了 起点 红:学习率只大 5% 每弹一次幅度更大 —— 收不住 ↑ 竖直 = 陡的那个参数 → 横向 = 平的那个参数 蓝:学习率只小 5% 不炸了 —— 可你看那串锯齿: 每一步大半的力气花在左右横跳上 真正朝谷底去的只有那一丁点横向位移 两条轨迹的学习率只差一成 —— 一条炸了,一条慢得让人着急。你能下手的区间就这么窄。 Ⓑ 那取中间那个值呢 —— 问题就在这儿:中间也没有好的 下面三行不是估的,是算出来的 再大 5% 陡的方向直接发散 门槛卡死在 2 ÷ 陡峭程度 —— 跟你想不想快无关 就取最大的安全值 平的方向走完一半 0.37 × 陡峭比 步 —— 图上这个谷(比 25)= 9 步 而真实模型呢 梯度尺度差几个数量级 比值按 1000 算,同一条公式给出 364 步 —— 而这已经是最快的了 所以这不是「参数没调好」 —— 是这个谷的形状决定的,而形状是模型给你的,不是你能选的。 出路只有一条:别再找那个「最好的全局学习率」了,它不存在 给每个参数配它自己的那一个。—— 而「它自己的」该是多少?下一格给标准答案(一阶导 ÷ 二阶导),再往后讲为什么真实训练里只能它。 代价也在这儿:「每个参数一个」意味着要为每个参数存东西 —— 这一讲开头那笔 12 字节的优化器状态,根子就是这张图逼出来的。
⭐⭐⭐ 这张图替掉的是一句纯文字的断言 —— 「梯度尺度差好几个数量级」听着像句套话,画出来就是红色那条弹着弹着飞出了画面。
⭐⭐ Ⓑ 才是重点:中间那个「折中」根本不存在。在不发散的前提下能用的最大学习率,让平的那个方向走完一半也要几百步 —— 这不是调参没调好,是谷的形状决定的。
两条轨迹都是真跑出来的,三个数都有 assert 盯着。
出处与口径

📌 「同一个学习率,大了第二次更新就飞出地图之外、小了更新一百次还走不到谷底,所以不同参数应该有不同的 learning rate」这个两难取自 李宏毅《Training Tip》投影片(2025 秋 GenAI-ML 课程)。

⛔ 但他那两个具体数字我们没照抄 —— 抄一个别人挑出来的学习率,等于把别人的地形当成自己的。这里自己搭了谷、自己跑了两遍,每条轨迹都有 assert 盯着。

⚠️ 画面用的陡峭比是 25,不是正文说的「几个数量级」。第一版按 1,000 画,物理没错但画面废了:平方向几十步只挪几十像素,锯齿退化成一根竖线。⭐ 判据:要让人「看见」某个动态,参数按「看得见」选,真实量级交给公式外推。Ⓑ 第三行就是那条外推。

⭐ Ⓑ 那两个步数是闭式算的:平方向每步收缩 (1 − 1.9 ÷ 陡峭比),走到一半就是 log½ ÷ log(那个收缩率)。脚本里拿真跑一遍的结果对过账,差不超过 1 步。

⭐⭐⭐ 把学习率本身做成一根能走的轴。 滑块从左往右 = η 从小调到大,上面那条轨迹跟着连续变形:
贴着谷底慢慢蹭 → 走得刚好 → 开始左右横跳 → 炸出画面。
⭐ 注意标尺上那三条挨在一起的线:中间红的是发散门槛 (闭式的,2 ÷ 陡方向曲率),两侧细线是它的 ±5% ——  只隔标尺的一成,而跨过去就从收敛翻成发散。
⛔ 而左边那头也不是好消息:η 小到不横跳时, 平缓方向二十六步只走掉三分之一两头都不满意,这才是「一个学习率」的困境。 (15 秒无声循环,Manim 渲染,脚本在 tools/manim/。 门槛、±5% 的收敛/发散、以及左端那「三分之一」都是脚本当场跑出来并有断言钉住的。)

⭐ 自适应那条线,核心想法只有一句

别用梯度的大小,只用梯度的方向 ——  把每个参数自己的尺度先除掉。

可凭什么是「除」? —— 这个问题有一个标准答案,而且小到能当场验。

那一步该迈多大 —— 它其实是有标准答案的 前面只讲了「除完之后是多少」—— 这一格回答「为什么要除」 最优步长 = |一阶导| ÷ 二阶导 Ⓐ 在一条抛物线上,「一阶导 ÷ 二阶导」正好就是「离底还有多远」 所以理论上一步就能到底 —— 不用试、不用调 你在这儿 还差 0.5 最低点 陡 谷 一阶导 2 · 二阶导 4 2 ÷ 4 = 0.5 跟「还差 0.5」一模一样 你在这儿 还差 2 最低点 平 谷 一阶导 1 · 二阶导 0.5 1 ÷ 0.5 = 2 跟「还差 2」一模一样 Ⓑ 而这就解释了一件反直觉的事:梯度大,不代表离得远 所以只看梯度大小,跨参数就会比错 只看一阶导(梯度) 陡谷 2.0 > 平谷 1.0 → 「陡谷那个离得更远,该迈大步」 错了。它其实离底更近(0.5 对 2.0) 除以二阶导之后 0.5 和 2.0 → 两个都正好对 除完之后,两个参数才可比 Ⓒ 那为什么没人真去算二阶导?因为算不起 这一格把上面那条漂亮的式子,接回我们这一讲的现实 二阶导 每两个参数之间都有一个 几千亿参数 → 根本存不下,更别说算 那就估它 一阶导的历史 梯度一直很大 → 多半在陡的地方 落到实现 Adagrad:累加平方和 Adam:滑动平均 分母上那一坨,就是这么来的 所以自适应优化器分母上那一项,不是「一个技巧」—— 它是二阶导的替身 顺着这条线再看一眼本讲 3.3 那个结论就通了:Adam 把梯度的尺度除掉之后,每步大约挪 0.2 × 学习率 —— 它之所以敢把尺度除掉,是因为那个分母本来就在替二阶导干活 也要知道这个替身只是替身:「梯度一直大 = 曲率大」只是个经验假设,它不总成立 —— 这正是为什么调参这件事至今还是手艺。
⭐⭐⭐ 这张图回答的是「为什么要除 —— Ⓐ 在抛物线上,一阶导 ÷ 二阶导正好等于「离底还有多远」,所以最优的一步是算得出来的。
⭐⭐ Ⓑ 最反直觉的一格:梯度大不代表离得远 —— 陡谷那个点梯度是 2.0,却只差 0.5。不除以二阶导,两个参数根本不可比。
Ⓒ 收回现实:二阶导在几千亿参数上算不起,所以自适应优化器拿一阶导的历史去估它 —— Adam 分母上那一坨就是这么来的。
出处与口径

📌 「最优步长 = |一阶导| ÷ 二阶导」与「用一阶导去估二阶导」取自李宏毅(台大)Gradient Descent 课程投影片的讲法 —— 图是我们自己重画的

⭐ 图上四个数都能当场验:抛物线 f = a·x² 上,一阶导 2ax ÷ 二阶导 2a = x,正好是到底的距离。脚本里带 assert。

⭐⭐ 看完上面那张图,下面这三个名字就不是三个技巧,是同一件事的三次改良 —— 它们都在凑那个分母

⭐⭐⭐ 现在关键的一步:Adam 的更新量大约是 「动量 ÷ 二阶矩的根号」—— 分子分母的量纲互相抵消

也就是说,Adam 把梯度的量纲给除掉了 —— 这一步的大小,变成了一个跟梯度尺度无关的常数

⚠️ 那个常数是多少?—— 不是 1,是 0.2。 这一条我原来写错了,而且错得很典型:「量纲抵消了」只说明它是常数, 没说明它等于 1。

实测它稳定在 0.2~0.3,而且不同尺寸的模型都一样

⭐⭐⭐ 而它有一个漂亮到不像话的闭式解 ——  在「梯度基本是噪声」这个近似下:
update RMS ≈ √[(1−β₁) / (1+β₁)]
代入 β₁ = 0.9:0.2294。跟实测对上了。

📌 出处:苏剑林《为什么 Adam 的 Update RMS 是 0.2?》(科学空间 kexue.fm/archives/11267)——  文中给了数值模拟(纯高斯梯度模拟出 0.225)与平均场近似两条路, 两边都落在同一个数上。

⭐⭐ 注意这个式子里只有 β₁,没有 β₂—— 两个 β 就此分工清楚:β₂ 管「记多久」,β₁ 管「每步迈多大」 (3.6 接着讲)。

⭐⭐ 所以学习率在 Adam 里的身份变了 ——  它不再是折算系数,而是直接乘在一个常数上: 每一步每个参数大约挪 0.2 × 学习率

这才是「Adam 的学习率好调」的真正原因: 它跟 loss 的尺度、跟梯度的大小都脱钩了, 只剩下一个由 β₁ 定死的系数。

AdamW 改的那一处很小,但很重要

原来 Adam 把权重衰减混进梯度里一起算, 于是它也被那个自适应分母除了一遍 ——  结果梯度大的参数,正则反而弱。

⭐ AdamW 就是把它从梯度里拿出来,直接加到更新那一步上。 论文标题里那个 decoupled,说的就是这件事 —— 整篇论文主要就干了这一件事。

⭐ 而 0.2 这个常数,顺手解决了「换优化器要不要重调学习率」

既然 Adam 每步挪的是 0.2 × 学习率, 那别的优化器只要也把自己的 update 对齐到 0.2 —— Adam 那套调好的学习率和权重衰减就能直接搬过去用。

⭐⭐ 这不是纸上谈兵:Kimi K2 从 Adam 迁到 Muon,用的就是这一招 —— 把 Muon 的 update RMS 统一成 0.2,其余超参照抄。

📌 同一篇(苏剑林,科学空间)。 ⭐ 这是「量纲」这条线最实用的一个落点先把某个量做成常数,常数就能当接口用。

Muon —— 同一个动作「扔掉大小,只留方向」,在三种形状上的三个版本 三列并排看,就会发现它们是一件事 · 分开讲就变成三个知识点 标量 对角阵 一般矩阵 Ⓐ 为什么它配叫「矩阵版的符号函数」 每一列都是:左边原样,右边只剩方向 标量 −1 0 +1 0.42 不管它原来是 0.42 还是 42 出来都是 +1 对角阵 +0.8 +1 -0.3 -1 +0.5 +1 对角线上 各管各的 = 拍平成向量 也没差 这一格是「向量做法」的极限 一般矩阵 原来的奇异值 全部拉到 1 方向全留下,大小全扔掉 Ⓑ 「矩阵和向量不都是一堆数字吗」—— 一个例子就能说清 拍平成向量,抹掉的正是这件事 一个矩阵 对角线那几个是特殊的 —— 迹就是它们的和, 而迹在相似变换下不变,还等于所有特征值之和 拍平 拍平成一个大向量之后 所有位置长得一模一样 —— 淡紫小竖线标的是原来在对角线上的那六个,优化器已经看不出来了 SGD / Adam 都是这么干的(逐元素) —— 而 Muon 拒绝拍平。 Ⓒ 换一把尺子,就换一个优化器 —— 而尺子的差别就是这两个形状 同一个问题:在「这一步不许迈太大」的区域里,沿梯度方向走最远 σ₁ σ₂ 梯度方向 这一步落在这儿 F 范数 = √(σ₁² + σ₂²) (把矩阵拍平,算欧氏长度) 最远点在梯度方向的延长线上 → 比例原封不动:大的还是大,小的还是小 得到的就是 SGD σ₁ σ₂ 梯度方向 这一步落在这儿 1 1 谱范数 = max(σ₁, σ₂) (由「矩阵乘向量」这个动作诱导出来) 最远点是那个角 → 两个奇异值都变成 1 —— 这就是 msign 得到的正是 Muon 同一个梯度 同一个问题 只换了区域的形状 Ⓓ 「每步多做十几次矩阵乘,不会很慢吗」—— 它塞进了一段本来就空着的时间 作者自己算过这笔账:FLOP 开销低于 1% —— 但那是 FLOP 口径,不是墙钟口径 算这一步的梯度 算力满载 空窗 算力几乎闲着 算下一步的梯度 算力满载 msign 就干在这儿 而且这些矩阵乘尺寸固定、可以并行 —— 不是那种会卡住流水线的活 所以它花的是本来就闲着的算力,换来的是每参数少 4 字节(16 → 12)
⭐⭐⭐ Ⓐ 的三列要并排看 —— 分开讲是三个知识点,并排放你自己就看出来了:它们是同一个动作。
⭐⭐ Ⓑ 用「迹」回答了「矩阵和向量到底差在哪」;Ⓒ 说明换一把量尺就换一个优化器;Ⓓ 解释了那十几次矩阵乘为什么几乎不要钱
📌 四条讲法都取自苏剑林《Muon 优化器赏析》(科学空间)—— 图是我们自己重画的。
出处与口径

讲法来源:苏剑林《Muon 优化器赏析:从向量到矩阵的本质跨越》,科学空间 spaces.ac.cn/archives/10592。Ⓐ 的三个特例、Ⓑ 的「迹」那个例子、Ⓒ 的范数视角、Ⓓ 的空窗解释,四条都取自该文 —— ⛔ 图是我们自己重画的,不是搬运

⚠️ Ⓒ 那个范数视角该文另引《Old Optimizer, New Norm: An Anthology》;⭐ 该文还考了一笔源流:2015 年的 Stochastic Spectral Descent 已经提出过大致相同的算法

⛔ 本图不含任何收敛速度对照 —— 本课没有这些优化器的对照实测,跟 fig4-optimizers 同一条规矩

⚠️ Ⓓ 那个「低于 1%」是作者按 Llama 405B 的宽度和每批 token 数算出来的 FLOP 开销,不是我们量的墙钟时间。⛔ Newton-Schulz 是一串互相依赖的小矩阵乘,FLOP 少不代表时间短 —— 真要用请在目标配置上自己量(这正是本讲那条「不能跨规模照抄」

⛔ 本图早先写的是「5% 以内 / 作者称 2%」,那个 2% 查无出处 —— 作者原文给的是 FLOP 开销低于 1%。2026-09-17 订正

Muon 更激进一层

⭐ 上面那张图基本讲完了。 这里只补图上没说的三句。

⭐ Ⓐ 那三列里,中间那一列值得单独想一下: Muon 作用在对角阵上会退化成逐元素取 sign —— 也就是说,「向量做法」本来就是「矩阵做法」的一个特例。 ⛔ 我们以前把它们当两种东西讲,那是把特例和一般情形讲反了。

📌 公式里不写、但真实存在的那几道

梯度算完,到真正更新之间,还要过: 多卡 all-reduce 求平均 → 按全局范数裁剪一次(防止偶发大梯度炸掉) → 按 schedule 取当前学习率 → 权重衰减 → 才是那一步更新。

⭐⭐⭐ 这一整条线可以用一句话收: 梯度只告诉你往哪走,优化器决定走多远

而几十年的演化,全都在回答同一个问题 ——  这个「多远」该由谁说了算:

  • SGD 说 —— 说了算,你去调。
  • AdaGrad / Adam 说 —— 梯度自己的历史尺度说了算。
  • Muon 说 —— 矩阵的几何说了算。

⭐ 说到底就这么一条线索。

3.4 换个优化器,账单就跟着变

⭐ 优化器的选择是一个显存决策,不只是收敛速度的决策。 —— 这一点常被忽略,而它恰恰是这一讲要立的那条判据。

⛔ Muon 那一份不是白省的 —— 三条限制都是作者自己写明的。

  • 只管二维参数。标量、向量,以及 embedding 和最后那个输出头,仍然走 AdamW ——  作者明说输入输出层用 AdamW 效果才最好。 所以整模型省不到那四分之一。
  • 每一步更慢(原文写明)。⛔ 所以它不是「又快又省」 ——  它把一笔账挪到了另一笔账上:显存那栏减了,单步时间那栏加了。

⚠️ 这跟 fig-muon Ⓓ 那句「它塞进了一段本来就空着的时间」不矛盾, 但两句话一定要一起读: 「更慢」说的是方向(多做的那几轮矩阵乘不是免费的), 「塞进空窗」说的是幅度(作者按 FLOP 算下来低于 1%)。
⛔ 而幅度是你的配置说了算的 —— 那几轮是一串互相依赖的小矩阵乘, FLOP 少不代表墙钟时间短真要用,在目标配置上自己量一遍。

⭐⭐ 而它的证据方式值得单独一提:Muon 在 NanoGPT 那个刷速度的公开竞赛里 把记录提了 35%,此后十二次破纪录、七个不同的人,全都还在用它。 —— 要是有人能把 AdamW 调到一样好,换回去就能破纪录,可没人换。

⭐ 顺带一句跟上一讲连起来:DeepSeek-V4 就是用 Muon 训的 —— 论文摘要里列的三大升级之一。

⛔ 在比账单之前,先把这几个名字串成一条线 ——  不然它们只是几个并列的选项,只能靠背。

从 AdaGrad 到 Adam —— 这不是四个选项,是一条因果链 同一串恒定梯度,看各自给出的「有效学习率」—— 健康的应该是一条平线 AdaGrad:掉到 0 RMSProp:开头冲太高 Adam:平 Ⓐ 每一环补上一环的洞,同时留下一个新的 每个节点里那条小曲线,就是它给出的有效学习率(真跑的 SGD 全局一个学习率 尺度差大的参数伺候不了 除以「它自己历史梯度的均方根」 → 每个参数一个尺度 AdaGrad 2011 分母只增不减 → 有效学习率单调掉到 0 累加换成滑动平均 → 分母会遗忘,不再单调掉 RMSProp 2012 没动量,也没偏差校正 → 头几步冲得极高 +动量 +偏差校正 → 方向也平滑了,开头也稳了 Adam 2014 恒定梯度下从第一步起就是平的 Ⓑ AdaGrad 的病在后期 —— 分母只增不减,学习率自己走向 0 横轴是步数,纵轴是有效学习率(稳态 = 1) 1 0 第 600 步 SGD(对照) AdaGrad 跑到第 600 步只剩稳态的 4.1% 注意这跟数据无关 —— 喂进去的梯度自始至终是同一个数,是那个只增不减的分母把它自己压死的。 Ⓒ 而 RMSProp 的病在初期 —— 滑动平均从 0 起步,头几步严重低估 分母被低估 → 更新被放大。Adam 的偏差校正就是来修这一下的 第 400 步 RMSProp(无校正) Adam 若不校正 Adam(有校正) 第一步被放大 32 倍(冲出画面) 而 Adam 若不校正,峰值只有 6.6 倍 —— 因为动量那一项的偏差部分抵消了它 绿线从第一步起就是平的 —— 偏差校正不是小修小补,它把开头那几十步从「不可用」变成了「可用」。 这条链最强的一个证据:后面那个能把前面那个当成特例含进去 Adam 原文自己证了一件事:AdaGrad 就是 Adam 的一个特例(β₁ 取 0、(1−β₂) 取无穷小、学习率按 t 的负二分之一次方退火)。—— 所以这真的是一条链,不是四个并列 的选项。 而每一环都还欠着下一环:Adam 留下的那个洞是权重衰减被卷进了自适应的分母里(AdamW 来修),以及每个参数要挂两份状态(Muon 来省)。这两环本讲后面都有。
⭐⭐⭐ 这不是四个并列的选项,是一条因果链 —— 每一环补上一环的洞,同时留下一个新的。
⭐⭐ 同一串恒定的梯度进去,健康的优化器应当给出一条平线;不平的就是它自己的毛病。AdaGrad 的病在后期(分母只增不减,有效学习率自己掉到稳态的 4%);RMSProp 的病在初期(滑动平均从 0 起步,头一步被放大 32 倍);Adam 从第一步起就是平的
这条链最强的证据来自 Adam 原文自己:AdaGrad 就是 Adam 的一个特例(β₁ 取 0、(1−β₂) 取无穷小、学习率按 t 的负二分之一次方退火)。
出处与口径

📌 这一格每一环都出自 Adam 原文(arXiv 1412.6980,已读原文):AdaGrad 的更新式原文逐字写作 θ_{t+1} = θ_t − α·g_t / √(Σ g²) —— 分母是累加的平方和,不除以 t,所以「单调衰减」这个病是从式子里长出来的;Adam 的定位原文写作「combine the advantages of AdaGrad(稀疏梯度)and RMSProp(在线与非平稳)」。

⭐ 偏差校正那一条,原文的说法是:滑动平均从 0 起步,于是估计 biased towards zero,especially during the initial timesteps,而且 β 越接近 1 越严重。同一段还点名了 RMSProp:without bias correction, like in RMSProp —— Ⓒ 那条橙线画的就是这句话。

⚠️ 李宏毅课上那页写的 Adagrad 是「均方根」形式(分母里除了 t+1),跟这里画的原版不一样。⛔ 本讲按原版画,因为「有效学习率单调衰减」这个病只在原版(累加不除 t)上成立 —— 这个区别值得知道,不然两边对不上会以为自己算错了。

📌 三则出处,放在一起看很有意思(完整的故事在讲义里):AdaGrad 是正经的 JMLR 长文(Duchi / Hazan / Singer,12(61):2121–2159, 2011,整整 39 页),⭐ 而且题目是《Adaptive Subgradient Methods for Online Learning and Stochastic Optimization》—— 它根本不是为深度学习写的

⭐⭐⭐ RMSProp 没有论文。它唯一的「出版物」是 Hinton 那门 Coursera 网课的一页 slide,标准引用写作 Tieleman & Hinton (2012), Lecture 6.5-rmsprop。⭐ 而那页 slide 的副标题本身就是算法的完整定义:「Divide the gradient by a running average of its recent magnitude」—— 把梯度除以它近期幅度的滑动平均。一句话,一页片子,成了这条链的中间一环。

⭐⭐ Adam 那篇(ICLR 2015)有两处特别好玩,都在论文里:首页脚注写着「Author ordering determined by coin flip over a Google Hangout」—— 作者顺序是在 Hangout 上掷硬币定的;致谢里写着「special thanks to Ivo Danihelka, and Tom Schaul for coining the name Adam」—— Adam 这个名字不是作者自己起的。

⭐ 图上的 4.1%、32 倍、6.6 倍三个数都是脚本算的,五条 assert 盯着:AdaGrad 必须单调衰减、跑到第 600 步必须掉到很小、不校正时第一步必须被放大很多倍、做了校正必须从第一步起恒为 1、以及「Adam 不校正的峰值必须低于 RMSProp 那条」(因为动量那一项的偏差会部分抵消)。

⭐ 省状态的三条岔路 —— 它们各砍掉了那 16 字节里的不同部分

⭐ 都是在动那 16 字节里的某一格,但动的格子不一样 —— 下面的字节数按 3.1 那套口径当场算:

做法动了哪一格每参数字节代价
AdamW(基准)——2+2+4+4+4 = 16——
Muon二阶矩整份拿掉2+2+4+4 = 12 只管二维参数;改用正交化补回来
Lion同样只留一份动量2+2+4+4 = 12 更新只用梯度的符号
Adafactor把二阶矩分解成一行 + 一列 那一份 ≈ 0;不带动量时总共 8 丢掉了 v 的逐元素信息
8-bit Adam两份状态各量化到 1 字节 2+2+4+1+1 = 10 算法不变,换的是存法

⭐⭐ 注意这三条走的是三个不同的思路Muon 和 Lion 是「不要那一份」;Adafactor 是「换个更省的表示」; 8-bit 是「同一份东西存得更小」。 —— 所以 8-bit 可以跟前两者叠加,而前两者互斥。

📌 8-bit Adam 那条值得多说一句,因为它的做法很干净。

它把状态张量切成小块、逐块独立量化, 再加上一种对大值小值都精确的非线性量化。 论文报告:在一系列任务上保持 32 位的效果,而且不用改任何优化器超参 —— 两行代码的 drop-in 替换。

📌 Dettmers 等,arXiv 2110.02861(ICLR 2022 spotlight)。

⚠️ 但「不改算法只改存法」这句话不完全成立,摘要里就写着。 论文把三件事绑在一起才达到 32 位的效果:分块量化、非线性量化, 以及第三件 —— 一个「稳定 embedding 层」

第三件是动网络结构的,不是存法。 它存在的理由很具体:语言模型的输入 token 分布极不均匀, embedding 那一层的梯度方差特别大 ——  而量化最怕的就是方差大。

⭐ 这一条本身就是个判据: 一个号称「drop-in 替换」的东西,先去摘要里数一数它到底改了几处 —— 「两行代码」说的是调用方改两行,不是它内部只改了一处。

⚠️ Muon 要在大模型上真跑起来,还得补两样

3.3 那段讲的是 Muon 的想法。 可想法好用不等于能直接上规模 ——  2025 年有一篇专门做这件事的论文,识别出两条必需的补丁:

⭐ 补上这两样之后,它才能「开箱即用」地跑大规模训练。 论文的规模律实验报告:在算力最优的设定下,Muon 的计算效率约为 AdamW 的两倍。

📌 arXiv 2502.16982(Moonlight,3B/16B MoE,5.7T token)。 ⚠️ 「两倍计算效率」是论文自己的规模律结论,不是我们的实测 ——  而且回想 2.6 那条判据:这类结论换规模要重验。至今仍未核实的一条:Muon 下主权重是否仍需 fp32。

3.5 ⭐ 学习率:那条曲线到底怎么画

在钻进「曲线怎么画」之前,先给一把随时能用的尺子 —— 它不告诉你最优解,但能告诉你「有没有离谱」。

「学习率设对了吗」 —— 有一条能用眼睛对齐的线 量的是更新量 ÷ 参数本身 —— 是更新量,不是梯度(原文特意强调的) 太小 刚好 太大 黑线 = 1e−3 Ⓐ 三个学习率,同一个网络 —— 看它们各自落在黑线的哪一侧 纵轴对数;黑线是 CS231n 给的经验值 1e−3 —— 一条粗略的经验线,不是定律 10⁻5 10⁻4 10⁻3 10⁻2 10⁻1 10⁰ 步数 → 1e−3 CS231n 的经验线 太小 η=0.000961 刚好 η=0.0288 太大 η=0.865 原文的读法就三句:贴着线 = 大致合适低太多 = 学习率可能偏小高太多 = 学习率可能偏大 它的好处是跟模型大小无关 —— 分子分母同量纲,约掉了。换个模型这条线还在原地。 Ⓑ 可「太大」那条的 loss 反而更低 —— 代价在别处 跑出来就是这样,不藏 —— 所以这条比值不是「让 loss 最低」的指标 0 5 10 15 第一层权重的范数 ‖W‖ 刚好/太小 两条几乎重合 太大 太大那条把权重撑到 4.2 倍,loss 的逐步抖动大了约 9 倍 —— 而学习率再乘 2.1 就直接 NaN 判据:这个比值不是「loss 最低」的指标,是「你离悬崖多远」的指标。 —— 诊断指标不是优化目标。 这是个玩具问题 —— 所以「loss 更低」在这儿本来就不算数 数据是随机 X → 随机 Y,没有验证集 —— 网络在记忆噪声,谁记得快谁 loss 低。 真实训练里,撑大权重 + 抖动加剧通常换来的是更差的泛化。 但这一格想讲的那条不依赖这个问题的好坏:那条比值线只回答「你的步子相对参数本身是不是一个合理的量级」,它是个体温计,不是治疗方案
⭐⭐⭐ 「学习率设对了吗」有一条能用眼睛对齐的线。量的是更新量 ÷ 参数本身(是更新量,不是梯度 —— 原文特意强调),CS231n 给的经验值是 1e−3
⭐⭐ Ⓐ 三个学习率跑同一个网络:太小那条贴在黑线下、「刚好」那条骑在黑线上、太大那条远在黑线上方。
Ⓑ 是这张图最该看的一格:太大那条的 loss 反而更低 —— 可它把权重撑到 4.2 倍、抖动大了约 9 倍,而学习率再乘 2.1 就 NaN
⭐⭐⭐ 所以这条比值不是「loss 最低」的指标,是「你离悬崖多远」的指标—— 诊断指标不是优化目标。

⚠️ 曲线是我们自己跑的两层 MLP(纯 numpy、固定种子);「刚好」那个学习率是扫出来的,「再乘 2.1 就 NaN」是二分搜出来的
出处与口径

📌 CS231n《Neural Networks Part 3》"Ratio of weights:updates":「A rough heuristic is that this ratio should be somewhere around 1e-3.」⭐ 原文还特意强调 "Note: updates, not the raw gradients"

⚠️ 曲线是我们自己跑的:两层 tanh MLP(8→16→4),纯 numpy 手写前向反向,固定随机种子,跑 400 步。「刚好」那个学习率是**扫出来的**(让稳态比值最贴 1e−3),不是挑的;「再乘 2.1 就 NaN」是二分搜出来的

⭐ 整条学习率曲线只有三段:升上去、稳住、降下来

两种学习率曲线 —— 形状差很多,落点却撞在同一个数上 纵轴是占峰值的百分比,不是绝对学习率 · warmup 那一小段按真实比例画就是一条竖线,图上放宽了 GPT-3 · 余弦 V3 · 平台+衰减 峰值的 10% Ⓐ 横轴是训练进度(占总 token 数的比例) 两条都归一化了,所以可以直接叠在一起比 0% 25% 50% 75% 100% 训练进度(已消耗的 token 占总量的比例) 100% 50% 10% 占峰值 峰值的 10% GPT-3 余弦 2,600 亿 token 处降到 10% DeepSeek-V3 长平台 10T token 才开始降 warmup 画宽了 GPT-3 真实只占 0.125% Ⓑ WSD 这两年流行起来的理由特别实在 不是「它收敛更好」—— 是它把一个决定往后推了 余弦的硬约束 曲线形状依赖终点 所以你必须一开始就知道总共训多少步 中途想多训一段,整条曲线都得重来 平台段的自由 恒定段可以随时截断 接一小段快速衰减,就能出一个能用的 checkpoint 想加数据接着训,从平台段续上就行 一句话:它把「训多久」这个决定,从开局推迟到了随时 两个落点,一个巧合 warmup 的量级是「总量的千分之几」,不是需要精调的东西 —— GPT-3 占 0.125%,V3 是头 2,000 步。 而且有论文实测:目标学习率固定的话,warmup 拉长基本没收益 决定效果的是峰值本身。 两条曲线的落点都在峰值的十分之一附近 —— GPT-3 明写「降到 10%」;V3 是 2.2e-4 → 2.2e-5,也正好十分之一(最后一小段再往下到 7.3e-6,约 3.3%)。 两个样 本不构成定律,但足以说明「降到零」不是默认做法
⭐⭐ 这张图比的是形状,不是高度 —— 两个模型的峰值差 3.7 倍,不归一化根本叠不到一起。峰值的绝对值另有一份七档对照,在正文里。
看两件事:V3 那条长平台,以及两条曲线都落在同一根「10%」虚线上
出处与口径

GPT-3:arXiv 2005.14165 §2.3 —— 「头 3.75 亿 token 线性 warmup」「2,600 亿 token 内余弦降到 10%,之后继续以 10% 训练」;总量 3,000 亿 token

DeepSeek-V3:arXiv 2412.19437 §4.2 —— 2,000 步线性升到 2.2e-4 → 恒定到 10T → 4.3T 内余弦降到 2.2e-5 → 末 500B 里前 333B 保持、后 167B 换 7.3e-6;梯度裁剪范数 1.0

⭐ warmup 的真实机制(不是「样本太少估不准」):arXiv 2406.09405(NeurIPS 2024)—— 主要好处是让网络能承受更大的目标学习率。这个形状的系统分析:arXiv 2410.05192

⛔ 本图不含绝对学习率(两者峰值差 3.7 倍,叠在一起就看不出形状了)—— 七档规模的峰值对照表写在正文里

⭐ 形状就这样。下面把三段拆开 ——  每一段都有公开配置可以抄,不用猜。

第一段 · warmup —— 流行的那个解释是错的

大部分教程说:Adam 一开始样本太少、二阶矩估不准,所以得慢慢来。

⚠️ 2024 年有一篇专门做系统实验的论文,明确说这不是主因。

原文的判断是:warmup 压倒性的好处,在于 让网络能承受一个更大的目标学习率 ——  它把网络推到 loss 曲面上条件更好的区域去。 而「样本太少估不准」这个说法,他们直接反驳了: 真正的问题是预条件曲率一开始就很高,大 batch 也一样高。

⭐⭐ 所以 warmup 真正的价值是让你敢把峰值设高 —— 顺带让这个超参更好调(可选区间更宽)。

📌 出处:Why Warmup the Learning Rate? Underlying Mechanisms and Improvements,arXiv 2406.09405(NeurIPS 2024)。

⛔⛔ 但「明确说这不是主因」这句说得太满,得收一格。 那篇论文的 Limitations 自己写着:实验是在 相对小规模的数据集和模型上做的,能否推广到更大规模还需进一步研究。

⚠️ 具体有多小:Transformer 只有 4 层、上下文长度 64, 另一组是 CIFAR-10 上的 WRN。—— 这跟我们整讲在谈的尺度差着好几个数量级。

⭐⭐⭐ 而还有第三个说法 —— 它自带一个前提

苏剑林在讨论「Transformer 怎么解决梯度消失」时,顺带回答了这个问题。 他的切入点跟上面两个都不同:Adam 让「梯度消失」这件事换了个意思。

⭐ 先看 Adam 的更新量长什么样(原文的写法):

Δθ = −η · E[g] ⁄ √E[g²]

分子分母同量纲,所以那个分式是 O(1),更新量就是 O(η)。 —— 这跟 SGD 完全不同:SGD 的更新量正比于梯度,梯度小,步子就小

⭐⭐ 所以在 Adam 底下,梯度消失不再意味着「学不动」 —— 只要梯度还大于随机误差,参数照样拿到常数量级的更新。

⛔⛔ 可这恰恰是麻烦的开始。 原文的推理链是这样的(前提:模型确实有明显的梯度消失):

  1. 不做 warmup,一上来就快学。越靠后的层越敏感,学得越快。
  2. 可后面的层是拿前面的层的输出当输入的,而前面的层还没学好 —— 它是在糟糕的输入上快速前进。
  3. 很快它到达一个糟糕的局部最优,学习放缓, 反传回前面的梯度进一步变弱
  4. 梯度不准了,可 Adam 的更新量还是常数量级 —— 于是那就是一个常数量级的随机噪声。 方向开始乱,前面崩,后面跟着崩。

⭐⭐⭐ 这条推理给了一个可以拿去对日志的现象 —— 比任何机制解释都好用:

Post-LN 的模型不做 warmup,你会看到: loss 先快速收敛到一个常数附近,再训一段,开始发散,直至 NaN。

⭐ 而 warmup 做的事,用原文的话说是 「抑制了后面的层的学习速度,并且给了前面的层更多的优化时间, 以促进每个层的同步优化」

⛔⛔⛔ 但这一段有一个必须一起记住的前提,原文自己写得很清楚:

「这里的讨论前提是梯度消失,如果是 Pre Norm 之类的结构, 没有明显的梯度消失现象,那么不加 Warmup 往往也可以成功训练。

⭐⭐⭐ 于是三个说法各就各位了 —— 它们不是在抢同一个位置。

说法前提现在怎么看它
二阶矩样本太少、估不准没说前提 ⛔ 已被 2024 那篇系统实验驳掉
让网络能承受更大的峰值学习率一般情况 ⭐ 成立 —— 但那篇的实验只到 4 层
压住后面的层,等前面的层跟上 有明显梯度消失
(也就是 Post-LN)
⭐ 成立,而且给了可观察的现象

⭐⭐⭐ 所以该问的不是「哪个解释对」,是「在什么前提下」。 你的模型是 Pre-LN 还是 Post-LN,决定了哪一条在起作用 ——  现代 LLM 基本都是 Pre-LN(见 §6.5),所以第二条更相关; 但如果你在训 Post-LN 的东西、又看到「loss 先平后 NaN」,那就是第三条。

📌 出处:苏剑林《模型优化漫谈:BERT 的初始标准差为什么是 0.02?》, kexue.fm/archives/8747 —— 上面引号里的两句为原文原话。 ⚠️ 这是作者本人的分析文章,不是同行评议论文; 我们按「一个有前提的解释」收下它,不当定论。

⭐⭐⭐ 所以准确的说法是那个尺度上,他们把「估不准」这个解释证伪了, 并给出了一个更有解释力的机制。而这个结论能不能搬到千亿参数上, 他们自己没说能。

⛔ 顺带自曝一处,而且很讽刺: 这一讲在 2.6 立了一条判据 ——  小规模上验过的结论,到目标规模必须重验。 结果我转述这篇论文的时候自己就没守。 —— 判据写在纸上容易,用在自己引的那篇论文上才难。

实际用多长?GPT-3 是头 3.75 亿个 token 线性升上去 —— 总训练量 3000 亿,占 0.125%。DeepSeek-V3 是头 2000 步

⭐ 换句话说,这一段短到可以不当成超参看

第二段 · 峰值设多少 —— 有一张现成的表可以抄

GPT-3 论文的表 2.1八个规模的峰值学习率和 batch 一起列了出来 —— 这是一张免费的先验,不用自己扫:

📌 GPT-3 全系列(同一套数据、同一套训练流程,只有规模在变):

  • 125M → 6.0×10⁻⁴,batch 0.5M token | 350M → 3.0×10⁻⁴,batch 0.5M | 760M → 2.5×10⁻⁴,batch 0.5M
  • 1.3B → 2.0×10⁻⁴,batch 1M | 2.7B → 1.6×10⁻⁴,batch 1M
  • 6.7B → 1.2×10⁻⁴,batch 2M | 13B → 1.0×10⁻⁴,batch 2M
  • 175B → 0.6×10⁻⁴,batch 3.2M

⭐⭐⭐ 规律肉眼可见:模型越大,学习率越小、batch 越大。 而且幅度很温和 —— 1,400 倍的规模差,学习率只差 10 倍(125M 的 6.0e-4 → 175B 的 0.6e-4)。

论文自己的话是:「更大的模型通常能用更大的 batch, 但需要更小的学习率」。他们用梯度噪声尺度来定 batch。

⛔ 注意这张表里 batch 和学习率是一起变的 ——  这是最容易漏的一条:你不能只调学习率不看 batch,它俩绑在一起。 GPT-3 还额外把 batch 从 32K token 线性爬到目标值(头 40–120 亿 token 内)。 📌 出处:arXiv 2005.14165 表 2.1 与 §2.3

⛔⛔ 而「batch 越大、学习率就越能调大」这句,还有一层很多人不知道的天花板

⭐ 先把常见的那两条摆出来:

  • 平方根缩放:batch 扩大 n 倍、学习率扩大 √n 倍。 —— 推导的出发点是让更新量的噪声强度保持不变 (增量协方差 ∝ η²/B,要它不变就得 η ∝ √B)。
  • OpenAI 用梯度噪声尺度 Bnoise 给出 η* ≈ ηmax /(1 + Bnoise/B) —— 注意这个形式是单调递增且有上界的: batch 再大,学习率也只能逼近 ηmax,不会一直涨。

⭐⭐ 光这一条就够破一个想当然了: 「batch 翻倍,学习率就该翻倍」是错的 —— 它有上界,而且越往后收益越小。

⛔⛔⛔ 但在 Adam 底下,还可能更反直觉一点: 超过某个 batch 之后,最优学习率不升反降 —— 这就是所谓的 Surge 现象

⭐⭐⭐ 为什么会这样? 直观的说法是:这本质上是「自适应学习率」本身次优的体现。

把 Adam 粗略看成 sign(g):batch 越大,估出来的 sign(g̃) 就越接近真的 sign(g)。 ⛔ 可 sign(g) 就是最好的更新方向吗?—— 不一定, 尤其到了训练后期。
⭐ 于是 batch 取中等的时候,那点噪声反而在替你修正这种次优; batch 再大,噪声没了,修正的机会也没了, 所以只好更谨慎地把学习率降下来。

⛔⛔ 但这条有前提,而且原作者自己特意收了一格

「(当然这里还有一个限制,β 是始终小于 1 的, 如果 βnoise ≥ 1,那么最优学习率与 Batch Size 的关系依旧是单调递增的。)」

⭐ 苏剑林还顺手批评了那篇原论文原论文在一个较强的近似下得出「Surge 几乎总会出现」,他认为不大科学; 他自己换了个更合理的近似,结论是 「即使 Hessian 的非对角元不可忽略,Surge 现象也不一定会出现」

📌 苏剑林《当 Batch Size 增大时,学习率该如何随之变化?》 kexue.fm/archives/10542; 所评论的原论文为《Surge Phenomenon in Optimal Learning Rate and Batch Size Scaling》。 ⚠️ 引号里为原文原话。这是作者本人的分析,不是同行评议结论 —— 我们收的是「别想当然地外推」这条警告,不是一个可以照抄的公式。

一个现代参考值,它看着完全不合上面那条规律: DeepSeek-V3 是 671B 的 MoE,峰值 2.2×10⁻⁴ —— 比 GPT-3 的 175B 大得多,学习率却高出近四倍

我先给过一个解释,而它被上面那张表当场推翻了

那个解释是:「V3 每个 token 只激活 37B,所以该按 37B 看」。 ⛔ 可你往上一行看:GPT-3 的 13B 是 1.0×10⁻⁴,175B 是 0.6×10⁻⁴, 37B 夹在中间,按那张表应该是 0.8×10⁻⁴ 上下 —— 而 V3 是 2.2×10⁻⁴。
换个口径没有把差距抹平,只是从「差 3.7 倍」变成「差 2.8 倍」。

⭐⭐⭐ 真正的答案是:这两个数本来就不该放在一起比

  • batch 完全不是一个量级。GPT-3 175B 是 3.2M token, V3 是 12.6M 起、爬到 62.9M —— 差二十倍。 而这一节反复说:学习率和 batch 是绑着动的。
  • 形状不一样。GPT-3 走余弦、早早就开始降; V3 是 WSD,峰值那个数要一路恒定扛过 10T token。 —— 「峰值」在两条曲线上不是同一个身份的东西。
  • 连参数都不是同一批。MoE 里绝大多数参数每一步只被一小部分 token 碰到。它挨的更新次数和稠密模型不是一个数量级。

⭐⭐ 所以这一格的教训,跟 2.5 那张图是同一条: GPT-3 那张表的价值,在于它内部八个规模是同一套配方、只有一个旋钮在动 —— 而这既是它能当先验的原因,也是它出了那张表就不能当先验的原因
判据:一条趋势只在「产生它的那次受控实验」内部可用。

⛔ 顺带自曝:我原来那句「看激活规模不看总参数」听起来很顺, 而它正是最危险的那一类断言 —— 听着像常识,还刚好能圆上场。 它错在拿一个没验过的换算去救一个跨实验的比较, 而跨实验的比较本身就不成立。

第三段 · 怎么降 —— 两个流派,终点却撞在同一个数上

⭐ 老派 · cosine(GPT-3 / Llama / Mistral 一系)

GPT-3 的做法:在 2600 亿 token 内按余弦曲线降到峰值的 10%, 之后就一直保持在这个十分之一不动

⭐ 新派 · warmup-stable-decay(DeepSeek-V3 就是这个形状)

V3 的完整曲线:2000 步线性升到 2.2×10⁻⁴ →  一直恒定,跑到 10T token 才开始降 →  用 4.3T token 余弦降到 2.2×10⁻⁵ →  最后 500B token 里,前 333B 恒定 2.2×10⁻⁵,后 167B 换成 7.3×10⁻⁶。

📌 arXiv 2412.19437 §4.2;梯度裁剪范数 1.0, batch 从 3072 爬到 15360(头 469B token)后保持。 这个形状的系统分析见 arXiv 2410.05192

⭐⭐⭐ 两派的差别不在「哪个收敛更好」, 而在它们各自要求你什么时候做决定 —— 图 Ⓑ 那两栏讲的就是这件事。

⭐⭐ 而两派的终点撞在同一个数上 ——  图里那条横贯全场的虚线。

⭐ 落点:怎么调,三条就够

  • ① 别从零猜。找一个规模和架构都接近的公开配置抄过来当起点 —— 上面那张表和 V3 的报告都是公开的,这是免费的先验
  • ② 跨规模不能直接照抄。回到 3.3 那条: 学习率带量纲。小规模上扫出来的值,要么用可迁移的参数化, 要么在目标规模上小步重验
  • ③ 判据看两条曲线。loss 炸了、或者 grad norm 出突刺 —— 太大; loss 下得慢、grad norm 一直贴着地板 —— 太小

⛔ 梯度裁剪是安全带,不是调参手段。 V3 设的是 1.0。靠调裁剪阈值来压住一个太大的学习率, 是在掩盖问题不是在解决问题。

⛔⛔ 最后一条,也是最容易被忽略的: 单报一个学习率数字,是没有意义的。

它必须连着 batch、schedule 形状、总 token 数 一起报 —— 换掉其中任何一个,这个数就不再适用。 ⭐ 上面那两套配置之所以能抄,正是因为它们四样都写全了

3.6 ⭐⭐ 另外那两个超参:β₁ 和 β₂ 到底在干什么

学习率讲完了,可 Adam 还有两个数没讲β₁ 和 β₂。 ⭐ 它们通常被当成「用默认值就行」的东西 —— 但它们的含义其实特别清楚。

⭐⭐ 一句话:β 决定「老的贡献能活多久」。

滑动平均每步乘一个 β 再掺新值,所以有效窗口大约是 1/(1−β)

  • β₁ = 0.9 → 动量的窗口约 10 步。它只负责把方向抖动抹平。
  • β₂ = 0.95 → 20 步β₂ = 0.99 → 约 100 步

⭐ 所以 β₂ 调大 = 分母更平稳但反应更慢;调小 = 跟得紧但噪声大。 —— 这不是玄学,就是一个「记多久」的旋钮

⭐⭐⭐ 而 PaLM 在这里做了一件反常的事:让 β₂ 随步数变

它用的是 β₂ = 1 − k−0.8k 是当前步数

代几个数进去看看它在干嘛:

⭐⭐ 也就是说,窗口是跟着训练一起长的(大约按 k0.8)。 而固定的 β₂ = 0.99,意味着从头到尾都只记 100 步

❓ 为什么要这么做?论文给的理由很具体,而且很好懂。

因为罕见的 embedding token ——  一个生僻字可能几千步才出现一次。在一个 100 步的窗口里, 它的二阶矩估计基本是垃圾:大部分时候是 0,偶尔来一个大值。

⭐ 而窗口跟着训练一起长,这个问题就自己消失了。 论文原话:「我们发现这比标准的 β₂ = 0.99 更稳定」。

📌 arXiv 2204.02311 的超参段。同段还有: β₁ = 0.9、学习率前 1 万步固定 10⁻²、之后按 1/√k 衰减。

⚠️ 这里有个顺带推出来的结论,而它跟 3.2 那条判据打架 ——  值得说清楚。

3.2 说:动量和二阶矩是滑动平均,老的贡献会被衰减掉,所以 bf16 扛得住。
可如果 β₂ 一路涨到 0.9999,那个「衰减」就慢得快要不衰减了 —— v 越来越像一个真的累加器。

⭐⭐ 所以那条判据要带上前提:它成立的条件是「β 不随步数趋近 1」。 ⛔ 这一条是我们自己推的,不是论文说的 —— PaLM 用的是 Adafactor 系、也没报精度配置, 能不能真观察到这个影响,没有证据。当猜想看。

⭐ 还捡到一条:PaLM 把 GPT-3 手调的那张表自动化了

PaLM 用的其实是 Adafactor,但关掉了分解 —— 论文说这等价于「带参数缩放的 Adam」: 按参数矩阵的均方根去缩放学习率。

⭐⭐⭐ 而论文自己点破了这跟 3.5 那张表的关系: 「这么做的效果,类似于 GPT-3 那样手工把学习率随规模调小」。

⭐ 换句话说:3.5 那张七档表是手动解法,参数缩放是自动解法。 而且论文说自动那版还多一个好处 ——  尺度本来就不同的那些参数(embedding、layer norm 的缩放) 不会被按同一个比例一起压下去。

3.7 ⭐ checkpoint 要存的,正好就是那 12 字节

⭐ 这一节能用 3.1 那张账单直接推出来, 不用查任何文档。

那 16 字节里,哪些必须进 checkpoint?一格一格问「丢了能不能重建」:

⭐⭐⭐ 所以 checkpoint 的大小 = 那 12 字节 × 参数量。 对 671B:约 7.3 TiB 一份。

⚠️ 但别把这个数安到 V3 头上 ——  3.2 刚说过它的一二阶矩是 bf16,那就是 4+2+2 = 8 字节, 一份 4.9 TiB,比经典口径小三分之一。 —— 省状态的那些做法,省的同时也在省 checkpoint。

这一下解释了两件工程上的事:

  • 为什么 checkpoint 这么大、存这么慢。 —— 它存的不是模型,是优化器。模型本身只占其中的三分之一。
  • 为什么「发布的模型权重」比 checkpoint 小得多。 —— 发布只要那 2 字节的推理权重,而 checkpoint 里那 12 字节全是训练用的 —— 一比六

⛔⛔ 但只存这 12 字节还不够 —— 还有三样,漏一样都会让「恢复」变成「重来」。

  • ① 步数 —— 因为学习率、β₂、weight decay 都可能是步数的函数(见 3.5、3.6)。
  • ② 数据读到哪儿了 —— ⭐ 这一条最容易漏,而且漏了最隐蔽: 恢复之后又把同一批数据重训一遍,loss 曲线看着完全正常
  • ③ 随机数状态 —— dropout、数据打散都要用。

⭐⭐ 而 §6.2 那个「回滚 + 跳数据」的救火办法,前提正是 ①②。 —— 数据位置说不清,你连「跳掉哪几批」都没法讲。

顺带一个真实的省法:优化器那部分不一定要留在加速器上DeepSeek-V3 就把权重的指数滑动平均放在 CPU 内存里、每步异步更新 —— 既不占显存也不占训练时间,用途是提前估计「学习率衰减之后会有多好」。 📌 arXiv 2412.19437 §3.2.3。

3.8 ⭐⭐⭐ LoRA:大多数人真正会碰到的那种训练

⭐ 前面整讲算的都是从零开始训—— 可绝大多数人这辈子跑的训练,是拿别人训好的模型接着调

⭐⭐ 好消息是:这件事不需要任何新概念它就是上面那张 16 字节的账,换一个算法重算一遍

📌 LoRA 一句话:底座冻住不动,只在旁边挂一对很瘦的小矩阵。

原来那个 4096 × 4096 的大矩阵不训了; 旁边加一个 4096 × 16 和一个 16 × 4096只训这两个

关键在那个 16两个瘦矩阵加起来 2 × 4096 × 16,只有原矩阵的 0.8%

那三本账各自怎么变

全量微调LoRA省了吗
常驻
权重 + 梯度 + 优化器状态
16 B × 全部参数 2 B × 全部(冻住,只读)
+ 16 B × 那 0.1%
⭐⭐⭐ 省得最狠
算力3×(前向 1 + 反向 2) 约 2×⭐ 省了约三分之一
激活整条链都要在场 整条链还是都要在场基本不省

⭐⭐⭐ 第一行拿 70 亿参数代进去,数字很夸张。

  • 全量微调6.74e9 × 16 B = 100.4 GiB
  • LoRA(每层给 Q 和 V 各挂一对,秩 16): 冻结底座 12.55 GiB + adapter 0.125 GiB12.68 GiB

⭐⭐ 相差 7.9 倍 ——  而 adapter 一共只有 839 万个参数,占全模型 0.124%

为什么差这么多那 16 字节里有 14 字节 (梯度 + 主权重 + m + v)是只有「要被更新的参数」才需要的。 冻住的参数只留 2 字节权重 —— 它退回成了推理的样子。

⭐⭐ 第二行那个「3× 变 2×」值得单独说,因为它直接来自 1.3。

1.3 说过反向要付两笔乘法: ① 算输入的梯度(往下一层传),② 算权重的梯度(给优化器用)。

⭐⭐⭐ 而冻住的那些层,第 ② 笔可以整个不算 ——  它的权重根本不更新,算出梯度也没人用。

于是:前向 1 + 只剩输入梯度那一笔 1 = 约 2×,而不是 3×。

⭐ 这一条是从 1.3 那两笔账直接推出来的,不是新知识 —— 这正是把 LoRA 放在这里讲的理由。

⛔⛔ 而第三行才是最容易被误解、也最值钱的一行。

很多人以为「只训 0.1% 的参数 = 训练变得很轻」。显存上不是。

⭐⭐⭐ 因为 adapter 挂在每一层上。 —— 最前面那层的 adapter 也要梯度, 而梯度只能从 loss 一路传回去(1.5)。 所以整条链的激活,该在场的还得在场。

⭐ 结论很具体:LoRA 之后,激活那一块往往成了新的大头 —— 所以 LoRA 还是要开重算第二节)。 两件事不是替代关系,是叠加关系。

⚠️ 严格说激活会少一点:冻结的线性层不再需要为「权重梯度」 留住它的输入。但非线性那些算子照样要,所以只是少一点,不是少一个数量级。 ⛔ 这一句是按 1.3 那两笔账推的,我们没有实测。

⭐⭐ 一句话收口:LoRA 改的是「有多少参数带着那 16 字节」, 它没有改变「反向要穿过整个网络」这件事

而这正好是这一讲的主轴在小尺度上的一次复现参数那一块可以按比例砍,而链条本身砍不动 —— 因为它来自反向传播的数学,不来自你训多少参数。

第 四 节

这张账单直接决定了并行策略长什么样

这一节是通向专题五的桥。 ⭐ 把上面几项按大小排一遍,你会发现 ——  ZeRO 那三级,正是照着这个顺序来的。

4.1 先把账按大小排一遍

把前面三节的东西摆在一起,按占多少排个序 ——  这个顺序本身就是这一节的全部内容。

4.2 ZeRO 的三级,正是照着这个顺序来的

⭐ ZeRO 的分级不用背。 把上面那个顺序倒过来读一遍,你就把它推出来了。

ZeRO 的三级 —— 同一根条,被削掉三次 顺序不用背:谁最大、谁最少被用到,就先切谁 权重 梯度 优化器状态 Ⓐ 横条的长度 = 每张卡真正要装下的量 75 亿参数、64 路数据并行算 —— 论文 Figure 1 用的就是这组 基线 纯数据并行 每张卡都存一整份 权重 梯度 优化器状态 120.0 GB 每张卡 ZeRO-1 切优化器状态 Pos 权重 梯度 切走了 31.4 GB 每张卡 跟基线一样 ZeRO-2 再切梯度 Pos+g 权重 切走了 16.6 GB 每张卡 跟基线一样 ZeRO-3 最后才切权重 Pos+g+p ≈ FSDP 切走了 1.9 GB 每张卡 1.5 倍 还要装多少 通信量 三段的长度按同一个换算常数画,所以长短可以直接比 这四个数不是抄的,是算的 —— 脚本里带 assert 跟论文 Figure 1 对账,对不上就不让构建。能算的就别抄。 Ⓑ 顺序不是历史巧合 —— 是「多大」和「多久用一次」排出来的 所以这三级不能换顺序,也不用背 每步用几次 → ↑ 每参数几字节 1 次 10 100 2 B 12 B 1|优化器状态 每步只用一次 通信不变 2|梯度 反向结束汇总一次 通信不变 3|权重 每层前向都要 通信 ×1.5 这一角是空的:又大又忙的东西 顺序不用背 —— 沿着「先大后忙」这条线走一遍,排出来的就是 ZeRO 的三级。 而这条判据出了 ZeRO 还能用:先动又大又闲的那一项,最后才碰又小又忙的。 顺带把「基线的 2Ψ」拆开看 —— 它解释了为什么切优化器状态是白捡的 纯数据并行每步那笔 all-reduce,其实是两步:先 reduce-scatter(Ψ)把梯度归约并散开,再 all-gather(Ψ)把结果收回来 —— 合起来 2Ψ。 而 ZeRO-1 要的正好是「散开」那个中间状态:每张卡只更新自己那一份优化器状态。它没有多要一次通信,它只是没把中间结果扔掉。
⭐⭐⭐ Ⓐ 是同一根条被削掉三次 —— 虚线框就是切走的部分,右边那个 GB 是每张卡真正还要装下的量。
⭐⭐ Ⓑ 解释了顺序:「多大」乘以「多久用一次」。优化器状态最大、又只在更新那一瞬间用 —— 所以它先被切。
能算的就别抄 —— 抄来的数对不上时,你不知道是自己抄错了还是人家印错了;算出来的对不上,当场就报错。
出处与口径

📌 arXiv 1910.02054 §3(16 字节口径与 Figure 1 的四个数)、§7.2(通信量 2Ψ / 2Ψ / 3Ψ)。

⭐ 图上四个 GB 数由 75 亿参数 × 每参数字节数算出,脚本内 assert 与论文 Figure 1 对账(120 / 31.4 / 16.6 / 1.9 GB)。

⭐⭐ 切的顺序只有一条原则:谁最大、谁最少被用到,就先切谁。

  • ZeRO-1 切优化器状态。最大,而且只在更新时用一次 —— 切了几乎不增加通信。先切它,性价比最高。
  • ZeRO-2 再切梯度。它跟权重一样大,但反向结束前它可以是散着的 —— 在那之前它可以是散着的。
  • ZeRO-3(≈ FSDP)最后才切参数。因为每一层前向都要用它 —— 一切开,每层都得先 all-gather 把它拼回来。 通信代价是三级里最高的。

⭐⭐⭐ 所以顺序不是随意的,也不是历史巧合 ——  它是「大小」和「被用到的频率」这两个量排出来的。

📌 出处:Rajbhandari 等,ZeRO: Memory Optimizations Toward Training Trillion Parameter Models,arXiv 1910.02054⭐ 上一节那个「每参数 16 字节」也出自它 ——  同一篇论文既给了账单,也给了切法。

4.3 同一张账单,其它几招各在切哪一项

把 ZeRO 之外那几个常听见的名词摆回这张账单上,它们各自的位置立刻就清楚了

⭐⭐ 所以这几个名词不是「几种优化技巧」, 是同一张账单上的四个不同栏目各自的对策—— 而专题五整讲要做的,就是把「切」这件事本身摊开算。

第 五 节

一个完整 step 的总账 —— 以及峰值出现在哪一刻

把前向(专题一)+ 反向 + 更新合成一张表。 ⭐ 而其中最有用的一问不是「总共多少」,是 「峰值出现在哪一个时刻」

一个 step 的总账 —— 「峰值在哪一刻」要先问「哪一项的峰」 这里算的是全局总量,不是单卡 · 横轴只标阶段,不是真实时长比例 权重 · 不动 优化器 · 不动 激活 · 前向堆 梯度 · 反向堆 Ⓐ 两条不动的,加两条形状正好相反 下面那四层按真实比例画(global batch = 1 条)—— 所以优化器那一块最厚 开始 前向结束 反向结束 更新完 一个 step 的时间轴 权重 1.22 TiB 优化器状态 7.32 TiB 全程不动,而且是四项里最大的一块 —— 却只在最后那一瞬间被用一次 梯度 1.22 TiB + 激活 106.75 GiB —— 薄成这样是真的 ↑ 上面那两层放大 20 倍 —— 真实比例下它们太薄,两个峰看不出来 激活的峰 前向刚结束,反向还没开始 梯度的峰 反向刚结束,还没更新 两个峰不在同一时刻 先问「哪一项的峰」,再问「在哪一刻」 Ⓑ 「激活大还是优化器状态大」—— 这个问题没有固定答案 因为常驻那块不随 batch 变,而激活线性地随 batch 涨 常驻那一块(不随 batch 变) 9.76 TiB 权重 2 + 梯度 2 + 优化器 12 = 16 B/参数,× 671B 激活 · global batch = 1 条 0.10 TiB 开了全量重算的 128K 序列 激活 · global batch = 94 条 9.76 TiB 到这里才追平 —— 这个数就是分水岭 所以 「谁最大」不是模型的属性,是这次训练配置的属性 —— 换个 global batch,答案就换了。
⭐⭐⭐ 全专题的收口图 —— 前面四节各自那一块,第一次摆到同一条时间轴上。
Ⓐ 要看的是两条形状相反的带子:蓝的前向堆、橙的反向堆 —— 它们的高点差了大半个 step
⭐⭐ Ⓑ 三根条全是当场算的,脚本里带 assert —— 改任何一个输入,对不上会自己报错。
📌 梯度累积到底省什么、不省什么,在 1.8 有完整的五条
出处与口径

⭐ Ⓑ 三行都是当场算的:671e9 × 16 B ÷ 1024⁴ = 9.76 TiB;9.76 TiB × 1024 ÷ 106.75 GiB = 93.7 条 —— 脚本里有 assert,改数会自己报错

⚠️ 16 B/参数是 ZeRO 论文的经典口径(arXiv 1910.02054);DeepSeek-V3 实际把一二阶矩放成了 bf16,它的每参数账跟这个数不一样(见 3.2)

⚠️ 激活那 106.75 GiB 是自己按算子推的估算,没有第三方背书 —— 所以那个「94 条」也是量级,不是阈值

⛔ 全图算的是全局总量,没有任何并行切分 —— 怎么把它摊到多少张卡上,是专题五整讲的事

⭐⭐⭐ 同一笔账,让它跑一遍。 上面那张图画的是这座山的形状;这段画的是它怎么被堆起来、又怎么塌下去
常驻(权重 + 优化器状态,14 B/参数)—— 始终不动 |  激活 —— 前向堆高,反向释放 |  梯度(2 B/参数)—— 反向才出现,反向结束时最全 |  最下面那排小格 = 61 层,走到哪一层就点到哪一格
⭐⭐ 盯住那根竖线:最高点出现在前向刚走完、反向还没开始的那一刻 —— 不是训练的某个阶段,是一个瞬间。 (三条带的高度按真实字节数算,跟上面那张图同一套口径。 14.6 秒无声循环,Manim 渲染,脚本在 tools/manim/。)

5.1 显存:四项加起来是多少

⭐ 前面四节各算各的,现在把它们摆到一起这一节全部按全局总量算 —— 不分卡,怎么切是下一讲的事。

📌 常驻那一块的算式,短到可以口算:

671e9 参数 × 16 字节 ÷ 1024⁴ = 9.76 TiB。 拆开是:权重 1.22 + 梯度 1.22 + 优化器状态 7.32

换算一下就知道这是个什么量级 ——  按每张卡 80 GiB 算,光这一块就要 125 张,还没算激活、没算任何冗余。 (80 GiB 是这里随手设的换算基准,不是在说某款具体的卡。)

⭐⭐ 注意这三项里,只有优化器状态那 7.32 TiB 是纯开销 —— 它既不参与前向,也不参与反向, 只在每一步的最后那一瞬间被用一次。

⭐⭐⭐ 顺着这个数往下问一句,会撞上一件反直觉的事。

常驻那 9.76 TiB 是定值 —— 跟你喂多长、喂多少条都无关。 那激活要涨到多大才追得上它?

按一条 128K 序列 106.75 GiB 算:94 条 —— 折合大约 1,230 万 token

⭐⭐ 而这里有个容易被跳过的前提:那 1,230 万 token 怎么凑出来的,显存不在乎。

  • 94 条 128K,还是 3,000 条 4K ——  对显存是同一件事
  • 为什么激活账里没有随 S² 涨的项 —— 注意力分数矩阵根本不落显存(1.2 那三条边界的第二条)。 剩下的每一项都是「每个 token 一份」,所以它只认 token 总数。

但算力不是这样。

同样 1,230 万 token,拆成长序列拆成短序列, 要算的浮点数差很多 —— 因为 attention 那部分随 涨, 而它不在显存账上、却在算力账上

⭐⭐⭐ 所以序列长度这个旋钮,在两本账上的行为不一样显存那本只看总量,算力那本还要看你怎么切。 —— 下一节那个「6ND 严重低估」,讲的就是这件事。

⚠️ 前提说清楚:以上都建立在用了 FlashAttention 这一条上。换成把分数矩阵写进显存的朴素实现, S² 就回到显存账里,上面整段都不成立。

5.2 算力:那个人人都在用的 6ND,在长上下文下严重低估

估训练算力最常用的一条是 6ND每个参数、每个 token,大约 6 次浮点运算(前向 2、反向 4 —— 正好对上第一节那个 1+2)。

代进去(MoE 要用激活参数,不是总参数): 6 × 37e9 × 131,072 ≈ 29.1 PFLOP,一条 128K 序列。

⛔⛔ 但这个数在 128K 上是错的 —— 而且错得很多。

6ND 数的是「跟权重相乘」那部分。 而 attention 里 QKᵀ 和 AV 这两步没有权重参与 —— 它们根本不在这个公式的账里。

附录 7.2 那一行说:在 128K 上, attention 占一层前向算力的 82.1%。 也就是说 6ND 数到的只是剩下那 17.9%

⭐⭐⭐ 反推一下:真实算力大约是 6ND五到六倍 —— 29.1 PFLOP 实际接近 160 PFLOP

⚠️ 这个倍数是由 82.1% 反推的(1 ÷ 0.179), 而 82.1% 本身是自己按算子推的 —— 当量级看。 ⭐ 要记的不是这个倍数,是那个公式什么时候会骗你:序列一长就会

⭐ 反过来说,6ND 在短序列上是很好用的 —— GPT-3 那个 2,048,attention 只占百分之一点几, 漏掉它完全无所谓。又是同一件事:不是公式错了,是前提变了。

5.3 峰值:先问「哪一项的峰」

第一节那张山形图给过一个答案:峰值在前向末尾。 那句话只在谈激活的时候成立。

把四项一起画到时间轴上(就是上面那张图), 会看到一件当时看不出来的事:不同的项,峰在不同的时刻。

⭐⭐⭐ 推广出去,这是一条到处能用的判据:

把几条形状不同的曲线叠起来之后, 总和的峰,未必落在任何一条单独的峰上。 —— 所以「什么时候最挤」这个问题, 必须连着「挤的是哪一项」一起问。

📌 工程上的直接后果: 你去查 OOM,光知道「峰值 xx GiB」没用 ——  要知道是哪一项在那一刻最大,才知道该去动哪个开关。 激活大就开重算 / 上 CP;优化器状态大就上 ZeRO ——  这两条路走反了,一点用都没有。

5.4 ⭐⭐⭐ 同一张账,缩到一台你摸得到的机器

⭐ 前面那些 TiB 级的数字都是对的,但它们没有分母—— 所以这一小节把同一张账,缩到一个你能自己算一遍的模型上。

📌 设定:一个 1.25 亿参数的模型,一张卡,batch 8 × 序列 1,024。

12 层、宽度 768、MLP 宽 3,072 ——  这就是 GPT-3 论文表 2.1 的第一行3.5 那张八档表最上面那档)。

第一步 · 常驻那一块

照搬 3.1 那个 16 字节:

每参数1.25 亿参数
bf16 权重2 B0.23 GiB
bf16 梯度2 B0.23 GiB
优化器状态(fp32 主权重 + m + v)12 B1.40 GiB
小计16 B1.86 GiB

第一个可以自己验的结论这个模型的权重只有 0.23 GiB,而它要占掉 1.86 GiB ——  八倍。多出来的全是训练才要的。

第二步 · 激活那一块

一层里每个 token 留下的东西,一项一项数(宽度 768):

留下的张量宽度
两个 norm 的输出2 × 768 = 1,536
Q / K / V3 × 768 = 2,304
attention 输出768
MLP 第一层输出3,072
激活函数的输出3,072
MLP 第二层输出768
一层合计11,520 个数 = 22.5 KiB(bf16)

× 12 层 = 每个 token 270 KiB。 × 8,192 个 token = 2.11 GiB

⛔⛔ 停一下 —— 这里出了一件值得吃惊的事。

⭐⭐⭐ 激活 2.11 GiB > 常驻 1.86 GiB。 一个只有 1.25 亿参数的小模型, 激活已经比权重、梯度、优化器状态加起来还大了。

⭐ 而它不需要长上下文,序列才 1,024,光是 batch 开到 8 就够了。 —— 所以重算不是大模型的专利。

第三步 · 开了重算之后

2.2 那一档:每层只留入口那一份(768 宽), 外加当前正在重算的那一层的完整一份。

0.32 GiB —— 2.11 掉到 0.32,省了 6.7 倍, 代价是总算力从 3× 变 4×。

不开重算开了重算
常驻1.86 GiB1.86 GiB
激活2.11 GiB0.32 GiB
合计3.97 GiB2.18 GiB
算力4×(+33%)

⭐⭐ 这张小表,把这一讲的三条主线全串上了。

  • 16 字节那张账3.1)—— 1.86 GiB 是这么来的。
  • 激活随 token 总数长(5.1,就在本节)——  2.11 GiB 里唯一的变量就是那 8,192 个 token。
  • 拿算力换显存2.2)——  多付 33%,省 6.7 倍。比例跟 671B 上那笔几乎一样。

⭐⭐⭐ 而最后这一句才是这一小节存在的理由: 换了 5,000 倍的规模,那些比例基本没变,变的只是绝对值。 —— 所以这一讲教的是比例,不是那些数。

⚠️ 三条边界,跟 1.2 那张表同一套。

  • 这些数是按算子逐项推的,不是从哪份报告抄的 —— 框架的算子融合会让它小一些。当量级看。
  • 没算注意力分数矩阵(假设用了 FlashAttention), 也没算临时缓冲、碎片、通信缓冲 —— 真跑起来还要留富余。
  • 「开了重算 0.32 GiB」假设每层都设了存档点 —— 2.2 那张图的 Ⓒ 说了,这不一定是最省的密度。

5.5 这张账单交给下一讲的是什么

专题一算完前向,结论是「装不下」这一讲把账补全之后,那句话变得具体了 ——  不是「装不下」,是「四项各自装不下,而且各有各的治法」。

⭐⭐ 四项 + 四种对策,一一对上:

  • 优化器状态(最大,最少被用到)→ (ZeRO-1)
  • 梯度(跟权重一样大,但反向结束前可以散着)→ (ZeRO-2)
  • 权重(每层都要用)→ 最后才切(ZeRO-3/FSDP),因为通信最贵
  • 激活(随 batch × 序列长度涨)→ 重算,或者按序列(CP)

⭐⭐⭐ 你会注意到四条里有三条写着同一个字:—— 那就是专题五整讲要做的事:切这件事本身,也是要算账的。

第 六 节

训不崩 —— loss 飞了怎么办,以及怎么让它别飞

前面整整一讲都在算可账算得再清楚,也回答不了训练现场最常见的那一句: 「loss 突然飞上去了,怎么办?」

训不崩 —— 把它当成一条河,三段各治各的 这一节最容易写成一张技巧清单 —— 列完了还是不知道该先动哪个 上游 · 改河道 中游 · 装闸门 下游 · 捞船 Ⓐ 把它想成一条河 —— 同一场水患,三个位置都能下手 从上游到下游:越靠上越治本,越靠下越应急 上游 改河道 把河挖宽、挖直 —— 水本来就不容易漫出来 · 归一化放进残差块里(Pre-LN) · QK-norm —— 别让注意力分数长疯 · 初始化按深度缩一下 一次性:开工前做完,之后不用管 中游 装闸门 水位一高就关闸 —— 不让它冲下去 · z-loss —— 摁住 softmax 前的 logits · 梯度裁剪 —— 按全局范数,常设 1.0 · 关键处用高精度算 常驻:一直开着,便宜,多数不伤质量 下游 捞船 已经翻了 —— 把船捞回来重新开 · 回滚到 spike 之前的 checkpoint · 跳掉那一段数据,再往下跑 · 把出事的时间点和数据位置记下来 每出一次事就得干一次 —— 而且没修好 水往这边流 → loss 飞了 越往下游,越是在救火 —— 而上游那几样开工前做一次就完了,下游那几样每出一次事就得再干一次 Ⓑ PaLM 做过一个消融,结论跟所有人的第一反应相反 这个消融只花了一次重训,却把一整类猜测排除掉了 拿「坏数据」这个假设验一验 同一批数据 + 换一个时刻 结果:不飞。 所以不是数据坏 两个条件同时满足才出事 缺一样都不会飞 这解释了为什么回滚 + 跳数据管用 Ⓒ 这一节真正的难点:让它稳住很容易,难的是稳住而不牺牲质量 ST-MoE 论文 Table 4 的原数 —— 三种做法,三次独立训练 基线 -1.755 稳定 4/6 六次跑崩两次 收紧 update clipping -4.206 稳定 3/3 条短了一大截 —— 看这里 router z-loss -1.741 稳定 3/3 条最长,而且一次没崩 横条越长质量越好(这是个越大越好的指标)。中间那根「稳定 3/3」看着完美,可它把质量打穿了。
⭐⭐ Ⓐ 按「治在链条的哪一步」排,不按「有哪些技巧」排 —— 技巧清单列完了,读者还是不知道该先动哪个
⭐⭐⭐ Ⓑ 是 PaLM 那个消融:同一批数据换个时刻喂就不飞,所以 spike 是数据和参数状态撞出来的,不是数据本身坏。
Ⓒ 中间那根条看着完美(稳定 3/3),可它把质量打穿了
出处与口径

PaLM 540B 的 spike 与消融:arXiv 2204.02311 §5.1 —— 「大约 20 次」「回滚约 100 步」「跳 200–500 批」「不认为是数据本身坏」都是原文

router z-loss 与 Ⓒ 那张表:arXiv 2202.08906(ST-MoE)。⭐ 论文自己的小标题就是「很多方法能稳住稀疏模型,但代价是质量变差

QK-norm:arXiv 2302.05442(ViT-22B)§2 —— 在约 80 亿参数处观察到训练发散,根因是注意力 logits 变得极大、注意力权重塌成几乎 one-hot(熵接近 0)

归一化的位置:arXiv 2002.04745(ICML 2020)—— 证明 Post-LN 在初始化时靠近输出层的梯度很大,所以必须靠 warmup 压住;换成 Pre-LN 则初始化时梯度就是良态的

⛔ 本图不含任何我们自己的实测 —— 四条全部来自公开论文,Ⓒ 那三个数是照抄 Table 4

6.1 先看一个真实的画面

⛔⛔ PaLM 训 540B 的时候,loss 飞了大约 20 次

而且 —— 梯度裁剪是开着的。 这些 spike 出现在极不规则的时刻,有时候训到很晚才来; 而更小的模型上根本没观察到。

📌 arXiv 2204.02311 §5.1。 「大约 20 次」「尽管梯度裁剪是开着的」都是论文原话。

loss spike —— 它长这样,而且没人能说清为什么 曲线是按论文文字描述画的示意图,不是实测数据 —— 所以纵轴故意不给刻度 540B:约 20 次 更小的模型:没观察到 裁剪开着 Ⓐ 同一张纸上放两条曲线,差别一眼就看出来 横轴是训练步数,纵轴是 loss —— 越低越好,所以主干是往下走的 loss 步数 → 540B 约 20 次 更小的模型 一次都没有 下面这条另起一根纵轴 —— 这里比的是形状,不是谁低 图上只有三件事,而这三件都是论文明写的: 时刻毫无规律 —— 挤在一起的有,隔很久才来一个的也有 很晚了还会来 —— 不是「开头不稳,稳了就没事」 小模型是平的 —— 所以小规模验过不代表大规模安全 Ⓑ 「保险开着还是飞了」听着像悖论 —— 其实是两个量 把上面那段放大来看:裁剪摁住的那条被削平了,飞掉的那条照样飞 ① 裁剪管的是这个量:梯度范数 (纵轴为对数,跨了一个量级) 裁剪阈值 这一截被砍掉 原本要迈这么大 削平的顶 红线每一处都贴着阈值、没有一处越过去 —— 保险 100% 生效了,一次都没漏 ② 可是飞掉的是另一个量:loss (同一段时间、同一批时刻; 纵轴也按这一段放大了) 六个时刻,一次不落 —— 摁住了「迈多大」,没摁住 loss 所以这不矛盾 保险装在 「这一步迈多大」上 事故发生在 「loss」上 不是失效了, 保的不是这件事 那到底为什么会飞? 论文说:没找到 所以 §6 后面讲的全部是「绕过去」和「降低概率」,没有一条是「根治」 知道一件事至今无解,跟知道它的三种治法一样重要 —— 否则你会以为是自己配错了参数。 判据:把「这是开放问题」写出来,比把它糊过去更有用。 糊过去的代价,是学员在自己的训练里反复找一个不存在的配置错误。
⭐⭐⭐ Ⓐ 只要看两条线的形状差别 —— 上面那条扎满了尖,下面那条是平的。而下面那条是同一套配方的小模型。
⭐⭐ Ⓑ 是这一节存在的理由:保险是开着的,事故照样发生。—— 所以后面那些办法全是「绕过去」,没有一条是「根治」。
⚠️ Ⓐ 那条曲线纵轴上一个刻度都没有,这是故意的 —— 没有刻度的曲线是示意,有刻度的才是数据。论文给了前者的描述,没给后者的数。
出处与口径

📌 arXiv 2204.02311(PaLM)§5.1 —— 「大约 20 次」「尽管梯度裁剪是开着的」「更小的模型上没有观察到」均出自该节原文。

⚠️ 曲线为示意:尖峰的位置按「看不出规律」这一条挑出来,高度与宽度论文未给数,故纵轴不设刻度。没有刻度的曲线是示意,有刻度的曲线才是数据。

⭐ 先记住这个规模效应:小模型上训得好好的,不代表大模型上不会飞—— 这跟 2.6 那条「不能跨规模照抄」是同一件事, 只不过这次翻车的不是收益,是训练本身。

6.2 ⭐⭐⭐ 那个最反直觉的消融:不是数据坏

第一反应都一样:肯定是那批数据有问题。 PaLM 团队去验了这件事,做法很干净:

把 spike 前后那几批数据单独拎出来,从另一个更早的 checkpoint 重新喂一遍 —— 结果不飞。

⭐⭐⭐ 所以 spike 不是「坏数据」造成的, 是这批数据当时那个参数状态 撞在一起才出的事。

换个时刻喂同样的数据,什么都不会发生。

⭐ 这一下就解释了他们那个看起来很土的办法为什么管用: 回滚到 spike 之前约 100 步的 checkpoint,跳掉那 200–500 批数据,继续跑 —— 之后同一个点就不再飞了。

⛔ 注意这不是「修好了」,是「绕过去了」。 论文自己也写着:由于训练成本太高,他们没能找到一个有原则的缓解办法。 ⭐ 这句坦白值得原样转述 —— 这一行至今仍是开放问题。

6.3 ⭐⭐ 这一节真正的难点:稳住很容易,难的是稳住而不掉质量

「让它别飞」有一堆办法:学习率调到极小、裁剪收到极紧…… 都能稳,代价是模型变差。

ST-MoE 那篇论文把这件事量出来了 —— 同一个配置换不同随机种子反复跑,看几次能跑完、以及跑完的质量:

做法稳定性质量(越大越好)结论
基线4 / 6−1.755六次里两次训崩
收紧 update clipping(0.1)3 / 3−4.206 ⛔ 稳了,可质量被打穿了
router z-loss3 / 3−1.741 ⭐ 稳了,质量还略好一点

⭐⭐⭐ 中间那一行是这张表的全部价值: 「稳定 3/3」看着完美,可它是拿质量换来的。 —— 评价一个稳定性手段,必须同时看这两栏。

⚠️ 两个口径要说清楚,不然这张表会被读得太重。

  • 分母不一样。基线跑的是 6 个种子,每个稳定性手段各跑 3 个种子 —— 所以「4/6」和「3/3」不是同一个分母下的比较。 ⛔ 3/3 也就是三次,样本小到不该当成「保证稳定」。
  • 这个实验是刻意挑出来的。作者写得很直白:小模型很少不稳定, 大模型又贵到跑不起足够的种子。于是他们特意选了一个约三分之一概率会崩的配置, 还特意换到多语种数据上 —— 因为那会让不稳定更频繁。

⭐ 这不是黑点,是好的实验设计 ——  研究一个罕见故障,你必须先把它变得不罕见。 但读数的时候要记住:这个概率是被调高过的。

📌 arXiv 2202.08906(ST-MoE)Table 4。 论文自己的小标题就是「很多方法能稳住稀疏模型,但代价是质量变差」。

6.4 治法一 · 治数值:不让某些量长太大

z-loss —— 摁住 softmax 前面那个 logits

做法出人意料地简单:在总 loss 上加一项,惩罚 softmax 归一化因子的对数的平方。 —— 它逼着那个 log Z 待在 0 附近,也就是不让 logits 长得太大

⭐ 两个地方都能加,而且都有人用:

为什么 z-loss 特别划算:ST-MoE 的解释是 ——  logits 后面接的是指数函数,而指数对输入误差极度敏感。 把输入范围压小,等于直接降低了这一步的数值风险, 而模型并不真的需要那么大的动态范围。

梯度裁剪 —— 它是安全带,不是调参手段

全局范数裁:把所有参数的梯度当成一个大向量, 范数超过阈值就整体等比例缩回去。常见设 1.0(PaLM、DeepSeek-V3 都是 1.0)。

⛔ 但 6.1 那句话要记住:PaLM 开着裁剪,照样飞了 20 次。 —— 裁剪拦得住单步的爆炸,拦不住「参数状态已经走到了一个坏地方」。

6.5 治法二 · 治结构:让它天生不容易飞

⭐⭐ 归一化放在哪儿 —— 这件事顺带回答了 3.5 那个悬案

原始 Transformer 把 LayerNorm 放在残差块之间(Post-LN)。 后来大家改成放在残差块里面(Pre-LN)。 这不是风格问题。

⭐⭐⭐ 2020 年有一篇论文用理论证明了这件事, 而它的结论直接接上了3.5

Post-LN 在初始化的那一刻,靠近输出层的梯度就很大 ——  这时候用大学习率必然不稳。而 warmup 正是在实践中压住这个问题的办法。
反过来,Pre-LN 在初始化时梯度就是良态的 ——  所以论文说:Pre-LN 可以把 warmup 那一段去掉。

⭐⭐ 所以「为什么要 warmup」这个问题,历史上有过两个答案: 2020 年那篇说是被 Post-LN 逼的3.5 引的 2024 年那篇说它让网络能承受更大的峰值学习率。 ⭐ 两个都对,而且不矛盾 ——  前者解释了「为什么当年非要不可」,后者解释了「为什么今天还留着」。

📌 arXiv 2002.04745(ICML 2020),用平均场理论证的。

⛔⛔ 不过上面那段只说了「Post-LN 的梯度不好」,没说为什么 而机制其实一句话就能讲完,讲完之后还能顺手回答一个更扎人的问题: 既然它不好,那为什么不干脆去掉?

⭐⭐⭐ Post-LN 是 xt+1 = Norm(xt + Ft(xt))

初始化那一刻,x 和 F(x) 可以看成两个独立随机向量, 各自方差是 1,加起来方差就是 2 —— 而 Norm 要把方差拉回 1, 所以它在这一刻相当于「除以 √2」。

⭐⭐ 一层一层递归下去: 最初那个输入 x₀,在第 l 层输出里的系数就是 2−l/2 —— 残差那条「直通路」被指数地削掉了,越靠近输入削得越狠。

⭐ 所以苏剑林那句话说得很重,但它有依据: 「在 Post Norm 的 BERT 模型中,LN 不仅不能缓解梯度消失, 它还是梯度消失的『元凶』之一。」

Post-LN 里,归一化本身就是梯度消失的元凶之一 —— 可它还留着,而且有道理 下面画的是前向的残差直通项 —— 「所以梯度也消失」那一步由正文交代,图上不冒充 Pre-LN:平权 Post-LN:2^(−l/2) 同一条曲线,两种价值 左右两边是同一张图 —— 变的只有底下那个结论 横轴层号,纵轴(对数)= 最初那个输入,在这一层的输出里还剩多少 还剩多少 层 → 1 10⁻1 10⁻2 10⁻3 10⁻4 Pre-LN Post-LN Ⓐ 预训练的时候 残差「名存实亡」 越靠近输入,信号被削得越狠 前面的层几乎学不到东西 还剩多少 层 → 1 10⁻1 10⁻2 10⁻3 10⁻4 Pre-LN Post-LN Ⓑ 而拿去微调的时候 这正好是你要的 微调只想动靠近输出的那几层 它替你把前面的层按住了 曲线一模一样,连坐标都没动 —— 换的只是「你拿它来干什么」。到第 24 层,Post-LN 的直通项只剩 Pre-LN 的 1/4096 那为什么不干脆把 LN 去掉? —— 因为方差会一路涨上去 去掉之后 x + F(x) 的方差就是 2,残差越多方差越大,所以还是得加一个 Norm —— 问题从来不是「加不加」,是加在哪儿 Pre-LN 的加法是 x + F(Norm(x)),最后总输出再加一个 Norm —— 这样每个残差分支是平权的,就没有上面那条指数衰减了。
⭐⭐⭐ 左右两边是同一张图 —— 连坐标都没动。变的只有底下那个结论框。
⭐⭐ 红线是 Post-LN:最初那个输入在第 l 层输出里只剩 2−l/2,到第 24 层是 1/4096;绿线是 Pre-LN,一直是 1。
Ⓐ 说这叫「残差名存实亡」,Ⓑ 说这正好替你按住了前面的层—— 同一条曲线,预训练时是毛病,微调时是功能。
画的是前向的残差直通项,不是梯度倍率—— 两者相关,但不是一回事。
出处与口径

📌 苏剑林《模型优化漫谈:BERT 的初始标准差为什么是 0.02?》kexue.fm/archives/8747 —— 「在 Post Norm 的 BERT 模型中,LN 不仅不能缓解梯度消失,它还是梯度消失的『元凶』之一」为原文原话,2^(−l/2) 出自该文公式的递归展开。

⚠️ 这是作者本人的分析文章,不是同行评议论文。⛔ 图上画的是前向的残差直通项,不是梯度倍率 —— 两者相关,但不是一回事

⛔⛔⛔ 那为什么不去掉它? —— 因为去掉之后 x + F(x) 的方差就是 2,残差越多方差越大。 问题从来不是「加不加 Norm」,是「加在哪儿」: Pre-LN 改成 x + F(Norm(x))、最后总输出再加一个 Norm, 这样每个残差分支是平权的,就没有那条指数衰减了。

⭐⭐⭐ 而最反直觉的一条是: 这个「毛病」,换个场合就变成了功能。

Finetune 的时候,我们本来就希望优先调靠近输出的参数, 不要过度动靠近输入的 —— 免得把预训练学到的东西破坏掉。 而梯度消失的意思恰恰是「越靠近输入,它对最终输出的影响越弱」, 这正好是 finetune 想要的。

⭐⭐ 所以预训练好的 Post-LN 模型, 往往比 Pre-LN 有更好的 finetune 效果⛔ 同一条曲线,预训练时是毛病,微调时是功能 —— 这就是上面那张图左右两半在说的事。

📌 出处:苏剑林《模型优化漫谈:BERT 的初始标准差为什么是 0.02?》 kexue.fm/archives/8747。 ⚠️ 这是作者本人的分析文章,不是同行评议论文 —— 我们按「一个讲得通的机制解释」收下,不当定论。

顺带把 3.5 那三个 warmup 解释的前提坐实了: 第三条说「压住后面的层、等前面的层跟上」, 它要的正是这里这条指数衰减 ——  Pre-LN 没有这条线,所以那一条对它也就不适用。

QK-norm —— 注意力分数长疯了

ViT 往上做到约 80 亿参数的时候,训了几千步就发散。 根因查出来很具体:注意力 logits 变得极大, 于是 softmax 之后的注意力权重塌成几乎 one-hot(熵接近 0)。

⭐ 治法也很直接:在做点积之前,给 Q 和 K 各做一次 LayerNorm。 —— 这就是今天很多大模型标配的 QK-norm 的来历。

📌 arXiv 2302.05442(ViT-22B)§2。 ⚠️ 出处要说准:QK-norm 不是 ViT-22B 发明的。 它的原文写的是「我们采用 Gilmer 等(2023)的做法」 —— ViT-22B 是把它用到 220 亿参数上并留下那张对照图的地方, 这才是它值得引的原因。 ⭐ 注意这又是一次「只在大规模上才出现」的故障 ——  跟 6.1 那条是同一个模式。

初始化 —— 一句话带过,但别漏

深网络里,如果每一层的残差分支都按同样的尺度初始化, 信号会随层数累积放大。常规做法是按深度把残差分支的初始化缩一下 —— 它不解决 spike,但它决定了你从一个多好的起点出发。

6.6 顺带清掉两个常见误解

❓ 「bf16 是不是也要 loss scaling?」—— 不用。

loss scaling 是 fp16 时代的必需品:fp16 的指数位只有 5 位, 小梯度会直接下溢成 0,所以要先把 loss 放大若干倍再算、更新前再缩回去。
⭐ 而 bf16 的指数位跟 fp32 一样是 8 位 ——  动态范围没缩,缩的是尾数精度。所以它不下溢,也就不需要 loss scaling。

fp16 的毛病不是「不够细」,是「下不去」 这一格只有一根真数轴和几条真边界 —— 每个数都是从 IEEE 754 的位模式当场算的 fp16 正规数下界 fp16 归零线 bf16 下界 Ⓐ 同一根对数轴上,两种 16 位格式的下界差了 33 个数量级 轴上越往左数越小;过了红线,那个数在 fp16 里就是 0 这一整段:bf16 存得下,fp16 存不下 33 个数量级 10⁻40 10⁻35 10⁻30 10⁻25 10⁻20 10⁻15 10⁻10 10⁻5 10⁰ 大 → bf16 1.18e-38 fp16 归零线 5.96e-08 fp16 正规数下界 6.10e-05 所以 fp16 的问题从来不是「16 位不够细」 —— bf16 也是 16 位,而且尾数比它还少。问题是它的指数位只有 5 位,下界抬得太高 Ⓑ 那 loss scaling 在干什么 —— 就是把整条分布在这根轴上往右推 乘 1024 = 在对数轴上平移 3.01 格 —— 够不够,看它跨没跨过红线 一个梯度 2e-08 fp16 存 → 0 × 1024 变成 2.05e-05 活下来了 诚实一句:放大之后它落在 fp16 的次正规区(红线右、橙线左)—— 存得下了,但精度打折。要进正规区得用更大的 scale。 而 bf16 根本不用做这件事 —— 它的下界在左边 33 个数量级之外,那些梯度本来就在范围内 所以「bf16 要不要 loss scaling」这个问题,答案在指数位上,不在位数上 bf16 和 fp16 都是 16 位,可 bf16 把位数分给了指数(8 位,跟 fp32 一样),fp16 分给了尾数(10 位)。 于是 bf16 更粗但更宽 判据:「精度不够」是个形容词,「这条线左边全没了」是可数的损失。 —— 把一个笼统的担心换成一根能指的线,误解自己就散了。
⭐⭐⭐ fp16 的毛病不是「不够细」,是「下不去」。
⭐⭐ Ⓐ 一根对数轴上三条真边界:fp16 正规数下界 6.10e−5、fp16 归零线 5.96e−8(再往左就真是 0)、bf16 下界 1.18e−38。中间那一大片绿 —— 33 个数量级,全是 bf16 存得下、fp16 存不下的。
Ⓑ loss scaling 就是在这根轴上整体右移:×1024 = 平移 3.01 格,一个本来归零的梯度就跨过红线活了。
每个边界都是从 numpy.float16 的位模式当场算的,「归零」也是真舍了一遍验的;没有直方图—— 那份数据我们没有,不编。
出处与口径

📌 图上每个边界都是从 numpy.float16 的位模式当场算的:最小正规数 0x0400 = 2⁻¹⁴,最小次正规数 0x0001 = 2⁻²⁴;bf16 与 fp32 同为 8 位指数,下界 2⁻¹²⁶。「低于红线会被舍成 0」是真舍了一遍验出来的。

⭐ 这张图的讲法取自混合精度那篇经典图(梯度直方图 + 一条 FP16 下界竖线,见 Narang & Micikevicius et al., 2018) —— ⛔ 但我们没有那份直方图数据,所以这里不画直方图,只画真边界。

⭐ 这跟 3.2 那条判据是配套的: bf16 换来的是「范围够、精度差」,所以怕的是累加,不是下溢

❓ 「训练是不是应该做到可复现?」—— 能做到,而且很值。

PaLM 做到了逐比特可复现:从任何一个 checkpoint 重启, 后续每一步的结果跟原来那次一模一样

靠的是两件事:确定性的计算框架, 以及确定性的数据管线 ——  后者的关键是「第几批数据」只由步数决定,跟进程数、恢复次数都无关。

⛔ 而 6.2 那个「回滚 + 跳数据」的办法,前提就是这个。 —— 数据管线不确定,你连「跳掉哪几批」都说不清。

📌 同样出自 arXiv 2204.02311

 ⭐⭐⭐ 全课落点:四条判据

这一讲报了几十个数字。但真正想留下的不是那些数。

  • ① 要不要高精度,比的是两个数「这一项每一步挪动的相对幅度」,对上「尾数能分辨的最小格」 —— bf16 那一格约是 1/256
    滑动平均每步都在遗忘,挪得动;主权重跨十万步累加、单步只挪千分之几,挪不动。 (3.2
  • ② 该不该重算,看「每省一字节要付多少 FLOPs」。 线性层的这个数恰好等于输入宽度;attention 的随序列长度线性上升 —— 两条斜率不同,必然相交。(2.3
  • ③ 两个量如果对同一个变量的增长速度不同, 那它们的排序必然在某一点翻转—— 所以任何「A 比 B 划算」的结论,都得带上它成立的那一段区间; 脱离区间的排序是没有意义的。
    ⭐ 本讲里它翻转过好几处,而且每处的那个变量都不一样: 序列长度(2.4)、模型规模(6.1,就在本节)、 以及同一个开关换规模连收益的正负号都变(2.6)。
  • ④ 评价任何一个稳定性手段,要同时看「稳没稳」和「质量掉没掉」。 只看前者的话,把学习率设成零是最优解。(本节 6.3)

⭐ 数字会过期,这四条不会。

这一节刻意一个我们自己的实测都没有 ——  四条全部来自公开论文并逐条核过原文。 理由写在 §七:这一讲已经有太多「自己推的」数了, 稳定性这一块不该再加。

附 录

这一讲的数是怎么核的

⭐⭐ 这一讲报了几十个数字。它们不是同一种东西有的能查到论文原文,有的是我们自己按算子推的 —— 后者可能差一倍,前者不会。

7.0 ⭐⭐ 这些东西在框架里叫什么

⭐ 这一讲讲了一堆概念。可回去打开框架,该搜哪个词?

这一讲里的说法在框架里叫什么核的地方
全量重算2.2 --recompute-granularity full Megatron-LM
training/arguments.py
core/transformer/transformer_config.py
选择性重算2.3 --recompute-granularity selective--recompute-modules
候选项:core_attn默认)/ layernorm / moe_act / mla_up_proj / mlp
梯度裁剪1.8 --clip-grad(默认 1.0,正好是 V3 用的那个值)
重算,手写的那种 torch.utils.checkpoint.checkpoint(...) PyTorch
torch/utils/checkpoint.py
选择性重算,按张量挑 create_selective_checkpoint_contextsCheckpointPolicy
⭐ 这就是 2.3 那条判据在 PyTorch 里的落点: 你自己写规则决定哪个算子存、哪个重算
重算(JAX 那边) jax.checkpoint(别名 jax.remat JAX
jax/_src/ad_checkpoint.py
选择性重算的现成策略 jax.checkpoint_policies.dots_with_no_batch_dims_saveable
⭐ 名字直接说出了判据:「矩阵乘的结果留着」 —— 正是 2.3 那张表最贵的那几行
梯度累积1.8 第 ④ 道) ddp.no_sync() —— 上下文管理器
进去之后梯度只在本地累加、不跨卡同步退出后的第一次 forward-backward 才同步一次(原文档语)
PyTorch
torch/nn/parallel/distributed.py
「跨卡汇总能藏进计算里」1.8 第 ② 道) bucket_cap_mb(默认 25 MiB
⭐ 梯度攒够一桶就发一次 ——  所以前面的层还在算,后面的桶已经在路上了。 那一道「✅ 能藏」,藏在这个桶里
ZeRO 的三级4.2 配置里的 zero_optimization.stage1 / 2 / 3 DeepSpeed
runtime/zero/config.py
8-bit 优化器3.4 bitsandbytes.optim.Adam8bit / AdamW8bit bitsandbytes
optim/__init__.py

⭐⭐ 这张表里有两行值得单独看一眼 —— 它们是这一讲的判据被写进 API 名字里了。

  • dots_with_no_batch_dims_saveable ——  「矩阵乘的结果存下来」。为什么是矩阵乘? 因为它每省一字节要付的 FLOPs 最多 —— 正是 2.3 那条闭式解。
  • ddp.no_sync() ——  ⭐⭐ 把这两行放一起看就清楚了:no_sync() 干的事, 正是把第 ② 道整个关掉。 累积期间一次不通信,攒完再一次性同步 ——  所以「梯度累积顺带省通信」不是副作用,它就是这个 API 的用途。
  • core_attn 是 Megatron 的默认 ——  而按 2.4,长上下文下它恰恰是最不该重算的那一个。 ⭐ 默认值是有年代的,这条在 2.6 展开过。

⚠️ 顺带承认这一讲少讲了一档。 PyTorch 的 CheckpointPolicy 其实有六个枚举值, 不是两个:{MUST,PREFER}_SAVE{MUST,PREFER}_RECOMPUTE、 还有 {MUST,PREFER}_CPU_OFFLOAD
⭐ 也就是说,框架眼里「这块激活怎么办」是三选一存在显存里 / 扔了重算 / 卸到 CPU 再搬回来 ——  而本讲只讲了前两个。第三条是拿 PCIe 带宽换显存, 跟本讲那条「拿算力换显存」是同一个形状、不同的货币。
⛔ 至于 MUSTPREFER 的区别,文档写得很直接: MUST_* 表示这条不许被 torch.compile 那类子系统覆盖

⭐ 所以这张表不只是「查词表」,它是个反向检查 ——  如果你理解的判据跟框架提供的选项对不上, 多半是你的判据错了,或者那个默认过期了。

⚠️ 这张表最容易过期,所以用法要说清楚。

上面每一个名字都是 2026 年 9 月在各自主干上当场 grep 到的, 不是凭印象写的 —— 而框架改名是常事。

判据:给了名字,就得给「在哪个文件里能查到」。 —— 否则名字一改,这张表就从「帮忙」变成「误导」,而且不会有任何东西报错

7.1 查得到出处的(放心抄)

数字 / 说法出处
每参数 16 字节(2+2+12)的经典口径;ZeRO 三级的切法 Rajbhandari 等,arXiv 1910.02054
GPT-3 七档峰值学习率与 batch;375M token 线性 warmup; 2,600 亿 token 内余弦降到 10%;weight decay 0.1;上下文 2,048 arXiv 2005.14165 表 2.1 与 §2.3
DeepSeek-V3 的 完整 LR schedule(2K 步 → 10T 恒定 → 4.3T 余弦 → 末段两级);梯度裁剪 1.0;batch 3072→15360; bf16 的一/二阶矩 + fp32 主权重和梯度;重算 RMSNorm 与 MLA 上投影 arXiv 2412.19437 §3.2.3(重算)/ §3.3.3(低精度)/ §4.2(超参)
选择性重算的判据原文;GPT-3 省 70% 付 2.7%;MT-NLG 省 65% 付 1.6% Korthikanti 等,arXiv 2205.05198
warmup 的真实机制(让网络能承受更大的目标学习率; 「样本太少估不准」不是主因;warmup 拉长收益边际) arXiv 2406.09405(NeurIPS 2024)
Adam / AdamW / Adafactor / Lion 的原始定义 1412.6980 / 1711.05101 / 1804.04235 / 2302.06675
Muon 只作用于二维参数、embedding 与输出头仍用 AdamW、 per-step 墙钟比 AdamW 慢 作者 writeup(三条都是原文

7.2 ⚠️ 自己推的(当量级看,别当阈值)

数字推导链它可能错在哪
4.15 TiB / 106.75 GiB 激活 按算子逐项推(1.7 那两张表)。输入只有 V3 的 config + 官方参考实现的 MLA 前向 框架的算子融合程度、MoE 派发是否真的物化九份
attention 占一层算力 82.1% 按 causal 折半口径算。专题一独立算出 81.8%,两边对上了 口径(折没折半);换头维度就变
交叉点 S ≈ 5,734 1.25·S = 7,168。1.25 来自 V3 头维度 192/128、causal 折半 ⛔ 换一条线性层做对照,交点就不同 —— 它不是工程阈值
6ND 低估五到六倍 由上面那个 82.1% 反推(1 ÷ 0.179) 它继承了 82.1% 的全部不确定性;目前没有第三方实测佐证
9.76 TiB 常驻块 · 约 94 条序列的分水岭 671e9 × 16 B ÷ 1024⁴;再除以 106.75 GiB。图脚本里带 assert 16 B 是经典口径 —— V3 自己就不是这么配的(见 3.2)
「按每张卡 80 GiB 算要 125 张」 9.76 TiB ÷ 80 GiB ⛔ 80 GiB 是随手设的换算基准,不是某款具体的卡; 真实部署还要算冗余、通信缓冲、碎片

⭐⭐⭐ 为什么要专门写这一节。

因为这两类数字长得一模一样 —— 都是几位有效数字加一个单位。 而读者会把它们一起抄走

⭐ 判据:一个自己推出来的数,必须连着推导一起给; 否则它冒充的是事实。

📌 顺带一条方法:上面 82.1% 那一行是这一讲最该学的做法 —— 它之所以可信,不是因为算得仔细, 而是因为专题一用完全独立的另一套算法给出了 81.8%。 ⭐ 孤立的估算没法验证;能对上一个外部锚点的估算,一次就把整套公式验了。

← 回 课程总纲 · 前向那半张账在 专题一 · 一个 Token 的一生 · 这张账怎么切,在 专题五 · 并行策略(未上线)
本页由 Courses/tools/topic04-build.py 生成 ——  正文写在那个脚本里