Skip to content

注意力机制变体

标准自注意力的计算复杂度随序列长度 N 呈平方增长,且 KV 缓存在推理时占用大量显存。本页梳理工程实践中最重要的注意力变体:从 Multi-Query、Grouped-Query 到 Multi-head Latent Attention、Sparse、Linear、Sliding Window,再到 FlashAttention——它们分别从精度、效率、显存三个维度优化注意力。前置阅读:注意力机制、Transformer 架构。

把自注意力想象成一场会议室讨论会:

  • 标准多头注意力(MHA)= 每个人(query)都有自己专属的笔记本(KV),开会前要翻所有人的笔记——笔记越多,翻得越慢,存放笔记的柜子也越大。
  • Multi-Query Attention(MQA)= 全会议室共享一本笔记本。翻得快了,但信息可能不够精细——推理速度飞升,精度略有下降。
  • Grouped-Query Attention(GQA)= 折中方案:分成几个小组,每组共享一本笔记。兼顾速度和精度,LLaMA-2/3 就用这个。
  • Multi-head Latent Attention(MLA)= 把每个人的笔记压缩成一份”摘要卡片”(低秩潜变量),推理时只存卡片、用的时候再还原——DeepSeek-V2 的方案,压缩率高达 93%。
  • Sparse Attention= 不让每个人翻所有人的笔记,只看”重要的几个”(局部窗口 + 少量全局点),像只读同事的而不是全公司的。
  • Linear Attention= 用数学技巧把计算顺序调换,把 N 平方的开销降到 N 线性——代价是近似,精度有损。
  • Sliding Window Attention= 每个人只看左右邻居(局部窗口),靠多层堆叠让信息间接传遍全场。
  • FlashAttention= 不是改变算法本身,而是优化 GPU 读写顺序——同样的计算,更快更省显存(减少 HBM 读写)。FlashAttention-3(2024)进一步利用 Hopper GPU 的异步计算能力,FlashAttention-4(2025)扩展到 Blackwell GPU。

注意力公式中最关键的运算之一是 softmax——它把一组任意实数分数转换为概率分布(非负数,且总和为 1)。对向量 s = (s₁, s₂, …, sₙ),softmax 的定义为:

softmax(si)=exp⁡(si)∑jexp⁡(sj)\text{softmax}(s_i) = \frac{\exp(s_i)}{\sum_j \exp(s_j)}
  • 为什么用指数 exp? 因为指数函数永远为正,保证了输出非负;同时它放大了较大分数的差距——一个得分为 5 的 token 相比得分为 1 的 token,注意力权重会高出 e⁴ ≈ 55 倍,这就是注意力”聚焦”在相关 token 上的数学基础。
  • 数值稳定性:直接计算 exp(sᵢ) 在 sᵢ 较大时会导致数值溢出。实际实现中先减去最大值再做 softmax:softmax(sᵢ) = exp(sᵢ - max(s)) / Σⱼ exp(sⱼ - max(s)),结果在数学上完全等价但数值更稳定。FlashAttention 的”在线 softmax”算法正是利用了这一性质来实现分块计算。

给定输入序列 X(维度 N×d),通过三个可学习的权重矩阵 W_Q、W_K、W_V 将其投影为三个矩阵:

Q=X⋅WQK=X⋅WKV=X⋅WVQ = X \cdot W_Q \quad K = X \cdot W_K \quad V = X \cdot W_V

注意力输出为:

Attention(Q,K,V)=softmax(Q⋅KTdk)⋅V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{Q \cdot K^T}{\sqrt{d_k}}\right) \cdot V
  • 注意力矩阵(attention matrix)= Q · Kᵀ / √d_k,维度 N×N,第 (i,j) 个元素表示第 i 个 token 对第 j 个 token 的关注程度(尚未归一化)。
  • 为什么除以 √d_k? 当 d_k 较大时,Q·Kᵀ 的值会很大(它是 d_k 个项的求和),导致 softmax 进入梯度饱和区。除以 √d_k(即缩放因子,scaling factor)将方差控制在 1 附近,使梯度更稳定。这一技巧由 Transformer 原论文(Vaswani et al., 2017)提出。
  • 多头(Multi-Head):将 d 维向量切分成 h 个头,每个头独立做注意力,再拼接。这让模型同时关注不同子空间的信息。

