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 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