Image representation interpretation

#71
by kanishka - opened

Hi! I am interested in using the image representations from the Llama-3.2-11B-Vision model, and am unsure how exactly to interpret the output.

This is what I ran:

from PIL import Image
from transformers import MllamaForConditionalGeneration, AutoProcessor

model_id = "meta-llama/Llama-3.2-11B-Vision"

model = MllamaForConditionalGeneration.from_pretrained(
    model_id,
    torch_dtype=torch.bfloat16,
    device_map="auto",
)
processor = AutoProcessor.from_pretrained(model_id)

image_paths = ["../data/aardvark/aardvark_01b.jpg", "../data/aardvark/aardvark_02s.jpg", "../data/aardvark/aardvark_10s.jpg"] # change accordingly

img = [Image.open(open(ip, "rb")) for ip in image_paths]

encoded = processors(images=img, text=[" ", " ", ""], return_tensors='pt')
encoded.pop("input_ids")
encoded.pop("attention_mask")
encoded.pop("cross_attention_mask")
encoded = encoded.to(torch.bfloat16)
encoded = encoded.to("cuda:0")

output = model.vision_model(**encoded)

output.last_hidden_state.shape

# output: torch.Size([1, 3, 4, 1025, 7680])

I'm wondering how I should interpret this shape -- the 3 clearly is the batch; and 7680 is clearly the representation dimension. But I am unsure what the other shapes represent. I'd love to get some clarity on this

Thanks!

Sign up or log in to comment