Files
SirRoboGarage/SAC_LSTM_Bot_garage/src/SAC_LSTM_Bot/state.nim
T

118 lines
4.3 KiB
Nim
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
## 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))