在自回归解码(autoregressive decoding)时,每生成一个新 token,需要用到之前所有 token 的 K 和 V。为避免重复计算,将这些 K、V 存入 KV 缓存(KV Cache)。其显存占用公式为:

KV 缓存大小=2×num_layers×num_kv_heads×head_dim×seq_len×batch_size×dtype_size\text{KV 缓存大小} = 2 \times \text{num\_layers} \times \text{num\_kv\_heads} \times \text{head\_dim} \times \text{seq\_len} \times \text{batch\_size} \times \text{dtype\_size}

其中 2 是因为同时缓存 K 和 V,dtype_size 是每个元素的字节数(float32 = 4 字节,float16/bfloat16 = 2 字节)。

举例:LLaMA-2 70B 有 80 层、64 个 Q 头、8 个 KV 头(GQA)、head_dim=128,batch_size=1,float16,序列长度 4096:

KV 缓存=2×80×8×128×4096×1×2=1,342,177,280 字节≈1.25 GB\text{KV 缓存} = 2 \times 80 \times 8 \times 128 \times 4096 \times 1 \times 2 = 1{,}342{,}177{,}280 \text{ 字节} \approx 1.25 \text{ GB}

序列扩展到 32K 时,单条请求的 KV 缓存就达到约 10 GB——这就是长上下文推理的核心瓶颈。

理解 FlashAttention 的关键在于 GPU 的存储层次:

层次容量带宽说明
HBM(High Bandwidth Memory)40-80 GB(A100/H100)~2-3 TB/s显卡主显存,存放模型权重、激活值,容量大但相对慢
SRAM(Static RAM,片上缓存)~20 MB(A100),~228 MB(H100)~19-28 TB/sGPU 许算核心旁边的快速缓存,容量极小但带宽极高(约 HBM 的 10 倍)

标准注意力在计算时,需要把 N×N 的中间矩阵先写入 HBM,再读回做 softmax——大量时间花在 HBM 的读写上。FlashAttention 的核心思想是:把 Q、K、V 分成小块(tiling),加载到 SRAM 中完成全部计算,只把最终结果写回 HBM。

标准自注意力的计算为:Attention(Q, K, V) = softmax(Q · Kᵀ / √d_k) · V。其中 Q、K、V 的维度均为 N×d。Q · Kᵀ 这一步产生 N×N 的注意力矩阵(attention matrix),因此:

  • 计算量:Q·Kᵀ 需要 O(N²·d) 次乘加,softmax 对 N×N 矩阵逐行操作需 O(N²),再乘 V 又需 O(N²·d)。总计 O(N²·d)。
  • 显存:注意力矩阵本身占 O(N²),加上 Q、K、V 的 O(N·d),总量为 O(N² + N·d)。

当 N=8192、d=128 时,注意力矩阵有 6700 万个元素——这在现代 LLM 中是常态。

在推理(解码)时,每生成一个 token 要保存已生成的所有 K、V(称为 KV 缓存)。序列越长,KV 缓存越大,成为推理瓶颈(详见上方 KV 缓存公式)。

MQA 的做法:所有 query 头共享同一组 K 和 V(只有 1 个 KV 头)。假设有 h 个头,KV 缓存从 h 份降为 1 份,推理时 KV 读取量减少到 1/h。代价是精度通常略有下降(约 1-2 个点)。MQA 由 Shazeer 在 2019 年提出,GPT-J、PaLM、Falcon 等模型采用。

数学表示:标准 MHA 中,第 i 个头的注意力为 headᵢ = Attention(Q·WᵢQ, K·WᵢK, V·WᵢV),有 h 组独立的 (WᵢK, WᵢV) 投影。MQA 将所有头的 K、V 投影替换为共享的 (WK, WV):

