You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何创建匹配指定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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.07 03:11:22