Files
SirRoboGarage/common_libs/tests/sweep_tm_pattern.nim
T
SirStone ca82053a11 TM gun: the discrete-target diagnosis was RIGHT - it learns now. Still loses to Linear.
The user's goal: a TM gun that is the best 1v1 gun, starting from scratch every
battle but quickly overfitting the current enemy. The previous attempt (knob
tuning) failed: NO configuration beat its own shuffled-feedback control, and the
TM-off ablation scored the same as TM-on, i.e. the TM's correction was
near-zero-mean noise. Diagnosis then: a Tsetlin Machine is a CLASSIFIER, and we
were asking it for an absolute aim point - a regression target. So this attempt
gave it a DISCRETE target (multi-class over guess-factor buckets) with 40
binary/bucketed motion features, and measured it against Linear, the default
Tsetlin gun, and a MANDATORY shuffled control.

THE DIAGNOSIS IS CONFIRMED - THE TM LEARNS, DECISIVELY:
  online class accuracy     46.0%  vs shuffled control 20.0%   (2.3x chance)
  raw ungated argmax        21.2%/18.6% vs shuffled 15.2%/8.3% (18/18, p<0.0001)
  TMPattern > its shuffled control, overall   17/1 runs, p=0.0001
Compare the previous attempt, which could not beat shuffled feedback at all.
TMPattern also beats the default Tsetlin gun early (17/1, p=0.0001), so it is a
strictly better TM gun than the one in the rack.

BUT IT IS NOT COMPETITIVE WITH LINEAR ON REAL SURFERS:
  real DrussGT, bmPath (the shipped metric), 3 seeds, pooled early/overall
    Linear            34.0% (6358/18715)    24.3% (58297/239943)
    TMPattern (gated) 27.9% (15514/55535)   22.0% (158658/719681)
    TMPatternShuf     28.7%                 19.4%
  Linear > TMPattern: 15/18 early p=0.0075, 15/18 overall p=0.0075
  bmPoint: neutral (7.2%/4.6% vs Linear 7.2%/4.7%)
  synthetic controlled motion: matches/edges Linear (66.8%/60.6% vs 66.4%/59.6%,
    shuffled 55.7%/50.1%) - the mechanism works when motion is predictable.

So: the representation fix moved this from "learns nothing" to "learns strongly
but applies its knowledge badly". INFERRED reason for the residual loss: the
linear lead is already the modal GF bucket (the label histogram is centred), so
corrective excursions away from it are net-negative. The measured deficit lives
in the BASELINE and in RANGE, not in the TM knobs - which is why further knob
tuning was never going to work.

Best config: gated hard K=5, TM_CONF_MARGIN=0.25, TM_SHRINK=0.5.
NOT TRIED (time-boxed): the binary-reversal target, and a RADIAL (range-holding)
target - the latter is the top next step.

Adds `common_libs/guns/tm_pattern.nim` (NOT registered in the rack),
`common_libs/tests/sweep_tm_pattern.nim`, and a durable writeup at
`common_libs/tests/tm_pattern_sweep_results.md`.
2026-09-22 00:56:01 +02:00

350 lines
14 KiB
Nim
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
## 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..<TM_CLASSES:
dst.st.labHist[c] += src.st.labHist[c]
dst.st.choHist[c] += src.st.choHist[c]
proc replayRound(states: seq[WorldState], lastSeen: seq[int], enemyId, baseTick: int,
driver: GunDriver, metric: BulletMetric,
obsBefore: GunStats, obsCount: proc(): GunStats): Adapt =
var tracker = initTracker(1, metric)
var res: Adapt
inc res.rounds
for si in 0..<states.len:
let state = states[si]
var preds: array[len(PowerBins), GunPrediction]
for i in 0..<len(PowerBins):
preds[i] = driver.predictCb(state, bulletSpeed(PowerBins[i]))
let ready = if driver.readyCb == nil: true else: driver.readyCb()
if ready:
tracker.spawnBullets(0, preds, state, enemyId)
var enemyPositions: Table[int, tuple[x, y: float, lastSeenTick: int, alive: bool]]
var lst = state.tick
if si < lastSeen.len and lastSeen[si] >= 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..<TM_CLASSES:
res.st.labHist[c] = after.labHist[c] - obsBefore.labHist[c]
res.st.choHist[c] = after.choHist[c] - obsBefore.choHist[c]
result = res
proc replayFixture(fx: Fixture, path: string, driver: GunDriver, metric: BulletMetric,
obsBefore: GunStats, obsCount: proc(): GunStats,
maxRounds = 0): Adapt =
var spans =
if fx.meta.source == "synthetic": @[(start: 0, count: fx.states.len)]
else: loadRounds(path)
if maxRounds > 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..<fx.states.len:
let t = fx.states[i].tick
if t >= 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..<TM_CLASSES:
ls.add &"{a.st.labHist[c]},"
cs.add &"{a.st.choHist[c]},"
echo &"\n# class histogram {variantName(v)}: labels=[{ls}] chosen=[{cs}] " &
&"onlineAcc={a.st.classCorrect}/{a.st.classTotal}"
# ── per-run distributions (a "run" = one fixture × one seed) ──
# Linear is deterministic: replicate its one row per fixture across seeds so a
# paired comparison against a stochastic variant has a partner per run.
type Key = tuple[fixture: string, seed: int]
var byVariant = initTable[string, Table[Key, float]]() # early rate
var byVariantAll = initTable[string, Table[Key, float]]() # overall rate
for v in variants:
byVariant[variantName(v)] = initTable[Key, float]()
byVariantAll[variantName(v)] = initTable[Key, float]()
for row in rows:
let early = if row.r.n100 > 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..<variants.len:
for j in 0..<variants.len:
if i == j: continue
let aName = variantName(variants[i])
let bName = variantName(variants[j])
for (label, tab) in [("early", byVariant), ("overall", byVariantAll)]:
var winsA, winsB, ties, n = 0
for k, va in tab[aName].pairs:
if k notin tab[bName]: continue
let vb = tab[bName][k]
inc n
if va > 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()