feat(PPO_Bot): command abstraction layer — goto/aimTo controllers (#24)

- Add gotoTick/aimToTick controller functions (#25)
- Update network dims: actor 5→6, state 42→44 (#26)
- Rewrite mapActions for 6-dim command space (#27)
- Delete stale weight files (shape mismatch)
- Fix existing tests for new signatures

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-08-17 19:08:31 +02:00
parent 56e0b306c9
commit cdde60d79f
39 changed files with 490 additions and 78 deletions
BIN
View File
Binary file not shown.
+69
View File
@@ -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)
# 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)
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)
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)
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)
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)
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)
let actsMax = mapActions(rawMax, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.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"
BIN
View File
Binary file not shown.
+54
View File
@@ -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"
+11 -14
View File
@@ -11,16 +11,16 @@ func isNaNF(x: float32): bool = classify(x) == fcNan
when isMainModule:
# ---- MLP forward shape ----
let mlp = initMLP(42, 64, 5)
let inp = zeros[float32](42)
let mlp = initMLP(STATE_DIM, 64, ACTION_DIM)
let inp = zeros[float32](STATE_DIM)
let mlpOut = mlp.forward(inp)
assert mlpOut.shape[0] == 5, &"MLP output shape wrong: {mlpOut.shape}"
assert mlpOut.shape[0] == ACTION_DIM, &"MLP output shape wrong: {mlpOut.shape}"
# ---- ActorCritic actorForward ----
let ac = initActorCritic()
let state = zeros[float32](42)
let state = zeros[float32](STATE_DIM)
let (acts, logP) = ac.actorForward(state)
assert acts.shape[0] == 5, &"actorForward actions shape wrong: {acts.shape}"
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}"
@@ -31,23 +31,20 @@ when isMainModule:
# ---- logStd floor: collapsing logStd should not break actorForward ----
var ac2 = initActorCritic()
for i in 0..<5: ac2.logStd[i] = -10.0'f32
for i in 0..<ACTION_DIM: ac2.logStd[i] = -10.0'f32
let (acts2, logP2) = ac2.actorForward(state)
assert acts2.shape[0] == 5, "acts2 shape wrong after logStd=-10"
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](5)
let raw = randomNormalTensor[float32](ACTION_DIM)
let speed = 4.0'f32
let botActs = mapActions(raw, speed, 0.0'f32) # gunHeat=0 → fire allowed
# 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) # gunHeat=0 → fire allowed
assert botActs.targetSpeed >= -8.0'f32 and botActs.targetSpeed <= 8.0'f32,
&"targetSpeed out of range: {botActs.targetSpeed}"
let maxTurn = 10.0'f32 - 0.75'f32 * abs(speed) # = 7.0
assert botActs.turnRate >= -maxTurn and botActs.turnRate <= maxTurn,
&"turnRate out of range: {botActs.turnRate}"
assert botActs.gunTurnRate >= -20.0'f32 and botActs.gunTurnRate <= 20.0'f32,
&"gunTurnRate out of range: {botActs.gunTurnRate}"
@@ -55,7 +52,7 @@ when isMainModule:
&"firePower out of range: {botActs.firePower}"
# shouldFire=false when gunHeat > 0
let noFire = mapActions(raw, speed, 1.0'f32)
let noFire = mapActions(raw, 1.0'f32, 800.0, 600.0, 400.0, 300.0, 0.0, speed.float, 0.0)
assert not noFire.shouldFire, "shouldFire should be false when gunHeat > 0"
echo "All tests passed"
+4 -4
View File
@@ -10,16 +10,16 @@ template check(cond: bool, msg: string) =
# Bot at arena center; stationary enemy due north (same x, higher y).
# Tank Royale: y increases northward.
# True bearing to enemy = 0° (north = 0° in game coords).
# 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 = 0.0 # north
let trueBearing = 90.0 # north in math convention (0=east, CCW+)
# Radar starts pointing at the enemy (radarDirection = 0°, due north).
var radarDir = 0.0
# 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
Binary file not shown.
+5 -2
View File
@@ -84,7 +84,7 @@ block testStateVectorLength:
arenaWidth: 800.0, arenaHeight: 600.0,
)
let sv = buildStateVector(bot, t)
check sv.shape == [42], "state vector has 42 elements"
check sv.shape == [44], "state vector has 44 elements"
block testStateVectorRange:
var t = initEnemyTracker()
@@ -95,7 +95,7 @@ block testStateVectorRange:
arenaWidth: 800.0, arenaHeight: 600.0,
)
let sv = buildStateVector(bot, t)
for i in 0 ..< 42:
for i in 0 ..< 44:
check sv[i] >= -2.0f32 and sv[i] <= 2.0f32,
&"sv[{i}]={sv[i]} out of [-2,2] range"
@@ -155,5 +155,8 @@ block testHistoryPaddedWhenEmpty:
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"
echo "All tests passed"
+4 -4
View File
@@ -29,7 +29,7 @@ block testBuffer:
var buf = initTrajectoryBuffer()
check buf.len == 0, "empty buffer len == 0"
let t1 = Transition(state: zeros[float32](42), action: zeros[float32](5),
let t1 = Transition(state: zeros[float32](STATE_DIM), action: zeros[float32](ACTION_DIM),
logProb: -1.0'f32, reward: 0.5'f32, value: 0.3'f32)
buf.add(t1)
buf.add(t1)
@@ -80,14 +80,14 @@ block testPpoUpdate:
var buf = initTrajectoryBuffer()
for _ in 0..<10:
let s = randomNormalTensor[float32](42)
let a = randomNormalTensor[float32](5)
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, action: a, logProb: lp, reward: 0.1'f32, value: v))
var adam: ACAdamStates
ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam, epochs = 2, miniBatchSize = 5)
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]