Download runtime/jevlike/predict.py from Cem13/kodama-core: direct link, hf CLI and curl.
- Browser
- Download file 12.9 kB
-
https://huggingface.co/Cem13/kodama-core/resolve/main/runtime/jevlike/predict.py
- Command line
-
hf download hf://Cem13/kodama-core/runtime/jevlike/predict.py
-
curl -L -o predict.py https://huggingface.co/Cem13/kodama-core/resolve/main/runtime/jevlike/predict.py
12.9 kB
| """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 | |
| def max_len(self) -> int: | |
| return self.ser.max_len | |
| 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] | |
| 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") | |