PyTorch torch.nn.attention.bias 深度解析:CausalBias 因果注意力偏置的原理与实战
本篇指南围绕 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_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.functional 的 is_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.Tensor(torch/nn/attention/bias.py),但它并不持有真实的注意力偏置数据,而是一个只保存形状与变体信息的描述对象。这带来一个重要工程收益:(seq_len_q, seq_len_kv) 大小的掩码不需要在构造时物化占用显存,只有真正需要(比如回退到 math 内核)时才调用 _materialize 生成。
构造函数参数
CausalBias.__init__ 接受三个参数(torch/nn/attention/bias.py):
| 参数 | 类型 | 说明 |
|---|---|---|
variant |
CausalVariant |
偏置变体,必须是 UPPER_LEFT 或 LOWER_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:
causal_upper_left(*size)(torch/nn/attention/bias.py):创建左上对齐的因果偏置,等价于is_causal=True。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_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 约束:
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)。- 布尔掩码语义:True 表示该位置参与注意力,与
nn.MultiheadAttention的key_padding_mask(True 表示被屏蔽)语义相反,这一点在 torch/nn/functional.py 的文档中有专门说明。 - dropout 行为:SDPA 会始终按
dropout_p应用 dropout,评估阶段应显式传0.0。 - 源码注释与测试均标注
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_bias 是 CausalBias 实例时,PyTorch 的函数覆盖协议会自动将调用劫持到 CausalBias._dispatch 静态方法,由它决定走哪条执行路径。
三条分派路径
_dispatch(torch/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)——由于CausalVariant是IntEnum,LOWER_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.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 的做法是:
- 将
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_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 稳定性:源码对
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 到内核分派的完整心智模型。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0631
MiniCPM5-2BMiniCPM5-2B 是一款面向端侧、本地部署和资源受限场景的 2B 稠密 Transformer,能够达到同尺寸开源模型 SOTA 水平。Markdown00
video-shotcraftAI宣传片skill,使用 Remotion 制作电影级产品视频:提供106 张镜头配方卡和可复用的视频魔板。适用于 Claude Code 与 Codex以及所有其他智能体Markdown00
HivisionIDPhotos⚡️HivisionIDPhotos: a lightweight and efficient AI ID photos tools. 一个轻量级的AI证件照制作算法。Python09
DragonOSDragonOS is an operating system developed from scratch using Rust, with Linux compatibility. It is designed for **Serverless** scenarios. 使用Rust从0自研内核,具有Linux兼容性的操作系统,面向云计算Serverless场景而设计。Rust00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00