Attention 机制哪家强?SDPA、FlashAttention、xFormers、手动实现全面对比
目录
收起
1. 四种 Attention 实现方式简介
2. 性能对比实验
实验配置
2.1 batch size =[1, 2, 4, 8, 16, 32, 64, 128, 256, 512]
显存占用对比(单位:MB)
数值精度对比(L2 误差和最大误差)
2.2 seq_lengths = [64, 128, 256, 512, 1024, 2048, 4096]
计算耗时对比(单位:秒)
显存占用对比(单位:MB)
2.3 num_heads = [1, 2, 4, 8, 16, 32, 64, 128]
计算耗时对比(单位:秒)
显存占用对比(单位:MB)
2.4 dims = [8, 16, 32, 64, 128, 256, 512, 1024, 2048]
计算耗时对比(单位:秒)
显存占用对比(单位:MB)
建议:
a. 对于大 dim(如 dim ≥ 512):
3. 关键结论
✅ SDPA(PyTorch 官方)综合最优
⚡ FlashAttention 适合长序列
⏳ xFormers 灵活性高
手动实现最慢
4. 优化建议
5. 测试代码
References
近年来,Transformer 模型在 NLP、CV 等领域大放异彩,而 Attention(注意力机制) 是其核心组件。不同的 Attention 实现方式(如 PyTorch 官方的 scaled_dot_product_attention(SDPA)、FlashAttention、xFormers 和手动实现)在计算效率、显存占用、数值精度等方面表现如何?
本文基于实测数据(实验代码见文末),对比不同 Attention 实现在不同 batch_size(批量数据),sequence_length(序列长度),num_heads(注意力头的数量),head_dim(注意力头的维度)下的性能,并分析其背后的优化原理。
1. 四种 Attention 实现方式简介
| 实现方式 | 特点 | 适用场景 |
|---|---|---|
| 手动实现 | 纯 PyTorch 代码,逐步骤计算 QK^T → softmax → mask → matmul(V) | 教学、调试 |
| xFormers | Meta 开源的优化库,支持多种 Attention 变体(如 Linformer、Local Attention) | 研究、自定义 Attention |
| FlashAttention | 斯坦福提出的 IO 优化 Attention,减少显存访问 | 长序列训练(如 LLM) |
| SDPA(PyTorch 官方) | PyTorch 2.0+ 内置的 F.scaled_dot_product_attention,自动选择最优后端 | 生产环境首选 |
2. 性能对比实验
实验配置
- GPU: NVIDIA H800 (80GB)
- CUDA Version: 12.2
- flash-attn: 2.7.4.post1
- torch: 2.3.1
- torchvision: 0.18.1
- xformers: 0.0.27
我们基于 Vit-large 的基本配置来做对比实验,假设图像分辨率为512 x 512,patch_size 大小为16,此时序列长度为 512/16 x 512/16 = 1024 作为参照,batch_size 以 64 作为参照,num_heads 以 16 作为参照,embedding_dim 为 1024,此时 head_dim = embedding_dim/num_heads = 64 作为参照。
即我们设定【B,N, H, D】 = 【64, 1024, 16, 64】,当我们分析其中一个变量的影响时,我们固定另外3个参数的数值不变。我们的评价指标为不同 attention 方法的耗时,GPU Memory 峰值,以及精度对比。
2.1 batch size =[1, 2, 4, 8, 16, 32, 64, 128, 256, 512]
计算耗时对比(单位:秒)
| batch_size | 手动实现 | SDPA | FlashAttention | xFormers |
|---|---|---|---|---|
| 1 | 0.00012 | 0.000037 | 0.000048 | 0.000084 |
| 2 | 0.00024 | 0.000047 | 0.000057 | 0.000091 |
| 4 | 0.00044 | 0.000076 | 0.000088 | 0.000122 |
| 8 | 0.00083 | 0.000134 | 0.000164 | 0.000186 |
| 16 | 0.00158 | 0.000244 | 0.000269 | 0.000302 |
| 32 | 0.00311 | 0.000467 | 0.000497 | 0.000529 |
| 64 | 0.00617 | 0.000897 | 0.000952 | 0.000980 |
| 128 | 0.01233 | 0.001741 | 0.001849 | 0.001876 |
| 256 | 0.02463 | 0.003452 | 0.003714 | 0.003795 |
| 512 | 0.04906 | 0.006859 | 0.007751 | 0.007593 |

