Spaces:
Running on Zero
Running on Zero
Download app.py from akhaliq/clef-decision-workflow: direct link, hf CLI and curl.
- Browser
- Download file 9.92 kB
-
https://huggingface.co/spaces/akhaliq/clef-decision-workflow/resolve/main/app.py
- Command line
-
hf download hf://spaces/akhaliq/clef-decision-workflow/app.py
-
curl -L -o app.py https://huggingface.co/spaces/akhaliq/clef-decision-workflow/resolve/main/app.py
9.92 kB
| 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( | |
| '<div style="display:grid;grid-template-columns:minmax(120px,40%) 1fr 56px;' | |
| 'gap:8px;align-items:center;margin:3px 0;font-size:13px;">' | |
| f'<div style="overflow:hidden;text-overflow:ellipsis;white-space:nowrap;' | |
| f'font-weight:{weight};">{html.escape(str(label))}</div>' | |
| '<div style="height:12px;border-radius:6px;background:#f3f4f6;' | |
| 'border:1px solid #e5e7eb;overflow:hidden;">' | |
| f'<div style="height:100%;width:{value * 100:.1f}%;background:{fill};"></div></div>' | |
| f'<div style="text-align:right;font-family:ui-monospace,monospace;">' | |
| f'{value * 100:.1f}%</div></div>' | |
| ) | |
| extra = "" | |
| if question["type"] == "score": | |
| expected = sum(i * p[str(i)] for i in range(len(question["criteria"]))) | |
| extra = f' · expected score <b>{expected:.2f}</b>' | |
| instr = question.get("instructions") or question_id | |
| cards.append( | |
| '<div style="border:1px solid #e5e7eb;border-radius:10px;padding:12px 14px;' | |
| 'background:#ffffff;">' | |
| '<div style="display:flex;gap:8px;align-items:center;margin-bottom:2px;">' | |
| f'<span style="font-weight:700;font-family:ui-monospace,monospace;">' | |
| f'{html.escape(question_id)}</span>' | |
| f'<span style="font-size:11px;padding:1px 8px;border-radius:999px;' | |
| f'background:#fdf2f8;color:#be185d;">{question["type"]}</span></div>' | |
| f'<div style="opacity:0.75;font-size:13px;">{html.escape(str(instr))}</div>' | |
| f'<div style="margin:6px 0 8px;">→ <b>{html.escape(str(best))}</b>{extra}</div>' | |
| + "".join(rows) | |
| + "</div>" | |
| ) | |
| footer = ( | |
| f'<div style="font-size:12px;opacity:0.65;">{n_tokens} input tokens · ' | |
| f'forward pass {latency_ms:.0f} ms</div>' | |
| ) | |
| return ( | |
| '<!DOCTYPE html><html><head><meta charset="utf-8"></head>' | |
| '<body style="margin:0;padding:16px;font-family:system-ui,-apple-system,sans-serif;' | |
| 'background:#f9fafb;color:#111827;display:flex;flex-direction:column;gap:12px;">' | |
| + "".join(cards) | |
| + footer | |
| + "</body></html>" | |
| ) | |
| def _gpu_duration(record: dict) -> int: | |
| # Measured on xlarge: ~0.4 s for a short schema, ~2-3 s with an image, ~12 s at ~13k tokens. | |
| # A 15 s floor is needed so a cold worker (streaming 55 GB of weights) is not aborted. | |
| chars = len(json.dumps(record["state"], ensure_ascii=False)) + len(json.dumps(record["questions"])) | |
| seconds = 15 + chars / 2000 + (4 if record.get("images") else 0) | |
| return int(min(40, seconds)) | |
| def _forward(record: dict): | |
| encoded = encode_record(processor.tokenizer, record, processor=processor) | |
| batch = collate_records([encoded], processor.tokenizer.pad_token_id, torch.device("cuda")) | |
| torch.cuda.synchronize() | |
| start = time.perf_counter() | |
| with torch.inference_mode(): | |
| logits = model(batch)[0] | |
| torch.cuda.synchronize() | |
| latency_ms = (time.perf_counter() - start) * 1000 | |
| probs = { | |
| q.question_id: dict(zip(q.option_ids, ql.float().softmax(-1).cpu().tolist())) | |
| for q, ql in zip(encoded.questions, logits) | |
| } | |
| return probs, latency_ms, len(encoded.input_ids) | |
| def decide(state: str, questions: str, image=None): | |
| """Answer a typed question schema about a state with Cloudflare Clef. | |
| Args: | |
| state: The situation to decide on, as plain text or a JSON value. | |
| questions: JSON object mapping question IDs to questions. Each question has a `type` | |
| (`noul` = true/false, `choice` = named options, `score` = ordered levels), optional | |
| `instructions`, and `criteria` (object of option -> description for choice, | |
| list of level descriptions for score). | |
| image: Optional image that is part of the state. Arrives from the canvas as a | |
| {"path"/"url"} file dict (or None). | |
| Returns: | |
| A self-contained HTML report of per-option probabilities and a SystemOne-style | |
| JSON response with the answer, confidence and probabilities for every question. | |
| """ | |
| parsed_questions = _parse_questions(questions) | |
| record = {"state": _parse_state(state), "questions": parsed_questions} | |
| pil_image = _load_image(image) | |
| if pil_image is not None: | |
| record["images"] = [_prep_image(pil_image)] | |
| probs, latency_ms, n_tokens = _forward(record) | |
| response = { | |
| "model": "clef", | |
| "answers": {qid: systemone_answer(parsed_questions[qid], probs[qid]) for qid in parsed_questions}, | |
| "usage": {"input_tokens": n_tokens, "output_tokens": 0}, | |
| } | |
| return _render(parsed_questions, probs, latency_ms, n_tokens), response | |
| demo = gr.Workflow(graph="workflow.json", bind={"decide": decide}) | |
| if __name__ == "__main__": | |
| demo.launch(mcp_server=True) | |