""" Multi-enemy comparison: Linear+Hebbian vs WiSARD K=12 vs TsetlinBot (RTM). Uses 181-col decimal CSVs (4 enemies: crazy, spinbot, target, walls). Online learning, all predictors share state across battles of the same enemy. Rep power: p1.07 for learning signal; all powers reported in summary. """ import csv, math, os, random, glob from collections import defaultdict DATA_DIR = os.path.join(os.path.dirname(__file__), "../data") OUT_PATH = os.path.join(os.path.dirname(__file__), "battle_comparison.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"] REP_PS, REP_POWER = "p1.07", 1.07 MAX_DIST = 1414.0 HIT_RADIUS = 18.0 # approx half tank width in pixels (0-99 scale ≈ arena/14) # ── encoding ─────────────────────────────────────────────────────────────────── def to_bits(v, n): g = int(v) ^ (int(v) >> 1) return [(g >> (n - 1 - i)) & 1 for i in range(n)] def row_to_binary(row, n_frames=4): """Build n_frames*69-bit input. 181-col CSVs have x/y as 0-99 int.""" bits = [] for fi in range(n_frames): p = f"f{fi}_" bits += to_bits(int(row[p+"bearing_sin"]), 8) bits += to_bits(int(row[p+"bearing_cos"]), 8) bits += to_bits(int(row[p+"distance"]), 7) bits += to_bits(int(row[p+"velocity"]), 5) bits += to_bits(int(row[p+"heading_sin"]), 8) bits += to_bits(int(row[p+"heading_cos"]), 8) bits += to_bits(int(row[p+"enemy_x"]), 7) bits += to_bits(int(row[p+"enemy_y"]), 7) bits += to_bits(min(int(row[p+"enemy_energy"]), 2046), 11) return bits # 4*69 = 276 bits # ── WiSARD K=12, 276 bits ────────────────────────────────────────────────────── class WiSARD: def __init__(self, n_bits=276, k=12, seed=42): self.k = k n_pad = math.ceil(n_bits / k) * k self.n_nodes = n_pad // k rng = random.Random(seed) idx = list(range(n_bits)) + [0] * (n_pad - n_bits) rng.shuffle(idx) self.perm = idx self.tables = [{} for _ in range(self.n_nodes)] # sparse dicts def _addrs(self, bits): out = [] for nd in range(self.n_nodes): a = 0 for b in range(self.k): pi = nd * self.k + b a = (a << 1) | (bits[self.perm[pi]] if pi < len(bits) else 0) out.append(a) return out def predict(self, bits): addrs = self._addrs(bits) sx = sy = cnt = 0.0 for nd, a in enumerate(addrs): e = self.tables[nd].get(a) if e and e[2] > 0: sx += e[0] / e[2]; sy += e[1] / e[2]; cnt += 1 return (sx/cnt, sy/cnt) if cnt else (0.0, 0.0) def learn(self, bits, dx, dy): for nd, a in enumerate(self._addrs(bits)): e = self.tables[nd].get(a) if e is None: self.tables[nd][a] = [dx, dy, 1] else: e[0] += dx; e[1] += dy; e[2] += 1 # ── Regression TM (one output dimension) ────────────────────────────────────── N_CLAUSES = 60 N_STATES = 15 S_SPEC = 3.0 T_THRESH = 30 RESID_MAX = 30.0 # residual clamp (0-99 scale) class RTM: def __init__(self): n_lit = 276 * 2 half = N_CLAUSES // 2 self.ta = [[N_STATES] * n_lit for _ in range(N_CLAUSES)] self.pol = [1] * half + [-1] * half self.n_lit = n_lit def _clause_out(self, c, x_aug): ta_c = self.ta[c] for l in range(self.n_lit): if ta_c[l] > N_STATES and x_aug[l] == 0: return 0 return 1 def predict(self, x): x_aug = x + [1 - b for b in x] v = sum(self.pol[c] * self._clause_out(c, x_aug) for c in range(N_CLAUSES)) v = max(-T_THRESH, min(T_THRESH, v)) return v / T_THRESH * RESID_MAX def learn(self, x, residual): x_aug = x + [1 - b for b in x] err = residual - self.predict(x) p_fb = min(1.0, abs(err) / (2 * RESID_MAX)) for c in range(N_CLAUSES): if random.random() >= p_fb: continue pol = self.pol[c]; o = self._clause_out(c, x_aug); ta_c = self.ta[c] if (err > 0 and pol > 0) or (err < 0 and pol < 0): if o == 1: for l in range(self.n_lit): if x_aug[l] == 1: if random.random() < (S_SPEC-1)/S_SPEC: if ta_c[l] < 2*N_STATES: ta_c[l] += 1 else: if random.random() < 1.0/S_SPEC: if ta_c[l] > 1: ta_c[l] -= 1 else: for l in range(self.n_lit): if random.random() < 1.0/S_SPEC: if ta_c[l] > 1: ta_c[l] -= 1 else: if o == 1: for l in range(self.n_lit): if x_aug[l] == 0 and ta_c[l] > N_STATES: ta_c[l] -= 1 # ── helpers ──────────────────────────────────────────────────────────────────── def bullet_speed(p): return 20.0 - 3.0 * p def flight_ticks(dist_enc, p): return (dist_enc / 99.0 * MAX_DIST) / bullet_speed(p) def euclid(ax, ay, bx, by): return math.sqrt((ax-bx)**2 + (ay-by)**2) def predict_linear(row, power): t = flight_ticks(row["f0_distance"], 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 heading_sector(hs, hc): s = hs/199*2-1; c = hc/199*2-1 return int(math.degrees(math.atan2(s,c)) % 360 / 45) % 8 def dist_band(d, lo, hi): return 0 if d <= lo else (1 if d <= hi else 2) def stats(errs): if not errs: return None n = len(errs) mae = sum(errs)/n rmse = math.sqrt(sum(e*e for e in errs)/n) hit = sum(1 for e in errs if e < HIT_RADIUS)/n*100 return {"n":n, "mae":mae, "rmse":rmse, "hit":hit} _REQUIRED_KEYS = ( [f"f{fi}_{col}" for fi in range(4) for col in ("bearing_sin","bearing_cos","distance","velocity","heading_sin","heading_cos","enemy_x","enemy_y","enemy_energy")] + [f"{ps}_enemy_x" for ps in ["p0.10","p0.42","p0.74","p1.07","p1.39","p1.71","p2.03","p2.36","p2.68","p3.00"]] + [f"{ps}_enemy_y" for ps in ["p0.10","p0.42","p0.74","p1.07","p1.39","p1.71","p2.03","p2.36","p2.68","p3.00"]] ) def load_csv(path): rows = [] with open(path) as f: for row in csv.DictReader(f): try: d = {k: float(v) for k,v in row.items() if v not in (None, "NA", "")} if all(k in d for k in _REQUIRED_KEYS): rows.append(d) except (ValueError, TypeError): pass return rows # ── per-enemy backtest ───────────────────────────────────────────────────────── def run_enemy(enemy_files, seed=42): """ Process all battles for one enemy type. Predictors share state across battles (cumulative online learning). Returns dict of results. """ random.seed(seed) # shared predictor state ws = WiSARD(seed=seed) rtm_x, rtm_y = RTM(), RTM() heb = [[[0.0, 0.0] for _ in range(3)] for _ in range(8)] # compute global band thresholds across all battles all_dists = [] all_rows_list = [] for path in sorted(enemy_files): rows = load_csv(path) all_rows_list.append(rows) all_dists.extend(r["f0_distance"] for r in rows) all_dists.sort() n_d = len(all_dists) band_lo, band_hi = all_dists[n_d//3], all_dists[2*n_d//3] errs = { "lin": defaultdict(list), "wis": defaultdict(list), "tm": defaultdict(list), } errs_last50_tm = defaultdict(list) all_processed = 0 for rows in all_rows_list: n = len(rows) for i, row in enumerate(rows): x = row_to_binary(row) sec = heading_sector(row["f0_heading_sin"], row["f0_heading_cos"]) band = dist_band(row["f0_distance"], band_lo, band_hi) cx_heb, cy_heb = heb[sec][band] cx_ws, cy_ws = ws.predict(x) rx_tm = rtm_x.predict(x) ry_tm = rtm_y.predict(x) rx_acc = ry_acc = 0.0 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) errs["lin"][ps].append(euclid(lx+cx_heb, ly+cy_heb, ax, ay)) errs["wis"][ps].append(euclid(lx+cx_ws, ly+cy_ws, ax, ay)) errs["tm"][ps].append( euclid(lx+rx_tm, ly+ry_tm, ax, ay)) if i >= n - 50: errs_last50_tm[ps].append(errs["tm"][ps][-1]) rx_acc += ax - lx; ry_acc += ay - ly # online learning rx_mean, ry_mean = rx_acc/len(POWER_STRS), ry_acc/len(POWER_STRS) ws.learn(x, rx_mean, ry_mean) rtm_x.learn(x, rx_mean); rtm_y.learn(x, ry_mean) lx_r, ly_r = predict_linear(row, REP_POWER) ax_r = row[f"{REP_PS}_enemy_x"]; ay_r = row[f"{REP_PS}_enemy_y"] heb[sec][band][0] += 0.1 * (ax_r - (lx_r + cx_heb)) heb[sec][band][1] += 0.1 * (ay_r - (ly_r + cy_heb)) all_processed += 1 return errs, errs_last50_tm, all_processed # ── main ────────────────────────────────────────────────────────────────────── def main(): random.seed(42) # Find all 181-col decimal CSVs grouped by enemy type all_files = sorted(glob.glob(os.path.join(DATA_DIR, "*_decimal.csv"))) by_enemy = defaultdict(list) for f in all_files: base = os.path.basename(f) enemy = base.split("_battle")[0].replace("_decimal","") # shouldn't be needed but safe by_enemy[enemy].append(f) lines = [] w = lines.append w("=" * 72) w("PREDICTOR COMPARISON: BNNBot (Linear+Hebbian) vs WiSARD K=12 vs TsetlinBot (RTM)") w(f"Rep power: {REP_PS} | Hit threshold: {HIT_RADIUS} (0-99 scale, ~18px)") w(f"Online learning, cumulative per-enemy | Scores measured across all powers") w("=" * 72) w("") global_errs = {"lin": [], "wis": [], "tm": [], "tm_warm": []} enemy_summaries = {} for enemy in sorted(by_enemy.keys()): files = by_enemy[enemy] print(f" {enemy}: {len(files)} battles...") errs, errs_last50_tm, n_rows = run_enemy(files) # avg MAE across all powers def avg_mae(d): return sum(stats(d[ps])["mae"] for ps in POWER_STRS)/len(POWER_STRS) def avg_hit(d): return sum(stats(d[ps])["hit"] for ps in POWER_STRS)/len(POWER_STRS) s_lin = {"mae": avg_mae(errs["lin"]), "hit": avg_hit(errs["lin"])} s_wis = {"mae": avg_mae(errs["wis"]), "hit": avg_hit(errs["wis"])} s_tm = {"mae": avg_mae(errs["tm"]), "hit": avg_hit(errs["tm"])} s_tm_warm = { "mae": sum(stats(errs_last50_tm[ps])["mae"] for ps in POWER_STRS)/len(POWER_STRS) if errs_last50_tm["p1.07"] else float("nan"), } enemy_summaries[enemy] = (s_lin, s_wis, s_tm, s_tm_warm) winner = min([("BNNBot(Lin+Heb)", s_lin["mae"]), ("WiSARD", s_wis["mae"]), ("TsetlinBot", s_tm["mae"])], key=lambda x:x[1])[0] # accumulate globals for ps in POWER_STRS: global_errs["lin"].extend(errs["lin"][ps]) global_errs["wis"].extend(errs["wis"][ps]) global_errs["tm"].extend(errs["tm"][ps]) global_errs["tm_warm"].extend(errs_last50_tm[ps]) w(f"Enemy: {enemy:12s} battles={len(files)} rows={n_rows}") w(f" Avg MAE (all powers): BNNBot={s_lin['mae']:5.2f} WiSARD={s_wis['mae']:5.2f} TsetlinBot={s_tm['mae']:5.2f} (TM-warm={s_tm_warm['mae']:5.2f})") w(f" Avg Hit% (all powers): BNNBot={s_lin['hit']:5.1f}% WiSARD={s_wis['hit']:5.1f}% TsetlinBot={s_tm['hit']:5.1f}%") w(f" Winner by MAE: {winner}") w("") # per-power detail for this enemy w(f" Per-power breakdown (MAE, {enemy}):") w(f" {'Power':<8} {'BNNBot':>8} {'WiSARD':>8} {'TsetlinBot':>10} {'Best':>12}") w(" " + "-" * 52) for ps in POWER_STRS: ml = stats(errs["lin"][ps])["mae"] mw = stats(errs["wis"][ps])["mae"] mt = stats(errs["tm"][ps])["mae"] best = min([("BNNBot",ml),("WiSARD",mw),("TsetlinBot",mt)], key=lambda x:x[1])[0] w(f" {ps:<8} {ml:>8.2f} {mw:>8.2f} {mt:>10.2f} {best:>12s}") w("") # global summary w("=" * 72) w("OVERALL SUMMARY (all enemies, all powers combined)") w("=" * 72) n_all = len(global_errs["lin"]) if n_all: g_lin = sum(global_errs["lin"])/n_all g_wis = sum(global_errs["wis"])/n_all g_tm = sum(global_errs["tm"])/n_all g_tm_warm = sum(global_errs["tm_warm"])/len(global_errs["tm_warm"]) if global_errs["tm_warm"] else float("nan") g_hit_lin = sum(1 for e in global_errs["lin"] if e