""" WiSARD / WNN backtest against battle CSV data. Stdlib only: csv, math, random. """ import csv import math import random import os CSV_PATH = os.path.join(os.path.dirname(__file__), "../data/target_battle_1_decimal.csv") OUT_PATH = os.path.join(os.path.dirname(__file__), "backtest_wisard_report.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 # Per-frame field layout (from binary_encoding.nim) # Fields: bearing_sin(8), bearing_cos(8), distance(7), velocity(5), # heading_sin(8), heading_cos(8), enemy_x(7), enemy_y(7), enemy_energy(11) FIELD_WIDTHS = [8, 8, 7, 5, 8, 8, 7, 7, 11] # = 69 bits total FRAME_BITS = 69 # sum(FIELD_WIDTHS) FRAME_FIELDS = ["bearing_sin", "bearing_cos", "distance", "velocity", "heading_sin", "heading_cos", "enemy_x", "enemy_y", "enemy_energy"] # --------------------------------------------------------------------------- # Bit encoding helpers (mirrors binary_encoding.nim) # --------------------------------------------------------------------------- def _to_gray(v): return v ^ (v >> 1) def _from_gray(g): r = g m = r >> 1 while m: r ^= m m >>= 1 return r def int_to_bits(value, nbits): """Gray-encode value and return list of nbits bits (MSB first).""" gray = _to_gray(value) return [(gray >> (nbits - 1 - i)) & 1 for i in range(nbits)] def encode_frame(row, prefix): """Return 69-bit list for one frame given row dict and prefix like 'f0_'.""" bits = [] bits += int_to_bits(int(row[prefix + "bearing_sin"]), 8) bits += int_to_bits(int(row[prefix + "bearing_cos"]), 8) bits += int_to_bits(int(row[prefix + "distance"]), 7) bits += int_to_bits(int(row[prefix + "velocity"]), 5) bits += int_to_bits(int(row[prefix + "heading_sin"]), 8) bits += int_to_bits(int(row[prefix + "heading_cos"]), 8) bits += int_to_bits(int(row[prefix + "enemy_x"]), 7) bits += int_to_bits(int(row[prefix + "enemy_y"]), 7) bits += int_to_bits(int(row[prefix + "enemy_energy"]),11) return bits # len == 69 def build_input(row, use_xor=False): """Build 276 or 345 bit input from frames 0-3.""" frames = [encode_frame(row, f"f{i}_") for i in range(4)] bits = [] for f in frames: bits += f if use_xor: # frame0 XOR frame1 = 69 extra bits bits += [a ^ b for a, b in zip(frames[0], frames[1])] return bits # --------------------------------------------------------------------------- # Core maths # --------------------------------------------------------------------------- def bullet_speed(power): return 20.0 - 3.0 * power def flight_ticks(dist_enc, power): dist_px = dist_enc / 99.0 * MAX_DIST return dist_px / 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 def stats(errors): if not errors: return {} n = len(errors) mae = sum(errors) / n rmse = math.sqrt(sum(e * e for e in errors) / n) med = sorted(errors)[n // 2] mx = max(errors) pct5 = sum(1 for e in errors if e <= 5) / n * 100 return {"mae": mae, "rmse": rmse, "median": med, "max": mx, "pct5": pct5, "n": n} # --------------------------------------------------------------------------- # WiSARD network # --------------------------------------------------------------------------- class WiSARD: """ LUT-based WNN for regression (2D correction output). Each LUT entry stores (sum_x, sum_y, count). Inference = mean of (sum_x/count, sum_y/count) across nodes with count>0. Learning = accumulate residuals at addressed entries. """ def __init__(self, n_bits, k, seed=42): """ n_bits: total input bits k: bits per LUT node (tuple size) """ self.k = k n_bits_pad = math.ceil(n_bits / k) * k # pad to multiple of k self.n_nodes = n_bits_pad // k # Fixed random permutation of bit indices (pad extra with 0-index) rng = random.Random(seed) indices = list(range(n_bits)) + [0] * (n_bits_pad - n_bits) self.perm = indices[:] rng.shuffle(self.perm) # fixed at init, same for all rows # LUT tables: list of dicts {addr: [sum_x, sum_y, count]} # Using dicts — sparse; most addresses never seen with small datasets self.tables = [{} for _ in range(self.n_nodes)] def _addresses(self, bits): """Return list of n_nodes integer addresses.""" addrs = [] for node in range(self.n_nodes): addr = 0 for bit_idx in range(self.k): perm_idx = node * self.k + bit_idx b = bits[self.perm[perm_idx]] if perm_idx < len(bits) else 0 addr = (addr << 1) | b addrs.append(addr) return addrs def predict_correction(self, bits): """Return (cx, cy) WiSARD correction, or (0,0) if no data yet.""" 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] > 0: 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, residual_x, residual_y): """Accumulate residuals at addressed entries.""" for node, addr in enumerate(self._addresses(bits)): entry = self.tables[node].get(addr) if entry is None: self.tables[node][addr] = [residual_x, residual_y, 1] else: entry[0] += residual_x entry[1] += residual_y entry[2] += 1 # --------------------------------------------------------------------------- # Backtest runner # --------------------------------------------------------------------------- def run_wisard(rows, k, use_xor, rep_power_str="p1.07", rep_power=1.07): n_bits = 276 + (69 if use_xor else 0) net = WiSARD(n_bits, k) errors_by_power = {ps: [] for ps in POWER_STRS} errors_first50 = {ps: [] for ps in POWER_STRS} errors_last50 = {ps: [] for ps in POWER_STRS} n = len(rows) for i, row in enumerate(rows): bits = build_input(row, use_xor) cx, cy = net.predict_correction(bits) for ps, power in zip(POWER_STRS, POWER_LEVELS): t = flight_ticks(row["f0_distance"], power) lx, ly = predict_linear(row, t) px = lx + cx py = ly + cy ax, ay = row[f"{ps}_enemy_x"], row[f"{ps}_enemy_y"] e = euclid(px, py, ax, ay) errors_by_power[ps].append(e) if i < 50: errors_first50[ps].append(e) if i >= n - 50: errors_last50[ps].append(e) # Learn: residual for representative power (same choice as backtest.py) t_rep = flight_ticks(row["f0_distance"], rep_power) lx_r, ly_r = predict_linear(row, t_rep) ax_r = row[f"{rep_power_str}_enemy_x"] ay_r = row[f"{rep_power_str}_enemy_y"] rx = ax_r - (lx_r + cx) ry = ay_r - (ly_r + cy) net.learn(bits, rx, ry) return errors_by_power, errors_first50, errors_last50, net # --------------------------------------------------------------------------- # Main # --------------------------------------------------------------------------- def load_csv(): with open(CSV_PATH) as f: reader = csv.DictReader(f) rows = [] for row in reader: try: rows.append({k: float(v) for k, v in row.items()}) except ValueError: pass return rows def main(): rows = load_csv() lines = [] w = lines.append w("=" * 72) w("WiSARD / WNN BACKTEST REPORT") w(f"Rows: {len(rows)}") w("=" * 72) # ---- Baseline P1 for comparison ---- w("\n--- BASELINE: Pure Linear Extrapolation (P1) ---\n") w(f"{'Power':<8} {'MAE':>7} {'RMSE':>7} {'Median':>8} {'Max':>8} {'%<5u':>7}") w("-" * 52) p1_mae = {} for ps, power in zip(POWER_STRS, POWER_LEVELS): errs = [] for row in rows: t = flight_ticks(row["f0_distance"], power) px, py = predict_linear(row, t) ax, ay = row[f"{ps}_enemy_x"], row[f"{ps}_enemy_y"] errs.append(euclid(px, py, ax, ay)) s = stats(errs) p1_mae[ps] = s["mae"] w(f"{ps:<8} {s['mae']:>7.2f} {s['rmse']:>7.2f} {s['median']:>8.2f} {s['max']:>8.2f} {s['pct5']:>6.1f}%") # ---- WiSARD configurations ---- configs = [ (8, False, "K=8, 276 bits (no XOR)"), (8, True, "K=8, 345 bits (with XOR)"), (12, False, "K=12, 276 bits (no XOR)"), (12, True, "K=12, 345 bits (with XOR)"), ] all_results = [] for k, use_xor, label in configs: n_bits = 276 + (69 if use_xor else 0) n_nodes = math.ceil(n_bits / k) print(f"Running {label} ({n_nodes} LUT nodes × 2^{k} entries)...") ebp, ef50, el50, net = run_wisard(rows, k, use_xor) all_results.append((label, k, use_xor, ebp, ef50, el50)) w(f"\n--- WiSARD: {label} | {n_nodes} nodes ---\n") w(f"{'Power':<8} {'MAE':>7} {'RMSE':>7} {'Median':>8} {'Max':>8} {'%<5u':>7} {'vs P1':>8}") w("-" * 62) for ps in POWER_STRS: s = stats(ebp[ps]) delta = s["mae"] - p1_mae[ps] w(f"{ps:<8} {s['mae']:>7.2f} {s['rmse']:>7.2f} {s['median']:>8.2f} {s['max']:>8.2f} {s['pct5']:>6.1f}% {delta:>+8.2f}") # Learning curve for p1.07 ps_r = "p1.07" mae_f = stats(ef50[ps_r])["mae"] if ef50[ps_r] else float("nan") mae_l = stats(el50[ps_r])["mae"] if el50[ps_r] else float("nan") w(f"\n Learning curve (p1.07): first-50 MAE={mae_f:.2f} last-50 MAE={mae_l:.2f}") # LUT fill stats total_entries = sum(len(t) for t in net.tables) total_capacity = len(net.tables) * (2 ** k) fill_pct = total_entries / total_capacity * 100 w(f" LUT fill: {total_entries} / {total_capacity} entries ({fill_pct:.3f}%)") # ---- Summary comparison ---- w("\n\n" + "=" * 72) w("SUMMARY — Average MAE across all power levels") w("=" * 72) w(f"\n{'Config':<32} {'Avg MAE':>9} {'vs P1':>8}") w("-" * 52) p1_avg = sum(p1_mae.values()) / len(p1_mae) w(f"{'P1 (linear baseline)':<32} {p1_avg:>9.3f}") for label, k, use_xor, ebp, ef50, el50 in all_results: avg = sum(stats(ebp[ps])["mae"] for ps in POWER_STRS) / len(POWER_STRS) delta = avg - p1_avg w(f"{label:<32} {avg:>9.3f} {delta:>+8.3f}") w("\n") w("Note: negative 'vs P1' = improvement; positive = worse than linear baseline.") w("Learning is online: each row trains BEFORE the next prediction.") w("WiSARD learns the residual correction on top of linear extrapolation.") report = "\n".join(lines) print(report) with open(OUT_PATH, "w") as f: f.write(report + "\n") print(f"\n[saved to {OUT_PATH}]") if __name__ == "__main__": main()