PyTorch Flex Attention 完全指南:从 score_mod 到 BlockMask 的官方 API 深度解析
2026/9/10 1:51:39 网站建设 项目流程

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及其配套的BlockMaskFlexKernelOptions、掩码工具函数与辅助输出接口,并结合仓库源码(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 整理):

参数类型说明
queryTensorQuery 张量,形状(B, Hq, L, E)。FP8 数据类型下建议行主序(row-major)内存布局以获得最佳性能
keyTensorKey 张量,形状(B, Hkv, S, E)。FP8 下建议行主序
valueTensorValue 张量,形状(B, Hkv, S, Ev)。FP8 下建议列主序(column-major)
score_modCallable,可选修改注意力分数的函数,默认不应用任何修改(_identity
block_maskBlockMask,可选控制注意力块稀疏模式的 BlockMask 对象,默认生成一个覆盖全长度的空块掩码
scalefloat,可选softmax 前应用的缩放因子,默认为1/sqrt(E)
enable_gqabool设为 True 时启用 Grouped Query Attention(GQA),将 key/value 头广播到 query 头
return_lsebool是否返回注意力的 logsumexp,已废弃,请改用return_aux=AuxRequest(lse=True)
kernel_optionsFlexKernelOptions,可选控制底层 Triton kernel 行为的选项
return_auxAuxRequest,可选指定要计算并返回的辅助输出

返回值:注意力输出张量,形状(B, Hq, L, Ev);当指定return_aux时返回(output, aux),其中auxAuxOutput

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 输入,需先转为稠密张量;
  • 维度必须为 4Dquery、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)。

二、辅助输出:AuxRequestAuxOutput

从源码(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 均可用

需要特别注意两个兼容性事实:

  1. 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_lsereturn_aux会抛出ValueError
  2. 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_lengthstuple[int, int](Q 长度, KV 长度)
kv_num_blocksTensor每一行(Q 块行)中存在的 KV 块数量
kv_indicesTensorkv_indices[i]是第 i 行的块位置序列,kv_indices[i][kv_num_blocks[i]]之后的值为未定义
full_kv_num_blocks/full_kv_indicesTensor,可选记录“满块”(无需逐元素 mask 的块),可跳过对满块应用mask_mod,对因果掩码约有 15% 加速(源码注释所述)
q_num_blocks/q_indicesTensor,可选反向传播所需(计算 dKV 需沿 Q 维迭代),由 KV 侧信息自动转置生成
full_q_num_blocks/full_q_indicesTensor,可选同上,面向满块
dq_write_order/dq_write_order_fullTensor,可选块稀疏 FLASH 反向的确定性 dQ 写入顺序元数据
dq_kv_order/dq_kv_order_sptTensor / bool,可选生成写入顺序所用的 KV 调度顺序
BLOCK_SIZEtuple[int, int]块大小(Q_BLOCK_SIZE, KV_BLOCK_SIZE),默认(128, 128)
mask_modCallable生成该掩码所用的 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]] = 1

3.3 常用方法

  • from_kv_blocks(...)(类方法):从 KV 块信息直接构造 BlockMask,可传入BLOCK_SIZEmask_modseq_lengths,并自动通过_transpose_ordered生成 Q 侧块信息(flex_attention.py)。compute_q_blocks=False可跳过 Q 块生成。
  • as_tuple(flatten=True):返回 BlockMask 全部属性的元组(tensor 与上下文),flatten=True时展开BLOCK_SIZEseq_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 顺序 ) -> BlockMask

mask_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)

底层实现流程(源码可见):

  1. _get_mod_type校验函数必须是 4 参数的mask_mod,否则报错;
  2. 通过create_mask在 (B, H, Q_LEN, KV_LEN) 网格上物化布尔掩码(内部用torch.vmapmask_mod广播到 batch/head/序列维,见_vmap_for_bhqkv);
  3. _convert_mask_to_block_mask将稠密布尔掩码 padding 到块大小的整数倍,并按块聚合:separate_full_blocks=True时区分满块与部分块;
  4. _dense_to_ordered/_ordered_to_dense完成稠密与有序块列表之间的转换;
  5. 最终调用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_modBlockMask构造器与from_kv_blocks的默认值)。


