diff --git a/PPO_Bot/PPO_Bot.nim b/PPO_Bot/PPO_Bot.nim index 7f5aed8..97c4d66 100644 --- a/PPO_Bot/PPO_Bot.nim +++ b/PPO_Bot/PPO_Bot.nim @@ -57,7 +57,7 @@ proc trainingThreadProc(args: TrainingArgs) {.thread.} = var localAc = args.ac var localAdam = args.adamStates let m = ppoUpdate(localAc, args.buffer, lastValue = args.lastValue, adamStates = localAdam) - saveCheckpoint(localAc, args.weightsRoot, args.roundNum) + saveCheckpoint(localAc, localAdam, args.weightsRoot, args.roundNum) resultChan.send(TrainingResult(ac: localAc, adamStates: localAdam, metrics: m)) # ── Bot methods ─────────────────────────────────────────────────────────────── @@ -219,7 +219,9 @@ when isMainModule: resultChan.open() createDir(weightsRoot) cleanStaleTempDirs(weightsRoot) - discard loadBestAvailable(ac, weightsRoot) + let loadResult = loadBestAvailable(ac, gAdamStates, weightsRoot) + if loadResult.loaded: + roundCounter = loadResult.roundNum var bot = PPOBot( tracker: initEnemyTracker(), diff --git a/PPO_Bot/tests/test_weights.nim b/PPO_Bot/tests/test_weights.nim index f65cc8a..b2143dc 100644 --- a/PPO_Bot/tests/test_weights.nim +++ b/PPO_Bot/tests/test_weights.nim @@ -4,6 +4,7 @@ import std/[os, math] import arraymancer import "../network" +import "../training" import "../weights" template check(cond: bool, msg: string) = @@ -67,18 +68,21 @@ block testLoadBest: # Try with no weights — should return false var acEmpty = initActorCritic() - check not loadBestAvailable(acEmpty, root), "should return false with no weights" + var adamEmpty: ACAdamStates + check not loadBestAvailable(acEmpty, adamEmpty, root).loaded, "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/" + var adam1: ACAdamStates + check loadBestAvailable(ac1, adam1, root).loaded, "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/" + var adam2: ACAdamStates + check loadBestAvailable(ac2, adam2, root).loaded, "should load from checkpoint_1/" removeDir(root) diff --git a/PPO_Bot/training.nim b/PPO_Bot/training.nim index 5e8dcf4..d06cc97 100644 --- a/PPO_Bot/training.nim +++ b/PPO_Bot/training.nim @@ -71,9 +71,9 @@ proc computeGAE*(rewards, values: seq[float32]; # ── Manual Adam state ───────────────────────────────────────────────────────── type - AdamState = object - m, v: Tensor[float32] - t: int + AdamState* = object + m*, v*: Tensor[float32] + t*: int proc initAdamState(like: Tensor[float32]): AdamState = result.m = zeros_like(like) @@ -130,9 +130,9 @@ proc mlpBackward(mlp: MLP; fwd: MLPFwd; x: Tensor[float32]; type ACAdamStates* = object ## One AdamState per learnable tensor in ActorCritic. - aw1, ab1, aw2, ab2, aw3, ab3: AdamState # actor MLP - cw1, cb1, cw2, cb2, cw3, cb3: AdamState # critic MLP - logStd: AdamState + aw1*, ab1*, aw2*, ab2*, aw3*, ab3*: AdamState # actor MLP + cw1*, cb1*, cw2*, cb2*, cw3*, cb3*: AdamState # critic MLP + logStd*: AdamState initialized*: bool proc initACAdamStates*(ac: ActorCritic): ACAdamStates = diff --git a/PPO_Bot/weights.nim b/PPO_Bot/weights.nim index 9929b34..bab6ee7 100644 --- a/PPO_Bot/weights.nim +++ b/PPO_Bot/weights.nim @@ -3,6 +3,7 @@ import std/[os, times, strutils, algorithm, sequtils] import arraymancer import ./network +import ./training # ── Tensor names — order must match save/load ───────────────────────────────── @@ -61,17 +62,105 @@ proc saveWeightsAtomic*(ac: ActorCritic, targetDir: string) = removeDir(targetDir) moveDir(tmpDir, targetDir) -proc saveCheckpoint*(ac: ActorCritic, weightsRoot: string, roundNum: int) = - ## Always saves to weightsRoot/latest/. +# ── Adam state file names ────────────────────────────────────────────────────── + +const adamMFiles = [ + "adam_aw1_m.npy", "adam_ab1_m.npy", "adam_aw2_m.npy", "adam_ab2_m.npy", + "adam_aw3_m.npy", "adam_ab3_m.npy", + "adam_cw1_m.npy", "adam_cb1_m.npy", "adam_cw2_m.npy", "adam_cb2_m.npy", + "adam_cw3_m.npy", "adam_cb3_m.npy", + "adam_logstd_m.npy", +] +const adamVFiles = [ + "adam_aw1_v.npy", "adam_ab1_v.npy", "adam_aw2_v.npy", "adam_ab2_v.npy", + "adam_aw3_v.npy", "adam_ab3_v.npy", + "adam_cw1_v.npy", "adam_cb1_v.npy", "adam_cw2_v.npy", "adam_cb2_v.npy", + "adam_cw3_v.npy", "adam_cb3_v.npy", + "adam_logstd_v.npy", +] + +proc saveAdamStates*(adam: ACAdamStates, dir: string) = + ## Write Adam m/v tensors and t counters to dir/. + adam.aw1.m.write_npy(dir / "adam_aw1_m.npy"); adam.aw1.v.write_npy(dir / "adam_aw1_v.npy") + adam.ab1.m.write_npy(dir / "adam_ab1_m.npy"); adam.ab1.v.write_npy(dir / "adam_ab1_v.npy") + adam.aw2.m.write_npy(dir / "adam_aw2_m.npy"); adam.aw2.v.write_npy(dir / "adam_aw2_v.npy") + adam.ab2.m.write_npy(dir / "adam_ab2_m.npy"); adam.ab2.v.write_npy(dir / "adam_ab2_v.npy") + adam.aw3.m.write_npy(dir / "adam_aw3_m.npy"); adam.aw3.v.write_npy(dir / "adam_aw3_v.npy") + adam.ab3.m.write_npy(dir / "adam_ab3_m.npy"); adam.ab3.v.write_npy(dir / "adam_ab3_v.npy") + adam.cw1.m.write_npy(dir / "adam_cw1_m.npy"); adam.cw1.v.write_npy(dir / "adam_cw1_v.npy") + adam.cb1.m.write_npy(dir / "adam_cb1_m.npy"); adam.cb1.v.write_npy(dir / "adam_cb1_v.npy") + adam.cw2.m.write_npy(dir / "adam_cw2_m.npy"); adam.cw2.v.write_npy(dir / "adam_cw2_v.npy") + adam.cb2.m.write_npy(dir / "adam_cb2_m.npy"); adam.cb2.v.write_npy(dir / "adam_cb2_v.npy") + adam.cw3.m.write_npy(dir / "adam_cw3_m.npy"); adam.cw3.v.write_npy(dir / "adam_cw3_v.npy") + adam.cb3.m.write_npy(dir / "adam_cb3_m.npy"); adam.cb3.v.write_npy(dir / "adam_cb3_v.npy") + adam.logStd.m.write_npy(dir / "adam_logstd_m.npy") + adam.logStd.v.write_npy(dir / "adam_logstd_v.npy") + # t counters (all stepped in lockstep; store each for safety) + writeFile(dir / "adam_t.txt", + [$adam.aw1.t, $adam.ab1.t, $adam.aw2.t, $adam.ab2.t, + $adam.aw3.t, $adam.ab3.t, $adam.cw1.t, $adam.cb1.t, + $adam.cw2.t, $adam.cb2.t, $adam.cw3.t, $adam.cb3.t, + $adam.logStd.t].join("\n")) + +proc loadAdamStates*(adam: var ACAdamStates, dir: string) = + ## Load Adam m/v tensors and t counters from dir/. Called only when files exist. + template lm(dest: untyped, path: string) = + dest = read_npy[float32](path) + lm(adam.aw1.m, dir / "adam_aw1_m.npy"); lm(adam.aw1.v, dir / "adam_aw1_v.npy") + lm(adam.ab1.m, dir / "adam_ab1_m.npy"); lm(adam.ab1.v, dir / "adam_ab1_v.npy") + lm(adam.aw2.m, dir / "adam_aw2_m.npy"); lm(adam.aw2.v, dir / "adam_aw2_v.npy") + lm(adam.ab2.m, dir / "adam_ab2_m.npy"); lm(adam.ab2.v, dir / "adam_ab2_v.npy") + lm(adam.aw3.m, dir / "adam_aw3_m.npy"); lm(adam.aw3.v, dir / "adam_aw3_v.npy") + lm(adam.ab3.m, dir / "adam_ab3_m.npy"); lm(adam.ab3.v, dir / "adam_ab3_v.npy") + lm(adam.cw1.m, dir / "adam_cw1_m.npy"); lm(adam.cw1.v, dir / "adam_cw1_v.npy") + lm(adam.cb1.m, dir / "adam_cb1_m.npy"); lm(adam.cb1.v, dir / "adam_cb1_v.npy") + lm(adam.cw2.m, dir / "adam_cw2_m.npy"); lm(adam.cw2.v, dir / "adam_cw2_v.npy") + lm(adam.cb2.m, dir / "adam_cb2_m.npy"); lm(adam.cb2.v, dir / "adam_cb2_v.npy") + lm(adam.cw3.m, dir / "adam_cw3_m.npy"); lm(adam.cw3.v, dir / "adam_cw3_v.npy") + lm(adam.cb3.m, dir / "adam_cb3_m.npy"); lm(adam.cb3.v, dir / "adam_cb3_v.npy") + lm(adam.logStd.m, dir / "adam_logstd_m.npy"); lm(adam.logStd.v, dir / "adam_logstd_v.npy") + let ts = readFile(dir / "adam_t.txt").strip().splitLines() + if ts.len >= 13: + adam.aw1.t = parseInt(ts[0]); adam.ab1.t = parseInt(ts[1]) + adam.aw2.t = parseInt(ts[2]); adam.ab2.t = parseInt(ts[3]) + adam.aw3.t = parseInt(ts[4]); adam.ab3.t = parseInt(ts[5]) + adam.cw1.t = parseInt(ts[6]); adam.cb1.t = parseInt(ts[7]) + adam.cw2.t = parseInt(ts[8]); adam.cb2.t = parseInt(ts[9]) + adam.cw3.t = parseInt(ts[10]); adam.cb3.t = parseInt(ts[11]) + adam.logStd.t = parseInt(ts[12]) + adam.initialized = true + +proc adamStateFilesExist(dir: string): bool = + ## Check that the minimum set of Adam files is present. + for f in adamMFiles: + if not fileExists(dir / f): return false + for f in adamVFiles: + if not fileExists(dir / f): return false + fileExists(dir / "adam_t.txt") + +proc saveCheckpoint*(ac: ActorCritic, adam: ACAdamStates, + weightsRoot: string, roundNum: int) = + ## Always saves weights + Adam state to weightsRoot/latest/. ## Every 50 rounds also saves to checkpoint_{1,2,3} in round-robin. - saveWeightsAtomic(ac, weightsRoot / "latest") + ## Round counter is saved to weightsRoot/round_counter.txt (outside checkpoint dirs). + let latestDir = weightsRoot / "latest" + saveWeightsAtomic(ac, latestDir) + if adam.initialized: + saveAdamStates(adam, latestDir) + writeFile(weightsRoot / "round_counter.txt", $roundNum) 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)) + let ckDir = weightsRoot / ("checkpoint_" & $slot) + saveWeightsAtomic(ac, ckDir) + if adam.initialized: + saveAdamStates(adam, ckDir) -proc loadBestAvailable*(ac: var ActorCritic, weightsRoot: string): bool = +proc loadBestAvailable*(ac: var ActorCritic, adam: var ACAdamStates, + weightsRoot: string): tuple[loaded: bool, roundNum: int] = ## Try latest/ first, then checkpoints sorted newest-first by mtime. - ## Returns true if weights loaded, false if all fail (random init stays). + ## Returns (true, roundNum) if weights loaded, (false, 0) if all fail. + ## Adam state is loaded if present alongside weights; otherwise left uninitialised. + ## Round counter is read from weightsRoot/round_counter.txt if present. let checkpoints = [weightsRoot / "checkpoint_1", weightsRoot / "checkpoint_2", weightsRoot / "checkpoint_3"] @@ -93,8 +182,12 @@ proc loadBestAvailable*(ac: var ActorCritic, weightsRoot: string): bool = break if ok: ac.loadWeights(candidate) - return true - result = false + if adamStateFilesExist(candidate): + adam.loadAdamStates(candidate) + let rcPath = weightsRoot / "round_counter.txt" + let roundNum = if fileExists(rcPath): parseInt(readFile(rcPath).strip()) else: 0 + return (loaded: true, roundNum: roundNum) + result = (loaded: false, roundNum: 0) proc cleanStaleTempDirs*(weightsRoot: string) = ## Delete any dirs inside weightsRoot whose name contains "_tmp_".