headi=Attention(Q⋅WiQ,  K⋅WK,  V⋅WV)\text{head}_i = \text{Attention}(Q \cdot W_i^Q, \; K \cdot W^K, \; V \cdot W^V)

GQA 是 MHA 和 MQA 的中间态:把 h 个 query 头分成 g 组,每组共享一个 KV 头。当 g=1 时退化为 MQA,当 g=h 时退化为标准 MHA。LLaMA-2(70B)和 LLaMA-3 默认使用 GQA(如 8 个 KV 头服务 64 个 query 头)。Ainslie et al. (2023) 的研究表明,GQA 在保持接近 MHA 精度的同时获得 MQA 的推理加速。

数学表示:将 h 个 query 头索引按 h/g 分组,第 j 组的所有 query 头共享第 j 个 KV 头:

headi=Attention(Q⋅WiQ,  K⋅WjK,  V⋅WjV),j=⌊ih/g⌋\text{head}_i = \text{Attention}(Q \cdot W_i^Q, \; K \cdot W_j^K, \; V \cdot W_j^V), \quad j = \left\lfloor \frac{i}{h/g} \right\rfloor

KV 缓存缩减:与 MHA 相比,GQA 的 KV 缓存缩减为 g/h。例如 64 个 Q 头 + 8 个 KV 头 → KV 缓存缩减为 8/64 = 12.5%,即减少 87.5%。

GQA 的上采样(GQA-Upcasting):推理时,如果需要将 GQA 模型转为 MHA 以获得更高质量,可以将每个 KV 头复制 h/g 份——代价是恢复完整的 KV 缓存开销。

MLA 是 DeepSeek-V2(2024)提出的创新注意力机制,旨在同时降低 KV 缓存和保持精度。核心思想是低秩压缩(low-rank compression):

  1. 将 KV 投影到一个低维潜变量(latent vector)c_KV(维度 d_c ≪ d × num_heads),只缓存这个潜变量。
  2. 推理时通过一个”上投影”矩阵将潜变量还原为完整的 K、V。
cKV=WDKV⋅htKt=WUK⋅cKVVt=WUV⋅cKVc_{KV} = W_{DKV} \cdot h_t \quad\quad K_t = W_{UK} \cdot c_{KV} \quad\quad V_t = W_{UV} \cdot c_{KV}

DeepSeek-V2 通过 MLA 将 KV 缓存减少了 93.3%(相比 MHA),同时性能几乎不受影响。MLA 还结合了 解耦 RoPE(decoupled RoPE)策略——将 RoPE 应用于一个额外的低维子空间,避免旋转位置编码干扰低秩压缩的数学性质。

MLA 代表了注意力机制设计的新范式:不再通过减少头数(如 GQA/MQA)来压缩缓存,而是从信息瓶颈角度出发做低秩降维。DeepSeek-V3(2024)和 DeepSeek-R1(2025)都沿用了 MLA 架构。

核心思想:将 N×N 的全注意力矩阵替换为稀疏模式,只计算部分 query-key 对。常见策略:

  • 局部窗口(Local/Window):每个 token 只 attends 前后 w 个 token(如 w=256)。
  • 膨胀窗口(Dilated):类似膨胀卷积(dilated convolution),以步长跳着看,扩大感受野而不增加计算量。
  • 全局 token(Global):设置少量全局 token,所有位置都 attends 它们(如 BigBird 的策略)。
  • 随机连接(Random):随机选少量 token 做 attention,理论上保证图的连通性(connectivity)。

BigBird(Zaheer et al., 2020)结合以上四种模式,理论上是 Turing 完备的(Turing complete,即能逼近任意序列函数),而计算量从 O(N²) 降到 O(N)。Longformer 则采用局部窗口 + 少量全局 token 的实用组合。

标准注意力的 softmax 包含指数运算,无法简单重排计算顺序。Linear Attention 用核函数(kernel function)k(x,y) 近似 softmax,将 Attention(Q,K,V) 改写为:

