多分类场景下SHAP值形状异常及多类别SHAP可视化疑问
SHAP多分类场景下的形状问题与多类别可视化方法
问题背景
我有一个包含816行、8个数值特征的小型数据集,目标变量为包含4个类别的分类变量。在该数据集上训练XGB模型后,使用SHAP分析特征重要性的代码如下:
explainer = shap.TreeExplainer(model, X_test) shap_values = explainer(X_test_scaled_df)
运行后得到的shap_values形状为(816,8,4)(816为样本数、8为特征数、4为类别数),与文档/线上示例的预期形状不符。必须执行以下代码才能绘制类别0的特征重要性:
shap.plots.bar(shap_values[:,:,0])
但文档及线上示例显示只需执行shap.plots.bar(shap_values)即可。使用SHAP版本0.45.1,更新版本后结果一致,请问这是为什么?另外,如何在单个图中展示所有类别的SHAP值?
问题解答
1. 为什么shap_values是三维数组?
这是多分类任务下的正常行为:
- 当模型是二分类时,
TreeExplainer默认只返回正类的SHAP值,所以形状是(n_samples, n_features),这也是大部分基础示例展示的场景。 - 当模型是多分类(比如你的4类别任务),
TreeExplainer会为每个类别单独计算SHAP值,因此输出形状为(n_samples, n_features, n_classes),第三维对应每个类别的SHAP贡献。
文档中的示例大多基于二分类或回归任务,所以没有展示多分类的三维数组情况,这才导致你产生认知差异。
2. 如何在单个图中展示所有类别的SHAP值?
有三种常用方法:
方法一:用shap.summary_plot直接生成多类别柱状图
指定plot_type="bar"并传入类别名称,即可在同一图中展示每个特征在所有类别中的平均绝对SHAP值:
# 假设你有类别名称列表,比如class_names = ["类别0", "类别1", "类别2", "类别3"] shap.summary_plot(shap_values, X_test_scaled_df, plot_type="bar", class_names=class_names)
方法二:手动汇总特征全局重要性后绘制
先计算每个特征在所有类别中的平均绝对SHAP值总和,再转换为SHAP可识别的对象绘制:
import numpy as np # 计算每个特征在所有类别中的平均绝对SHAP值总和 feature_importance = np.mean(np.abs(shap_values), axis=(0, 2)) # 转换为Explanation对象适配bar图 shap_exp = shap.Explanation(values=feature_importance, feature_names=X_test_scaled_df.columns) shap.plots.bar(shap_exp)
方法三:绘制分组柱状图对比各类别特征重要性
如果需要对比每个类别单独的特征重要性,可以手动用Matplotlib生成分组柱状图:
import matplotlib.pyplot as plt fig, ax = plt.subplots(figsize=(10,6)) bar_width = 0.2 x = np.arange(len(X_test_scaled_df.columns)) # 遍历每个类别,绘制对应特征的平均绝对SHAP值 for i in range(4): class_shap_mean = np.mean(np.abs(shap_values[:,:,i]), axis=0) ax.bar(x + i*bar_width, class_shap_mean, width=bar_width, label=f"类别{i}") ax.set_xticks(x + bar_width*1.5) ax.set_xticklabels(X_test_scaled_df.columns, rotation=45) ax.legend() plt.title("各类别特征SHAP重要性对比") plt.tight_layout() plt.show()
内容的提问来源于stack exchange,提问作者Mobius88
相关产品推荐
相关产品推荐

