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("""
This app uses AI to find NASA space images based on what you describe.