标准:Attention=softmax(Q⋅KT)⋅V\text{标准:} \quad \text{Attention} = \text{softmax}(Q \cdot K^T) \cdot V 线性近似:Attention=φ(Q)⋅(φ(K)T⋅V)\text{线性近似:} \quad \text{Attention} = \varphi(Q) \cdot (\varphi(K)^T \cdot V)

其中 φ(·) 是一个非线性特征映射(feature map),将 softmax 的 exp 近似替代。关键在于 φ(K)ᵀ·V 可以先算,得到一个 d×d 矩阵(与 N 无关),再乘以 φ(Q)。这样计算量变为 O(N·d²),当 d 远小于 N 时为线性复杂度。

  • 核函数(kernel function):在机器学习中,核函数 k(x,y) 衡量两个向量之间的相似度。标准注意力的”核”是 softmax,即 k(x,y) ∝ exp(x·y);Linear Attention 用更简单的核替代,如 ELU+1(Linear Transformer)或随机傅里叶特征(Performer)。
  • 代表工作有 Linear Transformer(Katharopoulos et al., 2020)和 Performer(Choromanski et al., 2021),但近似精度在大规模模型上仍有差距。

Longformer(Beltagy et al., 2020)和 Mistral 系列采用。每个 token 只 attends 窗口大小 w 内的 token,计算量为 O(N·w)。通过堆叠 L 层,有效感受野为 L·w——例如 32 层 × 窗口 4096 = 131072 的理论感受野。Mistral-7B 用滑动窗口 + 全局注意力实现了高效的长序列建模。

滑动窗口 + GQA 的组合是当前长上下文模型的高效配方:GQA 压缩 KV 缓存的”宽度”(头数),滑动窗口压缩”深度”(序列长度),两者叠加可以将 KV 缓存降至传统 MHA 的极小比例。

FlashAttention(Dao et al., 2022)不是改变注意力公式,而是优化 GPU 内存访问模式(I/O-aware computation)。传统实现先把 N×N 的注意力矩阵写到 HBM(高带宽显存),再读回做 softmax——大量 I/O 开销。FlashAttention 把 Q、K、V 分块加载到 SRAM(片上快速缓存),用**分块算法(tiling)和在线 softmax(online softmax)**逐步计算,避免中间矩阵落盘。

在线 softmax 的原理:传统 softmax 需要先遍历所有元素求最大值和求和,再做归一化。FlashAttention 利用了一个数学性质——可以在分块处理时增量更新最大值 m 和归一化因子 l(即 sum of exp),每处理一个新块就修正之前的累积值,最终结果与标准 softmax 完全一致。这就是”在线”的含义:不需要一次性看到全部数据。

效果:同样的精度(exact attention,非近似),速度提升 2-4 倍,显存从 O(N²) 降到 O(N)。FlashAttention-2(Dao, 2023)进一步优化并行度和 warp 级效率(GPU 的最小执行单元级别)。

FlashAttention-3(Shah et al., 2024)针对 NVIDIA Hopper 架构(H100 GPU)的新硬件特性进行了深度优化,引入三项关键技术:

  1. 异步计算重叠(Overlap via Warp-Specialization):利用 Hopper 的 Tensor Memory Accelerator(TMA,张量内存加速器)异步加载数据,同时用专门的 warp 做 softmax,实现计算与数据搬移的流水线重叠——GPU 不再”等数据”。
  2. 块状 matmul 与 softmax 交织(Interleaving):在 GEMM(矩阵乘)还在执行时就开始做 softmax 的归约操作,进一步隐藏延迟。
  3. FP8 低精度支持:利用 Hopper 硬件原生支持的 FP8 格式和 block-wise quantization(分块量化),配合 incoherent processing(非相干处理,通过随机旋转降低量化误差)。

