diff --git a/PPO_Bot/PPO_Bot.nim b/PPO_Bot/PPO_Bot.nim index 081fc49..666eca7 100644 --- a/PPO_Bot/PPO_Bot.nim +++ b/PPO_Bot/PPO_Bot.nim @@ -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(), diff --git a/PPO_Bot/tests/test_weights.nim b/PPO_Bot/tests/test_weights.nim new file mode 100644 index 0000000..f65cc8a --- /dev/null +++ b/PPO_Bot/tests/test_weights.nim @@ -0,0 +1,101 @@ +## test_weights.nim — assert-based tests for weights.nim +## Run: nim c tests/test_weights.nim && ./tests/test_weights + +import std/[os, math] +import arraymancer +import "../network" +import "../weights" + +template check(cond: bool, msg: string) = + if not cond: + quit("FAIL: " & msg, 1) + +const tmpBase = "/tmp/test_weights_nim" + +# ── saveWeights / loadWeights roundtrip ─────────────────────────────────────── + +block testRoundtrip: + let dir = tmpBase & "_roundtrip" + removeDir(dir) + + let ac1 = initActorCritic() + saveWeights(ac1, dir) + + var ac2 = initActorCritic() + loadWeights(ac2, dir) + + # Verify a sample of tensors + template tensorEq(a, b: Tensor[float32]) = + check a.shape == b.shape, "shape mismatch" + let diff = abs(a - b) + var maxDiff = 0.0'f32 + for v in diff: maxDiff = max(maxDiff, v) + check maxDiff < 1e-6'f32, "tensor values differ by " & $maxDiff + + tensorEq(ac1.actor.w1, ac2.actor.w1) + tensorEq(ac1.actor.b1, ac2.actor.b1) + tensorEq(ac1.actor.w3, ac2.actor.w3) + tensorEq(ac1.critic.w1, ac2.critic.w1) + tensorEq(ac1.critic.b3, ac2.critic.b3) + tensorEq(ac1.logStd, ac2.logStd) + + removeDir(dir) + +# ── saveWeightsAtomic ───────────────────────────────────────────────────────── + +block testAtomic: + let dir = tmpBase & "_atomic" + removeDir(dir) + + let ac = initActorCritic() + saveWeightsAtomic(ac, dir) + + check dirExists(dir), "targetDir should exist after atomic save" + for f in ["actor_w1.npy", "critic_w1.npy", "log_std.npy"]: + check fileExists(dir / f), "missing file: " & f + + removeDir(dir) + +# ── loadBestAvailable ───────────────────────────────────────────────────────── + +block testLoadBest: + let root = tmpBase & "_loadbest" + removeDir(root) + createDir(root) + + let ac0 = initActorCritic() + + # Try with no weights — should return false + var acEmpty = initActorCritic() + check not loadBestAvailable(acEmpty, root), "should return false with no weights" + + # Save to latest/; should load + saveWeights(ac0, root / "latest") + var ac1 = initActorCritic() + check loadBestAvailable(ac1, root), "should load from latest/" + + # Remove latest/, save to checkpoint_1/ — should fall back + removeDir(root / "latest") + saveWeights(ac0, root / "checkpoint_1") + var ac2 = initActorCritic() + check loadBestAvailable(ac2, root), "should load from checkpoint_1/" + + removeDir(root) + +# ── cleanStaleTempDirs ──────────────────────────────────────────────────────── + +block testClean: + let root = tmpBase & "_clean" + removeDir(root) + createDir(root) + + let stale = root / "latest_tmp_12345" + createDir(stale) + check dirExists(stale), "stale dir should exist before clean" + + cleanStaleTempDirs(root) + check not dirExists(stale), "stale dir should be gone after clean" + + removeDir(root) + +echo "All tests passed" diff --git a/PPO_Bot/training.nim b/PPO_Bot/training.nim index 313f590..9e84bc6 100644 --- a/PPO_Bot/training.nim +++ b/PPO_Bot/training.nim @@ -149,9 +149,11 @@ proc initACAdamStates(ac: ActorCritic): ACAdamStates = result.cb3 = initAdamState(ac.critic.b3) result.logStd = initAdamState(ac.logStd) -# Persistent Adam state — survives across ppoUpdate calls (lives in training module) -var gAdamStates: ACAdamStates -var gAdamInit = false +# 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 ───────────────────────────────────────────────────────── @@ -172,7 +174,7 @@ proc ppoUpdate*(ac: var ActorCritic; entropyCoeff: float32 = 0.01'f32; valueLossCoeff: float32 = 0.5'f32; lr: float32 = 3e-4'f32; - maxGradNorm: float32 = 0.5'f32) = + maxGradNorm: float32 = 0.5'f32) {.gcsafe.} = if buffer.len == 0: return # Initialise Adam states once (persists across rounds) diff --git a/PPO_Bot/weights.nim b/PPO_Bot/weights.nim new file mode 100644 index 0000000..a2f1abd --- /dev/null +++ b/PPO_Bot/weights.nim @@ -0,0 +1,95 @@ +## weights.nim — save/load ActorCritic weights as .npy files. + +import std/[os, times, strutils] +import arraymancer +import ./network + +# ── Tensor names — order must match save/load ───────────────────────────────── + +const weightFiles = [ + "actor_w1.npy", "actor_b1.npy", "actor_w2.npy", "actor_b2.npy", + "actor_w3.npy", "actor_b3.npy", + "critic_w1.npy", "critic_b1.npy", "critic_w2.npy", "critic_b2.npy", + "critic_w3.npy", "critic_b3.npy", + "log_std.npy", +] + +proc saveWeights*(ac: ActorCritic, dir: string) = + ## Write all weight tensors to dir/ as .npy files. + createDir(dir) + ac.actor.w1.write_npy(dir / "actor_w1.npy") + ac.actor.b1.write_npy(dir / "actor_b1.npy") + ac.actor.w2.write_npy(dir / "actor_w2.npy") + ac.actor.b2.write_npy(dir / "actor_b2.npy") + ac.actor.w3.write_npy(dir / "actor_w3.npy") + ac.actor.b3.write_npy(dir / "actor_b3.npy") + ac.critic.w1.write_npy(dir / "critic_w1.npy") + ac.critic.b1.write_npy(dir / "critic_b1.npy") + ac.critic.w2.write_npy(dir / "critic_w2.npy") + ac.critic.b2.write_npy(dir / "critic_b2.npy") + ac.critic.w3.write_npy(dir / "critic_w3.npy") + ac.critic.b3.write_npy(dir / "critic_b3.npy") + ac.logStd.write_npy(dir / "log_std.npy") + +proc loadWeights*(ac: var ActorCritic, dir: string) = + ## Load all weight tensors from dir/. Asserts shapes match. + template loadAndCheck(dest: untyped, path: string) = + let loaded = read_npy[float32](path) + doAssert loaded.shape == dest.shape, + "Shape mismatch loading " & path & ": got " & $loaded.shape & " want " & $dest.shape + dest = loaded + + loadAndCheck(ac.actor.w1, dir / "actor_w1.npy") + loadAndCheck(ac.actor.b1, dir / "actor_b1.npy") + loadAndCheck(ac.actor.w2, dir / "actor_w2.npy") + loadAndCheck(ac.actor.b2, dir / "actor_b2.npy") + loadAndCheck(ac.actor.w3, dir / "actor_w3.npy") + loadAndCheck(ac.actor.b3, dir / "actor_b3.npy") + loadAndCheck(ac.critic.w1, dir / "critic_w1.npy") + loadAndCheck(ac.critic.b1, dir / "critic_b1.npy") + loadAndCheck(ac.critic.w2, dir / "critic_w2.npy") + loadAndCheck(ac.critic.b2, dir / "critic_b2.npy") + loadAndCheck(ac.critic.w3, dir / "critic_w3.npy") + loadAndCheck(ac.critic.b3, dir / "critic_b3.npy") + loadAndCheck(ac.logStd, dir / "log_std.npy") + +proc saveWeightsAtomic*(ac: ActorCritic, targetDir: string) = + ## Write to a temp dir, then rename atomically over targetDir. + let tmpDir = targetDir & "_tmp_" & $int(epochTime()) + saveWeights(ac, tmpDir) + if dirExists(targetDir): + removeDir(targetDir) + moveDir(tmpDir, targetDir) + +proc saveCheckpoint*(ac: ActorCritic, weightsRoot: string, roundNum: int) = + ## Always saves to weightsRoot/latest/. + ## Every 50 rounds also saves to checkpoint_{1,2,3} in round-robin. + saveWeightsAtomic(ac, weightsRoot / "latest") + if roundNum mod 50 == 0: + let slot = ((roundNum div 50 - 1) mod 3) + 1 # 50→1, 100→2, 150→3, 200→1, … + saveWeightsAtomic(ac, weightsRoot / ("checkpoint_" & $slot)) + +proc loadBestAvailable*(ac: var ActorCritic, weightsRoot: string): bool = + ## Try latest/, then checkpoint_3/, checkpoint_2/, checkpoint_1/. + ## Returns true if weights loaded, false if all fail (random init stays). + for candidate in [weightsRoot / "latest", + weightsRoot / "checkpoint_3", + weightsRoot / "checkpoint_2", + weightsRoot / "checkpoint_1"]: + if dirExists(candidate): + var ok = true + for f in weightFiles: + if not fileExists(candidate / f): + ok = false + break + if ok: + ac.loadWeights(candidate) + return true + result = false + +proc cleanStaleTempDirs*(weightsRoot: string) = + ## Delete any dirs inside weightsRoot whose name contains "_tmp_". + if not dirExists(weightsRoot): return + for kind, path in walkDir(weightsRoot): + if kind == pcDir and "_tmp_" in lastPathPart(path): + removeDir(path)