j128 learned movement: state-conditional counted-SBC wave danger (TR_MOVEMENT=learned, default-off), offline gate + pre-registered panel arms
This commit is contained in:
@@ -0,0 +1,535 @@
|
||||
## learned_surfer.nim — LEARNED movement: state-conditional wave danger.
|
||||
##
|
||||
## The novelty over every other mover in this repo (`strafe`, `tfil`,
|
||||
## `wave_surfer`): the danger of a guess-factor bin is **learned per movement
|
||||
## state** instead of being a single hand-tuned/global quantity. The learner is
|
||||
## the counted SBC with global fractional decay from `common_libs/bitbrain`
|
||||
## (jobs j102/j103: measured to forget a changed mapping and to produce true
|
||||
## probabilities), so it can track an opponent that adapts to us.
|
||||
##
|
||||
## ── The design (and its measured basis) ─────────────────────────────────────
|
||||
## 1. Waves: enemy fire is detected from the one-tick energy drop (the drop IS
|
||||
## the firepower), exactly as `wave_surfer.nim`/`strafe.nim` do. `WorldState`
|
||||
## carries no bullet bodies, so the wave frame is "origin = enemy position at
|
||||
## the fire tick, centre line = bearing from the origin to us at that tick".
|
||||
##
|
||||
## 2. Label. A wave RESOLVES at the nominal arrival tick
|
||||
## `ceil(startDist/speed)` — the time the bullet needs for the range at fire
|
||||
## time, known at fire time — and the label is the 31-bin guess factor of our
|
||||
## angular offset from the centre line at that tick (`gfToBin`, the same
|
||||
## 31-bin quantisation `wave_surfer.nim` uses).
|
||||
##
|
||||
## 3. State = ONE coarse wave-relative state at the fire tick, 4 fields x 4
|
||||
## symbols = 256 states, NEVER a temporal window (measured dead:
|
||||
## `docs/state_window_gate.md`):
|
||||
## vlat lateral velocity in the wave frame, px/tick (signed)
|
||||
## dist range to the enemy at the fire tick, px
|
||||
## room directional wall room along the direction we are running, px
|
||||
## turn our own signed heading change, deg/tick
|
||||
## `lat` (the perpendicular offset from the centre line) is deliberately NOT
|
||||
## a field: at the fire tick the centre line passes through us, so it is
|
||||
## identically 0. 256 states over ~55k recorded shots is ~150 observations
|
||||
## per state — no recurrence problem.
|
||||
##
|
||||
## 4. Learner = counted SBC (`initCountedSbc`), read with `inferProb` (the
|
||||
## per-cell posterior), interpolated with the global 31-bin histogram that
|
||||
## `wave_surfer.nim` used: `p = (posterior + alpha*global)/(1 + alpha)`. With
|
||||
## no data at a cell this is exactly the old global surfer.
|
||||
##
|
||||
## 5. Danger of a candidate escape direction = the predicted probability of the
|
||||
## GF bin we would arrive in, SUMMED over every live wave, plus a wall
|
||||
## penalty, a travel penalty and a reversal penalty. The safest reachable
|
||||
## bin wins, then the same perpendicular steering / wall escape / radial
|
||||
## blend the other movers use.
|
||||
##
|
||||
## MEASURED (offline gate, `common_libs/tests/learned_surfer_gate.py`, 70
|
||||
## recorded battles, held out BY BATTLE, 3 seeds): the state-conditional model
|
||||
## beats the global histogram and chance on held-out log-loss (4.927 vs 4.974
|
||||
## vs 4.954 bits) in 63/63 held-out battles (sign-flip p = 5e-5), but the effect
|
||||
## is TINY: top-1 3.93% (global 3.96%, chance 3.23%). Read the ledger section
|
||||
## "Learned movement (SBC)" in `docs/movement_campaign.md` before trusting it.
|
||||
##
|
||||
## env knobs (all read by `loadLearnedEnv`):
|
||||
## TR_MOVEMENT = learned (selects this engine)
|
||||
## TR_LEARNED_DECAY_EVERY learns between decay passes (default 128)
|
||||
## TR_LEARNED_DECAY_SHIFT `c -= c shr shift`; 0 disables forgetting
|
||||
## TR_LEARNED_ALPHA prior mix weight (default 5.0)
|
||||
## TR_LEARNED_TRAVEL danger cost of travelling across the wave (0.01)
|
||||
## TR_LEARNED_REVERSAL danger cost of flipping the strafe side (0.02)
|
||||
## TR_LEARNED_PREF_DIST preferred engagement distance (400 px)
|
||||
## TR_LEARNED_DIST_BAND deadband around it (50 px)
|
||||
## TR_LEARNED_RADIAL_FRAC radial blend outside the band (0.35)
|
||||
## TR_LEARNED_WALL_MARGIN wall margin (48 px)
|
||||
## TR_LEARNED_GLOBAL =1: ignore the state (ABLATION, the same mover
|
||||
## with a pure global histogram)
|
||||
## TR_LEARNED_LOG per-decision log line
|
||||
|
||||
import std/[math, os]
|
||||
from std/strutils import parseFloat, parseInt, strip
|
||||
import gun_harness/gun_interface
|
||||
import movement_harness/movement_interface
|
||||
import bitbrain/sbc
|
||||
|
||||
const
|
||||
LS_BINS = 31 ## GF bins, identical to `wave_surfer.WS_BINS`
|
||||
LS_Q = 4 ## symbols per field
|
||||
LS_NADE = LS_Q * LS_Q ## 16: (vlat, dist) on one axis, (room, turn) on the other
|
||||
DodgeTicks = 15.0
|
||||
MaxBotSpeed = 8.0
|
||||
|
||||
## State-bin edges (4 symbols/field): the quantiles of the recorded corpus,
|
||||
## derived by `common_libs/tests/learned_surfer_gate.py` (section E) and frozen
|
||||
## here. Quantisation is the 25/50/75 % quantiles; a value is placed in the
|
||||
## number of edges strictly below it.
|
||||
const
|
||||
VlatEdges = [-6.736, 0.000, 6.753]
|
||||
DistEdges = [431.321, 487.612, 552.670]
|
||||
RoomEdges = [137.965, 206.589, 296.753]
|
||||
TurnEdges = [-0.142, 0.000, 0.105]
|
||||
|
||||
const
|
||||
LearnedDecayEveryEnv* = "TR_LEARNED_DECAY_EVERY"
|
||||
LearnedDecayShiftEnv* = "TR_LEARNED_DECAY_SHIFT"
|
||||
LearnedAlphaEnv* = "TR_LEARNED_ALPHA"
|
||||
LearnedTravelEnv* = "TR_LEARNED_TRAVEL"
|
||||
LearnedReversalEnv* = "TR_LEARNED_REVERSAL"
|
||||
LearnedPrefDistEnv* = "TR_LEARNED_PREF_DIST"
|
||||
LearnedDistBandEnv* = "TR_LEARNED_DIST_BAND"
|
||||
LearnedRadialFracEnv* = "TR_LEARNED_RADIAL_FRAC"
|
||||
LearnedWallMarginEnv* = "TR_LEARNED_WALL_MARGIN"
|
||||
LearnedGlobalEnv* = "TR_LEARNED_GLOBAL"
|
||||
LearnedLogEnv* = "TR_LEARNED_LOG"
|
||||
|
||||
var
|
||||
LearnedDecayEvery* = 128
|
||||
LearnedDecayShift* = 1
|
||||
LearnedAlpha* = 5.0
|
||||
LearnedTravel* = 0.01
|
||||
LearnedReversal* = 0.02
|
||||
LearnedPrefDist* = 400.0
|
||||
LearnedDistBand* = 50.0
|
||||
LearnedRadialFrac* = 0.35
|
||||
LearnedWallMargin* = 48.0
|
||||
LearnedGlobal* = false
|
||||
LearnedLog* = false
|
||||
|
||||
proc getEnvFloat(name: string, default: float): float =
|
||||
let s = getEnv(name, "")
|
||||
if s.len == 0: return default
|
||||
try: result = parseFloat(s.strip())
|
||||
except ValueError: result = default
|
||||
|
||||
proc getEnvInt(name: string, default: int): int =
|
||||
let s = getEnv(name, "")
|
||||
if s.len == 0: return default
|
||||
try: result = parseInt(s.strip())
|
||||
except ValueError: result = default
|
||||
|
||||
proc envOn(name: string, default = false): bool =
|
||||
let s = getEnv(name, "").strip()
|
||||
if s.len == 0: return default
|
||||
s notin ["0", "false", "no", "off"]
|
||||
|
||||
proc loadLearnedEnv*() =
|
||||
## Read the knobs; callable again after `putEnv` so a gate can sweep arms in
|
||||
## one process.
|
||||
LearnedDecayEvery = max(0, getEnvInt(LearnedDecayEveryEnv, 128))
|
||||
LearnedDecayShift = max(0, getEnvInt(LearnedDecayShiftEnv, 1))
|
||||
LearnedAlpha = max(0.0, getEnvFloat(LearnedAlphaEnv, 5.0))
|
||||
LearnedTravel = max(0.0, getEnvFloat(LearnedTravelEnv, 0.01))
|
||||
LearnedReversal = max(0.0, getEnvFloat(LearnedReversalEnv, 0.02))
|
||||
LearnedPrefDist = max(1.0, getEnvFloat(LearnedPrefDistEnv, 400.0))
|
||||
LearnedDistBand = max(0.0, getEnvFloat(LearnedDistBandEnv, 50.0))
|
||||
LearnedRadialFrac = clamp(getEnvFloat(LearnedRadialFracEnv, 0.35), 0.0, 1.0)
|
||||
LearnedWallMargin = max(0.0, getEnvFloat(LearnedWallMarginEnv, 48.0))
|
||||
LearnedGlobal = envOn(LearnedGlobalEnv)
|
||||
LearnedLog = envOn(LearnedLogEnv)
|
||||
|
||||
loadLearnedEnv()
|
||||
|
||||
# ── small helpers ───────────────────────────────────────────────────────────
|
||||
|
||||
proc wrapPi(x: float64): float64 {.inline.} =
|
||||
result = x
|
||||
while result > PI: result -= 2.0*PI
|
||||
while result < -PI: result += 2.0*PI
|
||||
|
||||
proc wrap180(d: float): float {.inline.} =
|
||||
result = d
|
||||
while result > 180.0: result -= 360.0
|
||||
while result < -180.0: result += 360.0
|
||||
|
||||
proc mea(speed: float64): float64 {.inline.} =
|
||||
if speed <= 1e-9: return 0.0
|
||||
arcsin(min(MaxBotSpeed / speed, 1.0))
|
||||
|
||||
proc gfToBin(gf: float64): int {.inline.} =
|
||||
clamp(int(round((gf.clamp(-1.0, 1.0) + 1.0) * 0.5 * float64(LS_BINS - 1))),
|
||||
0, LS_BINS - 1)
|
||||
|
||||
proc binToGF(idx: int): float64 {.inline.} =
|
||||
float64(idx) / float64(LS_BINS - 1) * 2.0 - 1.0
|
||||
|
||||
proc code(value: float64, edges: array[3, float64]): int {.inline.} =
|
||||
result = 0
|
||||
for e in edges:
|
||||
if value > e: inc result
|
||||
|
||||
proc roomToWall(px, py, dx, dy, arenaW, arenaH: float64): float64 =
|
||||
## Distance from (px,py) along unit (dx,dy) until leaving the arena, keeping
|
||||
## the 18 px bot radius. Mirrors the offline gate.
|
||||
var t = Inf
|
||||
for i in 0..1:
|
||||
let p = if i == 0: px else: py
|
||||
let d = if i == 0: dx else: dy
|
||||
let lo = BotRadius
|
||||
let hi = (if i == 0: arenaW else: arenaH) - BotRadius
|
||||
if abs(d) > 1e-9:
|
||||
let cand = if d > 0.0: (hi - p) / d else: (lo - p) / d
|
||||
if cand < t: t = cand
|
||||
if t == Inf: return 0.0
|
||||
max(0.0, t)
|
||||
|
||||
# ── module types ────────────────────────────────────────────────────────────
|
||||
|
||||
type
|
||||
LSWave = object
|
||||
originX, originY: float64
|
||||
bearing: float64 ## enemy -> us at the fire tick (centre line)
|
||||
speed: float64
|
||||
startDist: float64
|
||||
power: float64
|
||||
ticksLeft: int ## ticks until the nominal arrival
|
||||
fresh: bool ## created this tick: do not age it yet
|
||||
stateRow: int ## (vlat, dist) code
|
||||
stateCol: int ## (room, turn) code
|
||||
|
||||
LearnedSurferModule* = object
|
||||
sbc*: Sbc
|
||||
glob*: array[LS_BINS, int]
|
||||
glc: int ## learns since the last global-histogram decay
|
||||
scores: seq[float]
|
||||
waves: seq[LSWave]
|
||||
prevEnergy: seq[tuple[id: int, energy: float]]
|
||||
strafeDir: float64
|
||||
dir: float64 ## direction commanded last tick (+-1)
|
||||
prevX, prevY: float64 ## our position one tick ago
|
||||
prevHeading: float64
|
||||
debugGraphics*: bool
|
||||
decisions*: int ## decisions taken (diagnostic)
|
||||
|
||||
proc resetRound*(m: var LearnedSurferModule) =
|
||||
## Per-ROUND reset: the waves and the smoothed global prior are per round, but
|
||||
## the counted SBC memory is deliberately NOT wiped — the whole point of the
|
||||
## counted+decay mode is that the memory is bounded and decays on its own, so
|
||||
## it survives a round boundary without becoming a battle-long static average
|
||||
## (the j115 defect). Use `resetBattle` for a hard wipe.
|
||||
m.waves = @[]
|
||||
m.prevEnergy = @[]
|
||||
m.strafeDir = 1.0
|
||||
m.dir = 1.0
|
||||
m.prevX = 0.0
|
||||
m.prevY = 0.0
|
||||
m.prevHeading = 0.0
|
||||
m.decisions = 0
|
||||
for i in 0..<LS_BINS: m.glob[i] = 0
|
||||
m.glc = 0
|
||||
|
||||
proc initLearnedSurfer*(): LearnedSurferModule =
|
||||
result.debugGraphics = false
|
||||
result.sbc = initCountedSbc(LS_NADE, LS_BINS,
|
||||
max(0, LearnedDecayEvery),
|
||||
max(0, LearnedDecayShift))
|
||||
result.sbc.decayEvery = max(0, LearnedDecayEvery)
|
||||
result.sbc.decayShift = max(0, LearnedDecayShift)
|
||||
result.scores = newSeq[float](LS_BINS)
|
||||
result.resetRound()
|
||||
|
||||
proc resetBattle*(m: var LearnedSurferModule) =
|
||||
## Hard wipe of the learned memory (round the SBC's counters and `clear`).
|
||||
m.sbc.clear()
|
||||
m.resetRound()
|
||||
|
||||
proc clearGraphics*(m: var LearnedSurferModule) {.inline.} = discard
|
||||
proc removeBulletNear*(m: var LearnedSurferModule, x, y: float) {.inline.} = discard
|
||||
|
||||
# ── the learner ─────────────────────────────────────────────────────────────
|
||||
|
||||
proc prior(m: LearnedSurferModule, k: int): float =
|
||||
var tot = 0
|
||||
for g in m.glob: tot += g
|
||||
(float(m.glob[k]) + 1.0) / (float(tot) + float(LS_BINS))
|
||||
|
||||
proc learnWave(m: var LearnedSurferModule, row, col, bin: int) =
|
||||
## One supervised sample: (state) -> (GF bin at arrival), in the counted SBC.
|
||||
## The SBC applies its own global decay every `decayEvery` learns; the global
|
||||
## histogram below is aged on the same schedule so the two stay comparable.
|
||||
discard m.sbc.learn([row.int32], [col.int32], bin)
|
||||
if m.glob[bin] < 255: inc m.glob[bin]
|
||||
inc m.glc
|
||||
if LearnedDecayEvery > 0 and LearnedDecayShift > 0 and
|
||||
m.glc >= LearnedDecayEvery:
|
||||
for k in 0..<LS_BINS:
|
||||
m.glob[k] = m.glob[k] - (m.glob[k] shr LearnedDecayShift)
|
||||
m.glc = 0
|
||||
|
||||
proc predictState*(m: var LearnedSurferModule, row, col: int,
|
||||
p: var array[LS_BINS, float]) =
|
||||
## p(bin | state) — the counted SBC's per-cell posterior interpolated with the
|
||||
## global histogram. With no data at the cell this is exactly the old surfer.
|
||||
## `p` is filled with the prior FIRST, so the fallback costs nothing.
|
||||
var tot = 0
|
||||
for g in m.glob: tot += g
|
||||
let den = float(tot) + float(LS_BINS)
|
||||
for k in 0..<LS_BINS: p[k] = (float(m.glob[k]) + 1.0) / den
|
||||
if LearnedGlobal: return
|
||||
for k in 0..<LS_BINS: m.scores[k] = 0.0
|
||||
m.sbc.inferProb([row.int32], [col.int32], m.scores)
|
||||
var st = 0.0
|
||||
for k in 0..<LS_BINS: st += m.scores[k]
|
||||
if st <= 0.0: return
|
||||
for k in 0..<LS_BINS:
|
||||
p[k] = (m.scores[k] + LearnedAlpha * p[k]) / (1.0 + LearnedAlpha)
|
||||
|
||||
# ── fire detection + state ──────────────────────────────────────────────────
|
||||
|
||||
proc prevEnergyGet(m: LearnedSurferModule, id: int): float =
|
||||
for e in m.prevEnergy:
|
||||
if e.id == id: return e.energy
|
||||
100.0
|
||||
|
||||
proc prevEnergySet(m: var LearnedSurferModule, id: int, energy: float) =
|
||||
for i in 0..<m.prevEnergy.len:
|
||||
if m.prevEnergy[i].id == id:
|
||||
m.prevEnergy[i].energy = energy
|
||||
return
|
||||
m.prevEnergy.add((id: id, energy: energy))
|
||||
|
||||
proc detectFire(m: var LearnedSurferModule, id: int, ex, ey, eenergy: float,
|
||||
ws: WorldState) =
|
||||
## One enemy's energy sample. A plausible one-tick firepower drop IS a wave;
|
||||
## the state is the wave-relative state at THIS tick (the fire tick).
|
||||
let prev = m.prevEnergyGet(id)
|
||||
let drop = prev - eenergy
|
||||
m.prevEnergySet(id, eenergy)
|
||||
if drop < 0.09 or drop > 3.01: return
|
||||
|
||||
let botX = ws.selfX
|
||||
let botY = ws.selfY
|
||||
let bspeed = 20.0 - 3.0 * drop
|
||||
let d = hypot(botX - ex, botY - ey)
|
||||
let bearing = arctan2(botY - ey, botX - ex) # centre line
|
||||
let ux = cos(bearing)
|
||||
let uy = sin(bearing)
|
||||
|
||||
# vlat: lateral velocity in the wave frame. The centre line passes through us
|
||||
# at this tick, so lat(now) == 0 and vlat == -lat(prev).
|
||||
let dxp = m.prevX - ex
|
||||
let dyp = m.prevY - ey
|
||||
var vlat = -(dxp * (-uy) + dyp * ux)
|
||||
if m.prevX == 0.0 and m.prevY == 0.0 and m.prevHeading == 0.0: vlat = 0.0
|
||||
let roomDx = if vlat >= 0.0: -uy else: uy
|
||||
let roomDy = if vlat >= 0.0: ux else: -ux
|
||||
let room = roomToWall(botX, botY, roomDx, roomDy, ws.arenaWidth, ws.arenaHeight)
|
||||
let turn = wrap180(float(ws.selfHeading) - float(m.prevHeading))
|
||||
|
||||
let row = code(vlat, VlatEdges) * LS_Q + code(d, DistEdges)
|
||||
let col = code(room, RoomEdges) * LS_Q + code(turn, TurnEdges)
|
||||
|
||||
m.waves.add LSWave(
|
||||
originX: ex, originY: ey, bearing: bearing, speed: bspeed,
|
||||
startDist: d, power: drop,
|
||||
ticksLeft: max(1, int(ceil(d / max(bspeed, 1e-9)))),
|
||||
fresh: true,
|
||||
stateRow: row, stateCol: col,
|
||||
)
|
||||
|
||||
# ── the mover ───────────────────────────────────────────────────────────────
|
||||
|
||||
proc computeMove*(m: var LearnedSurferModule, ws: WorldState): MoveCommand =
|
||||
let botX = ws.selfX
|
||||
let botY = ws.selfY
|
||||
|
||||
# ── 1. fire detection (every alive enemy, per-enemy energy) ───────────────
|
||||
var seen = 0
|
||||
for ei in ws.enemies:
|
||||
inc seen
|
||||
m.detectFire(ei.id, ei.x, ei.y, ei.energy, ws)
|
||||
if seen == 0 and (ws.enemyX != 0.0 or ws.enemyY != 0.0):
|
||||
m.detectFire(-1, ws.enemyX, ws.enemyY, ws.enemyEnergy, ws)
|
||||
|
||||
# ── 2. advance + resolve waves; every resolution is a training sample ────
|
||||
var i = 0
|
||||
while i < m.waves.len:
|
||||
if m.waves[i].fresh:
|
||||
m.waves[i].fresh = false # created this tick: not one tick old yet
|
||||
else:
|
||||
dec m.waves[i].ticksLeft
|
||||
if m.waves[i].ticksLeft <= 0:
|
||||
let w = m.waves[i]
|
||||
let maxA = mea(w.speed)
|
||||
if maxA >= 1e-9:
|
||||
let off = wrapPi(arctan2(botY - w.originY, botX - w.originX) - w.bearing)
|
||||
let bin = gfToBin(clamp(off / maxA, -1.0, 1.0))
|
||||
m.learnWave(w.stateRow, w.stateCol, bin)
|
||||
if LearnedLog:
|
||||
echo "[learned] resolve bin=", bin, " state=", w.stateRow, "/",
|
||||
w.stateCol, " d=", w.startDist.int, " e=", w.originX.int, ",",
|
||||
w.originY.int
|
||||
m.waves.del(i)
|
||||
else:
|
||||
inc i
|
||||
|
||||
# ── 3. danger of every candidate bin, summed over every live wave ────────
|
||||
var ps: seq[array[LS_BINS, float]]
|
||||
ps.setLen(m.waves.len)
|
||||
for j in 0..<m.waves.len:
|
||||
m.predictState(m.waves[j].stateRow, m.waves[j].stateCol, ps[j])
|
||||
|
||||
# nearest wave = largest radius/distance ratio (the one about to arrive)
|
||||
var nearest = -1
|
||||
var bestRatio = -1.0
|
||||
for j in 0..<m.waves.len:
|
||||
let dd = max(1e-6, hypot(botX - m.waves[j].originX, botY - m.waves[j].originY))
|
||||
let done = 1.0 - float(m.waves[j].ticksLeft) /
|
||||
float(max(1, int(ceil(m.waves[j].startDist /
|
||||
max(m.waves[j].speed, 1e-9)))))
|
||||
let ratio = done / dd
|
||||
if ratio > bestRatio:
|
||||
bestRatio = ratio
|
||||
nearest = j
|
||||
|
||||
var perpAngle = 0.0
|
||||
|
||||
if nearest >= 0:
|
||||
let w = m.waves[nearest]
|
||||
let dx = botX - w.originX
|
||||
let dy = botY - w.originY
|
||||
let d = max(1.0, hypot(dx, dy))
|
||||
let toBot = arctan2(dy, dx)
|
||||
let maxA = mea(w.speed)
|
||||
let curGF = if maxA >= 1e-9:
|
||||
clamp(wrapPi(toBot - w.bearing) / maxA, -1.0, 1.0)
|
||||
else: 0.0
|
||||
let dodgeDist = max(MaxBotSpeed, ws.selfSpeed) * DodgeTicks
|
||||
|
||||
var bestBin = gfToBin(curGF)
|
||||
var bestDanger = Inf
|
||||
var bestDirs: array[LS_BINS, float]
|
||||
|
||||
for j in 0..<LS_BINS:
|
||||
let gfJ = binToGF(j)
|
||||
let angleJ = w.bearing + gfJ * maxA
|
||||
let px = w.originX + cos(angleJ) * d
|
||||
let py = w.originY + sin(angleJ) * d
|
||||
var ux = px - botX
|
||||
var uy = py - botY
|
||||
let ul = hypot(ux, uy)
|
||||
if ul < 1e-6:
|
||||
ux = cos(perpAngle); uy = sin(perpAngle)
|
||||
else:
|
||||
ux /= ul; uy /= ul
|
||||
let futureX = botX + ux * dodgeDist
|
||||
let futureY = botY + uy * dodgeDist
|
||||
bestDirs[j] = if gfJ >= curGF: 1.0 else: -1.0
|
||||
|
||||
var danger = ps[nearest][j]
|
||||
# every OTHER live wave contributes its own predicted mass at the bin the
|
||||
# candidate direction would land in for THAT wave.
|
||||
for k in 0..<m.waves.len:
|
||||
if k == nearest: continue
|
||||
let w2 = m.waves[k]
|
||||
let d2 = max(1.0, hypot(futureX - w2.originX, futureY - w2.originY))
|
||||
let mA2 = mea(w2.speed)
|
||||
if mA2 < 1e-9: continue
|
||||
let off2 = wrapPi(arctan2(futureY - w2.originY, futureX - w2.originX) -
|
||||
w2.bearing)
|
||||
let b2 = gfToBin(clamp(off2 / mA2, -1.0, 1.0))
|
||||
danger += ps[k][b2]
|
||||
|
||||
let wallHit = futureX < LearnedWallMargin or
|
||||
futureX > ws.arenaWidth - LearnedWallMargin or
|
||||
futureY < LearnedWallMargin or
|
||||
futureY > ws.arenaHeight - LearnedWallMargin
|
||||
if wallHit: danger *= 5.0
|
||||
danger += LearnedTravel * abs(gfJ - curGF)
|
||||
if bestDirs[j] != m.dir: danger += LearnedReversal
|
||||
|
||||
if danger < bestDanger:
|
||||
bestDanger = danger
|
||||
bestBin = j
|
||||
|
||||
let bestGF = binToGF(bestBin)
|
||||
m.strafeDir = bestDirs[bestBin]
|
||||
m.dir = m.strafeDir
|
||||
inc m.decisions
|
||||
if LearnedLog:
|
||||
var pAvg = 0.0
|
||||
for k in 0..<LS_BINS: pAvg += ps[nearest][k] / float(LS_BINS)
|
||||
echo "[learned] wave d=", d.int, " curGF=", (curGF * 100.0).int,
|
||||
" bestGF=", (bestGF * 100.0).int, " dir=", m.strafeDir.int,
|
||||
" pSafe=", (ps[nearest][bestBin] * 1000.0).int,
|
||||
" pCur=", (ps[nearest][gfToBin(curGF)] * 1000.0).int,
|
||||
" pAvg=", (pAvg * 1000.0).int,
|
||||
" state=", m.waves[nearest].stateRow, "/", m.waves[nearest].stateCol
|
||||
# perpendicular in the wave's own frame: +90 deg from origin->bot increases
|
||||
# the GF (CCW), -90 decreases it.
|
||||
perpAngle = if m.strafeDir >= 0.0: toBot + PI * 0.5
|
||||
else: toBot - PI * 0.5
|
||||
elif ws.enemyX != 0.0 or ws.enemyY != 0.0:
|
||||
let toBot = arctan2(botY - ws.enemyY, botX - ws.enemyX)
|
||||
perpAngle = if m.strafeDir >= 0.0: toBot + PI * 0.5
|
||||
else: toBot - PI * 0.5
|
||||
else:
|
||||
perpAngle = degToRad(ws.selfHeading)
|
||||
|
||||
# ── 4. never drive into a wall ───────────────────────────────────────────
|
||||
let nearLeft = botX < LearnedWallMargin
|
||||
let nearRight = botX > ws.arenaWidth - LearnedWallMargin
|
||||
let nearBottom = botY < LearnedWallMargin
|
||||
let nearTop = botY > ws.arenaHeight - LearnedWallMargin
|
||||
if nearLeft or nearRight or nearBottom or nearTop:
|
||||
let px = cos(perpAngle)
|
||||
let py = sin(perpAngle)
|
||||
if (nearLeft and px < 0.0) or (nearRight and px > 0.0) or
|
||||
(nearBottom and py < 0.0) or (nearTop and py > 0.0):
|
||||
m.strafeDir = -m.strafeDir
|
||||
m.dir = m.strafeDir
|
||||
perpAngle = perpAngle + PI
|
||||
let escapeAngle = arctan2(ws.arenaHeight * 0.5 - botY,
|
||||
ws.arenaWidth * 0.5 - botX)
|
||||
let ex = cos(escapeAngle) + cos(perpAngle)
|
||||
let ey = sin(escapeAngle) + sin(perpAngle)
|
||||
perpAngle = arctan2(ey, ex)
|
||||
|
||||
# ── 5. distance control outside the deadband ────────────────────────────
|
||||
let enemyDist = hypot(ws.enemyX - botX, ws.enemyY - botY)
|
||||
let distErr = enemyDist - LearnedPrefDist
|
||||
let radialFrac =
|
||||
if distErr > LearnedDistBand: LearnedRadialFrac # too far -> approach
|
||||
elif distErr < -LearnedDistBand: -LearnedRadialFrac # too close -> retreat
|
||||
else: 0.0
|
||||
if abs(radialFrac) > 1e-9:
|
||||
let radialAngle = arctan2(ws.enemyY - botY, ws.enemyX - botX) +
|
||||
(if radialFrac < 0.0: PI else: 0.0)
|
||||
let rx = cos(perpAngle) * (1.0 - abs(radialFrac)) +
|
||||
cos(radialAngle) * abs(radialFrac)
|
||||
let ry = sin(perpAngle) * (1.0 - abs(radialFrac)) +
|
||||
sin(radialAngle) * abs(radialFrac)
|
||||
perpAngle = arctan2(ry, rx)
|
||||
|
||||
# ── 6. steer the body, full speed ───────────────────────────────────────
|
||||
let desiredDeg = radToDeg(perpAngle)
|
||||
var delta = desiredDeg - ws.selfHeading
|
||||
while delta > 180.0: delta -= 360.0
|
||||
while delta < -180.0: delta += 360.0
|
||||
let goForward = abs(delta) <= 90.0
|
||||
if not goForward:
|
||||
delta = if delta >= 0.0: delta - 180.0 else: delta + 180.0
|
||||
|
||||
m.prevX = botX
|
||||
m.prevY = botY
|
||||
m.prevHeading = float(ws.selfHeading)
|
||||
|
||||
(speed: (if goForward: 8.0 else: -8.0),
|
||||
turnRate: delta.clamp(-10.0, 10.0))
|
||||
+1134
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,85 @@
|
||||
# Learned-surfer Gate A — offline prediction quality (VETO ONLY)
|
||||
|
||||
corpus : /tmp/tfil_ab2/out
|
||||
battles : 70
|
||||
shots used : 54923
|
||||
state fields: vlat, dist, room, turn (5 fields, ONE state, no window)
|
||||
label : 31-bin guess factor at wave resolution (wave_surfer.gfToBin), mode=nominal
|
||||
learner : counted SBC (saturating uint8 + `c -= c shr shift` every decayEvery learns), inferProb readout, alpha=5 prior mix
|
||||
|
||||
## A. held-out prediction quality (mean over 3 battle splits)
|
||||
|
||||
| config | states | log-loss(state) | log-loss(global) | Δ | top-1 st | top-1 glob | top-3 st | top-3 glob |
|
||||
|---|---:|---:|---:|---:|---:|---:|---:|---:|
|
||||
| Q4 decay128/1 (primary) | 256 | 4.9272 | 4.9739 | -0.0467 | 0.0393 | 0.0396 | 0.1224 | 0.1215 |
|
||||
| Q4 decay32/1 | 256 | 4.9594 | 5.0133 | -0.0539 | 0.0370 | 0.0312 | 0.1102 | 0.0970 |
|
||||
| Q4 NO DECAY | 256 | 4.8408 | 4.9221 | -0.0813 | 0.0633 | 0.0463 | 0.1572 | 0.1255 |
|
||||
| Q4 decay128/2 | 256 | 4.8991 | 4.9435 | -0.0444 | 0.0464 | 0.0463 | 0.1234 | 0.1240 |
|
||||
| Q3 decay128/1 | 81 | 4.9401 | 4.9739 | -0.0338 | 0.0396 | 0.0396 | 0.1244 | 0.1215 |
|
||||
| Q5 decay128/1 | 625 | 4.9200 | 4.9739 | -0.0539 | 0.0402 | 0.0396 | 0.1240 | 0.1215 |
|
||||
|
||||
## B. floors (same held-out test sets, primary config splits)
|
||||
|
||||
| predictor | log-loss (bits) | top-1 | top-3 |
|
||||
|---|---:|---:|---:|
|
||||
| chance (uniform 31) | 4.9542 | 0.0323 | 0.0968 |
|
||||
| majority bin | 4.9739 | 0.0463 | n/a |
|
||||
| global 31-bin histogram (the OLD surfer) | 4.9739 | 0.0396 | 0.1215 |
|
||||
| state-conditional counted SBC | 4.9272 | 0.0393 | 0.1224 |
|
||||
|
||||
NOTE: the 'unconditional average' and the '31-bin global histogram of the
|
||||
old surfer' are the SAME estimator by construction (both are the train
|
||||
marginal over bins); they are therefore reported as one row. The majority
|
||||
predictor is the degenerate top-1 version of the same marginal.
|
||||
|
||||
## C. primary config per split (recurrence + paired per-battle stats)
|
||||
|
||||
| seed | train shots | test shots | distinct states | mean count/state | test shots with a SEEN state | Δlog-loss (state-global) |
|
||||
|---|---:|---:|---:|---:|---:|---:|
|
||||
| 0 | 38172 | 16751 | 256 | 149.1 | 100.0% | -0.0454 |
|
||||
| 1 | 38312 | 16611 | 256 | 149.7 | 100.0% | -0.0465 |
|
||||
| 2 | 38373 | 16550 | 256 | 149.9 | 100.0% | -0.0482 |
|
||||
|
||||
### paired per-battle statistics (primary config, all splits pooled)
|
||||
|
||||
| metric | n battles | mean Δ | SD | 95% CI | sign | p(sign) | p(sign-flip) | MDE |
|
||||
|---|---:|---:|---:|---|---:|---:|---:|---:|
|
||||
| logloss(state-global), pooled | 63 | -0.0467 | 0.0058 | [-0.0481, -0.0453] | 0/63 | 2.168e-19 | 5e-05 | 0.0021 |
|
||||
| logloss(state-global) seed0 | 21 | -0.0454 | 0.0048 | [-0.0475, -0.0433] | 0/21 | 9.537e-07 | 5e-05 | 0.0029 |
|
||||
| logloss(state-global) seed1 | 21 | -0.0465 | 0.0060 | [-0.0490, -0.0439] | 0/21 | 9.537e-07 | 5e-05 | 0.0037 |
|
||||
| logloss(state-global) seed2 | 21 | -0.0482 | 0.0065 | [-0.0510, -0.0454] | 0/21 | 9.537e-07 | 5e-05 | 0.0040 |
|
||||
(negative Δ = the state-conditional model predicts better)
|
||||
|
||||
## D. label-shuffle control (same states, train labels permuted)
|
||||
|
||||
| arm | log-loss(state) | log-loss(global) | Δ | top-1(state) |
|
||||
|---|---:|---:|---:|---:|
|
||||
| shuffled labels | 4.9480 | 4.9535 | -0.0056 | 0.0384 |
|
||||
| real labels (Q4 decay128) | 4.9272 | 4.9739 | -0.0467 | 0.0393 |
|
||||
|
||||
(chance log-loss floor = 4.9542 bits, target entropy = the
|
||||
global-histogram log-loss above; a shuffled-label state model must fall
|
||||
back to it.)
|
||||
|
||||
## E. canonical state-bin edges (4 symbols/field), hard-coded into the
|
||||
module `common_libs/movements/learned_surfer.nim` (derived from the whole
|
||||
corpus; the gate numbers above use TRAIN-only edges per split, so the
|
||||
report is not conditioned on these)
|
||||
|
||||
```nim
|
||||
vlat : -6.736, 0.000, 6.753
|
||||
dist : 431.321, 487.612, 552.670
|
||||
room : 137.965, 206.589, 296.753
|
||||
turn : -0.142, 0.000, 0.105
|
||||
```
|
||||
|
||||
cross-check with the canonical edges: log-loss(state) 4.9272, log-loss(global) 4.9739, Δ -0.0467, top-1 0.0393
|
||||
|
||||
## MEASURED vs INFERRED
|
||||
|
||||
* MEASURED: every number in this file, produced by the command in the
|
||||
module docstring on the recorded corpus.
|
||||
* INFERRED: that this offline prediction-quality result transfers to the
|
||||
LIVE closed loop. It cannot: the recorded trajectory was produced while
|
||||
the enemy gun reacted to a DIFFERENT mover (see
|
||||
docs/offline_harness_trust.md §0/§4). This gate is a veto only.
|
||||
@@ -0,0 +1,598 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Gate A (OFFLINE VETO) for the state-conditional LEARNED MOVEMENT module.
|
||||
|
||||
Question: on the recorded corpus, does a *state-conditional* model predict the
|
||||
GUESS-FACTOR BIN at which an incoming wave crosses the dodging bot **better**
|
||||
than (i) the unconditional average / (ii) the 31-bin global histogram that the
|
||||
old crude surfer (`common_libs/movements/wave_surfer.nim`, job j115) used, and
|
||||
(iii) chance?
|
||||
|
||||
Per `docs/offline_harness_trust.md` this harness is VETO-ONLY: it may kill the
|
||||
design; it can never select it. The live panel is the decider.
|
||||
|
||||
--------------------------------------------------------------------------------
|
||||
WHAT IS MEASURED (all per shot, no replay, no simulator)
|
||||
--------------------------------------------------------------------------------
|
||||
Corpus: `/tmp/tfil_ab2/out/<A..E>/runN.jsonl{,.events.jsonl,.rounds.json}` — 70
|
||||
recorded battles. Per-shot geometry is re-derived with the validated instrument
|
||||
`common_libs/tests/analyze_drussgt_dodge_vs_power.py` (its `Run` class), so the
|
||||
attribution of a fire event to a side is the one already documented in
|
||||
`docs/drussgt_dodge_vs_power.md`.
|
||||
|
||||
Frame: the module cannot see bullet bodies (`WorldState` has none), so the wave
|
||||
frame is the one `wave_surfer.nim` actually uses — origin = the shooter's
|
||||
position at the fire tick, "centre line" = the bearing from the origin to the
|
||||
TARGET at the fire tick. The bullet's own direction is NOT used (using it would
|
||||
leak the enemy's aim).
|
||||
|
||||
Label (what `wave_surfer.gfToBin` computes at wave resolution): walk ticks
|
||||
after the fire; the wave resolves on the first tick where `speed*k >=
|
||||
dist(origin, target)`; the label is the 31-bin quantisation of
|
||||
`wrap180(bearing(origin -> target at that tick) - centreBearing) / maxEA`.
|
||||
|
||||
State at the FIRE tick (5 fields, each quantised into Q symbols; the same
|
||||
wave-relative state family that survived `docs/state_window_gate.md` §1):
|
||||
lat signed perpendicular offset from the centre line (px)
|
||||
vlat lat(t) - lat(t-1) (px/tick, signed)
|
||||
dist range to the shooter (px)
|
||||
room directional wall room along sign(vlat)*normal (px)
|
||||
turn signed heading change (deg/tick)
|
||||
A SINGLE state is used -- never a temporal window of states (measured dead in
|
||||
`docs/state_window_gate.md`).
|
||||
|
||||
Learner: the COUNTED SBC with global fractional decay that the module uses
|
||||
(`common_libs/bitbrain/sbc.nim`, counted mode): one saturating uint8 counter per
|
||||
(state, bin), `c -= c shr shift` every `decayEvery` learns, read out with the
|
||||
per-cell posterior (`inferProb`) interpolated with the global prior, i.e.
|
||||
p = (counts[state]/n_state + alpha*prior) / (1 + alpha). This Python model is
|
||||
the same estimator the Nim module runs; `--selftest` cross-checks it.
|
||||
|
||||
Split: BY BATTLE (never by tick inside a round), 70/30, 3 seeds.
|
||||
|
||||
Run:
|
||||
python3 common_libs/tests/learned_surfer_gate.py \
|
||||
--corpus /tmp/tfil_ab2/out \
|
||||
--report common_libs/tests/fixtures/learned_surfer_gate_report.txt \
|
||||
--json common_libs/tests/fixtures/learned_surfer_gate.json
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
import statistics
|
||||
import sys
|
||||
from collections import Counter
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
import analyze_drussgt_dodge_vs_power as adp # validated per-shot geometry
|
||||
|
||||
NBINS = 31
|
||||
MAX_FLIGHT = 220
|
||||
NBOT = 18.0
|
||||
ARENA_W, ARENA_H = 800.0, 600.0
|
||||
FIELDS = ["vlat", "dist", "room", "turn"] # `lat` at the fire tick is identically 0 (the centre line passes through us), so it is NOT a field
|
||||
|
||||
|
||||
# --------------------------------------------------------------------- helpers
|
||||
def wrap180(a):
|
||||
return ((a + 180.0) % 360.0) - 180.0
|
||||
|
||||
|
||||
def room_to_wall(px, py, dx, dy):
|
||||
"""Distance from (px,py) along unit (dx,dy) until leaving the arena."""
|
||||
t = float("inf")
|
||||
for p, d, lo, hi in ((px, dx, NBOT, ARENA_W - NBOT),
|
||||
(py, dy, NBOT, ARENA_H - NBOT)):
|
||||
if abs(d) > 1e-9:
|
||||
cand = (hi - p) / d if d > 0 else (lo - p) / d
|
||||
if cand < t:
|
||||
t = cand
|
||||
return 0.0 if t == float("inf") else max(0.0, t)
|
||||
|
||||
|
||||
def gf_to_bin(gf):
|
||||
"""Exactly `wave_surfer.gfToBin` (31 bins over [-1, +1])."""
|
||||
v = max(-1.0, min(1.0, gf))
|
||||
return max(0, min(NBINS - 1, int(round((v + 1.0) * 0.5 * (NBINS - 1)))))
|
||||
|
||||
|
||||
def bin_to_gf(idx):
|
||||
return idx / (NBINS - 1) * 2.0 - 1.0
|
||||
|
||||
|
||||
# ------------------------------------------------------------------ extraction
|
||||
def extract(run, label_mode="nominal"):
|
||||
"""One record per shot: state fields at the fire tick + the GF-bin label.
|
||||
|
||||
label_mode:
|
||||
`nominal` -- GF of the target's position at the NOMINAL arrival tick
|
||||
`t0 + ceil(startDist/speed)` (the standard "wave hit GF":
|
||||
the time the bullet needs for the range at fire time, known
|
||||
at fire time).
|
||||
`resolve` -- GF at the tick the expanding wave circle catches the target
|
||||
(`speed*k >= dist(origin,target)`), i.e. exactly
|
||||
`wave_surfer.nim`'s resolution rule.
|
||||
"""
|
||||
out = []
|
||||
for s in run.shots(): # shots fired by side "s" at side "e"
|
||||
t0, rnd = s["tick"], s["rnd"]
|
||||
speed = s["speed"]
|
||||
p0x, p0y = s["_x"], s["_y"]
|
||||
rs = run.start[rnd]
|
||||
re = rs + run.count[rnd] - 1
|
||||
r0 = run.by_tick.get(t0)
|
||||
rp = run.by_tick.get(t0 - 1)
|
||||
if r0 is None or rp is None or t0 <= rs:
|
||||
continue
|
||||
|
||||
# centre line = origin -> target at the fire tick (the wave frame)
|
||||
tb0 = math.atan2(r0["ey"] - p0y, r0["ex"] - p0x)
|
||||
ux, uy = math.cos(tb0), math.sin(tb0)
|
||||
|
||||
def lat_of(r):
|
||||
dx, dy = r["ex"] - p0x, r["ey"] - p0y
|
||||
return dx * (-uy) + dy * ux # cross(u, d)
|
||||
|
||||
lat = lat_of(r0)
|
||||
vlat = lat - lat_of(rp)
|
||||
dist = math.hypot(r0["ex"] - p0x, r0["ey"] - p0y)
|
||||
sg = 1.0 if vlat >= 0 else -1.0
|
||||
room = room_to_wall(r0["ex"], r0["ey"], -uy * sg, ux * sg)
|
||||
turn = wrap180(r0["eh"] - rp["eh"])
|
||||
|
||||
# label: the module's own rule, two candidate conventions
|
||||
maxea = math.asin(min(8.0 / speed, 1.0))
|
||||
lab = None
|
||||
if label_mode == "nominal":
|
||||
kfix = max(1, int(math.ceil(dist / speed)))
|
||||
r = run.by_tick.get(t0 + kfix)
|
||||
if r is not None and t0 + kfix <= re:
|
||||
off = wrap180(math.atan2(r["ey"] - p0y, r["ex"] - p0x) - tb0)
|
||||
gf = max(-1.0, min(1.0, off / maxea)) if maxea > 1e-9 else 0.0
|
||||
lab = gf_to_bin(gf)
|
||||
else:
|
||||
for k in range(1, MAX_FLIGHT):
|
||||
if t0 + k > re:
|
||||
break
|
||||
r = run.by_tick.get(t0 + k)
|
||||
if r is None:
|
||||
break
|
||||
if speed * k >= math.hypot(r["ex"] - p0x, r["ey"] - p0y):
|
||||
off = wrap180(math.atan2(r["ey"] - p0y, r["ex"] - p0x) - tb0)
|
||||
gf = max(-1.0, min(1.0, off / maxea)) if maxea > 1e-9 else 0.0
|
||||
lab = gf_to_bin(gf)
|
||||
break
|
||||
if lab is None:
|
||||
continue
|
||||
|
||||
out.append(dict(
|
||||
battle=os.path.basename(os.path.dirname(run.cap_path)) + "/" +
|
||||
os.path.basename(run.cap_path).replace(".jsonl", ""),
|
||||
rnd=rnd, tick=t0, label=lab,
|
||||
lat=lat, vlat=vlat, dist=dist, room=room, turn=turn,
|
||||
))
|
||||
return out
|
||||
|
||||
|
||||
# ------------------------------------------------------------------ quantiser
|
||||
def edges_for(vals, q):
|
||||
xs = sorted(vals)
|
||||
n = len(xs)
|
||||
return [xs[min(n - 1, int((qi / q) * n))] for qi in range(1, q)]
|
||||
|
||||
|
||||
def edges_of(samples, q):
|
||||
return {f: edges_for([s[f] for s in samples], q) for f in FIELDS}
|
||||
|
||||
|
||||
def code_of(sample, edges, q):
|
||||
code = 0
|
||||
for f in FIELDS:
|
||||
e = edges[f]
|
||||
v = sample[f]
|
||||
c = 0
|
||||
for edge in e:
|
||||
if v > edge:
|
||||
c += 1
|
||||
code = code * q + c
|
||||
return code
|
||||
|
||||
|
||||
# -------------------------------------------------------------------- learner
|
||||
class CountedModel:
|
||||
"""Counted SBC (one cell per state) + global prior, exactly as the module
|
||||
runs it: saturating uint8 counters, `c -= c shr shift` every decayEvery
|
||||
learns, per-cell posterior interpolated with the global histogram."""
|
||||
|
||||
def __init__(self, nstates, decay_every=0, decay_shift=1, alpha=5.0):
|
||||
self.n = nstates
|
||||
self.d = decay_every
|
||||
self.sh = decay_shift
|
||||
self.alpha = alpha
|
||||
self.c = [[0] * NBINS for _ in range(nstates)]
|
||||
self.g = [0] * NBINS
|
||||
self.lc = 0
|
||||
self.glc = 0
|
||||
|
||||
def _decay(self, row):
|
||||
sh = self.sh
|
||||
if sh <= 0:
|
||||
return
|
||||
for k in range(NBINS):
|
||||
row[k] -= row[k] >> sh
|
||||
|
||||
def learn(self, state, label):
|
||||
row = self.c[state]
|
||||
if row[label] < 255:
|
||||
row[label] += 1
|
||||
self.g[label] += 1
|
||||
self.lc += 1
|
||||
self.glc += 1
|
||||
if self.d > 0 and self.lc >= self.d:
|
||||
for r in self.c:
|
||||
self._decay(r)
|
||||
self._decay(self.g)
|
||||
self.lc = 0
|
||||
if self.d > 0 and self.glc >= self.d:
|
||||
self.glc = 0
|
||||
|
||||
def prior(self):
|
||||
tot = sum(self.g)
|
||||
if tot == 0:
|
||||
return [1.0 / NBINS] * NBINS
|
||||
return [(self.g[k] + 1.0) / (tot + NBINS) for k in range(NBINS)]
|
||||
|
||||
def predict(self, state):
|
||||
pr = self.prior()
|
||||
n = sum(self.c[state])
|
||||
if n == 0:
|
||||
return pr
|
||||
a = self.alpha
|
||||
return [((self.c[state][k] / n) + a * pr[k]) / (1.0 + a)
|
||||
for k in range(NBINS)]
|
||||
|
||||
|
||||
def log2(x):
|
||||
return math.log2(max(x, 1e-12))
|
||||
|
||||
|
||||
def evaluate(model, samples, edges, q):
|
||||
lls, t1, t3 = [], 0, 0
|
||||
for s in samples:
|
||||
p = model.predict(code_of(s, edges, q))
|
||||
y = s["label"]
|
||||
lls.append(-log2(p[y]))
|
||||
order = sorted(range(NBINS), key=lambda k: -p[k])
|
||||
if order[0] == y:
|
||||
t1 += 1
|
||||
if y in order[:3]:
|
||||
t3 += 1
|
||||
n = len(samples)
|
||||
return dict(logloss=statistics.fmean(lls) if lls else float("nan"),
|
||||
top1=t1 / n if n else float("nan"),
|
||||
top3=t3 / n if n else float("nan"))
|
||||
|
||||
|
||||
def by_battle_ll(samples, edges, q, predict):
|
||||
"""per-battle mean log-loss (paired statistics need one number per battle)."""
|
||||
acc = {}
|
||||
for s in samples:
|
||||
p = predict(code_of(s, edges, q))
|
||||
acc.setdefault(s["battle"], []).append(-log2(p[s["label"]]))
|
||||
return {b: statistics.fmean(v) for b, v in acc.items() if v}
|
||||
|
||||
|
||||
def sign_flip_p(deltas, reps=20000, seed=7):
|
||||
"""Two-sided sign-flip permutation on paired per-battle deltas."""
|
||||
obs = statistics.fmean(deltas)
|
||||
rng = random.Random(seed)
|
||||
hits = 0
|
||||
for _ in range(reps):
|
||||
s = statistics.fmean([d if rng.random() < 0.5 else -d for d in deltas])
|
||||
if abs(s) >= abs(obs) - 1e-12:
|
||||
hits += 1
|
||||
return (hits + 1) / (reps + 1)
|
||||
|
||||
|
||||
def sign_test_p(deltas):
|
||||
pos = sum(1 for d in deltas if d > 0)
|
||||
neg = sum(1 for d in deltas if d < 0)
|
||||
n = pos + neg
|
||||
if n == 0:
|
||||
return 1.0, 0, 0
|
||||
tail = sum(math.comb(n, k) for k in range(min(pos, neg) + 1)) / 2 ** n
|
||||
return min(1.0, 2 * tail), pos, n
|
||||
|
||||
|
||||
def paired_stats(deltas, name, unit="battles"):
|
||||
n = len(deltas)
|
||||
m = statistics.fmean(deltas)
|
||||
sd = statistics.stdev(deltas) if n > 1 else 0.0
|
||||
se = sd / math.sqrt(n) if n else float("nan")
|
||||
ps, pos, nd = sign_test_p(deltas)
|
||||
return dict(metric=name, n=n, mean=m, sd=sd, se=se,
|
||||
ci=[m - 1.96 * se, m + 1.96 * se],
|
||||
sign=f"{pos}/{nd}", p_sign=ps, p_signflip=sign_flip_p(deltas),
|
||||
mde=2.8 * se, unit=unit)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------- driver
|
||||
def split_battles(battles, seed, frac=0.7):
|
||||
bs = sorted(battles)
|
||||
rng = random.Random(seed)
|
||||
rng.shuffle(bs)
|
||||
k = int(len(bs) * frac)
|
||||
return set(bs[:k]), set(bs[k:])
|
||||
|
||||
|
||||
def run_config(samples, seeds, q, decay_every, decay_shift, alpha,
|
||||
fixed_edges=None, shuffle_labels=False):
|
||||
per_seed = []
|
||||
model_dump = None
|
||||
for seed in seeds:
|
||||
battles = {s["battle"] for s in samples}
|
||||
tr_b, te_b = split_battles(battles, seed)
|
||||
tr = [s for s in samples if s["battle"] in tr_b]
|
||||
te = [s for s in samples if s["battle"] in te_b]
|
||||
if not tr or not te:
|
||||
continue
|
||||
edges = fixed_edges or edges_of(tr, q)
|
||||
|
||||
if shuffle_labels:
|
||||
# control: same states, labels permuted inside the train set
|
||||
rng = random.Random(1000 + seed)
|
||||
labs = [s["label"] for s in tr]
|
||||
rng.shuffle(labs)
|
||||
tr = [dict(s, label=l) for s, l in zip(tr, labs)]
|
||||
|
||||
nstates = q ** len(FIELDS)
|
||||
mod = CountedModel(nstates, decay_every, decay_shift, alpha)
|
||||
for s in tr:
|
||||
mod.learn(code_of(s, edges, q), s["label"])
|
||||
|
||||
glob = CountedModel(nstates, decay_every, decay_shift, alpha)
|
||||
for s in tr:
|
||||
glob.learn(0, s["label"]) # state forced to 0 = global only
|
||||
|
||||
maj = Counter(s["label"] for s in tr).most_common(1)[0][0]
|
||||
# majority as a degenerate top-1 predictor; log-loss from the prior
|
||||
gp = glob.prior()
|
||||
|
||||
m_state = evaluate(mod, te, edges, q)
|
||||
m_glob = evaluate(glob, te, edges, q)
|
||||
chance = dict(logloss=log2(NBINS), top1=1.0 / NBINS, top3=3.0 / NBINS)
|
||||
m_maj = dict(logloss=statistics.fmean([-log2(gp[s["label"]]) for s in te]),
|
||||
top1=sum(1 for s in te if s["label"] == maj) / len(te),
|
||||
top3=float("nan"))
|
||||
# paired per-battle deltas (state - global), lower log-loss is better
|
||||
ll_s = by_battle_ll(te, edges, q, mod.predict)
|
||||
ll_g = by_battle_ll(te, edges, q, glob.predict)
|
||||
common = sorted(set(ll_s) & set(ll_g))
|
||||
deltas = [ll_s[b] - ll_g[b] for b in common]
|
||||
|
||||
seen = sum(1 for s in te if sum(mod.c[code_of(s, edges, q)]) > 0)
|
||||
recur = Counter(code_of(s, edges, q) for s in tr)
|
||||
per_seed.append(dict(
|
||||
seed=seed, n_train=len(tr), n_test=len(te),
|
||||
state=m_state, glob=m_glob, chance=chance, maj=m_maj,
|
||||
delta_logloss=statistics.fmean(deltas),
|
||||
stats=paired_stats(deltas, "logloss(state-global)"),
|
||||
seen_frac=seen / len(te),
|
||||
distinct_states=len(recur), mean_count=len(tr) / len(recur),
|
||||
top_count=max(recur.values()), marginal_entropy_bits=entropy(tr),
|
||||
deltas=deltas,
|
||||
))
|
||||
if model_dump is None:
|
||||
model_dump = dict(edges=edges, q=q)
|
||||
agg = {}
|
||||
if per_seed:
|
||||
for key in ("logloss", "top1", "top3"):
|
||||
agg["state_" + key] = statistics.fmean(p["state"][key] for p in per_seed)
|
||||
agg["glob_" + key] = statistics.fmean(p["glob"][key] for p in per_seed)
|
||||
agg["chance_" + key] = statistics.fmean(p["chance"][key] for p in per_seed)
|
||||
agg["maj_" + key] = statistics.fmean(p["maj"][key] for p in per_seed)
|
||||
agg["delta_logloss"] = statistics.fmean(p["delta_logloss"] for p in per_seed)
|
||||
agg["seen_frac"] = statistics.fmean(p["seen_frac"] for p in per_seed)
|
||||
agg["distinct_states"] = statistics.fmean(p["distinct_states"] for p in per_seed)
|
||||
agg["mean_count"] = statistics.fmean(p["mean_count"] for p in per_seed)
|
||||
agg["entropy_bits"] = statistics.fmean(p["marginal_entropy_bits"] for p in per_seed)
|
||||
# paired per-battle stats, pooled across every (seed, battle) delta
|
||||
all_d = [d for p in per_seed for d in p["deltas"]]
|
||||
agg["stats"] = paired_stats(all_d, "logloss(state-global)")
|
||||
return agg, per_seed
|
||||
|
||||
|
||||
def entropy(samples):
|
||||
c = Counter(s["label"] for s in samples)
|
||||
n = sum(c.values())
|
||||
return -sum((v / n) * math.log2(v / n) for v in c.values())
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--corpus", default="/tmp/tfil_ab2/out")
|
||||
ap.add_argument("--report", default=None)
|
||||
ap.add_argument("--json", default=None)
|
||||
ap.add_argument("--seeds", type=int, default=3)
|
||||
ap.add_argument("--label", default="nominal",
|
||||
choices=["nominal", "resolve", "both"])
|
||||
args = ap.parse_args()
|
||||
|
||||
if args.label == "both":
|
||||
for mode in ("nominal", "resolve"):
|
||||
sub = [a for a in sys.argv if not a.startswith("--label")]
|
||||
sub += ["--label", mode]
|
||||
print(f"\n########## label mode: {mode} ##########")
|
||||
os.execv(sys.executable, [sys.executable] + sub)
|
||||
return
|
||||
|
||||
lines = []
|
||||
|
||||
def out(s=""):
|
||||
print(s)
|
||||
lines.append(s)
|
||||
|
||||
runs = adp.discover_tfil(args.corpus)
|
||||
samples = []
|
||||
for r in runs:
|
||||
samples.extend(extract(r, args.label))
|
||||
out("# Learned-surfer Gate A — offline prediction quality (VETO ONLY)")
|
||||
out()
|
||||
out(f"corpus : {args.corpus}")
|
||||
out(f"battles : {len(runs)}")
|
||||
out(f"shots used : {len(samples)}")
|
||||
out(f"state fields: {', '.join(FIELDS)} (5 fields, ONE state, no window)")
|
||||
out(f"label : 31-bin guess factor at wave resolution "
|
||||
f"(wave_surfer.gfToBin), mode={args.label}")
|
||||
out(f"learner : counted SBC (saturating uint8 + `c -= c shr shift` "
|
||||
f"every decayEvery learns), inferProb readout, alpha=5 prior mix")
|
||||
out()
|
||||
|
||||
seeds = list(range(args.seeds))
|
||||
configs = [
|
||||
("Q4 decay128/1 (primary)", 4, 128, 1, 5.0),
|
||||
("Q4 decay32/1", 4, 32, 1, 5.0),
|
||||
("Q4 NO DECAY", 4, 0, 1, 5.0),
|
||||
("Q4 decay128/2", 4, 128, 2, 5.0),
|
||||
("Q3 decay128/1", 3, 128, 1, 5.0),
|
||||
("Q5 decay128/1", 5, 128, 1, 5.0),
|
||||
]
|
||||
results = {}
|
||||
out("## A. held-out prediction quality (mean over 3 battle splits)")
|
||||
out()
|
||||
out("| config | states | log-loss(state) | log-loss(global) | Δ | top-1 st | "
|
||||
"top-1 glob | top-3 st | top-3 glob |")
|
||||
out("|---|---:|---:|---:|---:|---:|---:|---:|---:|")
|
||||
for name, q, de, ds, al in configs:
|
||||
agg, per = run_config(samples, seeds, q, de, ds, al)
|
||||
results[name] = dict(agg=agg, per=per)
|
||||
out(f"| {name} | {agg['distinct_states']:.0f} | {agg['state_logloss']:.4f} | "
|
||||
f"{agg['glob_logloss']:.4f} | {agg['delta_logloss']:+.4f} | "
|
||||
f"{agg['state_top1']:.4f} | {agg['glob_top1']:.4f} | "
|
||||
f"{agg['state_top3']:.4f} | {agg['glob_top3']:.4f} |")
|
||||
out()
|
||||
|
||||
# reference floors from the primary config's splits
|
||||
pr = results[configs[0][0]]["per"]
|
||||
out("## B. floors (same held-out test sets, primary config splits)")
|
||||
out()
|
||||
out("| predictor | log-loss (bits) | top-1 | top-3 |")
|
||||
out("|---|---:|---:|---:|")
|
||||
out(f"| chance (uniform 31) | {log2(NBINS):.4f} | {1/NBINS:.4f} | {3/NBINS:.4f} |")
|
||||
out(f"| majority bin | "
|
||||
f"{statistics.fmean(p['maj']['logloss'] for p in pr):.4f} | "
|
||||
f"{statistics.fmean(p['maj']['top1'] for p in pr):.4f} | n/a |")
|
||||
out(f"| global 31-bin histogram (the OLD surfer) | "
|
||||
f"{statistics.fmean(p['glob']['logloss'] for p in pr):.4f} | "
|
||||
f"{statistics.fmean(p['glob']['top1'] for p in pr):.4f} | "
|
||||
f"{statistics.fmean(p['glob']['top3'] for p in pr):.4f} |")
|
||||
out(f"| state-conditional counted SBC | "
|
||||
f"{statistics.fmean(p['state']['logloss'] for p in pr):.4f} | "
|
||||
f"{statistics.fmean(p['state']['top1'] for p in pr):.4f} | "
|
||||
f"{statistics.fmean(p['state']['top3'] for p in pr):.4f} |")
|
||||
out()
|
||||
out("NOTE: the 'unconditional average' and the '31-bin global histogram of the")
|
||||
out("old surfer' are the SAME estimator by construction (both are the train")
|
||||
out("marginal over bins); they are therefore reported as one row. The majority")
|
||||
out("predictor is the degenerate top-1 version of the same marginal.")
|
||||
out()
|
||||
|
||||
# per-seed detail + recurrence for the primary config
|
||||
out("## C. primary config per split (recurrence + paired per-battle stats)")
|
||||
out()
|
||||
out("| seed | train shots | test shots | distinct states | mean count/state | "
|
||||
"test shots with a SEEN state | Δlog-loss (state-global) |")
|
||||
out("|---|---:|---:|---:|---:|---:|---:|")
|
||||
for p in pr:
|
||||
out(f"| {p['seed']} | {p['n_train']} | {p['n_test']} | {p['distinct_states']} | "
|
||||
f"{p['mean_count']:.1f} | {p['seen_frac']*100:.1f}% | "
|
||||
f"{p['delta_logloss']:+.4f} |")
|
||||
out()
|
||||
out("### paired per-battle statistics (primary config, all splits pooled)")
|
||||
out()
|
||||
out("| metric | n battles | mean Δ | SD | 95% CI | sign | p(sign) | "
|
||||
"p(sign-flip) | MDE |")
|
||||
out("|---|---:|---:|---:|---|---:|---:|---:|---:|")
|
||||
st = results[configs[0][0]]["agg"]["stats"]
|
||||
out(f"| logloss(state-global), pooled | {st['n']} | {st['mean']:+.4f} | "
|
||||
f"{st['sd']:.4f} | [{st['ci'][0]:+.4f}, {st['ci'][1]:+.4f}] | "
|
||||
f"{st['sign']} | {st['p_sign']:.4g} | {st['p_signflip']:.4g} | "
|
||||
f"{st['mde']:.4f} |")
|
||||
for p in pr:
|
||||
s2 = p["stats"]
|
||||
out(f"| logloss(state-global) seed{p['seed']} | {s2['n']} | {s2['mean']:+.4f} | "
|
||||
f"{s2['sd']:.4f} | [{s2['ci'][0]:+.4f}, {s2['ci'][1]:+.4f}] | "
|
||||
f"{s2['sign']} | {s2['p_sign']:.4g} | {s2['p_signflip']:.4g} | "
|
||||
f"{s2['mde']:.4f} |")
|
||||
out("(negative Δ = the state-conditional model predicts better)")
|
||||
out()
|
||||
|
||||
# label-shuffle control
|
||||
out("## D. label-shuffle control (same states, train labels permuted)")
|
||||
out()
|
||||
aggq, _ = run_config(samples, seeds, 4, 128, 1, 5.0, shuffle_labels=True)
|
||||
out("| arm | log-loss(state) | log-loss(global) | Δ | top-1(state) |")
|
||||
out("|---|---:|---:|---:|---:|")
|
||||
out(f"| shuffled labels | {aggq['state_logloss']:.4f} | {aggq['glob_logloss']:.4f} | "
|
||||
f"{aggq['delta_logloss']:+.4f} | {aggq['state_top1']:.4f} |")
|
||||
out(f"| real labels (Q4 decay128) | "
|
||||
f"{results['Q4 decay128/1 (primary)']['agg']['state_logloss']:.4f} | "
|
||||
f"{results['Q4 decay128/1 (primary)']['agg']['glob_logloss']:.4f} | "
|
||||
f"{results['Q4 decay128/1 (primary)']['agg']['delta_logloss']:+.4f} | "
|
||||
f"{results['Q4 decay128/1 (primary)']['agg']['state_top1']:.4f} |")
|
||||
out()
|
||||
out(f"(chance log-loss floor = {log2(NBINS):.4f} bits, target entropy = the")
|
||||
out("global-histogram log-loss above; a shuffled-label state model must fall")
|
||||
out("back to it.)")
|
||||
out()
|
||||
|
||||
# canonical edges for the MODULE
|
||||
edges = edges_of(samples, 4)
|
||||
out("## E. canonical state-bin edges (4 symbols/field), hard-coded into the")
|
||||
out("module `common_libs/movements/learned_surfer.nim` (derived from the whole")
|
||||
out("corpus; the gate numbers above use TRAIN-only edges per split, so the")
|
||||
out("report is not conditioned on these)")
|
||||
out()
|
||||
out("```nim")
|
||||
for f in FIELDS:
|
||||
out(f"{f:5s}: {', '.join(f'{e:.3f}' for e in edges[f])}")
|
||||
out("```")
|
||||
out()
|
||||
fixed_agg, _ = run_config(samples, seeds, 4, 128, 1, 5.0, fixed_edges=edges)
|
||||
out(f"cross-check with the canonical edges: log-loss(state) "
|
||||
f"{fixed_agg['state_logloss']:.4f}, log-loss(global) "
|
||||
f"{fixed_agg['glob_logloss']:.4f}, Δ {fixed_agg['delta_logloss']:+.4f}, "
|
||||
f"top-1 {fixed_agg['state_top1']:.4f}")
|
||||
out()
|
||||
|
||||
out("## MEASURED vs INFERRED")
|
||||
out()
|
||||
out("* MEASURED: every number in this file, produced by the command in the")
|
||||
out(" module docstring on the recorded corpus.")
|
||||
out("* INFERRED: that this offline prediction-quality result transfers to the")
|
||||
out(" LIVE closed loop. It cannot: the recorded trajectory was produced while")
|
||||
out(" the enemy gun reacted to a DIFFERENT mover (see")
|
||||
out(" docs/offline_harness_trust.md §0/§4). This gate is a veto only.")
|
||||
|
||||
if args.report:
|
||||
os.makedirs(os.path.dirname(args.report), exist_ok=True)
|
||||
with open(args.report, "w") as f:
|
||||
f.write("\n".join(lines) + "\n")
|
||||
if args.json:
|
||||
os.makedirs(os.path.dirname(args.json), exist_ok=True)
|
||||
with open(args.json, "w") as f:
|
||||
json.dump(dict(
|
||||
corpus=args.corpus, battles=len(runs), shots=len(samples),
|
||||
results={k: v["agg"] for k, v in results.items()},
|
||||
per_seed={k: [{kk: vv for kk, vv in p.items() if kk != "deltas"}
|
||||
for p in v["per"]] for k, v in results.items()},
|
||||
shuffle_control=aggq, edges=edges,
|
||||
), f, indent=1, default=str)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user