当前位置:首页>python>Python绘制shap交互特征作用网络图

Python绘制shap交互特征作用网络图

  • 2026-10-11 06:17:52
Python绘制shap交互特征作用网络图

代码绘制成果展示

根据读者建议,对之前的那张网络图做了一些改进。左上角的网络线图例用线条粗细来表示交互作用强度,线条越粗说明影响越大,数值范围表示最大、最小、中间值;左下角的节点的图例通过节点的大小来表示特征重要性,点越大影响越强,用来表示最大、最小、中间值;右上的颜色条表示网络连接线的颜色,用来表示交互作用的影响的方向,就是平均值,由深紫色(负向)过渡至深绿色(正向);右下的颜色条用来表示节点的颜色,代表单一特征作用的影响方向,平均值,由深蓝色过渡至深红色。
多种配色方案
多种形状标记方案

代码解释

第一部分

库的导入以及字体设置
# =========================================================================================# ====================================== 1. 环境设置 =======================================# =========================================================================================import numpy as npimport pandas as pdimport xgboost as xgbimport shapimport matplotlib.pyplot as pltimport matplotlib.colors as mcolorsimport networkx as nximport warningsfrom sklearn.model_selection import train_test_split, GridSearchCVfrom matplotlib.lines import Line2Dwarnings.filterwarnings("ignore", category=DeprecationWarning)warnings.filterwarnings("ignore", category=UserWarning)import matplotlibmatplotlib.rcParams['pdf.fonttype'] = 42matplotlib.rcParams['ps.fonttype'] = 42plt.rcParams['font.family'] = 'serif'plt.rcParams['font.serif'] = ['Times New Roman']plt.rcParams['axes.unicode_minus'] = False

第二部分

颜色库的设置以及配色方案的选择
# =========================================================================================# ======================================2.颜色库=======================================# =========================================================================================COLOR_SCHEMES = {    1: {'nodes': plt.cm.RdBu_r, 'edges': plt.cm.PRGn},}scheme_index = 20  # 颜色方案选择# 获取当前颜色方案current_color_scheme = COLOR_SCHEMES.get(scheme_index, COLOR_SCHEMES[1])

第三部分

形状标记库的设置以及配色方案的选择
# =========================================================================================# ======================================3.形状标记库=======================================# =========================================================================================STYLE_SCHEMES = {    1: {'marker': 'o', 'linestyle': '-'},}style_index = 1 # 形状标记方案# 获取样式方案current_style_scheme = STYLE_SCHEMES.get(style_index, STYLE_SCHEMES[1])

第四部分

数据的读取以及目标变量与特征变量的分离
# =========================================================================================# ======================================4.数据加载=======================================# =========================================================================================# 原始数据路径file_path = r'mock_data.xlsx'# 读取数据df = pd.read_excel(file_path)# 目标变量y = df.iloc[:, -1]# 特征变量X = df.iloc[:, :-1]# 获取特征列的名称并转换为列表features = X.columns.tolist()print(f"特征: {features}")print(f"数据类型: {X.shape}")

第五部分

数据集划分,超参数网格的设置,模型的训练,最佳模型的获取
# =========================================================================================# ======================================5.数据划分及模型构建=======================================# =========================================================================================# 划分训练集和测试集X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)# 超参数网格param_grid = {    'max_depth': [4, 6, 8],    # 'learning_rate': [0.05, 0.1, 0.2],    # 'n_estimators': [50, 100, 150]}# 初始化XGBoost 回归模型xgb_model = xgb.XGBRegressor(random_state=42, n_jobs=-1)# 初始化网格搜索对象grid_search = GridSearchCV(estimator=xgb_model, param_grid=param_grid, cv=5, scoring='neg_mean_squared_error',verbose=1)# 在训练集上拟合grid_search.fit(X_train, y_train)print(f"最佳参数: {grid_search.best_params_}")# 获取最佳模型best_model = grid_search.best_estimator_

第六部分

