""" Tsetlin Machine hyperparameter sweep — smart coarse-then-zoom sampling. Reuses RTM logic from backtest_tsetlin.py, parameterized. Strategy: Phase 1: Pre-captured coarse grid (20 configs, all axes explored). Captured from background run on 2500 rows; avoids re-running slow 200/500-clause configs that take 15-40s each in pure Python. Phase 2: Zoom — vary one axis at a time from top-5, cap clauses ≤ 100 to stay fast. Run on full 2500 rows. Phase 3: Final eval of top-3 on full data (clauses ≤ 200). Output: stdout table + sweep_tsetlin_results.txt Baseline: MAE 12.24 (linear extrapolation, from backtest report). """ import csv import math import os import random import sys import time # ── paths ───────────────────────────────────────────────────────────────────── DATA_DIR = os.path.join(os.path.dirname(__file__), "..") CSV_FILES = [ os.path.join(DATA_DIR, "data/target_battle_1_decimal.csv"), os.path.join(DATA_DIR, "out/data/target_battle_2_decimal.csv"), os.path.join(DATA_DIR, "out/data/target_battle_3_decimal.csv"), os.path.join(DATA_DIR, "out/data/target_battle_4_decimal.csv"), os.path.join(DATA_DIR, "out/data/target_battle_5_decimal.csv"), ] OUT_PATH = os.path.join(os.path.dirname(__file__), "sweep_tsetlin_results.txt") POWER_LEVELS = [0.10, 0.42, 0.74, 1.07, 1.39, 1.71, 2.03, 2.36, 2.68, 3.00] POWER_STRS = ["p0.10","p0.42","p0.74","p1.07","p1.39","p1.71","p2.03","p2.36","p2.68","p3.00"] MAX_DIST = 1414.0 FIELD_BITS = [ ("bearing_sin", 199, 8), ("bearing_cos", 199, 8), ("distance", 99, 7), ("velocity", 31, 5), ("heading_sin", 199, 8), ("heading_cos", 199, 8), ("enemy_x", 99, 7), ("enemy_y", 99, 7), ("enemy_energy", 1000, 10), ] BITS_PER_FRAME = sum(b for _, _, b in FIELD_BITS) # 68 N_FRAMES = 4 N_FEATURES = BITS_PER_FRAME * N_FRAMES # 272 N_LITERALS = N_FEATURES * 2 # 544 RESID_MAX = 30.0 LINEAR_BASELINE_MAE = 12.24 # ── pre-captured coarse results (from background run, 2500 rows, seed=42) ──── # format: (clauses, T, s, states, weighted, mae_all, mae_last, curve, mem_kb) COARSE_PRE = [ # axis: clauses (T=50 s=3.0 states=64 unweighted) (10, 50, 3.0, 64, False, 3.85, 0.18, 8.4, 10.6), (20, 50, 3.0, 64, False, 4.04, 0.29, 8.9, 21.2), (50, 50, 3.0, 64, False, 4.22, 0.37, 8.9, 53.1), (100, 50, 3.0, 64, False, 4.50, 0.62, 9.5, 106.2), (200, 50, 3.0, 64, False, 4.72, 0.69, 9.9, 212.5), (500, 50, 3.0, 64, False, 7.68, 0.74, 18.3, 531.2), # axis: T (clauses=100 s=3.0 states=64 unweighted) (100, 10, 3.0, 64, False, 4.59, 0.58, 9.5, 106.2), (100, 25, 3.0, 64, False, 4.50, 0.62, 9.5, 106.2), (100,100, 3.0, 64, False, 4.50, 0.62, 9.5, 106.2), (100,200, 3.0, 64, False, 4.50, 0.62, 9.5, 106.2), # axis: s (clauses=100 T=50 states=64 unweighted) (100, 50, 1.5, 64, False, 4.48, 0.77, 8.5, 106.2), (100, 50, 5.0, 64, False, 4.57, 0.45, 10.0, 106.2), (100, 50,10.0, 64, False, 5.07, 1.22, 9.8, 106.2), (100, 50,15.0, 64, False, 5.16, 0.69, 10.8, 106.2), # axis: states (clauses=100 T=50 s=3.0 unweighted) (100, 50, 3.0, 32, False, 4.90, 0.90, 10.1, 106.2), (100, 50, 3.0, 128, False, 4.30, 0.47, 8.9, 106.2), (100, 50, 3.0, 256, False, 3.98, 0.34, 8.4, 106.2), # axis: weighted (T=50 s=3.0 states=64) ( 50, 50, 3.0, 64, True, 3.88, 0.23, 8.4, 53.1), (100, 50, 3.0, 64, True, 3.82, 0.21, 8.3, 106.2), (200, 50, 3.0, 64, True, 3.63, 0.21, 8.0, 212.5), # extra 500-clause captures from background zoom (500, 50, 1.5, 32, True, 2.84, 0.27, 6.1, 531.2), (500, 25, 5.0, 128, True, 3.93, 0.21, 8.9, 531.2), (500, 25, 3.0, 32, False, 8.59, 1.51, 18.7, 531.2), (100,100, 1.5, 64, True, 3.45, 0.15, 7.5, 106.2), (200, 25, 1.5, 128, True, 3.54, 0.17, 7.7, 212.5), (200, 25, 1.5, 256, True, 3.63, 0.11, 8.0, 212.5), ] # ── grid axes for zoom ──────────────────────────────────────────────────────── CLAUSE_OPTS = [10, 20, 50, 100] # cap at 100 for zoom (pure Python speed) T_OPTS = [10, 25, 50, 100, 200] S_OPTS = [1.5, 3.0, 5.0, 10.0, 15.0] STATE_OPTS = [32, 64, 128, 256] # ── encoding ────────────────────────────────────────────────────────────────── def to_bits(value, n_bits): v = max(0, min(int(round(value)), (1 << n_bits) - 1)) return [(v >> i) & 1 for i in range(n_bits - 1, -1, -1)] def row_to_binary(row): bits = [] for fi in range(N_FRAMES): for fname, _, nbits in FIELD_BITS: bits.extend(to_bits(row[f"f{fi}_{fname}"], nbits)) return bits def bullet_speed(power): return 20.0 - 3.0 * power def predict_linear(row, power): t = (row["f0_distance"] / 99.0 * MAX_DIST) / bullet_speed(power) x0, y0 = row["f0_enemy_x"], row["f0_enemy_y"] x1, y1 = row["f1_enemy_x"], row["f1_enemy_y"] return x0 + (x0 - x1) * t, y0 + (y0 - y1) * t def euclid(ax, ay, bx, by): return math.sqrt((ax - bx) ** 2 + (ay - by) ** 2) # ── Parameterised RTM ───────────────────────────────────────────────────────── class RTM: def __init__(self, n_clauses, n_states, s, t_thresh, weighted=False): self.n_clauses = n_clauses self.n_states = n_states self.s = s self.t_thresh = t_thresh self.weighted = weighted half = n_clauses // 2 self.ta = [[n_states] * N_LITERALS for _ in range(n_clauses)] self.polarity = [1] * half + [-1] * half self.w = [1] * n_clauses def _clause_out(self, c, x_aug): ta_c = self.ta[c] ns = self.n_states for l in range(N_LITERALS): if ta_c[l] > ns and x_aug[l] == 0: return 0 return 1 def predict(self, x): x_aug = x + [1 - b for b in x] vote = 0 for c in range(self.n_clauses): vote += self.polarity[c] * self.w[c] * self._clause_out(c, x_aug) max_vote = self.t_thresh * (max(self.w) if self.weighted else 1) vote = max(-max_vote, min(max_vote, vote)) return vote / max_vote * RESID_MAX def learn(self, x, residual): x_aug = x + [1 - b for b in x] pred = self.predict(x) error = residual - pred p_fb = min(1.0, abs(error) / (2 * RESID_MAX)) s, ns = self.s, self.n_states two_ns = 2 * ns for c in range(self.n_clauses): if random.random() >= p_fb: continue pol = self.polarity[c] o = self._clause_out(c, x_aug) ta_c = self.ta[c] if (error > 0 and pol > 0) or (error < 0 and pol < 0): if o == 1: for l in range(N_LITERALS): if x_aug[l] == 1: if random.random() < (s - 1) / s and ta_c[l] < two_ns: ta_c[l] += 1 else: if random.random() < 1.0 / s and ta_c[l] > 1: ta_c[l] -= 1 if self.weighted: self.w[c] = min(self.w[c] + 1, 2 * self.t_thresh) else: for l in range(N_LITERALS): if random.random() < 1.0 / s and ta_c[l] > 1: ta_c[l] -= 1 else: if o == 1: for l in range(N_LITERALS): if x_aug[l] == 0 and ta_c[l] > ns: ta_c[l] -= 1 if self.weighted and self.w[c] > 1: self.w[c] -= 1 # ── data loading ────────────────────────────────────────────────────────────── def load_csvs(): rows = [] for path in CSV_FILES: if not os.path.exists(path): continue with open(path) as f: for row in csv.DictReader(f): try: rows.append({k: float(v) for k, v in row.items()}) except ValueError: pass return rows # ── single config evaluation ────────────────────────────────────────────────── def evaluate(rows, n_clauses, n_states, s, t_thresh, weighted, seed=42): random.seed(seed) rtm_x = RTM(n_clauses, n_states, s, t_thresh, weighted) rtm_y = RTM(n_clauses, n_states, s, t_thresh, weighted) errs_first = [] errs_last = [] errs_all = [] n = len(rows) third = n // 3 for i, row in enumerate(rows): x = row_to_binary(row) rx_tm = rtm_x.predict(x) ry_tm = rtm_y.predict(x) rx_sum = ry_sum = 0.0 batch_errs = [] for ps, power in zip(POWER_STRS, POWER_LEVELS): ax = row[f"{ps}_enemy_x"] ay = row[f"{ps}_enemy_y"] lx, ly = predict_linear(row, power) e = euclid(lx + rx_tm, ly + ry_tm, ax, ay) batch_errs.append(e) rx_sum += ax - lx ry_sum += ay - ly avg_e = sum(batch_errs) / len(batch_errs) errs_all.append(avg_e) if i < third: errs_first.append(avg_e) if i >= n - third: errs_last.append(avg_e) rtm_x.learn(x, rx_sum / len(POWER_LEVELS)) rtm_y.learn(x, ry_sum / len(POWER_LEVELS)) mae_all = sum(errs_all) / len(errs_all) mae_first = sum(errs_first) / len(errs_first) if errs_first else float("nan") mae_last = sum(errs_last) / len(errs_last) if errs_last else float("nan") mem_bytes = 2 * n_clauses * N_LITERALS # logical uint8 storage return { "mae_all": mae_all, "mae_first": mae_first, "mae_last": mae_last, "curve": mae_first - mae_last, "mem_kb": mem_bytes / 1024, } # ── zoom grid ───────────────────────────────────────────────────────────────── def zoom_configs(top5_cfgs, exclude_set): """One-axis-at-a-time neighbours, cap clauses ≤ 100 for speed.""" def nb(v, opts): i = opts.index(v) if v in opts else 0 return [opts[max(0, i-1)], opts[min(len(opts)-1, i+1)]] configs = set() for (c, t, s, st, w) in top5_cfgs: for nc in nb(c, CLAUSE_OPTS): configs.add((nc, t, s, st, w)) for nt in nb(t, T_OPTS): configs.add((c, nt, s, st, w)) for ns_ in nb(s, S_OPTS): configs.add((c, t, ns_, st, w)) for nst in nb(st, STATE_OPTS): configs.add((c, t, s, nst, w)) configs.add((c, t, s, st, not w)) # filter: exclude pre-captured and already-run; cap clauses ≤ 100 return [cfg for cfg in configs if cfg not in exclude_set and cfg[0] <= 100] # ── main ────────────────────────────────────────────────────────────────────── def fmt_cfg(c, t, s, st, w): return f"clauses={c:<3d} T={t:<3d} s={s:<5.1f} states={st:<3d} {'W' if w else ' '}" def main(): print("Loading CSV data...") rows = load_csvs() if not rows: print("ERROR: no CSV data found", file=sys.stderr) sys.exit(1) print(f" {len(rows)} rows from {sum(1 for p in CSV_FILES if os.path.exists(p))} files\n") # Phase 1: coarse is pre-captured (avoid re-running slow 200/500-clause) print("=== PHASE 1: COARSE GRID (pre-captured from background run) ===") print(f" {len(COARSE_PRE)} configs loaded") for kr in COARSE_PRE: c, t, s, st, w, mae_all, mae_last, curve, mem_kb = kr flag = " <-- best" if mae_all < LINEAR_BASELINE_MAE else "" print(f" [PRE] {fmt_cfg(c,t,s,st,w)} | MAE={mae_all:6.2f} last={mae_last:6.2f}" f" curve={curve:+5.1f} mem={mem_kb:5.1f}KB{flag}") # top-5 from coarse (by mae_all), only those with clauses ≤ 100 for zoom coarse_sorted = sorted(COARSE_PRE, key=lambda x: x[5]) top5_raw = [r for r in coarse_sorted if r[0] <= 100][:5] top5_cfgs = [(r[0], r[1], r[2], r[3], r[4]) for r in top5_raw] print(f"\n Top-5 coarse (clauses ≤ 100, for zoom):") for r in top5_raw: print(f" {fmt_cfg(r[0],r[1],r[2],r[3],r[4])} MAE={r[5]:.2f}") # Phase 2: zoom — run new configs pre_set = set((r[0], r[1], r[2], r[3], r[4]) for r in COARSE_PRE) zoom = zoom_configs(top5_cfgs, pre_set) print(f"\n=== PHASE 2: ZOOM ({len(zoom)} new configs, clauses ≤ 100) ===") t_start = time.time() zoom_results = [] for i, (c, t, s, st, w) in enumerate(zoom): t0 = time.time() r = evaluate(rows, c, t, s, st, w) elapsed = time.time() - t0 zoom_results.append((c, t, s, st, w, r["mae_all"], r["mae_first"], r["mae_last"], r["curve"], r["mem_kb"], elapsed)) flag = " <-- best" if r["mae_all"] < LINEAR_BASELINE_MAE else "" print(f" [Z {i+1:2d}] {fmt_cfg(c,t,s,st,w)} | " f"MAE={r['mae_all']:6.2f} last={r['mae_last']:6.2f} " f"curve={r['curve']:+5.1f} mem={r['mem_kb']:5.1f}KB {elapsed:.1f}s{flag}") total_t = time.time() - t_start # Merge all results # Pre-captured: (c, t, s, st, w, mae_all, mae_last, curve, mem_kb) # Zoom: (c, t, s, st, w, mae_all, mae_first, mae_last, curve, mem_kb, elapsed) all_results = [] for r in COARSE_PRE: c, t, s, st, w, mae_all, mae_last, curve, mem_kb = r all_results.append((c, t, s, st, w, mae_all, float("nan"), mae_last, curve, mem_kb, "(pre)")) for r in zoom_results: c, t, s, st, w, mae_all, mae_f, mae_l, curve, mem_kb, elapsed = r all_results.append((c, t, s, st, w, mae_all, mae_f, mae_l, curve, mem_kb, f"{elapsed:.1f}s")) all_results.sort(key=lambda x: x[5]) # ── report ──────────────────────────────────────────────────────────────── n_run = len(zoom_results) lines = [] a = lines.append a("=" * 100) a("TSETLIN MACHINE HYPERPARAMETER SWEEP RESULTS") a(f"Data: {len(rows)} rows | Pre-captured: {len(COARSE_PRE)} | Run: {n_run} | Zoom time: {total_t:.1f}s") a(f"Baseline (linear extrapolation): MAE = {LINEAR_BASELINE_MAE:.2f}") a(f"Note: 200/500-clause pre-captured (pure Python: ~15-40s/config). " f"Zoom capped at clauses≤100.") a("=" * 100) a("") a(f"{'#':<3} {'Clauses':>7} {'T':>4} {'s':>5} {'States':>6} {'W':>2} " f"{'MAE(all)':>9} {'MAE(last⅓)':>11} {'Curve':>7} {'Mem(KB)':>8} {'vs baseline':>12} {'time':>6}") a("-" * 100) a(f"{'':3} {'':7} {'':4} {'':5} {'':6} {'':2} " f" BASELINE {'12.24':>11} {'0.00':>12}") a("-" * 100) for rank, r in enumerate(all_results, 1): c, t, s, st, w, mae_all, mae_f, mae_l, curve, mem_kb, timing = r vs = mae_all - LINEAR_BASELINE_MAE flag = " ***" if vs < -0.5 else (" **" if vs < 0 else "") a(f"{rank:<3} {c:>7} {t:>4} {s:>5.1f} {st:>6} {'Y' if w else 'N':>2} " f"{mae_all:>9.2f} {mae_l:>11.2f} {curve:>+7.2f} {mem_kb:>8.1f} " f"{vs:>+10.2f}{flag} {timing:>6}") a("") a(" *** = beats baseline by >0.5 ** = beats baseline") a(" Curve = MAE(first⅓) - MAE(last⅓), positive = converging") a(" Mem = logical TA storage for 2 RTMs at 1 byte/state (uint8)") a(" (pre) = pre-captured from background run, not re-run this session") a("") a("--- TOP-3 CONFIGS ---") for rank, r in enumerate(all_results[:3], 1): c, t, s, st, w, mae_all, mae_f, mae_l, curve, mem_kb, timing = r a(f" #{rank}: clauses={c} T={t} s={s} states={st} weighted={'Y' if w else 'N'}") a(f" MAE(all)={mae_all:.2f} MAE(last⅓)={mae_l:.2f} " f"curve={curve:+.2f} mem={mem_kb:.1f}KB") a("") a("--- KEY FINDINGS ---") a(" 1. ALL configs beat the linear baseline (MAE 12.24) — even 10 clauses.") a(" 2. Weighted clauses consistently outperform unweighted at same clause count.") a(" 3. Lower clause counts (10-50) often match 100-200 clause accuracy — cold-start" " dominates all-rows MAE.") a(" 4. Last-⅓ MAE (warm TM) is nearly 0 across all configs: TM memorises the") a(" small dataset. In a real 1000-tick battle, last-⅓ is the relevant metric.") a(" 5. States: higher (256) helps — TAs move more slowly, more stable features.") a(" 6. s (specificity): 1.5-3.0 optimal. High s (10-15) = too sparse clauses.") a(" 7. T (threshold): nearly no effect — vote clamping is rarely active here.") a(" 8. 500-clause weighted s=1.5: best MAE(all)=2.84, but mem=531KB and slow.") a(" Practical recommendation: clauses=100 T=100 s=1.5 weighted=Y (MAE=3.45," " mem=106KB).") report = "\n".join(lines) print("\n" + report) with open(OUT_PATH, "w") as f: f.write(report + "\n") print(f"\n[saved to {OUT_PATH}]") if __name__ == "__main__": main()