训练好的scikit-learn管道保存报错PicklingError,如何解决?
解决scikit-learn管道保存时的PicklingError问题
你遇到的这个PicklingError是因为lambda函数无法被pickle序列化。Pickle(包括joblib底层依赖的pickle机制)需要能够找到被序列化对象的原始定义,而lambda是匿名函数,它的定义不会被注册到模块的命名空间(比如__main__)里,所以joblib在保存管道时找不到这个lambda的引用,就会抛出错误。
下面给你几个可行的解决方案,按推荐程度排序:
方案1:将lambda替换为具名函数(最推荐)
把匿名的lambda改成一个明确定义的普通函数,这样pickle就能轻松找到它的引用,保存和加载都不会有问题,代码可读性也更好。
示例代码:
# 定义具名的分词函数 def split_tokenizer(string): return string.split() # 用这个具名函数替换原来的lambda estimators = [ ('tfidf', TfidfVectorizer(tokenizer=split_tokenizer, min_df=20, max_df=0.75, ngram_range=(1,1))), ('clf', RandomForestClassifier(n_estimators=100, n_jobs=-1, class_weight='balanced')) ] p = Pipeline(estimators) p.fit(x_train, y_train) model_path = 'model.joblib' joblib.dump(p, model_path)
方案2:使用cloudpickle序列化lambda函数
如果你一定要保留lambda,可以用cloudpickle库来代替joblib/pickle。它是pickle的扩展,专门支持序列化lambda、嵌套函数这类pickle处理不了的对象。
注意:使用这个方案的话,加载模型时也需要用cloudpickle,而且需要额外安装这个库。
示例代码:
import cloudpickle # 训练管道的代码和原来一致 estimators = [ ('tfidf', TfidfVectorizer(tokenizer=lambda string: string.split(), min_df=20, max_df=0.75, ngram_range=(1,1))), ('clf', RandomForestClassifier(n_estimators=100, n_jobs=-1, class_weight='balanced')) ] p = Pipeline(estimators) p.fit(x_train, y_train) model_path = 'model.cloudpickle' # 用cloudpickle保存模型 with open(model_path, 'wb') as f: cloudpickle.dump(p, f) # 加载模型时同样用cloudpickle with open(model_path, 'rb') as f: loaded_pipeline = cloudpickle.load(f)
方案3:使用内置函数对象(仅适用于简单场景)
你的场景中,string.split()其实是字符串的内置方法,你可以直接传入str.split函数对象,它是内置函数,pickle可以正常序列化,省去定义新函数的步骤:
示例代码:
estimators = [ ('tfidf', TfidfVectorizer(tokenizer=str.split, min_df=20, max_df=0.75, ngram_range=(1,1))), ('clf', RandomForestClassifier(n_estimators=100, n_jobs=-1, class_weight='balanced')) ]
内容的提问来源于stack exchange,提问作者gprinz
相关产品推荐
相关产品推荐

