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:
@@ -1,5 +1,5 @@
|
|||||||
## ModularBot — plugin gun architecture tracer bullet.
|
## ModularBot — plugin gun architecture tracer bullet.
|
||||||
## Guns: HeadOnGun (0), LinearGun (1), TsetlinGun (2), CircularGun (3), GFGun (4), PatternMatcherGun (5), WallBounceGun (6), AccelGun (7), StopShotGun (8), DisplacementGun (9), AveragedLeadGun (10), DecayGFGun (11), KNNGun (12), TmSelectorGun (13), TmPatternGun (14) via GunHarness.
|
## Guns: HeadOnGun (0), LinearGun (1), TsetlinGun (2), CircularGun (3), GFGun (4), PatternMatcherGun (5), WallBounceGun (6), AccelGun (7), StopShotGun (8), DisplacementGun (9), AveragedLeadGun (10), DecayGFGun (11), KNNGun (12), TmSelectorGun (13), TmPatternGun (14), TmHorizonGun (15) via GunHarness.
|
||||||
## Radar: RadarLockModule (1v1) / AdaptiveMeleeRadarModule (2+ enemies), auto-switched per tick.
|
## Radar: RadarLockModule (1v1) / AdaptiveMeleeRadarModule (2+ enemies), auto-switched per tick.
|
||||||
## Movement: OscillatorModule (perpendicular strafing).
|
## Movement: OscillatorModule (perpendicular strafing).
|
||||||
|
|
||||||
@@ -26,6 +26,7 @@ import guns/decay_gf
|
|||||||
import guns/knn_gun
|
import guns/knn_gun
|
||||||
import guns/tm_selector
|
import guns/tm_selector
|
||||||
import guns/tm_pattern
|
import guns/tm_pattern
|
||||||
|
import guns/tm_horizon
|
||||||
import movements/phantom_meteor
|
import movements/phantom_meteor
|
||||||
import movements/rammer
|
import movements/rammer
|
||||||
import movements/ram_decision
|
import movements/ram_decision
|
||||||
@@ -119,7 +120,7 @@ let ShotLogPath = getEnv("GUN_SHOTLOG_PATH", "/tmp/shot_log.jsonl")
|
|||||||
## shows why the cap moved. The policy itself lives in the shared gun harness
|
## shows why the cap moved. The policy itself lives in the shared gun harness
|
||||||
## (`applyPowerPolicy`), so both live and offline paths see the same rule.
|
## (`applyPowerPolicy`), so both live and offline paths see the same rule.
|
||||||
let PowerLog = existsEnv("TR_POWER_LOG")
|
let PowerLog = existsEnv("TR_POWER_LOG")
|
||||||
const GunNames = ["HeadOn", "Linear", "Tsetlin", "Circular", "GuessFactor", "Pattern", "WallBounce", "Accel", "StopShot", "Displace", "AvgLead", "DecayGF", "KNN", "TMSelect", "TMPattern"]
|
const GunNames = ["HeadOn", "Linear", "Tsetlin", "Circular", "GuessFactor", "Pattern", "WallBounce", "Accel", "StopShot", "Displace", "AvgLead", "DecayGF", "KNN", "TMSelect", "TMPattern", "TMHorizon"]
|
||||||
|
|
||||||
## Rack id of the new TM pattern gun. It defaults to `TR_RACK_TMPATTERN=off`;
|
## Rack id of the new TM pattern gun. It defaults to `TR_RACK_TMPATTERN=off`;
|
||||||
## unlike the other guns, its virtual-bullet spawn is gated on rack admission
|
## unlike the other guns, its virtual-bullet spawn is gated on rack admission
|
||||||
@@ -128,6 +129,12 @@ const GunNames = ["HeadOn", "Linear", "Tsetlin", "Circular", "GuessFactor", "Pat
|
|||||||
## the default selection sequence — is byte-for-byte unchanged.
|
## the default selection sequence — is byte-for-byte unchanged.
|
||||||
const TmPatternId = 14
|
const TmPatternId = 14
|
||||||
|
|
||||||
|
## Rack id of the horizon-based TM corrector. Like TMPATTERN it defaults to
|
||||||
|
## `TR_RACK_TMHORIZON=off` and its virtual-bullet spawn is gated on rack
|
||||||
|
## admission, so the shipped default never spawns it and the shared tracker ring
|
||||||
|
## head — hence every other gun's learning order — is byte-for-byte unchanged.
|
||||||
|
const TmHorizonId = 15
|
||||||
|
|
||||||
const
|
const
|
||||||
CLR_GUN = "\e[33m" # yellow
|
CLR_GUN = "\e[33m" # yellow
|
||||||
CLR_MOVE = "\e[36m" # cyan
|
CLR_MOVE = "\e[36m" # cyan
|
||||||
@@ -174,6 +181,7 @@ type
|
|||||||
knnGun: KNNGun
|
knnGun: KNNGun
|
||||||
tmSelector: TmSelectorGun
|
tmSelector: TmSelectorGun
|
||||||
tmPattern: TmPatternGun
|
tmPattern: TmPatternGun
|
||||||
|
tmHorizon: TmHorizonGun
|
||||||
mover: TFILModule
|
mover: TFILModule
|
||||||
ringMover: TFILRingModule
|
ringMover: TFILRingModule
|
||||||
rammer: RammerModule
|
rammer: RammerModule
|
||||||
@@ -199,18 +207,18 @@ type
|
|||||||
roundNumber: int
|
roundNumber: int
|
||||||
realShotsFired: int
|
realShotsFired: int
|
||||||
realHits: int
|
realHits: int
|
||||||
gunRealShots: array[15, int]
|
gunRealShots: array[16, int]
|
||||||
gunRealHits: array[15, int]
|
gunRealHits: array[16, int]
|
||||||
# Same real-shot accounting split by the rack in force at fire time, so a
|
# Same real-shot accounting split by the rack in force at fire time, so a
|
||||||
# later data-driven pass can rank guns per mode (1v1 vs melee).
|
# later data-driven pass can rank guns per mode (1v1 vs melee).
|
||||||
gunRealShotsByMode: array[vb.RackMode, array[15, int]]
|
gunRealShotsByMode: array[vb.RackMode, array[16, int]]
|
||||||
gunRealHitsByMode: array[vb.RackMode, array[15, int]]
|
gunRealHitsByMode: array[vb.RackMode, array[16, int]]
|
||||||
pendingFires: seq[PendingShot] ## FIFO of fired shots awaiting onBulletFired bulletId stamp
|
pendingFires: seq[PendingShot] ## FIFO of fired shots awaiting onBulletFired bulletId stamp
|
||||||
bulletGun: Table[int, int] ## bulletId -> gun id, filled on onBulletFired, drained on resolution
|
bulletGun: Table[int, int] ## bulletId -> gun id, filled on onBulletFired, drained on resolution
|
||||||
bulletMode: Table[int, vb.RackMode] ## bulletId -> rack at fire time
|
bulletMode: Table[int, vb.RackMode] ## bulletId -> rack at fire time
|
||||||
bulletShot: Table[int, PendingShot] ## bulletId -> shot metadata (Task A shot log)
|
bulletShot: Table[int, PendingShot] ## bulletId -> shot metadata (Task A shot log)
|
||||||
pendingHitBullets: HashSet[int] ## hit bulletIds seen before their onBulletFired stamp
|
pendingHitBullets: HashSet[int] ## hit bulletIds seen before their onBulletFired stamp
|
||||||
gunSelectionCount: array[15, int]
|
gunSelectionCount: array[16, int]
|
||||||
lastPowerLogKey: string ## change detector for the TR_POWER_LOG line
|
lastPowerLogKey: string ## change detector for the TR_POWER_LOG line
|
||||||
lastKnownTargetId: int ## persists through death, used for round-end stats
|
lastKnownTargetId: int ## persists through death, used for round-end stats
|
||||||
# Radar measurement instrumentation (only touched when RadarScanLog is set).
|
# Radar measurement instrumentation (only touched when RadarScanLog is set).
|
||||||
@@ -569,6 +577,9 @@ method onRoundEnded*(bot: ModularBot, e: RoundEndedEventForBot) =
|
|||||||
## Dump per-gun virtual bullet stats to /tmp/gun_stats.jsonl (one line per round).
|
## Dump per-gun virtual bullet stats to /tmp/gun_stats.jsonl (one line per round).
|
||||||
if RecordWorldState:
|
if RecordWorldState:
|
||||||
bot.finishWorldStateRecord()
|
bot.finishWorldStateRecord()
|
||||||
|
# Per-round TM-horizon summary (change-gated behind TR_TMHORIZON_LOG=1) so the
|
||||||
|
# user can watch it learn across the round.
|
||||||
|
bot.tmHorizon.roundSummary()
|
||||||
# Use lastKnownTargetId: currentTargetId is -1 if enemy died before round end
|
# Use lastKnownTargetId: currentTargetId is -1 if enemy died before round end
|
||||||
let targetId = if bot.currentTargetId >= 0: bot.currentTargetId else: bot.lastKnownTargetId
|
let targetId = if bot.currentTargetId >= 0: bot.currentTargetId else: bot.lastKnownTargetId
|
||||||
# Per-target fitness when known, else a deterministic recency-weighted aggregate
|
# Per-target fitness when known, else a deterministic recency-weighted aggregate
|
||||||
@@ -577,7 +588,7 @@ method onRoundEnded*(bot: ModularBot, e: RoundEndedEventForBot) =
|
|||||||
let fit = bot.tracker.fitnessFor(targetId)
|
let fit = bot.tracker.fitnessFor(targetId)
|
||||||
|
|
||||||
var gunsArr = newJArray()
|
var gunsArr = newJArray()
|
||||||
for gid in 0..<15:
|
for gid in 0..<16:
|
||||||
var totalShots = 0
|
var totalShots = 0
|
||||||
var totalHits = 0
|
var totalHits = 0
|
||||||
for binIdx in 0..<len(vb.PowerBins):
|
for binIdx in 0..<len(vb.PowerBins):
|
||||||
@@ -690,12 +701,12 @@ method onRoundStarted*(bot: ModularBot, e: RoundStartedEvent) =
|
|||||||
bot.radarAcquireTicks = 0
|
bot.radarAcquireTicks = 0
|
||||||
bot.radarTrackTicks = 0
|
bot.radarTrackTicks = 0
|
||||||
bot.radarMeleeActive = false
|
bot.radarMeleeActive = false
|
||||||
for i in 0..<15:
|
for i in 0..<16:
|
||||||
bot.gunSelectionCount[i] = 0
|
bot.gunSelectionCount[i] = 0
|
||||||
bot.gunRealShots[i] = 0
|
bot.gunRealShots[i] = 0
|
||||||
bot.gunRealHits[i] = 0
|
bot.gunRealHits[i] = 0
|
||||||
for m in vb.RackMode:
|
for m in vb.RackMode:
|
||||||
for i in 0..<15:
|
for i in 0..<16:
|
||||||
bot.gunRealShotsByMode[m][i] = 0
|
bot.gunRealShotsByMode[m][i] = 0
|
||||||
bot.gunRealHitsByMode[m][i] = 0
|
bot.gunRealHitsByMode[m][i] = 0
|
||||||
# Reset per-round integrity counters so each /tmp/gun_stats.jsonl line reports
|
# Reset per-round integrity counters so each /tmp/gun_stats.jsonl line reports
|
||||||
@@ -720,6 +731,14 @@ method onRoundStarted*(bot: ModularBot, e: RoundStartedEvent) =
|
|||||||
# its clause teams and motion history explicitly, so a same-id opponent in the
|
# its clause teams and motion history explicitly, so a same-id opponent in the
|
||||||
# next round cannot inherit the previous round's overfit net.
|
# next round cannot inherit the previous round's overfit net.
|
||||||
bot.tmPattern.resetLearning()
|
bot.tmPattern.resetLearning()
|
||||||
|
# The horizon TM corrector KEEPS its machines across rounds: only the
|
||||||
|
# per-round observation ring and deferred labels are cleared (bots teleport
|
||||||
|
# between rounds, so old positions are meaningless). The machines are wiped on
|
||||||
|
# a NEW BATTLE in onGameStarted, with a round-1 fallback here in case that
|
||||||
|
# callback does not fire in this harness.
|
||||||
|
bot.tmHorizon.resetRoundState()
|
||||||
|
if e.roundNumber <= 1:
|
||||||
|
bot.tmHorizon.resetLearning("round1")
|
||||||
bot.isRamming = false
|
bot.isRamming = false
|
||||||
bot.ramDurationTicks = 0
|
bot.ramDurationTicks = 0
|
||||||
bot.ramCooldownTicks = 0
|
bot.ramCooldownTicks = 0
|
||||||
@@ -773,6 +792,11 @@ method onBotDeath*(bot: ModularBot, e: BotDeathEvent) =
|
|||||||
bot.currentTargetId = -1
|
bot.currentTargetId = -1
|
||||||
|
|
||||||
method onGameStarted*(bot: ModularBot, e: GameStartedEventForBot) =
|
method onGameStarted*(bot: ModularBot, e: GameStartedEventForBot) =
|
||||||
|
# A NEW BATTLE begins: wipe the horizon TM's machines and learned statistics
|
||||||
|
# (no persistence across battles). `onRoundStarted` only clears per-round
|
||||||
|
# state, so this is the only place the learning is dropped at a battle
|
||||||
|
# boundary. The round-1 fallback in `onRoundStarted` covers a missed callback.
|
||||||
|
bot.tmHorizon.resetLearning("game_start")
|
||||||
# minNumberOfParticipants == maxNumberOfParticipants for fixed battles; self is -1
|
# minNumberOfParticipants == maxNumberOfParticipants for fixed battles; self is -1
|
||||||
bot.initialEnemyCount = e.gameSetup.minNumberOfParticipants - 1
|
bot.initialEnemyCount = e.gameSetup.minNumberOfParticipants - 1
|
||||||
# The custom event must be registered AFTER `start()` ran `initGlobals()`,
|
# The custom event must be registered AFTER `start()` ran `initGlobals()`,
|
||||||
@@ -863,6 +887,10 @@ method run*(bot: ModularBot) =
|
|||||||
bot.currentTargetId = candidateId
|
bot.currentTargetId = candidateId
|
||||||
bot.targetSwitchTick = bot.tick
|
bot.targetSwitchTick = bot.tick
|
||||||
bot.cfgDirty = true # emit at tick end, after gun selection
|
bot.cfgDirty = true # emit at tick end, after gun selection
|
||||||
|
# The horizon TM learns ONE enemy at a time: wipe its machines when the
|
||||||
|
# target changes to a different bot id (gated by
|
||||||
|
# TR_TMHORIZON_RESET_ON_TARGET, default on; no-op in 1v1).
|
||||||
|
discard bot.tmHorizon.targetChanged(candidateId)
|
||||||
if bot.currentTargetId >= 0:
|
if bot.currentTargetId >= 0:
|
||||||
bot.lastKnownTargetId = bot.currentTargetId
|
bot.lastKnownTargetId = bot.currentTargetId
|
||||||
|
|
||||||
@@ -998,6 +1026,7 @@ method run*(bot: ModularBot) =
|
|||||||
var knnPreds: array[len(PowerBins), GunPrediction]
|
var knnPreds: array[len(PowerBins), GunPrediction]
|
||||||
var tmselPreds: array[len(PowerBins), GunPrediction]
|
var tmselPreds: array[len(PowerBins), GunPrediction]
|
||||||
var tmpPreds: array[len(PowerBins), GunPrediction]
|
var tmpPreds: array[len(PowerBins), GunPrediction]
|
||||||
|
var tmhPreds: array[len(PowerBins), GunPrediction]
|
||||||
# ── TR_VBULLET_ADMIT_ONLY gate ──────────────────────────────────────
|
# ── TR_VBULLET_ADMIT_ONLY gate ──────────────────────────────────────
|
||||||
# Rack membership used to filter only SELECTION, so every unselected gun
|
# Rack membership used to filter only SELECTION, so every unselected gun
|
||||||
# still ran predict()+spawnBullets() each tick to feed a fitness table the
|
# still ran predict()+spawnBullets() each tick to feed a fitness table the
|
||||||
@@ -1007,12 +1036,13 @@ method run*(bot: ModularBot) =
|
|||||||
# thinning to 1v1). Bullets already in flight are NOT cancelled: they are
|
# thinning to 1v1). Bullets already in flight are NOT cancelled: they are
|
||||||
# resolved by `tickBullets` and the feedback `case` below still calls the
|
# resolved by `tickBullets` and the feedback `case` below still calls the
|
||||||
# owning gun's onResult, so attribution survives a mid-round rack change.
|
# owning gun's onResult, so attribution survives a mid-round rack change.
|
||||||
# TMPATTERN keeps its own admission gate even when the knob is 0, so
|
# TMPATTERN and TMHORIZON keep their own admission gate even when the knob
|
||||||
# TR_VBULLET_ADMIT_ONLY=0 reproduces the exact pre-change rack.
|
# is 0, so TR_VBULLET_ADMIT_ONLY=0 reproduces the exact pre-change rack.
|
||||||
var admit: array[15, bool]
|
var admit: array[16, bool]
|
||||||
for gi in 0..<15:
|
for gi in 0..<16:
|
||||||
admit[gi] = vBulletAdmitted(gi, bot.rackMode, ActiveRackMembership,
|
admit[gi] = vBulletAdmitted(gi, bot.rackMode, ActiveRackMembership,
|
||||||
VBulletAdmitOnly or gi == TmPatternId)
|
VBulletAdmitOnly or gi == TmPatternId or
|
||||||
|
gi == TmHorizonId)
|
||||||
for i in 0..<len(PowerBins):
|
for i in 0..<len(PowerBins):
|
||||||
if admit[0]: headsUp[i] = bot.headOn.predict(bot.lastState, bulletSpeed(PowerBins[i]))
|
if admit[0]: headsUp[i] = bot.headOn.predict(bot.lastState, bulletSpeed(PowerBins[i]))
|
||||||
if admit[1]: linPreds[i] = bot.linear.predict(bot.lastState, bulletSpeed(PowerBins[i]))
|
if admit[1]: linPreds[i] = bot.linear.predict(bot.lastState, bulletSpeed(PowerBins[i]))
|
||||||
@@ -1032,6 +1062,9 @@ method run*(bot: ModularBot) =
|
|||||||
# TMPATTERN keeps its own rack-admission gate (see `admit` above).
|
# TMPATTERN keeps its own rack-admission gate (see `admit` above).
|
||||||
if admit[TmPatternId]:
|
if admit[TmPatternId]:
|
||||||
tmpPreds[i] = bot.tmPattern.predict(bot.lastState, bulletSpeed(PowerBins[i]))
|
tmpPreds[i] = bot.tmPattern.predict(bot.lastState, bulletSpeed(PowerBins[i]))
|
||||||
|
# TMHORIZON likewise: Pattern base + horizon TM correction.
|
||||||
|
if admit[TmHorizonId]:
|
||||||
|
tmhPreds[i] = bot.tmHorizon.predict(bot.lastState, bulletSpeed(PowerBins[i]))
|
||||||
|
|
||||||
if admit[0]: bot.tracker.spawnBullets(0, headsUp, bot.lastState, tid)
|
if admit[0]: bot.tracker.spawnBullets(0, headsUp, bot.lastState, tid)
|
||||||
if admit[1] and not gunDisabled(1): bot.tracker.spawnBullets(1, linPreds, bot.lastState, tid)
|
if admit[1] and not gunDisabled(1): bot.tracker.spawnBullets(1, linPreds, bot.lastState, tid)
|
||||||
@@ -1051,6 +1084,8 @@ method run*(bot: ModularBot) =
|
|||||||
bot.tracker.spawnBullets(13, tmselPreds, bot.lastState, tid)
|
bot.tracker.spawnBullets(13, tmselPreds, bot.lastState, tid)
|
||||||
if admit[TmPatternId] and not gunDisabled(TmPatternId):
|
if admit[TmPatternId] and not gunDisabled(TmPatternId):
|
||||||
bot.tracker.spawnBullets(TmPatternId, tmpPreds, bot.lastState, tid)
|
bot.tracker.spawnBullets(TmPatternId, tmpPreds, bot.lastState, tid)
|
||||||
|
if admit[TmHorizonId] and not gunDisabled(TmHorizonId):
|
||||||
|
bot.tracker.spawnBullets(TmHorizonId, tmhPreds, bot.lastState, tid)
|
||||||
|
|
||||||
# Build slim enemy table for tickBullets
|
# Build slim enemy table for tickBullets
|
||||||
var enemyPositions: Table[int, tuple[x, y: float, lastSeenTick: int, alive: bool]]
|
var enemyPositions: Table[int, tuple[x, y: float, lastSeenTick: int, alive: bool]]
|
||||||
@@ -1080,6 +1115,7 @@ method run*(bot: ModularBot) =
|
|||||||
of 12: bot.knnGun.onResult(fe)
|
of 12: bot.knnGun.onResult(fe)
|
||||||
of 13: bot.tmSelector.onResult(fe)
|
of 13: bot.tmSelector.onResult(fe)
|
||||||
of TmPatternId: bot.tmPattern.onResult(fe)
|
of TmPatternId: bot.tmPattern.onResult(fe)
|
||||||
|
of TmHorizonId: bot.tmHorizon.onResult(fe)
|
||||||
else: discard
|
else: discard
|
||||||
if fe.hit: inc bot.virtualHits else: inc bot.virtualMiss
|
if fe.hit: inc bot.virtualHits else: inc bot.virtualMiss
|
||||||
when DebugVBullets:
|
when DebugVBullets:
|
||||||
@@ -1122,6 +1158,7 @@ method run*(bot: ModularBot) =
|
|||||||
of 12: setTurretColor("#CC00CC"); setBulletColor("#FF44FF")
|
of 12: setTurretColor("#CC00CC"); setBulletColor("#FF44FF")
|
||||||
of 13: setTurretColor("#00FFCC"); setBulletColor("#66FFDD")
|
of 13: setTurretColor("#00FFCC"); setBulletColor("#66FFDD")
|
||||||
of TmPatternId: setTurretColor("#AAFF00"); setBulletColor("#CCFF66")
|
of TmPatternId: setTurretColor("#AAFF00"); setBulletColor("#CCFF66")
|
||||||
|
of TmHorizonId: setTurretColor("#00AAFF"); setBulletColor("#66CCFF")
|
||||||
else: discard
|
else: discard
|
||||||
|
|
||||||
let pred = case selectedGun
|
let pred = case selectedGun
|
||||||
@@ -1139,6 +1176,7 @@ method run*(bot: ModularBot) =
|
|||||||
of 12: bot.knnGun.predict(bot.lastState, bulletSpeed(power))
|
of 12: bot.knnGun.predict(bot.lastState, bulletSpeed(power))
|
||||||
of 13: bot.tmSelector.predict(bot.lastState, bulletSpeed(power))
|
of 13: bot.tmSelector.predict(bot.lastState, bulletSpeed(power))
|
||||||
of TmPatternId: bot.tmPattern.predict(bot.lastState, bulletSpeed(power))
|
of TmPatternId: bot.tmPattern.predict(bot.lastState, bulletSpeed(power))
|
||||||
|
of TmHorizonId: bot.tmHorizon.predict(bot.lastState, bulletSpeed(power))
|
||||||
else: bot.headOn.predict(bot.lastState, bulletSpeed(power))
|
else: bot.headOn.predict(bot.lastState, bulletSpeed(power))
|
||||||
let aimTarget = aimAngle(getX(), getY(), pred.x, pred.y)
|
let aimTarget = aimAngle(getX(), getY(), pred.x, pred.y)
|
||||||
|
|
||||||
@@ -1193,7 +1231,7 @@ proc seedSelectorRng() =
|
|||||||
|
|
||||||
when isMainModule:
|
when isMainModule:
|
||||||
var bot = ModularBot(
|
var bot = ModularBot(
|
||||||
tracker: vb.initTracker(15), # 0: HeadOn, 1: Linear, 2: Tsetlin, 3: Circular, 4: GuessFactor, 5: Pattern, 6: WallBounce, 7: Accel, 8: StopShot, 9: Displace, 10: AvgLead, 11: DecayGF, 12: KNN, 13: TMSelect, 14: TMPattern
|
tracker: vb.initTracker(16), # 0: HeadOn, 1: Linear, 2: Tsetlin, 3: Circular, 4: GuessFactor, 5: Pattern, 6: WallBounce, 7: Accel, 8: StopShot, 9: Displace, 10: AvgLead, 11: DecayGF, 12: KNN, 13: TMSelect, 14: TMPattern, 15: TMHorizon
|
||||||
headOn: HeadOnGun(),
|
headOn: HeadOnGun(),
|
||||||
linear: LinearGun(),
|
linear: LinearGun(),
|
||||||
circular: CircularGun(),
|
circular: CircularGun(),
|
||||||
@@ -1209,6 +1247,7 @@ when isMainModule:
|
|||||||
knnGun: initKNNGun(),
|
knnGun: initKNNGun(),
|
||||||
tmSelector: initTmSelectorGun(),
|
tmSelector: initTmSelectorGun(),
|
||||||
tmPattern: initTmRadialGun(),
|
tmPattern: initTmRadialGun(),
|
||||||
|
tmHorizon: initTmHorizonGun(),
|
||||||
radar: RadarLockModule(),
|
radar: RadarLockModule(),
|
||||||
meleeRadar: initAdaptiveMeleeRadar(),
|
meleeRadar: initAdaptiveMeleeRadar(),
|
||||||
mover: TFILModule(debugGraphics: true),
|
mover: TFILModule(debugGraphics: true),
|
||||||
|
|||||||
@@ -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