训练RFECV模型后如何持久化保存以实现新数据快速分类?
解决RFECV对象无法Pickle的问题
问题根源
你遇到的TypeError: cannot pickle 'generator' object错误,是因为直接将logo.split(X, y, groups=trial_number)的返回值传给了RFECV的cv参数——这个方法返回的是生成器对象,而生成器无法被序列化(pickle),RFECV会保存这个cv参数,导致整个对象无法被pickle。
方法1:将CV生成器转换为列表
把生成器转换成可序列化的列表,再传入RFECV,这样整个对象就能正常pickle了:
from sklearn.model_selection import LeaveOneGroupOut from sklearn.feature_selection import RFECV from sklearn.discriminant_analysis import LinearDiscriminantAnalysis import pickle logo = LeaveOneGroupOut() # 将生成器转换为列表,避免保存不可序列化的生成器 cv_splits = list(logo.split(X, y, groups=trial_number)) model = RFECV(LinearDiscriminantAnalysis(), step=1, cv=cv_splits) model.fit(X, y) # 现在可以正常保存模型 with open('rfecv_lda_model.pkl', 'wb') as f: pickle.dump(model, f) # 加载模型使用 with open('rfecv_lda_model.pkl', 'rb') as f: loaded_model = pickle.load(f) # 对新数据直接预测(自动完成特征筛选) y_pred = loaded_model.predict(X_new)
方法2:单独保存核心组件(更轻量)
如果你不需要保存整个RFECV对象,只保存筛选后的特征掩码和训练好的LDA模型即可,这样更节省内存,也避免了序列化问题:
# 训练完成后,提取核心组件 lda_estimator = model.estimator_ # 训练好的LDA模型 feature_support = model.support_ # 筛选后保留的特征掩码(布尔数组) # 分别保存 with open('trained_lda.pkl', 'wb') as f: pickle.dump(lda_estimator, f) with open('feature_support.pkl', 'wb') as f: pickle.dump(feature_support, f) # 加载后使用流程 with open('trained_lda.pkl', 'rb') as f: lda = pickle.load(f) with open('feature_support.pkl', 'rb') as f: support = pickle.load(f) # 先对新数据筛选特征,再预测 X_new_filtered = X_new[:, support] y_pred = lda.predict(X_new_filtered)
注意事项
- 方法1中,转换生成器为列表时,要确保内存足够——列表的长度等于你的
trial_number的分组数量,一般癌症诊断数据集的分组数不会过大,不会有内存问题。 - 方法2的优势是保存的文件更小,且只保留了预测所需的核心内容,适合只需要做预测的场景。
内容的提问来源于stack exchange,提问作者Jim
相关产品推荐
相关产品推荐

