File size: 9,920 Bytes
04336c7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
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)