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

如何实现带if-else分支的双随机森林模型单Pipeline以导出PMML

解决方案

实现思路

Scikit-learn 原生 Pipeline 不支持直接写 if-else 分支逻辑,要实现条件执行同时兼容 PMML 导出,需要自定义继承了BaseEstimator和TransformerMixin的转换类来封装分支规则,这种自定义组件可以被sklearn2pmml正常识别导出。

代码实现

1. 导入依赖

from sklearn.base import BaseEstimator, TransformerMixin
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import StandardScaler
from sklearn.ensemble import RandomForestClassifier, RandomForestRegressor
import numpy as np

2. 定义条件执行类

class ConditionalExecutionTransformer(BaseEstimator, TransformerMixin):
    def __init__(self, clf_pipeline, reg_pipeline, trigger_label=2):
        # 传入已训练好的分类、回归Pipeline,以及触发回归的分类标签
        self.clf_pipeline = clf_pipeline
        self.reg_pipeline = reg_pipeline
        self.trigger_label = trigger_label
    
    def fit(self, X, y=None):
        # 两个子Pipeline已经提前训练完成,无需重复fit,直接返回实例即可
        return self
    
    def transform(self, X):
        # 第一步:执行分类预测
        clf_result = self.clf_pipeline.predict(X)
        # 初始化回归结果,非触发样本默认填充NaN,可按需修改为其他默认值
        reg_result = np.full(clf_result.shape, np.nan)
        # 筛选出分类结果等于触发标签的样本
        trigger_idx = clf_result == self.trigger_label
        # 对符合条件的样本执行回归预测
        if np.any(trigger_idx):
            reg_result[trigger_idx] = self.reg_pipeline.predict(X[trigger_idx])
        # 返回结果可按需调整,这里同时返回分类结果和回归结果
        return np.column_stack([clf_result, reg_result])

3. 组装总Pipeline

# 你已提前训练好的两个独立Pipeline
# 分类Pipeline
pipe_clf = Pipeline([('classifier', RandomForestClassifier())])
# 回归Pipeline
pipe_reg = Pipeline([('scaler', StandardScaler()),('regressor', RandomForestRegressor())])

# 组装为单个总Pipeline
total_pipeline = Pipeline([
    ('conditional_step', ConditionalExecutionTransformer(
        clf_pipeline=pipe_clf,
        reg_pipeline=pipe_reg,
        trigger_label=2
    ))
])

4. 导出为PMML文件

使用sklearn2pmml导出即可,注意自定义类需要放在脚本顶层作用域,不要嵌套在其他函数中避免识别失败:

from sklearn2pmml import sklearn2pmml

sklearn2pmml(total_pipeline, "combined_model.pmml", with_repr=True)

注意事项

  • 如果需要非触发样本返回其他默认值,修改transform方法中reg_result的初始化逻辑即可
  • 若需要更高的PMML兼容性,可以改用sklearn2pmml内置的规则集组件替代自定义类

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 03:24:03