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

预训练BERT二分类任务中如何提取SHAP汇总图?

BERT二分类任务SHAP汇总图问题解决方案

问题1:调用shap.summary_plot触发断言错误(提示需要矩阵而非向量)

  • 核心原因:BERT是token级输入,二分类任务下SHAP返回的shap_values是包含两个数组的列表(分别对应负类、正类的SHAP值),如果错误传入一维向量(比如单样本的SHAP值)或维度不匹配的数组,就会触发断言。
  • 解决办法:
    1. 明确取目标类的SHAP矩阵:二分类场景下,shap_values[1]是正类的SHAP值,形状为[样本数, token数],这是summary_plot需要的二维矩阵格式。
    2. 过滤无效token:去掉[CLS]、[SEP]、[PAD]这类无意义token对应的SHAP值,避免引入无效维度。
    3. 检查输入格式:确保传入summary_plot的shap_values是所有样本的SHAP矩阵,而非单个样本的一维向量。

问题2:绘图仅显示蓝色点,调整维度后总数量固定20

  • 仅显示蓝色点的原因:
    1. 可能传入了负类的SHAP值,导致贡献值全为负;
    2. 直接用token id作为特征值,离散的id值无法形成合理数值分布,SHAP无法区分正负贡献的颜色映射;
    3. 未过滤padding token,无效token的SHAP值干扰了结果。
  • 总数量固定20的原因:shap.summary_plot默认只显示重要性前20的特征,需手动调整参数。
  • 解决办法:
    1. 映射token id为单词:用tokenizer把输入的token id转换成对应单词,传入summary_plot的feature_names参数,让特征更直观。
    2. 指定目标类的SHAP值:确保传入的是你关注类别(比如正类)的SHAP矩阵,这样颜色会根据token的贡献正负呈现不同色调。
    3. 调整显示数量:设置max_display参数(比如max_display=30),自定义显示的特征数量。
    4. 聚合相同单词的贡献:如果要按单词而非token位置汇总重要性,可以把相同单词的SHAP值聚合后再绘图,避免重复显示同一单词。

修正后的最小可复现代码

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 17:57:03