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:
2026-08-20 23:37:56 +02:00
parent a0a3840980
commit f130bf1254
2 changed files with 204 additions and 0 deletions
+117
View File
@@ -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))
+87
View File
@@ -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"