Files
SirRoboGarage/common_libs/tests/measure_bitbrain_gate.nim

649 lines
26 KiB
Nim

## OFFLINE GATE TEST (read-only fixtures, no Java, no battle, no server).
##
## QUESTION: can the generic BitBrain (ADE + SBC) library at
## `common_libs/bitbrain/` predict a FINE-GRAINED AIM CORRECTION better than
## (a) the straight-line naive extrapolation, (b) an always-the-same-answer
## fixed correction, and (c) the shipped Pattern gun's own prediction?
##
## INPUT = the SAME 53 bits the TMHorizon gun uses. They are harvested from the
## LIVE `TmHorizonGun` itself (the per-tick cached-bits + 4-bit horizon
## one-hot path), not re-derived, so the comparison to the TM gun we
## already measured is apples-to-apples.
## OUTPUT = a fine-grained angular-correction class N over a fixed range. The
## KEY readout is the COUNT-WEIGHTED MEAN of the per-class set-bit
## counts (a soft continuous predictor); the argmax is reported too.
## LABEL = the +h-tick FACT exactly as TMHorizon defines it: h = round(dist /
## speed), speed = 20 - 3*power; at tick t the enemy's ACTUAL angular
## offset from Pattern's prediction is looked up at t+h. Never crosses
## a round boundary.
## METRIC = arrival aim error in PIXELS and DEGREES, plus a simulated hit
## (within the 18 px bot radius). Arrival geometry is the harness
## `bmPoint` formula (fireDist = |Pattern - self|; arrivalTick =
## T + ceil(fireDist/v) - 1), the same relation measure_aim_vs_power.nim
## validated against the harness resolver.
## PROTOCOL = PREQUENTIAL (predict-then-learn, streaming). Two regimes:
## retained across all rounds of one battle, and reset every round.
##
## Run:
## nim c -r -d:release --nimcache:/tmp/nc_j92 --path:common_libs \
## common_libs/tests/measure_bitbrain_gate.nim
import std/[os, strformat, strutils, math, tables, algorithm, random, json, times]
import gun_harness/gun_interface
import gun_harness/virtual_bullets
import gun_harness/offline_range
import guns/tm_horizon
import bitbrain/bitbrain
const
repoRoot = currentSourcePath().parentDir.parentDir.parentDir
fixturesDir = repoRoot / "tools" / "fixtures"
metaDir = fixturesDir / "drussgt_meta"
NIN = TMH_N_BITS ## 53
BotR = BotRadius ## 18 px
TMPM = PI / 180.0
# ── configuration (env-overridable) ───────────────────────────────────────────
proc envF(name: string, d: float): float =
let v = getEnv(name, "")
if v.len == 0: return d
try: parseFloat(v.strip()) except ValueError: d
proc envI(name: string, d: int): int =
let v = getEnv(name, "")
if v.len == 0: return d
try: parseInt(v.strip()) except ValueError: d
proc envList(name, d: string): seq[int] =
let v = getEnv(name, d)
for p in v.split(','):
let t = p.strip()
if t.len > 0: result.add parseInt(t)
let
EVAL_POWER = envF("BB_POWER", 2.0)
PMAX_DEG = envF("BB_PMAX", 40.0)
AD_PASSES = envI("BB_PASSES", 2)
AD_STEP = envI("BB_STEP", 1)
AD_CHUNK = envI("BB_CHUNK", 1000)
AD_STRIDE = envI("BB_STRIDE", 3)
AD_INIT = envI("BB_AD_INIT", 1)
AD_TARGET = envF("BB_TARGET", 0.01)
NADE_LIST = envList("BB_NADE", "128,256")
NLIST = envList("BB_NLIST", "4,8,16,32,64")
SEED = 20240921'i64
# ── small math ────────────────────────────────────────────────────────────────
proc wrapRad(a: float): float {.inline.} =
result = a
while result > PI: result -= 2.0 * PI
while result < -PI: result += 2.0 * PI
proc medianOf(xs: var seq[float]): float =
if xs.len == 0: return NaN
xs.sort()
let n = xs.len
if n mod 2 == 1: xs[n div 2]
else: 0.5 * (xs[n div 2 - 1] + xs[n div 2])
proc pctOf(xs: var seq[float], q: float): float =
if xs.len == 0: return NaN
xs.sort()
xs[min(xs.len - 1, int(ceil(q * xs.len.float)) - 1)]
# ── fixture round metadata ────────────────────────────────────────────────────
type
RoundInfo = object
endTick: int
roundId: int
proc loadRoundInfo(path: string, states: seq[WorldState]): Table[int, RoundInfo] =
let rp = metaDir / (extractFilename(path) & ".rounds.json")
if not fileExists(rp):
for s in states: result[s.tick] = RoundInfo(endTick: states[^1].tick, roundId: 0)
return
var idxAt = initTable[int, int]()
for i, s in states: idxAt[s.tick] = i
let j = parseFile(rp)
var rid = 0
for r in j["rounds"]:
let s0 = r["startTick"].getInt()
let c = r["count"].getInt()
if s0 notin idxAt: continue
let i0 = idxAt[s0]
let lastIdx = min(states.len - 1, i0 + c - 1)
for k in i0 .. lastIdx:
result[states[k].tick] = RoundInfo(endTick: states[lastIdx].tick, roundId: rid)
inc rid
# ── harvested sample ──────────────────────────────────────────────────────────
type
GSample = object
bits: array[NIN, uint8]
selfX, selfY: float
baseX, baseY: float
baseBearing: float
fireTick: int
horizon: int
speed: float
fxId: int
roundId: int
labelErr: float ## radians, fact at t+h
arrTick: int
actualArrX, actualArrY: float
actualArrBearing: float
slX, slY: float ## straight-line prediction at arrTick
slBearing: float
fireDist: float
# ── harvest: drive the LIVE TmHorizonGun, capture its own bits ───────────────
proc harvest(fx: Fixture, fxId: int, power: float, rinfo: Table[int, RoundInfo],
samples: var seq[GSample]): int =
var g = initTmHorizonGun()
g.setShift(0.0) # pure predict arm; we read Pattern's base
let speed = bulletSpeed(power)
if speed <= 0.0: return 0
var idxAt = initTable[int, int]()
for i, s in fx.states: idxAt[s.tick] = i
var curRound = -1
for state in fx.states:
let T = state.tick
if T notin rinfo: continue
let ri = rinfo[T]
if ri.roundId != curRound:
g.resetRoundState() # mirror the live onRoundStarted wipe
curRound = ri.roundId
let pred = g.predict(state, speed)
# Reuse the gun's OWN exported builder: tmhBaseBits (49 draft bits) + the
# 4-bit horizon one-hot via tmhLits, exactly the cachedBits per-tick path.
let dist = hypot(state.enemyX - state.selfX, state.enemyY - state.selfY)
let h = tmhHorizonFor(dist, speed)
let bucket = tmhHorizonBucket(h)
let baseBits = g.tmhBaseBits(state)
let lits = tmhLits(baseBits, bucket)
if T + h > ri.endTick: continue
let fireDist = hypot(pred.x - state.selfX, pred.y - state.selfY)
if fireDist < 1e-6: continue
let arrTick = T + max(1, int(ceil(fireDist / speed))) - 1
if arrTick > ri.endTick: continue
if arrTick notin idxAt or (T + h) notin idxAt: continue
let actH = fx.states[idxAt[T + h]]
let sx = state.selfX
let sy = state.selfY
let baseBearing = arctan2(pred.y - sy, pred.x - sx)
let labelErr = wrapRad(arctan2(actH.enemyY - sy, actH.enemyX - sx) - baseBearing)
let actA = fx.states[idxAt[arrTick]]
let arrBearing = arctan2(actA.enemyY - sy, actA.enemyX - sx)
let hr = degToRad(state.enemyHeading)
let vx = cos(hr) * state.enemySpeed
let vy = sin(hr) * state.enemySpeed
let el = float(arrTick - T)
let slX = state.enemyX + vx * el
let slY = state.enemyY + vy * el
let slBearing = arctan2(slY - sy, slX - sx)
var bits: array[NIN, uint8]
for k in 0 ..< NIN: bits[k] = lits[k]
samples.add GSample(
bits: bits, selfX: sx, selfY: sy, baseX: pred.x, baseY: pred.y,
baseBearing: baseBearing, fireTick: T, horizon: h, speed: speed,
fxId: fxId, roundId: ri.roundId, labelErr: labelErr,
arrTick: arrTick, actualArrX: actA.enemyX, actualArrY: actA.enemyY,
actualArrBearing: arrBearing, slX: slX, slY: slY, slBearing: slBearing,
fireDist: fireDist)
inc result
# ── geometry: rotate the base aim point around the shooter ────────────────────
proc rotAim(s: GSample, corrRad: float): tuple[x, y: float] =
let dx = s.baseX - s.selfX
let dy = s.baseY - s.selfY
let b = arctan2(dy, dx) + corrRad
(s.selfX + cos(b) * s.fireDist, s.selfY + sin(b) * s.fireDist)
proc missOf(s: GSample, x, y: float): tuple[px, angRad, rng: float] =
let rng = hypot(s.actualArrX - s.selfX, s.actualArrY - s.selfY)
(hypot(x - s.actualArrX, y - s.actualArrY),
abs(wrapRad(arctan2(y - s.selfY, x - s.selfX) - s.actualArrBearing)),
rng)
# ── metrics ───────────────────────────────────────────────────────────────────
type
Metrics = object
n: int
sumPx, sumDeg: float
hits: int
angHits: int
px: seq[float]
deg: seq[float]
proc add(m: var Metrics, missPx, angRad, rng: float) =
inc m.n
m.sumPx += missPx
m.sumDeg += angRad * 180.0 / PI
m.px.add missPx
m.deg.add angRad * 180.0 / PI
if missPx < BotR: inc m.hits
if angRad < arctan(BotR / max(rng, 1e-6)): inc m.angHits
proc merge(dst: var Metrics, src: Metrics) =
dst.n += src.n
dst.sumPx += src.sumPx
dst.sumDeg += src.sumDeg
dst.hits += src.hits
dst.angHits += src.angHits
dst.px.add src.px
dst.deg.add src.deg
proc line(name: string, m0: Metrics): string =
var m = m0
if m.n == 0: return fmt"{name:<24} n=0"
let meanPx = m.sumPx / float(m.n)
let medPx = medianOf(m.px)
let p90Px = pctOf(m.px, 0.9)
let meanDeg = m.sumDeg / float(m.n)
let medDeg = medianOf(m.deg)
let hit = 100.0 * float(m.hits) / float(m.n)
let ahit = 100.0 * float(m.angHits) / float(m.n)
fmt"{name:<24} n={m.n:<6} meanPx={meanPx:>6.2f} medPx={medPx:>6.2f} p90Px={p90Px:>6.2f} " &
fmt"meanDeg={meanDeg:>6.2f} medDeg={medDeg:>6.2f} pxHit%={hit:>5.1f} angHit%={ahit:>5.1f}"
# ── BitBrain readout ──────────────────────────────────────────────────────────
proc binOf(err: float, nClasses: int, maxDeg: float): int =
let x = (err * 180.0 / PI)
var k = int((x + maxDeg) / (2.0 * maxDeg) * float(nClasses))
if k < 0: k = 0
if k >= nClasses: k = nClasses - 1
k
proc centersOf(nClasses: int, maxDeg: float): seq[float] =
let w = 2.0 * maxDeg / float(nClasses)
for k in 0 ..< nClasses:
result.add (-maxDeg + (float(k) + 0.5) * w) * TMPM
proc weightedMean(counts: openArray[int], centers: openArray[float]): float =
var num = 0.0
var den = 0.0
for k in 0 ..< counts.len:
num += float(counts[k]) * centers[k]
den += float(counts[k])
if den <= 0.0: return 0.0
num / den
# ── regime ────────────────────────────────────────────────────────────────────
type Regime = enum rgRetained, rgPerRound
proc regimeName(r: Regime): string =
case r
of rgRetained: "retained"
of rgPerRound: "perRound"
type
ConfigResult = object
nade, nClasses: int
regime: Regime
wm: Metrics
am: Metrics
# ── one prequential run of a config ───────────────────────────────────────────
proc runConfig(ads: seq[AddressDecoder], nClasses: int, samples: seq[GSample],
labels: seq[int], regime: Regime,
maxDeg: float): ConfigResult =
result.nade = if ads.len > 0: ads[0].nAde else: 0
result.nClasses = nClasses
result.regime = regime
var bb = initBitBrain(ads, crossPairs(ads.len), nClasses)
let centers = centersOf(nClasses, maxDeg)
var lists: seq[seq[int32]]
var counts = newSeq[int](nClasses)
var lastFx = -1
var lastRound = -1
for i, s in samples:
if s.fxId != lastFx:
bb.resetLearning()
lastFx = s.fxId
lastRound = -1
if regime == rgPerRound and s.roundId != lastRound:
bb.resetLearning()
lastRound = s.roundId
# ---- predict (before learning this sample) ----
bb.fireInto(s.bits, lists)
for k in 0 ..< counts.len: counts[k] = 0
var total = 0
for sl in 0 ..< bb.sbcs.len:
let spec = bb.specs[sl]
bb.sbcs[sl].infer(lists[spec.row], lists[spec.col], counts)
for k in 0 ..< counts.len: total += counts[k]
var best = 0
for k in 1 ..< nClasses:
if counts[k] > counts[best]: best = k
let wm = weightedMean(counts, centers)
var (wx, wy) = rotAim(s, wm)
var (mpx, mang, mrg) = missOf(s, wx, wy)
result.wm.add(mpx, mang, mrg)
let amc = if total > 0: centers[best] else: 0.0
(wx, wy) = rotAim(s, amc)
(mpx, mang, mrg) = missOf(s, wx, wy)
result.am.add(mpx, mang, mrg)
# ---- learn ----
for sl in 0 ..< bb.sbcs.len:
let spec = bb.specs[sl]
discard bb.sbcs[sl].learn(lists[spec.row], lists[spec.col], labels[i])
# ── AD synthesis + homeostasis ────────────────────────────────────────────────
proc countGE(sc: openArray[int], t: int): int =
## number of elements >= t in an ascending-sorted array
var a = 0
var b = sc.len
while a < b:
let m = (a + b) div 2
if sc[m] >= t: b = m else: a = m + 1
sc.len - a
proc initADs(nAde: int, widths: seq[int], train: seq[GSample],
seed: int64, pctInit: bool): seq[AddressDecoder] =
var rng = initRand(seed)
for w in widths:
# center = 0 because our inputs are BINARY (0/1). The reference 127 is the
# midpoint of 0..255; centring binary inputs at 127 would make every synapse
# contribute ~-127 and collapse the ADE code to a polarity count.
result.add initRandomAddressDecoder(nAde, w, NIN, rng,
scale = DefaultScale, center = 0,
threshold = 0'i32)
let stride = max(1, AD_STRIDE)
# Unsupervised percentile init: put each ADE's threshold at the 99th
# percentile of its own score over this data, i.e. the 1% operating point the
# paper's controller aims for. The deterministic controller below then refines
# it (and would find it on its own, but would need thousands of `step=1`
# intervals, far more than a battle of 500-2000 ticks provides).
if pctInit:
for ai in 0 ..< result.len:
var sc = newSeq[int]()
for e in 0 ..< result[ai].nAde:
sc.setLen(0)
var s = 0
while s < train.len:
sc.add result[ai].score(train[s].bits, e)
s += stride
sc.sort()
if sc.len == 0: continue
# Choose the integer threshold whose firing count is CLOSEST to 1% of
# the observed scores. A plain 99th percentile lands on a large atom
# (binary inputs make the score distribution discrete and sparse), so
# its realised rate can be several percent; picking the closest count
# pins the realised rate near 1%.
let target = AD_TARGET * float(sc.len)
var lo = sc[0] - 1 # count>=lo == n (> target)
var hi = sc[^1] + 1 # count>=hi == 0 (<= target)
while hi - lo > 1:
let mid = (lo + hi) div 2
if float(countGE(sc, mid)) > target: lo = mid else: hi = mid
let tA = lo
let tB = hi
let cA = countGE(sc, tA)
let cB = countGE(sc, tB)
result[ai].thresholds[e] =
if abs(float(cA) - target) <= abs(float(cB) - target): int32(tA)
else: int32(tB)
proc homeostasis(ads: var seq[AddressDecoder], train: seq[GSample]) =
let stride = max(1, AD_STRIDE)
for _ in 0 ..< AD_PASSES:
for ad in ads.mitems: ad.resetFiringCounts()
var i = 0
while i < train.len:
let j = min(i + AD_CHUNK * stride, train.len)
var k = i
var cnt = 0
while k < j:
for ad in ads.mitems: ad.accumulateFiring(train[k].bits)
inc cnt
k += stride
for ad in ads.mitems:
ad.adaptThresholds(interval = max(1, cnt), targetRate = AD_TARGET, step = AD_STEP)
i = j
proc fireRates(ads: seq[AddressDecoder], train: seq[GSample],
sampleN: int): seq[float] =
## Fire rate per AD measured on the SAME stride the thresholds were fitted on
## (the whole stream, every AD_STRIDE-th sample), so the reported number is the
## rate the controller actually achieved.
discard sampleN
if train.len == 0: return
let stride = max(1, AD_STRIDE)
for ad in ads:
var f = 0
var i = 0
var m = 0
while i < train.len:
f += ad.fireCount(train[i].bits)
inc m
i += stride
result.add float(f) / (float(m) * float(ad.nAde))
const Widths = [6, 8, 10, 12]
# ── baselines ────────────────────────────────────────────────────────────────
proc computeBaselines(samples: seq[GSample]):
tuple[pattern, straight, fixed: Metrics, meanErr: float] =
var sumErr = 0.0
for s in samples: sumErr += s.labelErr
let meanErr = if samples.len > 0: sumErr / float(samples.len) else: 0.0
for s in samples:
var (px, ang, rng) = missOf(s, s.baseX, s.baseY)
result.pattern.add(px, ang, rng)
(px, ang, rng) = missOf(s, s.slX, s.slY)
result.straight.add(px, ang, rng)
let (fx, fy) = rotAim(s, meanErr)
(px, ang, rng) = missOf(s, fx, fy)
result.fixed.add(px, ang, rng)
result.meanErr = meanErr
# ── main ─────────────────────────────────────────────────────────────────────
proc main() =
var files: seq[string]
let only = getEnv("BB_FILES", "")
if only.len > 0:
for p in only.split(','):
let t = p.strip()
if t.len > 0: files.add fixturesDir / t
else:
for f in walkFiles(fixturesDir / "tr_drussgt_vs_*.jsonl"): files.add f
files.sort()
if files.len == 0:
stderr.writeLine("no tr_drussgt fixtures found"); quit(1)
let t0 = epochTime()
var samples: seq[GSample]
var perFx: seq[(string, int)]
for fi, path in files:
let fx = loadFixture(path)
let rinfo = loadRoundInfo(path, fx.states)
let got = harvest(fx, fi, EVAL_POWER, rinfo, samples)
perFx.add (extractFilename(path), got)
stderr.writeLine(fmt"[harvest] {extractFilename(path)} ticks={fx.states.len} samples={got}")
stderr.writeLine(fmt"[harvest] total samples={samples.len} in {epochTime()-t0:.1f}s")
echo ""
echo "================================================================="
echo "BitBrain gate test - fine-grained aim correction"
echo "================================================================="
echo fmt"fixtures : {files.len} tr_drussgt_vs_* (open-loop replay)"
echo fmt"power={EVAL_POWER} speed={bulletSpeed(EVAL_POWER):.1f} class half-range=+-{PMAX_DEG:.1f} deg"
echo fmt"AD: widths {Widths} nAde={NADE_LIST} target={AD_TARGET*100.0:.1f}% passes={AD_PASSES} step={AD_STEP} stride={AD_STRIDE}"
echo ""
echo "== dataset (effective sample counts) =="
for (nm, n) in perFx: echo fmt" {nm:<40} samples={n}"
var abst: seq[float]
var roundCount = initTable[int, int]()
for s in samples:
abst.add abs(s.labelErr) * 180.0 / PI
roundCount[s.fxId] = max(roundCount.getOrDefault(s.fxId, 0), s.roundId + 1)
var a2 = abst
var aSum = 0.0
for v in a2: aSum += v
echo fmt" |label err| deg: mean={aSum/float(a2.len):.2f} " &
fmt"p50={a2.pctOf(0.5):.2f} p90={a2.pctOf(0.9):.2f} p99={a2.pctOf(0.99):.2f} " &
fmt"max={a2.pctOf(1.0):.2f}"
var totRounds = 0
for _, r in roundCount: totRounds += r
echo fmt" rounds={totRounds} samples/round={float(samples.len)/float(max(1,totRounds)):.0f}"
echo ""
# ── baselines ────────────────────────────────────────────────────────────
let bl = computeBaselines(samples)
echo "== baselines (same arrival geometry, same sample set) =="
var pm = bl.pattern
var sm = bl.straight
var fm = bl.fixed
echo line("Pattern (zero corr)", pm)
echo line("straight-line naive", sm)
echo line(fmt"fixed corr={bl.meanErr*180.0/PI:+.2f} deg", fm)
echo ""
# ── label bin histogram for the largest N (effective sizes per class) ─────
let nBig = NLIST[^1]
var hist = newSeq[int](nBig)
for s in samples: inc hist[binOf(s.labelErr, nBig, PMAX_DEG)]
var maxH = 0
for h in hist: maxH = max(maxH, h)
echo fmt"== label histogram (N={nBig}) =="
for k in 0 ..< nBig:
echo fmt" class {k:>2} [{(-PMAX_DEG + float(k)*2*PMAX_DEG/float(nBig)):>+6.1f},{(-PMAX_DEG + float(k+1)*2*PMAX_DEG/float(nBig)):>+6.1f}) n={hist[k]:<6}"
echo ""
# ── AD fire rates ────────────────────────────────────────────────────────
echo "== AD layer (synthesised on this data, homeostasis to ~1% target) =="
var adsByNade = initTable[int, seq[AddressDecoder]]()
for nade in NADE_LIST:
var ads = initADs(nade, @Widths, samples, SEED + int64(nade), AD_INIT != 0)
let fr0 = fireRates(ads, samples, 3000)
var f0str = ""
var f0sum = 0.0
for i, r in fr0:
f0str.add fmt"w{Widths[i]}={r*100.0:.2f}% "
f0sum += r
echo fmt" nAde={nade:<4} pct-init: {f0str} (mean={f0sum/float(fr0.len)*100.0:.2f}%)"
homeostasis(ads, samples)
adsByNade[nade] = ads
let fr = fireRates(ads, samples, 3000)
var fstr = ""
var fSum = 0.0
for i, r in fr:
fstr.add fmt"w{Widths[i]}={r*100.0:.2f}% "
fSum += r
echo fmt" nAde={nade:<4} +homeostasis:{fstr} (mean={fSum/float(fr.len)*100.0:.2f}%)"
echo ""
# ── prequential sweep ────────────────────────────────────────────────────
echo "== prequential sweep (predict-then-learn) =="
echo " BitBrain wm = count-weighted mean of class centers (KEY readout)"
echo " BitBrain arg = argmax class center"
echo " vs Pattern / straight-line / fixed baselines above"
echo ""
# precompute labels per N
var results: seq[ConfigResult]
for nade in NADE_LIST:
let ads = adsByNade[nade]
for nc in NLIST:
var labels = newSeq[int](samples.len)
for i, s in samples: labels[i] = binOf(s.labelErr, nc, PMAX_DEG)
for regime in [rgRetained, rgPerRound]:
let cr = runConfig(ads, nc, samples, labels, regime, PMAX_DEG)
echo fmt" nAde={nade:<4} N={nc:<3} {regimeName(regime):<9} | " & line("wm", cr.wm)
echo " argmax | " & line("argmax", cr.am)
var r = cr
r.nade = nade
results.add r
echo ""
# ── per-fixture breakdown for the winning config ────────────────────────
echo "== per-fixture breakdown: perRound, N=32, vs Pattern =========="
let winNade = NADE_LIST[^1]
let winAds = adsByNade[winNade]
let winN = 32
var winLabels = newSeq[int](samples.len)
for i, s in samples: winLabels[i] = binOf(s.labelErr, winN, PMAX_DEG)
for fi in 0 ..< files.len:
var sub: seq[GSample]
var subLab: seq[int]
for i, s in samples:
if s.fxId == fi:
sub.add s
subLab.add winLabels[i]
if sub.len == 0: continue
var pat: Metrics
for s in sub:
let (px, ang, rng) = missOf(s, s.baseX, s.baseY)
pat.add(px, ang, rng)
let cr = runConfig(winAds, winN, sub, subLab, rgPerRound, PMAX_DEG)
echo " " & align(extractFilename(files[fi]), 38) & " n=" & align($sub.len, 6)
echo " " & line("Pattern", pat)
echo " " & line("BitBrain wm", cr.wm)
echo " " & line("BitBrain arg", cr.am)
echo ""
# ── compact comparison for the verdict ───────────────────────────────────
echo "== compact: mean px error (lower is better) =="
echo " " & align("config",24) & " " & align("wm meanPx",9) & " " &
align("arg meanPx",10) & " " & align("wm hit%",8) & " " & align("arg hit%",8)
echo " " & align("Pattern",24) & " " & align(fmt"{pm.sumPx/float(pm.n):.2f}",9) & " " &
align("-",10) & " " & align(fmt"{100.0*float(pm.hits)/float(pm.n):.1f}",8) & " " & align("-",8)
echo " " & align("straight-line",24) & " " & align(fmt"{sm.sumPx/float(sm.n):.2f}",9) & " " &
align("-",10) & " " & align(fmt"{100.0*float(sm.hits)/float(sm.n):.1f}",8) & " " & align("-",8)
echo " " & align("fixed corr",24) & " " & align(fmt"{fm.sumPx/float(fm.n):.2f}",9) & " " &
align("-",10) & " " & align(fmt"{100.0*float(fm.hits)/float(fm.n):.1f}",8) & " " & align("-",8)
for r in results:
let nm = fmt"nAde{r.nade}/N{r.nClasses}/{regimeName(r.regime)}"
echo " " & align(nm,24) & " " & align(fmt"{r.wm.sumPx/float(r.wm.n):.2f}",9) & " " &
align(fmt"{r.am.sumPx/float(r.am.n):.2f}",10) & " " &
align(fmt"{100.0*float(r.wm.hits)/float(r.wm.n):.1f}",8) & " " &
align(fmt"{100.0*float(r.am.hits)/float(r.am.n):.1f}",8)
echo ""
# ── shuffled-label control (must collapse to the fixed baseline) ─────────
echo "== shuffled-label control (permuted labels, same input distribution) =="
echo " (the weighted mean shrinks to the label mean; the argmax keeps making a"
echo " confident pick, so its null is random-correction worse-than-baseline)"
let nadeBig = NADE_LIST[^1]
let adsBig = adsByNade[nadeBig]
var baseLabels = newSeq[int](samples.len)
for i, s in samples: baseLabels[i] = binOf(s.labelErr, nBig, PMAX_DEG)
for regime in [rgRetained, rgPerRound]:
var wPx, aPx: float
var wHit, aHit, shN: int
for rep in 0 ..< 3:
var labs = baseLabels
var rng = initRand(SEED + 777 + int64(rep))
for i in countdown(labs.len - 1, 1):
let j = rng.rand(i)
swap(labs[i], labs[j])
let cr = runConfig(adsBig, nBig, samples, labs, regime, PMAX_DEG)
wPx += cr.wm.sumPx; wHit += cr.wm.hits
aPx += cr.am.sumPx; aHit += cr.am.hits
shN += cr.wm.n
echo fmt" {regimeName(regime):<9} shuffled mean: wm meanPx={wPx/float(shN):.2f} pxHit%={100.0*float(wHit)/float(shN):.1f} " &
fmt"arg meanPx={aPx/float(shN):.2f} pxHit%={100.0*float(aHit)/float(shN):.1f}"
echo ""
stderr.writeLine(fmt"[total] {epochTime()-t0:.1f}s")
when isMainModule:
main()