如何为无批量维度的已训练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
相关产品推荐
相关产品推荐

