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.
|
||||
## 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 tankroyale_botapi
|
||||
import network
|
||||
@@ -35,6 +35,7 @@ var
|
||||
|
||||
# Wire logStd tunable params into network module vars (read before initActorCritic)
|
||||
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)
|
||||
|
||||
# ── Structured log output ─────────────────────────────────────────────────────
|
||||
@@ -59,11 +60,17 @@ proc hyperparmSnapshot(): string =
|
||||
&"\"entropyCoeff\":{hpEntropyCoeff},\"valueLossCoeff\":{hpValueLossCoeff}," &
|
||||
&"\"maxGradNorm\":{hpMaxGradNorm},\"gamma\":{hpGamma},\"lam\":{hpLam}," &
|
||||
&"\"epochs\":{hpEpochs},\"miniBatchSize\":{hpMiniBatchSize}," &
|
||||
&"\"logStdFloor\":{logStdFloor},\"initialLogStd\":{initialLogStd}"
|
||||
&"\"logStdFloor\":{logStdFloor},\"logStdCeiling\":{logStdCeiling},\"initialLogStd\":{initialLogStd}"
|
||||
|
||||
const botJsonPath = currentSourcePath().parentDir / "PPO_Bot.json"
|
||||
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
|
||||
tracker: EnemyTracker
|
||||
buffer: TrajectoryBuffer
|
||||
@@ -77,6 +84,12 @@ type PPOBot = ref object of Bot
|
||||
lastActions: BotActions # previous tick's decoded actions (for state vector)
|
||||
roundRewardSum: float32 # cumulative reward this round (for live display)
|
||||
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 gAdamStates: ACAdamStates # persists across rounds
|
||||
@@ -98,6 +111,7 @@ method onRoundStarted*(bot: PPOBot, e: RoundStartedEvent) =
|
||||
bot.hasLastTrans = false
|
||||
bot.roundRewardSum = 0.0'f32
|
||||
bot.roundTicks = 0
|
||||
bot.bulletCount = 0
|
||||
|
||||
method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
||||
inc roundCounter
|
||||
@@ -184,6 +198,53 @@ method run(bot: PPOBot) =
|
||||
while isRunning():
|
||||
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()))
|
||||
|
||||
let botData = BotStateData(
|
||||
@@ -203,7 +264,11 @@ method run(bot: PPOBot) =
|
||||
let remainingGunAngle = abs(normalizeRelativeAngle(
|
||||
directionTo(botData.x, botData.y, bot.lastActions.aimToX, bot.lastActions.aimToY) -
|
||||
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 value = ac.criticForward(state)
|
||||
let ex = if bot.tracker.hasContact: bot.tracker.current.x else: botData.arenaWidth / 2.0
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
#!/bin/sh
|
||||
# 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")"
|
||||
exec "./PPO_Bot"
|
||||
|
||||
@@ -27,7 +27,13 @@ proc update*(tracker: var EnemyTracker;
|
||||
## Call on ScannedBotEvent. Detects enemy fire from energy delta.
|
||||
|
||||
# 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):
|
||||
tracker.history[i] = tracker.history[i - 1]
|
||||
tracker.history[0] = (tracker.current.x, tracker.current.y,
|
||||
|
||||
+4
-3
@@ -4,11 +4,12 @@ import arraymancer
|
||||
import std/[math, random]
|
||||
|
||||
const
|
||||
STATE_DIM* = 44
|
||||
STATE_DIM* = 56
|
||||
ACTION_DIM* = 6
|
||||
|
||||
var
|
||||
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
|
||||
|
||||
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.
|
||||
let mean = ac.actor.forward(state)
|
||||
# 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))
|
||||
|
||||
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 =
|
||||
## Log-probability of action under current policy (no sampling).
|
||||
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))
|
||||
var logP = 0.0'f32
|
||||
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.
|
||||
|
||||
import std/math
|
||||
@@ -16,12 +16,20 @@ type
|
||||
gunHeat*: 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;
|
||||
remainingGotoDistance: float64 = 0.0;
|
||||
remainingGunAngle: float64 = 0.0): Tensor[float32] =
|
||||
## Build the 44-float normalized state tensor.
|
||||
remainingGunAngle: float64 = 0.0;
|
||||
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.
|
||||
result = zeros[float32](44)
|
||||
result = zeros[float32](56)
|
||||
|
||||
let aW = bot.arenaWidth
|
||||
let aH = bot.arenaHeight
|
||||
@@ -91,3 +99,21 @@ proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker;
|
||||
# --- Goto controller inputs (indices 42-43) ---
|
||||
result[42] = float32(remainingGotoDistance / diag)
|
||||
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,
|
||||
)
|
||||
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:
|
||||
var t = initEnemyTracker()
|
||||
@@ -95,7 +95,7 @@ block testStateVectorRange:
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
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,
|
||||
&"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
|
||||
check sv[42] == 0.0f32, "sv[42] (remainingGotoDistance) 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"
|
||||
|
||||
@@ -149,4 +149,27 @@ block testPpoUpdateNormalReward:
|
||||
check m.valueLoss == m.valueLoss, "valueLoss 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"
|
||||
|
||||
+33
-10
@@ -218,8 +218,12 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
var totalGradNorm = 0.0'f32
|
||||
var totalMiniBatches = 0
|
||||
|
||||
# Initialise Adam states once; caller persists them across rounds
|
||||
if not adamStates.initialized:
|
||||
# Initialise Adam states once; caller persists them across rounds.
|
||||
# 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)
|
||||
|
||||
# 1. GAE
|
||||
@@ -253,7 +257,7 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
# Treat as full-batch; breaks the loop unconditionally.
|
||||
let mbSizeCap = if miniBatchSize > 0: miniBatchSize else: bufLen
|
||||
|
||||
for _ in 1..epochs:
|
||||
for epochNum in 1..epochs:
|
||||
# Shuffle indices
|
||||
var indices = toSeq(0..<bufLen)
|
||||
shuffle(indices)
|
||||
@@ -263,7 +267,6 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
let mbEnd = min(mbStart + mbSizeCap, bufLen)
|
||||
if mbEnd <= mbStart: break
|
||||
let mbSize = mbEnd - mbStart
|
||||
|
||||
# Accumulators for gradients (zero-init)
|
||||
var dActorW1 = zeros[float32](ac.actor.w1.shape)
|
||||
var dActorB1 = zeros[float32](ac.actor.b1.shape)
|
||||
@@ -292,9 +295,13 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
|
||||
# ── Actor forward ──
|
||||
let actorFwd = mlpForwardCached(ac.actor, x)
|
||||
# Critic forward
|
||||
|
||||
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))
|
||||
|
||||
# New log prob
|
||||
@@ -337,13 +344,18 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
# Also: d(newLogP)/d(logStd_i) when logStd_i > -3:
|
||||
# = (action_i - mean_i)^2/std_i^2 - 1
|
||||
for i in 0..<ACTION_DIM:
|
||||
let isClamped = (ac.logStd[i] <= logStdFloor)
|
||||
if not isClamped:
|
||||
let isFloorClamped = (ac.logStd[i] <= logStdFloor)
|
||||
let isCeilingClamped = (ac.logStd[i] >= logStdCeiling)
|
||||
if not isFloorClamped:
|
||||
let s = std[i]
|
||||
let diff = (tr.action[i] - newMean[i]) / s
|
||||
let dLogP_dLogStdI = diff * diff - 1.0'f32
|
||||
dLogStd[i] += dLoss_dNewLogP * dLogP_dLogStdI -
|
||||
entropyCoeff / mbSize.float32 # entropy: d(-entropyCoeff*H)/d(logStd_i) = -entropyCoeff
|
||||
let ppoGrad = dLoss_dNewLogP * dLogP_dLogStdI
|
||||
# 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
|
||||
let gradActorOut = dLoss_dNewLogP *. dLogP_dMean # [5]
|
||||
@@ -380,6 +392,13 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
dCriticW1, dCriticB1, dCriticW2, dCriticB2, dCriticW3, dCriticB3
|
||||
]
|
||||
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
|
||||
inc totalMiniBatches
|
||||
if norm > maxGradNorm:
|
||||
@@ -403,13 +422,17 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
adamStep(ac.actor.w3, dActorW3, adamStates.aw3, lr)
|
||||
adamStep(ac.actor.b3, dActorB3, adamStates.ab3, 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.b1, dCriticB1, adamStates.cb1, lr)
|
||||
adamStep(ac.critic.w2, dCriticW2, adamStates.cw2, lr)
|
||||
adamStep(ac.critic.b2, dCriticB2, adamStates.cb2, lr)
|
||||
adamStep(ac.critic.w3, dCriticW3, adamStates.cw3, lr)
|
||||
adamStep(ac.critic.b3, dCriticB3, adamStates.cb3, lr)
|
||||
|
||||
mbStart = mbEnd
|
||||
|
||||
let totalSamples = (epochs * bufLen).float32
|
||||
|
||||
+33
-19
@@ -33,26 +33,30 @@ proc saveWeights*(ac: ActorCritic, dir: string) =
|
||||
ac.logStd.write_npy(dir / "log_std.npy")
|
||||
|
||||
proc loadWeights*(ac: var ActorCritic, dir: string) =
|
||||
## Load all weight tensors from dir/. Asserts shapes match.
|
||||
template loadAndCheck(dest: untyped, path: string) =
|
||||
## Load all weight tensors from dir/.
|
||||
## 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)
|
||||
doAssert loaded.shape == dest.shape,
|
||||
"Shape mismatch loading " & path & ": got " & $loaded.shape & " want " & $dest.shape
|
||||
if loaded.shape == dest.shape:
|
||||
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")
|
||||
loadAndCheck(ac.actor.b1, dir / "actor_b1.npy")
|
||||
loadAndCheck(ac.actor.w2, dir / "actor_w2.npy")
|
||||
loadAndCheck(ac.actor.b2, dir / "actor_b2.npy")
|
||||
loadAndCheck(ac.actor.w3, dir / "actor_w3.npy")
|
||||
loadAndCheck(ac.actor.b3, dir / "actor_b3.npy")
|
||||
loadAndCheck(ac.critic.w1, dir / "critic_w1.npy")
|
||||
loadAndCheck(ac.critic.b1, dir / "critic_b1.npy")
|
||||
loadAndCheck(ac.critic.w2, dir / "critic_w2.npy")
|
||||
loadAndCheck(ac.critic.b2, dir / "critic_b2.npy")
|
||||
loadAndCheck(ac.critic.w3, dir / "critic_w3.npy")
|
||||
loadAndCheck(ac.critic.b3, dir / "critic_b3.npy")
|
||||
loadAndCheck(ac.logStd, dir / "log_std.npy")
|
||||
loadOrSkip(ac.actor.w1, dir / "actor_w1.npy")
|
||||
loadOrSkip(ac.actor.b1, dir / "actor_b1.npy")
|
||||
loadOrSkip(ac.actor.w2, dir / "actor_w2.npy")
|
||||
loadOrSkip(ac.actor.b2, dir / "actor_b2.npy")
|
||||
loadOrSkip(ac.actor.w3, dir / "actor_w3.npy")
|
||||
loadOrSkip(ac.actor.b3, dir / "actor_b3.npy")
|
||||
loadOrSkip(ac.critic.w1, dir / "critic_w1.npy")
|
||||
loadOrSkip(ac.critic.b1, dir / "critic_b1.npy")
|
||||
loadOrSkip(ac.critic.w2, dir / "critic_w2.npy")
|
||||
loadOrSkip(ac.critic.b2, dir / "critic_b2.npy")
|
||||
loadOrSkip(ac.critic.w3, dir / "critic_w3.npy")
|
||||
loadOrSkip(ac.critic.b3, dir / "critic_b3.npy")
|
||||
loadOrSkip(ac.logStd, dir / "log_std.npy")
|
||||
|
||||
proc saveWeightsAtomic*(ac: ActorCritic, targetDir: string) =
|
||||
## 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) =
|
||||
## 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) =
|
||||
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.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")
|
||||
@@ -185,7 +195,11 @@ proc loadBestAvailable*(ac: var ActorCritic, adam: var ACAdamStates,
|
||||
if adamStateFilesExist(candidate):
|
||||
adam.loadAdamStates(candidate)
|
||||
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)
|
||||
result = (loaded: false, roundNum: 0)
|
||||
|
||||
|
||||
@@ -1 +1 @@
|
||||
9500
|
||||
0
|
||||
|
||||
@@ -70,22 +70,25 @@ public class RunTraining {
|
||||
);
|
||||
|
||||
var owner = new Object();
|
||||
int[] prevTotal = { 0 };
|
||||
|
||||
try (var handle = runner.startBattleAsync(setup, bots)) {
|
||||
|
||||
handle.getOnRoundEnded().on(owner, event -> {
|
||||
int round = event.getRoundNumber();
|
||||
int ticks = event.getTurnNumber();
|
||||
int score = 0;
|
||||
int totalScore = 0;
|
||||
boolean win = false;
|
||||
boolean found = false;
|
||||
for (var r : event.getResults()) {
|
||||
if (r.getName().equals("PPO_Bot")) {
|
||||
found = true;
|
||||
score = r.getTotalScore();
|
||||
totalScore = r.getTotalScore();
|
||||
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
|
||||
// loop resumes from round_counter instead of grinding dummy rounds.
|
||||
// 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
|
||||
String line = String.format(
|
||||
"{\"type\":\"game\",\"round\":%d,\"ticks\":%d,\"score\":%d,\"win\":%b,\"opponent\":\"%s\"}",
|
||||
round, ticks, score, win, opponent
|
||||
"{\"type\":\"game\",\"round\":%d,\"ticks\":%d,\"score\":%d,\"total_score\":%d,\"win\":%b,\"opponent\":\"%s\"}",
|
||||
round, ticks, score, totalScore, win, opponent
|
||||
);
|
||||
try {
|
||||
appendLine(logFile, line);
|
||||
|
||||
Reference in New Issue
Block a user