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

基于SHAP值生成违约概率模型Top3原因码:Pandas适配求助

解决SHAP生成原因码时feature_names索引报错的问题

核心问题分析

报错本质是sorted_indices中的索引值超出了feature_names列表长度,或是feature_names的顺序/数量和模型训练时的输入特征不匹配——原教程用ndarray时特征顺序固定,但Pandas DataFrame的列顺序若和模型训练时不一致,就会导致索引对应错误。

分步解决方案

1. 确保feature_names和模型输入特征完全匹配

直接从训练用的DataFrame中提取特征名,避免手动编写列表导致的顺序/数量错误:

# 假设df_train是训练模型用的DataFrame,target是标签列
feature_names = df_train.drop('target', axis=1).columns.tolist()

如果训练时只用到了部分列,要严格对应训练代码中传入模型的列顺序,比如:

# 训练时指定的特征列
used_features = ['age', 'income', 'loan_amount', 'credit_score']
feature_names = used_features

2. 修正SHAP值排序与索引匹配逻辑

针对生成Top3负面原因码(即SHAP值为负、拉低违约概率的特征),调整代码逻辑确保索引不会越界:

import numpy as np
import shap

# 假设model是训练好的GMB模型,df_test是测试数据集
explainer = shap.TreeExplainer(model)
# 二分类模型的shap_values通常是(2, n_samples, n_features),取对应违约类的SHAP值(一般是索引1)
shap_values = explainer.shap_values(df_test.drop('target', axis=1))[1]

# 对单个样本生成Top3负面原因码(以第0个样本为例)
sample_idx = 0
sample_shap = shap_values[sample_idx]

# 筛选出SHAP值为负的特征(负面贡献)
negative_mask = sample_shap < 0
negative_indices = np.where(negative_mask)[0]
negative_shap_values = sample_shap[negative_mask]

# 按SHAP值从小到大排序(最负的特征对违约概率的拉低作用最强)
sorted_indices = negative_indices[np.argsort(negative_shap_values)]

# 取Top3,避免样本负面特征不足3个的情况
top3_indices = sorted_indices[:3] if len(sorted_indices) >=3 else sorted_indices

# 映射到特征名
top3_negative_reasons = [feature_names[i] for i in top3_indices]

3. 处理特征工程后的特殊情况

如果训练时用了编码器(如OneHotEncoder),需要从编码器中获取编码后的特征名,而非原始列名:

from sklearn.preprocessing import OneHotEncoder

# 假设encoder是训练好的OneHotEncoder
encoded_feature_names = encoder.get_feature_names_out(input_features=used_features)
feature_names = encoded_feature_names.tolist()

关键注意事项

  • 永远保证feature_names的顺序、数量和模型训练时输入的特征完全一致,这是索引不越界的核心前提。
  • 二分类模型的SHAP值是二维数组,必须指定对应类别(违约类)的索引,否则会因维度不匹配导致后续索引错误。

内容的提问来源于stack exchange,提问作者bunti papu

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 08:40:13