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

使用transformers-interpret调用RoBERTa模型出现IndexError报错求助

问题描述

使用transformers-interpret库调用facebook/bart-large-mnli做零样本分类解释可正常运行,但调用roberta-large-mnli时触发IndexError。

BART模型可运行代码

from transformers import AutoModelForSequenceClassification, AutoTokenizer
from transformers_interpret import ZeroShotClassificationExplainer

tokenizer_zeroshot = AutoTokenizer.from_pretrained("facebook/bart-large-mnli")
model_zeroshot = AutoModelForSequenceClassification.from_pretrained("facebook/bart-large-mnli")

zero_shot_explainer_1 = ZeroShotClassificationExplainer(model_zeroshot, tokenizer_zeroshot)

word_attributions = zero_shot_explainer_1(
    "reporter : mr . trump , how are you planning on making america great again ? trump : it 's simple ... # trump # parody  # meme # funny # rt emoji_1942",
    labels = ["ironic", "non-ironic"],
)

html = zero_shot_explainer_1.visualize()

RoBERTa模型报错代码

from transformers import AutoModelForSequenceClassification, AutoTokenizer
from transformers_interpret import ZeroShotClassificationExplainer

tokenizer = AutoTokenizer.from_pretrained("roberta-large-mnli")
model = AutoModelForSequenceClassification.from_pretrained("roberta-large-mnli")

zero_shot_explainer = ZeroShotClassificationExplainer(model, tokenizer)

word_attributions = zero_shot_explainer(
    "reporter : mr . trump , how are you planning on making america great again ? trump : it 's simple ... # trump # parody  # meme # funny # rt emoji_1942",
    labels = ["ironic", "non-ironic"],
)

完整报错堆栈

---------------------------------------------------------------------------
IndexError                                Traceback (most recent call last)
<ipython-input-14-a7052a056405> in <module>()
      1 word_attributions = zero_shot_explainer(
      2     "reporter : mr . trump , how are you planning on making america great again ? trump : it 's simple ... # trump # parody  # meme # funny # rt emoji_1942",
----> 3     labels = ["ironic", "non-ironic"],
      4 )

11 frames
/usr/local/lib/python3.7/dist-packages/transformers_interpret/explainers/zero_shot_classification.py in __call__(self, text, labels, embedding_type, hypothesis_template, include_hypothesis, internal_batch_size, n_steps)
    290         self.hypothesis_labels = [hypothesis_template.format(label) for label in labels]
    291 
--> 292         predicted_text_idx = self._get_top_predicted_label_idx(text, self.hypothesis_labels)
    293 
    294         for i, _ in enumerate(self.labels):

/usr/local/lib/python3.7/dist-packages/transformers_interpret/explainers/zero_shot_classification.py in _get_top_predicted_label_idx(self, text, hypothesis_labels)
    135             token_type_ids, _ = self._make_input_reference_token_type_pair(input_ids, sep_idx)
    136             attention_mask = self._make_attention_mask(input_ids)
--> 137             preds = self._get_preds(input_ids, token_type_ids, position_ids, attention_mask)
    138             entailment_outputs.append(float(torch.sigmoid(preds[0])[0][self.entailment_idx]))
    139 

/usr/local/lib/python3.7/dist-packages/transformers_interpret/explainers/question_answering.py in _get_preds(self, input_ids, token_type_ids, position_ids, attention_mask)
    204                 token_type_ids=token_type_ids,
    205                 position_ids=position_ids,
--> 206                 attention_mask=attention_mask,
    207             )
    208 

/usr/local/lib/python3.7/dist-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs)
   1128         if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks
   1129                 or _global_forward_hooks or _global_forward_pre_hooks):
-> 1130             return forward_call(*input, **kwargs)
   1131         # Do not call functions when jit is used
   1132         full_backward_hooks, non_full_backward_hooks = [], []

/usr/local/lib/python3.7/dist-packages/transformers/models/roberta/modeling_roberta.py in forward(self, input_ids, attention_mask, token_type_ids, position_ids, head_mask, inputs_embeds, labels, output_attentions, output_hidden_states, return_dict)
   1213             output_attentions=output_attentions,
   1214             output_hidden_states=output_hidden_states,
-> 1215             return_dict=return_dict,
   1216         )
   1217         sequence_output = outputs[0]

/usr/local/lib/python3.7/dist-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs)
   1128         if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks
   1129                 or _global_forward_hooks or _global_forward_pre_hooks):
