背景
接上文 《大语言模型的基石:Transformer 入坑笔记(三) - 注意力机制和 Transformer》。
通过之前的内容,我们大致了解了大语言模型的基石:Transformer 的基础原理。 但随之而来有一个问题:随着参数量和上下文长度的增长,计算量太大扛不住了。
现在很多大语言模型,经过了多轮迭代,终于陆陆续续能在参数量持续提高的前提下同时支撑 1M 上下文。
实际上各家LLM对长上下文的解决方案并不完全一致。最近 Kimi K3 发布,我自己体验下来效果相当好,跻身国际第一梯队。 在我实际的复杂服务端工程里,它对长链路、异步调用场景中的边界、时序、一致性问题的分析能力,能和 GPT 最新模型互补;甚至我个人感觉,它比 GPT 给出的分析报告还要详细完整。
而 Kimi 的方案,早先也通过论文 《Kimi Linear: An Expressive, Efficient Attention Architecture》 公开了。 所以接下来这篇,我们先从线性注意力的基础讲起,入门级地过一遍大致的方案、思路和原理;从基础一路优化到 Kimi 这篇 KDA 设计的路线,就留到后面单独写了。
整个过程涉及的论文比较多,我也是零散地抽时间看和理解,所以整个阅读周期拖得比较长。 特别是优化相关的内容,可能会有一些前后用词和表达方式跟着那段时间看的论文走了,没完全修订统一成一样的表达形式,还请见谅。
复杂度估算
我们先回顾一下 《Attention Is All You Need》 里的注意力计算公式:
$$\mathrm{Attention}(Q, K, V) = \mathrm{Softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$其中对于向量 $z = (z_1, z_2, z_3, ..., z_n)$,标准 Softmax 函数定义是 $\text{Softmax}(z_i) = \frac{exp(z_i)}{\sum_{j=1}^{n} exp(z_j)}$。 实现上为了防止数值溢出,会先减去最大值,数学上等价:$\text{Softmax}(z_i) = \frac{exp(z_i - \max(z))}{\sum_{j=1}^{n} exp(z_j - \max(z))}$。
这里的 $d_k$ 是 Key 向量的维度。除以 $\sqrt{d_k}$ 是为了把点积结果拉回更合适的数值范围,避免维度变大后 Softmax 太容易进入饱和区。$W^Q$、$W^K$、$W^V$ 这些矩阵也是训练得到的。
$$ \begin{aligned} X &\in \mathbb{R}^{N \times d_{model}} \\ W^Q &\in \mathbb{R}^{d_{model} \times d_k} \\ W^K &\in \mathbb{R}^{d_{model} \times d_k} \\ W^V &\in \mathbb{R}^{d_{model} \times d_v} \end{aligned} $$$$ \begin{aligned} Q &= XW^Q \in \mathbb{R}^{N \times d_k} \\ K &= XW^K \in \mathbb{R}^{N \times d_k} \\ V &= XW^V \in \mathbb{R}^{N \times d_v} \end{aligned} $$附注:前面几篇里用 L 表示输入 token 长度,本文为了贴合复杂度分析的惯用记法,统一用 N。
X 是输入矩阵,N 是输入 token 长度。所以单是计算 $QK^T \in \mathbb{R}^{N \times N}$,复杂度就是 $O(N^2d)$。更确切的时间复杂度如下:
| 步骤 | 主要算术量 | 对 $N$ 的增长 |
|---|---|---|
| Q/K/V 投影 | $O(3Nd^2)$ | 线性 |
| $QK^T$ | $O(N^2d)$ | 平方 |
| Softmax | $O(hN^2)$ | 平方,但通常不是算术量最大的部分 |
| 注意力权重乘 $V$ | $O(N^2d)$ | 平方 |
| 输出投影 | $O(Nd^2)$ | 线性 |
其中,$h d_k\approx h d_v\approx d$,单层的时间复杂度可以粗略写成 $O(3Nd^2+N^2(2d+h))$。当 $N\gg d$ 时,两个注意力矩阵乘法占大头(Softmax 本身是逐元素的指数运算加归一化,单个元素的开销是常数)。
而在空间复杂度方面,两个注意力矩阵和 Softmax 结果都要占 $O(N^2)$ 级别的显存。按单层、$h=16$、fp16 的 $hN^2$ 估算:
| 上下文长度 | 朴素 Softmax:$N\times N$ 矩阵 $O(hN^2)$ |
|---|---|
| 4K | 512 MiB |
| 64K | 128 GiB |
| 256K | 2 TiB |
显然,上下文稍微增大一点,时间和空间开销都会爆炸式增长,肯定扛不住。 于是就有了后面线性注意力这条路线。
高效注意力的基础
高效注意力最早来自视觉领域的论文 《Efficient Attention: Attention with Linear Complexities》。
先把点积注意力表达为矩阵形式:$D(Q,K,V) = \rho\left(QK^T\right)V$。归一化函数可以选:
- 缩放(Scaling):$\rho(Y) = \frac{Y}{n}$(直接除以位置数)
- Softmax:$\rho(Y) = \sigma_{\text{row}}(Y)$(对每行做 Softmax)
而高效注意力,首先各特征向量仍经三个线性层得到 Q、K、V;但不再把键看作 n 个 $d_k$ 维向量,而看作 $d_k$ 张单通道特征图——每张图作为对所有位置的一套权重,对值特征加权求和,得到一个全局上下文向量。之所以叫“全局”,是因为该向量不对应任何具体位置,而是对整幅输入特征的某种全局描述。scaling 版本与对应的点积注意力严格等价;factorized-Softmax 版本则是新的可分解算子。
高效注意力则可以表达成:$E(Q,K,V) = \rho_q(Q)\left(\rho_k(K)^T V\right)$。
其中 $\rho_q$、$\rho_k$ 分别为查询与键的归一化函数。与点积注意力相同的两种归一化实现为:
- 缩放:$\rho_q(Y)=\rho_k(Y)=\frac{Y}{\sqrt{n}}$
- Softmax:$\rho_q(Y)=\sigma_{\text{row}}(Y)$,$\rho_k(Y)=\sigma_{\text{col}}(Y)$(Q 沿行、K 沿列分别做 Softmax)
原版先算 $N\times N$ 的 $QK^T$,再用它混合 V;新版先用 K 把 V 压成 $d_k$ 份全局摘要($K^TV$,尺寸为 $d_k\times d_v$),再让每个位置用 Q 组合这些摘要。结合律只保证 scaling 版本严格等价。factorized-Softmax 分别对 Q 的特征维、K 的位置维做 Softmax,得到的是新的低秩注意力矩阵。
Scaling 下严格等价:$D(Q,K,V) = \frac{QK^T}{n}V = \frac{Q}{\sqrt{n}}\left(\frac{K^T}{\sqrt{n}}V\right) = E(Q,K,V)$。而 Softmax 下的新算法在视觉任务中也能拿到和原来相近的指标。
这样,复杂度(忽略常数)就从 $O(N^2d)$ 降为 $O(Nd^2)$。在上下文很长、也就是 N 远大于维度 d 的时候,这能大幅降低计算量和空间需求。
FlashAttention
有多篇论文都提到了 FlashAttention。它虽然不属于线性注意力路线,不过为了方便理解,这里也先简单介绍一下。
前面提到,朴素的 Softmax 计算对显存的消耗是巨大的。对于单头注意力:$Q,K,V\in\mathbb{R}^{N\times d}$,$S=\frac{QK^T}{\sqrt d}$,$P=\operatorname{Softmax}(S)$,$O=PV$。朴素的 GPU 实现通常分成几段程序:
- 读 Q、K,算出 $N\times N$ 的 S,写回 HBM;
- 再读 S,逐行做 Softmax,得到 P,写回 HBM;
- 再读 P、V,算出 O;
- 训练时还要为反向传播保留 S 或 P。
问题不只是 S、P 各有 $N^2$ 个元素,它们还会在几段 GPU 程序之间反复写入、读出 HBM。即便计算跟得上,$N\times N$ 矩阵的反复传输也会让 IO 成为瓶颈。 于是 《FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness》 就是为了缓解这个问题。它并没有减少计算量,但是通过大量减少 IO 而大幅提升整体性能。
对于前面 缩放(Scaling) 那种仅仅是矩阵乘法的计算,很容易把大矩阵拆分成几个小矩阵分治计算。
比如我们可以把整体矩阵 $P = \begin{bmatrix} 1 & 2 & 3 & 4 \\ 5 & 6 & 7 & 8 \\ 9 & 10 & 11 & 12 \\ 13 & 14 & 15 & 16 \end{bmatrix}$ 分块划分 $P = \left[ \begin{array}{cc|cc} 1 & 2 & 3 & 4 \\ 5 & 6 & 7 & 8 \\ \hline 9 & 10 & 11 & 12 \\ 13 & 14 & 15 & 16 \end{array} \right]$
那么就有了四个子矩阵 $P_{11} = \begin{bmatrix} 1 & 2 \\ 5 & 6 \end{bmatrix}, \quad P_{12} = \begin{bmatrix} 3 & 4 \\ 7 & 8 \end{bmatrix}, \quad P_{21} = \begin{bmatrix} 9 & 10 \\ 13 & 14 \end{bmatrix}, \quad P_{22} = \begin{bmatrix} 11 & 12 \\ 15 & 16 \end{bmatrix}$ 。整体矩阵可以写成分块形式:$P = \begin{bmatrix} P_{11} & P_{12} \\ P_{21} & P_{22} \end{bmatrix}$ 。
那么对于矩阵 A 和 B 的乘法,就可以按分块的形式相乘:
$$ \begin{bmatrix} A_{11} & A_{12} \ A_{21} & A_{22} \end{bmatrix} \begin{bmatrix} B_{11} & B_{12} \ B_{21} & B_{22} \end{bmatrix}
\begin{bmatrix} A_{11}B_{11} + A_{12}B_{21} & A_{11}B_{12} + A_{12}B_{22} \ A_{21}B_{11} + A_{22}B_{21} & A_{21}B_{12} + A_{22}B_{22} \end{bmatrix} $$
这样就可以逐块计算,峰值空间占用立减 75%。如果还不够,可以多嵌套几层分块,或者把子矩阵拆得更小。
前面提到过,为了降低溢出风险,实现上会用数值上等价的减最大值版本:$\text{Softmax}(z_i) = \frac{exp(z_i - \max(z))}{\sum_{j=1}^{n} exp(z_j - \max(z))}$。但这也让 Softmax 没法直接分块计算:它要除以整行权重的和 $\sum_{j=1}^{N}\exp(s_j)$ 做归一化,而 $\max(z)$ 在不读完整行的情况下也拿不到。
FlashAttention 的核心思想是按块直接计算 $\mathrm{Softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$,并在逐块推进的过程中不断修正 $\max(z)$ 和归一化分母 $\sum_{j=1}^{N}\exp(s_j)$ 的缩放。
令:
- $m_{\text{old}}$:已经看过的分数中的最大值;
- $\ell_{\text{old}}$:以这个最大值为基准缩放后的 Softmax 分母;
- $a_{\text{old}}$:按同一尺度累计、但尚未除以分母的 Value 加权和(也就是 Softmax 分子乘 Value 矩阵的结果)。
每次读入一个新块,分数记为 $s^{(b)}$,Value 记为 $V^{(b)}$。然后计算新的最大值:
$$ m_{\text{new}} = \max\left(m_{\text{old}},\max_j s_j^{(b)}\right) $$每次更新 $\ell$ 时,都要先把旧的 $\ell_{\text{old}}$ 从 $m_{\text{old}}$ 基准换算到 $m_{\text{new}}$ 基准,再累加新块的贡献。
$$ \ell_{\text{new}} = \exp(m_{\text{old}}-m_{\text{new}})\ell_{\text{old}} + \sum_j\exp\left(s_j^{(b)}-m_{\text{new}}\right) $$$$ a_{\text{new}} = \exp(m_{\text{old}}-m_{\text{new}})a_{\text{old}} + \sum_j\exp\left(s_j^{(b)}-m_{\text{new}}\right)v_j^{(b)} $$全部块处理完后:
$$ o=\frac{a_{\text{new}}}{\ell_{\text{new}}} $$只计算 Softmax 本身时,额外维护 m 和 $\ell$ 就够了。FlashAttention 还要紧接着乘 V,所以把尚未归一化的加权和 a 也一起累计,避免生成完整的概率矩阵 P。m、$\ell$ 和 a 的显存只随分块行数增长,不再需要 $N\times N$ 级别的空间。
拿一个 Value 只有单个数字的玩具例子说明。某个 Query 对四个 Key 的分数是 $s=[1,2,3,0]$,对应的 Value 是 $v=[10,20,30,40]$,一次性计算完整 Softmax,输出约为:
$$ \operatorname{softmax}(s)v^\top\approx 26.2089 $$现在把四项拆成两块。第一块是分数 $[1,2]$ 和 Value $[10,20]$。以块内最大值 2 为基准,先记录:
$$ m_1=2, $$$$ \ell_1=e^{1-2}+e^{2-2}\approx 1.3679, $$$$ a_1=e^{1-2}\times 10+e^{2-2}\times 20\approx 23.6788. $$$m_1$ 是当前最大分数,$\ell_1$ 是 Softmax 分母的累计值,$a_1$ 是尚未除以分母的加权 Value。
第二块的分数是 $[3,0]$。新的全局最大值变成 3,旧累计值原来以 2 为基准,必须先乘 $e^{2-3}$,才能与新块放在同一尺度:
$$ \ell_2=e^{2-3}\ell_1+e^{3-3}+e^{0-3}\approx 1.5530, $$$$ a_2=e^{2-3}a_1+e^{3-3}\times 30+e^{0-3}\times 40\approx 40.7024. $$最后相除:
$$ o=\frac{a_2}{\ell_2}\approx 26.2089. $$后续还有一些扩展内容:《FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning》、《FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision》 和 《FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling》,我暂时也还没看,以后有兴趣再深入学习。
线性注意力
《Efficient Attention: Attention with Linear Complexities》 这篇论文主要是针对视觉任务的,分解后的 Softmax 和 Transformer 里的不一样,显然不能直接用到 LLM 里。 但是思路是相近的。于是 《Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention》 针对自回归场景提出了线性复杂度的自注意力方案。
输入序列 $x$ 被三个矩阵 $W_Q \in \mathbb{R}^{F\times D}$、$W_K \in \mathbb{R}^{F\times D}$、$W_V \in \mathbb{R}^{F\times M}$ 投影为对应的 Q、K、V,自注意力计算如下:
$$ \begin{aligned} Q &= xW_Q,\quad K = xW_K,\quad V = xW_V \\ A_l(x) &= V' = \mathrm{softmax}\left(\frac{QK^T}{\sqrt{D}}\right)V \end{aligned} $$用下标 $i$ 表示取矩阵第 $i$ 行,则对任意相似度函数,广义注意力可写为:
$$ V'_i = \frac{\sum_{j=1}^{N} \mathrm{sim}(Q_i, K_j)\, V_j}{\sum_{j=1}^{N} \mathrm{sim}(Q_i, K_j)} $$其中,$\mathrm{sim}(q,k)=\exp\left(\frac{q^T k}{\sqrt{D}}\right)$。
为什么这里是 $q^T k$,而上面是 $QK^T$? 这是因为一个是列向量视角,另一个是行向量视角。在向量层面,习惯上单个向量默认是列向量。拿同一组向量对比:
- 列向量视角:$q, k \in \mathbb{R}^{D \times 1}, \quad \text{点积} = q^T k$
- 行向量视角:$Q_i, K_j \in \mathbb{R}^{1 \times D}, \quad \text{点积} = Q_i K_j^T$
举个 $D = 2$ 的例子:
- 列向量视角:$q = \begin{bmatrix} a \\ b \end{bmatrix},\ k = \begin{bmatrix} c \\ d \end{bmatrix} \Rightarrow q^T k = ac + bd$
- 行向量视角:$Q_i = \begin{bmatrix} a & b \end{bmatrix},\ K_j = \begin{bmatrix} c & d \end{bmatrix} \Rightarrow Q_i K_j^T = ac + bd$
受前面高效注意力的启发,如果我们能把 $\mathrm{sim}(q,k)$ 写成 $\phi(q)^T\phi(k)$ 的形式,就可以把注意力改写为:
$$ V'_i = \frac{\sum_{j=1}^{N} \phi(Q_i)^T \phi(K_j)\, V_j}{\sum_{j=1}^{N} \phi(Q_i)^T \phi(K_j)} $$再利用矩阵乘法结合律进一步化简:
$$ V'_i = \frac{\phi(Q_i)^T \sum_{j=1}^{N} \phi(K_j) V_j^T}{\phi(Q_i)^T \sum_{j=1}^{N} \phi(K_j)} $$其实我觉得这样更好理解:$\big(\phi(Q)\,\phi(K)^T\big)V = \phi(Q)\,\big(\phi(K)^T V\big)$。
接下来,这篇论文把 $\phi(x)$ 定义为:
$$ \begin{aligned} \phi(x) &= \mathrm{elu}(x) + 1 \\ \mathrm{elu}(x) &= \begin{cases} x, & x > 0 \\ e^x - 1, & x \le 0 \end{cases} \end{aligned} $$也就是用它代替了原来的 exp 函数。这个替换并非完全无损。exp 有一个 $\mathrm{elu}+1$ 没有的性质:锐化(sharpening)——指数会把大 logits 急剧放大、把小 logits 压到接近 0,让注意力分布更“尖”。线性注意力没有这个放大机制,权重分布天然更平滑,更偏向“平均检索”而非“精确点名”。这也是为什么后续工作(如 Performer 用随机特征逼近 exp、cosFormer 加位置重加权)都在设法把 Softmax 的尖锐性补回来。
再把前面的 $\mathrm{sim}(q, k) = \phi(q)^T \phi(k)$ 代入,就得到 $\mathrm{sim}(q, k) = \big(\mathrm{elu}(q)+1\big)^T \big(\mathrm{elu}(k)+1\big)$。
然后我们继续引入两个累积量:
$$ S_i = \sum_{j=1}^{i} \phi(K_j) V_j^T,\qquad Z_i = \sum_{j=1}^{i} \phi(K_j) $$这里考虑了 因果掩码:当我们访问到第 i 个位置时,要屏蔽掉未来的 token,所以求和只累计到 i,而不是上面公式里的 N。那么线性注意力公式可简写为:
$$ V'_i = \frac{\phi(Q_i)^T S_i}{\phi(Q_i)^T Z_i} $$到这里,就比较容易解释为什么论文标题说 Transformers are RNNs 了:$S_i$ 和 $Z_i$ 分别可以由 $S_{i-1}$ 和 $Z_{i-1}$ 递推得到,每步只依赖一份固定大小的状态——这正是 RNN 的递归结构。
$$ S_0 = 0, Z_0 = 0 $$$$ S_i = S_{i-1} + \phi(x_i W_K)\,(x_i W_V)^T $$$$ Z_i = Z_{i-1} + \phi(x_i W_K) $$$$ y_i = f_l\left(\frac{\phi(x_i W_Q)^T S_i}{\phi(x_i W_Q)^T Z_i} + x_i\right) $$直观解释:$S_i\in\mathbb{R}^{C\times M}$ 是内容状态,$Z_i\in\mathbb{R}^{C}$ 是归一化状态。每来一个 token,就先写入两份状态,再用当前查询读取。在本文这种有限维 $\mathrm{elu}+1$ 特征映射下,状态大小确实与序列长度无关。
到这里,线性注意力的基本原理就讲完了,我们来对比一下它和前面几种方案的时间与空间复杂度。
| 维度 | 标准注意力(Vanilla) | FlashAttention | Efficient Attention(高效注意力) | Linear Attention(线性注意力) |
|---|---|---|---|---|
| 是否精确 Softmax | 精确 | 精确(数学上完全等价) | 近似($Q$ 按行、$K$ 按列各自 Softmax,归一化解耦;仅 Scaling 归一化时与原注意力严格等价) | 近似(换成 $\phi(q)^T\phi(k)$) |
| 时间复杂度(FLOPs) | $O(N^2 d)$ | $O(N^2 d)$,不变 | $O(N d^2)$,对 $N$ 线性 | $O(N\,c\,d) = O(N d^2)$,对 $N$ 线性 |
| HBM 访存量(IO) | $\Theta(Nd + N^2)$ | $\Theta\!\left(\frac{N^2 d^2}{M}\right)$,$M$ 为 SRAM 大小 | $O(Nd)$(只物化 $d \times d$ 上下文矩阵) | $O(Nd)$ |
| 显存(注意力部分) | $O(N^2)$(实例化注意力矩阵) | $O(Nd)$,分块计算 + 反向重算 | $O(Nd + d^2)$,不存 $N{\times}N$ 矩阵 | $O(Nd)$,不存 $N{\times}N$ 矩阵 |
| 自回归推理 | KV cache,每步 $O(Nd)$,cache 占 $O(Nd)$ | 同标准注意力(KV cache),但常数更小 | 论文未针对因果/自回归设计(主要面向视觉等双向任务) | RNN 式状态,每步 $O(cd)$,与 $N$ 无关,状态仅占 $O(cd)$ |
| 主要代价 | 平方时间 + 平方显存 | 仍是平方 FLOPs,超长序列最终撑不住 | 归一化解耦带来近似;各 query 共享 $d$ 个“全局上下文”模板,逐查询自适应弱 | 精度近似、因果训练要 cumsum/分块 |
最后
线性注意力后续的优化涉及的论文很多,我自己也是抽空慢慢看,所以后面单独写一篇再聊吧。
我本人并非 AI 领域从业者,可能有理解不到位的地方,欢迎指正。