如何在Python中将Scikit-learn训练好的模型转换为ONNX格式?

星敏大大_9157

星敏大大_9157

2026-08-06

771人浏览

原创

scikit-learn模型导出onnx前必须满足:仅支持官方兼容的estimator(如linearregression、randomforestclassifier等),禁用gradientboostingclassifier等未注册模型;输入须为numpy.ndarray且dtype与initial_types声明一致(如float32训练则用floattensortype);pipeline需所有step均被skl2onnx支持,否则须拆解或手动注册converter。

如何在python中将scikit-learn训练好的模型转换为onnx格式?

scikit-learn模型导出ONNX前必须满足哪些条件?

不是所有 sklearn 模型都能直接转 ONNX。ONNX Runtime 官方支持的转换器(skl2onnx)只覆盖部分 estimator 类型,比如 LogisticRegression、RandomForestClassifier、SVC、LinearRegression 等,但不支持 GradientBoostingClassifier(除非用 skl2onnx.convert_sklearn + 自定义 converter)、也不支持含自定义 transformer 的 Pipeline(除非该 transformer 已被 skl2onnx 显式支持)。

  • 必须使用 sklearn 原生 estimator,不能是封装过的类(如继承自 BaseEstimator 但未注册 converter 的自定义类)
  • 输入数据类型需明确:训练时用的是 numpy.ndarray 或 pandas.DataFrame,但导出时建议统一为 numpy.ndarray,否则可能触发 shape 推断失败
  • 分类任务中,predict_proba 是否可用取决于模型和 converter 版本;例如 RandomForestClassifier 在较新 skl2onnx 中默认支持,但老版本可能只输出 predict

用skl2onnx完成转换的最小可行代码怎么写?

核心是三步:构造 converter → 调用 convert_sklearn → 保存为 .onnx 文件。注意不能直接用 onnx.save,得靠 convert_sklearn 返回的 onnx.ModelProto 对象。

from sklearn.ensemble import RandomForestClassifier
from sklearn.datasets import make_classification
from skl2onnx import convert_sklearn
from skl2onnx.common.shape_calculator import calculate_linear_classifier_output_shapes
from skl2onnx.common.data_types import FloatTensorType
<h1>训练一个简单模型</h1><p>X, y = make_classification(n_samples=1000, n_features=4, n_classes=2, random_state=42)
model = RandomForestClassifier(n_estimators=10, max_depth=3, random_state=42)
model.fit(X, y)</p><h1>定义输入类型:必须指定 batch_size=1 和特征数</h1><p>initial_type = [('float_input', FloatTensorType([None, X.shape[1]]))]</p><h1>转换(classifier=True 启用概率输出)</h1><p>onnx_model = convert_sklearn(model, initial_types=initial_type, options={id(model): {'zipmap': False}})</p><div class="aritcle_card flexRow artxards">
											<div class="artcardd flexRow">
												<a class="aritcle_card_img" rel="nofollow" href="/xiazai/skill5288" title="提示词大师-python版"><img
														src="https://img.php.cn/upload/skill/000/000/081/179042051830184.jpg" alt="提示词大师-python版" onerror="this.onerror='';this.src='/static/lhimages/moren/morentu.png'" ></a>
												<div class="aritcle_card_info flexColumn">
													<a rel="nofollow" href="/xiazai/skill5288" title="提示词大师-python版" class="overflowclass">提示词大师-python版</a>
													<p class="overflowclass">图片提示词生成器?不止如此。
马甲系统 —— 把脑海中的画面,翻译成AI能理解的专业表达。
用得越多,它越懂你:首次需要多问几句确认方向,用久了几乎一说就懂。
用得越多,它越快:缓存机制让后续对话越来越省。
RAG进化:成功案例持续入库,越跑越聪明。
输入「新手指南」查看完整功能介绍</p>
												</div>
												<a rel="nofollow" href="/xiazai/skill5288" title="提示词大师-python版" class="aritcle_card_btn flexRow flexcenter"><b></b><span>下载</span>
												</a>
											</div>
										</div><h1>保存</h1><p>with open('rf.onnx', 'wb') as f:
f.write(onnx_model.SerializeToString())</p>
  • initial_types 里 [None, X.shape[1]] 表示动态 batch size,别写成 [1, 4] —— 否则推理时输入 batch >1 就报错
  • options 中的 zipmap: False 是为了去掉默认添加的 ZipMap 后处理节点,让输出是 raw logits 或 proba 数组,更便于下游解析
  • 如果模型是回归类(如 LinearRegression),不用传 classifier=True,也无需 zipmap 相关配置

转换后验证ONNX模型是否能正确推理?

