
代码绘制成果展示

本套代码构建了一个从数据提取、模型调优到可解释性分析及分析出图的完整机器学习分析流程。首先读取数据集,划分训练集与测试集,通过网格搜索结合交叉验证自动探寻最佳超参数,从而构建出具备极强非线性捕捉能力的XGBoost分类模型;在完成初步的预测评估并绘制包含ROC曲线、混淆矩阵及多指标对比柱状图的综合性能面板后,切入核心的SHAP归因分析环节。分析利用TreeSHAP算法高效计算全局样本的主效应与交互SHAP值,绘制出融合宏观特征重要性与数据分布的组合图,以及直观揭示驱动因子间两两协同互作强度的气泡热力图;在单因素依赖图与双特征交互图中,内嵌了基于Bootstrap自举重采样的LOWESS非参数局部加权平滑算法,为散点趋势附上了严谨的95%置信区间,还能够自动探测并高亮标注出环境驱动因子作用于模型输出时触发阈值(由负转正)的交叉点。在多维特征交互时,该流程还能敏锐捕捉到次级因子状态改变所引发的共阈值突变现象,同时整套流程搭载了60种经典配色的定制颜色库,可实现出版级科研图表的批量输出。








代码解释


第一部分

# =========================================================================================# ====================================== 1. 环境设置 =======================================# =========================================================================================import pandas as pdimport numpy as npimport matplotlib.pyplot as pltimport osfrom statsmodels.nonparametric.smoothers_lowess import lowess

第二部分

# =========================================================================================# ======================================2.颜色库=======================================# =========================================================================================COLOR_SCHEMES = {0: {'train': 'blue', 'test': 'red', 'hist': '#4B0082', 'shap_scatter': '#00008B', 'lowess': '#9400D3','ci': '#D3D3D3', 'inter_low': 'blue', 'inter_low_fit': 'darkblue', 'inter_high': 'red','inter_high_fit': 'darkred', 'cmap': ["blue", "#4B0082", "red"]},}

第三部分

# =========================================================================================# ======================================5.拟合线和置信区间计算函数=======================================# =========================================================================================# 使用LOWESS拟合数据,通过Bootstrap生成大量拟合线,从而计算出95%的置信区间def bootstrap_lowess_ci(x, y, n_boot=200, frac=0.5, ci_level=0.95):sorted_indices_orig = np.argsort(x) # 获取原始x数据升序排列的索引x_sorted_orig, y_sorted_orig = x.iloc[sorted_indices_orig].values, y[sorted_indices_orig] # 对原始x和y数据进行排序main_smoothed = lowess(y_sorted_orig, x_sorted_orig, frac=frac) # LOWESS平滑lower_bound = np.quantile(boot_lines_arr, alpha, axis=0) # 置信下界upper_bound = np.quantile(boot_lines_arr, 1 - alpha, axis=0) # 置信上界return main_smoothed, (x_range, lower_bound, upper_bound) # 拟合曲线、x范围、置信上下界

第四部分

# =========================================================================================# ======================================6.阈值点寻找函数=======================================# =========================================================================================# 寻找曲线穿过y=0的所有交点/X坐标。通过检测y值正负号的变化,利用两点线性插值法精确计算出过零点的x坐标def find_roots(x_curve, y_curve):roots = [] # 存放根/零点的列表sign_changes = np.where(np.diff(np.sign(y_curve)))[0] # 计算y值的符号差,找出正负号发生变化的相邻点索引x_root = x1 - y1 * (x2 - x1) / (y2 - y1) # 计算出零点x坐标roots.append(x_root) # 保存return roots # 返回所有找到的零点列表

第五部分

# =========================================================================================# ======================================7.阈值点绘制以及标注函数=======================================# =========================================================================================# 算出零点,还在图表上画出垂直虚线,标上数值标签def find_and_plot_crossings(ax, x_curve, y_curve, color, x_range):ax.text(x_root, # xy_pos, # yf' {x_root:.2f} ', # 文本color='white', # 颜色backgroundcolor=color, # 颜色ha='center', # 水平va='top', # 垂直fontsize=18, # 字体大小fontweight='bold', # 加粗bbox=dict(facecolor=color, edgecolor='none', pad=1), # 文本框transform=ax.get_xaxis_transform()) # 设置坐标变换drawn_texts.append((x_root, y_pos)) # 保存

第六部分




