## Discrete-target Tsetlin gun sweep: tm_pattern vs its shuffled-feedback ## control, the current default Tsetlin gun, and Linear. ## ## Metric: EARLY virtual hit rate = resolutions in the first 100 ticks of each ## ROUND (the TM starts cold every round, so this is what "learns fast" means), ## plus first-300 and whole-round rates. Every fixture with a round sidecar is ## split into rounds and each round is replayed with a FRESH gun instance ## (cold every battle, overfit within the battle). ## ## Usage: ## nim c -r --path:common_libs -d:release common_libs/tests/sweep_tm_pattern.nim \ ## --set=real --seeds=3 ## Flags: ## --set=real|range|synthetic fixture set (default real) ## --seeds=N seeds for stochastic variants (default 3) ## --maxrounds=N cap rounds per fixture ## --variants=a,b,c subset of: linear,tsetlin,tmpat,tmpat_shuf ## ## Emits per-run CSV then a COMPARE section with per-run means, ranges and a ## paired sign test (exact binomial, two-sided) for every variant pair. import std/[os, strformat, strutils, json, tables, math, random, algorithm, sequtils] import gun_harness/offline_range import range_guns import guns/tm_pattern import guns/linear const repoRoot = currentSourcePath().parentDir.parentDir.parentDir const fixturesDir = repoRoot / "tools" / "fixtures" type GunStats = object obs, labelMiss, traceMiss: int labHist, choHist: array[TM_CLASSES, int] classCorrect, classTotal: int Adapt = object h100, n100, h300, n300, hall, nall, f100, m100: int rounds: int st: GunStats RoundSpan = tuple[start, count: int] Row = object variant, fixture: string seed: int r: Adapt VariantKind = enum vLinear, vTsetlin, vTmpat, vTmpatShuf proc variantName(v: VariantKind): string = case v of vLinear: "Linear" of vTsetlin: "Tsetlin" of vTmpat: "TMPattern" of vTmpatShuf: "TMPatternShuf" proc loadRounds(path: string): seq[RoundSpan] = let dir = path.parentDir let base = path.extractFilename var side = dir / "drussgt_meta" / (base & ".rounds.json") if not fileExists(side): side = dir / (base & ".rounds.json") if not fileExists(side): return @[] let node = parseJson(readFile(side)) if not node.hasKey("rounds"): return @[] for r in node["rounds"]: result.add (r["startTick"].getInt(), r["count"].getInt()) proc addAdapt(dst: var Adapt, src: Adapt) = inc dst.rounds, src.rounds dst.h100 += src.h100; dst.n100 += src.n100 dst.h300 += src.h300; dst.n300 += src.n300 dst.hall += src.hall; dst.nall += src.nall dst.f100 += src.f100; dst.m100 += src.m100 dst.st.obs += src.st.obs; dst.st.labelMiss += src.st.labelMiss dst.st.traceMiss += src.st.traceMiss dst.st.classCorrect += src.st.classCorrect dst.st.classTotal += src.st.classTotal for c in 0..= 0: lst = lastSeen[si] if state.enemies.len > 0: for e in state.enemies: enemyPositions[e.id] = (x: e.x, y: e.y, lastSeenTick: lst, alive: true) else: enemyPositions[enemyId] = (x: state.enemyX, y: state.enemyY, lastSeenTick: lst, alive: true) let localTick = state.tick - baseTick tracker.tickBullets(state, enemyPositions, proc(gunId: GunId, binIdx: int, e: FeedbackEvent) = inc res.nall if e.hit: inc res.hall if localTick < 100: inc res.n100 if e.hit: inc res.h100 if localTick < 300: inc res.n300 if e.hit: inc res.h300 let fireTick = e.fireTick - baseTick if fireTick < 100: inc res.m100 if e.hit: inc res.f100 driver.resultCb(e)) let after = obsCount() res.st.obs = after.obs - obsBefore.obs res.st.labelMiss = after.labelMiss - obsBefore.labelMiss res.st.traceMiss = after.traceMiss - obsBefore.traceMiss res.st.classCorrect = after.classCorrect - obsBefore.classCorrect res.st.classTotal = after.classTotal - obsBefore.classTotal for c in 0.. 0 and spans.len > maxRounds: spans.setLen(maxRounds) if spans.len == 0: return replayRound(fx.states, fx.lastSeen, fx.enemyId, 0, driver, metric, obsBefore, obsCount) for sp in spans: var st: seq[WorldState] var ls: seq[int] for i in 0..= sp.start and t < sp.start + sp.count: st.add fx.states[i] ls.add(if i < fx.lastSeen.len: fx.lastSeen[i] else: -1) if st.len == 0: continue addAdapt(result, replayRound(st, ls, fx.enemyId, sp.start, driver, metric, obsBefore, obsCount)) proc emptyStats(): GunStats = GunStats() proc makeTmpatDriver(seed: int, shuffle: bool): tuple[driver: GunDriver, gun: ref TmPatternGun] = let g = new(TmPatternGun) g[] = initTmPatternGun() g[].shuffleLabels = shuffle if seed >= 0: randomize(seed) result.gun = g result.driver = GunDriver( name: (if shuffle: "TMPatternShuf" else: "TMPattern"), predictCb: proc(state: WorldState, bulletSpeed: float): GunPrediction = g[].predict(state, bulletSpeed), resultCb: proc(e: FeedbackEvent) = g[].onResult(e), readyCb: proc(): bool = g[].isWarmedUp()) proc tmpatStats(g: ref TmPatternGun): GunStats = result.obs = g[].totalObs result.labelMiss = g[].labelMisses result.traceMiss = g[].traceMisses result.labHist = g[].labelHist result.choHist = g[].chosenHist result.classCorrect = g[].classCorrect result.classTotal = g[].classTotal proc fixtureSet(name: string): seq[string] = case name of "synthetic", "": for n in SyntheticFixtureNames: result.add n of "real": for n in ["drussgt_vs_crazy", "drussgt_vs_spinbot", "drussgt_vs_drussgt", "tr_drussgt_vs_crazy", "tr_drussgt_vs_spinbot", "tr_drussgt_vs_modularbot"]: result.add(fixturesDir / (n & ".jsonl")) of "range": for n in SyntheticFixtureNames: result.add n for n in ["drussgt_vs_crazy", "drussgt_vs_spinbot", "tr_drussgt_vs_crazy"]: result.add(fixturesDir / (n & ".jsonl")) else: discard proc resolve(name: string): tuple[fx: Fixture, path: string] = let p = if fileExists(name): name else: fixturesDir / (name & ".jsonl") (loadFixture(p), p) # ── stats ──────────────────────────────────────────────────────────────────── proc rateStr(h, n: int): string = if n == 0: " n/a " else: &"{h.float / n.float * 100.0:5.1f}%" proc binomPmf(k, n: int): float = if k < 0 or k > n: return 0.0 var lg = 0.0 for i in 1..k: lg += ln(float(n - k + i)) - ln(float(i)) exp(lg - float(n) * ln(2.0)) proc signTestP(wins, n: int): float = ## exact two-sided binomial p under p=0.5 if n == 0: return 1.0 let lo = min(wins, n - wins) var s = 0.0 for k in 0..lo: s += binomPmf(k, n) min(1.0, 2.0 * s) proc main() = var nSeeds = 3 var set = "real" var maxRounds = 0 var metricName = "path" var variants: seq[VariantKind] = @[vLinear, vTsetlin, vTmpat, vTmpatShuf] for i in 1..paramCount(): let a = paramStr(i) if a.startsWith("--seeds="): nSeeds = parseInt(a[8..^1]) elif a.startsWith("--set="): set = a[6..^1] elif a.startsWith("--metric="): metricName = a[9..^1] elif a.startsWith("--maxrounds="): maxRounds = parseInt(a[12..^1]) elif a.startsWith("--variants="): variants = @[] for tok in a[11..^1].split(','): case tok.strip() of "linear": variants.add vLinear of "tsetlin": variants.add vTsetlin of "tmpat": variants.add vTmpat of "tmpat_shuf": variants.add vTmpatShuf else: discard let metric = if metricName == "point": bmPoint else: bmPath let names = fixtureSet(set) echo "# set=", set, " seeds=", nSeeds, " metric=", metricName, " variants=", variants.mapIt(variantName(it)).join(",") echo "variant,fixture,seed,h100,n100,h300,n300,hall,nall,f100,m100,rounds,obs,labelMiss,traceMiss" var rows: seq[Row] for name in names: let (fx, path) = resolve(name) let fxName = path.extractFilename.replace(".jsonl", "") for v in variants: let nIter = if v == vLinear: 1 else: nSeeds for seed in 1..nIter: var drv: GunDriver var gun: ref TmPatternGun var obsCount: proc(): GunStats = emptyStats case v of vLinear: drv = makeDriver("Linear", LinearGun()) of vTsetlin: let pair = makeTsetlinDriver(seed = seed) drv = pair.driver of vTmpat, vTmpatShuf: let pair = makeTmpatDriver(seed = seed, shuffle = (v == vTmpatShuf)) drv = pair.driver gun = pair.gun obsCount = proc(): GunStats = tmpatStats(gun) let r = replayFixture(fx, path, drv, metric, emptyStats(), obsCount, maxRounds) rows.add Row(variant: variantName(v), fixture: fxName, seed: seed, r: r) echo &"{variantName(v)},{fxName},{seed},{r.h100},{r.n100},{r.h300},{r.n300}," & &"{r.hall},{r.nall},{r.f100},{r.m100},{r.rounds},{r.st.obs},{r.st.labelMiss},{r.st.traceMiss}" # ── per-variant pooled summary ── echo "\n# ── pooled summary ──" echo "variant,runs,h100,n100,early%,h300,n300,early300%,hall,nall,overall%,obs,labelMiss,traceMiss" var pooled = initTable[string, Adapt]() for v in variants: pooled[variantName(v)] = Adapt() for row in rows: addAdapt(pooled[row.variant], row.r) for v in variants: let a = pooled[variantName(v)] echo &"{variantName(v)},{a.rounds},{a.h100},{a.n100},{rateStr(a.h100, a.n100)}," & &"{a.h300},{a.n300},{rateStr(a.h300, a.n300)},{a.hall},{a.nall}," & &"{rateStr(a.hall, a.nall)},{a.st.obs},{a.st.labelMiss},{a.st.traceMiss}" # ── label vs chosen class histogram (TMPattern only) ── for v in [vTmpat, vTmpatShuf]: if v in variants: let a = pooled[variantName(v)] var ls, cs: string for c in 0.. 0: row.r.h100.float / row.r.n100.float else: 0.0 let overall = if row.r.nall > 0: row.r.hall.float / row.r.nall.float else: 0.0 byVariant[row.variant][(row.fixture, row.seed)] = early byVariantAll[row.variant][(row.fixture, row.seed)] = overall if vLinear in variants: for row in rows: if row.variant == "Linear": for s in 2..nSeeds: byVariant["Linear"][(row.fixture, s)] = byVariant["Linear"][(row.fixture, 1)] byVariantAll["Linear"][(row.fixture, s)] = byVariantAll["Linear"][(row.fixture, 1)] echo "\n# ── per-run early-rate distribution (mean / min / max, n runs) ──" echo "variant,earlyMean%,earlyMin%,earlyMax%,overallMean%,overallMin%,overallMax%,n" for v in variants: var es: seq[float] var os: seq[float] for r in byVariant[variantName(v)].values: es.add r for r in byVariantAll[variantName(v)].values: os.add r if es.len == 0: continue es.sort(); os.sort() echo &"{variantName(v)},{es.sum/float(es.len)*100:.2f},{es[0]*100:.2f},{es[^1]*100:.2f}," & &"{os.sum/float(os.len)*100:.2f},{os[0]*100:.2f},{os[^1]*100:.2f},{es.len}" echo "\n# ── pairwise paired sign tests (rows = fixture×seed) ──" echo "A,B,metric,nA>B,nB>A,ties,p" for i in 0.. vb: inc winsA elif vb > va: inc winsB else: inc ties if n == 0: continue echo &"{aName},{bName},{label},{n},{winsA},{winsB},{ties},{signTestP(winsA, n - ties):.4f}" when isMainModule: main()