这一讲把训练一步的账补全:反向要付的三倍算力、 那些不能算完就扔的中间激活、以及每参数 16 字节的优化器状态。 补完之后,ZeRO 那三级分法就不用背了 —— 它是从这张账单里长出来的。
专题一跟着一个 token 走完了前向,最后算出一句话:装不进任何一块卡。 ⛔ 可那只是账单的一小半。
真正训练一步,显存里是同时压着四样东西的:权重、梯度、 优化器状态,外加一大堆不能算完就扔的中间激活。 这一讲把这张账单补全。
⭐⭐⭐ 而补完之后会发现一件多数人想不到的事: 最大的那一块,既不是权重,也不是激活 —— 是优化器状态。
⭐ 这一讲按「谁最大」的顺序讲,而不是按「训练流程」的顺序。
流程的顺序是:前向 → 反向 → 更新。 但那个顺序会让最大的那一块最后才出场,而它恰恰是决定一切的那一块。
⭐⭐ 所以这一讲的每一节只回答同一个问题: 这一项有多大,能不能省,省它要拿什么去换。
反向传播不是「再跑一遍」。它要付两样东西: 大约两倍于前向的算力,以及 —— 更要命的 —— 一路累加、不能提前扔掉的中间激活。
⭐ 求导是高中的事。 这里难的是:要求快三千亿个导数,而且每一步都要重求一遍。
⛔ 不过「求导是高中的事」这句话会劝退一批人 —— 所以先花一格,把这一讲要用到的三个词说完。 没有极限、没有公式,三句话,一张图。
⚠️ Ⓐ 的 5.00 → 5.03、Ⓒ 的 ×2 / ×0.5 / ×3 都是编出来的示意数,唯一的作用是让「相除」和「相乘」这两件事看得见。⭐ 脚本里 assert 了两条:总兑换率必须真的是三个乘积,而且链条里要有一级是缩小的 —— 不然「乘起来」会被读成「越乘越大」,而那正是梯度消失/爆炸那一节要讲的反面。
⭐ 这一格刻意不碰极限:严格地说导数是「动的那一点点趋于 0 时的极限」,而这张图画的是差商。⛔ 对本讲够用 —— 而且 §1.1 那个「笨办法」用的**正是差商**,所以这里不严格反而接得更顺。
⭐⭐ 「兑换率」这个说法不是为了好听,它自带两个钩子,两个都是本讲自己的:一串数相乘,从哪头开始乘代价差一万年(§1.5);而单位对不上所以必须再乘一个折算系数(§3.3)。
先看看笨办法长什么样 —— 这一步不能跳过, 不知道笨办法有多笨,就不会觉得反向传播有多神。
⛔ 笨办法(有限差分): 我想知道某个参数对 loss 有多大影响,就把它动一丁点, 整个网络重跑一遍前向,看 loss 变了多少。
那三千亿个参数,就是三千亿次前向。 一次前向按一秒算,跑完要将近一万年 —— 而这还只是一步。
⭐⭐⭐ 所以反向传播真正解决的问题,不是「怎么求导」, 是怎么把三千亿次前向压成一次。
起点很朴素 —— loss 是一个数。 (这一点后面是全部关键,先记着。)
那第一个梯度是什么?拿最常见的交叉熵来说:网络最后吐出一个概率分布 —— 下一个字是「的」的概率 0.3、是「了」的概率 0.2…… 而正确答案是一个只有一格是 1、其余全是 0 的东西。
⭐⭐ 第一个梯度 = 你猜的,减去正确答案。
猜高了的地方是正数,该高没高的地方是负数。 就这么简单。这个差,就是整条反向链的种子。
📌 这是 softmax + 交叉熵这一对的经典结果 —— 它们凑在一起,导数会漂亮地约化成「预测减真值」。 不是所有 loss 都这么好看,但这一对是今天的默认组合。
⚠️ 一个说准了才不会误导的细节:
那个「减」不是对概率求导,是对 logits 求导
—— 也就是 softmax 之前那一层的输出。
⭐ 这个区别有实际后果:正因为约掉的是 softmax 那一步,
这个梯度才不会在 softmax 饱和的地方消失
—— softmax 和交叉熵总是成对出现,一半的原因就在这儿。
📌 「推一下(nudge),推多少跟差多远成正比」「想让一个神经元更亮有三条路」「改权重要按上游亮度成比例,回报最大」「众口难调,只能取平均;只听那张 2 的,网络会把所有图都判成 2」四个装置,取自 3Blue1Brown《What is backpropagation really doing?》官方讲义 —— 已逐条核过原文,图是我们自己重画的,例子换成了本讲一直在用的「下一个字」。
⚠️ 他在「按上游亮度成比例」那里顺带提了赫布理论(neurons that fire together wire together),但原文自己就说这个类比并不严格(未训练的网络并没有在「想」那个答案),所以只记在这儿,不画上图。
⭐ Ⓑ 与 §1.6、Ⓒ 与 §1.8 的那两处接头,原文都没有,是本讲自己的合题 —— 3B1B 讲的是「反向传播在干嘛」,本讲要的是「它为什么这么费显存」。
⛔ 很多人脑子里的画面是「一个梯度值一层一层传下去」。不是这样。
往回传的是一整个张量,形状跟这一层的输出一模一样。 它的含义是:「loss 对我这一层输出的每一个位置,各有多敏感」 —— 这一层输出有多少个数,它就有多少个数。
⭐⭐ 每一层拿到这个上游传来的敏感度之后,只做两件事:
⭐ 这两件事各是一次矩阵乘。前向一次、反向两次 —— 三倍算力就是从这儿来的,不是估的,是数出来的。
⚠️ 「3 倍」是矩阵乘口径的常用近似:只数 matmul,忽略 norm / 激活 / 偏置 / 通信。真实 step 里这些占比不大,但不是零
⛔ 图上第 ② 步那句「用前向存下来的输入」是整个专题的枢纽 —— 激活扔不掉、以及下一节那笔重算交易,全都挂在这一句上
📌 「电路图 + 三个门(分发 / 交换 / 路由)」这个讲法取自 CS231n 的反向传播讲义 —— 图是我们自己重画的,数也是自己挑的。
⭐ 图上每个数都可以自己验:脚本里带 assert —— a 和 b 的梯度都是 4、c 的梯度是 −1,对不上就不让构建。
⭐⭐⭐ 因为偏导数是局部的。
❓ 「偏」这个字,就是全部的关窍。
普通的导数问的是:这个东西变了,结果变多少。
偏导数问的是:其它全部按住不动,只动这一个,结果变多少。
⭐ 而正因为其它都按住了,这个问题就缩到了一个算子身上 —— 它不需要知道外面的世界长什么样。这就是「局部」的意思。
每一个算子只需要知道「我自己是怎么把输入变成输出的」, 就能写出自己的反向规则。 它完全不需要知道前面是什么、后面是什么、整个网络长什么样。
⭐⭐ 这就是为什么自动微分能做成一个通用库。
矩阵乘写一个 backward,softmax 写一个 backward,加法写一个 backward —— 然后框架只干一件事:按顺序倒着把它们串一遍。
⭐ 没有任何人需要手推整个网络的导数。
反过来想才知道这有多可怕: 如果你非要写出「loss 对第一层某个权重」的解析式, 那个式子要穿过 61 层展开 —— 项数是天文数字,写不出来。
⭐⭐ 链式法则真正的意思是:你永远不用把它展开。 你只要把 61 个局部的小导数,按顺序乘起来。
❓ —— 那「乘起来」之后会怎样? 这一问有一个非常出名的答案,而它是链式法则最直接的后果。
⭐ Ⓐ 三条曲线、Ⓑ 那两个端点值、Ⓒ 那个 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 步,换算过来就是一千层。
tools/manim/。
出框层与消失层都是脚本当场算的,并有断言钉住。)数学上两个方向都成立,算出来的结果一模一样。 ⭐ 区别只有一个 —— 你得把整条链走多少遍。
⭐ 这一格是结构,不是数据 —— 图上唯一的量是「三千亿」,那是模型规模,前面已经立过。
📌 「正向模式 / 反向模式」是自动微分的标准术语;反向传播是反向模式用在神经网络上的那个特例。
tools/manim/。)⭐⭐⭐ 所以反向传播成立的全部理由就一句话: 参数有几千亿个,而 loss 只有一个。
多输入、单输出。 ⛔ 如果 loss 不是一个数、而是一百万维的输出,这套就不划算了 —— 那时候反倒该用从前往后。
⭐⭐⭐ 所以 1.1 那个问题的答案是:那三千亿次前向, 压成的是一次反向。
⭐ 这条判据可以直接迁移: 看到任何一个「求一大堆偏导」的问题,先数一下输入多、还是输出多 —— 它决定了你该从哪头开始。
回头看 1.3 那两件事里的第①件: 算自己权重的梯度,要用「前向时存下来的输入」。
⛔⛔ 这就是激活必须留着的根本原因 —— 不是谁设计得不好,是反向的数学本身要求它在场。
前向每算出一个中间结果,都得一直挂在那儿, 等反向走回来的时候用。整条网络走完,它们全都还在。
⭐⭐ 这一节的两笔账到此都立住了: 算力 3×(前向 1 + 反向 2); 显存里多出一整条从头挂到尾的激活。 而下一节要做的,就是拿第一样去换第二样。
⚠️ 台阶画了 12 级只是为了看得清 —— 真实是 61 层,而且每层内部还有若干个中间张量,山坡比图上细密得多
⛔ 山形画成直上直下是简化:真实曲线会因为 MoE 派发、attention 那几个大中间量而有凸起,但「顶点在前向末尾」这个结论不受影响
⚠️ 4.15 TiB 与 106.75 GiB 两个数是自己按算子推的(输入:V3 的 config + 官方参考实现的 MLA 前向),没有第三方背书
📌 图 Ⓑ 那两张卡片的读法: 它们是同一条序列的两种配置,不是两个模型。
⚠️ 两个数都是自己按算子推的(输入是 V3 的 config + 官方参考实现的 MLA 前向),没有第三方背书。 ⭐ 但就算差一倍,结论也不变:这个规模上,装不下。
⭐ 这一小节可以整段跳过。 它是给要自己动手算的人看的 —— 上面那两个数从哪来,全在这儿。
基准单位先立住:一份 hidden 宽的张量
= 131,072 × 7,168 × 2 B = 1.75 GiB。
下面所有数都是它的倍数。(序列 131,072、batch 1、bf16。)
| 留下的张量 | 宽度 | 大小 |
|---|---|---|
| norm 的输入(= 残差入口,重算模式下唯一留的那份) | 7,168 | 1.75 GiB |
| norm 的输出 | 7,168 | 1.75 GiB |
| Q 降维结果(layernorm 前后各一份) | 1,536 | 0.75 GiB |
| KV 降维结果(含 RoPE 那 64 维) | 576 / 512 | 0.27 GiB |
| Q 展开后 | 128 头 × 192 = 24,576 | 6.00 GiB |
| K、V 解压后 | 128 头 × 256 = 32,768 | 8.00 GiB |
| attention 输出 | 128 头 × 128 = 16,384 | 4.00 GiB |
| logsumexp(fp32) | 128 | 0.06 GiB |
| 小计 | 约 22.6 GiB |
| 留下的张量 | 份数 × 宽度 | 大小 |
|---|---|---|
| norm 的输入 / 输出 | 2 × 7,168 | 3.50 GiB |
| 路由分数 | 256 | 0.06 GiB |
| 派发出去的激活 | 9 份 × 7,168 | 15.75 GiB |
| gate / up / SwiGLU 乘积 | 3 × 9 份 × 2,048 | 13.50 GiB |
| 专家输出(合并前) | 9 份 × 7,168 | 15.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。
⚠️ 这张表的三条边界,讲的时候必须说清楚。
⭐ 1.3 只讲到梯度算出来那一刻就停了。 可它离「被用掉」还隔着好几道 —— 而这几道每一道都是推理里不存在的。
⚠️ 图上不含任何量 —— 通信到底多大、能藏掉多少,是 fig-batch 和专题五的事。
⭐ 第 ⑤ 道那句「用完就扔」是本讲另一处的伏笔:正因为它一次性,它才敢用 bf16 存 —— 而一旦要做累积,它就变成累加量,得升回 fp32。
纯数据并行下,每一步要把整份梯度在所有卡之间汇总一遍。 量级很好算:参数量 × 每参数字节数。
⭐⭐⭐ 关键在这儿:这笔通信量跟 batch 完全无关。 —— 参数有多少就传多少,你喂一条序列和喂一千条,传的是同样多的字节。
⭐⭐ 而计算量是随 batch 线性涨的。所以:
batch 越大,这笔通信被摊得越薄。
⭐ 这一下就解释了两件原本看着不相干的事:
📌 「有平行运算时,大小 batch 跑一次的时间差别不大;而一个 epoch 反过来」出自李宏毅2021 年《类神经网络训练不起来怎么办(二):批次与动量》。⛔ 图是我们自己重画的。
⚠️ 图上不含任何实测数字 —— 只画两段的形状。拐点落在哪,取决于模型、卡、并行配置,得自己在目标配置上量。
⭐⭐ 「算」和「传」能叠在一起,是训练里最重要的一类优化 —— 而推理侧没有对应物:它没有反向, 也就没有这条可以边算边传的长尾巴。
⭐ 裁剪那一道为什么藏不了、以及「逐层裁剪」为什么是另一件事, 图下面的落点带写了。这一节只留一句: 「这一步总共迈多大」才是跟 3.3 对话的那个量。
⭐ 一句话:梯度累积改的是时间线,不是总量。
⭐ 这一节的边界要说清楚: 上面只讲了「这一步存在、它有多大、能不能藏起来」。 至于在不同的切法下它究竟是 all-reduce 还是别的形态、量各自差多少 —— 那是专题五整讲的事。
⚠️ 通信量那个「参数量 × 每参数字节」是量级口径: 真实的环形汇总要来回搬,实际搬运量大约是它的两倍; 而 MoE 的专家部分走的又是另一套。当量级看。
这是全课第一次出现真正的取舍: 不是「有没有更好的办法」,而是两样东西只能选一样,你选哪个。
前向每算出一个中间结果就留着,是因为上一节那条 —— 反向要用。重算说:我不留了。
反向要用的时候,我从这一层的入口再往前跑一遍,现算出来。 ⭐ 就像游戏存档 —— 不存整个过程,只存一个存档点, 要用的时候从那儿重打一遍。
⭐ 先把量级说清楚 —— 这才是要记住的东西。
⭐⭐⭐ 这个兑换比例,夸张到不像是个「权衡」。
所以在大模型训练里,它默认就是开着的 —— 值得讨论的从来不是开不开,而是开到哪一档(图 Ⓑ 那三档)。
⚠️ 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 上投影 —— 正是那一档
对每一个中间张量问同一个问题: 把它扔掉、反向时重算回来,每省一个字节要付多少次浮点运算?
比值越小越该扔。就这一个数,整张排序表自己就出来了。
⭐⭐ 对线性层,这个数有闭式解,而且漂亮得出乎意料:
一个 [S,k] × [k,n] 的矩阵乘 ——
重算代价 2·S·k·n FLOPs,产出张量 S·n·2 字节(bf16)。
一除,每字节代价 = k。
⭐⭐⭐ 重算一个线性层的输出,每字节要付的 FLOPs
恰好等于它的输入宽度。
S 和 n 全约掉了 ——
跟序列长度无关,跟输出多宽也无关。
⭐ 所以「哪个线性层最该重算」这个问题,答案是看谁的输入最窄 —— 不用算,扫一眼 config 就知道。
attention 不是线性层:代价随 S² 涨,产出只随 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,734 是个粗略交叉点,不是工程阈值。 它依赖 causal 折半的口径、依赖 V3 的头维度、依赖你拿哪个线性层做对照。 ⭐ 要记的是那句「两条线斜率不同所以必然相交」,不是这个数。
把 V3 在 128K 上排一遍 —— 一层 MoE 块,按「每 GiB 要付多少 TFLOP」升序:
| 张量 | 省显存 | 重算代价 | 每 GiB 付 | 结论 |
|---|---|---|---|---|
| MoE 派发(复制成 9 份) | 15.75 GiB | ~0 | 0 | 白捡 |
| RMSNorm 输出(每层 2 个) | 3.50 GiB | 0.01 TFLOP | ~0 | 白捡 |
| SwiGLU 乘积 9 份(逐元素) | 4.50 GiB | 0.01 TFLOP | ~0 | 白捡 |
| K/V 解压(输入宽 512) | 8.00 GiB | 4.40 TFLOP | 0.55 | 划算 |
| Q 展开(输入宽 1,536) | 6.00 GiB | 9.90 TFLOP | 1.65 | 划算 |
| 专家输出 9 份(输入宽 2,048) | 15.75 GiB | 34.63 TFLOP | 2.20 | 划算 |
| gate / up / 路由 / 降维(输入宽 7,168) | 9.14 GiB | 70.35 TFLOP | 7.70 | 边际 |
| attention 输出 | 4.00 GiB | 703.69 TFLOP | 175.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。
「选择性重算」这个概念出自 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 上按我们这张表,它恰恰是最不该重算的那一个 —— 但它还在那里当默认。
⭐ 更值得看的是名单上的其余几项:
layernorm、moe_act、mla_up_proj
—— 对照 2.3 那张表:RMSNorm 输出、SwiGLU 乘积、K/V 解压与 Q 展开。
正好是最便宜的那几行,一个不多一个不少。
我们那条闭式解排出来的名单,和框架实际提供的选项对上了。
📌 megatron/core/transformer/transformer_config.py
的 recompute_modules 字段(2026-09 主干)。
⛔ 判据:说「现代框架都不这么干了」之前,去 grep 一下那个框架。
这类断言听起来像常识,而它正好是最容易过期的一类。
⭐ 前面算的是两本能算的账。可重算还有第三样代价, 它不在任何一本账上 —— 而且它不报错。
重算的前提是:把同一段前向再跑一遍,跑出来的要一模一样。 —— 有三种情况会让它不一样。
⭐ ① 随机数 —— 这一条框架替你挡住了。
前向里有 dropout。如果重算时抽到的是另一套掩码, 那反向用的就不是前向那张网络了 —— 梯度直接是错的,而且没有任何报错。
⭐⭐ PyTorch 的做法很干脆:把前向那一刻的随机数状态存下来,
重算前先恢复。
—— torch.utils.checkpoint 的
preserve_rng_state,默认就是 True。
⭐ 而这一条正好解释了 6.4 那个清单里的第三项: checkpoint 里除了权重和优化器状态,还得存随机数状态 —— 原因是同一个:随机数是训练结果的一部分,不是「运行时的临时东西」。
⛔ ② 副作用 —— 这一条没人替你挡。
如果那段前向除了返回结果,还顺手改了别的东西 —— 更新了一个滑动统计量、累加了一个计数器、写了一行日志 —— 重算会把它再改一遍。
⛔ 框架不会替你回滚这些,因为它根本不知道你改了什么。 —— 判据很简单:被重算的那一段,最好是个纯函数。
⚠️ ③ 逐比特不一致 —— 通常无害,但它有个具体的受害者。
同一段计算跑两遍,结果可能差在最后几位 —— 归约的顺序不同、挑的 kernel 不同,都会这样。
⭐ 对训练本身,这个量级的差别基本没影响。 ⛔ 但它让「逐比特可复现」变难 —— 而那正是 6.7 里 PaLM 做到的那件事, 也是 6.2 那个「回滚 + 跳数据」的救火办法所依赖的前提。
⭐⭐ 一句话:重算把「算力 ↔ 显存」这笔交易谈成了, 但它同时悄悄引入了一个「结果要可重现」的要求。 —— 前两本账会告诉你划不划算,这一条不会:它只在出事的时候现身。
⛔⛔ 同一个模型、同一个开关,换一个规模, 收益可能从正的变成负的。
原因不神秘:重算改变的是计算与访存的配比, 而这个配比在不同并行配置、不同芯片数下本来就不同。
⭐⭐⭐ 由此推出一条通用规则(它不只对 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 看起来那么低: 分母里有一大块被重算吃了,它做了功,但不算进「有效算力」。
前面两节都在跟激活较劲。可把账摊开一看 —— 那个从头到尾一言不发的角色,才是最大的一块。
❓ 先花两句话把前提说清楚,不然下面那张账单看不懂。
训练的时候,同一个参数在显存里同时存着两份:
① 一份省地方的(bf16,2 字节)—— 前向和反向都拿它算;
② 一份精确的(fp32,4 字节)—— 专门用来记账,每一步的更新加在它身上。
⭐ 这套做法就叫「混合精度」。 —— 为什么非得存两份,下一小节(3.2)整节都在回答。
混合精度训练下,每一个参数身上挂着五样东西 —— 而其中只有第一样是「模型本身」。
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 | 主权重 | 合计 | |
|---|---|---|---|---|---|---|
| 纯 fp32 | 4 | 4 | 4 | 4 | — | 16 B |
| 混合精度 | 2 | 2 | 4 | 4 | 4 | 16 B |
⭐ 看那一列:权重和梯度各省了 2 字节, 可换来的是多出一份 4 字节的主权重 —— 2 + 2 正好被 4 吃掉。
⭐⭐ 那混合精度到底换来了什么?—— 两样,都不是常驻显存。
⛔ 判据:说「省显存」之前,先问省的是哪一块。 常驻和激活是两本账,一句「省一半」把它们混在一起, 结论就会在你最需要它的时候是错的(比如估一张卡装不装得下)。
图上那五格,为什么有三格是 fp32、两格是 bf16? ⭐ 不用一格一格记 —— 根子上只有一个画面。
📌 Ⓑ 里每一个数都是脚本用一个真的 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;会被衰减掉的、用完就扔的,低精度就够。
⚠️ 「必须 fp32」这句话要留个口子 —— 有人真的把它省掉了。
主权重之所以要 fp32,是因为四舍五入会把小更新整个抹掉。 ⭐ 而如果把它换成随机舍入 —— 按余数的大小决定进位的概率 —— 那么被抹掉的部分,在多步之后期望上会被找回来。
⭐ 那么判据其实比「累不累加」还要细一层: 真正致命的不是累加,是累加 + 每次都朝同一个方向抹零 —— 把这个偏差去掉,bf16 主权重就是可行的。
⛔ 这不是推荐做法 —— 它要额外的支持,而 fp32 主权重便宜又省心。 放在这里是因为:知道一条规则「为什么成立」,才知道它什么时候可以不成立。
📌 出处:DeepSeek-V3 技术报告,arXiv 2412.19437 §3.3.3(原文分节是「低精度存储与通信」,不是 §3.2.3 那个重算小节)。 ⭐ 原文这句话几乎是逐条印证上面那张清单: 「用 BF16 而非 FP32 追踪 AdamW 的一阶和二阶矩,未观察到性能下降; 但主权重、以及用于 batch 累积的梯度,仍保留 FP32。」 —— 连「梯度因为要累积所以升回 fp32」这一条都写在里面。 ⛔ 顺带说明这张 16 字节的表不是定律 —— V3 这套配下来,每参数的账跟经典口径已经不是一个数了。
⭐ 这条判据出了这一讲还能用:FP8 训练为什么难, 问的就是同一句 —— 精度往下压的时候,哪些量能压、哪些必须留住, 看的还是它累不累加。那是另一课的事。
📌 Ⓑ 末尾那句「最快是相对于你怎么量一步」取自苏剑林《为什么我们偏爱各向同性?基于最速下降的理解》—— 原话是「梯度反方向是损失下降最快的方向,但这结论是有前提的,最关键的前提是它选取的度量是欧氏范数,如果换一个范数,那么最速方向也就变了」。⭐ 这句前提几乎所有教程都略过,而略过它,Muon 就只能被当成「又一个新优化器」。
📌 「球滚下山」「步长 ∝ 斜率所以不会冲过头」「落在哪个谷取决于起点」「高维时看那一列数的正负与相对大小」四个讲法,取自 3Blue1Brown《Gradient descent, how neural networks learn》官方讲义 —— 已逐条核过原文,图是我们自己重画的。
⭐ Ⓒ 那一整列负梯度,正是反向传播一遍算出来的那一份 —— 所以这一讲的第一节和第三节,接头就在这儿。
tools/manim/。)⛔ 那张图末尾留了一句不太好听的话:落在哪个谷,取决于你从哪儿出发。 一维的图上这件事看着很吓人 —— 满眼都是坑,随便掉哪个都出不来。 下面这张就是来拆这个画面的。
📌 「critical point 分 local minima 与 saddle point」「判据是 Hessian 特征值全正与否」「沿负特征值方向可继续下降,但实践中很少这么做」「实测 never reach a real local minima」四条,取自李宏毅2021《类神经网络训练不起来怎么办(一):局部最小值与鞍点》投影片 —— 已逐条核过原文,图是我们自己重画的。
⭐ 他课上还引了《三体Ⅲ·死神永生》里那个能在高维取物的魔法师作比,投影片原话是「从三维空间看它是封死的,在更高维度里并不是」。⛔ 小说情节我们没核,所以只转述这一句,不展开 —— 图上的主比喻用我们自己的垭口。
⚠️ Ⓑ 刻意不给概率:「每个方向朝上朝下各半,所以 d 维全朝上的概率是 2 的 −d 次方」这句话听起来很顺,但特征值的符号并不独立,那个数是编的。只讲定义层面必然成立的部分。
⭐ 原片收尾那句也在本讲别处出现:更小的 batch 和动量都有助于逃离临界点 —— 噪声和惯性,正好是 `fig-batch` 与 Adam 那两格在讲的东西。
⛔ 可这张图是切了两刀给你看剖面 —— 而「一个方向上翘、另一个方向下沉」本来是个三维的形状。 ⭐ 下面让它转起来。
tools/manim/。)⛔ 上面那张图欠了半句:卡不住,可垭口附近坡太平,它会在那儿磨很久。 —— 下面这张是那半句的第一个答案。
⭐ Ⓐ 两条轨迹、Ⓑ 三个箭头的长度、Ⓒ 那排柱子,全是脚本算的。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」。⛔ 图是我们自己重画的,曲线和数都是自己跑的。
tools/manim/。)⭐ 这三张图把方向说完了 —— 顺着负梯度走 loss 就在降,维度一多基本卡不死,平地上还能靠惯性滑过去。 这一节只剩一个问题:那一步,到底该迈多大?
最朴素的那一版只有一行:新权重 = 老权重 − 学习率 × 梯度。 ⭐ 但你仔细想 —— 这一行凭什么成立?
⭐⭐ 两边的单位对不上。
梯度的含义是「这个参数变一点,loss 变多少」—— 单位是 loss ÷ 参数。 可你要的是「参数该挪多少」—— 单位是参数本身。
⛔ 所以中间必须乘一个东西把它折过来 —— 那就是学习率。 反推一下它的单位:参数² ÷ loss。
⭐⭐⭐ 所以学习率不是一个纯数字。 它跟你的 loss 有多大、你的权重有多大,全绑在一起 —— 这就是为什么换个模型、换个 batch,学习率就得重调。 它压根不是个通用常数。
全网络共用一个学习率,可不同参数的梯度尺度能差好几个数量级。 —— 下面两条轨迹用的是同一个谷、同一套规则,只改了学习率。
📌 「同一个学习率,大了第二次更新就飞出地图之外、小了更新一百次还走不到谷底,所以不同参数应该有不同的 learning rate」这个两难取自 李宏毅《Training Tip》投影片(2025 秋 GenAI-ML 课程)。
⛔ 但他那两个具体数字我们没照抄 —— 抄一个别人挑出来的学习率,等于把别人的地形当成自己的。这里自己搭了谷、自己跑了两遍,每条轨迹都有 assert 盯着。
⚠️ 画面用的陡峭比是 25,不是正文说的「几个数量级」。第一版按 1,000 画,物理没错但画面废了:平方向几十步只挪几十像素,锯齿退化成一根竖线。⭐ 判据:要让人「看见」某个动态,参数按「看得见」选,真实量级交给公式外推。Ⓑ 第三行就是那条外推。
⭐ Ⓑ 那两个步数是闭式算的:平方向每步收缩 (1 − 1.9 ÷ 陡峭比),走到一半就是 log½ ÷ log(那个收缩率)。脚本里拿真跑一遍的结果对过账,差不超过 1 步。
tools/manim/。
门槛、±5% 的收敛/发散、以及左端那「三分之一」都是脚本当场跑出来并有断言钉住的。)别用梯度的大小,只用梯度的方向 —— 把每个参数自己的尺度先除掉。
❓ 可凭什么是「除」? —— 这个问题有一个标准答案,而且小到能当场验。
📌 「最优步长 = |一阶导| ÷ 二阶导」与「用一阶导去估二阶导」取自李宏毅(台大)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 的尺度、跟梯度的大小都脱钩了, 只剩下一个由 β₁ 定死的系数。
原来 Adam 把权重衰减混进梯度里一起算, 于是它也被那个自适应分母除了一遍 —— 结果梯度大的参数,正则反而弱。
⭐ AdamW 就是把它从梯度里拿出来,直接加到更新那一步上。 论文标题里那个 decoupled,说的就是这件事 —— 整篇论文主要就干了这一件事。
既然 Adam 每步挪的是 0.2 × 学习率,
那别的优化器只要也把自己的 update 对齐到 0.2
—— Adam 那套调好的学习率和权重衰减就能直接搬过去用。
⭐⭐ 这不是纸上谈兵:Kimi K2 从 Adam 迁到 Muon,用的就是这一招 —— 把 Muon 的 update RMS 统一成 0.2,其余超参照抄。
📌 同一篇(苏剑林,科学空间)。 ⭐ 这是「量纲」这条线最实用的一个落点: 先把某个量做成常数,常数就能当接口用。
⭐ 讲法来源:苏剑林《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 作用在对角阵上会退化成逐元素取 sign —— 也就是说,「向量做法」本来就是「矩阵做法」的一个特例。 ⛔ 我们以前把它们当两种东西讲,那是把特例和一般情形讲反了。
梯度算完,到真正更新之间,还要过: 多卡 all-reduce 求平均 → 按全局范数裁剪一次(防止偶发大梯度炸掉) → 按 schedule 取当前学习率 → 权重衰减 → 才是那一步更新。
⭐⭐⭐ 这一整条线可以用一句话收: 梯度只告诉你往哪走,优化器决定走多远。
而几十年的演化,全都在回答同一个问题 —— 这个「多远」该由谁说了算:
⭐ 说到底就这么一条线索。
⭐ 优化器的选择是一个显存决策,不只是收敛速度的决策。 —— 这一点常被忽略,而它恰恰是这一讲要立的那条判据。
⛔ Muon 那一份不是白省的 —— 三条限制都是作者自己写明的。
⚠️ 这跟 fig-muon Ⓓ 那句「它塞进了一段本来就空着的时间」不矛盾,
但两句话一定要一起读:
「更慢」说的是方向(多做的那几轮矩阵乘不是免费的),
「塞进空窗」说的是幅度(作者按 FLOP 算下来低于 1%)。
⛔ 而幅度是你的配置说了算的 —— 那几轮是一串互相依赖的小矩阵乘,
FLOP 少不代表墙钟时间短。真要用,在目标配置上自己量一遍。
⭐⭐ 而它的证据方式值得单独一提:Muon 在 NanoGPT 那个刷速度的公开竞赛里 把记录提了 35%,此后十二次破纪录、七个不同的人,全都还在用它。 —— 要是有人能把 AdamW 调到一样好,换回去就能破纪录,可没人换。
⭐ 顺带一句跟上一讲连起来:DeepSeek-V4 就是用 Muon 训的 —— 论文摘要里列的三大升级之一。
⛔ 在比账单之前,先把这几个名字串成一条线 —— 不然它们只是几个并列的选项,只能靠背。
📌 这一格每一环都出自 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 字节里的某一格,但动的格子不一样 —— 下面的字节数按 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 替换」的东西,先去摘要里数一数它到底改了几处 —— 「两行代码」说的是调用方改两行,不是它内部只改了一处。
3.3 那段讲的是 Muon 的想法。 可想法好用不等于能直接上规模 —— 2025 年有一篇专门做这件事的论文,识别出两条必需的补丁:
⭐ 补上这两样之后,它才能「开箱即用」地跑大规模训练。 论文的规模律实验报告:在算力最优的设定下,Muon 的计算效率约为 AdamW 的两倍。
📌 arXiv 2502.16982(Moonlight,3B/16B MoE,5.7T token)。 ⚠️ 「两倍计算效率」是论文自己的规模律结论,不是我们的实测 —— 而且回想 2.6 那条判据:这类结论换规模要重验。 ⛔ 至今仍未核实的一条:Muon 下主权重是否仍需 fp32。
⭐ 在钻进「曲线怎么画」之前,先给一把随时能用的尺子 —— 它不告诉你最优解,但能告诉你「有没有离谱」。
📌 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」是二分搜出来的。
⭐ 整条学习率曲线只有三段:升上去、稳住、降下来。
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 倍,叠在一起就看不出形状了)—— 七档规模的峰值对照表写在正文里
⭐ 形状就这样。下面把三段拆开 —— 每一段都有公开配置可以抄,不用猜。
⛔ 大部分教程说: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 底下,梯度消失不再意味着「学不动」 —— 只要梯度还大于随机误差,参数照样拿到常数量级的更新。
⛔⛔ 可这恰恰是麻烦的开始。 原文的推理链是这样的(前提:模型确实有明显的梯度消失):
⭐⭐⭐ 这条推理给了一个可以拿去对日志的现象 —— 比任何机制解释都好用:
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 全系列(同一套数据、同一套训练流程,只有规模在变):
⭐⭐⭐ 规律肉眼可见:模型越大,学习率越小、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 越大、学习率就越能调大」这句,还有一层很多人不知道的天花板。
⭐ 先把常见的那两条摆出来:
η* ≈ η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 倍」。
⭐⭐⭐ 真正的答案是:这两个数本来就不该放在一起比。
⭐⭐ 所以这一格的教训,跟 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 设的是 1.0。靠调裁剪阈值来压住一个太大的学习率, 是在掩盖问题不是在解决问题。
⛔⛔ 最后一条,也是最容易被忽略的: 单报一个学习率数字,是没有意义的。
它必须连着 batch、schedule 形状、总 token 数 一起报 —— 换掉其中任何一个,这个数就不再适用。 ⭐ 上面那两套配置之所以能抄,正是因为它们四样都写全了。
学习率讲完了,可 Adam 还有两个数没讲:β₁ 和 β₂。 ⭐ 它们通常被当成「用默认值就行」的东西 —— 但它们的含义其实特别清楚。
⭐⭐ 一句话:β 决定「老的贡献能活多久」。
滑动平均每步乘一个 β 再掺新值,所以有效窗口大约是 1/(1−β) 步。
⭐ 所以 β₂ 调大 = 分母更平稳但反应更慢;调小 = 跟得紧但噪声大。 —— 这不是玄学,就是一个「记多久」的旋钮。
它用的是 β₂ = 1 − k−0.8,k 是当前步数。
代几个数进去看看它在干嘛:
⭐⭐ 也就是说,窗口是跟着训练一起长的(大约按 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 用的其实是 Adafactor,但关掉了分解 —— 论文说这等价于「带参数缩放的 Adam」: 按参数矩阵的均方根去缩放学习率。
⭐⭐⭐ 而论文自己点破了这跟 3.5 那张表的关系: 「这么做的效果,类似于 GPT-3 那样手工把学习率随规模调小」。
⭐ 换句话说:3.5 那张七档表是手动解法,参数缩放是自动解法。 而且论文说自动那版还多一个好处 —— 尺度本来就不同的那些参数(embedding、layer norm 的缩放) 不会被按同一个比例一起压下去。
⭐ 这一节能用 3.1 那张账单直接推出来, 不用查任何文档。
那 16 字节里,哪些必须进 checkpoint?一格一格问「丢了能不能重建」:
⭐⭐⭐ 所以 checkpoint 的大小 = 那 12 字节 × 参数量。 对 671B:约 7.3 TiB 一份。
⚠️ 但别把这个数安到 V3 头上 —— 3.2 刚说过它的一二阶矩是 bf16,那就是 4+2+2 = 8 字节, 一份 4.9 TiB,比经典口径小三分之一。 —— 省状态的那些做法,省的同时也在省 checkpoint。
⭐ 这一下解释了两件工程上的事:
⛔⛔ 但只存这 12 字节还不够 —— 还有三样,漏一样都会让「恢复」变成「重来」。
⭐⭐ 而 §6.2 那个「回滚 + 跳数据」的救火办法,前提正是 ①②。 —— 数据位置说不清,你连「跳掉哪几批」都没法讲。
⭐ 顺带一个真实的省法:优化器那部分不一定要留在加速器上。 DeepSeek-V3 就把权重的指数滑动平均放在 CPU 内存里、每步异步更新 —— 既不占显存也不占训练时间,用途是提前估计「学习率衰减之后会有多好」。 📌 arXiv 2412.19437 §3.2.3。
⭐ 前面整讲算的都是从零开始训。 —— 可绝大多数人这辈子跑的训练,是拿别人训好的模型接着调。
⭐⭐ 好消息是:这件事不需要任何新概念, 它就是上面那张 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 亿参数代进去,数字很夸张。
⭐⭐ 相差 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 那三级,正是照着这个顺序来的。
把前面三节的东西摆在一起,按占多少排个序 —— 这个顺序本身就是这一节的全部内容。
⭐ ZeRO 的分级不用背。 把上面那个顺序倒过来读一遍,你就把它推出来了。
📌 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)。
⭐⭐ 切的顺序只有一条原则:谁最大、谁最少被用到,就先切谁。
⭐⭐⭐ 所以顺序不是随意的,也不是历史巧合 —— 它是「大小」和「被用到的频率」这两个量排出来的。
📌 出处:Rajbhandari 等,ZeRO: Memory Optimizations Toward Training Trillion Parameter Models,arXiv 1910.02054。 ⭐ 上一节那个「每参数 16 字节」也出自它 —— 同一篇论文既给了账单,也给了切法。
把 ZeRO 之外那几个常听见的名词摆回这张账单上,它们各自的位置立刻就清楚了:
⭐⭐ 所以这几个名词不是「几种优化技巧」, 是同一张账单上的四个不同栏目各自的对策。 —— 而专题五整讲要做的,就是把「切」这件事本身摊开算。
把前向(专题一)+ 反向 + 更新合成一张表。 ⭐ 而其中最有用的一问不是「总共多少」,是 「峰值出现在哪一个时刻」。
⭐ Ⓑ 三行都是当场算的: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 条」也是量级,不是阈值
⛔ 全图算的是全局总量,没有任何并行切分 —— 怎么把它摊到多少张卡上,是专题五整讲的事
tools/manim/。)⭐ 前面四节各算各的,现在把它们摆到一起。 这一节全部按全局总量算 —— 不分卡,怎么切是下一讲的事。
📌 常驻那一块的算式,短到可以口算:
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 怎么凑出来的,显存不在乎。
⛔ 但算力不是这样。
同样 1,230 万 token,拆成长序列和拆成短序列, 要算的浮点数差很多 —— 因为 attention 那部分随 S² 涨, 而它不在显存账上、却在算力账上。
⭐⭐⭐ 所以序列长度这个旋钮,在两本账上的行为不一样: 显存那本只看总量,算力那本还要看你怎么切。 —— 下一节那个「6ND 严重低估」,讲的就是这件事。
⚠️ 前提说清楚:以上都建立在用了 FlashAttention 这一条上。换成把分数矩阵写进显存的朴素实现, S² 就回到显存账里,上面整段都不成立。
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 只占百分之一点几,
漏掉它完全无所谓。又是同一件事:不是公式错了,是前提变了。
⛔ 第一节那张山形图给过一个答案:峰值在前向末尾。 那句话只在谈激活的时候成立。
把四项一起画到时间轴上(就是上面那张图), 会看到一件当时看不出来的事:不同的项,峰在不同的时刻。
⭐⭐⭐ 推广出去,这是一条到处能用的判据:
把几条形状不同的曲线叠起来之后, 总和的峰,未必落在任何一条单独的峰上。 —— 所以「什么时候最挤」这个问题, 必须连着「挤的是哪一项」一起问。
📌 工程上的直接后果: 你去查 OOM,光知道「峰值 xx GiB」没用 —— 要知道是哪一项在那一刻最大,才知道该去动哪个开关。 激活大就开重算 / 上 CP;优化器状态大就上 ZeRO —— 这两条路走反了,一点用都没有。
⭐ 前面那些 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 B | 0.23 GiB |
| bf16 梯度 | 2 B | 0.23 GiB |
| 优化器状态(fp32 主权重 + m + v) | 12 B | 1.40 GiB |
| 小计 | 16 B | 1.86 GiB |
⭐ 第一个可以自己验的结论: 这个模型的权重只有 0.23 GiB,而它要占掉 1.86 GiB —— 八倍。多出来的全是训练才要的。
一层里每个 token 留下的东西,一项一项数(宽度 768):
| 留下的张量 | 宽度 |
|---|---|
| 两个 norm 的输出 | 2 × 768 = 1,536 |
| Q / K / V | 3 × 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 GiB | 1.86 GiB |
| 激活 | 2.11 GiB | 0.32 GiB |
| 合计 | 3.97 GiB | 2.18 GiB |
| 算力 | 3× | 4×(+33%) |
⭐⭐ 这张小表,把这一讲的三条主线全串上了。
⭐⭐⭐ 而最后这一句才是这一小节存在的理由: 换了 5,000 倍的规模,那些比例基本没变,变的只是绝对值。 —— 所以这一讲教的是比例,不是那些数。
⚠️ 三条边界,跟 1.2 那张表同一套。
专题一算完前向,结论是「装不下」。 这一讲把账补全之后,那句话变得具体了 —— 不是「装不下」,是「四项各自装不下,而且各有各的治法」。
⭐⭐ 四项 + 四种对策,一一对上:
⭐⭐⭐ 你会注意到四条里有三条写着同一个字:切。 —— 那就是专题五整讲要做的事:切这件事本身,也是要算账的。
前面整整一讲都在算账。可账算得再清楚,也回答不了训练现场最常见的那一句: 「loss 突然飞上去了,怎么办?」
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
⛔⛔ PaLM 训 540B 的时候,loss 飞了大约 20 次。
而且 —— 梯度裁剪是开着的。 这些 spike 出现在极不规则的时刻,有时候训到很晚才来; 而更小的模型上根本没观察到。
📌 arXiv 2204.02311 §5.1。 「大约 20 次」「尽管梯度裁剪是开着的」都是论文原话。
📌 arXiv 2204.02311(PaLM)§5.1 —— 「大约 20 次」「尽管梯度裁剪是开着的」「更小的模型上没有观察到」均出自该节原文。
⚠️ 曲线为示意:尖峰的位置按「看不出规律」这一条挑出来,高度与宽度论文未给数,故纵轴不设刻度。没有刻度的曲线是示意,有刻度的曲线才是数据。
⭐ 先记住这个规模效应:小模型上训得好好的,不代表大模型上不会飞。 —— 这跟 2.6 那条「不能跨规模照抄」是同一件事, 只不过这次翻车的不是收益,是训练本身。
第一反应都一样:肯定是那批数据有问题。 PaLM 团队去验了这件事,做法很干净:
把 spike 前后那几批数据单独拎出来,从另一个更早的 checkpoint 重新喂一遍 —— 结果不飞。
⭐⭐⭐ 所以 spike 不是「坏数据」造成的, 是这批数据 和 当时那个参数状态 撞在一起才出的事。
换个时刻喂同样的数据,什么都不会发生。
⭐ 这一下就解释了他们那个看起来很土的办法为什么管用: 回滚到 spike 之前约 100 步的 checkpoint,跳掉那 200–500 批数据,继续跑 —— 之后同一个点就不再飞了。
⛔ 注意这不是「修好了」,是「绕过去了」。 论文自己也写着:由于训练成本太高,他们没能找到一个有原则的缓解办法。 ⭐ 这句坦白值得原样转述 —— 这一行至今仍是开放问题。
「让它别飞」有一堆办法:学习率调到极小、裁剪收到极紧…… 都能稳,代价是模型变差。
⭐ ST-MoE 那篇论文把这件事量出来了 —— 同一个配置换不同随机种子反复跑,看几次能跑完、以及跑完的质量:
| 做法 | 稳定性 | 质量(越大越好) | 结论 |
|---|---|---|---|
| 基线 | 4 / 6 | −1.755 | 六次里两次训崩 |
| 收紧 update clipping(0.1) | 3 / 3 | −4.206 | ⛔ 稳了,可质量被打穿了 |
| router z-loss | 3 / 3 | −1.741 | ⭐ 稳了,质量还略好一点 |
⭐⭐⭐ 中间那一行是这张表的全部价值: 「稳定 3/3」看着完美,可它是拿质量换来的。 —— 评价一个稳定性手段,必须同时看这两栏。
⚠️ 两个口径要说清楚,不然这张表会被读得太重。
⭐ 这不是黑点,是好的实验设计 —— 研究一个罕见故障,你必须先把它变得不罕见。 但读数的时候要记住:这个概率是被调高过的。
📌 arXiv 2202.08906(ST-MoE)Table 4。 论文自己的小标题就是「很多方法能稳住稀疏模型,但代价是质量变差」。
做法出人意料地简单:在总 loss 上加一项,惩罚 softmax 归一化因子的对数的平方。
—— 它逼着那个 log Z 待在 0 附近,也就是不让 logits 长得太大。
⭐ 两个地方都能加,而且都有人用:
10⁻⁴ · log²Z,
论文说它提高了训练稳定性。⭐ 为什么 z-loss 特别划算:ST-MoE 的解释是 —— logits 后面接的是指数函数,而指数对输入误差极度敏感。 把输入范围压小,等于直接降低了这一步的数值风险, 而模型并不真的需要那么大的动态范围。
按全局范数裁:把所有参数的梯度当成一个大向量, 范数超过阈值就整体等比例缩回去。常见设 1.0(PaLM、DeepSeek-V3 都是 1.0)。
⛔ 但 6.1 那句话要记住:PaLM 开着裁剪,照样飞了 20 次。 —— 裁剪拦得住单步的爆炸,拦不住「参数状态已经走到了一个坏地方」。
原始 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 不仅不能缓解梯度消失, 它还是梯度消失的『元凶』之一。」
📌 苏剑林《模型优化漫谈: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 没有这条线,所以那一条对它也就不适用。
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,但它决定了你从一个多好的起点出发。
❓ 「bf16 是不是也要 loss scaling?」—— 不用。
loss scaling 是 fp16 时代的必需品:fp16 的指数位只有 5 位,
小梯度会直接下溢成 0,所以要先把 loss 放大若干倍再算、更新前再缩回去。
⭐ 而 bf16 的指数位跟 fp32 一样是 8 位 ——
动态范围没缩,缩的是尾数精度。所以它不下溢,也就不需要 loss scaling。
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。
这一讲报了几十个数字。但真正想留下的不是那些数。
⭐ 数字会过期,这四条不会。
⛔ 这一节刻意一个我们自己的实测都没有 —— 四条全部来自公开论文并逐条核过原文。 理由写在 §七:这一讲已经有太多「自己推的」数了, 稳定性这一块不该再加。
⭐⭐ 这一讲报了几十个数字。它们不是同一种东西。 有的能查到论文原文,有的是我们自己按算子推的 —— 后者可能差一倍,前者不会。
⭐ 这一讲讲了一堆概念。可回去打开框架,该搜哪个词?
| 这一讲里的说法 | 在框架里叫什么 | 核的地方 |
|---|---|---|
| 全量重算(2.2) | --recompute-granularity full |
Megatron-LMtraining/arguments.pycore/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(...) |
PyTorchtorch/utils/checkpoint.py |
| 选择性重算,按张量挑 | create_selective_checkpoint_contexts + CheckpointPolicy⭐ 这就是 2.3 那条判据在 PyTorch 里的落点: 你自己写规则决定哪个算子存、哪个重算 | |
| 重算(JAX 那边) | jax.checkpoint(别名 jax.remat) |
JAXjax/_src/ad_checkpoint.py |
| 选择性重算的现成策略 | jax.checkpoint_policies.dots_with_no_batch_dims_saveable⭐ 名字直接说出了判据:「矩阵乘的结果留着」 —— 正是 2.3 那张表最贵的那几行 | |
| 梯度累积(1.8 第 ④ 道) | ddp.no_sync() —— 上下文管理器进去之后梯度只在本地累加、不跨卡同步; 退出后的第一次 forward-backward 才同步一次(原文档语) |
PyTorchtorch/nn/parallel/distributed.py |
| 「跨卡汇总能藏进计算里」(1.8 第 ② 道) | bucket_cap_mb(默认 25 MiB)⭐ 梯度攒够一桶就发一次 —— 所以前面的层还在算,后面的桶已经在路上了。 那一道「✅ 能藏」,藏在这个桶里 | |
| ZeRO 的三级(4.2) | 配置里的 zero_optimization.stage = 1 / 2 / 3 |
DeepSpeedruntime/zero/config.py |
| 8-bit 优化器(3.4) | bitsandbytes.optim.Adam8bit / AdamW8bit |
bitsandbytesoptim/__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 带宽换显存,
跟本讲那条「拿算力换显存」是同一个形状、不同的货币。
⛔ 至于 MUST 和 PREFER 的区别,文档写得很直接:
MUST_* 表示这条不许被 torch.compile 那类子系统覆盖。
⭐ 所以这张表不只是「查词表」,它是个反向检查 —— 如果你理解的判据跟框架提供的选项对不上, 多半是你的判据错了,或者那个默认过期了。
⚠️ 这张表最容易过期,所以用法要说清楚。
上面每一个名字都是 2026 年 9 月在各自主干上当场 grep 到的, 不是凭印象写的 —— 而框架改名是常事。
⛔ 判据:给了名字,就得给「在哪个文件里能查到」。 —— 否则名字一改,这张表就从「帮忙」变成「误导」,而且不会有任何东西报错。
| 数字 / 说法 | 出处 |
|---|---|
| 每参数 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(三条都是原文) |
| 数字 | 推导链 | 它可能错在哪 |
|---|---|---|
| 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 生成 ——
正文写在那个脚本里。