shivangibithel/Text2ImageRetrieval
1
1import streamlit as st2st.set_page_config(page_title='T2I', page_icon="🧊", layout='centered')3st.title("Text To Image Retrieval for KaggleX BPIOC Mentorship Program")4import torch5from transformers import AutoTokenizer, AutoModel6import faiss7import numpy as np8from PIL import Image9from sentence_transformers import SentenceTransformer10import json11import zipfile12 13# Map the image ids to the corresponding image URLs14image_map_name = 'captions.json'15 16with open(image_map_name, 'r') as f:17 caption_dict = json.load(f)18 19image_list = list(caption_dict.keys())20caption_list = list(caption_dict.values())21zip_path = "Images.zip"22zip_file = zipfile.ZipFile(zip_path)23 24model_name = "sentence-transformers/all-distilroberta-v1"25tokenizer = AutoTokenizer.from_pretrained(model_name)26model = SentenceTransformer(model_name)27# vectors = model.encode(caption_list)28vectors = np.load("./sbert_text_features.npy")29vector_dimension = vectors.shape[1]30index = faiss.IndexFlatIP(vector_dimension)31faiss.normalize_L2(vectors)32index.add(vectors)33 34def search(query, k=4):35 # Encode the query36 query_embedding = model.encode(query)37 query_vector = np.array([query_embedding])38 faiss.normalize_L2(query_vector)39 index.nprobe = index.ntotal40 41 # Search for the nearest neighbors in the FAISS index42 D, I = index.search(query_vector, k)43 44 # Map the image ids to the corresponding image URLs45 image_urls = []46 for i in I[0]:47 text_id = i48 image_id = str(image_list[i])49 image_data = zip_file.open("Images/" +image_id)50 image = Image.open(image_data)51 st.image(image, width=600)52 53query = st.text_input("Enter your search query here:")54if st.button("Search"):55 if query:56 search(query)