feat(PPO_Bot): weight persistence + background training (#17)
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
+67
-5
@@ -1,16 +1,18 @@
|
|||||||
## 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 on round end.
|
## Training: trajectory collected per tick, PPO update in background thread.
|
||||||
|
|
||||||
import std/os
|
import std/[os, locks]
|
||||||
import arraymancer
|
import arraymancer
|
||||||
import tankroyale_botapi
|
import tankroyale_botapi
|
||||||
import network
|
import network
|
||||||
import actions
|
import actions
|
||||||
import training
|
import training
|
||||||
|
import weights
|
||||||
import ./enemy_tracker
|
import ./enemy_tracker
|
||||||
import ./state_vector
|
import ./state_vector
|
||||||
|
|
||||||
const botJsonPath = currentSourcePath().parentDir / "PPO_Bot.json"
|
const botJsonPath = currentSourcePath().parentDir / "PPO_Bot.json"
|
||||||
|
const weightsRoot = currentSourcePath().parentDir / "weights"
|
||||||
|
|
||||||
type PPOBot = ref object of Bot
|
type PPOBot = ref object of Bot
|
||||||
tracker: EnemyTracker
|
tracker: EnemyTracker
|
||||||
@@ -25,6 +27,32 @@ type PPOBot = ref object of Bot
|
|||||||
|
|
||||||
var ac = initActorCritic()
|
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) =
|
method onScannedBot*(bot: PPOBot, e: ScannedBotEvent) =
|
||||||
bot.tracker.update(e.x, e.y, e.direction, e.speed, e.energy)
|
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
|
bot.hasLastTrans = false
|
||||||
|
|
||||||
method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
||||||
|
inc roundCounter
|
||||||
|
|
||||||
# Add round-end score bonus to last transition (if any)
|
# Add round-end score bonus to last transition (if any)
|
||||||
let roundReward = computeRoundReward(e.results.totalScore.float32)
|
let roundReward = computeRoundReward(e.results.totalScore.float32)
|
||||||
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
|
||||||
|
|
||||||
# PPO update synchronous; background thread is issue #17
|
# Pick up result from previous training thread if done
|
||||||
# ponytail: blocking update per round; move to thread pool when #17 lands
|
withLock(trainingLock):
|
||||||
ppoUpdate(ac, bot.buffer, lastValue = 0.0'f32)
|
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.buffer.clear()
|
||||||
bot.hasLastTrans = false
|
bot.hasLastTrans = false
|
||||||
|
|
||||||
|
createThread(trainingThread, trainingThreadProc, args)
|
||||||
|
|
||||||
method run(bot: PPOBot) =
|
method run(bot: PPOBot) =
|
||||||
# Seed energy on first tick
|
# Seed energy on first tick
|
||||||
bot.prevEnergy = getEnergy().float32
|
bot.prevEnergy = getEnergy().float32
|
||||||
@@ -110,6 +166,12 @@ method run(bot: PPOBot) =
|
|||||||
go()
|
go()
|
||||||
|
|
||||||
when isMainModule:
|
when isMainModule:
|
||||||
|
initLock(trainingLock)
|
||||||
|
resultChan.open()
|
||||||
|
createDir(weightsRoot)
|
||||||
|
cleanStaleTempDirs(weightsRoot)
|
||||||
|
discard loadBestAvailable(ac, weightsRoot)
|
||||||
|
|
||||||
var bot = PPOBot(
|
var bot = PPOBot(
|
||||||
tracker: initEnemyTracker(),
|
tracker: initEnemyTracker(),
|
||||||
buffer: initTrajectoryBuffer(),
|
buffer: initTrajectoryBuffer(),
|
||||||
|
|||||||
@@ -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"
|
||||||
@@ -149,9 +149,11 @@ proc initACAdamStates(ac: ActorCritic): ACAdamStates =
|
|||||||
result.cb3 = initAdamState(ac.critic.b3)
|
result.cb3 = initAdamState(ac.critic.b3)
|
||||||
result.logStd = initAdamState(ac.logStd)
|
result.logStd = initAdamState(ac.logStd)
|
||||||
|
|
||||||
# Persistent Adam state — survives across ppoUpdate calls (lives in training module)
|
# Persistent Adam state — per-thread so training thread can call ppoUpdate safely.
|
||||||
var gAdamStates: ACAdamStates
|
# ponytail: threadvar resets Adam each new training thread; persist across calls
|
||||||
var gAdamInit = false
|
# within one thread. If Adam across rounds matters, embed state in TrainingArgs.
|
||||||
|
var gAdamStates {.threadvar.}: ACAdamStates
|
||||||
|
var gAdamInit {.threadvar.}: bool
|
||||||
|
|
||||||
# ── Gradient clipping ─────────────────────────────────────────────────────────
|
# ── Gradient clipping ─────────────────────────────────────────────────────────
|
||||||
|
|
||||||
@@ -172,7 +174,7 @@ proc ppoUpdate*(ac: var ActorCritic;
|
|||||||
entropyCoeff: float32 = 0.01'f32;
|
entropyCoeff: float32 = 0.01'f32;
|
||||||
valueLossCoeff: float32 = 0.5'f32;
|
valueLossCoeff: float32 = 0.5'f32;
|
||||||
lr: float32 = 3e-4'f32;
|
lr: float32 = 3e-4'f32;
|
||||||
maxGradNorm: float32 = 0.5'f32) =
|
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 (persists across rounds)
|
||||||
|
|||||||
@@ -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)
|
||||||
Reference in New Issue
Block a user