![【Bug已解决】[Feature Request] [WebGPU EP] Add GRU support 解决方案](http://pic.xiahunao.cn/yaotu/【Bug已解决】[Feature Request] [WebGPU EP] Add GRU support 解决方案)
【Bug已解决】[Feature Request] [WebGPU EP] Add GRU support 解决方案一、现象长什么样用 ONNX Runtime 跑一个含GRU门控循环单元算子的模型指定用WebGPU EP执行时要么直接报错“GRU 不支持”要么默默把 GRU 节点回退到其他 EP如 CPU导致本该在 GPU/WebGPU 上高效跑的循环层被丢到 CPU整体性能骤降且用户无感知。现象# 现象 AWebGPU EP 直接拒绝 GRU # NotImplementedError: GRU is not supported on WebGPU EP # 现象 B不报错但 GRU 被静默回退 CPU # 图优化把 GRU 节点分配到 CPU EPWebGPU EP 只跑其余节点 # 性能比全 CPU 好不了多少但没有任何提示 # 现象 C只在含 GRU 的模型触发 # 不含 GRU 的模型 WebGPU EP 正常一旦有 GRU 节点就退化最坑的是现象 B能跑、不报错但 GRU 这个通常计算量不小的循环层被丢到 CPU跨 EP 数据搬运的开销甚至让整体比全 WebGPU 还慢用户以为“WebGPU EP 没用”却查不出原因。二、背景ONNX 的GRU算子门控循环单元在序列模型语音、简单时序里常见。ORT 的多个 EP 实现了 GRU kernelCPU EP 很早就支持CUDA EP 也支持但WebGPU EP 一直没有实现 GRU kernel于是遇到 GRU 节点时要么报错现象 A要么更常见图分区器把 GRU 分到 CPU EP现象 B。WebGPU 实现 GRU 的难点在于GRU 是带门控reset/update/hidden 三个门的循环需要在 WebGPU 的 compute shader 里实现Wxh * x Whh * h的门控逻辑、处理linear_before_reset、双向directionbidirectional的展开等。ORT 的 WebGPU EP 早期只覆盖了 Transformer/Attention 类循环类RNN/GRU/LSTM滞后。这是 EP 功能对齐审查里典型的坑WebGPU EP 的算子覆盖落后于 CPU/CUDAGRU 缺失且缺失时可能静默回退 CPU 拖慢性能。三、根因WebGPU EP 无 GRU kernel算子注册表里没有GRU遇到节点无法在 WebGPU 上执行 → 现象 A/B。静默回退 CPU图分区默认把不支持的节点分到 CPU EP没有告警用户不知道 GRU 被丢到 CPU现象 B。缺少 EP 算子覆盖对拍CI 没断言“某模型在 WebGPU EP 上 GRU 确实在 WebGPU 执行”回退长期存在。本质是WebGPU EP 缺 GRU kernel、且缺失时静默回退 CPU 拖慢性能缺覆盖对拍。四、最小可运行复现下面用 Python 模拟“GRU 在 WebGPU EP 不被支持、被静默回退 CPU”class MockEp: def __init__(self, supported): self.supported supported def can_run(self, op): return op in self.supported webgpu MockEp(supported{MatMul, Add, Transpose}) # 无 GRU cpu MockEp(supported{GRU, MatMul, Add}) def partition_buggy(graph_ops): buggy: 不支持的节点静默分给 CPU无提示。 plan {webgpu: [], cpu: []} for op in graph_ops: if webgpu.can_run(op): plan[webgpu].append(op) else: plan[cpu].append(op) # GRU 静默回退 CPU return plan def partition_fixed(graph_ops): plan {webgpu: [], cpu: []} for op in graph_ops: if webgpu.can_run(op): plan[webgpu].append(op) else: if op GRU: raise NotImplementedError(GRU not supported on WebGPU EP) plan[cpu].append(op) return plan print(buggy:, partition_buggy([MatMul, GRU, Add])) # {webgpu: [MatMul, Add], cpu: [GRU]} - GRU 静默回退 try: partition_fixed([MatMul, GRU, Add]) except NotImplementedError as e: print(fixed raises:, e) # 明确报错逼出实现 GRUbuggy静默回退fixed明确报错。五、解决方案第一层最小直接修复最小修复在 WebGPU EP 实现 GRU kernelcompute shader 里做门控 双向展开并注册到算子表同时图分区对不支持的 GRU 显式报错而非静默回退// WebGPU compute shader 里的 GRU 单步简化 compute workgroup_size(64) fn gru_step(builtin(global_invocation_id) gid: vec3u32) { let b gid.y; let t gid.x; // reset/update/hidden 三个门 let z sigmoid(Wz * x Uz * h_prev); // update gate let r sigmoid(Wr * x Ur * h_prev); // reset gate let n tanh(Wn * x r * (Un * h_prev)); // new gate h_new (1.0 - z) * n z * h_prev; // 双向正向 反向两路分别算最后拼接 }// 注册WebGPU EP 算子表加入 GRU kernel_registry_.Register(GRU, GruKernel::Create);这一层改动最小实现 GRU kernel 注册GRU 在 WebGPU 上跑通。但依赖“功能对齐维护”下看第二层。六、解决方案第二层结构性改进把“各 EP 的算子覆盖能力声明”固化成单一事实来源并强制不支持的关键算子显式报错/告警。下面这个 dataclass 集中管理from dataclasses import dataclass, field from typing import Dict, Set dataclass class WebGpuGruPolicy: 单一事实来源WebGPU EP 算子覆盖能力契约。 _supported: Dict[str, Set[str]] field(default_factorylambda: { CPU: {GRU, LSTM, RNN, MatMul}, CUDA: {GRU, LSTM, RNN, MatMul}, WebGPU: {MatMul, Add, Transpose}, # 初始缺 GRU }) def enable(self, ep: str, op: str) - None: self._supported.setdefault(ep, set()).add(op) def partition(self, ep: str, graph_ops: list, strict: bool True) - dict: plan {webgpu: [], cpu: []} for op in graph_ops: if op in self._supported.get(ep, set()): plan[ep.lower()].append(op) else: if strict and op in (GRU, LSTM, RNN): raise NotImplementedError(f{op} not supported on {ep} EP) plan[cpu].append(op) return plan def assert_covered(self, ep: str, ops: set) - None: missing ops - self._supported.get(ep, set()) if missing: raise AssertionError(f{ep} EP missing operators: {missing})用法policy WebGpuGruPolicy() policy.enable(WebGPU, GRU) # 补齐 GRU kernel 后登记 policy.assert_covered(WebGPU, {GRU, MatMul}) # 通过这一层的关键收益能力声明集中各 EP 算子覆盖在_supported缺失一目了然严格分区关键循环算子GRU/LSTM/RNN不支持时显式报错逼出实现杜绝静默回退覆盖断言assert_covered确保补齐后确实覆盖单一事实来源所有 EP 算子覆盖约定收口在WebGpuGruPolicy。七、解决方案第三层断言 / CI 守护把第二层钉成 pytest挂进 CI确保 GRU 覆盖且不再静默回退import pytest from your_package.webgpu_gru import WebGpuGruPolicy def test_gru_runs_on_webgpu_after_fix(): # 断言 1补齐 GRU 后 WebGPU 能分区到 GRU p WebGpuGruPolicy() p.enable(WebGPU, GRU) plan p.partition(WebGPU, [MatMul, GRU, Add]) assert GRU in plan[webgpu] def test_gru_missing_raises_strict(): # 断言 2未补齐且 strict 时 GRU 必须报错防静默回退 p WebGpuGruPolicy() with pytest.raises(NotImplementedError): p.partition(WebGPU, [GRU], strictTrue) def test_assert_covered_catches_gap(): # 断言 3WebGPU 缺 GRU 时覆盖断言报错 p WebGpuGruPolicy() with pytest.raises(AssertionError): p.assert_covered(WebGPU, {GRU}) def test_cpu_has_gru(): # 断言 4CPU 有 GRU对齐基准 p WebGpuGruPolicy() assert GRU in p._supported[CPU]四条断言从“补齐后跑通”“缺失报错”“覆盖断言抓缺口”“CPU 基准”四面把功能缺口钉死在 CI。八、排查清单WebGPU EP 跑含 GRU 模型性能差/报错时报GRU not supported确认 WebGPU EP 是否实现了 GRU kernel现象 A。不报错但整体比预期慢查 GRU 是否被静默回退 CPU跨 EP 搬运拖慢现象 B。是否只在含 GRU 触发查 WebGPU 算子覆盖是否落后于 CPU/CUDA。用第二层WebGpuGruPolicy能力声明集中 关键算子不支持即报错 覆盖断言。加第三层 pytest断言“补齐后跑通、缺失报错、覆盖断言抓缺口、CPU 基准”。任一 EP 缺失关键算子时应显式报错而非静默回退否则性能退化无声。九、小结WebGPU EP 缺 GRU 支持本质是WebGPU EP 的算子覆盖落后于 CPU/CUDA遇到 GRU 节点要么报错、要么静默回退 CPU导致循环层在 CPU 上跑、跨 EP 搬运拖慢整体且无提示且缺覆盖对拍。修复分三层——第一层在 WebGPU 实现 GRU compute shader kernel 并注册第二层用WebGpuGruPolicy这个 dataclass 把各 EP 算子覆盖收口成单一事实来源关键循环算子不支持即显式报错第三层用四条 pytest 把“补齐后跑通、缺失报错、覆盖断言抓缺口、CPU 基准”钉死在 CI。核心心法任一 EP 缺失关键算子时必须显式报错而非静默回退 CPU否则性能退化无声算子覆盖能力应集中声明并断言对齐。