aea0724d3a
- enemy_tracker: toggle lastOvershootDir each tick; make getRadarTurnRate take var tracker - training: remove threadvar Adam globals; pass adamStates as var param to ppoUpdate; export ACAdamStates - PPO_Bot: carry ACAdamStates through TrainingArgs/TrainingResult; drop trainingDone bool and Lock — use resultChan.tryRecv() directly as synchronisation - weights: sort checkpoint dirs newest-first by mtime instead of hardcoded order - tests/test_training: pass explicit ACAdamStates to ppoUpdate Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
105 lines
4.2 KiB
Nim
105 lines
4.2 KiB
Nim
## 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)
|