shap分析,特征重要性数据获取,交互作用数据获取,用于设置节点的大小、网络线的粗细,节点的颜色,网络线的颜色
# =========================================================================================# ======================================6.SHAP分析=======================================# =========================================================================================# 使用最佳模型创建SHAP树解释器explainer = shap.TreeExplainer(best_model)# 测试集的SHAP交互值shap_interaction_values = explainer.shap_interaction_values(X_test)# 测试集的SHAP值shap_values = explainer.shap_values(X_test)#节点数据#特征重要性,平均绝对值,点大小feature_importance_abs = np.abs(shap_values.mean(axis=0))#特征影响方向性,实际值平均,点颜色feature_importance_signed = shap_values.mean(axis=0)#边线数据#交互强度,平均绝对值,用于控制连线粗细mean_interaction_matrix_abs = np.abs(shap_interaction_values.mean(axis=0))np.fill_diagonal(mean_interaction_matrix_abs, 0)  # 忽略自身交互#交互影响方向,实际值平均,用于控制连线颜色mean_interaction_matrix_signed = shap_interaction_values.mean(axis=0)np.fill_diagonal(mean_interaction_matrix_signed, 0)  # 忽略自身交互

第七部分

创建画布,获取配色方案、形状标记
def plot_circular_interaction(features, importance_abs, importance_signed,interaction_matrix_abs, interaction_matrix_signed):    #获取颜色方案    cmap_nodes = current_color_scheme['nodes']    cmap_edges = current_color_scheme['edges']    # 获取节点形状标记    node_marker = current_style_scheme['marker']    # 获取连线样式    edge_linestyle = current_style_scheme['linestyle']    # 创建画布    fig, ax = plt.subplots(figsize=(12, 10), subplot_kw={'aspect': 'equal'})    # 获取特征的数量    n_features = len(features)

第八部分

利用 NetworkX 库来确定每个特征节点在画面上的位置
    # 向图中添加节点    G.add_nodes_from(features)    # 生成节点的环形布局坐标    pos = nx.circular_layout(G)    # 标签的坐标    label_pos = {k: (v * 1.1) for k, v in pos.items()}    # 颜色归一化    norm_edges = mcolors.Normalize(vmin=interaction_matrix_signed.min(),                                   vmax=interaction_matrix_signed.max())    #宽度/大小归一化基准    max_interaction_abs = np.max(interaction_matrix_abs)    max_importance_abs = np.max(importance_abs)

第九部分

利用 NetworkX 库来确定每个特征节点在画面上的位置
# 初始化交互列表    interactions = []    # 遍历特征    for i in range(n_features):        for j in range(i + 1, n_features):            # 如果绝对强度大于 0 (显示阈值)            if strength_abs > 0:                # 将交互对、绝对强度、实际强度添加到列表中                interactions.append((features[i], features[j], strength_abs, strength_signed))    # 根据绝对强度对交互列表进行排序    interactions.sort(key=lambda x: x[2])

第十部分

绘制网络连线
    # 遍历排序后的交互列表    for u, v, strength_abs, strength_signed in interactions:        # 修正点:根据实际值和当前边颜色方案获取线的颜色        color = cmap_edges(norm_edges(strength_signed))        # 根据绝对值计算线的粗细        width = 0.5 + (strength_abs / max_interaction_abs) * 8        # 线的透明度        alpha = 0.3 + (strength_abs / max_interaction_abs) * 0.7        # 绘制线        nx.draw_networkx_edges(G,                               pos,                               edgelist=[(u, v)],                               width=width,                               edge_color=[color],                               style=edge_linestyle,                               alpha=alpha, ax=ax)

第十一部分

绘制节点
    # --- 节点的处理 ---    # 节点颜色归一化    norm_nodes = mcolors.Normalize(vmin=importance_signed.min(),                                   vmax=importance_signed.max())    # 遍历每个特征    for i, feat in enumerate(features):        # 获取该特征的实际值        imp_sign = importance_signed[i]        # 获取该特征的绝对重要性        imp_abs = importance_abs[i]        #计算并添加节点颜色        node_colors.append(cmap_nodes(norm_nodes(imp_sign)))

