首页
/ PyTorch torch.nn.attention.bias 深度解析:CausalBias 因果注意力偏置的原理与实战

PyTorch torch.nn.attention.bias 深度解析:CausalBias 因果注意力偏置的原理与实战

2026-09-09 15:07:09作者:谭伦延

本篇指南围绕 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_attentionis_causal=True 参数在查询与键值序列长度相等(方阵)时语义是明确的:按主对角线做下三角掩码。但在 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.functionalis_causal 参数文档也明确引用了这个模块来定义非方阵时的行为,见 torch/nn/functional.py:当掩码为非方阵时,is_causal=True 采用的是 upper-left 对齐 的因果偏置形态。

CausalVariant:两种因果掩码的几何语义

CausalVariant 是一个 IntEnum,定义两种因果变体(见 torch/nn/attention/bias.py):

UPPER_LEFT:左上对齐

等价于 is_causal=True 的标准因果注意力,构造代码为:

torch.tril(torch.ones(size, dtype=torch.bool))

shape=(3,4) 为例,物化后的布尔掩码为:

[[1, 0, 0, 0],
 [1, 1, 0, 0],
 [1, 1, 1, 0]]

即第 i 行的 query 只能看到前 i+1 个 key。这是自回归解码中 query 位于序列开头一侧时的自然语义。

LOWER_RIGHT:右下对齐

其包含值(True)对齐到矩阵右下角,等价构造代码为:

diagonal_offset = size[1] - size[0]
torch.tril(
    torch.ones(size, dtype=torch.bool),
    diagonal=diagonal_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):

参数 类型 说明
variant CausalVariant 偏置变体,必须是 UPPER_LEFTLOWER_RIGHT,否则抛出 AssertionError
seq_len_q int query 的序列长度
seq_len_kv int key/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!",
        stacklevel=2,
    )

即使用 LOWER_RIGHT 变体时,若 seq_len_q > seq_len_kv,掩码的第一行会全为 False(没有任何 token 可见),softmax 对全 -inf 行求值会产生 NaN。测试代码中同样显式跳过这种组合,见 test/test_transformers.py

if 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!"
    )

两个工厂函数

推荐通过模块级工厂函数构造,而不是直接实例化 CausalBias

两者都要求恰好两个尺寸参数(分别对应 seq_len_qseq_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_RIGHT:query 是 KV 后缀
bias = causal_upper_left(128, 256)    # UPPER_LEFT:query 是 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, device="cuda", dtype=torch.float16
)
k = torch.randn(
    bsz, num_heads, seqlen_kv, head_dim, device="cuda", dtype=torch.float16
)
v = torch.randn(
    bsz, num_heads, seqlen_kv, head_dim, device="cuda", dtype=torch.float16
)

# 直接把 CausalBias 作为 attn_mask 传入,无需手动物化
out = F.scaled_dot_product_attention(q, k, v, attn_bias)

需要注意的 API 约束:

  1. is_causal 与 CausalBias 互斥_dispatch 的开头即检查(torch/nn/attention/bias.py),两者同时为真会抛出 ValueError: CausalBias should not be used with causal=True。测试 test_is_causal_and_mask_fails 验证了该错误信息(test/test_transformers.py)。
  2. 布尔掩码语义:True 表示该位置参与注意力,与 nn.MultiheadAttentionkey_padding_mask(True 表示被屏蔽)语义相反,这一点在 torch/nn/functional.py 的文档中有专门说明。
  3. dropout 行为:SDPA 会始终按 dropout_p 应用 dropout,评估阶段应显式传 0.0
  4. 源码注释与测试均标注 CausalBias 是 prototype/beta API,接口可能随版本变化。

内核分派机制:torch_function 与 _dispatch

CausalBias 能"直接喂给" SDPA 的关键在于它重写了 __torch_function__torch/nn/attention/bias.py):

@classmethod
def __torch_function__(cls, func, types, args=(), kwargs=None):
    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_biasCausalBias 实例时,PyTorch 的函数覆盖协议会自动将调用劫持到 CausalBias._dispatch 静态方法,由它决定走哪条执行路径。

三条分派路径

_dispatchtorch/nn/attention/bias.py)的逻辑可归纳为三条路径:

路径一:等价于 is_causal=True,直接复用融合内核的因果模式

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_mask=None, dropout_p=dropout_p, is_causal=True,
        scale=scale, enable_gqa=enable_gqa,
    )

