"""SystemOne: the public inference API (see SPEC.md). m = SystemOne.load("checkpoints/jevlike-base") m.predict(state, {"dept": Choice(...), "urgency": Score(...), "churn": Noul(...)}) All questions of a call (or of a whole ``predict_batch``) are scored in one batched forward pass (split only by a token budget). Probabilities are temperature-scaled with ``calibration.json`` when present (``calibrated=False`` returns raw T=1 outputs). Escalate semantics: the escalate head predicts whether the argmax answer is correct; after Platt scaling (fit on val by ``calibrate.py``) ``escalate_prob = 1 - P(argmax correct)`` and ``escalate = escalate_prob > threshold`` (``escalate_threshold`` in config.json, default 0.5, i.e. hand off when the model thinks its own answer is more likely wrong than right). Long states ----------- ``SystemOne.load(path, max_len=4096)`` overrides the checkpoint's context length (ModernBERT handles up to 8192 positions; checkpoints trained at 512 have not seen longer inputs). States longer than the state budget are handled per call by ``long_state`` (default set at load): * ``"truncate"`` (default): cut in the middle, keeping ~25% head / 75% tail (serialize.py). * ``"chunk"``: split into overlapping windows (``chunk_overlap`` tokens, at most ``max_chunks`` windows), score every (window, question) pair in the same batched forward pass, and pool the per-window *calibrated* distributions per question with ``chunk_agg`` (see ``CHUNK_AGG``): - ``"mean"`` log-linear pool: p_k ∝ exp(mean_c log p_c,k) (geometric mean, renormalized; = softmax of the mean window logits). Flat windows do not change which answer wins, but evidence found in only one of n windows is diluted (≈ its log-odds / n): conservative, best when the answer depends on the document as a whole (topic, tone). - ``"max"`` noul: P(true) = max_c P_c(true) ("true if any window shows it", a noisy-OR that does not compound across windows). choice/score: p_k ∝ max_c p_c,k. - ``"noisy_or"`` noul only: P(true) = 1 - prod_c (1 - P_c(true)) (compounds; overconfident with many windows). - ``"linear"`` arithmetic mean of the window distributions. ``chunk_agg="auto"`` = ``{"choice": "mean", "score": "mean", "noul": "max"}``; a dict overrides per type. ``escalate_prob`` of a pooled answer is that of the window giving the pooled answer the highest probability (the deciding window). Results then carry an extra top-level ``"chunks": {question: n_windows}``. """ from __future__ import annotations import json import math import os import time from typing import Mapping, Sequence import torch from transformers import AutoTokenizer from jevlike.calibrate import temperature from jevlike.model import DecisionModel from jevlike.serialize import Serializer, TokenBudgetBatchSampler, collate from jevlike.types import PublicQuestion, Question Questions = Mapping[str, "PublicQuestion | Question"] LONG_STATE = ("truncate", "chunk") CHUNK_AGG = {"choice": ("mean", "max", "linear"), "score": ("mean", "max", "linear"), "noul": ("max", "mean", "noisy_or", "linear")} # first = "auto" def resolve_chunk_agg(agg: "str | Mapping[str, str]") -> dict[str, str]: """``"auto"`` / one rule for all types / {type: rule} -> {type: rule}, validated.""" rules = {t: opts[0] for t, opts in CHUNK_AGG.items()} if isinstance(agg, str): if agg != "auto": rules = {t: agg for t in CHUNK_AGG} if agg == "noisy_or": # noul-only rule: keep the defaults for choice/score rules.update({"choice": "mean", "score": "mean"}) else: rules.update(agg) for t, r in rules.items(): if t not in CHUNK_AGG or r not in CHUNK_AGG[t]: raise ValueError(f"chunk_agg: {r!r} is not a rule for {t!r} (options: {CHUNK_AGG.get(t)})") return rules def pool_chunks(logp: torch.Tensor, qtype: str, rule: str) -> torch.Tensor: """Pool per-window log-probabilities [C, K] of one question into a distribution [K].""" if logp.shape[0] == 1: return logp[0].exp() if rule == "mean": return logp.mean(0).softmax(-1) if rule == "linear": return logp.exp().mean(0) if qtype == "noul" and rule in ("max", "noisy_or"): p_true = logp[:, 1].exp().max() if rule == "max" else -torch.expm1(logp[:, 0].sum()) return torch.stack([1 - p_true, p_true]) if rule == "max": return logp.max(0).values.softmax(-1) raise ValueError(f"unknown chunk aggregation rule {rule!r} for {qtype}") class SystemOne: def __init__(self, model: DecisionModel, tokenizer, calibration: dict | None, device: torch.device, max_tokens: int = 16384, long_state: str = "truncate", chunk_agg: "str | dict" = "auto", chunk_overlap: int = 128, max_chunks: int = 64): self.model, self.cal, self.device = model.eval(), calibration, device cfg = model.cfg self.max_tokens = max(max_tokens, cfg.max_len) # a micro-batch always fits one full-length sequence self.ser = Serializer(tokenizer, cfg.max_len, cfg.head_max_len) self.name = cfg.name self.threshold = cfg.escalate_threshold if long_state not in LONG_STATE: raise ValueError(f"long_state must be one of {LONG_STATE}, got {long_state!r}") self.long_state, self.chunk_agg = long_state, resolve_chunk_agg(chunk_agg) self.chunk_overlap, self.max_chunks = chunk_overlap, max_chunks @property def max_len(self) -> int: return self.ser.max_len @classmethod def load(cls, path: str, device: str | torch.device | None = None, max_len: int | None = None, **kw) -> "SystemOne": """``max_len`` overrides the checkpoint's context length (default: keep it; at most the backbone's ``max_position_embeddings``, 8192 for ModernBERT). Other keyword arguments (``max_tokens``, ``long_state``, ``chunk_agg``, ``chunk_overlap``, ``max_chunks``) go to the constructor.""" device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu")) cal_path = os.path.join(path, "calibration.json") cal = json.load(open(cal_path)) if os.path.exists(cal_path) else None model = DecisionModel.from_pretrained(path, device) if max_len: set_max_len(model.cfg, max_len) return cls(model, AutoTokenizer.from_pretrained(path), cal, device, **kw) def predict(self, state: str, questions: Questions, calibrated: bool = True, long_state: str | None = None, chunk_agg: "str | dict | None" = None) -> dict: return self.predict_batch([(state, questions)], calibrated=calibrated, long_state=long_state, chunk_agg=chunk_agg)[0] @torch.inference_mode() def predict_batch(self, items: Sequence[tuple[str, Questions]], calibrated: bool = True, long_state: str | None = None, chunk_agg: "str | dict | None" = None) -> list[dict]: """``latency_ms`` of each result is the wall time of the whole batched call. ``long_state`` / ``chunk_agg`` default to the values given at load (see module doc).""" t0 = time.perf_counter() long_state = long_state or self.long_state if long_state not in LONG_STATE: raise ValueError(f"long_state must be one of {LONG_STATE}, got {long_state!r}") rules = self.chunk_agg if chunk_agg is None else resolve_chunk_agg(chunk_agg) flat: list[tuple[int, str, Question]] = [] for i, (_, qs) in enumerate(items): for name, q in qs.items(): q = q if isinstance(q, Question) else Question.from_public(q) try: q.validate() except AssertionError as e: raise ValueError(f"question {name!r}: {e}") from None flat.append((i, name, q)) pairs = [(items[i][0], q) for i, _, q in flat] if long_state == "chunk": groups = self.ser.encode_chunked(pairs, self.chunk_overlap, self.max_chunks) else: groups = [[e] for e in self.ser.encode_many(pairs)] encs = [e for g in groups for e in g] logits, esc = self._forward(encs) cal = self.cal if calibrated else None if cal: esc = cal["escalate"]["a"] * esc + cal["escalate"]["b"] p_wrong = 1 - torch.sigmoid(esc) results = [{"answers": {}, "model": self.name, "latency_ms": 0.0} for _ in items] if long_state == "chunk": for r in results: r["chunks"] = {} start = 0 for (i, name, q), g in zip(flat, groups): sl = slice(start, start + len(g)) start += len(g) logp = torch.stack(logits[sl]) / temperature(cal, q.type, q.n_options) logp = logp.log_softmax(-1) p = pool_chunks(logp, q.type, rules[q.type]) deciding = int(logp[:, int(p.argmax())].argmax()) # window giving the answer the most mass results[i]["answers"][name] = self._answer(q, p, float(p_wrong[sl][deciding])) if long_state == "chunk": results[i]["chunks"][name] = len(g) ms = (time.perf_counter() - t0) * 1000 for r in results: r["latency_ms"] = ms return results def batches(self, lengths: Sequence[int]) -> list[list[int]]: """Length-sorted micro-batches with B * T_max <= max_tokens (a longer single sequence goes alone).""" return TokenBudgetBatchSampler(lengths, self.max_tokens, max_batch=256, shuffle=False).batches() def _forward(self, encs: list) -> tuple[list[torch.Tensor], torch.Tensor]: """Raw per-sequence option logits and escalate logits, token-budgeted micro-batches.""" logits: list[torch.Tensor] = [None] * len(encs) # type: ignore[list-item] esc = torch.zeros(len(encs)) for idx in self.batches([len(e) for e in encs]) if encs else []: b = collate([encs[j] for j in idx], None, self.ser.pad) k_max = int(b["n_options"].max()) b = {k: v.to(self.device, non_blocking=True) for k, v in b.items()} with torch.autocast(self.device.type, dtype=torch.bfloat16, enabled=self.device.type == "cuda"): out = self.model(**b, k_max=k_max) lg, e = out["logits"].float().cpu(), out["escalate_logit"].float().cpu() for r, j in enumerate(idx): logits[j] = lg[r, :len(encs[j].markers)] esc[idx] = e return logits, esc def _answer(self, q: Question, p: torch.Tensor, p_wrong: float) -> dict: probs = p.tolist() top = int(p.argmax()) esc = {"escalate": p_wrong > self.threshold, "escalate_prob": p_wrong} if q.type == "choice": return {"type": "choice", "choice": q.options[top].key, "probabilities": {o.key: pi for o, pi in zip(q.options, probs)}, "confidence": probs[top], **esc} if q.type == "score": k = len(probs) expected = sum(i * pi for i, pi in enumerate(probs)) return {"type": "score", "level": top, "label": q.options[top].text or q.options[top].key, "score": 10.0 * expected / (k - 1), "probabilities": probs, "confidence": probs[top], **esc} return {"type": "noul", "noul": probs[1], "confidence": max(probs), **esc} def set_max_len(cfg, max_len: int) -> None: """Set a DecisionConfig's context length (validated against the backbone's position limit); shrinks ``head_max_len`` to half of it if it no longer fits.""" limit = int(cfg.backbone_config.get("max_position_embeddings") or 8192) if not 64 <= max_len <= limit: raise ValueError(f"max_len must be in [64, {limit}] (backbone max_position_embeddings), got {max_len}") cfg.max_len = int(max_len) if cfg.head_max_len >= cfg.max_len: cfg.head_max_len = cfg.max_len // 2 def _demo(path: str) -> None: from jevlike.types import Choice, Noul, Score m = SystemOne.load(path) res = m.predict("Hi, I was charged twice this month and the app keeps crashing. Fix this or I'm leaving.", {"dept": Choice("Which team should handle this?", {"billing": "payments, invoices, refunds", "technical": "bugs, crashes, errors"}), "urgency": Score("How urgent is this?", ["not urgent", "soon", "blocking"]), "churn": Noul("The user threatens to cancel")}) print(json.dumps(res, indent=1)) assert all(math.isfinite(a["confidence"]) for a in res["answers"].values()) if __name__ == "__main__": import sys _demo(sys.argv[1] if len(sys.argv) > 1 else "checkpoints/jevlike-base")