## weights.nim — save/load ActorCritic weights as .npy files. import std/[os, times, strutils, algorithm, sequtils] import arraymancer import ./network import ./training # ── 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/. ## If a tensor's shape doesn't match (e.g. STATE_DIM changed), keep the ## freshly-initialised value and print a warning — other tensors still load. template loadOrSkip(dest: untyped, path: string) = let loaded = read_npy[float32](path) if loaded.shape == dest.shape: dest = loaded else: echo "weights: shape mismatch for " & path & " (got " & $loaded.shape & " want " & $dest.shape & ") — keeping fresh init" loadOrSkip(ac.actor.w1, dir / "actor_w1.npy") loadOrSkip(ac.actor.b1, dir / "actor_b1.npy") loadOrSkip(ac.actor.w2, dir / "actor_w2.npy") loadOrSkip(ac.actor.b2, dir / "actor_b2.npy") loadOrSkip(ac.actor.w3, dir / "actor_w3.npy") loadOrSkip(ac.actor.b3, dir / "actor_b3.npy") loadOrSkip(ac.critic.w1, dir / "critic_w1.npy") loadOrSkip(ac.critic.b1, dir / "critic_b1.npy") loadOrSkip(ac.critic.w2, dir / "critic_w2.npy") loadOrSkip(ac.critic.b2, dir / "critic_b2.npy") loadOrSkip(ac.critic.w3, dir / "critic_w3.npy") loadOrSkip(ac.critic.b3, dir / "critic_b3.npy") loadOrSkip(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) # ── 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. ## Shape mismatch (e.g. STATE_DIM changed) → keep zero-initialised state (safe fresh start). template lm(dest: untyped, path: string) = let loaded = read_npy[float32](path) if loaded.shape == dest.shape: dest = loaded else: echo "weights: Adam shape mismatch for " & path & " (got " & $loaded.shape & " want " & $dest.shape & ") — resetting Adam state" 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. ## 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, … let ckDir = weightsRoot / ("checkpoint_" & $slot) saveWeightsAtomic(ac, ckDir) if adam.initialized: saveAdamStates(adam, ckDir) 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, 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"] # Sort checkpoints newest-first by modification time var existing: seq[tuple[mtime: Time, path: string]] for p in checkpoints: if dirExists(p): existing.add((getLastModificationTime(p), p)) existing.sort(proc(a, b: tuple[mtime: Time, path: string]): int = cmp(b.mtime, a.mtime)) # descending let candidates = @[weightsRoot / "latest"] & existing.mapIt(it.path) for candidate in candidates: if dirExists(candidate): var ok = true for f in weightFiles: if not fileExists(candidate / f): ok = false break if ok: ac.loadWeights(candidate) if adamStateFilesExist(candidate): adam.loadAdamStates(candidate) let rcPath = weightsRoot / "round_counter.txt" # Torn/empty file (e.g. after a crash) must not abort startup → treat as 0 let roundNum = if fileExists(rcPath): try: parseInt(readFile(rcPath).strip()) except ValueError: 0 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_". if not dirExists(weightsRoot): return for kind, path in walkDir(weightsRoot): if kind == pcDir and "_tmp_" in lastPathPart(path): removeDir(path)