aed579b3af
The user's requirement: "every battle i means from round 1 to round end-battle, so retain all learning until the enemy change." What was built wiped the Tsetlin machines EVERY ROUND, in two places (`onRoundStarted` and the gun's own tick-regression self-reset), so in a 7-round battle each round started cold, trained ~360 samples and threw them away - discarding most of its one chance to do what was asked: overfit the current enemy over the whole battle. THE FIX - two kinds of state, two triggers: - **`resetRoundState` (per ROUND)**: the observation ring, pending/deferred labels, the bullet proxy, motion history, per-tick caches, `roundStartTrained`. These MUST clear every round, because bots teleport back to the starting corners between rounds - an old position would build a garbage label. (That exact class of bug shipped 36-58% wrong labels in the old gun.) - **`resetLearning` (per BATTLE / per ENEMY)**: both Tsetlin machines, `trained`, `sideCorrect/sideTotal`, all histograms, the magnitude median, `pendingDropped`, `observedTargetId`. These now SURVIVE round boundaries. Triggers for the machine wipe: `onGameStarted` (primary) plus a redundant `roundNumber <= 1` fallback in `onRoundStarted`; and a TARGET CHANGE (`targetChanged`, knob `TR_TMHORIZON_RESET_ON_TARGET` default on - a no-op in 1v1, fires on melee target switches; first acquisition never wipes). The tick-regression self-reset now clears ONLY per-round state. Still NO cross-battle persistence: grep for file I/O in the gun finds none. PROOF IT WORKS (live 2-round battle, `TR_RACK_PATTERN=off TR_RACK_TMHORIZON=both`): [tmh-reset] reason=game_start trained_was=0 [tmh-reset] reason=round1 trained_was=0 ...exactly TWO reset lines in the whole battle, both at battle start, and NONE at the round-2 boundary. And the per-round summaries: [tmh-round] trained=1249 thisRound=1249 ... sideAcc=711/1177 (60.4%) [tmh-round] trained=2226 thisRound=977 ... sideAcc=1361/2154 (63.2%) `trained` CLIMBED 1249 -> 2226 across the boundary, and side accuracy rose 60.4% -> 63.2% in round 2 (one battle - suggestive, not proof). Unit tests: `test_tm_horizon` 79 (was 54), including "trained SURVIVES the boundary", "clause states SURVIVE", "trained climbs round1->round2", "game-start wipes and all clauses end Exclude", "different enemy wipes / same enemy does not / knob-off does not", and crucially "a label CANNOT be built across a round boundary" (ringValidCount==0, ringHas(oldTick)==false, pendingCount==0) - the single most dangerous interaction of this change. Guards: test_tm_horizon 79, test_rack_membership 48, test_gun_harness 39, test_vbullet_metric 11, test_power_selection 3, test_adaptive_radar 41, test_tfil_ring_weights 24, test_power_policy 26, test_ram_decision 40, test_selector_tiebreak 19, test_tm_pattern_registration 20, test_vbullet_admit_gate 12, test_tm_diag 48, test_tm_automata_diag 55, test_tm_clause_shape 66. acceptance_offline_vs_online 12/12 VERDICT PASS. Shipped rack unchanged: DefaultRackMembership is still Pattern-only, TMHORIZON off. Residual (pre-existing, out of scope, stated): the internal base `PatternMatcherGun` has a rolling move-history buffer that is NOT cleared at round boundaries - it never was, and the SHIPPED Pattern gun carries history across rounds too. It cannot affect label correctness (labels come from `g.ring`), only base-prediction quality in a round's first ticks.
341 lines
16 KiB
Nim
341 lines
16 KiB
Nim
## Pure unit guard for the horizon-based TM gun (common_libs/guns/tm_horizon).
|
|
##
|
|
## No Java, no server, no battle. Covers what can be tested without a battle:
|
|
## * the horizon-from-flight-time maths and its [10,50] clamp;
|
|
## * the 4-bit horizon bucket boundaries;
|
|
## * feature extraction shapes: 49 draft bits with exactly one-hot blocks, and
|
|
## the 53-bit literal vector with pos/neg complementary literals;
|
|
## * the applied-shift geometry (0 deg = identity, +90 = CCW);
|
|
## * label lookup: a sample resolves h ticks later, is DROPPED when the
|
|
## observation is stale, the last h ticks of a round never resolve, and a
|
|
## round boundary wipes the pending queue AND the observation ring (so a
|
|
## label can never be built from a position across a round boundary);
|
|
## * the reset SPLIT: a round boundary keeps the machines (learning
|
|
## accumulates), a game start wipes them, and a target change wipes them
|
|
## unless TR_TMHORIZON_RESET_ON_TARGET is off;
|
|
## * a full round trains the two binary heads.
|
|
##
|
|
## Run with plain:
|
|
## nim c -r common_libs/tests/test_tm_horizon.nim
|
|
|
|
import std/[math]
|
|
import gun_harness/gun_interface
|
|
import gun_harness/virtual_bullets
|
|
import gun_harness/offline_range
|
|
import guns/pattern_matcher
|
|
import guns/tm_horizon
|
|
|
|
var failures = 0
|
|
proc check(name: string, ok: bool) =
|
|
if ok: echo "PASS: ", name
|
|
else: echo "FAIL: ", name; inc failures
|
|
|
|
proc approx(a, b, tol: float): bool {.inline.} = abs(a - b) <= tol
|
|
|
|
proc sideTeamCopy(g: TmHorizonGun): seq[int16] =
|
|
## Snapshot of the learned clause states, to prove the machines survive or are
|
|
## wiped by a given reset.
|
|
g.sideClauseStates()
|
|
|
|
proc drive(g: var TmHorizonGun, fx: Fixture) =
|
|
for state in fx.states:
|
|
for b in 0..<len(PowerBins):
|
|
discard g.predict(state, bulletSpeed(PowerBins[b]))
|
|
|
|
# ── horizon maths ─────────────────────────────────────────────────────────────
|
|
|
|
proc testHorizonMaths() =
|
|
check "horizon: dist 220 / speed 11 -> 20 ticks",
|
|
tmhHorizonFor(220.0, 11.0) == 20
|
|
check "horizon: dist 100 / speed 17 -> clamped to H_MIN (10)",
|
|
tmhHorizonFor(100.0, 17.0) == TMH_H_MIN
|
|
check "horizon: dist 800 / speed 11 -> clamped to H_MAX (50)",
|
|
tmhHorizonFor(800.0, 11.0) == TMH_H_MAX
|
|
check "horizon: zero bullet speed falls back to H_MIN",
|
|
tmhHorizonFor(300.0, 0.0) == TMH_H_MIN
|
|
check "horizon: rounding is to the nearest tick",
|
|
tmhHorizonFor(250.0, 10.0) == 25
|
|
|
|
proc testHorizonBuckets() =
|
|
check "bucket: 10..19 -> 0",
|
|
tmhHorizonBucket(10) == 0 and tmhHorizonBucket(19) == 0
|
|
check "bucket: 20..29 -> 1",
|
|
tmhHorizonBucket(20) == 1 and tmhHorizonBucket(29) == 1
|
|
check "bucket: 30..39 -> 2",
|
|
tmhHorizonBucket(30) == 2 and tmhHorizonBucket(39) == 2
|
|
check "bucket: 40..50 -> 3",
|
|
tmhHorizonBucket(40) == 3 and tmhHorizonBucket(50) == 3
|
|
|
|
proc testQuadrantNames() =
|
|
check "quadrant: cold model names COLD", tmhQuadrantName(-1, -1) == "COLD"
|
|
check "quadrant: LARGE-LEFT", tmhQuadrantName(1, 1) == "LARGE-LEFT"
|
|
check "quadrant: SMALL-RIGHT", tmhQuadrantName(0, 0) == "SMALL-RIGHT"
|
|
|
|
# ── feature shapes ────────────────────────────────────────────────────────────
|
|
|
|
proc blockSum(b: array[TMH_N_BASE, uint8], lo, hi: int): int =
|
|
for i in lo..hi: result += int(b[i])
|
|
|
|
proc testFeatureShapes() =
|
|
var g = initTmHorizonGun()
|
|
# A straight, moving enemy so every history-dependent block is populated.
|
|
let fx = synthesizeConstantVelocity(ticks = 40, speed = 4.0)
|
|
for state in fx.states:
|
|
discard g.predict(state, bulletSpeed(PowerBins[0]))
|
|
let state = fx.states[^1]
|
|
let b = g.tmhBaseBits(state)
|
|
check "bits: the draft vector is exactly 49 bits", b.len == TMH_N_BASE
|
|
# Every one-hot block must carry exactly one set bit.
|
|
check "bits: dist-to-nearest-wall is one-hot", blockSum(b, 0, 3) == 1
|
|
check "bits: which-wall-nearest is one-hot", blockSum(b, 4, 7) == 1
|
|
check "bits: dist-from-us is one-hot", blockSum(b, 8, 13) == 1
|
|
check "bits: enemy-heading-vs-line-to-us is one-hot", blockSum(b, 14, 16) == 1
|
|
check "bits: ticks-since-reversal is one-hot", blockSum(b, 20, 24) == 1
|
|
check "bits: turn-consistency-10 is one-hot", blockSum(b, 25, 27) == 1
|
|
check "bits: distance-moved-10 is one-hot", blockSum(b, 28, 30) == 1
|
|
check "bits: speed-trend-10 is one-hot", blockSum(b, 31, 33) == 1
|
|
check "bits: turn-rate-change-5 is one-hot", blockSum(b, 34, 36) == 1
|
|
check "bits: time-until-bullet is one-hot", blockSum(b, 37, 41) == 1
|
|
check "bits: bullet-lateral-offset is one-hot", blockSum(b, 42, 48) == 1
|
|
|
|
proc testLiteralLayout() =
|
|
var base: array[TMH_N_BASE, uint8]
|
|
base[0] = 1'u8
|
|
base[8] = 1'u8
|
|
for bucket in 0..<TMH_NH:
|
|
let lits = tmhLits(base, bucket)
|
|
check "lits: length is 2 * 53", lits.len == TMH_NLITS
|
|
check "lits: horizon bucket " & $bucket & " sets exactly one of the 4 raw bits",
|
|
lits[TMH_N_BASE + bucket] == 1'u8 and
|
|
lits[TMH_N_BASE + ((bucket + 1) mod TMH_NH)] == 0'u8
|
|
var complement = true
|
|
for i in 0..<TMH_N_BITS:
|
|
if int(lits[i]) + int(lits[i + TMH_N_BITS]) != 1: complement = false
|
|
check "lits: every literal has its complementary negation (bucket " & $bucket & ")",
|
|
complement
|
|
|
|
# ── shift geometry ────────────────────────────────────────────────────────────
|
|
|
|
proc testShiftGeometry() =
|
|
let p0 = tmhApplyShift(100.0, 100.0, 200.0, 100.0, 0.0)
|
|
check "shift: 0 deg is the identity", approx(p0.x, 200.0, 1e-9) and approx(p0.y, 100.0, 1e-9)
|
|
# East point rotated +90 deg (CCW) -> north (+y in the math convention).
|
|
let p90 = tmhApplyShift(100.0, 100.0, 200.0, 100.0, 90.0)
|
|
check "shift: +90 deg rotates east to +y (CCW)",
|
|
approx(p90.x, 100.0, 1e-9) and approx(p90.y, 200.0, 1e-9)
|
|
let pm90 = tmhApplyShift(100.0, 100.0, 200.0, 100.0, -90.0)
|
|
check "shift: -90 deg rotates east to -y (CW)",
|
|
approx(pm90.x, 100.0, 1e-9) and approx(pm90.y, 0.0, 1e-9)
|
|
check "shift: the aim distance is preserved",
|
|
approx(hypot(p90.x - 100.0, p90.y - 100.0), 100.0, 1e-9)
|
|
|
|
# ── label resolution ──────────────────────────────────────────────────────────
|
|
|
|
proc testRoundTrains() =
|
|
var g = initTmHorizonGun()
|
|
g.setShift(0.0) # pure-predict arm: still trains
|
|
let fx = synthesizeCircular(ticks = 240)
|
|
drive(g, fx)
|
|
check "train: a full round resolves samples (trained > 0)", g.trained > 0
|
|
check "train: the model warmed past the cold gate", g.trained >= TMH_MIN_OBS
|
|
check "train: the last h ticks are still pending at round end",
|
|
g.pendingCount > 0
|
|
|
|
proc testRoundReset() =
|
|
var g = initTmHorizonGun()
|
|
g.setShift(0.0)
|
|
drive(g, synthesizeCircular(ticks = 240))
|
|
let t1 = g.trained
|
|
g.resetLearning()
|
|
check "reset: trained is wiped", g.trained == 0
|
|
check "reset: the pending queue is wiped (no cross-round labels)", g.pendingCount == 0
|
|
check "reset: side accuracy counters are wiped", g.sideTotal == 0
|
|
drive(g, synthesizeCircular(ticks = 240))
|
|
check "reset: the fresh round trains again", g.trained > 0 and t1 > 0
|
|
|
|
proc testRoundBoundaryKeepsMachines() =
|
|
## THE SPLIT: a plain round boundary (`resetRoundState`) must keep the learned
|
|
## machines and drop only the per-round observation/label state.
|
|
var g = initTmHorizonGun()
|
|
g.setShift(0.0)
|
|
drive(g, synthesizeCircular(ticks = 240))
|
|
let trainedBefore = g.trained
|
|
let sigBefore = sideTeamCopy(g)
|
|
check "round-boundary: samples trained before the boundary", trainedBefore > 0
|
|
check "round-boundary: unresolved labels exist at the boundary", g.pendingCount > 0
|
|
g.resetRoundState()
|
|
check "round-boundary: trained SURVIVES the boundary", g.trained == trainedBefore
|
|
check "round-boundary: the clause states SURVIVE the boundary",
|
|
sideTeamCopy(g) == sigBefore
|
|
check "round-boundary: the pending label queue is cleared", g.pendingCount == 0
|
|
check "round-boundary: the observation ring is cleared", g.ringValidCount() == 0
|
|
|
|
proc testLearningAccumulatesAcrossRounds() =
|
|
## The whole point of the fix: round 2 keeps round 1's learning and keeps
|
|
## training on top of it, so `trained` CLIMBS across the battle.
|
|
var g = initTmHorizonGun()
|
|
g.setShift(0.0)
|
|
drive(g, synthesizeCircular(ticks = 240))
|
|
let r1 = g.trained
|
|
g.resetRoundState() # exactly what onRoundStarted now does
|
|
drive(g, synthesizeCircular(ticks = 240))
|
|
check "accumulate: round 2 starts from round 1's total (trained climbs)",
|
|
g.trained > r1
|
|
|
|
proc testGameStartWipesMachines() =
|
|
## A new BATTLE (`resetLearning`) must wipe machines + stats + round state.
|
|
var g = initTmHorizonGun()
|
|
g.setShift(0.0)
|
|
drive(g, synthesizeCircular(ticks = 240))
|
|
check "game-start: samples trained before the wipe", g.trained > 0
|
|
g.resetLearning("game_start")
|
|
check "game-start: trained is wiped", g.trained == 0
|
|
check "game-start: side accuracy counters are wiped", g.sideTotal == 0
|
|
check "game-start: the pending label queue is cleared", g.pendingCount == 0
|
|
check "game-start: the observation ring is cleared", g.ringValidCount() == 0
|
|
check "game-start: every clause is back to the Exclude boundary",
|
|
g.sideClausesAllExclude()
|
|
|
|
proc testTargetChangeResetsMachines() =
|
|
## Reset on TARGET change to a different bot id, gated by the knob.
|
|
var g = initTmHorizonGun()
|
|
g.setShift(0.0)
|
|
g.setResetOnTarget(true)
|
|
discard g.targetChanged(7) # first acquisition: never wipes
|
|
drive(g, synthesizeCircular(ticks = 240))
|
|
check "target: trained before the change", g.trained > 0
|
|
discard g.targetChanged(7) # same enemy: no wipe
|
|
check "target: the same enemy does not wipe the machines", g.trained > 0
|
|
let wiped = g.targetChanged(9) # different enemy: wipe
|
|
check "target: a DIFFERENT enemy wipes the machines", wiped and g.trained == 0
|
|
|
|
# Knob OFF: a different enemy must NOT wipe the machines.
|
|
var h = initTmHorizonGun()
|
|
h.setShift(0.0)
|
|
h.setResetOnTarget(false)
|
|
discard h.targetChanged(7)
|
|
drive(h, synthesizeCircular(ticks = 240))
|
|
check "target: trained before the change (knob off)", h.trained > 0
|
|
let wipedOff = h.targetChanged(9)
|
|
check "target: knob off leaves the machines intact",
|
|
(not wipedOff) and h.trained > 0
|
|
|
|
proc testLabelCannotCrossBoundary() =
|
|
## The most dangerous interaction: a deferred label must never be resolved
|
|
## against a position from the previous round. At a round boundary BOTH the
|
|
## pending queue and the observation ring are cleared, so the lookup at
|
|
## `fireTick + h` cannot find an old position (and no old pending survives).
|
|
var g = initTmHorizonGun()
|
|
g.setShift(0.0)
|
|
let fx = synthesizeCircular(ticks = 240)
|
|
drive(g, fx)
|
|
let endTick = fx.states[^1].tick
|
|
check "cross-boundary: unresolved labels exist at round end", g.pendingCount > 0
|
|
check "cross-boundary: the round-end observation is in the ring", g.ringHas(endTick)
|
|
g.resetRoundState() # the round boundary
|
|
check "cross-boundary: the old observation is gone", not g.ringHas(endTick)
|
|
check "cross-boundary: the deferred label queue is gone", g.pendingCount == 0
|
|
check "cross-boundary: the ring is empty after the boundary", g.ringValidCount() == 0
|
|
|
|
proc testTickRegressionKeepsMachines() =
|
|
## The tick-regression self-reset (a missed onRoundStarted) must clear only the
|
|
## per-round state and must NOT wipe the machines.
|
|
var g = initTmHorizonGun()
|
|
g.setShift(0.0)
|
|
drive(g, synthesizeCircular(ticks = 120))
|
|
let t = g.trained
|
|
let sig = sideTeamCopy(g)
|
|
check "regression: samples trained before the new round", t > 0
|
|
# Simulate the server resetting the tick counter to 0 for a new round.
|
|
let fx = synthesizeCircular(ticks = 5)
|
|
discard g.predict(fx.states[0], bulletSpeed(PowerBins[0]))
|
|
check "regression: a tick regression does NOT wipe the machines", g.trained == t
|
|
check "regression: the clause states survive the regression",
|
|
sideTeamCopy(g) == sig
|
|
check "regression: the old observation ring is cleared (one fresh tick only)",
|
|
g.ringValidCount() == 1
|
|
|
|
proc testStaleObservationsDropped() =
|
|
## Build a fixture whose `lastSeenTick` is frozen far in the past: every
|
|
## resolved label must be dropped, never trained on.
|
|
var g = initTmHorizonGun()
|
|
g.setShift(0.0)
|
|
var states: seq[WorldState]
|
|
for t in 0..<120:
|
|
var s = WorldState(
|
|
enemyX: 400.0 + 3.0 * t.float, enemyY: 300.0,
|
|
enemyHeading: 0.0, enemySpeed: 3.0, enemyEnergy: 100.0,
|
|
selfX: 100.0, selfY: 300.0, selfEnergy: 100.0,
|
|
arenaWidth: 800.0, arenaHeight: 600.0, tick: t,
|
|
enemies: @[EnemyInfo(id: 1, x: 400.0 + 3.0 * t.float, y: 300.0,
|
|
heading: 0.0, speed: 3.0, energy: 100.0,
|
|
lastSeenTick: 0)]) # frozen -> always stale
|
|
states.add s
|
|
for state in states:
|
|
for b in 0..<len(PowerBins):
|
|
discard g.predict(state, bulletSpeed(PowerBins[b]))
|
|
check "stale: no sample with a stale observation is ever trained", g.trained == 0
|
|
check "stale: the dropped samples are counted", g.pendingDropped > 0
|
|
|
|
proc testColdModelEmitsNoShift() =
|
|
## A cold machine must return Pattern's prediction UNCHANGED. Constant-velocity
|
|
## motion is predicted perfectly by Pattern, so no sample trains and the model
|
|
## stays cold for the whole fixture.
|
|
var g = initTmHorizonGun()
|
|
g.setShift(3.0)
|
|
var pm = PatternMatcherGun()
|
|
let fx = synthesizeConstantVelocity(ticks = 60, speed = 4.0)
|
|
for state in fx.states:
|
|
for b in 0..<len(PowerBins):
|
|
let p = g.predict(state, bulletSpeed(PowerBins[b]))
|
|
let q = pm.predict(state, bulletSpeed(PowerBins[b]))
|
|
if not approx(p.x, q.x, 1e-9) or not approx(p.y, q.y, 1e-9):
|
|
check "cold: prediction must equal Pattern byte-for-byte", false
|
|
return
|
|
check "cold: prediction equals Pattern byte-for-byte while cold", true
|
|
|
|
proc testWarmShiftMovesAim() =
|
|
## Force the model warm with an all-Exclude head (votes 0/0 -> RIGHT) and check
|
|
## the correction actually rotates Pattern's base aim by the configured -3 deg.
|
|
var g = initTmHorizonGun()
|
|
g.setShift(3.0, 1.0)
|
|
g.trained = 100 # force warm; the fresh head votes 0/0 -> RIGHT
|
|
let fx = synthesizeCircular(ticks = 6)
|
|
let state = fx.states[2]
|
|
let speed = bulletSpeed(PowerBins[0])
|
|
let p = g.predict(state, speed)
|
|
# Pattern caches per tick, so this is the exact base point g used.
|
|
let q = g.pattern.predict(state, speed)
|
|
check "warm: the corrected aim differs from the Pattern base",
|
|
not (approx(p.x, q.x, 1e-9) and approx(p.y, q.y, 1e-9))
|
|
let b0 = arctan2(q.y - state.selfY, q.x - state.selfX)
|
|
let b1 = arctan2(p.y - state.selfY, p.x - state.selfX)
|
|
var d = radToDeg(b1 - b0)
|
|
while d > 180.0: d -= 360.0
|
|
while d < -180.0: d += 360.0
|
|
check "warm: the applied rotation is the configured -3.0 deg",
|
|
approx(d, -3.0, 1e-6)
|
|
|
|
when isMainModule:
|
|
testHorizonMaths()
|
|
testHorizonBuckets()
|
|
testQuadrantNames()
|
|
testFeatureShapes()
|
|
testLiteralLayout()
|
|
testShiftGeometry()
|
|
testRoundTrains()
|
|
testRoundReset()
|
|
testRoundBoundaryKeepsMachines()
|
|
testLearningAccumulatesAcrossRounds()
|
|
testGameStartWipesMachines()
|
|
testTargetChangeResetsMachines()
|
|
testLabelCannotCrossBoundary()
|
|
testTickRegressionKeepsMachines()
|
|
testStaleObservationsDropped()
|
|
testColdModelEmitsNoShift()
|
|
testWarmShiftMovesAim()
|
|
if failures > 0:
|
|
echo "\n", failures, " check(s) FAILED"
|
|
quit(1)
|
|
echo "\nAll tm-horizon checks passed."
|