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:
2026-08-20 14:23:08 +02:00
parent 75e32e3315
commit 6ad51148f4
11 changed files with 239 additions and 52 deletions
+68 -3
View File
@@ -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
+3
View File
@@ -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"
+7 -1
View File
@@ -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
View File
@@ -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:
+30 -4
View File
@@ -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)
+25 -2
View File
@@ -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"
+23
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -1 +1 @@
9500
0
+7 -4
View File
@@ -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);