## weights.nim — save/load ActorCritic weights as .npy files. import std/[os, times, strutils, algorithm, sequtils] 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/ first, then checkpoints sorted newest-first by mtime. ## Returns true if weights loaded, false if all fail (random init stays). 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) 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)