| """ |
| Multimodal RAG Demo with Nemotron Embed VL and Rerank VL (ZeroGPU-friendly) |
| |
| Key ZeroGPU rule: |
| - DO NOT load GPU models at import time. |
| - Lazy-load models INSIDE the @spaces.GPU function (or inside helpers called from it). |
| |
| Models: |
| - Embed: nvidia/llama-nemotron-embed-vl-1b-v2 |
| - Rerank: nvidia/llama-nemotron-rerank-vl-1b-v2 |
| - Gen (preferred): nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-FP8 (text-only summary, trust_remote_code) |
| - If it fails, fallback to a smaller text-only model. |
| |
| Attention: |
| - Default: SDPA (most stable on Spaces) |
| - Optional: FlashAttention-2 if USE_FA2=1 and flash-attn is installed & compatible. |
| """ |
|
|
| import os |
| import time |
| import spaces |
| import torch |
| import gradio as gr |
| from PIL import Image |
| from datasets import load_dataset |
| from safetensors.torch import load_file |
| from transformers import ( |
| AutoModel, |
| AutoModelForSequenceClassification, |
| AutoProcessor, |
| AutoTokenizer, |
| AutoModelForCausalLM, |
| ) |
|
|
| |
| |
| |
|
|
| DEVICE_CPU = torch.device("cpu") |
|
|
| EMBED_MODEL_PATH = "nvidia/llama-nemotron-embed-vl-1b-v2" |
| EMBED_COMMIT_HASH = "5b5ca69c35bf6ec1484d2d5ff238626e67a745e2" |
|
|
| RERANK_MODEL_PATH = "nvidia/llama-nemotron-rerank-vl-1b-v2" |
| RERANK_COMMIT_HASH = "47e5a355d1a050c3e5f69d53f14964b1d34bcd9d" |
|
|
| GENERATION_MODEL_ID = "nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-FP8" |
| FALLBACK_GEN_MODEL_ID = os.getenv("FALLBACK_GEN_MODEL_ID", "Qwen/Qwen2.5-7B-Instruct") |
|
|
| |
|
|
| modality_to_tokens = {"image": 2048, "image_text": 10240, "text": 8192} |
|
|
| PATH_TO_EMBEDDING_FILE = os.getenv("EMBEDDINGS_FILE", "image_text_embeddings_10k.safetensors") |
|
|
|
|
| def check_flash_attention(): |
| import torch |
| from transformers.utils import is_flash_attn_2_available |
|
|
| print(f"--- Flash Attention Check ---") |
| print(f"PyTorch version: {torch.__version__}") |
| print(f"CUDA available: {torch.cuda.is_available()}") |
| |
| |
| fa2_available = is_flash_attn_2_available() |
| print(f"Transformers reports FA2 available: {fa2_available}") |
|
|
| if torch.cuda.is_available(): |
| capability = torch.cuda.get_device_capability() |
| print(f"GPU Compute Capability: {capability}") |
| if capability[0] < 8: |
| print("Note: FA2 requires Compute Capability 8.0+ (Ampere or newer).") |
| |
| return fa2_available |
|
|
| |
| if check_flash_attention(): |
| ATTN_IMPL = "flash_attention_2" |
| else: |
| ATTN_IMPL = "sdpa" |
|
|
| print(f"[INFO] Using {ATTN_IMPL} for model loading.") |
|
|
| |
| |
| |
| |
| |
| |
|
|
| |
| |
|
|
| |
| |
| |
|
|
| print("[INFO] Loading dataset (CPU)...") |
| dataset = load_dataset("mrdbourke/recipe-synthetic-images-10k") |
| train_split = dataset["train"] |
| print(f"[INFO] Dataset loaded with {len(train_split)} samples") |
|
|
| |
| PREFERRED_TEXT_COL = "recipe_markdown" |
| FALLBACK_TEXT_COLS = ["markdown", "text", "recipe_text", "content"] |
|
|
| if PREFERRED_TEXT_COL in train_split.column_names: |
| TEXT_COL = PREFERRED_TEXT_COL |
| else: |
| found = None |
| for c in FALLBACK_TEXT_COLS: |
| if c in train_split.column_names: |
| found = c |
| break |
| if found is None: |
| raise RuntimeError( |
| f"Could not find a recipe text column. Available columns: {train_split.column_names}" |
| ) |
| TEXT_COL = found |
| print(f"[WARN] '{PREFERRED_TEXT_COL}' not found. Using '{TEXT_COL}' instead.") |
|
|
| if "image" not in train_split.column_names: |
| raise RuntimeError(f"Dataset does not contain 'image' column. Columns: {train_split.column_names}") |
|
|
| print(f"[INFO] Using TEXT_COL='{TEXT_COL}'") |
|
|
| print("[INFO] Loading embeddings (CPU)...") |
| emb = load_file(PATH_TO_EMBEDDING_FILE) |
| if "image_text_embeddings" not in emb: |
| raise RuntimeError(f"'{PATH_TO_EMBEDDING_FILE}' missing key 'image_text_embeddings'. Keys: {list(emb.keys())}") |
|
|
| image_text_embeddings = emb["image_text_embeddings"].to(DEVICE_CPU) |
| print(f"[INFO] Embeddings loaded: {tuple(image_text_embeddings.shape)} | device={image_text_embeddings.device}") |
|
|
| |
| |
| |
|
|
| _embed_model = None |
| _embed_processor = None |
| _rerank_model = None |
| _rerank_processor = None |
|
|
| _gen_model = None |
| _gen_tokenizer = None |
|
|
| _embeddings_gpu = None |
|
|
|
|
| def _cuda() -> torch.device: |
| return torch.device("cuda") |
|
|
|
|
| def _load_embed_and_rerank_on_gpu(): |
| global _embed_model, _embed_processor, _rerank_model, _rerank_processor |
|
|
| device = _cuda() |
|
|
| if _embed_model is None or _embed_processor is None: |
| print("[INFO] Lazy-loading EMBED model on GPU...") |
| _embed_model = AutoModel.from_pretrained( |
| EMBED_MODEL_PATH, |
| revision=EMBED_COMMIT_HASH, |
| trust_remote_code=True, |
| torch_dtype=torch.bfloat16, |
| attn_implementation=ATTN_IMPL, |
| ).to(device).eval() |
|
|
| _embed_processor = AutoProcessor.from_pretrained( |
| EMBED_MODEL_PATH, |
| revision=EMBED_COMMIT_HASH, |
| trust_remote_code=True, |
| max_input_tiles=6, |
| use_thumbnail=True, |
| p_max_length=modality_to_tokens["image_text"], |
| ) |
|
|
| if _rerank_model is None or _rerank_processor is None: |
| print("[INFO] Lazy-loading RERANK model on GPU...") |
| _rerank_model = AutoModelForSequenceClassification.from_pretrained( |
| RERANK_MODEL_PATH, |
| revision=RERANK_COMMIT_HASH, |
| trust_remote_code=True, |
| torch_dtype=torch.bfloat16, |
| attn_implementation=ATTN_IMPL, |
| ).to(device).eval() |
|
|
| _rerank_processor = AutoProcessor.from_pretrained( |
| RERANK_MODEL_PATH, |
| revision=RERANK_COMMIT_HASH, |
| trust_remote_code=True, |
| max_input_tiles=6, |
| use_thumbnail=True, |
| rerank_max_length=modality_to_tokens["image_text"], |
| ) |
|
|
| return _embed_model, _embed_processor, _rerank_model, _rerank_processor |
|
|
|
|
| def _load_generation_model_on_gpu(): |
| """ |
| Try Nemotron 30B FP8 first. |
| If it fails, fall back to a smaller text model. |
| """ |
| global _gen_model, _gen_tokenizer |
| if _gen_model is not None and _gen_tokenizer is not None: |
| return _gen_model, _gen_tokenizer |
|
|
| device = _cuda() |
| |
| |
| try: |
| print("[INFO] Lazy-loading GENERATION model (Nemotron 30B FP8) on GPU...") |
| _gen_tokenizer = AutoTokenizer.from_pretrained( |
| GENERATION_MODEL_ID, |
| trust_remote_code=True, |
| use_fast=True, |
| ) |
| |
| |
| |
| |
| gen_attn_impl = ATTN_IMPL if ATTN_IMPL == "flash_attention_2" else "eager" |
|
|
| _gen_model = AutoModelForCausalLM.from_pretrained( |
| GENERATION_MODEL_ID, |
| trust_remote_code=True, |
| torch_dtype="auto", |
| device_map="auto", |
| attn_implementation=gen_attn_impl, |
| ).eval() |
| |
| print(f"[INFO] Nemotron generation model loaded OK with {gen_attn_impl}") |
| return _gen_model, _gen_tokenizer |
|
|
| except Exception as e: |
| print(f"[WARN] Nemotron FP8 load failed: {repr(e)}") |
| print(f"[WARN] Falling back to: {FALLBACK_GEN_MODEL_ID}") |
| |
| _gen_tokenizer = AutoTokenizer.from_pretrained( |
| FALLBACK_GEN_MODEL_ID, |
| trust_remote_code=True, |
| use_fast=True, |
| ) |
| |
| _gen_model = AutoModelForCausalLM.from_pretrained( |
| FALLBACK_GEN_MODEL_ID, |
| trust_remote_code=True, |
| torch_dtype=torch.bfloat16, |
| device_map="auto", |
| attn_implementation="sdpa", |
| ).eval() |
| |
| return _gen_model, _gen_tokenizer |
|
|
|
|
| |
| |
| |
|
|
| def _l2_normalize(x: torch.Tensor, eps: float = 1e-12) -> torch.Tensor: |
| return x / (x.norm(p=2, dim=-1, keepdim=True) + eps) |
|
|
|
|
| def match_query_to_embeddings( |
| query: str | Image.Image, |
| target_embeddings_to_match: torch.Tensor, |
| top_k: int = 50 |
| ) -> tuple[torch.Tensor, torch.Tensor]: |
| with torch.inference_mode(): |
| if isinstance(query, Image.Image): |
| q = _embed_model.encode_documents(images=[query]) |
| else: |
| q = _embed_model.encode_queries([query]) |
|
|
| sim = _l2_normalize(q) @ _l2_normalize(target_embeddings_to_match).T |
| sim = sim.flatten() |
| idx = torch.argsort(sim, descending=True)[:top_k] |
| scores = sim[idx] |
| return scores, idx |
|
|
|
|
| def rerank_samples( |
| query_text: str, |
| sorted_indices: torch.Tensor, |
| num_samples_to_rerank: int = 20, |
| ) -> tuple: |
| device = _cuda() |
| top_idx = sorted_indices[:num_samples_to_rerank] |
| subset = dataset["train"].select(top_idx.tolist()) |
|
|
| texts = subset[TEXT_COL] |
| images = subset["image"] |
|
|
| pairs = [{"question": query_text, "doc_text": t, "doc_image": im} for t, im in zip(texts, images)] |
|
|
| batch = _rerank_processor.process_queries_documents_crossencoder(pairs) |
| batch = {k: (v.to(device) if isinstance(v, torch.Tensor) else v) for k, v in batch.items()} |
|
|
| with torch.inference_mode(): |
| out = _rerank_model(**batch, return_dict=True) |
|
|
| logits = out.logits.squeeze(-1) |
| rerank_sorted = torch.argsort(logits, descending=True) |
| return subset, rerank_sorted |
|
|
|
|
| def generate_recipe_summary(recipe_texts: list[str], max_new_tokens: int = 384) -> str: |
| model, tok = _load_generation_model_on_gpu() |
|
|
| combined = "" |
| for i, r in enumerate(recipe_texts[:3], 1): |
| combined += f"\n\n--- RECIPE {i} ---\n{r}" |
|
|
| prompt = ( |
| "You are a helpful culinary assistant.\n" |
| "Summarize the following recipes in Markdown.\n\n" |
| "Return:\n" |
| "- 1β2 sentence overview of each\n" |
| "- key ingredients\n" |
| "- difficulty (Easy/Medium/Hard)\n" |
| "- which is best for a quick weeknight dinner\n\n" |
| f"{combined}\n\n" |
| "## Summary:\n" |
| ) |
|
|
| inputs = tok(prompt, return_tensors="pt").to(model.device) |
|
|
| with torch.inference_mode(): |
| out = model.generate( |
| **inputs, |
| max_new_tokens=max_new_tokens, |
| do_sample=True, |
| temperature=0.7, |
| top_p=0.9, |
| pad_token_id=tok.eos_token_id, |
| ) |
|
|
| gen = tok.decode(out[0][inputs["input_ids"].shape[-1]:], skip_special_tokens=True) |
| return gen.strip() |
|
|
|
|
| def _markdown_to_simple_html(markdown_text: str, max_reviews: int = 1) -> str: |
| |
| lines = (markdown_text or "").strip().split("\n") |
|
|
| title = "" |
| description = "" |
| cook_time = "" |
| num_ratings = "" |
| ingredients = [] |
| steps = [] |
| reviews = [] |
|
|
| current_section = None |
| in_ingredients = False |
| in_steps = False |
| in_reviews = False |
| review_count = 0 |
|
|
| for raw in lines: |
| line = raw.strip() |
|
|
| if line.startswith("# ") and not title: |
| title = line[2:].strip() |
| continue |
|
|
| if line.startswith("**Time:**"): |
| cook_time = line.replace("**Time:**", "").strip() |
| continue |
| if line.startswith("**Number of Ratings:**"): |
| num_ratings = line.replace("**Number of Ratings:**", "").strip() |
| continue |
|
|
| if line.startswith("## "): |
| section = line[3:].strip().lower() |
| current_section = section |
| in_ingredients = (section == "ingredients") |
| in_steps = section.startswith("steps") |
| in_reviews = (section == "reviews") |
| continue |
|
|
| if current_section == "description" and line and not line.startswith("#"): |
| description = line |
| continue |
|
|
| if in_ingredients and line.startswith("- "): |
| ingredients.append(line[2:].strip()) |
| continue |
|
|
| if in_steps and line and line[0].isdigit(): |
| step_text = line.split(". ", 1)[-1] if ". " in line else line |
| steps.append(step_text.strip()) |
| continue |
|
|
| if in_reviews and line.startswith("> ") and review_count < max_reviews: |
| reviews.append(line[2:].strip()) |
| review_count += 1 |
| continue |
|
|
| html = f""" |
| <div style="border: 1px solid #ddd; border-radius: 8px; padding: 16px; margin: 4px; background: #fff; |
| font-family: system-ui, -apple-system, sans-serif; font-size: 12px; height: 380px; overflow-y: auto;"> |
| <div style="font-weight: 700; font-size: 14px; margin-bottom: 8px;">{title or "Recipe"}</div> |
| <div style="display:flex; gap:12px; font-size:11px; color:#666; margin-bottom:10px; flex-wrap:wrap;"> |
| {f"<span>β±οΈ {cook_time}</span>" if cook_time else ""} |
| {f"<span>β {num_ratings} ratings</span>" if num_ratings else ""} |
| </div> |
| <div style="color:#555; margin-bottom:12px; font-style:italic; line-height:1.4;"> |
| {(description[:160] + "β¦") if len(description) > 160 else description} |
| </div> |
| <div style="margin-bottom: 10px;"> |
| <div style="font-weight:600; font-size:11px; margin-bottom:4px;">π Ingredients</div> |
| <div style="color:#444; line-height:1.5;"> |
| {", ".join(ingredients[:8])}{("β¦" if len(ingredients) > 8 else "")} |
| </div> |
| </div> |
| <div style="margin-bottom: 10px;"> |
| <div style="font-weight:600; font-size:11px; margin-bottom:4px;">π¨βπ³ Steps ({len(steps)})</div> |
| <ol style="margin:0; padding-left:18px; color:#444; line-height:1.5;"> |
| {"".join(f"<li>{(s[:90] + 'β¦') if len(s) > 90 else s}</li>" for s in steps[:4])} |
| {f"<li style='color:#999;'>β¦and {len(steps)-4} more</li>" if len(steps) > 4 else ""} |
| </ol> |
| </div> |
| {f"<div style='border-top:1px solid #eee; padding-top:10px; margin-top:10px;'><div style='font-weight:600; font-size:11px; margin-bottom:4px;'>π¬ Review</div><div style='color:#555; background:#f9f9f9; padding:8px; border-radius:6px; font-style:italic;'>{(reviews[0][:220] + 'β¦') if len(reviews[0]) > 220 else reviews[0]}</div></div>" if reviews else ""} |
| </div> |
| """ |
| return html |
|
|
|
|
| def create_recipe_cards_html(items: list[dict], num_results: int = 3) -> str: |
| cards = [] |
| for it in items[:num_results]: |
| sample = it["sample"] |
| md = sample.get(TEXT_COL, "") or "" |
| cards.append(f"<div style='flex:1; min-width:0;'>{_markdown_to_simple_html(md)}</div>") |
|
|
| return f""" |
| <div style="margin-top: 16px;"> |
| <h3 style="font-family: system-ui, -apple-system, sans-serif; font-size: 16px; font-weight: 600; margin-bottom: 12px;"> |
| Retrieved Texts |
| </h3> |
| <div style="display:flex; gap:12px; width:100%;">{''.join(cards)}</div> |
| </div> |
| """ |
|
|
|
|
| |
| |
| |
|
|
| @spaces.GPU |
| def retrieve(query_text, query_image, rerank_option, generate_summary_option): |
| global _embeddings_gpu |
|
|
| |
| _load_embed_and_rerank_on_gpu() |
|
|
| if _embeddings_gpu is None: |
| print("[INFO] Moving embeddings to GPU (cached)...") |
| _embeddings_gpu = image_text_embeddings.to(_cuda(), non_blocking=True) |
|
|
| |
| if query_text and str(query_text).strip(): |
| input_query = str(query_text).strip() |
| query_is_text = True |
| elif query_image is not None: |
| input_query = query_image |
| query_is_text = False |
| else: |
| raise gr.Error("Please provide either a text query or an image query.") |
|
|
| |
| t0 = time.time() |
| scores, idx = match_query_to_embeddings(input_query, _embeddings_gpu, top_k=20) |
| t1 = time.time() |
|
|
| top = dataset["train"].select(idx.tolist()) |
| scored = [{"score": float(s.item()), "sample": smp} for s, smp in zip(scores, top)] |
|
|
| gallery = [(it["sample"]["image"], f"Score: {it['score']:.4f}") for it in scored[:3]] |
| cards_html = create_recipe_cards_html(scored, num_results=3) |
|
|
| |
| if rerank_option == "True" and query_is_text: |
| r0 = time.time() |
| subset, rerank_sorted = rerank_samples(input_query, idx, num_samples_to_rerank=20) |
| r1 = time.time() |
|
|
| reranked = subset.select(rerank_sorted.tolist()) |
| scored = [{"score": None, "sample": smp} for smp in reranked] |
|
|
| gallery = [(it["sample"]["image"], f"Reranked: {i}") for i, it in enumerate(scored[:3])] |
| cards_html = create_recipe_cards_html(scored, num_results=3) |
| rerank_time = round(r1 - r0, 4) |
| elif rerank_option == "True" and not query_is_text: |
| rerank_time = "Reranking only supported for text queries" |
| else: |
| rerank_time = "Reranking turned off" |
|
|
| |
| if generate_summary_option == "True": |
| g0 = time.time() |
| recipe_texts = [it["sample"].get(TEXT_COL, "") for it in scored[:3]] |
| summary = generate_recipe_summary(recipe_texts) |
| summary = summary.replace("```markdown", "").replace("```", "").strip() |
| g1 = time.time() |
| gen_time = round(g1 - g0, 4) |
| else: |
| summary = "Generation turned off, no summary created" |
| gen_time = "Generation turned off" |
|
|
| timing = { |
| "retrieve_time": round(t1 - t0, 4), |
| "rerank_time": rerank_time, |
| "generation_time": gen_time, |
| "attn_impl": ATTN_IMPL, |
| "text_col": TEXT_COL, |
| } |
|
|
| return gallery, cards_html, summary, timing |
|
|
|
|
| |
| |
| |
|
|
| with gr.Blocks(title="Multimodal RAG Demo") as demo: |
| gr.Markdown(f"""# ποΈπ Multimodal RAG Demo (ZeroGPU-friendly) |
| |
| - Dataset: `mrdbourke/recipe-synthetic-images-10k` |
| - Text field used: `{TEXT_COL}` |
| - Embed: `nvidia/llama-nemotron-embed-vl-1b-v2` |
| - Rerank: `nvidia/llama-nemotron-rerank-vl-1b-v2` |
| - Gen (preferred): `nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-FP8` (text-only) |
| - Attention backend: `{ATTN_IMPL}` (set `USE_FA2=1` to try FA2) |
| """) |
|
|
| with gr.Row(): |
| with gr.Column(scale=1): |
| query_text = gr.Textbox(label="Text Query", placeholder="e.g. 'dinner recipes with tomatoes'", lines=2) |
| query_image = gr.Image(label="Image Query (optional)", type="pil", height=200) |
|
|
| generate_summary_option = gr.Radio(["True", "False"], value="False", label="Generate recipe summary") |
| rerank_option = gr.Radio(["True", "False"], value="False", label="Rerank initial results? (text only)") |
|
|
| search_btn = gr.Button("Search", variant="primary") |
|
|
| with gr.Column(scale=2): |
| gallery_output = gr.Gallery(label="Retrieved Recipe Images", columns=3, height="auto", object_fit="cover") |
| recipes_html = gr.HTML(label="Retrieved Recipe Texts") |
| summary_generation = gr.Markdown(label="Generated Summary") |
| timing_output = gr.JSON(label="Timings") |
|
|
| gr.Examples( |
| examples=[ |
| ["best omelette recipes", None, "False", "False"], |
| ["best omelette recipes", None, "False", "True"], |
| ["eggplant dip", None, "True", "True"], |
| ], |
| inputs=[query_text, query_image, rerank_option, generate_summary_option], |
| label="Example Queries", |
| ) |
|
|
| search_btn.click( |
| fn=retrieve, |
| inputs=[query_text, query_image, rerank_option, generate_summary_option], |
| outputs=[gallery_output, recipes_html, summary_generation, timing_output], |
| ) |
|
|
| if __name__ == "__main__": |
| demo.launch() |