PyTorch 变长注意力详解:torch.nn.attention.varlen 的 Flash Attention 与 cuDNN 双后端实现
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 被调用的完整链路。一、模块定位与公开 API1.1 模块文档与实际源码关联文档 docs/source/nn.attention.varlen.md 是一份 Sphinxautomodule/autofunction/autoclass存根它把文档渲染指向模块的真实实现varlen_attn变长注意力主入口autofunctionvarlen_attn_out带预分配输出张量的变长注意力autofunctionAuxRequest请求计算辅助输出如 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(N1,)查询的累计序列位置cumulative sequence positionscu_seq_k(N1,)或None键/值的累计序列位置为None时部分路径复用cu_seq_qmax_q标量批内最大查询序列长度max_k标量批内最大键/值序列长度形状图例模块 docstring 原文语义N批大小T_q批内查询 token 总数所有查询序列长度之和T_k批内键/值 token 总数H_q查询注意力头数H_kv键/值注意力头数非 GQA 时等于H_qD注意力头维度。cu_seq的构造方式在模块 docstring 的示例中给出cu_seq[0] 0cu_seq[1:] seq_lengths.cumsum(0)即第i个样本占据[cu_seq[i], cu_seq[i1])区间的 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_auxNone, scaleNone, window_size(-1, -1), enable_gqaFalse, seqused_kNone, block_tableNone, num_splitsNone, ) - Tensor | tuple[Tensor, Tensor]3.2 参数逐项说明return_auxAuxRequest | None请求辅助输出。AuxRequest是一个NamedTuple目前只包含一个布尔字段lse是否计算 log-sum-exp。当return_aux is not None and return_aux.lse为真时函数额外返回形状为(H_q, T_q)的 logsumexp 张量否则只返回输出张量。注意 lse 在反向传播中被标记为不可微见下文_setup_context。scalefloat | 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_gqabool默认False启用 Grouped Query Attention允许H_kv H_q。每个 KV 头被一组H_q / H_kv个查询头共享因此要求H_q能被H_kv整除不满足时抛出ValueErrorExpect number of query heads to be a multiple of kv heads for GQA...。若未启用 GQA 但头数不等同样抛错并提示Try setting enable_gqaTrue。seqused_kTensor, (N,)可选每个批元素的有效 KV token 数。设置后第i个样本只有前seqused_k[i]个 KV token 参与注意力。典型用途是 KV Cache 解码缓存槽比实际序列长。仅限推理——_setup_context中明确raise RuntimeError(seqused_k is an inference-only parameter.)不允许反向传播。block_tableTensor, (N, max_pages_per_seq)int32可选分页 KV Cache 的块表。此时key/value是「页池」物理页各序列任意交错block_table把每个序列的逻辑块映射回池中的物理页seqused_k[i]告诉 kernel 序列i实际有效的 token 数最后一页通常只填充一部分。必须与seqused_k同时提供且同样仅限推理。num_splitsint可选split-KV 的切分数。num_splits1表示禁用 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, devicecuda) total_tokens seq_lengths.sum().item() # 打包的 query / key / value query torch.randn(total_tokens, num_heads, head_dim, dtypetorch.float16, devicecuda) key torch.randn(total_tokens, num_heads, head_dim, dtypetorch.float16, devicecuda) value torch.randn(total_tokens, num_heads, head_dim, dtypetorch.float16, devicecuda) # 构造累计序列张量 cu_seq torch.zeros(batch_size 1, devicecuda, dtypetorch.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源码注释提醒该计算相对实际开销是高估的因为它把每个样本近似成(batch1, heads, seq_len, dim)的稠密形状。五、后端选择机制cuDNN 优先Flash Attention 兜底varlen_attn是「双后端」实现优先选择 cuDNN AttentionSDPBackend.CUDNN_ATTENTION否则回退 Flash AttentionSDPBackend.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_priorityTrue)可覆盖见 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 或 SM100ROCm 上恒不使用 cuDNNmax_q必须 128dtypequery必须是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 Cachewindow_size仅接受(-1, -1)或(-1, 0)普通滑动窗口走 FlashGQA /num_splits均不支持KV Cacheblock_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_priorityTrue时优先级可被覆盖、以及无可用后端时的报错行为。这为「在torch.nn.attention.sdpa_kernel(SDPBackend.FLASH_ATTENTION)上下文内强制 Flash」的用法提供了依据。六、split-KV 与批不变性batch invariancenum_splits是值得单独强调的精度相关参数。模块 docstring 的说明split-KV 把键/值序列维度切分到多个线程块并行计算再合并部分结果切分决策依赖max_k批内最长序列因此同一序列在不同批次组成下归约顺序可能不同浮点结果可能产生微小差异设置num_splits1禁用 split-KV 后给定序列无论与什么其他序列同批都能获得逐位一致的输出代价是查询数较少时 GPU 利用率下降None默认由 kernel 自动决策。测试test_batch_invariancetest/test_varlen_attention.py用固定种子构造「单独推理」与「拼接成批推理」两组cu_seq对比同一目标序列的输出是否逐位一致并覆盖num_splits与window_size的组合注释特别指出fa4 and cuDNN are batch invariant by default即 cuDNN 后端默认具备批不变性而 Flash 后端需要num_splits1才能保证。七、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_cachetest/test_varlen_attention.py验证了该路径。7.2 block_table分页 KV CachePagedAttentionkey/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_backwardFlash 走torch.ops.aten._flash_attention_backward最后返回(dq, dk, dv)以及 12 个None对应其余非张量参数三个自定义算子均注册了register_fakemeta 实现与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_splitscuDNN 路径调用torch.ops.aten._cudnn_attention_forward从返回元组中取output、softmax_lse与第 6 项rng_state两条路径的dropout_p均被硬编码为0.0返回的rng_state也是硬编码的全零(2,)uint64张量——即当前实现不支持 dropout这是使用前需要明确的限制。九、实践要点与限制汇总主题结论输入构造扁平打包 cu_seqint32、cumsum生成dtype 建议float16/bfloat16后端启用在torch.nn.attention.sdpa_kernel(...)上下文内运行varlen_attn_out仅支持 FlashcuDNN 前提cuDNN ≥ 9.18、SM90/SM100、max_q 128、头维度 8 对齐、不支持 GQA/num_splits/普通窗口dropout当前版本硬编码为 0不支持随机丢弃训练支持反向seqused_k/block_table仅限推理批不变性需要逐位一致时用num_splits1cuDNN 默认满足辅助输出需要 logsumexp 时传AuxRequest(lseTrue)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),仅供参考