AI-Workshop 文献分享 2026-08-08 21:30

AI-Workshop 文献分享 #16 — KDA:给长文本建立可更新的记忆

ltw 零基础讲解 KDA(Key-Value Delta Attention)——以固定形状的状态矩阵为长文本建立可更新的记忆:从标准注意力的二次复杂度出发,介绍状态型线性注意力的写入与读取、细粒度门控与 Delta Rule 纠偏机制,并分析其近线性效率优势与有限状态压缩的局限。

AI-Workshop 文献分享 KDA 线性注意力 长上下文

AI-Workshop 文献分享 #16 — KDA:给长文本建立可更新的记忆

分享人:ltw | 日期:2026-08-08

主题:KDA(Key-Value Delta Attention)零基础详解——为长文本建立可更新的记忆


本讲义解决三个问题:

  1. KDA 怎样让模型更高效地处理长文本?
  2. MoE 怎样让模型拥有很大容量,却不用每次运行全部参数?
  3. Attention Residuals 怎样让很深的网络按需取回早期层信息?

三者分别优化三个不同方向:序列长度、模型宽度、网络深度


1. 先看全局:三个技术各管什么?

一段输入进入大模型后,大致会反复经过两类计算:

输入 token
   ↓
注意力模块:不同 token 之间交换信息
   ↓
前馈网络:对每个 token 进行更复杂的特征变换
   ↓
残差连接:把前后层的信息连接起来
   ↓
重复很多层
   ↓
输出下一个 token

KDA、MoE 和注意力残差并不是三个互相替代的方案,而是作用在不同位置:

KDA
解决“序列太长”的问题
优化 token 与历史信息之间的交互

MoE
解决“模型容量太大”的问题
优化前馈网络中的参数使用方式

Attention Residuals
解决“网络太深”的问题
优化不同网络层之间的信息传递

可以先记住一句话:

KDA 管时间方向,MoE 管专家方向,注意力残差管深度方向。

这里的“时间方向”指 token 按顺序到来;“专家方向”指同一层有许多可选专家;“深度方向”指模型从浅层一直计算到深层。


2. 必要基础:token、向量和 Transformer 层

2.1 token 是什么?

模型不会直接把一句话当成一个整体处理,而是先切分成 token。例如:

“小明今天去了北京”
        ↓
“小明” “今天” “去了” “北京”

实际分词结果由 tokenizer 决定,可能比上面的示例更细。

每个 token 会被转换成一个向量:

\[ x_t\in\mathbb{R}^{d} \]

其中:

  • \(t\) 表示 token 在序列中的位置;
  • \(d\) 表示隐藏维度;
  • \(x_t\) 不是一个人能直接阅读的词义表,而是一组由模型学习的数值特征。

2.2 一个 Transformer 层通常做什么?

教学上可以把一个标准 Transformer 层拆成:

输入表示
  ├── 注意力:让 token 读取其他 token 的信息
  ├── 残差:保留进入模块前的信息
  ├── 前馈网络:进一步变换每个 token
  └── 残差:再次保留原信息

简化公式是:

\[ h'_l=h_{l-1}+\operatorname{Attention}(h_{l-1}) \]
\[ h_l=h'_l+\operatorname{FFN}(h'_l) \]

KDA 主要改造注意力部分,MoE 主要改造 FFN 部分,Attention Residuals 主要改造层与层之间的残差通路。


第一部分:KDA——给长文本建立可更新的记忆

3. 标准注意力为什么越来越贵?

3.1 Q、K、V 的直观含义

标准注意力会从输入向量生成三组表示:

\[ Q=XW_Q,\qquad K=XW_K,\qquad V=XW_V \]

可以把它们理解成:

Q(Query):我现在想找什么?
K(Key):我这条信息适合被什么问题找到?
V(Value):如果找到我,应该取走什么内容?

例如历史中有一句“小明住在北京”:

  • Key 可能带有“小明、居住地”等地址特征;
  • Value 可能带有“北京”等内容特征;
  • 当后面出现“小明住在哪里”时,Query 会尽量匹配这个 Key。

3.2 全注意力怎样读取历史?

\(t\) 个 token 的输出可以简化为:

\[ y_t=\sum_{i=1}^{t}\operatorname{softmax} \left(\frac{q_t^Tk_i}{\sqrt{d_k}}\right)v_i \]

它的工作方式是:

  1. 保存每个历史 token 的 \(k_i\)\(v_i\)
  2. 当前 \(q_t\) 与所有历史 \(k_i\) 逐个比较;
  3. 根据匹配程度对所有 \(v_i\) 加权求和。

大白话理解:每问一个新问题,都把历史资料重新翻一遍。

3.3 二次复杂度从哪里来?

长度为 \(n\) 的序列需要建立大量 token 两两关系:

\[ QK^T\in\mathbb{R}^{n\times n} \]

因此关系计算通常包含:

\[ O(n^2d) \]

如果长度翻倍:

\[ (2n)^2=4n^2 \]

