PyTorch torch.nn.attention.bias 深度解析:CausalBias 因果注意力偏置的原理与实战
PyTorch torch.nn.attention.bias 深度解析CausalBias 因果注意力偏置的原理与实战【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch本篇指南围绕 PyTorch 的torch.nn.attention.bias模块展开讲解CausalBias类、causal_upper_left/causal_lower_right工厂函数与CausalVariant枚举的设计动机与内部实现。读完本文后你将理解非方阵场景下两种因果掩码upper-left 与 lower-right的几何语义差异、CausalBias如何通过__torch_function__钩子拦截F.scaled_dot_product_attention并分派到 Flash / Efficient 融合内核以及在实际代码中如何安全地构造和使用这类注意力偏置。为什么需要独立的因果偏置模块F.scaled_dot_product_attention的is_causalTrue参数在查询与键值序列长度相等方阵时语义是明确的按主对角线做下三角掩码。但在seq_len_q ≠ seq_len_kv 的非方阵场景例如解码阶段对完整 KV cache 做增量注意力、或 cross-attention 中 query 比 key/value 短时因果的语义出现了歧义三角形掩码应该对齐左上角还是对齐右下角torch.nn.attention.bias模块就是为解决这个问题而生的。从源码 torch/nn/attention/bias.py 的模块 docstring 可以看到其定位Defines bias subclasses that work with scaled_dot_product_attention它定义了可被 SDPA 直接消费的偏置子类bias subclasses核心导出在 torch/nn/attention/bias.py__all__ [causal_upper_left, causal_lower_right, CausalVariant, CausalBias]torch.nn.functional的is_causal参数文档也明确引用了这个模块来定义非方阵时的行为见 torch/nn/functional.py当掩码为非方阵时is_causalTrue采用的是upper-left 对齐的因果偏置形态。CausalVariant两种因果掩码的几何语义CausalVariant是一个IntEnum定义两种因果变体见 torch/nn/attention/bias.pyUPPER_LEFT左上对齐等价于is_causalTrue的标准因果注意力构造代码为torch.tril(torch.ones(size, dtypetorch.bool))以shape(3,4)为例物化后的布尔掩码为[[1, 0, 0, 0], [1, 1, 0, 0], [1, 1, 1, 0]]即第i行的 query 只能看到前i1个 key。这是自回归解码中 query 位于序列开头一侧时的自然语义。LOWER_RIGHT右下对齐其包含值True对齐到矩阵右下角等价构造代码为diagonal_offset size[1] - size[0] torch.tril( torch.ones(size, dtypetorch.bool), diagonaldiagonal_offset, )以shape(3,4)为例[[1, 1, 0, 0], [1, 1, 1, 0], [1, 1, 1, 1]]这种变体适用于 query 是 KV 序列的后缀例如自回归模型中 query 是最新的几个 token而 KV cache 包含全部历史 token的场景——每个 query 只能看到自己及之前的 token而最新的 query 能看到全部 KV。文档特别指出当 query 与 key/value 序列长度相等时两种变体完全等价因为此时三角形矩阵恰好是方阵左上对齐与右下对齐产生同一张下三角矩阵。源码中两种变体分别由_upper_left与_lower_right两个私有方法实现torch/nn/attention/bias.py都基于torch.tril生成 bool 张量_materialize(device)方法则根据variant分派到对应构造未指定 device 时默认落在 CPU。CausalBias 类懒物化的 Tensor 子类CausalBias继承自torch.Tensortorch/nn/attention/bias.py但它并不持有真实的注意力偏置数据而是一个只保存形状与变体信息的描述对象。这带来一个重要工程收益(seq_len_q, seq_len_kv)大小的掩码不需要在构造时物化占用显存只有真正需要比如回退到 math 内核时才调用_materialize生成。构造函数参数CausalBias.__init__接受三个参数torch/nn/attention/bias.py参数类型说明variantCausalVariant偏置变体必须是UPPER_LEFT或LOWER_RIGHT否则抛出AssertionErrorseq_len_qintquery 的序列长度seq_len_kvintkey/value 的序列长度构造时有一条重要的安全性检查if seq_len_q seq_len_kv and variant CausalVariant.LOWER_RIGHT: warn( Lower right causal bias will produce NaNs in the output when seq_len_q seq_len_kv!, stacklevel2, )即使用LOWER_RIGHT变体时若seq_len_q seq_len_kv掩码的第一行会全为 False没有任何 token 可见softmax 对全-inf行求值会产生 NaN。测试代码中同样显式跳过这种组合见 test/test_transformers.pyif causal_variant CausalVariant.LOWER_RIGHT and seq_len_q seq_len_kv: self.skipTest( Lower right causal mask will produce NaNs in the output when seq_len_q seq_len_kv! )两个工厂函数推荐通过模块级工厂函数构造而不是直接实例化CausalBiascausal_upper_left(*size)torch/nn/attention/bias.py创建左上对齐的因果偏置等价于is_causalTrue。causal_lower_right(*size)torch/nn/attention/bias.py创建右下对齐的因果偏置。两者都要求恰好两个尺寸参数分别对应seq_len_q和seq_len_kv传入其他数量会抛出AssertionError如causal_lower_right only supports 2D tensors。from torch.nn.attention.bias import causal_upper_left, causal_lower_right # 128 长度的 query 对 256 长度的 KV cache 做因果注意力 bias causal_lower_right(128, 256) # LOWER_RIGHTquery 是 KV 后缀 bias causal_upper_left(128, 256) # UPPER_LEFTquery 是 KV 前缀由于CausalBias重写了__repr__为_materialize().__repr__()torch/nn/attention/bias.py在交互式环境中打印 bias 对象会看到完整物化后的布尔矩阵便于调试时直观确认掩码形状。实战示例与 scaled_dot_product_attention 配合使用CausalBias类 docstring 中给出了完整示例torch/nn/attention/bias.py这里完整保留并补充注释from torch.nn.attention.bias import causal_lower_right import torch.nn.functional as F import torch bsz, num_heads, seqlen_q, seqlen_kv, head_dim 32, 8, 4, 12, 8 # 创建右下对齐的因果偏置query 是 KV 序列的最后 4 个位置 attn_bias causal_lower_right(seqlen_q, seqlen_kv) q torch.randn( bsz, num_heads, seqlen_q, head_dim, devicecuda, dtypetorch.float16 ) k torch.randn( bsz, num_heads, seqlen_kv, head_dim, devicecuda, dtypetorch.float16 ) v torch.randn( bsz, num_heads, seqlen_kv, head_dim, devicecuda, dtypetorch.float16 ) # 直接把 CausalBias 作为 attn_mask 传入无需手动物化 out F.scaled_dot_product_attention(q, k, v, attn_bias)需要注意的 API 约束is_causal与 CausalBias 互斥_dispatch的开头即检查torch/nn/attention/bias.py两者同时为真会抛出ValueError: CausalBias should not be used with causalTrue。测试test_is_causal_and_mask_fails验证了该错误信息test/test_transformers.py。布尔掩码语义True 表示该位置参与注意力与nn.MultiheadAttention的key_padding_maskTrue 表示被屏蔽语义相反这一点在 torch/nn/functional.py 的文档中有专门说明。dropout 行为SDPA 会始终按dropout_p应用 dropout评估阶段应显式传0.0。源码注释与测试均标注CausalBias是 prototype/beta API接口可能随版本变化。内核分派机制torch_function与 _dispatchCausalBias能直接喂给 SDPA 的关键在于它重写了__torch_function__torch/nn/attention/bias.pyclassmethod def __torch_function__(cls, func, types, args(), kwargsNone): if kwargs is None: kwargs {} if func is torch.nn.functional.scaled_dot_product_attention: return cls._dispatch(*args, **kwargs) return super().__torch_function__(func, types, args, kwargs)也就是说当你调用F.scaled_dot_product_attention(q, k, v, attn_bias)且attn_bias是CausalBias实例时PyTorch 的函数覆盖协议会自动将调用劫持到CausalBias._dispatch静态方法由它决定走哪条执行路径。三条分派路径_dispatchtorch/nn/attention/bias.py的逻辑可归纳为三条路径路径一等价于 is_causalTrue直接复用融合内核的因果模式if ( attn_mask.seq_len_q attn_mask.seq_len_kv or attn_mask.variant CausalVariant.UPPER_LEFT ): return F.scaled_dot_product_attention( query, key, value, attn_maskNone, dropout_pdropout_p, is_causalTrue, scalescale, enable_gqaenable_gqa, )只要序列等长或者变体是UPPER_LEFT就没有必要物化掩码——直接委托给 SDPA 原生的is_causalTrue快速路径让各后端自己用硬件友好的方式实现因果掩码。UPPER_LEFT在方阵与非方阵下都与is_causalTrue语义一致所以无论形状如何都走这条路径。测试test_is_causal_equals_upper_left对多种非方阵形状验证了两者输出逐元素一致test/test_transformers.py。路径二LOWER_RIGHT Flash Attentionelif attn_mask.variant CausalVariant.LOWER_RIGHT: _validate_sdpa_input(query, key, value, None, dropout_p, is_causal, scale) sdpa_params SDPAParams(query, key, value, None, dropout_p, is_causal, enable_gqa) if can_use_flash_attention(sdpa_params): alignment 1 if query.device.type xpu else 8 og_head_size query.size(-1) og_scale _calculate_scale(og_head_size, scale) needs_padding og_head_size % alignment ! 0 if needs_padding: pad_len alignment - (og_head_size % alignment) query torch.nn.functional.pad(query, (0, pad_len)) key torch.nn.functional.pad(key, (0, pad_len)) value torch.nn.functional.pad(value, (0, pad_len)) out torch.ops.aten._scaled_dot_product_flash_attention( query, key, value, dropout_p, is_causalTrue, # TODO: Flash accepts causal True and for this particular op it means lower right return_debug_maskFalse, scaleog_scale, )[0] return _postprocess_flash_output(out, og_head_size)这段实现有几处值得注意的细节head_dim 对齐 paddingCUDA 上 Flash 内核要求 head size 为 8 的倍数XPU 上为 1不满足时会临时给 q/k/v 的最后一维 pad 到对齐长度计算完再通过_postprocess_flash_output裁回原宽度。Flash 内核对is_causalTrue的解释源码中的 TODO 注释指出对于_scaled_dot_product_flash_attention这个底层算子is_causalTrue的实际语义恰好就是lower-right掩码——这与上层F.scaled_dot_product_attention(is_causalTrue)的 upper-left 语义不同因此_dispatch才能以一行is_causalTrue调用精确表达 LOWER_RIGHT 语义无需任何掩码张量。GQA 说明这条 Flash 路径构造SDPAParams时透传了enable_gqa但_scaled_dot_product_flash_attention调用本身没有传递 enable_gqa 参数而 upper-left 路径路径一则完整透传了enable_gqa。从源码结构看LOWER_RIGHT GQA 的组合支持情况取决于底层算子版本使用前可结合sdpa_kernel上下文确认实际选中的后端。路径三LOWER_RIGHT Efficient Attention 或回退物化if can_use_efficient_attention(sdpa_params): compute_log_sumexp False if _input_requires_grad(query, key, value): compute_log_sumexp True return torch.ops.aten._efficient_attention_forward( query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2), biasNone, ... custom_mask_typeint(attn_mask.variant), compute_log_sumexpcompute_log_sumexp, scalescale, ... )[0].transpose(1, 2) else: _raise_kernel_warnings(sdpa_params) # We cant use efficient attention the only support for lower right is via materialization return F.scaled_dot_product_attention( query, key, value, attn_maskattn_mask._materialize(query.device), dropout_pdropout_p, is_causalFalse, scalescale, enable_gqaenable_gqa, )Efficient Attention 路径传入custom_mask_typeint(attn_mask.variant)——由于CausalVariant是IntEnumLOWER_RIGHT直接映射为一个整型枚举值由 C 侧的 xformers 风格内核解释为右下因果掩码同时若输入需要求梯度则开启compute_log_sumexp以支持反向。若两个融合内核都不可用例如 CPU、MPS 或不满足融合内核的输入约束最后回退到物化路径真正调用_materialize生成布尔掩码再以普通attn_mask形式交给 SDPA 的 math 内核处理。源码注释也直白地说明the only support for lower right is via materialization。_raise_kernel_warnings配合 torch/nn/attention/init.py 的WARN_FOR_UNFUSED_KERNELS全局开关可让使用者在设置torch.nn.attention.WARN_FOR_UNFUSED_KERNELS True后看到融合内核不可用的具体原因。值得强调的是无论走哪条路径CausalBias都避免了为 (seq_len_q, seq_len_kv) 生成显式掩码张量除最后回退路径外这正是把因果掩码做成 Tensor 子类而非普通 bool 张量的核心价值。与 torch.compile 的集成CausalBias与torch.compile的兼容性在源码顶部就做了声明torch/nn/attention/bias.pytorch._dynamo.allow_in_graph(is_flash_attention_available) torch._dynamo.allow_in_graph(can_use_flash_attention) torch._dynamo.allow_in_graph(can_use_efficient_attention) torch._dynamo.allow_in_graph(SDPAParams)这些allow_in_graph调用让 Dynamo 追踪时不将内核可用性判断视为图断点从而保证包含CausalBias的注意力调用可以被整体编译。测试test_causal_variants_compile使用CompileCounterWithBackend(aot_eager)验证了带 CausalBias 的 SDPA 在torch.compile下只产生一个编译帧即没有发生意外的图断裂test/test_transformers.pycnts CompileCounterWithBackend(aot_eager) ... self.assertEqual(cnts.frame_count, 1, Compiled graph should have 1 frame!)正确性验证测试用例给出的参照实现test/test_transformers.py 中的TestAttnBias测试类约 L6895-L7051是理解该模块语义的可靠参照其run_test的做法是将attn_bias._materialize(device)物化出的普通布尔掩码走一遍 SDPA 作为参考输出将CausalBias原样作为attn_mask再走一遍可选经torch.compile对前向输出与 q/k/v 的梯度分别torch.testing.assert_close。参数化形状覆盖了(16,16,128,128,16)方阵、(16,16,128,256,32)query 短于 KV、(16,16,256,128,32)query 长于 KV以及非 2 的幂形状(1,1,23,56,15)并对 float16 使用Tolerances(1e-3, 1e-3)前向 /Tolerances(5e-3, 5e-3)反向的容差test/test_transformers.py。此外SDPA 泛型测试里还用causal_lower_right作为数学参照来校验融合内核在is_causal场景下的结果test/test_transformers.py。使用建议与适用前提综合文档与源码使用torch.nn.attention.bias时的要点优先用causal_upper_left它与is_causalTrue完全等价在所有后端上都有融合支持且代码路径最简单只有当你的 query 是 KV 序列的后缀、需要最新 query 可见全部历史语义时才使用causal_lower_right。避免LOWER_RIGHTseq_len_q seq_len_kv会触发 NaN 警告softmax 遇到整行不可见应改用causal_upper_left或调整序列组织方式。不要同时传is_causalTrue会直接抛ValueError。关注融合内核可用性LOWER_RIGHT 的 Flash/Efficient 快速路径主要在 CUDAXPU 上 head 对齐要求为 1等支持融合内核的设备上生效在不可用时会回退为物化掩码 math 路径此时显存上会出现一个(seq_len_q, seq_len_kv)的 bool 掩码。可用 torch/nn/attention/init.py 提供的sdpa_kernel上下文管理器和SDPBackend枚举显式约束后端例如只允许[SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION, SDPBackend.MATH]。API 稳定性源码对CausalBias与CausalVariant均标注 prototype 警告升级 PyTorch 版本时建议回归TestAttnBias相关测试确认行为未变。小结torch.nn.attention.bias用极小的 API 面一个枚举、一个 Tensor 子类、两个工厂函数解决了非方阵因果注意力的语义歧义问题并借助__torch_function__协议把懒描述透明地接入F.scaled_dot_product_attention的分发体系能在 Flash/Efficient 融合内核中以is_causal/custom_mask_type表达就绝不物化掩码不能时再优雅回退。文档入口 docs/source/nn.attention.bias.md 与实现 torch/nn/attention/bias.py、测试 test/test_transformers.py 三者对照阅读可以快速建立从 API 到内核分派的完整心智模型。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
上一篇/下一篇内容由系统自动关联
返回资讯列表 →