69debbe347
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.
719 lines
26 KiB
Nim
719 lines
26 KiB
Nim
## 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()
|