如何获取scikit-learn中Random Forest与Decision Tree的类别特定特征重要性
从scikit-learn决策树/随机森林中获取单类别特征重要性的方案
无需训练n_class个一对Rest二分类模型,直接通过解析已训练完成的模型内部节点参数即可得到对应类别的特征重要性,实现逻辑如下:
核心原理
scikit-learn的树模型在训练时,每个分裂节点都会记录本次分裂带来的不纯度减少量(即基尼系数/熵的增益),以及该节点覆盖的所有样本的类别分布。全局特征重要性是所有分裂节点的增益总和按特征统计的结果,我们只需要将每次分裂的增益按当前节点内目标类的样本占比进行拆分,即可得到对应类别的特征重要性。
实现代码
import numpy as np from sklearn.tree import DecisionTreeClassifier from sklearn.ensemble import RandomForestClassifier def get_class_specific_feature_importance(model, target_class): """ 从已训练的决策树/随机森林中提取指定类别的特征重要性 :param model: 已完成拟合的DecisionTreeClassifier / RandomForestClassifier实例 :param target_class: 目标类别在model.classes_数组中的索引值 :return: 长度为特征数的数组,各元素对应特征对目标类的重要性,已归一化 """ n_features = model.n_features_in_ class_importance = np.zeros(n_features, dtype=np.float64) # 兼容决策树和随机森林两种输入 trees = model.estimators_ if isinstance(model, RandomForestClassifier) else [model] for tree in trees: tree_struct = tree.tree_ # 遍历树的所有节点 for node_idx in range(tree_struct.node_count): # 跳过叶子节点(无分裂操作) if tree_struct.children_left[node_idx] == tree_struct.children_right[node_idx]: continue # 本次分裂使用的特征索引 split_feature = tree_struct.feature[node_idx] # 当前节点的各类样本计数 node_class_counts = tree_struct.value[node_idx][0] total_samples = node_class_counts.sum() if total_samples == 0: continue # 目标类在当前节点的样本占比,用于拆分分裂增益 target_ratio = node_class_counts[target_class] / total_samples # 计算本次分裂的总不纯度减少量 gain = tree_struct.weighted_n_node_samples[node_idx] * tree_struct.impurity[node_idx] \ - tree_struct.weighted_n_node_samples[tree_struct.children_left[node_idx]] * tree_struct.impurity[tree_struct.children_left[node_idx]] \ - tree_struct.weighted_n_node_samples[tree_struct.children_right[node_idx]] * tree_struct.impurity[tree_struct.children_right[node_idx]] # 累加对应特征的单类别重要性 class_importance[split_feature] += gain * target_ratio # 归一化处理 total = class_importance.sum() if total > 0: class_importance /= total return class_importance
使用示例
from sklearn.datasets import load_iris from sklearn.model_selection import train_test_split # 加载测试数据集 iris = load_iris() X, y = iris.data, iris.target X_train, X_test, y_train, y_test = train_test_split(X, y, random_state=42) # 训练随机森林(决策树使用逻辑完全一致) rf_model = RandomForestClassifier(n_estimators=100, random_state=42) rf_model.fit(X_train, y_train) # 获取类别索引为1的类别的特征重要性 class_1_importance = get_class_specific_feature_importance(rf_model, target_class=1) print(f"类别{rf_model.classes_[1]}的特征重要性:", class_1_importance)
注意事项
- 该方法仅依赖已训练完成的模型参数,不会增加额外的训练开销,计算效率远高于训练多个二分类器的方案
- 适配基尼系数、熵、对数损失三种常用的树分裂准则,不需要修改代码逻辑
- 如果需要批量获取所有类别的特征重要性,遍历
model.classes_的索引循环调用函数即可
内容的提问来源于stack exchange,提问作者Adrian
相关产品推荐
相关产品推荐

