onnx2torch
一行代码将 ONNX 模型转换为可推理、可微调的 PyTorch fx.GraphModule
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
一行代码将 ONNX 模型转换为可推理、可微调的 PyTorch fx.GraphModule
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
图1:onnx2torch 项目 Logo
2021年,随着深度学习模型部署场景日益复杂,AI 工程师们面临一个普遍困境:手中的模型来自不同框架——TensorFlow、PyTorch、ONNX——但部署环境往往只支持其中某一种。ONNX(Open Neural Network Exchange)作为通用交换格式,本应解决这个问题,但现实是:能导出 ONNX,不等于能无缝加载到 PyTorch 里进行推理或微调。
ENOT-AutoDL 团队在内部模型部署实践中频繁遇到这个问题,于是干脆动手写了一个转换器,这就是 onnx2torch 的由来。该项目于 2021 年 12 月开源,目前已在 GitHub 收获超过 700 颗星、92 个 Fork,成为 ONNX→PyTorch 转换领域最活跃的工具之一。
onnx2torch 的设计哲学是极简至上。用户只需要调用一个 convert() 函数,就可以把 ONNX 模型文件或 ModelProto 对象转换为 PyTorch 的 fx.GraphModule:
from onnx2torch import convert
import onnx
# 方式一:直接传文件路径
torch_model = convert("/path/to/mobile_net_v2.onnx")
# 方式二:先加载 ONNX,再转换
onnx_model = onnx.load("/path/to/mobile_net_v2.onnx")
torch_model = convert(onnx_model)
转换完成后,返回的 PyTorch 模型可以直接用 torch.forward() 推理,也可以用 torch.onnx.export() 再导出回 ONNX——这在需要对比验证转换精度时非常有用。
更重要的是,转换结果与原始 ONNX 模型的数值误差可以控制在 1e-7 量级,满足绝大多数生产级精度需求。
onnx2torch 的 README 列出了一份相当全面的已测试模型清单,覆盖了业界主流的计算机视觉模型:
分类网络(来自 TorchVision):ResNet-18/50、MobileNetV2/V3、EfficientNet-B0/B1/B2/B3、WideResNet-50、ResNeXt-50、VGG-16、GoogLeNet、MnasNet、RegNet,基本涵盖了最常用的 backbone。
分割网络:DeepLabV3+、DeepLabV3(ResNet-50 backbone)、HRNet、UNet、FCN ResNet-50、LRASPP MobileNetV3,在医学影像和自动驾驶场景中常用的模型都有覆盖。
检测网络(来自 MMdetection 和 Ultralytics):SSDLite MobileNetV2、RetinaNet R50、SSD300 VGG、YOLOv3、YOLOv5,涵盖了从轻量到重型的不同检测需求。
Transformer 模型:ViT、Swin Transformer、GPT-J,表明 onnx2torch 对非卷积架构也有良好支持。
操作符层面,当前版本(v1.5.15)测试了 opset 9 到 16,推荐使用 opset 13。项目维护了一份详细的 operators.md,列出了每个操作符的支持状态和限制条件。
onnx2torch 的架构设计了一个精巧的装饰器注册机制。开发者如果遇到暂不支持的操作符,只需用 @add_converter 装饰器注册一个新的转换函数即可:
from onnx2torch import add_converter, OnnxNode, OnnxGraph
from onnx2torch.utils.common import OperationConverterResult
@add_converter(operation_type="Relu", version=6)
@add_converter(operation_type="Relu", version=13)
@add_converter(operation_type="Relu", version=14)
def _(node: OnnxNode, graph: OnnxGraph) -> OperationConverterResult:
return OperationConverterResult(
torch_module=nn.ReLU(),
onnx_mapping=onnx_mapping_from_node(node=node),
)
每个转换器接收 ONNX 节点和计算图,返回一个包含 PyTorch 模块和 ONNX 映射信息的 OperationConverterResult。这种模式让扩展成本降到最低,社区可以不断丰富支持的操作符集。
从代码层面看,onnx2torch 的核心在 converter.py 中的 convert() 函数,它完成以下关键步骤:
safe_shape_inference 对 ONNX 模型进行预处理,解决动态 shape 问题_remove_initializers_from_input() 将常量权重从输入列表移到 InitializersContainer 中管理get_converter() 从注册表查找对应操作符的转换函数,生成 PyTorch 子模块fx.GraphModule整个架构借助了 PyTorch FX(Functional eXchange)的动态图能力,而非简单的权重拷贝。这使得转换后的模型天然继承了 PyTorch 的 JIT 优化、量化、剪枝等生态工具链。
项目代码组织清晰:
onnx2torch/converter.py:核心转换引擎onnx2torch/node_converters/:73 个操作符转换器(按功能分类:激活函数、卷积、池化、矩阵运算等)onnx2torch/onnx_graph.py:ONNX 计算图的 Python 封装onnx2torch/utils/:公共工具(数据类型转换、padding 处理、shape 推导等)tests/node_converters/:每个转换器的单元测试onnx2torch 在工程规范化方面做得相当扎实:
代码质量评分公司内部评分:85/100。架构设计合理,模块边界清晰,注册机制优雅,但文档中缺少架构设计文档,且部分复杂操作符(如 ROIAlign)的转换注释较少,对新加入的贡献者有一定门槛。
onnx2torch 并非万能,以下几点需要特别注意:
不支持的操作符:Acosh、Asinh、Atanh、ArgMax、ArgMin 等尚不支持,转换前需要确认模型中所有操作符均在支持列表中。
opset 版本限制:推荐使用 opset 13,最低支持 opset 9,过旧的 ONNX 模型可能需要先用工具升级 opset 版本。
动态 shape 有限支持:对于高度动态的场景(如序列长度变化极大的 NLP 模型),可能需要额外的适配工作。
纯 CPU 转换:转换过程本身在 CPU 上运行,大模型转换可能较慢,但转换后模型的推理可在 GPU 上正常进行。
onnx2torch 的价值在于打通了 ONNX 生态与 PyTorch 生态之间的最后一公里。在实际应用中,这个工具解决了几个关键痛点:
训练-部署解耦:研究员在 PyTorch 中训练模型,通过 onnx2torch 转换后可直接部署到 ONNX Runtime、TFLite 等平台,或继续在 PyTorch 生态内做量化蒸馏。
跨框架模型复用:许多开源预训练模型以 ONNX 格式发布,onnx2torch 让这些模型可以无缝进入 PyTorch 工作流,无需重写网络结构代码。
模型溯源与调试:当需要对比 ONNX Runtime 推理结果和 PyTorch 推理结果时,onnx2torch 提供了逐算子级别的对齐能力。
从项目趋势看,2021 年至今持续活跃更新(最新版本 v1.5.15,2024年8月),社区 Issue 反馈积极,随着 Transformer 模型在 CV/NLP 领域的全面渗透,对 ViT、Swin 等架构的支持将成为 onnx2torch 下一阶段的重要方向。