关键观察:
a. SDPA 始终最快:在所有 batch_size 下,PyTorch 的 scaled_dot_product_attention(SDPA)都比手动实现快 5~10 倍。
b. FlashAttention 略慢于 SDPA:FlashAttention 计算时间比 SDPA 稍长(约 5%~10%),但仍远快于手动。
c. xFormers 稍慢:xFormers 的实现比 FlashAttention 略慢,但差距不大。
d. 手动实现最慢:随着 batch_size 增大,所有方法的显存占用实现的耗时线性增长。
显存占用对比(单位:MB)
| batch_size | 手动实现 | SDPA | FlashAttention | xFormers |
|---|---|---|---|---|
| 1 | 106 | 44 | 46 | 48 |
| 2 | 192 | 62 | 64 | 66 |
| 4 | 352 | 92 | 96 | 100 |
| 8 | 672 | 152 | 160 | 168 |
| 16 | 1312 | 273 | 289 | 305 |
| 32 | 2592 | 514 | 546 | 578 |
| 64 | 5152 | 996 | 1060 | 1124 |
| 128 | 10272 | 1960 | 2088 | 2216 |
| 256 | 20512 | 3888 | 4144 | 4400 |
| 512 | 40992 | 7744 | 8256 | 8768 |

关键观察:
a.SDPA 显存占用最低:在所有 batch_size 下,SDPA 的显存占用比手动实现低 50%~80%。
b. FlashAttention 和 xFormers 显存优化显著:虽然比 SDPA 稍高,但仍远低于手动实现.
c. 显存占用随 batch_size 线性增长:所有方法的显存占用都与 batch_size 成正比。
数值精度对比(L2 误差和最大误差)
| batch_size | SDPA (avg_l2) | SDPA (max_err) | Flash (avg_l2) | Flash (max_err) | xFormers (avg_l2) | xFormers (max_err) |
|---|---|---|---|---|---|---|
| 1 | 0.000227 | 0.000488 | 0.000227 | 0.000488 | 0.000227 | 0.000488 |
| 2 | 0.000226 | 0.000488 | 0.000226 | 0.000488 | 0.000226 | 0.000488 |
| … | … | … | … | … | … | … |
| 512 | 0.000227 | 0.000488 | 0.000227 | 0.000488 | 0.000227 | 0.000488 |
关键观察:
- 所有方法的数值误差几乎相同,说明优化实现并未牺牲计算精度。所有实验中我们发现精度都相差不大,因此后续我们不再单独列出精度。
2.2 seq_lengths = [64, 128, 256, 512, 1024, 2048, 4096]
计算耗时对比(单位:秒)
| seq_len | 手动实现 | SDPA | FlashAttention | xFormers |
|---|---|---|---|---|
| 64 | 0.00011 | 0.000042 | 0.000054 | 0.000088 |
| 128 | 0.00024 | 0.000049 | 0.000058 | 0.000090 |
| 256 | 0.00060 | 0.000096 | 0.000104 | 0.000138 |
| 512 | 0.00177 | 0.000259 | 0.000281 | 0.000313 |
| 1024 | 0.00616 | 0.000902 | 0.000959 | 0.001000 |
| 2048 | 0.02335 | 0.003334 | 0.003585 | 0.003606 |
| 4096 | 0.08889 | 0.013295 | 0.014418 | 0.014291 |

