""" WiSARD hyperparameter sweep. Stdlib only: csv, math, random. """ import csv import math import random import os DATA_DIR = os.path.join(os.path.dirname(__file__), "../out/data") OUT_PATH = os.path.join(os.path.dirname(__file__), "sweep_wisard_results.txt") # Use all target battle decimal CSVs for more data CSV_FILES = [ os.path.join(DATA_DIR, f"target_battle_{i}_decimal.csv") for i in range(1, 6) ] 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 HIT_RADIUS = 18.0 # px, bot half-width for hit-rate calc REP_POWER = 1.07 REP_POWER_STR = "p1.07" K_VALUES = [4, 6, 8, 10, 12, 14, 16, 20] XOR_OPTIONS = [False, True] BLEACH_VALS = [0, 1, 2, 3] SEEDS = [42, 137, 2718] # --------------------------------------------------------------------------- # Bit encoding (mirrors backtest_wisard.py) # --------------------------------------------------------------------------- def _to_gray(v): return v ^ (v >> 1) def int_to_bits(value, nbits): gray = _to_gray(int(value)) return [(gray >> (nbits - 1 - i)) & 1 for i in range(nbits)] def encode_frame(row, prefix): bits = [] bits += int_to_bits(row[prefix + "bearing_sin"], 8) bits += int_to_bits(row[prefix + "bearing_cos"], 8) bits += int_to_bits(row[prefix + "distance"], 7) bits += int_to_bits(row[prefix + "velocity"], 5) bits += int_to_bits(row[prefix + "heading_sin"], 8) bits += int_to_bits(row[prefix + "heading_cos"], 8) bits += int_to_bits(row[prefix + "enemy_x"], 7) bits += int_to_bits(row[prefix + "enemy_y"], 7) bits += int_to_bits(row[prefix + "enemy_energy"],11) return bits # 69 bits def build_input(row, use_xor=False): frames = [encode_frame(row, f"f{i}_") for i in range(4)] bits = [] for f in frames: bits += f if use_xor: bits += [a ^ b for a, b in zip(frames[0], frames[1])] return bits # 276 or 345 bits # --------------------------------------------------------------------------- # Physics helpers # --------------------------------------------------------------------------- def bullet_speed(power): return 20.0 - 3.0 * power def flight_ticks(dist_enc, power): return (dist_enc / 99.0 * MAX_DIST) / bullet_speed(power) def euclid(ax, ay, bx, by): return math.sqrt((ax - bx) ** 2 + (ay - by) ** 2) def predict_linear(row, t): 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 # --------------------------------------------------------------------------- # WiSARD with bleaching # --------------------------------------------------------------------------- class WiSARD: def __init__(self, n_bits, k, seed=42): self.k = k n_bits_pad = math.ceil(n_bits / k) * k self.n_nodes = n_bits_pad // k rng = random.Random(seed) indices = list(range(n_bits)) + [0] * (n_bits_pad - n_bits) self.perm = indices[:] rng.shuffle(self.perm) self.tables = [{} for _ in range(self.n_nodes)] def _addresses(self, bits): addrs = [] for node in range(self.n_nodes): addr = 0 for bi in range(self.k): pi = node * self.k + bi b = bits[self.perm[pi]] if pi < len(bits) else 0 addr = (addr << 1) | b addrs.append(addr) return addrs def predict(self, bits, bleach=0): addrs = self._addresses(bits) cx_sum = cy_sum = 0.0 count = 0 for node, addr in enumerate(addrs): entry = self.tables[node].get(addr) if entry and entry[2] > bleach: cx_sum += entry[0] / entry[2] cy_sum += entry[1] / entry[2] count += 1 if count == 0: return 0.0, 0.0 return cx_sum / count, cy_sum / count def learn(self, bits, rx, ry): for node, addr in enumerate(self._addresses(bits)): entry = self.tables[node].get(addr) if entry is None: self.tables[node][addr] = [rx, ry, 1] else: entry[0] += rx entry[1] += ry entry[2] += 1 # --------------------------------------------------------------------------- # Data loading # --------------------------------------------------------------------------- def load_rows(): 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 # --------------------------------------------------------------------------- # Run one config # --------------------------------------------------------------------------- def run_config(rows, k, use_xor, bleach, seed): n_bits = 276 + (69 if use_xor else 0) net = WiSARD(n_bits, k, seed) n = len(rows) all_errors = [] errors_first50 = [] errors_last50 = [] hits = 0 for i, row in enumerate(rows): bits = build_input(row, use_xor) cx, cy = net.predict(bits, bleach) t = flight_ticks(row["f0_distance"], REP_POWER) lx, ly = predict_linear(row, t) px = lx + cx py = ly + cy ax, ay = row[f"{REP_POWER_STR}_enemy_x"], row[f"{REP_POWER_STR}_enemy_y"] e = euclid(px, py, ax, ay) all_errors.append(e) if e <= HIT_RADIUS: hits += 1 if i < 50: errors_first50.append(e) if i >= n - 50: errors_last50.append(e) # Online learn: residual on top of current prediction rx = ax - px ry = ay - py net.learn(bits, rx, ry) mae = sum(all_errors) / n hit_rate = hits / n * 100 mae_f50 = sum(errors_first50) / len(errors_first50) if errors_first50 else float("nan") mae_l50 = sum(errors_last50) / len(errors_last50) if errors_last50 else float("nan") return mae, hit_rate, mae_f50, mae_l50 # --------------------------------------------------------------------------- # Main # --------------------------------------------------------------------------- def main(): rows = load_rows() print(f"Loaded {len(rows)} rows from {sum(os.path.exists(p) for p in CSV_FILES)} files") # Baseline: pure linear extrapolation (no correction) baseline_errors = [] baseline_hits = 0 for row in rows: t = flight_ticks(row["f0_distance"], REP_POWER) px, py = predict_linear(row, t) ax, ay = row[f"{REP_POWER_STR}_enemy_x"], row[f"{REP_POWER_STR}_enemy_y"] e = euclid(px, py, ax, ay) baseline_errors.append(e) if e <= HIT_RADIUS: baseline_hits += 1 baseline_mae = sum(baseline_errors) / len(baseline_errors) baseline_hitrate = baseline_hits / len(rows) * 100 print(f"Baseline linear extrapolation: MAE={baseline_mae:.2f}px hit%={baseline_hitrate:.1f}%") print(f"\nRunning sweep: {len(K_VALUES)} K × {len(XOR_OPTIONS)} XOR × {len(BLEACH_VALS)} bleach × {len(SEEDS)} seeds = " f"{len(K_VALUES)*len(XOR_OPTIONS)*len(BLEACH_VALS)*len(SEEDS)} configs...") results = [] # (avg_mae, hit_rate, mae_f50, mae_l50, std_mae, k, use_xor, bleach) total = len(K_VALUES) * len(XOR_OPTIONS) * len(BLEACH_VALS) done = 0 for k in K_VALUES: for use_xor in XOR_OPTIONS: for bleach in BLEACH_VALS: done += 1 seed_maes = [] seed_hits = [] seed_f50 = [] seed_l50 = [] for seed in SEEDS: mae, hit_rate, mae_f50, mae_l50 = run_config(rows, k, use_xor, bleach, seed) seed_maes.append(mae) seed_hits.append(hit_rate) seed_f50.append(mae_f50) seed_l50.append(mae_l50) avg_mae = sum(seed_maes) / len(seed_maes) avg_hit = sum(seed_hits) / len(seed_hits) avg_f50 = sum(seed_f50) / len(seed_f50) avg_l50 = sum(seed_l50) / len(seed_l50) mean_m = avg_mae std_mae = math.sqrt(sum((x - mean_m) ** 2 for x in seed_maes) / len(seed_maes)) results.append((avg_mae, avg_hit, avg_f50, avg_l50, std_mae, k, use_xor, bleach)) print(f" [{done:3d}/{total}] K={k:2d} xor={int(use_xor)} bleach={bleach} MAE={avg_mae:.2f}±{std_mae:.2f} hit%={avg_hit:.1f}%") results.sort(key=lambda r: r[0]) # sort by MAE ascending # ---- Format output ---- lines = [] w = lines.append w("=" * 90) w("WiSARD HYPERPARAMETER SWEEP — sorted by MAE (rep power=1.07, hit_radius=18px)") w(f"Dataset: {len(rows)} rows (target battles 1-5), 3 seeds per config") w("=" * 90) w("") w(f"{'Config':<30} {'MAE':>8} {'±std':>6} {'hit%':>7} {'MAE f50':>8} {'MAE l50':>8} {'vs base':>8}") w("-" * 90) w(f"{'P1 linear baseline':<30} {baseline_mae:>8.2f} {'':>6} {baseline_hitrate:>7.1f}% {'':>8} {'':>8} {'':>8}") w("-" * 90) for avg_mae, avg_hit, avg_f50, avg_l50, std_mae, k, use_xor, bleach in results: n_bits = 276 + (69 if use_xor else 0) xor_tag = "+XOR" if use_xor else " " delta = avg_mae - baseline_mae label = f"K={k:2d} {xor_tag} bleach={bleach} ({n_bits}b)" w(f"{label:<30} {avg_mae:>8.2f} {std_mae:>6.2f} {avg_hit:>7.1f}% {avg_f50:>8.2f} {avg_l50:>8.2f} {delta:>+8.2f}") w("") w("Columns: MAE=mean abs error(px) ±std=across 3 seeds hit%=shots within 18px") w(" MAE f50/l50=first/last 50 samples (learning curve) vs base=delta to P1") 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()