专栏Kimi K3 模型结构·系统8 / 11
17 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 基础设施和集群调度放在第 8 到 10 篇。

文本 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零件
这一篇覆盖三个新模块各自逼出来的系统工作。

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 的 q,k,vq, k, v 和 chunkwise 中间量在 10610^6 个 token 上就是几十 GB。这时序列必须切到 PP 张卡上,每张只拿一段,这就是上下文并行(CP)。通用的 CP 是怎么做的,见并行总览里的上下文并行一节;这里只讲 KDA 的状态跨卡怎么传。

softmax 注意力和线性注意力的 CP 通信不是一个量级。 softmax 注意力里每个 token 要看到前面所有 token 的 K,VK, V,所以 ring attention 让 PP 张卡轮转 KV 块,每张卡前向收发约 4bSd4bSd 字节,随序列长度 SS 线性增长。线性注意力把前面的上下文压进一个固定大小的状态 S∈Rdk×dvS \in \mathbb{R}^{d_k \times d_v},跨卡只需要传这个状态,通信量与序列长度无关。这是线性注意力做 CP 的先天优势,KCP 要保住它。

普通线性注意力:本地状态相加就够。 更新是 St=St−1+ktvt⊤S_t = S_{t-1} + k_t v_t^\top,段尾状态等于段内所有写入之和再加进入状态。于是每张卡从 S=0S = 0 出发算本地 token 产生的状态 S~[i]\widetilde S_{[i]},进入第 i+1i+1 段的状态就是 S~[1]+⋯+S~[i]\widetilde S_{[1]} + \cdots + \widetilde S_{[i]},一次 all-gather 各自求和即可。之前的线性注意力 CP 方法(LASP 等)都是这样做的。

KDA 不行。 回顾第 1 篇的公式 (1):

St=MtSt−1+β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 作用在进入的状态上,再加当前写入。前面几段写进状态的东西,会被本段每个 token 的 MtM_t 衰减和擦除,所以本段的效果取决于进入本段时的状态;从 S=0S = 0 算出来的 S~[i]\widetilde S_{[i]} 只包含本段自己的写入,丢掉了「进入状态被本段怎么改」这一半。

普通线性注意力,每段只是平移rank 1 那一段rank 2 那一段rank 3 那一段KDA随 token 变,先作用再写入rank 1 那一段rank 2 那一段rank 3 那一段三段的贡献线性叠加,交换顺序结果不变,进入状态本地状态之和。每一段是仿射映射,两段复合仍是仿射映射,可结合但不可交换。
普通线性注意力的段间状态可以直接相加;KDA 的每一段先用 作用进入状态再加本地写入,前面各段的贡献要穿过后面所有段的转移矩阵。两行里的 与 都只用第 段的 token 就能算出。

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]t←1 S[i]Ti,M[i+1]t←1:=∏r←1tMr∈Rdk×dk.(17)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}. \tag{17}

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

这个分解你在第 2 篇已经见过:那里把一个 chunk 的效果写成 SC=PCS0+HCS^C = P^C S^0 + H^C,PCP^C 是块内转移的累乘,HCH^C 是从零出发的块内写入。KCP 的 M[i]M_{[i]} 和 S~[i]\widetilde S_{[i]} 就是把 PP 和 HH 从 16 个 token 的 chunk 放大到一整段的版本。这也意味着算它们不需要新的数学:把本段按 chunk 跑一遍 chunkwise 递推,起点取 S=0S = 0 就得到 S~[i]\widetilde S_{[i]};M[i]M_{[i]} 是各 chunk 的 PP 依次相乘。论文只说这两个量「可本地计算」,这个对应关系是我的解读,FLA 里的实现是一个单独的预处理内核。

公式 (17) 的推导,以及为什么可以前缀扫描

第 1 步:归纳出段内的仿射形式。 对进入状态 S0S^0 和本段前 tt 个 token,断言 St=S~t+Mt←1S0S^t = \widetilde S^t + M^{t\leftarrow 1} S^0,其中 S~t\widetilde S^t 是同一递推从 00 出发的结果。t=0t = 0 时 S~0=0\widetilde S^0 = 0、M0←1=IM^{0 \leftarrow 1} = I(空乘积),成立。假设对 t−1t-1 成立,代入公式 (1):

St=MtSt−1+βtktvt⊤=MtS~t−1+βtktvt⊤+MtMt−1←1S0.S^t = M_t S^{t-1} + \beta_t k_t v_t^\top = M_t \widetilde S^{t-1} + \beta_t k_t v_t^\top + M_t M^{t-1 \leftarrow 1} S^0 .

前两项恰好是从零出发的递推走到第 tt 步,等于 S~t\widetilde S^t;第三项用矩阵乘的结合律合并成 Mt←1S0M^{t \leftarrow 1} S^0。归纳完成。这一步只用了递推对 SS 是线性的,没有用 MtM_t 的任何结构,所以对任何形如 St=MtSt−1+BtS_t = M_t S_{t-1} + B_t 的递推都成立。

