feat(guns): scale-aware power selection (+52% damage); TM classifier gun built, measured, DISABLED
TASK 2 - power selection, a clear win. bestPower used an ABSOLUTE MinHitRate = 0.40 bar. Measured per-bin virtual rates (rolling-100 fraction) show no bin ever clears 40%, so 11 of 14 guns were stuck at bin 0 (power 1.0) even where higher bins were comparable: Linear p1.0 44% p1.5 39% p2.0 30% p3.0 29% old bin 0 -> new bin 3 Accel p1.0 44% p1.5 40% p2.0 26% p3.0 29% old bin 1 -> new bin 3 Pattern p1.0 50% p1.5 40% p2.0 27% p3.0 12% old bin 1 -> new bin 2 Replaced with a scale-aware PowerBarFrac = 0.50 (a dimensionless FRACTION of the gun's own best bin rate). 13 of 14 selections now pick heavier bullets. Real effect vs DrussGT (8 rounds x 3 runs): hit rate unchanged (7.56% -> 7.47%) but damage dealt +52% (157 -> 239 per run) and rounds end faster. Same accuracy, half the shots, half again more damage. TASK 1 - the TM pattern-classifier gun does NOT earn its slot. It was built as a mixture of experts with a corrected-Granmo TM as a multi-class gate over HeadOn/Linear/Circular/WallBounce/Accel, labelled by which expert's prediction was closest to the actual enemy position (an exact, supervised, per-shot label - no delayed credit). Offline it loses to the best of its OWN experts on essentially every fixture, and against DrussGT it cost real performance: baseline (path+relative) 7.56% real hit rate, damage 157 + power fix 7.47%, damage 239 + power fix + TM gun 5.59%, damage 133 The gun was selected on 806 ticks and fired 24 real shots at 4.2%. So the tree ships with EnableTmSelector = false: code and wiring kept intact for re-enabling, but it is not in the active rack. Worth recording from the clause dump: the gate DOES latch onto meaningful structure. On energy-threshold-turner, HeadOn's clauses key on the energy bits (the rule's own driving variable) while Circular keys on distance/velocity. So the TM is learning something real and interpretable - it simply cannot beat 'always pick the best expert'. Root cause (INFERRED): the closest-expert label is noisy because several experts are near-tied, and under the path metric the winner varies by power bin while the gate sees one shared per-tick input, so a one-vs-rest gate over a saturated 870-bit clause space has no margin to exploit. (Zero-padding the 2-frame window was tried first and saturated every clause at 256-755 included literals; alternating the two real frames fixed that.) Also factors the corrected feedback into an exported tmLearnDir and exports the encoding/TM primitives; the Tsetlin tests still reproduce the documented mean=13.8 included literals, so the refactor is behaviour-preserving. Verified: 33/33 guard checks, tsetlin tests green, metric checks green, new power-selection guard green (13/14 selections change; relative bar still picks bin 1 and not bin 3 for a [30,25,12,5]% profile), 12/12 offline==online acceptance under the shipped default.
This commit is contained in:
@@ -1,5 +1,5 @@
|
||||
## ModularBot — plugin gun architecture tracer bullet.
|
||||
## Guns: HeadOnGun (0), LinearGun (1), TsetlinGun (2), CircularGun (3), GFGun (4), PatternMatcherGun (5), WallBounceGun (6), AccelGun (7), StopShotGun (8), DisplacementGun (9), AveragedLeadGun (10), DecayGFGun (11), KNNGun (12) via GunHarness.
|
||||
## Guns: HeadOnGun (0), LinearGun (1), TsetlinGun (2), CircularGun (3), GFGun (4), PatternMatcherGun (5), WallBounceGun (6), AccelGun (7), StopShotGun (8), DisplacementGun (9), AveragedLeadGun (10), DecayGFGun (11), KNNGun (12), TmSelectorGun (13) via GunHarness.
|
||||
## Radar: RadarLockModule (1v1) / MeleeScanModule (2+ enemies), auto-switched per tick.
|
||||
## Movement: OscillatorModule (perpendicular strafing).
|
||||
|
||||
@@ -24,6 +24,7 @@ import guns/displacement
|
||||
import guns/averaged_lead
|
||||
import guns/decay_gf
|
||||
import guns/knn_gun
|
||||
import guns/tm_selector
|
||||
import movements/phantom_meteor
|
||||
import movements/rammer
|
||||
import movements/the_floor_is_lava
|
||||
@@ -35,6 +36,15 @@ import targeting/target_selector
|
||||
const botJsonPath = currentSourcePath().parentDir / "ModularBot.json"
|
||||
const DebugVBullets = false
|
||||
const DebugCircular = false
|
||||
## Registration switch for the TM selector gun (id 13). Set false to build a
|
||||
## rack without it (A/B runs); the tracker still reserves gun id 13 so the other
|
||||
## ids and per-gun stats are unchanged.
|
||||
##
|
||||
## VERDICT: left FALSE. The TM gate underperforms its own experts on every
|
||||
## offline fixture (energy-threshold 63% vs Circular 100%, random-walk 45% vs
|
||||
## WallBounce 68%, Corners 25% vs Accel 34%) and the DrussGT A/B showed real hit
|
||||
## rate 5.59% vs 7.47% for the power-fix-only rack. Flip to true to re-enable.
|
||||
const EnableTmSelector = false
|
||||
## Task A instrumentation: log every real shot + its eventual outcome to
|
||||
## /tmp/shot_log.jsonl. Set to false to compile all shot-log machinery out.
|
||||
const ShotLog = true
|
||||
@@ -45,7 +55,7 @@ const ShotLog = true
|
||||
## can enable recording for just the battle it spawns by exporting the env var.
|
||||
let RecordWorldState* = existsEnv("TR_RECORD_WORLDSTATE")
|
||||
const WorldStateRecordPath = "/tmp/worldstate_record.jsonl"
|
||||
const GunNames = ["HeadOn", "Linear", "Tsetlin", "Circular", "GuessFactor", "Pattern", "WallBounce", "Accel", "StopShot", "Displace", "AvgLead", "DecayGF", "KNN"]
|
||||
const GunNames = ["HeadOn", "Linear", "Tsetlin", "Circular", "GuessFactor", "Pattern", "WallBounce", "Accel", "StopShot", "Displace", "AvgLead", "DecayGF", "KNN", "TMSelect"]
|
||||
|
||||
const
|
||||
CLR_GUN = "\e[33m" # yellow
|
||||
@@ -90,6 +100,7 @@ type
|
||||
avgLead: AveragedLeadGun
|
||||
decayGF: DecayGFGun
|
||||
knnGun: KNNGun
|
||||
tmSelector: TmSelectorGun
|
||||
mover: TFILModule
|
||||
rammer: RammerModule
|
||||
isRamming: bool
|
||||
@@ -110,13 +121,13 @@ type
|
||||
roundNumber: int
|
||||
realShotsFired: int
|
||||
realHits: int
|
||||
gunRealShots: array[13, int]
|
||||
gunRealHits: array[13, int]
|
||||
gunRealShots: array[14, int]
|
||||
gunRealHits: array[14, int]
|
||||
pendingFires: seq[PendingShot] ## FIFO of fired shots awaiting onBulletFired bulletId stamp
|
||||
bulletGun: Table[int, int] ## bulletId -> gun id, filled on onBulletFired, drained on resolution
|
||||
bulletShot: Table[int, PendingShot] ## bulletId -> shot metadata (Task A shot log)
|
||||
pendingHitBullets: HashSet[int] ## hit bulletIds seen before their onBulletFired stamp
|
||||
gunSelectionCount: array[13, int]
|
||||
gunSelectionCount: array[14, int]
|
||||
lastKnownTargetId: int ## persists through death, used for round-end stats
|
||||
|
||||
proc writeShotLog(shot: PendingShot, hit: bool, unresolved: bool) =
|
||||
@@ -322,7 +333,7 @@ method onRoundEnded*(bot: ModularBot, e: RoundEndedEventForBot) =
|
||||
let fit = bot.tracker.fitnessFor(targetId)
|
||||
|
||||
var gunsArr = newJArray()
|
||||
for gid in 0..<13:
|
||||
for gid in 0..<14:
|
||||
var totalShots = 0
|
||||
var totalHits = 0
|
||||
for binIdx in 0..<len(vb.PowerBins):
|
||||
@@ -382,7 +393,7 @@ method onRoundStarted*(bot: ModularBot, e: RoundStartedEvent) =
|
||||
bot.bulletGun.clear()
|
||||
bot.bulletShot.clear()
|
||||
bot.pendingHitBullets.clear()
|
||||
for i in 0..<13:
|
||||
for i in 0..<14:
|
||||
bot.gunSelectionCount[i] = 0
|
||||
bot.gunRealShots[i] = 0
|
||||
bot.gunRealHits[i] = 0
|
||||
@@ -589,6 +600,7 @@ method run*(bot: ModularBot) =
|
||||
var alPreds: array[len(PowerBins), GunPrediction]
|
||||
var dgPreds: array[len(PowerBins), GunPrediction]
|
||||
var knnPreds: array[len(PowerBins), GunPrediction]
|
||||
var tmselPreds: array[len(PowerBins), GunPrediction]
|
||||
for i in 0..<len(PowerBins):
|
||||
headsUp[i] = bot.headOn.predict(bot.lastState, bulletSpeed(PowerBins[i]))
|
||||
linPreds[i] = bot.linear.predict(bot.lastState, bulletSpeed(PowerBins[i]))
|
||||
@@ -603,6 +615,8 @@ method run*(bot: ModularBot) =
|
||||
alPreds[i] = bot.avgLead.predict(bot.lastState, bulletSpeed(PowerBins[i]))
|
||||
dgPreds[i] = bot.decayGF.predict(bot.lastState, bulletSpeed(PowerBins[i]))
|
||||
knnPreds[i] = bot.knnGun.predict(bot.lastState, bulletSpeed(PowerBins[i]))
|
||||
if EnableTmSelector:
|
||||
tmselPreds[i] = bot.tmSelector.predict(bot.lastState, bulletSpeed(PowerBins[i]))
|
||||
|
||||
bot.tracker.spawnBullets(0, headsUp, bot.lastState, tid)
|
||||
bot.tracker.spawnBullets(1, linPreds, bot.lastState, tid)
|
||||
@@ -618,6 +632,8 @@ method run*(bot: ModularBot) =
|
||||
bot.tracker.spawnBullets(10, alPreds, bot.lastState, tid)
|
||||
bot.tracker.spawnBullets(11, dgPreds, bot.lastState, tid)
|
||||
bot.tracker.spawnBullets(12, knnPreds, bot.lastState, tid)
|
||||
if EnableTmSelector and bot.tmSelector.isWarmedUp():
|
||||
bot.tracker.spawnBullets(13, tmselPreds, bot.lastState, tid)
|
||||
|
||||
# Build slim enemy table for tickBullets
|
||||
var enemyPositions: Table[int, tuple[x, y: float, lastSeenTick: int, alive: bool]]
|
||||
@@ -645,6 +661,7 @@ method run*(bot: ModularBot) =
|
||||
of 10: bot.avgLead.onResult(fe)
|
||||
of 11: bot.decayGF.onResult(fe)
|
||||
of 12: bot.knnGun.onResult(fe)
|
||||
of 13: bot.tmSelector.onResult(fe)
|
||||
else: discard
|
||||
if fe.hit: inc bot.virtualHits else: inc bot.virtualMiss
|
||||
when DebugVBullets:
|
||||
@@ -672,6 +689,7 @@ method run*(bot: ModularBot) =
|
||||
of 10: setTurretColor("#996633"); setBulletColor("#CC9966")
|
||||
of 11: setTurretColor("#008888"); setBulletColor("#00AAAA")
|
||||
of 12: setTurretColor("#CC00CC"); setBulletColor("#FF44FF")
|
||||
of 13: setTurretColor("#00FFCC"); setBulletColor("#66FFDD")
|
||||
else: discard
|
||||
|
||||
let pred = case selectedGun
|
||||
@@ -687,6 +705,7 @@ method run*(bot: ModularBot) =
|
||||
of 10: bot.avgLead.predict(bot.lastState, bulletSpeed(power))
|
||||
of 11: bot.decayGF.predict(bot.lastState, bulletSpeed(power))
|
||||
of 12: bot.knnGun.predict(bot.lastState, bulletSpeed(power))
|
||||
of 13: bot.tmSelector.predict(bot.lastState, bulletSpeed(power))
|
||||
else: bot.headOn.predict(bot.lastState, bulletSpeed(power))
|
||||
let aimTarget = aimAngle(getX(), getY(), pred.x, pred.y)
|
||||
|
||||
@@ -718,7 +737,7 @@ method run*(bot: ModularBot) =
|
||||
|
||||
when isMainModule:
|
||||
var bot = ModularBot(
|
||||
tracker: vb.initTracker(13), # 0: HeadOn, 1: Linear, 2: Tsetlin, 3: Circular, 4: GuessFactor, 5: Pattern, 6: WallBounce, 7: Accel, 8: StopShot, 9: Displace, 10: AvgLead, 11: DecayGF, 12: KNN
|
||||
tracker: vb.initTracker(14), # 0: HeadOn, 1: Linear, 2: Tsetlin, 3: Circular, 4: GuessFactor, 5: Pattern, 6: WallBounce, 7: Accel, 8: StopShot, 9: Displace, 10: AvgLead, 11: DecayGF, 12: KNN, 13: TMSelect
|
||||
headOn: HeadOnGun(),
|
||||
linear: LinearGun(),
|
||||
circular: CircularGun(),
|
||||
@@ -732,6 +751,7 @@ when isMainModule:
|
||||
avgLead: initAveragedLeadGun(),
|
||||
decayGF: initDecayGFGun(),
|
||||
knnGun: initKNNGun(),
|
||||
tmSelector: initTmSelectorGun(),
|
||||
radar: RadarLockModule(),
|
||||
meleeScan: initMeleeScan(),
|
||||
mover: TFILModule(debugGraphics: true),
|
||||
|
||||
@@ -18,7 +18,13 @@ const
|
||||
## full-map long shot (~90 ticks) need ~4700 slots;
|
||||
## 8192 wraps only after ~157 ticks. Each VirtualBullet
|
||||
## is ~120 bytes, so this array costs ~960 KiB.
|
||||
MinHitRate* = 0.40 ## 40% threshold for acceptable power selection
|
||||
MinHitRate* = 0.40 ## LEGACY absolute bar; no longer used by bestPower
|
||||
## (no bin on the live path-metric scale cleared it, so
|
||||
## once every bin had data bestPower fell to power 1.0).
|
||||
PowerBarFrac* = 0.50 ## RELATIVE power bar (dimensionless): a bin is
|
||||
## acceptable when its virtual hit rate is at least this
|
||||
## FRACTION of the same gun's best bin rate. Scales with
|
||||
## the metric instead of assuming a ~40% hit rate.
|
||||
MinObsBeforeCompete* = 50 ## min observations before a gun×bin enters competition
|
||||
TieMargin* = 0.02 ## ABSOLUTE mode: guns within this hit-rate margin of best are tied
|
||||
MinHitRateFloor* = 0.10 ## ABSOLUTE mode: if best gun < this, fall back to gun 0 (HeadOn)
|
||||
@@ -418,26 +424,32 @@ proc fitnessFor*(t: VirtualTracker, targetId: int): seq[GunFitness] =
|
||||
result[gunId].bins[binIdx].record(src.hits[k])
|
||||
|
||||
proc bestPower*(t: VirtualTracker, gunId: GunId, targetId: int = -1): (int, float) =
|
||||
## Returns (binIdx, power) with highest power that has >= MinHitRate.
|
||||
## Falls back to lowest power bin if nothing qualifies yet.
|
||||
## Uses per-enemy fitness when targetId >= 0 and data exists; else aggregate.
|
||||
## Returns (binIdx, power). Prefers the HIGHEST power bin whose virtual hit
|
||||
## rate is acceptable, where "acceptable" is measured RELATIVE to the same
|
||||
## gun's best bin (`rate >= PowerBarFrac * bestBinRate`, dimensionless) — not
|
||||
## against the legacy absolute `MinHitRate`. On the live path-metric scale a
|
||||
## gun's rates sit around 3-40%, so the absolute 40% bar never fired once every
|
||||
## bin had data and bestPower silently collapsed to power 1.0; the relative bar
|
||||
## discriminates between bins at any scale.
|
||||
##
|
||||
## An EMPTY bin is still handed out (highest power first) so every bin keeps
|
||||
## getting sampled, and a fully cold gun (no data anywhere) returns the lowest
|
||||
## power bin. Uses per-enemy fitness when targetId >= 0 and data exists; else
|
||||
## the deterministic aggregate.
|
||||
let fit = t.fitnessFor(targetId)
|
||||
result = (0, PowerBins[0])
|
||||
# Cold gun (zero observations in every bin): fall back to the lowest power bin,
|
||||
# as documented. Without this the countdown loop below would hit the empty
|
||||
# highest bin first and wrongly return power 3.0.
|
||||
var anyObs = false
|
||||
var bestRate = 0.0
|
||||
for binIdx in 0..<len(PowerBins):
|
||||
if fit[gunId].bins[binIdx].count > 0:
|
||||
anyObs = true
|
||||
break
|
||||
bestRate = max(bestRate, fit[gunId].bins[binIdx].hitRate())
|
||||
if not anyObs:
|
||||
return (0, PowerBins[0])
|
||||
# Warm gun: unchanged — return the highest power bin clearing MinHitRate
|
||||
# (or an empty bin, which the existing logic treats as acceptable).
|
||||
let bar = PowerBarFrac * bestRate
|
||||
for binIdx in countdown(len(PowerBins) - 1, 0):
|
||||
let rate = fit[gunId].bins[binIdx].hitRate()
|
||||
if rate >= MinHitRate or fit[gunId].bins[binIdx].count == 0:
|
||||
let fw = fit[gunId].bins[binIdx]
|
||||
if fw.count == 0 or fw.hitRate() >= bar:
|
||||
return (binIdx, PowerBins[binIdx])
|
||||
|
||||
proc chooseFromFit*(fit: seq[GunFitness], diag: ptr SelectorDiag = nil,
|
||||
|
||||
@@ -0,0 +1,328 @@
|
||||
## Tsetlin-machine pattern-classifier gun — a MIXTURE OF EXPERTS with a learned
|
||||
## TM gate.
|
||||
##
|
||||
## Why this exists (and why it is not the existing `tsetlin.nim` gun): the old
|
||||
## Tsetlin gun uses the TM as a pixel-correction REGRESSOR on a single baseline
|
||||
## and ranks 12th of 13. Here the TM does what Granmo's machine is actually
|
||||
## strong at — supervised multi-class classification — on the frame-stacked
|
||||
## binary encoding that measures 99.35-99.48% held-out accuracy at 2 frames.
|
||||
##
|
||||
## Architecture
|
||||
## ------------
|
||||
## experts : HeadOn, Linear, Circular, WallBounce, Accel (the existing
|
||||
## analytic guns, reused unchanged and cheap).
|
||||
## gate : the TM. Every tick it sees the same 2-frame Gray-coded binary
|
||||
## vector as the old Tsetlin gun and votes for the expert most
|
||||
## likely to be right.
|
||||
## label : EXACTLY observable and supervised — at virtual-bullet
|
||||
## resolution the FeedbackEvent carries the enemy's actual
|
||||
## position, so the label is simply which expert's stored
|
||||
## prediction was CLOSEST. No delayed credit, no eligibility trace.
|
||||
## output : the winning class's expert prediction for the requested
|
||||
## bulletSpeed (so the gun implements the ordinary
|
||||
## `predict(state, bulletSpeed)` interface and drops into the rack).
|
||||
##
|
||||
## TM reuse: encoding (`tmEncodeFrame`/`tmEncodeSelf`/`tmEncodeFullVector`) and
|
||||
## ALL learning primitives come from `guns/tsetlin.nim` — in particular the
|
||||
## CORRECTED Granmo feedback (`tmLearnDir`/`tmLearnOne`: Type I conditioned on
|
||||
## the clause output, reachable Type II, the (T - clip(v,-T,T))/(2T) resource
|
||||
## allocation, and the Eq. 6 all-Exclude bootstrap). Nothing is re-derived here.
|
||||
##
|
||||
## Multi-class formulation: Granmo's standard one-clause-team-per-class. Class c
|
||||
## is a binary TM (`d = +1` for the winning expert, `-1` for the rest) and the
|
||||
## predicted class is the argmax of the class votes. `TM_N_OUT = 2` independent
|
||||
## clause teams already live in one `TmNet`, so 5 classes fit in 3 nets.
|
||||
|
||||
import std/[math, random, strformat, algorithm]
|
||||
import gun_harness/gun_interface
|
||||
import gun_harness/virtual_bullets as vb
|
||||
import guns/tsetlin
|
||||
import guns/head_on
|
||||
import guns/linear
|
||||
import guns/circular
|
||||
import guns/wall_bounce
|
||||
import guns/accel_predictor
|
||||
|
||||
const
|
||||
N_EXPERTS* = 5
|
||||
N_NETS = (N_EXPERTS + 1) div 2 ## 3 nets × 2 outputs = 6 class teams
|
||||
SEL_TRACE_SLOTS = 1024 ## exact (fireTick,powerBin) ring, as tsetlin.nim
|
||||
SEL_MIN_OBS = 50 ## observations before the TM outvotes the bootstrap expert
|
||||
DebugSelector* = false
|
||||
|
||||
ExpertNames*: array[N_EXPERTS, string] =
|
||||
["HeadOn", "Linear", "Circular", "WallBounce", "Accel"]
|
||||
|
||||
type
|
||||
TmSelTrace = object
|
||||
fireTick: int
|
||||
powerBin: int
|
||||
preds: array[N_EXPERTS, GunPrediction] ## fire-time expert predictions
|
||||
votes: array[N_NETS * 2, float] ## fire-time class votes (clamped)
|
||||
cache: array[N_NETS, TmClauseCache] ## fire-time LEARNING clause outputs
|
||||
input: TmBinaryVector ## fire-time encoded input
|
||||
alive: bool
|
||||
|
||||
TmSelectorGun* = object
|
||||
nets: array[N_NETS, TmNet]
|
||||
frameBuff: array[2, TmFrameEncoded] ## [0]=newest, [1]=previous
|
||||
frameCount: int
|
||||
lastTick: int
|
||||
input: TmBinaryVector
|
||||
curCaches: array[N_NETS, TmClauseCache]
|
||||
votes: array[N_NETS * 2, float]
|
||||
winCount*: array[N_EXPERTS, int] ## cumulative winners (bootstrap only)
|
||||
totalObs*: int
|
||||
traces: array[SEL_TRACE_SLOTS, TmSelTrace]
|
||||
# experts
|
||||
headOn: HeadOnGun
|
||||
linear: LinearGun
|
||||
circular: CircularGun
|
||||
wallBounce: WallBounceGun
|
||||
accel: AccelGun
|
||||
# instrumentation
|
||||
predictCalls*: int
|
||||
trainCalls*: int
|
||||
traceMisses*: int
|
||||
lastChosen*: int
|
||||
debugGraphics*: bool
|
||||
|
||||
# ── helpers ──────────────────────────────────────────────────────────────────
|
||||
|
||||
proc selBinForSpeed(spd: float): int {.inline.} =
|
||||
for i in 0..<len(vb.PowerBins):
|
||||
if abs(spd - bulletSpeed(vb.PowerBins[i])) < 1e-6:
|
||||
return i
|
||||
-1
|
||||
|
||||
proc selTraceSlot(fireTick, binIdx: int): int {.inline.} =
|
||||
((fireTick * len(vb.PowerBins)) + binIdx) mod SEL_TRACE_SLOTS
|
||||
|
||||
proc expertPred*(g: var TmSelectorGun, idx: int, state: WorldState,
|
||||
bulletSpeed: float): GunPrediction =
|
||||
case idx
|
||||
of 0: g.headOn.predict(state, bulletSpeed)
|
||||
of 1: g.linear.predict(state, bulletSpeed)
|
||||
of 2: g.circular.predict(state, bulletSpeed)
|
||||
of 3: g.wallBounce.predict(state, bulletSpeed)
|
||||
of 4: g.accel.predict(state, bulletSpeed)
|
||||
else: g.headOn.predict(state, bulletSpeed)
|
||||
|
||||
proc initTmSelectorGun*(): TmSelectorGun =
|
||||
# states start at the Exclude boundary (0); one Type I step crosses into Include.
|
||||
for k in 0..<N_NETS:
|
||||
for s in result.nets[k].states.mitems: s = 0'i16
|
||||
result.lastTick = -1
|
||||
result.lastChosen = 2 # Circular
|
||||
randomize()
|
||||
result.debugGraphics = false
|
||||
|
||||
proc isWarmedUp*(g: TmSelectorGun): bool {.inline.} =
|
||||
## The gun needs the 2-frame window to encode; experts handle colder states.
|
||||
g.frameCount >= 2
|
||||
|
||||
# ── class selection ──────────────────────────────────────────────────────────
|
||||
|
||||
proc chooseClass(g: TmSelectorGun): int =
|
||||
## Bootstrap to the empirically best expert until the TM has SEL_MIN_OBS
|
||||
## labels; afterwards take the argmax class vote (ties broken at random, so no
|
||||
## index-0 bias toward HeadOn).
|
||||
if g.totalObs < SEL_MIN_OBS:
|
||||
var bestCount = -1
|
||||
for c in 0..<N_EXPERTS:
|
||||
if g.winCount[c] > bestCount:
|
||||
bestCount = g.winCount[c]
|
||||
result = c
|
||||
if bestCount <= 0: return 2 # Circular — sensible cold default
|
||||
return
|
||||
var bestV = -Inf
|
||||
for c in 0..<N_EXPERTS:
|
||||
if g.votes[c] > bestV: bestV = g.votes[c]
|
||||
var tied: seq[int]
|
||||
for c in 0..<N_EXPERTS:
|
||||
if g.votes[c] >= bestV - 1e-9: tied.add c
|
||||
result = tied[rand(tied.len - 1)]
|
||||
|
||||
# ── Gun interface ────────────────────────────────────────────────────────────
|
||||
|
||||
proc predict*(g: var TmSelectorGun, state: WorldState, bulletSpeed: float): GunPrediction =
|
||||
inc g.predictCalls
|
||||
|
||||
# Encode the current frame and refresh the TM votes at most once per tick.
|
||||
# The harness calls predict() once per power bin (4×/tick); the TM input does
|
||||
# not depend on bulletSpeed, so the forward pass is tick-guarded exactly like
|
||||
# the old Tsetlin gun's window shift.
|
||||
if state.tick != g.lastTick:
|
||||
g.lastTick = state.tick
|
||||
g.frameBuff[1] = g.frameBuff[0]
|
||||
let dist = hypot(state.enemyX - state.selfX, state.enemyY - state.selfY)
|
||||
let bearing = radToDeg(arctan2(state.enemyY - state.selfY, state.enemyX - state.selfX))
|
||||
g.frameBuff[0] = tmEncodeFrame(
|
||||
bearing, dist, state.enemySpeed, state.enemyHeading,
|
||||
state.arenaHeight - state.enemyY, state.enemyY,
|
||||
state.arenaWidth - state.enemyX, state.enemyX,
|
||||
state.enemyEnergy)
|
||||
if g.frameCount < 2: inc g.frameCount
|
||||
|
||||
# 2-frame window (the measured sweet spot: 39.6 effective literals/clause vs
|
||||
# 152.8 at 10 frames). All TM_WINDOW_SIZE slots are filled with a REAL frame
|
||||
# (alternating newest/previous) rather than zero-padding: constant-zero
|
||||
# literals have an always-true negation, which Type I then includes en masse
|
||||
# and saturates every clause (measured mean 256-755 included literals).
|
||||
# Duplicating the two real frames keeps every literal variable. Reuses the
|
||||
# shared 10-frame encoder/vector so the corrected TM primitives apply
|
||||
# unchanged.
|
||||
var window: array[TM_WINDOW_SIZE, TmFrameEncoded]
|
||||
let prev = if g.frameCount >= 2: g.frameBuff[1] else: g.frameBuff[0]
|
||||
for i in 0..<TM_WINDOW_SIZE:
|
||||
window[i] = if (i and 1) == 0: g.frameBuff[0] else: prev
|
||||
let selfState = tmEncodeSelf(
|
||||
state.arenaHeight - state.selfY, state.selfY,
|
||||
state.arenaWidth - state.selfX, state.selfX,
|
||||
state.selfEnergy, true)
|
||||
g.input = tmEncodeFullVector(window, selfState)
|
||||
for k in 0..<N_NETS:
|
||||
let (_, _, vx, vy) = tmForwardWithCache(g.nets[k], g.input, g.curCaches[k])
|
||||
g.votes[2 * k] = vx
|
||||
g.votes[2 * k + 1] = vy
|
||||
|
||||
var preds: array[N_EXPERTS, GunPrediction]
|
||||
for c in 0..<N_EXPERTS:
|
||||
preds[c] = g.expertPred(c, state, bulletSpeed)
|
||||
|
||||
let chosen = g.chooseClass()
|
||||
g.lastChosen = chosen
|
||||
when DebugSelector:
|
||||
echo fmt"[sel-dbg] tick={state.tick} bin={selBinForSpeed(bulletSpeed)} chosen={ExpertNames[chosen]} " &
|
||||
fmt"votes=[{g.votes[0]:.1f},{g.votes[1]:.1f},{g.votes[2]:.1f},{g.votes[3]:.1f},{g.votes[4]:.1f}] obs={g.totalObs}"
|
||||
|
||||
let binIdx = selBinForSpeed(bulletSpeed)
|
||||
if binIdx >= 0:
|
||||
let slot = selTraceSlot(state.tick, binIdx)
|
||||
g.traces[slot] = TmSelTrace(
|
||||
fireTick: state.tick, powerBin: binIdx,
|
||||
preds: preds, votes: g.votes, cache: g.curCaches, input: g.input, alive: true)
|
||||
|
||||
preds[chosen]
|
||||
|
||||
proc onResult*(g: var TmSelectorGun, e: FeedbackEvent) =
|
||||
let binIdx =
|
||||
if e.powerBin >= 0 and e.powerBin < len(vb.PowerBins): e.powerBin
|
||||
else: selBinForSpeed(bulletSpeed(e.bulletPower))
|
||||
if binIdx < 0:
|
||||
inc g.traceMisses
|
||||
return
|
||||
let slot = selTraceSlot(e.fireTick, binIdx)
|
||||
var t = addr g.traces[slot]
|
||||
if not t.alive or t.fireTick != e.fireTick or t.powerBin != binIdx:
|
||||
inc g.traceMisses
|
||||
return
|
||||
|
||||
# The label: which expert's fire-time prediction was closest to the actual
|
||||
# enemy position the virtual bullet resolved against. Exact and supervised.
|
||||
var winner = 0
|
||||
var bestD = Inf
|
||||
for c in 0..<N_EXPERTS:
|
||||
let d = hypot(t.preds[c].x - e.actualX, t.preds[c].y - e.actualY)
|
||||
if d < bestD:
|
||||
bestD = d
|
||||
winner = c
|
||||
inc g.winCount[winner]
|
||||
inc g.totalObs
|
||||
inc g.trainCalls
|
||||
|
||||
let lits = tmMakeLiterals(t.input)
|
||||
for c in 0..<N_EXPERTS:
|
||||
let k = c div 2
|
||||
let o = c mod 2
|
||||
let d = if c == winner: 1.0 else: -1.0
|
||||
g.nets[k].tmLearnDir(o, lits, t.cache[k], t.votes[c], d)
|
||||
t.alive = false
|
||||
|
||||
# ── interpretability ─────────────────────────────────────────────────────────
|
||||
|
||||
const
|
||||
FrameBitNames: array[TM_FRAME_BITS, string] = block:
|
||||
var a: array[TM_FRAME_BITS, string]
|
||||
for i in 0..<8: a[i] = "bearSin"
|
||||
for i in 0..<8: a[8+i] = "bearCos"
|
||||
for i in 0..<7: a[16+i] = "dist"
|
||||
for i in 0..<5: a[23+i] = "vel"
|
||||
for i in 0..<8: a[28+i] = "headSin"
|
||||
for i in 0..<8: a[36+i] = "headCos"
|
||||
for i in 0..<7: a[44+i] = "wallN"
|
||||
for i in 0..<7: a[51+i] = "wallS"
|
||||
for i in 0..<7: a[58+i] = "wallE"
|
||||
for i in 0..<7: a[65+i] = "wallW"
|
||||
for i in 0..<11: a[72+i] = "energy"
|
||||
a
|
||||
|
||||
proc bitName(bit: int): string =
|
||||
## Human name for a base vector bit index (0..TM_N_IN-1).
|
||||
if bit < TM_FRAME_BITS * TM_WINDOW_SIZE:
|
||||
let frame = bit div TM_FRAME_BITS
|
||||
let off = bit mod TM_FRAME_BITS
|
||||
result = fmt"f{frame}.{FrameBitNames[off]}"
|
||||
else:
|
||||
let off = bit - TM_FRAME_BITS * TM_WINDOW_SIZE
|
||||
if off < 7: result = "self.wallN"
|
||||
elif off < 14: result = "self.wallS"
|
||||
elif off < 21: result = "self.wallE"
|
||||
elif off < 28: result = "self.wallW"
|
||||
elif off < 39: result = "self.energy"
|
||||
else: result = "self.canFire"
|
||||
|
||||
proc litName(lit: int): string =
|
||||
## Literal `lit` is positive when `lit < TM_N_IN`, negated otherwise.
|
||||
if lit < TM_N_IN: bitName(lit) & "=1"
|
||||
else: bitName(lit - TM_N_IN) & "=0"
|
||||
|
||||
proc selectorClauseStats*(g: TmSelectorGun): tuple[nClauses, nActive: int,
|
||||
meanIncluded: float] =
|
||||
## Include-count over all class clause teams (5 × 50 clauses).
|
||||
var total = 0
|
||||
result.nClauses = N_EXPERTS * TM_N_CLAUSES
|
||||
for c in 0..<N_EXPERTS:
|
||||
let k = c div 2
|
||||
let o = c mod 2
|
||||
for cl in 0..<TM_N_CLAUSES:
|
||||
var inc = 0
|
||||
for lit in 0..<TM_N_LITERALS:
|
||||
if g.nets[k].states[tmStateIdx(o, cl, lit)] > 0: inc += 1
|
||||
total += inc
|
||||
if inc > 0: inc result.nActive
|
||||
if result.nActive > 0:
|
||||
result.meanIncluded = total.float / result.nActive.float
|
||||
|
||||
proc describeClauses*(g: TmSelectorGun, topN = 10): string =
|
||||
## Per expert class, the feature literals included most often (and only in
|
||||
## slots f0/f1 — slots 2..9 are the zero padding). One line per top feature:
|
||||
## `feature=value × count` where count is how many of the class's 50 clauses
|
||||
## include it. This is the interpretability payoff: it shows WHAT the gate
|
||||
## switches on.
|
||||
for c in 0..<N_EXPERTS:
|
||||
let k = c div 2
|
||||
let o = c mod 2
|
||||
# count clauses per literal
|
||||
var counts: array[TM_N_LITERALS, int]
|
||||
for cl in 0..<TM_N_CLAUSES:
|
||||
for lit in 0..<TM_N_LITERALS:
|
||||
if g.nets[k].states[tmStateIdx(o, cl, lit)] > 0: inc counts[lit]
|
||||
var idx: seq[int]
|
||||
for lit in 0..<TM_N_LITERALS:
|
||||
# Frames alternate f0 (newest) / f1 (previous); slots 2..9 are duplicates,
|
||||
# so the unique selection signal lives in the first two frames. Report
|
||||
# only those to keep the dump readable.
|
||||
let base = if lit < TM_N_IN: lit else: lit - TM_N_IN
|
||||
if base >= TM_FRAME_BITS * 2 and base < TM_FRAME_BITS * TM_WINDOW_SIZE: continue
|
||||
if counts[lit] > 0: idx.add lit
|
||||
idx.sort(proc(a, b: int): int = counts[b] - counts[a])
|
||||
result.add fmt"class {c} ({ExpertNames[c]}):"
|
||||
if idx.len == 0:
|
||||
result.add " <no active literals>\n"
|
||||
continue
|
||||
result.add " "
|
||||
for i in 0..<min(topN, idx.len):
|
||||
result.add fmt"{litName(idx[i])}×{counts[idx[i]]}"
|
||||
if i < min(topN, idx.len) - 1: result.add ", "
|
||||
result.add "\n"
|
||||
@@ -95,28 +95,28 @@ proc tmEncodeFullVector*(window: array[TM_WINDOW_SIZE, TmFrameEncoded],
|
||||
# ── Tsetlin Machine (adapted from BNNBot_garage/src/tsetlin_predictor.nim) ───
|
||||
|
||||
const
|
||||
TM_N_IN = TM_TOTAL_BITS # 870
|
||||
TM_N_OUT = 2 # cx, cy pixel corrections
|
||||
TM_N_LITERALS = TM_N_IN * 2 # 1740
|
||||
TM_N_CLAUSES = 50 # per output; issue #184 default
|
||||
TM_HALF = TM_N_CLAUSES div 2
|
||||
TM_N_STATES = 32 # automaton range [-32..32]
|
||||
TM_T = float(TM_HALF) # vote clamped to [-T, T]
|
||||
TM_S = 1.5 # specificity
|
||||
TM_N_IN* = TM_TOTAL_BITS # 870
|
||||
TM_N_OUT* = 2 # cx, cy pixel corrections
|
||||
TM_N_LITERALS* = TM_N_IN * 2 # 1740
|
||||
TM_N_CLAUSES* = 50 # per output; issue #184 default
|
||||
TM_HALF* = TM_N_CLAUSES div 2
|
||||
TM_N_STATES* = 32 # automaton range [-32..32]
|
||||
TM_T* = float(TM_HALF) # vote clamped to [-T, T]
|
||||
TM_S* = 1.5 # specificity
|
||||
TM_RESID_MAX = 80.0 # pixel correction range
|
||||
# ponytail: TM_N_STATES=32 needs int16 (int8 only fits ≤127, fine here); raise N_CLAUSES if underfitting
|
||||
|
||||
type
|
||||
TmClauseCache = array[TM_N_OUT * TM_N_CLAUSES, uint8]
|
||||
TmClauseCache* = array[TM_N_OUT * TM_N_CLAUSES, uint8]
|
||||
|
||||
TmNet = object
|
||||
states: array[TM_N_OUT * TM_N_CLAUSES * TM_N_LITERALS, int16]
|
||||
TmNet* = object
|
||||
states*: array[TM_N_OUT * TM_N_CLAUSES * TM_N_LITERALS, int16]
|
||||
# ponytail: int16 to safely hold [-32..32]; TM_N_STATES=32 fits int8 too but int16 is safer
|
||||
|
||||
proc tmStateIdx(outIdx, clause, lit: int): int {.inline.} =
|
||||
proc tmStateIdx*(outIdx, clause, lit: int): int {.inline.} =
|
||||
(outIdx * TM_N_CLAUSES + clause) * TM_N_LITERALS + lit
|
||||
|
||||
proc tmPolarity(clause: int): float {.inline.} =
|
||||
proc tmPolarity*(clause: int): float {.inline.} =
|
||||
if clause < TM_HALF: 1.0 else: -1.0
|
||||
|
||||
type
|
||||
@@ -134,12 +134,12 @@ type
|
||||
meanIncluded*: float ## mean include count over ACTIVE clauses
|
||||
meanIncludedAll*: float ## mean include count over ALL clauses (incl. empty)
|
||||
|
||||
proc tmMakeLiterals(input: TmBinaryVector): array[TM_N_LITERALS, uint8] =
|
||||
proc tmMakeLiterals*(input: TmBinaryVector): array[TM_N_LITERALS, uint8] =
|
||||
for i in 0..<TM_N_IN:
|
||||
result[i] = input[i]
|
||||
result[i + TM_N_IN] = 1'u8 - input[i]
|
||||
|
||||
proc tmEvalClause(net: TmNet, outIdx, clause: int,
|
||||
proc tmEvalClause*(net: TmNet, outIdx, clause: int,
|
||||
lits: array[TM_N_LITERALS, uint8],
|
||||
learning = false): uint8 =
|
||||
var hasIncluded = false
|
||||
@@ -157,7 +157,7 @@ proc tmEvalClause(net: TmNet, outIdx, clause: int,
|
||||
# clause at empty forever.
|
||||
return if learning: 1'u8 else: 0'u8
|
||||
|
||||
proc tmForwardWithCache(net: TmNet, input: TmBinaryVector,
|
||||
proc tmForwardWithCache*(net: TmNet, input: TmBinaryVector,
|
||||
cache: var TmClauseCache): (float, float, float, float) =
|
||||
## Returns (correctionX, correctionY, voteX, voteY). `cache` receives the
|
||||
## clause outputs under LEARNING semantics (empty clause = 1) for tmLearnOne;
|
||||
@@ -176,23 +176,19 @@ proc tmForwardWithCache(net: TmNet, input: TmBinaryVector,
|
||||
vy = clamp(vy, -TM_T, TM_T)
|
||||
(vx / TM_T * TM_RESID_MAX, vy / TM_T * TM_RESID_MAX, vx, vy)
|
||||
|
||||
proc tmLearnOne(net: var TmNet, outIdx: int, lits: array[TM_N_LITERALS, uint8],
|
||||
cache: TmClauseCache, vote: float, residual: float) =
|
||||
## One faithful Granmo Table 2/3 update against a continuous residual target.
|
||||
proc tmLearnDir*(net: var TmNet, outIdx: int, lits: array[TM_N_LITERALS, uint8],
|
||||
cache: TmClauseCache, vote: float, d: float) =
|
||||
## One faithful Granmo Table 2/3 clause update with an EXPLICIT desired vote
|
||||
## direction `d` in {-1, +1}. This is the exact corrected core the Tsetlin gun
|
||||
## uses; `tmLearnOne` is the regression wrapper that derives `d` from a
|
||||
## continuous residual, and the multi-class selector passes the class label
|
||||
## directly.
|
||||
##
|
||||
## `cache` holds clause outputs under LEARNING semantics (empty = 1, Eq. 6).
|
||||
## `vote` is the classification-semantics clause sum at prediction time,
|
||||
## already clamped to [-TM_T, TM_T]. `residual` is the correction target
|
||||
## delta = actual - linear baseline (see onResult), NOT actual - prediction.
|
||||
##
|
||||
## Regression adaptation: the "label" direction is d = sign(residual -
|
||||
## predicted), i.e. which way the correction must move. The resource
|
||||
## allocation of Granmo Eq. 8-11 collapses to a single probability
|
||||
## p = (T - d*clip(v,-T,T)) / (2T), applied both to the aligned clauses
|
||||
## (Type I) and the opposed ones (Type II), exactly as Algorithm 1 lines 11-22.
|
||||
let predicted = vote / TM_T * TM_RESID_MAX
|
||||
let error = residual - predicted
|
||||
let d = if error > 0.0: 1.0 elif error < 0.0: -1.0 else: return
|
||||
## already clamped to [-TM_T, TM_T]. The Granmo resource allocation collapses
|
||||
## to a single probability p = (T - d*clip(v,-T,T)) / (2T): it is high when the
|
||||
## vote opposes `d` and falls to 0 once the class is already won.
|
||||
let pFeedback = (TM_T - d * vote) / (2.0 * TM_T)
|
||||
if pFeedback <= 0.0: return
|
||||
|
||||
@@ -228,6 +224,25 @@ proc tmLearnOne(net: var TmNet, outIdx: int, lits: array[TM_N_LITERALS, uint8],
|
||||
if net.states[si] <= 0:
|
||||
net.states[si] = int16(min(int(net.states[si]) + 1, TM_N_STATES))
|
||||
|
||||
proc tmLearnOne*(net: var TmNet, outIdx: int, lits: array[TM_N_LITERALS, uint8],
|
||||
cache: TmClauseCache, vote: float, residual: float) =
|
||||
## One faithful Granmo Table 2/3 update against a continuous residual target.
|
||||
##
|
||||
## `cache` holds clause outputs under LEARNING semantics (empty = 1, Eq. 6).
|
||||
## `vote` is the classification-semantics clause sum at prediction time,
|
||||
## already clamped to [-TM_T, TM_T]. `residual` is the correction target
|
||||
## delta = actual - linear baseline (see onResult), NOT actual - prediction.
|
||||
##
|
||||
## Regression adaptation: the "label" direction is d = sign(residual -
|
||||
## predicted), i.e. which way the correction must move. The resource
|
||||
## allocation of Granmo Eq. 8-11 collapses to a single probability
|
||||
## p = (T - d*clip(v,-T,T)) / (2T), applied both to the aligned clauses
|
||||
## (Type I) and the opposed ones (Type II), exactly as Algorithm 1 lines 11-22.
|
||||
let predicted = vote / TM_T * TM_RESID_MAX
|
||||
let error = residual - predicted
|
||||
let d = if error > 0.0: 1.0 elif error < 0.0: -1.0 else: return
|
||||
net.tmLearnDir(outIdx, lits, cache, vote, d)
|
||||
|
||||
# ── TsetlinGun public type ────────────────────────────────────────────────────
|
||||
|
||||
const
|
||||
|
||||
@@ -34,6 +34,9 @@ const
|
||||
runnerJar = "/home/davide/Projects/tank-royale/runner/examples/lib/robocode-tankroyale-runner.jar"
|
||||
|
||||
const TsetlinId = 2
|
||||
const TmSelectorId = 13 ## also stochastic (rand() in Gate choose + TM feedback)
|
||||
|
||||
proc isStochastic(id: int): bool = id == TsetlinId or id == TmSelectorId
|
||||
|
||||
proc lastOnlineRound(path: string): JsonNode =
|
||||
result = nil
|
||||
@@ -78,12 +81,12 @@ proc main() =
|
||||
let reports = replayFixture(fx, buildAllGunDrivers(), liveActual = true)
|
||||
|
||||
# Map online stats by gun id.
|
||||
var onShots: array[13, int]
|
||||
var onHits: array[13, int]
|
||||
var onNames: array[13, string]
|
||||
var onShots: array[14, int]
|
||||
var onHits: array[14, int]
|
||||
var onNames: array[14, string]
|
||||
for g in online["guns"]:
|
||||
let id = g["id"].getInt()
|
||||
if id >= 0 and id < 13:
|
||||
if id >= 0 and id < 14:
|
||||
onShots[id] = g["vShots"].getInt()
|
||||
onHits[id] = g["vHits"].getInt()
|
||||
onNames[id] = g["name"].getStr()
|
||||
@@ -96,11 +99,11 @@ proc main() =
|
||||
echo "-----------------------------------------------------------------------"
|
||||
var matches = 0
|
||||
var deterministic = 0
|
||||
for id in 0..<13:
|
||||
for id in 0..<14:
|
||||
let r = reports[id]
|
||||
let match = r.hits == onHits[id] and r.shots == onShots[id]
|
||||
var verdict: string
|
||||
if id == TsetlinId:
|
||||
if isStochastic(id):
|
||||
verdict = if match: "MATCH (stochastic)" else: "differs (stochastic, expected)"
|
||||
else:
|
||||
inc deterministic
|
||||
@@ -117,7 +120,7 @@ proc main() =
|
||||
echo "VERDICT: FAIL — offline range does NOT reproduce the live metric."
|
||||
quit(1)
|
||||
echo "VERDICT: PASS — offline == online for all 12 deterministic guns."
|
||||
echo "(Tsetlin is stochastic and is allowed to differ.)"
|
||||
echo "(Tsetlin and TMSelect are stochastic and are allowed to differ.)"
|
||||
|
||||
when isMainModule:
|
||||
main()
|
||||
|
||||
@@ -17,12 +17,16 @@ import guns/displacement
|
||||
import guns/averaged_lead
|
||||
import guns/decay_gf
|
||||
import guns/knn_gun
|
||||
import guns/tm_selector
|
||||
|
||||
proc buildAllGunDrivers*(seed = -1): seq[GunDriver] =
|
||||
## seed >= 0 re-seeds the global RNG after constructing Tsetlin so the
|
||||
## stochastic gun's learning is reproducible for offline runs. (Its
|
||||
## constructor calls randomize(); we override that seed afterwards.)
|
||||
## seed >= 0 re-seeds the global RNG after constructing the stochastic guns
|
||||
## (Tsetlin and the TM selector both call randomize() in their constructors),
|
||||
## so their learning is reproducible for offline runs.
|
||||
##
|
||||
## Order matches ModularBot's gun ids exactly (TMSelect appended at 13).
|
||||
var tsetlin = initTsetlinGun()
|
||||
var tmSelector = initTmSelectorGun()
|
||||
if seed >= 0:
|
||||
randomize(seed)
|
||||
result = @[
|
||||
@@ -39,6 +43,7 @@ proc buildAllGunDrivers*(seed = -1): seq[GunDriver] =
|
||||
makeDriver("AvgLead", initAveragedLeadGun()),
|
||||
makeDriver("DecayGF", initDecayGFGun()),
|
||||
makeDriver("KNN", initKNNGun()),
|
||||
makeDriver("TMSelect", tmSelector),
|
||||
]
|
||||
|
||||
proc makeTsetlinDriver*(seed = -1): tuple[driver: GunDriver, gun: ref TsetlinGun] =
|
||||
@@ -57,3 +62,19 @@ proc makeTsetlinDriver*(seed = -1): tuple[driver: GunDriver, gun: ref TsetlinGun
|
||||
resultCb: proc(e: FeedbackEvent) = g[].onResult(e),
|
||||
readyCb: proc(): bool = g[].isWarmedUp(),
|
||||
)
|
||||
|
||||
proc makeTmSelectorDriver*(seed = -1): tuple[driver: GunDriver, gun: ref TmSelectorGun] =
|
||||
## Same as makeDriver("TMSelect", ...) but keeps a handle to the concrete gun
|
||||
## so a test can inspect its votes / clause interpretability after a replay.
|
||||
let g = new(TmSelectorGun)
|
||||
g[] = initTmSelectorGun()
|
||||
if seed >= 0:
|
||||
randomize(seed)
|
||||
result.gun = g
|
||||
result.driver = GunDriver(
|
||||
name: "TMSelect",
|
||||
predictCb: proc(state: WorldState, bulletSpeed: float): GunPrediction =
|
||||
g[].predict(state, bulletSpeed),
|
||||
resultCb: proc(e: FeedbackEvent) = g[].onResult(e),
|
||||
readyCb: proc(): bool = g[].isWarmedUp(),
|
||||
)
|
||||
|
||||
@@ -0,0 +1,105 @@
|
||||
## Guard + measurement for the scale-aware power bar in `bestPower`.
|
||||
##
|
||||
## Before the fix `bestPower` required an ABSOLUTE virtual hit rate >= MinHitRate
|
||||
## (0.40). On the shipped path-metric scale a gun's per-bin rates sit around
|
||||
## 3-40%, so once every bin had data no bin cleared the bar and bestPower
|
||||
## silently collapsed to bin 0 (power 1.0). The fix compares each bin's rate to
|
||||
## PowerBarFrac * (the same gun's best bin rate) — dimensionless, so it
|
||||
## discriminates at any scale.
|
||||
##
|
||||
## Run: nim c -r common_libs/tests/test_power_selection.nim
|
||||
|
||||
import std/[strformat, math, random, tables, os]
|
||||
import gun_harness/gun_interface
|
||||
import gun_harness/virtual_bullets
|
||||
import gun_harness/offline_range
|
||||
import range_guns
|
||||
|
||||
const fixturesDir = currentSourcePath().parentDir.parentDir.parentDir / "tools" / "fixtures"
|
||||
|
||||
var failures = 0
|
||||
proc check(name: string, ok: bool) =
|
||||
if ok: echo "PASS: ", name
|
||||
else: echo "FAIL: ", name; inc failures
|
||||
|
||||
proc recordHit(fw: var FitnessWindow, hit: bool) =
|
||||
fw.hits[fw.head] = hit
|
||||
fw.head = (fw.head + 1) mod WindowSize
|
||||
inc fw.count
|
||||
|
||||
proc seedFromReport(t: var VirtualTracker, targetId, gunId: int, r: GunReport) =
|
||||
if targetId notin t.fitness:
|
||||
t.fitness[targetId] = newSeq[GunFitness](t.numGuns)
|
||||
var fw = addr t.fitness[targetId][gunId].bins
|
||||
for b in 0..<len(PowerBins):
|
||||
for _ in 0..<r.bins[b].hits: recordHit(fw[][b], true)
|
||||
for _ in 0..<(r.bins[b].shots - r.bins[b].hits): recordHit(fw[][b], false)
|
||||
|
||||
proc oldBestPower(fit: GunFitness): int =
|
||||
## The pre-fix rule, reproduced for the A/B comparison only.
|
||||
var anyObs = false
|
||||
for b in 0..<len(PowerBins):
|
||||
if fit.bins[b].count > 0: anyObs = true
|
||||
if not anyObs: return 0
|
||||
for b in countdown(len(PowerBins) - 1, 0):
|
||||
if fit.bins[b].hitRate() >= MinHitRate or fit.bins[b].count == 0: return b
|
||||
0
|
||||
|
||||
proc main() =
|
||||
# A real closed-loop fixture where every gun's rates sit below the old 40%
|
||||
# absolute bar — the regime the bug lives in. Fall back to a synthetic fixture
|
||||
# if the committed capture is missing.
|
||||
let realPath = fixturesDir / "drussgt_vs_crazy.jsonl"
|
||||
let fx = if fileExists(realPath): loadFixture(realPath)
|
||||
else: synthesizeByName("energy-threshold-turner")
|
||||
echo "fixture: ", fx.meta.adversary, " (source=", fx.meta.source, ")"
|
||||
let reports = replayFixture(fx, buildAllGunDrivers(seed = 1), metric = bmPath)
|
||||
|
||||
var t = initTracker(reports.len)
|
||||
let tid = fx.enemyId
|
||||
for gi, r in reports:
|
||||
seedFromReport(t, tid, gi, r)
|
||||
|
||||
echo "per-gun per-bin virtual hit rate (measured) and selected bin (old | new):"
|
||||
var changed = 0
|
||||
var anyOldFellToZero = false
|
||||
for gi, r in reports:
|
||||
let fit = t.fitnessFor(tid)
|
||||
var rates = ""
|
||||
for b in 0..<len(PowerBins):
|
||||
rates.add fmt" p{PowerBins[b]:.1f}={fit[gi].bins[b].hitRate()*100:4.1f}%({r.bins[b].shots})"
|
||||
let ob = oldBestPower(fit[gi])
|
||||
let (nb, _) = t.bestPower(gi, tid)
|
||||
if ob != nb: inc changed
|
||||
# The bug: old rule returns bin 0 even though a higher bin is as good/better.
|
||||
if ob == 0 and nb > 0: anyOldFellToZero = true
|
||||
echo fmt" {r.name:<11}{rates} old=bin{ob} new=bin{nb}"
|
||||
|
||||
echo ""
|
||||
echo fmt"selections changed by the fix: {changed}/{reports.len}"
|
||||
check "the scale-aware bar changes at least one gun's power selection",
|
||||
changed > 0
|
||||
check "at least one gun the old absolute 0.40 bar sent to power 1.0 now uses a heavier bullet",
|
||||
anyOldFellToZero
|
||||
|
||||
# Relative bar still picks the best bin when it is the lowest power (no
|
||||
# pathological power-3 bias).
|
||||
block:
|
||||
var t2 = initTracker(1)
|
||||
# bin0 30%, bin1 25%, bin2 12%, bin3 5% -> bar=15%, only bins 0 and 1 clear
|
||||
# it; the HIGHEST acceptable is bin 1, not bin 0 and not bin 3.
|
||||
seedFromReport(t2, 7, 0, GunReport(name: "x", bins: [
|
||||
BinStat(shots: 100, hits: 30), BinStat(shots: 100, hits: 25),
|
||||
BinStat(shots: 100, hits: 12), BinStat(shots: 100, hits: 5)]))
|
||||
let (b, p) = t2.bestPower(0, 7)
|
||||
echo fmt" synthetic [30,25,12,5]% -> bin {b} (power {p})"
|
||||
check "relative bar picks the highest bin clearing 50% of the best (bin 1)",
|
||||
b == 1
|
||||
|
||||
if failures > 0:
|
||||
echo "\n", failures, " check(s) FAILED"
|
||||
quit(1)
|
||||
echo "\nAll power-selection checks passed."
|
||||
|
||||
when isMainModule:
|
||||
main()
|
||||
@@ -0,0 +1,37 @@
|
||||
## TM selector gun checks + interpretability dump.
|
||||
## Run: nim c -r common_libs/tests/test_tm_selector.nim
|
||||
|
||||
import std/[strformat, math, random]
|
||||
import gun_harness/offline_range
|
||||
import range_guns
|
||||
import guns/tm_selector
|
||||
import guns/head_on
|
||||
import guns/linear
|
||||
import guns/circular
|
||||
import guns/wall_bounce
|
||||
import guns/accel_predictor
|
||||
|
||||
proc main() =
|
||||
let names = ["decel-before-turn", "energy-threshold-turner", "oscillator", "random-walk"]
|
||||
for name in names:
|
||||
let fx = synthesizeByName(name)
|
||||
let (drv, gun) = makeTmSelectorDriver(seed = 1)
|
||||
let reps = replayFixture(fx, @[
|
||||
makeDriver("HeadOn", HeadOnGun()),
|
||||
makeDriver("Linear", LinearGun()),
|
||||
makeDriver("Circular", CircularGun()),
|
||||
makeDriver("WallBounce", initWallBounceGun()),
|
||||
makeDriver("Accel", initAccelGun()),
|
||||
drv], metric = bmPath)
|
||||
echo "=== ", name
|
||||
for r in reps: echo " ", formatReportRow(r)
|
||||
let st = gun[].selectorClauseStats()
|
||||
var wc = ""
|
||||
for c in 0..<N_EXPERTS: wc.add fmt" {ExpertNames[c]}={gun[].winCount[c]}"
|
||||
echo fmt" TMSelect trainCalls={gun[].trainCalls} traceMisses={gun[].traceMisses} " &
|
||||
fmt"clauses active={st.nActive}/{st.nClauses} meanIncl={st.meanIncluded:.1f}"
|
||||
echo " winner counts:", wc
|
||||
echo gun[].describeClauses(topN = 5)
|
||||
|
||||
when isMainModule:
|
||||
main()
|
||||
Reference in New Issue
Block a user