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

无法加载经dill序列化的含自定义estimator的Sklearn Pipeline

解决Sklearn Pipeline序列化后加载时的AttributeError问题

问题场景

使用自定义列转换器、estimator及lambda函数构建Sklearn Pipeline,因Pickle无法序列化lambda函数改用dill序列化。自定义estimator代码如下:

class customOLS(BaseEstimator):
    def __init__(self, ols):
        self.estimator_ols = ols

    def fit(self, X, y):
        X = pd.DataFrame(X)
        y = pd.DataFrame(y)
        print('---- Training OLS')
        self.estimator_ols = self.estimator_ols(y,X).fit()
        #print('---- Training LR')
        #self.estimator_lr = self.estimator_lr.fit(X,y)
        return self

    def get_estimators(self):
        return self.estimator_ols #, self.estimator_lr
                
    def predict_ols(self, X):
        res = self.estimator_ols.predict(X)
        return res

pipeline2 = Pipeline(
        steps=[
            ('dropper', drop_cols),
            ('remover',feature_remover),
            ("preprocessor", preprocess_ppl),
            ("estimator", customOLS(sm.OLS))
            ]
    )

序列化代码:

with open('data/baseModel_LR.joblib',"wb") as f:
        dill.dump(pipeline2, f)

加载序列化对象时执行以下代码:

with open('data/baseModel_LR.joblib',"rb") as f:
        model = dill.load(f)
model

出现错误:

AttributeError: 'customOLS' object has no attribute 'ols'

原因分析

Sklearn的序列化机制(dill序列化Pipeline时会遵循Sklearn的规则)会检查自定义estimator的__init__参数是否在实例对象中有对应的属性。你的customOLS类__init__方法接收的参数是ols,但实例中只保存了self.estimator_ols,没有保留self.ols属性,导致加载时Sklearn尝试读取ols属性失败,抛出AttributeError。

另外,你在fit方法中覆盖了self.estimator_ols的初始值(从传入的sm.OLS类变成了拟合后的模型实例),这进一步破坏了序列化时参数与属性的对应关系。

解决方案

修改customOLS类,在__init__中保留传入的ols参数作为实例属性,同时将拟合后的模型存到另一个属性中,确保__init__的参数与实例属性一一对应:

class customOLS(BaseEstimator):
    def __init__(self, ols):
        # 保留__init__参数对应的属性
        self.ols = ols
        # 初始化拟合后的模型属性为None
        self.estimator_ols = None

    def fit(self, X, y):
        X = pd.DataFrame(X)
        y = pd.DataFrame(y)
        print('---- Training OLS')
        # 使用self.ols(传入的OLS类)来创建并拟合模型
        self.estimator_ols = self.ols(y, X).fit()
        return self

    def get_estimators(self):
        return self.estimator_ols
                
    def predict_ols(self, X):
        res = self.estimator_ols.predict(X)
        return res

修改后重新序列化Pipeline,加载时就能正常读取属性,不会再抛出AttributeError。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.23 00:45:12