也就是说,长度增加两倍,核心关系计算大约增加到四倍。FlashAttention 可以大幅改善实际速度和显存占用,但不会从数学上消除所有 token 两两交互的二次规模。


4. 线性注意力的核心转变

全注意力的思路是:

保存每份历史原文
→ 查询时重新比较全部历史

状态型线性注意力换了一种思路:

历史到来时就整理进固定状态
→ 查询时直接读取整理后的状态

可以把固定状态记作:

\[ S_t\in\mathbb{R}^{d_k\times d_v} \]

这里 \(S_t\) 像一本由模型自己维护的“关联记忆本”。它不是人类语言摘要,而是一块数值矩阵。

最基础的写入和读取可以表示为:

\[ S_t=S_{t-1}+k_tv_t^T \]
\[ y_t=q_t^TS_t \]

其中:

  • \(k_t\):写入地址;
  • \(v_t\):写入内容;
  • \(k_tv_t^T\):把地址和内容建立关联;
  • \(q_t\):读取地址;
  • \(q_t^TS_t\):从记忆中读出匹配内容。

这使得历史长度增加时,不必让 KV 历史在 KDA 层中一直线性膨胀。


5. 为什么记忆状态是矩阵?

假设:

\[ k_t,q_t\in\mathbb{R}^{d_k},\qquad v_t\in\mathbb{R}^{d_v} \]

外积:

\[ k_tv_t^T \]

会产生一个 \(d_k\times d_v\) 的矩阵,正好能写入 \(S_t\)

5.1 一个二维地址例子

为了看清矩阵操作,假设地址只有两个维度:

\[ k_{\text{地点}}= \begin{bmatrix}1\\0\end{bmatrix}, \qquad k_{\text{职业}}= \begin{bmatrix}0\\1\end{bmatrix} \]

把地点“北京”和职业“医生”写入记忆:

\[ S= k_{\text{地点}}v_{\text{北京}}^T +k_{\text{职业}}v_{\text{医生}}^T \]

于是:

\[ S= \begin{bmatrix} v_{\text{北京}}^T\\ v_{\text{医生}}^T \end{bmatrix} \]

查询地点时:

\[ q_{\text{地点}}= \begin{bmatrix}1\\0\end{bmatrix} \]
\[ q_{\text{地点}}^TS=v_{\text{北京}}^T \]

这就是“Key 决定写到哪里,Value 决定写什么,Query 决定查哪里”。

5.2 真实模型不是固定字典

真实模型不会规定:

第 1 维永远代表地点
第 2 维永远代表职业

真实的 Key 和 Query 是稠密向量:

  • 一个向量可以同时包含人物、地点、时间等特征;
  • 相近含义可能落在相近方向;
  • 一次查询可能读取多个关联的软组合;
  • 地址空间由训练自动形成,而不是人工设计。

这带来了泛化能力,也带来了干扰问题:两个信息如果使用相近地址,写入同一块有限状态时可能互相影响。


6. 只会累加为什么不够?

最简单的状态更新是:

\[ S_t=S_{t-1}+k_tv_t^T \]

它只会增加内容,不会删除或修正内容。

假设模型先读到:

小明住在北京。

后来又读到:

小明已经搬到上海。

如果状态只会累加,就可能得到:

小明居住地 → 北京 + 上海

模型需要的不是盲目追加,而是:

  1. 判断旧信息哪些还应该保留;
  2. 查询旧记忆在当前地址会返回什么;
  3. 计算旧答案与新答案的差值;
  4. 只写入需要修正的部分。

KDA 的门控和 Delta Rule 就在解决这个问题。


7. KDA 的门控:决定忘多少、写多重

教学化简可以写成:

\[ S_t=A_tS_{t-1}+\beta_tk_tv_t^T \]

其中:

A_t:旧记忆保留多少
β_t:当前新内容写入多强

7.1 为什么需要遗忘门?

并不是所有旧信息都应该永久保留:

  • 人物姓名可能长期有效;
  • “当前所在位置”可能频繁变化;
  • 一段临时推理可能完成后就不再重要;
  • 旧版本代码信息可能被新修改覆盖。

如果模型永不遗忘,有限状态会不断受到旧内容干扰。

7.2 细粒度门控是什么意思?

粗粒度门控像一个总开关:

整本记忆统一保留 80%

更细粒度的门控像一排旋钮:

通道 A:保留 95%
通道 B:保留 20%
通道 C:保留 70%
通道 D:保留 5%

这些通道并没有人工名称。模型通过训练学习哪些数值方向该长期保留,哪些应该快速衰减。

KDA 的重要特点之一,就是比单一标量遗忘门拥有更细的通道级控制。实际 KDA 的状态转移比这里的 \(A_t\) 更具体,教学上先把复杂的衰减和状态转移统一放进 \(A_t\)


8. Delta Rule:不是重复抄写,而是纠正旧答案

8.1 先读取旧答案

按照当前 Key 查询旧状态:

\[ \hat v_t=S_{t-1}^Tk_t \]

