learning-report / learning_report.py
zeechimp's picture
Upload learning_report.py
7b9a704 verified
Raw History Blame Contribute Delete
37.4 kB
#!/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()