240 lines
8.7 KiB
Nim
240 lines
8.7 KiB
Nim
## Learned movement (SBC) — module unit / smoke test.
|
|
##
|
|
## No Java, no battle: a synthetic enemy fires at us on a fixed clock while our
|
|
## own bot integrates the commands the module returns, so the wave machinery,
|
|
## the state coding, the counted-SBC learner and the danger ranking are all
|
|
## exercised end to end.
|
|
##
|
|
## nim c -r --nimcache:/tmp/nc_j128 --path:../../common_libs \
|
|
## tests/test_learned_surfer.nim # from ModularBot_garage/
|
|
##
|
|
## Checks:
|
|
## 1. the module satisfies the MovementModule concept
|
|
## 2. waves are detected from the energy drop and RESOLVE at the nominal
|
|
## arrival tick with a valid 31-bin label
|
|
## 3. the counted SBC accumulates evidence, and the danger map is a proper
|
|
## probability distribution (sums to 1)
|
|
## 4. the state code is in range and moves with our movement state
|
|
## 5. TR_LEARNED_DECAY_SHIFT=0 keeps the counters (no forgetting), the default
|
|
## decays them
|
|
## 6. the run is deterministic
|
|
|
|
import std/[math, os]
|
|
import movements/learned_surfer
|
|
import gun_harness/gun_interface
|
|
import movement_harness/movement_interface
|
|
|
|
var checks = 0
|
|
var failures = 0
|
|
|
|
proc check(what: string, ok: bool) =
|
|
inc checks
|
|
if ok: echo "PASS ", what
|
|
else: echo "FAIL ", what
|
|
if not ok: inc failures
|
|
|
|
type Sim = object
|
|
x, y, heading, speed: float
|
|
enemyX, enemyY: float
|
|
tick: int
|
|
fireTick: int
|
|
enemyEnergy: float
|
|
|
|
proc step(m: var LearnedSurferModule, s: var Sim): MoveCommand =
|
|
let ws = WorldState(
|
|
enemyX: s.enemyX, enemyY: s.enemyY, enemyEnergy: s.enemyEnergy,
|
|
selfX: s.x, selfY: s.y, selfSpeed: s.speed, selfHeading: s.heading,
|
|
arenaWidth: 800.0, arenaHeight: 600.0, tick: s.tick,
|
|
enemies: @[EnemyInfo(id: 1, x: s.enemyX, y: s.enemyY,
|
|
heading: 180.0, speed: 0.0, energy: s.enemyEnergy)],
|
|
)
|
|
result = m.computeMove(ws)
|
|
# integrate our own motion (max turn 10 deg/tick, speed 8 px/tick)
|
|
s.heading += result.turnRate.clamp(-10.0, 10.0)
|
|
let v = result.speed.clamp(-8.0, 8.0)
|
|
s.x = min(780.0, max(20.0, s.x + v * cos(degToRad(s.heading))))
|
|
s.y = min(580.0, max(20.0, s.y + v * sin(degToRad(s.heading))))
|
|
s.speed = v
|
|
inc s.tick
|
|
# the enemy fires a power-1 bullet every 24 ticks (a 1.0 energy drop)
|
|
if s.tick mod 24 == 0:
|
|
s.enemyEnergy -= 1.0
|
|
s.fireTick = s.tick
|
|
if s.tick mod 24 == 1:
|
|
s.enemyEnergy += 1.0 # energy is restored by the harness so the next
|
|
# drop is measurable again (synthetic stream only)
|
|
|
|
proc run(ticks: int, decayShift: int): LearnedSurferModule =
|
|
putEnv(LearnedDecayShiftEnv, $decayShift)
|
|
putEnv(LearnedDecayEveryEnv, "16")
|
|
loadLearnedEnv()
|
|
result = initLearnedSurfer()
|
|
result.sbc.decayEvery = 16
|
|
result.sbc.decayShift = decayShift
|
|
var s = Sim(x: 400.0, y: 300.0, heading: 0.0, speed: 8.0,
|
|
enemyX: 400.0, enemyY: 60.0, enemyEnergy: 100.0)
|
|
for _ in 0..<ticks:
|
|
discard result.step(s)
|
|
|
|
proc runOutcome(ticks: int, decayShift: int): LearnedSurferModule =
|
|
putEnv(LearnedLabelEnv, "outcome")
|
|
result = run(ticks, decayShift)
|
|
putEnv(LearnedLabelEnv, "")
|
|
loadLearnedEnv()
|
|
|
|
proc counters(m: LearnedSurferModule): int =
|
|
var n = 0
|
|
for c in m.sbc.counters:
|
|
if c != 0'u8: inc n
|
|
n
|
|
|
|
proc testRealEvents() =
|
|
## `TR_LEARNED_REAL_EVENTS`: a wave resolves on the REAL bullet event (exact
|
|
## origin->endpoint line) instead of the arrival deadline, and is dropped the
|
|
## moment it resolves (ghost cleanup). Default-off parity is also pinned.
|
|
const Ex = 400.0
|
|
const Ey = 100.0
|
|
const Ux = 400.0
|
|
const Uy = 300.0
|
|
|
|
proc wsAt(tick: int; eEnergy: float): WorldState =
|
|
WorldState(
|
|
enemyX: Ex, enemyY: Ey, enemyEnergy: eEnergy,
|
|
selfX: Ux, selfY: Uy, selfSpeed: 8.0, selfHeading: 0.0,
|
|
arenaWidth: 800.0, arenaHeight: 600.0, tick: tick,
|
|
enemies: @[EnemyInfo(id: 1, x: Ex, y: Ey,
|
|
heading: 180.0, speed: 0.0, energy: eEnergy)])
|
|
|
|
# ── parity: with the knob off a real event changes nothing ──────────────
|
|
putEnv(LearnedLabelEnv, "histogram")
|
|
putEnv(LearnedRealEventsEnv, "")
|
|
loadLearnedEnv()
|
|
var mOff = initLearnedSurfer()
|
|
discard mOff.computeMove(wsAt(0, 100.0)) # tick 0: baseline energy sample
|
|
discard mOff.computeMove(wsAt(1, 99.0)) # tick 1: a 1.0 firepower drop
|
|
check "RE default-off: the fire is detected as a live wave",
|
|
mOff.liveWaves == 1
|
|
check "RE default-off: resolveEnemyBullet is a no-op",
|
|
(not mOff.resolveEnemyBullet(Ux, Uy, degToRad(90.0), 1, 13, true)) and
|
|
mOff.liveWaves == 1
|
|
|
|
# ── ON: the exact origin->endpoint line resolves the wave immediately ────
|
|
putEnv(LearnedRealEventsEnv, "1")
|
|
loadLearnedEnv()
|
|
var mOn = initLearnedSurfer()
|
|
discard mOn.computeMove(wsAt(0, 100.0))
|
|
discard mOn.computeMove(wsAt(1, 99.0))
|
|
check "RE on: the fire is detected as a live wave", mOn.liveWaves == 1
|
|
# endpoint = our position on the centre line (GF 0 -> bin 15) at ~nominal
|
|
# (200 px at speed 17 -> ~12 ticks after the fire at tick 1)
|
|
check "RE on: the real hit endpoint resolves the wave",
|
|
mOn.resolveEnemyBullet(Ux, Uy, degToRad(90.0), 1, 13, true)
|
|
check "RE on: the wave is dropped at once (ghost cleanup)", mOn.liveWaves == 0
|
|
check "RE on: a real resolution is counted", mOn.resolvedReal == 1
|
|
check "RE on: the exact straight line lands in the centre bin",
|
|
mOn.glob[15] == 1
|
|
check "RE on: the real flight time is recorded", mOn.lastFlightErr > -20.0
|
|
|
|
# ── ON: an unmatched wave still resolves (as a WALL MISS) at the deadline ─
|
|
var mMiss = initLearnedSurfer()
|
|
discard mMiss.computeMove(wsAt(0, 100.0))
|
|
discard mMiss.computeMove(wsAt(1, 99.0))
|
|
check "RE on: a wave is pending before any event", mMiss.liveWaves == 1
|
|
for t in 2..<40: discard mMiss.computeMove(wsAt(t, 99.0))
|
|
check "RE on: the unmatched wall wave eventually resolves",
|
|
mMiss.liveWaves == 0 and mMiss.resolvedDead >= 1
|
|
|
|
putEnv(LearnedRealEventsEnv, "")
|
|
putEnv(LearnedLabelEnv, "")
|
|
loadLearnedEnv()
|
|
|
|
proc main() =
|
|
# 1. concept
|
|
check "MovementModule concept", isMovementModule(LearnedSurferModule)
|
|
|
|
# 2./3./4. a real run
|
|
var m = run(ticks = 600, decayShift = 1)
|
|
check "waves were consumed (learned something)", counters(m) > 0
|
|
check "the global histogram has the resolutions in it",
|
|
(block:
|
|
var t = 0
|
|
for g in m.glob: t += g
|
|
t > 5)
|
|
check "wave-driven decisions were taken", m.decisions > 100
|
|
echo " resolutions in the global histogram = ",
|
|
(block:
|
|
var t = 0
|
|
for g in m.glob: t += g
|
|
t)
|
|
|
|
# the danger map is a probability distribution for a populated cell
|
|
var p: array[31, float]
|
|
m.predictState(0, 0, p)
|
|
var tot = 0.0
|
|
for v in p: tot += v
|
|
check "danger map sums to 1", abs(tot - 1.0) < 1e-9
|
|
|
|
# 5. decay: with shift 0 the counters only grow
|
|
let md = run(ticks = 600, decayShift = 0)
|
|
check "decayShift=0 keeps a memory (counters present)", counters(md) > 0
|
|
check "no-decay counts every resolution (no forgetting)",
|
|
(block:
|
|
var t = 0
|
|
for g in md.glob: t += g
|
|
t >= 20) # ~1 fire per 24 ticks, ~15-tick flight
|
|
check "the decaying run keeps less mass than the no-decay run",
|
|
(block:
|
|
var a = 0
|
|
for g in m.glob: a += g
|
|
var b = 0
|
|
for g in md.glob: b += g
|
|
a <= b)
|
|
|
|
# 6. determinism
|
|
let a = run(ticks = 300, decayShift = 1)
|
|
let b = run(ticks = 300, decayShift = 1)
|
|
check "deterministic", counters(a) == counters(b)
|
|
|
|
# 7. ablation knob: TR_LEARNED_GLOBAL ignores the state
|
|
putEnv(LearnedGlobalEnv, "1")
|
|
loadLearnedEnv()
|
|
let g = run(ticks = 200, decayShift = 1)
|
|
check "TR_LEARNED_GLOBAL still moves (prior-only map)", g.decisions > 20
|
|
delEnv(LearnedGlobalEnv)
|
|
|
|
# 8. outcome label: the 2-class counted SBC learns and reads a probability
|
|
putEnv(LearnedGlobalEnv, "")
|
|
let o = runOutcome(ticks = 600, decayShift = 1)
|
|
check "TR_LEARNED_LABEL=outcome still moves", o.decisions > 100
|
|
check "outcome memory accumulated",
|
|
(block:
|
|
var n = 0
|
|
for c in o.outcome.counters:
|
|
if c != 0'u8: inc n
|
|
n > 0)
|
|
check "outcome hit prior is a probability in [0,1]",
|
|
(block:
|
|
let p = o.predictHit(0, 0, 10)
|
|
p >= 0.0 and p <= 1.0)
|
|
let o2 = runOutcome(ticks = 300, decayShift = 1)
|
|
let o3 = runOutcome(ticks = 300, decayShift = 1)
|
|
check "outcome mode deterministic",
|
|
(block:
|
|
var n2, n3 = 0
|
|
for c in o2.outcome.counters:
|
|
if c != 0'u8: inc n2
|
|
for c in o3.outcome.counters:
|
|
if c != 0'u8: inc n3
|
|
n2 == n3)
|
|
putEnv(LearnedDecayShiftEnv, "")
|
|
putEnv(LearnedDecayEveryEnv, "")
|
|
loadLearnedEnv()
|
|
|
|
echo ""
|
|
testRealEvents()
|
|
echo ""
|
|
echo "checks=", checks, " failures=", failures
|
|
if failures > 0: quit(1)
|
|
|
|
main()
|