获取独热编码后决策树的特征名并优化可视化显示
解决决策树特征名称展示与符号替换问题
针对你的需求,我们可以通过生成独热编码后的特征名称、自定义决策树节点文本来解决这两个问题,以下是修改后的完整代码及说明:
步骤说明
- 定义原始特征名称并生成独热编码后的特征名:将每个原始特征的类别展开为
特征名_类别名格式,方便在决策树中识别。 - 优化标签编码方式:使用
LabelEncoder替代手动替换,更规范可靠。 - 自定义决策树节点文本:将默认的
<=0.5替换为对应类别等于的判断(因为独热编码特征仅为0或1,阈值0.5对应是否属于该类别),提升可读性。 - 增强决策树可视化细节:添加类别名称、填充颜色、圆角等,让树结构更清晰。
修改后的代码
import numpy as np from sklearn import preprocessing from sklearn import tree import matplotlib.pyplot as plt # 1. 定义原始特征名称和数据 feature_names = ["Outlook", "Temperature", "Humidity", "Wind"] X = np.array([ ["sunny", "sunny", "overcast", "rain", "rain", "rain", "overcast", "sunny", "sunny", "rain", "sunny", "overcast", "overcast", "rain"], ["hot", "hot", "hot", "mild", "cool", "cool", "cool", "mild", "cold", "mild", "mild", "mild", "hot", "mild"], ["high", "high", "high", "high", "normal", "normal", "normal", "high", "normal", "normal", "normal", "high", "normal", "high"], ["weak", "strong", "weak", "weak", "weak", "strong", "weak", "weak", "weak", "strong", "strong", "strong", "weak", "strong"] ]) Y = np.array(["no", "no", "yes", "yes", "yes", "no", "yes", "no", "yes", "yes", "yes", "yes", "yes", "no"]) X = X.transpose() Y = Y.transpose() # 2. 独热编码特征 enc = preprocessing.OneHotEncoder() enc.fit(X) Xenc = enc.transform(X).toarray() # 生成独热编码后的特征名称(格式:特征名_类别名) onehot_feature_names = [] for feat, cats in zip(feature_names, enc.categories_): for cat in cats: onehot_feature_names.append(f"{feat}_{cat}") # 3. 标签编码(替代手动替换,更规范) label_enc = preprocessing.LabelEncoder() Yenc = label_enc.fit_transform(Y) class_names = label_enc.classes_ # 获取类别名称:['no', 'yes'] # 4. 训练决策树 clf = tree.DecisionTreeClassifier() clf.fit(Xenc, Yenc) # 5. 自定义节点文本生成函数,替换<=为=,并转换为直观的类别判断 def custom_node_text(node_id, tree_model, feature_names): feature_idx = tree_model.tree_.feature[node_id] threshold = tree_model.tree_.threshold[node_id] if feature_idx == tree._tree.TREE_UNDEFINED: # 叶子节点,返回类别统计 counts = tree_model.tree_.value[node_id][0] total = counts.sum() return f"class: {class_names[np.argmax(counts)]}\nsamples: {total}\nvalue: {counts.tolist()}" else: # 分裂节点,替换符号并转换为类别判断 feat_name = feature_names[feature_idx] # 拆分特征名和类别:比如Outlook_sunny -> Outlook, sunny original_feat, cat = feat_name.split("_", 1) if threshold <= 0.5: # <=0.5表示该类别不成立,即原始特征不等于该类别 return f"{original_feat} != {cat}" else: # >0.5表示该类别成立,即原始特征等于该类别 return f"{original_feat} = {cat}" # 6. 绘制决策树 plt.figure(figsize=(12, 8)) tree.plot_tree( clf, feature_names=onehot_feature_names, class_names=class_names, filled=True, rounded=True, impurity=False, node_ids=False, # 使用自定义节点文本 label=lambda node_id: custom_node_text(node_id, clf, onehot_feature_names) ) plt.show()
关键修改点说明
- 特征名称生成:通过遍历
enc.categories_和原始特征名,生成特征名_类别名格式的独热特征名,解决决策树无特征标识的问题。 - 节点文本替换:利用
label参数传入自定义函数,将基于独热特征的阈值判断转换为原始字符串类别的等于/不等于判断,彻底替换掉<=符号,让决策逻辑更直观。 - 可视化增强:开启
filled=True和rounded=True,添加类别名称,让树结构更易读。
内容的提问来源于stack exchange,提问作者taou
相关产品推荐
相关产品推荐

