akhaliq's picture
akhaliq HF Staff
gr.Workflow app: bind decide() as an fn node
04336c7 verified
Raw History Blame Contribute Delete
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))
@spaces.GPU(duration=_gpu_duration, size="xlarge")
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)