Kimi K3 模型结构(7):结构决定系统
连载收尾。前六篇的每个结构选择都逼着系统那边做了一件事:KDA 的三种内核和为什么它的上下文并行不能直接把状态加起来;Block AttnRes 的内存和通信处理;LatentMoE 的融合 GEMM 与 token 中心的解码内核;混合架构下 512-token 粒度的前缀缓存;投机解码时 KDA 状态怎么回滚。
目录7 节
对应论文 §5.1、§5.2.2 的一部分、§5.2.3 和 §5.4。这一篇不讲新的结构,讲结构的后果:一个模块换掉之后,训练和推理系统里哪些东西跟着变。只挑和结构直接相关的部分,MoonEP、显存、RL 基础设施和集群调度放在第 8 到 10 篇。
KDA:三种执行形态,三种内核
KDA 用固定大小的状态换掉了随长度增长的 KV cache。论文说这个状态有两面:串行的更新和 GPU 偏好的宽而均匀的并行相冲突;但它尺寸固定,便宜地传输和复用。§5.1 的所有设计都是在压第一面、用第二面。
训练和 prefill:FlashKDA。 chunkwise 形式(第 2 篇)块内并行、块间串行。朴素实现里两个阶段交替,状态在块间传播时 SM 空转。FlashKDA 是一个基于 CUTLASS 的 chunkwise 内核,把块内计算和跨块的状态传播重叠起来:工作分解成按 token 并行的阶段和按头并行的递推,各自独立调度和调优。它已经作为 flash-linear-attention 的一个后端自动分发,训练和 prefill 都用它。代码在 MoonshotAI/FlashKDA。
超长 prefill 的设备内上下文并行。 张量并行把头切到不同设备上,但不缩短递推。纯 TP 部署下 prefill 一条超长序列时,每个 rank 只有几个头,大部分 SM 闲着。关键观察是:一段序列的状态转移可以在不知道进入状态的情况下独立算出来,事后再精确组合。于是一个 SM 级的自动上下文并行规划器把序列切到同一个 rank 的各个 SM 上,并行算每段的转移,再合并出每段精确的初始状态。这是设备内的,没有跨设备通信。下一节的 KCP 是同一个观察的跨设备版本。
decode:状态原地更新与投机解码的回滚。 decode 时瓶颈从并行度变成管理每步原地更新的状态。MTP 投机解码让这件事变棘手:验证阶段如果拒绝了一部分 draft token,状态已经走过了最后一个被接受的 token,没法简单回滚。给每个 draft 位置存一份状态快照可以回滚,但状态流量成倍增长,在线服务的大 batch 下这个成本占主导。
解法:任何被接受的 draft 前缀之后的状态,完全由这些 draft token 投影后的输入决定,而投影输入比状态小得多。所以只缓存投影输入,在片上重放来重建被接受 token 的状态,写回验证通过的 token 和 bonus token 的状态。这和同期的 ReplaySSM 是同一个想法。重放的 token、bonus token 和下一个 draft 窗口共用一个融合内核里的一个递推循环,内核覆盖短卷积、输入归一化、门、KDA 递推和输出归一化。验证延迟随验证 token 数次线性增长,低于存快照的基线。投影缓存不离开 decode 阶段,所以前缀缓存和 prefill/decode 分离用的负载和非投机服务一样。
KDA 上下文并行:为什么不能直接把状态加起来
上一节的设备内 CP 解决的是「一个 rank 上 SM 闲着」。但训练一条百万 token 的序列时(§3.4 说的序列维度切分),问题变成一张卡根本放不下:激活随长度线性增长,一层 KDA 的 和 chunkwise 中间量在 个 token 上就是几十 GB。这时序列必须切到 张卡上,每张只拿一段,这就是上下文并行(CP)。通用的 CP 是怎么做的,见并行总览里的上下文并行一节;这里只讲 KDA 的状态跨卡怎么传。
softmax 注意力和线性注意力的 CP 通信不是一个量级。 softmax 注意力里每个 token 要看到前面所有 token 的 ,所以 ring attention 让 张卡轮转 KV 块,每张卡前向收发约 字节,随序列长度 线性增长。线性注意力把前面的上下文压进一个固定大小的状态 ,跨卡只需要传这个状态,通信量与序列长度无关。这是线性注意力做 CP 的先天优势,KCP 要保住它。
普通线性注意力:本地状态相加就够。 更新是 ,段尾状态等于段内所有写入之和再加进入状态。于是每张卡从 出发算本地 token 产生的状态 ,进入第 段的状态就是 ,一次 all-gather 各自求和即可。之前的线性注意力 CP 方法(LASP 等)都是这样做的。
KDA 不行。 回顾第 1 篇的公式 (1):
delta rule 先把一个依赖 token 的矩阵 作用在进入的状态上,再加当前写入。前面几段写进状态的东西,会被本段每个 token 的 衰减和擦除,所以本段的效果取决于进入本段时的状态;从 算出来的 只包含本段自己的写入,丢掉了「进入状态被本段怎么改」这一半。
KCP 的分解(论文公式 (17))。 出路在于「被本段怎么改」这件事本身不依赖进入状态的取值。记 rank 那一段有 个 token, 是段内第 个 token 之后的状态, 是同一个递推从 出发得到的状态。对任意进入 rank 的状态,本地 个 token 之后
第一项是本地 token 自己生成的状态,第二项是前面 rank 的上下文经过本地 KDA 更新传播过来的部分。到段尾 时,累积转移 和零初始状态 都只用本地 token 就能算,不需要等前面的状态。这两个矩阵就是各 rank 交换的「片段」。
这个分解你在第 2 篇已经见过:那里把一个 chunk 的效果写成 , 是块内转移的累乘, 是从零出发的块内写入。KCP 的 和 就是把 和 从 16 个 token 的 chunk 放大到一整段的版本。这也意味着算它们不需要新的数学:把本段按 chunk 跑一遍 chunkwise 递推,起点取 就得到 ; 是各 chunk 的 依次相乘。论文只说这两个量「可本地计算」,这个对应关系是我的解读,FLA 里的实现是一个单独的预处理内核。
公式 (17) 的推导,以及为什么可以前缀扫描
第 1 步:归纳出段内的仿射形式。 对进入状态 和本段前 个 token,断言 ,其中 是同一递推从 出发的结果。 时 、(空乘积),成立。假设对 成立,代入公式 (1):
前两项恰好是从零出发的递推走到第 步,等于 ;第三项用矩阵乘的结合律合并成 。归纳完成。这一步只用了递推对 是线性的,没有用 的任何结构,所以对任何形如 的递推都成立。
第 2 步:段是仿射映射,复合仍是仿射映射。 第 1 步说第 段把进入状态 映成 。连续过两段:
还是同一形状,新的转移是两个转移的乘积,新的零初始状态是「前一段的零初始状态穿过后一段的转移,再加后一段自己的」。把「片段对」 的复合定义成 ,它继承矩阵乘的结合律,但不可交换:。普通线性注意力是 的特例,复合退化成 ,可交换,所以才能「直接相加」。
第 3 步:展开得到公式 (17) 的求和形式。 从 开始,依次复合第 到第 段,进入 rank 的状态是
第 段写进去的东西要穿过第 到 段所有的转移。每一项都只由各段本地算出的片段构成。
第 4 步:为什么是前缀扫描。 有结合律的二元运算,前缀 对所有 的计算就是前缀扫描(prefix scan)问题,可以并行或按顺序做。KCP 选最简单的做法:all-gather 之后每个 rank 手里有全部 对片段,rank 从 出发顺序做 次 。每次是一个 的矩阵乘加一个加法,每头 次乘加, 路 CP 最多做 次,和一个 chunk 的计算量同量级,不值得再并行。
流程:
通信量:一个数字。 每头的 是 , 是 ,bf16 下两者合计 64 KB。K3 的 KDA 有 96 个头(config linear_attn_config.num_heads),一层的片段是 6 MB;all-gather 后每个 rank 收到其余 份, 时一层收 42 MB,再乘 KDA 层数。这个数和序列长度无关:1M token 和 32K token 传的一样多。对照 ring attention 的 :、 时一层每卡收发约 30 GB。FLA 的实现把总通信量写成 ,就是上面这笔账。计算方面每个 rank 只算自己那 个 token,加上 步扫描,随长度线性扩展。
文档边界。 预训练的序列是打包的,一个 CP 组里的 个 token 可能包含多个文档,状态在文档边界要清零。所以 rank 的扫描不是无条件走完前面所有片段,而是只按顺序合并同一文档的片段:文档从本 rank 内部开始的部分进入状态就是 ;跨了 rank 边界的文档才需要前面那些 rank 的片段。实现里这由 cu_seqlens 决定,每个 rank 要知道自己段首的那个文档是从哪个 rank 开始的。
backward、short conv 和 MLA 层。 论文只写了前向。下面这几条来自 FLA PR #691 的实现,论文没提:
- 反向要把状态的梯度沿序列反着传:每个 rank 先本地算自己那段对片段的梯度贡献,all-reduce 求和,进入状态的梯度用 rank 之间的点对点 send/recv 传回前一个 rank。
- KDA 前面的 short conv(核宽 4)也跨段:每个 rank 段首的 3 个 token 需要前一个 rank 段尾的 3 个 token,前向和反向各一次点对点收发。
- PR 给的数据:H800、32K 序列、,KDA 前向比 all-to-all 式 CP 快 76%,前向加反向快 86%;同一 PR 也给 Gated DeltaNet 做了同样的事,因为 GDN 的 有同样的结构。
- K3 是混合架构,每 4 层有一层 MLA。KCP 只管 KDA 层,MLA 层的 CP 仍然要交换 KV,论文没有写这部分怎么做。
这个构造建立在 DeltaNet 的上下文并行之上。上一节设备内的 SM 级 CP 用的是同一个分解,只是把「rank」换成「SM」、把 all-gather 换成片上合并。§3.4 说的「让百万 token 训练可行的序列维度切分」就是这一节。
Block AttnRes:内存和通信
第 4 篇已经列过,这里按训练、prefill、decode 归拢。
训练。 块代表在边界层生成一次、留在 GPU 上被后续所有层共享。AttnRes 的计算整体包进 activation checkpointing,每层为反向保存的激活和普通残差结构完全一样,块结构没有增加激活内存。流水线并行用 AttnRes 论文的 cache-based 通信:stage 之间只增量传新生成的块,micro-batch 结束就释放,达到内存占用的理论下界。
prefill。 在每个 TP rank 上都物化块代表会造成大量重复内存。做法是对激活做 sequence parallel:把 TP 的 all-reduce 拆成 reduce-scatter 和 all-gather,块内的 AttnRes 内核插在两个集合通信之间,作用在按序列切片的隐藏状态上,每个 token 的块代表只在一个 rank 上物化。
decode。 块间那一步(对已完成块的批量打分)放到 side stream 上,和主流的独立计算重叠。块内那一步不再单独起内核:AttnRes 输出的合并、部分和的更新、后面的 RMSNorm,一起融进前一个 TP all-reduce。
Stable LatentMoE:融合 GEMM 与 token 中心的解码
第 5 篇的结构把 expert 数和每 token 激活数都翻了倍,调度和协调开销随之上升,常规 MoE 内核撑不住利用率。
latent 投影的三个优化。 把 latent 下投影 和 router 融成一个 GEMM,两者输入都是全宽的 ;把 latent 权重矩阵切到各 rank 上,用 multimem store 指令把输出的 all-gather 融进 GEMM 的 epilogue;把这部分通信和其他算子(比如 shared expert 的计算)重叠。合起来消除了冗余的权重流量和重复计算,把通信延迟藏在计算后面。
routed expert 的解码内核。 小 batch 时 group GEMM 退化成对权重矩阵的访存受限流式读取,以 tile 为中心的常规内核为计算密集设计、预处理开销大,不适合。K3 的 MoE 解码内核建立在 WarpDecode 的 token 中心设计上:每个 warp 负责一个输出神经元,直接从内存流式读取对应的权重。再把每个 warp 细分成 lane team,各自处理一个不相交的 expert 子集,最后 warp 内归约。权重布局离线一次性重排,大幅减少运行时 MXFP4 反量化的开销。
训练侧一句话。 §5.2.1 的 MoonEP 做「完美平衡」的 expert 并行训练:从当前 micro-batch 的 router 输出规划冗余 expert 并预取,融合的 permute/unpermute 算子让 token 直接落到远端 rank 的 expert 分组位置,routed expert 的 GEMM 用感知负载的调度器,shared expert 的 GEMM 放到单独的流上重叠。这和 QB 是一对:QB 让负载在统计上平衡,MoonEP 让每一步的执行形状固定。它的冗余 expert 上界定理和整条执行路径在第 8 篇里展开。
混合架构的前缀缓存
这是结构选择在推理系统里最直接的后果。K3 的每个块有三层 KDA 和一层 MLA,两种 cache 完全不同:MLA 的 KV 随长度增长、按 token 分页;KDA 的状态固定大小、每个请求一份。一个缓存的前缀只有在两者都能在同一个边界恢复时才可复用。
统一的分页布局。 各管一套分配、驱逐、传输逻辑会重复。K3 把 KDA 状态打包进和 MLA KV 同一个分页块池,页大小统一到相同字节数,两种页共享一份分配、引用计数和驱逐实现。页内所有头的状态逐头连续存放,每个头的字节流自成一体,是跨节点传输的最小单元;prefill 和 decode 节点 TP 度不同时,重排在传输路径上完成,GPU 侧零 reshuffle。作者还提了一个副产品:两种页类型不对称,任何类型混淆的访问会得到垃圾数据而不是看似合理的数据,等于零开销的布局自检。
粒度问题。 基于块哈希的前缀缓存以物理块为粒度复用 KV:只有完整的块才哈希,只有块对齐的前缀才能复用(机制本身见推理服务的缓存与调度)。这在 K3 上崩了。块哈希要求所有层共用一个块大小,而一个前缀命中只有在命中边界处的 KDA 状态被持久化了才可用。KDA 每个序列只有一个大状态,快照只在稀疏的边界上负担得起,于是共享的块大小被迫到 1024 到 6144 个 token,哈希粒度跟着一样粗。这个粒度下缓存几乎没用:比一个块短的请求永远不能复用,chunked prefill 在跨过整块边界之前导出不了任何可缓存的前缀。
解耦两个粒度。 前缀哈希在 MLA 页内的细哈希块上跑(512 token),物理块仍是粗的分配单元。KDA 反向对齐:递推状态的 checkpoint 只存在 MLA 哈希端点的一个稀疏子集上,也就是查找唯一可能引用的位置。
- prefill 期间,半满的 MLA 页以它最后一个完整哈希块的链式哈希注册进前缀索引(每个哈希覆盖前面所有哈希块,匹配到一个端点就证明了到它为止的整个前缀),注册的端点随页填充推进。每次前向之后,KDA 内核在最后一个哈希对齐的位置持久化状态。checkpoint 很大,被请求推进所取代的中间 checkpoint 回收,对话轮次边界上的保留供跨请求复用。缓存的 checkpoint 是只读快照,命中时复制进请求私有的运行状态,新 checkpoint 写到新槽位。
- 查找分两级。MLA 级按链式哈希匹配整个物理块,在第一个缺失的块处退回到块内的哈希端点,所以半满的页也能命中。KDA 级要求候选边界在每一个 KDA cache 组里都有 checkpoint(每组维护独立的递推状态)。命中点是两级都满足的最长边界,永远是哈希块的倍数,不必是物理块的倍数。图里请求前 2800 个 token 和缓存前缀相同,命中在 ,在 6144 的物理块深处,从 2560 继续 prefill。
并发调度下的一致性。 三条规则各对应一个共享半满块时的具体故障:所有 cache 组从一个共享空闲列表取块,给一个组分配私有副本可能驱逐另一个组刚命中的块,所以每个命中块在任何分配之前先在所有组里 pin 住;私有副本的拷贝在前向之前才在 GPU 上执行,本调度步内分配或注册的块可能把前一个所有者的字节交给读者,所以这类块在拷贝落地前排除在匹配之外;一个 checkpoint 只有在每个 KDA 组里都存在时才能恢复请求,所以驱逐一个组的 checkpoint 会原子地使兄弟组的失效。结论是:混合 KDA–MLA 模型的前缀缓存达到了和全注意力模型同样的通用性,任何共享前缀都能在任意 512-token 边界复用,与请求长度、分块、调度交错无关。
视觉编码器的两件事
这两件事在第 8 篇的流水线时间线上有图,这里只记结论。
- 动态上下文并行。 长上下文多模态训练里,大图和长视频让编码器计算时间大增、设备间负载失衡。大图沿 patch 维切到多个设备,注意力通过跨 CP rank 收集 KV(gather-KV)来算;每个 CP 组再分成若干子组,多张大图负载均衡地分配过去,避免通信占比随规模增长。
- 把 ViT 算进流水线气泡。 K2.5 引入的 Decoupled Encoder Process 把 ViT 和文本训练分成独立阶段。K3 进一步观察到,交错 1F1B 下前几个 micro-batch 的文本前向全排在最开头,最后几个 micro-batch 的文本反向全排在最末尾。于是前几个 micro-batch 的 ViT 前向同步地提前做,其余塞进流水线气泡,反向类似。大部分 ViT 计算被藏进气泡,视觉编码器的有效开销基本消失。
连载收尾
七篇下来,K3 的每个矩阵形状都应该能默写了。回头看,这个模型的结构改动可以用一句话概括:在三个维度上,把「固定权重的累加」换成「数据相关的选择」。序列上,KDA 用 delta rule 和逐通道衰减选择记什么忘什么,MLA 每四层做一次无损的全局选择;深度上,AttnRes 用 softmax 在块之间选择读哪一层;宽度上,LatentMoE 在 896 个 expert 里选 16 个,并且为了让选择本身可训练,加了三样稳定化的东西。系统那边的每一项工作,都是在为这些选择付账。
结构部分到此为止。§5 剩下的并行、显存和状态管理,在系统篇的三篇里接着讲:第 8 篇训练并行与 MoonEP、第 9 篇显存账本、第 10 篇RL 与服务。本连载不覆盖数据与预训练配方(§3)、post-training(§4)和评测(§6)。
评论