From aea0724d3a815cccbfa425f4f4e46acdc2f40f95 Mon Sep 17 00:00:00 2001 From: Davide Cappellini Date: Sun, 16 Aug 2026 15:42:06 +0200 Subject: [PATCH] fix(PPO_Bot): radar oscillation, Adam persistence, checkpoint order, channel race MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 --- PPO_Bot/PPO_Bot.nim | 63 ++++++++++++++++++--------------- PPO_Bot/enemy_tracker.nim | 3 +- PPO_Bot/tests/test_training.nim | 3 +- PPO_Bot/training.nim | 46 +++++++++++------------- PPO_Bot/weights.nim | 21 +++++++---- 5 files changed, 74 insertions(+), 62 deletions(-) diff --git a/PPO_Bot/PPO_Bot.nim b/PPO_Bot/PPO_Bot.nim index 666eca7..8bd2634 100644 --- a/PPO_Bot/PPO_Bot.nim +++ b/PPO_Bot/PPO_Bot.nim @@ -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, locks] +import std/os import arraymancer import tankroyale_botapi import network @@ -26,30 +26,35 @@ type PPOBot = ref object of Bot hasLastTrans: bool var ac = initActorCritic() +var gAdamStates: ACAdamStates # persists across rounds # ── Background training state ───────────────────────────────────────────────── -type TrainingArgs = object - ac: ActorCritic - buffer: TrajectoryBuffer - lastValue: float32 - roundNum: int - weightsRoot: string +type + TrainingResult = object + ac: ActorCritic + adamStates: ACAdamStates + + TrainingArgs = object + ac: ActorCritic + adamStates: ACAdamStates + buffer: TrajectoryBuffer + lastValue: float32 + roundNum: int + weightsRoot: string var - trainingThread: Thread[TrainingArgs] - trainingDone: bool = true # true = no training running / last run finished - trainingLock: Lock - resultChan: Channel[ActorCritic] - roundCounter: int = 0 + trainingThread: Thread[TrainingArgs] + resultChan: Channel[TrainingResult] + threadLaunched: bool = false # true while training thread is running + roundCounter: int = 0 proc trainingThreadProc(args: TrainingArgs) {.thread.} = 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) - resultChan.send(localAc) - withLock(trainingLock): - trainingDone = true + resultChan.send(TrainingResult(ac: localAc, adamStates: localAdam)) # ── Bot methods ─────────────────────────────────────────────────────────────── @@ -71,12 +76,13 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) = if bot.hasLastTrans and bot.buffer.len > 0: bot.buffer.transitions[^1].reward += roundReward - # Pick up result from previous training thread if done - withLock(trainingLock): - if trainingDone and roundCounter > 1: - let (avail, trained) = resultChan.tryRecv() - if avail: - ac = trained + # Pick up result from previous training thread if available; channel IS the sync + if threadLaunched: + let (avail, trained) = resultChan.tryRecv() + if avail: + ac = trained.ac + gAdamStates = trained.adamStates + threadLaunched = false if bot.buffer.len == 0: 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 # ponytail: simple drop; queue if every round must train - withLock(trainingLock): - if not trainingDone: - bot.buffer.clear() - bot.hasLastTrans = false - return - trainingDone = false + if threadLaunched: + bot.buffer.clear() + bot.hasLastTrans = false + return let args = TrainingArgs( ac: ac, + adamStates: gAdamStates, buffer: bot.buffer, lastValue: 0.0'f32, roundNum: roundCounter, @@ -102,6 +107,7 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) = bot.hasLastTrans = false createThread(trainingThread, trainingThreadProc, args) + threadLaunched = true method run(bot: PPOBot) = # Seed energy on first tick @@ -166,7 +172,6 @@ method run(bot: PPOBot) = go() when isMainModule: - initLock(trainingLock) resultChan.open() createDir(weightsRoot) cleanStaleTempDirs(weightsRoot) diff --git a/PPO_Bot/enemy_tracker.nim b/PPO_Bot/enemy_tracker.nim index 0a83722..abf3fe2 100644 --- a/PPO_Bot/enemy_tracker.nim +++ b/PPO_Bot/enemy_tracker.nim @@ -69,7 +69,7 @@ proc normalizeRelative(angle: float64): float64 {.inline.} = if 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 = ## Returns radar turn rate (degrees/tick, positive = right). ## Before contact: full 45° sweep. @@ -92,3 +92,4 @@ proc getRadarTurnRate*(tracker: EnemyTracker; let overshoot = 10.0 let target = radarBearing + tracker.lastOvershootDir * overshoot result = target.clamp(-45.0, 45.0) + tracker.lastOvershootDir *= -1.0 diff --git a/PPO_Bot/tests/test_training.nim b/PPO_Bot/tests/test_training.nim index 90db3fd..e96dce3 100644 --- a/PPO_Bot/tests/test_training.nim +++ b/PPO_Bot/tests/test_training.nim @@ -86,7 +86,8 @@ block testPpoUpdate: let v = ac.criticForward(s) 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 let n = ac.actor.w1.shape[0] * ac.actor.w1.shape[1] diff --git a/PPO_Bot/training.nim b/PPO_Bot/training.nim index 9e84bc6..bf36684 100644 --- a/PPO_Bot/training.nim +++ b/PPO_Bot/training.nim @@ -128,13 +128,14 @@ proc mlpBackward(mlp: MLP; fwd: MLPFwd; x: Tensor[float32]; # ── Adam states for ActorCritic parameters ─────────────────────────────────── -type ACAdamStates = object +type ACAdamStates* = object ## One AdamState per learnable tensor in ActorCritic. aw1, ab1, aw2, ab2, aw3, ab3: AdamState # actor MLP cw1, cb1, cw2, cb2, cw3, cb3: AdamState # critic MLP logStd: AdamState + initialized*: bool -proc initACAdamStates(ac: ActorCritic): ACAdamStates = +proc initACAdamStates*(ac: ActorCritic): ACAdamStates = result.aw1 = initAdamState(ac.actor.w1) result.ab1 = initAdamState(ac.actor.b1) result.aw2 = initAdamState(ac.actor.w2) @@ -148,12 +149,7 @@ proc initACAdamStates(ac: ActorCritic): ACAdamStates = result.cw3 = initAdamState(ac.critic.w3) result.cb3 = initAdamState(ac.critic.b3) result.logStd = initAdamState(ac.logStd) - -# 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 + result.initialized = true # ── Gradient clipping ───────────────────────────────────────────────────────── @@ -168,6 +164,7 @@ proc globalNorm(grads: varargs[Tensor[float32]]): float32 = proc ppoUpdate*(ac: var ActorCritic; buffer: TrajectoryBuffer; lastValue: float32; + adamStates: var ACAdamStates; epochs: int = 4; miniBatchSize: int = 64; clipEpsilon: float32 = 0.2'f32; @@ -177,10 +174,9 @@ proc ppoUpdate*(ac: var ActorCritic; maxGradNorm: float32 = 0.5'f32) {.gcsafe.} = if buffer.len == 0: return - # Initialise Adam states once (persists across rounds) - if not gAdamInit: - gAdamStates = initACAdamStates(ac) - gAdamInit = true + # Initialise Adam states once; caller persists them across rounds + if not adamStates.initialized: + adamStates = initACAdamStates(ac) # 1. GAE let rewards = buffer.transitions.mapIt(it.reward) @@ -332,18 +328,18 @@ proc ppoUpdate*(ac: var ActorCritic; dCriticW3 = allGrads[11]; dCriticB3 = allGrads[12] # ── Adam updates ── - adamStep(ac.actor.w1, dActorW1, gAdamStates.aw1, lr) - adamStep(ac.actor.b1, dActorB1, gAdamStates.ab1, lr) - adamStep(ac.actor.w2, dActorW2, gAdamStates.aw2, lr) - adamStep(ac.actor.b2, dActorB2, gAdamStates.ab2, lr) - adamStep(ac.actor.w3, dActorW3, gAdamStates.aw3, lr) - adamStep(ac.actor.b3, dActorB3, gAdamStates.ab3, lr) - adamStep(ac.logStd, dLogStd, gAdamStates.logStd, lr) - adamStep(ac.critic.w1, dCriticW1, gAdamStates.cw1, lr) - adamStep(ac.critic.b1, dCriticB1, gAdamStates.cb1, lr) - adamStep(ac.critic.w2, dCriticW2, gAdamStates.cw2, lr) - adamStep(ac.critic.b2, dCriticB2, gAdamStates.cb2, lr) - adamStep(ac.critic.w3, dCriticW3, gAdamStates.cw3, lr) - adamStep(ac.critic.b3, dCriticB3, gAdamStates.cb3, lr) + adamStep(ac.actor.w1, dActorW1, adamStates.aw1, lr) + adamStep(ac.actor.b1, dActorB1, adamStates.ab1, lr) + adamStep(ac.actor.w2, dActorW2, adamStates.aw2, lr) + adamStep(ac.actor.b2, dActorB2, adamStates.ab2, lr) + adamStep(ac.actor.w3, dActorW3, adamStates.aw3, lr) + adamStep(ac.actor.b3, dActorB3, adamStates.ab3, lr) + adamStep(ac.logStd, dLogStd, adamStates.logStd, lr) + 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 diff --git a/PPO_Bot/weights.nim b/PPO_Bot/weights.nim index a2f1abd..9929b34 100644 --- a/PPO_Bot/weights.nim +++ b/PPO_Bot/weights.nim @@ -1,6 +1,6 @@ ## weights.nim — save/load ActorCritic weights as .npy files. -import std/[os, times, strutils] +import std/[os, times, strutils, algorithm, sequtils] import arraymancer import ./network @@ -70,12 +70,21 @@ proc saveCheckpoint*(ac: ActorCritic, weightsRoot: string, roundNum: int) = saveWeightsAtomic(ac, weightsRoot / ("checkpoint_" & $slot)) 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). - for candidate in [weightsRoot / "latest", - weightsRoot / "checkpoint_3", - weightsRoot / "checkpoint_2", - weightsRoot / "checkpoint_1"]: + let checkpoints = [weightsRoot / "checkpoint_1", + weightsRoot / "checkpoint_2", + weightsRoot / "checkpoint_3"] + # 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): var ok = true for f in weightFiles: