arXiv'26 | Kimi K3 的 KDA:1M context 不必把所有 KV 都留下
arXiv’26 | Kimi K3 的 KDA:1M context 不必把所有 KV 都留下
这是 Kimi K3 系列的第一篇,专门讲序列方向的改动:KDA(Kimi Delta Attention)和它怎样在 GPU 上运行。上一篇系列导读先把全局架构放在一起看;下一篇会讲深度方向的 AttnRes 和宽度方向的 Stable LatentMoE。
1. 1M context 的账单到底是什么
模型支持 1M context,听起来像是把“记忆容量”拉大了一百万个 token。但生成阶段真正要付的账单不是只有容量,还有每一步的读取。
Transformer 生成一个新 token 时,会把当前 token 变成 query,然后和历史 token 的 key 做匹配,再用匹配权重聚合 value。KV cache 保存了这些历史 key/value,避免重复计算;可每生成一个新 token,还是要把历史 KV 读出来。context 越长,单步读取越多。
这件事可以用一个很小的例子说明。假设历史有 8 个 token,当前 query 只真正需要其中 2 个:
历史位置: 1 2 3 4 5 6 7 8
相关性: 0.1 0.2 0.8 0.1 0.1 0.7 0.1 0.1
标准 attention 会读取 8 个位置,算完以后才发现 3 和 6 的权重最高。如果上下文是 1M,这个“先把全部读进来再决定谁重要”的过程会成为带宽瓶颈。
MLA(Multi-head Latent Attention)已经把每个历史位置的 KV 压缩成 latent 表示,降低了“每个位置有多大”的成本,但没有改变“要访问多少个位置”。K3 的 KDA 进一步问:能不能用一个更便宜的状态,把大部分历史压缩掉,而只留少数全局 attention 层?
K3 的 Hybrid Attention 采用 3:1 的比例:93 层中 69 层使用 KDA,24 层使用 MLA。KDA 负责让大多数层的状态大小不再随序列长度增长,MLA 保留少量全局读取能力,避免所有信息都被压进一个有限状态以后无法恢复。
2. 从 attention 到线性 attention:先把历史写成一张表
为了理解 KDA,先暂时去掉 softmax,只看最简单的线性 attention。当前 query、历史 key/value 分别记作 $\bm q_t$、$\bm k_j$、$\bm v_j$,输出是:
\[\bm o_t=\sum_{j\le t}(\bm q_t^\top\bm k_j)\bm v_j\]利用乘法结合律,可以先把所有历史的 $\bm k_j\bm v_j^\top$ 加起来,再用 query 去查:
\[\bm o_t=\bm q_t^\top\underbrace{\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_t$ 是一个 $[d_k,d_v]$ 的矩阵,可以把它想成一张固定大小的联想记忆表。每个新 token 到来时,把自己的 key/value 写进表里;query 到来时,用自己的 key 方向去查表。
这一步解决了 KV cache 随长度增长的问题:无论历史有 1K 还是 1M,状态 $\mathbf S$ 的形状都不变。可是它也带来一个明显缺陷:状态只增不减。
如果所有历史都无条件相加,早期的旧信息会和新信息互相干扰。更麻烦的是,三个月前的一条普通描述和刚刚出现的关键结论,在状态里没有自然的优先级。线性 attention 需要一种“记忆管理”机制,而不是一张永远只写不擦的表。
3. KDA:遗忘、擦除,再写入
最简单的记忆管理是给旧状态乘一个遗忘因子:
\[\mathbf S_t=\alpha_t\mathbf S_{t-1}+\bm k_t\bm v_t^\top\]这类 gated linear attention 会让旧信息逐步衰减,但整体打折仍然比较粗糙。新 token 如果和某条旧记忆使用的是相同的 key 方向,更合理的做法是把冲突的旧内容擦掉,再写入新内容;不冲突的方向则尽量保留。
KDA 的 delta rule 可以从一个很朴素的目标推出来:当前状态在被 $\bm k_t$ 查询时,应该返回新 token 的内容 $\bm v_t$。因此希望最小化:
\[\mathcal L(\mathbf S)=\frac12\|\mathbf S\bm k_t-\bm v_t\|^2\]从旧状态 $\mathbf S_{t-1}$ 出发做一步梯度下降,步长为 $\beta_t$:
\[\begin{aligned} \mathbf S_t &=\mathbf S_{t-1}-\beta_t(\mathbf S_{t-1}\bm k_t-\bm v_t)\bm k_t^\top\\ &=(\mathbf I-\beta_t\bm k_t\bm k_t^\top)\mathbf S_{t-1} +\beta_t\bm k_t\bm v_t^\top \end{aligned}\]拆开看就很直观:$(\mathbf I-\beta_t\bm k_t\bm k_t^\top)$ 会削弱旧状态里和当前 key 同方向的部分,$\beta_t\bm k_t\bm v_t^\top$ 再把新内容写进去。这个“先纠错、再写入”的来源,比把公式当作一个黑盒记住更重要。
KDA 再加上逐通道遗忘因子 $\bm\alpha_t$,得到论文中的核心更新:
\[\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\]读操作则是:
\[\tilde{\bm o}_t=\mathbf S_t^\top\bm q_t\]所以一次 KDA 更新的顺序是:
- 每个通道按自己的 $\alpha$ 给旧状态打折;
- 找出旧状态在当前 key 方向上的应答;
- 擦掉这条方向上的冲突部分;
- 用当前 value 写入新记忆;
- 用 query 从更新后的状态读出结果。
这就是 KDA 所谓的“固定大小状态”:它不是把历史简单平均,而是维护一张会遗忘、会覆盖的联想表。
4. 先写一个能跑的 KDA head
下面是一个教学版单 head 实现。它省略了 ShortConv、L2Norm 和输出门,但保留了 KDA 的状态更新。假设 batch size 是 $B$,序列长度是 $T$,输入维度是 $d_{model}$,状态大小是 $[d_k,d_v]$。
import torch
import torch.nn as nn
import torch.nn.functional as F
class KDAHead(nn.Module):
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)
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))
alpha = torch.sigmoid(self.w_alpha(x))
S = x.new_zeros(B, q.size(-1), v.size(-1)) # [B, d_k, d_v]
outputs = []
for t in range(T):
q_t, k_t, v_t = q[:, t], k[:, t], v[:, t]
Sa = alpha[:, t, :, None] * S
old = torch.einsum("bk,bkv->bv", k_t, Sa)
erase = torch.einsum("bk,bv->bkv", beta[:, t] * k_t, old)
write = torch.einsum("bk,bv->bkv", beta[:, t] * k_t, v_t)
S = Sa - erase + write
outputs.append(torch.einsum("bk,bkv->bv", q_t, S))
return torch.stack(outputs, dim=1) # [B, T, d_v]
读代码时可以把每一轮对应回上面的五步:Sa 是遗忘,old 是状态在当前 key 上的应答,erase 是擦除冲突,write 是写入,最后一个 einsum 是 query 读取状态。
这段代码的输入来自上一层的 hidden states,返回一个 [B,T,d_v] 的序列表示,交给后面的输出投影和残差连接。它最大的优点是状态 S 始终是 [B,d_k,d_v],不会因为 T 从 4K 变成 1M 而扩大。
但它还有一个 GPU 不喜欢的问题:for t in range(T)。每一步都要等上一步更新完 S,然后执行几个很小的矩阵运算。算法上没问题,硬件利用率却很差。
5. Chunkwise:同一个递推,换一种执行顺序
假设序列被切成每块 4 个 token:
chunk 0: [token 0, token 1, token 2, token 3]
chunk 1: [token 4, token 5, token 6, token 7]
chunk 2: [token 8, token 9, token 10, token 11]
每个 chunk 只从前一个 chunk 接收一个入口状态 $\mathbf S_{in}$。拿到入口状态以后,当前 chunk 内任意位置的输出都可以拆成:
token 6 的输出
├── 入口状态 S_in 经过 token 4、5、6 的遗忘后,被 q_6 读取
└── token 4、5、6 在当前 chunk 写入的内容,分别遗忘到 token 6 后被 q_6 读取
为了看清矩阵化,先忽略 KDA 的擦除项,只保留:
\[\mathbf S_i=\alpha_i\mathbf S_{i-1}+\bm k_i\bm v_i^\top\]展开 chunk 内第 $i$ 个状态:
\[\mathbf S_i= \underbrace{\left(\prod_{r=1}^{i}\alpha_r\right)\mathbf S_{in}}_{\text{chunk 之前的记忆}} + \underbrace{\sum_{j=1}^{i}\left(\prod_{r=j+1}^{i}\alpha_r\right)\bm k_j\bm v_j^\top}_{\text{当前 chunk 的写入}}\]假设 chunk 内 token 2 的衰减因子分别为 $\alpha_1,\alpha_2$,那么:
\[\mathbf S_2=\alpha_2\alpha_1\mathbf S_{in}+\alpha_2\bm k_1\bm v_1^\top+\bm k_2\bm v_2^\top\]第一项是旧记忆,第二项是 token 1 留下的内容,第三项是 token 2 刚写入的内容。token 3 不能出现在 token 2 的状态里,这就是 causal 约束。
接着定义累计衰减:
\[\Gamma_i=\prod_{r=1}^{i}\alpha_r\]从位置 $j$ 写入、到位置 $i$ 读取的衰减就是 $\Gamma_i/\Gamma_j$。于是可以把衰减分别乘到 query 和 key 上:
\[\left(\bm q_i\Gamma_i\right)^\top\left(\bm k_j/\Gamma_j\right) =\bm q_i^\top\bm k_j\frac{\Gamma_i}{\Gamma_j}\]这意味着整个 chunk 的 query 可以堆成矩阵 $Q$,key 堆成 $K$,一次 $QK^\top$ 就得到所有 $(i,j)$ 对的贡献;再把上三角屏蔽掉,只保留 $j\le i$ 的 causal 部分。
# C=4;每一行代表一个 token
# Q, K: [C, d_k], V: [C, d_v], S_in: [d_k, d_v]
Gamma = alpha.cumprod(dim=0) # [C, d_k]
from_past = (Q * Gamma) @ S_in # [C, d_v]
within = torch.tril((Q * Gamma) @ (K / Gamma).T) @ V
O = from_past + within # [C, d_v]
from_past 是入口状态对当前 chunk 的贡献,within 是当前 chunk 内历史写入的贡献。两者都变成了矩阵乘法,GPU 可以把一个 chunk 的多个 token 一起交给 Tensor Core。
这里并没有把顺序打乱:chunk 1 仍然要等待 chunk 0 产出 $\mathbf S_{out}$。变化的是串行依赖的粒度:原来 $T$ 个 token 需要 $T$ 次状态传递,现在大约只需要 $T/C$ 次 chunk 传递,chunk 内部则批量并行。
真实 KDA 还包含 delta rule 的擦除项,不能简单删掉。工程实现用 UT transform 等块内变换,把擦除、遗忘和写入一起整理成类似的 causal block 运算。上面的简化推导只负责解释为什么递推可以重排,不代表实现真的丢掉了擦除机制。
6. Lower-bounded decay:为什么要给遗忘设下限
Chunkwise 里出现了 K / Gamma。如果每一步的遗忘因子都小于 1,连续相乘以后 $\Gamma$ 会越来越小;它的倒数可能非常大。低精度训练使用 BF16 时,这种除法容易溢出,数值不稳定会直接变成训练或 kernel 的问题。
Kimi Linear 使用的负 Softplus 映射允许 log-decay 趋近 $-\infty$。K3 把 log-decay 限制在一个有限区间:
\[\bm g_t^h=g_{min}\cdot\operatorname{Sigmoid}(e^{A_h}\bm z_t^h), \qquad g_{min}=-5\]于是每个通道的遗忘因子不会低于 $e^{-5}\approx6.7\times10^{-3}$。对一个 16-token tile,累计 log-decay 被控制在大约 $(-80,0)$,倒数的范围也被限制住。
这个改动的价值不只是“训练不溢出”。数值范围受控以后,原本需要逐位置处理的对角 tile 也可以使用 dense Tensor Core 矩阵乘法,kernel 不必为对角区域保留一条慢路径。算法上只是给遗忘设了一个下限,系统上却换来了统一的 dense tile 执行。

