import os import random import numpy as np import torch import pandas as pd from PIL import Image from datasets import load_dataset from transformers import CLIPModel, CLIPProcessor from sentence_transformers import SentenceTransformer from sklearn.metrics.pairwise import cosine_similarity from sklearn.decomposition import PCA from sklearn.cluster import KMeans import gradio as gr # ── Config ─────────────────────────────────────────────────────────────────── SEED = 42 SAMPLE_SIZE = 3000 N_CLUSTERS = 6 random.seed(SEED) np.random.seed(SEED) device = "cuda" if torch.cuda.is_available() else "cpu" print(f"Using device: {device}") # ── Load Dataset ────────────────────────────────────────────────────────────── print("Loading dataset...") dataset = load_dataset("UniverseTBD/AstroLLaVA_convos", split="train") print(f"Dataset loaded: {len(dataset)} images") df = pd.DataFrame({ 'id': dataset['id'], 'caption': dataset['caption'], 'url': dataset['url'], 'corpus': dataset['corpus'], }) # ── Filter ──────────────────────────────────────────────────────────────────── astronomical_keywords = [ 'galaxy', 'star', 'nebula', 'planet', 'comet', 'telescope', 'space', 'cosmic', 'solar', 'lunar', 'asteroid', 'supernova', 'black hole', 'orbit', 'astronomical', 'universe', 'cosmos', 'hubble', 'nasa', 'eso', 'cluster', 'quasar', 'pulsar' ] blocklist = [ 'people', 'person', 'team', 'staff', 'scientist', 'astronomer', 'conference', 'meeting', 'forum', 'ceremony', 'award', 'portrait', 'building', 'office', 'campus', 'university', 'school', 'logo', 'poster', 'banner', 'screenshot', 'website', 'chart', 'graph', 'diagram', 'illustration', 'crowd', 'audience', 'group photo' ] def is_astronomical(caption): return any(k in caption.lower() for k in astronomical_keywords) def is_pure_space(caption): return not any(k in caption.lower() for k in blocklist) df_filtered = df[df['caption'].apply(is_astronomical)].reset_index(drop=True) df_filtered = df_filtered[df_filtered['caption'].apply(is_pure_space)].reset_index(drop=True) print(f"Filtered dataset: {len(df_filtered)} images") # ── Sample ──────────────────────────────────────────────────────────────────── random.seed(SEED) filtered_indices = df_filtered.index.tolist() sample_indices = random.sample(filtered_indices, min(SAMPLE_SIZE, len(filtered_indices))) # ── Load Models ─────────────────────────────────────────────────────────────── print("Loading CLIP model...") clip_model = CLIPModel.from_pretrained("openai/clip-vit-base-patch32").to(device) processor = CLIPProcessor.from_pretrained("openai/clip-vit-base-patch32") print("Loading text model...") text_model = SentenceTransformer('all-MiniLM-L6-v2') # ── Generate Embeddings ─────────────────────────────────────────────────────── print("Generating embeddings...") from tqdm import tqdm embeddings = [] sample_captions = [] batch_size = 32 for i in tqdm(range(0, len(sample_indices), batch_size)): batch_indices = sample_indices[i:i+batch_size] images, valid_indices = [], [] for idx in batch_indices: try: img = dataset[idx]['image'].convert("RGB") images.append(img) valid_indices.append(idx) except: continue if not images: continue inputs = processor(images=images, return_tensors="pt", padding=True).to(device) with torch.no_grad(): img_outputs = clip_model.vision_model(**inputs) img_features = img_outputs.pooler_output img_features = clip_model.visual_projection(img_features) img_features = img_features / img_features.norm(dim=-1, keepdim=True) embeddings.append(img_features.cpu().numpy()) sample_captions.extend([df_filtered.iloc[idx]['caption'] for idx in valid_indices]) embeddings = np.vstack(embeddings) print(f"Embeddings shape: {embeddings.shape}") # ── Clustering ──────────────────────────────────────────────────────────────── print("Clustering...") pca = PCA(n_components=50, random_state=SEED) embeddings_pca = pca.fit_transform(embeddings) kmeans = KMeans(n_clusters=N_CLUSTERS, random_state=SEED, n_init=10) cluster_labels = kmeans.fit_predict(embeddings_pca) # ── Caption Embeddings ──────────────────────────────────────────────────────── print("Generating caption embeddings...") caption_embeddings = text_model.encode(sample_captions, batch_size=256, show_progress_bar=True) print("Ready!") # ── Helpers ─────────────────────────────────────────────────────────────────── cluster_names_map = { 0: '🌌 Galaxies & Deep Space', 1: '☁️ Nebulae & Gas Clouds', 2: '🪐 Solar System & Planets', 3: '⭐ Stars & Star Clusters', 4: '🔭 Earth & Observatories', 5: '🌠 Mixed Astronomical' } def get_color_bar(similarity): filled = int(similarity * 20) empty = 20 - filled emoji = "🟢" if similarity >= 0.70 else "🟡" if similarity >= 0.50 else "🔴" return f"{emoji} Similarity: {similarity:.1%}\n{'█' * filled}{'░' * empty}" def get_top3(query_text): query_emb = text_model.encode([query_text]) sims = cosine_similarity(query_emb, caption_embeddings)[0] top_idx = sims.argsort()[::-1][:3] results = [] for idx in top_idx: results.append({ 'image': dataset[sample_indices[idx]]['image'].convert("RGB"), 'caption': sample_captions[idx], 'similarity': float(sims[idx]), 'cluster': cluster_names_map[cluster_labels[idx]] }) return results def get_top3_by_image(query_image): img = query_image.convert("RGB") inputs = processor(images=[img], return_tensors="pt", padding=True).to(device) with torch.no_grad(): img_outputs = clip_model.vision_model(**inputs) img_features = img_outputs.pooler_output img_features = clip_model.visual_projection(img_features) img_features = img_features / img_features.norm(dim=-1, keepdim=True) query_emb = img_features.cpu().numpy() sims = cosine_similarity(query_emb, embeddings)[0] top_idx = sims.argsort()[::-1][:3] results = [] for idx in top_idx: results.append({ 'image': dataset[sample_indices[idx]]['image'].convert("RGB"), 'caption': sample_captions[idx], 'similarity': float(sims[idx]), 'cluster': cluster_names_map[cluster_labels[idx]] }) return results def search_by_text(query_text): if not query_text.strip(): return None, "", None, "", None, "" import re query_text = re.sub(r'[^\x00-\x7F]+', '', query_text).strip() if not query_text: return None, "", None, "", None, "" results = get_top3(query_text) images = [np.array(r['image']) for r in results] captions = [f"{get_color_bar(r['similarity'])}\n{r['cluster']}\n\n{r['caption'][:1000]}..." for r in results] return images[0], captions[0], images[1], captions[1], images[2], captions[2] def search_by_image(query_image): if query_image is None: return None, "", None, "", None, "" results = get_top3_by_image(query_image) images = [np.array(r['image']) for r in results] captions = [f"{get_color_bar(r['similarity'])}\n{r['cluster']}\n\n{r['caption'][:1000]}..." for r in results] return images[0], captions[0], images[1], captions[1], images[2], captions[2] def surprise_me(): queries = ["spiral galaxy", "colorful nebula", "supernova explosion", "black hole", "star cluster", "planet surface", "solar flare", "Hubble deep field", "comet tail", "rings of saturn", "andromeda galaxy", "mars crater", "neutron star", "milky way panorama", "aurora borealis", "planetary nebula"] query = random.choice(queries) results = get_top3(query) images = [np.array(r['image']) for r in results] captions = [f"{get_color_bar(r['similarity'])}\n{r['cluster']}\n\n{r['caption'][:1000]}..." for r in results] return query, images[0], captions[0], images[1], captions[1], images[2], captions[2] # ── UI ──────────────────────────────────────────────────────────────────────── css = """ body { background-color: #0a0a1a !important; } .gradio-container { background-color: #0a0a1a !important; color: #ffffff !important; } .gr-button-primary { background: linear-gradient(135deg, #1a1aff, #7b2ff7) !important; border: none !important; color: #ffffff !important; font-weight: bold !important; } .gr-button { background: #1a1a3e !important; color: #ffffff !important; border: 1px solid #7b2ff7 !important; } label { color: #ffffff !important; font-weight: 600 !important; } p, span, div, h1, h2, h3 { color: #ffffff !important; } footer { display: none !important; } """ with gr.Blocks(title="🔭 NASA Space Image Recommender", css=css) as demo: gr.Markdown(""" # 🔭 NASA Space Image Recommender ### Search 5,000 curated NASA & Hubble images using natural language or upload your own image *Powered by CLIP embeddings + Sentence Transformers | Dataset: AstroLLaVA (ESO, APOD, Hubble)* """) gr.Markdown("### 🎥 Project Presentation Video") gr.HTML("""
""") gr.HTML("""
ℹ️ How it works — click to learn more

