RandomForestClassifier的SHAP可加性校验失败,SHAP值异常巨大
解决RandomForest+SHAP可加性校验失败及SHAP值异常问题
基于RandomForestClassifier训练文本分类模型,后端应用中使用SHAP进行可解释性分析时,遭遇SHAP可加性校验失败、生成的SHAP值异常巨大的问题。以下是相关代码、报错信息、已尝试步骤及解决方案:
模型训练代码
import pandas as pd import pickle from sklearn.feature_extraction.text import TfidfVectorizer from sklearn.ensemble import RandomForestClassifier from sklearn.model_selection import train_test_split, RandomizedSearchCV from sklearn.metrics import classification_report # Load and preprocess data news_df = pd.read_csv('../data/WELFake_Dataset.csv').fillna('') news_df['clean_text'] = news_df['text'].apply(preprocess_text) X = news_df['clean_text'] y = news_df['label'] X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42) # Vectorizer tfidf_vectorizer = TfidfVectorizer(stop_words='english', max_features=10000) tfidf_train = tfidf_vectorizer.fit_transform(X_train) tfidf_test = tfidf_vectorizer.transform(X_test) # Model training with hyperparameter tuning param_grid = { 'n_estimators': [100, 200, 300], 'max_depth': [None, 10, 20, 30], 'min_samples_split': [2, 5, 10], 'min_samples_leaf': [1, 2, 4], 'bootstrap': [True, False] } rf = RandomForestClassifier(random_state=42) random_search = RandomizedSearchCV(rf, param_distributions=param_grid, n_iter=20, cv=5, verbose=2, random_state=42, n_jobs=-1) random_search.fit(tfidf_train, y_train) best_rf_model = random_search.best_estimator_ # Save model and vectorizer with open('../backend/tfidf_vectorizer.pkl', 'wb') as f: pickle.dump(tfidf_vectorizer, f) with open('../backend/best_rf_model.pkl', 'wb') as f: pickle.dump(best_rf_model, f)
SHAP测试代码
import pickle import shap import numpy as np from utils.preprocessing import preprocess_text # Load vectorizer and model with open('../backend/tfidf_vectorizer.pkl', 'rb') as f: tfidf_vectorizer = pickle.load(f) with open('../backend/best_rf_model.pkl', 'rb') as f: classifier = pickle.load(f) # Initialize SHAP Explainer explainer_shap = shap.TreeExplainer(classifier) # Sample text sample_text = "The government has announced a new policy to combat fake news." preprocessed_text = preprocess_text(sample_text) feature_vector = tfidf_vectorizer.transform([preprocessed_text]).toarray() # Predict prediction = classifier.predict(feature_vector)[0] proba = classifier.predict_proba(feature_vector)[0][1] # Generate SHAP values shap_values = explainer_shap.shap_values(feature_vector) shap_values_positive = shap_values[1] if isinstance(shap_values, list) else shap_values # Sum SHAP values shap_sum = np.sum(shap_values_positive) print(f"Sum of SHAP values: {shap_sum}") print(f"Model output: {proba}") # Additivity Check if np.abs(shap_sum - proba) < 0.01: print("SHAP additivity check passed.") else: print("SHAP additivity check failed.")
报错信息
python test_shap_consistency.py [nltk_data] Downloading package punkt to [nltk_data] C:\Users\usr\AppData\Roaming\nltk_data... [nltk_data] Package punkt is already up-to-date! [nltk_data] Downloading package stopwords to [nltk_data] C:\Users\usr\AppData\Roaming\nltk_data... [nltk_data] Package stopwords is already up-to-date! [nltk_data] Downloading package wordnet to [nltk_data] C:\Users\usr\AppData\Roaming\nltk_data... [nltk_data] Package wordnet is already up-to-date! SHAP version: 0.46.1.dev82 Classifier type: <class 'sklearn.ensemble._forest.RandomForestClassifier'> Number of features in vectorizer: 10000 Number of features in model: 10000 Preprocessed text: government announced new policy combat fake news Feature vector shape: (1, 10000) First 10 features: [0. 0. 0. 0. 0. 0. 0. 0. 0. 0.] Prediction: Real Probability of being Real: 0.9291109563602599 Traceback (most recent call last): File "C:\Users\usr\Desktop\tema\backend\test_shap_consistency.py", line 50, in <module> shap_values = explainer_shap.shap_values(feature_vector_dense) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "C:\Users\usr\anaconda3\Lib\site-packages\shap\explainers\_tree.py", line 634, in shap_values self.assert_additivity(out, self.model.predict(X)) File "C:\Users\usr\anaconda3\Lib\site-packages\shap\explainers\_tree.py", line 816, in assert_additivity check_sum(self.expected_value[i] + phi[i].sum(-1), model_output[:, i]) File "C:\Users\usr\anaconda3\Lib\site-packages\shap\explainers\_tree.py", line 812, in check_sum raise ExplainerError(err_msg) shap.utils._exceptions.ExplainerError: Additivity check failed in TreeExplainer! Please ensure the data matrix you passed to the explainer is the same shape that the model was trained on. If your data shape is correct then please report this on GitHub. Consider retrying with the feature_perturbation='interventional' option. This check failed because for one of the samples the sum of the SHAP values was -71537078264766386643956768191945774872434602891202492950066515140319945066393991267812782394642882616075517846114722773991424.000000, while the model output was 0.070889. If this difference is acceptable you can set check_additivity=False to disable this check.
已尝试解决步骤
- 将SHAP版本更新至0.46;
- 确认TF-IDF向量器与RandomForest模型的特征数均为10000,匹配一致;
- 验证preprocess_text函数在所有脚本中保持一致;
- 重新训练并保存模型与向量器,确保完整性;
- 检查SHAP值与特征向量中无NaN或无穷值;
- 对齐库版本,确保环境一致性。
问题原因分析
- 特征扰动方式不匹配稀疏数据:TreeExplainer默认使用
tree_path_dependent模式,对于TF-IDF这种高稀疏特征(大部分值为0),树路径计算易出现数值溢出,导致SHAP值异常。 - 数据格式不一致:模型训练时用的是稀疏矩阵(csr_matrix),但测试时转为稠密矩阵(toarray()),存储格式差异引发SHAP内部计算逻辑冲突。
- SHAP开发版bug:使用的0.46.1.dev82是开发版本,存在高维稀疏树模型的数值计算bug,导致溢出。
解决方案
1. 切换到Interventional特征扰动模式
初始化Explainer时指定特征扰动方式,避免路径依赖的计算问题:
explainer_shap = shap.TreeExplainer(classifier, feature_perturbation="interventional")
2. 保持数据格式一致性(使用稀疏矩阵)
跳过toarray()转换,直接用TF-IDF生成的稀疏矩阵进行预测和SHAP计算:
# 保留稀疏矩阵格式 feature_vector = tfidf_vectorizer.transform([preprocessed_text]) # 预测和SHAP计算均使用该稀疏矩阵 prediction = classifier.predict(feature_vector)[0] proba = classifier.predict_proba(feature_vector)[0][1] shap_values = explainer_shap.shap_values(feature_vector)
3. 更换稳定版SHAP或降低特征维度
- 安装稳定版SHAP,避免开发版bug:
pip install shap==0.45.0
- 减少TF-IDF的
max_features参数(如从10000降至5000),降低特征稀疏性,减少计算压力。
4. 临时禁用可加性校验(不推荐长期使用)
若上述方法无效,可临时关闭校验,但会失去SHAP值正确性的保障:
shap_values = explainer_shap.shap_values(feature_vector, check_additivity=False)
内容的提问来源于stack exchange,提问作者idkrlly
相关产品推荐
相关产品推荐

