交叉验证计算SHAP值时numpy拼接应选axis=0还是axis=1?
代码正确性说明
你的实现整体思路成立,但交叉验证结果拼接部分的axis设置是旧版SHAP二分类场景下的巧合适配,换运行环境或任务类型就会出问题,具体逻辑和正确写法如下:
两个拼接操作axis不一致的原因
两个待拼接对象的维度结构完全不同,所以拼接轴不一样:
test_index是一维数组,每个元素对应测试集样本在原始数据集里的行号。拼接多折的测试集索引,本质是把多段样本的索引串成覆盖全量样本的索引列表,一维数组只有0轴,因此用axis=0拼接是完全正确的,这部分逻辑没有问题。- 0.40版本以前的SHAP库中,
TreeExplainer针对二分类树模型输出的shap_values是三维数组,维度顺序为[类别数量, 折内测试样本数, 特征数]。以你设置的2折交叉验证为例,每折测试集样本量是总样本的1/2,单折shap值形状为(2, 总样本数/2, 特征数),要把两折的样本维度拼合,对应操作的就是第1轴(样本所在轴),因此写axis=1在「二分类任务+旧版SHAP」的特定场景下刚好可以跑通,拼接后形状为(2, 总样本数, 特征数),和你后续取shap_values[1,:,:]计算正类SHAP均值的逻辑是匹配的。
现有代码的潜在问题
该写法通用性极差,遇到以下场景会直接报错或输出错误结果:
- 如果你安装的是SHAP 0.40及以上版本,二分类任务默认直接返回正类对应的二维SHAP数组,形状为
[样本数, 特征数],不存在类别维度,此时按axis=1拼接会直接触发维度不匹配报错。 - 如果是多分类任务,SHAP输出的数组维度顺序变化,按
axis=1拼接会混淆样本和特征维度,最终得到的特征重要性结果完全错误。 - 如果修改折数、或数据集样本量不能被折数整除,单折测试集样本量不一致时,拼接逻辑很容易出现维度错位。
通用稳妥实现
不要依赖SHAP的默认输出维度结构,每折计算完SHAP值后就提取目标类别的结果,统一整理为[样本数, 特征数]的二维数组,拼接时和测试集索引一样统一按样本轴(axis=0)拼接,从根源上避免维度混淆:
from sklearn.model_selection import KFold from sklearn.datasets import load_breast_cancer from sklearn.ensemble import RandomForestClassifier import shap import pandas as pd import numpy as np # 加载数据集 dataset = load_breast_cancer() X = dataset.data y = dataset.target feature_names = dataset.feature_names # 交叉验证拆分固定随机种子保证可复现 kf = KFold(n_splits=5, shuffle=True, random_state=0) shap_values_collection = [] test_indices_collection = [] for train_idx, test_idx in kf.split(X): X_train, X_test = X[train_idx], X[test_idx] y_train = y[train_idx] # 训练模型 clf = RandomForestClassifier(random_state=0) clf.fit(X_train, y_train) # 计算SHAP值 explainer = shap.TreeExplainer(clf) shap_vals = explainer.shap_values(X_test) # 兼容不同SHAP版本、不同分类任务,统一提取目标类别的二维SHAP矩阵 if isinstance(shap_vals, list): # 旧版SHAP返回列表,每个元素对应一个类别的(样本数,特征数)二维数组,二分类正类对应索引1 target_class_shap = shap_vals[1] else: # 新版SHAP二分类直接返回正类的二维SHAP数组 target_class_shap = shap_vals shap_values_collection.append(target_class_shap) test_indices_collection.append(test_idx) # 所有结果统一按样本维度(axis=0)拼接,逻辑一致不易出错 full_test_indices = np.concatenate(test_indices_collection, axis=0) full_shap_matrix = np.concatenate(shap_values_collection, axis=0) # 计算全局特征重要性(SHAP绝对值均值) feature_importance = pd.DataFrame({ "feature_name": feature_names, "mean_abs_shap": np.abs(full_shap_matrix).mean(axis=0) }).sort_values("mean_abs_shap", ascending=False).reset_index(drop=True) print(feature_importance)
其他细节修正
- 原代码中加载乳腺癌数据集却将变量命名为
iris属于笔误,不影响运行但容易造成误解,已修正。 - 原代码中
KFold未设置固定随机种子,每次运行拆分结果不一致,可复现性差,已补充random_state参数。 - 原代码中通过SHAP绝对值均值计算全局特征重要性的逻辑是正确的,符合SHAP官方推荐的全局重要性计算方式。
内容的提问来源于stack exchange,提问作者Slowat_Kela
相关产品推荐
相关产品推荐

