TM horizon: retain learning ACROSS ROUNDS, reset only when the ENEMY changes
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.
This commit is contained in:
@@ -0,0 +1,850 @@
|
||||
## tm_horizon.nim — a HORIZON-BASED Tsetlin Machine that CORRECTS the Pattern gun.
|
||||
##
|
||||
## Design (the exact brief this file implements):
|
||||
## * BASE — the shipped best gun, `PatternMatcherGun` (guns/pattern_matcher).
|
||||
## Pattern supplies the prediction; the TM only CORRECTS it. The old
|
||||
## mistake of layering a TM on a weak base (Linear) is not repeated.
|
||||
## * INPUT — the 49-bit draft spec (`tm_diag/feature_spec.draftTMSpec`) PLUS a
|
||||
## 4-bit horizon one-hot, i.e. 53 raw bits. The draft blocks are:
|
||||
## walls(8) / us(9) / motion(20) / bullets(12) = 49
|
||||
## The horizon block is derived from the ACTUAL bullet flight time,
|
||||
## `speed = 20 - 3*power`, `h = round(dist / speed)`, clamped to
|
||||
## [10, 50] (the measured 5..9 dead zone is excluded).
|
||||
## * OUTPUT — TWO BINARIES, four quadrants:
|
||||
## (a) is the enemy LEFT or RIGHT of the base prediction?
|
||||
## (b) is the angular correction bigger or smaller than the median?
|
||||
## A multi-class head created the majority-class trap before; two
|
||||
## balanced binaries cannot.
|
||||
## * LABEL — a FACT, not a correction: at tick t and horizon h, look up where
|
||||
## the enemy ACTUALLY was at t+h in our OWN observation ring. Never
|
||||
## cross a round boundary (pending samples are cleared on reset; the
|
||||
## last h ticks simply never resolve). Only samples whose enemy was
|
||||
## observed within `TMH_STALE_MAX` ticks are used — stale
|
||||
## observations are guesses and guessing in the answer key is what
|
||||
## shipped wrong labels before.
|
||||
## * TRAIN — the machines SURVIVE round boundaries: learning accumulates
|
||||
## across every round of the same battle/enemy ("every battle i means
|
||||
## from round 1 to round end-battle, so retain all learning until the
|
||||
## enemy change"). Two different triggers wipe two different kinds of
|
||||
## state:
|
||||
## - `resetRoundState` (per-ROUND) clears ONLY the observation ring,
|
||||
## the deferred-label queue and the motion/bullet/per-tick
|
||||
## history. It runs from `onRoundStarted` and on a tick
|
||||
## regression. Bots teleport back to their corners between
|
||||
## rounds, so old positions are meaningless and would poison the
|
||||
## labels.
|
||||
## - `resetLearning` (per-BATTLE / per-ENEMY) wipes the Tsetlin
|
||||
## machines and every learned statistic, then clears the round
|
||||
## state too. It runs on a NEW BATTLE (`onGameStarted`, with a
|
||||
## round-1 fallback) and when the TARGET changes to a different
|
||||
## bot id (`TR_TMHORIZON_RESET_ON_TARGET`, default on).
|
||||
## There is still NO persistence across battles: nothing is written
|
||||
## to disk and nothing carries into a different battle.
|
||||
## * APPLY — a SMALL hit-optimal-style correction (default a few degrees,
|
||||
## `TR_TMHORIZON_SHIFT`), NOT the conditional median. Gate 2b showed
|
||||
## the error median (4-16 deg) is catastrophically wrong when the
|
||||
## hit-optimal shift is only ~±2-3.5 deg. `TR_TMHORIZON_SHIFT=0`
|
||||
## disables the correction entirely (pure predict arm).
|
||||
##
|
||||
## EXPECTED OUTCOME: this design is PREDICTED TO LOSE. Gate 2b measured that it
|
||||
## needs ~80% side accuracy to break even on hits while the achievable signal is
|
||||
## ~60%, so it is expected to lose to Pattern. It is built anyway because the
|
||||
## user asked to see it in a real battle and because offline metrics have been
|
||||
## wrong before. See `common_libs/tests/gate2b_hit_optimal_results.txt`.
|
||||
##
|
||||
## COST: per-tick work is cached the way Pattern caches its path (base bits once
|
||||
## 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.
|
||||
##
|
||||
## Coordinate system: 0° = East, CCW positive (Tank Royale standard).
|
||||
|
||||
import std/[math, os, strutils, strformat]
|
||||
import gun_harness/gun_interface
|
||||
import guns/pattern_matcher
|
||||
import tm_diag/tm_core
|
||||
|
||||
const
|
||||
## ── feature geometry ────────────────────────────────────────────────────
|
||||
TMH_N_BASE* = 49 ## the draftTMSpec() bit count
|
||||
TMH_NH* = 4 ## horizon one-hot width
|
||||
TMH_N_BITS* = TMH_N_BASE + TMH_NH ## 53 raw bits
|
||||
TMH_NLITS* = 2 * TMH_N_BITS ## pos + neg literals
|
||||
## ── horizon range (measured: 5..9 is a dead zone with no signal) ────────
|
||||
TMH_H_MIN* = 10
|
||||
TMH_H_MAX* = 50
|
||||
## ── classifier shape (cheap; two binaries) ─────────────────────────────
|
||||
TMH_NCLAUSES* {.intdefine.} = 40
|
||||
TMH_NSTATES* {.intdefine.} = 64
|
||||
TMH_S_DEF {.strdefine.} = "3.0"
|
||||
TMH_S* = parseFloat(TMH_S_DEF)
|
||||
## Cold gate: below this many resolved samples the TM emits no correction.
|
||||
TMH_MIN_OBS* {.intdefine.} = 24
|
||||
## A resolved sample is only used when the enemy was observed this recently.
|
||||
TMH_STALE_MAX* {.intdefine.} = 8
|
||||
## ── rings ───────────────────────────────────────────────────────────────
|
||||
TMH_POS_RING* = 256 ## observation ring (> max horizon + lookback)
|
||||
TMH_PENDING_CAP* = 512 ## deferred-label queue (<= 4 spawns/tick * 50)
|
||||
TMH_BULLETS* = 32 ## our own in-flight bullets (proxy)
|
||||
## Magnitude-median histogram resolution. Fine enough that the median does not
|
||||
## collapse to 0.0 when the bulk of |err| sits in the first bin (which would
|
||||
## label every sample BIG and reintroduce a majority-class trap).
|
||||
TMH_ABS_RES* = 0.1
|
||||
TMH_ABS_BINS* = 640 ## 0.1 deg bins up to 64 deg
|
||||
## ── env knobs ───────────────────────────────────────────────────────────
|
||||
TMH_SHIFT_ENV* = "TR_TMHORIZON_SHIFT" ## deg; 0 = no correction arm
|
||||
TMH_BIG_MULT_ENV* = "TR_TMHORIZON_BIG_MULT" ## scale when magnitude is BIG
|
||||
TMH_LOG_ENV* = "TR_TMHORIZON_LOG" ## 1 = per-change thinking log
|
||||
## Reset the machines when the TARGET changes to a different bot id. Default
|
||||
## ON: in 1v1 there is one enemy so it is a no-op; in melee it fires on target
|
||||
## switches. Set 0 to observe melee without the per-switch wipe.
|
||||
TMH_RESET_ON_TARGET_ENV* = "TR_TMHORIZON_RESET_ON_TARGET"
|
||||
TMH_SHIFT_DEFAULT* = 2.0
|
||||
TMH_BIG_MULT_DEFAULT* = 1.5
|
||||
TMH_RESET_ON_TARGET_DEFAULT* = true
|
||||
|
||||
type
|
||||
TmhTick* = object
|
||||
tick*: int
|
||||
ex*, ey*, eh*, es*: float
|
||||
lastSeenTick*: int
|
||||
|
||||
TmhBullet = object
|
||||
## INFERRED proxy for one of OUR real bullets: detected from a self-energy
|
||||
## drop (0.05 < drop <= 3.1), direction proxied by self->enemy at fire time —
|
||||
## exactly the offline `buildBulletSeries` model.
|
||||
arrivalTick: int
|
||||
ux, uy: float
|
||||
sx, sy: float
|
||||
active: bool
|
||||
|
||||
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.
|
||||
fireTick*: int
|
||||
horizon*: int
|
||||
bucket*: int
|
||||
selfX*, selfY*: float
|
||||
baseBearing*: float
|
||||
sidePred*: int
|
||||
magPred*: int
|
||||
warm*: bool
|
||||
lits*: array[TMH_NLITS, uint8]
|
||||
|
||||
TmhEval = object
|
||||
valid: bool
|
||||
side: int
|
||||
sideConf: float
|
||||
mag: int
|
||||
|
||||
TmHorizonGun* = object
|
||||
pattern*: PatternMatcherGun
|
||||
sideMachine: TmMachine
|
||||
magMachine: TmMachine
|
||||
# per-class scratch caches (avoid per-sample allocation)
|
||||
sideScratch: array[2, seq[uint8]]
|
||||
magScratch: array[2, seq[uint8]]
|
||||
# observation ring
|
||||
ring: array[TMH_POS_RING, TmhTick]
|
||||
ringValid: array[TMH_POS_RING, bool]
|
||||
# deferred labels
|
||||
pending: array[TMH_PENDING_CAP, TmhPending]
|
||||
pendingCount*: int
|
||||
pendingDropped*: int
|
||||
# our own bullets (energy-drop proxy)
|
||||
bullets: array[TMH_BULLETS, TmhBullet]
|
||||
bulletCount: int
|
||||
# per-tick caches
|
||||
cachedBits: array[TMH_N_BASE, uint8]
|
||||
cachedBitsTick: int
|
||||
bitsValid: bool
|
||||
bucketEval: array[TMH_NH, TmhEval]
|
||||
# motion history
|
||||
lastTick: int
|
||||
hasPrev: bool
|
||||
sinceRev: int
|
||||
lastNonzeroSign: int
|
||||
selfPrevEnergy: float
|
||||
hasSelfPrev: bool
|
||||
# enqueue dedupe (predict is called once per power bin)
|
||||
lastEnqTick: int
|
||||
lastEnqBucket: int
|
||||
# logging
|
||||
lastLogKey: string
|
||||
lastLogTick: int
|
||||
shiftConfigured: bool
|
||||
shiftDeg*: float
|
||||
bigMult*: float
|
||||
logEnabled*: bool
|
||||
# round seeding
|
||||
seedCounter: int
|
||||
# per-battle config: reset the machines when the target changes enemy id
|
||||
resetOnTarget*: bool
|
||||
# the enemy id the machines currently represent (-1 = none). Used by
|
||||
# `targetChanged` so first acquisition never wipes and a new BATTLE clears it.
|
||||
observedTargetId*: int
|
||||
# cumulative `trained` at the start of the current round, so the per-round
|
||||
# summary can report `thisRound` while `trained` keeps climbing.
|
||||
roundStartTrained*: int
|
||||
# magnitude-median histogram (TMH_ABS_RES deg bins)
|
||||
absHist: array[TMH_ABS_BINS, int]
|
||||
absCount: int
|
||||
# ── instrumentation ─────────────────────────────────────────────────────
|
||||
trained*: int
|
||||
sideCorrect*, sideTotal*: int
|
||||
sidePredHist*: array[2, int]
|
||||
magPredHist*: array[2, int]
|
||||
sideLabelHist*: array[2, int]
|
||||
magLabelHist*: array[2, int]
|
||||
quadHist*: array[4, int]
|
||||
lastSidePred*: int
|
||||
lastMagPred*: int
|
||||
|
||||
# ── small pure helpers ───────────────────────────────────────────────────────
|
||||
|
||||
proc wrapDeg(d: float): float {.inline.} =
|
||||
result = d
|
||||
while result > 180.0: result -= 360.0
|
||||
while result < -180.0: result += 360.0
|
||||
|
||||
proc wrapRad(r: float): float {.inline.} =
|
||||
result = r
|
||||
while result > PI: result -= 2.0 * PI
|
||||
while result < -PI: result += 2.0 * PI
|
||||
|
||||
proc signf(x: float): int {.inline.} =
|
||||
if x > 1e-9: 1 elif x < -1e-9: -1 else: 0
|
||||
|
||||
proc envFloatT(name: string, default: float): float =
|
||||
let v = getEnv(name, "")
|
||||
if v.len == 0: return default
|
||||
try: parseFloat(v.strip())
|
||||
except ValueError: default
|
||||
|
||||
proc envBoolT(name: string, default: bool): bool =
|
||||
case getEnv(name, "").strip().toLowerAscii()
|
||||
of "1", "true", "yes", "on": true
|
||||
of "0", "false", "no", "off": false
|
||||
else: default
|
||||
|
||||
# ── horizon maths (pure, unit-tested) ────────────────────────────────────────
|
||||
|
||||
proc tmhHorizonFor*(dist, bulletSpeed: float): int =
|
||||
## Derive the horizon from the ACTUAL bullet flight time. `speed = 20-3*power`
|
||||
## so `bulletSpeed` already encodes the energy-aware power policy's power.
|
||||
## Clamped to the measured live band [10, 50].
|
||||
if bulletSpeed <= 0.0: return TMH_H_MIN
|
||||
result = int(round(dist / bulletSpeed))
|
||||
if result < TMH_H_MIN: result = TMH_H_MIN
|
||||
elif result > TMH_H_MAX: result = TMH_H_MAX
|
||||
|
||||
proc tmhHorizonBucket*(h: int): int =
|
||||
## 4 one-hot buckets over [10, 50]: 10-19, 20-29, 30-39, 40-50.
|
||||
if h < 20: 0
|
||||
elif h < 30: 1
|
||||
elif h < 40: 2
|
||||
else: 3
|
||||
|
||||
proc tmhQuadrantName*(side, mag: int): string =
|
||||
## `side` 1 = LEFT (positive angular error, CCW), 0 = RIGHT; `mag` 1 = BIG.
|
||||
## `side < 0` means the model is cold and emits no correction.
|
||||
if side < 0: return "COLD"
|
||||
let s = if side == 1: "LEFT" else: "RIGHT"
|
||||
let m = if mag == 1: "LARGE" else: "SMALL"
|
||||
m & "-" & s
|
||||
|
||||
proc tmhApplyShift*(selfX, selfY, px, py, shiftDeg: float): GunPrediction =
|
||||
## Rotate the base prediction point around the shooter by `shiftDeg` degrees
|
||||
## (positive = CCW). The aim DISTANCE is preserved; only the bearing moves.
|
||||
let dx = px - selfX
|
||||
let dy = py - selfY
|
||||
let d = hypot(dx, dy)
|
||||
if d < 1e-9: return GunPrediction(x: px, y: py)
|
||||
let b = arctan2(dy, dx) + degToRad(shiftDeg)
|
||||
GunPrediction(x: selfX + cos(b) * d, y: selfY + sin(b) * d)
|
||||
|
||||
proc tmhLits*(base: array[TMH_N_BASE, uint8],
|
||||
bucket: int): array[TMH_NLITS, uint8] =
|
||||
## Pack the 49 draft bits + the 4-bit horizon one-hot into the pos-then-neg
|
||||
## literal layout the TM core uses.
|
||||
var raw: array[TMH_N_BITS, uint8]
|
||||
for i in 0..<TMH_N_BASE: raw[i] = base[i]
|
||||
raw[TMH_N_BASE + bucket] = 1'u8
|
||||
for i in 0..<TMH_N_BITS:
|
||||
result[i] = raw[i]
|
||||
result[i + TMH_N_BITS] = 1'u8 - raw[i]
|
||||
|
||||
# ── construction / reset ─────────────────────────────────────────────────────
|
||||
|
||||
proc initTmHorizonGun*(): TmHorizonGun =
|
||||
result.pattern = PatternMatcherGun()
|
||||
result.sideMachine = newMachine(TMH_N_BITS, 2, TMH_NCLAUSES, TMH_NSTATES,
|
||||
TMH_S, seed = 1)
|
||||
result.magMachine = newMachine(TMH_N_BITS, 2, TMH_NCLAUSES, TMH_NSTATES,
|
||||
TMH_S, seed = 2)
|
||||
for c in 0..1:
|
||||
result.sideScratch[c] = newSeq[uint8](TMH_NCLAUSES)
|
||||
result.magScratch[c] = newSeq[uint8](TMH_NCLAUSES)
|
||||
result.lastTick = -1
|
||||
result.lastEnqTick = -1
|
||||
result.lastEnqBucket = -1
|
||||
result.shiftDeg = TMH_SHIFT_DEFAULT
|
||||
result.bigMult = TMH_BIG_MULT_DEFAULT
|
||||
result.resetOnTarget = TMH_RESET_ON_TARGET_DEFAULT
|
||||
result.observedTargetId = -1
|
||||
result.lastSidePred = -1
|
||||
result.lastMagPred = -1
|
||||
|
||||
proc ensureConfig(g: var TmHorizonGun) {.inline.} =
|
||||
## Lazily read the env knobs on first use (like Pattern's radial knobs), so a
|
||||
## unit test can poke the env in-process and an unset env cannot move the gun.
|
||||
if g.shiftConfigured: return
|
||||
g.shiftDeg = envFloatT(TMH_SHIFT_ENV, TMH_SHIFT_DEFAULT)
|
||||
g.bigMult = envFloatT(TMH_BIG_MULT_ENV, TMH_BIG_MULT_DEFAULT)
|
||||
g.logEnabled = envBoolT(TMH_LOG_ENV, false)
|
||||
g.resetOnTarget = envBoolT(TMH_RESET_ON_TARGET_ENV, TMH_RESET_ON_TARGET_DEFAULT)
|
||||
g.shiftConfigured = true
|
||||
|
||||
proc setShift*(g: var TmHorizonGun, shiftDeg: float, bigMult = TMH_BIG_MULT_DEFAULT) =
|
||||
## Explicit per-gun override (tests / offline sweeps). Writes the same fields
|
||||
## the env path writes, so the measured code path is identical.
|
||||
g.shiftDeg = shiftDeg
|
||||
g.bigMult = bigMult
|
||||
g.shiftConfigured = true
|
||||
|
||||
proc setResetOnTarget*(g: var TmHorizonGun, on: bool) =
|
||||
## Explicit per-gun override (tests). Writes the same field the env path does
|
||||
## and freezes config, so a later `ensureConfig` cannot move it.
|
||||
g.resetOnTarget = on
|
||||
g.shiftConfigured = true
|
||||
|
||||
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
|
||||
## once the bots teleport back to their starting corners. The Tsetlin machines
|
||||
## and their learned statistics are LEFT UNTOUCHED: learning accumulates across
|
||||
## rounds of the same battle/enemy. This is the ONLY reset that runs at a plain
|
||||
## round boundary.
|
||||
g.pendingCount = 0
|
||||
g.bulletCount = 0
|
||||
g.hasPrev = false
|
||||
g.hasSelfPrev = false
|
||||
g.sinceRev = 0
|
||||
g.lastNonzeroSign = 0
|
||||
g.lastTick = -1
|
||||
g.lastEnqTick = -1
|
||||
g.lastEnqBucket = -1
|
||||
g.bitsValid = false
|
||||
g.lastLogKey = ""
|
||||
g.lastLogTick = -1
|
||||
g.lastSidePred = -1
|
||||
g.lastMagPred = -1
|
||||
for i in 0..<TMH_NH:
|
||||
g.bucketEval[i] = TmhEval()
|
||||
for i in 0..<TMH_POS_RING: g.ringValid[i] = false
|
||||
# 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
|
||||
|
||||
proc resetLearning*(g: var TmHorizonGun, reason = "") =
|
||||
## PER-BATTLE / PER-ENEMY wipe. Wipes the Tsetlin machines and every learned
|
||||
## statistic, then clears the per-round state too (a new enemy makes the old
|
||||
## observations and pending labels meaningless). Called on a NEW BATTLE
|
||||
## (`onGameStarted`, with a round-1 fallback) and when the TARGET changes to a
|
||||
## different bot id. NEVER called at a plain round boundary — that is
|
||||
## `resetRoundState`. There is still no persistence across battles.
|
||||
let trainedWas = g.trained
|
||||
inc g.seedCounter
|
||||
g.sideMachine.resetMachine(seed = 1000'u64 + uint64(g.seedCounter))
|
||||
g.magMachine.resetMachine(seed = 2000'u64 + uint64(g.seedCounter))
|
||||
g.trained = 0
|
||||
g.sideCorrect = 0
|
||||
g.sideTotal = 0
|
||||
g.pendingDropped = 0
|
||||
g.observedTargetId = -1
|
||||
for i in 0..1:
|
||||
g.sidePredHist[i] = 0
|
||||
g.magPredHist[i] = 0
|
||||
g.sideLabelHist[i] = 0
|
||||
g.magLabelHist[i] = 0
|
||||
for i in 0..<4: g.quadHist[i] = 0
|
||||
for i in 0..<TMH_ABS_BINS: g.absHist[i] = 0
|
||||
g.absCount = 0
|
||||
g.resetRoundState()
|
||||
if reason.len > 0:
|
||||
g.ensureConfig()
|
||||
if g.logEnabled:
|
||||
echo fmt"[tmh-reset] reason={reason} trained_was={trainedWas}"
|
||||
|
||||
proc targetChanged*(g: var TmHorizonGun, enemyId: int): bool =
|
||||
## Per-ENEMY reset: when the selected target changes to a DIFFERENT bot id,
|
||||
## wipe the machines (the user's reset condition is literally "until the enemy
|
||||
## change"). First acquisition (`observedTargetId < 0`) never wipes, so the
|
||||
## round-start pick does not cold-start the gun every round. Gated by
|
||||
## `TR_TMHORIZON_RESET_ON_TARGET` (default on); returns true when it wiped.
|
||||
g.ensureConfig()
|
||||
if not g.resetOnTarget: return false
|
||||
if enemyId < 0: return false
|
||||
if g.observedTargetId >= 0 and enemyId != g.observedTargetId:
|
||||
g.resetLearning("target_change")
|
||||
g.observedTargetId = enemyId
|
||||
return true
|
||||
g.observedTargetId = enemyId
|
||||
false
|
||||
|
||||
# ── history / observation ring ───────────────────────────────────────────────
|
||||
|
||||
proc ringAt(g: TmHorizonGun, tick: int): tuple[ok: bool, t: TmhTick] =
|
||||
let slot = ((tick mod TMH_POS_RING) + TMH_POS_RING) mod TMH_POS_RING
|
||||
if g.ringValid[slot] and g.ring[slot].tick == tick:
|
||||
(true, g.ring[slot])
|
||||
else:
|
||||
(false, TmhTick())
|
||||
|
||||
proc ringValidCount*(g: TmHorizonGun): int =
|
||||
## Observability / test seam: how many observation-ring slots are valid.
|
||||
for i in 0..<TMH_POS_RING:
|
||||
if g.ringValid[i]: inc result
|
||||
|
||||
proc ringHas*(g: TmHorizonGun, tick: int): bool =
|
||||
## Observability / test seam: is there a valid observation exactly at `tick`?
|
||||
## Used to prove a round boundary drops the old positions (so a label can never
|
||||
## be built from a position across the boundary).
|
||||
g.ringAt(tick).ok
|
||||
|
||||
proc sideClauseStates*(g: TmHorizonGun): seq[int16] =
|
||||
## Test / observability seam: a snapshot of the side head's learned clause
|
||||
## states, so a test can prove the machines survive or are wiped by a reset.
|
||||
g.sideMachine.teams[0] & g.sideMachine.teams[1]
|
||||
|
||||
proc sideClausesAllExclude*(g: TmHorizonGun): bool =
|
||||
## Test seam: true when every side clause is at the Exclude boundary.
|
||||
for c in 0..<g.sideMachine.nClasses:
|
||||
for i in 0..<g.sideMachine.teams[c].len:
|
||||
if g.sideMachine.teams[c][i] != 0: return false
|
||||
true
|
||||
|
||||
proc addBullet(g: var TmHorizonGun, b: TmhBullet) =
|
||||
if g.bulletCount < TMH_BULLETS:
|
||||
g.bullets[g.bulletCount] = b
|
||||
inc g.bulletCount
|
||||
else:
|
||||
for i in 1..<TMH_BULLETS: g.bullets[i - 1] = g.bullets[i]
|
||||
g.bullets[TMH_BULLETS - 1] = b
|
||||
|
||||
proc bulletFeatures(g: var TmHorizonGun, state: WorldState): tuple[tta: int, lat: float] =
|
||||
## Nearest in-flight (proxy) bullet: ticks-to-arrival and the enemy's lateral
|
||||
## offset from that bullet's path. `tta < 0` means none.
|
||||
result.tta = -1
|
||||
result.lat = 0.0
|
||||
var w = 0
|
||||
for i in 0..<g.bulletCount:
|
||||
var b = g.bullets[i]
|
||||
if not b.active: continue
|
||||
if state.tick > b.arrivalTick:
|
||||
b.active = false
|
||||
continue
|
||||
g.bullets[w] = b
|
||||
inc w
|
||||
let ta = b.arrivalTick - state.tick
|
||||
if result.tta < 0 or ta < result.tta:
|
||||
result.tta = ta
|
||||
let vx = state.enemyX - b.sx
|
||||
let vy = state.enemyY - b.sy
|
||||
result.lat = b.ux * vy - b.uy * vx
|
||||
g.bulletCount = w
|
||||
|
||||
proc tmhUpdateHistory(g: var TmHorizonGun, state: WorldState) =
|
||||
## Once per tick: write the observation ring, advance the reversal counter and
|
||||
## the self-energy-drop bullet proxy.
|
||||
let slot = ((state.tick mod TMH_POS_RING) + TMH_POS_RING) mod TMH_POS_RING
|
||||
var lst = state.tick
|
||||
if state.enemies.len > 0: lst = state.enemies[0].lastSeenTick
|
||||
g.ring[slot] = TmhTick(tick: state.tick, ex: state.enemyX, ey: state.enemyY,
|
||||
eh: state.enemyHeading, es: state.enemySpeed,
|
||||
lastSeenTick: lst)
|
||||
g.ringValid[slot] = true
|
||||
|
||||
let sg = signf(state.enemySpeed)
|
||||
if sg != 0:
|
||||
if g.lastNonzeroSign != 0 and sg != g.lastNonzeroSign: g.sinceRev = 0
|
||||
else: inc g.sinceRev
|
||||
g.lastNonzeroSign = sg
|
||||
else:
|
||||
inc g.sinceRev
|
||||
|
||||
if g.hasSelfPrev:
|
||||
let drop = g.selfPrevEnergy - state.selfEnergy
|
||||
if drop > 0.05 and drop <= 3.1:
|
||||
let power = clamp(drop, 0.1, 3.0)
|
||||
let speed = 20.0 - 3.0 * power
|
||||
let dx = state.enemyX - state.selfX
|
||||
let dy = state.enemyY - state.selfY
|
||||
let nrm = max(1e-6, hypot(dx, dy))
|
||||
let flight = int(ceil(nrm / speed))
|
||||
g.addBullet(TmhBullet(arrivalTick: state.tick + flight,
|
||||
ux: dx / nrm, uy: dy / nrm,
|
||||
sx: state.selfX, sy: state.selfY, active: true))
|
||||
g.selfPrevEnergy = state.selfEnergy
|
||||
g.hasSelfPrev = true
|
||||
|
||||
# ── the 49 draft bits, live ──────────────────────────────────────────────────
|
||||
|
||||
proc tmhBaseBits*(g: var TmHorizonGun, state: WorldState): array[TMH_N_BASE, uint8] =
|
||||
## The draftTMSpec() blocks, computed from the live bot state + the observation
|
||||
## ring. Block layout (mirrors draftTMSpec):
|
||||
## 0..3 dist-to-nearest-wall | 4..7 which-wall-nearest
|
||||
## 8..13 dist-from-us | 14..16 enemy-heading-vs-line-to-us
|
||||
## 17..19 turn-direction t,t-1,t-2 | 20..24 ticks-since-reversal
|
||||
## 25..27 turn-consistency-10 | 28..30 distance-moved-10
|
||||
## 31..33 speed-trend-10 | 34..36 turn-rate-change-5
|
||||
## 37..41 time-until-bullet | 42..48 bullet-lateral-offset
|
||||
var b: array[TMH_N_BASE, uint8]
|
||||
template put1(idx: int) = b[idx] = 1'u8
|
||||
|
||||
let ex = state.enemyX
|
||||
let ey = state.enemyY
|
||||
|
||||
# walls (8)
|
||||
let dL = ex
|
||||
let dR = state.arenaWidth - ex
|
||||
let dT = state.arenaHeight - ey
|
||||
let dB = ey
|
||||
let dmin = min(min(dL, dR), min(dT, dB))
|
||||
if dmin < 50.0: put1(0)
|
||||
elif dmin < 100.0: put1(1)
|
||||
elif dmin < 200.0: put1(2)
|
||||
else: put1(3)
|
||||
var wb = 0
|
||||
let walls = [dL, dR, dT, dB]
|
||||
for w in 1..3:
|
||||
if walls[w] < walls[wb]: wb = w
|
||||
put1(4 + wb)
|
||||
|
||||
# us (9)
|
||||
let rng = hypot(ex - state.selfX, ey - state.selfY)
|
||||
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
|
||||
put1(8 + ub)
|
||||
|
||||
let lane = arctan2(state.selfY - ey, state.selfX - ex)
|
||||
let hdg = degToRad(state.enemyHeading)
|
||||
let perp = abs(sin(hdg - lane))
|
||||
var hb = 1
|
||||
if perp < 0.5: hb = 2
|
||||
elif perp > 0.866: hb = 0
|
||||
put1(14 + hb)
|
||||
|
||||
# motion (20)
|
||||
for k in 0..2:
|
||||
let cur = g.ringAt(state.tick - k)
|
||||
let prv = g.ringAt(state.tick - k - 1)
|
||||
if cur.ok and prv.ok:
|
||||
let d = wrapDeg(cur.t.eh - prv.t.eh)
|
||||
if d > 1e-6: put1(17 + k)
|
||||
|
||||
var rb = 4
|
||||
let sr = g.sinceRev
|
||||
if sr < 5: rb = 0
|
||||
elif sr < 10: rb = 1
|
||||
elif sr < 20: rb = 2
|
||||
elif sr < 40: rb = 3
|
||||
put1(20 + rb)
|
||||
|
||||
var pos = 0
|
||||
var neg = 0
|
||||
for k in 0..9:
|
||||
let cur = g.ringAt(state.tick - k)
|
||||
let prv = g.ringAt(state.tick - k - 1)
|
||||
if cur.ok and prv.ok:
|
||||
let d = wrapDeg(cur.t.eh - prv.t.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
|
||||
put1(25 + cb)
|
||||
|
||||
let j10 = g.ringAt(state.tick - 10)
|
||||
let dm = if j10.ok: hypot(ex - j10.t.ex, ey - j10.t.ey) else: 0.0
|
||||
var mb = 1
|
||||
if dm < 20.0: mb = 0
|
||||
elif dm > 50.0: mb = 2
|
||||
put1(28 + mb)
|
||||
|
||||
let spd10 = if j10.ok: abs(j10.t.es) else: abs(state.enemySpeed)
|
||||
let spdDiff = abs(state.enemySpeed) - spd10
|
||||
var sb = 1
|
||||
if spdDiff < -0.5: sb = 0
|
||||
elif spdDiff > 0.5: sb = 2
|
||||
put1(31 + sb)
|
||||
|
||||
var r1 = 0.0
|
||||
var n1 = 0
|
||||
for k in 0..4:
|
||||
let cur = g.ringAt(state.tick - k)
|
||||
let prv = g.ringAt(state.tick - k - 1)
|
||||
if cur.ok and prv.ok:
|
||||
r1 += abs(wrapDeg(cur.t.eh - prv.t.eh)); inc n1
|
||||
var r2 = 0.0
|
||||
var n2 = 0
|
||||
for k in 5..9:
|
||||
let cur = g.ringAt(state.tick - k)
|
||||
let prv = g.ringAt(state.tick - k - 1)
|
||||
if cur.ok and prv.ok:
|
||||
r2 += abs(wrapDeg(cur.t.eh - prv.t.eh)); inc n2
|
||||
let m1 = if n1 > 0: r1 / n1.float else: 0.0
|
||||
let m2 = if n2 > 0: r2 / n2.float else: 0.0
|
||||
let dtr = m1 - m2
|
||||
var tb = 1
|
||||
if dtr < -0.3: tb = 0
|
||||
elif dtr > 0.3: tb = 2
|
||||
put1(34 + tb)
|
||||
|
||||
# bullets (12) — INFERRED from self-energy drops (no gun heading live)
|
||||
let (tta, lat) = g.bulletFeatures(state)
|
||||
var b1 = 0
|
||||
if tta >= 0:
|
||||
if tta < 5: b1 = 1
|
||||
elif tta < 10: b1 = 2
|
||||
elif tta < 20: b1 = 3
|
||||
else: b1 = 4
|
||||
put1(37 + b1)
|
||||
|
||||
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
|
||||
put1(42 + lb)
|
||||
|
||||
result = b
|
||||
|
||||
# ── TM evaluation / training (no per-sample allocation) ──────────────────────
|
||||
|
||||
proc tmhEvalOne(m: TmMachine, lits: array[TMH_NLITS, uint8],
|
||||
caches: var array[2, seq[uint8]]): tuple[cls: int, conf: float] =
|
||||
var votes: array[2, float]
|
||||
for c in 0..1:
|
||||
votes[c] = tmForward(m, m.teams[c], lits, caches[c])
|
||||
let cls = if votes[1] > votes[0]: 1 else: 0
|
||||
let margin = abs(votes[1] - votes[0])
|
||||
let conf =
|
||||
if m.half > 0: clamp(margin / (2.0 * float(m.half)), 0.0, 1.0)
|
||||
else: 0.0
|
||||
(cls, conf)
|
||||
|
||||
proc tmhTrainOne(m: var TmMachine, lits: array[TMH_NLITS, uint8], label: int,
|
||||
caches: var array[2, 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)
|
||||
|
||||
# ── magnitude median (running histogram) ─────────────────────────────────────
|
||||
|
||||
proc recordAbs(g: var TmHorizonGun, a: float) =
|
||||
var idx = int(a / TMH_ABS_RES)
|
||||
if idx < 0: idx = 0
|
||||
if idx >= TMH_ABS_BINS: idx = TMH_ABS_BINS - 1
|
||||
inc g.absHist[idx]
|
||||
inc g.absCount
|
||||
|
||||
proc medianAbs(g: TmHorizonGun): float =
|
||||
if g.absCount == 0: return 0.0
|
||||
let half = g.absCount div 2
|
||||
var cum = 0
|
||||
for i in 0..<TMH_ABS_BINS:
|
||||
cum += g.absHist[i]
|
||||
if cum > half: return i.float * TMH_ABS_RES + TMH_ABS_RES * 0.5
|
||||
63.95
|
||||
|
||||
# ── deferred label resolution ────────────────────────────────────────────────
|
||||
|
||||
proc tmhResolveOne(g: var TmHorizonGun, state: WorldState,
|
||||
p: TmhPending): bool =
|
||||
## Resolve one due sample against our OWN observation at `fireTick + horizon`.
|
||||
## Returns false when the sample must be DROPPED (missing / stale
|
||||
## observation); a perfectly balanced (zero-error) sample is consumed without
|
||||
## training and counts as resolved.
|
||||
let slot = ((state.tick mod TMH_POS_RING) + TMH_POS_RING) mod TMH_POS_RING
|
||||
if not g.ringValid[slot] or g.ring[slot].tick != state.tick: return false
|
||||
if state.tick - g.ring[slot].lastSeenTick > TMH_STALE_MAX: return false
|
||||
let act = g.ring[slot]
|
||||
let ba = arctan2(act.ey - p.selfY, act.ex - p.selfX)
|
||||
let err = wrapRad(ba - p.baseBearing)
|
||||
let absE = abs(err)
|
||||
|
||||
if absE >= 1e-9:
|
||||
let med = g.medianAbs()
|
||||
let sideLabel = if err > 0.0: 1 else: 0
|
||||
let magLabel = if absE > med: 1 else: 0
|
||||
if p.warm:
|
||||
inc g.sideTotal
|
||||
if p.sidePred == sideLabel: inc g.sideCorrect
|
||||
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
|
||||
g.recordAbs(absE)
|
||||
true
|
||||
|
||||
proc tmhResolvePending(g: var TmHorizonGun, state: WorldState) =
|
||||
## Resolve every sample due at this tick, compacting the queue in place.
|
||||
## Samples whose horizon reaches past the round end are simply never due and
|
||||
## are discarded by the next `resetRoundState` (a plain round boundary).
|
||||
var w = 0
|
||||
for i in 0..<g.pendingCount:
|
||||
let p = g.pending[i]
|
||||
let due = p.fireTick + p.horizon
|
||||
if due > state.tick:
|
||||
g.pending[w] = p
|
||||
inc w
|
||||
elif due == state.tick:
|
||||
if not g.tmhResolveOne(state, p):
|
||||
inc g.pendingDropped
|
||||
else:
|
||||
inc g.pendingDropped
|
||||
g.pendingCount = w
|
||||
|
||||
proc tmhEnqueue(g: var TmHorizonGun, state: WorldState, baseBearing: float,
|
||||
h, bucket: int, lits: array[TMH_NLITS, uint8],
|
||||
sidePred, magPred: int, warm: bool) =
|
||||
if g.pendingCount >= TMH_PENDING_CAP:
|
||||
inc g.pendingDropped
|
||||
return
|
||||
g.pending[g.pendingCount] = TmhPending(
|
||||
fireTick: state.tick, horizon: h, bucket: bucket,
|
||||
selfX: state.selfX, selfY: state.selfY, baseBearing: baseBearing,
|
||||
sidePred: sidePred, magPred: magPred, warm: warm, lits: lits)
|
||||
inc g.pendingCount
|
||||
if warm:
|
||||
inc g.sidePredHist[sidePred]
|
||||
inc g.magPredHist[magPred]
|
||||
inc g.quadHist[sidePred * 2 + magPred]
|
||||
|
||||
# ── logging ──────────────────────────────────────────────────────────────────
|
||||
|
||||
proc tmhLog(g: var TmHorizonGun, state: WorldState, bulletSpeed: float,
|
||||
h, side, mag: int, sideConf, shift: float, aimDeg: float) =
|
||||
## ONE change-gated line (behind `TR_TMHORIZON_LOG=1`) so the user can see it
|
||||
## think per shot, not per tick. Format:
|
||||
## [tmh] h=22 p=2.0 pred=LARGE-LEFT shift=-3.0 aim=173.2 gun=Pattern conf=0.62 trained=1841
|
||||
if not g.logEnabled: return
|
||||
let q = tmhQuadrantName(side, mag)
|
||||
# Gate on the PREDICTION (quadrant) and the applied shift — NOT on the horizon,
|
||||
# which changes almost every tick as the range moves and would spam the log.
|
||||
let key = fmt"{q}|{shift:.1f}"
|
||||
if key == g.lastLogKey: return
|
||||
if state.tick == g.lastLogTick: return
|
||||
g.lastLogKey = key
|
||||
g.lastLogTick = state.tick
|
||||
let power = (20.0 - bulletSpeed) / 3.0
|
||||
echo fmt"[tmh] h={h} p={power:.1f} pred={q} shift={shift:.1f} " &
|
||||
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.
|
||||
g.ensureConfig()
|
||||
if not g.logEnabled: return
|
||||
let thisRound = g.trained - g.roundStartTrained
|
||||
let acc = if g.sideTotal > 0: g.sideCorrect.float / g.sideTotal.float * 100.0
|
||||
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"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]}] " &
|
||||
fmt"magPred=[S:{g.magPredHist[0]} B:{g.magPredHist[1]}] " &
|
||||
fmt"labels=[side R:{g.sideLabelHist[0]} L:{g.sideLabelHist[1]} " &
|
||||
fmt"mag S:{g.magLabelHist[0]} B:{g.magLabelHist[1]}]"
|
||||
|
||||
# ── Gun interface ────────────────────────────────────────────────────────────
|
||||
|
||||
proc isWarmedUp*(g: TmHorizonGun): bool {.inline.} = true
|
||||
|
||||
proc predict*(g: var TmHorizonGun, state: WorldState,
|
||||
bulletSpeed: float): GunPrediction =
|
||||
g.ensureConfig()
|
||||
|
||||
# Round boundary: a tick regression means a new round. Clear ONLY the
|
||||
# per-round state (ring + deferred labels); the machines SURVIVE so learning
|
||||
# accumulates across the whole battle.
|
||||
if state.tick < g.lastTick: g.resetRoundState()
|
||||
|
||||
# Once-per-tick: observe the world, then resolve any labels now due.
|
||||
if state.tick != g.lastTick:
|
||||
g.tmhUpdateHistory(state)
|
||||
g.tmhResolvePending(state)
|
||||
g.lastTick = state.tick
|
||||
for i in 0..<TMH_NH: g.bucketEval[i] = TmhEval()
|
||||
|
||||
# The base prediction is Pattern. The TM only corrects it.
|
||||
let base = g.pattern.predict(state, bulletSpeed)
|
||||
if bulletSpeed <= 0.0: return base
|
||||
|
||||
let dist = hypot(state.enemyX - state.selfX, state.enemyY - state.selfY)
|
||||
let h = tmhHorizonFor(dist, bulletSpeed)
|
||||
let bucket = tmhHorizonBucket(h)
|
||||
|
||||
# Base bits cached once per tick (Pattern-style per-tick caching).
|
||||
if not g.bitsValid or g.cachedBitsTick != state.tick:
|
||||
g.cachedBits = g.tmhBaseBits(state)
|
||||
g.cachedBitsTick = state.tick
|
||||
g.bitsValid = true
|
||||
let lits = tmhLits(g.cachedBits, bucket)
|
||||
|
||||
# Evaluate each horizon bucket at most once per tick.
|
||||
if not g.bucketEval[bucket].valid:
|
||||
let (sideCls, sideConf) = tmhEvalOne(g.sideMachine, lits, g.sideScratch)
|
||||
let (magCls, _) = tmhEvalOne(g.magMachine, lits, g.magScratch)
|
||||
g.bucketEval[bucket] = TmhEval(valid: true, side: sideCls,
|
||||
sideConf: sideConf, mag: magCls)
|
||||
let ev = g.bucketEval[bucket]
|
||||
|
||||
let warm = g.trained >= TMH_MIN_OBS
|
||||
let side = if warm: ev.side else: -1
|
||||
let mag = if warm: ev.mag else: -1
|
||||
g.lastSidePred = side
|
||||
g.lastMagPred = mag
|
||||
|
||||
let baseBearing = arctan2(base.y - state.selfY, base.x - state.selfX)
|
||||
|
||||
# Enqueue one sample per (tick, bucket); predict runs once per power bin.
|
||||
if g.lastEnqTick != state.tick or g.lastEnqBucket != bucket:
|
||||
g.tmhEnqueue(state, baseBearing, h, bucket, lits, side, mag, warm)
|
||||
g.lastEnqTick = state.tick
|
||||
g.lastEnqBucket = bucket
|
||||
|
||||
# Small hit-optimal-style correction (NOT the conditional median).
|
||||
var shift = 0.0
|
||||
if warm and g.shiftDeg != 0.0:
|
||||
let magScale = if mag == 1: g.bigMult else: 1.0
|
||||
let sgn = if side == 1: 1.0 else: -1.0
|
||||
shift = sgn * g.shiftDeg * magScale
|
||||
|
||||
let pred =
|
||||
if shift == 0.0: base
|
||||
else: tmhApplyShift(state.selfX, state.selfY, base.x, base.y, shift)
|
||||
|
||||
let aimDeg = radToDeg(arctan2(pred.y - state.selfY, pred.x - state.selfX))
|
||||
g.tmhLog(state, bulletSpeed, h, side, mag, ev.sideConf, shift, aimDeg)
|
||||
pred
|
||||
|
||||
proc onResult*(g: var TmHorizonGun, e: FeedbackEvent) =
|
||||
## The label comes from our own observation ring, not from virtual-bullet
|
||||
## feedback, so there is nothing to do here. The hook exists for the rack.
|
||||
discard
|
||||
@@ -0,0 +1,340 @@
|
||||
## 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."
|
||||
Reference in New Issue
Block a user