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:
+21
-25
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user