f130bf1254
35-dim normalized tensor (GameState → buildState). No history window — LSTM handles temporal context. Covers own-bot (7), enemy (7), derived (4), walls (4), bullets (12), scan staleness (1). All tests pass.
88 lines
2.6 KiB
Nim
88 lines
2.6 KiB
Nim
## Tests for state.nim — assert-based, no framework.
|
|
|
|
import std/math
|
|
import arraymancer
|
|
import SAC_LSTM_Bot/state
|
|
|
|
proc makeBase(): GameState =
|
|
result.arenaWidth = 1200.0
|
|
result.arenaHeight = 800.0
|
|
result.x = 600.0; result.y = 400.0
|
|
result.direction = 90.0; result.speed = 4.0
|
|
result.energy = 50.0
|
|
result.gunDirection = 90.0; result.gunHeat = 0.5
|
|
|
|
proc allInRange(t: Tensor[float32]): bool =
|
|
for v in t:
|
|
if v < -1.01f32 or v > 1.01f32: return false
|
|
true
|
|
|
|
proc hasNaN(t: Tensor[float32]): bool =
|
|
for v in t:
|
|
if v.float64.isNaN: return true
|
|
false
|
|
|
|
# 1. Correct shape
|
|
block:
|
|
let gs = makeBase()
|
|
let t = buildState(gs)
|
|
assert t.shape[0] == STATE_DIM, "wrong dim: " & $t.shape[0]
|
|
echo "PASS shape"
|
|
|
|
# 2. All values in [-1, 1] for typical input
|
|
block:
|
|
var gs = makeBase()
|
|
gs.hasContact = true
|
|
gs.enemy = EnemyData(x: 800.0, y: 300.0, direction: 45.0, speed: 6.0,
|
|
energy: 80.0, hasFired: true, lastFirePower: 2.0,
|
|
prevSpeed: 5.0, prevDirection: 40.0, hasPrevScan: true)
|
|
gs.bulletCount = 1
|
|
gs.bullets[0] = BulletData(x: 610.0, y: 410.0, power: 1.0)
|
|
gs.ticksSinceLastScan = 10
|
|
let t = buildState(gs)
|
|
assert not hasNaN(t), "NaN in tensor"
|
|
assert allInRange(t), "value out of [-1,1]"
|
|
echo "PASS range"
|
|
|
|
# 3. Missing scan → valid tensor, no NaN, zeros for enemy fields
|
|
block:
|
|
let gs = makeBase() # hasContact = false
|
|
let t = buildState(gs)
|
|
assert not hasNaN(t), "NaN with no scan"
|
|
for i in 7 .. 17:
|
|
assert t[i] == 0.0f32, "enemy slot " & $i & " should be 0"
|
|
echo "PASS no-scan zeros"
|
|
|
|
# 4. Bullet tracking: 0, 1, 2, 3 bullets
|
|
block:
|
|
for n in 0 .. 3:
|
|
var gs = makeBase()
|
|
gs.bulletCount = n
|
|
for i in 0 ..< n:
|
|
gs.bullets[i] = BulletData(x: float64(600 + i * 10), y: float64(400 + i * 5), power: 1.0)
|
|
let t = buildState(gs)
|
|
assert not hasNaN(t), "NaN with " & $n & " bullets"
|
|
# slots beyond bulletCount must be 0
|
|
for i in n ..< 3:
|
|
let base = 22 + i * 4
|
|
for j in 0 ..< 4:
|
|
assert t[base + j] == 0.0f32, "bullet slot " & $i & " slot " & $j & " should be 0"
|
|
echo "PASS bullet tracking 0-3"
|
|
|
|
# 5. Scan staleness increments and clamps
|
|
block:
|
|
var gs = makeBase()
|
|
gs.hasContact = true
|
|
gs.ticksSinceLastScan = 0
|
|
assert buildState(gs)[34] == 0.0f32, "staleness at 0"
|
|
gs.ticksSinceLastScan = 15
|
|
let mid = buildState(gs)[34]
|
|
assert mid > 0.0f32 and mid < 1.0f32, "staleness mid not in (0,1)"
|
|
gs.ticksSinceLastScan = 30
|
|
assert buildState(gs)[34] == 1.0f32, "staleness at 30"
|
|
gs.ticksSinceLastScan = 60
|
|
assert buildState(gs)[34] == 1.0f32, "staleness clamp at 60"
|
|
echo "PASS staleness"
|
|
|
|
echo "ALL TESTS PASSED"
|