#!/usr/bin/env python3 """ Needle 3 -- Gradio SDK Space, GPU-capable via ZeroGPU. WHY THIS FILE REPLACES THE DOCKER SPACE ZeroGPU is available only to Spaces using the **Gradio SDK** (HF docs: "Currently, ZeroGPU Spaces are exclusively compatible with the Gradio SDK"). A Space's SDK is immutable after creation, so enabling ZeroGPU required a new Space, not a settings change. WHY THE JAX PATH (unchanged from the Docker build) Needle 3 ships two runtimes. The native C engine (libneedle.so + needle3.cact) exposes only needle_init/complete/embed/load/reset -- no logits, and its confidence head returned a constant 1.0 in testing. The JAX path (needle/model/run.py) does: logits = decode_fn(params, buffer)[0, pos] so a genuine probability distribution is obtainable. That is the entire reason this endpoint can answer in Jev's shape (choice / noul / score) at all. SHAPE A Gradio Blocks app is the primary interface (so the Space is SDK-legal and the @spaces.GPU decorator applies), and the Jev-shaped FastAPI routes are MOUNTED onto the same server via `gr.mount_gradio_app` / the `app=` argument, so existing HTTP clients keep working. @spaces.GPU is applied to the inference entry points. Outside a ZeroGPU environment the decorator is effect-free, so the same code runs on cpu-basic. """ from __future__ import annotations import json import math import os import threading from typing import Any os.environ.setdefault("NEEDLE_TELEMETRY", "0") # --------------------------------------------------------------------- backend # Detect GPU availability. JAX preallocating 75% of a shared ZeroGPU slice is # hostile to the dynamic allocation model, so preallocation is disabled before # JAX is imported anywhere. On a non-GPU Space these are harmless. if os.environ.get("SPACES_ZERO_GPU"): os.environ.setdefault("XLA_PYTHON_CLIENT_PREALLOCATE", "false") os.environ.setdefault("JAX_PLATFORMS", "cuda") else: os.environ.setdefault("JAX_PLATFORMS", "cpu") import gradio as gr # noqa: E402 try: import spaces # provided by ZeroGPU Spaces; absent on plain cpu-basic HAS_SPACES = True except Exception: HAS_SPACES = False class _NoopSpaces: """Effect-free stand-in so the same code runs without ZeroGPU.""" @staticmethod def GPU(*dargs, **dkwargs): def deco(fn): return fn if len(dargs) == 1 and callable(dargs[0]) and not dkwargs: return dargs[0] return deco spaces = _NoopSpaces() MODEL_ID = "needle3-gpu-0.3.0" CKPT = os.environ.get("NEEDLE_CKPT", "needle3.safetensors") def _ensure_ckpt() -> str: """Fetch the JAX checkpoint (242 MB) if it is not already present. Not committed to this repo: it would bloat every clone, and it already lives in the model repo. `needle.model.tokenizer.get_tokenizer()` fetches its own `tokenizer.model` automatically, so only the checkpoint needs handling here. """ if os.path.exists(CKPT) and os.path.getsize(CKPT) > 0: return CKPT from huggingface_hub import hf_hub_download path = hf_hub_download(repo_id="Cactus-Compute/needle3", filename="checkpoints/needle3.safetensors", local_dir=".") # hf_hub_download preserves the repo-relative path; normalise to CKPT import shutil if os.path.abspath(path) != os.path.abspath(CKPT): shutil.copyfile(path, CKPT) return CKPT _LOCK = threading.Lock() _CACHE: dict[str, Any] = {} def _load() -> dict: """Load params/config/tokenizer once (JAX path, logits-capable).""" if _CACHE: return _CACHE _ensure_ckpt() import jax import jax.numpy as jnp from needle.model.architecture import SimpleAttentionNetwork, TransformerConfig from needle.model.checkpoints import read_checkpoint from needle.model.tokenizer import get_tokenizer ckpt = read_checkpoint(CKPT) saved = ckpt["config"] if isinstance(ckpt["config"], dict) else dict(vars(ckpt["config"])) config = TransformerConfig.from_saved(saved) raw_params = {k: v for k, v in ckpt["params"].items() if not k.startswith("mtp_")} # The published checkpoint declares dtype=bfloat16, and TransformerConfig builds # every layer with that dtype. JAX on this GPU aborts on bf16->f16 conversions: # "Unsupported conversion from bf16 to f16 / # LLVM ERROR: Unsupported rounding mode for conversion." # Casting the PARAMS alone does not help -- the model's own dtype must change too, # so set it on the config BEFORE constructing the network. # NEEDLE_CAST: "float32" (default) | "bfloat16" (original) | "none" _cast = os.environ.get("NEEDLE_CAST", "float32").lower() if _cast in ("float32", "float", "f32"): # TransformerConfig is a plain @dataclass (no .replace()), so assign directly. config.dtype = "float32" params = {k: (v.astype("float32") if hasattr(v, "astype") else v) for k, v in raw_params.items()} else: params = raw_params model = SimpleAttentionNetwork(config=config) tokenizer = get_tokenizer() @jax.jit def decode_fn(p, tokens): return model.apply({"params": p}, tokens) _CACHE.update({"jax": jax, "jnp": jnp, "params": params, "model": model, "tokenizer": tokenizer, "decode_fn": decode_fn, "config": config, "backend": jax.default_backend(), "cast": _cast, "devices": [str(d) for d in jax.devices()]}) return _CACHE def _encode(text: str) -> list[int]: c = _load() from needle.model.tokenizer import BOS_ID return [BOS_ID] + c["tokenizer"].encode(text) def _first_token_probs(prompt: str, candidates: list[str]) -> dict[str, float]: """Distribution over each candidate's FIRST token at the decision point.""" c = _load() jnp = c["jnp"] ids = _encode(prompt) cfg = c["config"] buf_len = min(cfg.max_seq_len, len(ids) + 1) buf = jnp.full((1, buf_len), 0, dtype=jnp.int32) buf = buf.at[0, :len(ids)].set(jnp.array(ids, dtype=jnp.int32)) logits = c["decode_fn"](c["params"], buf)[0, len(ids) - 1] logp = c["jax"].nn.log_softmax(logits.astype(jnp.float32), axis=-1) raw: dict[str, float] = {} for cand in candidates: toks = c["tokenizer"].encode(cand) raw[cand] = float(jnp.exp(logp[int(toks[0])])) if toks else 0.0 total = sum(raw.values()) return {k: (v / total if total > 0 else 0.0) for k, v in raw.items()} def _sequence_probs(prompt: str, candidates: list[str]) -> dict[str, float]: """Teacher-forced score of each candidate as a COMPLETE continuation, then softmax-normalised across candidates.""" c = _load() jnp = c["jnp"] tok = c["tokenizer"] cfg = c["config"] prompt_ids = _encode(prompt) scores: dict[str, float] = {} for cand in candidates: cand_ids = tok.encode(cand) if not cand_ids: scores[cand] = float("-inf") continue full = prompt_ids + cand_ids if len(full) > cfg.max_seq_len: scores[cand] = float("-inf") continue n = len(full) buf = jnp.full((1, n), 0, dtype=jnp.int32) buf = buf.at[0, :n].set(jnp.array(full, dtype=jnp.int32)) logits = c["decode_fn"](c["params"], buf)[0] logp = c["jax"].nn.log_softmax(logits.astype(jnp.float32), axis=-1) start = len(prompt_ids) - 1 scores[cand] = sum(float(logp[start + i, int(t)]) for i, t in enumerate(cand_ids)) finite = {k: v for k, v in scores.items() if v != float("-inf")} if not finite: return {k: 0.0 for k in scores} mx = max(finite.values()) exps = {k: math.exp(v - mx) for k, v in finite.items()} z = sum(exps.values()) out = {k: (exps[k] / z if z else 0.0) for k, v in scores.items()} for k in scores: out.setdefault(k, 0.0) return out # --------------------------------------------------------------- Jev primitives def _do_choice(state: str, instructions: str, criteria: dict) -> dict: """Choice. Uses the BARE prompt + first-token matching, which measurably beat the native + full-sequence method (see README method table).""" options = list(criteria.keys()) lines = [instructions, "", f"Text: {state}", "", "Choose exactly one of: " + ", ".join(options), "Answer: "] prompt = "\n".join(lines) dist = _first_token_probs(prompt, options) best = max(dist, key=dist.get) if dist else None out = {k: v for k, v in dist.items()} tot = sum(out.values()) out = {k: (v / tot if tot else 0.0) for k, v in out.items()} return {"type": "choice", "choice": best, "probabilities": out, "confidence": max(out.values()) if out else 0.0} def _do_noul(state: str, instructions: str) -> dict: """Noul: P(yes) as a single number, per Jev's shape (no confidence field).""" prompt = f"{instructions}\n\nText: {state}\n\nAnswer yes or no. Answer: " dist = _first_token_probs(prompt, ["Yes", "No"]) p_yes = dist.get("Yes", 0.0) p_no = dist.get("No", 0.0) tot = p_yes + p_no return {"type": "noul", "noul": (p_yes / tot) if tot else 0.0} def _do_score(state: str, instructions: str, levels: list[str]) -> dict: """Score: probability-weighted expectation over ordered levels, so the result can land BETWEEN levels, as Jev's does. Method (changed after measurement on the native engine): score each level's DESCRIPTION as a continuation, rather than reading the probability of a digit token. Measured 6/8 exact and 7/8 within-one using per-level descriptions with triggers, versus a collapsing digit approach. Jev's own docs note "each level is judged on its own against the state", which is what scoring each level description independently approximates. """ prompt = (f"{instructions}\n\nText: {state}\n\n" "Which of these descriptions fits the text? Choose one.\n" + "\n".join(f"- {d}" for d in levels) + "\nBest match: ") dist = _sequence_probs(prompt, list(levels)) tot = sum(dist.values()) dist = {k: (v / tot if tot else 0.0) for k, v in dist.items()} score = sum(float(i) * dist.get(d, 0.0) for i, d in enumerate(levels)) return {"type": "score", "score": score, "legend": {str(i): d for i, d in enumerate(levels)}, "probabilities": {str(i): dist.get(d, 0.0) for i, d in enumerate(levels)}, "confidence": max(dist.values()) if dist else 0.0} @spaces.GPU(duration=240) def answer_systemone(state: str, questions: dict, backend: str = "needle") -> dict: """Run a Jev-shaped request on the chosen backend. backend="needle" -> Cactus Needle 3 through the JAX logits path (default) backend="bart" -> facebook/bart-large-mnli zero-shot, the free rival implementation of the same three primitives BART is deliberately NOT GPU-decorated work: it runs on the CPU of this Space (device=-1) and is loaded lazily, so it does not change Needle's cold start. The decorator stays on this entry point because it is one endpoint; the BART branch simply does not need the GPU. """ if backend == "bart": from bart_backend import bart_choice, bart_score, bart_noul answers: dict[str, Any] = {} for qid, q in questions.items(): q = q or {} qtype = q.get("type") instr = q.get("instructions", "") if qtype == "choice": answers[qid] = bart_choice(state, instr, q.get("criteria") or {}) elif qtype == "noul": answers[qid] = bart_noul(state, instr, yes_text=q.get("yes_text"), no_text=q.get("no_text")) elif qtype == "score": answers[qid] = bart_score(state, instr, q.get("criteria") or []) else: raise ValueError(f"unsupported question type: {qtype!r}") return {"model": "facebook/bart-large-mnli", "answers": answers, "usage": {"input_tokens": None, "output_tokens": None}, "_backend_note": "BART-MNLI zero-shot; multi_label=False (softmax " "across options, matching the cited article's code)."} with _LOCK: _load() answers: dict[str, Any] = {} for qid, q in questions.items(): qtype = (q or {}).get("type") instr = (q or {}).get("instructions", "") if qtype == "choice": answers[qid] = _do_choice(state, instr, q.get("criteria") or {}) elif qtype == "noul": answers[qid] = _do_noul(state, instr) elif qtype == "score": answers[qid] = _do_score(state, instr, q.get("criteria") or []) else: raise ValueError(f"unsupported question type: {qtype!r}") return {"model": MODEL_ID, "answers": answers, "usage": {"input_tokens": None, "output_tokens": None}} # --------------------------------------------------------------- tuned adapter vocab # The trigger vocabulary from the TUNED adapter (eval_harness.TRIGGERS), copied here so # the Space can drive the tuned path without importing the eval harness. These are the # lists that produced 87.5% held out; see the README for why each design choice exists. TRIG = { "billing": r"\b(charge|charged|charges|invoice|refund|refunded|billing|billed|bill|" r"payment|paid|pay|subscription|price|pricing|duplicate|double charged|" r"credited|owing|owed|bank|card|cost)\b", "technical": r"\b(error|errors|bug|bugs|crash|crashes|broken|fails|failing|failed|" r"outage|down|exception|blank|screen|sync|search|upload|webhook|api|" r"json|csv|export|import|dashboard|app|version|update|fix|feature|" r"documentation)\b", "account": r"\b(log ?in|login|log into|sign ?in|sign in|password|reset|two-factor|" r"2fa|authentication|locked out|lock|unlock|access|permission|profile|" r"email address|account|session expired|erasure|delete my account)\b", } @spaces.GPU(duration=240) def compare_backends() -> dict: """Head-to-head: Needle 3 vs BART-MNLI on the SAME questions. Uses three unambiguous cases with known correct answers, plus one polarity case in each direction, so both the positive and negative side of `noul` are exercised -- a one-sided set is how a "yes to everything" bug hides. """ crit = {"billing": "Payment, charges, invoices, refunds, subscriptions.", "technical": "Bugs, errors, outages, broken integrations.", "account": "Login, password, account access, permissions."} levels = ["Calm and patient, no rush", "Not a calm case: a problem is being reported or the writer is upset"] cases = [ ("I was charged twice for order A-104. Please refund the duplicate.", {"department": "billing", "refund_requested": 1}), ("I cannot log in, my password is rejected.", {"department": "account", "refund_requested": 0}), ("The export button throws a 500 error.", {"department": "technical", "refund_requested": 0}), ] Q = {"department": {"type": "choice", "instructions": "Which team?", "criteria": crit}, "refund_requested": {"type": "noul", "instructions": "Is the writer asking for a refund?", "yes_text": "the writer is asking for money back", "no_text": "the writer is not asking for any money back"}} out: dict[str, Any] = {"questions": list(Q), "cases": [], "summary": {}} backends = ("tuned", "space", "bart") tally = {b: {"choice": 0, "noul": 0, "n": 0} for b in backends} for state, expect in cases: row = {"state": state, "expect": expect} for backend in backends: try: if backend == "tuned": # the tuned adapter declares only its own criteria; give it the # REQUEST portion for trigger scoping, as its design requires tq = {"department": {"type": "choice", "instructions": "Which team?", "criteria": {k: {"description": v, "triggers": TRIG[k]} for k, v in crit.items()}}, "refund_requested": {"type": "noul", "instructions": "Is the writer asking " "for a refund?"}, "_request_text": state} r = answer_tuned(state, tq) elif backend == "space": r = answer_systemone(state, Q, backend="needle") else: r = answer_systemone(state, Q, backend="bart") except Exception as e: row[backend] = {"error": f"{type(e).__name__}: {e}"} continue ch = r["answers"]["department"] nl = r["answers"]["refund_requested"] got_ch = ch.get("choice") got_nl = int(round(nl.get("noul", 0.0))) row[backend] = {"choice": got_ch, "choice_correct": got_ch == expect["department"], "choice_conf": ch.get("confidence"), "noul": round(nl.get("noul", 0.0), 4), "noul_correct": got_nl == expect["refund_requested"], "model": r.get("model")} tally[backend]["n"] += 2 tally[backend]["choice"] += int(got_ch == expect["department"]) tally[backend]["noul"] += int(got_nl == expect["refund_requested"]) out["cases"].append(row) for b in tally: n2 = max(tally[b]["n"], 1) tally[b]["accuracy"] = (tally[b]["choice"] + tally[b]["noul"]) / n2 out["summary"] = tally out["label_note"] = ( "tuned = the trigger-based adapter that scored 87.5% held out on 28 tickets. " "space = this Space's bare-prompt first-token path (a DIFFERENT, weaker design; " "do not read it as 'Needle'). bart = facebook/bart-large-mnli zero-shot.") return out def _ensure_tuned_assets() -> dict: """Fetch what the TUNED adapter path needs, at MODULE scope. The tuned adapter (jev_adapter.py) uses the NATIVE engine -- `needle.Needle(tools=...)` with no `weights=` -- which is the configuration that scored 87.5% held out. That is a DIFFERENT runtime from the JAX-logits path this Space normally uses: it needs the native `libneedle` shared library plus the `needle3.cact` archive, both of which the package auto-downloads into ~/.cache/cactus-needle/. Why module scope: network access inside a @spaces.GPU function is restricted on ZeroGPU, so anything that fetches must run at import. This is the same reason the safetensors checkpoint is prefetched at startup. """ info: dict[str, Any] = {"ok": False} try: import needle as _n # _load_base fetches/validates the .cact; the engine library is fetched lazily # by _lib_path on first real call, so trigger it here too. from needle import _base_weights_path, _lib_path lib = _lib_path(3) w = _base_weights_path(3) info.update({"ok": True, "lib": str(lib), "weights": str(w), "weights_bytes": os.path.getsize(w) if os.path.exists(w) else None}) except Exception as e: info["error"] = f"{type(e).__name__}: {e}" return info @spaces.GPU(duration=600) def answer_tuned(state: str, questions: dict) -> dict: """The TUNED adapter path: one tool per option, triggers, scoped to the request. This is the implementation that scored 87.5% held out on the 28-ticket set. It is exposed separately from `answer_systemone` because it is a different runtime (native engine vs JAX logits) and a different design (tools + triggers vs bare-prompt first-token matching). Keeping them distinct stops the comparison from silently attributing the naive path's score to the tuned one. """ from jev_adapter import run_systemone as _tuned # The adapter needs the request alone for trigger scoping. The caller may pass a # combined state; if it does not declare the request portion we use the whole thing # and the adapter warns about it. request_text = questions.pop("_request_text", None) if isinstance(questions, dict) else None return _tuned(state, questions, request_text=request_text) @spaces.GPU(duration=120) def run_method_comparison() -> dict: """The scoring-method comparison on three unambiguous cases. A first-token, bare prompt (best in testing: 2/3, margin +0.106) B whole-word, bare prompt C first-token, native D full tool-call payload, native """ with _LOCK: _load() crit = {"billing": "Payment, charges, invoices, refunds, subscriptions.", "technical": "Bugs, errors, outages, broken integrations.", "account": "Login, password, account access, permissions."} instr = "Which team should handle this?" cases = [("I was charged twice for order A-104. Please refund the duplicate.", "billing"), ("I cannot log in, my password is rejected.", "account"), ("The export button throws a 500 error.", "technical")] opts = list(crit.keys()) rows = {m: [] for m in "ABCD"} for state, exp in cases: bare = "\n".join([instr, "", f"Text: {state}", "", "Choose exactly one of: " + ", ".join(opts), "Answer: "]) dA = _first_token_probs(bare, opts) dB = _sequence_probs(bare, opts) dC = _first_token_probs(bare, opts) payloads = [f'{{"name": "classify", "arguments": {{"category": {json.dumps(o)}}}}}' for o in opts] dD = _sequence_probs(bare, payloads) dD = {o: dD.get(p, 0.0) for o, p in zip(opts, payloads)} for m, d in (("A", dA), ("B", dB), ("C", dC), ("D", dD)): best = max(d, key=d.get) if d else None rows[m].append({"expected": exp, "picked": best, "correct": best == exp, "p_expected": round(d.get(exp, 0.0), 4)}) summary = {} for m, r in rows.items(): summary[m] = {"correct": sum(1 for x in r if x["correct"]), "mean_p_expected": round(sum(x["p_expected"] for x in r) / len(r), 4)} return {"model": MODEL_ID, "summary": summary, "detail": rows, "legend": {"A": "first-token, bare (best in testing)", "B": "whole-word continuation, bare", "C": "first-token, native tool_call", "D": "full tool_call payload, native"}} def health() -> str: """Lightweight check. Deliberately does NOT import JAX. On ZeroGPU the GPU is only attached inside a @spaces.GPU function, so initialising JAX here fails with "No visible GPU devices". Use the GPU probe tab to test the actual JAX/CUDA path. """ ok_ckpt = os.path.exists(CKPT) and os.path.getsize(CKPT) > 0 return (f"ok=True ckpt_present={ok_ckpt} ckpt={CKPT} " f"spaces_module={HAS_SPACES} " f"SPACES_ZERO_GPU={os.environ.get('SPACES_ZERO_GPU')} " f"JAX_PLATFORMS={os.environ.get('JAX_PLATFORMS')}") @spaces.GPU(duration=120) def gpu_probe() -> dict: """DECISIVE TEST: does JAX initialise on CUDA *inside* a @spaces.GPU function? Outside the decorator there is no GPU at all (that is why health() does not touch JAX). Whether JAX can use the ZeroGPU slice at all is undocumented -- ZeroGPU is PyTorch-shaped, and JAX-on-ZeroGPU is not a supported combination. This endpoint answers it with the actual device list. """ out: dict[str, Any] = {"env_SPACES_ZERO_GPU": os.environ.get("SPACES_ZERO_GPU"), "env_JAX_PLATFORMS": os.environ.get("JAX_PLATFORMS"), "env_PREALLOCATE": os.environ.get("XLA_PYTHON_CLIENT_PREALLOCATE")} try: import jax out["backend"] = jax.default_backend() out["devices"] = [str(d) for d in jax.devices()] out["device_count"] = len(jax.devices()) # prove a real computation runs on the GPU import jax.numpy as jnp x = jnp.arange(8.0) out["computed"] = float(jnp.sum(x * 2)) out["ok"] = True except Exception as e: # noqa: BLE001 out["ok"] = False out["error"] = f"{type(e).__name__}: {e}" return out @spaces.GPU(duration=120) def run_selftest() -> dict: """Known-answer probes, comparing positive vs negative on matched pairs.""" with _LOCK: _load() pairs = [("Does the customer request a refund?", "I was charged twice. Please refund order A-104.", "The meeting is Tuesday at 3pm."), ("Does the message convey urgency?", "HELP! Everything is down and we are losing money NOW.", "Just wondering about your pricing tiers sometime.")] out = [] for instr, pos, neg in pairs: dp = _first_token_probs(f"{instr}\n\nText: {pos}\n\nAnswer yes or no. Answer: ", ["Yes", "No"]) dn = _first_token_probs(f"{instr}\n\nText: {neg}\n\nAnswer yes or no. Answer: ", ["Yes", "No"]) out.append({"instruction": instr, "p_yes_positive": round(dp.get("Yes", 0.0), 4), "p_yes_negative": round(dn.get("Yes", 0.0), 4), "discriminates": round(dp.get("Yes", 0.0) - dn.get("Yes", 0.0), 4)}) c = _load() return {"ok": True, "backend": c["backend"], "devices": c["devices"], "pairs": out, "mean_discrimination": round(sum(p["discriminates"] for p in out) / len(out), 4)} # --------------------------------------------------------------------- Gradio UI CARD = """ # needle3 logits (Gradio) A Jev-shaped (`/v1/systemone`) endpoint backed by [Cactus Needle 3](https://huggingface.co/Cactus-Compute/needle3), using its **JAX path** so real next-token logits are available. **Honest status:** distributions are real; accuracy is not yet good. On three unambiguous classification cases the best method scores **2/3**. `noul` (yes/no) did not discriminate on matched positive/negative pairs. See the *Method comparison* and *Selftest* tabs for the raw numbers rather than a claim. All three primitives read the next-token distribution over your candidate strings, with the bare prompt format that measured best (the model's native `` format scored *worse*). """ with gr.Blocks(title="needle3 logits") as demo: gr.Markdown(CARD) with gr.Tab("System One"): st = gr.Textbox(label="state", lines=4, value="I was charged twice for order A-104. Please refund the duplicate.") instr = gr.Textbox(label="instructions", value="Which team should handle this?") backend_sel = gr.Radio(choices=["needle", "bart"], value="needle", label="backend", info="needle = Cactus Needle 3 (JAX logits). " "bart = facebook/bart-large-mnli zero-shot.") with gr.Row(): b1 = gr.Button("choice: team", variant="primary") b2 = gr.Button("noul: asks for refund?") b3 = gr.Button("score: frustration") out_json = gr.JSON(label="response") def _choice(state, instructions, backend): return answer_systemone(state, {"department": { "type": "choice", "instructions": instructions, "criteria": {"billing": "Payment, charges, invoices, refunds, subscriptions.", "technical": "Bugs, errors, outages, broken integrations.", "account": "Login, password, account access, permissions."}}}, backend=backend) def _noul(state, instructions, backend): # Honour the instructions box. Previously this ignored it and asked a # hardcoded question, so a custom question was silently discarded. # For BART both sides of the pair are supplied: the article flags the # hand-written negative side as this approach's real weakness. q = {"type": "noul", "instructions": instructions or "Does the message ask for a refund?"} if backend == "bart": q["yes_text"] = (instructions or "asks for a refund") + ": yes" q["no_text"] = (instructions or "asks for a refund") + ": no" return answer_systemone(state, {"answer": q}, backend=backend) def _score(state, instructions, backend): return answer_systemone(state, {"frustration": { "type": "score", "instructions": instructions or "How frustrated is the customer?", "criteria": ["Calm", "Frustrated but civil", "Very angry"]}}, backend=backend) b1.click(_choice, [st, instr, backend_sel], out_json, api_name="choice") b2.click(_noul, [st, instr, backend_sel], out_json, api_name="noul") b3.click(_score, [st, instr, backend_sel], out_json, api_name="score") with gr.Tab("Compare backends"): gr.Markdown("**tuned** = the trigger-based adapter (87.5% held out) · " "**space** = this Space's bare-prompt path (weaker, different design) · " "**bart** = facebook/bart-large-mnli zero-shot. Both polarities of " "`noul` are exercised, so a 'yes to everything' bug cannot hide.") cmp_btn = gr.Button("run head-to-head", variant="primary") cmp_out = gr.JSON(label="comparison") cmp_btn.click(compare_backends, None, cmp_out, api_name="compare") tun_btn = gr.Button("tuned adapter only (single state)") tun_out = gr.JSON(label="tuned response") tun_q = gr.Code(label="questions (JSON) -- passed through verbatim", language="json", value=json.dumps({ "department": {"type": "choice", "instructions": "Which team?", "criteria": {k: {"description": v, "triggers": TRIG[k]} for k, v in {"billing": "Payment, charges, refunds.", "technical": "Bugs, errors, outages.", "account": "Login, password, access."}.items()}}, "refund_requested": {"type": "noul", "instructions": "Is the writer asking for a refund?"}}, indent=2)) def _run_tuned(state, questions_text): """Run the tuned adapter on the CALLER'S questions. This used to take only the state and hardcode its own two questions. That made the endpoint useless for comparison work -- a head-to-head that sent its own questions was silently answered on a different, fixed set, and the mismatch looked like a model result. Questions are now passed through. """ if isinstance(questions_text, str): try: qs = json.loads(questions_text) except Exception as e: return {"error": f"questions is not valid JSON: {e}"} else: qs = questions_text if not isinstance(qs, dict) or not qs: return {"error": "questions must be a non-empty JSON object"} qs = dict(qs) # RESPECT a caller-supplied _request_text. This line used to be an # unconditional `qs["_request_text"] = state`, which OVERWROTE the request # portion the client sent and forced the full state as the scoping text. # That re-introduced the reference-leak bug the adapter was built to avoid: # a refund-policy sentence in the state armed the yes-side and the answer # became yes for almost every policy question. The client knows which part # of the state is the request; do not overwrite it. qs.setdefault("_request_text", state) return answer_tuned(state, qs) tun_btn.click(_run_tuned, [st, tun_q], tun_out, api_name="tuned") with gr.Tab("Method comparison"): gr.Markdown("Which scoring method actually separates the three clear cases?") mc_btn = gr.Button("run comparison", variant="primary") mc_out = gr.JSON() mc_btn.click(run_method_comparison, None, mc_out, api_name="methods") with gr.Tab("Selftest"): gr.Markdown("Known-answer probes. `discriminates` near 0 or negative means the " "yes/no signal is not usable.") st_btn = gr.Button("run selftest", variant="primary") st_out = gr.JSON() st_btn.click(run_selftest, None, st_out, api_name="selftest") with gr.Tab("GPU probe"): gr.Markdown( "**Does JAX work on ZeroGPU at all?** Outside a `@spaces.GPU` function there is " "no GPU attached, so this is the only way to tell. JAX-on-ZeroGPU is not a " "documented-supported combination (ZeroGPU is PyTorch-shaped), so this may fail " "even though the hardware is allocated." ) gp_btn = gr.Button("probe GPU", variant="primary") gp_out = gr.JSON() gp_btn.click(gpu_probe, None, gp_out, api_name="gpu_probe") with gr.Tab("Health"): h_btn = gr.Button("check") h_out = gr.Textbox(label="status") h_btn.click(health, None, h_out, api_name="health") if __name__ == "__main__": # Fetch the checkpoint BEFORE launch. Gradio executes this file as __main__, # so a prefetch placed in an `else:` branch would never run. # # Why module scope, not inside @spaces.GPU: network access inside a # GPU-decorated function is restricted on ZeroGPU, and the documented pattern # is to place model assets at root module level. Failure is surfaced via # /health rather than breaking boot. try: _ensure_ckpt() print(f"[startup] checkpoint ready: {CKPT} " f"({os.path.getsize(CKPT) if os.path.exists(CKPT) else 0} bytes)") except Exception as _e: # noqa: BLE001 print(f"[startup] checkpoint prefetch failed: {type(_e).__name__}: {_e}") # The tuned adapter uses the NATIVE engine (libneedle + needle3.cact), which is a # DIFFERENT runtime from the JAX safetensors path above. Prefetch it here too, so the # first tuned call does not hit the network inside a GPU-decorated function. _tuned_asset_info = _ensure_tuned_assets() print(f"[startup] tuned assets: {_tuned_asset_info}") demo.queue(max_size=8).launch(server_name="0.0.0.0", server_port=7860) else: try: _ensure_ckpt() except Exception as _e: # noqa: BLE001 print(f"[startup] checkpoint prefetch failed: {type(_e).__name__}: {_e}") try: print(f"[startup] tuned assets: {_ensure_tuned_assets()}") except Exception as _e: # noqa: BLE001 print(f"[startup] tuned asset prefetch failed: {type(_e).__name__}: {_e}")