This app uses AI to find NASA space images based on what you describe.

  1. 🧠 CLIP Model — converts your text into a 512-dimensional embedding vector
  2. 📊 Sentence Transformers — converts image captions into 384-dimensional vectors
  3. 📐 Cosine Similarity — measures how close your query is to each image caption
  4. 🎯 Top 3 — returns the most similar images with their cluster category
  5. 🖼️ Image Search — upload any image and CLIP finds visually similar NASA photos
""") with gr.Tabs(): with gr.Tab("🔍 Search by Text"): with gr.Row(): query_input = gr.Textbox(label="🌌 Describe what you want to see", placeholder="e.g. spiral galaxy, colorful nebula, rings of saturn...", scale=5) search_btn = gr.Button("Search 🚀", variant="primary", scale=1) surprise_btn = gr.Button("🎲 Surprise me!", scale=1) gr.Examples( examples=[ ["🌀 spiral galaxy"], ["☁️ colorful nebula"], ["🪐 rings of saturn"], ["🕳️ black hole"], ["⭐ star cluster"], ["💥 supernova explosion"], ["🔭 Hubble deep field"], ["🌌 aurora borealis"], ["🌌 andromeda galaxy"], ["🔴 mars crater"], ["☄️ comet tail"], ["🌠 milky way panorama"], ["💫 planetary nebula"], ["💥 galaxy collision"], ["☀️ solar flare"] ], inputs=query_input, examples_per_page=15 ) with gr.Row(): with gr.Column(): img1 = gr.Image(label="🥇 Best Match", height=280) cap1 = gr.Textbox(label="📖 Description", lines=5) with gr.Column(): img2 = gr.Image(label="🥈 Second Match", height=280) cap2 = gr.Textbox(label="📖 Description", lines=5) with gr.Column(): img3 = gr.Image(label="🥉 Third Match", height=280) cap3 = gr.Textbox(label="📖 Description", lines=5) search_btn.click(fn=search_by_text, inputs=query_input, outputs=[img1, cap1, img2, cap2, img3, cap3]) surprise_btn.click(fn=surprise_me, inputs=[], outputs=[query_input, img1, cap1, img2, cap2, img3, cap3]) with gr.Tab("🖼️ Search by Image"): gr.Markdown("### Upload any space image to find visually similar NASA images") with gr.Row(): img_input = gr.Image(label="Upload your image", height=300, type="pil") img_search_btn = gr.Button("Find Similar 🔭", variant="primary") with gr.Row(): with gr.Column(): img4 = gr.Image(label="🥇 Best Match", height=280) cap4 = gr.Textbox(label="📖 Description", lines=5) with gr.Column(): img5 = gr.Image(label="🥈 Second Match", height=280) cap5 = gr.Textbox(label="📖 Description", lines=5) with gr.Column(): img6 = gr.Image(label="🥉 Third Match", height=280) cap6 = gr.Textbox(label="📖 Description", lines=5) img_search_btn.click(fn=search_by_image, inputs=img_input, outputs=[img4, cap4, img5, cap5, img6, cap6]) demo.launch()