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
+21 -25
View File
@@ -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