scikit-learn自定义类构建Pipeline序列化报错:无法pickle _thread.RLock对象
sklearn Pipeline序列化报
can't pickle _thread.RLock objects错误的解决方法 问题根源
该错误同时仅在带main函数的脚本中复现,核心原因有两点:
- 自定义估计器不符合scikit-learn开发约定:你在
CustomKMeans的__init__方法中直接实例化了内部的KMeans对象,同时__init__使用了可变默认参数cluster_parameters = {},会导致实例属性意外持有额外的作用域引用;此外如果ForecastPredictor类中的训练完成的LGBMRegressor在训练时启用了多线程,部分版本的lightgbm会在模型对象中保留线程锁引用,无法被序列化。 - 作用域不匹配:Jupyter环境中所有类、函数默认在全局作用域定义和实例化,而带main函数的脚本中如果自定义类定义在
if __name__ == "__main__"代码块内部,或者pipeline实例在main函数的局部作用域创建且类的导入路径不明确,pickle在序列化时无法解析类的全局引用,会连带序列化局部作用域的上下文对象,其中可能包含线程锁。
修复方案
1. 修正自定义估计器实现,符合scikit-learn约定
- 所有内部估计器的实例化放到
fit方法中执行,不要在__init__里直接实例化 - 不要使用可变对象作为
__init__的默认参数,改用None作为默认值再在方法内初始化
修正后的CustomKMeans示例如下:
class CustomKMeans(BaseEstimator, TransformerMixin): # 可变默认参数改用None,避免多个实例共享同一对象 def __init__(self, pretrained=False, force_transformation=False, cluster_parameters=None): self.pretrained = pretrained self.force_transformation = force_transformation self.cluster_parameters = cluster_parameters if cluster_parameters is not None else {} # 内部KMeans实例化移到fit方法中 def fit(self, X, y=None): # 内部估计器统一在fit方法实例化 self.KMeans = KMeans(**self.cluster_parameters) self.KMeans.fit(X) return self def transform(self, X): return self.KMeans.predict(X)
对PreprocessArticleInfo、ForecastPredictor两个类也做同步修改,所有内部的ColumnTransformer、LGBMRegressor都放到fit方法里实例化,不要在__init__中创建。
2. 清理lightgbm模型的线程锁引用
如果错误由LGBMRegressor带来,可在训练完成后重置模型线程参数,清除锁引用:
# 在ForecastPredictor的fit方法、训练完LGBMRegressor后执行 self.lgb_model.booster_.reset_parameter({'num_threads': 1})
如果不需要保留booster的训练状态,也可以通过导出参数的方式彻底避免锁问题:
# 训练完成后导出参数,清空模型对象 self.lgb_params = self.lgb_model.get_params() self.lgb_model = None # 预测前重新加载模型 self.lgb_model = LGBMRegressor(**self.lgb_params)
3. 修正脚本作用域结构
确保所有自定义类都定义在脚本的顶层作用域,不要嵌套在main函数或者if __name__ == "__main__"代码块内部,参考脚本结构:
import joblib import pandas as pd from sklearn.pipeline import Pipeline from sklearn.base import BaseEstimator, TransformerMixin, RegressorMixin # 所有自定义类统一定义在脚本顶层 class PreprocessSalesData(BaseEstimator, TransformerMixin): # 类实现逻辑 pass class PreprocessArticleInfo(BaseEstimator, TransformerMixin): # 类实现逻辑 pass class CustomKMeans(BaseEstimator, TransformerMixin): # 修正后的类实现逻辑 pass class ForecastPredictor(BaseEstimator, RegressorMixin): # 修正后的类实现逻辑 pass # main函数仅处理pipeline实例化、训练、序列化逻辑 def main(): pipeline = Pipeline( steps=[ ("preprocess_sales_data", PreprocessSalesData()), ("preprocess_article_info", PreprocessArticleInfo()), ("find_neighbours", CustomKMeans()), ("sales_regressor", ForecastPredictor()) ] ) # 此处执行pipeline.fit(训练数据)逻辑 joblib.dump(pipeline, 'pipe.joblib') if __name__=="__main__": main()
内容的提问来源于stack exchange,提问作者M. Moresi
相关产品推荐
相关产品推荐

