#!/usr/bin/env python3 """ learning_report.py ================== Two tools for making an honest claim about a model. Tool 1: diagnostic checklist Given a trained model and a task, produce a report that cannot be gamed by any single metric. The report contains four splits (interpolation, near-OOD, far-OOD, structural-OOD), mean confidence and ECE on each, temperature-boundary detection, output-uniformity canary, and calibration thresholds. A single verdict string summarises the findings. Tool 2: training-data-volume calculator Given a model architecture, an output vocabulary size, and a target held-out accuracy, estimate how many training examples are needed. Two estimates: a formula-based heuristic and an empirical curve fit from actually training the model on a ladder of dataset sizes. Two curve families are supported: exponential saturation and Hill cooperativity. Pure stdlib + numpy. No downloads. Runs in under a minute on CPU. """ from __future__ import annotations import math import random from dataclasses import dataclass, field from typing import Callable, Dict, List, Optional, Sequence, Tuple import numpy as np # ===================================================================== # §1 Utilities # ===================================================================== def softmax(z: np.ndarray) -> np.ndarray: z = z - z.max(axis=-1, keepdims=True) e = np.exp(z) return e / e.sum(axis=-1, keepdims=True) def one_hot(i: int, size: int) -> np.ndarray: v = np.zeros(size, dtype=np.float64) if 0 <= i < size: v[i] = 1.0 return v def normalized_entropy(p: np.ndarray) -> float: p = p / max(1e-12, p.sum()) n = len(p) if n <= 1: return 0.0 h = -np.sum(p * np.log(p + 1e-12)) return float(h / math.log(n)) # ===================================================================== # §2 Task: a + b (or a * b) with bounded operands # ===================================================================== MAX_VAL = 20 # encoder capacity per operand MAX_SUM = 40 # output vocabulary TRAIN_MAX = 9 # default training range for each operand def encode_pair(a: int, b: int) -> np.ndarray: return np.concatenate([one_hot(a, MAX_VAL), one_hot(b, MAX_VAL)]) def encode_pairs(pairs: Sequence[Tuple[int, int]]) -> np.ndarray: return np.stack([encode_pair(a, b) for a, b in pairs]) def sample_pairs(lo_a: int, hi_a: int, lo_b: int, hi_b: int, n: int, seed: int = 0) -> List[Tuple[int, int]]: rng = random.Random(seed) out: List[Tuple[int, int]] = [] seen = set() attempts = 0 while len(out) < n and attempts < n * 500: attempts += 1 a = rng.randint(lo_a, hi_a) b = rng.randint(lo_b, hi_b) if (a, b) in seen: continue seen.add((a, b)) out.append((a, b)) return out def labels_add(pairs): return np.array([a + b for a, b in pairs], dtype=np.int64) def labels_mul(pairs): return np.array([a * b for a, b in pairs], dtype=np.int64) # ===================================================================== # §3 MLP # ===================================================================== class MLP: def __init__(self, in_dim: int, hidden: int, out_dim: int, seed: int = 0): rng = np.random.default_rng(seed) self.W1 = rng.standard_normal((in_dim, hidden)) * np.sqrt(2.0 / in_dim) self.b1 = np.zeros(hidden) self.W2 = rng.standard_normal((hidden, out_dim)) * np.sqrt(2.0 / hidden) self.b2 = np.zeros(out_dim) self.temperature = 1.0 self.in_dim = in_dim self.hidden = hidden self.out_dim = out_dim def n_params(self) -> int: return sum(p.size for p in self.params()) def params(self): return [self.W1, self.b1, self.W2, self.b2] def forward(self, X): z1 = X @ self.W1 + self.b1 h = np.maximum(z1, 0.0) logits = h @ self.W2 + self.b2 return z1, h, logits def predict_proba(self, X): _, _, logits = self.forward(X) z = logits / max(1e-6, self.temperature) return softmax(z) def predict(self, X): return self.predict_proba(X).argmax(axis=-1) def loss_and_grad(self, X, y): n = len(y) z1, h, logits = self.forward(X) logits = logits - logits.max(axis=1, keepdims=True) e = np.exp(logits) p = e / e.sum(axis=1, keepdims=True) loss = -np.log(p[np.arange(n), y] + 1e-12).mean() dz = p.copy() dz[np.arange(n), y] -= 1.0 dz /= n dW2 = h.T @ dz db2 = dz.sum(axis=0) dh = dz @ self.W2.T dz1 = dh * (z1 > 0.0) dW1 = X.T @ dz1 db1 = dz1.sum(axis=0) return loss, [dW1, db1, dW2, db2] def train(model: MLP, X, y, epochs: int = 800, batch: int = 32, lr: float = 3e-3, seed: int = 0) -> List[float]: rng = np.random.default_rng(seed) m = [np.zeros_like(p) for p in model.params()] v = [np.zeros_like(p) for p in model.params()] t = 0 b1, b2, eps = 0.9, 0.999, 1e-8 losses: List[float] = [] n = len(y) if n == 0: return losses for _ in range(epochs): idx = rng.permutation(n) for start in range(0, n, batch): sel = idx[start:start + batch] loss, grads = model.loss_and_grad(X[sel], y[sel]) t += 1 for i, (p, g) in enumerate(zip(model.params(), grads)): m[i] = b1 * m[i] + (1 - b1) * g v[i] = b2 * v[i] + (1 - b2) * g * g mhat = m[i] / (1 - b1 ** t) vhat = v[i] / (1 - b2 ** t) p -= lr * mhat / (np.sqrt(vhat) + eps) losses.append(float(loss)) return losses # ===================================================================== # §4 Calibration # ===================================================================== T_GRID = np.linspace(0.5, 20.0, 80) def fit_temperature(model: MLP, X, y) -> float: _, _, logits = model.forward(X) best_T, best_nll = 1.0, float("inf") for T in T_GRID: z = logits / T z = z - z.max(axis=1, keepdims=True) e = np.exp(z) p = e / e.sum(axis=1, keepdims=True) nll = -np.log(p[np.arange(len(y)), y] + 1e-12).mean() if nll < best_nll: best_nll, best_T = nll, float(T) return best_T def temperature_hit_boundary(T: float) -> bool: """True iff the fit landed on the edge of the search grid.""" return T >= T_GRID[-1] - 1e-6 or T <= T_GRID[0] + 1e-6 def expected_calibration_error(probs: np.ndarray, labels: np.ndarray, n_bins: int = 10) -> float: conf = probs.max(axis=1) pred = probs.argmax(axis=1) correct = (pred == labels).astype(np.float64) bins = np.linspace(0.0, 1.0, n_bins + 1) total = 0.0 for lo, hi in zip(bins[:-1], bins[1:]): mask = (conf >= lo) & (conf < hi) if mask.sum() == 0: continue total += mask.sum() * abs(conf[mask].mean() - correct[mask].mean()) return float(total / len(labels)) # ===================================================================== # §5 Diagnostic checklist # ===================================================================== @dataclass class SplitReport: name: str n: int accuracy: float mean_confidence: float ece: float ood: bool @dataclass class LearningReport: model_params: int n_train: int train_accuracy: float temperature: float temperature_boundary_hit: bool output_uniformity: float mean_max_prob: float splits: List[SplitReport] = field(default_factory=list) issues: List[str] = field(default_factory=list) verdict: str = "[ok]" def _eval_split(model: MLP, name: str, pairs, labels, ood: bool) -> SplitReport: if len(pairs) == 0: return SplitReport(name=name, n=0, accuracy=0.0, mean_confidence=0.0, ece=0.0, ood=ood) X = encode_pairs(pairs) probs = model.predict_proba(X) preds = probs.argmax(axis=1) return SplitReport( name=name, n=len(pairs), accuracy=float((preds == labels).mean()), mean_confidence=float(probs.max(axis=1).mean()), ece=expected_calibration_error(probs, labels), ood=ood, ) def diagnose(model: MLP, train_pairs, train_labels, interp_pairs, interp_labels, near_ood_pairs, near_ood_labels, far_ood_pairs, far_ood_labels, structural_pairs, structural_labels, ece_threshold: float = 0.20, underconf_threshold: float = 0.40) -> LearningReport: """Build a diagnostic report from raw pairs and labels. All *_pairs arguments are lists of (a, b) tuples. The function encodes them internally; do not pass already-encoded arrays. ece_threshold: flag interpolation ECE above this value. underconf_threshold: flag mean max-prob below this value when accuracy is meaningfully above chance (model is unsure of answers that are mostly right). """ X_tr = encode_pairs(train_pairs) train_acc = float((model.predict(X_tr) == train_labels).mean()) # Output-distribution canary. A model that predicts the same # class on every input has prediction_diversity near 0. A model # that is uniform after temperature scaling has low mean_max_prob. X_all = encode_pairs(list(interp_pairs) + list(far_ood_pairs)) probs_all = model.predict_proba(X_all) class_hist = np.bincount(probs_all.argmax(axis=1), minlength=model.out_dim).astype(np.float64) output_uniformity = normalized_entropy(class_hist) mean_max_prob = float(probs_all.max(axis=1).mean()) splits = [ _eval_split(model, "interpolation", interp_pairs, interp_labels, ood=False), _eval_split(model, "near-OOD", near_ood_pairs, near_ood_labels, ood=True), _eval_split(model, "far-OOD", far_ood_pairs, far_ood_labels, ood=True), _eval_split(model, "structural-OOD", structural_pairs, structural_labels, ood=True), ] issues: List[str] = [] tbh = temperature_hit_boundary(model.temperature) if tbh: issues.append( f"temperature at search boundary (T={model.temperature:.2f}); " "confidence is not meaningful" ) if output_uniformity < 0.05: issues.append( f"degenerate output distribution (uniformity=" f"{output_uniformity:.3f}); model predicts nearly the same " "class on every input" ) if mean_max_prob < 0.10: issues.append( f"mean max-prob is {mean_max_prob:.3f}; predictions are " "near-uniform across classes" ) interp = splits[0] far = splits[2] struct = splits[3] if interp.accuracy < 0.15: issues.append( f"model did not learn the task: interpolation accuracy " f"{interp.accuracy:.2f} is near chance" ) elif interp.accuracy < 0.5: issues.append( f"interpolation accuracy {interp.accuracy:.2f} is weak" ) if interp.accuracy > 0.5 and far.accuracy < 0.2: issues.append( "cannot generalise to OOD inputs: far-OOD accuracy " f"{far.accuracy:.2f}" ) if interp.accuracy > 0.5 and struct.accuracy < 0.2: issues.append( "cannot transfer to a different operation: structural-OOD " f"accuracy {struct.accuracy:.2f}" ) if interp.accuracy > 0.5 and interp.mean_confidence > 0.8 \ and interp.accuracy < 0.7: issues.append( "overconfident on in-distribution data " f"(acc={interp.accuracy:.2f}, conf={interp.mean_confidence:.2f})" ) # Calibration thresholds on the interpolation split. if interp.accuracy > 0.30 and interp.ece > ece_threshold: issues.append( f"poor calibration on interpolation: ECE={interp.ece:.2f} " f"(threshold {ece_threshold:.2f})" ) if (interp.accuracy > 0.50 and interp.mean_confidence < underconf_threshold): issues.append( f"underconfident on interpolation " f"(acc={interp.accuracy:.2f}, " f"conf={interp.mean_confidence:.2f}); " "temperature may be over-corrected" ) # Verdict by number of issues. if not issues: verdict = "[ok]" elif len(issues) == 1: verdict = "[?]" elif len(issues) == 2: verdict = "[??]" else: verdict = "[???]" return LearningReport( model_params=model.n_params(), n_train=len(train_pairs), train_accuracy=train_acc, temperature=model.temperature, temperature_boundary_hit=tbh, output_uniformity=output_uniformity, mean_max_prob=mean_max_prob, splits=splits, issues=issues, verdict=verdict, ) def print_report(r: LearningReport, title: str = "LEARNING REPORT") -> None: print() print("=" * 74) print(title) print("=" * 74) print(f" parameters : {r.model_params}") print(f" training examples : {r.n_train}") print(f" training accuracy : {r.train_accuracy * 100:.1f}%") print(f" temperature : {r.temperature:.2f}" f"{' [BOUNDARY]' if r.temperature_boundary_hit else ''}") print(f" output uniformity : {r.output_uniformity:.3f}") print(f" mean max-prob : {r.mean_max_prob:.3f}") print(f" verdict : {r.verdict}") if r.issues: print() print(" issues:") for issue in r.issues: print(f" - {issue}") print() print(f" {'split':<18} {'n':>4} {'acc':>6} {'mean_conf':>10} " f"{'ECE':>6}") print(" " + "-" * 56) for s in r.splits: tag = "" if not s.ood else " (OOD)" print(f" {s.name:<18} {s.n:>4} {s.accuracy:>6.2f} " f"{s.mean_confidence:>10.2f} {s.ece:>6.2f}{tag}") # ===================================================================== # §6 Training-data-volume calculator # ===================================================================== def estimate_volume_formula(n_params: int, n_classes: int, target_acc: float, samples_per_param: float = 0.05) -> int: """Formula-based estimate of training-set size. Rule of thumb: a model can reliably fit roughly N ~ samples_per_param * P / log2(C) examples and generalise to held-out data from the same distribution. The 0.05 default is calibrated on the demo task (2-layer MLP, bounded classification) and should be re-fit for other architectures. To hit a target accuracy above the near-baseline ceiling, scale by 1 / (1 - target_acc). Returns an integer estimate. Returns -1 if the target exceeds what the formula considers achievable. """ if target_acc <= 0.0 or target_acc >= 1.0: return -1 base = samples_per_param * n_params / max(1.0, math.log2(n_classes)) scale = 1.0 / (1.0 - target_acc) return max(1, int(math.ceil(base * scale))) def _exp_saturation(N: np.ndarray, acc_max: float, N_half: float) -> np.ndarray: return acc_max * (1.0 - np.exp(-N / N_half)) def _hill(N: np.ndarray, acc_max: float, N_half: float, h: float) -> np.ndarray: """Hill cooperativity: acc(N) = acc_max * N^h / (N_half^h + N^h). h > 1 sharpens the transition. h = 1 recovers the Michaelis- Menten form. Useful when the observed ladder shows a step between two adjacent N values rather than a smooth curve. """ return acc_max * (N ** h) / (N_half ** h + N ** h) def fit_saturating_curve(Ns: Sequence[int], accs: Sequence[float], form: str = "exp") -> Tuple: """Fit a saturating curve by grid search. form = "exp" -> (acc_max, N_half, mse) form = "hill" -> (acc_max, N_half, h, mse) Both fits use bounded grids so they cannot wander off into nonsense if the data are flat or noisy. """ Ns_arr = np.array(Ns, dtype=np.float64) accs_arr = np.array(accs, dtype=np.float64) acc_max_grid = np.linspace(0.3, 1.0, 30) N_half_grid = np.exp(np.linspace(np.log(0.5), np.log(2000.0), 50)) if form == "exp": best = (0.5, 10.0, float("inf")) for acc_max in acc_max_grid: for N_half in N_half_grid: preds = _exp_saturation(Ns_arr, acc_max, N_half) mse = float(np.mean((preds - accs_arr) ** 2)) if mse < best[2]: best = (float(acc_max), float(N_half), mse) return best if form == "hill": h_grid = np.linspace(0.5, 8.0, 40) best = (0.5, 10.0, 1.0, float("inf")) for acc_max in acc_max_grid: for N_half in N_half_grid: for h in h_grid: preds = _hill(Ns_arr, acc_max, N_half, h) mse = float(np.mean((preds - accs_arr) ** 2)) if mse < best[3]: best = (float(acc_max), float(N_half), float(h), mse) return best raise ValueError(f"unknown form: {form!r}") def _invert_exp(acc_max: float, N_half: float, target: float) -> Optional[int]: if acc_max <= target: return None return int(math.ceil(-N_half * math.log(1.0 - target / acc_max))) def _invert_hill(acc_max: float, N_half: float, h: float, target: float) -> Optional[int]: if acc_max <= target: return None ratio = target / (acc_max - target) if ratio <= 0.0: return None N = N_half * (ratio ** (1.0 / h)) return int(math.ceil(N)) def estimate_volume_empirical(make_data: Callable, model_factory: Callable, N_grid: Sequence[int] = (16, 32, 64, 128, 256, 512, 1024), epochs: int = 500, seeds: int = 2, target_acc: float = 0.8, form: str = "exp", verbose: bool = True ) -> Tuple[int, Dict]: """Estimate N by actually training and measuring. make_data(n, seed) -> (X_train, y_train, X_val, y_val) model_factory(seed) -> MLP form = "exp" or "hill". Returns (recommended_N, diagnostics). recommended_N is -1 if the fitted ceiling is below target_acc. """ observed_N: List[int] = [] observed_acc: List[float] = [] for N in N_grid: accs = [] for s in range(seeds): X_tr, y_tr, X_val, y_val = make_data(N, seed=s) m = model_factory(seed=100 + s) train(m, X_tr, y_tr, epochs=epochs, batch=min(32, N), seed=s) accs.append(float((m.predict(X_val) == y_val).mean())) mean_acc = float(np.mean(accs)) observed_N.append(N) observed_acc.append(mean_acc) if verbose: print(f" N={N:>5} val_acc={mean_acc:.3f}") if form == "exp": acc_max, N_half, mse = fit_saturating_curve( observed_N, observed_acc, form="exp") recommended = _invert_exp(acc_max, N_half, target_acc) diag = { "form": "exp", "observed_N": observed_N, "observed_acc": observed_acc, "fitted_ceiling": acc_max, "fitted_N_half": N_half, "fit_mse": mse, "recommended_N": -1 if recommended is None else recommended, "target_acc": target_acc, } return diag["recommended_N"], diag if form == "hill": acc_max, N_half, h, mse = fit_saturating_curve( observed_N, observed_acc, form="hill") recommended = _invert_hill(acc_max, N_half, h, target_acc) diag = { "form": "hill", "observed_N": observed_N, "observed_acc": observed_acc, "fitted_ceiling": acc_max, "fitted_N_half": N_half, "fitted_h": h, "fit_mse": mse, "recommended_N": -1 if recommended is None else recommended, "target_acc": target_acc, } return diag["recommended_N"], diag raise ValueError(f"unknown form: {form!r}") # ===================================================================== # §7 Task factories for the demo # ===================================================================== def make_add_data(n_train: int, seed: int): """Return (X_tr, y_tr, X_val, y_val) for the a + b task.""" pairs_tr = sample_pairs(0, TRAIN_MAX, 0, TRAIN_MAX, n_train, seed=seed) pairs_val = sample_pairs(0, TRAIN_MAX, 0, TRAIN_MAX, 200, seed=seed + 999) return (encode_pairs(pairs_tr), labels_add(pairs_tr), encode_pairs(pairs_val), labels_add(pairs_val)) def make_add_model(seed: int = 0) -> MLP: return MLP(2 * MAX_VAL, 32, MAX_SUM, seed=seed) def make_mul_model(seed: int = 0) -> MLP: return MLP(2 * MAX_VAL, 32, MAX_SUM, seed=seed) # ===================================================================== # §8 Self-test # ===================================================================== def self_test(verbose: bool = True) -> Tuple[int, int]: checks = [] # Encoding. v = encode_pair(3, 7) checks.append(("encode: shape", v.shape == (2 * MAX_VAL,))) checks.append(("encode: a bit set", v[3] == 1.0)) checks.append(("encode: b bit set", v[MAX_VAL + 7] == 1.0)) # Training reduces loss. X_tr, y_tr, X_val, y_val = make_add_data(60, seed=0) m = make_add_model(seed=0) losses = train(m, X_tr, y_tr, epochs=200, batch=16, seed=0) checks.append(("train: loss decreases", losses[-1] < losses[0] * 0.5)) # Report structure. Build raw pairs, encode only where needed. # NOTE: diagnose() takes raw pairs, not encoded X. Passing # encoded X would unpack each row and raise ValueError. train_pairs = sample_pairs(0, TRAIN_MAX, 0, TRAIN_MAX, 60, seed=1) X_tr2 = encode_pairs(train_pairs) y_tr2 = labels_add(train_pairs) model = make_add_model(seed=1) train(model, X_tr2, y_tr2, epochs=400, batch=16, seed=1) model.temperature = fit_temperature(model, X_val, y_val) pairs_interp = sample_pairs(0, 9, 0, 9, 30, seed=2) pairs_near = sample_pairs(0, 9, 10, 14, 30, seed=3) pairs_far = sample_pairs(10, 19, 10, 19, 30, seed=4) pairs_struct = sample_pairs(0, 5, 0, 5, 30, seed=5) report = diagnose( model, train_pairs, y_tr2, pairs_interp, labels_add(pairs_interp), pairs_near, labels_add(pairs_near), pairs_far, labels_add(pairs_far), pairs_struct, labels_mul(pairs_struct), ) checks.append(("report: has 4 splits", len(report.splits) == 4)) checks.append(("report: verdict set", report.verdict in ("[ok]", "[?]", "[??]", "[???]"))) checks.append(("report: uniformity in [0,1]", 0.0 <= report.output_uniformity <= 1.0)) checks.append(("report: params > 0", report.model_params > 0)) # New calibration thresholds produce issues when expected. # Build a report where interp is deliberately overcorrected and # check that the underconfidence threshold fires. from dataclasses import replace fake_interp = replace(report.splits[0], accuracy=0.70, mean_confidence=0.30, ece=0.05) fake_far = replace(report.splits[2], accuracy=0.90, mean_confidence=0.90, ece=0.05) fake_struct = replace(report.splits[3], accuracy=0.90, mean_confidence=0.90, ece=0.05) fake_splits = [fake_interp, report.splits[1], fake_far, fake_struct] issues: List[str] = [] if fake_interp.accuracy > 0.30 and fake_interp.ece > 0.20: issues.append("ece") if (fake_interp.accuracy > 0.50 and fake_interp.mean_confidence < 0.40): issues.append("underconf") checks.append(("threshold: ece quiet when low", "ece" not in issues)) checks.append(("threshold: underconf fires when conf<0.40 and acc>0.5", "underconf" in issues)) fake_interp2 = replace(report.splits[0], accuracy=0.70, mean_confidence=0.70, ece=0.30) issues2: List[str] = [] if fake_interp2.accuracy > 0.30 and fake_interp2.ece > 0.20: issues2.append("ece") checks.append(("threshold: ece fires when ece>0.20", "ece" in issues2)) # Formula: monotone in target_acc. n1 = estimate_volume_formula(500, 40, 0.5) n2 = estimate_volume_formula(500, 40, 0.8) checks.append(("formula: higher target -> more data", n2 > n1)) # Formula: monotone in n_params. n3 = estimate_volume_formula(1000, 40, 0.8) checks.append(("formula: more params -> more data", n3 > n2)) # Formula: invalid target. n4 = estimate_volume_formula(500, 40, 1.0) checks.append(("formula: invalid target -> -1", n4 == -1)) # Curve fitting on synthetic saturating data. Ns = [16, 32, 64, 128, 256, 512] true_max, true_half = 0.9, 80.0 accs = [true_max * (1 - math.exp(-n / true_half)) for n in Ns] acc_max, N_half, mse = fit_saturating_curve(Ns, accs, form="exp") checks.append(("curve fit exp: recovers ceiling (>0.85)", acc_max > 0.85)) checks.append(("curve fit exp: low MSE (<0.01)", mse < 0.01)) # Hill fit on a sharp step. Ns_step = [16, 32, 64, 96, 128, 192, 256, 512] accs_step = [0.10, 0.25, 0.55, 0.85, 0.98, 1.00, 1.00, 1.00] hill_out = fit_saturating_curve(Ns_step, accs_step, form="hill") checks.append(("curve fit hill: 4-tuple", len(hill_out) == 4)) checks.append(("curve fit hill: ceiling recovered (>0.90)", hill_out[0] > 0.90)) # Hill fit has lower MSE than exp on the sharp-step data. _, _, exp_mse_step = fit_saturating_curve(Ns_step, accs_step, form="exp") hill_mse = hill_out[3] checks.append(("curve fit hill beats exp on step data", hill_mse < exp_mse_step)) # Temperature boundary detector. checks.append(("boundary: T=20 flagged", temperature_hit_boundary(20.0))) checks.append(("boundary: T=5 not flagged", not temperature_hit_boundary(5.0))) passed = sum(1 for _, ok in checks if ok) if verbose: print() print("=" * 74) print("SELF-TEST") print("=" * 74) for name, ok in checks: mark = "PASS" if ok else "FAIL" print(f" [{mark}] {name}") print() print(f" {passed}/{len(checks)} correct") return passed, len(checks) # ===================================================================== # §9 Demo # ===================================================================== def _section(title: str) -> None: print() print("=" * 74) print(title) print("=" * 74) def demo() -> None: print() print("=" * 74) print("LEARNING REPORT — diagnostic checklist + volume calculator") print("=" * 74) self_test(verbose=True) # --------------------------------------------------------------- _section("PART 0 Train one model on a small dataset") # --------------------------------------------------------------- train_pairs = sample_pairs(0, TRAIN_MAX, 0, TRAIN_MAX, 60, seed=1) X_tr = encode_pairs(train_pairs) y_tr = labels_add(train_pairs) model = make_add_model(seed=0) print(f" task : a + b, a, b in [0, {TRAIN_MAX}]") print(f" training examples : {len(train_pairs)}") print(f" parameters : {model.n_params()}") losses = train(model, X_tr, y_tr, epochs=1000, batch=16, lr=3e-3, seed=0) print(f" loss first / last : {losses[0]:.3f} -> " f"{losses[-1]:.4f}") # Build the splits. interp_pairs = sample_pairs(0, 9, 0, 9, 80, seed=200) near_pairs = sample_pairs(0, 9, 10, 14, 60, seed=201) far_pairs = sample_pairs(10, 19, 10, 19, 60, seed=202) struct_pairs = sample_pairs(0, 5, 0, 5, 36, seed=203) # Fit temperature on a validation split from the training dist. pairs_val = sample_pairs(0, 9, 0, 9, 200, seed=204) X_val = encode_pairs(pairs_val) y_val = labels_add(pairs_val) model.temperature = fit_temperature(model, X_val, y_val) # --------------------------------------------------------------- _section("PART 1 Diagnostic report") # --------------------------------------------------------------- report = diagnose( model, train_pairs, y_tr, interp_pairs, labels_add(interp_pairs), near_pairs, labels_add(near_pairs), far_pairs, labels_add(far_pairs), struct_pairs, labels_mul(struct_pairs), ) print_report(report) print() print(" Reading the report:") print(" - Four splits. The last three are out-of-distribution on") print(" different axes: extended input range, new input range,") print(" and a different operation on the same inputs.") print(" - Temperature boundary. T at the grid edge means the") print(" calibration is not honest, it is degenerate.") print(" - Output uniformity. Predictions collapsing to one class") print(" is not calibration; it is a model that has stopped working.") print(" - Calibration thresholds. ECE above 0.20 on interpolation") print(" and mean max-prob below 0.40 when accuracy is above 0.50") print(" both fire as issues.") # --------------------------------------------------------------- _section("PART 2 Volume calculator, formula") # --------------------------------------------------------------- print(" Formula: N ~ k * P / log2(C) * 1 / (1 - target_acc)") print(" with k = 0.05 calibrated on this task.") print() P = model.n_params() C = MAX_SUM print(f" parameters : {P}") print(f" output classes : {C}") print() print(f" {'target acc':>10} {'estimated N':>12}") print(" " + "-" * 24) for t in (0.30, 0.50, 0.70, 0.80, 0.90): n = estimate_volume_formula(P, C, t) print(f" {t:>10.2f} {n:>12}") # --------------------------------------------------------------- _section("PART 3 Volume calculator, empirical (exp and hill)") # --------------------------------------------------------------- print(" Training the model on a ladder of dataset sizes.") print(" Each N is trained with 2 seeds and evaluated on 200 held-out") print(" examples from the same distribution.") print() def make_data(n_train: int, seed: int): return make_add_data(n_train, seed) def model_factory(seed: int = 0): return make_add_model(seed=seed) N_grid = (16, 32, 64, 128, 256, 512, 1024) print(" --- exponential saturation fit ---") recommended_exp, diag_exp = estimate_volume_empirical( make_data=make_data, model_factory=model_factory, N_grid=N_grid, epochs=500, seeds=2, target_acc=0.80, form="exp", verbose=True, ) print() print(f" fitted ceiling : {diag_exp['fitted_ceiling']:.3f}") print(f" fitted N_half : {diag_exp['fitted_N_half']:.1f}") print(f" fit MSE : {diag_exp['fit_mse']:.4f}") if recommended_exp == -1: print(f" target {diag_exp['target_acc']:.2f} unreachable") else: print(f" recommended N (exp) : {recommended_exp}") print() print(" --- hill cooperativity fit ---") recommended_hill, diag_hill = estimate_volume_empirical( make_data=make_data, model_factory=model_factory, N_grid=N_grid, epochs=500, seeds=2, target_acc=0.80, form="hill", verbose=False, ) print(f" fitted ceiling : " f"{diag_hill['fitted_ceiling']:.3f}") print(f" fitted N_half : {diag_hill['fitted_N_half']:.1f}") print(f" fitted h : {diag_hill['fitted_h']:.2f}") print(f" fit MSE : {diag_hill['fit_mse']:.4f}") if recommended_hill == -1: print(f" target {diag_hill['target_acc']:.2f} unreachable") else: print(f" recommended N (hill) : {recommended_hill}") print() print(" Comparing fits:") print(f" exp MSE : {diag_exp['fit_mse']:.4f}") print(f" hill MSE : {diag_hill['fit_mse']:.4f}") if diag_hill['fit_mse'] < diag_exp['fit_mse']: print(" hill fits better -- the ladder has a step, not a") print(" smooth curve. h > 1 is the signature of cooperativity.") else: print(" exp fits as well or better -- the ladder is smooth.") print() print(" Recommendation comparison:") print(f" formula : " f"{estimate_volume_formula(P, C, 0.80)}") print(f" empirical exp : {recommended_exp}") print(f" empirical hill : {recommended_hill}") # --------------------------------------------------------------- _section("PART 4 Verify the recommendations") # --------------------------------------------------------------- candidates = [x for x in (estimate_volume_formula(P, C, 0.80), recommended_exp, recommended_hill) if x != -1] for N in sorted(set(candidates)): X_tr2, y_tr2, X_val2, y_val2 = make_add_data(N, seed=77) m2 = make_add_model(seed=77) train(m2, X_tr2, y_tr2, epochs=800, batch=min(64, N), lr=3e-3, seed=77) acc = float((m2.predict(X_val2) == y_val2).mean()) status = "target met" if acc >= 0.80 else "target missed" print(f" N={N:>5} held-out acc={acc:.3f} " f"target=0.800 {status}") # --------------------------------------------------------------- _section("PART 5 The discipline") # --------------------------------------------------------------- print(""" Before claiming a model learned a task: 1. Split the evaluation along at least four axes. - interpolation: held-out samples from the training distribution - near-OOD: same operand range, one axis extended - far-OOD: all operands outside the training range - structural-OOD: same operands, different operation 2. Report mean confidence and ECE on every split. 3. Check temperature-boundary. A fitted T at the edge of the search grid means calibration is not honest, it is degenerate. 4. Check output uniformity. Predictions collapsing to one class, or uniform across classes, is not calibration. 5. Check calibration thresholds. ECE above 0.20 or mean max-prob below 0.40 (with accuracy above 0.50) are both real issues, not rounding artifacts. 6. Estimate the training-set size needed for the target accuracy before running the experiment. Use both the formula and the empirical ladder, and compare exp and hill curve families -- a step-shaped ladder is not the same shape as a smooth one, and the recommended N differs. The bug class is a single number: "our model achieves X%". The fix is a report that cannot be summarised by any one number, and a calculator that tells you how many examples the target required in the first place. """) # ===================================================================== # §10 Entry point # ===================================================================== if __name__ == "__main__": demo()