观察:
- 指数级增长趋势:所有实现的耗时随seq_len呈O(n²)增长
- SDPA优势明显:
- 短序列(64)时比手动实现快2.6倍
- 长序列(4096)时优势扩大到6.7倍
显存占用对比(单位:MB)
| seq_len | 手动实现 | SDPA | FlashAttention | xFormers |
|---|---|---|---|---|
| 64 | 96 | 80 | 88 | 96 |
| 128 | 224 | 152 | 160 | 168 |
| 256 | 544 | 273 | 289 | 305 |
| 512 | 1568 | 514 | 546 | 578 |
| 1024 | 5152 | 996 | 1060 | 1124 |
| 2048 | 18464 | 1960 | 2088 | 2216 |
| 4096 | 69664 | 3888 | 4144 | 4400 |

内存优化分析:
a. SDPA始终最优:
- 64长度时节省16.7%显存
- 4096长度时显存节省达94.4%
b. FlashAttention表现:
- 在2048长度时比SDPA多占6.5%显存
- 但仍比手动实现节省88.7%
2.3 num_heads = [1, 2, 4, 8, 16, 32, 64, 128]
计算耗时对比(单位:秒)
| num_heads | 手动实现 | SDPA | FlashAttention | xFormers | SDPA加速比 |
|---|---|---|---|---|---|
| 1 | 0.00038 | 0.000076 | 0.000088 | 0.000127 | 5.0x |
| 2 | 0.00082 | 0.000133 | 0.000148 | 0.000182 | 6.2x |
| 4 | 0.00158 | 0.000244 | 0.000268 | 0.000300 | 6.5x |
| 8 | 0.00311 | 0.000466 | 0.000500 | 0.000527 | 6.7x |
| 16 | 0.00615 | 0.000892 | 0.000955 | 0.000973 | 6.9x |
| 32 | 0.01231 | 0.001742 | 0.001846 | 0.001866 | 7.1x |
| 64 | 0.02465 | 0.003460 | 0.003660 | 0.003813 | 7.1x |
| 128 | 0.04911 | 0.006860 | 0.007767 | 0.007584 | 7.2x |

关键观察:
a. 线性增长特性:
- 头数每翻倍,耗时近似翻倍
- 验证了Attention计算的O(n)复杂度特性
b. SDPA优势稳定:
- 加速比稳定在6-7倍区间
- 128头时仍保持7.2倍优势
显存占用对比(单位:MB)
| num_heads | 手动实现 | SDPA | FlashAttention | xFormers | SDPA节省率 |
|---|---|---|---|---|---|
| 1 | 328 | 80 | 88 | 96 | 75.6% |
| 2 | 672 | 152 | 160 | 168 | 77.4% |
| 4 | 1312 | 273 | 289 | 305 | 79.2% |
| 8 | 2592 | 514 | 546 | 578 | 80.2% |
| 16 | 5152 | 996 | 1060 | 1124 | 80.7% |
| 32 | 10272 | 1960 | 2088 | 2216 | 80.9% |
| 64 | 20512 | 3888 | 4144 | 4400 | 81.0% |
| 128 | 40992 | 7744 | 8256 | 8768 | 81.1% |

内存优化分析:
a. 稳定节省率:
- SDPA显存节省稳定在80%左右
- 说明优化实现的存储复杂度与头数无关
b. 增长斜率差异:
- 手动实现:斜率≈40MB/head
- SDPA:斜率≈7.6MB/head
- 优化实现的显存增长更平缓
2.4 dims = [8, 16, 32, 64, 128, 256, 512, 1024, 2048]
计算耗时对比(单位:秒)
| head_dim | 手动实现 | SDPA | FlashAttention | xFormers | SDPA加速比 |
|---|---|---|---|---|---|
| 8 | 0.00514 | 0.00073 | 0.00075 | 0.00075 | 7.0x |
| 16 | 0.00532 | 0.00074 | 0.00076 | 0.00076 | 7.2x |
| 32 | 0.00556 | 0.00073 | 0.00069 | 0.00073 | 7.6x |
| 64 | 0.00617 | 0.00089 | 0.00095 | 0.00098 | 6.9x |
| 128 | 0.00737 | 0.00162 | 0.00158 | 0.00159 | 4.6x |
| 256 | 0.00996 | 0.00337 | 0.00336 | 0.00341 | 3.0x |
| 512 | 0.01564 | 0.01762 | – | 0.01747 | 0.9x |
| 1024 | 0.02737 | 0.03540 | – | 0.03524 | 0.8x |
| 2048 | 0.05152 | 0.07396 | – | 0.07610 | 0.7x |

