发布于 

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 这篇技术报告里我最喜欢的地方:它没有假装“只要上下文拉长就行”,而是认真处理了长上下文模型最麻烦的部分——记忆怎么存、怎么忘、怎么读、怎么跑得动。