Download invoke.py from etanlightstone/ddm-medium-injection: direct link, hf CLI and curl.
- Browser
- Download file 9.53 kB
-
https://huggingface.co/etanlightstone/ddm-medium-injection/resolve/main/invoke.py
- Command line
-
hf download hf://etanlightstone/ddm-medium-injection/invoke.py
-
curl -L -o invoke.py https://huggingface.co/etanlightstone/ddm-medium-injection/resolve/main/invoke.py
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) | |
| # -------------------------------------------------------------------------------------- | |
| class Choice: | |
| instructions: Any | |
| criteria: dict[str, str] | |
| def to_dict(self): | |
| return {"type": "choice", "instructions": self.instructions, "criteria": self.criteria} | |
| class Score: | |
| instructions: Any | |
| criteria: list[str] | |
| def to_dict(self): | |
| return {"type": "score", "instructions": self.instructions, "criteria": self.criteria} | |
| 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) | |
| 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))), | |
| }, | |
| } | |
| 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() | |