# =========================================================================================# ==============================8.分类评估图=======================================# =========================================================================================def plot_classification_results(metrics, colors, output_folder, n_classes):y_train_true, y_train_prob = metrics['train']['true'], metrics['train']['prob'] #训练集真实标签和预测概率y_test_true, y_test_prob = metrics['test']['true'], metrics['test']['prob'] #测试集真实标签和预测概率ax_roc.plot([0, 1], [0, 1], 'k--', lw=2, label='Random Chance') #绘制参考对角线ax_roc.set_xlabel('False Positive Rate', fontsize=18) #x轴标题ax_roc.set_ylabel('True Positive Rate', fontsize=18) #y轴标题ax_roc.set_title('ROC Curve', fontsize=22) #主标题ax_roc.legend(loc='lower right', fontsize=14) #图例ax_roc.grid(True) #网格线apply_plot_styles(ax_roc) #坐标轴边框设置save_fig_dual(fig_roc, output_folder, 'classification_roc_curve') #保存ax_cm.set_xlabel('Predicted Label', fontsize=18) #x轴标题ax_cm.set_ylabel('True Label', fontsize=18) #y轴标题ax_cm.set_title('Confusion Matrix (Validation)', fontsize=22) #主标题ax_bar.legend(fontsize=14, loc='lower right') # 图例ax_bar.grid(axis='y', linestyle='--', alpha=0.7) #网格线apply_plot_styles(ax_bar) #坐标轴设置save_fig_dual(fig_bar, output_folder, 'classification_metrics_bar') #保存plt.close(fig_bar) #关闭

第七部分


# =========================================================================================# ======================================9.特征重要性条形图与SHAP蜂巢图组合图绘制函数=======================================# =========================================================================================def plot_shap_summary(shap_values, X_test, feature_names, shap_df, base_values, colors, output_folder):# 创建画布fig = plt.figure(figsize=(10, 10), dpi=300)ax_sw = fig.add_axes([0.32, 0.11, 0.59, 0.77]) # 定义主坐标轴在画布上的相对位置及大小ax_bar = ax_sw.twiny() # 创建共享Y轴apply_plot_styles(ax_sw) # 调整边框和刻度粗细apply_plot_styles(ax_bar) # 调整边框和刻度粗细colorbar_ax = fig.axes[-1] # 获取颜色条轴colorbar_ax.tick_params(labelsize=16) # 刻度设置colorbar_ax.set_ylabel("", labelpad=0) # 去掉原始标题# 保存save_fig_dual(fig, output_folder, 'combined_shap_summary_plot')plt.close(fig) # 关闭

第八部分


# =========================================================================================# ======================================10.特征交互强度气泡热力图绘制函数=======================================# =========================================================================================def plot_bubble_heatmap(shap_interaction, feature_names, colors, output_folder):n_features = len(feature_names) # 获取特征总数inter_matrix = np.zeros((n_features, n_features)) # 用于存储两两交互强度# y轴标题ax.set_ylabel("Trigger thresholds\nDriving factors", # 文本fontsize=18, # 字体大小fontweight='bold', # 加粗labelpad=-2) # 间距# 子图编号ax.text(0.08, # x0.84, # y'(a)', # 编号transform=ax.transAxes, # 坐标系fontsize=28, # 字体大小fontweight='bold') # 加粗cax = ax.inset_axes([0.05, 0.05, 0.03, 0.35]) # 创建颜色条轴cbar.ax.set_yticklabels([f'{v:.3f}' for v in cbar_ticks], fontsize=14, fontweight='bold')y_positions = [0.4, 0.283, 0.166, 0.05] # 气泡图例位置save_fig_dual(fig, output_folder, 'shap_interaction_bubble_heatmap_diag_zero') # 保存plt.close(fig) # 关闭

第九部分


