【Bug已解决】WebGPU: Graph capture mode should handle INT64 transparently without per-op opt-in 解决方案

📅 2026/8/13 12:23:24
【Bug已解决】WebGPU: Graph capture mode should handle INT64 transparently without per-op opt-in 解决方案
【Bug已解决】WebGPU: Graph capture mode should handle INT64 transparently without per-op opt-in 解决方案一、现象长什么样在 ONNX Runtime 的 WebGPU EP 上开启**图捕获graph capture**模式来加速推理时只要模型里有用到INT64的张量比如位置索引、token id、Gather 的下标、位置编码里的arange就会出问题要么捕获阶段直接报错提示某个 op “不支持 INT64 捕获”要么该 op 悄悄退回非捕获路径执行整条捕获图的优势丢失延迟反而变高要么在某些 INT64 输入形状下直接产生错误结果。最小触发概念性import onnxruntime as ort so ort.SessionOptions() so.graph_optimization_level ort.GraphOptimizationLevel.ORT_ENABLE_ALL # 开启 WebGPU 图捕获 so.add_session_config_entry(gpu_graph_capture, 1) sess ort.InferenceSession(model_with_int64.onnx, so, providers[WebGpuExecutionProvider]) # 模型里一旦有 INT64 的 Gather/Slice/Where捕获就失败或回退关键点是当前实现要求每个用到 INT64 的 op 都显式“opt-in”声明自己支持捕获否则捕获整条图就会因为这一个 op 而失败或回退。这对用户极不友好——明明只是普通的索引运算却要深入每个 op 去开开关。二、背景WebGPU底层是 wgpu的存储缓冲storage buffer对 64 位整数支持很受限很多 WebGPU 后端尤其早期 Metal / D3D 实现不允许在 storage buffer 里直接读写 64 位整数也不允许 64 位整数参与某些原子/索引操作。所以 ONNX Runtime 的 WebGPU EP 在做图捕获把一整条计算图录制成一个可重放的 GPU 命令序列时遇到 INT64 必须做特殊处理。现有的“特殊处理”是逐 op 开关EP 内部维护一份“哪些 op 在捕获模式下支持 INT64”的白名单只有白名单里的 op 才走 INT64 捕获路径不在白名单的 op要么报错要么回退。这意味着白名单覆盖不全 → 用户模型里随便一个 INT64 op 就破功新增 op 要手动加白名单 → 维护成本高、容易漏用户无法干预 → 只能改模型把 INT64 改成 INT32 规避。这正是“should handle INT64 transparently without per-op opt-in”的诉求INT64 应该在捕获框架层面被透明处理而不是让每个 op 各自 opt-in。三、根因根因是图捕获对 INT64 的处理被放在了 per-op 白名单里而不是放在捕获基础设施层面逐 op 白名单捕获逻辑在调度每个 op 时检查“这个 op 有没有声明支持 INT64 捕获”没有就 fail/回退。INT64 的处置权分散到每个 op 实现覆盖不全。缺少透明的 INT64 桥接捕获基础设施没有统一的“INT64 ↔ INT32”桥接层。实际上绝大多数用 INT64 的 op索引、位置、计数其数值范围在推理时都落在 INT32 内完全可以在捕获入口把 INT64 缓冲“位级 reinterpret 成两个 i32”或在 op 边界 cast 成 INT32执行完再 cast 回来整个过程对 op 透明。回退破坏整体捕获只要有一个 op 回退非捕获路径整条捕获图的优势就没了延迟回到逐 op 提交的水平。所以这不是某个 op 算错而是INT64 的捕获策略设计错了层级——它该是基础设施的责任却被推给了每个 op。四、最小可运行复现下面用 Python NumPy 模拟“INT64 缓冲在只支持 INT32 的捕获后端下的处理”用位级拆分把 INT64 当成两个 INT32 存储执行用 INT32取回时再拼回 INT64。这复现了“透明处理 INT64”的核心思路import numpy as np def int64_to_two_i32(x: np.ndarray) - np.ndarray: 把 INT64 缓冲透明地拆成两个 INT32捕获后端的存储表示。 x x.astype(np.int64) lo (x 0xFFFFFFFF).astype(np.int32) hi (x 32).astype(np.int32) # 交错存放模拟 storage buffer 的紧凑布局 return np.stack([lo, hi], axis-1).astype(np.int32) def two_i32_to_int64(packed: np.ndarray) - np.ndarray: lo packed[..., 0].astype(np.int64) hi packed[..., 1].astype(np.int64) return (hi 32) | (lo 0xFFFFFFFF) def captured_gather(indices_packed, data): 捕获后端里用 INT32 索引做 Gather透明的 INT64 处理。 indices two_i32_to_int64(indices_packed) return data[indices] # 后端只认 INT32 缓冲但语义上是 INT64 索引 if __name__ __main__: data np.arange(1000, dtypenp.float32) idx np.array([3, 17, 999, 42], dtypenp.int64) packed int64_to_two_i32(idx) # 透明拆成 INT32 out captured_gather(packed, data) # 捕获路径用 INT32 执行 assert np.array_equal(out, data[idx]) print(INT64 索引被透明地处理为 INT32 缓冲结果一致)跑出来会打印“结果一致”——说明 INT64 索引完全可以在捕获层透明转成 INT32 处理根本不需要每个 op 单独 opt-in。五、解决方案第一层最小直接修复最小修复在图捕获基础设施里加一个统一的 INT64 桥接层而不是让每个 op 自己声明。对使用者来说临时规避方案是把模型里的 INT64 输入/中间量在导出时降级成 INT32如果值域允许# 导出时把位置索引类 INT64 改成 INT32值域在 int32 内时安全 import torch class Int32FriendlyModel(torch.nn.Module): def forward(self, input_ids): # 原本用 arange 生成 INT64 位置这里直接转 INT32 pos torch.arange(input_ids.shape[1], dtypetorch.int32, deviceinput_ids.device).unsqueeze(0) return self.backbone(input_ids, position_idspos)对 ORT 仓库侧在捕获入口统一做 INT64→(i32,i32) 的 reinterpret 或 castop 边界再 cast 回来op 自身完全无感。这一层立刻让大多数 INT64 索引场景能走捕获无需改每个 op。六、解决方案第二层结构性改进把“INT64 在捕获中如何透明处理”收口成唯一的配置对象WebGpuGraphCaptureInt64Policy所有捕获逻辑读它from dataclasses import dataclass, field from typing import Tuple, Literal from enum import Enum class Int64Strategy(Enum): BITCAST_TO_I32 bitcast_to_i32 # 位级拆成两个 i32推荐零精度损失 CAST_TO_I32 cast_to_i32 # 直接 cast值域需在 int32 内 PER_OP_OPTIN per_op_optin # 旧行为逐 op 白名单 dataclass(frozenTrue) class WebGpuGraphCaptureInt64Policy: WebGPU 图捕获中 INT64 处理的单一事实来源。 # 默认透明处理不再逐 op opt-in strategy: Int64Strategy Int64Strategy.BITCAST_TO_I32 # 允许透明处理的 op 类别索引/位置/计数类 INT64 都涵盖 transparent_op_types: Tuple[str, ...] ( Gather, Scatter, Slice, Where, NonZero, Range, TopK, Squeeze, Unsqueeze, ) # 是否对未覆盖的 op 才回退而非整图回退 fallback_per_op_not_whole_graph: bool True # 捕获失败时是否给出明确错误而非静默回退 fail_loud: bool True def handles_transparently(self, op_type: str) - bool: return self.strategy ! Int64Strategy.PER_OP_OPTIN or \ op_type in self.transparent_op_types def describe(self) - str: return fINT64 在捕获基础设施层透明处理策略{self.strategy.value} POLICY WebGpuGraphCaptureInt64Policy() def plan_capture(op_types: list, policy: WebGpuGraphCaptureInt64Policy POLICY) - dict: return {op: policy.handles_transparently(op) for op in op_types}所有捕获逻辑读同一份POLICYINT64 处置从“逐 op 白名单”升级为“基础设施透明处理”新增 op 自动覆盖。七、解决方案第三层断言 / CI 守护把“INT64 透明处理、不全图回退、结果一致”做成断言。下面用 pytest 风格守护复用第四节的位级拆分逻辑import numpy as np def test_int64_packed_roundtrip(): idx np.array([3, 17, 999, 42], dtypenp.int64) packed int64_to_two_i32(idx) assert packed.dtype np.int32 assert np.array_equal(two_i32_to_int64(packed), idx) def test_capture_handles_int64_transparently(policy): assert policy.handles_transparently(Gather) is True assert policy.handles_transparently(Slice) is True assert policy.strategy ! Int64Strategy.PER_OP_OPTIN def test_whole_graph_not_fallen_back(policy): assert policy.fallback_per_op_not_whole_graph is True def test_gather_result_matches_int64(policy): data np.arange(1000, dtypenp.float32) idx np.array([3, 17, 999], dtypenp.int64) out captured_gather(int64_to_two_i32(idx), data) assert np.array_equal(out, data[idx])这四组断言锁住(1) INT64↔i32 打包无损失(2) INT64 op 被透明处理、不再逐 op opt-in(3) 单个 op 回退而非整图回退(4) 捕获结果与原 INT64 语义一致。CI 跑通即代表 INT64 捕获透明可用。八、排查清单遇到 WebGPU 图捕获遇 INT64 失败/回退确认是否开了 graph captureadd_session_config_entry(gpu_graph_capture, 1)。找模型里的 INT64 opGather/Slice/Where/Range/TopK的索引/位置输入是不是 INT64。看错误是不是“op 不支持 INT64 捕获”是的话就是逐 op 白名单没覆盖。临时规避导出时把值域在 int32 内的 INT64 降级成 INT32。根本修复在捕获基础设施层加 INT64→i32 透明桥接而不是逐 op opt-in。统一策略对象用WebGpuGraphCaptureInt64Policy固化策略新增 op 自动覆盖。CI 守护断言 INT64 透明处理、结果一致、不全图回退。九、小结WebGPU graph capture mode should handle INT64 transparently without per-op opt-in的根因是ONNX Runtime 的 WebGPU 捕获逻辑把 INT64 的处置放在了逐 op 白名单里只有显式 opt-in 的 op 才支持 INT64 捕获导致模型里任意一个 INT64 op 就让整条捕获失败或回退。最小修复是导出时把值域允许的 INT64 降级成 INT32结构性改进是用唯一的WebGpuGraphCaptureInt64Policy在捕获基础设施层透明处理 INT64位级拆成 i32新增 op 自动覆盖CI 用四组断言守护“透明处理、不全图回退、结果一致”。记住INT64 捕获是基础设施的责任不该让每个 op 各自报名。