\(\hat v_t\) 表示旧记忆在当前地址原本会返回的内容。

8.2 再计算误差

新内容是 \(v_t\),旧答案是 \(\hat v_t\),差值为:

\[ \Delta v_t=v_t-\hat v_t \]

8.3 只写入纠偏量

教学化简后的更新为:

\[ S_t=A_tS_{t-1} +\beta_tk_t(\Delta v_t)^T \]

展开得到:

\[ S_t=A_tS_{t-1} +\beta_tk_t \left(v_t-S_{t-1}^Tk_t\right)^T \]

大白话翻译:

先保留仍然有用的旧记忆
+
看看旧记忆在当前地址答了什么
+
计算它距离新答案还差多少
+
只把这个差值写回去

8.4 一维数字例子

旧状态返回:

\[ \hat v_t=0.7 \]

新目标是:

\[ v_t=0.9 \]

误差为:

\[ \Delta v_t=0.9-0.7=0.2 \]

如果写入强度为:

\[ \beta_t=0.5 \]

那么更新后大约是:

\[ 0.7+0.5\times0.2=0.8 \]

这不是把 \(0.9\) 再叠加一次,而是朝正确答案移动一半距离。

8.5 Delta Rule 有什么好处?

  • 旧答案已经正确时,误差很小,不必重复写入;
  • 旧答案过时时,可以朝新答案纠偏;
  • 比无条件累加更不容易让状态无限堆积;
  • 可以学习不断变化的关联关系。

但它也不是完美数据库。有限状态中的地址仍可能冲突,压缩仍可能损失逐字细节。

注意:本节公式用于建立直觉,不是 KDA 完整实现的逐行复现。真实实现还包含归一化、具体门控参数化、多头结构和高效块算法。


9. KDA 一次处理 token 的完整数据流

当前 token 表示 x_t
        │
        ├── 生成 q_t:现在想查什么
        ├── 生成 k_t:当前内容写到什么地址
        ├── 生成 v_t:当前内容是什么
        ├── 生成 A_t:旧状态如何衰减
        └── 生成 β_t:纠偏写入多强
                 │
                 ↓
用 k_t 查询旧状态的旧答案
                 │
                 ↓
计算 v_t - 旧答案
                 │
                 ↓
衰减旧状态并写入差量
                 │
                 ↓
得到新状态 S_t
                 │
                 ↓
用 q_t 读取 S_t
                 │
                 ↓
得到当前输出 y_t

记忆口诀:

q:想查什么
k:写到哪里
v:写入什么
A:旧的留多少
β:这次改多大
S:整块记忆状态
y:最终读出的内容

10. KDA 为什么能更高效?

10.1 序列长度方向近似线性增长

全注意力需要比较大量 token 对:

\[ O(n^2d) \]

状态型注意力逐 token 更新固定形状状态,长度相关部分通常近似:

\[ O(nd_kd_v) \]

这里不能机械地比较两个公式中的常数,因为具体维度、头数、GPU Kernel 和混合层比例都会影响真实性能。关键区别是:

全注意力:长度 n 出现在二次项中
状态型注意力:长度 n 主要出现在一次项中

10.2 自回归推理时不需要无限增长的 KDA KV 历史

标准注意力为了生成下一个 token,通常要保留历史 Key 和 Value。上下文越长,KV Cache 越大。

KDA 层主要维护固定形状状态 \(S_t\)

读入新 token
→ 更新 S_t
→ 丢弃这一步不再需要的中间量

因此 KDA 层的主要记忆状态不会随历史 token 数量无限增长。

10.3 递推公式为什么还能适合 GPU?

公式中 \(S_t\) 依赖 \(S_{t-1}\),看起来必须逐 token 串行。实际工程会使用分块算法:

逻辑上:token 状态按顺序传递
工程上:序列切成 chunk,块内尽量并行,块间传递汇总状态

这就像统计全书累计数据:可以先并行统计每一章,再按顺序合并章节状态。

KDA 使用专门设计的状态转移和 chunkwise 算法,让递推表达与 GPU 块状计算尽量兼容。高效 Kernel,例如 FlashKDA,则负责把数学优势转化为实际速度。


11. KDA 的局限:压缩不是无损保存

固定大小的状态必须把越来越长的历史压缩进去,因此存在明确取舍:

KDA 更擅长

  • 持续追踪变化的状态;
  • 维护常见关联;
  • 汇总长文本中的模式;
  • 以较稳定的状态成本处理更长序列。

KDA 可能较弱

  • 逐字复制很早以前的一段罕见文本;
  • 区分大量地址极为相近的细节;
  • 无损保留每个历史 token;
  • 精确比较任意两个远距离 token。

所以更合理的架构不是“永远只用 KDA”,而是混合使用:

多数层使用 KDA
负责高效连续记忆

周期性插入全注意力或 Gated MLA
负责精确 token 级全局交互

类比:

KDA = 平时整理和查询笔记
全注意力 = 必要时重新翻阅原始资料