Views
No views yet

1import io
2import torch
3from PIL import Image
4from peft import PeftModel
5from transformers import SiglipModel, SiglipProcessor
6from datasets import load_dataset, Features, Image, Value
7
8features = Features({
9 "image": Image(decode=True),
10 "image_filename": Value("string"),
11 "keyword": Value("string"),
12 "broad_topical_query": Value("string"),
13 "broad_topical_explanation": Value("string"),
14 "specific_detail_query": Value("string"),
15 "specific_detail_explanation": Value("string"),
16 "visual_element_query": Value("string"),
17 "visual_element_explanation": Value("string")
18})
19
20ds = load_dataset(
21 "parquet",
22 data_files={
23 "train": ["wiki_dataset-train.parquet"],
24 "test": ["wiki_dataset-test.parquet"]
25 },
26 features=features
27)
28
29train_ds = ds["train"]
30test_ds = ds["test"]
31
32base_model_id = "google/siglip-so400m-patch14-384"
33ft_model_id = "dj86/siglip-ft-enpedia"
34
35model = SiglipModel.from_pretrained(base_model_id)
36
37model = PeftModel.from_pretrained(model, ft_model_id)
38
39processor = SiglipProcessor.from_pretrained(base_model_id)
40
41images = [train_ds[0]["image"], train_ds[1]["image"], train_ds[2]["image"]]
42inputs = processor(images=images, return_tensors="pt")
43
44texts = ["an image of "+train_ds[0]["keyword"], "an image of "+train_ds[1]["keyword"], "an image of "+train_ds[2]["keyword"]]
45text_inputs = processor(text=texts, return_tensors="pt", padding=True)
46
47with torch.no_grad():
48 image_embeds = model.get_image_features(**inputs)
49 text_embeds = model.get_text_features(**text_inputs)
50
51image_embeds = image_embeds / image_embeds.norm(dim=-1, keepdim=True)
52text_embeds = text_embeds / text_embeds.norm(dim=-1, keepdim=True)
53
54similarity = torch.matmul(text_embeds, image_embeds.T)
55
56print("Similarity:", similarity)
