Re-adaptation: the sliding WINDOW is a control-validated +9.3pp on late accuracy
The user's first-hand diagnosis: "once the pattern is learnt enough we hit DrussGT, but as soon as it adapts we are not fast enough to re-adapt again." The gun kept EVERY sample for the whole battle (which the user explicitly asked for), so stale evidence weighed the same as new evidence - accumulation without forgetting. TASK 1 - THE CURVE, MEASURED FIRST (prequential side accuracy, 15 rounds x 4 horizons, deciles of the eval stream): arm early% late% decay frozen (early only) 64.9 61.4 -3.5 accum (shipped) 76.4 75.3 -1.1 window N=150 84.7 84.6 -0.0 resetdrop 5pp 83.2 84.3 +1.0 **The decay IS real but modest and LOCALIZED IN THE LAST ~30% of the round** (last-third vs middle-third: frozen -7.7pp, accum -3.7, window -1.3, resetdrop -1.3). TASK 3 - WHICH FIX HELPS (late-half accuracy, shuffled control in parens): accum (shipped) 75.3 **window N=150 84.6 (51)** +9.3pp **resetdrop 5pp 84.3 (51)** +9.0pp rehearse-all (retrain, NO forgetting) 79.5 - **Sliding window: +9.3pp late (within-round), +9.7 (cross-round), +10.3 (shield).** The shuffled control stays ~50-51%, so it learns the ENEMY, not noise. - **Forgetting is the essential ingredient**: the same periodic retrain WITHOUT forgetting reaches only ~79.5%, so roughly half the gain is the retraining mechanics and half is the forgetting. - Change detection ties the window on late accuracy and gives the best decay, but at a 5pp threshold it fired 300-500 times in the offline stream - noisy. - INERTIA IS REDUNDANT WITH THE WINDOW: lower inertia helps the keep-everything model a lot (accum late 75.3 -> 83.3 at N=8) but leaves the window FLAT (84.6-84.7 at every N). Low inertia and forgetting are SUBSTITUTES, and forgetting is the robust one - `TR_TMHORIZON_NSTATES=8` is NOT the primary fix. VERDICT: SHIPPING CANDIDATE = `TR_TMHORIZON_WINDOW=150`, kept at default 0 until a live A/B confirms. HONEST CAVEATS THAT SET EXPECTATIONS: - The fixtures are OPEN-LOOP (DrussGT does not react to our bullets), so a true mid-round adaptation is NOT present; the dominant measured effect is the LEVEL gap, not the decay magnitude. The user's "it adapts" magnitude is still INFERRED. - The harness uses a STRAIGHT-LINE base while the live gun uses Pattern's prediction, and it ignores the h-tick label delay, so its absolute side accuracy (75-85%) is INFLATED: the same gun measured ~52% = chance live against DrussGT. So +9.3pp is a real ARM DELTA, not a promise that the gun now clears the ~80% accuracy wall that hits need. A live A/B must decide. Adds `measure_tm_readapt.nim` (prequential harness with the mandatory shuffled control, within-round + cross-round protocols, inertia sweep) and its captured results. test_tm_horizon 79 -> 104 (five new groups). test_tm_diag 48, test_tm_automata_diag 55, test_tm_clause_shape 66, test_rack_membership 48 pass. ModularBot compiles. .gitignore: switched from a broad `measure_*` pattern to EXPLICIT binary names. The broad rule was too blunt - it also excluded the `.txt` results file, which made `git add` refuse the whole commit twice. Sources stay tracked; binaries do not.
This commit is contained in:
+15
@@ -51,3 +51,18 @@ common_libs/tests/measure_power_policy
|
||||
common_libs/tests/test_tsetlin_gun
|
||||
common_libs/tests/test_tsetlin_live
|
||||
common_libs/tests/test_tm_pattern_learning
|
||||
|
||||
# measurement tool binaries (no extension); sources are measure_*.nim
|
||||
tr_bots/
|
||||
|
||||
# Measurement-tool binaries (extensionless; their .nim/.py/.txt sources stay tracked)
|
||||
common_libs/tests/measure_tm_readapt
|
||||
common_libs/tests/measure_tm_pattern_cost
|
||||
common_libs/tests/measure_cornering_guns
|
||||
common_libs/tests/measure_vbullet_admit_gate
|
||||
common_libs/tests/measure_tm_miss_shrink
|
||||
common_libs/tests/audit_wave_pairing
|
||||
common_libs/tests/compare_pairing
|
||||
common_libs/tests/diag_synthetic
|
||||
common_libs/tests/diag_automata_validation
|
||||
common_libs/tests/diag_tm_pattern_offline
|
||||
|
||||
@@ -56,6 +56,31 @@
|
||||
## per tick, one TM evaluation per horizon bucket), so it stays in the
|
||||
## neighbourhood of the 0.36 ms/tick TM gun rather than Tsetlin's ~5.3 ms/tick.
|
||||
##
|
||||
## RE-ADAPTATION (all knobs default to the behaviour above — nothing regresses):
|
||||
## The user's failure mode is "it learns the enemy, the enemy adapts, and we
|
||||
## are too slow to re-adapt". The cause is that the machines accumulate EVERY
|
||||
## resolved sample for the whole battle, so old evidence weighs as much as new.
|
||||
## Three independent, off-by-default fixes:
|
||||
## * TR_TMHORIZON_WINDOW=N (>0): every TR_TMHORIZON_RETRAIN_EVERY samples
|
||||
## (default 50) rebuild BOTH heads from scratch and retrain on the last N
|
||||
## resolved samples from a bounded ring. Stale evidence ages out.
|
||||
## * TR_TMHORIZON_RESET_DROP=pp (>0): track the rolling-100 side accuracy;
|
||||
## when it falls more than `pp` below its own recent peak, treat it as
|
||||
## "the enemy changed", rebuild the heads and re-learn from the recent
|
||||
## window. The rolling-100/300 accuracy and its per-round min/mean/end are
|
||||
## always available in `roundSummary`; TR_TMHORIZON_ACCURVE=1 logs a
|
||||
## `[tmh-acc]` curve line every 25 warm samples so the decay is visible.
|
||||
## * TR_TMHORIZON_NSTATES=K: runtime automata state count (inertia), so the
|
||||
## fast/slow trade-off can be swept without a rebuild. Default = the
|
||||
## compile-time TMH_NSTATES.
|
||||
## Offline measurement (`tests/measure_tm_readapt.nim`, prequential side
|
||||
## accuracy on the DrussGT fixtures): on `tr_drussgt_vs_modularbot` the
|
||||
## keep-everything arm sits at ~76% late accuracy, the sliding window at ~85%,
|
||||
## the change-detection re-learn at ~84%, and a shuffled-label control stays at
|
||||
## ~50%. Lowering the state count from 64 to 8 lifts the keep-everything arm to
|
||||
## ~83% but leaves the window flat — i.e. forgetting and low inertia are
|
||||
## substitutes, and forgetting is the stronger one.
|
||||
##
|
||||
## Coordinate system: 0° = East, CCW positive (Tank Royale standard).
|
||||
|
||||
import std/[math, os, strutils, strformat]
|
||||
@@ -101,6 +126,34 @@ const
|
||||
TMH_SHIFT_DEFAULT* = 2.0
|
||||
TMH_BIG_MULT_DEFAULT* = 1.5
|
||||
TMH_RESET_ON_TARGET_DEFAULT* = true
|
||||
## ── re-adaptation knobs (all default to today's behaviour) ─────────────
|
||||
## Sliding-window / forgetting: >0 rebuilds both heads on the most recent N
|
||||
## resolved training samples every `TMH_RETRAIN_EVERY` new samples, so stale
|
||||
## evidence ages out. 0 = keep everything (current shipped behaviour).
|
||||
TMH_WINDOW_ENV* = "TR_TMHORIZON_WINDOW"
|
||||
## Change detection: when the rolling side accuracy falls more than this many
|
||||
## percentage points below its own recent peak, treat it as "the enemy
|
||||
## changed" and rebuild + re-learn from the recent window. 0 = off.
|
||||
TMH_RESET_DROP_ENV* = "TR_TMHORIZON_RESET_DROP"
|
||||
## Runtime automata state count (inertia). Default = the compile-time value.
|
||||
TMH_NSTATES_ENV* = "TR_TMHORIZON_NSTATES"
|
||||
## 1 = log the rolling accuracy every `TMH_ACC_LOG_EVERY` warm samples.
|
||||
TMH_ACCURVE_ENV* = "TR_TMHORIZON_ACCURVE"
|
||||
## How often the sliding/window mode does its full retrain, and how many
|
||||
## epochs that retrain runs. Defaults keep the retrain cheap.
|
||||
TMH_RETRAIN_EVERY_ENV* = "TR_TMHORIZON_RETRAIN_EVERY"
|
||||
TMH_EPOCHS_ENV* = "TR_TMHORIZON_EPOCHS"
|
||||
## ── fixed instrumentation geometry ──────────────────────────────────────
|
||||
TMH_ACC_CAP* = 512 ## rolling-accuracy ring (> the 300 window)
|
||||
TMH_ACC_WIN1* = 100 ## fast rolling window (samples)
|
||||
TMH_ACC_WIN2* = 300 ## slow rolling window (samples)
|
||||
TMH_ACC_LOG_EVERY* = 25 ## ACCURVE: one line every N warm samples
|
||||
TMH_RETRAIN_EVERY_DEF* = 50
|
||||
TMH_EPOCHS_DEF* = 1
|
||||
## When only RESET_DROP is on (WINDOW == 0) the re-learn uses this window.
|
||||
TMH_RESET_WINDOW_DEF* = 200
|
||||
TMH_RESET_MIN_SAMPLES* = 100 ## never judge a drop before this many samples
|
||||
TMH_RESET_COOLDOWN* = 100 ## samples between two change-triggered resets
|
||||
|
||||
type
|
||||
TmhTick* = object
|
||||
@@ -117,6 +170,13 @@ type
|
||||
sx, sy: float
|
||||
active: bool
|
||||
|
||||
TmhSample = object
|
||||
## One resolved training sample kept in the bounded ring that feeds the
|
||||
## sliding-window rebuild and the change-detection re-learn.
|
||||
lits: array[TMH_NLITS, uint8]
|
||||
side: uint8
|
||||
mag: uint8
|
||||
|
||||
TmhPending* = object
|
||||
## One deferred training sample. `lits` is the exact literal vector the TM
|
||||
## saw at fire time; the label is resolved h ticks later.
|
||||
@@ -188,6 +248,29 @@ type
|
||||
# magnitude-median histogram (TMH_ABS_RES deg bins)
|
||||
absHist: array[TMH_ABS_BINS, int]
|
||||
absCount: int
|
||||
# ── re-adaptation: knobs + bounded sample ring ──────────────────────────
|
||||
windowN*: int ## TR_TMHORIZON_WINDOW (0 = keep everything)
|
||||
resetDrop*: float ## TR_TMHORIZON_RESET_DROP (0 = off), percentage points
|
||||
retrainEvery*: int
|
||||
retrainEpochs*: int
|
||||
nStates*: int ## runtime automata state count (inertia)
|
||||
buffer: seq[TmhSample]
|
||||
bufCap: int
|
||||
bufCount: int ## total resolved samples ever appended
|
||||
sinceRetrain: int
|
||||
# ── rolling side accuracy (change detection + curve) ────────────────────
|
||||
accRing: array[TMH_ACC_CAP, uint8]
|
||||
accPos: int
|
||||
accCount*: int ## total warm samples recorded in the ring
|
||||
accPeak*: float ## peak rolling-100 accuracy (percent)
|
||||
sinceResetDrop: int
|
||||
resetDrops*: int ## change-triggered re-learns
|
||||
accurveEnabled*: bool
|
||||
accSnapMin*: float
|
||||
accSnapMax*: float
|
||||
accSnapSum*: float
|
||||
accSnapN*: int
|
||||
lastAccLogCount: int
|
||||
# ── instrumentation ─────────────────────────────────────────────────────
|
||||
trained*: int
|
||||
sideCorrect*, sideTotal*: int
|
||||
@@ -226,6 +309,12 @@ proc envBoolT(name: string, default: bool): bool =
|
||||
of "0", "false", "no", "off": false
|
||||
else: default
|
||||
|
||||
proc envIntT(name: string, default: int): int =
|
||||
let v = getEnv(name, "")
|
||||
if v.len == 0: return default
|
||||
try: parseInt(v.strip())
|
||||
except ValueError: default
|
||||
|
||||
# ── horizon maths (pure, unit-tested) ────────────────────────────────────────
|
||||
|
||||
proc tmhHorizonFor*(dist, bulletSpeed: float): int =
|
||||
@@ -275,11 +364,30 @@ proc tmhLits*(base: array[TMH_N_BASE, uint8],
|
||||
|
||||
# ── construction / reset ─────────────────────────────────────────────────────
|
||||
|
||||
proc ensureBuffer(g: var TmHorizonGun) =
|
||||
## Grow the resolved-sample ring to hold whichever window the configured
|
||||
## forgetting/re-learn mode needs. Never shrinks, so switching knobs at
|
||||
## runtime cannot lose already-buffered samples.
|
||||
let want = max(g.windowN,
|
||||
(if g.resetDrop > 0.0: TMH_RESET_WINDOW_DEF else: 0))
|
||||
if want > g.bufCap:
|
||||
g.bufCap = want
|
||||
g.buffer.setLen(want)
|
||||
|
||||
proc initTmHorizonGun*(): TmHorizonGun =
|
||||
result.pattern = PatternMatcherGun()
|
||||
result.sideMachine = newMachine(TMH_N_BITS, 2, TMH_NCLAUSES, TMH_NSTATES,
|
||||
# Runtime inertia: read once here so the state count can be swept with an env
|
||||
# knob instead of a rebuild. Default is the compile-time TMH_NSTATES.
|
||||
result.nStates = clamp(envIntT(TMH_NSTATES_ENV, TMH_NSTATES), 2, 4096)
|
||||
result.windowN = max(0, envIntT(TMH_WINDOW_ENV, 0))
|
||||
result.resetDrop = max(0.0, envFloatT(TMH_RESET_DROP_ENV, 0.0))
|
||||
result.retrainEvery = max(1, envIntT(TMH_RETRAIN_EVERY_ENV, TMH_RETRAIN_EVERY_DEF))
|
||||
result.retrainEpochs = max(1, envIntT(TMH_EPOCHS_ENV, TMH_EPOCHS_DEF))
|
||||
result.accurveEnabled = envBoolT(TMH_ACCURVE_ENV, false)
|
||||
result.ensureBuffer()
|
||||
result.sideMachine = newMachine(TMH_N_BITS, 2, TMH_NCLAUSES, result.nStates,
|
||||
TMH_S, seed = 1)
|
||||
result.magMachine = newMachine(TMH_N_BITS, 2, TMH_NCLAUSES, TMH_NSTATES,
|
||||
result.magMachine = newMachine(TMH_N_BITS, 2, TMH_NCLAUSES, result.nStates,
|
||||
TMH_S, seed = 2)
|
||||
for c in 0..1:
|
||||
result.sideScratch[c] = newSeq[uint8](TMH_NCLAUSES)
|
||||
@@ -317,6 +425,23 @@ proc setResetOnTarget*(g: var TmHorizonGun, on: bool) =
|
||||
g.resetOnTarget = on
|
||||
g.shiftConfigured = true
|
||||
|
||||
proc setWindow*(g: var TmHorizonGun, n: int) =
|
||||
## Explicit per-gun override (tests / offline sweeps): sliding-window width.
|
||||
g.windowN = max(0, n)
|
||||
g.ensureBuffer()
|
||||
|
||||
proc setResetDrop*(g: var TmHorizonGun, pp: float) =
|
||||
## Explicit per-gun override (tests): change-detection threshold in pp.
|
||||
g.resetDrop = max(0.0, pp)
|
||||
g.ensureBuffer()
|
||||
|
||||
proc setRetrainConfig*(g: var TmHorizonGun, every, epochs: int) =
|
||||
g.retrainEvery = max(1, every)
|
||||
g.retrainEpochs = max(1, epochs)
|
||||
|
||||
proc setAccurve*(g: var TmHorizonGun, on: bool) =
|
||||
g.accurveEnabled = on
|
||||
|
||||
proc resetRoundState*(g: var TmHorizonGun) =
|
||||
## PER-ROUND wipe ONLY. Clears the observation ring, the deferred-label queue
|
||||
## and every piece of motion / bullet / per-tick history that is meaningless
|
||||
@@ -344,6 +469,12 @@ proc resetRoundState*(g: var TmHorizonGun) =
|
||||
# A fresh round starts here: remember the cumulative count so the per-round
|
||||
# summary can report `thisRound` while `trained` keeps climbing.
|
||||
g.roundStartTrained = g.trained
|
||||
# Per-round accuracy-curve snapshots (the rolling ring itself SURVIVES the
|
||||
# boundary: change detection must see the whole battle).
|
||||
g.accSnapMin = 0.0
|
||||
g.accSnapMax = 0.0
|
||||
g.accSnapSum = 0.0
|
||||
g.accSnapN = 0
|
||||
|
||||
proc resetLearning*(g: var TmHorizonGun, reason = "") =
|
||||
## PER-BATTLE / PER-ENEMY wipe. Wipes the Tsetlin machines and every learned
|
||||
@@ -369,6 +500,16 @@ proc resetLearning*(g: var TmHorizonGun, reason = "") =
|
||||
for i in 0..<4: g.quadHist[i] = 0
|
||||
for i in 0..<TMH_ABS_BINS: g.absHist[i] = 0
|
||||
g.absCount = 0
|
||||
# Re-adaptation state: the bounded sample ring and the rolling-accuracy ring
|
||||
# are part of "learning", so a new battle/enemy wipes them.
|
||||
g.bufCount = 0
|
||||
g.sinceRetrain = 0
|
||||
g.accCount = 0
|
||||
g.accPos = 0
|
||||
g.accPeak = 0.0
|
||||
g.sinceResetDrop = 0
|
||||
g.resetDrops = 0
|
||||
g.lastAccLogCount = 0
|
||||
g.resetRoundState()
|
||||
if reason.len > 0:
|
||||
g.ensureConfig()
|
||||
@@ -416,6 +557,10 @@ proc sideClauseStates*(g: TmHorizonGun): seq[int16] =
|
||||
## states, so a test can prove the machines survive or are wiped by a reset.
|
||||
g.sideMachine.teams[0] & g.sideMachine.teams[1]
|
||||
|
||||
proc headStates*(g: TmHorizonGun): int =
|
||||
## Test seam: the runtime automata state count both heads were built with.
|
||||
g.sideMachine.nStates
|
||||
|
||||
proc sideClausesAllExclude*(g: TmHorizonGun): bool =
|
||||
## Test seam: true when every side clause is at the Exclude boundary.
|
||||
for c in 0..<g.sideMachine.nClasses:
|
||||
@@ -651,6 +796,80 @@ proc tmhTrainOne(m: var TmMachine, lits: array[TMH_NLITS, uint8], label: int,
|
||||
let d = if c == label: 1.0 else: -1.0
|
||||
tmLearnDir(m, m.teams[c], lits, caches[c], votes[c], d)
|
||||
|
||||
# ── sliding window / change detection ────────────────────────────────────────
|
||||
|
||||
proc bufferedCount*(g: TmHorizonGun): int =
|
||||
## Test/observability seam: how many resolved samples the bounded ring has
|
||||
## seen since the last battle/enemy reset.
|
||||
g.bufCount
|
||||
|
||||
proc bufferCapacity*(g: TmHorizonGun): int =
|
||||
## Test/observability seam: the allocated ring capacity (0 = buffering off).
|
||||
g.bufCap
|
||||
|
||||
proc rollingAcc*(g: TmHorizonGun, n: int): float =
|
||||
## Side accuracy over the last `min(n, accCount)` recorded warm samples.
|
||||
## This is the measurement that shows whether accuracy decays when the
|
||||
## enemy changes. Returns 0.0 when nothing has been recorded yet.
|
||||
if g.accCount == 0: return 0.0
|
||||
let k = min(n, g.accCount)
|
||||
if k <= 0: return 0.0
|
||||
var cor = 0
|
||||
for i in 0..<k:
|
||||
let idx = ((g.accPos - 1 - i) mod TMH_ACC_CAP + TMH_ACC_CAP) mod TMH_ACC_CAP
|
||||
cor += int(g.accRing[idx])
|
||||
cor.float / k.float
|
||||
|
||||
proc retrainFromBuffer*(g: var TmHorizonGun, n: int) =
|
||||
## FULL rebuild of both heads from scratch, then train on the most recent
|
||||
## `min(n, buffered)` resolved samples (oldest -> newest). `trained` is NOT
|
||||
## reset: it counts resolved samples over the battle, so the cold gate and the
|
||||
## per-round summary stay monotone while the MACHINES forget. The rebuild is
|
||||
## DETERMINISTIC given the buffer (fixed seed), so an A/B rerun is repeatable.
|
||||
g.sideMachine.resetMachine(seed = 1000'u64)
|
||||
g.magMachine.resetMachine(seed = 2000'u64)
|
||||
if g.bufCap <= 0: return
|
||||
let k = min(n, min(g.bufCount, g.bufCap))
|
||||
if k <= 0: return
|
||||
for _ in 0..<g.retrainEpochs:
|
||||
for j in 0..<k:
|
||||
let idx = ((g.bufCount - k + j) mod g.bufCap + g.bufCap) mod g.bufCap
|
||||
let s = g.buffer[idx]
|
||||
tmhTrainOne(g.sideMachine, s.lits, int(s.side), g.sideScratch)
|
||||
tmhTrainOne(g.magMachine, s.lits, int(s.mag), g.magScratch)
|
||||
|
||||
proc recordWarm(g: var TmHorizonGun, correct: bool) =
|
||||
## Record one warm sample's side correctness, update the rolling windows and
|
||||
## the change-detection peak, and fire the re-learn trigger when the rolling
|
||||
## accuracy has fallen far enough below its own recent peak.
|
||||
g.accRing[g.accPos] = (if correct: 1'u8 else: 0'u8)
|
||||
g.accPos = (g.accPos + 1) mod TMH_ACC_CAP
|
||||
inc g.accCount
|
||||
let a100 = g.rollingAcc(TMH_ACC_WIN1) * 100.0
|
||||
if g.accSnapN == 0:
|
||||
g.accSnapMin = a100
|
||||
g.accSnapMax = a100
|
||||
else:
|
||||
if a100 < g.accSnapMin: g.accSnapMin = a100
|
||||
if a100 > g.accSnapMax: g.accSnapMax = a100
|
||||
g.accSnapSum += a100
|
||||
inc g.accSnapN
|
||||
if a100 > g.accPeak: g.accPeak = a100
|
||||
inc g.sinceResetDrop
|
||||
if g.resetDrop > 0.0 and g.accCount >= TMH_RESET_MIN_SAMPLES and
|
||||
(g.accPeak - a100) > g.resetDrop and g.sinceResetDrop >= TMH_RESET_COOLDOWN:
|
||||
let n = if g.windowN > 0: g.windowN else: TMH_RESET_WINDOW_DEF
|
||||
g.retrainFromBuffer(n)
|
||||
g.accPeak = a100
|
||||
g.sinceResetDrop = 0
|
||||
inc g.resetDrops
|
||||
if g.accurveEnabled and (g.accCount mod TMH_ACC_LOG_EVERY == 0) and
|
||||
g.accCount != g.lastAccLogCount:
|
||||
g.lastAccLogCount = g.accCount
|
||||
let a300 = g.rollingAcc(TMH_ACC_WIN2) * 100.0
|
||||
echo fmt"[tmh-acc] trained={g.trained} warm={g.accCount} acc100={a100:.1f} " &
|
||||
fmt"acc300={a300:.1f} peak={g.accPeak:.1f} resets={g.resetDrops}"
|
||||
|
||||
# ── magnitude median (running histogram) ─────────────────────────────────────
|
||||
|
||||
proc recordAbs(g: var TmHorizonGun, a: float) =
|
||||
@@ -691,12 +910,26 @@ proc tmhResolveOne(g: var TmHorizonGun, state: WorldState,
|
||||
let magLabel = if absE > med: 1 else: 0
|
||||
if p.warm:
|
||||
inc g.sideTotal
|
||||
if p.sidePred == sideLabel: inc g.sideCorrect
|
||||
let correct = p.sidePred == sideLabel
|
||||
if correct: inc g.sideCorrect
|
||||
# Rolling accuracy + change detection (the re-adaptation measurement).
|
||||
g.recordWarm(correct)
|
||||
inc g.sideLabelHist[sideLabel]
|
||||
inc g.magLabelHist[magLabel]
|
||||
tmhTrainOne(g.sideMachine, p.lits, sideLabel, g.sideScratch)
|
||||
tmhTrainOne(g.magMachine, p.lits, magLabel, g.magScratch)
|
||||
inc g.trained
|
||||
# Bounded ring feeding the sliding-window rebuild / change-detection
|
||||
# re-learn. Appended AFTER the online update so a retrain triggered this
|
||||
# tick sees this sample too.
|
||||
if g.bufCap > 0:
|
||||
g.buffer[g.bufCount mod g.bufCap] =
|
||||
TmhSample(lits: p.lits, side: uint8(sideLabel), mag: uint8(magLabel))
|
||||
inc g.bufCount
|
||||
inc g.sinceRetrain
|
||||
if g.windowN > 0 and g.sinceRetrain >= g.retrainEvery:
|
||||
g.sinceRetrain = 0
|
||||
g.retrainFromBuffer(g.windowN)
|
||||
g.recordAbs(absE)
|
||||
true
|
||||
|
||||
@@ -755,16 +988,25 @@ proc tmhLog(g: var TmHorizonGun, state: WorldState, bulletSpeed: float,
|
||||
fmt"aim={aimDeg:.1f} gun=Pattern conf={sideConf:.2f} trained={g.trained}"
|
||||
|
||||
proc roundSummary*(g: var TmHorizonGun) =
|
||||
## Per-round summary on round end (behind the same log switch), so the user can
|
||||
## watch it learn across the round.
|
||||
## Per-round summary on round end (behind `TR_TMHORIZON_LOG=1`, or
|
||||
## `TR_TMHORIZON_ACCURVE=1` for the accuracy curve alone), so the user can
|
||||
## watch it learn across the round. The rolling windows are the re-adaptation
|
||||
## measurement: `acc100` = last 100 warm samples, `acc300` = last 300, and
|
||||
## `accCurve[min/mean/end]` the per-round spread of the rolling-100 snapshot.
|
||||
g.ensureConfig()
|
||||
if not g.logEnabled: return
|
||||
if not g.logEnabled and not g.accurveEnabled: return
|
||||
let thisRound = g.trained - g.roundStartTrained
|
||||
let acc = if g.sideTotal > 0: g.sideCorrect.float / g.sideTotal.float * 100.0
|
||||
else: 0.0
|
||||
let a100 = g.rollingAcc(TMH_ACC_WIN1) * 100.0
|
||||
let a300 = g.rollingAcc(TMH_ACC_WIN2) * 100.0
|
||||
let snapMean = if g.accSnapN > 0: g.accSnapSum / float(g.accSnapN) else: 0.0
|
||||
echo fmt"[tmh-round] trained={g.trained} thisRound={thisRound} pending={g.pendingCount} " &
|
||||
fmt"dropped={g.pendingDropped} sideAcc={g.sideCorrect}/{g.sideTotal} " &
|
||||
fmt"({acc:.1f}%) " &
|
||||
fmt"acc100={a100:.1f}% acc300={a300:.1f}% " &
|
||||
fmt"accCurve[min={g.accSnapMin:.1f} mean={snapMean:.1f} end={a100:.1f} n={g.accSnapN}] " &
|
||||
fmt"resets={g.resetDrops} window={g.windowN} resetDrop={g.resetDrop:.1f} " &
|
||||
fmt"quad=[SR:{g.quadHist[0]} LR:{g.quadHist[1]} " &
|
||||
fmt"SL:{g.quadHist[2]} LL:{g.quadHist[3]}] " &
|
||||
fmt"sidePred=[R:{g.sidePredHist[0]} L:{g.sidePredHist[1]}] " &
|
||||
|
||||
@@ -0,0 +1,718 @@
|
||||
## OFFLINE re-adaptation measurement for the horizon TM (tm_horizon).
|
||||
##
|
||||
## The user's live observation: "once the pattern is learnt enough we hit
|
||||
## DrussGT, but as soon as it adapts we are not fast enough to re-adapt". This
|
||||
## harness makes that failure mode measurable WITHOUT a Java battle, using the
|
||||
## committed DrussGT fixtures (READ-ONLY) and the same fact-label pipeline as
|
||||
## `measure_tm_miss_shrink.nim` / `measure_tm_hit_optimal.nim`.
|
||||
##
|
||||
## It answers three questions:
|
||||
## 1. Does the rolling SIDE accuracy decay as the battle progresses?
|
||||
## 2. Does a candidate fix (sliding window / change-detection re-learn /
|
||||
## lower inertia) improve the LATE half, where re-adaptation matters?
|
||||
## 3. Do the fixes only "help" with SHUFFLED labels? (mandatory control)
|
||||
##
|
||||
## Metrics are PREQUENTIAL: at each streamed sample the model predicts the
|
||||
## side label BEFORE it is updated with it (exactly what the live gun's
|
||||
## `warm` accuracy counters record). Early = first half of the eval stream,
|
||||
## Late = second half. Decay = Late - Early (negative = it is losing the enemy).
|
||||
##
|
||||
## Protocols:
|
||||
## * within — train on the first `warmFrac` of a round, then stream the rest
|
||||
## of the SAME round. The nearest no-Java proxy for "train early,
|
||||
## adapt later". On real open-loop fixtures there may be little
|
||||
## enemy change, so a flat curve here is a NEGATIVE result to report.
|
||||
## * cross — warm on round R, then stream round R+1. A real distribution
|
||||
## shift (new start geometry). Directly tests "the enemy changed,
|
||||
## re-adapt now", which is what the retain-across-rounds gun faces.
|
||||
##
|
||||
## Run: nim c -r -d:release --path:common_libs \
|
||||
## common_libs/tests/measure_tm_readapt.nim
|
||||
## [--fixtures=a,b] [--window=150] [--resetdrop=5.0] [--warmfrac=0.4]
|
||||
## [--retrain=50] [--clauses=40] [--states=64] [--epochs=5] [--protocol=both]
|
||||
|
||||
import std/[json, os, strformat, strutils, math, tables]
|
||||
import tm_diag/tm_core
|
||||
|
||||
# ── configuration ────────────────────────────────────────────────────────────
|
||||
|
||||
const repoRoot* = currentSourcePath().parentDir.parentDir.parentDir
|
||||
const fixturesDir* = repoRoot / "tools" / "fixtures"
|
||||
const metaDir* = fixturesDir / "drussgt_meta"
|
||||
|
||||
const HORIZONS = [15, 20, 25, 30]
|
||||
const NH = 4
|
||||
const N_BASE = 49
|
||||
const N_BITS = N_BASE + 4
|
||||
const MIN_I = 12
|
||||
const NRBIN = 10 ## relative-eval-position bins for the curve
|
||||
const ACC_CAP = 512
|
||||
|
||||
var
|
||||
cfgClauses = 40
|
||||
cfgStates = 64
|
||||
cfgS = 3.0
|
||||
cfgWarmEpochs = 5
|
||||
cfgWindow = 150
|
||||
cfgResetDrop = 5.0
|
||||
cfgRetrainEvery = 50
|
||||
cfgWarmFrac = 0.40
|
||||
cfgResetWindow = 200
|
||||
|
||||
const PRIMARY = ["tr_drussgt_vs_modularbot.jsonl",
|
||||
"tr_drussgt_vs_modularbot_shield.jsonl"]
|
||||
|
||||
# ── small helpers ────────────────────────────────────────────────────────────
|
||||
|
||||
proc wrap180(a: float): float {.inline.} =
|
||||
var r = a
|
||||
while r > 180.0: r -= 360.0
|
||||
while r <= -180.0: r += 360.0
|
||||
r
|
||||
|
||||
proc signf(x: float): int {.inline.} =
|
||||
if x > 1e-9: 1 elif x < -1e-9: -1 else: 0
|
||||
|
||||
proc jf(d: JsonNode, k: string): float =
|
||||
let n = d[k]
|
||||
case n.kind
|
||||
of JFloat: n.getFloat
|
||||
of JInt: float(n.getInt)
|
||||
else: parseFloat(n.getStr)
|
||||
|
||||
proc pct(c, t: int): float =
|
||||
if t <= 0: return NaN
|
||||
100.0 * float(c) / float(t)
|
||||
|
||||
# ── data model ───────────────────────────────────────────────────────────────
|
||||
|
||||
type
|
||||
Tick* = object
|
||||
tick*: int
|
||||
ex*, ey*, eh*, es*, ee*: float
|
||||
sx*, sy*, sh*, ss*, se*: float
|
||||
|
||||
Rnd* = object
|
||||
roundNo*: int
|
||||
st*: seq[Tick]
|
||||
|
||||
BulletSeries* = object
|
||||
tta*: seq[int]
|
||||
lat*: seq[float]
|
||||
|
||||
Sample = object
|
||||
lits: seq[uint8]
|
||||
side: int
|
||||
|
||||
ArmKind = enum akFrozen, akAccum, akWindow, akResetDrop, akRehearse
|
||||
|
||||
ArmState = object
|
||||
m: TmMachine
|
||||
caches: seq[seq[uint8]]
|
||||
predCache: seq[uint8]
|
||||
buf: seq[seq[uint8]]
|
||||
buflab: seq[int]
|
||||
bufCap, bufCount: int
|
||||
sinceRetrain: int
|
||||
accRing: array[ACC_CAP, uint8]
|
||||
accPos, accCount: int
|
||||
accPeak: float
|
||||
sinceResetDrop: int
|
||||
resets: int
|
||||
|
||||
ArmResult = object
|
||||
earlyCor, earlyTot, lateCor, lateTot, totCor, totTot: int
|
||||
binCor, binTot: array[NRBIN, int]
|
||||
resets: int
|
||||
rollVals: seq[float] ## fine rolling curve (this run only)
|
||||
|
||||
const ArmNames: array[ArmKind, string] =
|
||||
["frozen", "accum", "window", "resetdrop", "rehearse-all"]
|
||||
|
||||
# ── fixture loading ──────────────────────────────────────────────────────────
|
||||
|
||||
proc loadTicks(path: string): seq[Tick] =
|
||||
for line in lines(path):
|
||||
let ln = line.strip()
|
||||
if ln.len == 0: continue
|
||||
let d = parseJson(ln)
|
||||
if not d.hasKey("tick"): continue
|
||||
result.add Tick(tick: d["tick"].getInt,
|
||||
ex: jf(d, "ex"), ey: jf(d, "ey"), eh: jf(d, "eh"),
|
||||
es: jf(d, "es"), ee: jf(d, "ee"),
|
||||
sx: jf(d, "sx"), sy: jf(d, "sy"), sh: jf(d, "sh"),
|
||||
ss: jf(d, "ss"), se: jf(d, "se"))
|
||||
|
||||
proc loadRounds(path: string, ticks: seq[Tick]): seq[Rnd] =
|
||||
let rp = metaDir / (extractFilename(path) & ".rounds.json")
|
||||
var spans: seq[(int, int)]
|
||||
if fileExists(rp):
|
||||
let j = parseFile(rp)
|
||||
for r in j["rounds"]:
|
||||
spans.add (r["startTick"].getInt, r["count"].getInt)
|
||||
elif ticks.len > 0:
|
||||
spans.add (ticks[0].tick, ticks.len)
|
||||
var idxByTick = initTable[int, int]()
|
||||
for i, t in ticks: idxByTick[t.tick] = i
|
||||
for sp in spans:
|
||||
let (s0, c) = sp
|
||||
if not idxByTick.hasKey(s0): continue
|
||||
let i0 = idxByTick[s0]
|
||||
var st: seq[Tick]
|
||||
for k in 0..<c:
|
||||
if i0 + k < ticks.len: st.add ticks[i0 + k]
|
||||
if st.len > 0:
|
||||
result.add Rnd(roundNo: result.len + 1, st: st)
|
||||
|
||||
proc buildBulletSeries(r: Rnd): BulletSeries =
|
||||
let L = r.st.len
|
||||
result.tta = newSeq[int](L)
|
||||
for i in 0..<L: result.tta[i] = -1
|
||||
result.lat = newSeq[float](L)
|
||||
for t0 in 1..<L:
|
||||
let drop = r.st[t0 - 1].se - r.st[t0].se
|
||||
if drop <= 0.05 or drop > 3.1: continue
|
||||
var power = drop
|
||||
if power < 0.1: power = 0.1
|
||||
if power > 3.0: power = 3.0
|
||||
let speed = 20.0 - 3.0 * power
|
||||
let rng = hypot(r.st[t0].ex - r.st[t0].sx, r.st[t0].ey - r.st[t0].sy)
|
||||
let flight = int(ceil(rng / speed))
|
||||
let dx = r.st[t0].ex - r.st[t0].sx
|
||||
let dy = r.st[t0].ey - r.st[t0].sy
|
||||
let nrm = max(1e-6, hypot(dx, dy))
|
||||
let ux = dx / nrm
|
||||
let uy = dy / nrm
|
||||
for k in 0..flight:
|
||||
let t = t0 + k
|
||||
if t >= L: break
|
||||
let ta = t0 + flight - t
|
||||
if result.tta[t] < 0 or ta < result.tta[t]:
|
||||
result.tta[t] = ta
|
||||
let vx = r.st[t].ex - r.st[t0].sx
|
||||
let vy = r.st[t].ey - r.st[t0].sy
|
||||
result.lat[t] = ux * vy - uy * vx
|
||||
|
||||
# ── 49 draft bits (causal, at tick i) — identical to the measure_tm_* pipeline ─
|
||||
|
||||
proc buildBase(r: Rnd, bi: BulletSeries, i: int,
|
||||
sinceRev: seq[int]): array[N_BASE, int] =
|
||||
let s = r.st
|
||||
let cur = s[i]
|
||||
let dL = cur.ex
|
||||
let dR = 800.0 - cur.ex
|
||||
let dT = 600.0 - cur.ey
|
||||
let dBottom = cur.ey
|
||||
let dmin = min(min(dL, dR), min(dT, dBottom))
|
||||
var wallBin = 3
|
||||
if dmin < 50.0: wallBin = 0
|
||||
elif dmin < 100.0: wallBin = 1
|
||||
elif dmin < 200.0: wallBin = 2
|
||||
result[wallBin] = 1
|
||||
var wb = 0
|
||||
let walls = [dL, dR, dT, dBottom]
|
||||
for w in 1..3:
|
||||
if walls[w] < walls[wb]: wb = w
|
||||
result[4 + wb] = 1
|
||||
|
||||
let rng = hypot(cur.ex - cur.sx, cur.ey - cur.sy)
|
||||
var ub = 5
|
||||
if rng < 100.0: ub = 0
|
||||
elif rng < 200.0: ub = 1
|
||||
elif rng < 300.0: ub = 2
|
||||
elif rng < 400.0: ub = 3
|
||||
elif rng < 600.0: ub = 4
|
||||
result[8 + ub] = 1
|
||||
|
||||
let lane = arctan2(cur.sy - cur.ey, cur.sx - cur.ex)
|
||||
let hdg = cur.eh * PI / 180.0
|
||||
let perp = abs(sin(hdg - lane))
|
||||
var hb = 1
|
||||
if perp < 0.5: hb = 2
|
||||
elif perp > 0.866: hb = 0
|
||||
result[14 + hb] = 1
|
||||
|
||||
for k in 0..2:
|
||||
if i - 1 - k >= 0:
|
||||
let d = wrap180(s[i - k].eh - s[i - 1 - k].eh)
|
||||
if d > 1e-6: result[17 + k] = 1
|
||||
|
||||
var rb = 4
|
||||
let sr = sinceRev[i]
|
||||
if sr < 5: rb = 0
|
||||
elif sr < 10: rb = 1
|
||||
elif sr < 20: rb = 2
|
||||
elif sr < 40: rb = 3
|
||||
result[20 + rb] = 1
|
||||
|
||||
var pos = 0
|
||||
var neg = 0
|
||||
for k in 0..9:
|
||||
if i - 1 - k < 0: break
|
||||
let d = wrap180(s[i - k].eh - s[i - 1 - k].eh)
|
||||
if d > 1e-6: inc pos
|
||||
elif d < -1e-6: inc neg
|
||||
let tot = pos + neg
|
||||
let cons = if tot > 0: max(pos, neg).float / tot.float else: 0.0
|
||||
var cb = 0
|
||||
if cons > 0.8: cb = 2
|
||||
elif cons >= 0.5: cb = 1
|
||||
result[25 + cb] = 1
|
||||
|
||||
let j0 = max(0, i - 10)
|
||||
let dm = hypot(cur.ex - s[j0].ex, cur.ey - s[j0].ey)
|
||||
var mb = 1
|
||||
if dm < 20.0: mb = 0
|
||||
elif dm > 50.0: mb = 2
|
||||
result[28 + mb] = 1
|
||||
|
||||
let st10 = abs(s[max(0, i - 10)].es)
|
||||
let spdDiff = abs(cur.es) - st10
|
||||
var sb = 1
|
||||
if spdDiff < -0.5: sb = 0
|
||||
elif spdDiff > 0.5: sb = 2
|
||||
result[31 + sb] = 1
|
||||
|
||||
var r1 = 0.0
|
||||
var n1 = 0
|
||||
for k in 0..4:
|
||||
if i - 1 - k >= 0:
|
||||
r1 += abs(wrap180(s[i - k].eh - s[i - 1 - k].eh)); inc n1
|
||||
var r2 = 0.0
|
||||
var n2 = 0
|
||||
for k in 5..9:
|
||||
if i - 1 - k >= 0:
|
||||
r2 += abs(wrap180(s[i - k].eh - s[i - 1 - k].eh)); inc n2
|
||||
let m1 = if n1 > 0: r1 / float(n1) else: 0.0
|
||||
let m2 = if n2 > 0: r2 / float(n2) else: 0.0
|
||||
let dtr = m1 - m2
|
||||
var tb = 1
|
||||
if dtr < -0.3: tb = 0
|
||||
elif dtr > 0.3: tb = 2
|
||||
result[34 + tb] = 1
|
||||
|
||||
let tta = bi.tta[i]
|
||||
var b1 = 0
|
||||
if tta >= 0:
|
||||
if tta < 5: b1 = 1
|
||||
elif tta < 10: b1 = 2
|
||||
elif tta < 20: b1 = 3
|
||||
else: b1 = 4
|
||||
result[37 + b1] = 1
|
||||
|
||||
let lat = bi.lat[i]
|
||||
var lb = 3
|
||||
if lat < -72.0: lb = 0
|
||||
elif lat < -36.0: lb = 1
|
||||
elif lat < -18.0: lb = 2
|
||||
elif lat <= 18.0: lb = 3
|
||||
elif lat <= 36.0: lb = 4
|
||||
elif lat <= 72.0: lb = 5
|
||||
else: lb = 6
|
||||
result[42 + lb] = 1
|
||||
|
||||
proc toLits(base: array[N_BASE, int], h: int): seq[uint8] =
|
||||
var raw: array[N_BITS, int]
|
||||
for i in 0..<N_BASE: raw[i] = base[i]
|
||||
for k, hh in HORIZONS:
|
||||
if hh == h: raw[N_BASE + k] = 1
|
||||
result = newSeq[uint8](2 * N_BITS)
|
||||
for i in 0..<N_BITS:
|
||||
let v = uint8(if raw[i] != 0: 1 else: 0)
|
||||
result[i] = v
|
||||
result[i + N_BITS] = 1'u8 - v
|
||||
|
||||
proc sinceRevSeries(r: Rnd): seq[int] =
|
||||
let L = r.st.len
|
||||
result = newSeq[int](L)
|
||||
var lastFlip = -1
|
||||
var prevSg = 0
|
||||
for i in 0..<L:
|
||||
let sg = signf(r.st[i].es)
|
||||
if sg != 0:
|
||||
if prevSg != 0 and sg != prevSg: lastFlip = i
|
||||
prevSg = sg
|
||||
result[i] = if lastFlip < 0: i + 1000 else: i - lastFlip
|
||||
|
||||
# ── sample collection (fact labels, never across a round boundary) ───────────
|
||||
|
||||
proc collectSamples(r: Rnd, bi: BulletSeries,
|
||||
base: seq[array[N_BASE, int]], hidx: int): seq[Sample] =
|
||||
let s = r.st
|
||||
let L = s.len
|
||||
let h = HORIZONS[hidx]
|
||||
for i in MIN_I..<L:
|
||||
let j = i + h
|
||||
if j >= L: continue
|
||||
let cur = s[i]
|
||||
let gx = cur.ex + cur.es * cos(cur.eh * PI / 180.0) * float(h)
|
||||
let gy = cur.ey + cur.es * sin(cur.eh * PI / 180.0) * float(h)
|
||||
let ba = arctan2(s[j].ey - cur.sy, s[j].ex - cur.sx)
|
||||
let bg = arctan2(gy - cur.sy, gx - cur.sx)
|
||||
let err = radToDeg(arctan2(sin(ba - bg), cos(ba - bg)))
|
||||
if abs(err) < 1e-9: continue
|
||||
result.add Sample(lits: toLits(base[i], h),
|
||||
side: (if err > 0.0: 1 else: 0))
|
||||
|
||||
proc buildRoundSamples(r: Rnd): seq[seq[Sample]] =
|
||||
## samples[hidx] for one round.
|
||||
let bi = buildBulletSeries(r)
|
||||
let L = r.st.len
|
||||
var base = newSeq[array[N_BASE, int]](L)
|
||||
let sr = sinceRevSeries(r)
|
||||
for i in 0..<L: base[i] = buildBase(r, bi, i, sr)
|
||||
result = newSeq[seq[Sample]](NH)
|
||||
for hi in 0..<NH:
|
||||
result[hi] = collectSamples(r, bi, base, hi)
|
||||
|
||||
# ── arm mechanics ────────────────────────────────────────────────────────────
|
||||
|
||||
proc trainOne(m: var TmMachine, lits: openArray[uint8], label: int,
|
||||
caches: var seq[seq[uint8]]) =
|
||||
var votes: array[2, float]
|
||||
for c in 0..1:
|
||||
votes[c] = tmForward(m, m.teams[c], lits, caches[c])
|
||||
for c in 0..1:
|
||||
let d = if c == label: 1.0 else: -1.0
|
||||
tmLearnDir(m, m.teams[c], lits, caches[c], votes[c], d)
|
||||
|
||||
proc predictOne(m: TmMachine, lits: openArray[uint8],
|
||||
cache: var seq[uint8]): int =
|
||||
let v0 = tmForward(m, m.teams[0], lits, cache)
|
||||
let v1 = tmForward(m, m.teams[1], lits, cache)
|
||||
if v1 > v0: 1 else: 0
|
||||
|
||||
proc newArm(cap: int): ArmState =
|
||||
result.m = newMachine(N_BITS, 2, cfgClauses, cfgStates, cfgS, seed = 12345)
|
||||
result.caches = newSeq[seq[uint8]](2)
|
||||
for c in 0..1: result.caches[c] = newSeq[uint8](cfgClauses)
|
||||
result.predCache = newSeq[uint8](cfgClauses)
|
||||
result.bufCap = max(1, cap)
|
||||
result.buf = newSeq[seq[uint8]](result.bufCap)
|
||||
result.buflab = newSeq[int](result.bufCap)
|
||||
|
||||
proc trainEpochs(arm: var ArmState, litsList: seq[seq[uint8]],
|
||||
lab: seq[int], lo, hi, epochs: int) =
|
||||
if hi <= lo: return
|
||||
var order = newSeq[int](hi - lo)
|
||||
for e in 0..<epochs:
|
||||
for i in 0..<order.len: order[i] = lo + i
|
||||
for i in countdown(order.len - 1, 1):
|
||||
let j = int(arm.m.rng.nextU64() mod uint64(i + 1))
|
||||
swap(order[i], order[j])
|
||||
for idx in order:
|
||||
trainOne(arm.m, litsList[idx], lab[idx], arm.caches)
|
||||
|
||||
proc pushBuffer(arm: var ArmState, lits: seq[uint8], lab: int) =
|
||||
if arm.bufCap <= 0: return
|
||||
let slot = arm.bufCount mod arm.bufCap
|
||||
arm.buf[slot] = lits
|
||||
arm.buflab[slot] = lab
|
||||
inc arm.bufCount
|
||||
|
||||
proc rebuild(arm: var ArmState, n: int) =
|
||||
arm.m.resetMachine(seed = 777'u64)
|
||||
let k = min(n, min(arm.bufCount, arm.bufCap))
|
||||
if k <= 0: return
|
||||
for j in 0..<k:
|
||||
let idx = ((arm.bufCount - k + j) mod arm.bufCap + arm.bufCap) mod arm.bufCap
|
||||
trainOne(arm.m, arm.buf[idx], arm.buflab[idx], arm.caches)
|
||||
|
||||
proc armRolling(arm: ArmState, n: int): float =
|
||||
if arm.accCount == 0: return 0.0
|
||||
let k = min(n, arm.accCount)
|
||||
var cor = 0
|
||||
for i in 0..<k:
|
||||
let idx = ((arm.accPos - 1 - i) mod ACC_CAP + ACC_CAP) mod ACC_CAP
|
||||
cor += int(arm.accRing[idx])
|
||||
cor.float / k.float
|
||||
|
||||
proc recordAcc(arm: var ArmState, correct: bool, resetDrop: float): bool =
|
||||
arm.accRing[arm.accPos] = (if correct: 1'u8 else: 0'u8)
|
||||
arm.accPos = (arm.accPos + 1) mod ACC_CAP
|
||||
inc arm.accCount
|
||||
let a100 = armRolling(arm, 100) * 100.0
|
||||
if a100 > arm.accPeak: arm.accPeak = a100
|
||||
inc arm.sinceResetDrop
|
||||
if resetDrop > 0.0 and arm.accCount >= 100 and
|
||||
(arm.accPeak - a100) > resetDrop and arm.sinceResetDrop >= 100:
|
||||
arm.accPeak = a100
|
||||
arm.sinceResetDrop = 0
|
||||
inc arm.resets
|
||||
return true
|
||||
false
|
||||
|
||||
proc runArm(kind: ArmKind, warmLits: seq[seq[uint8]], warmLab: seq[int],
|
||||
streamLits: seq[seq[uint8]], streamLab: seq[int],
|
||||
emitCurve: bool): ArmResult =
|
||||
let cap =
|
||||
case kind
|
||||
of akRehearse: 1400
|
||||
of akWindow: max(1, cfgWindow)
|
||||
of akResetDrop: max(cfgWindow, cfgResetWindow)
|
||||
else: 1
|
||||
var arm = newArm(cap)
|
||||
# ── warmup ──
|
||||
case kind
|
||||
of akWindow:
|
||||
let lo = max(0, warmLits.len - cfgWindow)
|
||||
for i in lo..<warmLits.len: arm.pushBuffer(warmLits[i], warmLab[i])
|
||||
arm.trainEpochs(warmLits, warmLab, lo, warmLits.len, cfgWarmEpochs)
|
||||
of akFrozen, akAccum, akResetDrop, akRehearse:
|
||||
for i in 0..<warmLits.len: arm.pushBuffer(warmLits[i], warmLab[i])
|
||||
arm.trainEpochs(warmLits, warmLab, 0, warmLits.len, cfgWarmEpochs)
|
||||
|
||||
# ── stream ──
|
||||
let n = streamLits.len
|
||||
let half = n div 2
|
||||
for k in 0..<n:
|
||||
let lits = streamLits[k]
|
||||
let lab = streamLab[k]
|
||||
let pred = predictOne(arm.m, lits, arm.predCache)
|
||||
let correct = pred == lab
|
||||
let trigger = arm.recordAcc(correct,
|
||||
(if kind == akResetDrop: cfgResetDrop else: 0.0))
|
||||
if k < half:
|
||||
inc result.earlyTot
|
||||
if correct: inc result.earlyCor
|
||||
else:
|
||||
inc result.lateTot
|
||||
if correct: inc result.lateCor
|
||||
inc result.totTot
|
||||
if correct: inc result.totCor
|
||||
let bi = if n > 0: min(NRBIN - 1, (k * NRBIN) div n) else: 0
|
||||
inc result.binTot[bi]
|
||||
if correct: inc result.binCor[bi]
|
||||
if kind != akFrozen:
|
||||
if trigger and kind == akResetDrop:
|
||||
let rn = if cfgWindow > 0: cfgWindow else: cfgResetWindow
|
||||
arm.rebuild(rn)
|
||||
trainOne(arm.m, lits, lab, arm.caches)
|
||||
arm.pushBuffer(lits, lab)
|
||||
if kind == akWindow:
|
||||
inc arm.sinceRetrain
|
||||
if arm.sinceRetrain >= cfgRetrainEvery:
|
||||
arm.sinceRetrain = 0
|
||||
arm.rebuild(cfgWindow)
|
||||
elif kind == akRehearse:
|
||||
inc arm.sinceRetrain
|
||||
if arm.sinceRetrain >= cfgRetrainEvery:
|
||||
arm.sinceRetrain = 0
|
||||
arm.rebuild(arm.bufCount) # ALL buffered samples: no forgetting
|
||||
if emitCurve and arm.accCount >= 100 and arm.accCount mod 50 == 0:
|
||||
result.rollVals.add armRolling(arm, 100) * 100.0
|
||||
result.resets = arm.resets
|
||||
|
||||
proc addRes(dst: var ArmResult, src: ArmResult) =
|
||||
dst.earlyCor += src.earlyCor; dst.earlyTot += src.earlyTot
|
||||
dst.lateCor += src.lateCor; dst.lateTot += src.lateTot
|
||||
dst.totCor += src.totCor; dst.totTot += src.totTot
|
||||
dst.resets += src.resets
|
||||
for b in 0..<NRBIN:
|
||||
dst.binCor[b] += src.binCor[b]
|
||||
dst.binTot[b] += src.binTot[b]
|
||||
dst.rollVals.add src.rollVals
|
||||
|
||||
proc shuffleLabels(lab: seq[int], seed: uint64): seq[int] =
|
||||
result = lab
|
||||
var rng = seedRng(seed)
|
||||
for i in countdown(result.len - 1, 1):
|
||||
let j = int(rng.nextU64() mod uint64(i + 1))
|
||||
swap(result[i], result[j])
|
||||
|
||||
# ── reporting ────────────────────────────────────────────────────────────────
|
||||
|
||||
proc printArmRow(h: int, name: string, r: ArmResult) =
|
||||
if r.earlyTot + r.lateTot == 0: return
|
||||
let early = pct(r.earlyCor, r.earlyTot)
|
||||
let late = pct(r.lateCor, r.lateTot)
|
||||
let tot = pct(r.totCor, r.totTot)
|
||||
var curve = ""
|
||||
for b in 0..<NRBIN:
|
||||
curve.add &"{pct(r.binCor[b], r.binTot[b]):>5.0f}"
|
||||
echo &"{HORIZONS[h]:>3} {name:<16} {r.earlyTot + r.lateTot:>6} " &
|
||||
&"{early:>7.1f} {late:>7.1f} {late - early:>+7.1f} {tot:>7.1f} [{curve}] resets={r.resets}"
|
||||
|
||||
|
||||
# ── main ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
proc main() =
|
||||
var names = @PRIMARY
|
||||
for i in 1..paramCount():
|
||||
let a = paramStr(i)
|
||||
if a.startsWith("--fixtures="): names = a[11..^1].split(',')
|
||||
elif a.startsWith("--window="): cfgWindow = parseInt(a[9..^1])
|
||||
elif a.startsWith("--resetdrop="): cfgResetDrop = parseFloat(a[12..^1])
|
||||
elif a.startsWith("--retrain="): cfgRetrainEvery = parseInt(a[10..^1])
|
||||
elif a.startsWith("--warmfrac="): cfgWarmFrac = parseFloat(a[11..^1])
|
||||
elif a.startsWith("--clauses="): cfgClauses = parseInt(a[10..^1])
|
||||
elif a.startsWith("--states="): cfgStates = parseInt(a[9..^1])
|
||||
elif a.startsWith("--epochs="): cfgWarmEpochs = parseInt(a[9..^1])
|
||||
elif a.startsWith("--resetwindow="): cfgResetWindow = parseInt(a[14..^1])
|
||||
|
||||
echo "=" .repeat(100)
|
||||
echo "TM HORIZON RE-ADAPTATION (offline, prequential side accuracy)"
|
||||
echo "=" .repeat(100)
|
||||
echo &"fixtures : {names.join(\", \")}"
|
||||
echo &"TM : {cfgClauses} clauses, {cfgStates} states, s={cfgS}, warmEpochs={cfgWarmEpochs}"
|
||||
echo &"arms : frozen(none) | accum(keep all) | window(N={cfgWindow}, retrain {cfgRetrainEvery}) | " &
|
||||
&"resetdrop(drop<{cfgResetDrop:.1f}pp -> retrain last {cfgWindow})"
|
||||
echo &"protocol : within-round warmFrac={cfgWarmFrac:.2f}; cross-round (warm R, stream R+1)"
|
||||
echo &"metric : prequential; early=1st half of eval, late=2nd half; decay=late-early"
|
||||
echo &"control : labels shuffled within each set (mandatory integrity check)"
|
||||
echo "=" .repeat(100)
|
||||
|
||||
for name in names:
|
||||
let path = if name.endsWith(".jsonl"): fixturesDir / name
|
||||
else: fixturesDir / (name & ".jsonl")
|
||||
if not fileExists(path):
|
||||
echo &"# SKIP missing fixture {path}"
|
||||
continue
|
||||
let ticks = loadTicks(path)
|
||||
let rounds = loadRounds(path, ticks)
|
||||
var rs: seq[seq[seq[Sample]]]
|
||||
for r in rounds: rs.add buildRoundSamples(r)
|
||||
echo &"\n## FIXTURE {name}: {rounds.len} rounds"
|
||||
|
||||
for protocol in ["within", "cross"]:
|
||||
var pooled: array[ArmKind, ArmResult]
|
||||
var pooledShuf: array[ArmKind, ArmResult]
|
||||
var examples: seq[ArmResult]
|
||||
var exH = 0
|
||||
for hidx in 0..<NH:
|
||||
if protocol == "within":
|
||||
for rIdx in 0..<rounds.len:
|
||||
let sam = rs[rIdx][hidx]
|
||||
let n = sam.len
|
||||
if n < 30: continue
|
||||
let warmN = max(10, int(cfgWarmFrac * float(n)))
|
||||
if warmN >= n: continue
|
||||
var warmLits: seq[seq[uint8]]
|
||||
var warmLab: seq[int]
|
||||
var streamLits: seq[seq[uint8]]
|
||||
var streamLab: seq[int]
|
||||
for i in 0..<warmN:
|
||||
warmLits.add sam[i].lits; warmLab.add sam[i].side
|
||||
for i in warmN..<n:
|
||||
streamLits.add sam[i].lits; streamLab.add sam[i].side
|
||||
let sw = shuffleLabels(warmLab, 1000'u64 + uint64(rIdx))
|
||||
let ss = shuffleLabels(streamLab, 2000'u64 + uint64(rIdx))
|
||||
for kind in ArmKind:
|
||||
let emit = (rIdx == 0)
|
||||
let res = runArm(kind, warmLits, warmLab, streamLits, streamLab, emit)
|
||||
addRes(pooled[kind], res)
|
||||
if kind != akFrozen:
|
||||
let rs2 = runArm(kind, warmLits, sw, streamLits, ss, false)
|
||||
addRes(pooledShuf[kind], rs2)
|
||||
if emit and examples.len <= 12:
|
||||
examples.add res
|
||||
exH = hidx
|
||||
else:
|
||||
for rIdx in 0..<rounds.len - 1:
|
||||
let warmSam = rs[rIdx][hidx]
|
||||
let streamSam = rs[rIdx + 1][hidx]
|
||||
if warmSam.len < 20 or streamSam.len < 20: continue
|
||||
var warmLits: seq[seq[uint8]]
|
||||
var warmLab: seq[int]
|
||||
var streamLits: seq[seq[uint8]]
|
||||
var streamLab: seq[int]
|
||||
for smp in warmSam:
|
||||
warmLits.add smp.lits; warmLab.add smp.side
|
||||
for smp in streamSam:
|
||||
streamLits.add smp.lits; streamLab.add smp.side
|
||||
let sw = shuffleLabels(warmLab, 3000'u64 + uint64(rIdx))
|
||||
let ss = shuffleLabels(streamLab, 4000'u64 + uint64(rIdx))
|
||||
for kind in ArmKind:
|
||||
let emit = (rIdx == 0)
|
||||
let res = runArm(kind, warmLits, warmLab, streamLits, streamLab, emit)
|
||||
addRes(pooled[kind], res)
|
||||
if kind != akFrozen:
|
||||
let rs2 = runArm(kind, warmLits, sw, streamLits, ss, false)
|
||||
addRes(pooledShuf[kind], rs2)
|
||||
if emit and examples.len <= 12:
|
||||
examples.add res
|
||||
exH = hidx
|
||||
|
||||
echo &"\n## {name} protocol={protocol} (pooled over rounds/horizons)"
|
||||
echo " h arm N early% late% decay total% curve(deciles of eval) resets"
|
||||
for kind in ArmKind:
|
||||
printArmRow(0, ArmNames[kind], pooled[kind])
|
||||
if kind != akFrozen:
|
||||
printArmRow(0, ArmNames[kind] & "-shuf", pooledShuf[kind])
|
||||
|
||||
# distinct-horizon view for the key arms
|
||||
echo &" -- by horizon (accum / window / resetdrop, TRUE labels only) --"
|
||||
for hidx in 0..<NH:
|
||||
var per: array[ArmKind, ArmResult]
|
||||
if protocol == "within":
|
||||
for rIdx in 0..<rounds.len:
|
||||
let sam = rs[rIdx][hidx]
|
||||
let n = sam.len
|
||||
if n < 30: continue
|
||||
let warmN = max(10, int(cfgWarmFrac * float(n)))
|
||||
if warmN >= n: continue
|
||||
var wl: seq[seq[uint8]]
|
||||
var wla: seq[int]
|
||||
var sl: seq[seq[uint8]]
|
||||
var sla: seq[int]
|
||||
for i in 0..<warmN: wl.add sam[i].lits; wla.add sam[i].side
|
||||
for i in warmN..<n: sl.add sam[i].lits; sla.add sam[i].side
|
||||
for kind in ArmKind:
|
||||
if kind == akRehearse: continue # pooled table only (expensive)
|
||||
addRes(per[kind], runArm(kind, wl, wla, sl, sla, false))
|
||||
else:
|
||||
for rIdx in 0..<rounds.len - 1:
|
||||
let warmSam = rs[rIdx][hidx]
|
||||
let streamSam = rs[rIdx + 1][hidx]
|
||||
if warmSam.len < 20 or streamSam.len < 20: continue
|
||||
var wl: seq[seq[uint8]]
|
||||
var wla: seq[int]
|
||||
var sl: seq[seq[uint8]]
|
||||
var sla: seq[int]
|
||||
for smp in warmSam: wl.add smp.lits; wla.add smp.side
|
||||
for smp in streamSam: sl.add smp.lits; sla.add smp.side
|
||||
for kind in ArmKind:
|
||||
if kind == akRehearse: continue # pooled table only (expensive)
|
||||
addRes(per[kind], runArm(kind, wl, wla, sl, sla, false))
|
||||
for kind in ArmKind:
|
||||
printArmRow(hidx, ArmNames[kind], per[kind])
|
||||
echo ""
|
||||
|
||||
# ── inertia sweep (fix #3): does a LOWER state count help re-adaptation? ──
|
||||
if name == names[0]:
|
||||
echo &"\n## INERTIA SWEEP ({name}, within protocol, pooled over rounds x horizons)"
|
||||
echo " states arm early%/late% (TRUE labels)"
|
||||
for ns in [8, 16, 64, 256]:
|
||||
cfgStates = ns
|
||||
var ap: array[ArmKind, ArmResult]
|
||||
for hidx in 0..<NH:
|
||||
for rIdx in 0..<rounds.len:
|
||||
let sam = rs[rIdx][hidx]
|
||||
let n = sam.len
|
||||
if n < 30: continue
|
||||
let warmN = max(10, int(cfgWarmFrac * float(n)))
|
||||
if warmN >= n: continue
|
||||
var wl: seq[seq[uint8]]
|
||||
var wla: seq[int]
|
||||
var sl: seq[seq[uint8]]
|
||||
var sla: seq[int]
|
||||
for i in 0..<warmN: wl.add sam[i].lits; wla.add sam[i].side
|
||||
for i in warmN..<n: sl.add sam[i].lits; sla.add sam[i].side
|
||||
for kind in [akAccum, akWindow]:
|
||||
addRes(ap[kind], runArm(kind, wl, wla, sl, sla, false))
|
||||
for kind in [akAccum, akWindow]:
|
||||
let e = pct(ap[kind].earlyCor, ap[kind].earlyTot)
|
||||
let l = pct(ap[kind].lateCor, ap[kind].lateTot)
|
||||
echo &" {ns:>5} {ArmNames[kind]:<12} {e:>6.1f}/{l:>6.1f}"
|
||||
cfgStates = 64
|
||||
|
||||
echo "\n" & "=".repeat(100)
|
||||
echo "## INTEGRITY: shuffled-label controls should hover at chance (~50%)"
|
||||
echo " If a fix only improves the shuffled rows, it is NOT learning the enemy."
|
||||
echo " rehearse-all = same periodic full retrain as window but WITHOUT forgetting"
|
||||
echo " (isolates the effect of the sliding buffer from the effect of retraining)."
|
||||
echo &"## VERDICT INPUTS: late-half accuracy, TRUE vs shuffled (pp)"
|
||||
echo "=" .repeat(100)
|
||||
|
||||
when isMainModule:
|
||||
main()
|
||||
@@ -0,0 +1,192 @@
|
||||
TM HORIZON RE-ADAPTATION — offline measurement results
|
||||
Generated by: nim c -r -d:release --path:common_libs common_libs/tests/measure_tm_readapt.nim
|
||||
Fixtures: tr_drussgt_vs_modularbot.jsonl (15 rounds), tr_drussgt_vs_modularbot_shield.jsonl (10 rounds)
|
||||
NOTE: the primary-fixture by-horizon view omits the expensive rehearse-all arm (shown in the pooled table).
|
||||
|
||||
/home/davide/Projects/SirRoboGarage/common_libs/tests/measure_tm_readapt.nim(34, 50) Warning: imported and not used: 'algorithm' [UnusedImport]
|
||||
====================================================================================================
|
||||
TM HORIZON RE-ADAPTATION (offline, prequential side accuracy)
|
||||
====================================================================================================
|
||||
fixtures : tr_drussgt_vs_modularbot.jsonl
|
||||
TM : 40 clauses, 64 states, s=3.0, warmEpochs=5
|
||||
arms : frozen(none) | accum(keep all) | window(N=150, retrain 50) | resetdrop(drop<5.0pp -> retrain last 150)
|
||||
protocol : within-round warmFrac=0.40; cross-round (warm R, stream R+1)
|
||||
metric : prequential; early=1st half of eval, late=2nd half; decay=late-early
|
||||
control : labels shuffled within each set (mandatory integrity check)
|
||||
====================================================================================================
|
||||
|
||||
## FIXTURE tr_drussgt_vs_modularbot.jsonl: 15 rounds
|
||||
|
||||
## tr_drussgt_vs_modularbot.jsonl protocol=within (pooled over rounds/horizons)
|
||||
h arm N early% late% decay total% curve(deciles of eval) resets
|
||||
15 frozen 46842 64.9 61.4 -3.5 63.1 [ 61. 65. 63. 69. 66. 68. 58. 63. 59. 59.] resets=0
|
||||
15 accum 46842 76.4 75.3 -1.1 75.8 [ 74. 78. 73. 77. 80. 76. 71. 75. 76. 78.] resets=0
|
||||
15 accum-shuf 46842 50.9 50.7 -0.2 50.8 [ 50. 50. 51. 52. 51. 51. 51. 50. 51. 50.] resets=0
|
||||
15 window 46842 84.7 84.6 -0.0 84.6 [ 78. 86. 86. 86. 86. 85. 85. 85. 84. 84.] resets=0
|
||||
15 window-shuf 46842 50.5 51.0 +0.4 50.7 [ 50. 51. 50. 50. 52. 51. 52. 52. 50. 50.] resets=0
|
||||
15 resetdrop 46842 83.2 84.3 +1.0 83.7 [ 74. 84. 86. 87. 86. 84. 84. 86. 83. 84.] resets=301
|
||||
15 resetdrop-shuf 46842 50.1 50.6 +0.5 50.4 [ 50. 50. 50. 50. 50. 51. 51. 50. 51. 50.] resets=338
|
||||
15 rehearse-all 46842 81.2 79.5 -1.6 80.4 [ 77. 83. 81. 83. 83. 79. 77. 80. 80. 81.] resets=0
|
||||
15 rehearse-all-shuf 46842 50.6 50.7 +0.1 50.6 [ 50. 50. 51. 51. 51. 51. 51. 51. 50. 51.] resets=0
|
||||
-- by horizon (accum / window / resetdrop, TRUE labels only) --
|
||||
15 frozen 11778 63.0 59.1 -3.8 61.0 [ 61. 63. 60. 66. 65. 62. 60. 60. 56. 57.] resets=0
|
||||
15 accum 11778 74.8 73.9 -0.9 74.4 [ 73. 78. 69. 77. 77. 75. 71. 73. 73. 76.] resets=0
|
||||
15 window 11778 83.6 83.7 +0.1 83.6 [ 78. 86. 83. 85. 86. 84. 84. 84. 83. 82.] resets=0
|
||||
15 resetdrop 11778 82.2 83.1 +0.9 82.7 [ 73. 83. 84. 86. 85. 83. 83. 85. 82. 82.] resets=77
|
||||
20 frozen 11733 64.2 59.9 -4.3 62.0 [ 61. 65. 61. 68. 65. 68. 58. 60. 56. 58.] resets=0
|
||||
20 accum 11733 76.4 74.4 -2.0 75.4 [ 74. 78. 73. 77. 80. 75. 72. 75. 74. 75.] resets=0
|
||||
20 window 11733 84.7 84.3 -0.5 84.5 [ 78. 87. 87. 86. 86. 84. 86. 85. 83. 83.] resets=0
|
||||
20 resetdrop 11733 83.0 83.9 +0.8 83.4 [ 74. 83. 87. 86. 86. 84. 85. 86. 81. 84.] resets=75
|
||||
25 frozen 11688 65.0 63.4 -1.7 64.2 [ 61. 64. 66. 69. 66. 71. 57. 67. 63. 60.] resets=0
|
||||
25 accum 11688 77.1 75.8 -1.3 76.4 [ 74. 77. 75. 78. 81. 77. 71. 74. 77. 80.] resets=0
|
||||
25 window 11688 85.3 84.7 -0.6 85.0 [ 79. 87. 88. 87. 87. 86. 85. 85. 83. 85.] resets=0
|
||||
25 resetdrop 11688 83.3 84.2 +1.0 83.8 [ 74. 83. 85. 87. 87. 84. 84. 85. 84. 85.] resets=72
|
||||
30 frozen 11643 67.3 63.2 -4.0 65.2 [ 63. 69. 64. 72. 69. 70. 57. 67. 61. 61.] resets=0
|
||||
30 accum 11643 77.4 77.0 -0.4 77.2 [ 75. 80. 74. 78. 81. 77. 69. 78. 80. 81.] resets=0
|
||||
30 window 11643 85.0 85.9 +0.9 85.4 [ 78. 85. 87. 87. 87. 87. 85. 87. 85. 86.] resets=0
|
||||
30 resetdrop 11643 84.4 85.8 +1.5 85.1 [ 75. 84. 88. 88. 87. 85. 85. 88. 85. 86.] resets=77
|
||||
|
||||
|
||||
## tr_drussgt_vs_modularbot.jsonl protocol=cross (pooled over rounds/horizons)
|
||||
h arm N early% late% decay total% curve(deciles of eval) resets
|
||||
15 frozen 73164 65.4 68.8 +3.4 67.1 [ 58. 65. 67. 69. 67. 72. 70. 66. 71. 65.] resets=0
|
||||
15 accum 73164 73.7 75.6 +2.0 74.6 [ 71. 73. 76. 74. 75. 76. 77. 75. 76. 75.] resets=0
|
||||
15 accum-shuf 73164 49.6 50.6 +1.0 50.1 [ 49. 50. 50. 49. 49. 50. 51. 51. 51. 51.] resets=0
|
||||
15 window 73164 85.2 85.3 +0.2 85.3 [ 82. 87. 87. 84. 86. 86. 87. 85. 86. 84.] resets=0
|
||||
15 window-shuf 73164 50.3 50.2 -0.1 50.3 [ 50. 50. 51. 50. 50. 49. 50. 51. 50. 51.] resets=0
|
||||
15 resetdrop 73164 82.9 84.7 +1.8 83.8 [ 74. 86. 86. 83. 86. 85. 86. 84. 85. 83.] resets=466
|
||||
15 resetdrop-shuf 73164 50.5 50.7 +0.2 50.6 [ 50. 51. 51. 50. 51. 49. 51. 51. 51. 51.] resets=528
|
||||
15 rehearse-all 73164 77.5 78.7 +1.3 78.1 [ 74. 76. 80. 76. 81. 79. 80. 77. 79. 79.] resets=0
|
||||
15 rehearse-all-shuf 73164 50.3 50.7 +0.4 50.5 [ 49. 51. 51. 51. 50. 50. 51. 51. 51. 51.] resets=0
|
||||
-- by horizon (accum / window / resetdrop, TRUE labels only) --
|
||||
15 frozen 18396 64.2 66.5 +2.3 65.3 [ 56. 62. 67. 68. 68. 70. 68. 66. 68. 60.] resets=0
|
||||
15 accum 18396 72.0 74.0 +2.0 73.0 [ 68. 70. 74. 73. 75. 73. 76. 75. 74. 71.] resets=0
|
||||
15 window 18396 84.0 84.4 +0.4 84.2 [ 81. 85. 85. 84. 85. 85. 86. 84. 85. 82.] resets=0
|
||||
15 resetdrop 18396 81.5 83.4 +2.0 82.4 [ 71. 84. 84. 83. 85. 84. 86. 84. 83. 81.] resets=121
|
||||
20 frozen 18326 64.4 68.8 +4.4 66.6 [ 56. 62. 67. 69. 67. 72. 69. 67. 71. 64.] resets=0
|
||||
20 accum 18326 73.5 75.8 +2.3 74.6 [ 69. 72. 75. 76. 76. 76. 76. 78. 76. 73.] resets=0
|
||||
20 window 18326 85.1 85.5 +0.3 85.3 [ 82. 85. 87. 85. 87. 87. 86. 85. 86. 83.] resets=0
|
||||
20 resetdrop 18326 82.7 84.7 +2.0 83.7 [ 73. 85. 86. 84. 86. 86. 86. 84. 85. 83.] resets=116
|
||||
25 frozen 18256 66.7 69.8 +3.1 68.3 [ 60. 66. 68. 71. 68. 73. 71. 66. 74. 66.] resets=0
|
||||
25 accum 18256 73.9 76.1 +2.2 75.0 [ 71. 74. 75. 74. 75. 77. 77. 73. 76. 77.] resets=0
|
||||
25 window 18256 85.3 85.6 +0.3 85.5 [ 83. 87. 88. 84. 85. 85. 87. 86. 85. 84.] resets=0
|
||||
25 resetdrop 18256 83.6 85.3 +1.7 84.4 [ 74. 86. 87. 84. 87. 85. 87. 84. 85. 84.] resets=115
|
||||
30 frozen 18186 66.2 70.0 +3.8 68.1 [ 61. 68. 67. 69. 66. 72. 71. 66. 72. 69.] resets=0
|
||||
30 accum 18186 75.2 76.7 +1.5 75.9 [ 74. 76. 78. 74. 74. 78. 77. 73. 78. 78.] resets=0
|
||||
30 window 18186 86.3 86.0 -0.3 86.1 [ 84. 89. 88. 84. 87. 86. 87. 84. 87. 86.] resets=0
|
||||
30 resetdrop 18186 84.0 85.4 +1.4 84.7 [ 77. 88. 86. 83. 86. 86. 86. 84. 87. 85.] resets=114
|
||||
|
||||
|
||||
## INERTIA SWEEP (tr_drussgt_vs_modularbot.jsonl, within protocol, pooled over rounds x horizons)
|
||||
states arm early%/late% (TRUE labels)
|
||||
8 accum 84.3/ 83.3
|
||||
8 window 85.3/ 84.7
|
||||
16 accum 79.5/ 79.3
|
||||
16 window 84.6/ 84.6
|
||||
64 accum 76.4/ 75.3
|
||||
64 window 84.7/ 84.6
|
||||
256 accum 76.4/ 75.3
|
||||
256 window 84.7/ 84.6
|
||||
|
||||
====================================================================================================
|
||||
## INTEGRITY: shuffled-label controls should hover at chance (~50%)
|
||||
If a fix only improves the shuffled rows, it is NOT learning the enemy.
|
||||
rehearse-all = same periodic full retrain as window but WITHOUT forgetting
|
||||
(isolates the effect of the sliding buffer from the effect of retraining).
|
||||
## VERDICT INPUTS: late-half accuracy, TRUE vs shuffled (pp)
|
||||
====================================================================================================
|
||||
|
||||
########################################################################
|
||||
|
||||
/home/davide/Projects/SirRoboGarage/common_libs/tests/measure_tm_readapt.nim(34, 50) Warning: imported and not used: 'algorithm' [UnusedImport]
|
||||
====================================================================================================
|
||||
TM HORIZON RE-ADAPTATION (offline, prequential side accuracy)
|
||||
====================================================================================================
|
||||
fixtures : tr_drussgt_vs_modularbot_shield.jsonl
|
||||
TM : 40 clauses, 64 states, s=3.0, warmEpochs=5
|
||||
arms : frozen(none) | accum(keep all) | window(N=150, retrain 50) | resetdrop(drop<5.0pp -> retrain last 150)
|
||||
protocol : within-round warmFrac=0.40; cross-round (warm R, stream R+1)
|
||||
metric : prequential; early=1st half of eval, late=2nd half; decay=late-early
|
||||
control : labels shuffled within each set (mandatory integrity check)
|
||||
====================================================================================================
|
||||
|
||||
## FIXTURE tr_drussgt_vs_modularbot_shield.jsonl: 10 rounds
|
||||
|
||||
## tr_drussgt_vs_modularbot_shield.jsonl protocol=within (pooled over rounds/horizons)
|
||||
h arm N early% late% decay total% curve(deciles of eval) resets
|
||||
15 frozen 28614 64.4 63.0 -1.5 63.7 [ 67. 58. 65. 72. 59. 72. 67. 57. 61. 58.] resets=0
|
||||
15 accum 28614 75.3 75.9 +0.7 75.6 [ 74. 71. 75. 79. 78. 78. 75. 68. 79. 79.] resets=0
|
||||
15 accum-shuf 28614 50.8 50.9 +0.1 50.8 [ 51. 51. 51. 50. 52. 50. 52. 51. 50. 51.] resets=0
|
||||
15 window 28614 85.5 86.2 +0.7 85.9 [ 81. 87. 87. 85. 88. 86. 86. 87. 86. 87.] resets=0
|
||||
15 window-shuf 28614 50.9 50.6 -0.3 50.7 [ 51. 51. 50. 51. 51. 49. 52. 50. 50. 52.] resets=0
|
||||
15 resetdrop 28614 83.2 85.3 +2.1 84.3 [ 74. 82. 87. 84. 89. 85. 83. 87. 85. 86.] resets=177
|
||||
15 resetdrop-shuf 28614 51.4 50.5 -0.9 51.0 [ 51. 52. 51. 51. 52. 50. 51. 51. 50. 50.] resets=202
|
||||
15 rehearse-all 28614 82.0 80.2 -1.7 81.1 [ 78. 83. 82. 82. 85. 80. 79. 78. 80. 83.] resets=0
|
||||
15 rehearse-all-shuf 28614 50.4 50.7 +0.3 50.6 [ 51. 51. 51. 50. 50. 49. 51. 50. 52. 51.] resets=0
|
||||
-- by horizon (accum / window / resetdrop, TRUE labels only) --
|
||||
15 frozen 7185 62.4 62.8 +0.4 62.6 [ 61. 58. 63. 74. 56. 68. 69. 57. 61. 59.] resets=0
|
||||
15 accum 7185 75.2 75.8 +0.6 75.5 [ 71. 73. 74. 79. 79. 75. 74. 73. 78. 78.] resets=0
|
||||
15 window 7185 84.2 85.3 +1.1 84.7 [ 80. 84. 87. 84. 86. 85. 84. 87. 86. 84.] resets=0
|
||||
15 resetdrop 7185 82.0 83.6 +1.6 82.8 [ 71. 81. 86. 85. 87. 83. 83. 86. 83. 83.] resets=45
|
||||
20 frozen 7164 62.7 62.5 -0.2 62.6 [ 63. 59. 62. 69. 60. 71. 67. 58. 58. 58.] resets=0
|
||||
20 accum 7164 74.4 75.2 +0.9 74.8 [ 71. 73. 74. 78. 76. 77. 76. 65. 77. 80.] resets=0
|
||||
20 window 7164 85.4 85.7 +0.3 85.6 [ 80. 87. 86. 85. 89. 85. 85. 87. 84. 88.] resets=0
|
||||
20 resetdrop 7164 83.5 85.3 +1.8 84.4 [ 71. 84. 88. 85. 90. 85. 83. 88. 84. 86.] resets=42
|
||||
25 frozen 7143 65.7 62.2 -3.5 64.0 [ 72. 56. 66. 73. 61. 73. 66. 56. 61. 56.] resets=0
|
||||
25 accum 7143 76.1 75.8 -0.3 76.0 [ 78. 69. 75. 80. 79. 79. 76. 65. 81. 78.] resets=0
|
||||
25 window 7143 86.1 86.9 +0.8 86.5 [ 81. 88. 88. 85. 89. 87. 86. 88. 86. 88.] resets=0
|
||||
25 resetdrop 7143 84.1 85.7 +1.6 84.9 [ 78. 83. 86. 84. 89. 86. 83. 87. 86. 87.] resets=46
|
||||
30 frozen 7122 66.9 64.4 -2.5 65.6 [ 74. 60. 70. 71. 60. 75. 65. 56. 65. 60.] resets=0
|
||||
30 accum 7122 75.4 77.0 +1.6 76.2 [ 77. 68. 75. 78. 78. 81. 75. 69. 80. 80.] resets=0
|
||||
30 window 7122 86.5 87.0 +0.6 86.7 [ 82. 89. 87. 86. 89. 87. 87. 87. 86. 88.] resets=0
|
||||
30 resetdrop 7122 83.3 86.8 +3.5 85.0 [ 77. 81. 87. 83. 89. 86. 85. 87. 86. 89.] resets=44
|
||||
|
||||
|
||||
## tr_drussgt_vs_modularbot_shield.jsonl protocol=cross (pooled over rounds/horizons)
|
||||
h arm N early% late% decay total% curve(deciles of eval) resets
|
||||
15 frozen 46902 59.6 62.9 +3.3 61.2 [ 57. 63. 63. 57. 59. 64. 65. 65. 58. 62.] resets=0
|
||||
15 accum 46902 72.5 74.7 +2.2 73.6 [ 73. 71. 73. 75. 71. 73. 75. 77. 71. 77.] resets=0
|
||||
15 accum-shuf 46902 50.3 50.4 +0.1 50.4 [ 50. 50. 51. 51. 48. 50. 51. 50. 50. 51.] resets=0
|
||||
15 window 46902 85.4 86.2 +0.9 85.8 [ 82. 87. 85. 87. 86. 87. 87. 85. 85. 86.] resets=0
|
||||
15 window-shuf 46902 50.4 50.2 -0.2 50.3 [ 50. 51. 51. 51. 49. 51. 50. 50. 50. 51.] resets=0
|
||||
15 resetdrop 46902 83.7 85.8 +2.1 84.8 [ 76. 86. 85. 87. 86. 86. 87. 85. 85. 86.] resets=294
|
||||
15 resetdrop-shuf 46902 50.6 50.3 -0.3 50.4 [ 50. 51. 50. 50. 51. 51. 49. 50. 50. 51.] resets=351
|
||||
15 rehearse-all 46902 78.4 79.5 +1.1 78.9 [ 77. 77. 79. 81. 78. 79. 83. 78. 77. 81.] resets=0
|
||||
15 rehearse-all-shuf 46902 50.2 50.9 +0.7 50.5 [ 50. 52. 50. 50. 49. 50. 51. 51. 51. 51.] resets=0
|
||||
-- by horizon (accum / window / resetdrop, TRUE labels only) --
|
||||
15 frozen 11793 58.2 61.7 +3.4 60.0 [ 55. 62. 62. 55. 58. 64. 63. 64. 58. 60.] resets=0
|
||||
15 accum 11793 71.9 73.2 +1.3 72.6 [ 71. 70. 73. 76. 70. 71. 74. 77. 70. 75.] resets=0
|
||||
15 window 11793 84.0 84.6 +0.6 84.3 [ 82. 84. 84. 86. 84. 86. 86. 84. 83. 84.] resets=0
|
||||
15 resetdrop 11793 82.1 84.3 +2.2 83.2 [ 74. 85. 83. 85. 84. 85. 86. 83. 84. 83.] resets=76
|
||||
20 frozen 11748 59.4 63.3 +3.9 61.3 [ 57. 62. 63. 57. 58. 64. 65. 67. 59. 62.] resets=0
|
||||
20 accum 11748 72.6 74.7 +2.2 73.6 [ 74. 71. 73. 75. 70. 74. 74. 78. 69. 78.] resets=0
|
||||
20 window 11748 85.7 86.1 +0.4 85.9 [ 82. 87. 85. 88. 87. 88. 87. 85. 84. 86.] resets=0
|
||||
20 resetdrop 11748 83.3 85.5 +2.3 84.4 [ 77. 85. 84. 86. 85. 87. 85. 85. 84. 86.] resets=74
|
||||
25 frozen 11703 60.8 63.2 +2.4 62.0 [ 57. 65. 64. 57. 61. 64. 67. 63. 58. 63.] resets=0
|
||||
25 accum 11703 73.2 74.8 +1.7 74.0 [ 72. 73. 74. 76. 71. 75. 75. 78. 69. 77.] resets=0
|
||||
25 window 11703 85.9 87.1 +1.2 86.5 [ 82. 87. 86. 88. 87. 88. 89. 85. 87. 87.] resets=0
|
||||
25 resetdrop 11703 84.6 86.5 +1.9 85.5 [ 75. 88. 86. 88. 86. 87. 87. 86. 86. 86.] resets=72
|
||||
30 frozen 11658 59.8 63.5 +3.7 61.7 [ 59. 62. 61. 57. 60. 63. 66. 65. 59. 64.] resets=0
|
||||
30 accum 11658 72.3 75.9 +3.6 74.1 [ 74. 71. 74. 72. 71. 74. 77. 77. 74. 78.] resets=0
|
||||
30 window 11658 86.0 87.2 +1.3 86.6 [ 82. 89. 85. 87. 88. 88. 89. 87. 87. 86.] resets=0
|
||||
30 resetdrop 11658 84.9 87.0 +2.1 85.9 [ 77. 88. 85. 87. 88. 87. 88. 86. 87. 88.] resets=72
|
||||
|
||||
|
||||
## INERTIA SWEEP (tr_drussgt_vs_modularbot_shield.jsonl, within protocol, pooled over rounds x horizons)
|
||||
states arm early%/late% (TRUE labels)
|
||||
8 accum 85.0/ 84.4
|
||||
8 window 86.0/ 86.2
|
||||
16 accum 79.5/ 79.8
|
||||
16 window 85.5/ 86.2
|
||||
64 accum 75.3/ 75.9
|
||||
64 window 85.5/ 86.2
|
||||
256 accum 75.3/ 75.9
|
||||
256 window 85.5/ 86.2
|
||||
|
||||
====================================================================================================
|
||||
## INTEGRITY: shuffled-label controls should hover at chance (~50%)
|
||||
If a fix only improves the shuffled rows, it is NOT learning the enemy.
|
||||
rehearse-all = same periodic full retrain as window but WITHOUT forgetting
|
||||
(isolates the effect of the sliding buffer from the effect of retraining).
|
||||
## VERDICT INPUTS: late-half accuracy, TRUE vs shuffled (pp)
|
||||
====================================================================================================
|
||||
@@ -18,7 +18,7 @@
|
||||
## Run with plain:
|
||||
## nim c -r common_libs/tests/test_tm_horizon.nim
|
||||
|
||||
import std/[math]
|
||||
import std/[math, os]
|
||||
import gun_harness/gun_interface
|
||||
import gun_harness/virtual_bullets
|
||||
import gun_harness/offline_range
|
||||
@@ -316,6 +316,94 @@ proc testWarmShiftMovesAim() =
|
||||
check "warm: the applied rotation is the configured -3.0 deg",
|
||||
approx(d, -3.0, 1e-6)
|
||||
|
||||
proc testReadaptDefaultKnobs() =
|
||||
## The shipped defaults must be EXACTLY today's behaviour: no buffering, no
|
||||
## change detection, compile-time state count, rolling curve off.
|
||||
var g = initTmHorizonGun()
|
||||
check "readapt: WINDOW defaults to 0 (keep everything)", g.windowN == 0
|
||||
check "readapt: RESET_DROP defaults to 0 (off)", g.resetDrop == 0.0
|
||||
check "readapt: the sample buffer is disabled by default", g.bufferCapacity() == 0
|
||||
check "readapt: NSTATES defaults to the compile-time TMH_NSTATES",
|
||||
g.nStates == TMH_NSTATES and g.headStates() == TMH_NSTATES
|
||||
check "readapt: the accuracy curve is off by default", not g.accurveEnabled
|
||||
drive(g, synthesizeCircular(ticks = 240))
|
||||
check "readapt: no reset fires when RESET_DROP is off", g.resetDrops == 0
|
||||
check "readapt: rolling accuracy is still recorded (measurement only)",
|
||||
g.rollingAcc(100) >= 0.0 and g.rollingAcc(100) <= 1.0
|
||||
|
||||
proc testRuntimeStatesKnob() =
|
||||
## Inertia is sweepable WITHOUT a rebuild via TR_TMHORIZON_NSTATES.
|
||||
putEnv("TR_TMHORIZON_NSTATES", "16")
|
||||
var g = initTmHorizonGun()
|
||||
delEnv("TR_TMHORIZON_NSTATES")
|
||||
check "readapt: NSTATES env sets the runtime state count", g.nStates == 16
|
||||
check "readapt: both heads get the runtime state count",
|
||||
g.headStates() == 16
|
||||
var d = initTmHorizonGun()
|
||||
check "readapt: an unset env cannot move the state count",
|
||||
d.nStates == TMH_NSTATES
|
||||
|
||||
proc testWindowBuffersAndRebuilds() =
|
||||
## Sliding-window mode must allocate a bounded ring, keep buffering resolved
|
||||
## samples, and rebuild both heads deterministically from that ring.
|
||||
var g = initTmHorizonGun()
|
||||
g.setShift(0.0)
|
||||
g.setWindow(50)
|
||||
g.setRetrainConfig(50, 1)
|
||||
check "window: the ring is allocated to the window width", g.bufferCapacity() == 50
|
||||
drive(g, synthesizeCircular(ticks = 240))
|
||||
check "window: more than the window width of samples were buffered",
|
||||
g.bufferedCount() > 50 and g.trained > 0
|
||||
let before = sideTeamCopy(g)
|
||||
g.retrainFromBuffer(50)
|
||||
let a = sideTeamCopy(g)
|
||||
g.retrainFromBuffer(50)
|
||||
let b = sideTeamCopy(g)
|
||||
check "window: a rebuild is deterministic from the buffer", a == b
|
||||
check "window: a rebuild actually resets + retrains the machine", a != before
|
||||
check "window: the resolved-sample count keeps climbing across rebuilds",
|
||||
g.trained > 0
|
||||
|
||||
proc testWindowRetrainKeepsDeferredState() =
|
||||
## A window rebuild / re-learn must NOT touch the deferred-label queue or the
|
||||
## observation ring; only `resetRoundState`/`resetLearning` may clear those.
|
||||
var g = initTmHorizonGun()
|
||||
g.setShift(0.0)
|
||||
g.setWindow(50)
|
||||
g.setRetrainConfig(50, 1)
|
||||
drive(g, synthesizeCircular(ticks = 240))
|
||||
let pend = g.pendingCount
|
||||
let ringN = g.ringValidCount()
|
||||
check "window: deferred labels exist before the rebuild", pend > 0
|
||||
g.retrainFromBuffer(50)
|
||||
check "window: retrain keeps the deferred label queue", g.pendingCount == pend
|
||||
check "window: retrain keeps the observation ring", g.ringValidCount() == ringN
|
||||
g.resetLearning()
|
||||
check "window: a battle reset drops the buffered samples", g.bufferedCount() == 0
|
||||
check "window: a battle reset drops the rolling accuracy ring", g.accCount == 0
|
||||
|
||||
proc testResetDropTriggersRelearn() =
|
||||
## Change detection: with the knob on, a rolling-accuracy drop below its own
|
||||
## peak fires a re-learn; with it off, nothing fires. The circular fixture
|
||||
## mixes a cold start with warm tracking, so the peak/current gap is real.
|
||||
var off = initTmHorizonGun()
|
||||
off.setShift(0.0)
|
||||
drive(off, synthesizeCircular(ticks = 240))
|
||||
check "resetdrop: off by default fires nothing", off.resetDrops == 0
|
||||
|
||||
var g = initTmHorizonGun()
|
||||
g.setShift(0.0)
|
||||
g.setResetDrop(0.1)
|
||||
g.setRetrainConfig(50, 1)
|
||||
check "resetdrop: enabling it allocates the re-learn ring",
|
||||
g.bufferCapacity() >= TMH_RESET_WINDOW_DEF
|
||||
drive(g, synthesizeCircular(ticks = 240))
|
||||
check "resetdrop: a rolling-accuracy drop triggers at least one re-learn",
|
||||
g.resetDrops >= 1
|
||||
check "resetdrop: the accuracy ring is populated", g.accCount == g.sideTotal
|
||||
check "resetdrop: rolling accuracy stays a valid fraction",
|
||||
g.rollingAcc(100) >= 0.0 and g.rollingAcc(100) <= 1.0
|
||||
|
||||
when isMainModule:
|
||||
testHorizonMaths()
|
||||
testHorizonBuckets()
|
||||
@@ -334,6 +422,11 @@ when isMainModule:
|
||||
testStaleObservationsDropped()
|
||||
testColdModelEmitsNoShift()
|
||||
testWarmShiftMovesAim()
|
||||
testReadaptDefaultKnobs()
|
||||
testRuntimeStatesKnob()
|
||||
testWindowBuffersAndRebuilds()
|
||||
testWindowRetrainKeepsDeferredState()
|
||||
testResetDropTriggersRelearn()
|
||||
if failures > 0:
|
||||
echo "\n", failures, " check(s) FAILED"
|
||||
quit(1)
|
||||
|
||||
Reference in New Issue
Block a user