如何基于HistGradientBoostingClassifier绘制决策树?
我有一个已拟合好的HistGradientBoostingClassifier模型(命名为RF_90),想绘制其中一棵或多棵决策树,但没找到原生实现的函数。我能访问TreePredictor对象及其节点,但sklearn.tree.plot_tree只支持DecisionTree类型的对象。我试了下面的代码:
from sklearn.tree import plot_tree plot_tree(RF_90._predictors[0][0])
结果触发了错误:
InvalidParameterError: The 'decision_tree' parameter of plot_tree must
be an instance of 'sklearn.tree._classes.DecisionTreeClassifier' or an
instance of 'sklearn.tree._classes.DecisionTreeRegressor'. Got
<sklearn.ensemble._hist_gradient_boosting.predictor.TreePredictor
object at 0x7f676ebf0310> instead.
TreePredictor是HistGradientBoosting系列模型的内部预测器,并不兼容plot_tree的输入要求,你可以通过以下两种方式实现树的可视化:
方法1:手动解析TreePredictor节点绘制
TreePredictor的nodes_属性存储了树的核心节点数据(特征索引、阈值、子节点索引、样本数等),可以遍历这些数据,用matplotlib手动绘制树结构:
import matplotlib.pyplot as plt from matplotlib.pylab import rcParams rcParams['figure.figsize'] = 15, 20 def plot_tree_predictor(tree_predictor, feature_names=None, ax=None): if ax is None: ax = plt.gca() nodes = tree_predictor.nodes_ def draw_node(node_idx, x, y, width): node = nodes[node_idx] # 绘制叶子节点 if node.is_leaf: ax.text(x, y, f"叶子节点\n输出值: {node.value[0]:.2f}", ha='center', va='center', bbox=dict(boxstyle='round', fc='white')) return # 绘制分裂节点 feature = feature_names[node.feature_idx] if feature_names else f"特征 {node.feature_idx}" ax.text(x, y, f"{feature} ≤ {node.threshold:.2f}\n样本数: {node.n_samples}", ha='center', va='center', bbox=dict(boxstyle='round', fc='lightblue')) # 绘制子节点连线 left_x = x - width/2 right_x = x + width/2 next_y = y - 1 ax.plot([x, left_x], [y-0.2, next_y+0.2], 'k-') ax.plot([x, right_x], [y-0.2, next_y+0.2], 'k-') # 递归绘制子节点 draw_node(node.left_child_idx, left_x, next_y, width/2) draw_node(node.right_child_idx, right_x, next_y, width/2) draw_node(0, 0.5, 1, 0.5) ax.axis('off') plt.show() # 使用示例,替换your_feature_names为你的特征列表 plot_tree_predictor(RF_90._predictors[0][0], feature_names=your_feature_names)
方法2:转换为兼容的DecisionTree对象
手动构建DecisionTreeClassifier对象,将TreePredictor的节点数据映射到该对象的内部结构,即可直接用plot_tree绘图:
from sklearn.tree import DecisionTreeClassifier, plot_tree def tree_predictor_to_decision_tree(tree_predictor, feature_names=None): dt = DecisionTreeClassifier() dt.n_features_in_ = tree_predictor.n_features_ dt.feature_names_in_ = feature_names if feature_names else [f"特征{i}" for i in range(dt.n_features_in_)] # 模拟DecisionTree的tree_属性结构 dt.tree_ = type('MockTree', (), {})() nodes = tree_predictor.nodes_ dt.tree_.node_count = len(nodes) dt.tree_.children_left = [node.left_child_idx if not node.is_leaf else -1 for node in nodes] dt.tree_.children_right = [node.right_child_idx if not node.is_leaf else -1 for node in nodes] dt.tree_.feature = [node.feature_idx if not node.is_leaf else -2 for node in nodes] dt.tree_.threshold = [node.threshold if not node.is_leaf else -2 for node in nodes] dt.tree_.value = [[node.value] for node in nodes] dt.tree_.n_node_samples = [node.n_samples for node in nodes] dt.tree_.max_depth = tree_predictor.max_depth_ return dt # 转换并绘图 dt = tree_predictor_to_decision_tree(RF_90._predictors[0][0], feature_names=your_feature_names) plot_tree(dt, feature_names=your_feature_names, filled=True) plt.show()
注意:这种转换依赖sklearn内部属性结构,不同版本可能存在兼容性问题,使用前建议测试。
内容的提问来源于stack exchange,提问作者Alejandro Gonzalez

