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,117 @@
|
||||
## State vector module — produces a 35-dimensional normalized tensor for SAC+LSTM policy.
|
||||
## No bot API imports; takes plain data structs populated from game events.
|
||||
## The LSTM handles temporal context, so no explicit history window here.
|
||||
|
||||
import std/math
|
||||
import arraymancer
|
||||
|
||||
const STATE_DIM* = 35
|
||||
|
||||
type
|
||||
BulletData* = object
|
||||
## Enemy bullet in flight (absolute arena coords + fire power).
|
||||
x*, y*: float64
|
||||
power*: float64 # fire power in [0.1, 3.0]; speed = 20 - 3*power
|
||||
|
||||
EnemyData* = object
|
||||
## Current enemy state, from the most recent onScannedBot event.
|
||||
x*, y*: float64
|
||||
direction*: float64
|
||||
speed*: float64
|
||||
energy*: float64
|
||||
hasFired*: bool
|
||||
lastFirePower*: float64
|
||||
prevSpeed*: float64 # speed from the previous scan (for acceleration)
|
||||
prevDirection*: float64 # direction from the previous scan (for turn rate)
|
||||
hasPrevScan*: bool # true once we have at least two scans
|
||||
|
||||
GameState* = object
|
||||
## Accumulates data from bot events. Populate fields before calling buildState.
|
||||
# Own bot
|
||||
x*, y*: float64
|
||||
direction*: float64
|
||||
speed*: float64
|
||||
energy*: float64
|
||||
gunDirection*: float64
|
||||
gunHeat*: float64
|
||||
arenaWidth*, arenaHeight*: float64
|
||||
# Enemy
|
||||
hasContact*: bool
|
||||
enemy*: EnemyData
|
||||
ticksSinceLastScan*: int
|
||||
# Bullets in flight (up to 3 tracked)
|
||||
bullets*: array[3, BulletData]
|
||||
bulletCount*: int
|
||||
|
||||
proc buildState*(gs: GameState): Tensor[float32] =
|
||||
## Build the 35-float normalized state tensor.
|
||||
##
|
||||
## Layout:
|
||||
## [0-6] own bot: x/aW, y/aH, dir/360, speed/8, energy/100, gunDir/360, gunHeat/1.8
|
||||
## [7-13] enemy: x/aW, y/aH, dir/360, speed/8, energy/100, hasFired, lastFirePower/3
|
||||
## [14-17] derived: enemyAccel/8, enemyTurnRate/180, relBearing/180, distance/diag
|
||||
## [18-21] walls: top, bottom, left, right — each / max(aW,aH)
|
||||
## [22-33] bullets: up to 3 × (relX/aW, relY/aH, speed/20, ticksToImpact clamped to 1)
|
||||
## [34] scan staleness: ticksSinceLastScan/30 clamped to 1
|
||||
result = zeros[float32](STATE_DIM)
|
||||
|
||||
let aW = gs.arenaWidth
|
||||
let aH = gs.arenaHeight
|
||||
let diag = sqrt(aW * aW + aH * aH)
|
||||
let wMax = max(aW, aH)
|
||||
|
||||
# --- Own bot (0-6) ---
|
||||
result[0] = float32(gs.x / aW)
|
||||
result[1] = float32(gs.y / aH)
|
||||
result[2] = float32(gs.direction / 360.0)
|
||||
result[3] = float32(gs.speed / 8.0)
|
||||
result[4] = float32(gs.energy / 100.0)
|
||||
result[5] = float32(gs.gunDirection / 360.0)
|
||||
result[6] = float32(gs.gunHeat / 1.8)
|
||||
|
||||
# --- Enemy current (7-13) ---
|
||||
if gs.hasContact:
|
||||
result[7] = float32(gs.enemy.x / aW)
|
||||
result[8] = float32(gs.enemy.y / aH)
|
||||
result[9] = float32(gs.enemy.direction / 360.0)
|
||||
result[10] = float32(gs.enemy.speed / 8.0)
|
||||
result[11] = float32(gs.enemy.energy / 100.0)
|
||||
result[12] = float32(if gs.enemy.hasFired: 1.0 else: 0.0)
|
||||
result[13] = float32(gs.enemy.lastFirePower / 3.0)
|
||||
|
||||
# --- Derived (14-17) ---
|
||||
if gs.hasContact:
|
||||
if gs.enemy.hasPrevScan:
|
||||
result[14] = float32((gs.enemy.speed - gs.enemy.prevSpeed) / 8.0)
|
||||
let dDir = ((gs.enemy.direction - gs.enemy.prevDirection) + 540.0) mod 360.0 - 180.0
|
||||
result[15] = float32(dDir / 180.0)
|
||||
let dx = gs.enemy.x - gs.x
|
||||
let dy = gs.enemy.y - gs.y
|
||||
let absDir = (180.0 * arctan2(dx, dy) / PI + 360.0) mod 360.0
|
||||
let relBearing = ((absDir - gs.direction) + 540.0) mod 360.0 - 180.0
|
||||
result[16] = float32(relBearing / 180.0)
|
||||
result[17] = float32(sqrt(dx * dx + dy * dy) / diag)
|
||||
|
||||
# --- Wall distances (18-21): top, bottom, left, right ---
|
||||
result[18] = float32((aH - gs.y) / wMax)
|
||||
result[19] = float32(gs.y / wMax)
|
||||
result[20] = float32(gs.x / wMax)
|
||||
result[21] = float32((aW - gs.x) / wMax)
|
||||
|
||||
# --- Bullet tracking (22-33): up to 3 bullets × 4 floats ---
|
||||
# Per slot: relX/aW, relY/aH, speed/20, ticksToImpact/diag (clamped to 1)
|
||||
for i in 0 ..< min(gs.bulletCount, 3):
|
||||
let b = gs.bullets[i]
|
||||
let bSpd = 20.0 - 3.0 * b.power
|
||||
let bdx = b.x - gs.x
|
||||
let bdy = b.y - gs.y
|
||||
let bdist = sqrt(bdx * bdx + bdy * bdy)
|
||||
let ticks = if bSpd > 0.0: min(bdist / bSpd / diag, 1.0) else: 0.0
|
||||
let base = 22 + i * 4
|
||||
result[base + 0] = float32(bdx / aW)
|
||||
result[base + 1] = float32(bdy / aH)
|
||||
result[base + 2] = float32(bSpd / 20.0)
|
||||
result[base + 3] = float32(ticks)
|
||||
|
||||
# --- Scan staleness (34) ---
|
||||
result[34] = float32(min(gs.ticksSinceLastScan.float64 / 30.0, 1.0))
|
||||
@@ -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