专栏Kimi K3 模型结构·系统10 / 11
9 min学习

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,并行度是滑块。

文本 token图像 / 视频视觉 token 与文本 token 交错后进入同一条主干残差流(prefix sum)重复 23 次3 KDA : 1 Gated MLA第 2 – 92 层每 12 层是一个 AttnRes 块AttnRes 来源(最多 9 个)每个 α 前都有这一组来源α = softmax(wₗ · RMSNorm(来源))wₗ 是每个子层各一个的可学习向量logits输出前再聚合一次全部块最终隐藏状态 + 下一个 token 的 embeddingAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和α词表 163840,隐藏维 7168Token Embedding163840 × 7168从零训练的视觉编码器,hidden 1024,12 头MoonViT-V227 层 · 0.4B · patch 14merge_type = sd2_tpool2×2 Pixel-shuffle+ 时间池化 · token ÷4PatchMergerMLPV2,GELU,RMSNormProjector4096 → 4096 → 7168词嵌入永远是一个来源= Embedding每个块的输出是块内所有子层输出之和已完成的块(每块 12 层之和)本块里已经算完的子层之和,相当于块内的普通残差流当前块 partial sumKimi Delta Attention:96 头 × 128 维,状态 128×128/头KDA第 1 层第一层不用 MoE,用一个稠密 SiTU-GLU FFNDense FFN仅第 1 层 · 中间维 33792三层 KDA,每层后接一个 Stable LatentMoEKDA×37168 → 3584 latent → 16 个 routed expert → RMSNorm → 7168Stable LatentMoE896 选 16 + 2 sharedDeepSeek 式 MLA,无位置编码,满秩 sigmoid 输出门Gated MLA×1 · NoPE同上Stable LatentMoE896 选 16 + 2 shared主干末尾额外放一层全局注意力Gated MLA第 93 层 · 收尾同上Stable LatentMoE第 93 层最终归一化RMSNorm不与 embedding 共享权重LM Head→ 163840预训练时 1 层多 token 预测;post-training 微调成投机解码的 draftMTP 层 ×1镜像主干 block · 部署时做 EAGLE-3 draft
序列 · KDA序列 · Gated MLA深度 · AttnRes宽度 · Stable LatentMoE输入 · MoonViT-V2零件
显存的大头是 MoE 层:expert 权重、expert 的中间激活、以及 AttnRes 每 12 层一份的块代表。

账本:每张卡上住着什么

训练时一张卡上的显存分五类:权重、梯度、优化器状态、激活、通信缓冲。通用的记账方法在显存账本那篇,这里直接按 K3 的形状填。先把参数量按 config 算清楚,第 0 篇算过总数,这里要按「expert 和非 expert」拆开,因为两者被不同的并行切:

参数怎么算数量
routed expert92 层 × 896 个 × 3×3584×30723\times 3584\times 3072约 2.72T
KDA 层69 层,每层 q/k/v/gate/o 五个 7168×122887168\times 12288 加低秩的 α\alpha约 30.6B
MLA 层24 层,每层低秩 q/kv 加 gate/o约 5.6B
MoE 层的非 expert 部分92 层 × (WW^\downarrowWW^\uparrow、router、两个 shared expert)约 17.5B
embedding + LM head2×163840×71682\times 163840\times 7168约 2.3B
第 0 层 dense FFN3×7168×337923\times 7168\times 33792约 0.7B

非 expert 部分加起来约 57B,被 PP 切成 PP 份,在 EP 和 DP 的所有 rank 上都有副本;expert 部分 2.72T 被 PP 和 EP 一起切成 P×RP\times R 份。所以每张卡的 BF16 权重是

57BP×2 B+2.72TPR×2 B.\frac{57\text{B}}{P}\times 2\ \text{B} + \frac{2.72\text{T}}{P R}\times 2\ \text{B}.

梯度和优化器状态跟着权重走。K3 用 Muon,优化器状态是 FP32 主权重加 FP32 动量,每个参数 8 字节;ZeRO-1 把它按权重的副本数 NN 切开(分级的账见 ZeRO 那篇)。梯度按 FP32 算(论文 §5.3 提到「策略模型的 FP32 梯度缓冲」),不切的话每个参数 4 字节。

激活是唯一随 micro-batch 长度 SS 增长的一项。逐层数一下反向要保存的张量:KDA 层每个 token 大约 7 万个值(输入、卷积后的 q/k/v、门、输出),MLA 层大约 7.5 万,MoE 层大约 24.5 万,其中 16 个 expert 的输入副本和输出各占 5.7 万。93 层加起来每个 token 约 2900 万个值,BF16 下约 58 MB。这个数字要乘 SS,再乘上这张卡同时攒着几份 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 rr 峰值时攒着 PrP-r 份 micro-batch 的激活,rank 0 最满、rank P1P-1 最空,虚拟段会让这个阶梯更陡。推导在流水线并行那篇里,那篇的时间线切到「在途激活」视图能看到每个 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 ee 的计算是

ae=SiTU(Wegatexe, Weupxe),oe=Wedownae,y=etop-kpeoe.a_e = \operatorname{SiTU}(W^{\text{gate}}_e x_e,\ W^{\text{up}}_e x_e), \qquad o_e = W^{\text{down}}_e a_e, \qquad y = \sum_{e\in\text{top-}k} p_e\, o_e .

aea_e 是 expert 的中间激活(代码里的 act_output),oeo_e 是 expert 输出(output),pep_e 是 router 归一化后的概率。反向要算 pep_e 的梯度:

