shrewd 0.1.0__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.
- shrewd/__init__.py +72 -0
- shrewd/_autotune.py +138 -0
- shrewd/acquire.py +54 -0
- shrewd/calibrate.py +375 -0
- shrewd/data.py +75 -0
- shrewd/decide.py +363 -0
- shrewd/decisions.py +1364 -0
- shrewd/encoder.py +237 -0
- shrewd/evaluate.py +395 -0
- shrewd/judge.py +364 -0
- shrewd/optimize.py +236 -0
- shrewd/project.py +711 -0
- shrewd/students.py +423 -0
- shrewd/teacher.py +270 -0
- shrewd/zeroshot.py +170 -0
- shrewd-0.1.0.dist-info/METADATA +626 -0
- shrewd-0.1.0.dist-info/RECORD +19 -0
- shrewd-0.1.0.dist-info/WHEEL +4 -0
- shrewd-0.1.0.dist-info/licenses/LICENSE +21 -0
shrewd/__init__.py
ADDED
|
@@ -0,0 +1,72 @@
|
|
|
1
|
+
"""Turn LLM judgments into a small, fast, local text model for one fixed task."""
|
|
2
|
+
|
|
3
|
+
from typing import TYPE_CHECKING
|
|
4
|
+
|
|
5
|
+
if TYPE_CHECKING:
|
|
6
|
+
from shrewd._autotune import autotune as autotune
|
|
7
|
+
from shrewd.decide import Choice as Choice
|
|
8
|
+
from shrewd.decide import Noul as Noul
|
|
9
|
+
from shrewd.decide import Score as Score
|
|
10
|
+
from shrewd.decisions import Decisions as Decisions
|
|
11
|
+
from shrewd.evaluate import DistillResult as DistillResult
|
|
12
|
+
from shrewd.evaluate import Finding as Finding
|
|
13
|
+
from shrewd.project import Project as Project
|
|
14
|
+
from shrewd.students import load as load
|
|
15
|
+
|
|
16
|
+
__version__ = "0.1.0"
|
|
17
|
+
|
|
18
|
+
__all__ = [
|
|
19
|
+
"Choice", "Decisions", "DistillResult", "Finding", "Noul", "Project", "Score",
|
|
20
|
+
"autotune", "load",
|
|
21
|
+
]
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def __dir__():
|
|
25
|
+
return sorted(set(globals()) | set(__all__))
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class _needs_teacher_extra:
|
|
29
|
+
def __enter__(self):
|
|
30
|
+
return self
|
|
31
|
+
|
|
32
|
+
def __exit__(self, exc_type, exc, tb):
|
|
33
|
+
if isinstance(exc, ImportError) and exc.name in ("litellm", "gepa"):
|
|
34
|
+
raise ImportError(
|
|
35
|
+
f"training needs {exc.name}, which is not installed: "
|
|
36
|
+
'pip install "shrewd[teacher]" (inference-only installs can skip it)'
|
|
37
|
+
) from None
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
# Lazy so that `from shrewd import load` works in inference environments
|
|
41
|
+
# where litellm/gepa are not installed.
|
|
42
|
+
def __getattr__(name):
|
|
43
|
+
if name == "Project":
|
|
44
|
+
with _needs_teacher_extra():
|
|
45
|
+
from shrewd.project import Project
|
|
46
|
+
|
|
47
|
+
return Project
|
|
48
|
+
if name == "autotune":
|
|
49
|
+
# the submodule is named _autotune so this attribute never collides with it
|
|
50
|
+
# (a same-named submodule would shadow the function on from-imports)
|
|
51
|
+
with _needs_teacher_extra():
|
|
52
|
+
from shrewd._autotune import autotune
|
|
53
|
+
|
|
54
|
+
return autotune
|
|
55
|
+
if name == "load":
|
|
56
|
+
from shrewd.students import load
|
|
57
|
+
|
|
58
|
+
return load
|
|
59
|
+
if name in ("Choice", "Noul", "Score"):
|
|
60
|
+
from shrewd import decide
|
|
61
|
+
|
|
62
|
+
return getattr(decide, name)
|
|
63
|
+
if name == "Decisions":
|
|
64
|
+
with _needs_teacher_extra():
|
|
65
|
+
from shrewd.decisions import Decisions
|
|
66
|
+
|
|
67
|
+
return Decisions
|
|
68
|
+
if name in ("DistillResult", "Finding"):
|
|
69
|
+
from shrewd import evaluate
|
|
70
|
+
|
|
71
|
+
return getattr(evaluate, name)
|
|
72
|
+
raise AttributeError(f"module 'shrewd' has no attribute {name!r}")
|
shrewd/_autotune.py
ADDED
|
@@ -0,0 +1,138 @@
|
|
|
1
|
+
"""Budget-capped controller over the Project pipeline - cheapest teacher first, escalating
|
|
2
|
+
only while the dev-split score falls short of `target` and the next tier fits the
|
|
3
|
+
remaining budget. The locked test set is evaluated once, on the winner.
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
import json
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
|
|
9
|
+
from shrewd import data, evaluate
|
|
10
|
+
from shrewd.optimize import build_seed_prompt
|
|
11
|
+
from shrewd.project import Project
|
|
12
|
+
from shrewd.students import STUDENTS, check_students
|
|
13
|
+
from shrewd.teacher import estimate_calls, resolve_model
|
|
14
|
+
|
|
15
|
+
# multiplier on the estimated labeling cost of a tier, to leave room for
|
|
16
|
+
# optimize() and the final test labeling before committing to the tier
|
|
17
|
+
ESTIMATE_MARGIN = 1.5
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def _spent(projects):
|
|
21
|
+
return sum(s["cost_usd"] for p in projects for s in p._manifest["stages"])
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def _dev_probe(kept, dev, student_name, classes, seed):
|
|
25
|
+
"""Macro-F1 on gold dev of a student trained on pool rows alone (no dev leak)."""
|
|
26
|
+
probe = STUDENTS[student_name](seed=seed)
|
|
27
|
+
probe.fit(kept["text"].tolist(), kept["label"].tolist())
|
|
28
|
+
return evaluate.compute_metrics(
|
|
29
|
+
dev["label"], probe.predict(dev["text"].tolist()), classes
|
|
30
|
+
)["macro_f1"]
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def autotune(dir, instructions, labels, seed_df, pool_df, budget_usd, target=None,
|
|
34
|
+
teachers=("anthropic/claude-haiku-4-5", "anthropic/claude-sonnet-5"),
|
|
35
|
+
students=("tfidf",), votes=1, optimize_budget=600, seed=42):
|
|
36
|
+
"""Run the escalation ladder under a spending cap and return the winning pipeline.
|
|
37
|
+
|
|
38
|
+
One sub-project per teacher tier under `dir` (same seed data, so identical splits). Per
|
|
39
|
+
tier: optimize -> label -> a free sweep of `students` x min_confidence on the dev split.
|
|
40
|
+
Escalates while the best dev macro-F1 is below `target` (always if None) and the
|
|
41
|
+
estimated cost fits.
|
|
42
|
+
|
|
43
|
+
Returns {"teacher", "student", "min_confidence", "dir", "dev_macro_f1", "spent_usd",
|
|
44
|
+
"result", "trail"}. The trail is also written to `dir`/autotune_trail.json. Re-running
|
|
45
|
+
with the same `dir` resumes from caches.
|
|
46
|
+
"""
|
|
47
|
+
check_students(students) # fail on a missing optional extra before any spend
|
|
48
|
+
parent = Path(dir)
|
|
49
|
+
parent.mkdir(parents=True, exist_ok=True)
|
|
50
|
+
confidence_rungs = (0.0,) if votes == 1 else (0.0, 0.67, 1.0)
|
|
51
|
+
pool_texts = pool_df["text"].astype(str).tolist()
|
|
52
|
+
trail, projects = [], []
|
|
53
|
+
best = None
|
|
54
|
+
|
|
55
|
+
for model in (resolve_model(m) for m in teachers):
|
|
56
|
+
tier_dir = parent / model.split("/")[-1]
|
|
57
|
+
project = Project(
|
|
58
|
+
tier_dir, instructions=instructions, labels=labels, teacher=model, seed=seed
|
|
59
|
+
)
|
|
60
|
+
if not (project.dir / "seed_test.csv").exists():
|
|
61
|
+
project.add_seed(seed_df)
|
|
62
|
+
|
|
63
|
+
# budget guard - price the tier's labeling (seed prompt as proxy) before spending
|
|
64
|
+
_, estimate = estimate_calls(
|
|
65
|
+
pool_texts, build_seed_prompt(instructions, labels), model, votes,
|
|
66
|
+
project._cache(),
|
|
67
|
+
)
|
|
68
|
+
remaining = round(budget_usd - _spent(projects), 2)
|
|
69
|
+
if estimate is None:
|
|
70
|
+
trail.append({"rung": f"teacher:{model}", "decision":
|
|
71
|
+
"no litellm price data, budget guard waived for this tier"})
|
|
72
|
+
else:
|
|
73
|
+
planned = round(estimate * ESTIMATE_MARGIN, 2)
|
|
74
|
+
if planned > remaining:
|
|
75
|
+
if best is None:
|
|
76
|
+
raise ValueError(
|
|
77
|
+
f"budget_usd={budget_usd} cannot cover the first tier "
|
|
78
|
+
f"({model}: ~${planned}); nothing was spent"
|
|
79
|
+
)
|
|
80
|
+
trail.append({"rung": f"teacher:{model}", "decision":
|
|
81
|
+
f"skipped: needs ~${planned}, ${remaining} left"})
|
|
82
|
+
break
|
|
83
|
+
|
|
84
|
+
project.optimize(budget=optimize_budget, target=target)
|
|
85
|
+
project.label(pool_df, votes=votes)
|
|
86
|
+
projects.append(project)
|
|
87
|
+
|
|
88
|
+
labeled = data.read_csv(project.dir / "pool_labeled.csv")
|
|
89
|
+
dev = data.read_csv(project.dir / "seed_dev.csv")
|
|
90
|
+
test = data.read_csv(project.dir / "seed_test.csv")
|
|
91
|
+
valid, _ = data.usable_pool(labeled, dev, test)
|
|
92
|
+
classes = list(project.labels)
|
|
93
|
+
for min_confidence in confidence_rungs:
|
|
94
|
+
kept = valid[valid["confidence"] >= min_confidence]
|
|
95
|
+
if kept["label"].nunique() < 2:
|
|
96
|
+
trail.append({
|
|
97
|
+
"rung": f"teacher:{model}", "min_confidence": min_confidence,
|
|
98
|
+
"decision": "probe skipped: under two labels left after filtering",
|
|
99
|
+
})
|
|
100
|
+
continue
|
|
101
|
+
for student in students:
|
|
102
|
+
dev_f1 = _dev_probe(kept, dev, student, classes, seed)
|
|
103
|
+
trail.append({
|
|
104
|
+
"rung": f"teacher:{model}", "student": student,
|
|
105
|
+
"min_confidence": min_confidence, "dev_macro_f1": round(dev_f1, 4),
|
|
106
|
+
"spent_usd": round(_spent(projects), 2),
|
|
107
|
+
})
|
|
108
|
+
if best is None or dev_f1 > best["dev_f1"]:
|
|
109
|
+
best = {"project": project, "student": student,
|
|
110
|
+
"min_confidence": min_confidence, "dev_f1": dev_f1}
|
|
111
|
+
if target is not None and best is not None and best["dev_f1"] >= target:
|
|
112
|
+
trail.append({"rung": f"teacher:{model}",
|
|
113
|
+
"decision": f"dev target {target} reached, stopping early"})
|
|
114
|
+
break
|
|
115
|
+
|
|
116
|
+
if best is None:
|
|
117
|
+
(parent / "autotune_trail.json").write_text(json.dumps(trail, indent=2))
|
|
118
|
+
raise RuntimeError(
|
|
119
|
+
"no configuration could be probed: every rung left fewer than two distinct "
|
|
120
|
+
"labels (the teacher may be collapsing to one class). See "
|
|
121
|
+
f"{parent / 'autotune_trail.json'} and the per-tier pool_labeled.csv files."
|
|
122
|
+
)
|
|
123
|
+
# the single test-set evaluation, on the winning configuration only
|
|
124
|
+
result = best["project"].distill(
|
|
125
|
+
student=best["student"], min_confidence=best["min_confidence"]
|
|
126
|
+
)
|
|
127
|
+
summary = {
|
|
128
|
+
"teacher": best["project"].teacher,
|
|
129
|
+
"student": best["student"],
|
|
130
|
+
"min_confidence": best["min_confidence"],
|
|
131
|
+
"dir": str(best["project"].dir),
|
|
132
|
+
"dev_macro_f1": round(best["dev_f1"], 4),
|
|
133
|
+
"spent_usd": round(_spent(projects), 2),
|
|
134
|
+
}
|
|
135
|
+
trail.append({"rung": "final", **summary,
|
|
136
|
+
"test_macro_f1": round(result.metrics["student"]["macro_f1"], 4)})
|
|
137
|
+
(parent / "autotune_trail.json").write_text(json.dumps(trail, indent=2))
|
|
138
|
+
return {**summary, "result": result, "trail": trail}
|
shrewd/acquire.py
ADDED
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
"""Pick which pool rows are worth a teacher call. The ones the current tfidf probe is least
|
|
2
|
+
sure about, spread across clusters. Measured at 1.1-2.3x fewer labels than file order for
|
|
3
|
+
the same student accuracy. The probe's picks were about as useful to the embedding
|
|
4
|
+
student as random order, so the probe does not commit you to shipping tfidf.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
from sklearn.cluster import KMeans
|
|
9
|
+
from sklearn.decomposition import TruncatedSVD
|
|
10
|
+
|
|
11
|
+
from shrewd.students import TfidfStudent
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def default_batch(n_pool):
|
|
15
|
+
return int(min(500, max(50, n_pool // 10)))
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def fit_probe(train_texts, train_labels, seed):
|
|
19
|
+
probe = TfidfStudent(seed=seed)
|
|
20
|
+
probe.fit(train_texts, train_labels)
|
|
21
|
+
return probe
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def select(probe, candidates, batch, seed):
|
|
25
|
+
"""Return indices into `candidates` for the next batch.
|
|
26
|
+
|
|
27
|
+
Rank by margin between the top two probabilities (small margin = unsure), take the
|
|
28
|
+
4x most uncertain as a shortlist, cluster the shortlist and keep the most uncertain
|
|
29
|
+
row from each cluster.
|
|
30
|
+
"""
|
|
31
|
+
if batch >= len(candidates):
|
|
32
|
+
return list(range(len(candidates)))
|
|
33
|
+
proba = probe.predict_proba(candidates)
|
|
34
|
+
top2 = np.sort(proba, axis=1)[:, -2:]
|
|
35
|
+
margin = top2[:, 1] - top2[:, 0]
|
|
36
|
+
order = np.argsort(margin)
|
|
37
|
+
shortlist = order[: min(len(candidates), 4 * batch)]
|
|
38
|
+
if len(shortlist) <= batch:
|
|
39
|
+
return [int(i) for i in shortlist]
|
|
40
|
+
features = probe._pipe.named_steps["features"].transform([candidates[i] for i in shortlist])
|
|
41
|
+
dims = min(64, features.shape[1] - 1, len(shortlist) - 1)
|
|
42
|
+
if dims < 2:
|
|
43
|
+
return [int(i) for i in order[:batch]]
|
|
44
|
+
z = TruncatedSVD(dims, random_state=seed).fit_transform(features)
|
|
45
|
+
clusters = KMeans(n_clusters=batch, n_init=3, random_state=seed).fit_predict(z)
|
|
46
|
+
picked = []
|
|
47
|
+
for c in range(batch):
|
|
48
|
+
members = np.where(clusters == c)[0]
|
|
49
|
+
if len(members):
|
|
50
|
+
picked.append(int(shortlist[members[np.argmin(margin[shortlist[members]])]]))
|
|
51
|
+
if len(picked) < batch: # empty clusters: top up by uncertainty
|
|
52
|
+
taken = set(picked)
|
|
53
|
+
picked += [int(i) for i in order if int(i) not in taken][: batch - len(picked)]
|
|
54
|
+
return picked
|
shrewd/calibrate.py
ADDED
|
@@ -0,0 +1,375 @@
|
|
|
1
|
+
"""Post-hoc calibration of a student's probabilities.
|
|
2
|
+
|
|
3
|
+
With more than two options the calibration map runs on the top score, so the predicted
|
|
4
|
+
option never changes. For yes/no questions the calibrated quantity is P(yes) itself and it
|
|
5
|
+
may cross 0.5: a question with a 3% base rate should not sit at 0.5.
|
|
6
|
+
See `cross_fit_proba` for which rows the calibrator is fit on.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
from scipy.optimize import minimize_scalar
|
|
11
|
+
from sklearn.isotonic import IsotonicRegression
|
|
12
|
+
from sklearn.linear_model import LogisticRegression
|
|
13
|
+
from sklearn.model_selection import StratifiedKFold
|
|
14
|
+
|
|
15
|
+
EPS = 1e-12
|
|
16
|
+
METHODS = ("temperature", "isotonic", "platt", "none")
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def _logits(proba):
|
|
20
|
+
return np.log(np.clip(np.asarray(proba, dtype=float), EPS, 1.0))
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def _softmax(z):
|
|
24
|
+
z = z - z.max(axis=1, keepdims=True)
|
|
25
|
+
e = np.exp(z)
|
|
26
|
+
return e / e.sum(axis=1, keepdims=True)
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def _onehot(y_idx, k):
|
|
30
|
+
out = np.zeros((len(y_idx), k))
|
|
31
|
+
out[np.arange(len(y_idx)), y_idx] = 1.0
|
|
32
|
+
return out
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def encode(labels, classes):
|
|
36
|
+
"""Label strings to column indices. Labels outside `classes` come back as -1."""
|
|
37
|
+
index = {c: i for i, c in enumerate(classes)}
|
|
38
|
+
return np.array([index.get(str(y), -1) for y in labels])
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
# ---------------------------------------------------------------- scoring rules
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def nll(proba, y_idx):
|
|
45
|
+
"""Mean negative log likelihood of the true class (log loss)."""
|
|
46
|
+
proba = np.clip(np.asarray(proba, dtype=float), EPS, 1.0)
|
|
47
|
+
return float(-np.mean(np.log(proba[np.arange(len(y_idx)), y_idx])))
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def brier(proba, y_idx):
|
|
51
|
+
"""Multiclass Brier score: mean squared error over the whole probability vector.
|
|
52
|
+
|
|
53
|
+
Ranges 0 (perfect) to 2 (confidently wrong). Unlike ECE this is a proper scoring
|
|
54
|
+
rule, so it can't be gamed by reporting the base rate on every row.
|
|
55
|
+
"""
|
|
56
|
+
proba = np.asarray(proba, dtype=float)
|
|
57
|
+
return float(np.mean(((proba - _onehot(y_idx, proba.shape[1])) ** 2).sum(axis=1)))
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def ece(confidence, correct, bins=10):
|
|
61
|
+
"""Expected calibration error with equal-mass bins, so the number does not swing with the
|
|
62
|
+
bin count when most rows sit in one confidence band.
|
|
63
|
+
"""
|
|
64
|
+
confidence = np.asarray(confidence, dtype=float)
|
|
65
|
+
correct = np.asarray(correct, dtype=float)
|
|
66
|
+
if len(confidence) == 0:
|
|
67
|
+
return float("nan")
|
|
68
|
+
order = np.argsort(confidence, kind="mergesort")
|
|
69
|
+
conf, corr = confidence[order], correct[order]
|
|
70
|
+
total = 0.0
|
|
71
|
+
for chunk in np.array_split(np.arange(len(conf)), min(bins, len(conf))):
|
|
72
|
+
if len(chunk):
|
|
73
|
+
total += len(chunk) / len(conf) * abs(conf[chunk].mean() - corr[chunk].mean())
|
|
74
|
+
return float(total)
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def reliability_table(confidence, correct, bins=10):
|
|
78
|
+
"""Per-bin stated confidence vs. measured accuracy: the reliability diagram as rows."""
|
|
79
|
+
confidence = np.asarray(confidence, dtype=float)
|
|
80
|
+
correct = np.asarray(correct, dtype=float)
|
|
81
|
+
order = np.argsort(confidence, kind="mergesort")
|
|
82
|
+
conf, corr = confidence[order], correct[order]
|
|
83
|
+
rows = []
|
|
84
|
+
for chunk in np.array_split(np.arange(len(conf)), min(bins, max(len(conf), 1))):
|
|
85
|
+
if not len(chunk):
|
|
86
|
+
continue
|
|
87
|
+
rows.append(
|
|
88
|
+
{
|
|
89
|
+
"low": round(float(conf[chunk].min()), 4),
|
|
90
|
+
"high": round(float(conf[chunk].max()), 4),
|
|
91
|
+
"n": int(len(chunk)),
|
|
92
|
+
"stated": round(float(conf[chunk].mean()), 4),
|
|
93
|
+
"actual": round(float(corr[chunk].mean()), 4),
|
|
94
|
+
}
|
|
95
|
+
)
|
|
96
|
+
return rows
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def aurc(confidence, correct):
|
|
100
|
+
"""Area under the risk-coverage curve: mean error rate over every prefix of rows
|
|
101
|
+
sorted most-confident-first. Lower is better. It rewards a confidence signal that
|
|
102
|
+
puts the mistakes at the bottom, independently of whether the scale is calibrated.
|
|
103
|
+
"""
|
|
104
|
+
confidence = np.asarray(confidence, dtype=float)
|
|
105
|
+
correct = np.asarray(correct, dtype=float)
|
|
106
|
+
if len(confidence) == 0:
|
|
107
|
+
return float("nan")
|
|
108
|
+
order = np.argsort(-confidence, kind="mergesort")
|
|
109
|
+
errors = 1.0 - correct[order]
|
|
110
|
+
return float(np.mean(np.cumsum(errors) / np.arange(1, len(errors) + 1)))
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
def calibration_metrics(proba, labels, classes, bins=10):
|
|
114
|
+
"""Every calibration number for one set of predictions, as a plain dict."""
|
|
115
|
+
proba = np.asarray(proba, dtype=float)
|
|
116
|
+
y_idx = encode(labels, classes)
|
|
117
|
+
keep = y_idx >= 0
|
|
118
|
+
proba, y_idx = proba[keep], y_idx[keep]
|
|
119
|
+
if len(y_idx) == 0:
|
|
120
|
+
return {}
|
|
121
|
+
# for a yes/no question the reliability curve that matters runs on P(yes) across
|
|
122
|
+
# every row, not on the winning side's confidence: with a 3% base rate the latter
|
|
123
|
+
# is dominated by easy "no"s and looks excellent while P(yes) is off by 10x
|
|
124
|
+
if len(classes) == 2:
|
|
125
|
+
confidence, correct = proba[:, 1], (y_idx == 1).astype(float)
|
|
126
|
+
else:
|
|
127
|
+
confidence = proba.max(axis=1)
|
|
128
|
+
correct = (proba.argmax(axis=1) == y_idx).astype(float)
|
|
129
|
+
accuracy = float((proba.argmax(axis=1) == y_idx).mean())
|
|
130
|
+
return {
|
|
131
|
+
"n": int(len(y_idx)),
|
|
132
|
+
"accuracy": round(accuracy, 4),
|
|
133
|
+
"mean_confidence": round(float(confidence.mean()), 4),
|
|
134
|
+
"overconfidence": round(float(confidence.mean() - correct.mean()), 4),
|
|
135
|
+
"ece": round(ece(confidence, correct, bins), 4),
|
|
136
|
+
"brier": round(brier(proba, y_idx), 4),
|
|
137
|
+
"nll": round(nll(proba, y_idx), 4),
|
|
138
|
+
"aurc": round(aurc(confidence, correct), 4),
|
|
139
|
+
"reliability": reliability_table(confidence, correct, bins),
|
|
140
|
+
}
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
# ---------------------------------------------------------------- the calibrator
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
class Calibrator:
|
|
147
|
+
"""A fitted, monotone rescaling of a student's probabilities. Serializes to JSON.
|
|
148
|
+
|
|
149
|
+
`temperature` (one scalar), `platt` (two parameters on the log-odds, the right choice
|
|
150
|
+
for a lopsided yes/no question) or `isotonic` (a step function, needs thousands of rows
|
|
151
|
+
and collapses on a few dozen).
|
|
152
|
+
"""
|
|
153
|
+
|
|
154
|
+
def __init__(self, method, params, n_fit=0, classes=None):
|
|
155
|
+
self.method = method
|
|
156
|
+
self.params = params
|
|
157
|
+
self.n_fit = n_fit
|
|
158
|
+
self.classes = classes
|
|
159
|
+
|
|
160
|
+
# -- fitting ---------------------------------------------------------------
|
|
161
|
+
|
|
162
|
+
@classmethod
|
|
163
|
+
def fit(cls, proba, labels, classes, method="auto", seed=0):
|
|
164
|
+
"""Fit on (probabilities, true labels). `method="auto"` picks by cross-validated
|
|
165
|
+
log loss among temperature/platt/isotonic/none, which keeps a small or degenerate
|
|
166
|
+
calibration set from choosing a flexible method it cannot support."""
|
|
167
|
+
proba = np.asarray(proba, dtype=float)
|
|
168
|
+
y_idx = encode(labels, classes)
|
|
169
|
+
keep = y_idx >= 0
|
|
170
|
+
proba, y_idx = proba[keep], y_idx[keep]
|
|
171
|
+
if len(y_idx) < 20 or len(np.unique(y_idx)) < 2:
|
|
172
|
+
return cls("none", {}, n_fit=int(len(y_idx)), classes=list(classes))
|
|
173
|
+
if method == "auto":
|
|
174
|
+
method = cls._pick(proba, y_idx, classes, seed)
|
|
175
|
+
params = cls._fit_params(method, proba, y_idx, len(classes))
|
|
176
|
+
return cls(method, params, n_fit=int(len(y_idx)), classes=list(classes))
|
|
177
|
+
|
|
178
|
+
@staticmethod
|
|
179
|
+
def _fit_params(method, proba, y_idx, k):
|
|
180
|
+
"""Fit one method's parameters. Multi-class maps the top score and keeps the argmax;
|
|
181
|
+
two-class maps P(positive) directly and may cross 0.5.
|
|
182
|
+
"""
|
|
183
|
+
if method == "none":
|
|
184
|
+
return {}
|
|
185
|
+
if method == "temperature":
|
|
186
|
+
return {"temperature": _fit_temperature(_logits(proba), _onehot(y_idx, k))}
|
|
187
|
+
binary = k == 2
|
|
188
|
+
if binary:
|
|
189
|
+
score, target = proba[:, 1], (y_idx == 1).astype(float)
|
|
190
|
+
else:
|
|
191
|
+
score, target = proba.max(axis=1), (proba.argmax(axis=1) == y_idx).astype(float)
|
|
192
|
+
if method == "platt":
|
|
193
|
+
a, b = _fit_platt(score, target)
|
|
194
|
+
return {"a": a, "b": b, "binary": binary}
|
|
195
|
+
if method == "isotonic":
|
|
196
|
+
iso = IsotonicRegression(out_of_bounds="clip", y_min=0.0, y_max=1.0)
|
|
197
|
+
iso.fit(score, target)
|
|
198
|
+
return {
|
|
199
|
+
"x": [round(float(v), 6) for v in iso.X_thresholds_],
|
|
200
|
+
"y": [round(float(v), 6) for v in iso.y_thresholds_],
|
|
201
|
+
"binary": binary,
|
|
202
|
+
}
|
|
203
|
+
raise ValueError(f"unknown calibration method: {method!r} (expected one of {METHODS})")
|
|
204
|
+
|
|
205
|
+
@classmethod
|
|
206
|
+
def _pick(cls, proba, y_idx, classes, seed, n_splits=4):
|
|
207
|
+
"""Cross-validated log loss over the calibration rows themselves."""
|
|
208
|
+
n_splits = max(2, min(n_splits, int(np.bincount(y_idx).min()) if len(y_idx) else 2))
|
|
209
|
+
try:
|
|
210
|
+
folds = list(StratifiedKFold(n_splits, shuffle=True, random_state=seed).split(
|
|
211
|
+
proba, y_idx
|
|
212
|
+
))
|
|
213
|
+
except ValueError:
|
|
214
|
+
return "temperature"
|
|
215
|
+
scores = {}
|
|
216
|
+
for method in METHODS:
|
|
217
|
+
total, ok = 0.0, True
|
|
218
|
+
for train, test in folds:
|
|
219
|
+
try:
|
|
220
|
+
params = cls._fit_params(method, proba[train], y_idx[train], len(classes))
|
|
221
|
+
held = cls(method, params, classes=list(classes)).transform(proba[test])
|
|
222
|
+
except (ValueError, FloatingPointError):
|
|
223
|
+
ok = False
|
|
224
|
+
break
|
|
225
|
+
total += nll(held, y_idx[test]) * len(test)
|
|
226
|
+
if ok:
|
|
227
|
+
scores[method] = total / len(y_idx)
|
|
228
|
+
if not scores:
|
|
229
|
+
return "temperature"
|
|
230
|
+
best = min(scores, key=scores.get)
|
|
231
|
+
# a tie goes to the simpler method: temperature generalizes off-distribution
|
|
232
|
+
# better than isotonic and there is no reason to pay for flexibility you don't use
|
|
233
|
+
for method in METHODS:
|
|
234
|
+
if method in scores and scores[method] <= scores[best] + 1e-4:
|
|
235
|
+
return method
|
|
236
|
+
return best
|
|
237
|
+
|
|
238
|
+
# -- applying --------------------------------------------------------------
|
|
239
|
+
|
|
240
|
+
def transform(self, proba):
|
|
241
|
+
"""Rescale probabilities. Row-stochastic in, row-stochastic out."""
|
|
242
|
+
proba = np.clip(np.asarray(proba, dtype=float), EPS, 1.0)
|
|
243
|
+
if self.method == "none":
|
|
244
|
+
return proba / proba.sum(axis=1, keepdims=True)
|
|
245
|
+
if self.method == "temperature":
|
|
246
|
+
return _softmax(_logits(proba) / self.params["temperature"])
|
|
247
|
+
binary = bool(self.params.get("binary")) and proba.shape[1] == 2
|
|
248
|
+
score = proba[:, 1] if binary else proba.max(axis=1)
|
|
249
|
+
if self.method == "platt":
|
|
250
|
+
mapped = _apply_platt(score, self.params["a"], self.params["b"])
|
|
251
|
+
elif self.method == "isotonic":
|
|
252
|
+
mapped = np.interp(score, self.params["x"], self.params["y"])
|
|
253
|
+
else:
|
|
254
|
+
raise ValueError(f"unknown calibration method: {self.method!r}")
|
|
255
|
+
mapped = np.clip(mapped, EPS, 1.0 - EPS)
|
|
256
|
+
if binary:
|
|
257
|
+
return np.column_stack([1.0 - mapped, mapped])
|
|
258
|
+
return _rescale_top(proba, mapped)
|
|
259
|
+
|
|
260
|
+
# -- persistence -----------------------------------------------------------
|
|
261
|
+
|
|
262
|
+
def to_dict(self):
|
|
263
|
+
return {
|
|
264
|
+
"method": self.method,
|
|
265
|
+
"params": self.params,
|
|
266
|
+
"n_fit": self.n_fit,
|
|
267
|
+
"classes": self.classes,
|
|
268
|
+
}
|
|
269
|
+
|
|
270
|
+
@classmethod
|
|
271
|
+
def from_dict(cls, blob):
|
|
272
|
+
return cls(
|
|
273
|
+
blob["method"], blob.get("params", {}), blob.get("n_fit", 0), blob.get("classes")
|
|
274
|
+
)
|
|
275
|
+
|
|
276
|
+
def __repr__(self):
|
|
277
|
+
detail = ""
|
|
278
|
+
if self.method == "temperature":
|
|
279
|
+
detail = f" T={self.params['temperature']:.3f}"
|
|
280
|
+
return f"<Calibrator {self.method}{detail} fit on {self.n_fit} rows>"
|
|
281
|
+
|
|
282
|
+
|
|
283
|
+
def _rescale_top(proba, new_top):
|
|
284
|
+
"""Set the winning class to `new_top` and spread the rest proportionally. `new_top` is
|
|
285
|
+
floored where the runner-up would draw level, so the winning option never changes.
|
|
286
|
+
"""
|
|
287
|
+
out = np.array(proba, dtype=float, copy=True)
|
|
288
|
+
rows = np.arange(len(out))
|
|
289
|
+
win = out.argmax(axis=1)
|
|
290
|
+
top = out[rows, win]
|
|
291
|
+
rest = out.sum(axis=1) - top
|
|
292
|
+
other_max = np.where(out.shape[1] > 1, (out - np.eye(out.shape[1])[win] * out).max(axis=1), 0.0)
|
|
293
|
+
floor = np.where(rest > EPS, other_max / np.maximum(rest + other_max, EPS), 0.0)
|
|
294
|
+
new_top = np.maximum(new_top, np.minimum(floor + 1e-9, 1.0 - EPS))
|
|
295
|
+
scale = np.where(rest > EPS, (1.0 - new_top) / np.maximum(rest, EPS), 0.0)
|
|
296
|
+
out *= scale[:, None]
|
|
297
|
+
out[rows, win] = new_top
|
|
298
|
+
total = out.sum(axis=1, keepdims=True)
|
|
299
|
+
return out / np.maximum(total, EPS)
|
|
300
|
+
|
|
301
|
+
|
|
302
|
+
def _fit_temperature(logits, onehot, bounds=(-3.0, 3.0)):
|
|
303
|
+
"""The temperature minimizing log loss. Optimized over log T so it stays positive."""
|
|
304
|
+
|
|
305
|
+
def objective(log_t):
|
|
306
|
+
z = logits / np.exp(log_t)
|
|
307
|
+
z = z - z.max(axis=1, keepdims=True)
|
|
308
|
+
log_p = z - np.log(np.exp(z).sum(axis=1, keepdims=True))
|
|
309
|
+
return -np.mean((onehot * log_p).sum(axis=1))
|
|
310
|
+
|
|
311
|
+
result = minimize_scalar(objective, bounds=bounds, method="bounded")
|
|
312
|
+
return round(float(np.exp(result.x)), 6)
|
|
313
|
+
|
|
314
|
+
|
|
315
|
+
def _fit_platt(confidence, correct):
|
|
316
|
+
"""Two-parameter logistic on the top score's log-odds.
|
|
317
|
+
|
|
318
|
+
Falls back to the identity map when the calibration rows are all right or all
|
|
319
|
+
wrong: there is no slope to estimate from one class, and a logistic fit would
|
|
320
|
+
either refuse outright or run its coefficients off to infinity.
|
|
321
|
+
"""
|
|
322
|
+
correct = correct.astype(int)
|
|
323
|
+
if len(np.unique(correct)) < 2:
|
|
324
|
+
return 1.0, 0.0
|
|
325
|
+
x = np.log(np.clip(confidence, EPS, 1 - 1e-9) / np.clip(1 - confidence, EPS, 1.0))
|
|
326
|
+
model = LogisticRegression(C=1e6, solver="lbfgs", max_iter=1000)
|
|
327
|
+
model.fit(x.reshape(-1, 1), correct)
|
|
328
|
+
return round(float(model.coef_[0][0]), 6), round(float(model.intercept_[0]), 6)
|
|
329
|
+
|
|
330
|
+
|
|
331
|
+
def _apply_platt(confidence, a, b):
|
|
332
|
+
x = np.log(np.clip(confidence, EPS, 1 - 1e-9) / np.clip(1 - confidence, EPS, 1.0))
|
|
333
|
+
return 1.0 / (1.0 + np.exp(-(a * x + b)))
|
|
334
|
+
|
|
335
|
+
|
|
336
|
+
# ---------------------------------------------------------------- the right rows
|
|
337
|
+
|
|
338
|
+
|
|
339
|
+
def cross_fit_proba(make_student, texts, labels, classes, n_splits=5, seed=0, desc=None):
|
|
340
|
+
"""Out-of-fold probabilities over the teacher-labeled pool: K refits, each scoring the
|
|
341
|
+
rows it did not train on. No API calls.
|
|
342
|
+
|
|
343
|
+
Fitting on the student's own training rows gives a temperature near 1 (memorized
|
|
344
|
+
scores). Fitting on the seed dev split (80-200 easy rows) made two of three datasets
|
|
345
|
+
worse. The labels are the teacher's, so this calibrates to agreement with the teacher;
|
|
346
|
+
the report says so when the teacher is the weak link.
|
|
347
|
+
"""
|
|
348
|
+
texts, labels = list(texts), [str(y) for y in labels]
|
|
349
|
+
index = {c: i for i, c in enumerate(classes)}
|
|
350
|
+
y = np.array([index.get(label, -1) for label in labels])
|
|
351
|
+
keep = np.where(y >= 0)[0]
|
|
352
|
+
if len(keep) < 40:
|
|
353
|
+
return None, None
|
|
354
|
+
counts = np.bincount(y[keep], minlength=len(classes))
|
|
355
|
+
n_splits = int(max(2, min(n_splits, counts[counts > 0].min())))
|
|
356
|
+
if n_splits < 2:
|
|
357
|
+
return None, None
|
|
358
|
+
|
|
359
|
+
sub_texts = [texts[i] for i in keep]
|
|
360
|
+
sub_labels = [labels[i] for i in keep]
|
|
361
|
+
oof = np.zeros((len(keep), len(classes)))
|
|
362
|
+
splitter = StratifiedKFold(n_splits, shuffle=True, random_state=seed)
|
|
363
|
+
folds = list(splitter.split(np.zeros(len(keep)), y[keep]))
|
|
364
|
+
for fold, (train, test) in enumerate(folds, 1):
|
|
365
|
+
if desc:
|
|
366
|
+
print(f" {desc}: calibration fold {fold}/{len(folds)}", flush=True)
|
|
367
|
+
student = make_student()
|
|
368
|
+
student.fit([sub_texts[i] for i in train], [sub_labels[i] for i in train])
|
|
369
|
+
proba = student.predict_proba([sub_texts[i] for i in test])
|
|
370
|
+
for j, cls in enumerate(student.classes_):
|
|
371
|
+
if cls in index:
|
|
372
|
+
oof[test, index[cls]] = proba[:, j]
|
|
373
|
+
total = oof.sum(axis=1, keepdims=True)
|
|
374
|
+
oof = np.divide(oof, total, out=np.full_like(oof, 1.0 / len(classes)), where=total > EPS)
|
|
375
|
+
return oof, sub_labels
|