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)教学、调试
xFormersMeta 开源的优化库,支持多种 Attention 变体(如 Linformer、Local Attention)研究、自定义 Attention
FlashAttention斯坦福提出的 IO 优化 Attention,减少显存访问长序列训练(如 LLM)
SDPA(PyTorch 官方)PyTorch 2.0+ 内置的 F.scaled_dot_product_attention,自动选择最优后端生产环境首选

2. 性能对比实验

实验配置

我们基于 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手动实现SDPAFlashAttentionxFormers
10.000120.0000370.0000480.000084
20.000240.0000470.0000570.000091
40.000440.0000760.0000880.000122
80.000830.0001340.0001640.000186
160.001580.0002440.0002690.000302
320.003110.0004670.0004970.000529
640.006170.0008970.0009520.000980
1280.012330.0017410.0018490.001876
2560.024630.0034520.0037140.003795
5120.049060.0068590.0077510.007593
耗时 VS batch_size

关键观察:

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手动实现SDPAFlashAttentionxFormers
1106444648
2192626466
43529296100
8672152160168
161312273289305
322592514546578
64515299610601124
12810272196020882216
25620512388841444400
51240992774482568768
GPU Memory vs batch_size

关键观察:

a.SDPA 显存占用最低:在所有 batch_size 下,SDPA 的显存占用比手动实现低 50%~80%。

b. FlashAttention 和 xFormers 显存优化显著:虽然比 SDPA 稍高,但仍远低于手动实现.

c. 显存占用随 batch_size 线性增长:所有方法的显存占用都与 batch_size 成正比。

数值精度对比(L2 误差和最大误差)

batch_sizeSDPA (avg_l2)SDPA (max_err)Flash (avg_l2)Flash (max_err)xFormers (avg_l2)xFormers (max_err)
10.0002270.0004880.0002270.0004880.0002270.000488
20.0002260.0004880.0002260.0004880.0002260.000488
…………………
5120.0002270.0004880.0002270.0004880.0002270.000488

关键观察:

  • 所有方法的数值误差几乎相同,说明优化实现并未牺牲计算精度。所有实验中我们发现精度都相差不大,因此后续我们不再单独列出精度。

2.2 seq_lengths = [64, 128, 256, 512, 1024, 2048, 4096]

计算耗时对比(单位:秒)

seq_len手动实现SDPAFlashAttentionxFormers
640.000110.0000420.0000540.000088
1280.000240.0000490.0000580.000090
2560.000600.0000960.0001040.000138
5120.001770.0002590.0002810.000313
10240.006160.0009020.0009590.001000
20480.023350.0033340.0035850.003606
40960.088890.0132950.0144180.014291
耗时 VS Seq Len

观察:

  1. 指数级增长趋势:所有实现的耗时随seq_len呈O(n²)增长
  2. SDPA优势明显:
  • 短序列(64)时比手动实现快2.6倍
  • 长序列(4096)时优势扩大到6.7倍

显存占用对比(单位:MB)

seq_len手动实现SDPAFlashAttentionxFormers
6496808896
128224152160168
256544273289305
5121568514546578
1024515299610601124
204818464196020882216
409669664388841444400
GPU Memory vs Seq Len

内存优化分析:

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手动实现SDPAFlashAttentionxFormersSDPA加速比
10.000380.0000760.0000880.0001275.0x
20.000820.0001330.0001480.0001826.2x
40.001580.0002440.0002680.0003006.5x
80.003110.0004660.0005000.0005276.7x
160.006150.0008920.0009550.0009736.9x
320.012310.0017420.0018460.0018667.1x
640.024650.0034600.0036600.0038137.1x
1280.049110.0068600.0077670.0075847.2x
耗时 VS Num Heads

关键观察:

a. 线性增长特性:

  • 头数每翻倍,耗时近似翻倍
  • 验证了Attention计算的O(n)复杂度特性

b. SDPA优势稳定:

  • 加速比稳定在6-7倍区间
  • 128头时仍保持7.2倍优势

显存占用对比(单位:MB)

num_heads手动实现SDPAFlashAttentionxFormersSDPA节省率
132880889675.6%
267215216016877.4%
4131227328930579.2%
8259251454657880.2%
1651529961060112480.7%
321027219602088221680.9%
642051238884144440081.0%
1284099277448256876881.1%
GPU Memory VS Num Heads

内存优化分析:

a. 稳定节省率:

  • SDPA显存节省稳定在80%左右
  • 说明优化实现的存储复杂度与头数无关

b. 增长斜率差异:

  • 手动实现:斜率≈40MB/head
  • SDPA:斜率≈7.6MB/head
  • 优化实现的显存增长更平缓

2.4 dims = [8, 16, 32, 64, 128, 256, 512, 1024, 2048]

计算耗时对比(单位:秒)

head_dim手动实现SDPAFlashAttentionxFormersSDPA加速比
80.005140.000730.000750.000757.0x
160.005320.000740.000760.000767.2x
320.005560.000730.000690.000737.6x
640.006170.000890.000950.000986.9x
1280.007370.001620.001580.001594.6x
2560.009960.003370.003360.003413.0x
5120.015640.01762–0.017470.9x
10240.027370.03540–0.035240.8x
20480.051520.07396–0.076100.7x
耗时 VS Dim

关键观察:

a. 性能拐点:

  • dim≤64:优化效果显著(加速比>5x)
  • dim=128:加速比降至4.6x
  • dim≥512:优化实现反而更慢

b. FlashAttention表现:

  • 在dim≤256时与SDPA相当
  • 大dim时因分块策略失效退出竞争

c. xFormers稳定性:

  • 始终略慢于SDPA
  • 但支持超大dim计算

显存占用对比(单位:MB)

head_dim手动实现SDPAFlashAttentionxFormersSDPA节省率
8422413214816496.9%
16438427629230893.7%
32464051654858088.9%
6451529961060112480.7%
128617619562084221268.3%
256822438764132438852.9%
512123209760–976020.8%
10241948818464–194885.3%
20483484836896–38944-5.9%
GPU Memory VS Dim

存储模式分析:

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 优化
xFormersC++/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

  1. 如何优化transformer的attention?
  2. torch.nn.functional.scaled_dot_product_attention – PyTorch 2.7 documentation
  3. https://github.com/facebookresearch/xformers
  4. 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有三种实现:

  1. 基于FlashAttention-2的实现;
  2. Memory-Efficient Attention(facebook xformers);
  3. Pytorch版本对上述伪代码的c++实现(对应MATH后端)。

针对CUDA后端,SDPA可能会调用经过优化的内核以提高性能。对于所有其他后端,将使用PyTorch实现。所有实现方式默认都是启用的,SDPA会尝试根据输入自动选择最优的实现方式。为了对使用哪种实现方式提供更细粒度的控制,torch提供了以下函数来启用和禁用各种实现方式:

  1. torch.nn.attention.sdpa_kernel:一个上下文管理器,用于启用或禁用任何一种实现方式;
  2. torch.backends.cuda.enable_flash_sdp:全局启用或禁用FlashAttention;
  3. torch.backends.cuda.enable_mem_efficient_sdp:全局启用或禁用memory efficient attention;
  4. 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后端。

5. 参考

  1. FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning
  2. Memory-Efficient Attention
  3. Grouped-Query Attention
  4. Attention Is All You Need

发表回复

您的邮箱地址不会被公开。 必填项已用 * 标注