预训练BERT二分类任务中如何提取SHAP汇总图?
BERT二分类任务SHAP汇总图问题解决方案
问题1:调用shap.summary_plot触发断言错误(提示需要矩阵而非向量)
- 核心原因:BERT是token级输入,二分类任务下SHAP返回的
shap_values是包含两个数组的列表(分别对应负类、正类的SHAP值),如果错误传入一维向量(比如单样本的SHAP值)或维度不匹配的数组,就会触发断言。 - 解决办法:
- 明确取目标类的SHAP矩阵:二分类场景下,
shap_values[1]是正类的SHAP值,形状为[样本数, token数],这是summary_plot需要的二维矩阵格式。 - 过滤无效token:去掉
[CLS]、[SEP]、[PAD]这类无意义token对应的SHAP值,避免引入无效维度。 - 检查输入格式:确保传入
summary_plot的shap_values是所有样本的SHAP矩阵,而非单个样本的一维向量。
- 明确取目标类的SHAP矩阵:二分类场景下,
问题2:绘图仅显示蓝色点,调整维度后总数量固定20
- 仅显示蓝色点的原因:
- 可能传入了负类的SHAP值,导致贡献值全为负;
- 直接用token id作为特征值,离散的id值无法形成合理数值分布,SHAP无法区分正负贡献的颜色映射;
- 未过滤padding token,无效token的SHAP值干扰了结果。
- 总数量固定20的原因:
shap.summary_plot默认只显示重要性前20的特征,需手动调整参数。 - 解决办法:
- 映射token id为单词:用tokenizer把输入的token id转换成对应单词,传入
summary_plot的feature_names参数,让特征更直观。 - 指定目标类的SHAP值:确保传入的是你关注类别(比如正类)的SHAP矩阵,这样颜色会根据token的贡献正负呈现不同色调。
- 调整显示数量:设置
max_display参数(比如max_display=30),自定义显示的特征数量。 - 聚合相同单词的贡献:如果要按单词而非token位置汇总重要性,可以把相同单词的SHAP值聚合后再绘图,避免重复显示同一单词。
- 映射token id为单词:用tokenizer把输入的token id转换成对应单词,传入
修正后的最小可复现代码
import torch from transformers import BertTokenizer, BertForSequenceClassification import shap # 加载模型和tokenizer tokenizer = BertTokenizer.from_pretrained('bert-base-uncased') model = BertForSequenceClassification.from_pretrained('bert-base-uncased', num_labels=2) model.eval() # 测试数据集 texts = ["This is great!", "Terrible service", "I love this product", "Worst experience ever"] # 预处理输入 inputs = tokenizer(texts, return_tensors="pt", padding=True, truncation=True) # 创建SHAP解释器 explainer = shap.Explainer(model, tokenizer) # 获取二分类的SHAP值:[0]负类,[1]正类 shap_values = explainer(inputs) # 过滤无效token,提取有效SHAP值和对应单词 valid_shap_list = [] valid_token_list = [] for idx in range(len(shap_values[1])): # 获取当前样本的所有token tokens = tokenizer.convert_ids_to_tokens(inputs['input_ids'][idx]) # 过滤特殊token valid_indices = [i for i, t in enumerate(tokens) if t not in ['[CLS]', '[SEP]', '[PAD]']] # 提取有效SHAP值和token valid_shap = shap_values[1][idx][valid_indices].numpy() valid_tokens = [tokens[i] for i in valid_indices] valid_shap_list.append(valid_shap) valid_token_list.append(valid_tokens) # 转换为SHAP要求的矩阵格式 valid_shap_matrix = torch.tensor(valid_shap_list).numpy() # 生成特征名(这里按token位置,也可以聚合相同单词) feature_names = valid_token_list[0] # 绘制汇总图,指定显示10个特征 shap.summary_plot(valid_shap_matrix, inputs['input_ids'][:, 1:-1].numpy(), feature_names=feature_names, max_display=10)
内容的提问来源于stack exchange,提问作者Nemo
相关产品推荐
相关产品推荐

