"""Sample-efficiency study: how learning efficiency (dataset size, time) depends on (1) the amount of structure the type provides and (2) the complexity of the true function. A. Type ablation (target fixed = 'legal'): learn with atom libraries derived from progressively less of the type structure: full = cell + count + lineprof + profcount + cmp no-cmp = cell + count + lineprof + profcount no-lines = cell + count + cmp (no line family at all) cells = cell only (bare positions; pure pattern DL) Learning curves: held-out accuracy and certificate bits vs n (IID samples). B. Function complexity (structure fixed = full): planted random certificates of k rules; measure the n needed to reach 99% and the Occam product err(n) * n against the planted description length. All learning uses the deterministic order-invariant batch learner, so each (sample, domain) pair maps to exactly one certificate. Run: python3 experiments/exp_curves.py """ import json import os import random import sys import time import numpy as np sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src")) from certlang.core import Certificate, Rule, learn_batch, make_rule from certlang.tictactoe import (ALL_FAMILIES, N_FIELDS, RESULTS, TTTData, build_domain, certificate_predictions, labels_legal) OUT = os.path.join(os.path.dirname(__file__), "..", "results") os.makedirs(OUT, exist_ok=True) ABLATIONS = { "full": ALL_FAMILIES, "no-cmp": ("cell", "count", "lineprof", "profcount"), "no-lines": ("cell", "count", "cmp"), "cells": ("cell",), } # poorer structure => bigger certificates => recompression ~R^2 gets slow, so # the low-structure ablations get smaller grids / fewer seeds GRID = { "full": [(n, (0, 1, 2)) for n in (125, 250, 500, 1000, 2000, 4000, 8000)], "no-cmp": [(125, (0, 1, 2)), (250, (0, 1, 2)), (500, (0, 1, 2)), (1000, (0, 1, 2)), (2000, (0,)), (4000, (0,))], "no-lines": [(125, (0, 1, 2)), (250, (0, 1, 2)), (500, (0, 1, 2)), (1000, (0, 1, 2)), (2000, (0,)), (4000, (0,))], "cells": [(125, (0, 1, 2)), (250, (0, 1, 2)), (500, (0,)), (1000, (0,)), (2000, (0,))], } def run_one(domain, labels, n, seed, node_budget=300_000): rng = random.Random(seed * 100003 + n) idx = rng.sample(range(N_FIELDS), n) obs = {int(x): int(labels[x]) for x in idx} t0 = time.time() learner = learn_batch(domain, obs, max_atoms=4, node_budget=node_budget) dt = time.time() - t0 mask = np.zeros(N_FIELDS, dtype=bool) mask[idx] = True pred = certificate_predictions(learner.cert) acc = float((pred[~mask] == labels[~mask]).mean()) return {"n": n, "seed": seed, "acc": acc, "bits": learner.cert.cost(), "rules": len(learner.cert.rules), "repairs": learner.n_repairs, "seconds": round(dt, 2)} def plant_certificate(domain, k, rng): """Random k-rule certificate over the full atom library; rejected if any class or rule is degenerate on the domain.""" while True: cert = Certificate(0, 4) for _ in range(k): m = rng.randrange(1, 4) atoms = tuple(rng.sample(domain.atoms, m)) cert.rules.append(make_rule(atoms, rng.randrange(4), domain.universe)) pred = certificate_predictions(cert) sizes = np.bincount(pred, minlength=4) touched = sum(1 for r in cert.rules if r.mask) if touched == k and sizes.max() < 0.995 * N_FIELDS: return cert, pred def main(): d = TTTData() labels = labels_legal(d) report = {"ablation": {}, "planted": []} out_path = os.path.join(OUT, "curves_results.json") def dump(): with open(out_path, "w") as fh: json.dump(report, fh, indent=2) print("=== A. type ablation, target 'legal', IID sampling ===") for name, fams in ABLATIONS.items(): domain = build_domain(d, fams) print(f"\n-- structure '{name}': {len(domain.atoms)} atoms") rows = [] for n, seeds in GRID[name]: for seed in seeds: r = run_one(domain, labels, n, seed, node_budget=300_000 if name == "full" else 120_000) rows.append(r) accs = [r["acc"] for r in rows if r["n"] == n] bits = [r["bits"] for r in rows if r["n"] == n] secs = [r["seconds"] for r in rows if r["n"] == n] print(f" n={n:5d}: acc {100*np.mean(accs):6.2f}% " f"(+-{100*np.std(accs):.2f}) bits {np.mean(bits):8.1f} " f"time {np.mean(secs):6.1f}s") report["ablation"].setdefault(name, { "n_atoms": len(domain.atoms), "rows": rows}) dump() report["ablation"][name] = {"n_atoms": len(domain.atoms), "rows": rows} dump() print("\n=== B. planted certificates of graded complexity (full type) ===") domain = build_domain(d) rng = random.Random(99) for k in (1, 2, 4, 8, 16): for draw in range(2): cert, pred = plant_certificate(domain, k, rng) kappa = cert.cost() # what the learner itself would compress the full table to is the # fair complexity reference; the planted cost is an upper bound entry = {"k": k, "draw": draw, "kappa_planted": kappa, "curve": []} for n in (250, 1000, 4000): r = run_one(domain, pred, n, seed=7) r["occam_err_times_n"] = (1 - r["acc"]) * n entry["curve"].append(r) report["planted"].append(entry) dump() cs = ", ".join(f"n={c['n']}: {100*c['acc']:.1f}%" f" ({c['bits']:.0f}b)" for c in entry["curve"]) print(f" k={k:2d} draw {draw}: planted {kappa:6.1f} bits | {cs}") dump() print(f"\nwrote {out_path}") if __name__ == "__main__": main()