内容提要
本文详解Kimi K3模型训练中的显存优化策略,通过账本形式分析权重、激活等内存占用,并介绍六项关键技术:均衡PP rank激活、统一激活管理器、MoE反向改写、AttnRes块复用、Pipeline ZeRO-2及P2P Muon,以显存换效率,实现2.8T参数的高效训练。
延伸解读
显存优化的核心:以“货币”换空间
文章将显存优化比作“账本”,每一项技术都对应“搬走某一行,付一种货币”。例如,激活重算消耗额外FLOPs,FP8量化牺牲精度,offload占用PCIe或网络带宽。这种视角强调优化并非免费,而是将显存压力转移到其他资源上,并利用计算与通信的空闲时段来隐藏开销。理解这一点有助于评估不同优化策略的适用场景。
1F1B流水线的显存不均衡问题
在1F1B流水线并行中,rank 0因warmup阶段需保存最多P份在途激活,而rank P−1仅需1份,导致显存占用严重不均。K3通过远程offload将前序rank的激活转移至空闲rank,实现负载均衡。这揭示了流水线并行中“气泡”与显存压力的权衡,以及虚拟段(V>1)虽减少气泡但进一步加剧rank 0负担的特性。
MoE反向传播的数学技巧:省去expert输出存储
MoE层反向计算router概率梯度时,原本需保存expert输出o_e。通过线性变换将W_down转置并作用于上游梯度,可将梯度计算转化为仅依赖中间激活a_e和上游梯度的形式,从而无需存储o_e。这一数学改写受SonicMoE启发,以少量逐元素计算为代价,显著减少激活内存,是MoE训练显存优化的关键手段。
Pipeline ZeRO-2与P2P Muon:梯度与优化器的显存削减
梯度方面,Pipeline ZeRO-2将梯度分片并存储于CPU,GPU仅保留两个VPP chunk大小的double buffer,通过交替使用隐藏reduce通信。优化器方面,Muon的P2P通信替代全量all-gather,每个rank仅接收1/N的参数分片,消除了完整参数缓冲。两者分别以PCIe带宽和通信流水化为代价,大幅降低显存占用。
Q&A
Kimi K3 模型训练时,每张 GPU 上的显存主要被哪几类数据占用?
训练时一张卡上的显存分五类:权重、梯度、优化器状态、激活、通信缓冲。前三类由参数量和并行度决定,激活由 micro-batch 长度和流水线的在途份数决定。
在 1F1B 流水线并行中,为什么 rank 0 的激活占用最高?
在 1F1B 的 warmup 阶段,rank r 先做 P-1-r 个 micro-batch 的前向,然后才开始一前一反交替。前向做完、反向还没来的 micro-batch,其激活必须留在卡上。因此 rank r 峰值时攒着 P-r 份激活,rank 0 攒 P 份,rank P-1 只攒 1 份,所以 rank 0 最满。
K3 如何均衡不同 PP rank 的激活占用?
K3 使用 Mooncake Transfer Engine 将前面 rank 的激活远程 offload 到其他 PP rank 的显存中,使各 rank 的激活占用拉平。这相当于把在途份数从 rank 0 的峰值换成各 rank 的平均值,代价是节点间的传输带宽。
统一激活管理器中的重算、量化、offload 分别节省什么资源?
重算节省显存,代价是额外的前向 FLOPs;FP8 量化节省一半字节,代价是精度损失;offload/远程 offload 节省整个张量的显存,代价是 PCIe 或网络带宽。三种策略可组合使用。
MoE 反向传播中,为什么可以不保存 expert 的输出?
因为 expert 输出 o_e = W_down * a_e 是线性的,可以将 W_down 移到内积的另一边,使得 router 概率 p_e 的梯度只依赖中间激活 a_e 和上游梯度,而不需要 o_e。这样就不必为反向保存 expert 输出,节省显存。
Pipeline ZeRO-2 如何减少梯度显存?
Pipeline ZeRO-2 将梯度按 DP 副本数分片,并将分片存储在 CPU 内存中,GPU 上只保留两个 VPP chunk 大小的 double grad buffer。反向算出的梯度先进入 buffer,reduce 后累加到 CPU 分片,从而将 GPU 上的梯度显存从全量参数大小降到两个 chunk 大小。
P2P Muon 相比全量 all-gather 节省了什么?
P2P Muon 按矩阵分工,每个 rank 只负责一部分矩阵的正交化,通过 P2P 从 owner rank 拉取所需分片,每 rank 接收量从整份参数降到整份参数的 1/N,且不再需要构建完整参数缓冲,节省了显存和通信量。