feat(SAC_LSTM_Bot): state vector module (#42)
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.
This commit is contained in:
@@ -0,0 +1,87 @@
|
||||
## 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"
|
||||
Reference in New Issue
Block a user