diff --git a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/state.nim b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/state.nim new file mode 100644 index 0000000..6015b78 --- /dev/null +++ b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/state.nim @@ -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)) diff --git a/SAC_LSTM_Bot/tests/test_state.nim b/SAC_LSTM_Bot/tests/test_state.nim new file mode 100644 index 0000000..0eb256c --- /dev/null +++ b/SAC_LSTM_Bot/tests/test_state.nim @@ -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"