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
+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"