feat(PPO_Bot): weight persistence + background training (#17)

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-08-16 15:35:42 +02:00
parent eadd177d3b
commit 473d67f644
4 changed files with 269 additions and 9 deletions
+67 -5
View File
@@ -1,16 +1,18 @@
## PPO_Bot — enemy tracker + state vector wired into the game loop.
## Training: trajectory collected per tick, PPO update on round end.
## Training: trajectory collected per tick, PPO update in background thread.
import std/os
import std/[os, locks]
import arraymancer
import tankroyale_botapi
import network
import actions
import training
import weights
import ./enemy_tracker
import ./state_vector
const botJsonPath = currentSourcePath().parentDir / "PPO_Bot.json"
const weightsRoot = currentSourcePath().parentDir / "weights"
type PPOBot = ref object of Bot
tracker: EnemyTracker
@@ -25,6 +27,32 @@ type PPOBot = ref object of Bot
var ac = initActorCritic()
# ── Background training state ─────────────────────────────────────────────────
type TrainingArgs = object
ac: ActorCritic
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
proc trainingThreadProc(args: TrainingArgs) {.thread.} =
var localAc = args.ac
ppoUpdate(localAc, args.buffer, lastValue = args.lastValue)
saveCheckpoint(localAc, args.weightsRoot, args.roundNum)
resultChan.send(localAc)
withLock(trainingLock):
trainingDone = true
# ── Bot methods ───────────────────────────────────────────────────────────────
method onScannedBot*(bot: PPOBot, e: ScannedBotEvent) =
bot.tracker.update(e.x, e.y, e.direction, e.speed, e.energy)
@@ -36,17 +64,45 @@ method onRoundStarted*(bot: PPOBot, e: RoundStartedEvent) =
bot.hasLastTrans = false
method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
inc roundCounter
# Add round-end score bonus to last transition (if any)
let roundReward = computeRoundReward(e.results.totalScore.float32)
if bot.hasLastTrans and bot.buffer.len > 0:
bot.buffer.transitions[^1].reward += roundReward
# PPO update synchronous; background thread is issue #17
# ponytail: blocking update per round; move to thread pool when #17 lands
ppoUpdate(ac, bot.buffer, lastValue = 0.0'f32)
# 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
if bot.buffer.len == 0:
bot.hasLastTrans = false
return
# 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
let args = TrainingArgs(
ac: ac,
buffer: bot.buffer,
lastValue: 0.0'f32,
roundNum: roundCounter,
weightsRoot: weightsRoot,
)
bot.buffer.clear()
bot.hasLastTrans = false
createThread(trainingThread, trainingThreadProc, args)
method run(bot: PPOBot) =
# Seed energy on first tick
bot.prevEnergy = getEnergy().float32
@@ -110,6 +166,12 @@ method run(bot: PPOBot) =
go()
when isMainModule:
initLock(trainingLock)
resultChan.open()
createDir(weightsRoot)
cleanStaleTempDirs(weightsRoot)
discard loadBestAvailable(ac, weightsRoot)
var bot = PPOBot(
tracker: initEnemyTracker(),
buffer: initTrajectoryBuffer(),