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:
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)
|
||||
# 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"
|
||||
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"
|
||||
@@ -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"
|
||||
|
||||
@@ -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.
@@ -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"
|
||||
|
||||
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user