etanlightstone's picture
Upload Domino Decision Model checkpoint
1f9a4df verified
Raw History Blame Contribute Delete
9.53 kB
"""
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()