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
287 lines
10 KiB
Python
287 lines
10 KiB
Python
"""
|
||
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()
|