import torch import torch.nn.functional as F import numpy as np import pandas as pd import gradio as gr from transformers import CLIPModel, CLIPTokenizer from sklearn.metrics.pairwise import cosine_similarity from datasets import load_dataset from PIL import Image import traceback MODEL_ID = "openai/clip-vit-base-patch32" tokenizer = CLIPTokenizer.from_pretrained(MODEL_ID) model = CLIPModel.from_pretrained(MODEL_ID) model.eval() embeddings = np.load("car_embeddings.npz")["embeddings"] metadata = pd.read_csv("car_metadata.csv") ds = load_dataset("DamianBoborzi/car_images", split="train").select(range(10000)) top3_default_indices = ( metadata .nlargest(3, "aesthetic_score")["dataset_index"] .astype(int) .tolist() ) def get_default_images(): imgs = [] for ds_idx in top3_default_indices: img = ds[ds_idx]["image"] if not isinstance(img, Image.Image): img = Image.fromarray(img) imgs.append(img.convert("RGB")) return imgs def get_text_vector(user_query: str) -> np.ndarray: inputs = tokenizer(user_query, return_tensors="pt", padding=True, truncation=True, max_length=77) with torch.no_grad(): text_out = model.text_model(**inputs) feat = model.text_projection(text_out.pooler_output) return F.normalize(feat, p=2, dim=-1).cpu().numpy() def recommend(user_query: str): if not user_query or not user_query.strip(): defaults = get_default_images() rows = [metadata[metadata["dataset_index"] == i].iloc[0] for i in top3_default_indices] results = [] for rank, (img, row) in enumerate(zip(defaults, rows)): score = float(row["aesthetic_score"]) car_name = str(row["text"]).split(",")[0].split("(")[0].strip()[:55] label = f"#{rank+1} Aesthetic Score: {score:.2f}\n{car_name}" results += [img, label] return results[0], results[1], results[2], results[3], results[4], results[5] try: query_vec = get_text_vector(user_query) scores = cosine_similarity(query_vec, embeddings).flatten() top3 = scores.argsort()[-3:][::-1] results = [] for rank, df_idx in enumerate(top3): row = metadata.iloc[int(df_idx)] score = float(scores[int(df_idx)]) ds_idx = int(row["dataset_index"]) img = ds[ds_idx]["image"] if not isinstance(img, Image.Image): img = Image.fromarray(img) img = img.convert("RGB") car_name = str(row["text"]).split(",")[0].split("(")[0].strip()[:55] label = f"#{rank+1} Score: {score:.4f} | Cluster {int(row['cluster'])}\n{car_name}" results += [img, label] return results[0], results[1], results[2], results[3], results[4], results[5] except Exception as e: traceback.print_exc() msg = f"Error: {e}" return None, msg, None, msg, None, msg default_imgs = get_default_images() default_rows = [metadata[metadata["dataset_index"] == i].iloc[0] for i in top3_default_indices] default_labels = [ f"#{r+1} Aesthetic Score: {float(row['aesthetic_score']):.2f}\n" f"{str(row['text']).split(',')[0].split('(')[0].strip()[:55]}" for r, row in enumerate(default_rows) ] with gr.Blocks(theme=gr.themes.Soft()) as demo: gr.Markdown("# 🏎️ Car Recommendation Engine") gr.Markdown( "Describe the car you're looking for and CLIP will find the 3 most similar cars " "from a dataset of **10,000 real car images**.\n\n" "*Default view shows the top 3 highest-rated cars by aesthetic score.*" ) with gr.Row(): query_box = gr.Textbox( label="Describe your ideal car", placeholder="e.g. A red aggressive supercar with wide body kit", scale=5 ) search_btn = gr.Button("🔍 Find Matches", variant="primary", scale=1) gr.Markdown("### 🏆 Top 3 Matches") with gr.Row(): with gr.Column(): img1 = gr.Image(label="Match 1", type="pil", value=default_imgs[0]) lbl1 = gr.Textbox(show_label=False, interactive=False, value=default_labels[0]) with gr.Column(): img2 = gr.Image(label="Match 2", type="pil", value=default_imgs[1]) lbl2 = gr.Textbox(show_label=False, interactive=False, value=default_labels[1]) with gr.Column(): img3 = gr.Image(label="Match 3", type="pil", value=default_imgs[2]) lbl3 = gr.Textbox(show_label=False, interactive=False, value=default_labels[2]) search_btn.click(fn=recommend, inputs=query_box, outputs=[img1, lbl1, img2, lbl2, img3, lbl3]) gr.Markdown("---") gr.Markdown("## 📽️ Project Presentation Video") with gr.Row(): gr.HTML("""