kodama-core / runtime /jevlike /predict.py
Cem13's picture
Release Kodama Core with weights, runtime, attribution, and evaluation
c7893fa verified
Raw History Blame Contribute Delete
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
@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")