首页
学习
活动
专区
圈层
工具
发布
社区首页 >专栏 >大语言模型的基石:Transformer 入坑笔记(四) - 线性注意力(Linear Attention) 的基础

大语言模型的基石:Transformer 入坑笔记(四) - 线性注意力(Linear Attention) 的基础

作者头像
owent
发布2026-09-07 09:04:19
发布2026-09-07 09:04:19
160
举报
文章被收录于专栏:owentowent

背景

接上文 《大语言模型的基石: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^QW^KW^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 NQK^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 实现通常分成几段程序:

  1. 读 Q、K,算出 N\times N 的 S,写回 HBM;
  2. 再读 S,逐行做 Softmax,得到 P,写回 HBM;
  3. 再读 P、V,算出 O;
  4. 训练时还要为反向传播保留 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_{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}

那么对于矩阵 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_iZ_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 领域从业者,可能有理解不到位的地方,欢迎指正。

本文参与 腾讯云自媒体同步曝光计划,分享自作者个人站点/博客。
原始发表:2026-09-06,如有侵权请联系 cloudcommunity@tencent.com 删除
目录
  • 背景
  • 复杂度估算
  • 高效注意力的基础
  • FlashAttention
  • $$ \begin{bmatrix} A_{11} & A_{12} \ A_{21} & A_{22} \end{bmatrix} \begin{bmatrix} B_{11} & B_{12} \ B_{21} & B_{22} \end{bmatrix}
    • 线性注意力
    • 最后
问题归档专栏文章快讯文章归档关键词归档开发者手册归档开发者手册 Section 归档