Pipeline封装决策树调用plot_tree报无tree_属性错误解决方法
问题描述
使用scikit-learn绘制决策树时触发如下报错:
'Pipeline' object has no attribute 'tree_'
复现流程
- 构建包含列预处理器、决策树分类器的Pipeline,其中预处理器分别对类别特征做独热编码、对数值特征做标准化:
preprocessor = ColumnTransformer([ ('one-hot-encoder', categorical_preprocessor, categorical_columns), ('standard_scaler', numerical_preprocessor, numerical_columns)]) model3 = make_pipeline(preprocessor, DecisionTreeClassifier()) - 拟合模型并生成预测结果:
model3 = model3.fit(data_train, target_train) y_pred3 = model3.predict(data_test) - 直接传入Pipeline对象调用绘图方法:
tree.plot_tree(model3)
执行后触发AttributeError,完整报错栈如下:
AttributeError Traceback (most recent call last) ~\AppData\Local\Temp/ipykernel_22012/3111274197.py in <module> ----> 1 tree.plot_tree(model3) ~\anaconda3\lib\site-packages\sklearn\tree\_export.py in plot_tree(decision_tree, max_depth, feature_names, class_names, label, filled, impurity, node_ids, proportion, rounded, precision, ax, fontsize) 193 fontsize=fontsize, 194 ) ---> 195 return exporter.export(decision_tree, ax=ax) 196 197 ~\anaconda3\lib\site-packages\sklearn\tree\_export.py in export(self, decision_tree, ax) 654 ax.clear() 655 ax.set_axis_off() ---> 656 my_tree = self._make_tree(0, decision_tree.tree_, decision_tree.criterion) 657 draw_tree = buchheim(my_tree) 658 AttributeError: 'Pipeline' object has no attribute 'tree_'
核心疑问:
- 报错的根本原因是什么?
- 使用Pipeline封装模型后是否无法实现决策树可视化?
- 正确的决策树绘制方式是什么?
解决方案
报错原因
tree.plot_tree() 方法要求传入的第一个参数是训练完成的决策树估计器实例,而传入的model3是包含预处理、模型训练两步的Pipeline对象,Pipeline本身不存在tree_这个决策树专属属性,因此触发属性不存在的错误。Pipeline封装不会阻碍决策树可视化,只需要传入正确的对象即可。
正确操作步骤
- 从拟合完成的Pipeline中提取最后一步的决策树模型
make_pipeline会自动按照「类名小写」的规则给每一步命名,可以通过named_steps属性按名称索引对应步骤的对象:# 提取决策树模型 dt_clf = model3.named_steps["decisiontreeclassifier"] - (推荐)提取预处理后的特征名称
由于预处理器中的独热编码会扩展原始类别特征,直接绘图会出现特征名和节点分裂规则不匹配、默认显示X[0]/X[1]这类无意义标识的问题,可以通过预处理器的get_feature_names_out()方法拿到处理后的全量特征名:# 提取预处理后的特征名 feature_names = model3.named_steps["columntransformer"].get_feature_names_out() - 调用绘图接口完成可视化
import matplotlib.pyplot as plt # 设置合适的画布大小,避免节点文字重叠 plt.figure(figsize=(20, 16)) tree.plot_tree( dt_clf, feature_names=feature_names, class_names=dt_clf.classes_, # 传入类别名,节点显示更直观 filled=True, rounded=True ) plt.show()
内容的提问来源于stack exchange,提问作者PawelKinczyk
相关产品推荐
相关产品推荐

