1ed7797cb6
- 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
408 lines
18 KiB
Python
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()
|