CoolFace
Apppublic

shivangibithel/Text2ImageRetrieval

sourceHugging Facemitupdated 2y agoView on Hugging Face
1likes
app.py56 linesDownload Raw Back to root
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)