## 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 PPO_Bot/network import PPO_Bot/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..= -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, 400.0, 300.0) assert not noFire.shouldFire, "shouldFire should be false when gunHeat > 0" echo "All tests passed"