feat(PPO_Bot): persist Adam optimizer state and round counter across restarts (#35)
Save ACAdamStates (m/v tensors + t counters) as .npy files alongside network weights in latest/ and checkpoint dirs; save round counter to round_counter.txt. loadBestAvailable restores both on startup; fresh start works unchanged when files are absent. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
+4
-2
@@ -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(),
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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 =
|
||||
|
||||
+101
-8
@@ -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_".
|
||||
|
||||
Reference in New Issue
Block a user