【Bug已解决】[Performance] [WebGPU] Proposal: Add an option to use online softmax instead of naive softm…
【Bug已解决】[Performance] [WebGPU] Proposal: Add an option to use online softmax instead of naive softmax 解决方案一、现象长什么样在 WebGPU EP 上跑注意力尤其是长序列如长上下文 LLM、长音频发现 softmax 用的是朴素实现naive softmax先算整行最大值做减法防止溢出再 exp、再求和归一。这套实现需要两遍遍历分数矩阵一遍求 max/sum一遍归一在 WebGPU 上意味着两次独立的 compute dispatch 两次全局内存读写长序列下成为性能瓶颈。现象# 现象 A长序列注意力比预期慢 # 朴素 softmax 两遍遍历序列越长、两次 dispatch 的开销越明显 # 注意力成为整图热点 # 现象 B显存带宽压力大 # 两遍都要把整个 score 矩阵读回/写回带宽翻倍WebGPU 上尤其明显 # 现象 C只有开选项/长序列才暴露 # 短序列两遍开销可忽略长序列几 K 以上性能落差大最坑的是现象 A用户只看到“WebGPU 注意力不如预期快”想不到是 softmax 实现方式的问题——朴素 softmax 数值对、结果对只是慢且慢在“多一遍遍历”属于纯性能问题而非正确性问题。二、背景Softmax 的**在线算法online softmax也叫 flash softmax**由 Tri Dao 等提出维护 running max 和 running sum一遍遍历就能同时完成 max/scale/exp/normalize避免朴素实现的两遍。这对注意力尤其关键——score 矩阵可能很大seq×seq两遍意味着两倍的全局内存流量。WebGPU EP 的注意力 kernel 早期用朴素 softmax实现简单、数值直观没有在线版本。提案要求加一个选项让用户选择 online softmax尤其是长序列场景。难点在于 WebGPU 的 workgroup 共享内存workgroup memory要用来跨线程归约 running max/sum且要保持数值与朴素 softmax 一致避免不同实现结果微差导致对拍失败。这是性能优化审查里典型的坑用朴素两遍 softmax 而非在线单遍长序列下算力/带宽翻倍浪费且可通过选项开放更优实现。三、根因朴素 softmax 两遍遍历attention kernel 先做m max(scores)、d sum(exp(scores-m))再out exp(scores-m)/d两次 dispatch 两次全局内存读写。无 online softmax 选项用户无法选更优的单遍实现长序列只能忍受两遍开销。缺少性能对拍回归CI 只测正确性没测“长序列注意力耗时”性能退化无感知。本质是WebGPU 注意力用朴素两遍 softmax、无 online 选项长序列性能浪费且缺性能回归。四、最小可运行复现下面用 Python 模拟“朴素 softmax两遍vs online softmax单遍”的遍历次数差异import math def naive_softmax_buggy(scores): buggy: 两遍遍历求 max再归一。 passes 0 m max(scores); passes 1 exps [math.exp(s - m) for s in scores]; passes 1 d sum(exps); passes 1 return [e / d for e in exps], passes def online_softmax_fixed(scores): fixed: 单遍在线归约running max/sum。 passes 0 m_prev, d_prev -math.inf, 0.0 out [] for s in scores: m max(m_prev, s) d d_prev * math.exp(m_prev - m) math.exp(s - m) m_prev, d_prev m, d passes 1 out [math.exp(s - m_prev) / d_prev for s in scores] return out, passes s [1.0, 2.0, 0.5, 3.0] _, p_naive naive_softmax_buggy(s) _, p_online online_softmax_fixed(s) print(naive passes:, p_naive, online passes:, p_online) # 3 vs 1naive需多遍遍历含两次全局读online单遍长序列下带宽省一倍。五、解决方案第一层最小直接修复最小修复在 attention kernel 里提供 online softmax 路径用 workgroup 共享内存做 running max/sum 的单遍归约并由选项开关// WebGPU attention: online softmax 单遍简化 varworkgroup running_max: f32; varworkgroup running_sum: f32; compute workgroup_size(64) fn attention_online(builtin(global_invocation_id) gid: vec3u32) { let q load_q(gid); var m -3.4e38; var d 0.0; // 单遍遍历每个 key维护 running max/sum for (var i 0u; i seq; i) { let s dot(q, load_k(i)) * scale; let m_new max(m, s); d d * exp(m - m_new) exp(s - m_new); // 在线修正 m m_new; acc exp(s - m) * load_v(i); } out[gid] acc / d; // 一遍完成无需第二遍 }这一层改动最小用在线归约替代两遍长序列带宽省一倍。但依赖“选项与数值一致性维护”下看第二层。六、解决方案第二层结构性改进把“softmax 实现策略naive/online可选项 数值一致性”固化成单一事实来源。下面这个 dataclass 集中管理from dataclasses import dataclass, field from typing import Dict dataclass class WebGpuOnlineSoftmaxPolicy: 单一事实来源WebGPU 注意力 softmax 策略契约。 option_enabled: bool False # 用户是否开启 online softmax _supported_modes: tuple (naive, online) def select_kernel(self, seq_len: int) - str: # 长序列或用户显式开启时优先 online if self.option_enabled and seq_len 1024: return attention_online return attention_naive def assert_numerically_equal(self, out_naive, out_online, atol1e-3) - None: # 两种实现数值必须一致避免对拍失败 if any(abs(a - b) atol for a, b in zip(out_naive, out_online)): raise AssertionError(online softmax disagrees with naive)这一层的关键收益可选项option_enabled让用户选 online长序列自动优选数值一致assert_numerically_equal保证 online 与 naive 结果一致避免对拍失败单一事实来源所有 softmax 策略约定收口在WebGpuOnlineSoftmaxPolicy。七、解决方案第三层断言 / CI 守护把第二层钉成 pytest挂进 CI覆盖在线 softmax 的正确性与性能选择import math import pytest from your_package.webgpu_online_softmax import WebGpuOnlineSoftmaxPolicy def _softmax(xs): m max(xs); e [math.exp(s - m) for s in xs]; d sum(e) return [v / d for v in e] def test_online_matches_naive(): # 断言 1online 与 naive 数值一致 p WebGpuOnlineSoftmaxPolicy(option_enabledTrue) xs [1.0, 2.0, 0.5, 3.0] naive _softmax(xs) # 用同类在线算法算此处复用 fixed 逻辑示意 m, d -math.inf, 0.0 for s in xs: m_new max(m, s); d d * math.exp(m - m_new) math.exp(s - m_new); m m_new online [math.exp(s - m) / d for s in xs] p.assert_numerically_equal(naive, online) def test_long_seq_picks_online(): # 断言 2开启且长序列时选 online kernel p WebGpuOnlineSoftmaxPolicy(option_enabledTrue) assert p.select_kernel(4096) attention_online def test_short_seq_uses_naive(): # 断言 3短序列仍用 naive避免 workgroup 开销不划算 p WebGpuOnlineSoftmaxPolicy(option_enabledTrue) assert p.select_kernel(64) attention_naive def test_option_off_uses_naive(): # 断言 4未开启选项一律 naive p WebGpuOnlineSoftmaxPolicy(option_enabledFalse) assert p.select_kernel(8192) attention_naive四条断言从“数值一致”“长序列选 online”“短序列用 naive”“未开启用 naive”四面把策略回归钉死在 CI。八、排查清单WebGPU 注意力长序列性能差时长序列是否比预期慢很多查 softmax 是否朴素两遍遍历现象 A。是否提供 online softmax 选项没有就确认 attention kernel 只有 naive 路径。切换 online 后数值是否和 naive 一致不一致会导致对拍失败必须保证一致。用第二层WebGpuOnlineSoftmaxPolicy可选项 数值一致校验 长短序列自动选。加第三层 pytest断言“数值一致、长序列选 online、短序列用 naive、未开启用 naive”。在线 softmax 用 workgroup 共享内存归约 running max/sum长序列带宽省一倍但短序列 overhead 不划算应按序列长选择。九、小结WebGPU 注意力用朴素 softmax 性能差本质是注意力 kernel 用朴素两遍 softmax求 max、再归一长序列下两倍的全局内存流量与 dispatch 开销成为瓶颈且无 online softmax 选项开放更优实现且缺性能对拍。修复分三层——第一层用 workgroup 共享内存做 online softmax 单遍归约并由选项开关第二层用WebGpuOnlineSoftmaxPolicy这个 dataclass 把“策略可选项 数值一致性 长短序列自动选择”收口成单一事实来源第三层用四条 pytest 把“数值一致、长序列选 online、短序列用 naive、未开启用 naive”钉死在 CI。核心心法长序列注意力应提供 online softmax 选项单遍归约省一倍带宽且必须保证与朴素 softmax 数值一致长短序列应按需自动选择。