线性回归中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
相关产品推荐
相关产品推荐

