如何用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
相关产品推荐
相关产品推荐