# =========================================================================================# ======================================11.SHAP单因素依赖图绘制函数=======================================# =========================================================================================def plot_dependence(X_test, shap_values, feature_names, colors, save_folder):n_features = min(9, len(feature_names)) # 只选择前9个最重要特征进行绘制# 创建画布fig, axes = plt.subplots(3, 3, figsize=(15, 12))axes = axes.flatten() # 展平为一维数组,便于按顺序循环遍历调用labels = [chr(97 + i) for i in range(n_features)] # 生成子图编号ax1.set_ylim(0, counts.max() * 1.1) # 左侧y轴范围# 绘制散点ax2.scatter(x_values, # xshap_vals, # yalpha=0.7, # 透明度s=25, # 散点大小color=colors['shap_scatter'], # 颜色label='Sample', # 图例标签zorder=2) #find_and_plot_crossings(ax2, main_fit[:, 0], main_fit[:, 1], 'black', x_range) # 寻找阈值并绘制ax1.set_xlabel(f'{feature_name}', fontsize=18) # x轴标题h1, l1 = ax1.get_legend_handles_labels() # 提取直方图的图例句柄和标签h2, l2 = ax2.get_legend_handles_labels() # 提取散点、拟合线图例句柄和标签# 添加图例

第十部分


# =========================================================================================# ======================================12.特征交互效应依赖图绘制函数=======================================# =========================================================================================def plot_interaction(X_test, shap_interaction, feature_names, colors, save_folder):color='black', # 颜色linestyle='--', # 样式lw=2, # 粗细zorder=0) # 层y_lim = max(np.abs(shap_vals).max() * 1.1, 0.1) # 右侧最大值ax2.set_ylim(-y_lim, y_lim) # 右侧y轴范围# 颜色条轴cax_auto = ax2.inset_axes([1.24, 0.0, 0.04, 1.0])# 创建颜色条cbar = fig.colorbar(points, cax=cax_auto)# 标题cbar.set_label(s_name, size=18, labelpad=5)# 刻度设置cbar.ax.tick_params(labelsize=18)cbar.ax.tick_params(axis='y', width=2, length=4, direction='in')# 控制颜色条自身的外边框线宽for spine in cbar.ax.spines.values():spine.set_linewidth(1.5)ax2.legend(h2 + h1, l2 + l1, loc='lower right', fontsize=13)plt.tight_layout() # 调整布局save_fig_dual(fig, save_folder, 'interaction_top9_grid') # 保存plt.close(fig) # 关闭

第十一部分

记录训练集和测试集的性能指标。调用 shap.Explainer 结合树模型的高效TreeSHAP算法,计算测试集全部样本的特征主效应SHAP值以及二阶交互作用。计算SHAP绝对均值并按重要性从大到小对特征名称、数据矩阵、SHAP矩阵和交互矩阵进行重排序,调用上面的函数分析绘图。# =========================================================================================# ======================================13.执行部分 =======================================# =========================================================================================if __name__ == '__main__':output_folder = r'F:\公众号素材\20260728shap重要性+依赖图+交互效应图-分类任务' # 结果输出路径test_size = 0.3 # 测试集比例random_state = 0 # 随机种子os.makedirs(output_folder, exist_ok=True) # 是否存在,若无则自动创建一个新文件夹# 划分数据X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=test_size, random_state=random_state)# 超参数param_grid = {'n_estimators': [50, 100, 200],'max_depth': [2, 3, 5],'learning_rate': [0.05, 0.1]}# 实例化XGBoost分类模型xgb_model = xgb.XGBClassifier(random_state=random_state)# 配置网格搜索grid_search = GridSearchCV(estimator=xgb_model, param_grid=param_grid, scoring='f1_macro', cv=3, n_jobs=-1)print("\n正在计算SHAP 值") ascending=False) #特征重要性表并降序排列sorted_features = shap_df["feature"].values.tolist() #提取排好序的特征名列表sorted_indices = [feature_names.index(f) for f in sorted_features] #获取重排后的特征原索引X_test_sorted = X_test[sorted_features] #依重要性重排测试集特征shap_values_sorted = shap_values[:, sorted_indices] #依重要性重排SHAP值shap_interaction_sorted = shap_interaction[:, sorted_indices][:, :, sorted_indices] #依重要性重排交互值

如何应用到你自己的数据

1.设置文件夹地址,执行部分:
output_folder = r'分类任务'2.设置测试集比例,执行部分:
test_size = 0.3 # 测试集比例3.设置随机种子,执行部分:
random_state = 0 # 随机种子4.设置目标变量,执行部分:
target_column_name = 'Vegetation_Anomaly' # 目标变量名5.设置超参数网格,执行部分:
param_grid = {'n_estimators': [50, 100, 200],'max_depth': [2, 3, 5],'learning_rate': [0.05, 0.1]}
6.设置要解释的类别,执行部分:
EXPLAIN_CLASSES = 'all'7.设置是否进行批量绘图,执行部分:
plot_all = True
推荐


获取方式
