feat(PPO_Bot): bot-relative bullets + scan staleness (STATE_DIM=57)
- Bullet state (indices 44-55): enemy-relative → bot-relative frame (bot needs threat vectors to itself for dodging, not to enemy) - New index 56: scan staleness = min(ticksSinceLastScan / 30, 1.0) (gives policy a confidence signal for enemy data freshness) - warm_start.py updated: 44→57 dim expansion, TARGET_DIM variable - Tests updated for new state layout Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
+1
-1
@@ -4,7 +4,7 @@ import arraymancer
|
||||
import std/[math, random]
|
||||
|
||||
const
|
||||
STATE_DIM* = 56
|
||||
STATE_DIM* = 57
|
||||
ACTION_DIM* = 6
|
||||
|
||||
var
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
## State vector builder — produces 56-float normalized tensor for PPO policy.
|
||||
## State vector builder — produces 57-float normalized tensor for PPO policy.
|
||||
## No bot API imports; takes plain BotState + EnemyTracker structs.
|
||||
|
||||
import std/math
|
||||
@@ -26,10 +26,11 @@ proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker;
|
||||
remainingGunAngle: float64 = 0.0;
|
||||
bullets: openArray[BulletData] = [];
|
||||
bulletCount: int = 0): Tensor[float32] =
|
||||
## Build the 56-float normalized state tensor.
|
||||
## Build the 57-float normalized state tensor.
|
||||
## Indices 0-43: existing features. Indices 44-55: up to 3 bullet slots (4 floats each).
|
||||
## Index 56: scan staleness (ticksSinceLastScan / 30, clamped to 1).
|
||||
## All values clipped to roughly [-1, 1] via division by physical maxima.
|
||||
result = zeros[float32](56)
|
||||
result = zeros[float32](57)
|
||||
|
||||
let aW = bot.arenaWidth
|
||||
let aH = bot.arenaHeight
|
||||
@@ -102,14 +103,12 @@ proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker;
|
||||
|
||||
# --- Bullet tracking (indices 44-55): up to 3 enemy bullets, 4 floats each ---
|
||||
# Per bullet: relX/aW, relY/aH, speed/20, ticksToImpact/diag
|
||||
# Positions are relative to enemy (threat vector). Slots beyond bulletCount stay 0.
|
||||
let ex = if enemy.hasContact: enemy.current.x else: bot.arenaWidth / 2.0
|
||||
let ey = if enemy.hasContact: enemy.current.y else: bot.arenaHeight / 2.0
|
||||
# Positions are relative to bot (useful for dodging). Slots beyond bulletCount stay 0.
|
||||
for i in 0 ..< min(bulletCount, 3):
|
||||
let b = bullets[i]
|
||||
let bSpeed = 20.0 - 3.0 * b.power # Tank Royale bullet speed formula
|
||||
let bdx = b.x - ex
|
||||
let bdy = b.y - ey
|
||||
let bdx = b.x - bot.x
|
||||
let bdy = b.y - bot.y
|
||||
let bdist = sqrt(bdx * bdx + bdy * bdy)
|
||||
let ticks = if bSpeed > 0.0: bdist / bSpeed else: 0.0
|
||||
let base = 44 + i * 4
|
||||
@@ -117,3 +116,6 @@ proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker;
|
||||
result[base + 1] = float32(bdy / bot.arenaHeight)
|
||||
result[base + 2] = float32(bSpeed / 20.0)
|
||||
result[base + 3] = float32(ticks / diag)
|
||||
|
||||
# --- Scan staleness (index 56) ---
|
||||
result[56] = float32(min(enemy.current.ticksSinceLastScan.float64 / 30.0, 1.0))
|
||||
|
||||
@@ -84,7 +84,7 @@ block testStateVectorLength:
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
check sv.shape == [56], "state vector has 56 elements"
|
||||
check sv.shape == [57], "state vector has 57 elements"
|
||||
|
||||
block testStateVectorRange:
|
||||
var t = initEnemyTracker()
|
||||
@@ -95,7 +95,7 @@ block testStateVectorRange:
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
for i in 0 ..< 56:
|
||||
for i in 0 ..< 57:
|
||||
check sv[i] >= -2.0f32 and sv[i] <= 2.0f32,
|
||||
&"sv[{i}]={sv[i]} out of [-2,2] range"
|
||||
|
||||
@@ -161,6 +161,8 @@ block testHistoryPaddedWhenEmpty:
|
||||
# indices 44-55 (bullet slots) default to 0 when no bullets provided
|
||||
for i in 44 ..< 56:
|
||||
check sv[i] == 0.0f32, &"bullet slot {i} should be 0 when no bullets"
|
||||
# index 56 (staleness): no contact so ticksSinceLastScan=0 → 0/30 = 0
|
||||
check sv[56] == 0.0f32, "sv[56] staleness should be 0 when no contact"
|
||||
|
||||
block testBulletSlots:
|
||||
var t = initEnemyTracker()
|
||||
@@ -171,15 +173,15 @@ block testBulletSlots:
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
# Bullet at (300,250), power 1.0 → speed = 20-3 = 17
|
||||
# relX = 300-400 = -100, relY = 250-300 = -50
|
||||
# dist to enemy = sqrt(100^2+50^2) ≈ 111.8, ticks ≈ 111.8/17 ≈ 6.58
|
||||
# relX = 300-200 = 100, relY = 250-200 = 50 (relative to bot, not enemy)
|
||||
# dist to bot = sqrt(100^2+50^2) ≈ 111.8, ticks ≈ 111.8/17 ≈ 6.58
|
||||
let b = BulletData(x: 300.0, y: 250.0, power: 1.0)
|
||||
let sv = buildStateVector(bot, t, 0.0, 0.0, [b], 1)
|
||||
check abs(sv[44] - (-100.0/800.0).float32) < 0.001f32, "bullet relX"
|
||||
check abs(sv[45] - (-50.0/600.0).float32) < 0.001f32, "bullet relY"
|
||||
check abs(sv[44] - (100.0/800.0).float32) < 0.001f32, "bullet relX"
|
||||
check abs(sv[45] - (50.0/600.0).float32) < 0.001f32, "bullet relY"
|
||||
check abs(sv[46] - (17.0/20.0).float32) < 0.001f32, "bullet speed norm"
|
||||
# second slot should be zero-padded
|
||||
for i in 48 ..< 56:
|
||||
for i in 48 ..< 57:
|
||||
check sv[i] == 0.0f32, &"unused bullet slot {i} should be 0"
|
||||
|
||||
echo "All tests passed"
|
||||
|
||||
Reference in New Issue
Block a user