Spaces:
Running on Zero
Running on Zero
| 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") | |
| 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() | |