PyTorch Flex Attention 完全指南:从 score_mod 到 BlockMask 的官方 API 深度解析
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
导读
torch.nn.attention.flex_attention是 PyTorch 中的原型(prototype)特性,实现了一种可编程的注意力机制:在标准缩放点积注意力(SDPA)基础上,通过用户自定义的score_mod函数任意修改注意力分数,并通过BlockMask以块稀疏(block-sparse)方式跳过无关计算。本指南以官方 API 参考文档为骨架,逐项剖析flex_attention及其配套的BlockMask、FlexKernelOptions、掩码工具函数与辅助输出接口,并结合仓库源码(torch/nn/attention/flex_attention.py)与测试用例(test/inductor/test_flex_attention.py)给出可直接运行的示例。读完本文,你将能够熟练编写score_mod/mask_mod、构造与切片BlockMask、调优 Triton kernel 选项,并在训练与推理场景中正确使用 GQA、辅助输出与多后端支持。
官方 API 参考入口:docs/source/nn.attention.flex_attention.md。
一、核心入口:flex_attention函数
1.1 函数签名与参数
flex_attention是模块的用户入口,其完整签名如下:
flex_attention( query: Tensor, key: Tensor, value: Tensor, score_mod: Callable[[Tensor, Tensor, Tensor, Tensor, Tensor], Tensor] | None = None, block_mask: BlockMask | None = None, scale: float | None = None, enable_gqa: bool = False, return_lse: bool = False, kernel_options: FlexKernelOptions | None = None, *, return_aux: AuxRequest | None = None, ) -> Tensor | tuple[Tensor, Tensor] | tuple[Tensor, AuxOutput]各参数说明(依据源码 torch/nn/attention/flex_attention.py 的 docstring 整理):
| 参数 | 类型 | 说明 |
|---|---|---|
query | Tensor | Query 张量,形状(B, Hq, L, E)。FP8 数据类型下建议行主序(row-major)内存布局以获得最佳性能 |
key | Tensor | Key 张量,形状(B, Hkv, S, E)。FP8 下建议行主序 |
value | Tensor | Value 张量,形状(B, Hkv, S, Ev)。FP8 下建议列主序(column-major) |
score_mod | Callable,可选 | 修改注意力分数的函数,默认不应用任何修改(_identity) |
block_mask | BlockMask,可选 | 控制注意力块稀疏模式的 BlockMask 对象,默认生成一个覆盖全长度的空块掩码 |
scale | float,可选 | softmax 前应用的缩放因子,默认为1/sqrt(E) |
enable_gqa | bool | 设为 True 时启用 Grouped Query Attention(GQA),将 key/value 头广播到 query 头 |
return_lse | bool | 是否返回注意力的 logsumexp,已废弃,请改用return_aux=AuxRequest(lse=True) |
kernel_options | FlexKernelOptions,可选 | 控制底层 Triton kernel 行为的选项 |
return_aux | AuxRequest,可选 | 指定要计算并返回的辅助输出 |
返回值:注意力输出张量,形状(B, Hq, L, Ev);当指定return_aux时返回(output, aux),其中aux为AuxOutput。
1.2score_mod:分数修改函数的契约
score_mod在 query 与 key 计算出的注意力分数之后、softmax 之前被调用,其签名为:
def score_mod( score: Tensor, # 标量张量,与 q/k/v 同 dtype、同设备 batch: Tensor, # 批索引,torch.int,与 score 同设备 head: Tensor, # query 头索引,torch.int q_idx: Tensor, # query 索引,torch.int k_idx: Tensor, # key/value 索引,torch.int ) -> Tensor:源码中通过_get_mod_type检查函数位置参数个数来区分两类修改函数(见 flex_attention.py):
- 5 个位置参数→ 判定为
score_mod(SCORE_MOD 类型); - 4 个位置参数→ 判定为
mask_mod(MASK_MOD 类型); - 其他个数直接抛出
AssertionError。
例如一个最简的因果分数修改函数:
def causal_score_mod(score, batch, head, q_idx, k_idx): return torch.where(q_idx >= k_idx, score, float("-inf")) output = flex_attention(q, k, v, score_mod=causal_score_mod)1.3 输入校验与限制
源码在进入 kernel 前执行了严格的输入校验(flex_attention.py):
- 不支持 NestedTensor:
_validate_no_nested_tensors会拒绝 jagged NestedTensor 输入,需先转为稠密张量; - 维度必须为 4D:
query、key、value 必须均为 4D 张量; - 头数一致:未开启 GQA 时
Hq必须等于Hkv,否则提示设置enable_gqa=True; - GQA 校验:开启后要求
Hq % Hkv == 0; - 批大小:
query.size(0)与key.size(0)不一致时,必须有非 None 的block_mask且其批维度与 query 一致; - 设备:支持 CUDA、CPU、XPU、HPU、MPS;CPU 与 MPS 仅支持推理(inference),不支持反向传播(flex_attention.py);
- BlockMask 形状匹配:
_validate_block_mask_shape保证 block_mask 的长度与 q/kv 长度一致,不一致时会给出“mask 太小/太大”的明确报错(flex_attention.py)。
二、辅助输出:AuxRequest与AuxOutput
从源码(flex_attention.py)可以看到,辅助输出使用一对 NamedTuple 描述“请求什么”与“返回什么”:
class AuxRequest(NamedTuple): lse: bool = False # 是否计算 logsumexp max_scores: bool = False # 是否计算每行的最大分数 class AuxOutput(NamedTuple): lse: Tensor | None = None # 形状 (B, Hq, L) max_scores: Tensor | None = None # 未请求时为 None使用示例:
from torch.nn.attention.flex_attention import flex_attention, AuxRequest out, aux = flex_attention(q, k, v, block_mask=block_mask, return_aux=AuxRequest(lse=True, max_scores=True)) # aux.lse 与 aux.max_scores 均可用需要特别注意两个兼容性事实:
return_lse已废弃:源码中的@deprecated装饰器明确提示 "return_lse is deprecated and will be removed in v2.10. Use return_aux=AuxRequest(lse=True) instead."(flex_attention.py)。同时指定return_lse与return_aux会抛出ValueError。BACKEND="FLASH"不支持max_scores:_apply_kernel_options中会抛出NotImplementedError;CPU 设备同样不支持返回 max scores(flex_attention.py)。
三、BlockMask:块稀疏掩码的数据结构
3.1 设计动机与格式
BlockMask是 Flex Attention 表示块稀疏注意力掩码的格式,介于 BCSR 与非稀疏格式之间。其核心思想是:只有当一个KV_BLOCK_SIZE × Q_BLOCK_SIZE块内的所有元素都稀疏时,该块才被视为稀疏——这与硬件期望连续加载与计算的行为高度一致(见 flex_attention.py)。
该格式针对“简洁”与“kernel 效率”优化,不针对体积优化:掩码总是被缩小KV_BLOCK_SIZE * Q_BLOCK_SIZE倍,若担心体积,可通过增大块大小来压缩。
3.2 核心字段
BlockMask的完整字段(flex_attention.py):
| 字段 | 类型 | 用途 |
|---|---|---|
seq_lengths | tuple[int, int] | (Q 长度, KV 长度) |
kv_num_blocks | Tensor | 每一行(Q 块行)中存在的 KV 块数量 |
kv_indices | Tensor | kv_indices[i]是第 i 行的块位置序列,kv_indices[i][kv_num_blocks[i]]之后的值为未定义 |
full_kv_num_blocks/full_kv_indices | Tensor,可选 | 记录“满块”(无需逐元素 mask 的块),可跳过对满块应用mask_mod,对因果掩码约有 15% 加速(源码注释所述) |
q_num_blocks/q_indices | Tensor,可选 | 反向传播所需(计算 dKV 需沿 Q 维迭代),由 KV 侧信息自动转置生成 |
full_q_num_blocks/full_q_indices | Tensor,可选 | 同上,面向满块 |
dq_write_order/dq_write_order_full | Tensor,可选 | 块稀疏 FLASH 反向的确定性 dQ 写入顺序元数据 |
dq_kv_order/dq_kv_order_spt | Tensor / bool,可选 | 生成写入顺序所用的 KV 调度顺序 |
BLOCK_SIZE | tuple[int, int] | 块大小(Q_BLOCK_SIZE, KV_BLOCK_SIZE),默认(128, 128) |
mask_mod | Callable | 生成该掩码所用的 mask_mod 函数 |
从该格式还原稠密掩码的等价逻辑(源码注释中的示例):
dense_mask = torch.zeros(ROWS, COLS) for row in range(ROWS): for block_idx in range(num_blocks_in_row[row]): dense_mask[row, col_indices[row, block_idx]] = 13.3 常用方法
from_kv_blocks(...)(类方法):从 KV 块信息直接构造 BlockMask,可传入BLOCK_SIZE、mask_mod、seq_lengths,并自动通过_transpose_ordered生成 Q 侧块信息(flex_attention.py)。compute_q_blocks=False可跳过 Q 块生成。as_tuple(flatten=True):返回 BlockMask 全部属性的元组(tensor 与上下文),flatten=True时展开BLOCK_SIZE与seq_lengths,这是传给底层 HOP 的序列化格式(flex_attention.py)。__getitem__(索引):最多支持三个索引,分别选择 batch、head 与Q 块行(注意:第三个索引选择的是块网格中的行,而非单个 query token)。整数索引会被归一化为长度为 1 的切片,从而保留维度。例如block_mask[:, :, 1]选择第 2 个 Q 块(覆盖Q_BLOCK_SIZE个 query token)。切片后的 BlockMask 是“打包”的,mask_mod会被替换为会报错的桩函数(_sliced_mask_mod_error),需要从原始 BlockMask 重新取回mask_mod(flex_attention.py)。_adjust(new_q_len, new_kv_len):将已有 BlockMask“裁剪”到左上角的新长度,适合复用掩码的场景,但源码注释明确说明并非对所有 mask_mod 都有效(flex_attention.py)。sparsity():返回稀疏块百分比,即“未计算块”的占比;numel()返回掩码元素总数(不剔除稀疏)。to_dense():还原为稠密块掩码张量。to_string(grid_size=(20, 20), limit=4):将掩码可视化为 ASCII 图(█满块、░部分块、空格为空块),grid_size=-1输出未压缩版本(可能非常大),limit控制最多打印多少个 batch/head(flex_attention.py)。to(device):将全部 tensor 属性迁移到目标设备,返回新对象,不修改原 BlockMask。shape属性:batch_dims + seq_lengths。
四、掩码工具函数:BlockMask Utilities
4.1create_block_mask:由 mask_mod 生成 BlockMask
这是最常用的掩码构造入口,签名与参数如下(flex_attention.py):
create_block_mask( mask_mod, # mask 函数:返回布尔张量,True 允许注意力连接 B, H, Q_LEN, KV_LEN, # 批大小、头数、Q 长度、KV 长度 device=None, # 默认取当前加速器或 "cpu" BLOCK_SIZE=128, # int 或 (Q_BLOCK_SIZE, KV_BLOCK_SIZE) _compile=False, # 已废弃,建议直接用 torch.compile(create_block_mask) separate_full_blocks=True, # True 时满块单独存储,kernel 可跳过 mask_mod compute_dq_write_order=False, # True 时预计算确定性 dQ 写入顺序元数据 dq_kv_order=True, # False=升序 n 块,True=降序/SPT 顺序 ) -> BlockMaskmask_mod的 4 参数签名:
def mask_mod(b, h, q_idx, kv_idx) -> Tensor: # 返回布尔张量 ...官方示例(因果掩码):
def causal_mask(b, h, q_idx, kv_idx): return q_idx >= kv_idx block_mask = create_block_mask(causal_mask, 1, 1, 8192, 8192, device="cuda") query = torch.randn(1, 1, 8192, 64, device="cuda", dtype=torch.float16) key = torch.randn(1, 1, 8192, 64, device="cuda", dtype=torch.float16) value = torch.randn(1, 1, 8192, 64, device="cuda", dtype=torch.float16) output = flex_attention(query, key, value, block_mask=block_mask)底层实现流程(源码可见):
_get_mod_type校验函数必须是 4 参数的mask_mod,否则报错;- 通过
create_mask在 (B, H, Q_LEN, KV_LEN) 网格上物化布尔掩码(内部用torch.vmap把mask_mod广播到 batch/head/序列维,见_vmap_for_bhqkv); _convert_mask_to_block_mask将稠密布尔掩码 padding 到块大小的整数倍,并按块聚合:separate_full_blocks=True时区分满块与部分块;_dense_to_ordered/_ordered_to_dense完成稠密与有序块列表之间的转换;- 最终调用
BlockMask.from_kv_blocks组装,必要时计算 dQ 写入顺序元数据。
注意create_block_mask的_compile参数已发出DeprecationWarning:建议直接torch.compile(create_block_mask)(...)。
4.2create_mask:直接生成稠密掩码张量
若你只需要一个稠密布尔掩码(而非块掩码),可调用create_mask(mod_fn, B, H, Q_LEN, KV_LEN, device),返回形状为(B, H, M, N)的掩码张量(flex_attention.py)。它同样自动识别 5 参数score_mod与 4 参数mask_mod:
- 对
score_mod:先在零分数上应用 mod 函数,再通过torch.where(torch.isneginf(out), False, True)将-inf分数转换为False; - 对
mask_mod:直接 vmap 求值返回布尔张量。
4.3and_masks/or_masks:掩码组合
两个函数接收任意多个mask_mod,返回一个新的mask_mod:
def and_masks(*mask_mods): # 交集:所有掩码同时为 True 才允许 def or_masks(*mask_mods): # 并集:任一掩码为 True 即允许实现(flex_attention.py)中会先校验所有参数可调用(callable),随后分别以&(初始new_ones)或|(初始new_zeros)逐项折叠。典型用法是组合因果与滑动窗口:
block_mask = create_block_mask( and_masks(causal_mask, sliding_window_mask), 1, 1, 4096, 4096, device="cuda" )4.4noop_mask:恒真掩码
返回一个值为True的标量布尔张量(flex_attention.py),作为默认的空操作mask_mod(BlockMask构造器与from_kv_blocks的默认值)。
五、FlexKernelOptions:kernel 级调优参数
FlexKernelOptions是一个TypedDict(total=False),控制底层 Triton kernel 的性能与数值行为(flex_attention.py)。绝大多数用户无需手动指定,默认 autotuning 已提供良好性能。所有选项均可加fwd_/bwd_前缀,使其只作用于前向或反向,例如fwd_BLOCK_M、bwd_BLOCK_M1。
5.1 常用选项速查表
| 选项 | 类型 | 说明 |
|---|---|---|
num_warps | int | CUDA kernel 使用的 warp 数,越高性能可能越好但寄存器压力增大,默认由 autotuning 决定 |
num_stages | int | kernel 流水线级数,越高性能可能越好但共享内存占用增大,默认 autotuning |
BLOCK_M | int | 前向 Q 序列维的线程块大小,必须为 2 的幂,常用 16/32/64/128 |
BLOCK_N | int | 前向 K/V 序列维的线程块大小,2 的幂 |
BLOCK_M1/BLOCK_N1/BLOCK_M2/BLOCK_N2 | int | 反向专用块大小,使用时须加bwd_前缀(如bwd_BLOCK_M1) |
PRESCALE_QK | bool | 是否将 QK 按1/sqrt(d)预缩放(含换底),稍快但可能引入更多数值误差,默认False |
ROWS_GUARANTEED_SAFE | bool | 跳过 softmax 行最大值的 sanitize 保护。要求比“每行至少注意一个 key”更强的条件:每行在其第一个被调度的 KV 块中就必须有未掩码的 key,否则运行中最大值仍是-inf,exp2会溢出为 NaN。因果注意力满足此条件,局部/窗口掩码往往不满足;eager 执行不会受影响,NaN 只在torch.compile下出现。默认False(源码有完整说明,flex_attention.py) |
BLOCKS_ARE_CONTIGUOUS | bool | 保证掩码中所有块连续,可优化块遍历。因果掩码满足,但 prefix_lm + sliding window 不满足,默认False |
WRITE_DQ | bool | 控制反向 DQ 迭代循环中是否做梯度 scatter;设为False会改在 DK 循环中完成,某些 score_mod/mask_mod 下可能更快,默认True |
FORCE_USE_FLEX_ATTENTION | bool | 强制使用 flex attention kernel 而非短序列下更优的 flex-decoding kernel,便于调试,默认False |
USE_TMA | bool | 是否在支持的硬件上使用 Tensor Memory Accelerator(TMA),实验性,目前针对 NVIDIA Hopper+ GPU,默认False |
kpack/matrix_instr_nonkdim/waves_per_eu | int | ROCm 专用参数(kernel packing、矩阵指令非 K 维、每 EU 波数) |
BACKEND | Literal | 选择具体 kernel 后端:"AUTO"(默认,按启发式在 flex_attention 与 flex_decoding 间选择)、"TRITON"(标准 Triton kernel)、"TRITON_DECODE"(短序列专用 flex_decoding)、"FLASH"(实验性,cute-dsl Flash kernel,需安装 flash)。不能与FORCE_USE_FLEX_ATTENTION等旧开关同时使用,否则抛错 |
5.2 使用示例
源码 docstring 给出了三种写法:
# 方式一:普通字典(向后兼容) kernel_opts = {"BLOCK_M": 64, "BLOCK_N": 64, "PRESCALE_QK": True} output = flex_attention(q, k, v, kernel_options=kernel_opts) # 方式二:TypedDict(类型安全,推荐) from torch.nn.attention.flex_attention import FlexKernelOptions kernel_opts: FlexKernelOptions = {"BLOCK_M": 64, "BLOCK_N": 64, "PRESCALE_QK": True} output = flex_attention(q, k, v, kernel_options=kernel_opts) # 方式三:前向/反向分别指定 kernel_opts: FlexKernelOptions = { "fwd_BLOCK_M": 64, "bwd_BLOCK_M1": 32, "PRESCALE_QK": False, } output = flex_attention(q, k, v, kernel_options=kernel_opts)5.3 后端选择与默认值
_apply_kernel_options(flex_attention.py)会做如下归一化:
- 校验
BACKEND必须是AUTO / TRITON / TRITON_DECODE / FLASH之一,非法值抛ValueError; - 与
FORCE_USE_FLEX_ATTENTION=True同时出现时抛RuntimeError; - 设置默认值:
BACKEND="AUTO"、PRESCALE_QK=False、ROWS_GUARANTEED_SAFE=False、BLOCKS_ARE_CONTIGUOUS=False、WRITE_DQ=True; - 内部强制写入
OUTPUT_LOGSUMEXP(前向是否输出 LSE 由torch.is_grad_enabled()决定;CPU/MPS 仅推理,恒为False)与OUTPUT_MAX(由return_aux.max_scores决定)。
六、GQA、缩放与内存布局等进阶用法
6.1 GQA(Grouped Query Attention)
设置enable_gqa=True后,key/value 头会广播到 query 头。源码要求Hq % Hkv == 0,否则抛出明确错误(flex_attention.py):
out = flex_attention(q, k, v, enable_gqa=True) # Hq 是 Hkv 的整数倍6.2scale参数
默认缩放为1/sqrt(E)(E 为嵌入维度),传入scale可覆盖,它会在 softmax 之前应用。
6.3 FP8 内存布局约束
对 FP8 dtype(torch.float8_e4m3fn、torch.float8_e5m2),且运行在 NVIDIA SM89(Ampere+FP8)至 SM100 之间、启用 CUDA 的硬件上,_enforce_mem_layouts会强制:
- query:行主序(左操作数
q @ k.T); - key:行主序(转置后为列主序右操作数);
- value:列主序(右操作数
softmax_scores @ v)。
SM100(Blackwell)支持 FP8 GEMM 的 TN/NT/TT/NN 全布局,因此该检查仅在更早架构上生效(flex_attention.py)。
6.4 未编译时的行为与调试开关
- 必须配合
torch.compile使用:直接调用flex_attention(未包裹在torch.compile中)时,源码会发出UserWarning,提示当前使用未融合的实现、会物化完整分数矩阵;推荐使用torch.compile(flex_attention)(...)(flex_attention.py)。 - 调试开关:设置
torch.nn.attention.flex_attention._FLEX_ATTENTION_DISABLE_COMPILE_DEBUG = True可临时关闭内部编译包装,从而在score_mod/mask_mod中设置断点或print。该开关仅影响直接调用 flex_attention 时的内部编译;若已用torch.compile包裹则无效,且不适用于反向传播(flex_attention.py)。 - 在 Dynamo 编译路径下,
flex_attention会把head_dim与num_heads标记为静态(torch._dynamo.mark_static),以利于 kernel 生成。
七、测试与验证
仓库为 Flex Attention 提供了大量测试(test/inductor/test_flex_attention.py),可作为理解行为与验证实现的参考:
TestFlexAttention(第 594 行起):核心前向/反向正确性测试,包含test_GQA、test_dependent_causal_bidirectional、test_padded_dense_causal、test_kv_batch_broadcast_causal_mask、test_kv_order_invariance_padded_causal等用例;TestBlockMask(第 7410 行起):BlockMask 的构造、索引与工具函数测试;TestLearnableBiases(第 9373 行起):test_relative_1d_bias、test_learnable_bias_global_compiled、test_comparison_vs_sdpa_with_learnable_bias等,展示如何用score_mod实现相对位置偏置并与标准 SDPA 对比。
测试中的精度校验逻辑也印证了实现要点:编译结果与参考实现比较时,float32 使用更宽松的容差(约 10 倍),其余 dtype 为 1.1 倍,源码注释解释这源于 online softmax 的数值特性(见_check_equal/_check_out)。
八、小结
torch.nn.attention.flex_attention通过四个核心抽象把注意力机制变得可编程且高效:
score_mod:以 5 参数函数任意改写注意力分数,可表达相对位置偏置、ALiBi、学习型偏置、软掩码等;mask_mod+create_block_mask:以 4 参数布尔函数描述任意稀疏模式,编译为块稀疏BlockMask后由 kernel 跳过无关块;BlockMask:BCSR 风格的有序块列表,配套from_kv_blocks、索引切片、sparsity()、to_dense()/to_string()等实用工具;FlexKernelOptions:以 TypedDict 透传 Triton kernel 级调优(块大小、warp/stage、数值开关、后端选择)。
使用时请牢记:这是 PyTorch 的原型特性(API 参考文档与源码 docstring 均标注 prototype),API 可能随版本演进;return_lse即将移除、create_block_mask(_compile=...)已废弃,新代码应使用return_aux=AuxRequest(...)与torch.compile(create_block_mask);生产使用务必用torch.compile(flex_attention)(...)获得融合 kernel,而非默认的未融合回退实现。
延伸阅读:模块其余实现位于 torch/nn/attention/flex_attention.py,配套的
_registry.py、bias.py、varlen.py、_fa3.py、_fa4.py提供了 bias 组合与 varlen(变长序列)等扩展能力,可在需要时进一步查阅。
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考