chore: rename libs→common_libs, all bot dirs to _garage suffix, fix all path refs
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
Executable
BIN
Binary file not shown.
@@ -0,0 +1,36 @@
|
||||
## diag_ppo_fullbuffer.nim — run ppoUpdate on a FULL 8192-transition buffer.
|
||||
## Verifies whether the round-10 death (first ppoUpdate at buffer capacity)
|
||||
## is a real crash in ppoUpdate or purely the saveAdamStates empty-shape bug.
|
||||
|
||||
import std/[math, random]
|
||||
import arraymancer
|
||||
import "../network"
|
||||
import "../training"
|
||||
|
||||
var ac = initActorCritic()
|
||||
var adam: ACAdamStates # uninitialised → ppoUpdate must reinit (fresh-process path)
|
||||
var buf = initTrajectoryBuffer()
|
||||
|
||||
randomize(1)
|
||||
while buf.len < MAX_TRANSITIONS:
|
||||
var t: Transition
|
||||
for i in 0..<STATE_DIM: t.state[i] = rand(1.0'f32) - 0.5'f32
|
||||
for i in 0..<ACTION_DIM: t.action[i] = rand(1.0'f32) - 0.5'f32
|
||||
t.logProb = rand(1.0'f32) - 1.0'f32
|
||||
t.reward = rand(0.02'f32) - 0.01'f32
|
||||
t.value = rand(0.1'f32)
|
||||
t.done = buf.len mod 300 == 299
|
||||
buf.add(t)
|
||||
|
||||
echo "buffer.len = ", buf.len, " (MAX=", MAX_TRANSITIONS, ")"
|
||||
var ac2 = ac
|
||||
var adam2: ACAdamStates
|
||||
try:
|
||||
let m = ppoUpdate(ac2, buf, lastValue = 0.0'f32, adamStates = adam2)
|
||||
echo "ppoUpdate OK: aLoss=", m.actorLoss, " vLoss=", m.valueLoss, " gNorm=", m.gradNorm
|
||||
except CatchableError as e:
|
||||
echo "CAUGHT CatchableError: ", e.msg
|
||||
echo getStackTrace(e)
|
||||
except Defect as e:
|
||||
echo "CAUGHT Defect: ", e.msg
|
||||
echo getStackTrace(e)
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,31 @@
|
||||
## diag_savecheckpoint.nim — reproduce the startup+saveCheckpoint path that
|
||||
## crashes with `io_npy.nim(143, 3) `0 < t.shape.len`` under the boot server.
|
||||
## Run: nim c -d:release tests/diag_savecheckpoint.nim && ./tests/diag_savecheckpoint
|
||||
|
||||
import std/[os]
|
||||
import arraymancer
|
||||
import "../network"
|
||||
import "../training"
|
||||
import "../weights"
|
||||
|
||||
const weightsRoot = currentSourcePath().parentDir.parentDir / "weights"
|
||||
|
||||
var ac = initActorCritic()
|
||||
var adam: ACAdamStates
|
||||
|
||||
let loadResult = loadBestAvailable(ac, adam, weightsRoot)
|
||||
echo "loaded=", loadResult.loaded, " roundNum=", loadResult.roundNum
|
||||
echo "adam.initialized=", adam.initialized
|
||||
echo "adam.aw1.m.shape=", adam.aw1.m.shape, " aw1.v.shape=", adam.aw1.v.shape
|
||||
echo "adam.logStd.m.shape=", adam.logStd.m.shape
|
||||
echo "adam.aw1.t=", adam.aw1.t
|
||||
|
||||
try:
|
||||
saveCheckpoint(ac, adam, weightsRoot, 1)
|
||||
echo "saveCheckpoint OK"
|
||||
except CatchableError as e:
|
||||
echo "CAUGHT CatchableError: ", e.msg
|
||||
echo getStackTrace(e)
|
||||
except Defect as e:
|
||||
echo "CAUGHT Defect: ", e.msg
|
||||
echo getStackTrace(e)
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,69 @@
|
||||
## Assert-based tests for actions.nim.
|
||||
## Run: nim c -r tests/test_actions.nim
|
||||
|
||||
import arraymancer
|
||||
import std/math
|
||||
import "../actions"
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
let arenaW = 1200.0
|
||||
let arenaH = 800.0
|
||||
|
||||
# Build a 6-element zero tensor and a helper to set individual values
|
||||
proc makeRaw(vals: array[6, float32]): Tensor[float32] =
|
||||
result = zeros[float32](6)
|
||||
for i in 0 ..< 6: result[i] = vals[i]
|
||||
|
||||
# --- 6-dim input produces a valid BotActions ---
|
||||
block basicDecode:
|
||||
let raw = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 0.0, 0.0])
|
||||
let acts = mapActions(raw, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
# sigmoid(0)*arenaW = 0.5*1200 = 600, sigmoid(0)*arenaH = 0.5*800 = 400
|
||||
check abs(acts.gotoX - 600.0) < 1e-6, "gotoX = sigmoid(0)*arenaW"
|
||||
check abs(acts.gotoY - 400.0) < 1e-6, "gotoY = sigmoid(0)*arenaH"
|
||||
check abs(acts.aimToX - 600.0) < 1e-6, "aimToX = sigmoid(0)*arenaW"
|
||||
check abs(acts.aimToY - 400.0) < 1e-6, "aimToY = sigmoid(0)*arenaH"
|
||||
|
||||
# --- Coordinates bounded to arena size ---
|
||||
block coordBounds:
|
||||
# Large positive raw → sigmoid ≈ 1 → close to arenaW/arenaH
|
||||
let rawHigh = makeRaw([100.0'f32, 100.0, 100.0, 100.0, 0.0, 0.0])
|
||||
let actsHigh = mapActions(rawHigh, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
check actsHigh.gotoX <= arenaW + 1e-9, "gotoX <= arenaWidth"
|
||||
check actsHigh.gotoY <= arenaH + 1e-9, "gotoY <= arenaHeight"
|
||||
check actsHigh.gotoX >= 0.0, "gotoX >= 0"
|
||||
# Large negative raw → sigmoid ≈ 0 → close to 0
|
||||
let rawLow = makeRaw([-100.0'f32, -100.0, -100.0, -100.0, 0.0, 0.0])
|
||||
let actsLow = mapActions(rawLow, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
check actsLow.gotoX >= -1e-9, "gotoX >= 0 (low raw)"
|
||||
check actsLow.gotoY >= -1e-9, "gotoY >= 0 (low raw)"
|
||||
|
||||
# --- Fire triggers correctly ---
|
||||
block fireTrigger:
|
||||
# tanh(positive) >= 0 → fire when gunHeat = 0
|
||||
let rawFire = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 1.0, 0.0])
|
||||
let actsFire = mapActions(rawFire, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
check actsFire.shouldFire, "positive tanh → should fire when gun cool"
|
||||
|
||||
# tanh(negative) < 0 → no fire
|
||||
let rawNoFire = makeRaw([0.0'f32, 0.0, 0.0, 0.0, -1.0, 0.0])
|
||||
let actsNoFire = mapActions(rawNoFire, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
check not actsNoFire.shouldFire, "negative tanh → no fire"
|
||||
|
||||
# gunHeat > 0 → no fire even with positive decision
|
||||
let actsHot = mapActions(rawFire, 0.5, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
check not actsHot.shouldFire, "positive tanh but gun hot → no fire"
|
||||
|
||||
# --- Fire power in [0.1, 3.0] ---
|
||||
block firePowerRange:
|
||||
let rawMin = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 0.0, -100.0])
|
||||
let rawMax = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 0.0, 100.0])
|
||||
let actsMin = mapActions(rawMin, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
let actsMax = mapActions(rawMax, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
check actsMin.firePower >= 0.1 - 1e-6, "firePower >= 0.1"
|
||||
check actsMax.firePower <= 3.0 + 1e-6, "firePower <= 3.0"
|
||||
|
||||
echo "test_actions: all passed"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,54 @@
|
||||
## Assert-based tests for controllers.nim.
|
||||
## Run: nim c -r tests/test_controllers.nim
|
||||
|
||||
import std/math
|
||||
import "../controllers"
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# gotoTick tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
block forwardMovement:
|
||||
# Bot at (0,0) facing East (geometric 0°), target due east at (100,0) → bearing=0 → forward
|
||||
let (spd2, turn2) = gotoTick(100.0, 0.0, 0.0, 0.0, 0.0, 0.0)
|
||||
check spd2 > 0.0, "forward: targetSpeed should be positive"
|
||||
check abs(turn2) < 1e-9, "forward: no turn needed when already aimed"
|
||||
|
||||
block reverseMovement:
|
||||
# Bot at (100,0) facing East (direction=0), target at (0,0) — directly behind
|
||||
# bearing = normalizeRelativeAngle(180 - 0) = 180 → |bearing|>90 → reverse
|
||||
let (spd, _) = gotoTick(0.0, 0.0, 100.0, 0.0, 0.0, 0.0)
|
||||
check spd < 0.0, "reverse: targetSpeed should be negative when target is behind"
|
||||
|
||||
block turnRateClamping:
|
||||
# Bot at (0,0) facing North (game north = geometric 90°, so direction=90 in geometric)
|
||||
# Target at (100,0) = East. bearing = normalizeRelativeAngle(0 - 90) = -90 → still ≤90
|
||||
# Use large perpendicular target so bearing is 89°, and high speed → small maxTurn
|
||||
# At speed=8, calcMaxTurnRate = 10 - 0.75*8 = 4°
|
||||
# Bot facing East (0°), target at angle 89° bearing (just under 90)
|
||||
let (_, turn) = gotoTick(100.0 * cos(89.0 * PI / 180.0), 100.0 * sin(89.0 * PI / 180.0), 0.0, 0.0, 0.0, 8.0)
|
||||
check abs(turn) <= 4.0 + 1e-9, "turn rate clamped to calcMaxTurnRate at speed=8 (max 4°)"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# aimToTick tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
block gunShortestArc:
|
||||
# Gun facing East (0°), target due north (geometric 90°) → turn left +90° but clamped to 20
|
||||
let rate = aimToTick(0.0, 100.0, 0.0, 0.0, 0.0)
|
||||
check rate > 0.0, "gun shortest arc: should turn toward target"
|
||||
# Gun facing East (0°), target due south (geometric 270° → normalised -90°)
|
||||
let rate2 = aimToTick(0.0, -100.0, 0.0, 0.0, 0.0)
|
||||
check rate2 < 0.0, "gun shortest arc: should turn other way for target behind"
|
||||
|
||||
block gunTurnRateClamping:
|
||||
# 180° away → clamped to ±20
|
||||
let rate = aimToTick(-100.0, 0.0, 0.0, 0.0, 0.0) # target west, gun east
|
||||
check abs(rate) <= 20.0 + 1e-9, "gun turn rate clamped to ±MAX_GUN_TURN_RATE"
|
||||
check abs(abs(rate) - 20.0) < 1e-9, "gun turn rate at max when 180° away"
|
||||
|
||||
echo "test_controllers: all passed"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,58 @@
|
||||
## test_network.nim — assert-based tests for network.nim and actions.nim.
|
||||
## Run: nim c --threads:on tests/test_network.nim && ./tests/test_network
|
||||
|
||||
import arraymancer
|
||||
import std/[math, strformat]
|
||||
import ../network
|
||||
import ../actions
|
||||
|
||||
func isFiniteF(x: float32): bool = classify(x) notin {fcNan, fcInf, fcNegInf}
|
||||
func isNaNF(x: float32): bool = classify(x) == fcNan
|
||||
|
||||
when isMainModule:
|
||||
# ---- MLP forward shape ----
|
||||
let mlp = initMLP(STATE_DIM, 64, ACTION_DIM)
|
||||
let inp = zeros[float32](STATE_DIM)
|
||||
let mlpOut = mlp.forward(inp)
|
||||
assert mlpOut.shape[0] == ACTION_DIM, &"MLP output shape wrong: {mlpOut.shape}"
|
||||
|
||||
# ---- ActorCritic actorForward ----
|
||||
let ac = initActorCritic()
|
||||
let state = zeros[float32](STATE_DIM)
|
||||
let (acts, logP) = ac.actorForward(state)
|
||||
assert acts.shape[0] == ACTION_DIM, &"actorForward actions shape wrong: {acts.shape}"
|
||||
assert not isNaNF(logP), "logProb is NaN"
|
||||
assert isFiniteF(logP), &"logProb not finite: {logP}"
|
||||
|
||||
# ---- criticForward ----
|
||||
let v = ac.criticForward(state)
|
||||
assert not isNaNF(v), "critic value is NaN"
|
||||
assert isFiniteF(v), &"critic value not finite: {v}"
|
||||
|
||||
# ---- logStd floor: collapsing logStd should not break actorForward ----
|
||||
var ac2 = initActorCritic()
|
||||
for i in 0..<ACTION_DIM: ac2.logStd[i] = -10.0'f32
|
||||
let (acts2, logP2) = ac2.actorForward(state)
|
||||
assert acts2.shape[0] == ACTION_DIM, "acts2 shape wrong after logStd=-10"
|
||||
assert isFiniteF(logP2), &"logP2 not finite with floored logStd: {logP2}"
|
||||
|
||||
# ---- action mapping ranges ----
|
||||
let raw = randomNormalTensor[float32](ACTION_DIM)
|
||||
let speed = 4.0'f32
|
||||
# arena 800×600, bot at centre, heading north, gun north
|
||||
let botActs = mapActions(raw, 0.0'f32, 800.0, 600.0, 400.0, 300.0, 0.0, speed.float, 0.0, 400.0, 300.0) # gunHeat=0 → fire allowed
|
||||
|
||||
assert botActs.targetSpeed >= -8.0'f32 and botActs.targetSpeed <= 8.0'f32,
|
||||
&"targetSpeed out of range: {botActs.targetSpeed}"
|
||||
|
||||
assert botActs.gunTurnRate >= -20.0'f32 and botActs.gunTurnRate <= 20.0'f32,
|
||||
&"gunTurnRate out of range: {botActs.gunTurnRate}"
|
||||
|
||||
assert botActs.firePower >= 0.1'f32 and botActs.firePower <= 3.0'f32,
|
||||
&"firePower out of range: {botActs.firePower}"
|
||||
|
||||
# shouldFire=false when gunHeat > 0
|
||||
let noFire = mapActions(raw, 1.0'f32, 800.0, 600.0, 400.0, 300.0, 0.0, speed.float, 0.0, 400.0, 300.0)
|
||||
assert not noFire.shouldFire, "shouldFire should be false when gunHeat > 0"
|
||||
|
||||
echo "All tests passed"
|
||||
@@ -0,0 +1,42 @@
|
||||
## Regression test: radar lock must hold on a stationary target.
|
||||
## Run: nim c -r tests/test_radar_lock.nim
|
||||
|
||||
import std/[math, strformat]
|
||||
import "../enemy_tracker"
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
# Bot at arena center; stationary enemy due north (same x, higher y).
|
||||
# Tank Royale: y increases northward.
|
||||
# Radar uses math convention (0°=east, CCW+); north = 90° in that system.
|
||||
|
||||
let botX = 400.0
|
||||
let botY = 300.0
|
||||
let enemyX = 400.0 # same x → dx = 0
|
||||
let enemyY = 500.0 # north of bot → dy > 0
|
||||
let trueBearing = 90.0 # north in math convention (0=east, CCW+)
|
||||
|
||||
# Radar starts pointing at the enemy (radarDirection = 90°, due north in math convention).
|
||||
var radarDir = 90.0
|
||||
|
||||
var tracker = initEnemyTracker()
|
||||
# Prime with contact at the known position
|
||||
tracker.update(enemyX, enemyY, 0.0, 0.0, 100.0)
|
||||
|
||||
echo "Tick | radarDir | trueBearing | bearingErr"
|
||||
for tick in 1 .. 20:
|
||||
let rate = tracker.getRadarTurnRate(botX, botY, 0.0, radarDir)
|
||||
radarDir = (radarDir + rate + 360.0) mod 360.0
|
||||
|
||||
# Simulate a successful scan every tick (enemy is stationary)
|
||||
tracker.update(enemyX, enemyY, 0.0, 0.0, 100.0)
|
||||
|
||||
# Bearing error: signed difference, wrapped to [-180, 180]
|
||||
let err = ((radarDir - trueBearing) + 540.0) mod 360.0 - 180.0
|
||||
echo &" {tick:2d} | {radarDir:8.3f}° | {trueBearing:8.3f}° | {err:+.3f}°"
|
||||
|
||||
check abs(err) <= 15.0, &"tick {tick}: radar {radarDir:.1f}° drifted > 15° from target {trueBearing:.1f}°"
|
||||
|
||||
echo "All radar lock tests passed"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,187 @@
|
||||
## Assert-based tests for EnemyTracker and StateVector.
|
||||
## Run: nim c -r tests/test_state.nim
|
||||
|
||||
import std/[math, strformat]
|
||||
import arraymancer
|
||||
|
||||
# Import from parent dir
|
||||
import "../enemy_tracker"
|
||||
import "../state_vector"
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# EnemyTracker tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
block testBasicUpdate:
|
||||
var t = initEnemyTracker()
|
||||
t.update(200.0, 300.0, 90.0, 5.0, 80.0)
|
||||
check t.hasContact, "hasContact after update"
|
||||
check t.current.x == 200.0, "x after update"
|
||||
check t.current.y == 300.0, "y after update"
|
||||
check t.current.direction == 90.0, "direction after update"
|
||||
check t.current.speed == 5.0, "speed after update"
|
||||
check t.current.energy == 80.0, "energy after update"
|
||||
check t.current.ticksSinceLastScan == 0, "ticksSinceLastScan reset"
|
||||
|
||||
block testFireDetection:
|
||||
var t = initEnemyTracker()
|
||||
# First update sets prevEnergy
|
||||
t.update(100.0, 100.0, 0.0, 0.0, 100.0)
|
||||
# Second update: energy drop of 3.0 → enemy fired power 3.0
|
||||
t.update(100.0, 100.0, 0.0, 0.0, 97.0)
|
||||
check t.current.hasFired, "hasFired when energy drops by 3.0"
|
||||
check abs(t.current.lastFirePower - 3.0) < 0.001, "lastFirePower == 3.0"
|
||||
|
||||
block testNoFireOnSmallDrop:
|
||||
var t = initEnemyTracker()
|
||||
t.update(100.0, 100.0, 0.0, 0.0, 100.0)
|
||||
# Drop of 0.05 — below MIN_FIRE_POWER threshold
|
||||
t.update(100.0, 100.0, 0.0, 0.0, 99.95)
|
||||
check not t.current.hasFired, "no fire on small energy drop"
|
||||
|
||||
block testDeadReckoning:
|
||||
var t = initEnemyTracker()
|
||||
# direction=0° in Tank Royale means north (y increases)
|
||||
t.update(100.0, 100.0, 0.0, 5.0, 100.0)
|
||||
t.deadReckon()
|
||||
# x unchanged (sin 0° = 0), y increases by speed (cos 0° = 1)
|
||||
check abs(t.current.x - 100.0) < 0.001, "dead reckon: x unchanged for dir=0"
|
||||
check abs(t.current.y - 105.0) < 0.001, "dead reckon: y += speed for dir=0"
|
||||
check t.current.ticksSinceLastScan == 1, "ticksSinceLastScan incremented"
|
||||
|
||||
block testDeadReckonEast:
|
||||
var t = initEnemyTracker()
|
||||
# direction=90° → east (sin 90° = 1, cos 90° = 0)
|
||||
t.update(100.0, 100.0, 90.0, 5.0, 100.0)
|
||||
t.deadReckon()
|
||||
check abs(t.current.x - 105.0) < 0.001, "dead reckon east: x += speed"
|
||||
check abs(t.current.y - 100.0) < 0.001, "dead reckon east: y unchanged"
|
||||
|
||||
block testHistoryWindow:
|
||||
var t = initEnemyTracker()
|
||||
# Feed 6 updates — history should hold last 5
|
||||
for i in 1 .. 6:
|
||||
t.update(float64(i) * 10.0, float64(i) * 20.0, 0.0, float64(i), 100.0)
|
||||
check t.historyCount == 5, "historyCount capped at 5"
|
||||
# history[0] should be the second-to-last scan (i=5)
|
||||
check abs(t.history[0].x - 50.0) < 0.001, "history[0].x == 50 (i=5)"
|
||||
check abs(t.history[4].x - 10.0) < 0.001, "history[4].x == 10 (i=1)"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# StateVector tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
block testStateVectorLength:
|
||||
var t = initEnemyTracker()
|
||||
t.update(400.0, 300.0, 45.0, 3.0, 80.0)
|
||||
let bot = BotStateData(
|
||||
x: 200.0, y: 200.0, direction: 90.0, speed: 4.0, energy: 50.0,
|
||||
gunDirection: 180.0, gunHeat: 0.5,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
check sv.shape == [57], "state vector has 57 elements"
|
||||
|
||||
block testStateVectorRange:
|
||||
var t = initEnemyTracker()
|
||||
t.update(400.0, 300.0, 180.0, 8.0, 100.0)
|
||||
let bot = BotStateData(
|
||||
x: 800.0, y: 600.0, direction: 360.0, speed: 8.0, energy: 100.0,
|
||||
gunDirection: 360.0, gunHeat: 1.8,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
for i in 0 ..< 57:
|
||||
check sv[i] >= -2.0f32 and sv[i] <= 2.0f32,
|
||||
&"sv[{i}]={sv[i]} out of [-2,2] range"
|
||||
|
||||
block testWallDistances:
|
||||
# Bot at (100, 200) in 800×600 arena
|
||||
# wallMax = max(800, 600) = 800
|
||||
# top = (600 - 200) / 800 = 400/800 = 0.5
|
||||
# bottom = 200 / 800 = 0.25
|
||||
# left = 100 / 800 = 0.125
|
||||
# right = (800 - 100) / 800 = 700/800 = 0.875
|
||||
var t = initEnemyTracker()
|
||||
let bot = BotStateData(
|
||||
x: 100.0, y: 200.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
||||
gunDirection: 0.0, gunHeat: 0.0,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
check abs(sv[18] - 0.5f32) < 0.001f32, "top wall = 0.5"
|
||||
check abs(sv[19] - 0.25f32) < 0.001f32, "bottom wall = 0.25"
|
||||
check abs(sv[20] - 0.125f32) < 0.001f32, "left wall = 0.125"
|
||||
check abs(sv[21] - 0.875f32) < 0.001f32, "right wall = 0.875"
|
||||
|
||||
block testRelativeBearing:
|
||||
# Game convention: north=0°, CW. arctan2(dx,dy) used.
|
||||
# Bot at (0,0) dir=0°. Enemy at (0,100) → due north → absDir=0°.
|
||||
# relBearing = (0 - 0 + 540) mod 360 - 180 = 0°. normalized = 0/180 = 0.0
|
||||
var t = initEnemyTracker()
|
||||
t.update(0.0, 100.0, 0.0, 0.0, 100.0)
|
||||
let bot = BotStateData(
|
||||
x: 0.0, y: 0.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
||||
gunDirection: 0.0, gunHeat: 0.0,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
check abs(sv[16] - 0.0f32) < 0.01f32, "relative bearing = 0.0 (due north), got " & $sv[16]
|
||||
|
||||
block testRelativeBearingEast:
|
||||
# Enemy at (100,0) → due east → absDir=90°.
|
||||
# relBearing = (90 - 0 + 540) mod 360 - 180 = 90°. normalized = 90/180 = 0.5
|
||||
var t = initEnemyTracker()
|
||||
t.update(100.0, 0.0, 0.0, 0.0, 100.0)
|
||||
let bot = BotStateData(
|
||||
x: 0.0, y: 0.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
||||
gunDirection: 0.0, gunHeat: 0.0,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
check abs(sv[16] - 0.5f32) < 0.01f32, "relative bearing east = 0.5, got " & $sv[16]
|
||||
|
||||
block testHistoryPaddedWhenEmpty:
|
||||
var t = initEnemyTracker()
|
||||
let bot = BotStateData(
|
||||
x: 400.0, y: 300.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
||||
gunDirection: 0.0, gunHeat: 0.0,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
for i in 22 ..< 42:
|
||||
check sv[i] == 0.0f32, &"history slot {i} should be 0 when no contact"
|
||||
# indices 42-43 (goto inputs) default to 0 when not provided
|
||||
check sv[42] == 0.0f32, "sv[42] (remainingGotoDistance) should default to 0"
|
||||
check sv[43] == 0.0f32, "sv[43] (remainingGunAngle) should default to 0"
|
||||
# indices 44-55 (bullet slots) default to 0 when no bullets provided
|
||||
for i in 44 ..< 56:
|
||||
check sv[i] == 0.0f32, &"bullet slot {i} should be 0 when no bullets"
|
||||
# index 56 (staleness): no contact so ticksSinceLastScan=0 → 0/30 = 0
|
||||
check sv[56] == 0.0f32, "sv[56] staleness should be 0 when no contact"
|
||||
|
||||
block testBulletSlots:
|
||||
var t = initEnemyTracker()
|
||||
t.update(400.0, 300.0, 0.0, 0.0, 100.0) # enemy at (400,300)
|
||||
let bot = BotStateData(
|
||||
x: 200.0, y: 200.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
||||
gunDirection: 0.0, gunHeat: 0.0,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
# Bullet at (300,250), power 1.0 → speed = 20-3 = 17
|
||||
# relX = 300-200 = 100, relY = 250-200 = 50 (relative to bot, not enemy)
|
||||
# dist to bot = sqrt(100^2+50^2) ≈ 111.8, ticks ≈ 111.8/17 ≈ 6.58
|
||||
let b = BulletData(x: 300.0, y: 250.0, power: 1.0)
|
||||
let sv = buildStateVector(bot, t, 0.0, 0.0, [b], 1)
|
||||
check abs(sv[44] - (100.0/800.0).float32) < 0.001f32, "bullet relX"
|
||||
check abs(sv[45] - (50.0/600.0).float32) < 0.001f32, "bullet relY"
|
||||
check abs(sv[46] - (17.0/20.0).float32) < 0.001f32, "bullet speed norm"
|
||||
# second slot should be zero-padded
|
||||
for i in 48 ..< 57:
|
||||
check sv[i] == 0.0f32, &"unused bullet slot {i} should be 0"
|
||||
|
||||
echo "All tests passed"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,175 @@
|
||||
## test_training.nim — assert-based tests for training.nim
|
||||
## Run: nim c tests/test_training.nim && ./tests/test_training
|
||||
|
||||
import std/[math, random]
|
||||
import arraymancer
|
||||
import "../network"
|
||||
import "../training"
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
# ── computeTickReward ─────────────────────────────────────────────────────────
|
||||
|
||||
block testTickReward:
|
||||
# I lost 2, enemy lost 10 → reward = -2 - (-10) = 8
|
||||
# + default closeness shaping 0.01*(1-0/maxDist) = 0.01 (gunBearingAbs=180 → 0)
|
||||
let r = computeTickReward(-2.0'f32, -10.0'f32)
|
||||
check abs(r - 8.01'f32) < 1e-6'f32, "computeTickReward(-2, -10) == 8.01, got " & $r
|
||||
|
||||
# ── computeRoundReward ────────────────────────────────────────────────────────
|
||||
|
||||
block testRoundReward:
|
||||
let r = computeRoundReward(350.0'f32)
|
||||
check abs(r - 7.0'f32) < 1e-6'f32, "computeRoundReward(350) == 7.0, got " & $r
|
||||
# bounded: long-battle cumulative scores must saturate, not blow the value scale
|
||||
check abs(computeRoundReward(89299.0'f32) - 8.0'f32) < 1e-6'f32,
|
||||
"computeRoundReward(89299) == 8.0 (capped), got " & $computeRoundReward(89299.0'f32)
|
||||
|
||||
# ── TrajectoryBuffer ──────────────────────────────────────────────────────────
|
||||
|
||||
block testBuffer:
|
||||
var buf = initTrajectoryBuffer()
|
||||
check buf.len == 0, "empty buffer len == 0"
|
||||
|
||||
let t1 = Transition(state: zeros[float32](STATE_DIM).stateToArr,
|
||||
action: zeros[float32](ACTION_DIM).actionToArr,
|
||||
logProb: -1.0'f32, reward: 0.5'f32, value: 0.3'f32)
|
||||
buf.add(t1)
|
||||
buf.add(t1)
|
||||
buf.add(t1)
|
||||
check buf.len == 3, "buffer len == 3 after 3 adds"
|
||||
|
||||
buf.clear()
|
||||
check buf.len == 0, "buffer len == 0 after clear"
|
||||
|
||||
# ── computeGAE — hand-calculated 3-step ──────────────────────────────────────
|
||||
|
||||
block testGAE:
|
||||
# rewards = [1.0, 0.0, 1.0], values = [0.5, 0.5, 0.5], lastValue = 0.0
|
||||
# gamma = 0.99, lam = 0.95
|
||||
# delta_2 = 1.0 + 0.99*0.0 - 0.5 = 0.5
|
||||
# adv_2 = 0.5
|
||||
# delta_1 = 0.0 + 0.99*0.5 - 0.5 = -0.005
|
||||
# adv_1 = -0.005 + 0.99*0.95*0.5 ≈ -0.005 + 0.47025 = 0.46525
|
||||
# delta_0 = 1.0 + 0.99*0.5 - 0.5 = 0.995
|
||||
# adv_0 = 0.995 + 0.99*0.95*0.46525 ≈ 0.995 + 0.43744 = 1.43244
|
||||
let (adv, ret) = computeGAE(
|
||||
rewards = @[1.0'f32, 0.0'f32, 1.0'f32],
|
||||
values = @[0.5'f32, 0.5'f32, 0.5'f32],
|
||||
lastValue = 0.0'f32,
|
||||
gamma = 0.99'f32,
|
||||
lam = 0.95'f32
|
||||
)
|
||||
|
||||
check abs(adv[2] - 0.5'f32) < 1e-4'f32,
|
||||
"adv[2] should be ~0.5, got " & $adv[2]
|
||||
check abs(adv[1] - 0.46525'f32) < 1e-3'f32,
|
||||
"adv[1] should be ~0.46525, got " & $adv[1]
|
||||
check abs(adv[0] - 1.43244'f32) < 1e-2'f32,
|
||||
"adv[0] should be ~1.43244, got " & $adv[0]
|
||||
|
||||
# returns = adv + values
|
||||
check abs(ret[2] - (0.5'f32 + 0.5'f32)) < 1e-4'f32, "ret[2] = adv[2] + 0.5"
|
||||
check abs(ret[0] - (adv[0] + 0.5'f32)) < 1e-4'f32, "ret[0] = adv[0] + 0.5"
|
||||
|
||||
# ── ppoUpdate runs without crash; weights change ──────────────────────────────
|
||||
|
||||
block testPpoUpdate:
|
||||
randomize(42)
|
||||
var ac = initActorCritic()
|
||||
|
||||
# Save a copy of w1 before update
|
||||
let w1Before = ac.actor.w1.clone()
|
||||
|
||||
var buf = initTrajectoryBuffer()
|
||||
for _ in 0..<10:
|
||||
let s = randomNormalTensor[float32](STATE_DIM)
|
||||
let a = randomNormalTensor[float32](ACTION_DIM)
|
||||
let lp = ac.computeLogProb(s, a)
|
||||
let v = ac.criticForward(s)
|
||||
buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp, reward: 0.1'f32, value: v))
|
||||
|
||||
var adam: ACAdamStates
|
||||
discard ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam, epochs = 2, miniBatchSize = 5)
|
||||
|
||||
# Weights should have changed — compare flattened
|
||||
let n = ac.actor.w1.shape[0] * ac.actor.w1.shape[1]
|
||||
let w1After = ac.actor.w1.reshape(n)
|
||||
let w1Flat = w1Before.reshape(n)
|
||||
var changed = false
|
||||
for i in 0..<n:
|
||||
if abs(w1After[i] - w1Flat[i]) > 1e-9'f32:
|
||||
changed = true
|
||||
break
|
||||
check changed, "actor w1 should change after ppoUpdate"
|
||||
|
||||
# ── ppoUpdate on constant-reward trajectory: zero-variance guard ─────────────
|
||||
# A passive round has near-constant per-tick rewards; with constant values the
|
||||
# GAE advantages are identical → zero variance. The normalization must not
|
||||
# amplify/NaN on this — update must complete with finite losses.
|
||||
|
||||
block testPpoUpdateConstantReward:
|
||||
randomize(43)
|
||||
var ac = initActorCritic()
|
||||
var buf = initTrajectoryBuffer()
|
||||
for _ in 0..<64:
|
||||
let s = randomNormalTensor[float32](STATE_DIM)
|
||||
let a = randomNormalTensor[float32](ACTION_DIM)
|
||||
let lp = ac.computeLogProb(s, a)
|
||||
buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp,
|
||||
reward: 0.05'f32, value: 0.5'f32)) # constant reward+value
|
||||
|
||||
var adam: ACAdamStates
|
||||
let m = ppoUpdate(ac, buf, lastValue = 0.5'f32, adamStates = adam,
|
||||
epochs = 2, miniBatchSize = 16)
|
||||
check m.actorLoss == m.actorLoss, "actorLoss NaN on constant-reward round"
|
||||
check m.valueLoss == m.valueLoss, "valueLoss NaN on constant-reward round"
|
||||
check m.gradNorm == m.gradNorm, "gradNorm NaN on constant-reward round"
|
||||
|
||||
# ── ppoUpdate on normal-reward trajectory: finite losses ─────────────────────
|
||||
|
||||
block testPpoUpdateNormalReward:
|
||||
randomize(44)
|
||||
var ac = initActorCritic()
|
||||
var buf = initTrajectoryBuffer()
|
||||
for i in 0..<64:
|
||||
let s = randomNormalTensor[float32](STATE_DIM)
|
||||
let a = randomNormalTensor[float32](ACTION_DIM)
|
||||
let lp = ac.computeLogProb(s, a)
|
||||
let v = ac.criticForward(s)
|
||||
let r = 0.05'f32 + 0.5'f32 * sin(float32(i))
|
||||
buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp, reward: r, value: v))
|
||||
|
||||
var adam: ACAdamStates
|
||||
let m = ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam,
|
||||
epochs = 2, miniBatchSize = 16)
|
||||
check m.actorLoss == m.actorLoss, "actorLoss NaN on normal-reward round"
|
||||
check m.valueLoss == m.valueLoss, "valueLoss NaN on normal-reward round"
|
||||
check m.gradNorm == m.gradNorm, "gradNorm NaN on normal-reward round"
|
||||
|
||||
# ── logStd ceiling: raw param must never drift above the collection clamp ─────
|
||||
# Regression for the train/collection std mismatch: logStd starting above the
|
||||
# ceiling (as in the warm-start snapshot) must be clamped back to [floor, ceiling]
|
||||
# by the first Adam step, so recomputed logP matches the acting policy's std.
|
||||
|
||||
block testLogStdCeilingClamp:
|
||||
randomize(45)
|
||||
var ac = initActorCritic()
|
||||
ac.logStd = newTensor[float32](ACTION_DIM).map(proc(v: float32): float32 = 3.0'f32)
|
||||
var buf = initTrajectoryBuffer()
|
||||
for _ in 0..<16:
|
||||
let s = randomNormalTensor[float32](STATE_DIM)
|
||||
let a = randomNormalTensor[float32](ACTION_DIM)
|
||||
let lp = ac.computeLogProb(s, a)
|
||||
buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp,
|
||||
reward: 0.1'f32, value: 0.5'f32))
|
||||
var adam: ACAdamStates
|
||||
discard ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam,
|
||||
epochs = 1, miniBatchSize = 16)
|
||||
for v in ac.logStd:
|
||||
check v <= logStdCeiling + 1e-6'f32, "logStd above ceiling after ppoUpdate"
|
||||
check v >= logStdFloor - 1e-6'f32, "logStd below floor after ppoUpdate"
|
||||
|
||||
echo "All tests passed"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,105 @@
|
||||
## test_weights.nim — assert-based tests for weights.nim
|
||||
## Run: nim c tests/test_weights.nim && ./tests/test_weights
|
||||
|
||||
import std/[os, math]
|
||||
import arraymancer
|
||||
import "../network"
|
||||
import "../training"
|
||||
import "../weights"
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
const tmpBase = "/tmp/test_weights_nim"
|
||||
|
||||
# ── saveWeights / loadWeights roundtrip ───────────────────────────────────────
|
||||
|
||||
block testRoundtrip:
|
||||
let dir = tmpBase & "_roundtrip"
|
||||
removeDir(dir)
|
||||
|
||||
let ac1 = initActorCritic()
|
||||
saveWeights(ac1, dir)
|
||||
|
||||
var ac2 = initActorCritic()
|
||||
loadWeights(ac2, dir)
|
||||
|
||||
# Verify a sample of tensors
|
||||
template tensorEq(a, b: Tensor[float32]) =
|
||||
check a.shape == b.shape, "shape mismatch"
|
||||
let diff = abs(a - b)
|
||||
var maxDiff = 0.0'f32
|
||||
for v in diff: maxDiff = max(maxDiff, v)
|
||||
check maxDiff < 1e-6'f32, "tensor values differ by " & $maxDiff
|
||||
|
||||
tensorEq(ac1.actor.w1, ac2.actor.w1)
|
||||
tensorEq(ac1.actor.b1, ac2.actor.b1)
|
||||
tensorEq(ac1.actor.w3, ac2.actor.w3)
|
||||
tensorEq(ac1.critic.w1, ac2.critic.w1)
|
||||
tensorEq(ac1.critic.b3, ac2.critic.b3)
|
||||
tensorEq(ac1.logStd, ac2.logStd)
|
||||
|
||||
removeDir(dir)
|
||||
|
||||
# ── saveWeightsAtomic ─────────────────────────────────────────────────────────
|
||||
|
||||
block testAtomic:
|
||||
let dir = tmpBase & "_atomic"
|
||||
removeDir(dir)
|
||||
|
||||
let ac = initActorCritic()
|
||||
saveWeightsAtomic(ac, dir)
|
||||
|
||||
check dirExists(dir), "targetDir should exist after atomic save"
|
||||
for f in ["actor_w1.npy", "critic_w1.npy", "log_std.npy"]:
|
||||
check fileExists(dir / f), "missing file: " & f
|
||||
|
||||
removeDir(dir)
|
||||
|
||||
# ── loadBestAvailable ─────────────────────────────────────────────────────────
|
||||
|
||||
block testLoadBest:
|
||||
let root = tmpBase & "_loadbest"
|
||||
removeDir(root)
|
||||
createDir(root)
|
||||
|
||||
let ac0 = initActorCritic()
|
||||
|
||||
# Try with no weights — should return false
|
||||
var acEmpty = initActorCritic()
|
||||
var adamEmpty: ACAdamStates
|
||||
check not loadBestAvailable(acEmpty, adamEmpty, root).loaded, "should return false with no weights"
|
||||
|
||||
# Save to latest/; should load
|
||||
saveWeights(ac0, root / "latest")
|
||||
var ac1 = initActorCritic()
|
||||
var adam1: ACAdamStates
|
||||
check loadBestAvailable(ac1, adam1, root).loaded, "should load from latest/"
|
||||
|
||||
# Remove latest/, save to checkpoint_1/ — should fall back
|
||||
removeDir(root / "latest")
|
||||
saveWeights(ac0, root / "checkpoint_1")
|
||||
var ac2 = initActorCritic()
|
||||
var adam2: ACAdamStates
|
||||
check loadBestAvailable(ac2, adam2, root).loaded, "should load from checkpoint_1/"
|
||||
|
||||
removeDir(root)
|
||||
|
||||
# ── cleanStaleTempDirs ────────────────────────────────────────────────────────
|
||||
|
||||
block testClean:
|
||||
let root = tmpBase & "_clean"
|
||||
removeDir(root)
|
||||
createDir(root)
|
||||
|
||||
let stale = root / "latest_tmp_12345"
|
||||
createDir(stale)
|
||||
check dirExists(stale), "stale dir should exist before clean"
|
||||
|
||||
cleanStaleTempDirs(root)
|
||||
check not dirExists(stale), "stale dir should be gone after clean"
|
||||
|
||||
removeDir(root)
|
||||
|
||||
echo "All tests passed"
|
||||
Reference in New Issue
Block a user