shap.plots.bar()报错及summary_plot合并剩余特征需求求助
问题解决:SHAP条形图展示前10特征并合并剩余项
一、shap.plots.bar()报错原因及修复
报错AssertionError: You must pass an Explanation object...是因为该方法要求传入SHAP Explanation对象,而你传入的shap_values[0]只是原始的SHAP值数组,不符合参数要求。
修复步骤:
- 用
shap.Explanation封装SHAP值、特征数据和特征名 - 传入
shap.plots.bar()并设置max_display=10,该参数会自动将排名10以后的所有特征重要性合并为名为"Other"的条形
修改后的代码:
explainer = shap.KernelExplainer(model=agent.policy.predict, data=state_df, link="identity") shap_values = explainer.shap_values(X = state_df.iloc[0:35,:]) # 创建Explanation对象 exp = shap.Explanation( values=shap_values[0], data=state_df.iloc[0:35,:], feature_names=state_df.columns ) # 自动合并剩余特征的条形图 shap.plots.bar(exp, max_display=10)
二、shap.summary_plot(..., plot_type="bar")实现合并剩余特征的替代方案
summary_plot本身没有自动合并剩余特征的参数,需要手动处理数据后绘图:
- 计算每个特征的平均绝对SHAP值(条形图默认指标)
- 按重要性排序后取前10个,剩余特征的重要性求和合并为"Other"
- 用matplotlib手动绘制条形图
示例代码:
import matplotlib.pyplot as plt import numpy as np # 计算每个特征的平均绝对SHAP值 mean_abs_shap = np.mean(np.abs(shap_values[0]), axis=0) # 配对特征名与重要性并降序排序 feature_shap_sorted = sorted( zip(state_df.columns, mean_abs_shap), key=lambda x: x[1], reverse=True ) # 提取前10特征,合并剩余项 top10 = feature_shap_sorted[:10] other_sum = sum([x[1] for x in feature_shap_sorted[10:]]) top10.append(("Other", other_sum)) # 拆分数据用于绘图 features = [x[0] for x in top10] shap_scores = [x[1] for x in top10] # 绘制条形图 plt.barh(features, shap_scores) plt.gca().invert_yaxis() # 让高重要性特征在上方 plt.xlabel("Mean Absolute SHAP Value") plt.title("Top 10 Feature Importance (Others Merged)") plt.show()
补充说明
- 多输出模型下,
shap_values是长度等于输出数的列表,你用shap_values[0]处理第一个输出的逻辑是正确的 KernelExplainer适配35个样本的场景,若后续样本量增大,建议根据模型类型改用TreeExplainer(树模型)或DeepExplainer(深度学习模型)提升效率
内容的提问来源于stack exchange,提问作者coar
相关产品推荐
相关产品推荐

