"""Active self-falsifying sampling (failure-mode study). The <100% failure mode of the learner is precisely "accidental predicates": cheap rules consistent with the observed sample whose extensions contain unobserved counterexamples. Since the learner *knows its own rules' extensions*, the natural fix is to sample the next labels from inside them: every accidental rule then gets falsified quickly, while every correct rule just gets confirmed (at no cost growth). Protocol: start from a small IID seed; in each round, learn, then query q unobserved points from each current rule's extension plus q random points (exploration). Compare with plain IID at the same total label budget. Run: python3 experiments/exp_active.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 learn_batch from certlang.tictactoe import (N_FIELDS, RESULTS, TTTData, build_domain, certificate_predictions, labels_legal, mask_to_bits) OUT = os.path.join(os.path.dirname(__file__), "..", "results") os.makedirs(OUT, exist_ok=True) def evaluate(cert, labels, obs): pred = certificate_predictions(cert) mask = np.zeros(N_FIELDS, dtype=bool) mask[list(obs)] = True test = ~mask per_class = {RESULTS[y]: float((pred[test & (labels == y)] == y).mean()) if (test & (labels == y)).sum() else None for y in range(4)} return float((pred[test] == labels[test]).mean()), per_class def active_run(domain, labels, seed_n=200, rounds=12, q=8, seed=0, discriminative=False): rng = random.Random(seed) obs = {int(x): int(labels[x]) for x in rng.sample(range(N_FIELDS), seed_n)} history = [] for t in range(rounds): learner = learn_batch(domain, obs, node_budget=300_000) acc, per_class = evaluate(learner.cert, labels, obs) history.append({"round": t, "labels": len(obs), "acc": acc, "bits": learner.cert.cost(), "rules": len(learner.cert.rules), "per_class": per_class}) # query inside each rule's extension (self-falsification) ... new = set() for r in learner.cert.rules: pts = np.nonzero(mask_to_bits(r.mask, N_FIELDS))[0] cand = [int(p) for p in pts if int(p) not in obs] rng.shuffle(cand) new.update(cand[:q]) # ... plus, in discriminative mode, points where a rule DISAGREES # with its own cheaper generalizations (drop-one-atom variants): # query-by-committee over the cost lattice. If the bolder rule is # true, these labels let broaden/merge adopt it; if not, they refute # it once instead of leaving it forever untested. if discriminative: for r in learner.cert.rules: if len(r.atoms) < 2: continue for k in range(len(r.atoms)): bmask = domain.universe for i, a in enumerate(r.atoms): if i != k: bmask &= a.mask diff = bmask & ~r.mask if diff == 0: continue pts = np.nonzero(mask_to_bits(diff, N_FIELDS))[0] cand = [int(p) for p in pts if int(p) not in obs] rng.shuffle(cand) new.update(cand[:max(2, q // 2)]) # ... plus random exploration pool = [x for x in rng.sample(range(N_FIELDS), 4 * q) if x not in obs] new.update(pool[:q]) for x in new: obs[x] = int(labels[x]) learner = learn_batch(domain, obs, node_budget=300_000) acc, per_class = evaluate(learner.cert, labels, obs) history.append({"round": rounds, "labels": len(obs), "acc": acc, "bits": learner.cert.cost(), "rules": len(learner.cert.rules), "per_class": per_class}) return history, learner def iid_run(domain, labels, n, seed=0): rng = random.Random(seed) obs = {int(x): int(labels[x]) for x in rng.sample(range(N_FIELDS), n)} learner = learn_batch(domain, obs, node_budget=300_000) acc, per_class = evaluate(learner.cert, labels, obs) return {"labels": n, "acc": acc, "bits": learner.cert.cost(), "rules": len(learner.cert.rules), "per_class": per_class} def main(): d = TTTData() domain = build_domain(d) labels = labels_legal(d) report = {"active": [], "iid": []} for disc in (False, True): mode = "active+discriminative" if disc else "active" print(f"=== {mode} sampling, target 'legal' ===") report[mode] = [] for seed in (0, 1, 2): t0 = time.time() history, learner = active_run(domain, labels, seed=seed, discriminative=disc) print(f"\n-- seed {seed} ({time.time()-t0:.0f}s):") for h in history: draw = h["per_class"].get("Draw") print(f" round {h['round']:2d}: labels={h['labels']:4d} " f"acc={100*h['acc']:6.2f}% bits={h['bits']:7.1f} " f"rules={h['rules']:2d} draw_acc=" f"{'-' if draw is None else f'{100*draw:.0f}%'}") report[mode].append(history) if seed == 0: print("\n final certificate:") for line in learner.cert.explain(RESULTS).split("\n"): print(" ", line) print("\n=== IID baseline at matched budgets ===") budgets = sorted({h["labels"] for mode in ("active", "active+discriminative") for hist in report[mode] for h in hist})[::2] for n in budgets: accs, draws = [], [] for seed in (0, 1, 2): r = iid_run(domain, labels, n, seed=seed) accs.append(r["acc"]) if r["per_class"].get("Draw") is not None: draws.append(r["per_class"]["Draw"]) report["iid"].append({"seed": seed, **r}) print(f" n={n:4d}: acc={100*np.mean(accs):6.2f}% " f"(draw {100*np.mean(draws) if draws else float('nan'):.0f}%)") with open(os.path.join(OUT, "active_results.json"), "w") as fh: json.dump(report, fh, indent=2) print(f"\nwrote {os.path.join(OUT, 'active_results.json')}") if __name__ == "__main__": main()