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

