ONNX Runtime Web 模型部署
PyTorch → ONNX 导出、模型量化、WebGPU 加速推理、性能 Benchmark
ONNX Runtime Web 概述
ONNX Runtime Web 是 ONNX Runtime 的 WebAssembly 构建版本,专为浏览器端模型推理设计。它允许开发者直接在浏览器中加载和运行 ONNX 格式的深度学习模型,无需依赖后端服务器,从而降低推理延迟并保护用户数据隐私。
ONNX Runtime Web 支持三种计算后端:
| 后端 | 加速技术 | 浏览器兼容性 | 适用场景 |
|---|---|---|---|
| WASM | CPU 通用计算 | 所有现代浏览器 | 兼容性优先,小模型推理 |
| WebGL | GPU 图形管线 | Chrome / Firefox / Safari | 图像类模型推理 |
| WebGPU | GPU 通用计算 | Chrome / Edge / Firefox Nightly | 高性能推理,大模型加速 |
WebGL 后端利用 GPU 的图形渲染管线进行张量计算,兼容性较广但受限于浮点精度和运算符支持。WASM 后端基于 CPU 计算,虽然速度不及 GPU,但兼容性最好,适合移动端和低端设备。WebGPU 作为下一代 Web 图形标准,提供对 GPU 计算单元的底层访问能力,在推理性能和精度上显著优于 WebGL,是浏览器端模型部署的未来方向。
PyTorch → ONNX 导出流程
将 PyTorch 模型导出为 ONNX 格式是在 Web 端部署的第一步。PyTorch 提供了 torch.onnx.export() 函数,通过追踪(tracing)或脚本化(scripting)方式将模型的计算图序列化为 ONNX 协议格式。
torch.onnx.export() 参数详解
torch.onnx.export(
model, # PyTorch 模型(需置于 eval 模式)
args, # 示例输入张量
f, # 输出文件路径(.onnx)
input_names=None, # 输入节点名称列表
output_names=None, # 输出节点名称列表
dynamic_axes=None, # 动态轴配置字典
opset_version=14, # ONNX opset 版本号
do_constant_folding=True, # 是否执行常量折叠优化
export_params=True, # 是否导出训练参数
verbose=False # 是否打印导出日志
)- input_names / output_names:为模型的输入输出张量指定名称,后续在推理时通过这些名称来绑定数据。建议命名为
"input"、"output"等语义化标识。 - dynamic_axes:指定哪些维度是动态的(即可变的),例如批处理大小和序列长度。这是 Web 端部署的关键配置,因为浏览器端的输入批次大小通常不固定。
- opset_version:ONNX 算子集的版本号。较高的版本支持更多算子,但需要考虑 ONNX Runtime Web 的支持范围。推荐使用 opset 14 或 15,兼容性和算子覆盖率较为均衡。
- do_constant_folding:开启后会在导出时预先计算常量子图,减小模型体积并提升推理速度。
动态轴配置
动态轴允许模型在推理时接受不同形状的输入。对于 NLP 模型(如 BERT),序列长度通常是可变的;对于视觉模型,则常需要支持可变批大小。
dynamic_axes = {
"input": {0: "batch_size", 1: "sequence_length"},
"attention_mask": {0: "batch_size", 1: "sequence_length"},
"output": {0: "batch_size"}
}上述配置表示 input 和 attention_mask 张量的第 0 维(批大小)和第 1 维(序列长度)是动态的,output 张量的第 0 维也是动态的。ONNX Runtime 在推理时会根据实际输入自动推断这些维度。
代码示例:导出 ResNet-50 到 ONNX
以下示例演示如何将预训练的 ResNet-50 模型导出为 ONNX 格式:
import torch
import torchvision.models as models
# 加载预训练模型并切换到评估模式
model = models.resnet50(pretrained=True)
model.eval()
# 创建示例输入(批大小为 1,3 通道,224x224)
dummy_input = torch.randn(1, 3, 224, 224)
# 配置动态轴——允许可变批大小
dynamic_axes = {
"input": {0: "batch_size"},
"output": {0: "batch_size"}
}
# 导出 ONNX 模型
torch.onnx.export(
model,
dummy_input,
"resnet50.onnx",
input_names=["input"],
output_names=["output"],
dynamic_axes=dynamic_axes,
opset_version=14,
do_constant_folding=True,
verbose=False
)
print("ONNX 模型导出成功:resnet50.onnx")导出后的 ONNX 模型文件可以直接在 ONNX Runtime Web 中加载和推理。
模型量化
模型量化通过降低权重和激活值的数值精度来减小模型体积并加速推理。ONNX Runtime 支持多种量化方式,其中动态量化和静态量化是最常用的两种。
量化原理
深度学习模型默认以 FP32(32 位浮点数)存储参数。量化将其压缩为 FP16(16 位浮点)或 INT8(8 位整型),从而将模型体积减少 50% 至 75%。推理时,量化模型在保持接近原始精度的前提下,显著降低内存带宽占用并加速矩阵运算。
| 精度类型 | 存储大小(相对 FP32) | 推理速度提升 | 精度损失 |
|---|---|---|---|
| FP32 | 1x(基准) | 1x(基准) | 无 |
| FP16 | 0.5x | 1.5–2x | 极小 |
| INT8 | 0.25x | 2–4x | 较小 |
ONNX Runtime 量化工具
ONNX Runtime 提供了 quantize_static 和 quantize_dynamic 两个核心量化函数:
- quantize_dynamic:仅对权重进行量化,无需校准数据,使用简单,适合快速部署。
- quantize_static:同时对权重和激活值进行量化,需要少量校准数据,精度优于动态量化。
代码示例:对导出的 ONNX 模型进行动态量化
from onnxruntime.quantization import quantize_dynamic, QuantType
# 输入:之前导出的 FP32 模型
input_model_path = "resnet50.onnx"
output_model_path = "resnet50_int8.onnx"
# 执行动态量化,权重降为 INT8
quantize_dynamic(
model_input=input_model_path,
model_output=output_model_path,
weight_type=QuantType.QInt8 # 可选 QUInt8 / QInt8
)
import os
fp32_size = os.path.getsize(input_model_path) / 1024 / 1024
int8_size = os.path.getsize(output_model_path) / 1024 / 1024
print(f"FP32 模型大小:{fp32_size:.2f} MB")
print(f"INT8 量化后大小:{int8_size:.2f} MB")
print(f"压缩比:{fp32_size / int8_size:.2f}x")以 ResNet-50 为例,FP32 模型约 98 MB,INT8 动态量化后约为 25 MB,体积缩减约 74%,Top-1 准确率通常下降不超过 1%。
WebGPU 加速推理
WebGPU 是 W3C 发布的新一代 GPU API,相比 WebGL 提供了对 GPU 计算单元的更底层、更灵活的控制能力。
WebGPU 的优势 vs WebGL
| 对比维度 | WebGL | WebGPU |
|---|---|---|
| 计算能力 | 仅支持片段着色器间接计算 | 原生 Compute Shader 支持 |
| 精度支持 | 多数实现仅 FP16 | FP32 / FP16 完整支持 |
| 运算符覆盖 | 有限,复杂算子需降级到 CPU | 覆盖 ONNX 核心算子集 |
| 内存管理 | 需手动管理纹理内存 | 自动管理,支持 Buffer 直接操作 |
| 多线程 | 不原生支持 | 配合 Web Worker + 共享内存 |
对于 Transformer、BERT 等包含大量矩阵乘法和注意力机制的模型,WebGPU 的性能优势尤为明显。
配置 ONNX Runtime Web 使用 WebGPU
// 配置 ONNX Runtime Web 使用 WebGPU 后端
ort.env.wasm.numThreads = 4; // WASM 线程数(用于 CPU 回退)
ort.env.wasm.simd = true; // 启用 SIMD 加速
// 创建 WebGPU 推理会话
const session = await ort.InferenceSession.create("resnet50_int8.onnx", {
executionProviders: ["webgpu"] // 优先使用 WebGPU
});executionProviders 可以设置为数组,ONNX Runtime 会按优先级依次尝试。例如 ["webgpu", "webgl", "wasm"] 表示优先使用 WebGPU,不可用时回退到 WebGL 或 WASM。
代码示例:在浏览器中加载模型进行推理
以下是一个完整的 HTML 页面,展示如何在浏览器中使用 ONNX Runtime Web 加载 ResNet-50 模型并执行推理:
<!DOCTYPE html>
<html>
<head>
<title>ONNX Runtime Web 推理示例</title>
<script src="https://cdn.jsdelivr.net/npm/onnxruntime-web/dist/ort.min.js"></script>
</head>
<body>
<h1>ONNX Runtime Web 推理</h1>
<p id="status">正在加载模型...</p>
<script>
async function runInference() {
try {
// 配置 ONNX Runtime 环境
ort.env.wasm.numThreads = 4;
ort.env.wasm.simd = true;
// 创建推理会话,优先使用 WebGPU
const session = await ort.InferenceSession.create(
"https://example.com/models/resnet50_int8.onnx",
{ executionProviders: ["webgpu", "webgl", "wasm"] }
);
document.getElementById("status").textContent =
`模型已加载,使用后端:${session.providers[0]}`;
// 创建输入张量(批大小为 1,3x224x224)
const inputTensor = new ort.Tensor(
"float32",
new Float32Array(1 * 3 * 224 * 224),
[1, 3, 224, 224]
);
// 执行推理
const feeds = { "input": inputTensor };
const results = await session.run(feeds);
const output = results.output.data;
// 获取 Top-1 预测结果
const maxIndex = output.indexOf(Math.max(...output));
document.getElementById("status").textContent =
`推理完成,预测类别索引:${maxIndex},置信度:${output[maxIndex].toFixed(4)}`;
} catch (error) {
document.getElementById("status").textContent =
`推理失败:${error.message}`;
}
}
runInference();
</script>
</body>
</html>上述代码首先加载 ONNX Runtime Web 的 CDN 库,然后创建推理会话并指定优先使用 WebGPU 后端。创建输入张量后调用 session.run() 执行前向推理,最后从输出张量中获取预测结果。
性能 Benchmark
以下基准测试数据展示了不同模型在 WASM、WebGL 和 WebGPU 三种后端下的推理延迟对比。测试环境为 Chrome 120 + Windows 11,GPU 为 NVIDIA RTX 3060。
| 模型 | 输入尺寸 | WASM (ms) | WebGL (ms) | WebGPU (ms) | WebGPU 加速比 |
|---|---|---|---|---|---|
| ResNet-50 | 1x3x224x224 | 245.3 | 89.7 | 28.4 | 8.6x |
| MobileNet-v2 | 1x3x224x224 | 68.2 | 28.1 | 12.5 | 5.5x |
| BERT-Base | 1x128 | 486.7 | 215.4 | 62.8 | 7.7x |
| TinyBERT | 1x128 | 142.5 | 67.3 | 22.1 | 6.4x |
从数据可以看出:
- WebGPU 相对于 WASM 的加速比在 5.5x 到 8.6x 之间,对于计算密集型模型(如 ResNet-50 和 BERT-Base)加速效果最为显著。
- WebGL 的性能介于 WASM 和 WebGPU 之间,对于图像模型表现尚可,但在处理 Transformer 类模型时受限于算子支持度,部分计算需回退到 CPU,导致延迟偏高。
- 量化后的 INT8 模型在 WebGPU 后端上可进一步获得 15%–30% 的额外加速,同时模型加载时间也大幅缩短。
综合来看,在生产环境中推荐采用 INT8 量化 + WebGPU 优先、WebGL 回退 的部署策略,在兼容性和性能之间取得最佳平衡。