第十二部分

绘制节点位置的特征名称标注
       # 绘制标签文本        plt.text(x,                 y,                 node,                 size=12,                 horizontalalignment=ha,                 verticalalignment='center')    # 关闭坐标轴    ax.axis('off')    # x轴显示范围    ax.set_xlim(-1.5, 1.5)    # y轴显示范围    ax.set_ylim(-1.5, 1.5)    # 标题    plt.title('(a) Green Ecological -> Agricultural Production', y=0.95, fontsize=16)

第十三部分

设置图例
    # ---------------------------左侧图例-----------------------    #线条粗细图例的三个等级数值    line_levels = [max_interaction_abs, max_interaction_abs * 0.5, max_interaction_abs * 0.1]    #将数值转换为字符串标签,用于图例显示具体数值    line_labels = [f"{val:.2f}" for val in line_levels]    legend1 = ax.legend(legend_lines,                        line_labels,  # 传入格式化后的数值标签列表                        loc='center left',  #左侧居中                        bbox_to_anchor=(-0.1, 0.8),  #精确位置                        title="Interaction Strength\n(Line Width)",  #图例标题                        frameon=False,  #去掉图例边框                        labelspacing=1.5)  #图例垂直间距    #添加到轴上    ax.add_artist(legend1)    #定义节点大小等级数值    node_levels = [max_importance_abs,                   max_importance_abs * 0.5,                   max_importance_abs * 0.1]    #用于图例显示具体数值    node_labels = [f"{val:.2f}" for val in node_levels]        legend_nodes.append(Line2D([0],                                   [0],                                   marker=node_marker,  #标记形状                                   color='w',  #线条颜色                                   markerfacecolor='black',  #点的填充颜色                                   markersize=s,  #标记点的大小                                   linestyle='None'))  #不绘制连接线,只显示点    #添加点图例    ax.legend(legend_nodes,              node_labels,  #数值标签列表              loc='center left',  #左侧居中              bbox_to_anchor=(-0.1, 0.32),  #精确位置              title="Feature Importance\n(Node Size)",  #图例标题              frameon=False,  #去掉图例边框              labelspacing=3)  #图例垂直间距

第十四部分

设置颜色条
    # --- 颜色条 ---    #定义边颜色条的位置,左,下,宽,高    cbar_edge_pos = [0.82, 0.55, 0.015, 0.25]    # 创建一个新的轴用于放颜色条    cax_edge = fig.add_axes(cbar_edge_pos)    #设置线的颜色条的标签    cbar_edge.set_label('Interaction Value (Signed)', rotation=270, labelpad=15, fontsize=10)    #去掉线的颜色条的轮廓线    cbar_edge.outline.set_visible(False)    # --- 颜色条---    #节点颜色条的位置    cbar_node_pos = [0.82, 0.20, 0.015, 0.25]    #绘制节点颜色条    cbar_node = plt.colorbar(sm_node, cax=cax_node)    #设置节点颜色条的标签    cbar_node.set_label('Feature Value (Signed)', rotation=270, labelpad=15, fontsize=10)    #去掉节点颜色条的轮廓线    cbar_node.outline.set_visible(False)    # 保存    save_path_png = fr"{style_index}_scheme{scheme_index}.png"    save_path_pdf = fr"{style_index}_scheme{scheme_index}.pdf"    plt.savefig(save_path_png, dpi=300, bbox_inches='tight')    plt.savefig(save_path_pdf, bbox_inches='tight')

第十五部分

