如何强制指定ONNX opset版本以适配模型量化需求
LightGBM转ONNX及量化问题解决方案
1. 先升级依赖解决opset版本上限问题
你当前使用的onnx 1.10.2、onnxmltools 1.10.0均为多年前的旧版本,本身最高仅支持base域opset 13、ai.onnx.ml域opset 1,要使用更高版本opset先更新依赖包:
- 卸载旧版本:
pip uninstall -y onnx onnxmltools onnxruntime onnxconverter-common - 安装适配新版本:
pip install onnx>=1.14 onnxmltools>=1.14 onnxruntime>=1.16 onnxconverter-common>=1.13
升级后默认支持base域opset 19+、ai.onnx.ml域opset 3+,满足高版本opset需求。
2. 模型转换环节修改
你遇到的两个核心问题:target_opset提示最高仅支持9、量化时ZipMap形状推断错误,都可以在转换环节解决:
The maximum opset needed by this model is only 9是旧版LightGBM转换器默认仅使用低版本算子的提示,不影响功能,转换后可通过官方版本转换工具强制升级opset- ZipMap算子是分类模型默认加的字典格式输出算子,本身不支持量化阶段的形状推断,转换时直接关闭即可,不影响推理结果还能提升推理性能
修改后的转换代码如下:
import numpy as np import lightgbm as lgb import onnx from onnxmltools.convert import convert_lightgbm from onnxconverter_common.data_types import FloatTensorType from onnx.version_converter import convert_version # 原有数据生成、模型训练代码不变 max_depth = 8 num_classes = 2 n_estimators = 1000 n_features = 3000 n_fit = 100 n_pred= 100000 X = np.random.rand(n_fit, n_features).astype(np.float32) y = np.random.randint(num_classes, size=n_fit) test_data = np.random.rand(n_pred, n_features).astype('float32') model = lgb.LGBMClassifier(n_estimators=n_estimators, max_depth=max_depth, pred_early_stop=False) model.fit(X, y) # 转换参数修改 input_types = [("input", FloatTensorType([None, n_features]))] # batch维度改为None支持动态输入,避免硬编码形状引发推断错误 onnx_ml_model = convert_lightgbm( model, initial_types=input_types, target_opset=16, # 按需求指定目标opset,新版本支持19+ zipmap=False # 关闭ZipMap输出,直接输出概率张量 ) # 强制升级到目标opset版本 target_opset_version = 16 onnx_ml_model = convert_version(onnx_ml_model, target_opset_version) # 校验模型合法性 onnx.checker.check_model(onnx_ml_model)
3. 量化代码修改
你原有代码的两个错误:
- 变量名错误:
new_model_path未定义 - 量化接口误用:
quantize_qat仅适用于感知量化训练得到的模型,普通训练的LightGBM模型用动态量化接口即可
修改后的量化代码如下:
from onnxruntime.quantization import quantize_dynamic, QuantType model_path = "ONNX_edge_deployment/src/APIs/YOLO_ONNX/lgbm.onnx" model_quant = 'ONNX_edge_deployment/src/APIs/YOLO_ONNX/lgbm_quant.onnx' onnx.save(onnx_ml_model, model_path) quantized_model = quantize_dynamic( model_input=model_path, model_output=model_quant, weight_type=QuantType.QUInt8, optimize_model=True )
完成以上修改后,你可以执行以下代码验证opset版本是否符合预期:
for opset in onnx_ml_model.opset_import: print(f"domain: '{opset.domain}', version: {opset.version}")
内容的提问来源于stack exchange,提问作者Luis Ramon Ramirez Rodriguez
相关产品推荐
相关产品推荐

