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

使用Captum获取文本分类解释时遇RuntimeError问题求助

解决BERT+Integrated Gradients的张量类型错误

问题背景

使用BERT训练文本二分类模型后,尝试用Captum的Integrated Gradients分析测试集中正确预测样本的关键影响词汇时,触发如下错误:

RuntimeError: Expected tensor for argument #1 'indices' to have one of the following scalar types: Long, Int; but got torch.cuda.FloatTensor instead (while checking arguments for embedding)

错误原因

Captum的Integrated Gradients在计算归因过程中,会自动将输入的input_ids张量转换为浮点类型,但BERT的词嵌入层要求输入的索引必须是整数类型(Long/Int),两者类型不匹配导致报错。

修复方案

核心修改点

  • 在forward_func中,将传入的浮点型input_ids强制转换为torch.long类型,适配BERT的embedding层要求
  • 替换过时的transformers.AdamW为PyTorch原生的torch.optim.AdamW
  • 为ig.attribute添加internal_batch_size参数,避免大张量导致的内存溢出

修改后的完整代码

import pandas as pd
import torch
from torch.utils.data import DataLoader
from transformers import BertTokenizer, BertForSequenceClassification
from sklearn.metrics import accuracy_score
from captum.attr import IntegratedGradients

# Loading data
train_df = pd.read_csv('train_dataset.csv')
test_df = pd.read_csv('test_dataset.csv')

# Tokenizer
tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')

def preprocess_data(df, tokenizer, max_len=128):
    inputs = tokenizer(list(df['text']), padding=True, truncation=True, max_length=max_len, return_tensors="pt")
    labels = torch.tensor(df['label'].values, dtype=torch.long)
    return inputs, labels

train_inputs, train_labels = preprocess_data(train_df, tokenizer)
test_inputs, test_labels = preprocess_data(test_df, tokenizer)

# DataLoader
train_dataset = torch.utils.data.TensorDataset(train_inputs['input_ids'], train_inputs['attention_mask'], train_labels)
train_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)

test_dataset = torch.utils.data.TensorDataset(test_inputs['input_ids'], test_inputs['attention_mask'], test_labels)
test_loader = DataLoader(test_dataset, batch_size=16)

# Model setup
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = BertForSequenceClassification.from_pretrained('bert-base-uncased', num_labels=2).to(device)

# Optimizer - 使用PyTorch原生AdamW替代transformers的过时实现
optimizer = torch.optim.AdamW(model.parameters(), lr=5e-5)

# Training Loop
model.train()
for epoch in range(3):  # Train for 3 epochs
    total_loss = 0.0
    for batch in train_loader:
        input_ids, attention_mask, labels = [x.to(device) for x in batch]
        optimizer.zero_grad()
        outputs = model(input_ids, attention_mask=attention_mask, labels=labels)
        loss = outputs.loss
        total_loss += loss.item()
        loss.backward()
        optimizer.step()
    print(f"Epoch {epoch+1} average loss: {total_loss/len(train_loader):.4f}")

# Evaluation
model.eval()
correct_predictions = []
all_preds = []
with torch.no_grad():
    for batch in test_loader:
        input_ids, attention_mask, labels = [x.to(device) for x in batch]
        outputs = model(input_ids, attention_mask=attention_mask)
        preds = torch.argmax(outputs.logits, dim=1)
        correct_predictions.extend((preds == labels).cpu().numpy().tolist())
        all_preds.extend(preds.cpu().numpy().tolist())
accuracy = accuracy_score(test_labels.numpy(), all_preds)
print(f"Test Accuracy: {accuracy:.2f}")

# Integrated Gradients
ig = IntegratedGradients(model)

def get_influential_words(input_text, model, tokenizer, ig, device):
    model.eval()
    # Tokenizing the input text
    inputs = tokenizer(input_text, return_tensors="pt", truncation=True, padding=True, max_length=128)
    input_ids = inputs['input_ids'].to(device)
    attention_mask = inputs['attention_mask'].to(device)

    # forward function for IG - 关键:将输入的浮点型input_ids转换回long类型
    def forward_func(inputs):
        input_ids_long = inputs.to(torch.long)
        outputs = model(input_ids_long, attention_mask=attention_mask)
        return outputs.logits

    # Applying Integrated Gradients
    true_label = test_df['label'].iloc[test_df[test_df['text'] == input_text].index[0]]
    attributions, delta = ig.attribute(input_ids, target=true_label, return_convergence_delta=True, internal_batch_size=1)
    tokens = tokenizer.convert_ids_to_tokens(input_ids[0].tolist())
    token_importances = attributions.sum(dim=2).squeeze(0).detach().cpu().numpy()

    return list(zip(tokens, token_importances))

# Analysing influential words for correctly predicted texts
for idx, correct in enumerate(correct_predictions):
    if correct:
        text = test_df['text'].iloc[idx]
        influential_words = get_influential_words(text, model, tokenizer, ig, device)
        print(f"\nInfluential words for text: {text}")
        # 按归因值绝对值降序排序,跳过特殊token
        sorted_influential = sorted(influential_words, key=lambda x: abs(x[1]), reverse=True)
        for token, score in sorted_influential:
            if token not in ['[CLS]', '[SEP]', '[PAD]']:
                print(f"{token}: {score:.4f}")

额外优化说明

  • 训练循环计算平均损失,更准确反映训练状态
  • 分析时跳过特殊token,只展示真实词汇的归因值
  • 按归因值绝对值降序排序,直观展示核心影响词汇

内容的提问来源于stack exchange,提问作者Nayantara Jeyaraj

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 21:44:53