讲师中心 微信公众号
AI工具推荐 视频效率加速

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

星敏大大_9157

星敏大大_9157

发布时间:2026-08-06 11:56:23

|

771人浏览过

|

来源于php中文网

原创

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><p><span>立即学习</span>“<a href="https://pan.quark.cn/s/00968c3c2c15" style="text-decoration: underline !important; color: blue; font-weight: bolder;" rel="nofollow" target="_blank">Python免费学习笔记(深入)</a>”;</p><div class="aritcle_card flexRow">
                                                        <div class="artcardd flexRow">
                                                                <a class="aritcle_card_img" href="/xiazai/skill4769" title="Python Testing"><img
                                                                                src="https://img.php.cn/upload/skill/000/000/081/179021887894914.jpg" alt="Python Testing"  onerror="this.onerror='';this.src='/static/lhimages/moren/morentu.png'" ></a>
                                                                <div class="aritcle_card_info flexColumn">
                                                                        <a href="/xiazai/skill4769" title="Python Testing">Python Testing</a>
                                                                        <p>Python 测试速查:运行 pytest、使用 mock/patch、参数化、fixtures、异步、覆盖率测试。</p>
                                                                </div>
                                                                <a href="/xiazai/skill4769" title="Python Testing" 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。

热门AI工具

更多
DeepSeek

DeepSeek是一款面向对话、写作、编程和推理场景的AI大模型工具。

音述AI
音述AI Hot

一款AI音频处理工具,主要用于音述AI是一个以“用声音述说故事”为核心的 AI 音乐创作与声音分享社区,适合需要提升相关任务效率的用户。

PixTV
PixTV Hot

PixTV是一款面向AIGC内容创作的AI视频生成工具。

LibLibAI
LibLibAI Hot

一款AI视频创作工具,主要用于国内领先的AI创意平台,以海量模型、低门槛操作与“创作-分享-商业化”生态,让小白与专业创作者都能高效实现图文乃至视频创意表达,适合需要提升相关任务效率的用户。

超级简历WonderCV

一款AI办公效率工具,主要用于免费求职简历模版下载制作,应届生职场人必备简历制作神器,适合需要提升相关任务效率的用户。

Seko
Seko Hot

一款AI视频创作工具,主要用于商汤科技推出的创编一体的AI短视频创作Agent,适合需要提升相关任务效率的用户。

WorkBuddy

一款AI办公效率工具,主要用于腾讯云推出的AI原生桌面智能体工作台,适合需要提升相关任务效率的用户。

豆包大模型

豆包大模型是一款由字节跳动推出的企业级大语言模型服务平台。

墨刀AI
墨刀AI Hot

一款AI图像与设计工具,主要用于产品经理的专属智能体,适合需要提升相关任务效率的用户。

相关专题

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

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

1631

2023.07.20

python能做什么
python能做什么

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

3984

2023.07.25

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

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

1629

2023.07.31

python教程
python教程

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

22997

2023.08.03

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

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

2827

2023.08.04

python eval
python eval

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

2867

2023.08.04

scratch和python区别
scratch和python区别

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

1123

2023.08.11

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

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

596

2023.08.10

LLVM自定义Pass怎么写
LLVM自定义Pass怎么写

本专题聚焦LLVM自定义Pass开发,整理Pass类结构、run()方法、PreservedAnalyses、CMake构建、插件注册、-load-pass-plugin加载和测试用例编写流程。

20

2026.09.30

热门下载

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

精品课程

更多
相关推荐
/
热门推荐
/
最新课程
关于我们 免责申明 举报中心 意见反馈 讲师合作 广告合作 最新更新
php中文网:公益在线php培训,帮助PHP学习者快速成长!
关注服务号
PHP中文网订阅号
每天精选资源文章推送

Copyright 2014-2026 https://www.php.cn/ All Rights Reserved | php.cn