259e67d7ff
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
70 lines
3.1 KiB
Nim
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"
|