Files
2026-08-27 18:42:03 +02:00

70 lines
3.1 KiB
Nim

## Assert-based tests for actions.nim.
## Run: nim c -r tests/test_actions.nim
import arraymancer
import std/math
import PPO_Bot/actions
template check(cond: bool, msg: string) =
if not cond:
quit("FAIL: " & msg, 1)
let arenaW = 1200.0
let arenaH = 800.0
# Build a 6-element zero tensor and a helper to set individual values
proc makeRaw(vals: array[6, float32]): Tensor[float32] =
result = zeros[float32](6)
for i in 0 ..< 6: result[i] = vals[i]
# --- 6-dim input produces a valid BotActions ---
block basicDecode:
let raw = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 0.0, 0.0])
let acts = mapActions(raw, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
# sigmoid(0)*arenaW = 0.5*1200 = 600, sigmoid(0)*arenaH = 0.5*800 = 400
check abs(acts.gotoX - 600.0) < 1e-6, "gotoX = sigmoid(0)*arenaW"
check abs(acts.gotoY - 400.0) < 1e-6, "gotoY = sigmoid(0)*arenaH"
check abs(acts.aimToX - 600.0) < 1e-6, "aimToX = sigmoid(0)*arenaW"
check abs(acts.aimToY - 400.0) < 1e-6, "aimToY = sigmoid(0)*arenaH"
# --- Coordinates bounded to arena size ---
block coordBounds:
# Large positive raw → sigmoid ≈ 1 → close to arenaW/arenaH
let rawHigh = makeRaw([100.0'f32, 100.0, 100.0, 100.0, 0.0, 0.0])
let actsHigh = mapActions(rawHigh, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
check actsHigh.gotoX <= arenaW + 1e-9, "gotoX <= arenaWidth"
check actsHigh.gotoY <= arenaH + 1e-9, "gotoY <= arenaHeight"
check actsHigh.gotoX >= 0.0, "gotoX >= 0"
# Large negative raw → sigmoid ≈ 0 → close to 0
let rawLow = makeRaw([-100.0'f32, -100.0, -100.0, -100.0, 0.0, 0.0])
let actsLow = mapActions(rawLow, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
check actsLow.gotoX >= -1e-9, "gotoX >= 0 (low raw)"
check actsLow.gotoY >= -1e-9, "gotoY >= 0 (low raw)"
# --- Fire triggers correctly ---
block fireTrigger:
# tanh(positive) >= 0 → fire when gunHeat = 0
let rawFire = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 1.0, 0.0])
let actsFire = mapActions(rawFire, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
check actsFire.shouldFire, "positive tanh → should fire when gun cool"
# tanh(negative) < 0 → no fire
let rawNoFire = makeRaw([0.0'f32, 0.0, 0.0, 0.0, -1.0, 0.0])
let actsNoFire = mapActions(rawNoFire, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
check not actsNoFire.shouldFire, "negative tanh → no fire"
# gunHeat > 0 → no fire even with positive decision
let actsHot = mapActions(rawFire, 0.5, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
check not actsHot.shouldFire, "positive tanh but gun hot → no fire"
# --- Fire power in [0.1, 3.0] ---
block firePowerRange:
let rawMin = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 0.0, -100.0])
let rawMax = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 0.0, 100.0])
let actsMin = mapActions(rawMin, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
let actsMax = mapActions(rawMax, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
check actsMin.firePower >= 0.1 - 1e-6, "firePower >= 0.1"
check actsMax.firePower <= 3.0 + 1e-6, "firePower <= 3.0"
echo "test_actions: all passed"