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

如何用Hugging Face库获取LXMERT模型的prediction logits?

解决LXMERT无法获取prediction_logits的问题

问题出在你加载的模型类上:

  • 你用的LxmertModel是基础编码器,只负责输出语言、视觉的中间表征和注意力,不带预训练任务头,所以自然没有prediction_logits和seq_relationship_logits。
  • 而VisualBERT你可能用的是带预训练头的类(比如VisualBertForPreTraining),所以能直接拿到那些logits。

解决步骤

  1. 改用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")
  1. 输入数据后,就能在输出里拿到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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.25 17:15:47