-> 1130             return forward_call(*input, **kwargs)
   1131         # Do not call functions when jit is used
   1132         full_backward_hooks, non_full_backward_hooks = [], []

/usr/local/lib/python3.7/dist-packages/transformers/models/roberta/modeling_roberta.py in forward(self, input_ids, attention_mask, token_type_ids, position_ids, head_mask, inputs_embeds, encoder_hidden_states, encoder_attention_mask, past_key_values, use_cache, output_attentions, output_hidden_states, return_dict)
    844             token_type_ids=token_type_ids,
    845             inputs_embeds=inputs_embeds,
-> 846             past_key_values_length=past_key_values_length,
    847         )
    848         encoder_outputs = self.encoder(

/usr/local/lib/python3.7/dist-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs)
   1128         if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks
   1129                 or _global_forward_hooks or _global_forward_pre_hooks):
-> 1130             return forward_call(*input, **kwargs)
   1131         # Do not call functions when jit is used
   1132         full_backward_hooks, non_full_backward_hooks = [], []

/usr/local/lib/python3.7/dist-packages/transformers/models/roberta/modeling_roberta.py in forward(self, input_ids, token_type_ids, position_ids, inputs_embeds, past_key_values_length)
    127         if inputs_embeds is None:
    128             inputs_embeds = self.word_embeddings(input_ids)
-> 129         token_type_embeddings = self.token_type_embeddings(token_type_ids)
    130 
    131         embeddings = inputs_embeds + token_type_embeddings

/usr/local/lib/python3.7/dist-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs)
   1128         if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks
   1129                 or _global_forward_hooks or _global_forward_pre_hooks):
-> 1130             return forward_call(*input, **kwargs)
   1131         # Do not call functions when jit is used
   1132         full_backward_hooks, non_full_backward_hooks = [], []

/usr/local/lib/python3.7/dist-packages/torch/nn/modules/sparse.py in forward(self, input)
    158         return F.embedding(
    159             input, self.weight, self.padding_idx, self.max_norm,
-> 160             self.norm_type, self.scale_grad_by_freq, self.sparse)
    161 
    162     def extra_repr(self) -> str:

/usr/local/lib/python3.7/dist-packages/torch/nn/functional.py in embedding(input, weight, padding_idx, max_norm, norm_type, scale_grad_by_freq, sparse)
   2197         # remove once script supports set_grad_enabled
   2198         _no_grad_embedding_renorm_(weight, input, max_norm, norm_type)
-> 2199     return torch.embedding(weight, input, padding_idx, scale_grad_by_freq, sparse)
   2200 
   2201 

IndexError: index out of range in self

报错原因分析

从报错堆栈的核心位置token_type_embeddings = self.token_type_embeddings(token_type_ids)可以定位问题:

  • RoBERTa模型没有设计多类型token的embedding层,它的token_type_embeddings权重仅包含1个维度(对应唯一的token类型),但transformers-interpret的ZeroShotClassificationExplainer在处理零样本分类时,会为"文本+假设"的拼接输入生成多组token type ID(比如0和1),导致传入的索引超出了RoBERTa模型token_type_embeddings的有效范围,触发IndexError。
  • BART模型原生支持多类型token区分,能正常处理这种拼接输入的token type ID,因此不会报错。

临时解决办法

可以通过包装RoBERTa模型,忽略传入的token_type_ids参数来规避问题:

import torch
from transformers import AutoModelForSequenceClassification, AutoTokenizer
from transformers_interpret import ZeroShotClassificationExplainer

class RoBERTaZeroShotWrapper(torch.nn.Module):
    def __init__(self, model):
        super().__init__()
        self.model = model
    
    def forward(self, input_ids, attention_mask, **kwargs):
        # 忽略token_type_ids参数,适配RoBERTa模型
        return self.model(input_ids=input_ids, attention_mask=attention_mask)

# 加载模型和分词器
tokenizer = AutoTokenizer.from_pretrained("roberta-large-mnli")
model = AutoModelForSequenceClassification.from_pretrained("roberta-large-mnli")

# 包装模型
model = RoBERTaZeroShotWrapper(model)

# 正常调用解释器
zero_shot_explainer = ZeroShotClassificationExplainer(model, tokenizer)
word_attributions = zero_shot_explainer(
    "reporter : mr . trump , how are you planning on making america great again ? trump : it 's simple ... # trump # parody  # meme # funny # rt emoji_1942",
    labels = ["ironic", "non-ironic"],
)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.23 02:09:13