Python3.10加载sklearn DecisionTreeClassifier pickle模型遇ValueError
问题原因
scikit-learn 0.23.1 到 1.3.2 版本间,DecisionTreeClassifier 的内部节点数组结构发生了变化:新版本新增了 missing_go_to_left 字段(用于定义缺失值的分支逻辑),但旧版本序列化的模型没有该字段,导致加载时出现 dtype 不兼容的错误。单纯使用高版本 pickle 协议重新保存无法解决问题,因为模型的核心结构仍为旧版本格式。
解决方案
方案一:在原环境转换为兼容格式(推荐)
回到原环境(Python 3.7.3 + scikit-learn 0.23.1),对模型进行格式转换后重新保存,确保新环境可直接加载:
import joblib from sklearn.tree import DecisionTreeClassifier # 加载旧模型 model = joblib.load("finalized_model_m.sav") # 提取树结构状态,补充新版本所需字段 tree_state = model.tree_.__getstate__() # 定义兼容新版本的dtype,添加missing_go_to_left字段 new_node_dtype = [ ('left_child', '<i8'), ('right_child', '<i8'), ('feature', '<i8'), ('threshold', '<f8'), ('impurity', '<f8'), ('n_node_samples', '<i8'), ('weighted_n_node_samples', '<f8'), ('missing_go_to_left', 'u1') ] # 转换数组格式并设置默认值(1表示缺失值走左分支,与旧版本默认逻辑一致) tree_state['nodes'] = tree_state['nodes'].astype(new_node_dtype) tree_state['nodes']['missing_go_to_left'] = 1 # 重建模型并设置参数 new_model = DecisionTreeClassifier(**model.get_params()) new_model.tree_.__setstate__(tree_state) # 保存兼容新版本的模型 joblib.dump(new_model, "model_compatible.joblib")
之后在新环境(Python 3.10 + scikit-learn 1.3.2)中,直接用 joblib.load 加载转换后的模型即可。
方案二:在新环境手动修复加载的模型
若无法访问原环境,可在新环境加载模型时,手动补全缺失的字段:
import pickle import numpy as np from sklearn.tree import DecisionTreeClassifier for m in models: file = 'finalized_model_' + m + '.sav' # 加载原始模型 with open(file, 'rb') as f: loaded_model = pickle.load(f) # 获取树结构状态并修复节点数组 tree_state = loaded_model.tree_.__getstate__() old_nodes = tree_state['nodes'] # 定义新版本要求的dtype new_dtype = np.dtype([ ('left_child', '<i8'), ('right_child', '<i8'), ('feature', '<i8'), ('threshold', '<f8'), ('impurity', '<f8'), ('n_node_samples', '<i8'), ('weighted_n_node_samples', '<f8'), ('missing_go_to_left', 'u1') ]) # 创建新数组并迁移旧数据,补全新字段默认值 new_nodes = np.zeros(old_nodes.shape, dtype=new_dtype) for field in old_nodes.dtype.names: new_nodes[field] = old_nodes[field] new_nodes['missing_go_to_left'] = 1 # 与旧版本默认逻辑对齐 # 更新树结构并使用模型 tree_state['nodes'] = new_nodes loaded_model.tree_.__setstate__(tree_state) df[m] = loaded_model.predict_proba(X)[:, 1] print(m)
注意事项
missing_go_to_left字段默认设为1,对应旧版本中缺失值走左分支的逻辑;若你的旧模型有特殊缺失值处理规则,需调整该值。- 后续推荐使用 scikit-learn 官方指定的
joblib工具序列化模型,而非原生 pickle,它对 sklearn 模型结构的兼容性更好。
内容的提问来源于stack exchange,提问作者user3369545
相关产品推荐
相关产品推荐

