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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 07:33:23