性能:在 H100 上,FP16 模式达到 740 TFLOPs/s(75% MFU,即模型算力利用率),相比 FA2 加速 1.5-2.0×;FP8 模式达到接近 1.2 PFLOPs/s,且数值精度比朴素 FP8 实现低 2.6× 的误差。

FlashAttention-4 使用 NVIDIA 新的 CuTeDSL(CUTLASS DSL,一种声明式 GPU kernel 语言)编写,同时优化 Hopper 和 Blackwell(B200 GPU)架构。这是 FlashAttention 系列首次官方支持最新的 Blackwell GPU,进一步提升了超长上下文训练的效率。

RoPE(Su et al., 2021)不是改变注意力结构,而是一种位置编码方式:通过旋转向量对的方式将位置信息注入 Q 和 K,使得两个 token 的注意力分数只依赖它们的相对距离。

数学直觉:将 d 维向量视为 d/2 个二维平面,每个平面按 token 的位置 m 做角度为 mθ 的旋转。对于位置 m 的 Q 和位置 n 的 K,旋转角度差为 (m-n)θ,因此 Q·Kᵀ 只依赖相对位置 m-n——这就是”相对位置编码”的本质。

x2k′=x2k⋅cos⁡(m⋅θk)−x2k+1⋅sin⁡(m⋅θk)x'_{2k} = x_{2k} \cdot \cos(m \cdot \theta_k) - x_{2k+1} \cdot \sin(m \cdot \theta_k) x2k+1′=x2k⋅sin⁡(m⋅θk)+x2k+1⋅cos⁡(m⋅θk)x'_{2k+1} = x_{2k} \cdot \sin(m \cdot \theta_k) + x_{2k+1} \cdot \cos(m \cdot \theta_k)

其中 θ_k = 10000^(-2k/d) 是不同维度对使用的基础频率(低维用高频、高维用低频,类似多尺度)。

RoPE 是 LLaMA、Qwen、Mistral 等主流大模型的标配,优势是天然支持相对位置和外推(extrapolation,即训练时用短序列、推理时扩展到长序列)。配合 YaRN 或 NTK-aware 缩放等技巧,可以将 4K 上下文外推到 128K 甚至更长。

PyTorch:对比 MHA、GQA、MQA、MLA 的 KV 缓存大小

Section titled “PyTorch:对比 MHA、GQA、MQA、MLA 的 KV 缓存大小”
# 配置:64 个 query 头,head_dim=128,序列长度 4096,float16
num_heads, head_dim, seq_len, num_layers = 64, 128, 4096, 80
dtype_bytes = 2 # float16
for name, num_kv in [("MHA", 64), ("GQA-8", 8), ("GQA-4", 4), ("MQA", 1)]:
# KV 缓存: 2(K+V) × num_layers × num_kv × seq_len × head_dim × dtype_bytes
kv_gigabytes = (2 * num_layers * num_kv * seq_len * head_dim * dtype_bytes
/ 1024**3)
ratio = num_kv / num_heads * 100 # 相对 MHA 的缓存比例
print(f"{name:6s}: {num_kv:2d} KV 头, KV 缓存 = {kv_gigabytes:.2f} GB "
f"({ratio:.1f}% of MHA)")
# MHA : 64 KV 头, KV 缓存 = 10.00 GB (100.0%)
# GQA-8 : 8 KV 头, KV 缓存 = 1.25 GB ( 12.5%)
# GQA-4 : 4 KV 头, KV 缓存 = 0.63 GB ( 6.3%)
# MQA : 1 KV 头, KV 缓存 = 0.16 GB ( 1.6%)
# MLA(DeepSeek-V2 风格)的压缩率可达 MHA 的 ~6.7%(减少 93.3%)

PyTorch 2.0+:使用内置 FlashAttention(SDPA)

