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

如何为无批量维度的已训练ONNX模型添加batch维度

ONNX模型新增batch维度实操方案
  • 先安装依赖包
    pip install onnx onnxruntime onnx-simplifier

  • 第一步:简化并推断模型形状
    首先对原有模型做简化和形状推断,补全所有节点的维度信息,避免修改过程中出现维度不匹配问题:

import onnx
from onnx.tools import update_model_dims
from onnxsim import simplify

# 加载原始模型
model = onnx.load("your_old_model.onnx")
# 简化模型,消除内部硬编码的常量维度
model_simp, check = simplify(model)
assert check, "模型简化失败"
# 执行形状推断
model_simp = onnx.shape_inference.infer_shapes(model_simp)
  • 第二步:确认输入名称和原始维度
    打印所有输入的信息,记录7个输入的名称和原有维度:
for idx, input_node in enumerate(model_simp.graph.input):
    dims = [d.dim_value for d in input_node.type.tensor_type.shape.dim]
    print(f"输入{idx} 名称:{input_node.name} 原始维度:{dims}")
  • 第三步:更新输入维度添加batch轴
    构造维度更新规则,给每个输入的最前面添加动态batch维度(用-1表示可变batch大小),比如原有输入维度是(256,),更新后就是(-1, 256):
# 替换成你自己的7个输入的名称和对应的新维度
dim_update_dict = {
    "input_name_1": [-1, 256],
    "input_name_2": [-1, 128],
    # 依次写完7个输入的配置
}

# 执行维度更新
updated_model = update_model_dims.update_inputs_outputs_dims(
    model_simp, 
    input_dims=dim_update_dict,
    output_dims={} # 输出维度会自动同步更新,不需要额外配置
)

# 验证模型合法性
onnx.checker.check_model(updated_model)
# 保存新模型
onnx.save(updated_model, "model_with_batch.onnx")
  • 验证效果
    取原有单样本输入,用np.expand_dims(input, 0)在最前面新增batch轴后送入新模型推理,对比和原模型的推理结果,误差小于1e-6即为修改成功。

内容的提问来源于stack exchange,提问作者Arjun

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 01:48:05