b0654d18eb
The user pushed back on "the TM can't be your best 1v1 gun", correctly, because two
decisive tests had never been run. Both are now run and they agree.
TASK 1 - THE GF HEAD vs ITS MAJORITY-CLASS BASELINE (offline, n=1,751,067):
label histogram [254286, 284578, 678879, 297055, 236269]
majority class = 2 (the CENTRE bucket) = 38.77%
RAW head accuracy = 36.69% -> margin **-2.08 pp, BELOW majority**
GATED head accuracy = 40.37% vs 38.75% majority -> +1.62 pp, BUT it predicts the
majority class on 62.4% of ticks and its minority recall is 13.6% / 12.9% - a
base-rate predictor wearing a classifier's clothes.
Shuffled control sits at its own majority (20.04% vs 20.12%), confirming chance.
**THE OLD "46% vs 20% CHANCE" FIGURE I QUOTED WAS WRONG ON TWO COUNTS:** the
baseline is 38.8%, not 20%, and the 46% predated the deferred-label fix. Against
the correct baseline the head is BELOW it.
TASK 2 - THE FIRST-EVER LIVE A/B OF THE TM GUN (7 runs x 7 rounds per arm, one
frozen binary from git archive HEAD = eb74f9b2, sha256 cb66d66b..., real DrussGT,
every arm forced alone with TR_RACK_<GUN>=both and all 14 others off, liveness
confirmed per run):
arm shots real % dmg/run round wins
onlyPattern 4610 10.74% 285 25/49
onlyTMPATTERN (radial) 3374 3.50% 71 0/49
onlyLinear 3218 3.23% 61 0/49
Pattern vs TM: +7.22 pp / +213.7 dmg, exact p=0.0006
TM vs Linear: +0.30 pp, p=0.659 (dmg p=0.438)
**The TM is statistically INDISTINGUISHABLE from its own Linear base live.** So it
is not "the TM works and we are aiming it wrong".
DIRECT ANSWER: **(c) It loses live AND sits at/below majority - the target carries
no learnable signal beyond the base rate, and that is the reason.** The reason is
not the machine, not the knobs, and not the application alone: the thing it was
asked to predict is dominated by the modal answer.
This closes the TM-as-gun thread. If a TM is wanted in the bot, a firing gate or a
movement decision is a better fit for a boolean-rule classifier than an aim point -
that is untested and is a different project.
A LIVE GF-MODE ARM WAS NOT RUN (stated as unmeasured): the task pinned one frozen
HEAD binary and HEAD registers the TM gun as radial only; Task 1 already makes GF
the unpromising candidate.
HARNESS FIX WORTH KEEPING: `tools/ab/which_gun_arm_env.sh` left the TARGET gun
unset, so with the now-Pattern-only default it silently fell back to the FULL rack
- an arm could appear to test a single gun while actually running the whole rack.
It now emits `TR_RACK_<GUN>=both` for the target and `=off` for all 14 others.
(Earlier which-gun results are unaffected: they ran before the Pattern-only default,
or - as in the melee/1v1 campaign - set the explicit `=both` themselves.)
tm_pattern.nim gains a per-class confusion matrix (warm samples only) to support the
majority baseline; no behaviour change. Adds Round 4 to
tm_pattern_sweep_results.md with both tasks and the interpretation rule.
719 lines
29 KiB
Nim
719 lines
29 KiB
Nim
## TM pattern gun — a DISCRETE-target Tsetlin Machine on top of a self-consistent
|
|
## linear forecast.
|
|
##
|
|
## Why this exists (and why it is not `guns/tsetlin.nim`): a previous sweep
|
|
## (common_libs/tests/sweep_tsetlin.nim, documented in the tsetlin.nim header)
|
|
## found that the old gun's TM used as a pixel-correction REGRESSOR learns
|
|
## nothing — every config was indistinguishable from its own shuffled-feedback
|
|
## control, and the whole deficit vs Linear lived in the old gun's one-shot
|
|
## internal baseline. The conclusion was that the REPRESENTATION and the TARGET
|
|
## were the problem, not the knobs.
|
|
##
|
|
## This gun attacks both:
|
|
##
|
|
## * BASE — `forecastLinear` from lead_forecast.nim, the exact self-consistent
|
|
## constant-velocity forecast LinearGun uses, so GF class 0 (center)
|
|
## reproduces the Linear gun byte-for-byte. Any measured difference
|
|
## is attributable to the TM, not to a weaker baseline.
|
|
## * TARGET — the discrete class is a GUESS-FACTOR BUCKET: which lateral
|
|
## escape sector (in max-escape-angle units) the enemy occupied at
|
|
## the tick our bullet would have arrived. A small multi-class
|
|
## classification, which is what a Tsetlin Machine is for.
|
|
## * LABEL — computed from the enemy position at the BASE forecast's arrival
|
|
## tick, looked up in a per-tick position ring. This is deliberate:
|
|
## under the shipped `bmPath` metric `FeedbackEvent.actualXY` is the
|
|
## closest-approach point on the gun's OWN aim ray, which biases the
|
|
## label toward the gun's own last output (a self-referential
|
|
## feedback loop). Reading our own recorded history at the base
|
|
## arrival tick gives a clean, metric-independent label.
|
|
## * FEATURES — binary/bucketed motion context (lateral-velocity sign over 3
|
|
## ticks, turn-rate sign, time since reversal, radial fraction,
|
|
## speed/distance/flight-time/wall/energy bands). 40 bits.
|
|
##
|
|
## The TM core is a compact, self-contained Granmo Table 2/3 implementation with
|
|
## the corrected feedback rules and Eq. 6 empty-clause bootstrap (the same
|
|
## corrected core as tsetlin.nim / tm_selector.nim, re-derived here at a small
|
|
## feature width so a 5-class team is cheap).
|
|
##
|
|
## Per-enemy specialisation: the net is FRESH per gun instance (battle) and the
|
|
## gun resets it if the target id changes mid-battle. Offline the range replays
|
|
## one round per fresh instance, which is exactly "cold every battle, overfit
|
|
## within the battle".
|
|
|
|
import std/[math, random, strutils]
|
|
import gun_harness/gun_interface
|
|
import gun_harness/virtual_bullets as vb
|
|
import guns/lead_forecast
|
|
|
|
const
|
|
## ── classifier shape ────────────────────────────────────────────────────
|
|
TM_CLASSES* {.intdefine.} = 5 ## GF buckets, centers -1, -0.5, 0, +0.5, +1
|
|
TM_NBITS* = 40
|
|
TM_NLITS* = TM_NBITS * 2
|
|
TM_NCLAUSES* {.intdefine.} = 40 ## per class
|
|
TM_HALF* = TM_NCLAUSES div 2
|
|
TM_NSTATES* {.intdefine.} = 64 ## automaton range [-NSTATES, NSTATES]
|
|
TM_T* = float(TM_HALF)
|
|
TM_S_DEF {.strdefine.} = "3.0"
|
|
TM_S* = parseFloat(TM_S_DEF)
|
|
TM_MIN_OBS* {.intdefine.} = 24 ## freshly-cold -> look straight ahead
|
|
## Confidence gate: only leave the centre (GF=0) bucket when the winning class
|
|
## beats the centre class by this fraction of TM_T. 0.0 = raw argmax.
|
|
TM_CONF_MARGIN_DEF {.strdefine.} = "0.0"
|
|
TM_CONF_MARGIN* = parseFloat(TM_CONF_MARGIN_DEF)
|
|
## Scale applied to the predicted GF when a correction is taken.
|
|
TM_SHRINK_DEF {.strdefine.} = "1.0"
|
|
TM_SHRINK* = parseFloat(TM_SHRINK_DEF)
|
|
## Readout: "hard" argmax (with the confidence gate) or "soft" vote-weighted.
|
|
TM_GF_MODE {.strdefine.} = "hard"
|
|
TM_SOFT_BETA_DEF {.strdefine.} = "4.0"
|
|
TM_SOFT_BETA* = parseFloat(TM_SOFT_BETA_DEF)
|
|
TM_TRACE_SLOTS = 1024
|
|
POS_RING = 512
|
|
## Deferred-label queue (Task 2): a virtual bullet whose radial correction
|
|
## aimed SHORT resolves BEFORE its BASE arrival tick, when the arrival-tick
|
|
## position is not yet in `posRing`. Instead of dropping the sample
|
|
## (`labelMisses`), the trace is copied here and resolved on the first later
|
|
## `predict` tick at which the base arrival tick's position exists, so every
|
|
## fired virtual bullet contributes an unbiased training sample.
|
|
TM_PENDING_SLOTS = 1024
|
|
DebugTMPattern* = false
|
|
## ── radial head (Task 2) ────────────────────────────────────────────────
|
|
## Radial label = (enemy radius at the BASE arrival tick) - (base fire
|
|
## distance), bucketed over +/-TM_RADIAL_RANGE px. Readout advances/retards
|
|
## the aim distance along the base bearing.
|
|
TM_RADIAL_RANGE_DEF {.strdefine.} = "60.0"
|
|
TM_RADIAL_RANGE* = parseFloat(TM_RADIAL_RANGE_DEF)
|
|
TM_RAD_MARGIN_DEF {.strdefine.} = "0.25"
|
|
TM_RAD_MARGIN* = parseFloat(TM_RAD_MARGIN_DEF)
|
|
## ── reversal head (Task 3) ──────────────────────────────────────────────
|
|
## Binary: did the enemy's heading turn direction over the flight oppose the
|
|
## direction it was turning at fire time?
|
|
TM_REV_TURN_DEG_DEF {.strdefine.} = "10.0"
|
|
TM_REV_TURN_DEG* = parseFloat(TM_REV_TURN_DEG_DEF)
|
|
TM_REV_MARGIN_DEF {.strdefine.} = "0.0"
|
|
TM_REV_MARGIN* = parseFloat(TM_REV_MARGIN_DEF)
|
|
TM_REV_GAIN_DEF {.strdefine.} = "1.0"
|
|
TM_REV_GAIN* = parseFloat(TM_REV_GAIN_DEF)
|
|
|
|
type
|
|
TmTargetMode* = enum
|
|
tmGF ## round-1: lateral guess-factor bucket
|
|
tmRadial ## Task 2: radial displacement bucket (aim-distance correction)
|
|
tmReversal ## Task 3: binary turn reversal; flips the GF correction sign
|
|
|
|
TmBits* = array[TM_NLITS, uint8]
|
|
|
|
TmPatternTrace = object
|
|
fireTick: int
|
|
powerBin: int
|
|
arrivalTick: int
|
|
baseBearing: float
|
|
fireX, fireY: float
|
|
fireHeading: float
|
|
fireTurn: int
|
|
fireDist: float
|
|
lits: TmBits
|
|
votes: array[TM_CLASSES, float]
|
|
cache: array[TM_CLASSES, array[TM_NCLAUSES, uint8]]
|
|
chosen: int
|
|
radVotes: array[TM_CLASSES, float]
|
|
radCache: array[TM_CLASSES, array[TM_NCLAUSES, uint8]]
|
|
radChosen: int
|
|
revVotes: array[2, float]
|
|
revCache: array[2, array[TM_NCLAUSES, uint8]]
|
|
revChosen: int
|
|
warm: bool
|
|
alive: bool
|
|
|
|
PosSample = object
|
|
tick: int
|
|
x, y: float
|
|
heading: float
|
|
valid: bool
|
|
|
|
PendingResolve = object
|
|
## A fired virtual bullet whose label was not yet resolvable at resolution
|
|
## time. `trace` is a COPY of the fire-time trace (features + clause
|
|
## caches), so the deferred training update is identical to an immediate
|
|
## one, just later.
|
|
arrivalTick: int
|
|
powerBin: int
|
|
power: float
|
|
trace: TmPatternTrace
|
|
|
|
TmPatternGun* = object
|
|
teams: array[TM_CLASSES, seq[int16]]
|
|
radTeams: array[TM_CLASSES, seq[int16]]
|
|
revTeams: array[2, seq[int16]]
|
|
targetMode*: TmTargetMode
|
|
traces: array[TM_TRACE_SLOTS, TmPatternTrace]
|
|
# ── deferred labels (Task 2) ──
|
|
pending: array[TM_PENDING_SLOTS, PendingResolve]
|
|
pendingCount: int
|
|
pendingDropped*: int
|
|
# ── history ──
|
|
posRing: array[POS_RING, PosSample]
|
|
lastTick: int
|
|
prevTick: int
|
|
prevX, prevY, prevHeading: float
|
|
hasPrev: bool
|
|
latSignHist: array[3, int]
|
|
turnSignHist: array[3, int]
|
|
lastNonzeroLat: int
|
|
sinceReversal: int
|
|
radialFracSm: float
|
|
latPersist: int
|
|
currentTarget: int
|
|
# ── instrumentation ──
|
|
totalObs*: int
|
|
predictCalls*: int
|
|
trainCalls*: int
|
|
traceMisses*: int
|
|
labelMisses*: int
|
|
chosenHist*: array[TM_CLASSES, int]
|
|
labelHist*: array[TM_CLASSES, int]
|
|
radChosenHist*: array[TM_CLASSES, int]
|
|
radLabelHist*: array[TM_CLASSES, int]
|
|
revChosenHist*: array[2, int]
|
|
revLabelHist*: array[2, int]
|
|
radCorrect*, radTotal*: int
|
|
revCorrect*, revTotal*: int
|
|
## Radial-target instrumentation (ablation support): raw label statistics
|
|
## and the mean APPLIED aim-distance shift, so a constant-offset control can
|
|
## be compared against the actual average readout the TM produces.
|
|
radDeltaSum*, radDeltaAbsSum*: float
|
|
radDeltaN*: int
|
|
radOffsetSum*: float
|
|
radOffsetN*: int
|
|
classCorrect*: int ## warm predictions whose class matched the eventual label
|
|
classTotal*: int ## warm predictions with a resolvable label
|
|
## Per-class confusion matrix for the GF head, indexed [true label][predicted
|
|
## class], counted over WARM predictions only (the same samples `classTotal`
|
|
## scores). This is what makes the majority-class baseline and per-class
|
|
## precision/recall measurable. Row sums = the warm label histogram; the
|
|
## diagonal sum = classCorrect.
|
|
confusion*: array[TM_CLASSES, array[TM_CLASSES, int]]
|
|
## Same for the radial head.
|
|
radConfusion*: array[TM_CLASSES, array[TM_CLASSES, int]]
|
|
lastChosen*: int
|
|
shuffleLabels*: bool ## control: replace the computed GF label with a random class
|
|
forceBase*: bool ## measurement: ignore the TM, emit the pure LinearGun base
|
|
debugGraphics*: bool
|
|
|
|
# ── TM core (Granmo Table 2/3, corrected resource allocation) ────────────────
|
|
|
|
proc tmPolarity(cl: int): float {.inline.} =
|
|
if cl < TM_HALF: 1.0 else: -1.0
|
|
|
|
proc tmNewTeam(): seq[int16] =
|
|
result = newSeq[int16](TM_NCLAUSES * TM_NLITS) # 0 = Exclude boundary
|
|
|
|
proc tmEval(team: seq[int16], lits: TmBits, cl: int, learning: bool): uint8 =
|
|
let base = cl * TM_NLITS
|
|
var hasInc = false
|
|
for lit in 0..<TM_NLITS:
|
|
if team[base + lit] > 0:
|
|
hasInc = true
|
|
if lits[lit] == 0'u8: return 0'u8
|
|
if hasInc: return 1'u8
|
|
# Eq. 6: the empty conjunction is vacuously true during learning, false in
|
|
# classification. Without this the all-Exclude init deadlocks.
|
|
return if learning: 1'u8 else: 0'u8
|
|
|
|
proc tmForward(team: seq[int16], lits: TmBits,
|
|
cache: var array[TM_NCLAUSES, uint8]): float =
|
|
var v = 0.0
|
|
for cl in 0..<TM_NCLAUSES:
|
|
let o = tmEval(team, lits, cl, learning = false)
|
|
cache[cl] = tmEval(team, lits, cl, learning = true)
|
|
v += tmPolarity(cl) * float(o)
|
|
clamp(v, -TM_T, TM_T)
|
|
|
|
proc tmLearnDir(team: var seq[int16], lits: TmBits,
|
|
cache: array[TM_NCLAUSES, uint8], vote, d: float) =
|
|
## One Granmo update of one class team with desired vote direction `d`.
|
|
let pFeedback = (TM_T - d * vote) / (2.0 * TM_T)
|
|
if pFeedback <= 0.0: return
|
|
for cl in 0..<TM_NCLAUSES:
|
|
if rand(1.0) >= pFeedback: continue
|
|
let pol = tmPolarity(cl)
|
|
let cOut = cache[cl]
|
|
let base = cl * TM_NLITS
|
|
if pol * d > 0.0:
|
|
# Type I (Table 2) collapsed to the resulting state move.
|
|
for lit in 0..<TM_NLITS:
|
|
var st = int(team[base + lit])
|
|
if lits[lit] == 1'u8:
|
|
if cOut == 1'u8:
|
|
if rand(1.0) < (TM_S - 1.0) / TM_S: st = min(st + 1, TM_NSTATES)
|
|
else:
|
|
if rand(1.0) < 1.0 / TM_S: st = max(st - 1, -TM_NSTATES)
|
|
else:
|
|
if rand(1.0) < 1.0 / TM_S: st = max(st - 1, -TM_NSTATES)
|
|
team[base + lit] = int16(st)
|
|
else:
|
|
# Type II (Table 3): penalise exclusion of a zero literal when firing.
|
|
if cOut == 1'u8:
|
|
for lit in 0..<TM_NLITS:
|
|
if lits[lit] == 0'u8:
|
|
if team[base + lit] <= 0:
|
|
team[base + lit] = int16(min(int(team[base + lit]) + 1, TM_NSTATES))
|
|
|
|
# ── geometry helpers ─────────────────────────────────────────────────────────
|
|
|
|
proc tmBinForSpeed(spd: float): int {.inline.} =
|
|
for i in 0..<len(vb.PowerBins):
|
|
if abs(spd - bulletSpeed(vb.PowerBins[i])) < 1e-6: return i
|
|
-1
|
|
|
|
proc tmTraceSlot(fireTick, binIdx: int): int {.inline.} =
|
|
((fireTick * len(vb.PowerBins)) + binIdx) mod TM_TRACE_SLOTS
|
|
|
|
proc gfToBucket(gf: float): int {.inline.} =
|
|
clamp(int(round((gf + 1.0) * 0.5 * float(TM_CLASSES - 1))), 0, TM_CLASSES - 1)
|
|
|
|
proc bucketToGF(c: int): float {.inline.} =
|
|
if TM_CLASSES <= 1: 0.0
|
|
else: float(c) / float(TM_CLASSES - 1) * 2.0 - 1.0
|
|
|
|
proc radToBucket(delta: float): int {.inline.} =
|
|
let u = clamp(delta / TM_RADIAL_RANGE, -1.0, 1.0)
|
|
clamp(int(round((u + 1.0) * 0.5 * float(TM_CLASSES - 1))), 0, TM_CLASSES - 1)
|
|
|
|
proc bucketToRadial(c: int): float {.inline.} =
|
|
if TM_CLASSES <= 1: 0.0
|
|
else: (float(c) / float(TM_CLASSES - 1) * 2.0 - 1.0) * TM_RADIAL_RANGE
|
|
|
|
proc normDeg(d: float): float {.inline.} =
|
|
result = d
|
|
while result > 180.0: result -= 360.0
|
|
while result < -180.0: result += 360.0
|
|
|
|
# ── public API ───────────────────────────────────────────────────────────────
|
|
|
|
proc initTmPatternGun*(): TmPatternGun =
|
|
for c in 0..<TM_CLASSES: result.teams[c] = tmNewTeam()
|
|
for c in 0..<TM_CLASSES: result.radTeams[c] = tmNewTeam()
|
|
for c in 0..<2: result.revTeams[c] = tmNewTeam()
|
|
result.targetMode = tmGF
|
|
result.lastTick = -1
|
|
result.prevTick = -1
|
|
result.currentTarget = -1
|
|
result.lastChosen = (TM_CLASSES - 1) div 2
|
|
randomize()
|
|
result.debugGraphics = false
|
|
|
|
proc initTmRadialGun*(): TmPatternGun =
|
|
## The RACK-REGISTERED instance: the RADIAL target mode, which is the
|
|
## control-validated winner under `bmPoint` (see tm_pattern_sweep_results.md,
|
|
## Round 2 Task 2). The gun type carries all three heads; the live rack only
|
|
## ever selects this radial-mode instance.
|
|
result = initTmPatternGun()
|
|
result.targetMode = tmRadial
|
|
|
|
proc isWarmedUp*(g: TmPatternGun): bool {.inline.} = true
|
|
|
|
proc resetLearning*(g: var TmPatternGun) =
|
|
## Fresh concept: wipe every clause team and the motion history. Called when
|
|
## the target id changes so a new opponent starts from a cold net.
|
|
for c in 0..<TM_CLASSES: g.teams[c] = tmNewTeam()
|
|
for c in 0..<TM_CLASSES: g.radTeams[c] = tmNewTeam()
|
|
for c in 0..<2: g.revTeams[c] = tmNewTeam()
|
|
g.totalObs = 0
|
|
g.hasPrev = false
|
|
for i in 0..<3:
|
|
g.latSignHist[i] = 0
|
|
g.turnSignHist[i] = 0
|
|
g.sinceReversal = 0
|
|
g.radialFracSm = 0.0
|
|
g.latPersist = 0
|
|
g.pendingCount = 0
|
|
|
|
proc tmUpdateHistory(g: var TmPatternGun, state: WorldState) =
|
|
if state.tick == g.lastTick: return
|
|
g.lastTick = state.tick
|
|
let slot = ((state.tick mod POS_RING) + POS_RING) mod POS_RING
|
|
g.posRing[slot] = PosSample(tick: state.tick, x: state.enemyX, y: state.enemyY,
|
|
heading: state.enemyHeading, valid: true)
|
|
if g.hasPrev and state.tick > g.prevTick:
|
|
let dx = state.enemyX - g.prevX
|
|
let dy = state.enemyY - g.prevY
|
|
let spd = hypot(dx, dy)
|
|
let lx = state.enemyX - state.selfX
|
|
let ly = state.enemyY - state.selfY
|
|
let ld = hypot(lx, ly)
|
|
var latSign = 0
|
|
var crossFrac = 0.0
|
|
var radialFrac = 0.0
|
|
if ld > 1e-6 and spd > 1e-6:
|
|
let cross = (lx * dy - ly * dx) / ld
|
|
crossFrac = abs(cross) / spd
|
|
radialFrac = abs((lx * dx + ly * dy) / ld) / spd
|
|
if cross > 0.5: latSign = 1
|
|
elif cross < -0.5: latSign = -1
|
|
for i in countdown(2, 1): g.latSignHist[i] = g.latSignHist[i - 1]
|
|
g.latSignHist[0] = latSign
|
|
if latSign != 0 and g.lastNonzeroLat != 0 and latSign != g.lastNonzeroLat:
|
|
g.sinceReversal = 0
|
|
else:
|
|
inc g.sinceReversal
|
|
if latSign != 0: g.lastNonzeroLat = latSign
|
|
g.latPersist = if latSign != 0 and latSign == g.latSignHist[1]: 1 else: 0
|
|
var dh = normDeg(state.enemyHeading - g.prevHeading)
|
|
let ts = if dh > 0.5: 1 elif dh < -0.5: -1 else: 0
|
|
for i in countdown(2, 1): g.turnSignHist[i] = g.turnSignHist[i - 1]
|
|
g.turnSignHist[0] = ts
|
|
g.radialFracSm = 0.8 * g.radialFracSm + 0.2 * radialFrac
|
|
g.prevX = state.enemyX
|
|
g.prevY = state.enemyY
|
|
g.prevHeading = state.enemyHeading
|
|
g.prevTick = state.tick
|
|
g.hasPrev = true
|
|
|
|
proc tmBuildBits(g: var TmPatternGun, state: WorldState, flightTicks: float):
|
|
TmBits =
|
|
## 40 binary/bucketed context features. Written as literals directly: bit i
|
|
## and its negation at i + TM_NBITS.
|
|
var bits: array[TM_NBITS, uint8]
|
|
var o = 0
|
|
template put(v: uint8) = (bits[o] = v; inc o)
|
|
template putSign(s: int) =
|
|
# ternary -> 2 bits: (+, -); 0 -> (0,0)
|
|
put(if s > 0: 1'u8 else: 0'u8)
|
|
put(if s < 0: 1'u8 else: 0'u8)
|
|
for i in 0..2: putSign(g.latSignHist[i]) # 6
|
|
for i in 0..2: putSign(g.turnSignHist[i]) # 6
|
|
|
|
# time since last lateral reversal: 3 one-hot
|
|
if g.sinceReversal <= 3: put 1'u8 else: put 0'u8
|
|
if g.sinceReversal > 3 and g.sinceReversal <= 10: put 1'u8 else: put 0'u8
|
|
if g.sinceReversal > 10: put 1'u8 else: put 0'u8
|
|
|
|
# lateral persistence / magnitude: 3 bits
|
|
put(if g.latPersist == 1: 1'u8 else: 0'u8)
|
|
|
|
let spd = state.enemySpeed
|
|
# speed band (uses signed velocity magnitude; classic fixtures may be negative)
|
|
let aspd = abs(spd)
|
|
if aspd < 1.0: put 1'u8 else: put 0'u8
|
|
if aspd >= 1.0 and aspd < 4.0: put 1'u8 else: put 0'u8
|
|
if aspd >= 4.0: put 1'u8 else: put 0'u8
|
|
|
|
# distance band
|
|
let d = hypot(state.enemyX - state.selfX, state.enemyY - state.selfY)
|
|
if d < 150.0: put 1'u8 else: put 0'u8
|
|
if d >= 150.0 and d < 350.0: put 1'u8 else: put 0'u8
|
|
if d >= 350.0: put 1'u8 else: put 0'u8
|
|
|
|
# flight-time band
|
|
if flightTicks < 10.0: put 1'u8 else: put 0'u8
|
|
if flightTicks >= 10.0 and flightTicks < 25.0: put 1'u8 else: put 0'u8
|
|
if flightTicks >= 25.0: put 1'u8 else: put 0'u8
|
|
|
|
# walls close (4 bits)
|
|
put(if state.enemyY < 60.0: 1'u8 else: 0'u8)
|
|
put(if state.arenaHeight - state.enemyY < 60.0: 1'u8 else: 0'u8)
|
|
put(if state.arenaWidth - state.enemyX < 60.0: 1'u8 else: 0'u8)
|
|
put(if state.enemyX < 60.0: 1'u8 else: 0'u8)
|
|
|
|
# radial fraction band (3 one-hot)
|
|
if g.radialFracSm < 0.35: put 1'u8 else: put 0'u8
|
|
if g.radialFracSm >= 0.35 and g.radialFracSm < 0.7: put 1'u8 else: put 0'u8
|
|
if g.radialFracSm >= 0.7: put 1'u8 else: put 0'u8
|
|
|
|
# enemy energy band (2 one-hot)
|
|
if state.enemyEnergy < 20.0: put 1'u8 else: put 0'u8
|
|
if state.enemyEnergy >= 20.0: put 1'u8 else: put 0'u8
|
|
|
|
# heading relative to LOS (toward / away)
|
|
let lx = state.enemyX - state.selfX
|
|
let ly = state.enemyY - state.selfY
|
|
let ld = hypot(lx, ly)
|
|
var dot = 0.0
|
|
if ld > 1e-6:
|
|
let hr = degToRad(state.enemyHeading)
|
|
dot = (cos(hr) * lx + sin(hr) * ly) / ld
|
|
put(if dot > 0.3: 1'u8 else: 0'u8)
|
|
put(if dot < -0.3: 1'u8 else: 0'u8)
|
|
|
|
# approach (closing / opening) relative to displacement direction
|
|
var closing = 0.0
|
|
if ld > 1e-6 and g.hasPrev:
|
|
let vx = state.enemyX - g.prevX
|
|
let vy = state.enemyY - g.prevY
|
|
closing = (vx * lx + vy * ly) / ld
|
|
put(if closing < -0.3: 1'u8 else: 0'u8)
|
|
put(if closing > 0.3: 1'u8 else: 0'u8)
|
|
|
|
# pack into literal vector
|
|
for i in 0..<TM_NBITS:
|
|
result[i] = bits[i]
|
|
result[i + TM_NBITS] = 1'u8 - bits[i]
|
|
|
|
proc tmSoftGF(votes: array[TM_CLASSES, float]): float =
|
|
## Vote-weighted expectation of the GF bucket centres. The multiclass votes are
|
|
## not calibrated probabilities, but a softmax over them gives a smooth,
|
|
## self-shrinking readout (flat votes -> GF 0; a confident tail -> near that
|
|
## tail), which avoids the up-to-half-bucket aim error of a hard argmax.
|
|
var mx = -Inf
|
|
for c in 0..<TM_CLASSES:
|
|
if votes[c] > mx: mx = votes[c]
|
|
var sum = 0.0
|
|
var w: array[TM_CLASSES, float]
|
|
for c in 0..<TM_CLASSES:
|
|
w[c] = exp((votes[c] - mx) * TM_SOFT_BETA / TM_T)
|
|
sum += w[c]
|
|
if sum <= 0.0: return 0.0
|
|
for c in 0..<TM_CLASSES:
|
|
result += (w[c] / sum) * bucketToGF(c)
|
|
|
|
proc tmChooseAt(votes: openArray[float], centre: int, margin: float,
|
|
nObs: int): int =
|
|
## Cold gun or no class beating the centre by `margin` (fraction of TM_T)
|
|
## -> the centre. Shared by the GF, radial and reversal heads.
|
|
if nObs < TM_MIN_OBS:
|
|
return centre
|
|
var best = 0
|
|
var bestV = -Inf
|
|
for c in 0..<votes.len:
|
|
if votes[c] > bestV: bestV = votes[c]; best = c
|
|
if best == centre:
|
|
return centre
|
|
if (votes[best] - votes[centre]) / TM_T < margin:
|
|
return centre
|
|
result = best
|
|
|
|
proc tmChooseClass(g: var TmPatternGun, votes: array[TM_CLASSES, float]): int =
|
|
tmChooseAt(votes, (TM_CLASSES - 1) div 2, TM_CONF_MARGIN, g.totalObs)
|
|
|
|
proc tmResolveTrace(g: var TmPatternGun, t: TmPatternTrace, power: float) =
|
|
## One label + one TM update for a fired virtual bullet, using the enemy
|
|
## position recorded at the BASE arrival tick. `t` is a value copy of the
|
|
## fire-time trace, so this is safe to call either from `onResult` (the label
|
|
## is already resolvable) or from `tmFlushPending` (the label was deferred
|
|
## because the bullet resolved before its base arrival tick).
|
|
let s = ((t.arrivalTick mod POS_RING) + POS_RING) mod POS_RING
|
|
let speed = bulletSpeed(power)
|
|
let mea = arcsin(clamp(8.0 / speed, -1.0, 1.0))
|
|
let actualBearing = arctan2(g.posRing[s].y - t.fireY, g.posRing[s].x - t.fireX)
|
|
var delta = actualBearing - t.baseBearing
|
|
while delta > PI: delta -= 2.0 * PI
|
|
while delta < -PI: delta += 2.0 * PI
|
|
let gf = if mea > 1e-10: clamp(delta / mea, -1.0, 1.0) else: 0.0
|
|
|
|
# The shuffled control randomises ONLY the head the current mode is claiming.
|
|
let shuffleGF = g.shuffleLabels and g.targetMode == tmGF
|
|
let shuffleRad = g.shuffleLabels and g.targetMode == tmRadial
|
|
let shuffleRev = g.shuffleLabels and g.targetMode == tmReversal
|
|
|
|
let winner = if shuffleGF: rand(TM_CLASSES - 1) else: gfToBucket(gf)
|
|
inc g.labelHist[winner]
|
|
if t.warm:
|
|
inc g.classTotal
|
|
if winner == t.chosen: inc g.classCorrect
|
|
inc g.confusion[winner][t.chosen]
|
|
|
|
# Radial label: enemy radius at the base arrival tick minus the base fire
|
|
# distance. Independent of our own aim, so it is a clean target.
|
|
let actualRadius = hypot(g.posRing[s].x - t.fireX, g.posRing[s].y - t.fireY)
|
|
let radDelta = actualRadius - t.fireDist
|
|
inc g.radDeltaN
|
|
g.radDeltaSum += radDelta
|
|
g.radDeltaAbsSum += abs(radDelta)
|
|
let radWinner = if shuffleRad: rand(TM_CLASSES - 1) else: radToBucket(radDelta)
|
|
inc g.radLabelHist[radWinner]
|
|
if t.warm:
|
|
inc g.radTotal
|
|
if radWinner == t.radChosen: inc g.radCorrect
|
|
inc g.radConfusion[radWinner][t.radChosen]
|
|
|
|
# Reversal label: net heading turn over the flight, opposite to the direction
|
|
# the enemy was turning at fire time.
|
|
let dh = normDeg(g.posRing[s].heading - t.fireHeading)
|
|
let netTurn = if dh > TM_REV_TURN_DEG: 1 elif dh < -TM_REV_TURN_DEG: -1 else: 0
|
|
let revWinner =
|
|
if shuffleRev: rand(1)
|
|
elif t.fireTurn != 0 and netTurn != 0 and netTurn != t.fireTurn: 1
|
|
else: 0
|
|
inc g.revLabelHist[revWinner]
|
|
if t.warm:
|
|
inc g.revTotal
|
|
if revWinner == t.revChosen: inc g.revCorrect
|
|
|
|
case g.targetMode
|
|
of tmGF:
|
|
for c in 0..<TM_CLASSES:
|
|
let d = if c == winner: 1.0 else: -1.0
|
|
g.teams[c].tmLearnDir(t.lits, t.cache[c], t.votes[c], d)
|
|
of tmRadial:
|
|
for c in 0..<TM_CLASSES:
|
|
let d = if c == radWinner: 1.0 else: -1.0
|
|
g.radTeams[c].tmLearnDir(t.lits, t.radCache[c], t.radVotes[c], d)
|
|
of tmReversal:
|
|
for c in 0..<TM_CLASSES:
|
|
let d = if c == winner: 1.0 else: -1.0
|
|
g.teams[c].tmLearnDir(t.lits, t.cache[c], t.votes[c], d)
|
|
for c in 0..<2:
|
|
let d = if c == revWinner: 1.0 else: -1.0
|
|
g.revTeams[c].tmLearnDir(t.lits, t.revCache[c], t.revVotes[c], d)
|
|
inc g.totalObs
|
|
inc g.trainCalls
|
|
|
|
proc tmFlushPending(g: var TmPatternGun) =
|
|
## Resolve every deferred trace whose BASE arrival tick is now recorded in
|
|
## `posRing`. Called once per `predict` right after `tmUpdateHistory`, so the
|
|
## just-written current tick is visible. Entries are compacted in place.
|
|
if g.pendingCount == 0: return
|
|
var w = 0
|
|
for i in 0..<g.pendingCount:
|
|
let p = addr g.pending[i]
|
|
let s = ((p.arrivalTick mod POS_RING) + POS_RING) mod POS_RING
|
|
if g.posRing[s].valid and g.posRing[s].tick == p.arrivalTick:
|
|
g.tmResolveTrace(p.trace, p.power)
|
|
else:
|
|
if w != i: g.pending[w] = g.pending[i]
|
|
inc w
|
|
g.pendingCount = w
|
|
|
|
proc predict*(g: var TmPatternGun, state: WorldState, bulletSpeed: float):
|
|
GunPrediction =
|
|
inc g.predictCalls
|
|
|
|
# target-change reset (per-enemy specialisation)
|
|
if state.enemies.len > 0:
|
|
let tid = state.enemies[0].id
|
|
if g.currentTarget != tid:
|
|
if g.currentTarget >= 0: g.resetLearning()
|
|
g.currentTarget = tid
|
|
|
|
g.tmUpdateHistory(state)
|
|
# Deferred-label flush (Task 2): resolve any fired bullet whose BASE arrival
|
|
# tick is now in the ring, before the cold-start gate reads totalObs.
|
|
g.tmFlushPending()
|
|
|
|
if bulletSpeed <= 0.0:
|
|
return GunPrediction(x: state.enemyX, y: state.enemyY)
|
|
|
|
let mea = arcsin(clamp(8.0 / bulletSpeed, -1.0, 1.0))
|
|
let f = forecastLinear(state, bulletSpeed)
|
|
let flightTicks = f.dist / bulletSpeed
|
|
|
|
let lits = g.tmBuildBits(state, flightTicks)
|
|
|
|
var votes: array[TM_CLASSES, float]
|
|
var caches: array[TM_CLASSES, array[TM_NCLAUSES, uint8]]
|
|
for c in 0..<TM_CLASSES:
|
|
votes[c] = tmForward(g.teams[c], lits, caches[c])
|
|
var radVotes: array[TM_CLASSES, float]
|
|
var radCaches: array[TM_CLASSES, array[TM_NCLAUSES, uint8]]
|
|
for c in 0..<TM_CLASSES:
|
|
radVotes[c] = tmForward(g.radTeams[c], lits, radCaches[c])
|
|
var revVotes: array[2, float]
|
|
var revCaches: array[2, array[TM_NCLAUSES, uint8]]
|
|
for c in 0..<2:
|
|
revVotes[c] = tmForward(g.revTeams[c], lits, revCaches[c])
|
|
|
|
let chosen = g.tmChooseClass(votes)
|
|
let radChosen = tmChooseAt(radVotes, (TM_CLASSES - 1) div 2, TM_RAD_MARGIN,
|
|
g.totalObs)
|
|
let revChosen = tmChooseAt(revVotes, 0, TM_REV_MARGIN, g.totalObs)
|
|
g.lastChosen = chosen
|
|
inc g.chosenHist[chosen]
|
|
inc g.radChosenHist[radChosen]
|
|
inc g.revChosenHist[revChosen]
|
|
|
|
var gf = 0.0
|
|
var radOffset = 0.0
|
|
if not (g.forceBase or g.totalObs < TM_MIN_OBS):
|
|
case g.targetMode
|
|
of tmGF:
|
|
if TM_GF_MODE == "soft": gf = TM_SHRINK * tmSoftGF(votes)
|
|
else: gf = TM_SHRINK * bucketToGF(chosen)
|
|
of tmRadial:
|
|
radOffset = bucketToRadial(radChosen)
|
|
g.radOffsetSum += radOffset
|
|
inc g.radOffsetN
|
|
of tmReversal:
|
|
if TM_GF_MODE == "soft": gf = TM_SHRINK * tmSoftGF(votes)
|
|
else: gf = TM_SHRINK * bucketToGF(chosen)
|
|
if revChosen == 1:
|
|
gf = -TM_REV_GAIN * gf
|
|
|
|
# An all-zero correction must reproduce `LinearGun` BYTE-FOR-BYTE, so use its
|
|
# exact aim point and its exact [0, arena] clamp rather than the
|
|
# BotRadius-inset clamp the corrective excursions use.
|
|
let usedBase = (gf == 0.0 and radOffset == 0.0)
|
|
var px, py: float
|
|
if usedBase:
|
|
px = f.x
|
|
py = f.y
|
|
else:
|
|
let aimAngle = f.bearing + gf * mea
|
|
let aimDist = f.dist + radOffset
|
|
px = state.selfX + cos(aimAngle) * aimDist
|
|
py = state.selfY + sin(aimAngle) * aimDist
|
|
|
|
let binIdx = tmBinForSpeed(bulletSpeed)
|
|
if binIdx >= 0:
|
|
let slot = tmTraceSlot(state.tick, binIdx)
|
|
# The virtual bullet advances one step on its own spawn tick, so it reaches
|
|
# the base fire distance after max(0, ceil(flightTicks)-1) further ticks.
|
|
let arrOff = max(0, int(ceil(flightTicks)) - 1)
|
|
g.traces[slot] = TmPatternTrace(
|
|
fireTick: state.tick, powerBin: binIdx,
|
|
arrivalTick: state.tick + arrOff,
|
|
baseBearing: f.bearing,
|
|
fireX: state.selfX, fireY: state.selfY,
|
|
fireHeading: state.enemyHeading, fireTurn: g.turnSignHist[0],
|
|
fireDist: f.dist,
|
|
lits: lits, votes: votes, cache: caches, chosen: chosen,
|
|
radVotes: radVotes, radCache: radCaches, radChosen: radChosen,
|
|
revVotes: revVotes, revCache: revCaches, revChosen: revChosen,
|
|
warm: (g.totalObs >= TM_MIN_OBS), alive: true)
|
|
|
|
when DebugTMPattern:
|
|
echo "tmp tick=", state.tick, " bin=", binIdx, " chosen=", chosen,
|
|
" radChosen=", radChosen, " revChosen=", revChosen,
|
|
" gf=", gf, " roff=", radOffset, " obs=", g.totalObs, " votes=", votes
|
|
|
|
GunPrediction(
|
|
x: if usedBase: clamp(px, 0.0, state.arenaWidth)
|
|
else: clamp(px, BotRadius, state.arenaWidth - BotRadius),
|
|
y: if usedBase: clamp(py, 0.0, state.arenaHeight)
|
|
else: clamp(py, BotRadius, state.arenaHeight - BotRadius),
|
|
)
|
|
|
|
proc onResult*(g: var TmPatternGun, e: FeedbackEvent) =
|
|
let binIdx =
|
|
if e.powerBin >= 0 and e.powerBin < len(vb.PowerBins): e.powerBin
|
|
else: tmBinForSpeed(bulletSpeed(e.bulletPower))
|
|
if binIdx < 0:
|
|
inc g.traceMisses
|
|
return
|
|
let slot = tmTraceSlot(e.fireTick, binIdx)
|
|
var t = addr g.traces[slot]
|
|
if not t.alive or t.fireTick != e.fireTick or t.powerBin != binIdx:
|
|
inc g.traceMisses
|
|
return
|
|
|
|
# Clean label: enemy position at the BASE arrival tick from our own history.
|
|
let s = ((t.arrivalTick mod POS_RING) + POS_RING) mod POS_RING
|
|
if g.posRing[s].valid and g.posRing[s].tick == t.arrivalTick:
|
|
g.tmResolveTrace(t[], e.bulletPower)
|
|
else:
|
|
# DEFER (Task 2): the bullet resolved BEFORE its BASE arrival tick, which
|
|
# happens whenever the radial correction aimed SHORT. The arrival-tick
|
|
# position is not recorded yet, so keep a COPY of the trace and train on it
|
|
# once that tick is in the ring (`tmFlushPending`). Dropping it here is what
|
|
# biased the training set toward only the resolvable (long/centre) aims.
|
|
if g.pendingCount < TM_PENDING_SLOTS:
|
|
g.pending[g.pendingCount] = PendingResolve(
|
|
arrivalTick: t.arrivalTick, powerBin: binIdx,
|
|
power: e.bulletPower, trace: t[])
|
|
inc g.pendingCount
|
|
else:
|
|
inc g.pendingDropped
|
|
inc g.labelMisses
|
|
t.alive = false
|