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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 13:59:53