Section titled “PyTorch 2.0+:使用内置 FlashAttention(SDPA)”
import torch
import torch.nn.functional as F
# PyTorch 2.0+ 的 scaled_dot_product_attention 自动选择最优后端
# (包括 FlashAttention、memory-efficient attention、math fallback)
q = torch.randn(1, 8, 4096, 128, device="cuda") # (batch, heads, seq, dim)
k = torch.randn(1, 8, 4096, 128, device="cuda")
v = torch.randn(1, 8, 4096, 128, device="cuda")
# 自动使用 FlashAttention(CUDA 环境 + 满足条件时)
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
print(f"输出形状: {out.shape}") # torch.Size([1, 8, 4096, 128])
# 检查实际调用了哪个后端
from torch.nn.attention import sdpa_kernel, SDPBackend
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
out_flash = F.scaled_dot_product_attention(q, k, v, is_causal=True)
print("FlashAttention 后端执行成功")
# GQA 示例:Q 有 8 头,K/V 只有 2 头(enable_gqa=True 从 PyTorch 2.5 起)
q_gqa = torch.randn(1, 8, 4096, 128, device="cuda")
k_gqa = torch.randn(1, 2, 4096, 128, device="cuda") # 2 个 KV 头
v_gqa = torch.randn(1, 2, 4096, 128, device="cuda")
out_gqa = F.scaled_dot_product_attention(q_gqa, k_gqa, v_gqa, is_causal=True,
enable_gqa=True)
print(f"GQA 输出形状: {out_gqa.shape}") # (1, 8, 4096, 128)
import numpy as np
N, w = 10, 3 # 序列长度 10,窗口大小 3(看前面 3 个 token)
mask = np.zeros((N, N), dtype=bool)
# 每个 token 只能看到自己前面 w 个 + 自己
for i in range(N):
start = max(0, i - w + 1) # 窗口左边界
mask[i, start:i+1] = True # 窗口内的位置可见
print(mask[:5, :5].astype(int))
# [1 0 0 0 0] token 0 只看自己
# [1 1 0 0 0] token 1 看 0-1
# [1 1 1 0 0] token 2 看 0-2
# [0 1 1 1 0] token 3 看 1-3(窗口滑动)
# [0 0 1 1 1] token 4 看 2-4

手写在线 Softmax(FlashAttention 的核心算法)

