"""Full tic-tac-toe experiment (deliverable E). - 19683 fields, two targets: 'simple' = extension of the task statement's hand-written certificate; 'legal' = reachable-terminal classification (X/O/Draw for finished reachable games, everything else Illegal). - Train on ~25% of fields (orbit-stratified and IID variants), learn a certificate incrementally, recompress, predict the rest. - Controls: majority class, 1-NN (Hamming on raw cells), random-label control (same pipeline on shuffled labels), memorization cost. Run: python3 experiments/exp_ttt.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 IncrementalLearner, choice from certlang.tictactoe import (N_FIELDS, RESULTS, TTTData, build_domain, certificate_predictions, labels_legal, labels_simple, orbit_stratified_sample, orbits) OUT = os.path.join(os.path.dirname(__file__), "..", "results") os.makedirs(OUT, exist_ok=True) def knn1(train_idx, labels, test_idx, fields): """1-NN with Hamming distance on the 9 raw cells (majority over ties).""" tr = fields[train_idx].astype(np.int16) trlab = labels[train_idx] preds = np.empty(len(test_idx), dtype=np.int8) chunk = 512 for s in range(0, len(test_idx), chunk): te = fields[test_idx[s:s + chunk]].astype(np.int16) dist = (te[:, None, :] != tr[None, :, :]).sum(axis=2) dmin = dist.min(axis=1) for i in range(len(te)): cand = trlab[dist[i] == dmin[i]] preds[s + i] = np.bincount(cand, minlength=4).argmax() return preds def run_learner(domain, train_idx, labels, order_seed, max_atoms=4, node_budget=400_000, recompress=True, tag=""): rng = random.Random(order_seed) order = list(train_idx) rng.shuffle(order) learner = IncrementalLearner(domain, max_atoms=max_atoms, node_budget=node_budget) t0 = time.time() for x in order: learner.observe(int(x), int(labels[x])) t_learn = time.time() - t0 raw_cost = learner.cert.cost() raw_rules = len(learner.cert.rules) t0 = time.time() if recompress: learner.recompress(orders=2, rng=random.Random(order_seed + 1)) t_rec = time.time() - t0 print(f" [{tag}] repairs={learner.n_repairs} raw_cost={raw_cost:.1f} " f"({raw_rules} rules) -> recompressed={learner.cert.cost():.1f} " f"({len(learner.cert.rules)} rules) " f"[learn {t_learn:.1f}s, recompress {t_rec:.1f}s]") return learner, {"repairs": learner.n_repairs, "raw_cost": raw_cost, "raw_rules": raw_rules, "cost": learner.cert.cost(), "rules": len(learner.cert.rules), "t_learn": t_learn, "t_recompress": t_rec} def evaluate(cert, labels, train_mask): pred = certificate_predictions(cert) test = ~train_mask acc_train = float((pred[train_mask] == labels[train_mask]).mean()) acc_test = float((pred[test] == labels[test]).mean()) per_class = {} for y in range(4): sel = test & (labels == y) per_class[RESULTS[y]] = { "n": int(sel.sum()), "acc": float((pred[sel] == labels[sel]).mean()) if sel.sum() else None} return pred, {"train_acc": acc_train, "test_acc": acc_test, "test_per_class": per_class, "n_test_errors": int((pred[test] != labels[test]).sum())} def main(): t0 = time.time() d = TTTData() domain = build_domain(d) print(f"|X| = {N_FIELDS}, atoms after semantic dedupe = {len(domain.atoms)}") targets = {"simple": labels_simple(d), "legal": labels_legal(d)} for name, lab in targets.items(): counts = {RESULTS[y]: int((lab == y).sum()) for y in range(4)} print(f"target '{name}' class counts: {counts}") # sanity: reachable-terminal tic-tac-toe counts are known legal_counts = {RESULTS[y]: int((targets['legal'] == y).sum()) for y in range(4)} assert legal_counts["Draw"] == 16, legal_counts expect = {"XWon": 626, "OWon": 316, "Draw": 16} for k, v in expect.items(): if legal_counts[k] != v: print(f" WARNING: {k}={legal_counts[k]} differs from literature {v}") else: print(f" sanity ok: {k} = {v} (matches known reachable-terminal count)") orbs = orbits(d) print(f"D4 orbits: {len(orbs)} (Burnside check: " f"{sum(len(o) for o in orbs)} fields total)") report = {"n_atoms": len(domain.atoms), "n_orbits": len(orbs), "targets": {}} for tname, labels in targets.items(): print(f"\n=== target '{tname}' ===") treport = {} train_frac = 0.25 samplings = ("orbit", "iid", "orbit+rare") if tname == "legal" \ else ("orbit", "iid") for sname in samplings: rng = random.Random(0) if sname == "orbit": train_idx = orbit_stratified_sample(d, labels, train_frac, rng) elif sname == "orbit+rare": # adequate representation of rare structural classes: sample # 75% of the orbits of any class with fewer than 100 members # (still holding some orbits out for testing) train_idx = set(orbit_stratified_sample(d, labels, train_frac, rng)) rare = [y for y in range(4) if (labels == y).sum() < 100] for y in rare: orbs = [o for o in orbits(d) if labels[o[0]] == y] rng.shuffle(orbs) for o in orbs[:max(1, round(0.75 * len(orbs)))]: train_idx |= set(o) train_idx = sorted(train_idx) else: train_idx = rng.sample(range(N_FIELDS), 5000) train_mask = np.zeros(N_FIELDS, dtype=bool) train_mask[train_idx] = True tr_counts = {RESULTS[y]: int((labels[train_idx] == y).sum()) for y in range(4)} print(f"-- sampling={sname}: {len(train_idx)} train fields, " f"class counts {tr_counts}") learner, lstats = run_learner(domain, train_idx, labels, order_seed=1, tag=f"{tname}/{sname}") pred, estats = evaluate(learner.cert, labels, train_mask) print(f" test accuracy = {estats['test_acc']*100:.2f}% " f"({estats['n_test_errors']} errors on " f"{N_FIELDS - len(train_idx)} unseen fields)") for cname, st in estats["test_per_class"].items(): if st["n"]: print(f" {cname:8s} n={st['n']:6d} acc={st['acc']*100:6.2f}%") # baselines maj = np.bincount(labels[train_idx], minlength=4).argmax() maj_acc = float((labels[~train_mask] == maj).mean()) test_idx = np.where(~train_mask)[0] t1 = time.time() knn_pred = knn1(np.array(train_idx), labels, test_idx, d.fields) knn_acc = float((knn_pred == labels[test_idx]).mean()) print(f" baselines: majority={maj_acc*100:.2f}% " f"1-NN={knn_acc*100:.2f}% [{time.time()-t1:.0f}s] " f"memorize-train={2.0*len(train_idx):.0f} bits, " f"memorize-all={2.0*N_FIELDS:.0f} bits, " f"certificate={learner.cert.cost():.1f} bits") treport[sname] = { "n_train": len(train_idx), "train_class_counts": tr_counts, "learner": lstats, **estats, "majority_acc": maj_acc, "knn_acc": knn_acc, "memorize_train_bits": 2.0 * len(train_idx), "memorize_all_bits": 2.0 * N_FIELDS, } if sname == "orbit": cert_text = learner.cert.explain(RESULTS) print("\n Learned certificate (recompressed):") for line in cert_text.split("\n"): print(" ", line) treport["certificate"] = cert_text report["targets"][tname] = treport # ---- train-size sweep on the legal target ------------------------------ print("\n=== train-size sweep (target 'legal', orbit sampling) ===") labels = targets["legal"] sweep = [] for frac in (0.025, 0.05, 0.125, 0.25): rng = random.Random(0) idx = orbit_stratified_sample(d, labels, frac, rng) tm = np.zeros(N_FIELDS, dtype=bool) tm[idx] = True learner, lstats = run_learner(domain, idx, labels, order_seed=1, tag=f"sweep {frac:.3f}") _, estats = evaluate(learner.cert, labels, tm) print(f" frac={frac:.3f} n={len(idx):5d}: cost={lstats['cost']:8.1f} bits " f"({lstats['cost']/len(idx):.3f} b/ex), test acc " f"{estats['test_acc']*100:.2f}%") sweep.append({"frac": frac, "n_train": len(idx), **lstats, "test_acc": estats["test_acc"]}) report["sweep_legal_orbit"] = sweep # ---- random-label control ---------------------------------------------- print("\n=== random-label control (target 'legal' marginals) ===") rng = random.Random(5) perm = np.array(rng.sample(range(N_FIELDS), N_FIELDS)) shuffled = targets["legal"][perm] # same marginals, structure destroyed controls = [] for nc in (250, 500, 1000): ctrl_idx = rng.sample(range(N_FIELDS), nc) learner, lstats = run_learner(domain, ctrl_idx, shuffled, order_seed=2, max_atoms=3, node_budget=60_000, recompress=True, tag=f"control {nc}") train_mask = np.zeros(N_FIELDS, dtype=bool) train_mask[ctrl_idx] = True pred, estats = evaluate(learner.cert, shuffled, train_mask) bits_per_example = learner.cert.cost() / nc maj_acc = float((shuffled[~train_mask] == np.bincount(shuffled[ctrl_idx]).argmax()).mean()) print(f" control n={nc}: {learner.cert.cost():.0f} bits " f"({bits_per_example:.2f} bits/example), test acc " f"{estats['test_acc']*100:.2f}% (majority {maj_acc*100:.2f}%)") controls.append({"n_train": nc, **lstats, **estats, "bits_per_example": bits_per_example, "majority_acc": maj_acc}) report["random_control"] = controls with open(os.path.join(OUT, "ttt_results.json"), "w") as fh: json.dump(report, fh, indent=2) print(f"\ntotal time {time.time()-t0:.0f}s; " f"wrote {os.path.join(OUT, 'ttt_results.json')}") if __name__ == "__main__": main()