"""Poker-hand classification: a harder problem, and the NN comparison. 2,598,960 hands, 10 classes with extreme imbalance (4 royal flushes vs 1.3M high cards) — a standard stress case where MLPs plateau and miss the rare classes entirely. We train: - the structural certificate learner, on (a) class-stratified samples (k examples per class) and (b) realistic random deals; - a numpy MLP (2x64 ReLU, one-hot 85 input as in the UCI poker-hand setup), on the same samples, with class-weighted loss for fairness; - an MLP on tic-tac-toe 'legal' as a second comparison point. Everything is evaluated on ALL remaining hands. Run: python3 experiments/exp_poker.py (~15-30 min, ~600 MB) """ 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.mlp import MLP from certlang import poker from certlang import tictactoe as ttt OUT = os.path.join(os.path.dirname(__file__), "..", "results") os.makedirs(OUT, exist_ok=True) def onehot_hands(hands: np.ndarray) -> np.ndarray: """UCI-style encoding: per card, one-hot rank (13) + one-hot suit (4).""" n = len(hands) X = np.zeros((n, 85), dtype=np.float32) for j in range(5): X[np.arange(n), j * 17 + hands[:, j] % 13] = 1 X[np.arange(n), j * 17 + 13 + hands[:, j] // 13] = 1 return X def eval_all(pred, labels, train_idx, nclass, names): mask = np.zeros(len(labels), dtype=bool) mask[train_idx] = True test = ~mask overall = float((pred[test] == labels[test]).mean()) per = {} for y in range(nclass): sel = test & (labels == y) per[names[y]] = {"n": int(sel.sum()), "acc": float((pred[sel] == y).mean()) if sel.sum() else None} return overall, per def print_per_class(per): parts = [] for name, st in per.items(): if st["n"]: parts.append(f"{name} {100*st['acc']:.1f}% (n={st['n']})") print(" " + "; ".join(parts)) def stratified_sample(labels, k, rng, nclass): idx = [] for y in range(nclass): pool = np.nonzero(labels == y)[0] take = min(k, len(pool)) idx.extend(int(x) for x in rng.sample(list(pool), take)) return idx def main(): t0 = time.time() print("building poker domain (2,598,960 hands) ...") d = poker.PokerData() labels = labels = poker.labels_poker(d) counts = np.bincount(labels, minlength=10).tolist() assert counts == poker.KNOWN_COUNTS, counts print(f" class counts match the combinatorial literature: {counts}") domain = poker.build_domain(d) print(f" {len(domain.atoms)} atoms after semantic dedupe " f"[{time.time()-t0:.0f}s]") report = {"class_counts": counts, "n_atoms": len(domain.atoms), "certificate": {}, "mlp": {}} # ---- structural learner -------------------------------------------------- for mode, sizes in (("stratified", (5, 20, 100)), ("random", (1000, 5000))): for k in sizes: rng = random.Random(k) if mode == "stratified": train_idx = stratified_sample(labels, k, rng, 10) else: train_idx = rng.sample(range(poker.N_HANDS), k) obs = {int(x): int(labels[x]) for x in train_idx} t1 = time.time() learner = learn_batch(domain, obs, max_atoms=4, node_budget=200_000) dt = time.time() - t1 pred = np.full(poker.N_HANDS, learner.cert.default, dtype=np.int8) for r in learner.cert.rules: pred[poker.mask_to_bits(r.mask)] = r.y overall, per = eval_all(pred, labels, train_idx, 10, poker.CLASSES) tag = f"{mode}-{k}" print(f"\n certificate [{tag}]: n={len(train_idx)}, " f"{learner.cert.cost():.1f} bits, " f"{len(learner.cert.rules)} rules, " f"test acc {100*overall:.3f}% [{dt:.0f}s]") print_per_class(per) report["certificate"][tag] = { "n_train": len(train_idx), "bits": learner.cert.cost(), "rules": len(learner.cert.rules), "acc": overall, "per_class": per, "seconds": dt} if tag == "stratified-20": print("\n learned certificate (stratified-20):") for line in learner.cert.explain(poker.CLASSES).split("\n"): print(" ", line) report["certificate"][tag]["text"] = \ learner.cert.explain(poker.CLASSES) # ---- MLP baseline -------------------------------------------------------- print("\n=== MLP baselines (2x64 ReLU, one-hot 85, class-weighted) ===") freq = np.array(counts, dtype=np.float64) cw = (freq.sum() / (10 * freq)).astype(np.float32) for tag, k, epochs in (("stratified-100", 100, 400), ("random-5000", 5000, 120), ("random-25000", 25000, 60)): rng = random.Random(1000 + k) if tag.startswith("stratified"): train_idx = stratified_sample(labels, k, rng, 10) weight = None # already balanced else: train_idx = rng.sample(range(poker.N_HANDS), k) weight = cw Xtr = onehot_hands(d.hands[train_idx]) ytr = labels[np.array(train_idx)].astype(np.int64) t1 = time.time() net = MLP(85, 10, hidden=64, seed=0) net.fit(Xtr, ytr, epochs=epochs, class_weight=weight) pred = np.empty(poker.N_HANDS, dtype=np.int8) chunk = 200_000 for s in range(0, poker.N_HANDS, chunk): pred[s:s + chunk] = net.predict(onehot_hands(d.hands[s:s + chunk])) overall, per = eval_all(pred, labels, train_idx, 10, poker.CLASSES) dt = time.time() - t1 print(f"\n MLP [{tag}]: n={len(train_idx)}, " f"{net.n_params} params (~{net.n_params*32} bits), " f"test acc {100*overall:.3f}% [{dt:.0f}s]") print_per_class(per) report["mlp"][tag] = {"n_train": len(train_idx), "params": net.n_params, "acc": overall, "per_class": per, "seconds": dt} # ---- MLP on tic-tac-toe legal, same split as the main experiment --------- print("\n=== MLP on tic-tac-toe 'legal' (orbit 25% train) ===") td = ttt.TTTData() tl = ttt.labels_legal(td) rng = random.Random(0) tidx = ttt.orbit_stratified_sample(td, tl, 0.25, rng) Xall = np.zeros((ttt.N_FIELDS, 27), dtype=np.float32) for p in range(9): for v in range(3): Xall[:, p * 3 + v] = td.fields[:, p] == v freq = np.bincount(tl[tidx], minlength=4).astype(np.float64) cw4 = (freq.sum() / (4 * np.maximum(freq, 1))).astype(np.float32) net = MLP(27, 4, hidden=64, seed=0) net.fit(Xall[tidx], tl[np.array(tidx)].astype(np.int64), epochs=200, class_weight=cw4) pred = net.predict(Xall) overall, per = eval_all(pred, tl, tidx, 4, ttt.RESULTS) print(f" MLP: test acc {100*overall:.3f}% " f"({net.n_params} params ~ {net.n_params*32} bits) " f"vs certificate 102.7 bits / 99.65%") print_per_class(per) report["mlp_ttt"] = {"acc": overall, "per_class": per, "params": net.n_params} with open(os.path.join(OUT, "poker_results.json"), "w") as fh: json.dump(report, fh, indent=2) print(f"\ntotal {time.time()-t0:.0f}s; wrote results/poker_results.json") if __name__ == "__main__": main()