注意力机制变体
标准自注意力的计算复杂度随序列长度 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:从分数到概率
Section titled “Softmax:从分数到概率”注意力公式中最关键的运算之一是 softmax——它把一组任意实数分数转换为概率分布(非负数,且总和为 1)。对向量 s = (s₁, s₂, …, sₙ),softmax 的定义为:
- 为什么用指数 exp? 因为指数函数永远为正,保证了输出非负;同时它放大了较大分数的差距——一个得分为 5 的 token 相比得分为 1 的 token,注意力权重会高出 e⁴ ≈ 55 倍,这就是注意力”聚焦”在相关 token 上的数学基础。
- 数值稳定性:直接计算
exp(sᵢ)在 sᵢ 较大时会导致数值溢出。实际实现中先减去最大值再做 softmax:softmax(sᵢ) = exp(sᵢ - max(s)) / Σⱼ exp(sⱼ - max(s)),结果在数学上完全等价但数值更稳定。FlashAttention 的”在线 softmax”算法正是利用了这一性质来实现分块计算。
注意力的完整公式
Section titled “注意力的完整公式”给定输入序列 X(维度 N×d),通过三个可学习的权重矩阵 W_Q、W_K、W_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 个头,每个头独立做注意力,再拼接。这让模型同时关注不同子空间的信息。
KV 缓存的显存计算
Section titled “KV 缓存的显存计算”在自回归解码(autoregressive decoding)时,每生成一个新 token,需要用到之前所有 token 的 K 和 V。为避免重复计算,将这些 K、V 存入 KV 缓存(KV Cache)。其显存占用公式为:
其中 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:
序列扩展到 32K 时,单条请求的 KV 缓存就达到约 10 GB——这就是长上下文推理的核心瓶颈。
GPU 存储层次:HBM 与 SRAM
Section titled “GPU 存储层次:HBM 与 SRAM”理解 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/s | GPU 许算核心旁边的快速缓存,容量极小但带宽极高(约 HBM 的 10 倍) |
标准注意力在计算时,需要把 N×N 的中间矩阵先写入 HBM,再读回做 softmax——大量时间花在 HBM 的读写上。FlashAttention 的核心思想是:把 Q、K、V 分成小块(tiling),加载到 SRAM 中完成全部计算,只把最终结果写回 HBM。
标准多头注意力的开销回顾
Section titled “标准多头注意力的开销回顾”标准自注意力的计算为: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 缓存公式)。
Multi-Query Attention(MQA)
Section titled “Multi-Query Attention(MQA)”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):
Grouped-Query Attention(GQA)
Section titled “Grouped-Query Attention(GQA)”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 头:
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 缓存开销。
Multi-head Latent Attention(MLA)
Section titled “Multi-head Latent Attention(MLA)”MLA 是 DeepSeek-V2(2024)提出的创新注意力机制,旨在同时降低 KV 缓存和保持精度。核心思想是低秩压缩(low-rank compression):
- 将 KV 投影到一个低维潜变量(latent vector)c_KV(维度 d_c ≪ d × num_heads),只缓存这个潜变量。
- 推理时通过一个”上投影”矩阵将潜变量还原为完整的 K、V。
DeepSeek-V2 通过 MLA 将 KV 缓存减少了 93.3%(相比 MHA),同时性能几乎不受影响。MLA 还结合了 解耦 RoPE(decoupled RoPE)策略——将 RoPE 应用于一个额外的低维子空间,避免旋转位置编码干扰低秩压缩的数学性质。
MLA 代表了注意力机制设计的新范式:不再通过减少头数(如 GQA/MQA)来压缩缓存,而是从信息瓶颈角度出发做低秩降维。DeepSeek-V3(2024)和 DeepSeek-R1(2025)都沿用了 MLA 架构。
Sparse Attention
Section titled “Sparse Attention”核心思想:将 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 的实用组合。
Linear Attention
Section titled “Linear Attention”标准注意力的 softmax 包含指数运算,无法简单重排计算顺序。Linear Attention 用核函数(kernel function)k(x,y) 近似 softmax,将 Attention(Q,K,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),但近似精度在大规模模型上仍有差距。
Sliding Window Attention
Section titled “Sliding Window Attention”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
Section titled “FlashAttention”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(2024)
Section titled “FlashAttention-3(2024)”FlashAttention-3(Shah et al., 2024)针对 NVIDIA Hopper 架构(H100 GPU)的新硬件特性进行了深度优化,引入三项关键技术:
- 异步计算重叠(Overlap via Warp-Specialization):利用 Hopper 的 Tensor Memory Accelerator(TMA,张量内存加速器)异步加载数据,同时用专门的 warp 做 softmax,实现计算与数据搬移的流水线重叠——GPU 不再”等数据”。
- 块状 matmul 与 softmax 交织(Interleaving):在 GEMM(矩阵乘)还在执行时就开始做 softmax 的归约操作,进一步隐藏延迟。
- 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(2025)
Section titled “FlashAttention-4(2025)”FlashAttention-4 使用 NVIDIA 新的 CuTeDSL(CUTLASS DSL,一种声明式 GPU kernel 语言)编写,同时优化 Hopper 和 Blackwell(B200 GPU)架构。这是 FlashAttention 系列首次官方支持最新的 Blackwell GPU,进一步提升了超长上下文训练的效率。
旋转位置编码(RoPE)
Section titled “旋转位置编码(RoPE)”RoPE(Su et al., 2021)不是改变注意力结构,而是一种位置编码方式:通过旋转向量对的方式将位置信息注入 Q 和 K,使得两个 token 的注意力分数只依赖它们的相对距离。
数学直觉:将 d 维向量视为 d/2 个二维平面,每个平面按 token 的位置 m 做角度为 mθ 的旋转。对于位置 m 的 Q 和位置 n 的 K,旋转角度差为 (m-n)θ,因此 Q·Kᵀ 只依赖相对位置 m-n——这就是”相对位置编码”的本质。
其中 θ_k = 10000^(-2k/d) 是不同维度对使用的基础频率(低维用高频、高维用低频,类似多尺度)。
RoPE 是 LLaMA、Qwen、Mistral 等主流大模型的标配,优势是天然支持相对位置和外推(extrapolation,即训练时用短序列、推理时扩展到长序列)。配合 YaRN 或 NTK-aware 缩放等技巧,可以将 4K 上下文外推到 128K 甚至更长。
注意力变体效率对比
Section titled “注意力变体效率对比”Sparse / Sliding Window 注意力模式
Section titled “Sparse / Sliding Window 注意力模式”FlashAttention 的分块计算流程
Section titled “FlashAttention 的分块计算流程”PyTorch:对比 MHA、GQA、MQA、MLA 的 KV 缓存大小
Section titled “PyTorch:对比 MHA、GQA、MQA、MLA 的 KV 缓存大小”# 配置:64 个 query 头,head_dim=128,序列长度 4096,float16num_heads, head_dim, seq_len, num_layers = 64, 128, 4096, 80dtype_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 torchimport 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, SDPBackendwith 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)numpy 手写滑动窗口 attention mask
Section titled “numpy 手写滑动窗口 attention mask”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(数值上等价)训练与推理技巧
Section titled “训练与推理技巧”- 始终启用 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)成为现实。
典型类库与工具
Section titled “典型类库与工具”| 类库 | 语言 | 说明 |
|---|---|---|
| flash-attn | Python | FlashAttention 官方 CUDA 实现(FA2),Dao 等维护 |
| flash-attn-3 | Python | FlashAttention-3,针对 Hopper GPU(H100)优化,支持 FP8 |
| flash-attn-4 | Python | FlashAttention-4,CuTeDSL 编写,支持 Hopper 和 Blackwell(B200) |
| xformers | Python | Meta 的高效 Transformer 库,内置 memory_efficient_attention |
| torch.nn.functional | Python | PyTorch 2.0+ 的 scaled_dot_product_attention,自动调用 FlashAttention |
| triton | Python | OpenAI 的 GPU kernel 语言,FlashAttention 和许多自定义注意力基于 Triton |
| vllm | Python | 高吞吐推理框架,PagedAttention 高效管理 KV 缓存 |
| sageattention | Python | 2024-2025 年兴起的量化注意力推理加速库,支持 INT8/FP8 KV 缓存 |
| 术语 | 英文 | 解释 |
|---|---|---|
| 多头注意力 | Multi-Head Attention, MHA | 标准 Transformer 注意力,每个头有独立的 Q/K/V 投影 |
| 多查询注意力 | Multi-Query Attention, MQA | 所有 query 头共享一组 K/V,大幅减少 KV 缓存 |
| 分组查询注意力 | Grouped-Query Attention, GQA | query 头分组共享 K/V,在 MHA 和 MQA 之间取得平衡 |
| 潜注意力 | Multi-head Latent Attention, MLA | DeepSeek-V2 提出,将 KV 缓存压缩为低维潜变量再按需还原,压缩率超 90% |
| 稀疏注意力 | Sparse Attention | 只计算部分 query-key 对的注意力,如局部窗口+全局+随机 |
| 线性注意力 | Linear Attention | 用核函数近似 softmax 并重排计算,将平方复杂度降为线性 |
| 滑动窗口注意力 | Sliding Window Attention | 每个 token 只 attends 固定窗口大小内的 token |
| KV 缓存 | KV Cache | 自回归推理时缓存历史 token 的 K/V 向量,避免重复计算 |
| FlashAttention | FlashAttention | 通过分块和在线 softmax 优化 GPU 内存访问的高效注意力实现 |
| 旋转位置编码 | Rotary Position Embedding, RoPE | 通过旋转注入相对位置信息,支持长度外推 |
| 膨胀注意力 | Dilated Attention | 类似膨胀卷积,以步长间隔取 attention 位置以扩大感受野 |
| 在线 softmax | Online Softmax | 不需一次性读取全部数据,可逐块增量更新的 softmax 算法,FlashAttention 的核心 |
| 核函数 | Kernel Function | 衡量两个向量相似度的函数,Linear Attention 用简单核近似 softmax 核 |
| HBM | High Bandwidth Memory | GPU 主显存,容量大(40-80GB)但带宽相对有限 |
| SRAM | Static RAM | GPU 片上缓存,容量极小(~20MB)但带宽极高(约 HBM 的 10 倍) |
| 注意力矩阵 | Attention Matrix | Q·Kᵀ 产生的 N×N 矩阵,元素值表示 token 间的原始关注分数 |
| TMA | Tensor Memory Accelerator | Hopper GPU 的异步数据搬运单元,FlashAttention-3 利用它重叠计算与 I/O |
| MFU | Model FLOPs Utilization | 模型算力利用率,衡量实际计算吞吐占 GPU 理论峰值的比例 |
| 解耦 RoPE | Decoupled RoPE | MLA 中将 RoPE 应用于独立子空间的技术,避免干扰低秩压缩 |
| PagedAttention | PagedAttention | vLLM 提出的分页式 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 完备性证明。