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:
+34
-29
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user