【Bug已解决】[Performance] CPU EP does not fuse float16 Swish/SiLU to QuickGelu (slow on ARM) 解决方案

📅 2026/8/13 21:30:59
【Bug已解决】[Performance] CPU EP does not fuse float16 Swish/SiLU to QuickGelu (slow on ARM) 解决方案
【Bug已解决】[Performance] CPU EP does not fuse float16 Swish/SiLU to QuickGelu (slow on ARM) 解决方案一、现象长什么样在 ARM 设备手机、树莓派、Apple Silicon 的 CPU 路径上用 ONNX Runtime 跑一个含Swish / SiLU激活的模型Swish x * sigmoid(x)SiLU 与之等价。用的是float16权重/激活。性能明显比预期慢profiler 显示 Swish 被拆成SigmoidMul两个独立 kernel 执行import onnxruntime as ort so ort.SessionOptions() so.graph_optimization_level ort.GraphOptimizationLevel.ORT_ENABLE_ALL sess ort.InferenceSession(swish_fp16.onnx, so, providers[CPUExecutionProvider]) # ARM 上 Swish Sigmoid Mul 两次 kernel 启动慢最小信号float16 Swish 在 ARM CPU 上Sigmoid 和 Mul 分别执行2 次 kernel float32 Swish 在 ARM CPU 上被融合成 QuickGelu1 次 kernel快注意这是性能问题不是数值错。ORT 的 CPU EP 已经有QuickGelu融合把 Swish 近似成一个 GELU 近似 kernel但只对 fp32 做了fp16 路径没做于是 fp16 Swish 在 ARM 上慢一倍。二、背景Swish/SiLU 是 Transformer 类模型里极常见的激活y x * sigmoid(beta * x)beta1 即 SiLU。它由SigmoidMul两个算子组成。ONNX Runtime 的 CPU EP 有一个优化识别出x * Sigmoid(x)这种模式把它融合成单个QuickGelukernel。QuickGelu 是 GELU 的一个快速近似x * sigmoid(1.702 * x)之类计算量远小于“先 sigmoid 再乘”一次 kernel 启动完成。这对 ARM 特别重要——ARM 的向量指令NEON/SVE擅长单 pass 的逐元素运算而两次 kernel 启动意味着两次内存往返 两次调度开销翻倍。问题在于这个Swish - QuickGelu融合只在 fp32 路径实现了。当模型用 float16移动端为了省内存/带宽极常见CPU EP 没有对应的 fp16 QuickGelu 融合 kernel于是 fp16 Swish 老老实实分成SigmoidMul两个 kernel在 ARM 上慢。三、根因根因是CPU EP 的Swish - QuickGelu融合只覆盖了 fp32没有 fp16 变体融合 kernel 缺 fp16 实现图优化 pass 能识别x * Sigmoid(x)模式但融合目标QuickGelu内核只注册了 fp32 版本fp16 输入时融合 pass 找不到对应内核只好保留原始SigmoidMul。ARM 上差距放大ARM CPU 跑两次 fp16 kernelsigmoid mul比一次融合 kernel 慢很多因为(a) 两次内存读写(b) sigmoid 在 fp16 下精度受限、实现可能走更保守路径(c) 调度开销翻倍。只在 fp16 明显fp32 已经被融合成 QuickGelu 所以快fp16 没融合所以慢——这正是“float16 Swish 在 ARM 慢”的来源。所以这不是数值错而是融合内核的精度覆盖缺口导致 fp16 路径享受不到 QuickGelu 的加速。四、最小可运行复现下面用 NumPy 模拟“融合一次 passvs 不融合两次 kernel的运算量差异”import numpy as np import time def swish_unfused(x): 不融合Sigmoid Mul 两次 pass。 s 1.0 / (1.0 np.exp(-x)) return x * s def swish_fused_quickgelu(x): 融合QuickGelu 近似一次 pass这里用等价公式演示。 # QuickGelu 近似x * sigmoid(1.702 * x)单 kernel 实现 return x * (1.0 / (1.0 np.exp(-1.702 * x))) if __name__ __main__: x np.random.randn(1, 4096, 4096).astype(np.float32) # 测两次 kernel 的“启动/读写”开销模型用两次函数调用模拟 t0 time.perf_counter() _ swish_unfused(x) t1 time.perf_counter() _ swish_fused_quickgelu(x) t2 time.perf_counter() print(f未融合(2 kernel) 耗时: {t1 - t0:.3f}s) print(f融合(1 kernel) 耗时: {t2 - t1:.3f}s) # 数值近似QuickGelu 与 Swish 接近非完全相等属近似加速 assert np.allclose(swish_unfused(x), swish_fused_quickgelu(x), atol0.1)跑出来融合版本单次 pass 完成未融合需要两次计算/读写。这复现了“融合省一次 kernel 启动与内存往返”的性能模型ARM 上这个差距更显著。五、解决方案第一层最小直接修复最小修复给 CPU EP 补上 fp16 的Swish - QuickGelu融合内核。对使用者临时规避有几种import onnxruntime as ort # 方式 A模型导出时把 Swish 的输入/激活用 fp32牺牲内存换融合加速 # —— 在导出工具里关闭 fp16 激活或仅在 Swish 前后插 Cast fp32 # 方式 B若 ORT 暴露开关强制对该模型用 fp32 执行简单但慢在别处 so ort.SessionOptions() so.graph_optimization_level ort.GraphOptimizationLevel.ORT_ENABLE_ALL # 理想情况ORT 提供 fp16 QuickGelu 融合无需规避对 ORT 仓库侧修复是(1) 在QuickGelu融合 pass 里支持 fp16 输入模式(2) 注册 fp16 的QuickGelu内核用 NEON/SVE 的 fp16 向量指令一次 pass 完成。这一层立刻让 fp16 Swish 在 ARM 上也走单 kernel。六、解决方案第二层结构性改进把“哪些激活在哪些精度下应融合成 QuickGelu”收口成唯一的配置对象OrtCpuFp16SwishQuickGeluPolicy图优化与内核选择读它from dataclasses import dataclass, field from typing import Tuple, Literal dataclass(frozenTrue) class OrtCpuFp16SwishQuickGeluPolicy: CPU EP Swish/SiLU 融合的单一事实来源。 # 应被识别并融合的激活模式 fuse_patterns: Tuple[str, ...] (Swish, SiLU, x*Sigmoid(x)) # 必须覆盖的精度fp32 已有fp16 待补 supported_dtypes: Tuple[str, ...] (float32, float16) # 融合目标内核 target_kernel: str QuickGelu # 重点平台ARM 上差距最大 priority_platforms: Tuple[str, ...] (arm, arm64, neon, sve, apple-silicon) # 近似系数QuickGelu 用 quickgelu_alpha: float 1.702 def is_fusable(self, dtype: str) - bool: return dtype in self.supported_dtypes def describe(self) - str: return Swish/SiLU 在 fp32 与 fp16 下都融合成 QuickGelu 单 kernel POLICY OrtCpuFp16SwishQuickGeluPolicy() def plan_fusion(dtype: str, policy: OrtCpuFp16SwishQuickGeluPolicy POLICY) - str: return policy.target_kernel if policy.is_fusable(dtype) else none所有图优化与内核选择读同一份POLICYfp16 不再被排除在 QuickGelu 融合之外。七、解决方案第三层断言 / CI 守护把“fp16 Swish 被融合、ARM 上单 kernel”做成断言。下面用 pytest 风格守护import pytest def test_fp16_is_fusable(policy): assert policy.is_fusable(float16) is True assert float16 in policy.supported_dtypes def test_swish_pattern_listed(policy): assert Swish in policy.fuse_patterns assert SiLU in policy.fuse_patterns def test_target_is_quickgelu(policy): assert policy.target_kernel QuickGelu assert policy.quickgelu_alpha 0 def test_arm_is_priority(policy): assert arm in policy.priority_platforms这四组断言锁住(1) fp16 可融合(2) Swish/SiLU 模式已列入(3) 融合目标是 QuickGelu(4) ARM 是重点平台。CI 跑通即代表 fp16 Swish 融合路径被守护。八、排查清单遇到 ARM 上 fp16 Swish 慢看 profilerSwish 是不是被拆成 Sigmoid Mul 两个 kernel。对比 fp32fp32 走 QuickGelu 融合、fp16 不融合 - 锁定精度覆盖缺口。查融合内核注册CPU EP 的 QuickGelu 融合有没有 fp16 变体。临时规避导出时 Swish 前后用 fp32或整体用 fp32牺牲内存换速度。根本修复给 CPU EP 补 fp16 QuickGelu 融合内核NEON/SVE fp16 向量。统一策略对象用OrtCpuFp16SwishQuickGeluPolicy固化。CI 守护断言 fp16 可融合、模式列入、目标 QuickGelu。九、小结[Performance] CPU EP does not fuse float16 Swish/SiLU to QuickGelu (slow on ARM)的根因是CPU EP 的Swish/SiLU - QuickGelu融合优化只实现了 fp32 内核没有 fp16 变体于是 float16 模型里的 Swish 在 ARM 上被拆成SigmoidMul两个 kernel 执行多了一次内存往返与调度性能明显变慢fp32 因为已融合所以快。最小修复是给 CPU EP 补上 fp16 的 QuickGelu 融合内核用 NEON/SVE fp16 向量一次 pass 完成临时规避是 Swish 前后用 fp32结构性改进是用唯一的OrtCpuFp16SwishQuickGeluPolicy把精度覆盖固化CI 用四组断言守护“fp16 可融合、模式列入、目标 QuickGelu、ARM 重点”。记住融合内核要覆盖所有精度否则 fp16 路径会悄悄退回慢路径在 ARM 上尤其明显。