Spaces:
Running on Zero
Running on Zero
Download app.py from google/embedding-draw-challenge: direct link, hf CLI and curl.
- Browser
- Download file 43.8 kB
-
https://huggingface.co/spaces/google/embedding-draw-challenge/resolve/main/app.py
- Command line
-
hf download hf://spaces/google/embedding-draw-challenge/app.py
-
curl -L -o app.py https://huggingface.co/spaces/google/embedding-draw-challenge/resolve/main/app.py
43.8 kB
| import os | |
| import sys | |
| import subprocess | |
| from huggingface_hub import hf_hub_download | |
| token = os.environ.get("HF_TOKEN") | |
| transformers_wheel = hf_hub_download( | |
| repo_id="gg-hf-em/embeddinggemma-2-eap-extras", | |
| filename="transformers-5.18.0.dev0-py3-none-any.whl", | |
| repo_type="dataset", | |
| token=token | |
| ) | |
| # Install custom local wheels | |
| subprocess.run([sys.executable, "-m", "pip", "install", transformers_wheel], check=False) | |
| try: | |
| import spaces | |
| except ImportError: | |
| class _SpacesFallback: | |
| def GPU(*args, **kwargs): | |
| def _decorator(fn): | |
| return fn | |
| if len(args) == 1 and callable(args[0]) and not kwargs: | |
| return args[0] | |
| return _decorator | |
| spaces = _SpacesFallback() | |
| import base64 | |
| import io | |
| import inspect | |
| import time | |
| from typing import Dict, List, Any, Tuple | |
| import gradio as gr | |
| import numpy as np | |
| import torch | |
| from PIL import Image | |
| from sentence_transformers import SentenceTransformer | |
| # Optimize CPU threading for low-latency inference | |
| torch.set_num_threads(min(16, os.cpu_count() or 8)) | |
| # Curated Theme & Object Database from GDD Section 3 | |
| THEME_DATABASE: Dict[str, Dict[str, Any]] = { | |
| "Fruit": { | |
| "target": "Banana", | |
| "distractors": [ | |
| "Apple", | |
| "Strawberry", | |
| "Orange", | |
| "Watermelon", | |
| "Pineapple", | |
| "Grapes", | |
| "Cherries", | |
| ], | |
| "icon": "🍌", | |
| "description": "Draw organic silhouettes and vibrant fruit hues.", | |
| "hints": { | |
| "Banana": "Emphasize the curved crescent silhouette and bright yellow hue (#FFE135) with brown tips.", | |
| "Apple": "Draw a rounded red or green fruit body with a top indent and small leaf.", | |
| "Strawberry": "Sketch a tapered heart/cone shape in bright red with green leafy cap and seeds.", | |
| "Orange": "Draw a round citrus sphere in bright orange (#F97316) with subtle texture.", | |
| "Watermelon": "Draw a green rind slice or wedge with bright red flesh and black seeds.", | |
| "Pineapple": "Sketch an oval golden-yellow body with crosshatch diamond lines and a spiky green crown.", | |
| "Grapes": "Draw a clustered bunch of round purple or green berries hanging from a vine stem.", | |
| "Cherries": "Sketch a pair of glossy bright red round berries joined at the top by twin green stems.", | |
| }, | |
| }, | |
| "Vehicle": { | |
| "target": "Sports Car", | |
| "distractors": [ | |
| "Bicycle", | |
| "Airplane", | |
| "Bus", | |
| "Sailboat", | |
| "Helicopter", | |
| "Train", | |
| "Rocket", | |
| ], | |
| "icon": "🏎️", | |
| "description": "Capture mechanical geometry, wheels, and aerodynamic profiles.", | |
| "hints": { | |
| "Sports Car": "Highlight a low aerodynamic chassis in red, sloped windshield, and two distinct wheels.", | |
| "Bicycle": "Draw two spoked wheels connected by a triangular diamond frame and handlebars.", | |
| "Airplane": "Sketch a central fuselage tube with swept wings and a tail fin.", | |
| "Bus": "Draw a long rectangular body with a row of passenger windows and wheels.", | |
| "Sailboat": "Draw a curved boat hull floating on blue waves with a tall mast and triangular sails.", | |
| "Helicopter": "Sketch a rounded cockpit cabin with wide horizontal top rotor blades and a tail boom.", | |
| "Train": "Draw a locomotive engine on rails with multiple wheels, a cab, smokestack, and front grill.", | |
| "Rocket": "Sketch a vertical cylindrical body with a pointed nose cone, round window, fins, and orange exhaust flame.", | |
| }, | |
| }, | |
| "Animal": { | |
| "target": "Elephant", | |
| "distractors": [ | |
| "Giraffe", | |
| "Lion", | |
| "Cat", | |
| "Penguin", | |
| "Butterfly", | |
| "Turtle", | |
| "Rabbit", | |
| ], | |
| "icon": "🐘", | |
| "description": "Sketch distinctive anatomical traits: trunks, ears, necks, and manes.", | |
| "hints": { | |
| "Elephant": "Focus on the large grey body, long curved trunk, wide floppy ear, and sturdy legs.", | |
| "Giraffe": "Emphasize an elongated vertical neck, spotted pattern, and long slender legs.", | |
| "Lion": "Draw a bold circular mane surrounding a feline face and golden body.", | |
| "Cat": "Sketch pointed triangular ears, whiskers, a curved tail, and agile feline silhouette.", | |
| "Penguin": "Draw an upright black body with a white belly oval, orange beak, side flippers, and orange feet.", | |
| "Butterfly": "Sketch large symmetrical colorful wings extending from a slender central body with two antennae.", | |
| "Turtle": "Draw a rounded dome shell with hexagonal patterns, four stubby legs, and an outstretched head.", | |
| "Rabbit": "Sketch a compact body with two long upright ears, a round nose, whiskers, and a fluffy tail.", | |
| }, | |
| }, | |
| "Country (Shape/Icon)": { | |
| "target": "Korea", | |
| "distractors": [ | |
| "Japan", | |
| "Brazil", | |
| "Australia", | |
| "Canada", | |
| "Italy", | |
| "United States", | |
| "France", | |
| ], | |
| "icon": "🇰🇷", | |
| "description": "Draw iconic national symbols, flags, or geographic silhouettes.", | |
| "hints": { | |
| "Korea": "Draw the iconic red/blue Taegeuk swirl circle with black corner trigrams (Taegeukgi) or peninsula shape.", | |
| "Japan": "Draw a crisp red sun disc centered on a white rectangular flag canvas.", | |
| "Brazil": "Sketch a green flag with a yellow diamond and blue celestial globe in the center.", | |
| "Australia": "Draw the distinctive Southern Cross stars, Union Jack corner, or continental island outline.", | |
| "Canada": "Draw a white flag center flanked by vertical red bars with a bold red maple leaf in the middle.", | |
| "Italy": "Sketch the vertical green, white, and red tricolor flag or the iconic boot-shaped peninsula.", | |
| "United States": "Draw horizontal red and white stripes with a blue top-left canton containing white stars.", | |
| "France": "Sketch the vertical blue, white, and red tricolor flag or the iconic Eiffel Tower silhouette.", | |
| }, | |
| }, | |
| } | |
| class EmbeddingDrawEngine: | |
| """Multimodal vector embedding engine powered by SentenceTransformers & EmbeddingGemma 2.""" | |
| def __init__(self, model_path: str = "./embeddinggemma-2", cache_path: str = ".cache/reference_vectors.npz"): | |
| print(f"[EmbeddingEngine] Initializing SentenceTransformer from {model_path}...") | |
| t0 = time.time() | |
| if not os.path.exists(model_path): | |
| model_path = "gg-hf-em/embeddinggemma-2" | |
| self.cache_path = cache_path | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| self.model = SentenceTransformer( | |
| model_path, | |
| truncate_dim=128, | |
| similarity_fn_name="dot", | |
| device=device, | |
| config_kwargs={"audio_config": None}, | |
| ) | |
| # Optimize image processor token length for real-time stroke evaluation (< 0.5s latency) | |
| if hasattr(self.model[0], "processor") and hasattr(self.model[0].processor, "image_processor"): | |
| self.model[0].processor.image_processor.max_soft_tokens = 70 | |
| self.model[0].processor.image_processor.image_seq_length = 70 | |
| # Hook to capture authentic pre-normalization L2 vector norm from Pooling layer | |
| self._last_raw_norm: float = 479.69 | |
| if len(self.model) > 1: | |
| self.model[1].register_forward_hook(self._capture_norm_hook) | |
| # Collect all unique objects across themes | |
| all_objects = set() | |
| for theme_data in THEME_DATABASE.values(): | |
| all_objects.add(theme_data["target"]) | |
| all_objects.update(theme_data["distractors"]) | |
| self.sorted_objects: List[str] = sorted(list(all_objects)) | |
| # Initialize attributes in main process so they always exist even before first GPU call | |
| self.text_embeddings: Dict[str, np.ndarray] = {} | |
| self.blank_embedding: np.ndarray | None = None | |
| self.blank_similarities: Dict[str, float] = {obj: 0.6100 for obj in self.sorted_objects} | |
| print(f"[EmbeddingEngine] Model initialized in {time.time() - t0:.2f}s (device={self.model.device})") | |
| def _compute_reference_vectors(self) -> Dict[str, Any]: | |
| """Encode reference text prompts and blank canvas baseline image.""" | |
| prompts = [f"A simple drawing of {obj}" for obj in self.sorted_objects] | |
| encoded_texts = self.model.encode(prompts, normalize_embeddings=True) | |
| text_embeddings = { | |
| obj: np.array(emb, dtype=np.float32) | |
| for obj, emb in zip(self.sorted_objects, encoded_texts) | |
| } | |
| blank_img = Image.new("RGB", (240, 240), color=(255, 255, 255)) | |
| blank_embedding = np.array( | |
| self.model.encode({"image": blank_img}, normalize_embeddings=True), | |
| dtype=np.float32, | |
| ) | |
| raw_norm = float(self._last_raw_norm) | |
| blank_similarities = { | |
| obj: float(np.dot(blank_embedding, text_emb)) | |
| for obj, text_emb in text_embeddings.items() | |
| } | |
| return { | |
| "text_embeddings": text_embeddings, | |
| "blank_embedding": blank_embedding, | |
| "blank_similarities": blank_similarities, | |
| "raw_norm": raw_norm, | |
| } | |
| def _load_cache_from_disk(self) -> bool: | |
| """Attempt to load precomputed reference vectors from disk cache.""" | |
| if not self.cache_path or not os.path.exists(self.cache_path): | |
| return False | |
| try: | |
| with np.load(self.cache_path) as data: | |
| if "blank_embedding" not in data or "blank_raw_norm" not in data: | |
| return False | |
| loaded_text_embs: Dict[str, np.ndarray] = {} | |
| for obj in self.sorted_objects: | |
| key = f"text_emb_{obj}" | |
| if key not in data: | |
| return False | |
| loaded_text_embs[obj] = np.array(data[key], dtype=np.float32) | |
| self.blank_embedding = np.array(data["blank_embedding"], dtype=np.float32) | |
| self._last_raw_norm = float(data["blank_raw_norm"]) | |
| self.text_embeddings = loaded_text_embs | |
| self.blank_similarities = { | |
| obj: float(np.dot(self.blank_embedding, text_emb)) | |
| for obj, text_emb in self.text_embeddings.items() | |
| } | |
| print(f"[EmbeddingEngine] Loaded cached reference vectors from {self.cache_path}") | |
| return True | |
| except Exception as e: | |
| print(f"[EmbeddingEngine] Failed to load reference cache ({e}), rebuilding...") | |
| return False | |
| def _save_cache_to_disk(self) -> None: | |
| """Save computed reference vectors to disk cache.""" | |
| if not self.cache_path or self.blank_embedding is None or not self.text_embeddings: | |
| return | |
| try: | |
| cache_dir = os.path.dirname(self.cache_path) | |
| if cache_dir: | |
| os.makedirs(cache_dir, exist_ok=True) | |
| save_dict: Dict[str, Any] = { | |
| "blank_embedding": self.blank_embedding, | |
| "blank_raw_norm": np.array(self._last_raw_norm, dtype=np.float32), | |
| } | |
| for obj, emb in self.text_embeddings.items(): | |
| save_dict[f"text_emb_{obj}"] = emb | |
| np.savez(self.cache_path, **save_dict) | |
| print(f"[EmbeddingEngine] Saved reference vectors cache to {self.cache_path}") | |
| except Exception as e: | |
| print(f"[EmbeddingEngine] Warning: could not write cache file {self.cache_path}: {e}") | |
| def ensure_reference_embeddings(self) -> None: | |
| """Build and cache reference text embeddings and blank canvas baseline in memory and on disk.""" | |
| if self.text_embeddings and self.blank_embedding is not None: | |
| return | |
| if self._load_cache_from_disk(): | |
| return | |
| t0 = time.time() | |
| print(f"[EmbeddingEngine] Building reference vectors for {len(self.sorted_objects)} objects...") | |
| try: | |
| res = _gpu_build_reference_embeddings() | |
| except Exception: | |
| res = self._compute_reference_vectors() | |
| self.text_embeddings = res["text_embeddings"] | |
| self.blank_embedding = res["blank_embedding"] | |
| self.blank_similarities = res["blank_similarities"] | |
| self._last_raw_norm = float(res["raw_norm"]) | |
| self._save_cache_to_disk() | |
| print(f"[EmbeddingEngine] Reference vectors built and cached in {time.time() - t0:.2f}s") | |
| def _capture_norm_hook(self, module, inputs, output): | |
| if hasattr(output, "__getitem__") and "sentence_embedding" in output: | |
| tensor = output["sentence_embedding"] | |
| self._last_raw_norm = float(torch.norm(tensor, p=2, dim=-1).mean().item()) | |
| def compute_ink_ratio(img: Image.Image) -> float: | |
| """Compute fraction of non-white pixels on canvas to detect blank or near-blank state.""" | |
| gray = img.convert("L") | |
| arr = np.array(gray, dtype=np.uint8) | |
| non_white = np.sum(arr < 248) | |
| return float(non_white) / float(arr.size) | |
| def evaluate_drawing( | |
| self, | |
| image_input: Image.Image | str | None, | |
| theme_name: str, | |
| target_object: str, | |
| ) -> Dict[str, Any]: | |
| """Evaluate canvas image against target and comparing objects using Cosine Similarity.""" | |
| t_start = time.time() | |
| if theme_name not in THEME_DATABASE: | |
| theme_name = "Fruit" | |
| theme_info = THEME_DATABASE[theme_name] | |
| all_theme_objects = [theme_info["target"]] + theme_info["distractors"] | |
| if target_object not in all_theme_objects: | |
| target_object = theme_info["target"] | |
| distractors = [o for o in all_theme_objects if o != target_object] | |
| ordered_objects = [target_object] + distractors | |
| if image_input is None: | |
| ink_ratio = 0.0 | |
| resized_img = None | |
| else: | |
| if isinstance(image_input, str): | |
| if "," in image_input: | |
| image_input = image_input.split(",", 1)[1] | |
| img_bytes = base64.b64decode(image_input) | |
| raw_img = Image.open(io.BytesIO(img_bytes)) | |
| else: | |
| raw_img = image_input | |
| # Composite onto solid white RGB canvas | |
| rgb_img = Image.new("RGB", raw_img.size, (255, 255, 255)) | |
| if raw_img.mode in ("RGBA", "LA") or (raw_img.mode == "P" and "transparency" in raw_img.info): | |
| rgba = raw_img.convert("RGBA") | |
| rgb_img.paste(rgba, mask=rgba.split()[3]) | |
| else: | |
| rgb_img.paste(raw_img.convert("RGB")) | |
| resized_img = rgb_img.resize((240, 240), Image.Resampling.BILINEAR) | |
| ink_ratio = self.compute_ink_ratio(resized_img) | |
| # If canvas is blank or virtually blank (< 0.05% ink), return clean zero state immediately without GPU call | |
| if resized_img is None or ink_ratio < 0.0005: | |
| latency_ms = int((time.time() - t_start) * 1000) | |
| scores_list = [ | |
| { | |
| "name": obj, | |
| "is_target": (obj == target_object), | |
| "percentage": 0.0, | |
| "raw_cosine": round(self.blank_similarities.get(obj, 0.61), 4), | |
| "delta": 0.0, | |
| } | |
| for obj in ordered_objects | |
| ] | |
| return { | |
| "theme": theme_name, | |
| "target": target_object, | |
| "scores": scores_list, | |
| "target_score": 0.0, | |
| "win": False, | |
| "telemetry": { | |
| "loss": 1.3863, | |
| "norm": round(self._last_raw_norm, 2), | |
| "cosine": round(self.blank_similarities.get(target_object, 0.61), 4), | |
| "latency_ms": latency_ms, | |
| }, | |
| "insight": f"Canvas ready. Start drawing **{target_object}**! {theme_info['hints'].get(target_object, '')}", | |
| } | |
| # Execute multimodal embedding inference on ZeroGPU worker | |
| gpu_result = _gpu_run_inference(resized_img, ordered_objects) | |
| self.blank_similarities.update(gpu_result["all_blank_similarities"]) | |
| raw_norm = float(gpu_result["raw_norm"]) | |
| self._last_raw_norm = raw_norm | |
| # Calculate Cosine Similarities and baseline deltas | |
| raw_cosines = np.array(gpu_result["raw_cosines"], dtype=np.float32) | |
| blank_cosines = np.array(gpu_result["blank_cosines"], dtype=np.float32) | |
| deltas = raw_cosines - blank_cosines | |
| # Contrastive Cross-Entropy Loss telemetry (temperature tau = 0.025) | |
| tau = 0.025 | |
| scaled_deltas = deltas / tau | |
| exp_deltas = np.exp(scaled_deltas - np.max(scaled_deltas)) | |
| softmax_probs = exp_deltas / np.sum(exp_deltas) | |
| contrastive_loss = float(-np.log(max(softmax_probs[0], 1e-7))) | |
| # Map cosine similarity & relative semantic dominance to 0-100% scale | |
| c_delta = deltas - np.mean(deltas) | |
| c_raw = raw_cosines - np.mean(raw_cosines) | |
| combined_signal = 0.65 * c_delta + 0.35 * c_raw | |
| logits = combined_signal / max(float(np.std(combined_signal)), 0.005) * 1.75 | |
| exp_l = np.exp(logits - np.max(logits)) | |
| probs = exp_l / np.sum(exp_l) | |
| # Positive gain boost over blank canvas | |
| gain_boost = np.clip((deltas - np.min(deltas)) / max(float(np.ptp(deltas)), 0.012), 0.0, 1.0) | |
| # Direct absolute delta reward so strong semantic matches reach 92-98% | |
| abs_gain_bonus = np.clip(deltas / 0.045, -0.2, 1.0) * 18.0 | |
| pcts = 14.0 + 54.0 * probs + 22.0 * gain_boost + abs_gain_bonus | |
| # Smooth ink warm-up ramp so single tiny dots don't spike to 90% | |
| ink_warmup = min(1.0, ink_ratio / 0.006) | |
| pcts = pcts * ink_warmup | |
| pcts = np.clip(pcts, 1.0, 99.2) | |
| scores_list = [] | |
| for idx, obj in enumerate(ordered_objects): | |
| scores_list.append( | |
| { | |
| "name": obj, | |
| "is_target": (idx == 0), | |
| "percentage": round(float(pcts[idx]), 1), | |
| "raw_cosine": round(float(raw_cosines[idx]), 4), | |
| "delta": round(float(deltas[idx]), 4), | |
| } | |
| ) | |
| target_score = scores_list[0]["percentage"] | |
| is_win = bool(target_score >= 90.0) | |
| # Determine dynamic Insight message | |
| best_distractor = max(scores_list[1:], key=lambda x: x["percentage"]) | |
| hint_text = theme_info["hints"].get(target_object, "Keep refining your drawing!") | |
| if is_win: | |
| insight_msg = ( | |
| f"🎯 Semantic Lock Achieved ({target_score:.1f}%)! EmbeddingGemma 2 strongly aligns your sketch " | |
| f"with 'A simple drawing of {target_object}'." | |
| ) | |
| elif best_distractor["percentage"] > target_score + 5.0: | |
| insight_msg = ( | |
| f"⚠️ Vector Drift Alert: Your sketch currently resembles **{best_distractor['name']}** " | |
| f"({best_distractor['percentage']:.1f}%) more than **{target_object}** ({target_score:.1f}%). " | |
| f"Tip: {hint_text}" | |
| ) | |
| elif target_score >= 70.0: | |
| insight_msg = ( | |
| f"🔥 High Semantic Alignment ({target_score:.1f}%)! You are very close to the 90% threshold. " | |
| f"Add a few more defining details or color fill to lock in **{target_object}**!" | |
| ) | |
| elif target_score >= 40.0: | |
| insight_msg = ( | |
| f"📈 Promising Trajectory ({target_score:.1f}%). **{target_object}** is emerging in the embedding space. " | |
| f"Tip: {hint_text}" | |
| ) | |
| else: | |
| insight_msg = f"✏️ Keep drawing! {hint_text}" | |
| latency_ms = int((time.time() - t_start) * 1000) | |
| return { | |
| "theme": theme_name, | |
| "target": target_object, | |
| "scores": scores_list, | |
| "target_score": target_score, | |
| "win": is_win, | |
| "telemetry": { | |
| "loss": round(contrastive_loss, 4), | |
| "norm": round(raw_norm, 2), | |
| "cosine": round(float(raw_cosines[0]), 4), | |
| "latency_ms": latency_ms, | |
| }, | |
| "insight": insight_msg, | |
| } | |
| # Instantiate global EmbeddingEngine singleton | |
| ENGINE = EmbeddingDrawEngine() | |
| def _gpu_build_reference_embeddings() -> Dict[str, Any]: | |
| """Standalone ZeroGPU worker task: builds reference text and blank canvas embeddings on GPU.""" | |
| return ENGINE._compute_reference_vectors() | |
| def _gpu_run_inference(resized_img: Image.Image, ordered_objects: List[str]) -> Dict[str, Any]: | |
| """Standalone ZeroGPU worker task: encodes player drawing and computes cosines on GPU.""" | |
| ENGINE.ensure_reference_embeddings() | |
| player_emb = np.array(ENGINE.model.encode({"image": resized_img}, normalize_embeddings=True), dtype=np.float32) | |
| raw_norm = float(ENGINE._last_raw_norm) | |
| raw_cosines = [float(np.dot(player_emb, ENGINE.text_embeddings[o])) for o in ordered_objects] | |
| blank_cosines = [float(ENGINE.blank_similarities[o]) for o in ordered_objects] | |
| return { | |
| "raw_cosines": raw_cosines, | |
| "blank_cosines": blank_cosines, | |
| "raw_norm": raw_norm, | |
| "all_blank_similarities": dict(ENGINE.blank_similarities), | |
| } | |
| def extract_pil_image(sketch_data: Any) -> Image.Image | None: | |
| """Extract composite PIL Image from Gradio Sketchpad output.""" | |
| if sketch_data is None: | |
| return None | |
| if isinstance(sketch_data, Image.Image): | |
| return sketch_data | |
| if isinstance(sketch_data, dict): | |
| img = sketch_data.get("composite") or sketch_data.get("background") | |
| if isinstance(img, Image.Image): | |
| return img | |
| if isinstance(img, np.ndarray): | |
| return Image.fromarray(img) | |
| return None | |
| if isinstance(sketch_data, np.ndarray): | |
| return Image.fromarray(sketch_data) | |
| return None | |
| def get_theme_objects(theme_name: str) -> List[str]: | |
| info = THEME_DATABASE.get(theme_name, THEME_DATABASE["Fruit"]) | |
| return [info["target"]] + info["distractors"] | |
| def format_scoreboard(result: Dict[str, Any]) -> Dict[str, float]: | |
| scores = result.get("scores", []) | |
| if not scores or max(round(item["percentage"]) for item in scores) <= 0: | |
| return {} | |
| return { | |
| f"🎯 {item['name'].upper()} (TARGET)" if item["is_target"] else item["name"]: round(item["percentage"] / 100.0, 4) | |
| for item in scores | |
| } | |
| def format_status_message(result: Dict[str, Any]) -> str: | |
| if result.get("win", False): | |
| return ( | |
| f"### 🎉 CHALLENGE CLEAR! ({result['target_score']:.1f}%)\n" | |
| f"You successfully captured the essence of **{result['target']}**! " | |
| f"The vector embedding recognizes your artistic intent.\n\n" | |
| f"{result['insight']}" | |
| ) | |
| return f"### 💡 Hint & Status\n{result['insight']}" | |
| SCOREBOARD_TARGET_CSS = """ | |
| <style> | |
| /* Highlight the Target Object row in the Similarity Scoreboard (light & dark mode) */ | |
| .confidence-set[data-testid^="🎯"] { | |
| --stat-background-fill: #10B981 !important; | |
| background: rgba(16, 185, 129, 0.10) !important; | |
| border: 1.5px solid rgba(16, 185, 129, 0.45) !important; | |
| border-radius: 10px !important; | |
| padding: 8px 12px !important; | |
| margin-bottom: 10px !important; | |
| box-shadow: 0 2px 8px rgba(16, 185, 129, 0.12) !important; | |
| } | |
| .confidence-set[data-testid^="🎯"] .bar { | |
| height: 8px !important; | |
| background: #10B981 !important; | |
| } | |
| .confidence-set[data-testid^="🎯"] .text, | |
| .confidence-set[data-testid^="🎯"] .confidence { | |
| font-weight: 800 !important; | |
| color: #047857 !important; | |
| font-size: 1.04em !important; | |
| } | |
| .confidence-set[data-testid^="🎯"] .line { | |
| border-color: rgba(16, 185, 129, 0.45) !important; | |
| } | |
| @media (prefers-color-scheme: dark) { | |
| :root:not(.light) .confidence-set[data-testid^="🎯"], | |
| body:not(.light) .confidence-set[data-testid^="🎯"] { | |
| background: rgba(16, 185, 129, 0.18) !important; | |
| border-color: rgba(52, 211, 153, 0.55) !important; | |
| box-shadow: 0 2px 12px rgba(16, 185, 129, 0.22) !important; | |
| } | |
| :root:not(.light) .confidence-set[data-testid^="🎯"] .text, | |
| :root:not(.light) .confidence-set[data-testid^="🎯"] .confidence, | |
| body:not(.light) .confidence-set[data-testid^="🎯"] .text, | |
| body:not(.light) .confidence-set[data-testid^="🎯"] .confidence { | |
| color: #34D399 !important; | |
| } | |
| } | |
| .dark .confidence-set[data-testid^="🎯"], | |
| :root.dark .confidence-set[data-testid^="🎯"], | |
| body.dark .confidence-set[data-testid^="🎯"] { | |
| background: rgba(16, 185, 129, 0.18) !important; | |
| border-color: rgba(52, 211, 153, 0.55) !important; | |
| box-shadow: 0 2px 12px rgba(16, 185, 129, 0.22) !important; | |
| } | |
| .dark .confidence-set[data-testid^="🎯"] .text, | |
| .dark .confidence-set[data-testid^="🎯"] .confidence, | |
| :root.dark .confidence-set[data-testid^="🎯"] .text, | |
| :root.dark .confidence-set[data-testid^="🎯"] .confidence, | |
| body.dark .confidence-set[data-testid^="🎯"] .text, | |
| body.dark .confidence-set[data-testid^="🎯"] .confidence { | |
| color: #34D399 !important; | |
| } | |
| </style> | |
| """ | |
| def format_popup_dialog(result: Dict[str, Any]) -> str: | |
| """Render a celebratory 'Challenge Clear' popup dialog (supporting light & dark mode) when similarity reaches >= 90%.""" | |
| if not result.get("win", False): | |
| return SCOREBOARD_TARGET_CSS | |
| target = result.get("target", "Target") | |
| score = float(result.get("target_score", 90.0)) | |
| cosine = result.get("telemetry", {}).get("cosine", 0.0) | |
| ts = time.time_ns() | |
| return SCOREBOARD_TARGET_CSS + f""" | |
| <style> | |
| @keyframes ccOverlayFadeIn {{ | |
| from {{ opacity: 0; }} | |
| to {{ opacity: 1; }} | |
| }} | |
| @keyframes ccModalPopIn {{ | |
| 0% {{ opacity: 0; transform: scale(0.88) translateY(16px); }} | |
| 70% {{ transform: scale(1.02) translateY(-2px); }} | |
| 100% {{ opacity: 1; transform: scale(1) translateY(0); }} | |
| }} | |
| #challenge-clear-overlay {{ | |
| --cc-overlay-bg: rgba(15, 23, 42, 0.60); | |
| --cc-card-bg: var(--block-background-fill, #FFFFFF); | |
| --cc-card-border: rgba(16, 185, 129, 0.30); | |
| --cc-card-shadow: 0 25px 50px -12px rgba(0, 0, 0, 0.28); | |
| --cc-text-primary: var(--body-text-color, #0F172A); | |
| --cc-text-secondary: #475569; | |
| --cc-text-muted: #64748B; | |
| --cc-surface-bg: var(--background-fill-secondary, #F8FAFC); | |
| --cc-surface-border: var(--border-color-primary, #E2E8F0); | |
| --cc-ring-track: #E2E8F0; | |
| --cc-accent: #10B981; | |
| --cc-accent-hover: #059669; | |
| --cc-accent-text: #059669; | |
| --cc-badge-bg: #ECFDF5; | |
| --cc-badge-text: #047857; | |
| --cc-badge-border: rgba(16, 185, 129, 0.25); | |
| --cc-btn-sec-bg: #F1F5F9; | |
| --cc-btn-sec-bg-hover: #E2E8F0; | |
| --cc-btn-sec-text: #1E293B; | |
| --cc-btn-sec-border: #CBD5E1; | |
| --cc-close-bg: #F1F5F9; | |
| --cc-close-bg-hover: #E2E8F0; | |
| --cc-close-text: #64748B; | |
| --cc-ghost-hover-bg: #F1F5F9; | |
| }} | |
| @media (prefers-color-scheme: dark) {{ | |
| :root:not(.light) #challenge-clear-overlay, | |
| body:not(.light) #challenge-clear-overlay {{ | |
| --cc-overlay-bg: rgba(2, 6, 23, 0.78); | |
| --cc-card-bg: var(--block-background-fill, #1E293B); | |
| --cc-card-border: rgba(52, 211, 153, 0.35); | |
| --cc-card-shadow: 0 25px 55px -12px rgba(0, 0, 0, 0.70); | |
| --cc-text-primary: var(--body-text-color, #F8FAFC); | |
| --cc-text-secondary: #CBD5E1; | |
| --cc-text-muted: #94A3B8; | |
| --cc-surface-bg: var(--background-fill-secondary, #0F172A); | |
| --cc-surface-border: var(--border-color-primary, #334155); | |
| --cc-ring-track: #334155; | |
| --cc-accent: #10B981; | |
| --cc-accent-hover: #059669; | |
| --cc-accent-text: #34D399; | |
| --cc-badge-bg: rgba(16, 185, 129, 0.16); | |
| --cc-badge-text: #6EE7B7; | |
| --cc-badge-border: rgba(52, 211, 153, 0.30); | |
| --cc-btn-sec-bg: #334155; | |
| --cc-btn-sec-bg-hover: #475569; | |
| --cc-btn-sec-text: #F8FAFC; | |
| --cc-btn-sec-border: #475569; | |
| --cc-close-bg: #334155; | |
| --cc-close-bg-hover: #475569; | |
| --cc-close-text: #CBD5E1; | |
| --cc-ghost-hover-bg: rgba(255, 255, 255, 0.08); | |
| }} | |
| }} | |
| .dark #challenge-clear-overlay, | |
| :root.dark #challenge-clear-overlay, | |
| body.dark #challenge-clear-overlay {{ | |
| --cc-overlay-bg: rgba(2, 6, 23, 0.78); | |
| --cc-card-bg: var(--block-background-fill, #1E293B); | |
| --cc-card-border: rgba(52, 211, 153, 0.35); | |
| --cc-card-shadow: 0 25px 55px -12px rgba(0, 0, 0, 0.70); | |
| --cc-text-primary: var(--body-text-color, #F8FAFC); | |
| --cc-text-secondary: #CBD5E1; | |
| --cc-text-muted: #94A3B8; | |
| --cc-surface-bg: var(--background-fill-secondary, #0F172A); | |
| --cc-surface-border: var(--border-color-primary, #334155); | |
| --cc-ring-track: #334155; | |
| --cc-accent: #10B981; | |
| --cc-accent-hover: #059669; | |
| --cc-accent-text: #34D399; | |
| --cc-badge-bg: rgba(16, 185, 129, 0.16); | |
| --cc-badge-text: #6EE7B7; | |
| --cc-badge-border: rgba(52, 211, 153, 0.30); | |
| --cc-btn-sec-bg: #334155; | |
| --cc-btn-sec-bg-hover: #475569; | |
| --cc-btn-sec-text: #F8FAFC; | |
| --cc-btn-sec-border: #475569; | |
| --cc-close-bg: #334155; | |
| --cc-close-bg-hover: #475569; | |
| --cc-close-text: #CBD5E1; | |
| --cc-ghost-hover-bg: rgba(255, 255, 255, 0.08); | |
| }} | |
| #challenge-clear-overlay .cc-close-btn:hover {{ | |
| background: var(--cc-close-bg-hover) !important; | |
| color: var(--cc-text-primary) !important; | |
| }} | |
| #challenge-clear-overlay .cc-btn-primary:hover {{ | |
| background: var(--cc-accent-hover) !important; | |
| transform: translateY(-1px); | |
| }} | |
| #challenge-clear-overlay .cc-btn-secondary:hover {{ | |
| background: var(--cc-btn-sec-bg-hover) !important; | |
| transform: translateY(-1px); | |
| }} | |
| #challenge-clear-overlay .cc-btn-ghost:hover {{ | |
| background: var(--cc-ghost-hover-bg) !important; | |
| color: var(--cc-text-primary) !important; | |
| }} | |
| </style> | |
| <div id="challenge-clear-overlay" data-ts="{ts}" role="dialog" aria-modal="true" aria-labelledby="challenge-clear-title" | |
| onclick="if (event.target === this) this.style.display = 'none';" | |
| style="position: fixed; inset: 0; z-index: 99999; background: var(--cc-overlay-bg); backdrop-filter: blur(6px); -webkit-backdrop-filter: blur(6px); display: flex; align-items: center; justify-content: center; padding: 20px; animation: ccOverlayFadeIn 0.2s ease-out;"> | |
| <div style="background: var(--cc-card-bg); color: var(--cc-text-primary); width: 100%; max-width: 440px; border-radius: 20px; padding: 32px 28px 26px; box-shadow: var(--cc-card-shadow), 0 0 0 1px var(--cc-card-border); text-align: center; position: relative; font-family: system-ui, -apple-system, sans-serif; animation: ccModalPopIn 0.32s cubic-bezier(0.16, 1, 0.3, 1);"> | |
| <button type="button" class="cc-close-btn" aria-label="Close dialog" | |
| onclick="document.getElementById('challenge-clear-overlay').style.display='none';" | |
| style="position: absolute; top: 14px; right: 14px; width: 32px; height: 32px; border-radius: 50%; border: none; background: var(--cc-close-bg); color: var(--cc-close-text); font-size: 16px; font-weight: 600; cursor: pointer; display: flex; align-items: center; justify-content: center; line-height: 1; transition: all 0.15s ease;"> | |
| ✕ | |
| </button> | |
| <div style="width: 112px; height: 112px; margin: 0 auto 18px; border-radius: 50%; background: conic-gradient(var(--cc-accent) {min(score, 100.0):.1f}%, var(--cc-ring-track) 0); display: flex; align-items: center; justify-content: center; box-shadow: 0 10px 25px -5px rgba(16, 185, 129, 0.35);"> | |
| <div style="width: 92px; height: 92px; border-radius: 50%; background: var(--cc-card-bg); display: flex; flex-direction: column; align-items: center; justify-content: center;"> | |
| <span style="font-size: 24px; font-weight: 800; color: var(--cc-accent-text); font-family: ui-monospace, SFMono-Regular, Menlo, monospace; line-height: 1.1;">{score:.1f}%</span> | |
| <span style="font-size: 10px; font-weight: 700; letter-spacing: 0.08em; text-transform: uppercase; color: var(--cc-text-muted); margin-top: 2px;">Similarity</span> | |
| </div> | |
| </div> | |
| <div style="display: inline-block; padding: 4px 12px; border-radius: 9999px; background: var(--cc-badge-bg); color: var(--cc-badge-text); border: 1px solid var(--cc-badge-border); font-size: 11px; font-weight: 700; letter-spacing: 0.1em; text-transform: uppercase; margin-bottom: 10px;"> | |
| 🎉 Goal Achieved (≥ 90%) | |
| </div> | |
| <h2 id="challenge-clear-title" style="margin: 0 0 10px; font-size: 26px; font-weight: 800; letter-spacing: -0.02em; color: var(--cc-text-primary); line-height: 1.2;"> | |
| CHALLENGE CLEAR! | |
| </h2> | |
| <p style="margin: 0 0 18px; font-size: 14.5px; line-height: 1.55; color: var(--cc-text-secondary);"> | |
| You successfully captured the essence of <strong style="color: var(--cc-text-primary); font-weight: 700;">{target}</strong>! | |
| The vector embedding recognizes your artistic intent. | |
| </p> | |
| <div style="background: var(--cc-surface-bg); border: 1px solid var(--cc-surface-border); border-radius: 12px; padding: 10px 14px; margin-bottom: 22px; font-size: 12.5px; color: var(--cc-text-secondary); font-family: ui-monospace, SFMono-Regular, Menlo, monospace;"> | |
| Target: <strong style="color: var(--cc-text-primary);">{target}</strong> • Raw Cosine: <strong style="color: var(--cc-text-primary);">{cosine:.4f}</strong> | |
| </div> | |
| <div style="display: flex; gap: 10px; justify-content: center; flex-wrap: wrap;"> | |
| <button type="button" class="cc-btn-primary" | |
| onclick="document.getElementById('challenge-clear-overlay').style.display='none'; if (window.editor && window.editor.reset_canvas) {{ window.editor.reset_canvas(); }} (document.querySelector('#next-target-btn button') || document.getElementById('next-target-btn'))?.click();" | |
| style="flex: 1; min-width: 140px; padding: 11px 16px; border-radius: 10px; border: none; background: var(--cc-accent); color: #FFFFFF; font-size: 14px; font-weight: 700; cursor: pointer; transition: all 0.15s ease; box-shadow: 0 4px 12px rgba(16, 185, 129, 0.25);"> | |
| 🎯 Next Target | |
| </button> | |
| <button type="button" class="cc-btn-secondary" | |
| onclick="document.getElementById('challenge-clear-overlay').style.display='none'; if (window.editor && window.editor.reset_canvas) {{ window.editor.reset_canvas(); }} (document.querySelector('#play-again-btn button') || document.getElementById('play-again-btn'))?.click();" | |
| style="flex: 1; min-width: 140px; padding: 11px 16px; border-radius: 10px; border: 1px solid var(--cc-btn-sec-border); background: var(--cc-btn-sec-bg); color: var(--cc-btn-sec-text); font-size: 14px; font-weight: 600; cursor: pointer; transition: all 0.15s ease;"> | |
| 🔄 Play Again | |
| </button> | |
| </div> | |
| <button type="button" class="cc-btn-ghost" | |
| onclick="document.getElementById('challenge-clear-overlay').style.display='none';" | |
| style="margin-top: 10px; width: 100%; padding: 8px 12px; border-radius: 8px; border: none; background: transparent; color: var(--cc-text-muted); font-size: 13px; font-weight: 500; cursor: pointer; transition: all 0.15s ease;"> | |
| Continue Viewing Drawing | |
| </button> | |
| </div> | |
| </div> | |
| """ | |
| _CANVAS_RESET_COUNTER = 0 | |
| def make_blank_canvas() -> Dict[str, Any]: | |
| """Create a fresh blank EditorValue dict with a unique layer URL so Gradio Sketchpad clears all strokes.""" | |
| global _CANVAS_RESET_COUNTER | |
| _CANVAS_RESET_COUNTER += 1 | |
| layer_img = Image.new("RGBA", (420, 420), (0, 0, 0, 0)) | |
| x = _CANVAS_RESET_COUNTER % 420 | |
| y = (_CANVAS_RESET_COUNTER // 420) % 420 | |
| layer_img.putpixel((x, y), (255, 255, 255, 1)) | |
| return { | |
| "background": None, | |
| "layers": [layer_img], | |
| "composite": None, | |
| } | |
| def evaluate_sketch(sketch_data: Any, theme_name: str, target_object: str) -> Tuple[Dict[str, float], str, str]: | |
| img = extract_pil_image(sketch_data) | |
| result = ENGINE.evaluate_drawing(img, theme_name, target_object) | |
| return format_scoreboard(result), format_status_message(result), format_popup_dialog(result) | |
| def on_theme_change(theme_name: str): | |
| objects = get_theme_objects(theme_name) | |
| default_target = objects[0] | |
| result = ENGINE.evaluate_drawing(None, theme_name, default_target) | |
| return ( | |
| gr.update(choices=objects, value=default_target), | |
| make_blank_canvas(), | |
| format_scoreboard(result), | |
| format_status_message(result), | |
| format_popup_dialog(result), | |
| ) | |
| def on_target_change(theme_name: str, target_object: str): | |
| result = ENGINE.evaluate_drawing(None, theme_name, target_object) | |
| return ( | |
| make_blank_canvas(), | |
| format_scoreboard(result), | |
| format_status_message(result), | |
| format_popup_dialog(result), | |
| ) | |
| def on_next_target(theme_name: str, current_target: str): | |
| objects = get_theme_objects(theme_name) | |
| idx = objects.index(current_target) if current_target in objects else 0 | |
| next_target = objects[(idx + 1) % len(objects)] | |
| result = ENGINE.evaluate_drawing(None, theme_name, next_target) | |
| return ( | |
| gr.update(value=next_target), | |
| make_blank_canvas(), | |
| format_scoreboard(result), | |
| format_status_message(result), | |
| format_popup_dialog(result), | |
| ) | |
| def create_app() -> gr.Blocks: | |
| """Construct minimal Gradio Blocks application using built-in Soft theme.""" | |
| # Build and cache reference embedding vectors while initializing the Gradio app | |
| ENGINE.ensure_reference_embeddings() | |
| theme = gr.themes.Soft() | |
| blocks_kwargs = {"title": "Embedding Draw Challenge — EmbeddingGemma 2"} | |
| if "theme" in inspect.signature(gr.Blocks.__init__).parameters: | |
| blocks_kwargs["theme"] = theme | |
| with gr.Blocks(**blocks_kwargs) as app: | |
| app.theme = theme | |
| gr.Markdown( | |
| "# 🎨 Embedding Draw Challenge\n" | |
| "Draw the target object below. **[EmbeddingGemma 2](https://ai.google.dev/gemma/docs/embeddinggemma)** guesses what you're drawing. Reach **90%** similarity on the target to clear the challenge!\n\n" | |
| "### 🕹️ How to Play\n" | |
| "1. **Choose One Target**: Select a **Theme** and pick **one single object** from the **Target Object** dropdown to draw.\n" | |
| "2. **Sketch Your Object**: Use the drawing canvas and color palette to sketch your chosen target.\n" | |
| "3. **Outscore Competing Distractors**: The other items listed in the dropdown act as competing distractors on the **Similarity Scoreboard**.\n" | |
| "4. **Reach 90% Similarity**: Refine your drawing until your chosen **Target Object** hits **90% or higher** to clear the challenge, then select another object from the dropdown to try a new round!\n" | |
| ) | |
| with gr.Row(): | |
| theme_dropdown = gr.Dropdown( | |
| choices=list(THEME_DATABASE.keys()), | |
| value="Fruit", | |
| label="Theme", | |
| info="Select a drawing category", | |
| scale=2, | |
| ) | |
| target_dropdown = gr.Dropdown( | |
| choices=get_theme_objects("Fruit"), | |
| value="Banana", | |
| label="Target Object", | |
| info="Draw ONLY this single selected object (other choices act as competing distractors)", | |
| scale=2, | |
| ) | |
| initial_result = ENGINE.evaluate_drawing(None, "Fruit", "Banana") | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| sketchpad = gr.Sketchpad( | |
| label="Drawing Canvas", | |
| type="pil", | |
| image_mode="RGBA", | |
| canvas_size=(420, 420), | |
| layers=False, | |
| brush=gr.Brush( | |
| colors=[ | |
| "#1A1A1B", | |
| "#FFE135", | |
| "#EF4444", | |
| "#2563EB", | |
| "#10B981", | |
| "#9333EA", | |
| "#94A3B8", | |
| "#F97316", | |
| "#78350F", | |
| "#FFFFFF", | |
| ], | |
| default_size=8, | |
| ), | |
| ) | |
| with gr.Column(scale=1): | |
| scoreboard = gr.Label( | |
| value=format_scoreboard(initial_result), | |
| num_top_classes=8, | |
| label="Similarity Scoreboard (Goal: ≥ 90%)", | |
| ) | |
| status_box = gr.Markdown( | |
| value=format_status_message(initial_result), | |
| ) | |
| with gr.Row(): | |
| play_again_btn = gr.Button( | |
| "🔄 Play Again", | |
| variant="secondary", | |
| elem_id="play-again-btn", | |
| ) | |
| next_target_btn = gr.Button( | |
| "🎯 Next Target", | |
| variant="primary", | |
| elem_id="next-target-btn", | |
| ) | |
| popup_dialog = gr.HTML( | |
| value=format_popup_dialog(initial_result), | |
| elem_id="challenge-clear-popup", | |
| apply_default_css=False, | |
| ) | |
| clear_canvas_js_1 = "(a) => { if (window.editor && window.editor.reset_canvas) { window.editor.reset_canvas(); } return a; }" | |
| clear_canvas_js_2 = "(a, b) => { if (window.editor && window.editor.reset_canvas) { window.editor.reset_canvas(); } return [a, b]; }" | |
| # Wire events | |
| sketchpad.change( | |
| fn=evaluate_sketch, | |
| inputs=[sketchpad, theme_dropdown, target_dropdown], | |
| outputs=[scoreboard, status_box, popup_dialog], | |
| ) | |
| theme_dropdown.change( | |
| fn=on_theme_change, | |
| inputs=[theme_dropdown], | |
| outputs=[target_dropdown, sketchpad, scoreboard, status_box, popup_dialog], | |
| js=clear_canvas_js_1, | |
| ) | |
| target_dropdown.change( | |
| fn=on_target_change, | |
| inputs=[theme_dropdown, target_dropdown], | |
| outputs=[sketchpad, scoreboard, status_box, popup_dialog], | |
| js=clear_canvas_js_2, | |
| ) | |
| play_again_btn.click( | |
| fn=on_target_change, | |
| inputs=[theme_dropdown, target_dropdown], | |
| outputs=[sketchpad, scoreboard, status_box, popup_dialog], | |
| js=clear_canvas_js_2, | |
| ) | |
| next_target_btn.click( | |
| fn=on_next_target, | |
| inputs=[theme_dropdown, target_dropdown], | |
| outputs=[target_dropdown, sketchpad, scoreboard, status_box, popup_dialog], | |
| js=clear_canvas_js_2, | |
| ) | |
| return app | |
| if __name__ == "__main__": | |
| app = create_app() | |
| launch_kwargs = { | |
| "server_name": "0.0.0.0", | |
| "server_port": 7860, | |
| } | |
| if "theme" in inspect.signature(app.launch).parameters: | |
| launch_kwargs["theme"] = gr.themes.Soft() | |
| app.launch(**launch_kwargs) | |