上图左侧对比了无界的负 Softplus 和 K3 的 bounded sigmoid;右侧展示了为什么 bounded decay 能让所有 causal tile 走矩阵乘法,而不是只在非对角区域快。
7. Gated MLA:为什么不能所有层都用 KDA
KDA 的状态固定大小,但它毕竟是压缩记忆:如果某个任务需要精确访问很久以前的一段原文,有限状态可能不如直接读取原始 KV。K3 因此保留 24 层 MLA,让模型在部分深度上仍然可以做全局 attention。
K3 对 MLA 做了两个配合性改动:
- MLA 层使用 NoPE,不再单独依赖 RoPE 的长上下文外推;KDA 的衰减提供了一部分顺序和远近信号,MLA 更专注于内容检索。
- MLA 使用与 KDA 类似的输出门,让每个 token 通过输入相关的 gate 调制全局 attention 的输出通道。
因此 Hybrid Attention 不是“69 层便宜层加 24 层昂贵层”的简单拼接,而是两个记忆系统分工:KDA 负责高频、低成本的状态更新,MLA 负责少量需要精确全局访问的场景。比例和层位置是模型设计与硬件成本之间的折中,不能直接推广成所有模型都应该采用 3:1。
8. 实验应该怎样读
论文在 H800 集群上比较 Kimi K2 和 K3 的推理 cost per million tokens。prefill 是一次性读入 prompt 的阶段,decode 是逐 token 生成的阶段。随着 token position 接近 128K,K3 的成本曲线增长更慢,decode 阶段的差距尤其明显。
这里要注意两个边界:第一,这个结果测的是完整系统,不是只测 KDA 状态更新;第二,indexer、chunkwise kernel、MLA 层比例和部署配置都会影响绝对数值。KDA 的核心结论是主状态不再随着长度保存完整 KV,但端到端收益仍然需要硬件实现来兑现。
9. 小结
KDA 的设计可以按一条因果链记住:
- 全量 KV cache 在超长 context 下带来持续增长的读取成本;
- 线性 attention 用固定矩阵状态代替历史,但只写不擦会产生干扰;
- delta rule 让写入变成“按 key 方向纠错”,再配合逐通道遗忘;
- 逐 token 更新无法喂饱 GPU,于是把递推重排成 chunk 间传状态、chunk 内矩阵并行;
- 衰减连乘会造成数值问题,K3 用 lower-bounded decay 换取统一的 Tensor Core kernel;
- 少量 Gated MLA 保留精确的全局访问,补上固定状态的记忆边界。
下一篇继续沿着“为什么会这样设计”的思路,讲 K3 的 AttnRes 和 Stable LatentMoE:一个处理层与层之间的信息流,一个处理 896 个专家如何稳定训练。
扯一句题外话:KDA 是推理系统层面的新 attention 结构,书里不会替你讲完 chunkwise kernel;《动手学 AutoML:从 NAS 到大语言模型优化实战》第 8、11 章整理了 LLM 剪枝、量化和后训练压缩,切入点不同,但同样是在模型能力和计算/存储成本之间找折中。
