This model is a fine-tuned version of
openai/clip-vit-large-patch14 as an image encoder and
microsoft/BiomedVLP-CXR-BERT-general as a text encoder on the
ROCO dataset.
It achieves the following results on the evaluation set:
Here is the heatmap of the similarity score of the first 30 samples on the test split of the ROCO dataset of images vs their captions:
heatmap
1import numpy as np
2from sklearn.metrics.pairwise import cosine_similarity
3from PIL import Image
4import pickle, torch, os
5from transformers import VisionTextDualEncoderModel, VisionTextDualEncoderProcessor
6
7# search a query in embeddings
8query = "Chest X-Ray photos"
9
10# embed the query
11inputs = processor(text=query, images=None, return_tensors="pt", padding=True)
12with torch.no_grad():
13 query_embedding = model.get_text_features(**inputs)[0].numpy()
14
15# load image embeddings
16with open("embeddings.pkl", 'rb') as f:
17 image_embeds = pickle.load(f)
18
19# find similar images indices
20def find_k_similar_images(query_embedding, image_embeds, k=2):
21 similarities = cosine_similarity(query_embedding.reshape(1, -1), image_embeds)
22 closest_indices = np.argsort(similarities[0])[::-1][:k]
23 return closest_indices
24similar_image_indices = find_k_similar_images(query_embedding, image_embeds, k=k)
25
26# TO-DO
27images_path = "/path/to/images/"
28images = [os.path.join(images_path,i) for i in os.listdir(images_path) if i.endswith(".jpg")]
29
30# get image paths
31similar_image_names = [images[index] for index in similar_image_indices]
32Image.open(similar_image_names[0])
This model can be effectively employed for zero-shot image classification, as exemplified below:
1import requests
2from PIL import Image
3import matplotlib.pyplot as plt
4
5from transformers import VisionTextDualEncoderModel, VisionTextDualEncoderProcessor
6
7model = VisionTextDualEncoderModel.from_pretrained("kaveh/rclip")
8processor = VisionTextDualEncoderProcessor.from_pretrained("kaveh/rclip")
9
10url = "https://huggingface.co/spaces/kaveh/radiology-image-retrieval/resolve/main/images/ROCO_09402.jpg"
11image = Image.open(requests.get(url, stream=True).raw)
12possible_class_names = ["Chest X-Ray", "Brain MRI", "Abdominal CT Scan", "Ultrasound", "OPG"]
13
14inputs = processor(text=possible_class_names, images=image, return_tensors="pt", padding=True)
15probs = model(**inputs).logits_per_image.softmax(dim=1).squeeze()
16
17print("".join([x[0] + ": " + x[1] + "\n" for x in zip(possible_class_names, [format(prob, ".4%") for prob in probs])]))
18image
1@misc{https://doi.org/10.57967/hf/0896,
2 doi = {10.57967/HF/0896},
3 url = {https://huggingface.co/kaveh/rclip},
4 author = {{Kaveh Shahhosseini}},
5 title = {rclip},
6 publisher = {Hugging Face},
7 year = {2023}
8}