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.
|
||||
## 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)
|
||||
|
||||
Reference in New Issue
Block a user