Kimi K3 技术报告解读 | 2.8T MoE 的架构与 Infra:KDA、AttnRes、Stable LatentMoE 是怎么落地到 GPU 上的
Kimi K3 技术报告解读 | 2.8T MoE 的架构与 Infra:KDA、AttnRes、Stable LatentMoE 是怎么落地到 GPU 上的
1. 前言:1M 上下文的 KV cache 账单
你有没有算过一笔账:一个 1M token 上下文的对话,KV cache 到底要多大?
以 K2 这种规模的模型为例,64 个 attention head、每层每 token 都要存一份 K 和 V,61 层下来,每个 token 的 KV cache 是 MB 级的。1M token 的会话就是 GB 级,而且这还只是单个请求。做 agent 应用的话,一个任务动辄几百次工具调用、上下文滚到几十万 token,显存和带宽全被这坨只增不减的 cache 吃掉了。
这就是长上下文 LLM 的核心矛盾:你为了不让模型”失忆”,被迫为每一个见过的 token 保留一份 KV,但其中绝大部分信息在 90% 的请求里根本不会被再用到。
Kimi K3 的答案是把宝押在”线性注意力”上:2.78T 总参数、104.2B 激活参数的 MoE 模型,93 层里有 69 层用的是 KDA(Kimi Delta Attention)——一种用固定大小状态代替 KV cache 的注意力机制,只留 24 层全局 MLA。配合深度方向的 Attention Residuals 和宽度方向的 Stable LatentMoE,整体 scaling 效率比 K2 提升约 2.5×,训练上下文直接拉到 1M。
架构换这么激进,infra 也得跟着重写:KV cache 没了,prefix cache 怎么做?状态是串行更新的,怎么并行?896 个专家怎么做到完美负载均衡?这份技术报告难得地把这些全摊开了。这篇文章我就按”架构 → 训练 infra → RL infra → serving infra”的顺序过一遍,公式会讲清楚”每一项是什么、从哪推出来的”,关键模块配 PyTorch 实现。训练配方和评测分数就略过了。
2. 全局架构:三个维度各换一个”默认设置”
先看全景。K3 的架构可以理解成:把 Transformer 里三个”默认的信息流动方式”各换成了一个新设计:
- 序列维度(token 之间怎么传信息):默认的 softmax attention → Hybrid Attention,每个 block 里 3 层 KDA + 1 层 Gated MLA,3:1 交替,最后一层补 MLA 保证收尾是全局 attention
- 深度维度(层与层之间怎么传信息):默认的残差连接 → Attention Residuals(AttnRes),每层可以”回头”加权检索前面所有层的输出
- 宽度维度(每个 token 的通道怎么变换):默认的 SwiGLU FFN → Stable LatentMoE,896 个 routed expert 选 16 个 + 2 个 shared expert

和 K2 的硬参数对比:

值得注意的几个点:
- hidden dim 纹丝不动(7,168),参数从 1.04T 涨到 2.78T 靠的是层数(61→93)、专家数(384→896)、每专家宽度(2,048→3,072)。为什么不敢动 hidden dim?后面讲 LatentMoE 会看到,宽度是激活值 scale 的锚,2.8T 规模下动它很容易训练爆炸
- 93 层 = 69 层 KDA + 24 层 MLA。KV cache 的开销大头被线性注意力吃掉了,MLA 只负责在关键位置提供全局注意力
- 新增 401M 的 ViT(27 层,patch size 14)做原生视觉输入;1 层 MTP(Multi-Token Prediction)层,部署时拿来当 speculative decoding 的 draft model
下面逐个拆。
3. KDA:把 attention 改写成一台”会遗忘的笔记机”
这是全文最核心的部分,我不直接甩公式,从 attention 本身开始推。
3.1 第一步:attention 其实是一种”状态累积”
标准 attention 里,位置 $t$ 的输出是对所有历史的加权求和:
\[\bm{o}_t = \sum_{j \le t} (\bm{q}_t^\top \bm{k}_j)\, \bm{v}_j\]$q$、$k$、$v$ 分别是 query/key/value,直觉上:$q_t$ 是”我想查什么”,$k_j$ 是”第 $j$ 条信息能被怎么查到”,$v_j$ 是”第 $j$ 条信息的内容”。现在做个小小的代数变换,把求和顺序换一下——先把所有的 k 和 v 外积累加起来,再用 q 去查:
\[\bm{o}_t = \bm{q}_t^\top \underbrace{\textstyle\sum_{j \le t} \bm{k}_j \bm{v}_j^\top}_{\mathbf{S}_t}, \qquad \mathbf{S}_t = \mathbf{S}_{t-1} + \bm{k}_t \bm{v}_t^\top\]注意 $\mathbf{S}$ 的形状是 $[d_k, d_v]$:它把整个历史压缩成了一个固定大小的矩阵,读的时候用 $q$ 查一次就行。这就是线性注意力的本质——attention 的”查表”操作可以拆成”维护一张表 + 查表”两步,维护表只需要一次外积累加,和序列长度无关。
到这一步有个致命问题:$\mathbf{S}$ 是只增不减的。每条 $k_j v_j^\top$ 都往里堆,堆多了互相干扰,早期的信息被后期的冲得稀烂,而且没有”优先级”——3 个月前的一条废话和上一句的关键结论,在 $\mathbf{S}$ 里权重一样。
3.2 第二步:加”遗忘”和”擦除”
最直接的修补是加遗忘:每步先把旧状态打个折再写入,$\mathbf{S}t = \alpha \mathbf{S}{t-1} + \bm{k}_t \bm{v}_t^\top$。这就是 Gated 线性注意力(比如 GDN、Mamba-2 这一类)。但”整体打折”还是太粗了——遗忘应该是有针对性的:新信息如果和某条旧记忆”查询方式”($k$ 方向)冲突,应该替换它,而不是无差别衰减。
KDA 用的 delta rule 就是干这个的。它有个很漂亮的推导:把”写入”看成一步梯度下降。我们希望状态 S 在被 $k_t$ 查的时候,返回的是 $v_t$,也就是希望最小化
\[\mathcal{L}(\mathbf{S}) = \tfrac{1}{2}\|\mathbf{S}\bm{k}_t - \bm{v}_t\|^2\]从当前状态 $S_{t-1}$ 出发走一步梯度下降(步长 $β_t$):
\[\mathbf{S}_t = \mathbf{S}_{t-1} - \beta_t \nabla_{\mathbf{S}}\mathcal{L} = \underbrace{(\mathbf{I} - \beta_t \bm{k}_t\bm{k}_t^\top)}_{\text{擦除冲突}}\mathbf{S}_{t-1} + \beta_t \bm{k}_t\bm{v}_t^\top\](梯度是 $(\mathbf{S}\bm{k}_t - \bm{v}_t)\bm{k}_t^\top$,展开整理一下就得到上面这行。)
逐项读一遍这个公式,每一项都有明确的身份:
- $\operatorname{Diag}(\bm{\alpha}t)\mathbf{S}{t-1}$:先遗忘。$\bm{\alpha}_t \in (0,1)^{d_k}$ 是逐通道的遗忘因子——$d_k$ 维的每一维有自己的遗忘速度,这是 KDA 和普通 delta rule 的关键区别(”channel-wise forget gate”),有的通道记性长,有的通道记性短
- $(\mathbf{I} - \beta_t \bm{k}_t\bm{k}_t^\top)$:再擦除。它会把旧状态里与 $\bm{k}_t$ 同方向的分量减掉 $\beta_t$ 倍——只动和新记忆冲突的那部分,其他方向原封不动
- $+\ \beta_t \bm{k}_t\bm{v}_t^\top$:最后写入新记忆
- 读操作:$\tilde{\bm{o}}_t = \mathbf{S}_t^\top \bm{q}_t$,和 3.1 一样
合起来就是 KDA 的核心公式(原论文 Eq.1):
\[\mathbf{S}_t = (\mathbf{I} - \beta_t \bm{k}_t \bm{k}_t^\top)\operatorname{Diag}(\bm{\alpha}_t)\mathbf{S}_{t-1} + \beta_t \bm{k}_t \bm{v}_t^\top, \quad \tilde{\bm{o}}_t = \mathbf{S}_t^\top \bm{q}_t\]一句话总结:$\mathbf{S}$ 是一张固定大小的”联想记忆表”,每个新 token 到来时,先按各自的遗忘曲线给旧记忆打折,再把和当前 key 冲突的旧记忆擦掉,写入新记忆。KV cache 随长度线性增长的问题,被这个 $[d_k, d_v]$ 的状态接管了。
3.3 class KDA:一个能跑的最简实现
光看公式不过瘾,直接写代码。下面是单个 KDA head 的最简 PyTorch 实现(省略了 ShortConv、多头并行这些工程细节,但数学是完整的,维度全部标了):
import torch
import torch.nn as nn
import torch.nn.functional as F
class KDAHead(nn.Module):
"""单头 KDA 的最简实现(教学用,省略 ShortConv / 输出门)"""
def __init__(self, d_model, d_k=128, d_v=128):
super().__init__()
self.wq = nn.Linear(d_model, d_k, bias=False)
self.wk = nn.Linear(d_model, d_k, bias=False)
self.wv = nn.Linear(d_model, d_v, bias=False)
self.w_beta = nn.Linear(d_model, 1, bias=False) # 写入强度
self.w_alpha = nn.Linear(d_model, d_k, bias=False) # 产生逐通道 decay logit
self.g_min = -5.0 # log-decay 下界,K3 特有,见 3.5
def forward(self, x): # x: [B, T, d_model]
B, T, _ = x.shape
q = F.silu(self.wq(x)) # [B, T, d_k] 查询
k = F.silu(self.wk(x)) # [B, T, d_k] 地址
v = F.silu(self.wv(x)) # [B, T, d_v] 内容
beta = torch.sigmoid(self.w_beta(x)) # [B, T, 1] 这次写入用多大力气
g = self.g_min * torch.sigmoid(self.w_alpha(x)) # [B, T, d_k] log-decay,被压在 (-5, 0)
alpha = g.exp() # [B, T, d_k] 遗忘因子,每通道一个
S = x.new_zeros(B, self.wk.out_features, self.wv.out_features) # [B, d_k, d_v] 初始空白记忆
outs = []
for t in range(T): # 教学用逐 token 循环;真实 kernel 是 chunkwise 并行
k_t, v_t, q_t = k[:, t], v[:, t], q[:, t] # [B,d_k], [B,d_v], [B,d_k]
Sa = alpha[:, t, :, None] * S # ① 遗忘:每通道打折 [B, d_k, d_v]
Sk = torch.einsum('bk,bkv->bv', k_t, Sa) # ② 旧状态在 k_t 上的"应答" [B, d_v]
S = Sa - torch.einsum('bk,bv->bkv', beta[:, t] * k_t, Sk) \
+ torch.einsum('bk,bv->bkv', beta[:, t] * k_t, v_t)
# ↑ ③ 擦除+写入:S ← (I - β k k^T) Sa + β k v^T,一次 einsum 搞定
outs.append(torch.einsum('bk,bkv->bv', q_t, S)) # ④ 读:[B, d_v]
return torch.stack(outs, dim=1) # [B, T, d_v]
对照 3.2 的公式看:Sa 对应 $\operatorname{Diag}(\bm{\alpha})\mathbf{S}$,Sk 是 $\bm{k}^\top \cdot Sa$(即旧状态对 $\bm{k}_t$ 这个”地址”的当前应答),减去 $\beta \cdot \bm{k} \cdot Sk$ 是擦除冲突、加上 $\beta \cdot \bm{k} \cdot \bm{v}$ 是写入。整个 forward 里没有任何随 $T$ 增长的存储——这就是”固定大小状态代替 KV cache”在代码层面的样子。
真实的 KDA 当然比这复杂:q/k/v 前面有 ShortConv(短卷积,让相邻 token 的投影混合一下)和 L2Norm,输出端有 RMSNorm + full-rank 输出门($y_t = \mathbf{W}_o[\operatorname{Sigmoid}(\mathbf{W}_g \bm{x}_t) \odot \operatorname{RMSNorm}(\tilde{\bm{o}}_t)]$,每个 token 自己决定状态的哪些通道放行),而且是 96 个 head 并行。但主干循环就是上面这几行。
3.4 chunkwise 形式:GPU 上真正跑起来的样子
上面的 for 循环没法吃满 GPU——$T$ 个时间步串行,每步就一次小矩阵乘,GPU 利用率惨不忍睹。KDA 能实用的关键在于:这个循环可以重排成”chunk 间串行、chunk 内并行”。
把序列切成块长为 $C$ 的 chunk,把循环展开(忽略擦除项,只看衰减和写入),位置 $i$ 的输出其实是两部分之和:
- chunk 内:同 chunk 里 $j \le i$ 的 $k_j v_j^\top$,经过”从 $j$ 衰减到 $i$”的贡献
- 跨 chunk:进入这个 chunk 时的状态 $\mathbf{S}_{\text{in}}$,经过”从 chunk 开头衰减到 $i$”的贡献
第一部分如果直接算,每个 $(i, j)$ 对都要单独乘一个衰减率乘积,又是逐位置对操作。KDA 的技巧是:给 query 端乘上自己的累积衰减 $\bm{\Gamma}_i$,给 key 端除以 chunk 末尾的累积衰减 $\bm{\Gamma}_C$,两者相除($\bm{\Gamma}_i / \bm{\Gamma}_j$)恰好还原出”从 $j$ 到 $i$”的衰减率——于是 chunk 内所有项变成了一个 $C \times C$ 的下三角矩阵乘:
# chunkwise 核心两行(简化,省略 UT transform 细节)
# Q, K: [C, d_k],V_tilde: [C, d_v],S_in: [d_k, d_v]
Gamma = alpha.cumprod(dim=0) # [C, d_k],每位置相对 chunk 开头的累积衰减
A = torch.tril((Q * Gamma) @ (K / Gamma).T) # [C, C] 下三角,chunk 内 attention 矩阵
O = (Gamma * Q) @ S_in + A @ V_tilde # 前项跨 chunk,后项 chunk 内,全是 matmul
串行的只剩 $\mathbf{S}_{\text{in}}$ 的传递(chunk 间一次 $[d_k, d_v]$ 矩阵的更新),chunk 内部全是 dense matmul——这才是能吃满 Tensor Core 的形态。
但注意 K / Gamma 这个除法——这就是数值炸弹。
3.5 lower-bounded decay:一行公式救了一个 kernel
$\bm{\Gamma}$ 是一串 $(0,1)$ 区间的数连乘,chunk 末尾的 $\bm{\Gamma}_C$ 可以小到任意接近 0,它的倒数 $1/\bm{\Gamma}_C$ 就可以大到任意离谱。BF16 的动态范围是有限的,除数一大直接溢出,训练就炸了。
Kimi Linear(KDA 的前一版工作)的应对是把衰减放到 log 空间算相对值,但即便如此,chunk 内部对角线附近的 tile(自己乘自己的部分)仍然躲不开逐位置对计算,是 kernel 的慢路径。
K3 的修法简单到有点朴素:给 log-decay 加个下界。原来 decay logit 经负 Softplus 映射,g 可以取到 -∞;K3 换成缩放的 sigmoid:
\[\bm{g}_t^h = g_{\min} \cdot \operatorname{Sigmoid}(e^{A_h} \bm{z}_t^h), \quad g_{\min} = -5\]$\bm{z}$ 是网络算出来的 decay logit,$A_h$ 是可学习的逐 head 缩放。效果:每步遗忘因子 $\bm{\alpha} > e^{-5} \approx 6.7 \times 10^{-3}$,16-token tile 内的累积 log-decay 被压在 $(-80, 0)$,$1/\bm{\Gamma}$ 最大不超过 $e^{80}$——BF16 装得下。

这个改动的连锁反应非常划算:数值范围有界之后,对角 tile 也能用 dense Tensor Core matmul 了,逐位置对的慢路径整个删掉,kernel 里不再有 if-else 分支。算法上只是把遗忘因子的下界从 0 抬到了 $e^{-5}$(模型最坏情况下”忘得慢一点点”),系统上换来了一个无慢路径的 kernel。这是典型的算法-系统协同设计,也是我觉得这份报告里最优雅的一处改动。
4. Gated MLA:剩下的 1/4 全局注意力
69 层 KDA 管效率,24 层 MLA 管质量。MLA(Multi-head Latent Attention,DeepSeek-V2 提出)的思路一句话:KV cache 太大,那就把每个 token 的 K、V 压缩成一个低维 latent 向量 $c_t = \mathbf{W}_c \bm{x}_t$ 存起来,算 attention 时再升维重建。用一点额外的计算换 cache 体积,对全局注意力层来说很划算。
K3 在 MLA 上做了两个改动:
- 全部 MLA 层不加位置编码(NoPE)。这听起来很反直觉——没有位置信息的 attention 怎么知道词序?答案是 KDA 层补上了:衰减机制天然编码了”远近”(越旧的记忆折扣越狠),位置敏感的任务交给 KDA,MLA 只管内容检索。这么做的实际好处在长上下文扩展时体现:不用重调 RoPE base、不用 YaRN 插值,位置编码参数根本不存在,从 64K 扩到 1M 少了一整个工程环节
- 加了和 KDA 同款的输出门:$y_t = \mathbf{W}_o[\operatorname{Sigmoid}(\mathbf{W}_g \bm{x}_t) \odot \tilde{\bm{o}}_t]$,满秩投影,每个 token 自己调制全局注意力读出来的通道
还有一个纯 infra 层面的细节:训练时 attention 输出保持 FP32(flash attention 的低精度舍入有一个有偏误差,长训练里会累积)。代价是片上输出 tile 显存翻倍,所以他们把训练 kernel 重写了,让输出 tile 和 KV staging buffer 共享 shared memory——一个精度选择背后是一次 kernel 重构,这种细节很能说明这个团队的工程深度。
5. Attention Residuals:把 attention 从时间轴搬到深度轴
残差连接有个很少被质疑的默认设定:第 $l$ 层拿到的信息 = 第 $l-1$ 层的输出。前面 90 层算出来的所有东西,都被”无差别累加”进一条隐状态里往下传。
你把视角转一下就会发现这很像 RNN:RNN 在时间轴上把所有历史压进一个状态,Transformer 用 attention 干掉了这个瓶颈;现在深度轴上,残差连接干着一模一样的事。93 层的模型,第 80 层想用第 10 层的某个特征?对不起,它已经被中间 70 层的累加搅过 70 遍了。
K3 的 AttnRes 把 attention 搬到了深度轴上:每层配一个可学习的伪 query $\bm{w}_l \in \mathbb{R}^d$,把前面所有层的输出当 K/V,算一组数据相关的权重做加权检索:
# AttnRes 核心(简化)
# H_prev: [l, d],前面所有层输出的堆叠(第 0 行是 token embedding)
attn = softmax(w_l @ RMSNorm(H_prev).T) # [l]:本层对每个前层输出的关注度
h_l = attn @ H_prev # [d]:加权检索结果
对 key 做 RMSNorm 是必须的——不同层的输出幅值差异很大,不做归一化的 softmax 会被幅值最大的那一两层垄断,检索就失效了。
Full 版本的代价是所有层输出都得活着:显存和流水线跨卡通信都是 $O(L \cdot d)$。所以 K3 用的是 Block 版:93 层切成 8 个 block(每块 12 层),block 内部各层输出直接求和压成一个 block 表示,加上 embedding 一共 9 个 K/V 源。每层只从这 9 个表示里检索,开销从 $O(L \cdot d)$ 降到 $O(N \cdot d)$。论文的消融显示 $N \approx 8$ 就能拿回 Full 版大部分收益——花 9 份 K/V 的代价,让每层都获得”翻旧账”的能力。
6. Stable LatentMoE:896 选 16 怎么训得稳
MoE(Mixture-of-Experts)的思路大家熟:把 FFN 复制成很多份”专家”,router 给每个 token 挑几个,激活参数远小于总参数。K3 的配置是 896 个 routed expert 每 token 选 16 个(稀疏度 56),外加 2 个所有 token 必经的 full-width shared expert。
K3 有个特有设计:routed 专家不在 7,168 维全宽空间里干活,而是先降到 3,584 维(0.5×)的 latent 空间。前向流程:
# Stable LatentMoE 前向(简化)
# x: [7168],latent 维度 ℓ=3584,专家中间层 3072
scores = torch.sigmoid(W_r @ x) # [896] router 打分
idx, p = qb_topk(scores + b, k=16) # 选 16 个专家(b 是负载均衡偏置,见 6.3)
u = sum(p_i * E_i(W_down @ x) for i in idx) # W_down: [7168→3584],专家在 latent 空间算
y = W_up @ RMSNorm(u) + Shared_1(x) + Shared_2(x) # W_up: [3584→7168],shared 全宽
为什么敢把专家砍到半宽?因为 shared 专家保住了全宽的信息通路,routed 分支只负责稀疏补充。896 个全宽专家的参数量是不可承受的,latent 化之后专家参数才压得住,同时 router 的搜索空间(896 选 16)还保留了大容量专家分化的可能。
但”896 选 16”这个极端稀疏度把两个老问题放大到了新的量级:
6.1 激活爆炸 → RMSNorm + SiTU-GLU
routed 路径是 W↓ → 专家 FFN → W↑ 一路连乘,本身条件数就不好;2.8T 规模下,这条路径的内部激活会爆掉。K3 的组合拳:
- W↑ 之前插 RMSNorm,把 16 个专家加权聚合后的 u 的 scale 拉回来再升维
- SwiGLU 换成 SiTU-GLU。SwiGLU 的两个乘法分支都是无界的,激活值想多大就多大,低精度训练里就是 overflow 隐患。SiTU-GLU 给两个分支各加了个 softcap——$\beta \tanh(x/\beta)$,原点附近近似线性(不伤拟合能力),大值处平滑封顶:
| K3 取 $\beta_1 = 4$、$\beta_2 = 25$,**输出被硬约束在 $ | f(x) | \le \beta_1 \beta_2 = 100$**。看曲线更直观——红线(SiTU-GLU)在原点附近贴着 SwiGLU 走,大值处压平: |

6.2 负载均衡 → Quantile Balancing
MoE 训练的经典难题:router 偏爱少数专家,热门专家的卡成为瓶颈,冷门专家”死掉”。DeepSeek 系的 auxiliary-loss-free 方案是给每个专家维护一个 bias $b_j$,只加在 Top-$k$ 选择上(不进梯度),按负载 $\pm\gamma$ 地调。但 sign 更新的粒度太粗了——384 个专家时还能用,896 个专家时要么收敛慢、要么负载来回震荡。
K3 的 Quantile Balancing(QB)把”慢慢调”换成了”一次算准”。核心观察:跑完一次 router 前向,每个 token 的 Top-$(k+1)$ 分数里就藏着答案——第 $k+1$ 名的分数就是”进入这个 token 的 Top-$k$ 需要跨过的门槛” $\alpha_i$。对专家 $j$,如果 bias 设成 $\hat{b}j$,它会被多少 token 选中?就是满足 $s{i,j} + \hat{b}j > \alpha_i$ 的 token 数。这个数随 $\hat{b}_j$ 单调递增,令它等于目标负载 $q = mk/n$,反解出来的 $\hat{b}_j$ 就是让专家 $j$ 恰好达标的那一个——数学上就是 margin $(s{i,j} - \alpha_i)$ 的 $(1 - k/n)$ 分位数:
# Quantile Balancing 更新(简化)
# s: [m, 896] router 分数,alpha: [m] 每个 token 的 Top-(k+1) 门槛
margin = alpha[:, None] - s # [m, 896]:专家进各 token Top-k 还差多少分
b_hat = -torch.quantile(margin, q=1 - k/n, dim=0) # [896]:逐专家取分位数,一步到位
b = b_hat - b_hat.mean() # 去公共偏移(整体平移不改变 Top-k)
论文里 Fig.5 的例子很直观:8 个 token、4 个专家、各选 1 个,naive Top-k 的负载是 (4,3,1,0)——E1 过热、E4 直接饿死;QB 调整 bias 后变成 (2,2,2,2):

工程上还有个坑:margin 有百万级、散在各卡上,直接 gather 不现实。K3 的做法是每专家本地统计 margin 直方图(1000 个 bin),all-reduce 直方图计数(整数加法,结果与数据怎么切分无关),再从合并直方图读分位数。通信量从百万个浮点数降到每专家 1000 个整数,误差只有 bin 宽度级(10⁻³ 量级)。这个”分布式的分位数估计”思路本身就很值得收藏。
7. 训练 Infra:3T 级 MoE 怎么跑满集群
架构讲完了,进入 infra。先补三个术语:TP(张量并行,把单个矩阵切到多卡)、EP(专家并行,不同专家放不同卡)、PP(流水线并行,把层切成阶段)。K3 的并行栈是 PP(带虚拟阶段)+ EP + ZeRO-1 DP + Pipeline ZeRO-2 梯度分片 + CP 的组合。3T 模型的训练系统要解决三件事:专家负载不均、显存装不下、KDA 的串行状态怎么并行。
7.1 MoonEP:把”负载均衡”做成数学保证
EP 的经典痛点:router 是数据相关的,每个 batch 各专家收到的 token 数天然不均。EP 下每个 rank 只持有自己的专家,token 要通过 all-to-all 通信寄到专家所在的卡上——最热的卡决定整批速度,其他卡干等。
现有方案(ECHO 这类)预设冗余专家数量或 token 上限,遇到极端分布可能直接无解、训练中断。K3 的 MoonEP(已开源)做到了完美均衡:每个 EP rank 恰好收到 $S \times K$ 个 token($S$ 序列长、$K$ 选专家数),所有卡计算量完全一致。两个核心机制:
- 冗余专家迁移 + 理论上界:热门专家复制一份到空闲卡分摊负载。K3 证明了:只要给每个 rank 预留 $E/R$ 个冗余槽位($E$ 专家数、$R$ 并行度),均衡方案一定存在,训练永远不会因为”找不到可行的迁移方案”而卡住。这个 bound 是紧的
- 零拷贝通信:planning kernel 预先算好每个 token 的目的地,token 直接写到远端卡上按专家分组好的位置,通信 buffer 以 view 形式直接交给计算。对比 DeepEP 的免拷贝路径要按最坏情况开 $S \times K \times R$ 的 buffer,MoonEP 只要固定的 $S \times K$
第三点是我认为最关键的:每层每卡 token 数恒等于 $S \times K$,所有 kernel 的 shape 编译期就定了。这意味着没有 host-device 同步(CPU 不用等 GPU 报数再决定 launch 什么 kernel),没有动态 shape 带来的重编译,CUDA graph 之类的优化随便上。做过大模型训练的人都知道,大规模下真正磨人的往往不是算力不够,而是这些微秒级的同步开销乘以巨大的层数×步数。
7.2 显存:张量粒度的”存哪”决策系统
2.8T 模型的激活、梯度、优化器状态都远超单卡显存。K3 的做法是把”这个张量存哪里、要不要留”做成每个张量独立声明的策略:重计算 / FP8 量化 / offload 到 CPU / offload 到其他卡的内存,全部可组合、与模型代码解耦。K3 的实际配置:大部分激活用 block-wise FP8 量化 + offload,逐元素算子用重计算。
几个有意思的点:
- MoE 反向省显存:把 permute 概率的梯度数学上重写成只依赖中间激活和上游梯度,前向 group-GEMM 只存 dispatch 输入,反向时重算,重算的通信和 GEMM 反向重叠——用一点计算换一大块显存
- AttnRes 反而省显存:block 表示只在 block 边界算一次、全层共享,整个 AttnRes 包进 checkpointing 后,每层保存的激活和标准残差架构一样多——深度轴上加了 attention,显存却几乎没涨
- Pipeline ZeRO-2:梯度按 DP rank 分片后存 CPU 内存,GPU 上只留双缓冲
- Muon 的 P2P 正交化:K3 用的 Per-Head Muon 优化器把 attention 投影矩阵按 head 切开分别做 Newton-Schulz 正交化(避免大梯度 head 垄断共享的更新方向),但优化器状态分片后正交化需要完整矩阵,naive 做法是全量 all-gather——K3 改成每个 rank 只通过 P2P 拉自己拥有的参数分片,通信和计算流水线重叠
7.3 KDA 的系统协同:固定状态是福也是祸
KDA 的固定大小状态 S 不用像 KV cache 那样随长度增长,这对通信是好事;但它的更新是串行的,切到多卡并行是难题。
- FlashKDA kernel:chunkwise 计算里”chunk 内并行”和”跨 chunk 传状态”是交替的,naive 实现在串行传状态阶段 SM 全部空转。FlashKDA(CUTLASS 实现)把两条路径拆开重叠调度,训练和 prefill 共用一个 kernel
- 单卡内并行:超长序列 prefill 时,TP 把 head 切薄了,每卡几个 head 根本喂不饱 SM。关键观察:每段的状态转移可以不依赖输入状态独立算出来,之后精确合并——于是在单卡 SM 层面切序列并行,零跨卡通信
- 跨卡 KDA Context Parallelism:普通线性注意力的 CP 可以各卡从零状态算本地状态再相加(因为更新是线性的),但 KDA 的转移算子 $(\mathbf{I} - \beta \bm{k}\bm{k}^\top)\operatorname{Diag}(\bm{\alpha})$ 依赖输入数据,不能直接加。KCP 把每段的贡献拆成两个本地可算的量——累积转移矩阵 $\mathbf{M}$ 和从零状态算的本地状态 $\tilde{\mathbf{S}}$——一次 all-gather 交换这些固定大小的 fragment,再用前缀扫描恢复每段真实的入状态。通信量与序列长度无关
8. 1M 上下文 Agentic RL:状态要”活得久”
K3 的 RL 是百万 token、成百上千次工具调用的长程任务,infra 的核心挑战从”算得快”变成了”状态活得久”。
外部 KV cache 池(write-back 设计)。RL 系统里 rollout 和 training 共存:正在解码的请求的 KV 在 GPU 上,而大量”暂时轮空但之后还要用”的前缀(比如一个多步任务中间等训练的状态)如果常驻 GPU 就爆了。K3 的策略:空闲前缀在被 GPU 驱逐时才写回 CPU DRAM,下次复用前预取——DRAM 带宽只为真正离开解码路径的前缀付费。KDA 状态跟着对应的 MLA KV block 一起 offload/预取(两者生命周期天然一致)。训练迭代时,模型权重和优化器状态进一步 offload 到 NVMe,把 DRAM 让给 cache 池。配合 partial rollout(长任务分段续跑),每个 1M 上下文的 RL 实验控制在几百张卡以内。
梯度 buffer 复用。RL 还需要 reference model 算 KL,但这模型太大放不进显存——权重常驻 CPU、按需搬入。搬进哪块显存?K3 的答案很妙:直接复用 policy 模型的 FP32 梯度 buffer。ZeRO-2 分片下每卡只有两个虚拟流水线 chunk 的梯度 buffer,reference 权重就逐 chunk 流水线式进出:一个 chunk 在算,另一个在预取。显存一点没多花。
沙箱 AgentENV(已开源)。Agent 会挂载磁盘、起容器、甚至跑内核操作,普通容器隔离不住(报告里明说遇到了内核 panic)。K3 用 Firecracker microVM,几个数字感受下这个系统的成熟度:
- 增量 checkpoint 只存脏页:checkpoint 133ms、resume 49ms。agent 等推理结果时沙箱可以整个暂停(这类等待最高占沙箱生命周期的 98%),不占任何内存 CPU
- Fork 能克隆出状态完全一致的沙箱且原沙箱继续跑——无副作用的 reward 评判靠这个实现(在克隆体上试错,不影响原环境)
- 秒级拉起数万个沙箱,内存超卖比 6.5×;K3 训练+评测累计创建 5120 万+ 个沙箱
9. Serving:KDA 时代重写一遍 Prefix Cache
推理侧最让我感慨的是:KDA 把 prefix cache 这个”已经解决好的问题”重新变成了未解决。
传统 KV cache 按 token 分页,任意前缀边界都能复用——这是 vLLM 们已经打磨得很成熟的能力。KDA 的状态是每请求一份的固定大小张量,只有存了 checkpoint 的边界才能恢复,而且 checkpoint 不能太密(存太多状态本身就成了新的开销)。一边是任意边界可复用的 KV,一边是稀疏边界才能复用的 KDA 状态,一个请求的前缀要么都能用、要么都不能用。
K3 的解法是粗细两级粒度解耦:
- 统一 cache 布局:KDA 状态直接打包进和 MLA KV 相同的分页 block 池,page 统一字节大小,分配/引用计数/驱逐一套代码全管——两套生命周期完全不同的缓存,共用一套管理机制
- 哈希在细粒度做:前缀匹配在 512-token 的 hash block 上做,物理 block 仍是粗分配单位(比如 6144-token 物理块 = 12 个 hash block)。KDA checkpoint 只在稀疏的 hash 边界存(通常是对话轮次边界)
- 命中流程:MLA 部分按链式哈希匹配到任意 hash 边界,KDA 部分恢复最近的 checkpoint,中间未覆盖的部分 copy-on-write 后继续 prefill

效果:任意共享前缀在任意 512-token 边界可复用,与请求长度、chunking、调度交错无关。
decode 侧也有新问题:MTP speculative decoding(用一个小的 draft 层一次猜多个 token,大模型再验证)碰上 KDA,被拒绝的 draft token 会让状态”多走了几步”,逐位置存状态快照的通信开销在大 batch 下是主导成本。K3 的 kernel 换了个思路:只缓存 draft token 的投影输入(q/k/v 这些,比状态小得多),被接受 token 的状态在片上重放重建,验证和写回在一个 fused kernel 里完成。
集群调度层面两个实用设计:典型 coding 请求是 400K token 前缀 + 4K token 增量,cache 命中比 miss 便宜几个数量级,所以请求要路由回持有它前缀的集群(每 session 绑主备两个集群,故障时备用集群重新 prefill,重填负载靠均匀散列摊到全集群);长短请求的单条成本差三个数量级(<2K 到 1M),按请求类别分独立资源预算,防止长上下文流量挤爆全系统首 token 延迟。
部署侧还有个容易被忽略的点:post-training 全程(SFT + RL)做了 MXFP4 量化感知训练——MoE 专家权重 MXFP4、激活 MXFP8,attention 投影和 router 保持高精度,rollout 和训练用同一套量化。训练时就按部署格式量化,train-inference mismatch 从源头上消掉了。
10. 部署:vLLM Day-0 支持是怎么把上面这些落地的
技术报告讲的是 Moonshot 自己的生产系统,对大多数人来说,真正能上手的是 vLLM 的 Day-0 支持(vLLM 官方博客)。这一节就讲讲前面那些设计在开源推理引擎里是怎么落地的,以及想自己跑的话需要什么。
最简启动命令长这样:
vllm serve moonshotai/Kimi-K3 \
--tensor-parallel-size 8 \
--trust-remote-code \
--load-format fastsafetensors \
--enable-prefix-caching \
--enable-auto-tool-choice \
--tool-call-parser kimi_k3 \
--reasoning-parser kimi_k3
最低门槛是一个 8×B300 节点(或 16×B200,AMD MI355X 也支持)。几个值得展开的点:
混合 cache 管理器落地了第 9 节的设计。一个 scheduler 下面管两种内存:full-attention 层的分页 KV block + KDA 层的固定大小循环状态 block。prefill 走 FlashKDA(Moonshot 开源的 CUTLASS 实现),decode 走融合 CUDA kernel。有个细节:K3 的 prefix caching 目前默认是关的,要显式传 --enable-prefix-caching——第 9 节那套”KDA checkpoint + MLA hash block”的机制还在快速迭代。
DSpark speculative decoding:单用户 118 → 370 tok/s(3.14×)。K3 自带的 MTP 层之外,Inferact 训练并开源了一个 DSpark draft 模型(block-diffusion 架构,一次并行生成多个 draft token,drafting 开销不随深度增长;还有个 confidence head 预测每个 draft 的接受率)。实测 coding 类低熵任务每步接受约 4.73 个 token,创意写作类高熵任务约 2.61 个。开起来就是加一行:
--speculative-config '{"model":"Inferact/Kimi-K3-DSpark","method":"dspark","num_speculative_tokens":7,"attention_backend":"FLASHINFER_MLA","draft_sample_method":"probabilistic","rejection_sample_method":"block"}'
MoE 权重原生跑 MXFP4。第 9 节说的量化感知训练在这里兑现:MoE 路径的权重直接以 MXFP4 执行。MoE backend 按拓扑选——TP 场景用 TRT-LLM-Gen,专家并行/分离式部署(DEP)用 MegaMoE,可选 EPLB(Expert-Parallel Load Balancing,专家并行负载均衡)让各 rank 计算量接近。
TEP prefill 的 sequence parallelism。prefill 用 attention TP + MoE EP 组合(TEP),比纯 TP 通信少、专家 GEMM 形状更规整。但 naive 实现每层要两次 all-reduce,每个 rank 都得物化完整 batch。vLLM 的解法是把 attention 输出的 all-reduce 换成 reduce-scatter,让每个 rank 只持有自己那份 token 分片——AttnRes 的跨层残差状态全程保持分片状态(这对 K3 尤其重要,AttnRes 让残差流变成了有独立计算和显存开销的持久跨层状态)。NCCL 的 reduce-scatter/all-gather 对 prefill 消息尺寸不友好,vLLM 自定义 kernel 比 NCCL 快 1.7×–4.5×。
P/D 分离部署。高吞吐场景把 prefill 和 decode 拆到不同机器(各自按自己的瓶颈配硬件),已验证的拓扑是 TEP8 prefill → DEP16 decode,KV 走 NIXL 传输。对混合模型这是道难题:MLA 的分页 KV、KDA 的循环状态、block table 三样都得正确到达。NIXL 把共享 cache page 拆成两个逻辑视图分别传,异构 TP 下还要维护逻辑-物理块映射、清零未传输的尾部区域防止脏数据泄漏。
Agent 时代的 cache 保留策略。KDA 状态不随长度增长(一层的状态约等于几千 token 的 MLA cache),这对几十万 token 的 agent 会话是巨大优势;但 checkpoint 不能每 token 都存(单个 KDA checkpoint 远大于一个 token 的 MLA)。vLLM 给了两个互补策略:
- Interval-based retention:固定间隔存 checkpoint(如每 32K token 一个),prompt 边界自动保留——多轮对话的下一轮通常从重放上一轮 prompt 开始,这个位置复用率最高。
VLLM_PREFIX_CACHE_RETENTION_INTERVAL可调,设 0 就只留 prompt 末尾(纯多轮对话场景的好选择) - Marconi-style selective retention(MLSys ‘25):规则很简单——第二次被命中才缓存。第一次出现只说明这个前缀存在,第二次出现才证明它是共享的。一次性的前缀不占 cache,反复出现的前缀自动升级,用户不用预测哪些前缀会变热
一些性能优化的数字,感受下开源社区的迭代速度:KDA metadata builder 从复用 GDN 的通用实现换成专用融合 Triton kernel,batch size 1 时准备延迟从 870μs 降到 34μs(-96%),DSpark 端到端延迟降 6%;低延迟 BF16 GEMM 绕过 Tensor Core 的 TMA setup 用 CUDA Core FMA 直做,kernel 快 8%-100%,小 batch 端到端约 -10%;LatentMoE 尾部融合(reduce-scatter 共享专家 + 列并行 up-projection + broadcast all-gather,替代两次 all-reduce)这一步延迟 -20%,端到端 7-8%。还有个开源协作的佳话:Moonshot 先开源 FlashKDA,社区开发者 Shikhar Mishra 一天内为 H100 优化出 Flash-Flash-KDA,vLLM 隔天就在 GB300 NVL72 上验证合入。
精度方面:vLLM 走 OpenAI 兼容 endpoint 全量评测通过,max reasoning effort 下 GSM8K 0.976、GPQA-Diamond 0.939、OCRBench 0.889、MMMU Pro Vision 0.818。一个实用提醒:K3 回答前思考量很大,评测分低往往是答案被截断而不是答错——把 max_tokens 给足、先查截断再查别的。
11. 结语
整份报告读下来,最大的感受是架构和系统是咬合着设计的,没有一处改动是孤立的:
- lower-bounded decay,一行公式,换来一个无慢路径的全 Tensor Core kernel(3.5)
- AttnRes 的 block 结构,深度轴上加了 attention,训练显存和 serving kernel 反而都有对应的简化(5、7.2)
- KDA 的固定状态,既是长上下文的算法优势,又逼着重写了一遍 prefix cache、speculative decoding、context parallelism(7.3、8、9)
- 连激活函数换成 SiTU-GLU 都有低精度量化的考量(6.1)
单独看每一处都是”小改动”,拼起来就是 2.5× 的 scaling 效率。对想做 infra 的同学,我建议按 7.1(MoonEP)→ 9(prefix cache)→ 3.4-3.5(KDA kernel)的顺序精读:第一块是所有大模型训练系统的通用功,后两块是线性注意力这种新架构带来的全新问题,比较有前瞻性。
顺带一提,MoE 稀疏化、MXFP4 这类低精度量化,本质上都属于 LLM 压缩这个大方向。我们之前把这个方向的方法论积累整理成了《动手学 AutoML:从 NAS 到大语言模型优化实战》,里面有 LLM 压缩(剪枝、量化)和后训练剪枝实战的章节。K3 这份报告讲的是系统层面的优化,和书里的模型压缩视角是同一个目标(让大模型跑得动、跑得便宜)下的不同角度,感兴趣可以搭配着看。
欢迎评论区交流,指出问题。
