Files
SirRoboGarage/BNNBot_garage/analysis/sweep_tsetlin.py
SirStone 1ed7797cb6 feat(ModularBot): 6 guns, pattern matcher, melee modules, adversarial bots
- New guns: guess-factor (GF histogram), pattern-matcher (movement tape replay)
- New modules: minimum-risk melee movement, spinning melee radar
- New test bots: PatternMover, RandomMover, WaveSurfer
- Fixed: FeedbackEvent now carries actualX/actualY for proper GF learning
- Fixed: TM gun warmup gating + directional residuals
- Fixed: circular gun integrated formula + multi-bin omega cache
- Fixed: oscillator wall-bounce lockout
- Fixed: phantom meteor perpendicular body orientation
- 6/6 battle wins across all enemy types
2026-09-20 00:59:53 +02:00

408 lines
18 KiB
Python

"""
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()