""" invoke.py - DominoDecisionClient: a local, Jev-style API wrapper around a trained Domino Decision Model (DominoDecisionModel in model.py). The system_one() method name is kept on purpose so code written for TypeSafe's SDK ports over with minimal changes. Unlike TypeSafe's TypeSafeClient (which calls their hosted API), this client loads a local checkpoint and runs it in-process. No network calls, no API key. Mirrors TypeSafe's request/response shape (docs.typesafe.ai/introduction/quickstart): state + typed questions in, typed answers with probabilities out. CLI python invoke.py --checkpoint runs/tiny/best.pt # runs a demo request python invoke.py --checkpoint runs/tiny/best.pt --request req.json cat req.json | python invoke.py --checkpoint runs/tiny/best.pt --request - Python (same feel as the typesafe-sdk) from invoke import DominoDecisionClient, Choice, Score, Noul client = DominoDecisionClient("runs/tiny/best.pt") response = client.system_one( state="Hi, my Stripe integration keeps failing and I'm losing sales. Help ASAP.", questions={ "department": Choice( instructions="Which team should handle this", criteria={"billing": "Payment or subscription issues", "technical": "Bugs or integration problems", "sales": "Pricing or account questions"}), "frustration": Score( instructions="How frustrated the customer appears", criteria=["Calm, just stating facts", "Frustrated but civil", "Very angry, strong language"]), "is_urgent": Noul(instructions="The message conveys urgency or time-sensitivity"), }, ) print(response["answers"]["department"]["choice"]) # or send the raw JSON body directly client.run({"state": "...", "questions": {...}}) """ from __future__ import annotations import argparse import json import os import sys from dataclasses import dataclass from typing import Any import torch from model import (RequestError, assemble, batch_to, collate, encode_branch, encode_state, get_tokenizer, load_checkpoint, normalize_questions, state_to_text) # -------------------------------------------------------------------------------------- # Question helpers (like typesafe_sdk.Choice / Score / Noul) # -------------------------------------------------------------------------------------- @dataclass class Choice: instructions: Any criteria: dict[str, str] def to_dict(self): return {"type": "choice", "instructions": self.instructions, "criteria": self.criteria} @dataclass class Score: instructions: Any criteria: list[str] def to_dict(self): return {"type": "score", "instructions": self.instructions, "criteria": self.criteria} @dataclass class Noul: instructions: Any def to_dict(self): return {"type": "noul", "instructions": self.instructions} # -------------------------------------------------------------------------------------- # Client # -------------------------------------------------------------------------------------- class DominoDecisionClient: def __init__(self, checkpoint: str, device: str | None = None, model_name: str | None = None): self.device = device or ("cuda" if torch.cuda.is_available() else "cpu") self.model, self.meta = load_checkpoint(checkpoint, map_location=self.device) self.model.eval() self.cfg = self.model.config self.tok = get_tokenizer(self.cfg.tokenizer) self.model_name = model_name or f"ddm-local:{os.path.basename(checkpoint)}" def system_one(self, state, questions: dict, model: str | None = None) -> dict: body = {"state": state, "questions": {k: v.to_dict() if hasattr(v, "to_dict") else v for k, v in questions.items()}} return self.run(body) @torch.inference_mode() def run(self, request: dict) -> dict: """Takes the raw JSON request body, returns the JSON response body. Raises RequestError for invalid requests (like an HTTP 422).""" if not isinstance(request, dict) or "state" not in request: raise RequestError("request must be an object with 'state' and 'questions'") qlist = normalize_questions(request.get("questions"), self.cfg) state_ids = encode_state(self.tok, state_to_text(request["state"])) branches = [encode_branch(self.tok, q) for q in qlist] # Budget checks + "fan-out": if all questions don't fit in one packed sequence, # split them into several rows that each repeat the state (still one forward pass). S = len(state_ids) for q, b in zip(qlist, branches): if S + len(b[0]) > self.cfg.max_branch_tokens: raise RequestError(f"question '{q['id']}': state + question exceeds " f"{self.cfg.max_branch_tokens} tokens") groups, cur, cur_len = [], [], S for i, b in enumerate(branches): if cur and cur_len + len(b[0]) > self.cfg.max_total_tokens: groups.append(cur) cur, cur_len = [], S cur.append(i) cur_len += len(b[0]) groups.append(cur) encoded = [assemble(state_ids, [branches[i] for i in g], [qlist[i]["type"] for i in g]) for g in groups] batch = batch_to(collate(encoded, self.tok.pad_id), self.device) opt_probs, noul_probs = self.model.predict(batch) opt_probs, noul_probs = opt_probs.float().cpu(), noul_probs.float().cpu() order = [i for g in groups for i in g] # flattened question order answers_by_idx = {} for flat, qi in enumerate(order): q = qlist[qi] K = len(q["keys"]) answers_by_idx[qi] = self._format(q, opt_probs[flat, :K].tolist(), noul_probs[flat].item()) answers = {qlist[i]["id"]: answers_by_idx[i] for i in range(len(qlist))} return { "model": self.model_name, "answers": answers, "usage": { "input_tokens": sum(len(e["ids"]) for e in encoded), # Like Jev, this is a billing-style figure computed from the serialised # answers after inference - nothing is generated token by token. "output_tokens": len(self.tok.encode(json.dumps(answers))), }, } @staticmethod def _format(q: dict, probs: list[float], noul_p: float) -> dict: r = lambda x: round(float(x), 2) if q["type"] == "noul": return {"type": "noul", "noul": r(noul_p)} K = len(probs) best = max(range(K), key=lambda i: probs[i]) if q["type"] == "choice": # Margin of the top option above uniform (same formula as TypeSafe's adapter). conf = 1.0 if K == 1 else (probs[best] - 1 / K) / (1 - 1 / K) return {"type": "choice", "choice": q["keys"][best], "confidence": r(conf), "probabilities": {k: r(p) for k, p in zip(q["keys"], probs)}} # score: probability-weighted level; confidence = 1 - normalised spread around mode expected = sum(i * p for i, p in enumerate(probs)) spread = sum(p * abs(i - best) for i, p in enumerate(probs)) / (K - 1) return {"type": "score", "score": r(expected), "confidence": r(1 - spread), "legend": {str(i): t.split(": ", 1)[1] for i, t in enumerate(q["texts"])}, "probabilities": {str(i): r(p) for i, p in enumerate(probs)}} DEMO_REQUEST = { "state": "Hi, I've been trying to connect my Stripe account for 3 days and the integration " "keeps failing. I'm losing sales. Please help ASAP.", "questions": { "department": {"type": "choice", "instructions": "Which team should handle this", "criteria": {"billing": "Payment or subscription issues", "technical": "Bugs or integration problems", "sales": "Pricing or account questions"}}, "frustration": {"type": "score", "instructions": "How frustrated the customer appears", "criteria": ["Calm, just stating facts", "Frustrated but civil", "Very angry, strong language"]}, "is_urgent": {"type": "noul", "instructions": "The message conveys urgency or time-sensitivity"}, }, } def main(): p = argparse.ArgumentParser(description="Run a Jev-style request against a local model") p.add_argument("--checkpoint", required=True) p.add_argument("--request", help="request JSON file, or '-' for stdin (default: demo)") p.add_argument("--device", default=None) args = p.parse_args() if args.request == "-": request = json.load(sys.stdin) elif args.request: with open(args.request, encoding="utf-8") as f: request = json.load(f) else: request = DEMO_REQUEST client = DominoDecisionClient(args.checkpoint, device=args.device) try: print(json.dumps(client.run(request), indent=2)) except RequestError as e: print(json.dumps({"error": {"status": 422, "message": str(e)}}, indent=2)) sys.exit(1) if __name__ == "__main__": main()