Section titled “手写在线 Softmax(FlashAttention 的核心算法)”
import numpy as np
def online_softmax(scores):
"""在线 softmax:逐块累积,无需一次性看到全部数据。
FlashAttention 用同样的原理在 SRAM 中分块计算。"""
m = -np.inf # 当前已知的最大值
l = 0.0 # 当前 exp 和的累积
result = np.zeros_like(scores)
for i, s in enumerate(scores):
m_new = max(m, s) # 更新最大值
l = l * np.exp(m - m_new) # 修正之前的累积和
l += np.exp(s - m_new) # 加入当前项
m = m_new
# 第二遍:用最终的 m 和 l 做归一化
m_final = m
l_final = l
for i, s in enumerate(scores):
result[i] = np.exp(s - m_final) / l_final
return result
# 验证:与标准 softmax 对比
scores = np.array([1.0, 3.0, 0.5, 2.0, 4.0])
online = online_softmax(scores)
standard = np.exp(scores) / np.exp(scores).sum()
print("在线 softmax:", np.round(online, 6))
print("标准 softmax:", np.round(standard, 6))
print("差异:", np.max(np.abs(online - standard))) # ~0.0(数值上等价)
  • 始终启用 FlashAttention:无论训练还是微调,FlashAttention/SDPA 几乎没有副作用(精度完全一致),却能让训练吞吐量提升 2-4 倍。PyTorch 2.0+ 的 F.scaled_dot_product_attention 已自动调用。
  • 用 GQA 从头训练:如果从头训练模型,直接使用 GQA(如 64 Q 头配 8 KV 头)。不要先训 MHA 再转 GQA(虽然可行,但增加了复杂度)。
  • 混合精度训练:注意力计算使用 bfloat16 或 float16。FlashAttention-3 的 FP8 模式在 H100 上能进一步加速,但需要配合 incoherent processing 控制量化误差。
  • 梯度检查点(Gradient Checkpointing)与 FlashAttention 兼容:FlashAttention 已将注意力显存降至 O(N),但 Transformer 其他层仍有显存压力。两者叠加使用效果最佳。
  • GQA + PagedAttention:vLLM 等推理引擎用 PagedAttention(分页式 KV 缓存管理)配合 GQA,将显存碎片化降到最低,支持高并发推理。
  • KV 缓存量化:将 KV 缓存从 float16 量化为 int8 或 int4,可将缓存大小再减半到四分之一。FP8(H100)和 KVCache 量化(如 KIVI、KVQuant)是 2024-2025 年的热点方向。
  • 滑动窗口用于流式推理:对话场景下上下文持续增长,滑动窗口可以让推理的显存占用恒定(只保留最近 w 个 token 的 KV),适合无限长度的流式场景。
  • RoPE 长度外推:训练时用 4K 上下文,推理时想扩展到 32K/128K,需配合 YaRN、NTK-aware 缩放或 Position Interpolation。
  • GQA 是当前大模型最优解:LLaMA-2/3、Mistral、Qwen 等主流开源模型默认使用 GQA。建议 num_kv_heads 设为 num_heads 的 1/4 到 1/8(如 64 头配 8 KV 头)。
  • MLA 是压缩极致的新选择:DeepSeek-V2/V3/R1 的实践证明,MLA 可以在精度几乎无损的前提下实现 93% 的 KV 缓存压缩,是超长上下文推理的有力候选。
  • FlashAttention 几乎零成本:不改模型结构、不改训练流程,只换一个实现就能提速 2-4 倍——没有理由不用。PyTorch 2.0+ 已内置(F.scaled_dot_product_attention)。在 H100 上,FlashAttention-3 可再提速 1.5-2.0×。
  • 长序列推理首选拓展 KV 缓存:推理阶段瓶颈不是计算而是 KV 缓存大小。GQA + 滑动窗口 + KV 量化可以大幅降低缓存压力。
  • Linear/Sparse Attention 精度仍有差距:在超长序列(32K 以上)任务上有价值,但在常规长度(4K-8K)场景不如标准注意力精确。选型时要在效率和质量间权衡。
  • RoPE 配合 YaRN/NTK 做长度外推:训练时用 4K 上下文,推理时想扩展到 32K,需配合位置插值或 NTK-aware 缩放。详见语言模型演进。
  • 大语言模型推理:GQA + FlashAttention-2/3 已成为 LLaMA、Mistral、Qwen 等所有主流大模型的标配组合。DeepSeek 系列则采用 MLA。详见语言模型演进。
  • 长文档处理:Longformer、BigBird 用于处理数千到数万 token 的长文档(法律合同、学术论文),详见NLP 基础。
  • 代码生成:代码上下文往往很长,滑动窗口注意力在代码补全模型中广泛应用。
  • 多模态模型:图像 token 数量大(如 1024 以上),稀疏注意力降低图文融合的计算开销。详见多模态模型。
  • 超长上下文(128K-1M):FlashAttention-3/4 + GQA + 滑动窗口的组合让百万级 token 的上下文窗口(如 Gemini 1.5 Pro 的 1M context)成为现实。