第 2 步:段是仿射映射,复合仍是仿射映射。 第 1 步说第 jj 段把进入状态 SS 映成 M[j]S+S~[j]M_{[j]} S + \widetilde S_{[j]}。连续过两段:

M[j+1](M[j]S+S~[j])+S~[j+1]=(M[j+1]M[j])S+(M[j+1]S~[j]+S~[j+1]),M_{[j+1]}\big(M_{[j]} S + \widetilde S_{[j]}\big) + \widetilde S_{[j+1]} = \big(M_{[j+1]} M_{[j]}\big) S + \big(M_{[j+1]}\widetilde S_{[j]} + \widetilde S_{[j+1]}\big),

还是同一形状,新的转移是两个转移的乘积,新的零初始状态是「前一段的零初始状态穿过后一段的转移,再加后一段自己的」。把「片段对」(M,S~)(M, \widetilde S) 的复合定义成 (M2,S~2)∘(M1,S~1):=(M2M1, M2S~1+S~2)(M_2, \widetilde S_2)\circ(M_1, \widetilde S_1) := (M_2 M_1,\ M_2\widetilde S_1 + \widetilde S_2),它继承矩阵乘的结合律,但不可交换:M2S~1≠M1S~2M_2 \widetilde S_1 \ne M_1 \widetilde S_2。普通线性注意力是 M≡IM \equiv I 的特例,复合退化成 S~1+S~2\widetilde S_1 + \widetilde S_2,可交换,所以才能「直接相加」。

第 3 步:展开得到公式 (17) 的求和形式。 从 S=0S = 0 开始,依次复合第 11 到第 ii 段,进入 rank i+1i+1 的状态是

S[i]Ti=∑j=1i(∏l←j+1iM[l])S~[j],S^{T_i}_{[i]} = \sum_{j=1}^{i} \Big(\prod_{l \leftarrow j+1}^{i} M_{[l]}\Big)\widetilde S_{[j]} ,

第 jj 段写进去的东西要穿过第 j+1j+1 到 ii 段所有的转移。每一项都只由各段本地算出的片段构成。

第 4 步:为什么是前缀扫描。 有结合律的二元运算,前缀 x1∘⋯∘xix_1 \circ \cdots \circ x_i 对所有 ii 的计算就是前缀扫描(prefix scan)问题,可以并行或按顺序做。KCP 选最简单的做法:all-gather 之后每个 rank 手里有全部 PP 对片段,rank i+1i+1 从 S=0S = 0 出发顺序做 ii 次 S←M[j]S+S~[j]S \leftarrow M_{[j]} S + \widetilde S_{[j]}。每次是一个 (dk×dk)(dk×dv)(d_k \times d_k)(d_k \times d_v) 的矩阵乘加一个加法,每头 128×128×128128 \times 128 \times 128 次乘加,PP 路 CP 最多做 P−1P - 1 次,和一个 chunk 的计算量同量级,不值得再并行。

流程:

每个 rank 本地,互不等待本地 token第段算本段的累积转移,每头 128×128算从出发的段尾状态一次 all-gather交换固定大小的两个矩阵rank的进入状态对前面的段按顺序做继续本地chunkwise 计算
KCP:每段只算两个固定大小的量,一次 all-gather 之后按顺序合并,通信量与序列长度无关,计算随长度线性扩展。

通信量:一个数字。 每头的 M[i]M_{[i]} 是 dk×dk=128×128d_k \times d_k = 128 \times 128,S~[i]\widetilde S_{[i]} 是 dk×dv=128×128d_k \times d_v = 128 \times 128,bf16 下两者合计 64 KB。K3 的 KDA 有 96 个头(config linear_attn_config.num_heads),一层的片段是 6 MB;all-gather 后每个 rank 收到其余 P−1P-1 份,P=8P = 8 时一层收 42 MB,再乘 KDA 层数。这个数和序列长度无关:1M token 和 32K token 传的一样多。对照 ring attention 的 4bSd4bSd:d=7168d = 7168、S=106S = 10^6 时一层每卡收发约 30 GB。FLA 的实现把总通信量写成 P×H×dk×(dk+dv)P \times H \times d_k \times (d_k + d_v),就是上面这笔账。计算方面每个 rank 只算自己那 S/PS/P 个 token,加上 P−1P-1 步扫描,随长度线性扩展。

文档边界。 预训练的序列是打包的,一个 CP 组里的 SS 个 token 可能包含多个文档,状态在文档边界要清零。所以 rank i+1i+1 的扫描不是无条件走完前面所有片段,而是只按顺序合并同一文档的片段:文档从本 rank 内部开始的部分进入状态就是 00;跨了 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 序列、P=4P = 4,KDA 前向比 all-to-all 式 CP 快 76%,前向加反向快 86%;同一 PR 也给 Gated DeltaNet 做了同样的事,因为 GDN 的 MtM_t 有同样的结构。
  • 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 下投影 W↓W^{\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 让每一步的执行形状固定。它的冗余 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 在跨过整块边界之前导出不了任何可缓存的前缀。

一个物理块 = 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 边界复用,与请求长度、分块、调度交错无关。

视觉编码器的两件事

这两件事在第 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)。

评论