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