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 +14 -0
- anyjev/backends/__init__.py +3 -0
- anyjev/backends/base.py +21 -0
- anyjev/backends/fake.py +94 -0
- anyjev/backends/hf.py +52 -0
- anyjev/backends/vllm.py +58 -0
- anyjev/calibrate/__init__.py +6 -0
- anyjev/calibrate/contextual.py +47 -0
- anyjev/calibrate/permute.py +58 -0
- anyjev/calibrate/posthoc.py +60 -0
- anyjev/decider.py +236 -0
- anyjev/question.py +90 -0
- anyjev/readout.py +115 -0
- anyjev/result.py +98 -0
- anyjev/state.py +20 -0
- anyjev-0.0.1.dist-info/METADATA +167 -0
- anyjev-0.0.1.dist-info/RECORD +20 -0
- anyjev-0.0.1.dist-info/WHEEL +5 -0
- anyjev-0.0.1.dist-info/licenses/LICENSE +202 -0
- anyjev-0.0.1.dist-info/top_level.txt +1 -0
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"
|
anyjev/backends/base.py
ADDED
|
@@ -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
|
+
...
|
anyjev/backends/fake.py
ADDED
|
@@ -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]
|
anyjev/backends/vllm.py
ADDED
|
@@ -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
|