使用HuggingFace的Flava模型获取contrastive_logits_per_image报错求助
问题解决方法
第一个错误原因及解决
FlavaModel是FLAVA的基础特征提取模型,仅输出文本/图像的编码特征,不包含图文对比学习(ITC)的logits输出,因此会出现'FlavaModelOutput' object has no attribute 'contrastive_logits_per_image'的错误。
要获取图文相似度分数,必须使用预训练任务模型FlavaForPreTraining,但需要正确调用以避免触发不必要的预训练分支。
第二个错误原因及解决
你遇到的AttributeError: 'NoneType' object has no attribute 'dim',是因为代码触发了FLAVA的掩码图像建模(MIM)任务分支,但未传入该任务所需的bool_masked_pos参数。
修正后的代码
from PIL import Image import requests from transformers import FlavaProcessor, FlavaForPreTraining model = FlavaForPreTraining.from_pretrained("facebook/flava-full") processor = FlavaProcessor.from_pretrained("facebook/flava-full") url = "http://images.cocodataset.org/val2017/000000039769.jpg" image = Image.open(requests.get(url, stream=True).raw) # 仅生成图文对比所需的输入,无需MIM/MLM相关参数 inputs = processor(text=["a photo of a cat"], images=image, return_tensors="pt") # 显式指定不计算损失,避免触发自动生成MIM/MLM标签的逻辑 outputs = model(**inputs, return_loss=False) logits_per_image = outputs.contrastive_logits_per_image # 图文相似度分数 probs = logits_per_image.softmax(dim=1) # 转换为概率分布
关键修正点
- 移除
return_codebook_pixels=True:该参数用于MIM任务,获取图文相似度时无需启用。 - 移除
input_ids_masked的手动传入:这是掩码语言建模(MLM)任务的参数,与图文对比无关。 - 显式设置
return_loss=False:避免模型自动生成MIM/MLM任务所需的标签和掩码参数,跳过相关分支的执行。
关于FutureWarning
该警告是transformers版本兼容提示,v5版本将移除device参数,不影响当前代码功能。若要消除警告,可升级transformers至最新稳定版本。
内容的提问来源于stack exchange,提问作者lazytux
相关产品推荐
相关产品推荐

