feat(PPO_Bot): enemy tracker + 42-float state vector (#15)
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -2,7 +2,6 @@
|
||||
## Forward pass uses a random ActorCritic policy (weights not yet trained).
|
||||
|
||||
import std/os
|
||||
import arraymancer
|
||||
import tankroyale_botapi
|
||||
import network
|
||||
import actions
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
## Enemy tracker — deterministic radar lock + dead reckoning for PPO_Bot.
|
||||
## No bot API imports; takes plain floats.
|
||||
|
||||
import std/math
|
||||
|
||||
type
|
||||
EnemyState* = object
|
||||
x*, y*: float64
|
||||
direction*: float64
|
||||
speed*: float64
|
||||
energy*: float64
|
||||
ticksSinceLastScan*: int
|
||||
hasFired*: bool
|
||||
lastFirePower*: float64
|
||||
|
||||
EnemyTracker* = object
|
||||
current*: EnemyState
|
||||
history*: array[5, tuple[x, y, direction, speed: float64]] # sliding window
|
||||
historyCount*: int # valid entries 0-5
|
||||
prevEnergy*: float64
|
||||
hasContact*: bool
|
||||
## For radar overshoot reversal
|
||||
lastOvershootDir*: float64 # +1 or -1
|
||||
|
||||
proc initEnemyTracker*(): EnemyTracker =
|
||||
result.lastOvershootDir = 1.0
|
||||
|
||||
proc update*(tracker: var EnemyTracker;
|
||||
scanX, scanY, scanDir, scanSpeed, scanEnergy: float64) =
|
||||
## Call on ScannedBotEvent. Detects enemy fire from energy delta.
|
||||
|
||||
# Shift history window
|
||||
if tracker.historyCount > 0:
|
||||
for i in countdown(min(tracker.historyCount, 4), 1):
|
||||
tracker.history[i] = tracker.history[i - 1]
|
||||
tracker.history[0] = (tracker.current.x, tracker.current.y,
|
||||
tracker.current.direction, tracker.current.speed)
|
||||
if tracker.historyCount < 5:
|
||||
inc tracker.historyCount
|
||||
|
||||
# Detect firing: energy drop in [0.1, 3.0] means enemy fired
|
||||
let delta = tracker.prevEnergy - scanEnergy
|
||||
if tracker.hasContact and delta >= 0.1 and delta <= 3.0:
|
||||
tracker.current.hasFired = true
|
||||
tracker.current.lastFirePower = delta
|
||||
else:
|
||||
tracker.current.hasFired = false
|
||||
|
||||
tracker.prevEnergy = scanEnergy
|
||||
tracker.current.x = scanX
|
||||
tracker.current.y = scanY
|
||||
tracker.current.direction = scanDir
|
||||
tracker.current.speed = scanSpeed
|
||||
tracker.current.energy = scanEnergy
|
||||
tracker.current.ticksSinceLastScan = 0
|
||||
tracker.hasContact = true
|
||||
|
||||
proc deadReckon*(tracker: var EnemyTracker) =
|
||||
## Call on missed ticks. Predict position from last known velocity.
|
||||
if not tracker.hasContact:
|
||||
return
|
||||
let rad = tracker.current.direction * PI / 180.0
|
||||
tracker.current.x += tracker.current.speed * sin(rad)
|
||||
tracker.current.y += tracker.current.speed * cos(rad)
|
||||
inc tracker.current.ticksSinceLastScan
|
||||
|
||||
proc normalizeRelative(angle: float64): float64 {.inline.} =
|
||||
result = angle mod 360.0
|
||||
if result >= 180.0: result -= 360.0
|
||||
elif result < -180.0: result += 360.0
|
||||
|
||||
proc getRadarTurnRate*(tracker: EnemyTracker;
|
||||
botX, botY, botDirection, radarDirection: float64): float64 =
|
||||
## Returns radar turn rate (degrees/tick, positive = right).
|
||||
## Before contact: full 45° sweep.
|
||||
## After contact: lock with overshoot; widen if stale.
|
||||
if not tracker.hasContact:
|
||||
return 45.0
|
||||
|
||||
if tracker.current.ticksSinceLastScan >= 2:
|
||||
# Lost lock — widen sweep proportional to staleness
|
||||
return 45.0
|
||||
|
||||
# Bearing from radar to enemy
|
||||
let dx = tracker.current.x - botX
|
||||
let dy = tracker.current.y - botY
|
||||
let absoluteDir = (180.0 * arctan2(dy, dx) / PI + 360.0) mod 360.0
|
||||
let radarBearing = normalizeRelative(absoluteDir - radarDirection)
|
||||
|
||||
# Overshoot by 10°, alternate direction each tick
|
||||
# ponytail: simple fixed overshoot; adaptive sweep if enemy is fast-turning
|
||||
let overshoot = 10.0
|
||||
let target = radarBearing + tracker.lastOvershootDir * overshoot
|
||||
result = target.clamp(-45.0, 45.0)
|
||||
@@ -0,0 +1,87 @@
|
||||
## State vector builder — produces 42-float normalized tensor for PPO policy.
|
||||
## No bot API imports; takes plain BotState + EnemyTracker structs.
|
||||
|
||||
import std/math
|
||||
import arraymancer
|
||||
import ./enemy_tracker
|
||||
|
||||
type
|
||||
BotStateData* = object
|
||||
## Plain data mirror of the bot's observable state.
|
||||
x*, y*: float64
|
||||
direction*: float64
|
||||
speed*: float64
|
||||
energy*: float64
|
||||
gunDirection*: float64
|
||||
gunHeat*: float64
|
||||
arenaWidth*, arenaHeight*: float64
|
||||
|
||||
proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker): Tensor[float32] =
|
||||
## Build the 42-float normalized state tensor.
|
||||
## All values clipped to roughly [-1, 1] via division by physical maxima.
|
||||
result = zeros[float32](42)
|
||||
|
||||
let aW = bot.arenaWidth
|
||||
let aH = bot.arenaHeight
|
||||
let diag = sqrt(aW * aW + aH * aH) # ≈ 1700 for 1200×800
|
||||
let wallMax = max(aW, aH)
|
||||
|
||||
# --- Current tick: own bot (indices 0-6) ---
|
||||
result[0] = float32(bot.x / aW)
|
||||
result[1] = float32(bot.y / aH)
|
||||
result[2] = float32(bot.direction / 360.0)
|
||||
result[3] = float32(bot.speed / 8.0)
|
||||
result[4] = float32(bot.energy / 100.0)
|
||||
result[5] = float32(bot.gunDirection / 360.0)
|
||||
result[6] = float32(bot.gunHeat / 1.8)
|
||||
|
||||
# --- Current tick: enemy (indices 7-13) ---
|
||||
if enemy.hasContact:
|
||||
result[7] = float32(enemy.current.x / aW)
|
||||
result[8] = float32(enemy.current.y / aH)
|
||||
result[9] = float32(enemy.current.direction / 360.0)
|
||||
result[10] = float32(enemy.current.speed / 8.0)
|
||||
result[11] = float32(enemy.current.energy / 100.0)
|
||||
result[12] = float32(if enemy.current.hasFired: 1.0 else: 0.0)
|
||||
result[13] = float32(enemy.current.lastFirePower / 3.0)
|
||||
# else: remain 0.0
|
||||
|
||||
# --- Derived features (indices 14-21) ---
|
||||
if enemy.hasContact:
|
||||
# Enemy acceleration (speed delta from last history entry)
|
||||
# Max delta is ±8 (stopped ↔ full speed); divide by 8 to normalize
|
||||
if enemy.historyCount >= 1:
|
||||
result[14] = float32((enemy.current.speed - enemy.history[0].speed) / 8.0)
|
||||
# Enemy turn rate (direction delta from last history entry, normalized to [-1,1])
|
||||
# Divide by 180 (max possible relative rotation) rather than 10 (max body turn rate)
|
||||
# ponytail: /180 covers all cases; use /10 if you want sensitivity to small turns
|
||||
if enemy.historyCount >= 1:
|
||||
let dirDelta = ((enemy.current.direction - enemy.history[0].direction) + 540.0) mod 360.0 - 180.0
|
||||
result[15] = float32(dirDelta / 180.0)
|
||||
# Relative bearing to enemy (signed, from bot perspective)
|
||||
let dx = enemy.current.x - bot.x
|
||||
let dy = enemy.current.y - bot.y
|
||||
let absDir = (180.0 * arctan2(dy, dx) / PI + 360.0) mod 360.0
|
||||
let relBearing = ((absDir - bot.direction) + 540.0) mod 360.0 - 180.0
|
||||
result[16] = float32(relBearing / 180.0)
|
||||
# Distance to enemy
|
||||
let dist = sqrt(dx * dx + dy * dy)
|
||||
result[17] = float32(dist / diag)
|
||||
|
||||
# Wall distances (indices 18-21): top, bottom, left, right
|
||||
# top = distance from bot to top wall (y=aH), bottom = distance to bottom (y=0)
|
||||
# left = distance to left (x=0), right = distance to right (x=aW)
|
||||
result[18] = float32((aH - bot.y) / wallMax) # top
|
||||
result[19] = float32(bot.y / wallMax) # bottom
|
||||
result[20] = float32(bot.x / wallMax) # left
|
||||
result[21] = float32((aW - bot.x) / wallMax) # right
|
||||
|
||||
# --- History: 5 ticks × 4 floats = 20 floats (indices 22-41) ---
|
||||
for i in 0 ..< 5:
|
||||
let base = 22 + i * 4
|
||||
if i < enemy.historyCount:
|
||||
result[base + 0] = float32(enemy.history[i].x / aW)
|
||||
result[base + 1] = float32(enemy.history[i].y / aH)
|
||||
result[base + 2] = float32(enemy.history[i].direction / 360.0)
|
||||
result[base + 3] = float32(enemy.history[i].speed / 8.0)
|
||||
# else: remain 0.0 (pad)
|
||||
@@ -0,0 +1,146 @@
|
||||
## Assert-based tests for EnemyTracker and StateVector.
|
||||
## Run: nim c -r tests/test_state.nim
|
||||
|
||||
import std/[math, strformat]
|
||||
import arraymancer
|
||||
|
||||
# Import from parent dir
|
||||
import "../enemy_tracker"
|
||||
import "../state_vector"
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# EnemyTracker tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
block testBasicUpdate:
|
||||
var t = initEnemyTracker()
|
||||
t.update(200.0, 300.0, 90.0, 5.0, 80.0)
|
||||
check t.hasContact, "hasContact after update"
|
||||
check t.current.x == 200.0, "x after update"
|
||||
check t.current.y == 300.0, "y after update"
|
||||
check t.current.direction == 90.0, "direction after update"
|
||||
check t.current.speed == 5.0, "speed after update"
|
||||
check t.current.energy == 80.0, "energy after update"
|
||||
check t.current.ticksSinceLastScan == 0, "ticksSinceLastScan reset"
|
||||
|
||||
block testFireDetection:
|
||||
var t = initEnemyTracker()
|
||||
# First update sets prevEnergy
|
||||
t.update(100.0, 100.0, 0.0, 0.0, 100.0)
|
||||
# Second update: energy drop of 3.0 → enemy fired power 3.0
|
||||
t.update(100.0, 100.0, 0.0, 0.0, 97.0)
|
||||
check t.current.hasFired, "hasFired when energy drops by 3.0"
|
||||
check abs(t.current.lastFirePower - 3.0) < 0.001, "lastFirePower == 3.0"
|
||||
|
||||
block testNoFireOnSmallDrop:
|
||||
var t = initEnemyTracker()
|
||||
t.update(100.0, 100.0, 0.0, 0.0, 100.0)
|
||||
# Drop of 0.05 — below MIN_FIRE_POWER threshold
|
||||
t.update(100.0, 100.0, 0.0, 0.0, 99.95)
|
||||
check not t.current.hasFired, "no fire on small energy drop"
|
||||
|
||||
block testDeadReckoning:
|
||||
var t = initEnemyTracker()
|
||||
# direction=0° in Tank Royale means north (y increases)
|
||||
t.update(100.0, 100.0, 0.0, 5.0, 100.0)
|
||||
t.deadReckon()
|
||||
# x unchanged (sin 0° = 0), y increases by speed (cos 0° = 1)
|
||||
check abs(t.current.x - 100.0) < 0.001, "dead reckon: x unchanged for dir=0"
|
||||
check abs(t.current.y - 105.0) < 0.001, "dead reckon: y += speed for dir=0"
|
||||
check t.current.ticksSinceLastScan == 1, "ticksSinceLastScan incremented"
|
||||
|
||||
block testDeadReckonEast:
|
||||
var t = initEnemyTracker()
|
||||
# direction=90° → east (sin 90° = 1, cos 90° = 0)
|
||||
t.update(100.0, 100.0, 90.0, 5.0, 100.0)
|
||||
t.deadReckon()
|
||||
check abs(t.current.x - 105.0) < 0.001, "dead reckon east: x += speed"
|
||||
check abs(t.current.y - 100.0) < 0.001, "dead reckon east: y unchanged"
|
||||
|
||||
block testHistoryWindow:
|
||||
var t = initEnemyTracker()
|
||||
# Feed 6 updates — history should hold last 5
|
||||
for i in 1 .. 6:
|
||||
t.update(float64(i) * 10.0, float64(i) * 20.0, 0.0, float64(i), 100.0)
|
||||
check t.historyCount == 5, "historyCount capped at 5"
|
||||
# history[0] should be the second-to-last scan (i=5)
|
||||
check abs(t.history[0].x - 50.0) < 0.001, "history[0].x == 50 (i=5)"
|
||||
check abs(t.history[4].x - 10.0) < 0.001, "history[4].x == 10 (i=1)"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# StateVector tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
block testStateVectorLength:
|
||||
var t = initEnemyTracker()
|
||||
t.update(400.0, 300.0, 45.0, 3.0, 80.0)
|
||||
let bot = BotStateData(
|
||||
x: 200.0, y: 200.0, direction: 90.0, speed: 4.0, energy: 50.0,
|
||||
gunDirection: 180.0, gunHeat: 0.5,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
check sv.shape == [42], "state vector has 42 elements"
|
||||
|
||||
block testStateVectorRange:
|
||||
var t = initEnemyTracker()
|
||||
t.update(400.0, 300.0, 180.0, 8.0, 100.0)
|
||||
let bot = BotStateData(
|
||||
x: 800.0, y: 600.0, direction: 360.0, speed: 8.0, energy: 100.0,
|
||||
gunDirection: 360.0, gunHeat: 1.8,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
for i in 0 ..< 42:
|
||||
check sv[i] >= -2.0f32 and sv[i] <= 2.0f32,
|
||||
&"sv[{i}]={sv[i]} out of [-2,2] range"
|
||||
|
||||
block testWallDistances:
|
||||
# Bot at (100, 200) in 800×600 arena
|
||||
# wallMax = max(800, 600) = 800
|
||||
# top = (600 - 200) / 800 = 400/800 = 0.5
|
||||
# bottom = 200 / 800 = 0.25
|
||||
# left = 100 / 800 = 0.125
|
||||
# right = (800 - 100) / 800 = 700/800 = 0.875
|
||||
var t = initEnemyTracker()
|
||||
let bot = BotStateData(
|
||||
x: 100.0, y: 200.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
||||
gunDirection: 0.0, gunHeat: 0.0,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
check abs(sv[18] - 0.5f32) < 0.001f32, "top wall = 0.5"
|
||||
check abs(sv[19] - 0.25f32) < 0.001f32, "bottom wall = 0.25"
|
||||
check abs(sv[20] - 0.125f32) < 0.001f32, "left wall = 0.125"
|
||||
check abs(sv[21] - 0.875f32) < 0.001f32, "right wall = 0.875"
|
||||
|
||||
block testRelativeBearing:
|
||||
# Bot at (0,0) dir=0°. Enemy at (0,100) → arctan2(100,0)=90° absolute.
|
||||
# relBearing = (90 - 0 + 540) mod 360 - 180 = 90°
|
||||
# normalized = 90 / 180 = 0.5
|
||||
var t = initEnemyTracker()
|
||||
t.update(0.0, 100.0, 0.0, 0.0, 100.0)
|
||||
let bot = BotStateData(
|
||||
x: 0.0, y: 0.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
||||
gunDirection: 0.0, gunHeat: 0.0,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
check abs(sv[16] - 0.5f32) < 0.01f32, "relative bearing = 0.5, got " & $sv[16]
|
||||
|
||||
block testHistoryPaddedWhenEmpty:
|
||||
var t = initEnemyTracker()
|
||||
let bot = BotStateData(
|
||||
x: 400.0, y: 300.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
||||
gunDirection: 0.0, gunHeat: 0.0,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
for i in 22 ..< 42:
|
||||
check sv[i] == 0.0f32, &"history slot {i} should be 0 when no contact"
|
||||
|
||||
echo "All tests passed"
|
||||
Reference in New Issue
Block a user