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
324 lines
11 KiB
Python
324 lines
11 KiB
Python
"""
|
||
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()
|