无法加载经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
相关产品推荐
相关产品推荐

