使用PyTorch Hub中RoBERTa模型时出现KeyError: 'mnli'错误的解决方法
解决PyTorch Hub加载RoBERTa时的KeyError: 'mnli'问题
我尝试使用PyTorch Hub中的RoBERTa模型运行MNLI任务,代码如下:
import torch roberta = torch.hub.load('pytorch/fairseq', 'roberta.large') tokens = roberta.encode('Roberta is a heavily optimized version of BERT.', 'Roberta is not very optimized.') roberta.predict('mnli', tokens).argmax() # 0: contradiction
运行时触发KeyError:
│ 458 │ │ 459 │ @_copy_to_script_wrapper │ 460 │ def __getitem__(self, key: str) -> Module: │ ❱ 461 │ │ return self._modules[key] │ 462 │ │ 463 │ def __setitem__(self, key: str, module: Module) -> None: │ 464 │ │ self.add_module(key, module) ╰────────────────────────────────────────────────────── KeyError: 'mnli'
问题原因
你加载的roberta.large是基础预训练RoBERTa模型,仅包含语言建模核心模块,没有附带MNLI任务对应的分类头,因此调用predict('mnli')时会找不到对应模块。
解决方法
需要加载专门针对MNLI任务微调过的RoBERTa模型,将模型加载语句替换为:
import torch # 加载带MNLI任务头的预训练模型 roberta = torch.hub.load('pytorch/fairseq', 'roberta.large.mnli') tokens = roberta.encode('Roberta is a heavily optimized version of BERT.', 'Roberta is not very optimized.') # 此时可正常执行MNLI预测 roberta.predict('mnli', tokens).argmax() # 0: contradiction
补充说明
Fairseq的PyTorch Hub提供了多种带任务头的预训练模型,格式为roberta.[size].[task],比如roberta.base.mnli、roberta.large.mnli等,选择对应任务的模型即可直接调用predict方法完成下游任务。
内容的提问来源于stack exchange,提问作者Afshin Oroojlooy
相关产品推荐
相关产品推荐