五、FlexKernelOptions:kernel 级调优参数

FlexKernelOptions是一个TypedDicttotal=False),控制底层 Triton kernel 的性能与数值行为(flex_attention.py)。绝大多数用户无需手动指定,默认 autotuning 已提供良好性能。所有选项均可加fwd_/bwd_前缀,使其只作用于前向或反向,例如fwd_BLOCK_Mbwd_BLOCK_M1

5.1 常用选项速查表

选项类型说明
num_warpsintCUDA kernel 使用的 warp 数,越高性能可能越好但寄存器压力增大,默认由 autotuning 决定
num_stagesintkernel 流水线级数,越高性能可能越好但共享内存占用增大,默认 autotuning
BLOCK_Mint前向 Q 序列维的线程块大小,必须为 2 的幂,常用 16/32/64/128
BLOCK_Nint前向 K/V 序列维的线程块大小,2 的幂
BLOCK_M1/BLOCK_N1/BLOCK_M2/BLOCK_N2int反向专用块大小,使用时须加bwd_前缀(如bwd_BLOCK_M1
PRESCALE_QKbool是否将 QK 按1/sqrt(d)预缩放(含换底),稍快但可能引入更多数值误差,默认False
ROWS_GUARANTEED_SAFEbool跳过 softmax 行最大值的 sanitize 保护。要求比“每行至少注意一个 key”更强的条件:每行在其第一个被调度的 KV 块中就必须有未掩码的 key,否则运行中最大值仍是-infexp2会溢出为 NaN。因果注意力满足此条件,局部/窗口掩码往往不满足;eager 执行不会受影响,NaN 只在torch.compile下出现。默认False(源码有完整说明,flex_attention.py)
BLOCKS_ARE_CONTIGUOUSbool保证掩码中所有块连续,可优化块遍历。因果掩码满足,但 prefix_lm + sliding window 不满足,默认False
WRITE_DQbool控制反向 DQ 迭代循环中是否做梯度 scatter;设为False会改在 DK 循环中完成,某些 score_mod/mask_mod 下可能更快,默认True
FORCE_USE_FLEX_ATTENTIONbool强制使用 flex attention kernel 而非短序列下更优的 flex-decoding kernel,便于调试,默认False
USE_TMAbool是否在支持的硬件上使用 Tensor Memory Accelerator(TMA),实验性,目前针对 NVIDIA Hopper+ GPU,默认False
kpack/matrix_instr_nonkdim/waves_per_euintROCm 专用参数(kernel packing、矩阵指令非 K 维、每 EU 波数)
BACKENDLiteral选择具体 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=FalseROWS_GUARANTEED_SAFE=FalseBLOCKS_ARE_CONTIGUOUS=FalseWRITE_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_e4m3fntorch.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_dimnum_heads标记为静态(torch._dynamo.mark_static),以利于 kernel 生成。

七、测试与验证

仓库为 Flex Attention 提供了大量测试(test/inductor/test_flex_attention.py),可作为理解行为与验证实现的参考:

  • TestFlexAttention(第 594 行起):核心前向/反向正确性测试,包含test_GQAtest_dependent_causal_bidirectionaltest_padded_dense_causaltest_kv_batch_broadcast_causal_masktest_kv_order_invariance_padded_causal等用例;
  • TestBlockMask(第 7410 行起):BlockMask 的构造、索引与工具函数测试;
  • TestLearnableBiases(第 9373 行起):test_relative_1d_biastest_learnable_bias_global_compiledtest_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通过四个核心抽象把注意力机制变得可编程且高效:

  1. score_mod:以 5 参数函数任意改写注意力分数,可表达相对位置偏置、ALiBi、学习型偏置、软掩码等;
  2. mask_mod+create_block_mask:以 4 参数布尔函数描述任意稀疏模式,编译为块稀疏BlockMask后由 kernel 跳过无关块;
  3. BlockMask:BCSR 风格的有序块列表,配套from_kv_blocks、索引切片、sparsity()to_dense()/to_string()等实用工具;
  4. 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.pybias.pyvarlen.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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询