【Bug已解决】WebGPU ConvTranspose produces incorrect results for fp16 models when output spatial dim > 2… 📅 2026/8/13 21:58:24 【Bug已解决】WebGPU ConvTranspose produces incorrect results for fp16 models when output spatial dim 2048 (index math done in f16) 解决方案一、现象长什么样在 WebGPU EP 上跑一个含ConvTranspose转置卷积 / 反卷积的fp16模型。当输出空间维度超过 2048时结果错误图像错位、数值乱输出维度 ≤ 2048 时正常const session await ort.InferenceSession.create(convtranspose_fp16.onnx, { executionProviders: [webgpu], }); // 输出空间 dim 2048 - 正常 // 输出空间 dim 2048 - 结果错索引错位最小信号fp16 输出空间 dim 2048 - ConvTranspose 结果错 fp16 输出空间 dim 2048 - 正常 fp32 路径 - 正常注意不是崩溃是大图时数值错。根因是索引运算用了 fp16大索引精度丢失。二、背景ConvTranspose计算输出特征图时需要为每个输出元素计算它对应输入/权重的索引输出位置o映射到输入位置、kernel偏移等。这些索引运算本质是整数/坐标数学。在 WebGPU 的 fp16 路径上为了和输入张量fp16保持一致或省内存内核可能把索引计算也放在 fp16里做比如用 fp16 的中间变量算o * stride - pad * dilation ...。这就出问题了fp16 能精确表示的整数是-2048 ~ 2048超过这个范围相邻整数间隔变大无法精确表示每一个整数。当输出空间维度 2048索引值超过 2048用 fp16 表示的索引会取整错误比如 2049 和 2050 在 fp16 下可能都变成同一个值于是输出元素的索引算错 - 数据被放到错误位置 - 结果错错位/乱。输出 dim ≤ 2048 时索引还在 fp16 精确整数范围内所以正常fp32 路径用 fp32 算索引精确整数到几百万所以也正常。只有fp16 输出 dim 2048触发精度丢失。三、根因根因是WebGPUConvTranspose的 fp16 内核把坐标/索引运算放在 fp16 里做输出空间维度 2048 时索引超出 fp16 精确整数范围取整错误导致输出元素错位索引数学用 fp16内核用 fp16 变量算输出位置o、kernel 偏移等坐标而不是用int32/fp32。fp16 整数精度上限 2048fp16 在 2048 的整数上无法一一精确表示索引算错。只影响 fp16 大输出fp32 路径用 fp32 算索引精确小输出≤2048索引在 fp16 精确范围内所以正常。不是结构错ConvTranspose 算法对只是坐标用错精度。所以这不是逻辑错而是坐标/索引运算用了错误精度fp16大图时精度丢失导致错位。四、最小可运行复现下面用 NumPy 模拟“fp16 索引在 2048 时取整错误”import numpy as np def index_in_fp16(o): 用 fp16 表示输出坐标索引大索引会取整错误。 return np.float16(o).astype(np.float32) # fp16 转回看是否还是原数 def convtranspose_index_ok(out_dim): 检查 0..out_dim 每个索引能否被 fp16 精确表示。 bad 0 for o in range(out_dim 1): if index_in_fp16(o) ! o: bad 1 return bad if __name__ __main__: bad_2048 convtranspose_index_ok(2048) bad_4096 convtranspose_index_ok(4096) print(输出 dim2048 时索引丢失个数:, bad_2048) # 0 print(输出 dim4096 时索引丢失个数:, bad_4096) # 0错位 assert bad_2048 0 assert bad_4096 0 # 大输出索引在 fp16 下不准 - 错位跑出来输出 dim2048 时索引零丢失dim4096 时有大量索引在 fp16 下取整错误。这复现了“fp16 索引 2048 精度丢失导致 ConvTranspose 错位”的机制。五、解决方案第一层最小直接修复最小修复让ConvTranspose的坐标/索引运算用int32或fp32绝不用 fp16。对使用者临时规避是对该算子用 fp32导出时 Cast 成 fp32 过 ConvTranspose 再 Cast 回或整体 fp32# 导出时把 ConvTranspose 的输入/输出用 fp32索引数学天然精确 # 在 ConvTranspose 前后插 Cast fp32对 ORT 仓库侧修复是改 WebGPU 的ConvTransposefp16 内核索引/坐标用int32或至少fp32计算只有权重/激活数据走 fp16输出位置映射等纯整数运算彻底脱离 fp16。这一层立刻让大输出维度结果正确。六、解决方案第二层结构性改进把“卷积类算子的索引数学必须用整型/ fp32”收口成唯一的配置对象OrtWebGpuConvTransposeFp16Policy内核选择读它from dataclasses import dataclass, field from typing import Tuple dataclass(frozenTrue) class OrtWebGpuConvTransposeFp16Policy: WebGPU ConvTranspose fp16 索引精度的单一事实来源。 # 坐标/索引运算的精度绝不用 fp16 index_math_dtype: str int32 # 受影响算子 affected_ops: Tuple[str, ...] (ConvTranspose, Conv, Pool) # fp16 精确整数上限超过必须用整型 fp16_exact_int_limit: int 2048 # 受影响 EP / 精度 affected: Tuple[str, ...] (WebGPUEexecutionProvider, fp16) def needs_int_index(self, out_dim: int) - bool: return out_dim self.fp16_exact_int_limit def describe(self) - str: return ConvTranspose 索引数学用 int32/fp32fp16 只用于数据不用坐标 POLICY OrtWebGpuConvTransposeFp16Policy() def plan_convtranspose(out_dim: int, policy: OrtWebGpuConvTransposeFp16Policy POLICY) - str: return int32_index if policy.needs_int_index(out_dim) else fp16_ok所有卷积类内核读同一份POLICY索引数学强制整型fp16 只用于数据。七、解决方案第三层断言 / CI 守护把“大输出 ConvTranspose 索引精确”做成断言。下面用 pytest 风格守护复用第四节逻辑import numpy as np def test_index_precise_upto_2048(): assert convtranspose_index_ok(2048) 0 def test_index_lost_above_2048(): assert convtranspose_index_ok(4096) 0 def test_int_index_for_large(policy): assert policy.needs_int_index(4096) is True assert policy.needs_int_index(1024) is False def test_index_dtype_not_fp16(policy): assert policy.index_math_dtype int32这四组断言锁住(1) ≤2048 索引精确(2) 2048 索引丢失证明必须用整型(3) 大输出触发整型索引(4) 索引精度非 fp16。CI 跑通即代表 ConvTranspose 索引精度被守护。八、排查清单遇到 WebGPU ConvTranspose fp16 大图结果错看输出维度2048 错、≤2048 对 - 锁定 fp16 索引精度。换 fp32 验证fp32 正常 - 确认索引精度问题。查内核索引运算坐标是不是用 fp16 算应改 int32/fp32。临时规避导出时 ConvTranspose 前后 Cast fp32或整体 fp32。根本修复内核索引用 int32/fp32fp16 只用于数据。统一策略对象用OrtWebGpuConvTransposeFp16Policy固化。CI 守护断言大输出索引精确、索引非 fp16。九、小结WebGPU ConvTranspose produces incorrect results for fp16 models when output spatial dim 2048 (index math done in f16)的根因是WebGPU 的ConvTransposefp16 内核把坐标/索引运算放在 fp16 里做而 fp16 只能精确表示到 ±2048 的整数输出空间维度超过 2048 时索引取整错误输出元素被放到错误位置结果错fp32 路径和小输出≤2048不受影响。最小修复是让 ConvTranspose 的索引/坐标运算用int32/fp32fp16 只用于权重/激活数据临时规避是 ConvTranspose 前后 Cast fp32结构性改进是用唯一的OrtWebGpuConvTransposeFp16Policy固化索引精度CI 用四组断言守护“≤2048 精确、2048 丢失、大输出用整型、索引非 fp16”。记住坐标和索引是整数运算绝不能用 fp16大图必错位。