import spaces # must be imported before torch import base64 import html import io import json import sys import time import urllib.parse from pathlib import Path import gradio as gr import httpx import torch from huggingface_hub import snapshot_download from PIL import Image MODEL_ID = "Cloudflare/clef" MODEL_PATH = Path(snapshot_download(MODEL_ID)) sys.path.insert(0, str(MODEL_PATH)) from joint_schema_model import ( # noqa: E402 ClefModel, JointSchemaHead, QUESTION_TYPES, collate_records, encode_record, systemone_answer, ) from safetensors.torch import load_file # noqa: E402 from transformers import AutoProcessor, Qwen3_5ForConditionalGeneration # noqa: E402 # Same as joint_schema_model.load_release_model, but loaded to CPU first and moved with # .to("cuda") so the ZeroGPU hijack can pack the weights (device_map would bypass it). backbone = Qwen3_5ForConditionalGeneration.from_pretrained(MODEL_PATH, dtype=torch.bfloat16) backbone.config.use_cache = False head = JointSchemaHead(**json.loads((MODEL_PATH / "joint_head_config.json").read_text())) head.load_state_dict(load_file(MODEL_PATH / "joint_head.safetensors"), strict=True) head = head.to(dtype=torch.bfloat16) model = ClefModel(backbone, head).eval().to("cuda") processor = AutoProcessor.from_pretrained(MODEL_PATH) MAX_IMAGE_SIDE = 1024 def _parse_state(state: str): text = (state or "").strip() if not text: raise gr.Error("Please provide a state (plain text or JSON).") try: return json.loads(text) except json.JSONDecodeError: return text def _parse_questions(questions: str) -> dict: try: parsed = json.loads(questions or "") except json.JSONDecodeError as error: raise gr.Error(f"Questions must be valid JSON: {error}") if not isinstance(parsed, dict) or not parsed: raise gr.Error("Questions must be a non-empty JSON object keyed by question ID.") for question_id, question in parsed.items(): if not isinstance(question, dict) or question.get("type") not in QUESTION_TYPES: raise gr.Error(f"{question_id}: type must be one of noul, choice, score.") if question["type"] == "choice" and not isinstance(question.get("criteria"), dict): raise gr.Error(f"{question_id}: choice criteria must be an object of option -> description.") if question["type"] == "score" and not isinstance(question.get("criteria"), list): raise gr.Error(f"{question_id}: score criteria must be a list of level descriptions.") if question["type"] != "noul" and not question.get("criteria"): raise gr.Error(f"{question_id}: criteria must not be empty.") return parsed def _load_image(image) -> Image.Image | None: """The workflow canvas hands media ports over as {"path": ...} / {"url": ...} dicts: uploads land on this server (path), remote URLs are passed through.""" if image is None or isinstance(image, Image.Image): return image src = None if isinstance(image, dict): src = image.get("path") or image.get("url") elif isinstance(image, str): src = image if not src or not isinstance(src, str): return None if src.startswith("data:"): _, _, b64 = src.partition(",") return Image.open(io.BytesIO(base64.b64decode(b64))) if src.startswith(("http://", "https://")): resp = httpx.get(src, timeout=30, follow_redirects=True) resp.raise_for_status() return Image.open(io.BytesIO(resp.content)) if src.startswith("/gradio_api/file="): src = urllib.parse.unquote(src.removeprefix("/gradio_api/file=")) return Image.open(src) def _prep_image(image: Image.Image) -> Image.Image: image = image.convert("RGB") scale = MAX_IMAGE_SIDE / max(image.size) if scale < 1: image = image.resize((round(image.width * scale), round(image.height * scale)), Image.LANCZOS) return image def _render(questions: dict, probs: dict, latency_ms: float, n_tokens: int) -> str: """Self-contained HTML report — the canvas renders it in a sandboxed iframe, so no Gradio CSS variables: everything is styled inline.""" cards = [] for question_id, question in questions.items(): p = probs[question_id] if question["type"] == "noul": options = [("true", "true", p["true"]), ("false", "false", p["false"])] elif question["type"] == "choice": options = [(k, f"{k} — {v}" if v else k, p[str(k)]) for k, v in question["criteria"].items()] else: options = [(str(i), f"{i} — {v}", p[str(i)]) for i, v in enumerate(question["criteria"])] best = max(options, key=lambda o: o[2])[0] rows = [] for key, label, value in options: is_best = key == best fill = "#ec4899" if is_best else "#9ca3af" weight = "600" if is_best else "400" rows.append( '