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

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 调度不在范围内。

文本 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词嵌入永远是一个来源b₀ = Embedding每个块的输出是块内所有子层输出之和b₁ … bₙ₋₁已完成的块(每块 12 层之和)本块里已经算完的子层之和,相当于块内的普通残差流bₙ⁽ⁱ⁾ 当前块 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零件
这一篇覆盖三个新模块各自逼出来的系统工作。

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 上下文并行:为什么不能直接把状态加起来

上下文并行把一条长序列切到 PP 个 rank 上。softmax 注意力的 CP 要在 rank 间交换 KV 块,大小随序列长度增长。线性注意力只需要传固定大小的状态。

普通线性注意力的 CP 很简单:每个 rank 从 S=0S = 0 出发算本地 token 产生的状态,把前面所有 rank 的本地状态加起来就是进入本 rank 的状态。因为更新是加法 St=St1+ktvtS_t = S_{t-1} + k_t v_t^\top,各段的贡献线性叠加。

KDA 不行。它的更新是

St=MtSt1+βtktvt,Mt:=(Iβtktkt)Diag(αt),S_t = M_t S_{t-1} + \beta_t k_t v_t^\top, \qquad M_t := (I - \beta_t k_t k_t^\top)\operatorname{Diag}(\alpha_t),

delta rule 把一个依赖 token 的矩阵 MtM_t 作用在进入的状态上。一段序列的效果取决于进入这段时的状态,从 S=0S = 0 算出来的本地状态不足以确定它。

KCP 的分解(论文公式 (17))。 记 rank ii 那一段有 TiT_i 个 token,S[i]tS^{t}_{[i]} 是段内第 tt 个 token 之后的状态,S~[i]t\widetilde S^{t}_{[i]} 是同一个递推从 S=0S = 0 出发得到的状态。对任意进入 rank i+1i+1 的状态,本地 tt 个 token 之后

S[i+1]t=S~[i+1]t+M[i+1]t1S[i]Ti,M[i+1]t1:=r1tMrRdk×dk.S^{t}_{[i+1]} = \widetilde S^{t}_{[i+1]} + M^{t \leftarrow 1}_{[i+1]}\, S^{T_i}_{[i]}, \qquad M^{t \leftarrow 1}_{[i+1]} := \prod_{r \leftarrow 1}^{t} M_r \in \mathbb{R}^{d_k \times d_k}.

第一项是本地 token 自己生成的状态,第二项是前面 rank 的上下文经过本地 KDA 更新传播过来的部分。到段尾 t=Ti+1t = T_{i+1} 时,累积转移 M[i+1]Ti+11M^{T_{i+1} \leftarrow 1}_{[i+1]} 和零初始状态 S~[i+1]Ti+1\widetilde S^{T_{i+1}}_{[i+1]} 都只用本地 token 就能算,不需要等前面的状态。这两个就是各 rank 交换的「片段」。

把递归展开,每个 rank 的进入状态都由这些片段组成:

S[i]Ti=j=1i(lj+1iM[l]Tl1)S~[j]Tj.S^{T_i}_{[i]} = \sum_{j=1}^{i} \Big(\prod_{l \leftarrow j+1}^{i} M^{T_l \leftarrow 1}_{[l]}\Big)\widetilde S^{T_j}_{[j]} .

rank 级的更新 SM[j]S+S~[j]S \leftarrow M_{[j]} S + \widetilde S_{[j]} 是可结合的,所以进入状态可以用前缀扫描恢复。流程:

代价是一次固定大小的 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 下投影 WW^{\downarrow} 和 router 融成一个 GEMM,两者输入都是全宽的 xx;把 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 在跨过整块边界之前导出不了任何可缓存的前缀。

一个物理块 = 6144 token = 12 个哈希块(每个 512 token)05121024153620482560307235844096460851205632○ 无 KDA checkpoint ● 有 checkpoint(通常在对话轮次边界) ● 本次命中命中边界 B = 2560MLA KV 已缓存(链式哈希匹配到这里)空,从 B 开始接着 prefill请求前 2800 个 token 与缓存前缀相同 → 恢复 B = 2560 处的 KDA checkpoint,写时复制半满的 MLA 块,从 2560 继续算,[0, 2560) 不重算。如果只按 6144 的物理块做哈希,这个请求什么都命中不了。

解耦两个粒度。 前缀哈希在 MLA 页内的细哈希块上跑(512 token),物理块仍是粗的分配单元。KDA 反向对齐:递推状态的 checkpoint 只存在 MLA 哈希端点的一个稀疏子集上,也就是查找唯一可能引用的位置。

  • prefill 期间,半满的 MLA 页以它最后一个完整哈希块的链式哈希注册进前缀索引(每个哈希覆盖前面所有哈希块,匹配到一个端点就证明了到它为止的整个前缀),注册的端点随页填充推进。每次前向之后,KDA 内核在最后一个哈希对齐的位置持久化状态。checkpoint 很大,被请求推进所取代的中间 checkpoint 回收,对话轮次边界上的保留供跨请求复用。缓存的 checkpoint 是只读快照,命中时复制进请求私有的运行状态,新 checkpoint 写到新槽位。
  • 查找分两级。MLA 级按链式哈希匹配整个物理块,在第一个缺失的块处退回到块内的哈希端点,所以半满的页也能命中。KDA 级要求候选边界在每一个 KDA cache 组里都有 checkpoint(每组维护独立的递推状态)。命中点是两级都满足的最长边界,永远是哈希块的倍数,不必是物理块的倍数。图里请求前 2800 个 token 和缓存前缀相同,命中在 B=2560=5×512B = 2560 = 5 \times 512,在 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)。其中并行和训练系统的部分会在「大模型笔记」专栏后面的章节里接着写。