使用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
相关产品推荐
相关产品推荐

