TM gun: the discrete-target diagnosis was RIGHT - it learns now. Still loses to Linear.
The user's goal: a TM gun that is the best 1v1 gun, starting from scratch every
battle but quickly overfitting the current enemy. The previous attempt (knob
tuning) failed: NO configuration beat its own shuffled-feedback control, and the
TM-off ablation scored the same as TM-on, i.e. the TM's correction was
near-zero-mean noise. Diagnosis then: a Tsetlin Machine is a CLASSIFIER, and we
were asking it for an absolute aim point - a regression target. So this attempt
gave it a DISCRETE target (multi-class over guess-factor buckets) with 40
binary/bucketed motion features, and measured it against Linear, the default
Tsetlin gun, and a MANDATORY shuffled control.
THE DIAGNOSIS IS CONFIRMED - THE TM LEARNS, DECISIVELY:
online class accuracy 46.0% vs shuffled control 20.0% (2.3x chance)
raw ungated argmax 21.2%/18.6% vs shuffled 15.2%/8.3% (18/18, p<0.0001)
TMPattern > its shuffled control, overall 17/1 runs, p=0.0001
Compare the previous attempt, which could not beat shuffled feedback at all.
TMPattern also beats the default Tsetlin gun early (17/1, p=0.0001), so it is a
strictly better TM gun than the one in the rack.
BUT IT IS NOT COMPETITIVE WITH LINEAR ON REAL SURFERS:
real DrussGT, bmPath (the shipped metric), 3 seeds, pooled early/overall
Linear 34.0% (6358/18715) 24.3% (58297/239943)
TMPattern (gated) 27.9% (15514/55535) 22.0% (158658/719681)
TMPatternShuf 28.7% 19.4%
Linear > TMPattern: 15/18 early p=0.0075, 15/18 overall p=0.0075
bmPoint: neutral (7.2%/4.6% vs Linear 7.2%/4.7%)
synthetic controlled motion: matches/edges Linear (66.8%/60.6% vs 66.4%/59.6%,
shuffled 55.7%/50.1%) - the mechanism works when motion is predictable.
So: the representation fix moved this from "learns nothing" to "learns strongly
but applies its knowledge badly". INFERRED reason for the residual loss: the
linear lead is already the modal GF bucket (the label histogram is centred), so
corrective excursions away from it are net-negative. The measured deficit lives
in the BASELINE and in RANGE, not in the TM knobs - which is why further knob
tuning was never going to work.
Best config: gated hard K=5, TM_CONF_MARGIN=0.25, TM_SHRINK=0.5.
NOT TRIED (time-boxed): the binary-reversal target, and a RADIAL (range-holding)
target - the latter is the top next step.
Adds `common_libs/guns/tm_pattern.nim` (NOT registered in the rack),
`common_libs/tests/sweep_tm_pattern.nim`, and a durable writeup at
`common_libs/tests/tm_pattern_sweep_results.md`.
This commit is contained in:
@@ -0,0 +1,493 @@
|
||||
## 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
|
||||
DebugTMPattern* = false
|
||||
|
||||
type
|
||||
TmBits* = array[TM_NLITS, uint8]
|
||||
|
||||
TmPatternTrace = object
|
||||
fireTick: int
|
||||
powerBin: int
|
||||
arrivalTick: int
|
||||
baseBearing: float
|
||||
fireX, fireY: float
|
||||
lits: TmBits
|
||||
votes: array[TM_CLASSES, float]
|
||||
cache: array[TM_CLASSES, array[TM_NCLAUSES, uint8]]
|
||||
chosen: int
|
||||
warm: bool
|
||||
alive: bool
|
||||
|
||||
PosSample = object
|
||||
tick: int
|
||||
x, y: float
|
||||
valid: bool
|
||||
|
||||
TmPatternGun* = object
|
||||
teams: array[TM_CLASSES, seq[int16]]
|
||||
traces: array[TM_TRACE_SLOTS, TmPatternTrace]
|
||||
# ── 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]
|
||||
classCorrect*: int ## warm predictions whose class matched the eventual label
|
||||
classTotal*: int ## warm predictions with a resolvable label
|
||||
lastChosen*: int
|
||||
shuffleLabels*: bool ## control: replace the computed GF label with a random class
|
||||
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 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()
|
||||
result.lastTick = -1
|
||||
result.prevTick = -1
|
||||
result.currentTarget = -1
|
||||
result.lastChosen = (TM_CLASSES - 1) div 2
|
||||
randomize()
|
||||
result.debugGraphics = false
|
||||
|
||||
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()
|
||||
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
|
||||
|
||||
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,
|
||||
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 tmChooseClass(g: var TmPatternGun, votes: array[TM_CLASSES, float]): int =
|
||||
## Cold gun or a flat vote vector -> straight ahead (GF = 0). Otherwise the
|
||||
## argmax class, but ONLY when it beats the centre class by TM_CONF_MARGIN
|
||||
## (fraction of TM_T); otherwise stay at the centre. This is what keeps the
|
||||
## gun from degrading to arbitrary buckets when the TM has no real evidence.
|
||||
let centre = (TM_CLASSES - 1) div 2
|
||||
if g.totalObs < TM_MIN_OBS:
|
||||
return centre
|
||||
var best = 0
|
||||
var bestV = -Inf
|
||||
for c in 0..<TM_CLASSES:
|
||||
if votes[c] > bestV: bestV = votes[c]; best = c
|
||||
if best == centre:
|
||||
return centre
|
||||
let margin = (votes[best] - votes[centre]) / TM_T
|
||||
if margin < TM_CONF_MARGIN:
|
||||
return centre
|
||||
result = best
|
||||
|
||||
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)
|
||||
|
||||
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])
|
||||
|
||||
let chosen = g.tmChooseClass(votes)
|
||||
g.lastChosen = chosen
|
||||
inc g.chosenHist[chosen]
|
||||
|
||||
var gf: float
|
||||
if g.totalObs < TM_MIN_OBS:
|
||||
gf = 0.0
|
||||
elif TM_GF_MODE == "soft":
|
||||
gf = TM_SHRINK * tmSoftGF(votes)
|
||||
else:
|
||||
gf = TM_SHRINK * bucketToGF(chosen)
|
||||
let aimAngle = f.bearing + gf * mea
|
||||
let px = state.selfX + cos(aimAngle) * f.dist
|
||||
let py = state.selfY + sin(aimAngle) * f.dist
|
||||
|
||||
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,
|
||||
lits: lits, votes: votes, cache: caches,
|
||||
chosen: chosen, warm: (g.totalObs >= TM_MIN_OBS), alive: true)
|
||||
|
||||
when DebugTMPattern:
|
||||
echo "tmp tick=", state.tick, " bin=", binIdx, " chosen=", chosen,
|
||||
" gf=", gf, " obs=", g.totalObs, " votes=", votes
|
||||
|
||||
GunPrediction(
|
||||
x: clamp(px, BotRadius, state.arenaWidth - BotRadius),
|
||||
y: 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 not g.posRing[s].valid or g.posRing[s].tick != t.arrivalTick:
|
||||
inc g.labelMisses
|
||||
t.alive = false
|
||||
return
|
||||
|
||||
let speed = bulletSpeed(e.bulletPower)
|
||||
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
|
||||
let winner = if g.shuffleLabels: rand(TM_CLASSES - 1) else: gfToBucket(gf)
|
||||
inc g.labelHist[winner]
|
||||
if t.warm:
|
||||
inc g.classTotal
|
||||
if winner == t.chosen: inc g.classCorrect
|
||||
|
||||
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)
|
||||
inc g.totalObs
|
||||
inc g.trainCalls
|
||||
t.alive = false
|
||||
Reference in New Issue
Block a user