PyTorch 变长注意力详解:torch.nn.attention.varlen 的 Flash Attention 与 cuDNN 双后端实现
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
导读
torch.nn.attention.varlen是 PyTorch 为「一批长度各不相同的序列」提供的高性能注意力接口。它不依赖 padding 对齐,而是把整批 token 打包成扁平张量、用累计序列长度张量(cu_seq)描述每个样本的边界,从而直接调用 Flash Attention 与 cuDNN 的融合 kernel。本文基于当前仓库源码,完整讲解varlen_attn、varlen_attn_out两个公开 API 与AuxRequest辅助输出请求,覆盖张量布局、参数语义、后端选择机制、KV Cache 解码(seqused_k/block_table)、滑动窗口与因果掩码、GQA、split-KV 与批不变性等实战要点。读完后你将能正确构造变长输入、按场景选择参数,并理解底层 kernel 被调用的完整链路。
一、模块定位与公开 API
1.1 模块文档与实际源码
关联文档 docs/source/nn.attention.varlen.md 是一份 Sphinxautomodule/autofunction/autoclass存根,它把文档渲染指向模块的真实实现:
varlen_attn:变长注意力主入口(autofunction);varlen_attn_out:带预分配输出张量的变长注意力(autofunction);AuxRequest:请求计算辅助输出(如 logsumexp)的配置类(autoclass)。
这三个符号的实际定义位于 torch/nn/attention/varlen.py,并通过__all__ = ["varlen_attn", "varlen_attn_out", "AuxRequest"]对外导出。模块 docstring 明确其定位:"Variable-length attention implementation using Flash Attention"——一个调用优化后 Flash Attention kernel 的高层 Python 接口。
1.2 与 scaled_dot_product_attention 的关系
varlen_attn的 docstring 指出它与scaled_dot_product_attention类似,但专门针对变长序列优化:不使用(batch, heads, seq_len, head_dim)的规则形状,而是采用「扁平打包 token + 累计序列位置」的描述方式。这意味着它天然适配 NestedTensor / jagged 数据、PagedAttention 等推理场景。
二、输入布局与形状约定
varlen_attn的核心输入是一组扁平张量与两个累计序列张量,形状约定(来自 varlen.py 的 docstring 与_varlen_attn_fake实现)如下:
| 参数 | 形状 | 说明 |
|---|---|---|
query | (T_q, H_q, D) | 全批查询 token 打包,T_q为各样本查询长度之和 |
key | (T_k, H_kv, D),或提供block_table时为(total_pages, page_size, H_kv, D) | 键张量 |
value | 同key | 值张量 |
cu_seq_q | (N+1,) | 查询的累计序列位置(cumulative sequence positions) |
cu_seq_k | (N+1,)或None | 键/值的累计序列位置;为None时(部分路径)复用cu_seq_q |
max_q | 标量 | 批内最大查询序列长度 |
max_k | 标量 | 批内最大键/值序列长度 |
形状图例(模块 docstring 原文语义):
N:批大小;T_q:批内查询 token 总数(所有查询序列长度之和);T_k:批内键/值 token 总数;H_q:查询注意力头数;H_kv:键/值注意力头数(非 GQA 时等于H_q);D:注意力头维度。
cu_seq的构造方式在模块 docstring 的示例中给出:cu_seq[0] = 0,cu_seq[1:] = seq_lengths.cumsum(0),即第i个样本占据[cu_seq[i], cu_seq[i+1])区间的 token。cu_seq通常为int32、位于 CUDA 上。
三、varlen_attn:完整参数语义
3.1 函数签名
varlen_attn( query, key, value, cu_seq_q, cu_seq_k, max_q, max_k, *, return_aux=None, scale=None, window_size=(-1, -1), enable_gqa=False, seqused_k=None, block_table=None, num_splits=None, ) -> Tensor | tuple[Tensor, Tensor]3.2 参数逐项说明
return_aux(AuxRequest | None):请求辅助输出。AuxRequest是一个NamedTuple,目前只包含一个布尔字段lse(是否计算 log-sum-exp)。当return_aux is not None and return_aux.lse为真时,函数额外返回形状为(H_q, T_q)的 logsumexp 张量,否则只返回输出张量。注意 lse 在反向传播中被标记为不可微(见下文_setup_context)。scale(float | None):注意力分数的正缩放因子。_validate_scale要求scale > 0,且该校验形式会拒绝NaN(源码注释特别说明:"This form also rejects NaN, unlike scale <= 0");传入非法值会抛出ValueError: scale must be greater than 0。测试 test/test_varlen_attention.py 中test_varlen_invalid_scale覆盖了该路径。window_size((left, right)):滑动窗口注意力窗口大小:(-1, -1):全注意力(默认);(-1, 0):因果注意力;(W, 0):窗口大小为W的因果滑动窗口注意力。 内部通过_normalize_window_size校验长度必须为 2,并把None归一化为[-1, -1];is_causal = (window_size == (-1, 0))由该参数推导,而非单独传布尔值。
enable_gqa(bool,默认False):启用 Grouped Query Attention,允许H_kv < H_q。每个 KV 头被一组H_q / H_kv个查询头共享,因此要求H_q能被H_kv整除;不满足时抛出ValueError("Expect number of query heads to be a multiple of kv heads for GQA...")。若未启用 GQA 但头数不等,同样抛错并提示Try setting enable_gqa=True。seqused_k(Tensor, (N,),可选):每个批元素的有效 KV token 数。设置后,第i个样本只有前seqused_k[i]个 KV token 参与注意力。典型用途是 KV Cache 解码:缓存槽比实际序列长。仅限推理——_setup_context中明确raise RuntimeError("seqused_k is an inference-only parameter."),不允许反向传播。block_table(Tensor, (N, max_pages_per_seq),int32,可选):分页 KV Cache 的块表。此时key/value是「页池」(物理页,各序列任意交错),block_table把每个序列的逻辑块映射回池中的物理页;seqused_k[i]告诉 kernel 序列i实际有效的 token 数(最后一页通常只填充一部分)。必须与seqused_k同时提供,且同样仅限推理。num_splits(int,可选):split-KV 的切分数。num_splits=1表示禁用 split-KV 以获得批不变性(batch invariance)。默认None由 kernel 自动决策。详见第六节。
3.3 返回值与辅助输出
- 默认只返回
output,形状(T_q, H_q, D); - 当
return_aux.lse为真时返回(output, lse),其中lse形状为(H_q, T_q)。
3.4 最小可用示例(来自模块 docstring,已标注需 CUDA 环境)
>>> batch_size, max_seq_len, embed_dim, num_heads = 2, 512, 1024, 16 >>> head_dim = embed_dim // num_heads >>> seq_lengths = [] >>> for _ in range(batch_size): ... length = torch.randint(1, max_seq_len // 64 + 1, (1,)).item() * 64 ... seq_lengths.append(min(length, max_seq_len)) >>> seq_lengths = torch.tensor(seq_lengths, device="cuda") >>> total_tokens = seq_lengths.sum().item() >>> >>> # 打包的 query / key / value >>> query = torch.randn(total_tokens, num_heads, head_dim, dtype=torch.float16, device="cuda") >>> key = torch.randn(total_tokens, num_heads, head_dim, dtype=torch.float16, device="cuda") >>> value = torch.randn(total_tokens, num_heads, head_dim, dtype=torch.float16, device="cuda") >>> >>> # 构造累计序列张量 >>> cu_seq = torch.zeros(batch_size + 1, device="cuda", dtype=torch.int32) >>> cu_seq[1:] = seq_lengths.cumsum(0) >>> max_len = seq_lengths.max().item() >>> >>> output = varlen_attn(query, key, value, cu_seq, cu_seq, max_len, max_len)该示例中 query/key/value 均为float16且位于 CUDA,与两个后端的最低 dtype 要求一致(见第五节)。
四、varlen_attn_out:预分配输出的变体
varlen_attn_out(out, query, key, value, cu_seq_q, cu_seq_k, max_q, max_k, *, ...)与varlen_attn行为相同,但把注意力输出写入调用方提供的out张量,避免内部重新分配(适合需要精确控制内存或复用缓冲区的场景)。
其实现要点:
- 内部调用自定义算子
torch_attn::_varlen_attn_out,该算子通过mutates_args={"out"}声明原地改写out; - 仅支持 Flash Attention 后端:进入函数后先检查
torch._C._get_flash_sdp_enabled(),未启用时抛RuntimeError("varlen_attn_out only supports SDPBackend.FLASH_ATTENTION; enable it with sdpa_kernel()."); - 底层调用
torch.ops.aten._flash_attention_forward_no_dropout_inplace,该算子被torch._dynamo.disallow_in_graph排除在 Dynamo 图外; return_aux.lse为真时同样返回(out, lse),其中lse由 in-place 前向返回,形状(H_q, T_q)。
此外,模块通过torch.utils.flop_counter的flop_registry为三个自定义算子注册了 FLOP 计数(见 torch/utils/flop_counter.py)。其中_varlen_attn_forward_flop利用_unpack_flash_attention_nested_shapes依据cu_seq将每个批元素的序列长度还原,再逐样本求和 FLOP(源码注释提醒该计算相对实际开销是高估的,因为它把每个样本近似成(batch=1, heads, seq_len, dim)的稠密形状)。
五、后端选择机制:cuDNN 优先,Flash Attention 兜底
varlen_attn是「双后端」实现:优先选择 cuDNN Attention(SDPBackend.CUDNN_ATTENTION),否则回退 Flash Attention(SDPBackend.FLASH_ATTENTION)。后端开关与优先级完全跟随torch.nn.attention.sdpa_kernel上下文管理器(模块 docstring 明确:"Backend enablement follows sdpa_kernel()")。
5.1 选择流程(_select_backend)
- 读取全局开关
torch._C._get_cudnn_sdp_enabled()与_get_flash_sdp_enabled(); - 若 cuDNN 开启,用
_cudnn_rejection_reasons收集不满足的约束,全部通过才视为cudnn_eligible; - 按
_get_sdp_priority_order()得到的优先级遍历(默认[CUDNN, FLASH];sdpa_kernel(..., set_priority=True)可覆盖,见 torch/nn/attention/init.py); - 选中最先满足条件的后端;若 cuDNN 开启但约束不满足,抛出带约束明细的
RuntimeError;若两个后端都未启用,提示用sdpa_kernel()启用其中之一。
_get_sdp_priority_order被torch.compiler.assume_constant_result装饰,会在 trace 时把后端优先级固化为常量,保证编译期决策一致性。
5.2 cuDNN 后端的约束清单
从 varlen.py 的_cudnn_rejection_reasons可以整理出 cuDNN varlen 的完整限制:
| 约束 | 说明 |
|---|---|
| 设备 | query必须在 CUDA 上 |
| 软件/硬件 | cuDNN ≥ 9.18,且设备算力主版本为 SM90 或 SM100(ROCm 上恒不使用 cuDNN) |
max_q | 必须> 128 |
| dtype | query必须是float16或bfloat16 |
| 头维度 | query.shape[-1]与value.shape[-1]必须能被 8 整除 |
| 特殊大头维度 | 头维度 ≤ 128 时通用;否则仅 SM100 + 特定维度组合(前向{(192,128),(192,192),(256,128),(256,256)}需 cuDNN ≥ 9.24,反向{(192,128)}需 cuDNN ≥ 9.19) |
| 因果 | (-1, 0)要求cu_seq_q is cu_seq_k(同一张量),且不允许 KV Cache |
window_size | 仅接受(-1, -1)或(-1, 0),普通滑动窗口走 Flash |
GQA /num_splits | 均不支持 |
| KV Cache | block_table必须搭配seqused_k |
这些约束在测试中都有对应覆盖,例如 test/test_varlen_attention.py 的test_cudnn_varlen_requires_shared_cu_seq、test_cudnn_varlen_unaligned_input_raises、test_cudnn_varlen_large_head_dims、test_cudnn_kv_cache_validation等。
5.3 与 sdpa_kernel 的联动测试
test_sdpa_kernel_backend_selection、test_sdpa_kernel_backend_priority、test_sdpa_kernel_backend_errors三个测试(test/test_varlen_attention.py)分别验证了:默认优先级下 cuDNN 优先、set_priority=True时优先级可被覆盖、以及无可用后端时的报错行为。这为「在torch.nn.attention.sdpa_kernel(SDPBackend.FLASH_ATTENTION)上下文内强制 Flash」的用法提供了依据。
六、split-KV 与批不变性(batch invariance)
num_splits是值得单独强调的精度相关参数。模块 docstring 的说明:
- split-KV 把键/值序列维度切分到多个线程块并行计算,再合并部分结果;
- 切分决策依赖
max_k(批内最长序列),因此同一序列在不同批次组成下,归约顺序可能不同,浮点结果可能产生微小差异; - 设置
num_splits=1禁用 split-KV 后,给定序列无论与什么其他序列同批,都能获得逐位一致的输出,代价是查询数较少时 GPU 利用率下降; None(默认)由 kernel 自动决策。
测试test_batch_invariance(test/test_varlen_attention.py)用固定种子构造「单独推理」与「拼接成批推理」两组cu_seq,对比同一目标序列的输出是否逐位一致,并覆盖num_splits与window_size的组合;注释特别指出"fa4 and cuDNN are batch invariant by default",即 cuDNN 后端默认具备批不变性,而 Flash 后端需要num_splits=1才能保证。
七、KV Cache 推理:seqused_k 与 block_table
变长注意力的重要落地场景是 KV Cache 解码,两个专用参数在模块 docstring 中有详细说明:
7.1 seqused_k:连续 KV Cache
当 KV Cache 槽位大于实际序列长度时,seqused_k[i]指定样本i真正有效的 token 数,kernel 只让前seqused_k[i]个 KV token 参与注意力。仅需把填充部分排除在注意力外,无需重新压缩张量。测试test_seqused_k_kv_cache(test/test_varlen_attention.py)验证了该路径。
7.2 block_table:分页 KV Cache(PagedAttention)
key/value退化为「页池」:形状(total_pages, page_size, H_kv, D),页与页之间、页与序列之间无固定顺序;block_table((N, max_pages_per_seq),int32)把每个序列的逻辑页映射回物理页;- 最后一页通常部分填充,因此必须同时提供
seqused_k说明各序列有效 token 数。
源码层面,block_table的存在会改变key/value的解析方式:num_heads_k = key.size(2) if block_table is not None else key.size(1)。测试test_block_table_kv_cache与 cuDNN 侧的test_cudnn_kv_cache(覆盖 paged、page_size、strided_table 等组合)共同验证了该功能。两个参数均被_setup_context标记为 inference-only:一旦进入反向传播即抛RuntimeError。
八、反向传播与自定义算子架构
varlen_attn的反向通过自定义算子torch_attn::_varlen_attn_backward实现,注册了完整的setup_context与 autograd 回调:
_setup_context保存前向中间量(query/key/value/out/lse/rng_state及cu_seq、max_q/max_k、is_causal/scale/window_size、backend),并把lse、rng_state标记为mark_non_differentiable(对应测试test_varlen_lse_is_not_differentiable);_backward依据前向选择的ctx.backend分发:cuDNN 走torch.ops.aten._cudnn_attention_backward,Flash 走torch.ops.aten._flash_attention_backward,最后返回(dq, dk, dv)以及 12 个None(对应其余非张量参数);- 三个自定义算子均注册了
register_fake(meta 实现)与register_autograd,保证在 FakeTensor/编译/元数据推理下形状正确。
前向内部实现(_varlen_attn)的关键细节:
- 自定义算子通过
@torch.library.custom_op("torch_attn::_varlen_attn", mutates_args={})注册;由于自定义算子 schema 不支持枚举参数,后端以int传递(SDPBackend.*.value,源码注释明确说明); - Flash 路径调用
torch.ops.aten._flash_attention_forward,传入window_size_left/right、seqused_k、block_table、num_splits; - cuDNN 路径调用
torch.ops.aten._cudnn_attention_forward,从返回元组中取output、softmax_lse与第 6 项rng_state; - 两条路径的
dropout_p均被硬编码为0.0,返回的rng_state也是硬编码的全零(2,)uint64张量——即当前实现不支持 dropout,这是使用前需要明确的限制。
九、实践要点与限制汇总
| 主题 | 结论 |
|---|---|
| 输入构造 | 扁平打包 +cu_seq(int32、cumsum生成),dtype 建议float16/bfloat16 |
| 后端启用 | 在torch.nn.attention.sdpa_kernel(...)上下文内运行;varlen_attn_out仅支持 Flash |
| cuDNN 前提 | cuDNN ≥ 9.18、SM90/SM100、max_q > 128、头维度 8 对齐、不支持 GQA/num_splits/普通窗口 |
| dropout | 当前版本硬编码为 0,不支持随机丢弃 |
| 训练 | 支持反向;seqused_k/block_table仅限推理 |
| 批不变性 | 需要逐位一致时用num_splits=1;cuDNN 默认满足 |
| 辅助输出 | 需要 logsumexp 时传AuxRequest(lse=True),lse 不可微,形状(H_q, T_q) |
| 调试建议 | 可参考 test/test_varlen_attention.py 中test_varlen_vs_sdpa(与scaled_dot_product_attention对拍)、test_batch_invariance、test_seqused_k_kv_cache、test_block_table_kv_cache构造验证用例 |
如需深入了解实现在 torch/nn/attention/varlen.py 中按上述流程逐段阅读;FLOP 计数细节见 torch/utils/flop_counter.py;后端启用与优先级 API 见 torch/nn/attention/init.py。集成到现有模块时,可结合 NestedTensor 相关基础设施(如 torch/nn/attention/flex_attention.py 中同类封装)理解 PyTorch 注意力族 API 的整体设计。
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考