如何创建匹配指定ONNX模型I/O信息的Dummy ONNX模型?
基于PyTorch生成Dummy ONNX模型的实现代码
首先安装依赖:
pip install torch onnx numpy
以下是完整代码,会自动读取原ONNX模型的输入输出元数据,生成完全匹配的Dummy模型:
import torch import onnx import numpy as np # 替换为你的原模型路径和目标Dummy模型输出路径 ORIGINAL_MODEL_PATH = "path/to/your/original_model.onnx" OUTPUT_DUMMY_MODEL_PATH = "dummy_model.onnx" # 1. 解析原ONNX模型的输入输出信息 original_model = onnx.load(ORIGINAL_MODEL_PATH) # 提取输入元数据:名称、形状、PyTorch数据类型 input_metadata = [] for in_tensor in original_model.graph.input: # 处理维度:原模型中0表示动态维度,转为-1 shape = [dim.dim_value if dim.dim_value != 0 else -1 for dim in in_tensor.type.tensor_type.shape.dim] # 转换ONNX dtype到PyTorch dtype np_dtype = onnx.helper.tensor_dtype_to_np_dtype(in_tensor.type.tensor_type.elem_type) torch_dtype = torch.from_numpy(np.array([0], dtype=np_dtype)).dtype input_metadata.append({ "name": in_tensor.name, "shape": shape, "dtype": torch_dtype }) # 提取输出元数据 output_metadata = [] for out_tensor in original_model.graph.output: shape = [dim.dim_value if dim.dim_value != 0 else -1 for dim in out_tensor.type.tensor_type.shape.dim] np_dtype = onnx.helper.tensor_dtype_to_np_dtype(out_tensor.type.tensor_type.elem_type) torch_dtype = torch.from_numpy(np.array([0], dtype=np_dtype)).dtype output_metadata.append({ "name": out_tensor.name, "shape": shape, "dtype": torch_dtype }) # 2. 定义Dummy模型:输出恒定值(示例用全0,可自行修改为其他固定值) class DummyModel(torch.nn.Module): def __init__(self, output_metadata): super().__init__() self.constant_outputs = [] for meta in output_metadata: # 处理动态维度:暂时固定为1,导出时会保留动态属性 fixed_shape = [1 if dim == -1 else dim for dim in meta["shape"]] constant_tensor = torch.zeros(fixed_shape, dtype=meta["dtype"]) self.constant_outputs.append(constant_tensor) def forward(self, *inputs): # 忽略输入,直接返回预定义的恒定输出 return tuple(self.constant_outputs) # 3. 初始化模型并导出为ONNX dummy_model = DummyModel(output_metadata) # 构造示例输入用于导出追踪 example_inputs = [] for meta in input_metadata: fixed_shape = [1 if dim == -1 else dim for dim in meta["shape"]] example_input = torch.randn(fixed_shape, dtype=meta["dtype"]) example_inputs.append(example_input) # 导出时严格匹配原模型的参数 torch.onnx.export( dummy_model, tuple(example_inputs), OUTPUT_DUMMY_MODEL_PATH, input_names=[meta["name"] for meta in input_metadata], output_names=[meta["name"] for meta in output_metadata], # 声明动态维度,和原模型保持一致 dynamic_axes={ meta["name"]: {i: f"dim_{i}" for i, dim in enumerate(meta["shape"]) if dim == -1} for meta in input_metadata }, # 必须匹配原模型的OPSET版本 opset_version=original_model.opset_import[0].version, do_constant_folding=True ) print(f"Dummy模型已生成:{OUTPUT_DUMMY_MODEL_PATH}")
额外需要注意的事项
- OPSET版本严格匹配:不同OPSET版本的算子语法和兼容性差异极大,必须使用原模型的OPSET版本导出,可通过
original_model.opset_import[0].version获取版本号。 - 动态维度声明:如果原模型存在动态维度(如batch_size设为-1),导出时必须在
dynamic_axes中明确标记这些维度,否则流水线可能因维度不兼容报错。 - 数据类型完全对齐:输入输出的数据类型(如float32、int64)必须和原模型完全一致,部分流水线会严格校验数据类型,不允许隐式转换。
- 输入输出顺序一致性:即使名称匹配,输入输出的顺序也必须和原模型完全一致,否则流水线会出现参数不匹配的问题。
- 元数据匹配(可选):如果流水线依赖ONNX模型的元数据(如作者、版本、自定义属性),需要手动给Dummy模型添加这些元数据,可通过
original_model.graph.metadata_props读取原数据后添加。 - 模型验证:生成Dummy模型后,建议用ONNX Runtime分别加载原模型和Dummy模型,输入相同测试数据,对比输出的名称、形状、类型和值,确保完全匹配。
- 算子兼容性:避免使用原流水线未支持的算子,虽然Dummy模型仅输出恒定值,但导出时PyTorch可能生成某些特定算子,需确保这些算子在目标流水线中可用。
内容的提问来源于stack exchange,提问作者aroyc
相关产品推荐
相关产品推荐

