SHAP瀑布图点垂直重叠难以解读,求解决方法
问题描述
我有一个包含6个输入特征、2个输出变量、1000条观测值的数据集,绘制Waterfall类型的SHAP图时出现了点垂直重叠的问题,导致很难解读输入对输出的影响。
我的代码如下:
# 数据集加载(原代码中dataset = 部分未完整,保留原样) dataset = x = df.iloc[:, 0:6] y = df.iloc[:, [6, 7]] x_train, x_test, y_train, y_test = train_test_split(x, y, test_size=0.2, random_state=0) rf_model = RandomForestRegressor() rf_pipeline = Pipeline([ ('scaler', StandardScaler()), ('regressor', rf_model) ]) rf_pipeline.fit(x_train, y_train) explainer = shap.TreeExplainer(rf_model) shap_values = explainer.shap_values(x_test) shap.summary_plot(shap_values[0], x_test) plt.title('SHAP Values for Output 1') shap.summary_plot(shap_values[1], x_test) plt.title('SHAP Values for Output 2')
问题大概率和数据范围有关,绘图时X轴范围约在-350到30之间,缩小X轴范围后第一个变量的点直接消失了,请问该怎么解决?
解决方法
- 调整点的透明度:在
shap.summary_plot中添加alpha参数(比如alpha=0.3),让重叠的点能透过彼此显示,不用缩小X轴范围也能大致判断数据分布。示例代码:shap.summary_plot(shap_values[0], x_test, alpha=0.3) - 切换为蜜蜂图(beeswarm)或小提琴图:Waterfall图本身在数据量大时容易出现点重叠,换成蜜蜂图可让点更分散,保留所有数据点的同时提升可读性。只需在
summary_plot中指定plot_type="beeswarm":shap.summary_plot(shap_values[0], x_test, plot_type="beeswarm") - 对SHAP值做截断处理:如果极端值由少数异常点导致,可用
numpy对SHAP值做截断,保留大部分数据的同时去掉极端值,比如截断到99%分位数:
这样既缩小了X轴范围,又不会让第一个变量的点消失,因为仅截断了极端的少数点。import numpy as np shap_vals_truncated = np.clip(shap_values[0], np.percentile(shap_values[0], 1), np.percentile(shap_values[0], 99)) shap.summary_plot(shap_vals_truncated, x_test) - 移除不必要的特征缩放:随机森林模型本身不需要特征缩放,你在Pipeline中加入的
StandardScaler属于多余操作,反而可能放大SHAP值的范围。可以去掉Scaler后重新训练模型,再计算SHAP值观察范围是否正常。修改后的Pipeline:rf_pipeline = Pipeline([ ('regressor', rf_model) ])
内容的提问来源于stack exchange,提问作者ziad361
相关产品推荐
相关产品推荐

