Files
SirRoboGarage/PPO_Bot/tests/test_network.nim
T
SirStone cdde60d79f 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>
2026-08-17 19:20:40 +02:00

59 lines
2.4 KiB
Nim
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
## 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) # 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)
assert not noFire.shouldFire, "shouldFire should be false when gunHeat > 0"
echo "All tests passed"