执行部分,打印分析结果,进行绘图
# =========================================================================================# ======================================8.主程序执行部分=======================================# =========================================================================================if __name__ == "__main__":    print("-" * 30)    print("特征重要性排序")    print("-" * 30)    #创建DataFrame对象,用于展示分析结果    df_importance = pd.DataFrame({        '特征': features,  #特征列        '重要性 (平均绝对值)': feature_importance_abs,  #重要性        '影响方向 (平均值)': feature_importance_signed  #影响方向    })    #根据重要性进行降序排序    df_importance = df_importance.sort_values(by='重要性 (平均绝对值)', ascending=False)    print(df_importance.to_string(index=False))    print("-" * 30)    print("SHAP 交互作用强度排序")    print("-" * 30)    # 初始化一个空列表,用于后续存储筛选出来的交互作用数据字典    interaction_list = []    # 获取特征的总数量,用于控制后续循环的次数    n_features = len(features)    # 开始外层循环,遍历每一个特征的索引 i,范围从 0 到 特征总数-1    for i in range(n_features):        # 开始内层循环,遍历i之后的每一个特征索引j,确保只计算组合(不重复计算且不含自身)        for j in range(i + 1, n_features):            # 从平均交互矩阵的绝对值中,获取第 i 个和第 j 个特征之间的交互强度            strength = mean_interaction_matrix_abs[i, j]            # 从平均交互矩阵的原始值中,获取第 i 个和第 j 个特征之间的交互方向(正负)            direction = mean_interaction_matrix_signed[i, j]            # 条件判断:如果交互强度大于0(即存在有效的交互作用),则执行以下代码块            if strength > 0:                #向interaction_list列表中追加一个包含当前交互对详细信息的字典                interaction_list.append({                    '特征 1': features[i],  # 记录第一个特征的名称                    '特征 2': features[j],  # 记录第二个特征的名称                    '交互作用强度 (平均绝对值)': strength,  # 记录该对特征的交互强度                    '交互作用影响方向 (平均值)': direction  # 记录该对特征的交互方向                })    # 将收集了所有交互信息的列表转换为一个Pandas DataFrame,方便后续处理    df_interactions = pd.DataFrame(interaction_list)    #如果df_interactions不为空,即找到了至少一个交互作用    if not df_interactions.empty:        # 根据交互作用强度这一列进行降序排序,让强交互排在前面        df_interactions = df_interactions.sort_values(by='交互作用强度 (平均绝对值)', ascending=False)        print(df_interactions.head(15).to_string(index=False))    else:        print("无显著交互作用。")    #调用绘图函数    plot_circular_interaction(features,#特征                              feature_importance_abs,  #节点大小依据                              feature_importance_signed,  #节点颜色依据                              mean_interaction_matrix_abs,  #连线粗细依据                              mean_interaction_matrix_signed)  #连线颜色依据

如何应用

1.选择你想要使用到的配色方案:

scheme_index = 1

2.选择你想要使用到的形状标记方案:

style_index = 1 # 形状标记方案

3.设置原始数据的路径:

file_path = r'mock_data.xlsx'

4.定义数据的目标变量和特征变量:

# 目标变量y = df.iloc[:, -1]# 特征变量X = df.iloc[:, :-1]

5.设置模型的超参数网格:

param_grid = {    'max_depth': [4, 6, 8],    'learning_rate': [0.05, 0.1, 0.2],   'n_estimators': [50, 100, 150]}

6.设置绘图结果的保存路径:

save_path_png = fr"{style_index}_scheme{scheme_index}.png"save_path_pdf = fr"{style_index}_scheme{scheme_index}.pdf"

推荐

期刊图片复现|Python绘制二维偏依赖PDP图
期刊复现|python绘制基于SHAP分析和GAM模型拟合的单特征依赖图
期刊图片复现|python绘制带有渐变颜色shap特征重要性组合图(条形图+蜂巢图)
期刊复现|用Python绘制SHAP特征重要性总览图、依赖图、双特征交互效应SHAP图,解锁XGBoost模型的终极奥秘
期刊图片复现|Python绘制shap重要性蜂巢图+单特征依赖图+交互效应强度气泡图+交互效应依赖图(回归+二分类+分类)

获取方式

需要的请后台私信我,注意只会分享练习数据和代码文件,不会提供答疑服务,代码文件中已经包含了每行代码的完整注释!!!

最新文章

随机文章