fix(PPO_Bot): SIGSEGV crash fixes + static buffers for thread safety
- bullets: seq[InFlightBullet] → array[4, InFlightBullet] + bulletCount (eliminates cross-thread heap realloc under ORC) - hasFired: edge-triggered (cleared after state build, not level-triggered) - round_counter parseInt: wrapped for empty/torn file → 0 - Static SVG + intent buffers to kill cross-thread heap realloc - Tick-local alive/bulletData also fixed arrays Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
+68
-3
@@ -1,7 +1,7 @@
|
|||||||
## PPO_Bot — enemy tracker + state vector wired into the game loop.
|
## PPO_Bot — enemy tracker + state vector wired into the game loop.
|
||||||
## Training: trajectory collected per tick, PPO update in background thread.
|
## Training: trajectory collected per tick, PPO update in background thread.
|
||||||
|
|
||||||
import std/[os, strformat, strutils, math, times]
|
import std/[os, strformat, strutils, math, times, algorithm]
|
||||||
import arraymancer
|
import arraymancer
|
||||||
import tankroyale_botapi
|
import tankroyale_botapi
|
||||||
import network
|
import network
|
||||||
@@ -35,6 +35,7 @@ var
|
|||||||
|
|
||||||
# Wire logStd tunable params into network module vars (read before initActorCritic)
|
# Wire logStd tunable params into network module vars (read before initActorCritic)
|
||||||
logStdFloor = getEnvFloat("PPOB_LOG_STD_FLOOR", -3.0'f32)
|
logStdFloor = getEnvFloat("PPOB_LOG_STD_FLOOR", -3.0'f32)
|
||||||
|
logStdCeiling = getEnvFloat("PPOB_LOG_STD_CEILING", 0.5'f32)
|
||||||
initialLogStd = getEnvFloat("PPOB_INITIAL_LOG_STD", 0.0'f32)
|
initialLogStd = getEnvFloat("PPOB_INITIAL_LOG_STD", 0.0'f32)
|
||||||
|
|
||||||
# ── Structured log output ─────────────────────────────────────────────────────
|
# ── Structured log output ─────────────────────────────────────────────────────
|
||||||
@@ -59,11 +60,17 @@ proc hyperparmSnapshot(): string =
|
|||||||
&"\"entropyCoeff\":{hpEntropyCoeff},\"valueLossCoeff\":{hpValueLossCoeff}," &
|
&"\"entropyCoeff\":{hpEntropyCoeff},\"valueLossCoeff\":{hpValueLossCoeff}," &
|
||||||
&"\"maxGradNorm\":{hpMaxGradNorm},\"gamma\":{hpGamma},\"lam\":{hpLam}," &
|
&"\"maxGradNorm\":{hpMaxGradNorm},\"gamma\":{hpGamma},\"lam\":{hpLam}," &
|
||||||
&"\"epochs\":{hpEpochs},\"miniBatchSize\":{hpMiniBatchSize}," &
|
&"\"epochs\":{hpEpochs},\"miniBatchSize\":{hpMiniBatchSize}," &
|
||||||
&"\"logStdFloor\":{logStdFloor},\"initialLogStd\":{initialLogStd}"
|
&"\"logStdFloor\":{logStdFloor},\"logStdCeiling\":{logStdCeiling},\"initialLogStd\":{initialLogStd}"
|
||||||
|
|
||||||
const botJsonPath = currentSourcePath().parentDir / "PPO_Bot.json"
|
const botJsonPath = currentSourcePath().parentDir / "PPO_Bot.json"
|
||||||
const weightsRoot = currentSourcePath().parentDir / "weights"
|
const weightsRoot = currentSourcePath().parentDir / "weights"
|
||||||
|
|
||||||
|
type
|
||||||
|
InFlightBullet = object
|
||||||
|
x, y: float64 # current position
|
||||||
|
vx, vy: float64 # velocity (pixels/tick)
|
||||||
|
power: float64
|
||||||
|
|
||||||
type PPOBot = ref object of Bot
|
type PPOBot = ref object of Bot
|
||||||
tracker: EnemyTracker
|
tracker: EnemyTracker
|
||||||
buffer: TrajectoryBuffer
|
buffer: TrajectoryBuffer
|
||||||
@@ -77,6 +84,12 @@ type PPOBot = ref object of Bot
|
|||||||
lastActions: BotActions # previous tick's decoded actions (for state vector)
|
lastActions: BotActions # previous tick's decoded actions (for state vector)
|
||||||
roundRewardSum: float32 # cumulative reward this round (for live display)
|
roundRewardSum: float32 # cumulative reward this round (for live display)
|
||||||
roundTicks: int # ticks this round
|
roundTicks: int # ticks this round
|
||||||
|
# Fixed-size bullet buffer — NO heap on the shared bot object. A seq here is
|
||||||
|
# allocated by the per-round bot thread and freed by the next round's thread
|
||||||
|
# (bot.bullets = @[] on round start) → foreign-heap free under --threads:on +
|
||||||
|
# ORC → SIGSEGV. N=4: the state vector only consumes the closest 3 slots.
|
||||||
|
bullets: array[4, InFlightBullet]
|
||||||
|
bulletCount: int
|
||||||
|
|
||||||
var ac = initActorCritic()
|
var ac = initActorCritic()
|
||||||
var gAdamStates: ACAdamStates # persists across rounds
|
var gAdamStates: ACAdamStates # persists across rounds
|
||||||
@@ -98,6 +111,7 @@ method onRoundStarted*(bot: PPOBot, e: RoundStartedEvent) =
|
|||||||
bot.hasLastTrans = false
|
bot.hasLastTrans = false
|
||||||
bot.roundRewardSum = 0.0'f32
|
bot.roundRewardSum = 0.0'f32
|
||||||
bot.roundTicks = 0
|
bot.roundTicks = 0
|
||||||
|
bot.bulletCount = 0
|
||||||
|
|
||||||
method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
||||||
inc roundCounter
|
inc roundCounter
|
||||||
@@ -184,6 +198,53 @@ method run(bot: PPOBot) =
|
|||||||
while isRunning():
|
while isRunning():
|
||||||
bot.tracker.deadReckon()
|
bot.tracker.deadReckon()
|
||||||
|
|
||||||
|
# Spawn bullet when enemy fired last scan tick. Edge-triggered: the flag is
|
||||||
|
# consumed later (after the state build) so a radar gap (deadReckon ticks)
|
||||||
|
# can't spawn k phantom bullets from one shot — exactly one bullet, and
|
||||||
|
# state index 12 still pulses 1 on the detection tick.
|
||||||
|
if bot.tracker.hasContact and bot.tracker.current.hasFired:
|
||||||
|
if bot.bulletCount < 4:
|
||||||
|
let power = bot.tracker.current.lastFirePower
|
||||||
|
let speed = 20.0 - 3.0 * power
|
||||||
|
# Approximate gun direction: bearing from enemy toward our position
|
||||||
|
let myX = getX(); let myY = getY()
|
||||||
|
let ang = arctan2(myY - bot.tracker.current.y, myX - bot.tracker.current.x)
|
||||||
|
bot.bullets[bot.bulletCount] = InFlightBullet(
|
||||||
|
x: bot.tracker.current.x,
|
||||||
|
y: bot.tracker.current.y,
|
||||||
|
vx: speed * cos(ang),
|
||||||
|
vy: speed * sin(ang),
|
||||||
|
power: power,
|
||||||
|
)
|
||||||
|
inc bot.bulletCount
|
||||||
|
|
||||||
|
# Advance in-flight bullets and prune those off-arena (in-place compaction
|
||||||
|
# into the fixed buffer — no per-tick heap churn).
|
||||||
|
let aW = float64(getArenaWidth()); let aH = float64(getArenaHeight())
|
||||||
|
var n = 0
|
||||||
|
for i in 0 ..< bot.bulletCount:
|
||||||
|
let b = bot.bullets[i]
|
||||||
|
let nx = b.x + b.vx; let ny = b.y + b.vy
|
||||||
|
if nx >= 0.0 and nx <= aW and ny >= 0.0 and ny <= aH:
|
||||||
|
bot.bullets[n] = InFlightBullet(x: nx, y: ny, vx: b.vx, vy: b.vy, power: b.power)
|
||||||
|
inc n
|
||||||
|
bot.bulletCount = n
|
||||||
|
|
||||||
|
# Convert to BulletData for state vector (closest 3 by distance to us)
|
||||||
|
let myX2 = getX(); let myY2 = getY()
|
||||||
|
var bulletData: array[4, BulletData]
|
||||||
|
var bulletDataCount = 0
|
||||||
|
for i in 0 ..< bot.bulletCount:
|
||||||
|
let b = bot.bullets[i]
|
||||||
|
bulletData[bulletDataCount] = BulletData(x: b.x, y: b.y, power: b.power)
|
||||||
|
inc bulletDataCount
|
||||||
|
# sort ascending by distance so the nearest threats fill slots 0-2
|
||||||
|
if bulletDataCount > 1:
|
||||||
|
bulletData.toOpenArray(0, bulletDataCount - 1).sort(proc(a, b: BulletData): int =
|
||||||
|
let da = hypot(a.x - myX2, a.y - myY2)
|
||||||
|
let db = hypot(b.x - myX2, b.y - myY2)
|
||||||
|
cmp(da, db))
|
||||||
|
|
||||||
setRadarTurnRate(bot.tracker.getRadarTurnRate(getX(), getY(), getDirection(), getRadarDirection()))
|
setRadarTurnRate(bot.tracker.getRadarTurnRate(getX(), getY(), getDirection(), getRadarDirection()))
|
||||||
|
|
||||||
let botData = BotStateData(
|
let botData = BotStateData(
|
||||||
@@ -203,7 +264,11 @@ method run(bot: PPOBot) =
|
|||||||
let remainingGunAngle = abs(normalizeRelativeAngle(
|
let remainingGunAngle = abs(normalizeRelativeAngle(
|
||||||
directionTo(botData.x, botData.y, bot.lastActions.aimToX, bot.lastActions.aimToY) -
|
directionTo(botData.x, botData.y, bot.lastActions.aimToX, bot.lastActions.aimToY) -
|
||||||
botData.gunDirection))
|
botData.gunDirection))
|
||||||
let state = buildStateVector(botData, bot.tracker, remainingGotoDistance, remainingGunAngle)
|
let state = buildStateVector(botData, bot.tracker, remainingGotoDistance, remainingGunAngle, bulletData, bulletDataCount)
|
||||||
|
# Consume the fired flag AFTER the state build: state index 12 saw the
|
||||||
|
# detection-tick pulse, and the next iteration's spawn check sees false —
|
||||||
|
# one shot → exactly one bullet, even across deadReckon gaps.
|
||||||
|
bot.tracker.current.hasFired = false
|
||||||
let (rawActs, logP) = ac.actorForward(state)
|
let (rawActs, logP) = ac.actorForward(state)
|
||||||
let value = ac.criticForward(state)
|
let value = ac.criticForward(state)
|
||||||
let ex = if bot.tracker.hasContact: bot.tracker.current.x else: botData.arenaWidth / 2.0
|
let ex = if bot.tracker.hasContact: bot.tracker.current.x else: botData.arenaWidth / 2.0
|
||||||
|
|||||||
@@ -1,4 +1,7 @@
|
|||||||
#!/bin/sh
|
#!/bin/sh
|
||||||
# PPO_Bot — PPO-trained RL bot (compiled native binary)
|
# PPO_Bot — PPO-trained RL bot (compiled native binary)
|
||||||
|
# OPENBLAS_NUM_THREADS=1: prevent OpenBLAS from spawning worker threads,
|
||||||
|
# which deadlock when called from within a multi-threaded Nim bot process.
|
||||||
|
export OPENBLAS_NUM_THREADS=1
|
||||||
cd -- "$(dirname -- "$0")"
|
cd -- "$(dirname -- "$0")"
|
||||||
exec "./PPO_Bot"
|
exec "./PPO_Bot"
|
||||||
|
|||||||
@@ -27,13 +27,19 @@ proc update*(tracker: var EnemyTracker;
|
|||||||
## Call on ScannedBotEvent. Detects enemy fire from energy delta.
|
## Call on ScannedBotEvent. Detects enemy fire from energy delta.
|
||||||
|
|
||||||
# Shift history window
|
# Shift history window
|
||||||
if tracker.historyCount > 0:
|
if tracker.historyCount == 0:
|
||||||
|
# Cold-start: pre-fill all slots with the incoming scan so indices 22-41
|
||||||
|
# are never zero-padded on tick 1. Accel/turn-rate correctly stay 0 (no delta yet).
|
||||||
|
for i in 0 ..< 5:
|
||||||
|
tracker.history[i] = (scanX, scanY, scanDir, scanSpeed)
|
||||||
|
tracker.historyCount = 5
|
||||||
|
else:
|
||||||
for i in countdown(min(tracker.historyCount, 4), 1):
|
for i in countdown(min(tracker.historyCount, 4), 1):
|
||||||
tracker.history[i] = tracker.history[i - 1]
|
tracker.history[i] = tracker.history[i - 1]
|
||||||
tracker.history[0] = (tracker.current.x, tracker.current.y,
|
tracker.history[0] = (tracker.current.x, tracker.current.y,
|
||||||
tracker.current.direction, tracker.current.speed)
|
tracker.current.direction, tracker.current.speed)
|
||||||
if tracker.historyCount < 5:
|
if tracker.historyCount < 5:
|
||||||
inc tracker.historyCount
|
inc tracker.historyCount
|
||||||
|
|
||||||
# Detect firing: energy drop in [0.1, 3.0] means enemy fired
|
# Detect firing: energy drop in [0.1, 3.0] means enemy fired
|
||||||
let delta = tracker.prevEnergy - scanEnergy
|
let delta = tracker.prevEnergy - scanEnergy
|
||||||
|
|||||||
+4
-3
@@ -4,11 +4,12 @@ import arraymancer
|
|||||||
import std/[math, random]
|
import std/[math, random]
|
||||||
|
|
||||||
const
|
const
|
||||||
STATE_DIM* = 44
|
STATE_DIM* = 56
|
||||||
ACTION_DIM* = 6
|
ACTION_DIM* = 6
|
||||||
|
|
||||||
var
|
var
|
||||||
logStdFloor*: float32 = -3.0'f32 # overridden by PPOB_LOG_STD_FLOOR
|
logStdFloor*: float32 = -3.0'f32 # overridden by PPOB_LOG_STD_FLOOR
|
||||||
|
logStdCeiling*: float32 = 0.5'f32 # overridden by PPOB_LOG_STD_CEILING
|
||||||
initialLogStd*: float32 = 0.0'f32 # overridden by PPOB_INITIAL_LOG_STD
|
initialLogStd*: float32 = 0.0'f32 # overridden by PPOB_INITIAL_LOG_STD
|
||||||
|
|
||||||
type
|
type
|
||||||
@@ -46,7 +47,7 @@ proc actorForward*(ac: ActorCritic, state: Tensor[float32]): tuple[actions: Tens
|
|||||||
## state: [STATE_DIM]. Returns sampled actions [ACTION_DIM] and sum log-prob.
|
## state: [STATE_DIM]. Returns sampled actions [ACTION_DIM] and sum log-prob.
|
||||||
let mean = ac.actor.forward(state)
|
let mean = ac.actor.forward(state)
|
||||||
# Floor logStd at logStdFloor before exp → min std ≈ exp(logStdFloor)
|
# Floor logStd at logStdFloor before exp → min std ≈ exp(logStdFloor)
|
||||||
var clampedLogStd = ac.logStd.map(proc(v: float32): float32 = max(v, logStdFloor))
|
var clampedLogStd = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
||||||
let std = clampedLogStd.map(proc(v: float32): float32 = exp(v))
|
let std = clampedLogStd.map(proc(v: float32): float32 = exp(v))
|
||||||
|
|
||||||
var actions = newTensor[float32](ACTION_DIM)
|
var actions = newTensor[float32](ACTION_DIM)
|
||||||
@@ -70,7 +71,7 @@ proc criticForward*(ac: ActorCritic, state: Tensor[float32]): float32 =
|
|||||||
proc computeLogProb*(ac: ActorCritic, state, action: Tensor[float32]): float32 =
|
proc computeLogProb*(ac: ActorCritic, state, action: Tensor[float32]): float32 =
|
||||||
## Log-probability of action under current policy (no sampling).
|
## Log-probability of action under current policy (no sampling).
|
||||||
let mean = ac.actor.forward(state)
|
let mean = ac.actor.forward(state)
|
||||||
let logStdClamped = ac.logStd.map(proc(v: float32): float32 = max(v, logStdFloor))
|
let logStdClamped = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
||||||
let std = logStdClamped.map(proc(v: float32): float32 = exp(v))
|
let std = logStdClamped.map(proc(v: float32): float32 = exp(v))
|
||||||
var logP = 0.0'f32
|
var logP = 0.0'f32
|
||||||
for i in 0..<ACTION_DIM:
|
for i in 0..<ACTION_DIM:
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
## State vector builder — produces 44-float normalized tensor for PPO policy.
|
## State vector builder — produces 56-float normalized tensor for PPO policy.
|
||||||
## No bot API imports; takes plain BotState + EnemyTracker structs.
|
## No bot API imports; takes plain BotState + EnemyTracker structs.
|
||||||
|
|
||||||
import std/math
|
import std/math
|
||||||
@@ -16,12 +16,20 @@ type
|
|||||||
gunHeat*: float64
|
gunHeat*: float64
|
||||||
arenaWidth*, arenaHeight*: float64
|
arenaWidth*, arenaHeight*: float64
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker;
|
proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker;
|
||||||
remainingGotoDistance: float64 = 0.0;
|
remainingGotoDistance: float64 = 0.0;
|
||||||
remainingGunAngle: float64 = 0.0): Tensor[float32] =
|
remainingGunAngle: float64 = 0.0;
|
||||||
## Build the 44-float normalized state tensor.
|
bullets: openArray[BulletData] = [];
|
||||||
|
bulletCount: int = 0): Tensor[float32] =
|
||||||
|
## Build the 56-float normalized state tensor.
|
||||||
|
## Indices 0-43: existing features. Indices 44-55: up to 3 bullet slots (4 floats each).
|
||||||
## All values clipped to roughly [-1, 1] via division by physical maxima.
|
## All values clipped to roughly [-1, 1] via division by physical maxima.
|
||||||
result = zeros[float32](44)
|
result = zeros[float32](56)
|
||||||
|
|
||||||
let aW = bot.arenaWidth
|
let aW = bot.arenaWidth
|
||||||
let aH = bot.arenaHeight
|
let aH = bot.arenaHeight
|
||||||
@@ -91,3 +99,21 @@ proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker;
|
|||||||
# --- Goto controller inputs (indices 42-43) ---
|
# --- Goto controller inputs (indices 42-43) ---
|
||||||
result[42] = float32(remainingGotoDistance / diag)
|
result[42] = float32(remainingGotoDistance / diag)
|
||||||
result[43] = float32(remainingGunAngle / 180.0)
|
result[43] = float32(remainingGunAngle / 180.0)
|
||||||
|
|
||||||
|
# --- 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
|
||||||
|
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 bdist = sqrt(bdx * bdx + bdy * bdy)
|
||||||
|
let ticks = if bSpeed > 0.0: bdist / bSpeed else: 0.0
|
||||||
|
let base = 44 + i * 4
|
||||||
|
result[base + 0] = float32(bdx / bot.arenaWidth)
|
||||||
|
result[base + 1] = float32(bdy / bot.arenaHeight)
|
||||||
|
result[base + 2] = float32(bSpeed / 20.0)
|
||||||
|
result[base + 3] = float32(ticks / diag)
|
||||||
|
|||||||
@@ -84,7 +84,7 @@ block testStateVectorLength:
|
|||||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||||
)
|
)
|
||||||
let sv = buildStateVector(bot, t)
|
let sv = buildStateVector(bot, t)
|
||||||
check sv.shape == [44], "state vector has 44 elements"
|
check sv.shape == [56], "state vector has 56 elements"
|
||||||
|
|
||||||
block testStateVectorRange:
|
block testStateVectorRange:
|
||||||
var t = initEnemyTracker()
|
var t = initEnemyTracker()
|
||||||
@@ -95,7 +95,7 @@ block testStateVectorRange:
|
|||||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||||
)
|
)
|
||||||
let sv = buildStateVector(bot, t)
|
let sv = buildStateVector(bot, t)
|
||||||
for i in 0 ..< 44:
|
for i in 0 ..< 56:
|
||||||
check sv[i] >= -2.0f32 and sv[i] <= 2.0f32,
|
check sv[i] >= -2.0f32 and sv[i] <= 2.0f32,
|
||||||
&"sv[{i}]={sv[i]} out of [-2,2] range"
|
&"sv[{i}]={sv[i]} out of [-2,2] range"
|
||||||
|
|
||||||
@@ -158,5 +158,28 @@ block testHistoryPaddedWhenEmpty:
|
|||||||
# indices 42-43 (goto inputs) default to 0 when not provided
|
# indices 42-43 (goto inputs) default to 0 when not provided
|
||||||
check sv[42] == 0.0f32, "sv[42] (remainingGotoDistance) should default to 0"
|
check sv[42] == 0.0f32, "sv[42] (remainingGotoDistance) should default to 0"
|
||||||
check sv[43] == 0.0f32, "sv[43] (remainingGunAngle) should default to 0"
|
check sv[43] == 0.0f32, "sv[43] (remainingGunAngle) should default to 0"
|
||||||
|
# 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"
|
||||||
|
|
||||||
|
block testBulletSlots:
|
||||||
|
var t = initEnemyTracker()
|
||||||
|
t.update(400.0, 300.0, 0.0, 0.0, 100.0) # enemy at (400,300)
|
||||||
|
let bot = BotStateData(
|
||||||
|
x: 200.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,
|
||||||
|
)
|
||||||
|
# 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
|
||||||
|
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[46] - (17.0/20.0).float32) < 0.001f32, "bullet speed norm"
|
||||||
|
# second slot should be zero-padded
|
||||||
|
for i in 48 ..< 56:
|
||||||
|
check sv[i] == 0.0f32, &"unused bullet slot {i} should be 0"
|
||||||
|
|
||||||
echo "All tests passed"
|
echo "All tests passed"
|
||||||
|
|||||||
@@ -149,4 +149,27 @@ block testPpoUpdateNormalReward:
|
|||||||
check m.valueLoss == m.valueLoss, "valueLoss NaN on normal-reward round"
|
check m.valueLoss == m.valueLoss, "valueLoss NaN on normal-reward round"
|
||||||
check m.gradNorm == m.gradNorm, "gradNorm NaN on normal-reward round"
|
check m.gradNorm == m.gradNorm, "gradNorm NaN on normal-reward round"
|
||||||
|
|
||||||
|
# ── logStd ceiling: raw param must never drift above the collection clamp ─────
|
||||||
|
# Regression for the train/collection std mismatch: logStd starting above the
|
||||||
|
# ceiling (as in the warm-start snapshot) must be clamped back to [floor, ceiling]
|
||||||
|
# by the first Adam step, so recomputed logP matches the acting policy's std.
|
||||||
|
|
||||||
|
block testLogStdCeilingClamp:
|
||||||
|
randomize(45)
|
||||||
|
var ac = initActorCritic()
|
||||||
|
ac.logStd = newTensor[float32](ACTION_DIM).map(proc(v: float32): float32 = 3.0'f32)
|
||||||
|
var buf = initTrajectoryBuffer()
|
||||||
|
for _ in 0..<16:
|
||||||
|
let s = randomNormalTensor[float32](STATE_DIM)
|
||||||
|
let a = randomNormalTensor[float32](ACTION_DIM)
|
||||||
|
let lp = ac.computeLogProb(s, a)
|
||||||
|
buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp,
|
||||||
|
reward: 0.1'f32, value: 0.5'f32))
|
||||||
|
var adam: ACAdamStates
|
||||||
|
discard ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam,
|
||||||
|
epochs = 1, miniBatchSize = 16)
|
||||||
|
for v in ac.logStd:
|
||||||
|
check v <= logStdCeiling + 1e-6'f32, "logStd above ceiling after ppoUpdate"
|
||||||
|
check v >= logStdFloor - 1e-6'f32, "logStd below floor after ppoUpdate"
|
||||||
|
|
||||||
echo "All tests passed"
|
echo "All tests passed"
|
||||||
|
|||||||
+33
-10
@@ -218,8 +218,12 @@ proc ppoUpdate*(ac: var ActorCritic;
|
|||||||
var totalGradNorm = 0.0'f32
|
var totalGradNorm = 0.0'f32
|
||||||
var totalMiniBatches = 0
|
var totalMiniBatches = 0
|
||||||
|
|
||||||
# Initialise Adam states once; caller persists them across rounds
|
# Initialise Adam states once; caller persists them across rounds.
|
||||||
if not adamStates.initialized:
|
# Also reinit if aw1.m has wrong shape (e.g. loaded from old checkpoint with
|
||||||
|
# different STATE_DIM, leaving a (0,) placeholder after shape-mismatch skip).
|
||||||
|
if not adamStates.initialized or
|
||||||
|
adamStates.aw1.m.shape.len == 0 or
|
||||||
|
adamStates.aw1.m.shape != ac.actor.w1.shape:
|
||||||
adamStates = initACAdamStates(ac)
|
adamStates = initACAdamStates(ac)
|
||||||
|
|
||||||
# 1. GAE
|
# 1. GAE
|
||||||
@@ -253,7 +257,7 @@ proc ppoUpdate*(ac: var ActorCritic;
|
|||||||
# Treat as full-batch; breaks the loop unconditionally.
|
# Treat as full-batch; breaks the loop unconditionally.
|
||||||
let mbSizeCap = if miniBatchSize > 0: miniBatchSize else: bufLen
|
let mbSizeCap = if miniBatchSize > 0: miniBatchSize else: bufLen
|
||||||
|
|
||||||
for _ in 1..epochs:
|
for epochNum in 1..epochs:
|
||||||
# Shuffle indices
|
# Shuffle indices
|
||||||
var indices = toSeq(0..<bufLen)
|
var indices = toSeq(0..<bufLen)
|
||||||
shuffle(indices)
|
shuffle(indices)
|
||||||
@@ -263,7 +267,6 @@ proc ppoUpdate*(ac: var ActorCritic;
|
|||||||
let mbEnd = min(mbStart + mbSizeCap, bufLen)
|
let mbEnd = min(mbStart + mbSizeCap, bufLen)
|
||||||
if mbEnd <= mbStart: break
|
if mbEnd <= mbStart: break
|
||||||
let mbSize = mbEnd - mbStart
|
let mbSize = mbEnd - mbStart
|
||||||
|
|
||||||
# Accumulators for gradients (zero-init)
|
# Accumulators for gradients (zero-init)
|
||||||
var dActorW1 = zeros[float32](ac.actor.w1.shape)
|
var dActorW1 = zeros[float32](ac.actor.w1.shape)
|
||||||
var dActorB1 = zeros[float32](ac.actor.b1.shape)
|
var dActorB1 = zeros[float32](ac.actor.b1.shape)
|
||||||
@@ -292,9 +295,13 @@ proc ppoUpdate*(ac: var ActorCritic;
|
|||||||
|
|
||||||
# ── Actor forward ──
|
# ── Actor forward ──
|
||||||
let actorFwd = mlpForwardCached(ac.actor, x)
|
let actorFwd = mlpForwardCached(ac.actor, x)
|
||||||
|
# Critic forward
|
||||||
|
|
||||||
let newMean = actorFwd.y # [ACTION_DIM]
|
let newMean = actorFwd.y # [ACTION_DIM]
|
||||||
|
|
||||||
let logStdClamped = ac.logStd.map(proc(v: float32): float32 = max(v, logStdFloor))
|
# Same clamp as collection (network.nim actorForward): train-time std must
|
||||||
|
# exactly match the std the acting policy used, or ratios are distorted.
|
||||||
|
let logStdClamped = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
||||||
let std = logStdClamped.map(proc(v: float32): float32 = exp(v))
|
let std = logStdClamped.map(proc(v: float32): float32 = exp(v))
|
||||||
|
|
||||||
# New log prob
|
# New log prob
|
||||||
@@ -337,13 +344,18 @@ proc ppoUpdate*(ac: var ActorCritic;
|
|||||||
# Also: d(newLogP)/d(logStd_i) when logStd_i > -3:
|
# Also: d(newLogP)/d(logStd_i) when logStd_i > -3:
|
||||||
# = (action_i - mean_i)^2/std_i^2 - 1
|
# = (action_i - mean_i)^2/std_i^2 - 1
|
||||||
for i in 0..<ACTION_DIM:
|
for i in 0..<ACTION_DIM:
|
||||||
let isClamped = (ac.logStd[i] <= logStdFloor)
|
let isFloorClamped = (ac.logStd[i] <= logStdFloor)
|
||||||
if not isClamped:
|
let isCeilingClamped = (ac.logStd[i] >= logStdCeiling)
|
||||||
|
if not isFloorClamped:
|
||||||
let s = std[i]
|
let s = std[i]
|
||||||
let diff = (tr.action[i] - newMean[i]) / s
|
let diff = (tr.action[i] - newMean[i]) / s
|
||||||
let dLogP_dLogStdI = diff * diff - 1.0'f32
|
let dLogP_dLogStdI = diff * diff - 1.0'f32
|
||||||
dLogStd[i] += dLoss_dNewLogP * dLogP_dLogStdI -
|
let ppoGrad = dLoss_dNewLogP * dLogP_dLogStdI
|
||||||
entropyCoeff / mbSize.float32 # entropy: d(-entropyCoeff*H)/d(logStd_i) = -entropyCoeff
|
# Entropy term pushes logStd up (update = param - lr*grad, grad is -entropyCoeff < 0).
|
||||||
|
# Gate it off at the ceiling to prevent runaway logStd.
|
||||||
|
let entropyGrad = if isCeilingClamped: 0.0'f32
|
||||||
|
else: -entropyCoeff / mbSize.float32
|
||||||
|
dLogStd[i] += ppoGrad + entropyGrad
|
||||||
|
|
||||||
# Backprop actor gradients
|
# Backprop actor gradients
|
||||||
let gradActorOut = dLoss_dNewLogP *. dLogP_dMean # [5]
|
let gradActorOut = dLoss_dNewLogP *. dLogP_dMean # [5]
|
||||||
@@ -380,6 +392,13 @@ proc ppoUpdate*(ac: var ActorCritic;
|
|||||||
dCriticW1, dCriticB1, dCriticW2, dCriticB2, dCriticW3, dCriticB3
|
dCriticW1, dCriticB1, dCriticW2, dCriticB2, dCriticW3, dCriticB3
|
||||||
]
|
]
|
||||||
let norm = globalNorm(allGrads)
|
let norm = globalNorm(allGrads)
|
||||||
|
# ponytail: NaN/Inf grad norm means something exploded this minibatch
|
||||||
|
# (extreme logprob ratios, poisoned initial weights, etc.). Skip the
|
||||||
|
# Adam update entirely — no-op is safer than writing NaN into weights,
|
||||||
|
# which corrupts all future inference and hangs the bot.
|
||||||
|
if norm != norm or norm > 1e15'f32:
|
||||||
|
mbStart = mbEnd
|
||||||
|
continue
|
||||||
totalGradNorm += norm
|
totalGradNorm += norm
|
||||||
inc totalMiniBatches
|
inc totalMiniBatches
|
||||||
if norm > maxGradNorm:
|
if norm > maxGradNorm:
|
||||||
@@ -403,13 +422,17 @@ proc ppoUpdate*(ac: var ActorCritic;
|
|||||||
adamStep(ac.actor.w3, dActorW3, adamStates.aw3, lr)
|
adamStep(ac.actor.w3, dActorW3, adamStates.aw3, lr)
|
||||||
adamStep(ac.actor.b3, dActorB3, adamStates.ab3, lr)
|
adamStep(ac.actor.b3, dActorB3, adamStates.ab3, lr)
|
||||||
adamStep(ac.logStd, dLogStd, adamStates.logStd, lr)
|
adamStep(ac.logStd, dLogStd, adamStates.logStd, lr)
|
||||||
|
# Collection clamps logStd to [floor, ceiling] at inference; clamp the raw
|
||||||
|
# param after the step so it can't drift above the ceiling (the old code
|
||||||
|
# only clamped at collection → train-time recompute used a bigger std than
|
||||||
|
# the policy that actually acted → distorted importance ratios).
|
||||||
|
ac.logStd = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
||||||
adamStep(ac.critic.w1, dCriticW1, adamStates.cw1, lr)
|
adamStep(ac.critic.w1, dCriticW1, adamStates.cw1, lr)
|
||||||
adamStep(ac.critic.b1, dCriticB1, adamStates.cb1, lr)
|
adamStep(ac.critic.b1, dCriticB1, adamStates.cb1, lr)
|
||||||
adamStep(ac.critic.w2, dCriticW2, adamStates.cw2, lr)
|
adamStep(ac.critic.w2, dCriticW2, adamStates.cw2, lr)
|
||||||
adamStep(ac.critic.b2, dCriticB2, adamStates.cb2, lr)
|
adamStep(ac.critic.b2, dCriticB2, adamStates.cb2, lr)
|
||||||
adamStep(ac.critic.w3, dCriticW3, adamStates.cw3, lr)
|
adamStep(ac.critic.w3, dCriticW3, adamStates.cw3, lr)
|
||||||
adamStep(ac.critic.b3, dCriticB3, adamStates.cb3, lr)
|
adamStep(ac.critic.b3, dCriticB3, adamStates.cb3, lr)
|
||||||
|
|
||||||
mbStart = mbEnd
|
mbStart = mbEnd
|
||||||
|
|
||||||
let totalSamples = (epochs * bufLen).float32
|
let totalSamples = (epochs * bufLen).float32
|
||||||
|
|||||||
+34
-20
@@ -33,26 +33,30 @@ proc saveWeights*(ac: ActorCritic, dir: string) =
|
|||||||
ac.logStd.write_npy(dir / "log_std.npy")
|
ac.logStd.write_npy(dir / "log_std.npy")
|
||||||
|
|
||||||
proc loadWeights*(ac: var ActorCritic, dir: string) =
|
proc loadWeights*(ac: var ActorCritic, dir: string) =
|
||||||
## Load all weight tensors from dir/. Asserts shapes match.
|
## Load all weight tensors from dir/.
|
||||||
template loadAndCheck(dest: untyped, path: string) =
|
## If a tensor's shape doesn't match (e.g. STATE_DIM changed), keep the
|
||||||
|
## freshly-initialised value and print a warning — other tensors still load.
|
||||||
|
template loadOrSkip(dest: untyped, path: string) =
|
||||||
let loaded = read_npy[float32](path)
|
let loaded = read_npy[float32](path)
|
||||||
doAssert loaded.shape == dest.shape,
|
if loaded.shape == dest.shape:
|
||||||
"Shape mismatch loading " & path & ": got " & $loaded.shape & " want " & $dest.shape
|
dest = loaded
|
||||||
dest = loaded
|
else:
|
||||||
|
echo "weights: shape mismatch for " & path &
|
||||||
|
" (got " & $loaded.shape & " want " & $dest.shape & ") — keeping fresh init"
|
||||||
|
|
||||||
loadAndCheck(ac.actor.w1, dir / "actor_w1.npy")
|
loadOrSkip(ac.actor.w1, dir / "actor_w1.npy")
|
||||||
loadAndCheck(ac.actor.b1, dir / "actor_b1.npy")
|
loadOrSkip(ac.actor.b1, dir / "actor_b1.npy")
|
||||||
loadAndCheck(ac.actor.w2, dir / "actor_w2.npy")
|
loadOrSkip(ac.actor.w2, dir / "actor_w2.npy")
|
||||||
loadAndCheck(ac.actor.b2, dir / "actor_b2.npy")
|
loadOrSkip(ac.actor.b2, dir / "actor_b2.npy")
|
||||||
loadAndCheck(ac.actor.w3, dir / "actor_w3.npy")
|
loadOrSkip(ac.actor.w3, dir / "actor_w3.npy")
|
||||||
loadAndCheck(ac.actor.b3, dir / "actor_b3.npy")
|
loadOrSkip(ac.actor.b3, dir / "actor_b3.npy")
|
||||||
loadAndCheck(ac.critic.w1, dir / "critic_w1.npy")
|
loadOrSkip(ac.critic.w1, dir / "critic_w1.npy")
|
||||||
loadAndCheck(ac.critic.b1, dir / "critic_b1.npy")
|
loadOrSkip(ac.critic.b1, dir / "critic_b1.npy")
|
||||||
loadAndCheck(ac.critic.w2, dir / "critic_w2.npy")
|
loadOrSkip(ac.critic.w2, dir / "critic_w2.npy")
|
||||||
loadAndCheck(ac.critic.b2, dir / "critic_b2.npy")
|
loadOrSkip(ac.critic.b2, dir / "critic_b2.npy")
|
||||||
loadAndCheck(ac.critic.w3, dir / "critic_w3.npy")
|
loadOrSkip(ac.critic.w3, dir / "critic_w3.npy")
|
||||||
loadAndCheck(ac.critic.b3, dir / "critic_b3.npy")
|
loadOrSkip(ac.critic.b3, dir / "critic_b3.npy")
|
||||||
loadAndCheck(ac.logStd, dir / "log_std.npy")
|
loadOrSkip(ac.logStd, dir / "log_std.npy")
|
||||||
|
|
||||||
proc saveWeightsAtomic*(ac: ActorCritic, targetDir: string) =
|
proc saveWeightsAtomic*(ac: ActorCritic, targetDir: string) =
|
||||||
## Write to a temp dir, then rename atomically over targetDir.
|
## Write to a temp dir, then rename atomically over targetDir.
|
||||||
@@ -104,8 +108,14 @@ proc saveAdamStates*(adam: ACAdamStates, dir: string) =
|
|||||||
|
|
||||||
proc loadAdamStates*(adam: var ACAdamStates, dir: string) =
|
proc loadAdamStates*(adam: var ACAdamStates, dir: string) =
|
||||||
## Load Adam m/v tensors and t counters from dir/. Called only when files exist.
|
## Load Adam m/v tensors and t counters from dir/. Called only when files exist.
|
||||||
|
## Shape mismatch (e.g. STATE_DIM changed) → keep zero-initialised state (safe fresh start).
|
||||||
template lm(dest: untyped, path: string) =
|
template lm(dest: untyped, path: string) =
|
||||||
dest = read_npy[float32](path)
|
let loaded = read_npy[float32](path)
|
||||||
|
if loaded.shape == dest.shape:
|
||||||
|
dest = loaded
|
||||||
|
else:
|
||||||
|
echo "weights: Adam shape mismatch for " & path &
|
||||||
|
" (got " & $loaded.shape & " want " & $dest.shape & ") — resetting Adam state"
|
||||||
lm(adam.aw1.m, dir / "adam_aw1_m.npy"); lm(adam.aw1.v, dir / "adam_aw1_v.npy")
|
lm(adam.aw1.m, dir / "adam_aw1_m.npy"); lm(adam.aw1.v, dir / "adam_aw1_v.npy")
|
||||||
lm(adam.ab1.m, dir / "adam_ab1_m.npy"); lm(adam.ab1.v, dir / "adam_ab1_v.npy")
|
lm(adam.ab1.m, dir / "adam_ab1_m.npy"); lm(adam.ab1.v, dir / "adam_ab1_v.npy")
|
||||||
lm(adam.aw2.m, dir / "adam_aw2_m.npy"); lm(adam.aw2.v, dir / "adam_aw2_v.npy")
|
lm(adam.aw2.m, dir / "adam_aw2_m.npy"); lm(adam.aw2.v, dir / "adam_aw2_v.npy")
|
||||||
@@ -185,7 +195,11 @@ proc loadBestAvailable*(ac: var ActorCritic, adam: var ACAdamStates,
|
|||||||
if adamStateFilesExist(candidate):
|
if adamStateFilesExist(candidate):
|
||||||
adam.loadAdamStates(candidate)
|
adam.loadAdamStates(candidate)
|
||||||
let rcPath = weightsRoot / "round_counter.txt"
|
let rcPath = weightsRoot / "round_counter.txt"
|
||||||
let roundNum = if fileExists(rcPath): parseInt(readFile(rcPath).strip()) else: 0
|
# Torn/empty file (e.g. after a crash) must not abort startup → treat as 0
|
||||||
|
let roundNum = if fileExists(rcPath):
|
||||||
|
try: parseInt(readFile(rcPath).strip())
|
||||||
|
except ValueError: 0
|
||||||
|
else: 0
|
||||||
return (loaded: true, roundNum: roundNum)
|
return (loaded: true, roundNum: roundNum)
|
||||||
result = (loaded: false, roundNum: 0)
|
result = (loaded: false, roundNum: 0)
|
||||||
|
|
||||||
|
|||||||
@@ -1 +1 @@
|
|||||||
9500
|
0
|
||||||
|
|||||||
@@ -70,22 +70,25 @@ public class RunTraining {
|
|||||||
);
|
);
|
||||||
|
|
||||||
var owner = new Object();
|
var owner = new Object();
|
||||||
|
int[] prevTotal = { 0 };
|
||||||
|
|
||||||
try (var handle = runner.startBattleAsync(setup, bots)) {
|
try (var handle = runner.startBattleAsync(setup, bots)) {
|
||||||
|
|
||||||
handle.getOnRoundEnded().on(owner, event -> {
|
handle.getOnRoundEnded().on(owner, event -> {
|
||||||
int round = event.getRoundNumber();
|
int round = event.getRoundNumber();
|
||||||
int ticks = event.getTurnNumber();
|
int ticks = event.getTurnNumber();
|
||||||
int score = 0;
|
int totalScore = 0;
|
||||||
boolean win = false;
|
boolean win = false;
|
||||||
boolean found = false;
|
boolean found = false;
|
||||||
for (var r : event.getResults()) {
|
for (var r : event.getResults()) {
|
||||||
if (r.getName().equals("PPO_Bot")) {
|
if (r.getName().equals("PPO_Bot")) {
|
||||||
found = true;
|
found = true;
|
||||||
score = r.getTotalScore();
|
totalScore = r.getTotalScore();
|
||||||
win = r.getRank() == 1;
|
win = r.getRank() == 1;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
int score = totalScore - prevTotal[0];
|
||||||
|
prevTotal[0] = totalScore;
|
||||||
// PPO_Bot's process died mid-battle: abort so run.sh's crash-restart
|
// PPO_Bot's process died mid-battle: abort so run.sh's crash-restart
|
||||||
// loop resumes from round_counter instead of grinding dummy rounds.
|
// loop resumes from round_counter instead of grinding dummy rounds.
|
||||||
// Frozen for a few harness rounds can be a healthy-but-lagging counter
|
// Frozen for a few harness rounds can be a healthy-but-lagging counter
|
||||||
@@ -113,8 +116,8 @@ public class RunTraining {
|
|||||||
}
|
}
|
||||||
// Append game-outcome JSON line
|
// Append game-outcome JSON line
|
||||||
String line = String.format(
|
String line = String.format(
|
||||||
"{\"type\":\"game\",\"round\":%d,\"ticks\":%d,\"score\":%d,\"win\":%b,\"opponent\":\"%s\"}",
|
"{\"type\":\"game\",\"round\":%d,\"ticks\":%d,\"score\":%d,\"total_score\":%d,\"win\":%b,\"opponent\":\"%s\"}",
|
||||||
round, ticks, score, win, opponent
|
round, ticks, score, totalScore, win, opponent
|
||||||
);
|
);
|
||||||
try {
|
try {
|
||||||
appendLine(logFile, line);
|
appendLine(logFile, line);
|
||||||
|
|||||||
Reference in New Issue
Block a user