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

💡 原文中文,约10700字,阅读约需26分钟。
📝

内容提要

本文探讨Transformer注意力机制的计算复杂度问题,指出标准注意力在长上下文下时间和空间开销呈平方增长。文章概述了高效注意力、FlashAttention和线性注意力等优化方案,其中线性注意力通过将Softmax替换为elu+1特征映射,将复杂度降至线性,并引入RNN式递归状态,为后续Kimi等长上下文模型奠定基础。

🔎

延伸解读

线性注意力的核心:用特征映射替代Softmax

线性注意力将Softmax替换为elu(x)+1的特征映射,使相似度函数可分解为φ(q)^Tφ(k),从而利用矩阵结合律将计算复杂度从O(N²d)降至O(Nd²)。这种替换并非无损,exp的锐化特性(放大差异)被丢失,导致注意力分布更平滑,偏向平均检索而非精确聚焦。后续工作如Performer、cosFormer等都在尝试弥补这一缺陷,恢复Softmax的尖锐性。

RNN式递归:状态与序列长度解耦

线性注意力引入两个累积量S_i和Z_i,分别存储键值对和键的加权和,通过递推公式S_i=S_{i-1}+φ(K_i)V_i^T更新,使每一步只依赖固定大小的状态,与序列长度无关。这解释了为何论文标题称“Transformers are RNNs”。在自回归推理时,无需存储不断增长的KV cache,状态仅占O(cd)空间,显著降低长上下文推理的内存开销。

FlashAttention:不降复杂度,但优化IO

FlashAttention并非线性注意力路线,它保持精确Softmax和O(N²d)的计算量,但通过分块计算和在线更新最大值、分母,避免实例化N×N的注意力矩阵,大幅减少HBM读写,使IO成为瓶颈的问题得到缓解。其核心是维护m、ℓ和a三个状态,逐块修正缩放,最终得到与标准Softmax完全一致的结果。

Q&A

标准注意力机制在长上下文场景下有什么主要问题?

标准注意力机制的时间和空间复杂度随序列长度N呈平方增长,即O(N^2d)的时间复杂度和O(N^2)的显存占用,导致在长上下文下计算和存储开销爆炸式增长。

线性注意力是如何将复杂度从平方级降到线性级的?

线性注意力通过将Softmax替换为elu+1特征映射,使得相似度函数可以分解为φ(q)^Tφ(k),从而利用矩阵乘法的结合律,将计算顺序从(QK^T)V改为φ(Q)(φ(K)^T V),避免了显式构造N×N的注意力矩阵,将复杂度降至O(Nd^2)。

为什么说线性注意力可以看作RNN?

因为线性注意力引入了两个累积状态S_i和Z_i,它们可以按递推公式S_i = S_{i-1} + φ(K_i)V_i^T和Z_i = Z_{i-1} + φ(K_i)更新,每一步只依赖固定大小的状态,与序列长度无关,这种递归结构类似于RNN。

FlashAttention是如何在不减少计算量的情况下提升性能的?

FlashAttention通过分块计算和在线Softmax技巧,避免了将完整的N×N注意力矩阵写入HBM,减少了IO访问,从而大幅提升了性能。它通过维护运行最大值m和归一化分母l,逐块更新,最终得到精确的Softmax结果。

线性注意力相比标准注意力有哪些优缺点?

优点:时间复杂度O(Nd^2)和显存O(Nd)均对N线性,适合长上下文;自回归推理时状态大小固定,与N无关。缺点:使用elu+1代替exp,缺乏锐化机制,导致注意力分布更平滑,可能降低对关键信息的聚焦能力,是一种近似方法。

高效注意力(Efficient Attention)与线性注意力有何关系?

高效注意力最早来自视觉领域,通过将注意力计算重写为ρ_q(Q)(ρ_k(K)^T V),将复杂度降为O(Nd^2)。线性注意力借鉴了其分解思想,但针对自回归场景,将Softmax替换为elu+1,并引入因果掩码和递归状态,使其适用于语言模型。

🏷️

标签

➡️

继续阅读