如何在HuggingFace Transformers的LLaVa实现中提取图像隐藏状态?
问题
我正在使用HuggingFace的transformers库提取LLaVa 1.5的所有隐藏单元。根据文档,可从视觉组件提取图像隐藏状态,但当前调用model.generate()得到的outputs对象仅包含['sequences', 'attentions', 'hidden_states', 'past_key_values']这些键,请问如何在现有输出基础上额外提取image_hidden_states?
以下是我的实现代码:
import torch from transformers import LlavaForConditionalGeneration, LlavaConfig, CLIPVisionConfig, LlamaConfig, AutoProcessor, LlavaProcessor from PIL import Image import requests from torchinfo import summary device = "cuda:0" if torch.cuda.is_available() else "cpu" model_id = 'llava-hf/llava-1.5-7b-hf' # Initializing a CLIP-vision config vision_config = CLIPVisionConfig(output_hidden_states=True, output_attentions=True, return_dict=True) # Initializing a Llama config text_config = LlamaConfig(output_hidden_states=True, output_attentions=True, return_dict=True) # Initializing a Llava llava-1.5-7b style configuration configuration = LlavaConfig(vision_config, text_config, output_hidden_states=True, output_attentions=True, return_dict=True) cfg=LlavaConfig(vision_config, text_config, output_hidden_states=True, output_attentions=True, return_dict=True) # Initializing a model from the llava-1.5-7b style configuration model = LlavaForConditionalGeneration(configuration).from_pretrained(model_id, output_hidden_states=True, output_attentions=True, return_dict=True) # Accessing the model configuration configuration = model.config model=model.to(device) print(summary(model)) processor = LlavaProcessor.from_pretrained("llava-hf/llava-1.5-7b-hf", output_hidden_states=True, output_attentions=True, return_dict=True) prompt = "USER: <image>\nIs there sun in the image? ASSISTANT:" url = "https://www.ilankelman.org/stopsigns/australia.jpg" image = Image.open(requests.get(url, stream=True).raw) inputs = processor(text=prompt, images=image, return_tensors="pt") inputs=inputs.to(device) with torch.no_grad(): outputs = model.generate(**inputs, output_hidden_states=True, return_dict_in_generate=True, max_new_tokens=1, min_new_tokens=1, return_dict=True) print(outputs.keys())
解决方案
model.generate()方法默认只返回文本生成相关的输出,不会包含视觉编码器的隐藏状态。要获取image_hidden_states,可以通过以下两种方式实现:
方法1:使用PyTorch钩子捕获视觉编码器输出
通过给视觉编码器的forward方法添加钩子,在模型运行时自动捕获隐藏状态:
import torch from transformers import LlavaForConditionalGeneration, LlavaProcessor from PIL import Image import requests device = "cuda:0" if torch.cuda.is_available() else "cpu" model_id = 'llava-hf/llava-1.5-7b-hf' # 加载模型和处理器,确保视觉编码器开启输出隐藏状态 model = LlavaForConditionalGeneration.from_pretrained( model_id, vision_config={"output_hidden_states": True}, output_hidden_states=True, return_dict=True ).to(device) processor = LlavaProcessor.from_pretrained(model_id) # 定义钩子函数,存储视觉编码器的隐藏状态 image_hidden_states = [] def vision_hook(module, input, output): image_hidden_states.append(output.hidden_states) # 给视觉编码器的最后一层添加钩子 model.vision_model.encoder.layers[-1].register_forward_hook(vision_hook) # 处理输入 prompt = "USER: <image>\nIs there sun in the image? ASSISTANT:" url = "https://www.ilankelman.org/stopsigns/australia.jpg" image = Image.open(requests.get(url, stream=True).raw) inputs = processor(text=prompt, images=image, return_tensors="pt").to(device) # 执行生成 with torch.no_grad(): outputs = model.generate( **inputs, return_dict_in_generate=True, max_new_tokens=1, min_new_tokens=1 ) # 查看捕获的视觉隐藏状态 print("视觉隐藏状态层数:", len(image_hidden_states[0])) print("最后一层视觉隐藏状态形状:", image_hidden_states[0][-1].shape) print("生成输出键:", outputs.keys())
方法2:手动调用视觉编码器和投影层
直接调用模型的视觉组件提取图像特征,无需依赖生成流程:
import torch from transformers import LlavaForConditionalGeneration, LlavaProcessor from PIL import Image import requests device = "cuda:0" if torch.cuda.is_available() else "cpu" model_id = 'llava-hf/llava-1.5-7b-hf' model = LlavaForConditionalGeneration.from_pretrained( model_id, vision_config={"output_hidden_states": True}, return_dict=True ).to(device) processor = LlavaProcessor.from_pretrained(model_id) # 单独处理图像输入 url = "https://www.ilankelman.org/stopsigns/australia.jpg" image = Image.open(requests.get(url, stream=True).raw) image_inputs = processor(images=image, return_tensors="pt").to(device) # 调用视觉编码器获取隐藏状态 with torch.no_grad(): vision_outputs = model.vision_model(**image_inputs) image_hidden_states = vision_outputs.hidden_states # 获取投影后的图像特征(LLaVA中用于和文本特征融合的部分) projected_image_features = model.multi_modal_projector(vision_outputs.last_hidden_state) print("视觉隐藏状态层数:", len(image_hidden_states)) print("最后一层视觉隐藏状态形状:", image_hidden_states[-1].shape) print("投影后图像特征形状:", projected_image_features.shape) # 继续执行文本生成(如果需要) prompt = "USER: <image>\nIs there sun in the image? ASSISTANT:" text_inputs = processor(text=prompt, return_tensors="pt").to(device) inputs = {**text_inputs, "image_features": projected_image_features} with torch.no_grad(): outputs = model.generate( **inputs, return_dict_in_generate=True, max_new_tokens=1, min_new_tokens=1 ) print("生成输出键:", outputs.keys())
关键说明
- 必须确保视觉编码器的配置
output_hidden_states=True,可以在加载模型时通过vision_config参数指定,无需手动重新初始化整个配置(原代码中重复初始化配置的步骤可简化)。 - 钩子方法适合在生成流程中同步捕获视觉状态,手动调用方法则更灵活,可单独提取图像特征再进行生成。
内容的提问来源于stack exchange,提问作者Mihir Mehta
相关产品推荐
相关产品推荐

