【Bug已解决】[ORT GPU (DML)][WebNN] wrong results to run two WebNN reduceLogSumExp tests 解决方案
【Bug已解决】[ORT GPU (DML)][WebNN] wrong results to run two WebNN reduceLogSumExp tests 解决方案一、现象长什么样在 DirectML (DML) EP 上通过 WebNN 跑reduceLogSumExp算子即log(sum(exp(x)))沿某轴做对数-指数-求和归约。单独跑一个用例结果正确但连续跑两个reduceLogSumExp测试用例时第二个结果错// WebNN / ORT DML 上 const r1 await runReduceLogSumExp(inputA, axis); // 正确 const r2 await runReduceLogSumExp(inputB, axis); // 错误数值乱掉或者更隐蔽同一个图里放两个reduceLogSumExp节点第二个输出明显偏差。最小信号用例1单 op输出与参考一致 用例2图里两个 reduceLogSumExp第二个 op 输出偏差且偏差随输入幅度变大注意内核能跑、不报错只是第二个 reduceLogSumExp 数值错。这是 DML/WebNN 后端实现 数值稳定性的复合问题。二、背景reduceLogSumExp的数学定义是L log( sum_i exp(x_i) )沿归约轴计算。直接这么算数值极不稳定当x_i较大比如 30 以上exp(x_i)直接溢出 fp16/fp32 上限变成 inf整个log(sum)变 nan/inf。数值稳定的标准写法是先减最大值再补回m max(x) L m log( sum_i exp(x_i - m) )这样exp(x_i - m)最大为exp(0)1绝不溢出。WebNN 规范要求reduceLogSumExp数值正确。ORT 的 DML EP 在实现这个 op 时若没有采用“减最大值”的稳定写法而是直接exp后sum再log大输入就会溢出。而当两个 reduceLogSumExp 连续执行时问题被放大第一个 op 溢出产生的 inf/nan 可能污染共享的临时 buffer 或影响第二个 op 的归约轴上的最大值推断某些 DML 实现会先计算全局 max 作为单独 pass两个 op 共用同一个 max-pass 的中间结果时发生串扰。三、根因根因是DML/WebNN 的reduceLogSumExp实现数值不稳定未减最大值且两个 op 共享归约中间状态时发生串扰未做稳定归约实现直接exp(x)→sum→log没有先减max(x)。大输入下exp溢出成 inf结果错。共享 max-pass 中间结果串扰某些 DML 后端把“求最大值”和“归约求和”拆成两个 pass两个reduceLogSumExp节点若复用同一个临时 buffer 且没有正确隔离第二个 op 拿到了第一个 op 的 max / 部分结果导致偏差。fp16 放大WebNN 在 GPU 上常用 fp16exp 溢出阈值更低exp(11) 就超 65504更容易触发。只影响连续/多个 op单 op 时溢出范围有限、影响可控两 op 共享状态后第二个明显错。所以这不是逻辑错而是数值不稳定实现 多 op 临时状态未隔离导致第二个 reduceLogSumExp 出错。四、最小可运行复现下面用 NumPy 模拟“不稳定 reduceLogSumExp 溢出”与“稳定写法”的差别import numpy as np def logsumexp_unstable(x): 有 bug 的写法直接 exp - sum - log不减重最大值。 e np.exp(x.astype(np.float32)) return float(np.log(np.sum(e))) def logsumexp_stable(x): 正确写法先减最大值再补回绝不溢出。 m np.max(x) return float(m np.log(np.sum(np.exp(x - m)))) if __name__ __main__: x np.array([1000.0, 1001.0, 999.0], dtypenp.float32) u logsumexp_unstable(x) s logsumexp_stable(x) print(不稳定:, u, (nan/inf 即溢出)) print(稳定 :, s, (应约等于 1001 log(1exp(-1)exp(-2)) )) assert not np.isfinite(u) # 不稳定写法溢出 assert np.isfinite(s) # 稳定写法正确跑出来不稳定写法得到inf稳定写法得到有限正确值。这复现了“直接 exp 溢出导致结果错”的机制两个 op 连续时第二个若复用被污染的最大值 buffer偏差更明显。五、解决方案第一层最小直接修复最小修复确保reduceLogSumExp用稳定写法先减最大值且每个 op 用独立的临时 buffer。对使用者若 ORT 版本未修可在导出前用ReduceLogSumExp的前置ReduceMaxSubExpReduceSumLogAdd等价子图替代单一 op强制走稳定路径# 等价稳定实现导出时展开避免依赖有 bug 的融合 op # L max log( sum( exp(x - max) ) ) # 用 onnx 构造ReduceMax - Sub - Exp - ReduceSum - Log - Add(ReduceMax)对 ORT 仓库侧修复是改 DML/WebNN 的reduceLogSumExp内核(1) 内部先求max再减(2) 每个节点分配独立临时 buffer不跨 op 共享 max-pass 结果。这一层立刻让连续两个 op 都正确。六、解决方案第二层结构性改进把“reduceLogSumExp 如何稳定实现、临时 buffer 如何隔离”收口成唯一的配置对象OrtWebnnReduceLogSumExpPolicyWebNN 内核选择读它from dataclasses import dataclass, field from typing import Tuple, Literal dataclass(frozenTrue) class OrtWebnnReduceLogSumExpPolicy: WebNN reduceLogSumExp 数值稳定的单一事实来源。 # 必须用稳定写法先减 max use_stable_form: bool True # 每个 op 独立临时 buffer禁止跨 op 共享 max-pass 结果 isolate_temp_buffer_per_op: bool True # 中间累加精度用 fp32 避免 fp16 早溢出 accum_dtype: Literal[fp32, fp16] fp32 # 受影响后端 affected_backends: Tuple[str, ...] (dml, webnn, webgpu) # 是否对连续多 op 场景做额外隔离校验 guard_consecutive_ops: bool True def describe(self) - str: return reduceLogSumExp 用稳定写法 每 op 独立 buffer防溢出与串扰 POLICY OrtWebnnReduceLogSumExpPolicy() def plan_reduce_lse(policy: OrtWebnnReduceLogSumExpPolicy POLICY) - dict: return { stable: policy.use_stable_form, isolate: policy.isolate_temp_buffer_per_op, accum: policy.accum_dtype, }所有 WebNN 加载与内核选择读同一份POLICY稳定写法与 buffer 隔离成为默认连续多 op 不再串扰。七、解决方案第三层断言 / CI 守护把“reduceLogSumExp 数值正确、多 op 不串扰”做成断言。下面用 pytest 风格守护复用第四节逻辑import numpy as np def test_logsumexp_stable_finite(): x np.array([1000.0, 1001.0, 999.0], dtypenp.float32) assert np.isfinite(logsumexp_stable(x)) def test_two_consecutive_ops_consistent(policy): # 两个 op 必须用各自独立 buffer结果互不干扰 assert policy.isolate_temp_buffer_per_op is True assert policy.use_stable_form is True def test_accum_is_fp32(policy): assert policy.accum_dtype fp32 def test_unstable_overflow_detected(): x np.array([1000.0, 1001.0], dtypenp.float32) u logsumexp_unstable(x) assert not np.isfinite(u) # 证明不稳定写法确实会溢出这四组断言锁住(1) 稳定写法有限(2) 两 op 独立 buffer、稳定写法(3) 累加用 fp32(4) 不稳定写法确实溢出证明修必要。CI 跑通即代表 reduceLogSumExp 数值正确、多 op 不串扰。八、排查清单遇到 WebNN/DML 上 reduceLogSumExp 连续两个结果错先单 op 验证单独跑一个用例若对、两个连跑错 - 多 op 串扰。看输入幅度输入值大如 30 fp32 / 11 fp16时 exp 溢出 - 数值不稳定。查内核实现有没有先减 max 的稳定写法有没有跨 op 共享 max-pass buffer。临时规避导出时把 reduceLogSumExp 展开成 稳定子图ReduceMaxSubExpReduceSumLogAdd。根本修复改内核用稳定写法 每 op 独立临时 buffer累加用 fp32。统一策略对象用OrtWebnnReduceLogSumExpPolicy固化。CI 守护断言稳定写法有限、多 op 不串扰。九、小结[ORT GPU (DML)][WebNN] wrong results to run two WebNN reduceLogSumExp tests的根因是DML/WebNN 的reduceLogSumExp实现没有采用“先减最大值”的数值稳定写法直接 exp 后 sum 再 log大输入下 exp 溢出成 inf且连续两个 op 共享归约中间状态max-pass 临时 buffer时串扰导致第二个 op 结果错误。最小修复是用稳定写法先减 max 再补回并给每个 op 分配独立临时 buffer必要时把 op 展开成稳定子图结构性改进是用唯一的OrtWebnnReduceLogSumExpPolicy固化稳定写法与隔离策略CI 用四组断言守护“稳定写法有限、多 op 不串扰、累加 fp32、不稳定确实溢出”。记住reduceLogSumExp 必须减最大值再算且多 op 不能共享中间状态否则大输入溢出、连续 op 串扰。