只要序列等长,或者变体是 UPPER_LEFT,就没有必要物化掩码——直接委托给 SDPA 原生的 is_causal=True 快速路径,让各后端自己用硬件友好的方式实现因果掩码。UPPER_LEFT 在方阵与非方阵下都与 is_causal=True 语义一致,所以无论形状如何都走这条路径。测试 test_is_causal_equals_upper_left 对多种非方阵形状验证了两者输出逐元素一致(test/test_transformers.py)。

路径二:LOWER_RIGHT + Flash Attention

elif 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_causal=True,  # TODO: Flash accepts causal = True and for this particular op it means lower right
            return_debug_mask=False,
            scale=og_scale,
        )[0]
        return _postprocess_flash_output(out, og_head_size)

这段实现有几处值得注意的细节:

  • head_dim 对齐 padding:CUDA 上 Flash 内核要求 head size 为 8 的倍数(XPU 上为 1),不满足时会临时给 q/k/v 的最后一维 pad 到对齐长度,计算完再通过 _postprocess_flash_output 裁回原宽度。
  • Flash 内核对 is_causal=True 的解释:源码中的 TODO 注释指出,对于 _scaled_dot_product_flash_attention 这个底层算子,is_causal=True 的实际语义恰好就是 lower-right 掩码——这与上层 F.scaled_dot_product_attention(is_causal=True) 的 upper-left 语义不同,因此 _dispatch 才能以一行 is_causal=True 调用精确表达 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),
            bias=None,
            ...
            custom_mask_type=int(attn_mask.variant),
            compute_log_sumexp=compute_log_sumexp,
            scale=scale,
            ...
        )[0].transpose(1, 2)
    else:
        _raise_kernel_warnings(sdpa_params)
        # We can't use efficient attention the only support for lower right is via materialization
        return F.scaled_dot_product_attention(
            query, key, value,
            attn_mask=attn_mask._materialize(query.device),
            dropout_p=dropout_p,
            is_causal=False,
            scale=scale,
            enable_gqa=enable_gqa,
        )
  • Efficient Attention 路径传入 custom_mask_type=int(attn_mask.variant)——由于 CausalVariantIntEnumLOWER_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.pyWARN_FOR_UNFUSED_KERNELS 全局开关,可让使用者在设置 torch.nn.attention.WARN_FOR_UNFUSED_KERNELS = True 后看到融合内核不可用的具体原因。

值得强调的是,无论走哪条路径,CausalBias避免了为 (seq_len_q, seq_len_kv) 生成显式掩码张量(除最后回退路径外),这正是把因果掩码做成 Tensor 子类而非普通 bool 张量的核心价值。

与 torch.compile 的集成

CausalBiastorch.compile 的兼容性在源码顶部就做了声明(torch/nn/attention/bias.py):

torch._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.py):

cnts = CompileCounterWithBackend("aot_eager")
...
self.assertEqual(cnts.frame_count, 1, "Compiled graph should have 1 frame!")

正确性验证:测试用例给出的参照实现

test/test_transformers.py 中的 TestAttnBias 测试类(约 L6895-L7051)是理解该模块语义的可靠参照,其 run_test 的做法是:

  1. attn_bias._materialize(device) 物化出的普通布尔掩码走一遍 SDPA 作为参考输出;
  2. CausalBias 原样作为 attn_mask 再走一遍(可选经 torch.compile);
  3. 对前向输出与 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_causal=True 完全等价,在所有后端上都有融合支持,且代码路径最简单;只有当你的 query 是 KV 序列的后缀、需要"最新 query 可见全部历史"语义时才使用 causal_lower_right
  • 避免 LOWER_RIGHT + seq_len_q > seq_len_kv:会触发 NaN 警告(softmax 遇到整行不可见),应改用 causal_upper_left 或调整序列组织方式。
  • 不要同时传 is_causal=True:会直接抛 ValueError
  • 关注融合内核可用性:LOWER_RIGHT 的 Flash/Efficient 快速路径主要在 CUDA(XPU 上 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 稳定性:源码对 CausalBiasCausalVariant 均标注 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 到内核分派的完整心智模型。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.76 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
860
1.35 K
docsdocs
暂无描述
Markdown
899
5.83 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
925
1.85 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.84 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
533
601
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.03 K
525
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.37 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
395