Kimi K3 模型结构(9):显存,把 2.8T 塞进 GPU
系统篇第二篇。先用 config.json 的形状列出训练时每张卡上住着什么,再把 §5.2.2 的六个技巧逐个放回账本:每一个都是把某一行搬走,付一种货币。推导 MoE 反向为什么可以不存 expert 输出,P2P Muon 比全量 all-gather 省在哪里。
目录8 节
对应论文 §5.2.2 "Memory-Efficient Training"。第 8 篇把时间账算平了,代价是显存:虚拟段让 rank 0 攒着更多份激活,MoonEP 让每个 rank 要给一倍的 expert 权重留位置。这一篇从一张账本开始,把 §5.2.2 的每个技巧放回账本里看它到底搬走了哪一行。
论文这一节和上一节一样没有任何绝对数字。下面账本里的形状全部来自 config.json,并行度是滑块。
账本:每张卡上住着什么
训练时一张卡上的显存分五类:权重、梯度、优化器状态、激活、通信缓冲。通用的记账方法在显存账本那篇,这里直接按 K3 的形状填。先把参数量按 config 算清楚,第 0 篇算过总数,这里要按「expert 和非 expert」拆开,因为两者被不同的并行切:
| 参数 | 怎么算 | 数量 |
|---|---|---|
| routed expert | 92 层 × 896 个 × | 约 2.72T |
| KDA 层 | 69 层,每层 q/k/v/gate/o 五个 加低秩的 | 约 30.6B |
| MLA 层 | 24 层,每层低秩 q/kv 加 gate/o | 约 5.6B |
| MoE 层的非 expert 部分 | 92 层 × (、、router、两个 shared expert) | 约 17.5B |
| embedding + LM head | 约 2.3B | |
| 第 0 层 dense FFN | 约 0.7B |
非 expert 部分加起来约 57B,被 PP 切成 份,在 EP 和 DP 的所有 rank 上都有副本;expert 部分 2.72T 被 PP 和 EP 一起切成 份。所以每张卡的 BF16 权重是
梯度和优化器状态跟着权重走。K3 用 Muon,优化器状态是 FP32 主权重加 FP32 动量,每个参数 8 字节;ZeRO-1 把它按权重的副本数 切开(分级的账见 ZeRO 那篇)。梯度按 FP32 算(论文 §5.3 提到「策略模型的 FP32 梯度缓冲」),不切的话每个参数 4 字节。
激活是唯一随 micro-batch 长度 增长的一项。逐层数一下反向要保存的张量:KDA 层每个 token 大约 7 万个值(输入、卷积后的 q/k/v、门、输出),MLA 层大约 7.5 万,MoE 层大约 24.5 万,其中 16 个 expert 的输入副本和输出各占 5.7 万。93 层加起来每个 token 约 2900 万个值,BF16 下约 58 MB。这个数字要乘 ,再乘上这张卡同时攒着几份 micro-batch 的激活,下一节讲这个「几份」。
把这些放进一张可以拖的表:
| 每张卡上住着什么 | 怎么算 | 朴素 | K3 | 付的货币 |
|---|---|---|---|---|
| 合计 |
形状来自 config.json:,93 层,69 层 KDA、24 层 MLA、92 层 MoE,896 个 expert 各 。激活按「反向要用的张量」逐层粗算:KDA 层约 70k 个值 / token,MLA 层约 75k,MoE 层约 245k(其中 16 个 expert 的分发副本和输出各 57k)。优化器按 Muon 记 FP32 主权重 + FP32 动量,梯度按 FP32。 是同一份权重的副本数,ZeRO 按它分片。数量级可信,个位数不可信。
「朴素」列是所有优化关掉的估算,「K3」列按开关算。默认参数下两列差一个数量级,差距几乎全在激活和 Muon 的完整参数缓冲上。下面按 §5.2.2 的六个小节,逐个看每个开关搬走了什么、付了什么。
跨 PP rank 均衡激活
先解释账本里「在途份数」这一项。1F1B 下 rank 峰值时攒着 份 micro-batch 的激活,rank 0 最满、rank 最空,虚拟段会让这个阶梯更陡。推导在流水线并行那篇里,那篇的时间线切到「在途激活」视图能看到每个 rank 的峰值份数。
这就是 §5.2.2「跨 PP rank 均衡激活」要解决的问题:显存压力集中在前面的 rank,后面的 rank 有大量空闲。K3 的做法直接:用 Mooncake Transfer Engine 把前面 rank 的激活远程 offload 到别的 PP rank 的显存里,让各 rank 的激活占用拉平。账本里对应的开关把「在途份数」从 rank 0 的峰值换成各 rank 的平均值。付的货币是节点间的传输带宽,藏在计算后面。
统一激活管理器:重算、量化、offload 都是存储策略
激活相关的其余三个技巧都建立在一个抽象上。每个为反向保存的张量绑一个可插拔的存储后端,重算、量化、本地 offload、远程 offload 不是四套机制,而是同一个接口下的四种策略,在张量粒度上可以自由组合,用张量上的轻量注解声明,和模型代码完全解耦。重算按函数粒度做,所以能跨层重算。实现上所有 GPU 显存在主计算流上分配、由一个内存池管理,避免多流造成的碎片和 host 侧开销;激活按层粒度预取回来,和计算重叠。
K3 的实际配置是:大部分激活用 block-wise FP8 量化加 offload 或远程 offload,逐元素算子用重算。三种策略各付什么货币、各适合哪类张量,显存账本那篇有一张对照表,这里不重复。账本里「激活 offload」开关打开后,卡上只驻留当前层和预取的下一层,其余在 CPU 或别的 rank 上。这是个粗模型,真实的驻留量取决于预取深度,但数量级就是这样:从「在途份数 × 全部层」变成「一两层」。
MoE 反向:推一下为什么可以不存 expert 输出
激活账本里最大的一行是 MoE 层。§5.2.2 在这一行上做了两处省,第一处值得把公式写出来。
一个 token 经过一个 routed expert 的计算是
是 expert 的中间激活(代码里的 act_output), 是 expert 输出(output), 是 router 归一化后的概率。反向要算 的梯度:
直接算需要 ,所以朴素实现要为反向保存每个 token 每个 expert 的输出,每 token 个值。但 是线性的,把 挪到内积的另一边:
右边的 正是反向传播到 时本来就要算的量(,差一个标量因子)。所以 的梯度可以用已经在算的东西加一次逐元素乘加得到,只依赖 和上游梯度, 不用存。论文说这是受 SonicMoE 启发,用一个数学变换消掉了反向对 output 的依赖,代价是一次轻量的逐元素计算。
第二处省在 dispatch 上。group GEMM 的输入是 token 按 expert 分组排列后的副本,朴素实现前向会把这个副本存下来。K3 只存 dispatch 之前的输入,反向时重新做一次 dispatch 把 GEMM 的输入恢复出来。dispatch 是通信,但它可以和 group GEMM 反向的一部分计算重叠(论文 Fig. 11 画的就是这个),所以这一份激活基本白省。两处合起来,MoE 层每 token 少存 个值,账本里的「MoE 反向改写 + 重算 dispatch」开关就是这个。
AttnRes:块代表只生成一次
第 4 篇讲过 Block AttnRes 的来源是 9 个块代表,第 7 篇讲过它在训练里的处理,这里只把它放进账本。块代表在边界层生成一次,之后所有层共享,直接留在 GPU 上;AttnRes 的计算整体包进 checkpointing,所以每层为反向保存的激活和普通残差结构完全一样,块结构没有额外的激活开销。流水线并行用 cache-based 通信,stage 之间只增量传新生成的块,micro-batch 结束就释放,达到内存占用的理论下界。账本里这一行是 个值每 token,乘在途份数,和其他激活比是小头。
梯度:Pipeline ZeRO-2 加 CPU 分片
激活之外,梯度是第二大的可变项。K3 用 Pipeline ZeRO-2 把梯度按 DP 的副本数 分片(ZeRO 的三级分片和流水线下的变体见 ZeRO 那篇),然后更进一步:分片存在 CPU 内存里,GPU 上只留一个 double grad buffer。反向算出的梯度先进这个 buffer,在 DP rank 之间 reduce 完之后累加到 CPU 上的分片。
这是「双槽位」模式第一次出现:GPU 上只有两个 chunk 大小的槽,交替使用,让 reduce 和累加藏在下一个 chunk 的反向后面。账本里「Pipeline ZeRO-2」开关把梯度从「每卡全部参数 × 4 字节」换成「两个 VPP chunk × 4 字节」。付的货币是 PCIe 带宽,和 offload 的激活共用。
Muon:P2P 代替全量 all-gather
最后一行是 K3 特有的。Muon 的 Newton–Schulz 正交化要对整个参数矩阵做,而分布式优化器把参数按 DP rank 切成了分片,所以每次更新前必须先把完整矩阵凑出来。
朴素做法是每个 rank 对整个参数缓冲做一次 all-gather。这有两个代价。显存上,每个 rank 要临时多出一份完整参数,账本里这一行和 BF16 权重同样大。通信上,每个 rank 都要收全部 个分片,每 rank 接收量约等于整份参数,规模一大通信成为主要瓶颈。
K3 的做法是按矩阵分工:每个 rank 只负责一部分矩阵的正交化,通过 P2P 从对应的 owner rank 拉取自己负责的那些矩阵的分片。每 rank 的接收量从「整份参数」降到「整份参数的 」,完整参数缓冲不再存在。通信和计算再按 model-chunk 缓冲的粒度流水化,把通信藏起来。账本里「P2P Muon」开关直接把这一行清零。
收尾:账本的前后两列
把六个开关全打开,再和「朴素」列并排看:
- 权重两行不动,它们由并行度决定,是这张卡的底。
- MoonEP 的冗余槽位是唯一变大的一行,用显存换第 8 篇的「规划永远有解」。
- 激活从「在途份数 × 全部层」变成「一两层」,FP8 再砍一半,MoE 改写砍掉每层最大的两块。
- 梯度从全量变成两个 chunk 的 double buffer,Muon 的完整参数缓冲消失。
每一行搬走的时候都付了一种货币:FLOPs、精度、PCIe、网络、或者一点代码复杂度。§5.2.2 没有一个技巧是免费的,它们的共同点是把付出的那部分藏进了第 8 篇里的某段空闲。
显存账本讲的是训练一步之内的事。第 10 篇把时间尺度拉长到 RL 的一轮迭代和线上服务的一次会话:当状态要活得比一步训练更久,它该住在哪里。