Hybrid KDA–MLA:Kimi K3 注意力机制解析
研究了下 K3 的注意力机制,不是把所有 attention 都换成某种线性 attention,而是用了一个很工程化的组合:
3 层 KDA + 1 层 Gated MLA
报告里的整体配置大致是:
69 层 KDA + 24 层 Gated MLA
这个比例很有意思。它背后其实是一个取舍:
KDA 负责便宜地处理超长上下文
MLA 负责周期性地保留全局 softmax attention 能力
如果只用普通 attention,1M token 的 KV cache 和计算成本会非常重。如果只用线性 attention,又可能损失一些精确的 token-to-token 全局交互。Kimi K3 的路线是:大部分层用 KDA 扛长上下文成本,少数层插入 MLA 做全局信息交换。
这篇只讲这件事。
先从普通 attention 的问题说起
普通 Transformer attention 做的是:
当前 token 的 query,去看历史所有 token 的 key/value
公式上是:
$$ \mathrm{Attention}(Q, K, V) = \mathrm{softmax} \left( \frac{QK^\top}{\sqrt{d}} \right)V $$
它的好处是表达力强。每个 token 都可以精确地和任意历史 token 建立关系。
但代价也很明显:
上下文越长,历史 K/V 越多
每生成一个 token,都要读越来越长的 KV cache
如果上下文是几千、几万 token,这还能扛。如果是 1M token,就麻烦了。
想象你在做一个 coding agent 任务:
系统提示 + 工具说明 + repo 文件 + grep 结果 + 测试日志 + 多轮修复记录
这些东西很容易堆成长上下文。普通 attention 要把每个历史 token 的 K/V 都留下来,后续每一步都要读它们。模型还没开始认真思考,显存和带宽先开始尖叫。
所以 Kimi K3 需要新的 attention 结构。
KDA:把历史压进一个会遗忘的状态
KDA,全称 Kimi Delta Attention。可以先别把它想成 attention,而是想成一本“会遗忘、会修正”的笔记本。
普通 attention 是:
保存所有历史 token
当前 token 去翻所有历史
KDA 是:
不保存所有历史 token
而是维护一个固定大小的状态 S
每来一个 token,就更新 S
当前 token 只读取 S
这个状态是:
$$ S_t \in \mathbb{R}^{d_k \times d_v} $$
你可以把它理解成模型当前对历史上下文的压缩记忆。
每来一个 token,KDA 大致做三步:
1. 先遗忘一部分旧状态
2. 看旧状态对当前 key 已经能预测出什么
3. 只把预测错的部分写进去
核心更新可以写成:
$$ \bar{S}t = \operatorname{Diag}(\alpha_t) S $$
$$ S_t = \bar{S}_t + \beta_t k_t \left( v_t - \bar{S}_t^\top k_t \right)^\top $$
输出是:
$$ \tilde{o}_t = S_t^\top q_t $$
这里几个量的含义是:
alpha_t:每个 channel 保留多少旧状态
beta_t:当前 token 写入力度
k_t:当前信息的地址
v_t:当前信息的内容
q_t:当前 token 读取状态时用的 query
最关键的是这一项:
$$ v_t - \bar{S}_t^\top k_t $$
它是 prediction error。
也就是说,KDA 不是简单把新信息硬塞进状态,而是先问:
旧状态用当前 key 查出来,已经能预测出什么 value?
如果预测得很准,那就少写一点。如果预测得不准,就把误差写进去。
这就是 delta rule 的味道。
举个例子。假设模型在读代码:
user = get_user()
print(user.name)
前面状态里可能已经记住了:
这里和 user 有关
后面又看到一段异常日志:
AttributeError: 'NoneType' object has no attribute 'name'
KDA 不需要把整段日志完整背下来,而是会把“user 可能是 None”这种对现有状态有增量的信息写进去。下一次再遇到相关 query,它就能从状态里读到这部分信息。
这就是 KDA 适合长上下文的地方:历史不是一长串 KV,而是持续压缩进状态。
KDA 和普通 linear attention 有什么区别?
这块很容易混在一起。KDA 属于 linear attention / recurrent attention 的大方向,但它不是最朴素的 linear attention。
普通 linear attention:通常是累加式记忆
很多普通 linear attention 可以粗略理解成:
$$ S_t = S_{t-1} + \phi(k_t)v_t^\top $$
输出类似:
$$ o_t = S_t^\top \phi(q_t) $$
它的思路是:
把历史 key/value 映射后累加进一个状态
当前 query 去读这个累加状态
这种结构很简单,也很快。但它有几个问题:
1. 历史信息不断累加,旧信息不容易主动忘掉
2. 新信息容易覆盖/污染旧状态
3. 状态更新通常是加法,表达能力有限
当然,不同 linear attention 变体会有 normalization、decay 或 gating,但最基础的版本可以先这样理解。
KDA:不是简单累加,而是“遗忘 + 误差修正”
KDA 的更新不是:
$$ S_t = S_{t-1} + k_t v_t^\top $$
而是:
$$ S_t = \bar{S}_t + \beta_t k_t \left( v_t - \bar{S}_t^\top k_t \right)^\top $$
差别有两个。
第一,KDA 有 channel-wise forgetting:
$$ \bar{S}t = \operatorname{Diag}(\alpha_t) S $$
每个 key channel 都有自己的保留系数:
$$ \alpha_t \in (0, 1)^{d_k} $$
所以它不是整块状态一起忘,而是不同维度以不同速度遗忘。
第二,KDA 是 delta-rule update。
它会先算:
$$ \mathrm{pred}_t = \bar{S}_t^\top k_t $$
然后只写入误差:
$$ \mathrm{error}_t = v_t - \mathrm{pred}_t $$
所以:
$$ S_t = \bar{S}_t + \beta_t k_t \mathrm{error}_t^\top $$
这比简单累加更像一个在线学习过程:
旧状态已经知道的,不重复写
旧状态不知道的,用误差修正
KDA 还有一个很工程化的 decay 设计
Kimi K3 里,遗忘系数不是随便预测的。它先预测 log-decay:
$$ g_t = g_{\min} \cdot \mathrm{Sigmoid}(e^A z_t) $$
然后:
$$ \alpha_t = \exp(g_t) $$
报告里:
$$ g_{\min} = -5 $$
所以:
$$ \alpha_t \in (e^{-5}, 1) $$
这意味着每一步不能忘得太狠。
这不是为了好看,而是为了 kernel。KDA 在 chunk 内并行时会用到连续遗忘系数的乘积:
$$ \gamma = \prod_t \alpha_t $$
也会用到:
$$ \frac{1}{\gamma} $$
如果 (\alpha_t) 太小,(\gamma) 会极小,(\frac{1}{\gamma}) 会数值爆炸。K3 给 decay 加下界后,一个 16-token tile 内的数值范围可控,很多原本需要特殊处理的 diagonal tile 也能走 Tensor Core GEMM。
这就是 KDA 和普通 linear attention 的一个重要区别:
KDA 不只是换了个线性公式
它的 recurrence、decay、数值范围和 kernel 是一起设计的
KDA kernel:怎么把递归做快?
KDA 公式看起来还是递归的:
$$ S_t = f(S_{t-1}, x_t) $$
如果逐 token 做:
for token in tokens:
S = update(S, token)
GPU 会很不开心。因为 GPU 喜欢大块矩阵乘,不喜欢一个 token 一个 token 串行。
所以 KDA 的 kernel 会把序列切成 chunk:
chunk 1 -> chunk 2 -> chunk 3 -> ...
chunk 之间传递状态:
$$ S_{\mathrm{out}}^{(i)} \rightarrow S_{\mathrm{in}}^{(i+1)} $$
但 chunk 内部尽量并行。
chunk 内输出可以拆成两部分:
$$ O = (\Gamma \odot Q) S_{\mathrm{in}} + A \tilde{V} $$
第一项:
$$ (\Gamma \odot Q) S_{\mathrm{in}} $$
表示当前 chunk 读取前面 chunk 传进来的历史状态。
第二项:
$$ A \tilde{V} $$
表示当前 chunk 内部 token 之间的因果交互。
其中:
$$ A = \operatorname{Tril} \left[ (Q \odot \Gamma) (K / \Gamma)^\top \right] $$
(\operatorname{Tril}) 表示只保留下三角,保证 causal。
你不用死盯这个公式,只需要抓住一点:
KDA 把 chunk 内的递归关系改写成矩阵乘,让 Tensor Core 可以跑起来。
FlashKDA 做的事情,就是把 chunk 内矩阵计算和 chunk 间状态传播重叠起来,减少 GPU 空转。
所以它不是魔法。它的工程思路是:
原始:token -> token -> token,太串行
优化:chunk -> chunk,chunk 内用 GEMM 并行
这也是为什么 lower-bounded decay 很重要:它让 chunk 内的矩阵计算数值稳定,可以更充分地使用 dense GEMM。
MLA:压缩 KV cache 的全局 attention
KDA 很适合长上下文,但它毕竟不是完整的 softmax attention。为了保留周期性的全局 token-to-token 交互,Kimi K3 插入了 Gated MLA。
MLA,全称 Multi-head Latent Attention。它的核心是:
仍然做 softmax attention
但不缓存完整 K/V
而是缓存低维 latent
普通 attention 对每个 token 缓存:
K_t, V_t
MLA 缓存:
c_t = W_c x_t
这个 (c_t) 是一个低维 latent。后续需要 key/value 时,再从 latent 中投影或融合出各个 head 需要的信息。
可以用一个很土但好懂的比喻:
普通 attention:每个历史 token 都保存完整档案
MLA:每个历史 token 只保存一个压缩包
KDA:不保存每个历史 token,只维护一本总笔记
MLA 的优势是,它仍然保留了全局 attention 的能力。当前 token 还是能和历史 token 做 softmax 交互,只是历史 token 的缓存变小了。
这和 KDA 的区别很大:
KDA:历史 token 被压进状态 S,不再逐个查历史 token
MLA:历史 token 还在,只是每个 token 的 KV 被压缩了
Gated MLA 和 NoPE:K3 里的两个改动
Kimi K3 里的 MLA 不是裸 MLA,而是 Gated MLA。
MLA 的原始输出记作:
$$ \tilde{o}_t $$
K3 会加一个输入相关的 gate:
$$ y_t = W_o \left[ \mathrm{Sigmoid}(W_g x_t) \odot \mathrm{RMSNorm}(\tilde{o}_t) \right] $$
这个 gate 可以理解成通道级阀门:
哪些 attention 输出通道要放大
哪些要压下去
由当前输入 x_t 决定
另外,K3 的 MLA 使用 NoPE,也就是没有显式位置编码。
这听起来有点反直觉。普通 attention 通常需要 RoPE 这类位置编码,否则 token 顺序会变模糊。
K3 能这么做,是因为它不是纯 MLA。大量 KDA 层已经通过 recurrent decay 提供了位置感和近因偏置。于是 MLA 可以更专注于内容级全局交互。
分工大概是:
KDA:提供顺序感、近因记忆、长上下文状态压缩
MLA:提供周期性的全局内容检索
这也是 Hybrid KDA–MLA 的精髓。
为什么这个组合适合 agent?
Hybrid KDA–MLA 特别适合长程 agent 任务。
比如一个 coding agent 正在修 bug:
1. 读 README
2. 搜索相关函数
3. 打开多个源码文件
4. 修改代码
5. 跑测试
6. 测试失败
7. 继续读日志
8. 再修改
这个过程里,上下文会越来越长,而且信息类型很多:
代码结构
错误日志
之前做过的尝试
工具返回结果
用户约束
KDA 像一个持续更新的工作记忆。它不会把每个历史 token 都完整留下,而是把重要信息不断压进状态。
MLA 则像定期全局回看。它让模型仍然有机会在压缩记忆之外,做更精确的 token-to-token 交互。
所以这个组合不是为了论文结构好看,而是服务于真实任务:
上下文很长
任务步骤很多
需要记住历史
也需要偶尔精确回查
一个简短总结
Kimi K3 的 Hybrid KDA–MLA 可以概括成:
KDA:固定状态、可遗忘、delta 更新,适合 1M 长上下文
MLA:压缩 KV cache 的 softmax attention,保留全局交互
Hybrid:多数层用 KDA 省成本,少数层用 MLA 补表达力
KDA 和普通 linear attention 的关键区别是:
普通 linear attention 更像累加式记忆
KDA 是 channel-wise forgetting + delta-rule error correction
再加上 lower-bounded decay 和 FlashKDA kernel,KDA 不只是一个公式改造,而是一套从模型结构到 GPU 执行都一起设计的长上下文方案。
这也是 Kimi K3 这篇技术报告里我最喜欢的地方:它没有假装“只要上下文拉长就行”,而是认真处理了长上下文模型最麻烦的部分——记忆怎么存、怎么忘、怎么读、怎么跑得动。