import spaces # must come before torch import json import sys import time import gradio as gr import torch from huggingface_hub import snapshot_download from safetensors.torch import load_file MODEL_ID = "IFM/K2-Type-0.9B" MODEL_DIR = snapshot_download(MODEL_ID) sys.path.insert(0, MODEL_DIR) # the repo ships its own inference package `jev/` from transformers import AutoTokenizer # noqa: E402 from jev.encode import Encoder, collate # noqa: E402 from jev.model import DecisionModel # noqa: E402 from jev.serve import answer, to_record # noqa: E402 CFG = json.load(open(f"{MODEL_DIR}/decision_config.json")) MAX_LEN = 8192 tokenizer = AutoTokenizer.from_pretrained(MODEL_DIR, trust_remote_code=True) model = DecisionModel(MODEL_DIR, head_dim=CFG.get("head_dim", 256)) model.head.load_state_dict(load_file(f"{MODEL_DIR}/pointer_head.safetensors")) model.temperature.fill_(CFG["temperature"]) model = model.eval().to("cuda") encoder = Encoder(tokenizer, MAX_LEN, MAX_LEN - 1024) PAD_ID = tokenizer.pad_token_id if tokenizer.pad_token_id is not None else tokenizer.eos_token_id TYPE_LABELS = {"Yes / No": "noul", "Choice": "choice", "Score (ordered)": "score"} TYPE_NAMES = {v: k for k, v in TYPE_LABELS.items()} def _run_request(req: dict) -> dict: """Run one /v1/systemone-style request through the model (single forward pass).""" t0 = time.perf_counter() try: rec = to_record(req) except Exception as e: # jev raises fastapi HTTPException raise gr.Error(getattr(e, "detail", str(e))) e = encoder.encode(rec) if e is None or len(e["decide"]) != len(rec["questions"]): raise gr.Error(f"Request does not fit in {MAX_LEN} tokens.") with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16): scores = model(collate([e], PAD_ID)) answers = {} for k, s in zip(e["qkeys"], scores): p = torch.softmax(s.float(), -1).tolist() answers[k] = answer(rec["questions"][k], p) torch.cuda.synchronize() return {"answers": answers, "model": CFG["name"], "input_tokens": len(e["ids"]), "latency_ms": round((time.perf_counter() - t0) * 1000, 1)} def _parse_state(state: str): s = (state or "").strip() if not s: raise gr.Error("Please enter a state (text or JSON).") if s[:1] in "{[": try: return json.loads(s) except json.JSONDecodeError: pass return s def _kv_lines(text: str) -> dict: out = {} for line in (text or "").splitlines(): line = line.strip() if not line: continue if ":" in line: k, v = line.split(":", 1) out[k.strip()] = v.strip() or None else: out[line] = None return out def _build_question(qtype: str, instructions: str, options: str): instructions = (instructions or "").strip() if not instructions: return None t = TYPE_LABELS.get(qtype, qtype) if t == "choice": crit = _kv_lines(options) if not crit: raise gr.Error(f"Choice question “{instructions}” needs at least one option (one per line).") return {"type": "choice", "instructions": instructions, "criteria": crit} if t == "score": levels = [l.strip() for l in (options or "").splitlines() if l.strip()] if len(levels) < 2: raise gr.Error(f"Score question “{instructions}” needs at least two levels (one per line, low → high).") return {"type": "score", "instructions": instructions, "criteria": levels} crit = {k.lower(): v for k, v in _kv_lines(options).items() if k.lower() in ("true", "false") and v} q = {"type": "noul", "instructions": instructions} if crit: q["criteria"] = crit return q def _label_for(ans: dict) -> dict: if ans["type"] == "noul": return {"true": ans["noul"], "false": round(1 - ans["noul"], 4)} if ans["type"] == "choice": return ans["probabilities"] return {f"{i}: {ans['legend'][i]}": p for i, p in ans["probabilities"].items()} def _summary(qid: str, q: dict, ans: dict) -> str: if ans["type"] == "noul": p = ans["noul"] return f"**{qid}** — {q['instructions']} → **{'YES' if p >= 0.5 else 'NO'}** (P(true) = {p:.3f})" if ans["type"] == "choice": return (f"**{qid}** — {q['instructions']} → **{ans['choice']}** " f"(p = {ans['probabilities'][ans['choice']]:.3f}, confidence {ans['confidence']:.2f})") lvl = ans["legend"][str(round(ans["score"]))] return (f"**{qid}** — {q['instructions']} → expected level **{ans['score']:.2f}** " f"(≈ {lvl}; confidence {ans['confidence']:.2f})") @spaces.GPU(duration=5) def decide( state: str, q1_type: str = "Choice", q1_instructions: str = "", q1_options: str = "", q2_type: str = "Yes / No", q2_instructions: str = "", q2_options: str = "", q3_type: str = "Score (ordered)", q3_instructions: str = "", q3_options: str = "", ): """Answer up to three typed questions about a state with K2-Type-0.9B in one forward pass. Args: state: The situation to decide about, as plain text or a JSON object. q1_type: "Yes / No", "Choice" or "Score (ordered)". q1_instructions: The question / statement. Leave empty to skip this question. q1_options: Choice: one option per line ("name: description"). Score: one level per line, low to high. Yes / No: optional "true: ..." and "false: ..." definitions. q2_type: Type of question 2. q2_instructions: Question 2 text (empty to skip). q2_options: Question 2 options. q3_type: Type of question 3. q3_instructions: Question 3 text (empty to skip). q3_options: Question 3 options. Returns: A markdown summary, one probability label per question, and the raw /v1/systemone response. """ slots = [(q1_type, q1_instructions, q1_options), (q2_type, q2_instructions, q2_options), (q3_type, q3_instructions, q3_options)] questions, keys = {}, [] for i, slot in enumerate(slots, 1): q = _build_question(*slot) keys.append(f"q{i}" if q else None) if q: questions[f"q{i}"] = q if not questions: raise gr.Error("Fill in at least one question.") req = {"state": _parse_state(state), "questions": questions} res = _run_request(req) lines = [_summary(k, questions[k], res["answers"][k]) for k in questions] lines.append(f"\n{res['input_tokens']} input tokens · {res['latency_ms']} ms on GPU") labels = [gr.update(value=_label_for(res["answers"][k]), visible=True) if k else gr.update(value=None, visible=False) for k in keys] return "\n\n".join(lines), *labels, {"request": req, "response": res} @spaces.GPU(duration=5) def systemone(request_json: str) -> dict: """Raw TypeSafe /v1/systemone call: {"state": ..., "questions": {id: {"type", "instructions", "criteria"}}}. Args: request_json: The request body as a JSON string. Returns: The /v1/systemone response with per-question answers and probabilities. """ try: req = json.loads(request_json) except json.JSONDecodeError as e: raise gr.Error(f"Invalid JSON: {e}") if not isinstance(req, dict) or "state" not in req or not req.get("questions"): raise gr.Error('Request must be an object with "state" and "questions".') return _run_request(req) TICKET = json.dumps({"subject": "Charged twice", "body": "You billed my card twice for March. Refund one or I cancel."}, indent=2) EXAMPLES = [ [TICKET, "Choice", "Which queue handles this?", "billing: Payments and refunds\ntechnical: Bugs and login\ngeneral: Anything else", "Yes / No", "The customer sounds angry.", "", "Score (ordered)", "How urgent is it?", "Low\nNormal\nHigh\nCritical"], ["The app crashes every time I try to log in with Google on my Android phone since yesterday's update. " "I have a client demo in two hours.", "Choice", "Which queue handles this?", "billing: Payments and refunds\ntechnical: Bugs and login\ngeneral: Anything else", "Yes / No", "The user mentions a time constraint.", "", "Score (ordered)", "How urgent is it?", "Low\nNormal\nHigh\nCritical"], ["Review: The hotel room was spotless and the staff were lovely, but the walls were paper-thin " "and we barely slept because of the party next door.", "Choice", "What is the overall sentiment of the review?", "positive\nnegative\nmixed", "Yes / No", "The reviewer would recommend this hotel to a light sleeper.", "true: they would recommend it\nfalse: they would not recommend it", "Score (ordered)", "Star rating the reviewer most likely gave.", "1 star\n2 stars\n3 stars\n4 stars\n5 stars"], ["Premise: A man is playing a guitar on a crowded street corner while people drop coins in his case.\n" "Hypothesis: A musician is performing in public.", "Choice", "Does the premise entail the hypothesis?", "entailment\nneutral\ncontradiction", "Yes / No", "The man is being paid for his music.", "", "Score (ordered)", "How confident can we be that the man is a professional musician?", "Not at all\nSlightly\nModerately\nVery"], ] RAW_EXAMPLE = json.dumps({ "state": {"subject": "Charged twice", "body": "You billed my card twice for March. Refund one or I cancel."}, "questions": { "queue": {"type": "choice", "instructions": "Which queue handles this?", "criteria": {"billing": "Payments and refunds", "technical": "Bugs and login", "general": "Anything else"}}, "angry": {"type": "noul", "instructions": "The customer sounds angry."}, "urgency": {"type": "score", "instructions": "How urgent is it?", "criteria": ["Low", "Normal", "High", "Critical"]}, }}, indent=2) CSS = """ #col-container { max-width: 1150px; margin: 0 auto; } .dark .gradio-container { color: var(--body-text-color); } """ with gr.Blocks(title="K2-Type-0.9B") as demo: with gr.Column(elem_id="col-container"): gr.Markdown( "# ⚖️ K2-Type-0.9B — typed decision model\n" "Give a **state** (text or JSON) and up to three typed questions — **yes/no**, **choice**, or " "**ordered score**. The model returns a calibrated probability for every option of every question from " "**one forward pass**; questions can't see each other, so adding one never changes another's answer. " "It never generates text.\n\n" "[Model card](https://huggingface.co/IFM/K2-Type-0.9B) · " "base: [IFM/K2-Horizon-0.9B](https://huggingface.co/IFM/K2-Horizon-0.9B)" ) with gr.Tab("Question builder"): with gr.Row(): with gr.Column(scale=5): state = gr.Textbox(label="State (text or JSON)", lines=7, value=TICKET) qboxes = [] defaults = EXAMPLES[0][1:] for i in range(3): with gr.Group(): with gr.Row(): qt = gr.Dropdown(list(TYPE_LABELS), value=defaults[3 * i], label=f"Question {i + 1} type", scale=1) qi = gr.Textbox(label=f"Question {i + 1} (leave empty to skip)", value=defaults[3 * i + 1], scale=3) qo = gr.Textbox( label="Options — Choice: one per line, `name: description` · Score: levels low→high · " "Yes/No: optional `true: …` / `false: …`", value=defaults[3 * i + 2], lines=3) qboxes += [qt, qi, qo] run = gr.Button("Decide", variant="primary") with gr.Column(scale=4): summary = gr.Markdown() labels = [gr.Label(label=f"Question {i + 1}", num_top_classes=10) for i in range(3)] with gr.Accordion("Raw request / response", open=False): raw = gr.JSON() run.click(decide, inputs=[state, *qboxes], outputs=[summary, *labels, raw], api_name="decide") gr.Examples( examples=EXAMPLES, inputs=[state, *qboxes], outputs=[summary, *labels, raw], fn=decide, cache_examples=False, run_on_click=True, ) with gr.Tab("Raw /v1/systemone"): gr.Markdown("Send any number of questions in TypeSafe's `/v1/systemone` wire format.") with gr.Row(): req_box = gr.Code(value=RAW_EXAMPLE, language="json", label="Request", lines=22) res_box = gr.JSON(label="Response") raw_btn = gr.Button("Send", variant="primary") raw_btn.click(systemone, inputs=req_box, outputs=res_box, api_name="systemone") demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)