try: # Must be imported before anything that touches CUDA/torch (e.g. # sentence-transformers below) -- required on HF's ZeroGPU hardware, # where importing torch first crashes the reload watcher with # "CUDA has been initialized before importing the `spaces` package". # Not installed/needed outside a ZeroGPU Space (e.g. running locally). import spaces _gpu = spaces.GPU except ImportError: def _gpu(fn): return fn import gradio as gr from app.rag import RAGRetriever from app.prompt import build_medgemma_prompt from app.inference_client import build_payload, encode_image_b64, invoke_and_wait, extract_answer, InferenceTimeout retriever = RAGRetriever(index_dir="rag_index") @_gpu def _zerogpu_startup_check(): """This app never runs model inference locally -- the real model runs on a remote SageMaker endpoint, called over HTTP. But ZeroGPU hardware refuses to start a Space with zero @spaces.GPU-decorated functions ("No @spaces.GPU function detected during startup"), so this dummy function exists purely to satisfy that check. It is never called. """ return True def _sanitize_history(history: list[dict]) -> list[dict]: """Gradio history entries can carry extra fields (metadata, options) and non-string content (file attachments become lists/dicts) that the server's strict message schema rejects. Reduce each turn to plain {role, content: str} before it goes into the prompt/payload. """ clean = [] for turn in history: role = turn.get("role") content = turn.get("content") if isinstance(content, str): text = content elif isinstance(content, (list, tuple)): parts = [] for item in content: if isinstance(item, str): parts.append(item) elif isinstance(item, dict) and item.get("type") == "text": parts.append(item.get("text", "")) else: parts.append("[attachment]") text = "\n".join(p for p in parts if p) else: text = "[attachment]" if content else "" if role in ("user", "assistant") and text: clean.append({"role": role, "content": text}) return clean def respond(message: dict, history: list[dict]): question = message.get("text", "") files = message.get("files", []) image_b64 = None if files: with open(files[0], "rb") as f: image_b64 = encode_image_b64(f.read()) # This RAG corpus (NIH MedQuAD, ~50 docs on rare genetic conditions) has no # radiology/imaging content, so it can't meaningfully inform an image # finding -- retrieving from it for image questions only forces in # mismatched context that burns reasoning tokens reconciling irrelevant # material instead of answering. Skip it when an image is attached. rag_chunks = [] if image_b64 else retriever.search(question, k=4, min_score=0.9) messages = build_medgemma_prompt( question=question, history=_sanitize_history(history), rag_chunks=rag_chunks, has_image=image_b64 is not None, ) # This model produces an internal reasoning span before the real answer, # and structured/bulleted answers run long -- 1024 tokens still wasn't # enough and was truncating mid-answer. Give it more headroom. payload = build_payload(messages, image_b64, max_new_tokens=2048) yield "Sending to model — if the endpoint scaled to zero, this can take 10-15 min to cold start..." try: result = invoke_and_wait(payload) except InferenceTimeout as e: yield f"Timed out waiting for a response: {e}" return except RuntimeError as e: yield str(e) return answer = extract_answer(result) if rag_chunks: sources = ", ".join(c["source"] for c in rag_chunks) yield f"{answer}\n\n---\n*Sources: {sources}*" else: yield f"{answer}\n\n---\n*No relevant reference material found for this question.*" demo = gr.ChatInterface( fn=respond, multimodal=True, title="MedGemma Medical VQA", description="Upload a medical image (optional) and ask a question. Answers are grounded in NIH MedQuAD reference material.", ) if __name__ == "__main__": demo.launch()