You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用Plotly复刻SHAP Summary Plot?数据与配色问题求助

用Plotly复刻SHAP Summary Plot的问题解答

1. Shap_FI.csv对应的DataFrame构建逻辑

这个DataFrame是**长格式(long-form)**结构,每一行对应「单个样本的单个特征」的SHAP相关数据,必须包含以下4列:

  • transaction_id:样本唯一标识(比如数据集里的样本ID,用来区分不同数据点)
  • feature:特征名称(比如"年龄""收入"这类)
  • shap_value:该样本对应特征的SHAP值(由SHAP库计算得出)
  • feature_value:该特征的原始取值(用来对应红蓝配色)

举个Python代码示例,假设你已经用SHAP算出了结果:

import pandas as pd
import shap

# 假设X是你的特征数据集(宽格式,每行一个样本,每列一个特征)
explainer = shap.TreeExplainer(model)
shap_values = explainer.shap_values(X)

# 转换成长格式DataFrame
shap_df = pd.DataFrame()
for i, feature in enumerate(X.columns):
    temp = pd.DataFrame({
        "transaction_id": X.index,  # 用原数据集索引当样本ID,或者自定义ID
        "feature": feature,
        "shap_value": shap_values[:, i],
        "feature_value": X[feature]
    })
    shap_df = pd.concat([shap_df, temp], ignore_index=True)

# 保存成csv就是Shap_FI.csv
shap_df.to_csv("Shap_FI.csv", index=False)

2. SHAP特有的红蓝配色实现

SHAP图里的红蓝是特征原始值的高低映射:红色对应特征值高,蓝色对应特征值低。在Plotly里可以这么实现:

  • 用散点图(px.scatter或者go.Scatter)绘制,把color参数绑定到feature_value
  • 设置颜色渐变尺度(colorscale)为从蓝色到红色的渐变,比如colorscale=['#4477AA', '#EE6677'](接近SHAP原生配色)
  • 可以添加颜色条(colorbar)并标注说明,让读者明确颜色对应特征值的高低

示例代码片段:

import plotly.express as px

fig = px.scatter(
    shap_df,
    x="shap_value",
    y="feature",
    color="feature_value",
    colorscale=['#4477AA', '#EE6677'],
    color_continuous_midpoint=shap_df["feature_value"].median(),  # 可选,让中点对应中位数
    title="SHAP Summary Plot (Plotly复刻版)"
)
fig.update_layout(coloraxis_colorbar=dict(title="特征原始值"))
fig.show()

3. transaction_id字段的作用

这个字段就是样本的唯一ID,核心作用是:

  • 在长格式DataFrame里,把同一个样本的所有特征行关联起来(比如一个样本有10个特征,就会拆成10行,这10行的transaction_id是一致的)
  • 绘图时其实不需要直接用到它,但构建DataFrame时必须有,用来保证每行是「唯一样本-唯一特征」的组合,避免数据混乱。如果你的数据集本身有样本ID,直接用就行;没有的话用原数据集的索引当ID也完全可以。

内容的提问来源于stack exchange,提问作者Julia Julchen

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.22 03:52:12