关键观察:
a. 性能拐点:
- dim≤64:优化效果显著(加速比>5x)
- dim=128:加速比降至4.6x
- dim≥512:优化实现反而更慢
b. FlashAttention表现:
- 在dim≤256时与SDPA相当
- 大dim时因分块策略失效退出竞争
c. xFormers稳定性:
- 始终略慢于SDPA
- 但支持超大dim计算
显存占用对比(单位:MB)
| head_dim | 手动实现 | SDPA | FlashAttention | xFormers | SDPA节省率 |
|---|---|---|---|---|---|
| 8 | 4224 | 132 | 148 | 164 | 96.9% |
| 16 | 4384 | 276 | 292 | 308 | 93.7% |
| 32 | 4640 | 516 | 548 | 580 | 88.9% |
| 64 | 5152 | 996 | 1060 | 1124 | 80.7% |
| 128 | 6176 | 1956 | 2084 | 2212 | 68.3% |
| 256 | 8224 | 3876 | 4132 | 4388 | 52.9% |
| 512 | 12320 | 9760 | – | 9760 | 20.8% |
| 1024 | 19488 | 18464 | – | 19488 | 5.3% |
| 2048 | 34848 | 36896 | – | 38944 | -5.9% |

存储模式分析:
a. 小dim优势:
- dim=8时节省96.9%显存
- 得益于不存储中间QK^T矩阵
b. 临界点现象:
- dim=512时节省率骤降至20.8%
- 说明优化策略发生本质变化
c. 负优化情况:
- dim=2048时SDPA显存反超
- 因需要额外内存处理大矩阵
总结:为什么 dim 增大时 SDPA/FlashAttention 变慢?
| 原因 | 影响 |
|---|---|
| 计算复杂度增加 | QK^T 的计算量 ( ^2⋅ ) 随 dim 线性增长 |
| 显存带宽压力 | Q, K, V 的显存占用增加,数据传输变慢 |
| 分块计算效率下降 | FlashAttention 的 tiling 策略在大 dim 时效果变差 |
| Tensor Core 利用率降低 | dim 过大时无法最优匹配 Tensor Core 计算模式 |
| 内核启动开销增加 | 大 dim 可能需要拆分成多个 CUDA 内核 |
建议:
a. 对于大 dim(如 dim ≥ 512):
- 尝试减少
num_heads,保持head_dim在64或128左右(如dim=512→num_heads=8,head_dim=64)。 - 使用混合精度(
FP16/BF16)以提升 Tensor Core 利用率。
b. 如果 dim 极大(如 2048):
- 考虑是否真的需要这么大的
dim,或改用低秩近似(如 Linformer)。
3. 关键结论
| 实现 | 优化层级 | 内存占用 | 速度 | 支持性 | 备注 |
|---|---|---|---|---|---|
| SDPA | 自动选择最优kernel | 最优 | 最快 | 强 | 官方主推 |
| FlashAttention | 手写 CUDA kernel | 很低 | 接近最快 | 限制多 | 专为 attention 优化 |
| xFormers | C++/CUDA 实现 | 中等 | 次之 | 通用 | 多 attention 变种 |
| 手动实现 | PyTorch Eager 模式 | 高 | 最慢 | 高 | 教学用,不建议生产 |
✅ SDPA(PyTorch 官方)综合最优
- 小
dim时最快(自动选择 FlashAttention 或 Tensor Core 优化)。 - 显存占用最低(适合训练大模型)。
- 生产环境首选(无需额外安装,兼容性强)。
⚡ FlashAttention 适合长序列
- 显存优化显著(适合
seq_len > 2048的场景)。 - 但
dim过大时(如≥1024),计算效率可能下降。
⏳ xFormers 灵活性高
- 支持多种 Attention 变体(如稀疏 Attention)。
- 但通用实现不如 SDPA/FlashAttention 快。
手动实现最慢
- 无算子融合,显存占用高。
- 仅建议用于调试或教学。
4. 优化建议
a. 优先使用 F.scaled_dot_product_attention(PyTorch ≥ 2.0)。
b. dim 不宜过大(推荐 64~128,太大时考虑减少 num_heads)。
c. 长序列训练用 FlashAttention(如 seq_len > 2048)。
d. 混合精度(FP16/BF16)提升速度(需 GPU 支持)。
5. 测试代码
import math
import time
import torch
import torch.nn.functional as F
from einops import rearrange
from flash_attn import flash_attn_func
from xformers.ops import memory_efficient_attention, LowerTriangularMask
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
# 自定义手写注意力
def custom_attention(q, k, v, causal=False):
score = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(q.size(-1))
if causal:
mask = torch.triu(torch.ones(score.shape[-2], score.shape[-1]), diagonal=1)
mask = mask.masked_fill(mask==1, torch.finfo(q.dtype).min)
mask = mask.to(q.device, q.dtype)
score = score + mask
attn = F.softmax(score, dim=-1)
o = torch.matmul(attn, v)
return o
def pytorch_func(q, k, v, causal=False):
o = F.scaled_dot_product_attention(q, k, v, is_causal=causal)
return o
# FlashAttention
def flash_attention(q, k, v, causal=False):
# FlashAttention 限制:D 不能超过 256
if q.size(-1) > 256:
raise ValueError("FlashAttention only supports head dimension up to 256.")
return flash_attn_func(q, k, v, causal=causal)
# xFormers Memory Efficient Attention
def xformers_attention(q, k, v, causal=False):
return memory_efficient_attention(q, k, v, attn_bias=None)
# 计算平均 L2 距离
def avg_l2_distance(tensor1, tensor2):
return torch.mean(torch.norm(tensor1 - tensor2, p=2, dim=-1))
# 计算最大绝对误差
def max_absolute_error(tensor1, tensor2):
return torch.max(torch.abs(tensor1 - tensor2))
# 测试函数
def test(func_name, q, k, v, *args, **kwargs):
# 确保使用 FlashAttention 时,D <= 256
if func_name == "flash_attention" and q.size(-1) > 256:
print("FlashAttention not supported for head dimension larger than 256. Skipping.")
return None, 0, 0
if func_name in ["custom_attention", "pytorch_func"]:
q = rearrange(q, "a b c d -> a c b d")
k = rearrange(k, "a b c d -> a c b d")
v = rearrange(v, "a b c d -> a c b d")
torch.cuda.reset_peak_memory_stats()
torch.cuda.synchronize()
for _ in range(5):
o = globals()[func_name](q, k, v, *args, **kwargs)
torch.cuda.synchronize()
st = time.time()
o = globals()[func_name](q, k, v, *args, **kwargs)
torch.cuda.synchronize()
tt = time.time() - st
max_memory = torch.cuda.max_memory_allocated() // 2**20 # 转换为MB
torch.cuda.empty_cache()
if func_name in ["custom_attention", "pytorch_func"]:
o = rearrange(o, "a c b d -> a b c d")
return o, tt, max_memory
# 主程序
if __name__ == "__main__":
# 定义所有的参数组合范围
# batch_sizes = [1, 2, 4, 8, 16, 32, 64, 128, 256, 512]
batch_sizes = [64]
# seq_lengths = [64, 128, 256, 512, 1024, 2048, 4096]
seq_lengths = [1024]
# num_heads = [1, 2, 4, 8, 16, 32, 64, 128]
num_heads = [16]
dims = [8, 16, 32, 64, 128, 256, 512, 1024, 2048]
# dims = [64]
# 存储每个测试的时间、内存和误差
results = []
# 遍历所有可能的超参数组合
for bsz in batch_sizes:
for sql in seq_lengths:
for nh in num_heads:
for hd in dims:
dtype = torch.float16
causal = False # 或根据需要设置为 False
print(f"===> Testing B={bsz}, N={sql}, H={nh}, D={hd}, dtype={dtype}, causal={causal}")
# 构造 q, k, v
q = torch.randn((bsz, sql, nh, hd)).to("cuda:0", dtype)
k = torch.rand_like(q)
v = torch.rand_like(q)
# 手写 Attention
o_ref, t_ref, m_ref = test("custom_attention", q, k, v, causal=causal)
# PyTorch Attention
o_sdpa, t_sdpa, m_sdpa = test("pytorch_func", q, k, v, causal=causal)
# Flash Attention
try:
o_flash, t_flash, m_flash = test("flash_attention", q, k, v, causal=causal)
except ValueError as e:
o_flash, t_flash, m_flash = None, None, None
# xFormers Memory Efficient Attention
o_xf, t_xf, m_xf = test("xformers_attention", q, k, v, causal=causal)
# 计算平均 L2 距离和最大绝对误差
avg_l2_sdpa = avg_l2_distance(o_ref, o_sdpa) if o_ref is not None and o_sdpa is not None else None
max_err_sdpa = max_absolute_error(o_ref, o_sdpa) if o_ref is not None and o_sdpa is not None else None
avg_l2_flash = avg_l2_distance(o_ref, o_flash) if o_ref is not None and o_flash is not None else None
max_err_flash = max_absolute_error(o_ref, o_flash) if o_ref is not None and o_flash is not None else None
avg_l2_xf = avg_l2_distance(o_ref, o_xf) if o_ref is not None and o_xf is not None else None
max_err_xf = max_absolute_error(o_ref, o_xf) if o_ref is not None and o_xf is not None else None
# 保存结果
results.append({
"batch_size": bsz,
"seq_len": sql,
"num_heads": nh,
"dim": hd,
"time_manual": t_ref,
"mem_manual": m_ref,
"time_sdpa": t_sdpa,
"mem_sdpa": m_sdpa,
"time_flash": t_flash if o_flash is not None else None,
"mem_flash": m_flash if o_flash is not None else None,
"time_xf": t_xf,
"mem_xf": m_xf,
"avg_l2_sdpa": avg_l2_sdpa.item() if avg_l2_sdpa is not None else None,
"max_err_sdpa": max_err_sdpa.item() if max_err_sdpa is not None else None,
"avg_l2_flash": avg_l2_flash.item() if avg_l2_flash is not None else None,
"max_err_flash": max_err_flash.item() if max_err_flash is not None else None,
"avg_l2_xf": avg_l2_xf.item() if avg_l2_xf is not None else None,
"max_err_xf": max_err_xf.item() if max_err_xf is not None else None,
})
# 将结果转为 DataFrame
df = pd.DataFrame(results)
df.to_csv("attention_test_results.csv", index=False)
References
- 如何优化transformer的attention?
- torch.nn.functional.scaled_dot_product_attention – PyTorch 2.7 documentation
- https://github.com/facebookresearch/xformers
- https://github.com/Dao-AILab/flash-attention
================
Scaled_dot_product_attention(SDPA)使用详解
在学习huggingFace的Transformer库时,我们不可避免会遇到scaled_dot_product_attention(SDPA)这个函数,它被用来加速大模型的Attention计算,本文就详细介绍一下它的使用方法,核心内容主要参考了torch.nn.functional中该函数的注释。
1. Attention计算公式
Attention的计算主要涉及三个矩阵:Q、K、V。我们先不考虑multi-head attention,只考虑one head的self attention。在大模型的prefill阶段,这三个矩阵的维度均为N x d,N即为上下文的长度;在decode阶段,Q的维度为1 x d, KV还是N x d。然后通过下面的公式计算attention矩阵:
在真正使用attention的时候,我们往往采用multi-head attention(MHA)。MHA的计算公式和one head attention基本一致,它改变了Q、K、V每一行的定义:将维度d的向量分成h组变成一个h x dk的矩阵,Q、K、V此时成为了 的三维矩阵(不考虑batch维)。分别将Q、K、V的第一和第二维进行转置得到三个维度为 的三维矩阵。此时的三个矩阵就是具有h个头的Q、K、V,我们就可以按照self attention的定义计算h个头的attention值。
不过,在真正进行大模型推理的时候就会发现KV Cache是非常占显存的,所以大家尝试各种手段压缩KV Cache,具体可以参考《大模型推理–KV Cache压缩》。一种手段就是将MHA替换成group query attention(GQA),这块在torch2.5以上的SDPA中也已经得到了支持。
2. SDPA伪代码
在SDPA的注释中,给出了伪代码:
def scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0.0,
is_causal=False, scale=None, enable_gqa=False) -> torch.Tensor:
L, S = query.size(-2), key.size(-2)
scale_factor = 1 / math.sqrt(query.size(-1)) if scale is None else scale
attn_bias = torch.zeros(L, S, dtype=query.dtype)
if is_causal:
assert attn_mask is None
temp_mask = torch.ones(L, S, dtype=torch.bool).tril(diagonal=0)
attn_bias.masked_fill_(temp_mask.logical_not(), float("-inf"))
attn_bias.to(query.dtype)
if attn_mask is not None:
if attn_mask.dtype == torch.bool:
attn_bias.masked_fill_(attn_mask.logical_not(), float("-inf"))
else:
attn_bias += attn_mask
if enable_gqa:
key = key.repeat_interleave(query.size(-3)//key.size(-3), -3)
value = value.repeat_interleave(query.size(-3)//value.size(-3), -3)
attn_weight = query @ key.transpose(-2, -1) * scale_factor
attn_weight += attn_bias
attn_weight = torch.softmax(attn_weight, dim=-1)
attn_weight = torch.dropout(attn_weight, dropout_p, train=True)
return attn_weight @ value
可以看出,我们实际在使用SDPA时除了query、key和value之外,还有另外几个参数:attn_mask、dropout_p、is_causal、scale和enable_gqa。scale就是计算Attention时的缩放因子,一般无需传递。dropout_p表示Dropout概率,在推理阶段也不需要传递,不过官方建议如下输入:dropout_p=(self.p if self.training else 0.0)。我们着重看一下另外三个参数在使用时该如何设置。
先看enable_gqa。前面提到GQA是一种KV Cache压缩方法,MHA的KV和Q一样,也会有h个头,GQA则将KV的h个头进行压缩来减小KV Cache的大小。比如Qwen2-7B-Instruct这个模型,Q的h等于28,KV的h等于4,相当于把KV Cache压缩到之前的七分之一。GQA虽然压缩了KV Cache,但是真正要计算Attention的时候还是需要对齐KV与Q的head数,所以我们可以看到HF Transformer库中的qwen2.py在Attention计算时会有一个repeat_kv的操作,目的就是将QKV的head数统一。在torch2.5以后的版本中,我们无需再手动去执行repeat_kv,直接将SDPA的enable_gqa设置为True即可自动完成repeat_kv,而且速度比自己去做repaet_kv还要更快。
attn_mask和is_causal两个参数的作用相同,目的都是要给softmax之前的QKT矩阵添加mask。只不过attn_mask是自己在外面构造mask矩阵,is_causal则是根据大模型推理的阶段属于prefill还是decode来进行设置。通过看伪代码可以看出,SDPA会首先构造一个L x S的零矩阵attn_bias,L表示Q的上下文长度,S表示KV Cache的长度。在prefill阶段,L和S相等,在decode阶段,L为1,S还是N。所以在prefill阶段,attn_bias就是一个N x N的矩阵,将is_causal设置为True时就会构造一个下三角为0,上三角为负无穷的矩阵作为attn_bias,然后将其加到QKT矩阵上,这样就实现了因果关系的Attention计算。在decode阶段,attn_bias就是一个1 x N的向量,此时可以将is_causal设置为False,attn_bias始终为0就不会对 行向量产生影响,表示KV Cache所有的行都参与计算,因果关系保持正确。
attn_mask作用和is_causal一样,但是需要我们自行构造,如果你对如何构造不了解建议就使用is_causal选项,prefill阶段设置为True,decode阶段设置为False,attn_mask设置为None。不过,如果prefill按照chunk来执行也即chunk_prefill阶段,我们会发现is_causal设置为True时的attn_bias设置的不正确,我们不是从左上角开始构造下三角矩阵,而是要从右下角开始构造下三角矩阵,这种情况下我们可以从外面自行构造attn_mask矩阵代替SDPA的构造。attn_mask有两种构造方式,一种是bool类型,True的位置会保持不变,False的位置会置为负无穷;一种是float类型,会直接将attn_mask加到SDPA内部的attn_bias上,和bool类型一样,我们一般是构造一个下三角为0上三角为负无穷的矩阵。总结来说,绝大多数情况下我们只需要设置is_causal选项,prefill阶段设置为True,decode阶段设置为False,attn_mask设置为None即可。如果推理阶段引入了chunk_prefill,则我们需要自行构造attn_mask,但是要注意构造的attn_mask矩阵是从右下角开始的下三角矩阵。
1. SDPA实现(翻译自SDPA注释)
目前SDPA有三种实现:
- 基于FlashAttention-2的实现;
- Memory-Efficient Attention(facebook xformers);
- Pytorch版本对上述伪代码的c++实现(对应MATH后端)。
针对CUDA后端,SDPA可能会调用经过优化的内核以提高性能。对于所有其他后端,将使用PyTorch实现。所有实现方式默认都是启用的,SDPA会尝试根据输入自动选择最优的实现方式。为了对使用哪种实现方式提供更细粒度的控制,torch提供了以下函数来启用和禁用各种实现方式:
- torch.nn.attention.sdpa_kernel:一个上下文管理器,用于启用或禁用任何一种实现方式;
- torch.backends.cuda.enable_flash_sdp:全局启用或禁用FlashAttention;
- torch.backends.cuda.enable_mem_efficient_sdp:全局启用或禁用memory efficient attention;
- torch.backends.cuda.enable_math_sdp:全局启用或禁用PyTorch的C++实现。
每个融合内核都有特定的输入限制。如果用户需要使用特定的融合实现方式,请使用torch.nn.attention.sdpa_kernel禁用PyTorch的C++实现。如果某个融合实现方式不可用,将会发出警告,说明该融合实现方式无法运行的原因。由于融合浮点运算的特性,此函数的输出可能会因所选择的后端内核而异。C++实现支持torch.float64,当需要更高精度时可以使用。对于math后端,如果输入是torch.half或torch.bfloat16类型,那么所有中间计算结果都会保持为torch.float类型。
4. SDPA使用示例
首先强调一点,灌入SDPA的QKV都是做过转置的,也即维度为batch x head x N x d,在老版本的torch中还需要QKV都是contiguous的,新版本下无此要求。SDPA注释中还给了两个示例,我们在此也给出:
# Optionally use the context manager to ensure one of the fused kernels is run
query = torch.rand(32, 8, 128, 64, dtype=torch.float16, device="cuda")
key = torch.rand(32, 8, 128, 64, dtype=torch.float16, device="cuda")
value = torch.rand(32, 8, 128, 64, dtype=torch.float16, device="cuda")
with sdpa_kernel(backends=[SDPBackend.FLASH_ATTENTION]):
F.scaled_dot_product_attention(query,key,value)
上述示例中,给定的输入为batch等于32,head等于8,上下文长度128,embedding维度64,然后通过sdpa_kernel选择使用FlashAttention。
示例二:
# Sample for GQA for llama3
query = torch.rand(32, 32, 128, 64, dtype=torch.float16, device="cuda")
key = torch.rand(32, 8, 128, 64, dtype=torch.float16, device="cuda")
value = torch.rand(32, 8, 128, 64, dtype=torch.float16, device="cuda")
with sdpa_kernel(backends=[SDPBackend.MATH]):
F.scaled_dot_product_attention(query,key,value,enable_gqa=True)
示例二演示了GQA的用法,给定的query head数为32,key和value均为8,此时我们可以通过enable_gqa选项来实现对GQA的支持,此外代码还通过sdpa_kernel选项使用了MATH后端。