Files
SirRoboGarage/PPO_Bot/tests/test_network.nim
T
SirStone aa4bc77068 feat(PPO_Bot): network forward pass + action mapping (#14)
Two-hidden-layer MLP actor-critic (42→64→64→5/1) with stochastic
actorForward, logStd floor at -3, and BotAction mapper wired into
the run() loop. Assert-based test suite covers shapes, finiteness,
logStd collapse, and all action range bounds.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-08-16 15:10:46 +02:00

62 lines
2.3 KiB
Nim

## 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(42, 64, 5)
let inp = zeros[float32](42)
let mlpOut = mlp.forward(inp)
assert mlpOut.shape[0] == 5, &"MLP output shape wrong: {mlpOut.shape}"
# ---- ActorCritic actorForward ----
let ac = initActorCritic()
let state = zeros[float32](42)
let (acts, logP) = ac.actorForward(state)
assert acts.shape[0] == 5, &"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..<5: ac2.logStd[i] = -10.0'f32
let (acts2, logP2) = ac2.actorForward(state)
assert acts2.shape[0] == 5, "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 speed = 4.0'f32
let botActs = mapActions(raw, speed, 0.0'f32) # 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}"
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, speed, 1.0'f32)
assert not noFire.shouldFire, "shouldFire should be false when gunHeat > 0"
echo "All tests passed"