基于3特征训练的分类决策树:2D/3D决策面绘制问题求助
解决3特征决策树的可视化问题
一、绘制所有两特征组合的2D决策面
咱们有3个特征,总共有3组特征组合:(特征0,特征1)、(特征0,特征2)、(特征1,特征2)。核心思路是固定第三个特征的取值(比如取训练集该特征的均值/中位数,或者你关心的特定值),生成前两个特征的网格后,把固定值填充进去组成模型要求的3特征输入,再预测得到决策面。
具体步骤&代码示例
import numpy as np import matplotlib.pyplot as plt from sklearn.tree import DecisionTreeClassifier # 假设你的classifier已训练完成,X_train是(70,3)数组,Y_train是标签 # 定义所有特征组合 feature_pairs = [(0,1), (0,2), (1,2)] feature_names = ["Feature 0", "Feature 1", "Feature 2"] # 设置子图布局 fig, axes = plt.subplots(1, 3, figsize=(18, 5)) for idx, (f1, f2) in enumerate(feature_pairs): ax = axes[idx] # 1. 固定第三个特征的取值(这里用训练集均值,可替换为中位数或自定义值) fixed_feature_idx = 3 - f1 - f2 fixed_val = X_train[:, fixed_feature_idx].mean() # 2. 生成前两个特征的二维网格 x_min, x_max = X_train[:, f1].min() - 0.1, X_train[:, f1].max() + 0.1 y_min, y_max = X_train[:, f2].min() - 0.1, X_train[:, f2].max() + 0.1 xx, yy = np.meshgrid(np.linspace(x_min, x_max, 100), np.linspace(y_min, y_max, 100)) # 3. 构建符合模型要求的3特征输入 input_grid = np.c_[xx.ravel(), yy.ravel()] # 插入固定特征列,确保顺序和训练时一致 input_grid = np.insert(input_grid, fixed_feature_idx, fixed_val, axis=1) # 4. 预测并将结果reshape为和xx匹配的二维形状 Z = classifier.predict(input_grid) Z = Z.reshape(xx.shape) # 5. 绘制决策面和样本点 ax.contourf(xx, yy, Z, alpha=0.3, cmap=plt.cm.Paired) ax.scatter(X_train[:, f1], X_train[:, f2], c=Y_train, edgecolor='k', cmap=plt.cm.Paired) ax.set_xlabel(feature_names[f1]) ax.set_ylabel(feature_names[f2]) ax.set_title(f"Decision Surface (Fixed {feature_names[fixed_feature_idx]} = {fixed_val:.2f})") plt.tight_layout() plt.show()
如果想观察不同固定值下的决策面变化,直接替换fixed_val为分位数或自定义数值即可。
二、解决plt.contourf的Z形状报错问题
你遇到的Input z must be a 2D array错误,核心原因是生成了多余的三维网格,导致xx是三维数组,predict后的一维Z无法匹配二维要求。
咱们梳理正确逻辑:
- 2D决策面只需要两个特征的二维网格(比如(100,100)形状的xx和yy),第三个特征是固定值,不需要生成网格;
- 把固定值重复成和
xx.ravel()长度一致的数组,和xx、yy的ravel结果拼接成3特征输入; - 预测得到的一维Z,直接reshape成xx的二维形状即可。
刚才第一部分的代码已经完全规避了这个问题,你可以对照检查自己的代码:是不是误生成了三个特征的三维网格?或者插入固定特征的顺序出错了?
三、3D可视化全3特征的决策面
决策树的决策面是轴对齐的分段平面(每个分裂都基于单个特征的阈值),可以用两种方式实现3D可视化:
方式1:用voxels展示类别区域
这种方式能直观看到每个类别占据的三维空间:
from mpl_toolkits.mplot3d import Axes3D # 生成三维网格 x_min, x_max = X_train[:,0].min()-0.1, X_train[:,0].max()+0.1 y_min, y_max = X_train[:,1].min()-0.1, X_train[:,1].max()+0.1 z_min, z_max = X_train[:,2].min()-0.1, X_train[:,2].max()+0.1 xx, yy, zz = np.meshgrid(np.linspace(x_min, x_max, 20), np.linspace(y_min, y_max, 20), np.linspace(z_min, z_max, 20)) # 构建输入并预测 input_grid = np.c_[xx.ravel(), yy.ravel(), zz.ravel()] Z = classifier.predict(input_grid) Z = Z.reshape(xx.shape) # 绘制voxels fig = plt.figure(figsize=(10,8)) ax = fig.add_subplot(111, projection='3d') # 为不同类别设置颜色 colors = np.empty(Z.shape, dtype=object) colors[Z==0] = 'red' colors[Z==1] = 'blue' # 多类别可继续添加颜色映射 ax.voxels(xx, yy, zz, Z != Z, facecolors=colors, alpha=0.3) ax.scatter(X_train[:,0], X_train[:,1], X_train[:,2], c=Y_train, edgecolor='k', s=50) ax.set_xlabel("Feature 0") ax.set_ylabel("Feature 1") ax.set_zlabel("Feature 2") ax.set_title("3D Decision Regions of Decision Tree") plt.show()
方式2:绘制决策边界平面
如果想清晰看到决策树的分裂规则,可以提取树的分裂节点,绘制对应的轴对齐平面:
fig = plt.figure(figsize=(10,8)) ax = fig.add_subplot(111, projection='3d') # 绘制样本点 ax.scatter(X_train[:,0], X_train[:,1], X_train[:,2], c=Y_train, edgecolor='k', s=50) # 提取决策树的分裂节点信息 tree = classifier.tree_ feature_indices = tree.feature thresholds = tree.threshold # 遍历分裂节点,绘制对应平面 for i in range(tree.node_count): if feature_indices[i] != -1: # 跳过叶子节点 f_idx = feature_indices[i] t = thresholds[i] # 根据特征索引生成平面网格 if f_idx == 0: yy_plane, zz_plane = np.meshgrid(np.linspace(y_min, y_max, 20), np.linspace(z_min, z_max, 20)) xx_plane = np.full_like(yy_plane, t) ax.plot_surface(xx_plane, yy_plane, zz_plane, color='gray', alpha=0.2) elif f_idx == 1: xx_plane, zz_plane = np.meshgrid(np.linspace(x_min, x_max, 20), np.linspace(z_min, z_max, 20)) yy_plane = np.full_like(xx_plane, t) ax.plot_surface(xx_plane, yy_plane, zz_plane, color='gray', alpha=0.2) elif f_idx == 2: xx_plane, yy_plane = np.meshgrid(np.linspace(x_min, x_max, 20), np.linspace(y_min, y_max, 20)) zz_plane = np.full_like(xx_plane, t) ax.plot_surface(xx_plane, yy_plane, zz_plane, color='gray', alpha=0.2) ax.set_xlabel("Feature 0") ax.set_ylabel("Feature 1") ax.set_zlabel("Feature 2") ax.set_title("3D Decision Boundaries of Decision Tree") plt.show()
这种方式能直观展示决策树如何通过一个个轴平面对空间进行划分。
内容的提问来源于stack exchange,提问作者Mymozaaa
相关产品推荐
相关产品推荐

