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

交叉验证计算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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 21:42:20