不能只看文件生成成功,必须用 onnxruntime 实际跑一次,并比对输出。常见失效点是 dtype 不匹配或输入 name 错误。

  • onnxruntime.InferenceSession 加载后,先查 session.get_inputs()[0].name,确保你 feed 的 key 和它一致(常是 'float_input',不是 'input')
  • 输入 numpy array 必须是 np.float32,哪怕训练时用的是 float64 —— ONNX 默认按 float32 解析,否则会静默截断或报 InvalidArgument
  • 分类模型输出有两个 blob:'probabilities'(当 zipmap=True)或 'label' + 'probabilities';若设 zipmap=False,则只有 'output',内容是 shape=(N, n_classes) 的概率数组
import onnxruntime as ort
import numpy as np
<p>sess = ort.InferenceSession('rf.onnx')
input_name = sess.get_inputs()[0].name
pred_onx = sess.run(None, {input_name: X.astype(np.float32)[:2]})[0]
pred_sk = model.predict_proba(X[:2])  # 注意:这里要和 ONNX 输出对齐维度
np.testing.assert_allclose(pred_onx, pred_sk, atol=1e-5)</p>

为什么Pipeline转换经常失败?

sklearn.pipeline.Pipeline 本身不被 skl2onnx 原生支持,除非每个 step 都是已注册 converter 的类。常见陷阱:

  • 包含 StandardScaler 是安全的(skl2onnx 支持),但包含自定义 TransformerMixin 子类就失败,除非你手动注册 converter
  • ColumnTransformer 支持有限:仅支持 OneHotEncoder、StandardScaler 等少数 transformer,且要求 remainder='passthrough' 或 remainder='drop',不能是 callable
  • 如果 pipeline 最后一步是 LogisticRegression,但前面有不支持的 transformer,整个 pipeline 无法转换 —— 此时得拆开:先转换 preprocessing 部分(用 convert_sklearn 单独处理 scaler),再拼接 ONNX 图(需用 onnx.compose 或手写 node),复杂度陡增

真正省事的做法是:训练完 pipeline 后,用 pipeline[:-1].transform(X) 提前处理好特征,再单独导出最后的 estimator。这样绕过 pipeline 转换限制,也更容易调试。

ONNX 导出不是“一键打包”,而是依赖 converter 实现的精确映射;一旦模型结构偏离标准 sklearn 接口,就得手动补 converter 或重构 pipeline。

Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!

相关专题

更多
python打包成可执行文件
python打包成可执行文件

本专题为大家带来python打包成可执行文件相关的文章,大家可以免费的下载体验。

2023.07.20

1631

4

python能做什么
python能做什么

python能做的有:可用于开发基于控制台的应用程序、多媒体部分开发、用于开发基于Web的应用程序、使用python处理数据、系统编程等等。本专题为大家提供python相关的各种文章、以及下载和课程。

2023.07.25

3964

7

format在python中的用法
format在python中的用法

Python中的format是一种字符串格式化方法,用于将变量或值插入到字符串中的占位符位置。通过format方法,我们可以动态地构建字符串,使其包含不同值。php中文网给大家带来了相关的教程以及文章,欢迎大家前来阅读学习。

2023.07.31

1629

3

python教程
python教程

Python已成为一门网红语言,即使是在非编程开发者当中,也掀起了一股学习的热潮。本专题为大家带来python教程的相关文章,大家可以免费体验学习。

2023.08.03

22877

23

python环境变量的配置
python环境变量的配置

Python是一种流行的编程语言,被广泛用于软件开发、数据分析和科学计算等领域。在安装Python之后,我们需要配置环境变量,以便在任何位置都能够访问Python的可执行文件。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2023.08.04

2807

5

python eval
python eval

eval函数是Python中一个非常强大的函数,它可以将字符串作为Python代码进行执行,实现动态编程的效果。然而,由于其潜在的安全风险和性能问题,需要谨慎使用。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2023.08.04

2847

5

scratch和python区别
scratch和python区别

scratch和python的区别:1、scratch是一种专为初学者设计的图形化编程语言,python是一种文本编程语言;2、scratch使用的是基于积木的编程语法,python采用更加传统的文本编程语法等等。本专题为大家提供scratch和python相关的文章、下载、课程内容,供大家免费下载体验。

2023.08.11

1123

5

python合并两个列表
python合并两个列表

Python是一种强大的编程语言,具有许多方便的功能和工具。在Python中,有多种方法可以合并两个列表。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2023.08.10

596

4

python是前端还是后端
python是前端还是后端

Python属于前端也属于后端,其灵活性和丰富的生态系统使得开发人员能够在不同的领域中灵活运用。本专题为大家提供python相关的文章、下载、课程内容,供大家免费下载体验。

2023.08.11

2203

5

热门下载

更多
网站特效
/
网站源码
/
网站素材
/
前端模板

精品课程

更多
相关推荐
/
热门推荐
/
最新课程