如何用Hugging Face库获取LXMERT模型的prediction logits?
解决LXMERT无法获取prediction_logits的问题
问题出在你加载的模型类上:
- 你用的
LxmertModel是基础编码器,只负责输出语言、视觉的中间表征和注意力,不带预训练任务头,所以自然没有prediction_logits和seq_relationship_logits。 - 而VisualBERT你可能用的是带预训练头的类(比如
VisualBertForPreTraining),所以能直接拿到那些logits。
解决步骤
- 改用
LxmertForPreTraining类加载模型,这个类包含了预训练所需的MLM(掩码语言模型)和句子关系预测任务头:
from transformers import LxmertForPreTraining, LxmertTokenizer tokenizer = LxmertTokenizer.from_pretrained("unc-nlp/lxmert-base-uncased") model = LxmertForPreTraining.from_pretrained("unc-nlp/lxmert-base-uncased")
- 输入数据后,就能在输出里拿到
prediction_logits和seq_relationship_logits了:
# 示例输入(文本+视觉特征) text = "A cat sitting on a mat." inputs = tokenizer(text, return_tensors="pt") # 视觉特征示例(实际需传入真实图像特征,这里用随机张量模拟) visual_feats = torch.randn(1, 36, 2048) visual_pos = torch.randn(1, 36, 4) outputs = model( input_ids=inputs["input_ids"], attention_mask=inputs["attention_mask"], visual_feats=visual_feats, visual_pos=visual_pos ) # 现在可以访问对应的logits print(outputs.prediction_logits.shape) # MLM预测logits print(outputs.seq_relationship_logits.shape) # 句子关系预测logits
补充说明
如果你是用LXMERT做下游任务(比如VQA),需要用对应的下游任务模型类(比如LxmertForQuestionAnswering),此时输出的是下游任务的预测logits,而非预训练阶段的prediction_logits。
内容的提问来源于stack exchange,提问作者lazytux
相关产品推荐
相关产品推荐

