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 基础设施、fleet 调度不在范围内。
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 上下文并行:为什么不能直接把状态加起来
上下文并行把一条长序列切到 个 rank 上。softmax 注意力的 CP 要在 rank 间交换 KV 块,大小随序列长度增长。线性注意力只需要传固定大小的状态。
普通线性注意力的 CP 很简单:每个 rank 从 出发算本地 token 产生的状态,把前面所有 rank 的本地状态加起来就是进入本 rank 的状态。因为更新是加法 ,各段的贡献线性叠加。
KDA 不行。它的更新是
delta rule 把一个依赖 token 的矩阵 作用在进入的状态上。一段序列的效果取决于进入这段时的状态,从 算出来的本地状态不足以确定它。
KCP 的分解(论文公式 (17))。 记 rank 那一段有 个 token, 是段内第 个 token 之后的状态, 是同一个递推从 出发得到的状态。对任意进入 rank 的状态,本地 个 token 之后
第一项是本地 token 自己生成的状态,第二项是前面 rank 的上下文经过本地 KDA 更新传播过来的部分。到段尾 时,累积转移 和零初始状态 都只用本地 token 就能算,不需要等前面的状态。这两个就是各 rank 交换的「片段」。
把递归展开,每个 rank 的进入状态都由这些片段组成:
rank 级的更新 是可结合的,所以进入状态可以用前缀扫描恢复。流程:
代价是一次固定大小的 all-gather,计算随长度线性扩展。这个构造建立在 DeltaNet 的上下文并行之上,KDA 版本在 flash-linear-attention 的 PR #691 里。§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 让每一步的执行形状固定。细节超出本连载范围。
混合架构的前缀缓存
这是结构选择在推理系统里最直接的后果。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 边界复用,与请求长度、分块、调度交错无关。
视觉编码器的两件事
- 动态上下文并行。 长上下文多模态训练里,大图和长视频让编码器计算时间大增、设备间负载失衡。大图沿 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 个,并且为了让选择本身可训练,加了三样稳定化的东西。系统那边的每一项工作,都是在为这些选择付账。
本连载没有覆盖的部分:数据与预训练配方(§3)、post-training 的 SFT、RL、多教师蒸馏(§4)、MoonEP 的细节、RL 与沙箱基础设施(§5.3)、评测(§6)。其中并行和训练系统的部分会在「大模型笔记」专栏后面的章节里接着写。