anyjev 0.0.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
anyjev/__init__.py ADDED
@@ -0,0 +1,14 @@
1
+ """AnyJev: turn any causal LLM into a Jev-style decision model.
2
+
3
+ Typed decisions (choice / score / noul) with probabilities, read from the
4
+ model's next-token logits in one prefill, no generation. Training-free
5
+ debiasing (L0) is on by default; post-hoc calibration (L1) when you have labels.
6
+
7
+ Not affiliated with, endorsed by, or derived from TypeSafe AI or Jev.
8
+ """
9
+ from anyjev.decider import Decider
10
+ from anyjev.question import Question
11
+ from anyjev.result import Decision, DecisionSet
12
+
13
+ __all__ = ["Question", "Decision", "DecisionSet", "Decider"]
14
+ __version__ = "0.0.1"
@@ -0,0 +1,3 @@
1
+ from anyjev.backends.base import Backend
2
+
3
+ __all__ = ["Backend"]
@@ -0,0 +1,21 @@
1
+ """Backend protocol. A backend does exactly one thing: given rendered prompts,
2
+ return the next-token log-probabilities of the requested token ids.
3
+ Everything else (debiasing, calibration, abstention) lives above it."""
4
+ from __future__ import annotations
5
+
6
+ from typing import Any, List, Protocol, Sequence
7
+
8
+ import numpy as np
9
+
10
+
11
+ class Backend(Protocol):
12
+ tokenizer: Any # must provide .encode(text, add_special_tokens=False) and optionally .chat_template
13
+ name: str
14
+
15
+ def next_token_logprobs(self, prompts: Sequence[str],
16
+ token_ids: Sequence[Sequence[int]]) -> List[np.ndarray]:
17
+ """For prompt i, return log p(token | prompt_i) for each id in token_ids[i],
18
+ taken from the full-vocabulary log-softmax at the last position.
19
+ Decision mode never samples; a backend that cannot expose restricted
20
+ next-token logits is not a supported backend."""
21
+ ...
@@ -0,0 +1,94 @@
1
+ """A synthetic backend with known, injectable biases. Used by the unit tests
2
+ and by the docs to show what L0 removes. It parses the prompts that
3
+ anyjev.readout builds, so it exercises the real prompt path."""
4
+ from __future__ import annotations
5
+
6
+ import re
7
+ from typing import Callable, Dict, List, Optional, Sequence
8
+
9
+ import numpy as np
10
+
11
+ from anyjev.calibrate.contextual import DEFAULT_PROBES
12
+
13
+ _OPT_LINE = re.compile(r"^([A-Z]|\d+)\. (.*)$")
14
+ _NOUL_LINE = re.compile(r"^Answer (Yes|No) or (Yes|No)\.$")
15
+
16
+
17
+ class FakeTokenizer:
18
+ chat_template = None
19
+
20
+ def __init__(self):
21
+ self._vocab: Dict[str, int] = {}
22
+
23
+ def encode(self, text: str, add_special_tokens: bool = False) -> List[int]:
24
+ # every label we care about is one token; anything else is "long"
25
+ if re.fullmatch(r" ?([A-Z]|\d|Yes|No)", text):
26
+ key = text.strip()
27
+ if key not in self._vocab:
28
+ self._vocab[key] = 1000 + len(self._vocab)
29
+ return [self._vocab[key]]
30
+ return [1, 2]
31
+
32
+ def id_to_label(self, tid: int) -> str:
33
+ for k, v in self._vocab.items():
34
+ if v == tid:
35
+ return k
36
+ raise KeyError(tid)
37
+
38
+
39
+ class FakeBackend:
40
+ """logit(position j, option o) = content(state, o) + position_bias[j] + label_prior[label_j]
41
+
42
+ content is 0 on content-free probes, so the prior is exactly the bias term.
43
+ """
44
+
45
+ def __init__(self, content: Callable[[str, str], float],
46
+ position_bias: Optional[Sequence[float]] = None,
47
+ label_prior: Optional[Dict[str, float]] = None,
48
+ temperature: float = 1.0):
49
+ self.name = "fake"
50
+ self.tokenizer = FakeTokenizer()
51
+ self.content = content
52
+ self.position_bias = list(position_bias or [])
53
+ self.label_prior = dict(label_prior or {})
54
+ self.temperature = temperature
55
+ self.calls = 0
56
+ self.prompts_seen = 0
57
+
58
+ def _parse(self, prompt: str):
59
+ state = prompt.split("State:\n", 1)[1].split("\n\nQuestion:", 1)[0]
60
+ if state == "(empty)":
61
+ state = ""
62
+ labels, options = [], []
63
+ for line in prompt.splitlines():
64
+ m = _OPT_LINE.match(line)
65
+ if m:
66
+ labels.append(m.group(1))
67
+ options.append(m.group(2))
68
+ m = _NOUL_LINE.match(line)
69
+ if m:
70
+ labels = [m.group(1), m.group(2)]
71
+ options = list(labels)
72
+ return state, labels, options
73
+
74
+ def next_token_logprobs(self, prompts: Sequence[str],
75
+ token_ids: Sequence[Sequence[int]]) -> List[np.ndarray]:
76
+ self.calls += 1
77
+ self.prompts_seen += len(prompts)
78
+ out = []
79
+ for prompt, ids in zip(prompts, token_ids):
80
+ state, labels, options = self._parse(prompt)
81
+ is_probe = state in DEFAULT_PROBES
82
+ logits = np.zeros(len(ids))
83
+ by_label = {self.tokenizer.id_to_label(t): k for k, t in enumerate(ids)}
84
+ for j, (lab, opt) in enumerate(zip(labels, options)):
85
+ z = 0.0 if is_probe else self.content(state, opt)
86
+ if j < len(self.position_bias):
87
+ z += self.position_bias[j]
88
+ z += self.label_prior.get(lab, 0.0)
89
+ logits[by_label[lab]] = z / self.temperature
90
+ # full-vocab log-softmax: pretend a bit of mass lives elsewhere
91
+ z = np.concatenate([logits, [-5.0]])
92
+ lp = z - np.log(np.exp(z - z.max()).sum()) - z.max()
93
+ out.append(lp[:-1])
94
+ return out
anyjev/backends/hf.py ADDED
@@ -0,0 +1,52 @@
1
+ """transformers backend: single forward per prompt batch, logits at the last position."""
2
+ from __future__ import annotations
3
+
4
+ from typing import List, Optional, Sequence
5
+
6
+ import numpy as np
7
+
8
+
9
+ class HFBackend:
10
+ def __init__(self, model_name: str, device: str = "cuda", dtype: str = "bfloat16",
11
+ batch_size: int = 16, trust_remote_code: bool = False, revision: Optional[str] = None):
12
+ import torch
13
+ from transformers import AutoModelForCausalLM, AutoTokenizer
14
+
15
+ self.name = model_name
16
+ self.batch_size = batch_size
17
+ self.device = device
18
+ self.revision = revision
19
+ self.tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=trust_remote_code,
20
+ revision=revision)
21
+ self.tokenizer.padding_side = "left"
22
+ if self.tokenizer.pad_token is None:
23
+ self.tokenizer.pad_token = self.tokenizer.eos_token
24
+ torch_dtype = getattr(torch, dtype) if dtype != "auto" else "auto"
25
+ self.model = AutoModelForCausalLM.from_pretrained(
26
+ model_name, torch_dtype=torch_dtype, device_map=device,
27
+ trust_remote_code=trust_remote_code, revision=revision)
28
+ self.model.eval()
29
+
30
+ def next_token_logprobs(self, prompts: Sequence[str],
31
+ token_ids: Sequence[Sequence[int]]) -> List[np.ndarray]:
32
+ import torch
33
+
34
+ n = len(prompts)
35
+ # sort by length to reduce padding, restore order at the end
36
+ lengths = [len(self.tokenizer.encode(p, add_special_tokens=False)) for p in prompts]
37
+ order = sorted(range(n), key=lambda i: lengths[i])
38
+ out: List[Optional[np.ndarray]] = [None] * n
39
+ for start in range(0, n, self.batch_size):
40
+ idx = order[start:start + self.batch_size]
41
+ enc = self.tokenizer([prompts[i] for i in idx], return_tensors="pt",
42
+ padding=True, add_special_tokens=False)
43
+ enc = {k: v.to(self.model.device) for k, v in enc.items()}
44
+ # explicit position ids so left padding does not shift positions
45
+ pos = (enc["attention_mask"].cumsum(-1) - 1).clamp(min=0)
46
+ with torch.no_grad():
47
+ logits = self.model(**enc, position_ids=pos).logits[:, -1, :].float()
48
+ lp = torch.log_softmax(logits, dim=-1)
49
+ for row, i in enumerate(idx):
50
+ ids = torch.as_tensor(list(token_ids[i]), device=lp.device)
51
+ out[i] = lp[row, ids].cpu().numpy().astype(np.float64)
52
+ return out # type: ignore[return-value]
@@ -0,0 +1,58 @@
1
+ """vLLM backend over the OpenAI-compatible server.
2
+
3
+ One request per prompt with `max_tokens=1`, `allowed_token_ids` restricted to
4
+ the label tokens and `logprobs=K`. vLLM reports logprobs after its logit
5
+ processors, so the K entries are exactly the labels, normalized over them.
6
+ Prefix caching on the server makes the K permutations of one state cheap.
7
+
8
+ vllm serve Qwen/Qwen3-8B --enable-prefix-caching
9
+ Decider(VLLMBackend("http://localhost:8000", "Qwen/Qwen3-8B"))
10
+ """
11
+ from __future__ import annotations
12
+
13
+ import concurrent.futures as cf
14
+ import json
15
+ import urllib.request
16
+ from typing import List, Sequence
17
+
18
+ import numpy as np
19
+
20
+
21
+ class VLLMBackend:
22
+ def __init__(self, base_url: str, model: str, tokenizer_name: str | None = None,
23
+ api_key: str = "EMPTY", workers: int = 16, timeout: float = 120.0):
24
+ from transformers import AutoTokenizer
25
+
26
+ self.base_url = base_url.rstrip("/")
27
+ self.name = model
28
+ self.api_key = api_key
29
+ self.workers = workers
30
+ self.timeout = timeout
31
+ self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_name or model)
32
+
33
+ def _one(self, prompt: str, ids: Sequence[int]) -> np.ndarray:
34
+ body = {
35
+ "model": self.name, "prompt": prompt, "max_tokens": 1, "temperature": 0.0,
36
+ "logprobs": len(ids), "allowed_token_ids": list(ids),
37
+ }
38
+ req = urllib.request.Request(
39
+ self.base_url + "/v1/completions", data=json.dumps(body).encode(),
40
+ headers={"Content-Type": "application/json", "Authorization": f"Bearer {self.api_key}"})
41
+ with urllib.request.urlopen(req, timeout=self.timeout) as r:
42
+ out = json.load(r)
43
+ top = out["choices"][0]["logprobs"]["top_logprobs"][0] # {token_str: logprob}
44
+ # map back by token id: decode each id the same way the server renders it
45
+ by_id = {}
46
+ for tid in ids:
47
+ tok = self.tokenizer.decode([tid])
48
+ conv = self.tokenizer.convert_ids_to_tokens(tid)
49
+ for key in (tok, conv):
50
+ if key in top:
51
+ by_id[tid] = float(top[key])
52
+ break
53
+ return np.array([by_id.get(tid, -30.0) for tid in ids], dtype=np.float64)
54
+
55
+ def next_token_logprobs(self, prompts: Sequence[str],
56
+ token_ids: Sequence[Sequence[int]]) -> List[np.ndarray]:
57
+ with cf.ThreadPoolExecutor(self.workers) as ex:
58
+ return list(ex.map(self._one, prompts, token_ids))
@@ -0,0 +1,6 @@
1
+ from anyjev.calibrate.contextual import apply_contextual, batch_prior, content_free_prior
2
+ from anyjev.calibrate.permute import cyclic_shifts, flip_rate_across_perms, marginalize
3
+ from anyjev.calibrate.posthoc import TemperatureScaler
4
+
5
+ __all__ = ["content_free_prior", "batch_prior", "apply_contextual", "cyclic_shifts", "marginalize",
6
+ "flip_rate_across_perms", "TemperatureScaler"]
@@ -0,0 +1,47 @@
1
+ """L0 prior estimation and correction.
2
+
3
+ Two ways to estimate the model's prior over the labels, both label-free:
4
+
5
+ batch (Zhou et al., ICLR 2024, default): the mean predicted distribution
6
+ over a batch of real inputs, kept running per question across calls.
7
+ Low variance in the bench: +1 to +2 accuracy points and a large ECE
8
+ improvement on every model; needs a handful of inputs before it kicks
9
+ in (min_prior_n), and assumes the label marginal is not extreme.
10
+ content_free (Zhao et al., ICML 2021, opt-in): the distribution on
11
+ content-free inputs ("N/A", "", "[MASK]"). Depends only on the question,
12
+ cached per question. High variance in the bench: +8 to +12 points on a
13
+ prompt-injection noul, but -3 on ordinal scores and -9 on one model's
14
+ noul questions, because for some questions the model's answer to an
15
+ empty input is an honest answer, not a label prior.
16
+
17
+ Correction is the same in both cases: divide by the prior and renormalize.
18
+ """
19
+ from __future__ import annotations
20
+
21
+ import numpy as np
22
+
23
+ DEFAULT_PROBES = ("N/A", "", "[MASK]")
24
+ EPS = 1e-8
25
+
26
+
27
+ def content_free_prior(cf_probs: np.ndarray) -> np.ndarray:
28
+ """cf_probs: [..., C, K] distributions over K positions for C probes.
29
+ Returns the mean prior [..., K]."""
30
+ cf = np.asarray(cf_probs, dtype=np.float64)
31
+ prior = cf.mean(axis=-2)
32
+ prior = np.clip(prior, EPS, None)
33
+ return prior / prior.sum(axis=-1, keepdims=True)
34
+
35
+
36
+ def batch_prior(p_batch: np.ndarray) -> np.ndarray:
37
+ """p_batch: [N, ..., K] distributions over real inputs. Returns [..., K]."""
38
+ p = np.asarray(p_batch, dtype=np.float64).mean(axis=0)
39
+ p = np.clip(p, EPS, None)
40
+ return p / p.sum(axis=-1, keepdims=True)
41
+
42
+
43
+ def apply_contextual(p_raw: np.ndarray, prior: np.ndarray) -> np.ndarray:
44
+ """p_raw, prior: [..., K]. Returns normalize(p_raw / prior)."""
45
+ p = np.asarray(p_raw, dtype=np.float64) / np.clip(np.asarray(prior, dtype=np.float64), EPS, None)
46
+ p = np.clip(p, EPS, None)
47
+ return p / p.sum(axis=-1, keepdims=True)
@@ -0,0 +1,58 @@
1
+ """L0: permutation marginalization for position bias (Zheng et al., ICLR 2024).
2
+
3
+ Show the options in K cyclic shifts so every option occupies every position
4
+ once, then combine the per-option probabilities across shifts.
5
+
6
+ Two ways to combine:
7
+ "logmean": average log-probabilities (geometric mean), then renormalize.
8
+ If the bias is additive in logit space, logit(i at pos j) = c_i + b_j,
9
+ the per-option average is c_i + mean(b) - mean(log Z_s), so the
10
+ result is softmax(c) exactly: the position bias is removed and the
11
+ result is invariant to how the options were originally listed.
12
+ "mean": average probabilities, the form used in the paper. Only invariant
13
+ to cyclic rotations of the original list; kept for comparison.
14
+ """
15
+ from __future__ import annotations
16
+
17
+ from typing import List, Optional, Sequence
18
+
19
+ import numpy as np
20
+
21
+ EPS = 1e-12
22
+
23
+
24
+ def cyclic_shifts(k: int, max_permutations: Optional[int] = None) -> List[List[int]]:
25
+ """perm[j] = original option index shown at position j."""
26
+ n = k if max_permutations is None else max(1, min(k, max_permutations))
27
+ return [[(j + s) % k for j in range(k)] for s in range(n)]
28
+
29
+
30
+ def marginalize(p_by_perm: np.ndarray, perms: Sequence[Sequence[int]],
31
+ combine: str = "logmean") -> np.ndarray:
32
+ """p_by_perm: [P, K] distributions indexed by *position*.
33
+ Returns [K] distribution indexed by original option."""
34
+ p_by_perm = np.asarray(p_by_perm, dtype=np.float64)
35
+ P, K = p_by_perm.shape
36
+ per_option = np.zeros((P, K))
37
+ for s, perm in enumerate(perms):
38
+ for j, i in enumerate(perm):
39
+ per_option[s, i] = p_by_perm[s, j]
40
+ if combine == "logmean":
41
+ z = np.log(np.clip(per_option, EPS, None)).mean(axis=0)
42
+ z = z - z.max()
43
+ out = np.exp(z)
44
+ elif combine == "mean":
45
+ out = per_option.mean(axis=0)
46
+ else:
47
+ raise ValueError("combine must be 'logmean' or 'mean'")
48
+ return out / out.sum()
49
+
50
+
51
+ def flip_rate_across_perms(p_by_perm: np.ndarray, perms: Sequence[Sequence[int]]) -> float:
52
+ """Fraction of permutations whose argmax (in option space) disagrees with
53
+ the first permutation's argmax. 0.0 means order-invariant on this item."""
54
+ p_by_perm = np.asarray(p_by_perm)
55
+ winners = [perm[int(np.argmax(p_by_perm[s]))] for s, perm in enumerate(perms)]
56
+ if len(winners) <= 1:
57
+ return 0.0
58
+ return float(np.mean([w != winners[0] for w in winners[1:]]))
@@ -0,0 +1,60 @@
1
+ """L1: temperature scaling (Guo et al., ICML 2017) fit on a labeled set.
2
+
3
+ Artifacts are small dicts keyed by (model, question key); reusing one across
4
+ models is a user error and is not detected here.
5
+ """
6
+ from __future__ import annotations
7
+
8
+ from dataclasses import dataclass
9
+ from typing import Dict, Sequence
10
+
11
+ import numpy as np
12
+
13
+ EPS = 1e-12
14
+
15
+
16
+ def _log_softmax(z: np.ndarray) -> np.ndarray:
17
+ z = z - z.max(axis=-1, keepdims=True)
18
+ return z - np.log(np.exp(z).sum(axis=-1, keepdims=True))
19
+
20
+
21
+ def nll(logits: np.ndarray, labels: Sequence[int]) -> float:
22
+ lp = _log_softmax(np.asarray(logits, dtype=np.float64))
23
+ return float(-lp[np.arange(len(labels)), np.asarray(labels)].mean())
24
+
25
+
26
+ @dataclass
27
+ class TemperatureScaler:
28
+ temperature: float = 1.0
29
+
30
+ @classmethod
31
+ def fit(cls, probs: np.ndarray, labels: Sequence[int],
32
+ log_t_range=(-3.0, 3.0), iters: int = 60) -> "TemperatureScaler":
33
+ """Golden-section search on log T minimizing NLL. probs: [N, K]."""
34
+ logits = np.log(np.clip(np.asarray(probs, dtype=np.float64), EPS, None))
35
+ lo, hi = log_t_range
36
+ phi = (np.sqrt(5) - 1) / 2
37
+ a, b = lo, hi
38
+ c, d = b - phi * (b - a), a + phi * (b - a)
39
+ fc, fd = nll(logits / np.exp(c), labels), nll(logits / np.exp(d), labels)
40
+ for _ in range(iters):
41
+ if fc < fd:
42
+ b, d, fd = d, c, fc
43
+ c = b - phi * (b - a)
44
+ fc = nll(logits / np.exp(c), labels)
45
+ else:
46
+ a, c, fc = c, d, fd
47
+ d = a + phi * (b - a)
48
+ fd = nll(logits / np.exp(d), labels)
49
+ return cls(temperature=float(np.exp((a + b) / 2)))
50
+
51
+ def apply(self, probs: np.ndarray) -> np.ndarray:
52
+ logits = np.log(np.clip(np.asarray(probs, dtype=np.float64), EPS, None))
53
+ return np.exp(_log_softmax(logits / self.temperature))
54
+
55
+ def to_dict(self) -> Dict[str, float]:
56
+ return {"method": "temperature", "temperature": self.temperature}
57
+
58
+ @classmethod
59
+ def from_dict(cls, d: Dict) -> "TemperatureScaler":
60
+ return cls(temperature=float(d["temperature"]))
anyjev/decider.py ADDED
@@ -0,0 +1,236 @@
1
+ """Decider: the pipeline from (state, questions) to leveled decisions.
2
+
3
+ raw : one prompt, one permutation, restricted softmax. What every clone does.
4
+ L0 : permutation marginalization + label-free prior correction (batch mean
5
+ by default, content-free probes optional). Zero labels.
6
+ L1 : L0, then a post-hoc calibrator fit on a labeled set for this question.
7
+ """
8
+ from __future__ import annotations
9
+
10
+ from typing import Any, Dict, List, Optional, Sequence, Tuple
11
+
12
+ import numpy as np
13
+
14
+ from anyjev.calibrate.contextual import (
15
+ DEFAULT_PROBES,
16
+ apply_contextual,
17
+ batch_prior,
18
+ content_free_prior,
19
+ )
20
+ from anyjev.calibrate.permute import cyclic_shifts, flip_rate_across_perms, marginalize
21
+ from anyjev.calibrate.posthoc import TemperatureScaler
22
+ from anyjev.question import Question
23
+ from anyjev.readout import (
24
+ DEFAULT_SYSTEM,
25
+ answer_labels,
26
+ build_prompt,
27
+ label_ids_for_perm,
28
+ map_label_tokens,
29
+ render_chat,
30
+ )
31
+ from anyjev.result import Decision, DecisionSet
32
+ from anyjev.state import render_state
33
+
34
+ LEVELS = ("raw", "L0", "L1")
35
+ PRIORS = ("batch", "content_free", "none")
36
+
37
+
38
+ def _softmax(lp: np.ndarray) -> np.ndarray:
39
+ z = np.asarray(lp, dtype=np.float64)
40
+ z = z - z.max()
41
+ p = np.exp(z)
42
+ return p / p.sum()
43
+
44
+
45
+ class Decider:
46
+ def __init__(self, backend, *, level: str = "L0", prior: str = "batch", min_prior_n: int = 8,
47
+ max_permutations: Optional[int] = None, combine: str = "logmean",
48
+ cf_probes: Sequence[str] = DEFAULT_PROBES, record_content_free: bool = False,
49
+ system: str = DEFAULT_SYSTEM):
50
+ if level not in LEVELS:
51
+ raise ValueError(f"level must be one of {LEVELS}")
52
+ if prior not in PRIORS:
53
+ raise ValueError(f"prior must be one of {PRIORS}")
54
+ self.backend = backend
55
+ self.level = level
56
+ self.prior = prior
57
+ self.min_prior_n = min_prior_n
58
+ self.max_permutations = max_permutations
59
+ self.combine = combine
60
+ self.cf_probes = tuple(cf_probes)
61
+ self.record_content_free = record_content_free
62
+ self.system = system
63
+ self._label_ids: Dict[tuple, List[int]] = {}
64
+ self._artifacts: Dict[str, TemperatureScaler] = {}
65
+ self._running: Dict[str, Tuple[np.ndarray, int]] = {} # q.key -> (sum p_pos_raw [P,K], n)
66
+ self._cf_cache: Dict[str, np.ndarray] = {} # q.key -> cf prior [P,K]
67
+
68
+ # ---- public -------------------------------------------------------
69
+ def decide(self, state: Any, questions: Sequence[Question], level: Optional[str] = None) -> DecisionSet:
70
+ level = level or self.level
71
+ decs = self._run([state], list(questions), level)
72
+ return DecisionSet([decs[(0, qi)] for qi in range(len(questions))], level)
73
+
74
+ def decide_batch(self, states: Sequence[Any], question: Question,
75
+ level: Optional[str] = None) -> List[Decision]:
76
+ """Many states, one question. The bench path, and the best path for
77
+ batch prior estimation."""
78
+ level = level or self.level
79
+ decs = self._run(list(states), [question], level)
80
+ return [decs[(si, 0)] for si in range(len(states))]
81
+
82
+ def calibrate(self, question: Question, states: Sequence[Any], labels: Sequence[int],
83
+ level: str = "L1") -> Dict[str, Any]:
84
+ """Fit an L1 artifact for this question on labeled states. labels are
85
+ option indices. Returns the artifact dict (store it; it is per model)."""
86
+ if level != "L1":
87
+ raise ValueError("only L1 calibration is implemented")
88
+ decs = self.decide_batch(states, question, level="L0")
89
+ probs = np.stack([d.probs for d in decs])
90
+ scaler = TemperatureScaler.fit(probs, labels)
91
+ self._artifacts[question.key] = scaler
92
+ return {"model": self.backend.name, "question": question.key, **scaler.to_dict()}
93
+
94
+ def load_artifact(self, question: Question, artifact: Dict[str, Any]) -> None:
95
+ if artifact.get("model") not in (None, self.backend.name):
96
+ raise ValueError(f"artifact was fit on {artifact['model']}, backend is {self.backend.name}")
97
+ self._artifacts[question.key] = TemperatureScaler.from_dict(artifact)
98
+
99
+ def export_artifacts(self) -> Dict[str, Any]:
100
+ """Every L1 artifact this decider holds, keyed by question hash. JSON-serializable."""
101
+ return {"model": self.backend.name,
102
+ "artifacts": {k: {"model": self.backend.name, "question": k, **v.to_dict()}
103
+ for k, v in self._artifacts.items()}}
104
+
105
+ def save_artifacts(self, path: str) -> None:
106
+ import json
107
+ with open(path, "w") as f:
108
+ json.dump(self.export_artifacts(), f, indent=1)
109
+
110
+ def load_artifacts(self, path_or_dict) -> int:
111
+ """Load artifacts saved by save_artifacts. Refuses artifacts fit on another model."""
112
+ import json
113
+ d = path_or_dict if isinstance(path_or_dict, dict) else json.load(open(path_or_dict))
114
+ if d.get("model") not in (None, self.backend.name):
115
+ raise ValueError(f"artifacts were fit on {d['model']}, backend is {self.backend.name}")
116
+ for key, art in d["artifacts"].items():
117
+ self._artifacts[key] = TemperatureScaler.from_dict(art)
118
+ return len(d["artifacts"])
119
+
120
+ def running_prior(self, question: Question) -> Optional[np.ndarray]:
121
+ """The batch prior accumulated so far for this question, [P, K] by position, or None."""
122
+ entry = self._running.get(question.key)
123
+ if entry is None or entry[1] < self.min_prior_n:
124
+ return None
125
+ return batch_prior(entry[0][None] / entry[1])
126
+
127
+ # ---- internals ----------------------------------------------------
128
+ def _ids_for(self, q: Question) -> List[int]:
129
+ key = (q.kind, q.k)
130
+ if key not in self._label_ids:
131
+ self._label_ids[key] = map_label_tokens(self.backend.tokenizer, answer_labels(q))
132
+ return self._label_ids[key]
133
+
134
+ def _perms(self, q: Question, level: str) -> List[List[int]]:
135
+ if level == "raw" or q.ordered:
136
+ return [list(range(q.k))]
137
+ if q.kind == "noul":
138
+ return [[0, 1], [1, 0]]
139
+ return cyclic_shifts(q.k, self.max_permutations)
140
+
141
+ def _run(self, states: List[Any], questions: List[Question], level: str) -> Dict[tuple, Decision]:
142
+ if level not in LEVELS:
143
+ raise ValueError(f"level must be one of {LEVELS}")
144
+ tok = self.backend.tokenizer
145
+ state_texts = [render_state(s) for s in states]
146
+ want_cf = level != "raw" and (self.prior == "content_free" or self.record_content_free)
147
+
148
+ # 1. collect every prompt once (content-free probes are shared across states)
149
+ prompt_index: Dict[str, int] = {}
150
+ prompt_ids: List[List[int]] = []
151
+
152
+ def add(text: str, ids: List[int]) -> int:
153
+ if text not in prompt_index:
154
+ prompt_index[text] = len(prompt_ids)
155
+ prompt_ids.append(ids)
156
+ return prompt_index[text]
157
+
158
+ plan = {}
159
+ for qi, q in enumerate(questions):
160
+ ids = self._ids_for(q)
161
+ perms = self._perms(q, level)
162
+ cf_rows = []
163
+ perm_ids = [label_ids_for_perm(q, ids, perm) for perm in perms]
164
+ if want_cf and q.key not in self._cf_cache:
165
+ for perm, pids in zip(perms, perm_ids):
166
+ cf_rows.append([add(render_chat(tok, build_prompt(probe, q, perm, self.system)), pids)
167
+ for probe in self.cf_probes])
168
+ for si, st in enumerate(state_texts):
169
+ real_rows = [add(render_chat(tok, build_prompt(st, q, perm, self.system)), pids)
170
+ for perm, pids in zip(perms, perm_ids)]
171
+ plan[(si, qi)] = (perms, real_rows, cf_rows)
172
+
173
+ # 2. one backend call
174
+ prompts = [None] * len(prompt_index)
175
+ for text, i in prompt_index.items():
176
+ prompts[i] = text
177
+ logprobs = self.backend.next_token_logprobs(prompts, prompt_ids)
178
+
179
+ # 3. per question: raw position-space distributions, priors
180
+ out: Dict[tuple, Decision] = {}
181
+ for qi, q in enumerate(questions):
182
+ perms, _, cf_rows = plan[(0, qi)]
183
+ p_pos_raw_all = []
184
+ for si in range(len(states)):
185
+ lp_real = np.stack([logprobs[r] for r in plan[(si, qi)][1]]) # [P, K]
186
+ p_pos_raw_all.append((lp_real, np.stack([_softmax(lp) for lp in lp_real])))
187
+
188
+ cf_prior = None
189
+ if want_cf:
190
+ if cf_rows:
191
+ self._cf_cache[q.key] = np.stack([
192
+ content_free_prior(np.stack([_softmax(logprobs[r]) for r in rows]))
193
+ for rows in cf_rows]) # [P, K]
194
+ cf_prior = self._cf_cache[q.key]
195
+ prior_used = b_prior = None
196
+ if level != "raw":
197
+ stack = np.stack([p for _, p in p_pos_raw_all]) # [N, P, K]
198
+ s, n = self._running.get(q.key, (np.zeros(stack.shape[1:]), 0))
199
+ self._running[q.key] = (s + stack.sum(axis=0), n + len(stack))
200
+ b_prior = self.running_prior(q)
201
+ if self.prior == "batch":
202
+ prior_used = b_prior
203
+ elif self.prior == "content_free":
204
+ prior_used = cf_prior
205
+
206
+ # 4. assemble per state
207
+ for si, (lp_real, p_pos_raw) in enumerate(p_pos_raw_all):
208
+ answer_mass = float(np.exp(lp_real).sum(axis=1).mean())
209
+ raw_probs = marginalize(p_pos_raw[:1], perms[:1])
210
+ diag: Dict[str, Any] = {"answer_mass": answer_mass, "raw_probs": raw_probs,
211
+ "permutations": len(perms), "perms": perms,
212
+ "p_pos_raw": p_pos_raw}
213
+ if level == "raw":
214
+ probs, achieved = raw_probs, "raw"
215
+ else:
216
+ p_pos = apply_contextual(p_pos_raw, prior_used) if prior_used is not None else p_pos_raw
217
+ probs = marginalize(p_pos, perms, self.combine)
218
+ achieved = "L0"
219
+ diag.update({
220
+ "prior_method": self.prior if prior_used is not None else "none",
221
+ "prior": prior_used,
222
+ "cf_prior": cf_prior,
223
+ "batch_prior": b_prior,
224
+ "order_flip_raw": flip_rate_across_perms(p_pos_raw, perms),
225
+ "order_flip_l0": flip_rate_across_perms(p_pos, perms),
226
+ "l0_probs": probs,
227
+ })
228
+ if level == "L1":
229
+ scaler = self._artifacts.get(q.key)
230
+ if scaler is None:
231
+ raise ValueError(f"no L1 artifact for question {q.id}; call calibrate() first")
232
+ probs = scaler.apply(probs)
233
+ achieved = "L1"
234
+ diag["temperature"] = scaler.temperature
235
+ out[(si, qi)] = Decision(q, np.asarray(probs), achieved, diag)
236
+ return out