如何修改SHAP依赖图代码添加回归线与直方图?
修改SHAP依赖图:添加回归线与独立直方图(匹配Kim等人2021年《Scientific reports》样式)
现有代码仅能生成SHAP散点依赖图,需修改为包含回归线、右侧独立特征直方图的样式,参考Kim等人2021年发表于《Scientific reports》的示例图。
原始代码
rf = RandomForestClassifier(n_estimators=500) rf.fit(X_train, y_train) explainer = shap.TreeExplainer(rf) shap_values = explainer.shap_values(X_test) shap.dependence_plot("var_1", shap_values[1], X_test, interaction_index=None, color='orange',show=False)
修改后的完整代码
import matplotlib.pyplot as plt from matplotlib.gridspec import GridSpec from sklearn.linear_model import LinearRegression import numpy as np # 原有模型训练与SHAP计算逻辑保留 rf = RandomForestClassifier(n_estimators=500) rf.fit(X_train, y_train) explainer = shap.TreeExplainer(rf) shap_values = explainer.shap_values(X_test) # 提取目标特征与对应SHAP值 feature_name = "var_1" x = X_test[feature_name].values.reshape(-1, 1) y_shap = shap_values[1] # 自定义布局:主散点图占80%宽度,直方图占20% gs = GridSpec(1, 2, width_ratios=[4, 1]) ax_main = plt.subplot(gs[0]) ax_hist = plt.subplot(gs[1], sharey=ax_main) # 绘制主图散点 ax_main.scatter(x, y_shap, color='orange', alpha=0.6, s=15) # 添加回归线 reg = LinearRegression().fit(x, y_shap) y_pred = reg.predict(x) ax_main.plot(x, y_pred, color='darkred', linewidth=2) # 绘制右侧特征分布直方图(水平方向) ax_hist.hist(x, bins=20, orientation='horizontal', color='lightgray', edgecolor='black') # 样式调整,匹配期刊图风格 ax_main.set_xlabel(feature_name, fontsize=12) ax_main.set_ylabel(f"SHAP value for {feature_name}", fontsize=12) ax_main.grid(axis='y', linestyle='--', alpha=0.7) # 隐藏直方图的y轴刻度,与主图对齐 ax_hist.set_yticklabels([]) ax_hist.set_xlabel("Count", fontsize=10) ax_hist.tick_params(axis='x', labelsize=8) # 调整子图间距 plt.tight_layout() plt.show()
关键修改说明
- 用
GridSpec手动划分绘图区域,摆脱shap.dependence_plot的默认布局限制,实现主图+右侧直方图的组合 - 手动提取特征与SHAP值绘制散点,灵活添加自定义回归线(通过
LinearRegression拟合) - 在右侧子图绘制目标特征的水平直方图,直观展示特征分布
- 调整坐标轴标签、网格、样式,对齐子图,匹配期刊图的简洁专业风格
内容的提问来源于stack exchange,提问作者Umer Mansoor
相关产品推荐
相关产品推荐

