fix(PPO_Bot): radar oscillation, Adam persistence, checkpoint order, channel race

- enemy_tracker: toggle lastOvershootDir each tick; make getRadarTurnRate take var tracker
- training: remove threadvar Adam globals; pass adamStates as var param to ppoUpdate; export ACAdamStates
- PPO_Bot: carry ACAdamStates through TrainingArgs/TrainingResult; drop trainingDone bool and Lock — use resultChan.tryRecv() directly as synchronisation
- weights: sort checkpoint dirs newest-first by mtime instead of hardcoded order
- tests/test_training: pass explicit ACAdamStates to ppoUpdate

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-08-16 15:42:06 +02:00
parent 473d67f644
commit aea0724d3a
5 changed files with 74 additions and 62 deletions
+34 -29
View File
@@ -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, locks] import std/os
import arraymancer import arraymancer
import tankroyale_botapi import tankroyale_botapi
import network import network
@@ -26,30 +26,35 @@ type PPOBot = ref object of Bot
hasLastTrans: bool hasLastTrans: bool
var ac = initActorCritic() var ac = initActorCritic()
var gAdamStates: ACAdamStates # persists across rounds
# ── Background training state ───────────────────────────────────────────────── # ── Background training state ─────────────────────────────────────────────────
type TrainingArgs = object type
ac: ActorCritic TrainingResult = object
buffer: TrajectoryBuffer ac: ActorCritic
lastValue: float32 adamStates: ACAdamStates
roundNum: int
weightsRoot: string TrainingArgs = object
ac: ActorCritic
adamStates: ACAdamStates
buffer: TrajectoryBuffer
lastValue: float32
roundNum: int
weightsRoot: string
var var
trainingThread: Thread[TrainingArgs] trainingThread: Thread[TrainingArgs]
trainingDone: bool = true # true = no training running / last run finished resultChan: Channel[TrainingResult]
trainingLock: Lock threadLaunched: bool = false # true while training thread is running
resultChan: Channel[ActorCritic] roundCounter: int = 0
roundCounter: int = 0
proc trainingThreadProc(args: TrainingArgs) {.thread.} = proc trainingThreadProc(args: TrainingArgs) {.thread.} =
var localAc = args.ac var localAc = args.ac
ppoUpdate(localAc, args.buffer, lastValue = args.lastValue) var localAdam = args.adamStates
ppoUpdate(localAc, args.buffer, lastValue = args.lastValue, adamStates = localAdam)
saveCheckpoint(localAc, args.weightsRoot, args.roundNum) saveCheckpoint(localAc, args.weightsRoot, args.roundNum)
resultChan.send(localAc) resultChan.send(TrainingResult(ac: localAc, adamStates: localAdam))
withLock(trainingLock):
trainingDone = true
# ── Bot methods ─────────────────────────────────────────────────────────────── # ── Bot methods ───────────────────────────────────────────────────────────────
@@ -71,12 +76,13 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
if bot.hasLastTrans and bot.buffer.len > 0: if bot.hasLastTrans and bot.buffer.len > 0:
bot.buffer.transitions[^1].reward += roundReward bot.buffer.transitions[^1].reward += roundReward
# Pick up result from previous training thread if done # Pick up result from previous training thread if available; channel IS the sync
withLock(trainingLock): if threadLaunched:
if trainingDone and roundCounter > 1: let (avail, trained) = resultChan.tryRecv()
let (avail, trained) = resultChan.tryRecv() if avail:
if avail: ac = trained.ac
ac = trained gAdamStates = trained.adamStates
threadLaunched = false
if bot.buffer.len == 0: if bot.buffer.len == 0:
bot.hasLastTrans = false bot.hasLastTrans = false
@@ -84,15 +90,14 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
# If last thread still running, drop this pass — start fresh with newer data # If last thread still running, drop this pass — start fresh with newer data
# ponytail: simple drop; queue if every round must train # ponytail: simple drop; queue if every round must train
withLock(trainingLock): if threadLaunched:
if not trainingDone: bot.buffer.clear()
bot.buffer.clear() bot.hasLastTrans = false
bot.hasLastTrans = false return
return
trainingDone = false
let args = TrainingArgs( let args = TrainingArgs(
ac: ac, ac: ac,
adamStates: gAdamStates,
buffer: bot.buffer, buffer: bot.buffer,
lastValue: 0.0'f32, lastValue: 0.0'f32,
roundNum: roundCounter, roundNum: roundCounter,
@@ -102,6 +107,7 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
bot.hasLastTrans = false bot.hasLastTrans = false
createThread(trainingThread, trainingThreadProc, args) createThread(trainingThread, trainingThreadProc, args)
threadLaunched = true
method run(bot: PPOBot) = method run(bot: PPOBot) =
# Seed energy on first tick # Seed energy on first tick
@@ -166,7 +172,6 @@ method run(bot: PPOBot) =
go() go()
when isMainModule: when isMainModule:
initLock(trainingLock)
resultChan.open() resultChan.open()
createDir(weightsRoot) createDir(weightsRoot)
cleanStaleTempDirs(weightsRoot) cleanStaleTempDirs(weightsRoot)
+2 -1
View File
@@ -69,7 +69,7 @@ proc normalizeRelative(angle: float64): float64 {.inline.} =
if result >= 180.0: result -= 360.0 if result >= 180.0: result -= 360.0
elif result < -180.0: result += 360.0 elif result < -180.0: result += 360.0
proc getRadarTurnRate*(tracker: EnemyTracker; proc getRadarTurnRate*(tracker: var EnemyTracker;
botX, botY, botDirection, radarDirection: float64): float64 = botX, botY, botDirection, radarDirection: float64): float64 =
## Returns radar turn rate (degrees/tick, positive = right). ## Returns radar turn rate (degrees/tick, positive = right).
## Before contact: full 45° sweep. ## Before contact: full 45° sweep.
@@ -92,3 +92,4 @@ proc getRadarTurnRate*(tracker: EnemyTracker;
let overshoot = 10.0 let overshoot = 10.0
let target = radarBearing + tracker.lastOvershootDir * overshoot let target = radarBearing + tracker.lastOvershootDir * overshoot
result = target.clamp(-45.0, 45.0) result = target.clamp(-45.0, 45.0)
tracker.lastOvershootDir *= -1.0
+2 -1
View File
@@ -86,7 +86,8 @@ block testPpoUpdate:
let v = ac.criticForward(s) let v = ac.criticForward(s)
buf.add(Transition(state: s, action: a, logProb: lp, reward: 0.1'f32, value: v)) buf.add(Transition(state: s, action: a, logProb: lp, reward: 0.1'f32, value: v))
ppoUpdate(ac, buf, lastValue = 0.0'f32, epochs = 2, miniBatchSize = 5) var adam: ACAdamStates
ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam, epochs = 2, miniBatchSize = 5)
# Weights should have changed — compare flattened # Weights should have changed — compare flattened
let n = ac.actor.w1.shape[0] * ac.actor.w1.shape[1] let n = ac.actor.w1.shape[0] * ac.actor.w1.shape[1]
+21 -25
View File
@@ -128,13 +128,14 @@ proc mlpBackward(mlp: MLP; fwd: MLPFwd; x: Tensor[float32];
# ── Adam states for ActorCritic parameters ─────────────────────────────────── # ── Adam states for ActorCritic parameters ───────────────────────────────────
type ACAdamStates = object type ACAdamStates* = object
## One AdamState per learnable tensor in ActorCritic. ## One AdamState per learnable tensor in ActorCritic.
aw1, ab1, aw2, ab2, aw3, ab3: AdamState # actor MLP aw1, ab1, aw2, ab2, aw3, ab3: AdamState # actor MLP
cw1, cb1, cw2, cb2, cw3, cb3: AdamState # critic MLP cw1, cb1, cw2, cb2, cw3, cb3: AdamState # critic MLP
logStd: AdamState logStd: AdamState
initialized*: bool
proc initACAdamStates(ac: ActorCritic): ACAdamStates = proc initACAdamStates*(ac: ActorCritic): ACAdamStates =
result.aw1 = initAdamState(ac.actor.w1) result.aw1 = initAdamState(ac.actor.w1)
result.ab1 = initAdamState(ac.actor.b1) result.ab1 = initAdamState(ac.actor.b1)
result.aw2 = initAdamState(ac.actor.w2) result.aw2 = initAdamState(ac.actor.w2)
@@ -148,12 +149,7 @@ proc initACAdamStates(ac: ActorCritic): ACAdamStates =
result.cw3 = initAdamState(ac.critic.w3) result.cw3 = initAdamState(ac.critic.w3)
result.cb3 = initAdamState(ac.critic.b3) result.cb3 = initAdamState(ac.critic.b3)
result.logStd = initAdamState(ac.logStd) result.logStd = initAdamState(ac.logStd)
result.initialized = true
# Persistent Adam state — per-thread so training thread can call ppoUpdate safely.
# ponytail: threadvar resets Adam each new training thread; persist across calls
# within one thread. If Adam across rounds matters, embed state in TrainingArgs.
var gAdamStates {.threadvar.}: ACAdamStates
var gAdamInit {.threadvar.}: bool
# ── Gradient clipping ───────────────────────────────────────────────────────── # ── Gradient clipping ─────────────────────────────────────────────────────────
@@ -168,6 +164,7 @@ proc globalNorm(grads: varargs[Tensor[float32]]): float32 =
proc ppoUpdate*(ac: var ActorCritic; proc ppoUpdate*(ac: var ActorCritic;
buffer: TrajectoryBuffer; buffer: TrajectoryBuffer;
lastValue: float32; lastValue: float32;
adamStates: var ACAdamStates;
epochs: int = 4; epochs: int = 4;
miniBatchSize: int = 64; miniBatchSize: int = 64;
clipEpsilon: float32 = 0.2'f32; clipEpsilon: float32 = 0.2'f32;
@@ -177,10 +174,9 @@ proc ppoUpdate*(ac: var ActorCritic;
maxGradNorm: float32 = 0.5'f32) {.gcsafe.} = maxGradNorm: float32 = 0.5'f32) {.gcsafe.} =
if buffer.len == 0: return if buffer.len == 0: return
# Initialise Adam states once (persists across rounds) # Initialise Adam states once; caller persists them across rounds
if not gAdamInit: if not adamStates.initialized:
gAdamStates = initACAdamStates(ac) adamStates = initACAdamStates(ac)
gAdamInit = true
# 1. GAE # 1. GAE
let rewards = buffer.transitions.mapIt(it.reward) let rewards = buffer.transitions.mapIt(it.reward)
@@ -332,18 +328,18 @@ proc ppoUpdate*(ac: var ActorCritic;
dCriticW3 = allGrads[11]; dCriticB3 = allGrads[12] dCriticW3 = allGrads[11]; dCriticB3 = allGrads[12]
# ── Adam updates ── # ── Adam updates ──
adamStep(ac.actor.w1, dActorW1, gAdamStates.aw1, lr) adamStep(ac.actor.w1, dActorW1, adamStates.aw1, lr)
adamStep(ac.actor.b1, dActorB1, gAdamStates.ab1, lr) adamStep(ac.actor.b1, dActorB1, adamStates.ab1, lr)
adamStep(ac.actor.w2, dActorW2, gAdamStates.aw2, lr) adamStep(ac.actor.w2, dActorW2, adamStates.aw2, lr)
adamStep(ac.actor.b2, dActorB2, gAdamStates.ab2, lr) adamStep(ac.actor.b2, dActorB2, adamStates.ab2, lr)
adamStep(ac.actor.w3, dActorW3, gAdamStates.aw3, lr) adamStep(ac.actor.w3, dActorW3, adamStates.aw3, lr)
adamStep(ac.actor.b3, dActorB3, gAdamStates.ab3, lr) adamStep(ac.actor.b3, dActorB3, adamStates.ab3, lr)
adamStep(ac.logStd, dLogStd, gAdamStates.logStd, lr) adamStep(ac.logStd, dLogStd, adamStates.logStd, lr)
adamStep(ac.critic.w1, dCriticW1, gAdamStates.cw1, lr) adamStep(ac.critic.w1, dCriticW1, adamStates.cw1, lr)
adamStep(ac.critic.b1, dCriticB1, gAdamStates.cb1, lr) adamStep(ac.critic.b1, dCriticB1, adamStates.cb1, lr)
adamStep(ac.critic.w2, dCriticW2, gAdamStates.cw2, lr) adamStep(ac.critic.w2, dCriticW2, adamStates.cw2, lr)
adamStep(ac.critic.b2, dCriticB2, gAdamStates.cb2, lr) adamStep(ac.critic.b2, dCriticB2, adamStates.cb2, lr)
adamStep(ac.critic.w3, dCriticW3, gAdamStates.cw3, lr) adamStep(ac.critic.w3, dCriticW3, adamStates.cw3, lr)
adamStep(ac.critic.b3, dCriticB3, gAdamStates.cb3, lr) adamStep(ac.critic.b3, dCriticB3, adamStates.cb3, lr)
mbStart = mbEnd mbStart = mbEnd
+15 -6
View File
@@ -1,6 +1,6 @@
## weights.nim — save/load ActorCritic weights as .npy files. ## weights.nim — save/load ActorCritic weights as .npy files.
import std/[os, times, strutils] import std/[os, times, strutils, algorithm, sequtils]
import arraymancer import arraymancer
import ./network import ./network
@@ -70,12 +70,21 @@ proc saveCheckpoint*(ac: ActorCritic, weightsRoot: string, roundNum: int) =
saveWeightsAtomic(ac, weightsRoot / ("checkpoint_" & $slot)) saveWeightsAtomic(ac, weightsRoot / ("checkpoint_" & $slot))
proc loadBestAvailable*(ac: var ActorCritic, weightsRoot: string): bool = proc loadBestAvailable*(ac: var ActorCritic, weightsRoot: string): bool =
## Try latest/, then checkpoint_3/, checkpoint_2/, checkpoint_1/. ## Try latest/ first, then checkpoints sorted newest-first by mtime.
## Returns true if weights loaded, false if all fail (random init stays). ## Returns true if weights loaded, false if all fail (random init stays).
for candidate in [weightsRoot / "latest", let checkpoints = [weightsRoot / "checkpoint_1",
weightsRoot / "checkpoint_3", weightsRoot / "checkpoint_2",
weightsRoot / "checkpoint_2", weightsRoot / "checkpoint_3"]
weightsRoot / "checkpoint_1"]: # Sort checkpoints newest-first by modification time
var existing: seq[tuple[mtime: Time, path: string]]
for p in checkpoints:
if dirExists(p):
existing.add((getLastModificationTime(p), p))
existing.sort(proc(a, b: tuple[mtime: Time, path: string]): int =
cmp(b.mtime, a.mtime)) # descending
let candidates = @[weightsRoot / "latest"] & existing.mapIt(it.path)
for candidate in candidates:
if dirExists(candidate): if dirExists(candidate):
var ok = true var ok = true
for f in weightFiles: for f in weightFiles: