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

线性回归中Train-Test拆分导致SHAP部分依赖图错位问题

SHAP部分依赖图错位问题的原因与解决方法

问题原因

错位的核心是基准值不匹配:

  • 用X_train初始化SHAP解释器时,SHAP计算的基准值(expected value)是训练集的预测结果均值。
  • 但调用shap.partial_dependence_plot时,若未显式指定model_expected_value,函数会默认计算传入的X_test的预测均值作为基准值。
  • 训练集和测试集的预测均值存在差异,导致SHAP值对应的黑点(样本实际贡献)与部分依赖线(基于测试集基准的预期值)在Y轴上错位。
  • 当用完整数据集X初始化解释器时,基准值是全数据集的预测均值,和绘图时默认计算的基准值一致,因此显示正常。

解决方法

要在Train-Test拆分场景下修正图表,需保证SHAP解释器的基准值与部分依赖图的基准值完全统一,以下是两种可行方案:

方案一:显式指定训练集的预测均值作为绘图基准值

计算训练集的预测均值,将其传入partial_dependence_plot的model_expected_value参数,强制对齐基准:

# 计算训练集的预测均值(即SHAP解释器的基准值)
train_pred_mean = model.predict(X_train).mean()

# 修正后的绘图代码
shap.partial_dependence_plot(
    "cement", model.predict, X_test,
    model_expected_value=train_pred_mean,  # 显式指定训练集基准值
    feature_expected_value=True, ice=False,
    shap_values=shap_values[idx:idx+1,:]
)

方案二:用训练集作为部分依赖图的背景数据集

部分依赖图的本质是展示特征在数据分布上的边际效应,用训练集(模型训练时的分布)作为背景更合理,同时基准值会自动匹配:

# 修正后的绘图代码,将X_test替换为X_train
shap.partial_dependence_plot(
    "cement", model.predict, X_train,  # 用训练集作为背景数据集
    model_expected_value=True, feature_expected_value=True, ice=False,
    shap_values=shap_values[idx:idx+1,:]
)

完整修正代码示例

import shap
from sklearn.model_selection import train_test_split
from sklearn.linear_model import LinearRegression
import pandas as pd
import matplotlib.pyplot as plt
import requests
import zipfile
from io import BytesIO

def load_data() -> pd.DataFrame:
    """
    Loads and returns the dataset from the given URL as a Pandas DataFrame.

    Returns:
        pd.DataFrame: The loaded dataset.
    """
    url = "https://archive.ics.uci.edu/static/public/165/concrete+compressive+strength.zip"

    r = requests.get(url)

    if r.ok:
        with zipfile.ZipFile(BytesIO(r.content)) as thezip:
            with thezip.open("Concrete_Data.xls") as thefile:
                return pd.read_excel(thefile, header=0)
    else:
        raise Exception("Something went wrong.")

df = load_data()

df = df.rename(
    columns={
        'Cement (component 1)(kg in a m^3 mixture)':'cement',
        'Blast Furnace Slag (component 2)(kg in a m^3 mixture)':'blast',
        'Fly Ash (component 3)(kg in a m^3 mixture)':'ash',
        'Water  (component 4)(kg in a m^3 mixture)':'water',
        'Superplasticizer (component 5)(kg in a m^3 mixture)':'superplasticizer',
        'Coarse Aggregate  (component 6)(kg in a m^3 mixture)':'coarse',
        'Fine Aggregate (component 7)(kg in a m^3 mixture)':'fine',
        'Age (day)':'age',
        'Concrete compressive strength(MPa, megapascals) ': 'strength'
    }
)
df = df.drop_duplicates()
X = df.drop(['strength'], axis=1) 
y = df['strength']

# Split the data
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)

# Train the model
model = LinearRegression()
model.fit(X_train, y_train)

# Initialize SHAP explainer and calculate values for the test set
explainer = shap.Explainer(model.predict, X_train)
shap_values = explainer(X_test)

# 方案一:显式指定训练集预测均值作为基准
train_pred_mean = model.predict(X_train).mean()
idx = 0
shap.partial_dependence_plot(
    "cement", model.predict, X_test,
    model_expected_value=train_pred_mean,
    feature_expected_value=True, ice=False,
    shap_values=shap_values[idx:idx+1,:]
)

# Save the plot
plt.tight_layout()
plt.savefig('shap_dependence_plot.png', dpi=300)

内容的提问来源于stack exchange,提问作者Simon Symhoven

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 01:27:42