类库语言说明
flash-attnPythonFlashAttention 官方 CUDA 实现(FA2),Dao 等维护
flash-attn-3PythonFlashAttention-3,针对 Hopper GPU(H100)优化,支持 FP8
flash-attn-4PythonFlashAttention-4,CuTeDSL 编写,支持 Hopper 和 Blackwell(B200)
xformersPythonMeta 的高效 Transformer 库,内置 memory_efficient_attention
torch.nn.functionalPythonPyTorch 2.0+ 的 scaled_dot_product_attention,自动调用 FlashAttention
tritonPythonOpenAI 的 GPU kernel 语言,FlashAttention 和许多自定义注意力基于 Triton
vllmPython高吞吐推理框架,PagedAttention 高效管理 KV 缓存
sageattentionPython2024-2025 年兴起的量化注意力推理加速库,支持 INT8/FP8 KV 缓存
术语英文解释
多头注意力Multi-Head Attention, MHA标准 Transformer 注意力,每个头有独立的 Q/K/V 投影
多查询注意力Multi-Query Attention, MQA所有 query 头共享一组 K/V,大幅减少 KV 缓存
分组查询注意力Grouped-Query Attention, GQAquery 头分组共享 K/V,在 MHA 和 MQA 之间取得平衡
潜注意力Multi-head Latent Attention, MLADeepSeek-V2 提出,将 KV 缓存压缩为低维潜变量再按需还原,压缩率超 90%
稀疏注意力Sparse Attention只计算部分 query-key 对的注意力,如局部窗口+全局+随机
线性注意力Linear Attention用核函数近似 softmax 并重排计算,将平方复杂度降为线性
滑动窗口注意力Sliding Window Attention每个 token 只 attends 固定窗口大小内的 token
KV 缓存KV Cache自回归推理时缓存历史 token 的 K/V 向量,避免重复计算
FlashAttentionFlashAttention通过分块和在线 softmax 优化 GPU 内存访问的高效注意力实现
旋转位置编码Rotary Position Embedding, RoPE通过旋转注入相对位置信息,支持长度外推
膨胀注意力Dilated Attention类似膨胀卷积,以步长间隔取 attention 位置以扩大感受野
在线 softmaxOnline Softmax不需一次性读取全部数据,可逐块增量更新的 softmax 算法,FlashAttention 的核心
核函数Kernel Function衡量两个向量相似度的函数,Linear Attention 用简单核近似 softmax 核
HBMHigh Bandwidth MemoryGPU 主显存,容量大(40-80GB)但带宽相对有限
SRAMStatic RAMGPU 片上缓存,容量极小(~20MB)但带宽极高(约 HBM 的 10 倍)
注意力矩阵Attention MatrixQ·Kᵀ 产生的 N×N 矩阵,元素值表示 token 间的原始关注分数
TMATensor Memory AcceleratorHopper GPU 的异步数据搬运单元,FlashAttention-3 利用它重叠计算与 I/O
MFUModel FLOPs Utilization模型算力利用率,衡量实际计算吞吐占 GPU 理论峰值的比例
解耦 RoPEDecoupled RoPEMLA 中将 RoPE 应用于独立子空间的技术,避免干扰低秩压缩
PagedAttentionPagedAttentionvLLM 提出的分页式 KV 缓存管理方案,类似操作系统的虚拟内存分页
  • Shazeer,「Fast Transformer Decoding: One Write-Head is All You Need」(2019):MQA 原始论文,首次提出共享 KV 头以加速推理。
  • Ainslie et al.,「GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints」(EMNLP 2023):GQA 论文,LLaMA-2 采用的方案,详细对比了 MHA/GQA/MQA 的精度-速度权衡。
  • Dao et al.,「FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness」(NeurIPS 2022):FlashAttention 原始论文,通过 IO 感知的分块计算大幅加速注意力。
  • Dao,「FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning」(ICLR 2024):FA2,优化并行度和 warp 级效率。
  • Shah et al.,「FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision」(2024):FA3,针对 Hopper GPU 利用异步计算和 FP8,H100 上达 740 TFLOPs/s。论文:arxiv.org/abs/2407.08608
  • DeepSeek-AI,「DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model」(2024):提出 MLA,KV 缓存减少 93.3%。论文:https://arxiv.org/abs/2405.04434
  • Beltagy et al.,「Longformer: The Long-Document Transformer」(2020):滑动窗口 + 全局注意力,处理长文档的经典工作。
  • Su et al.,「RoFormer: Enhanced Transformer with Rotary Position Embedding」(2021):RoPE 论文,旋转位置编码,当前主流大模型的标配。论文:https://arxiv.org/abs/2104.09864
  • Zaheer et al.,「Big Bird: Transformers for Longer Sequences」(NeurIPS 2020):稀疏注意力理论,Turing 完备性证明。