Files
SirRoboGarage/PPO_Bot_garage/weights.nim
T

212 lines
10 KiB
Nim

## 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)