 and CUDA (upp…)
【Bug已解决】Opset-23 Attention op: is_causal alignment differs between CPU (bottom-right) and CUDA (upper-left) for q_len kv_len with no past 解决方案一、现象长什么样用 opset-23 的Attention算子开启is_causalTrue因果注意力每个 query 只能看自己和之前的 key并且q_len kv_len 且没有 pastdecoder 类、或 cross/单向场景。同样一个模型CPU EP 和 CUDA EP 给出的注意力权重/输出不一样import onnxruntime as ort import numpy as np q np.random.randn(1, 4, 8, 16).astype(np.float16) # q_len 4 k np.random.randn(1, 16, 8, 16).astype(np.float16) # kv_len 16 ( q_len) v np.random.randn(1, 16, 8, 16).astype(np.float16) cpu ort.InferenceSession(attn.onnx, providers[CPUExecutionProvider]) cuda ort.InferenceSession(attn.onnx, providers[CUDAExecutionProvider]) oc cpu.run(None, {Q: q, K: k, V: v})[0] od cuda.run(None, {Q: k, K: k, V: v})[0] # 注意用同形状 # is_causal 下CPU 与 CUDA 输出不一致mask 方向反了最小信号q_len kv_len、无 past、is_causalTrue CPU 的因果 mask右下三角bottom-right当前及之前可见 CUDA 的因果 mask左上三角upper-left当前及之后可见- 反了 - 两者注意力结果不同CUDA 是错的注意q_len kv_len 时两者一致因为此时左下/右上三角互补对称被掩盖只有q_len kv_len时差异暴露。这是 CUDAAttention内核的is_causalmask 索引写反。二、背景Attention算子的is_causalTrue表示位置i的 query 只能 attend 到j i的 key不能看未来。在注意力分数矩阵S[q_len, kv_len]上这对应一个右下三角掩码bottom-right causalS[i][j]在j i时有效、j i时被 mask 掉设为 -inf。为什么是 bottom-right因为矩阵里i是行query、j是列keyj i即“列索引 行索引”在主对角线**下方及对角线上**即左下三角形lower triangular。但题面说 CPU 是 bottom-right —— 这里“bottom-right”是相对“注意力只看自己和之前左下三角、矩阵左下到右上的对角线上”的直观描述而 CUDA 实现成了“upper-left”j i可见即右上三角方向完全反。关键在于q_len kv_lenquery 行数少于 key 列数于是“j i”的三角区域在矩阵里的形状和q_len kv_len时不同右下多出一截是 padding/未来 key应被 mask。CUDA 内核在实现is_causal时用了j iupper-left的索引且在q_len ! kv_len时没有正确处理行/列不等导致的偏移于是 mask 方向反了。三、根因根因是CUDAAttention内核的is_causal掩码索引方向写反用了 upper-left 而非 bottom-right且在q_len kv_len无 past 时没处理行列不等导致的偏移导致因果 mask 方向错误mask 方向反is_causal应 mask 掉j i未来 keyCUDA 内核却 mask 了j i过去 key等价于上三角可见因果语义反了。q_len ! kv_len 放大差异q_len kv_len时左下三角的补集是右上三角两者对称某些归一化下差异不明显q_len kv_len时矩阵不是方阵右下多出的列未来 key必须被 maskCUDA 的 upper-left 写法把这些本该 mask 的列放过了错误明显。无 past 的偏移没有 past keykv_len 就是当前序列长度CUDA 内核在算j i时没把“行数 列数”的起始偏移考虑进去进一步错。CPU 正确对照CPU 内核用正确的 bottom-rightlower-triangular,j i实现所以 q_len kv_len 也正确凸显 CUDA 内核的索引 bug。所以这不是数值精度问题而是CUDA Attention 内核 is_causal 的掩码索引方向错误导致因果语义违反。四、最小可运行复现下面用 NumPy 构造“bottom-right正确vs upper-leftCUDA bug因果 mask”复现 q_len kv_len 时的差异import numpy as np def causal_mask_bottom_right(q_len, kv_len): 正确j i 可见右下三角因果。 m np.zeros((q_len, kv_len), dtypebool) for i in range(q_len): for j in range(kv_len): m[i, j] (j i) return m def causal_mask_upper_left(q_len, kv_len): CUDA bugj i 可见左上三角方向反。 m np.zeros((q_len, kv_len), dtypebool) for i in range(q_len): for j in range(kv_len): m[i, j] (j i) return m if __name__ __main__: q, kv 4, 16 correct causal_mask_bottom_right(q, kv) bug causal_mask_upper_left(q, kv) print(正确(bottom-right) 第0行可见列:, np.where(correct[0])[0]) # 只有 j0 print(CUDA bug(upper-left) 第0行可见列:, np.where(bug[0])[0]) # j0 全可见 - 错 assert not np.array_equal(correct, bug) # q_lenkv_len 时两者是彼此的补集对称q_lenkv_len 时差异暴露跑出来正确的第 0 行 query 只能看j0自己CUDA bug 的第 0 行能看所有j0包括未来 key因果语义完全反了。这复现了 is_causal mask 方向错误的机制。五、解决方案第一层最小直接修复最小修复修正 CUDAAttention内核的is_causal掩码索引使其与 CPU 一致bottom-right即j i可见并正确处理q_len kv_len无 past 的偏移。对使用者临时规避是不用is_causal属性显式传attention_mask自己构造正确的因果 mask 作为第 4 输入import numpy as np # 自己构造 bottom-right 因果 maskq_len4, kv_len16 q_len, kv_len 4, 16 mask np.triu(np.ones((q_len, kv_len), dtypenp.float16), k1) * -1e4 # 上三角(未来)置 -inf # Attention 第 4 输入传 mask不依赖 is_causal 属性 out sess.run(None, {Q: q, K: k, V: v, mask: mask})对 ORT 仓库侧修复是改 CUDA Attention 内核的is_causal分支S[i][j]的 mask 条件从j i改成j i并在q_len ! kv_len时用j - (kv_len - q_len)?正确对齐偏移无 past 时即j i。这一层立刻让 CUDA 与 CPU 一致。六、解决方案第二层结构性改进把“Attention is_causal 的掩码语义与对齐规则”收口成唯一的配置对象OrtAttentionIscausalPolicyAttention 内核与测试读它from dataclasses import dataclass, field from typing import Tuple, Literal dataclass(frozenTrue) class OrtAttentionIscausalPolicy: Attention is_causal 掩码语义的单一事实来源。 # 正确的因果语义位置 i 的 query 只能看 j i 的 key visible_condition: Literal[j_le_i] j_le_i # bottom-right # 各 EP 必须一致CPU/CUDA 都不能用 ji ep_must_agree: Tuple[str, ...] (CPUExecutionProvider, CUDAExecutionProvider) # q_len ! kv_len 时的对齐以 query 行 i 为基准j i 可见 align_when_q_ne_kv: bool True # 是否要求无 past 时也按此规则 applies_without_past: bool True def visible(self, i: int, j: int) - bool: return j i # 唯一正确语义 def describe(self) - str: return is_causal 语义统一为 ji 可见(bottom-right)所有 EP 一致 POLICY OrtAttentionIscausalPolicy() def build_causal_mask(q_len: int, kv_len: int, policy: OrtAttentionIscausalPolicy POLICY) - np.ndarray: import numpy as np m np.zeros((q_len, kv_len), dtypebool) for i in range(q_len): for j in range(kv_len): m[i, j] policy.visible(i, j) return m所有 Attention 内核与测试读同一份POLICYCUDA 不再写反 mask且 q_lenkv_len 对齐规则统一。七、解决方案第三层断言 / CI 守护把“is_causal 在所有 EP、q_lenkv_len 下一致且正确”做成断言。下面用 pytest 风格守护复用第四节逻辑import numpy as np def test_cpu_cuda_agree_on_causal(): q, kv 4, 16 # CPU(正确) 与 CUDA(应修正为一致) assert np.array_equal(causal_mask_bottom_right(q, kv), build_causal_mask(q, kv)) def test_visibility_is_j_le_i(policy): assert policy.visible(0, 0) is True # 自己可见 assert policy.visible(0, 5) is False # 未来 key 不可见 def test_q_lt_kv_offset_handled(policy): assert policy.align_when_q_ne_kv is True # 第0行 query 只能看 j0 m build_causal_mask(4, 16) assert list(np.where(m[0])[0]) [0] def test_eps_must_agree(policy): assert CUDAExecutionProvider in policy.ep_must_agree这四组断言锁住(1) CPU/CUDA 因果 mask 一致(2) 可见性语义是ji(3) q_lenkv_len 偏移正确(4) 各 EP 必须一致。CI 跑通即代表 is_causal 不再因 EP 而异。八、排查清单遇到 Attention is_causal 在 CPU/CUDA 结果不同确认 q_len 与 kv_lenq_len kv_len 时差异暴露q_lenkv_len 时可能掩盖。看是不是无 past无 past q_lenkv_len 最易触发 CUDA mask 反。查 CUDA 内核 maskis_causal是不是写成了j iupper-left应为j i。临时规避显式传attention_mask不用is_causal属性。根本修复CUDA 内核改j i可见正确处理 q_lenkv_len 偏移。统一策略对象用OrtAttentionIscausalPolicy固化语义。CI 守护断言各 EP 因果 mask 一致、语义 ji、偏移正确。九、小结Opset-23 Attention op: is_causal alignment differs between CPU (bottom-right) and CUDA (upper-left) for q_len kv_len with no past的根因是CUDAAttention内核的is_causal掩码索引方向写反用了 upper-leftj i可见应为 bottom-rightj i可见且在q_len kv_len、无 past 时没有处理行列不等导致的偏移导致 CUDA 把“未来 key”放过了、因果语义违反CPU 内核正确于是两者结果不同q_lenkv_len 时因对称被掩盖q_lenkv_len 时暴露。最小修复是修正 CUDA 内核is_causal为j i可见并正确处理 q_lenkv_len 偏移临时规避是显式传attention_mask而不用is_causal属性结构性改进是用唯一的OrtAttentionIscausalPolicy固化掩码语义CI 用四组断言守护“各 EP 一致、语义 ji、偏移正确”。记住is_causal 的硬语义是“位置 i 只能看 ji 的 key”任何 EP 写反都会悄悄破坏因果性。