File size: 20,753 Bytes
a3cb9f2
a34d6dd
a3cb9f2
a34d6dd
 
 
 
489bf3e
 
 
 
 
 
 
 
 
a3cb9f2
 
a34d6dd
a3cb9f2
a34d6dd
a3cb9f2
 
 
 
 
 
 
 
 
a34d6dd
 
a3cb9f2
 
489bf3e
 
 
a3cb9f2
a34d6dd
a3cb9f2
 
 
 
 
 
 
 
489bf3e
a34d6dd
80b6a6c
a34d6dd
489bf3e
a34d6dd
489bf3e
a3cb9f2
80b6a6c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
489bf3e
 
 
a3cb9f2
489bf3e
568874b
489bf3e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6b079df
489bf3e
 
 
 
a3cb9f2
a34d6dd
 
6b079df
489bf3e
a3cb9f2
489bf3e
 
 
 
 
 
a3cb9f2
a34d6dd
 
 
 
a3cb9f2
a34d6dd
 
a3cb9f2
489bf3e
 
568874b
6b079df
a34d6dd
 
 
489bf3e
a34d6dd
 
6b079df
a34d6dd
 
 
6b079df
 
 
 
489bf3e
6b079df
 
a34d6dd
 
 
 
 
 
 
 
 
 
 
 
6b079df
 
 
 
489bf3e
6b079df
 
a34d6dd
 
 
 
 
 
 
 
 
 
 
 
489bf3e
a34d6dd
489bf3e
 
b55773f
489bf3e
a34d6dd
 
 
 
6b079df
b55773f
489bf3e
a34d6dd
 
568874b
a34d6dd
568874b
 
 
b55773f
 
 
 
 
 
6b079df
 
 
489bf3e
 
b55773f
489bf3e
b55773f
 
a34d6dd
489bf3e
a34d6dd
 
489bf3e
b55773f
568874b
a34d6dd
568874b
 
 
b55773f
6b079df
 
 
489bf3e
 
b55773f
489bf3e
b55773f
a34d6dd
a3cb9f2
489bf3e
 
 
 
 
 
 
 
 
 
6b079df
489bf3e
 
6b079df
a3cb9f2
 
568874b
a3cb9f2
568874b
a3cb9f2
489bf3e
 
 
 
568874b
a3cb9f2
489bf3e
 
6b079df
 
489bf3e
6b079df
 
489bf3e
 
a3cb9f2
489bf3e
 
a3cb9f2
489bf3e
a3cb9f2
489bf3e
6b079df
a3cb9f2
568874b
 
 
 
489bf3e
 
568874b
489bf3e
 
6b079df
568874b
489bf3e
6b079df
489bf3e
a3cb9f2
a34d6dd
 
 
 
489bf3e
a34d6dd
 
 
489bf3e
a34d6dd
a3cb9f2
 
6b079df
a34d6dd
 
 
a3cb9f2
 
 
 
a34d6dd
6b079df
a3cb9f2
 
489bf3e
 
 
a3cb9f2
 
489bf3e
 
a3cb9f2
 
 
 
 
489bf3e
 
 
a3cb9f2
 
489bf3e
 
 
a3cb9f2
 
6b079df
 
a3cb9f2
568874b
a3cb9f2
 
 
568874b
 
a3cb9f2
568874b
 
a3cb9f2
 
568874b
6b079df
 
489bf3e
6b079df
489bf3e
a3cb9f2
 
568874b
a3cb9f2
 
 
568874b
a3cb9f2
 
 
 
568874b
a3cb9f2
 
 
568874b
a3cb9f2
 
 
 
568874b
489bf3e
 
 
 
6b079df
 
 
489bf3e
 
6b079df
489bf3e
 
 
 
a3cb9f2
6b079df
489bf3e
 
 
 
 
6b079df
 
489bf3e
 
568874b
a3cb9f2
 
489bf3e
 
568874b
489bf3e
 
 
 
568874b
 
489bf3e
 
6b079df
 
489bf3e
a3cb9f2
568874b
a3cb9f2
489bf3e
 
 
 
a3cb9f2
 
568874b
 
a3cb9f2
489bf3e
 
568874b
6b079df
 
489bf3e
a3cb9f2
489bf3e
568874b
 
a34d6dd
a3cb9f2
 
a34d6dd
a3cb9f2
 
 
568874b
 
489bf3e
568874b
a3cb9f2
489bf3e
 
6b079df
489bf3e
 
a3cb9f2
489bf3e
 
 
 
 
a3cb9f2
489bf3e
 
6b079df
489bf3e
 
 
 
 
 
 
6b079df
489bf3e
a3cb9f2
568874b
489bf3e
 
 
568874b
 
489bf3e
 
 
a3cb9f2
489bf3e
 
a3cb9f2
568874b
a34d6dd
489bf3e
a3cb9f2
 
489bf3e
a3cb9f2
 
489bf3e
 
 
a3cb9f2
489bf3e
 
 
 
 
 
 
 
 
 
a3cb9f2
 
 
568874b
 
a3cb9f2
489bf3e
 
 
 
a3cb9f2
 
568874b
a3cb9f2
489bf3e
a3cb9f2
 
 
 
 
 
 
 
 
6b079df
a3cb9f2
 
 
 
 
489bf3e
a3cb9f2
 
 
568874b
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
"""
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,
)

# -----------------------------------------------------------------------------
# Config
# -----------------------------------------------------------------------------

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")

# ATTN_IMPL = "flash_attention_2" if os.getenv("USE_FA2", "0") == "1" else "sdpa"

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()}")
    
    # Transformers helper check
    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

# Determine best implementation
if check_flash_attention():
    ATTN_IMPL = "flash_attention_2"
else:
    ATTN_IMPL = "sdpa" # Fallback to Scaled Dot Product Attention

print(f"[INFO] Using {ATTN_IMPL} for model loading.")

# model = AutoModelForCausalLM.from_pretrained(
#     model_id,
#     torch_dtype=torch.float16, # FA2 requires fp16 or bf16
#     attn_implementation=best_attn,
#     trust_remote_code=True
# ).to("cuda")

# Call it inside your setup or first GPU call
# FLASH_AVAILABLE = check_flash_attention()

# -----------------------------------------------------------------------------
# Load dataset + embeddings (CPU only)
# -----------------------------------------------------------------------------

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")

# Pick the main markdown field robustly
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}")

# -----------------------------------------------------------------------------
# Lazy GPU globals (must only be initialized inside @spaces.GPU)
# -----------------------------------------------------------------------------

_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()
    
    # 1) Try Nemotron FP8
    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,
        )
        
        # FIX: We force 'eager' or 'flash_attention_2' because this model 
        # doesn't support the default 'sdpa' implementation yet.
        # Since you installed the FA2 wheels, we'll try to use that first.
        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, # Changed here
        ).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", # Fallback model usually supports SDPA
        ).eval()
        
        return _gen_model, _gen_tokenizer


# -----------------------------------------------------------------------------
# Helpers
# -----------------------------------------------------------------------------

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:
    # Keep your original β€œcard” parsing lightweight.
    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>
    """


# -----------------------------------------------------------------------------
# Main GPU function (ZeroGPU allocation happens here)
# -----------------------------------------------------------------------------

@spaces.GPU
def retrieve(query_text, query_image, rerank_option, generate_summary_option):
    global _embeddings_gpu

    # Load VL models only now
    _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)

    # Choose query
    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.")

    # Retrieval
    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)

    # Rerank (text only)
    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"

    # Generation (optional)
    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


# -----------------------------------------------------------------------------
# UI
# -----------------------------------------------------------------------------

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()