news 2026/9/11 7:37:30

PyTorch 变长注意力详解:torch.nn.attention.varlen 的 Flash Attention 与 cuDNN 双后端实现

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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_attnvarlen_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)键张量
valuekey值张量
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] = 0cu_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_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整除;不满足时抛出ValueError("Expect number of query heads to be a multiple of kv heads for GQA...")。若未启用 GQA 但头数不等,同样抛错并提示Try setting enable_gqa=True
  • 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_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_counterflop_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

  1. 读取全局开关torch._C._get_cudnn_sdp_enabled()_get_flash_sdp_enabled()
  2. 若 cuDNN 开启,用_cudnn_rejection_reasons收集不满足的约束,全部通过才视为cudnn_eligible
  3. _get_sdp_priority_order()得到的优先级遍历(默认[CUDNN, FLASH]sdpa_kernel(..., set_priority=True)可覆盖,见 torch/nn/attention/init.py);
  4. 选中最先满足条件的后端;若 cuDNN 开启但约束不满足,抛出带约束明细的RuntimeError;若两个后端都未启用,提示用sdpa_kernel()启用其中之一。

_get_sdp_priority_ordertorch.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
dtypequery必须是float16bfloat16
头维度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 Cacheblock_table必须搭配seqused_k

这些约束在测试中都有对应覆盖,例如 test/test_varlen_attention.py 的test_cudnn_varlen_requires_shared_cu_seqtest_cudnn_varlen_unaligned_input_raisestest_cudnn_varlen_large_head_dimstest_cudnn_kv_cache_validation等。

5.3 与 sdpa_kernel 的联动测试

test_sdpa_kernel_backend_selectiontest_sdpa_kernel_backend_prioritytest_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_splitswindow_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_statecu_seqmax_q/max_kis_causal/scale/window_sizebackend),并把lserng_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/rightseqused_kblock_tablenum_splits
  • cuDNN 路径调用torch.ops.aten._cudnn_attention_forward,从返回元组中取outputsoftmax_lse与第 6 项rng_state
  • 两条路径的dropout_p均被硬编码为0.0,返回的rng_state也是硬编码的全零(2,)uint64张量——即当前实现不支持 dropout,这是使用前需要明确的限制。

九、实践要点与限制汇总

主题结论
输入构造扁平打包 +cu_seqint32cumsum生成),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_invariancetest_seqused_k_kv_cachetest_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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/11 7:36:20

移动优先索引时代,SEO网络公司如何系统做好移动端优化

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/11 7:30:32

LangChain高并发智能客服的流控、排队与降级协同治理

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/11 7:28:42

两栖动物活性肽Phyllomedusin与pENPNRFIGLM-NH₂的生物医学应用

1. Phyllomedusin与pENPNRFIGLM-NH₂&#xff1a;两栖动物活性肽的生物医学探秘在雨林深处的树蛙皮肤上&#xff0c;藏着自然界最精密的生物化学武器库。Phyllomedusin和pENPNRFIGLM-NH₂这两个看似晦涩的命名&#xff0c;实则是两栖动物防御系统中经过千万年进化锤炼的活性肽代…

作者头像 李华
网站建设 2026/9/11 7:28:38

UC3843AC开关电源设计核心:RT/CT振荡器、电流采样与反馈环路实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华