浏览器端模型推理:ONNX Runtime Web 与 WebGPU 后端
浏览器端模型推理ONNX Runtime Web 与 WebGPU 后端一、端侧推理的工程诉求隐私、延迟与离线能力大模型推理在云端运行是默认方案但有三类场景迫使推理下沉到浏览器端。隐私敏感场景。医疗影像分析、身份证 OCR、私有文档摘要用户不愿把数据上传服务器。端侧推理让数据不出设备从架构上消除泄露风险。延迟敏感场景。实时手势识别、表情驱动的虚拟形象、交互式图像编辑要求推理延迟低于 50ms。云端方案的网络往返通常 80-200ms端侧推理可把延迟压到 10-30ms。离线场景。弱网环境下的输入法预测、离线翻译、设备端语音唤醒必须脱离服务器运行。ONNX Runtime Web 是目前浏览器端推理的主流方案。它支持 WASM 与 WebGPU 两种后端前者兼容性广后者性能高。本文聚焦 WebGPU 后端的工程化落地。二、执行栈剖析WebGPU 后端的计算管线WebGPU 后端的完整执行栈如下从应用层到硬件层逐级下沉┌─────────────────────────────────────────────────┐ │ 应用层JS 调用 session.run() │ └──────────────────────┬──────────────────────────┘ │ ONNX tensor (TypedArray) ▼ ┌─────────────────────────────────────────────────┐ │ ORT Web 调度层 │ │ - 算子图拓扑排序 │ │ - 内存复用规划 │ │ - 算子到 backend 的分发 │ └──────────┬───────────────────────┬──────────────┘ │ WebGPU 算子 │ WASM 算子(fallback) ▼ ▼ ┌─────────────────────┐ ┌──────────────────────┐ │ WebGPU backend │ │ WASM backend │ │ - 算子编译为 WGSL │ │ - 算子编译为 wasm │ │ - GPUBuffer 传输 │ │ - SIMD/threads │ │ - compute pass 执行 │ │ - 纯 CPU 计算 │ └──────────┬──────────┘ └──────────────────────┘ ▼ ┌─────────────────────┐ │ 浏览器 WebGPU API │ │ - GPUDevice / Queue │ │ - ShaderModule │ └──────────┬──────────┘ ▼ ┌─────────────────────┐ │ GPU 硬件 (GPUAdapter)│ └─────────────────────┘WebGPU 后端的核心优势在于数据不经过 CPU。WASM 后端下每次推理都要把 tensor 数据从 GPU 显存拷到 CPU 内存计算完再拷回这个往返在中等模型上可占总耗时的 40%。WebGPU 后端把 Conv、MatMul 等算子编译为 WGSL即 WebGPU 着色器语言在 GPU 上直接执行数据全程驻留显存。后端选择策略如下表维度WASM 后端WebGPU 后端兼容性全平台2026 年覆盖率约 99%Chrome/Edge/Safari 17计算单元CPU多线程加 SIMDGPU数据拷贝每次推理 CPU 与 GPU 往返全程显存驻留中等模型延迟80-200ms15-50ms显存限制受限于内存受限于 GPU 显存初始化耗时500ms-1s2-5s含 shader 编译关键结论模型参数量小于 10M 时WASM 与 WebGPU 差距不大数据拷贝开销占比低。参数量超过 50M 时WebGPU 优势明显。参数量超过 300M 时WebGPU 受限于显存可能无法加载需用量化压缩。三、生产级推理封装会话管理与性能调优// ort-session-manager.ts // 会话管理是端侧推理的核心。 // 为什么不复用单个全局 session不同模型权重不同 // 切换时需要重新加载但频繁创建销毁 session 会 // 重复编译 shader耗时数秒。 import * as ort from onnxruntime-web/webgpu; interface SessionOptions { modelPath: string; maxConcurrency?: number; } export class ORTSessionManager { private sessions new Mapstring, ort.InferenceSession(); private loading new Mapstring, Promiseort.InferenceSession(); private locks new Mapstring, number(); private readonly maxConcurrency: number; constructor(maxConcurrency 1) { // WebGPU 不支持同一 session 的并发推理 // 多并发会导致 command buffer 乱序。 // 通过信号量串行化推理请求。 this.maxConcurrency maxConcurrency; } async getSession(key: string, opts: SessionOptions): Promiseort.InferenceSession { // 已加载直接返回 const cached this.sessions.get(key); if (cached) return cached; // 防止重复加载相同 key 的并发请求合并 if (this.loading.has(key)) { return this.loading.get(key)!; } const promise this.createSession(opts); this.loading.set(key, promise); try { const session await promise; this.sessions.set(key, session); return session; } finally { this.loading.delete(key); } } private async createSession(opts: SessionOptions): Promiseort.InferenceSession { // 超时保护模型加载可能因网络或 shader 编译卡住 const timeout new Promisenever((_, reject) setTimeout(() reject(new Error(模型加载超时)), 15_000) ); try { const session await Promise.race([ ort.InferenceSession.create(opts.modelPath, { executionProviders: [webgpu, wasm], // 优先 WebGPU回退 WASM graphOptimizationLevel: all, enableMemPattern: true, // 显存复用减少分配开销 enableCpuMemArena: false, // WebGPU 下不需要 CPU arena }), timeout, ]); return session; } catch (err) { // WebGPU 不可用时自动回退到 WASM console.warn([ORT] WebGPU 失败回退 WASM, err); return ort.InferenceSession.create(opts.modelPath, { executionProviders: [wasm], graphOptimizationLevel: all, }); } } async run( key: string, feeds: Recordstring, ort.Tensor, opts: SessionOptions ): PromiseRecordstring, ort.Tensor { const session await this.getSession(key, opts); // 信号量限制并发推理数 while ((this.locks.get(key) || 0) this.maxConcurrency) { await new Promise((r) setTimeout(r, 1)); } this.locks.set(key, (this.locks.get(key) || 0) 1); try { const results await session.run(feeds); return results; } catch (err) { // 推理失败可能是显存不足释放后由上层决定重试 console.error([ORT] 推理失败, err); throw err; } finally { this.locks.set(key, (this.locks.get(key) || 0) - 1); } } dispose(key?: string): void { if (key) { this.sessions.get(key)?.release(); this.sessions.delete(key); } else { this.sessions.forEach((s) s.release()); this.sessions.clear(); } } }// image-classify.ts // 图像分类推理含预处理、推理、后处理全链路。 // 为什么手动做预处理而非用 canvas.scale // canvas 的双线性插值质量不稳定不同浏览器结果不同 // 影响推理精度用 TypedArray 手动 resize 可保证一致性。 import * as ort from onnxruntime-web/webgpu; const IMAGENET_MEAN [0.485, 0.456, 0.406]; const IMAGENET_STD [0.229, 0.224, 0.225]; export async function classifyImage( bitmap: ImageBitmap, session: ort.InferenceSession ): Promise{ label: string; score: number }[] { // 预处理resize 到 224x224归一化 const tensor preprocessImage(bitmap); try { const feeds: Recordstring, ort.Tensor {}; const inputName session.inputNames[0]; feeds[inputName] tensor; const results await session.run(feeds); const outputName session.outputNames[0]; const output results[outputName]; // 后处理softmax 加 top-k const scores softmax(output.data as Float32Array); const topk getTopK(scores, 5); return topk.map(([idx, score]) ({ label: IMAGENET_LABELS[idx] ?? class_${idx}, score, })); } catch (err) { // 推理失败时返回空结果上层决定降级策略 console.error([classify] 推理失败, err); return []; } } function preprocessImage(bitmap: ImageBitmap): ort.Tensor { // 用 OffscreenCanvas 做 resize比 canvas 性能更好 const canvas new OffscreenCanvas(224, 224); const ctx canvas.getContext(2d)!; ctx.drawImage(bitmap, 0, 0, 224, 224); const imageData ctx.getImageData(0, 0, 224, 224); // CHW 排列加归一化ONNX 模型通常要求 CHW const { data } imageData; const tensor new Float32Array(3 * 224 * 224); for (let c 0; c 3; c) { for (let i 0; i 224 * 224; i) { const pixel data[i * 4 c]; tensor[c * 224 * 224 i] (pixel / 255 - IMAGENET_MEAN[c]) / IMAGENET_STD[c]; } } return new ort.Tensor(float32, tensor, [1, 3, 224, 224]); } function softmax(arr: Float32Array): Float32Array { let max -Infinity; for (const v of arr) if (v max) max v; let sum 0; const exp new Float32Array(arr.length); for (let i 0; i arr.length; i) { exp[i] Math.exp(arr[i] - max); sum exp[i]; } for (let i 0; i arr.length; i) exp[i] / sum; return exp; } function getTopK(arr: Float32Array, k: number): [number, number][] { return Array.from(arr) .map((v, i) [i, v] as [number, number]) .sort((a, b) b[1] - a[1]) .slice(0, k); } const IMAGENET_LABELS: string[] []; // 省略 1000 条标签四、端侧推理的边界显存、精度与兼容性WebGPU 端侧推理的边界清晰且不可逾越。显存上限。浏览器 WebGPU 的 maxStorageBufferBindingSize限制通常为 1-2GB但实际可用显存受设备 GPU 影响。移动端集显可能只有 256MB 可用。模型加载时权重需要全部驻留显存超过上限会抛出 GPUError。量化是主要手段FP32 转 INT8 可压缩 4 倍精度损失通常在 1-3%。但 INT8 在 WebGPU 下的算子支持不完整部分模型需要回退到 FP16。首次加载耗时。WebGPU 后端首次推理时需要编译所有 WGSL shader这个过程在中等模型上耗时 2-5 秒。后续推理会命中 shader 缓存。生产建议在页面加载后预热跑一次空输入推理把 shader 编译提前到用户交互之前。精度一致性。WebGPU 的浮点运算遵循 IEEE 754但不同 GPU 厂商的算子实现存在微小差异尤其是超越函数。同一模型在 NVIDIA 与 AMD GPU 上的输出可能在小数点后 4 位有差异。对分类任务无影响但对数值敏感任务如回归预测需评估容差。兼容性矩阵。截至 2026 年中WebGPU 在 Chrome 与 Edge桌面加 Android支持完整Safari 17 以上支持Firefox 仍处于实验阶段。iOS Safari 17.4 以上支持但性能受限。回退策略必须覆盖WebGPU 不可用时回退 WASM 后端再不可用时回退云端推理。功耗代价。GPU 满载推理的功耗远高于 CPU 空闲态。移动设备上持续推理会导致发热与电池快速消耗。实测 MobileViT-XS 在手机上连续推理 30 秒电池温度上升 4-6 摄氏度。对持续推理场景如实时摄像头分析建议降低推理频率至 10-15fps并监听 device.lost 事件做降级。禁用场景。对推理精度要求接近 100% 的场景医疗诊断、金融风控不应使用端侧推理因为量化与浮点差异不可控。对模型权重保密要求高的场景也不适合浏览器端模型可被用户下载权重泄露风险无法消除。五、总结ONNX Runtime Web 的 WebGPU 后端把 GPU 计算能力引入浏览器使端侧推理在中等模型上达到可用延迟。核心机制是算子编译为 WGSL、数据全程驻留显存、shader 缓存复用。落地步骤如下。第一步检测 navigator.gpu 可用性规划回退链路。第二步用 FP16 或 INT8 量化模型控制显存占用。第三步封装 SessionManager处理加载并发与超时。第四步实现预处理链路保证数值一致性。第五步页面加载后预热 session提前编译 shader。第六步监听 device.lost 与功耗信号设计降级策略。性能指标以首次推理延迟与稳态延迟为准。目标50M 参数模型预热后单次推理延迟低于 30ms移动端可接受 60ms。