aa4bc77068
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>
62 lines
2.3 KiB
Nim
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"
|