Lpe=Ly, oe.\frac{\partial L}{\partial p_e} = \Big\langle \frac{\partial L}{\partial y},\ o_e \Big\rangle .

直接算需要 oeo_e,所以朴素实现要为反向保存每个 token 每个 expert 的输出,每 token 16×358416\times 3584 个值。但 oe=Wedownaeo_e = W^{\text{down}}_e a_e 是线性的,把 WedownW^{\text{down}}_e 挪到内积的另一边:

Ly, Wedownae=(Wedown)Ly, ae.\Big\langle \frac{\partial L}{\partial y},\ W^{\text{down}}_e a_e \Big\rangle = \Big\langle (W^{\text{down}}_e)^{\top}\frac{\partial L}{\partial y},\ a_e \Big\rangle .

右边的 (Wedown)L/y(W^{\text{down}}_e)^{\top}\,\partial L/\partial y 正是反向传播到 aea_e 时本来就要算的量(L/ae=pe(Wedown)L/y\partial L/\partial a_e = p_e (W^{\text{down}}_e)^{\top}\partial L/\partial y,差一个标量因子)。所以 pep_e 的梯度可以用已经在算的东西加一次逐元素乘加得到,只依赖 aea_e 和上游梯度,oeo_e 不用存。论文说这是受 SonicMoE 启发,用一个数学变换消掉了反向对 output 的依赖,代价是一次轻量的逐元素计算。

等价dispatch 后的输入gate / up GEMM+ SiTU-GLUact_output,保存GEMMoutput,不再保存combinedoutput,上游梯度原来:改写:只要,多一次逐元素乘加
上面是一个 expert 的前向。router 概率 的梯度原本要用 expert 输出 ,改写之后只用中间激活 和上游梯度, 这一行从激活账本里消失。

第二处省在 dispatch 上。group GEMM 的输入是 token 按 expert 分组排列后的副本,朴素实现前向会把这个副本存下来。K3 只存 dispatch 之前的输入,反向时重新做一次 dispatch 把 GEMM 的输入恢复出来。dispatch 是通信,但它可以和 group GEMM 反向的一部分计算重叠(论文 Fig. 11 画的就是这个),所以这一份激活基本白省。两处合起来,MoE 层每 token 少存 2×16×35842\times 16\times 3584 个值,账本里的「MoE 反向改写 + 重算 dispatch」开关就是这个。

AttnRes:块代表只生成一次

第 4 篇讲过 Block AttnRes 的来源是 9 个块代表,第 7 篇讲过它在训练里的处理,这里只把它放进账本。块代表在边界层生成一次,之后所有层共享,直接留在 GPU 上;AttnRes 的计算整体包进 checkpointing,所以每层为反向保存的激活和普通残差结构完全一样,块结构没有额外的激活开销。流水线并行用 cache-based 通信,stage 之间只增量传新生成的块,micro-batch 结束就释放,达到内存占用的理论下界。账本里这一行是 9×71689\times 7168 个值每 token,乘在途份数,和其他激活比是小头。

梯度:Pipeline ZeRO-2 加 CPU 分片

激活之外,梯度是第二大的可变项。K3 用 Pipeline ZeRO-2 把梯度按 DP 的副本数 NN 分片(ZeRO 的三级分片和流水线下的变体见 ZeRO 那篇),然后更进一步:分片存在 CPU 内存里,GPU 上只留一个 double grad buffer。反向算出的梯度先进这个 buffer,在 DP rank 之间 reduce 完之后累加到 CPU 上的分片。

CPU 上的梯度分片按 VPP chunk 分块1234GPU 上的两个槽位buffer A接收本 chunk 的反向梯度buffer Breduce 完累加到 CPU 分片用完交换角色
double grad buffer:GPU 上只有两个 VPP chunk 大小的梯度缓冲,一个在被反向写入,另一个在 reduce 并累加到 CPU 分片。§5.3 说 RL 训练里每张卡正是只留两个 VPP chunk 的梯度缓冲。

这是「双槽位」模式第一次出现: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 都要收全部 N1N-1 个分片,每 rank 接收量约等于整份参数,规模一大通信成为主要瓶颈。

K3 的做法是按矩阵分工:每个 rank 只负责一部分矩阵的正交化,通过 P2P 从对应的 owner rank 拉取自己负责的那些矩阵的分片。每 rank 的接收量从「整份参数」降到「整份参数的 1/N1/N」,完整参数缓冲不再存在。通信和计算再按 model-chunk 缓冲的粒度流水化,把通信藏起来。账本里「P2P Muon」开关直接把这一行清零。

收尾:账本的前后两列

把六个开关全打开,再和「朴素」列并排看:

  • 权重两行不动,它们由并行度决定,是这张卡的底。
  • MoonEP 的冗余槽位是唯一变大的一行,用显存换第 8 篇的「规划永远有解」。
  • 激活从「在途份数 × 全部层」变成「一两层」,FP8 再砍一半,MoE 改写砍掉每层最大的两块。
  • 梯度从全量变成两个 chunk 的 double buffer,Muon 的完整参数缓冲消失。

每一行搬走的时候都付了一种货币:FLOPs、精度、PCIe、网络、或者一点代码复杂度。§5.2.2 没有一个技巧是免费的,它们的共同点是把付出的那部分藏进了第 8 篇里的某段空闲。

显存账本讲的是训练一步之内的事。第 10 篇把时间尺度拉长到 RL 的一轮迭代和线上服务的一次会话:当状态要活得比一步训练更久,它该住在哪里。