feat(SAC_LSTM_Bot): weight persistence module (#46)
Save/load all SAC-LSTM tensors (actor, 2 critics, 2 target critics, alpha, Adam states) into a single .zip of .npy files. Atomic write via temp path + rename. Adam types (AdamVar, SACAdamStates) defined here for training.nim to use. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -1,6 +1,9 @@
|
||||
# Static-link OpenBLAS for portable deployment
|
||||
# ponytail: adjust path per machine, or use pkg-config
|
||||
switch("passL", "-L/nix/store/v07svn2y92bvzjl51aj7c9ca1cwg7rw7-openblas-0.3.32/lib -lopenblas")
|
||||
# libzip for weight checkpoint zip files
|
||||
# ponytail: nix store path; adjust per machine, or use pkg-config
|
||||
switch("passL", "-L/nix/store/wqvz31s598bvj3zb747943xhl38hjc6h-libzip-1.11.4/lib -lzip")
|
||||
switch("threads", "on")
|
||||
# begin Nimble config (version 2)
|
||||
when withDir(thisDir(), system.fileExists("nimble.paths")):
|
||||
|
||||
@@ -0,0 +1,352 @@
|
||||
## weights.nim — save/load all SAC-LSTM network tensors as .npy inside a .zip.
|
||||
##
|
||||
## Strategy: write_npy writes to paths; zip/zipfiles.addFile reads from paths.
|
||||
## So we write each tensor to a temp .npy, add it to the zip, then delete temps.
|
||||
## Load reverses: extract each entry to a temp .npy, read_npy, delete.
|
||||
## Atomic save: build the zip in a temp path, then rename over the target.
|
||||
|
||||
import arraymancer except Linear
|
||||
import zip/zipfiles
|
||||
import std/[os, times, strutils]
|
||||
import SAC_LSTM_Bot/network
|
||||
|
||||
# ── Adam state types (used by training.nim) ───────────────────────────────────
|
||||
|
||||
type
|
||||
AdamVar* = object
|
||||
m*, v*: Tensor[float32]
|
||||
t*: int
|
||||
|
||||
## Adam states for one Linear layer (w and b).
|
||||
LinearAdam* = object
|
||||
w*, b*: AdamVar
|
||||
|
||||
## Adam states for one LSTMCell (wCombined and bCombined).
|
||||
LSTMCellAdam* = object
|
||||
wCombined*, bCombined*: AdamVar
|
||||
|
||||
## Adam states for one ActorNet.
|
||||
ActorAdam* = object
|
||||
fc1*, fc2*, muHead*, logStdHead*: LinearAdam
|
||||
lstm*: LSTMCellAdam
|
||||
|
||||
## Adam states for one CriticNet.
|
||||
CriticAdam* = object
|
||||
fc1*, fc2*, fc3*: LinearAdam
|
||||
lstm*: LSTMCellAdam
|
||||
|
||||
SACAdamStates* = object
|
||||
actor*: ActorAdam
|
||||
critic1*: CriticAdam
|
||||
critic2*: CriticAdam
|
||||
alpha*: AdamVar # scalar, shape [1]
|
||||
initialized*: bool
|
||||
|
||||
# ── Init helpers ──────────────────────────────────────────────────────────────
|
||||
|
||||
proc initAdamVar(t: Tensor[float32]): AdamVar =
|
||||
AdamVar(m: zeros[float32](t.shape), v: zeros[float32](t.shape), t: 0)
|
||||
|
||||
proc initLinearAdam*(l: Linear): LinearAdam =
|
||||
LinearAdam(w: initAdamVar(l.w), b: initAdamVar(l.b))
|
||||
|
||||
proc initLSTMCellAdam*(c: LSTMCell): LSTMCellAdam =
|
||||
LSTMCellAdam(
|
||||
wCombined: initAdamVar(c.wCombined),
|
||||
bCombined: initAdamVar(c.bCombined))
|
||||
|
||||
proc initActorAdam*(a: ActorNet): ActorAdam =
|
||||
ActorAdam(
|
||||
fc1: initLinearAdam(a.fc1),
|
||||
fc2: initLinearAdam(a.fc2),
|
||||
muHead: initLinearAdam(a.muHead),
|
||||
logStdHead: initLinearAdam(a.logStdHead),
|
||||
lstm: initLSTMCellAdam(a.lstm))
|
||||
|
||||
proc initCriticAdam*(c: CriticNet): CriticAdam =
|
||||
CriticAdam(
|
||||
fc1: initLinearAdam(c.fc1),
|
||||
fc2: initLinearAdam(c.fc2),
|
||||
fc3: initLinearAdam(c.fc3),
|
||||
lstm: initLSTMCellAdam(c.lstm))
|
||||
|
||||
proc initSACAdamStates*(actor: ActorNet; critic1, critic2: CriticNet): SACAdamStates =
|
||||
result.actor = initActorAdam(actor)
|
||||
result.critic1 = initCriticAdam(critic1)
|
||||
result.critic2 = initCriticAdam(critic2)
|
||||
result.alpha = initAdamVar(ones[float32](1))
|
||||
result.initialized = true
|
||||
|
||||
# ── Internal: temp dir per save ───────────────────────────────────────────────
|
||||
|
||||
proc tmpDir(): string =
|
||||
getTempDir() / ("sacw_" & $int(epochTime() * 1000))
|
||||
|
||||
# ── Save helpers ──────────────────────────────────────────────────────────────
|
||||
|
||||
template addT(z: var ZipArchive; name: string; t: Tensor[float32]; tmp: string) =
|
||||
## Write tensor to a temp file, add to zip, delete temp file.
|
||||
let p = tmp / name
|
||||
t.write_npy(p)
|
||||
z.addFile(name, p)
|
||||
|
||||
proc addLinear(z: var ZipArchive; prefix: string; l: Linear; tmp: string) =
|
||||
addT(z, prefix & "_w.npy", l.w, tmp)
|
||||
addT(z, prefix & "_b.npy", l.b, tmp)
|
||||
|
||||
proc addLSTMCell(z: var ZipArchive; prefix: string; c: LSTMCell; tmp: string) =
|
||||
addT(z, prefix & "_wc.npy", c.wCombined, tmp)
|
||||
addT(z, prefix & "_bc.npy", c.bCombined, tmp)
|
||||
|
||||
proc addActorNet(z: var ZipArchive; prefix: string; a: ActorNet; tmp: string) =
|
||||
addLinear(z, prefix & "_fc1", a.fc1, tmp)
|
||||
addLSTMCell(z, prefix & "_lstm", a.lstm, tmp)
|
||||
addLinear(z, prefix & "_fc2", a.fc2, tmp)
|
||||
addLinear(z, prefix & "_mu", a.muHead, tmp)
|
||||
addLinear(z, prefix & "_logstd", a.logStdHead, tmp)
|
||||
|
||||
proc addCriticNet(z: var ZipArchive; prefix: string; c: CriticNet; tmp: string) =
|
||||
addLinear(z, prefix & "_fc1", c.fc1, tmp)
|
||||
addLSTMCell(z, prefix & "_lstm", c.lstm, tmp)
|
||||
addLinear(z, prefix & "_fc2", c.fc2, tmp)
|
||||
addLinear(z, prefix & "_fc3", c.fc3, tmp)
|
||||
|
||||
proc addAdamVar(z: var ZipArchive; prefix: string; v: AdamVar; tmp: string) =
|
||||
addT(z, prefix & "_m.npy", v.m, tmp)
|
||||
addT(z, prefix & "_v.npy", v.v, tmp)
|
||||
|
||||
proc addLinearAdam(z: var ZipArchive; prefix: string; la: LinearAdam; tmp: string) =
|
||||
addAdamVar(z, prefix & "_w", la.w, tmp)
|
||||
addAdamVar(z, prefix & "_b", la.b, tmp)
|
||||
|
||||
proc addLSTMCellAdam(z: var ZipArchive; prefix: string; la: LSTMCellAdam; tmp: string) =
|
||||
addAdamVar(z, prefix & "_wc", la.wCombined, tmp)
|
||||
addAdamVar(z, prefix & "_bc", la.bCombined, tmp)
|
||||
|
||||
proc addActorAdam(z: var ZipArchive; prefix: string; a: ActorAdam; tmp: string) =
|
||||
addLinearAdam(z, prefix & "_fc1", a.fc1, tmp)
|
||||
addLSTMCellAdam(z, prefix & "_lstm", a.lstm, tmp)
|
||||
addLinearAdam(z, prefix & "_fc2", a.fc2, tmp)
|
||||
addLinearAdam(z, prefix & "_mu", a.muHead, tmp)
|
||||
addLinearAdam(z, prefix & "_logstd", a.logStdHead, tmp)
|
||||
|
||||
proc addCriticAdam(z: var ZipArchive; prefix: string; c: CriticAdam; tmp: string) =
|
||||
addLinearAdam(z, prefix & "_fc1", c.fc1, tmp)
|
||||
addLSTMCellAdam(z, prefix & "_lstm", c.lstm, tmp)
|
||||
addLinearAdam(z, prefix & "_fc2", c.fc2, tmp)
|
||||
addLinearAdam(z, prefix & "_fc3", c.fc3, tmp)
|
||||
|
||||
# ── Public API ────────────────────────────────────────────────────────────────
|
||||
|
||||
proc saveWeights*(path: string;
|
||||
actor: ActorNet;
|
||||
critic1, critic2: CriticNet;
|
||||
targetCritic1, targetCritic2: CriticNet;
|
||||
alpha: float32) =
|
||||
## Save network tensors (no Adam states) to `path` (.zip).
|
||||
## Atomic: writes to a temp path first, then renames.
|
||||
let tmp = tmpDir()
|
||||
createDir(tmp)
|
||||
let tmpZip = path & ".tmp"
|
||||
try:
|
||||
var z: ZipArchive
|
||||
if not z.open(tmpZip, fmWrite):
|
||||
raise newException(IOError, "cannot create zip: " & tmpZip)
|
||||
addActorNet(z, "actor", actor, tmp)
|
||||
addCriticNet(z, "c1", critic1, tmp)
|
||||
addCriticNet(z, "c2", critic2, tmp)
|
||||
addCriticNet(z, "tc1", targetCritic1,tmp)
|
||||
addCriticNet(z, "tc2", targetCritic2,tmp)
|
||||
# alpha: store as a 1-element tensor
|
||||
addT(z, "alpha.npy", [alpha].toTensor.asType(float32), tmp)
|
||||
z.close()
|
||||
createDir(path.parentDir)
|
||||
moveFile(tmpZip, path)
|
||||
finally:
|
||||
removeDir(tmp)
|
||||
if fileExists(tmpZip): removeFile(tmpZip)
|
||||
|
||||
proc saveCheckpoint*(path: string;
|
||||
actor: ActorNet;
|
||||
critic1, critic2: CriticNet;
|
||||
targetCritic1, targetCritic2: CriticNet;
|
||||
alpha: float32;
|
||||
adam: SACAdamStates) =
|
||||
## Save networks + Adam states to `path` (.zip). Atomic.
|
||||
let tmp = tmpDir()
|
||||
createDir(tmp)
|
||||
let tmpZip = path & ".tmp"
|
||||
try:
|
||||
var z: ZipArchive
|
||||
if not z.open(tmpZip, fmWrite):
|
||||
raise newException(IOError, "cannot create zip: " & tmpZip)
|
||||
addActorNet(z, "actor", actor, tmp)
|
||||
addCriticNet(z, "c1", critic1, tmp)
|
||||
addCriticNet(z, "c2", critic2, tmp)
|
||||
addCriticNet(z, "tc1", targetCritic1,tmp)
|
||||
addCriticNet(z, "tc2", targetCritic2,tmp)
|
||||
addT(z, "alpha.npy", [alpha].toTensor.asType(float32), tmp)
|
||||
if adam.initialized:
|
||||
addActorAdam(z, "adam_actor", adam.actor, tmp)
|
||||
addCriticAdam(z, "adam_c1", adam.critic1, tmp)
|
||||
addCriticAdam(z, "adam_c2", adam.critic2, tmp)
|
||||
addAdamVar(z, "adam_alpha", adam.alpha, tmp)
|
||||
# t counters (all in lockstep; store as text)
|
||||
writeFile(tmp / "adam_t.txt",
|
||||
$adam.actor.fc1.w.t & "\n" &
|
||||
$adam.critic1.fc1.w.t & "\n" &
|
||||
$adam.critic2.fc1.w.t & "\n" &
|
||||
$adam.alpha.t)
|
||||
z.addFile("adam_t.txt", tmp / "adam_t.txt")
|
||||
z.close()
|
||||
createDir(path.parentDir)
|
||||
moveFile(tmpZip, path)
|
||||
finally:
|
||||
removeDir(tmp)
|
||||
if fileExists(tmpZip): removeFile(tmpZip)
|
||||
|
||||
# ── Load helpers ──────────────────────────────────────────────────────────────
|
||||
|
||||
template loadT(name: string; tmp: string): Tensor[float32] =
|
||||
read_npy[float32](tmp / name)
|
||||
|
||||
proc loadLinear(z: var ZipArchive; prefix, tmp: string): Linear =
|
||||
z.extractFile(prefix & "_w.npy", tmp / (prefix & "_w.npy"))
|
||||
z.extractFile(prefix & "_b.npy", tmp / (prefix & "_b.npy"))
|
||||
result.w = read_npy[float32](tmp / (prefix & "_w.npy"))
|
||||
result.b = read_npy[float32](tmp / (prefix & "_b.npy"))
|
||||
|
||||
proc loadLSTMCell(z: var ZipArchive; prefix, tmp: string): LSTMCell =
|
||||
z.extractFile(prefix & "_wc.npy", tmp / (prefix & "_wc.npy"))
|
||||
z.extractFile(prefix & "_bc.npy", tmp / (prefix & "_bc.npy"))
|
||||
result.wCombined = read_npy[float32](tmp / (prefix & "_wc.npy"))
|
||||
result.bCombined = read_npy[float32](tmp / (prefix & "_bc.npy"))
|
||||
result.hiddenDim = result.bCombined.shape[0] div 4
|
||||
|
||||
proc loadActorNet(z: var ZipArchive; prefix, tmp: string): ActorNet =
|
||||
result.fc1 = loadLinear(z, prefix & "_fc1", tmp)
|
||||
result.lstm = loadLSTMCell(z, prefix & "_lstm", tmp)
|
||||
result.fc2 = loadLinear(z, prefix & "_fc2", tmp)
|
||||
result.muHead = loadLinear(z, prefix & "_mu", tmp)
|
||||
result.logStdHead = loadLinear(z, prefix & "_logstd", tmp)
|
||||
result.hiddenDim = result.lstm.hiddenDim
|
||||
|
||||
proc loadCriticNet(z: var ZipArchive; prefix, tmp: string): CriticNet =
|
||||
result.fc1 = loadLinear(z, prefix & "_fc1", tmp)
|
||||
result.lstm = loadLSTMCell(z, prefix & "_lstm", tmp)
|
||||
result.fc2 = loadLinear(z, prefix & "_fc2", tmp)
|
||||
result.fc3 = loadLinear(z, prefix & "_fc3", tmp)
|
||||
result.hiddenDim = result.lstm.hiddenDim
|
||||
|
||||
proc loadAdamVarFromZip(z: var ZipArchive; prefix, tmp: string): AdamVar =
|
||||
z.extractFile(prefix & "_m.npy", tmp / (prefix & "_m.npy"))
|
||||
z.extractFile(prefix & "_v.npy", tmp / (prefix & "_v.npy"))
|
||||
result.m = read_npy[float32](tmp / (prefix & "_m.npy"))
|
||||
result.v = read_npy[float32](tmp / (prefix & "_v.npy"))
|
||||
|
||||
proc loadLinearAdam(z: var ZipArchive; prefix, tmp: string): LinearAdam =
|
||||
result.w = loadAdamVarFromZip(z, prefix & "_w", tmp)
|
||||
result.b = loadAdamVarFromZip(z, prefix & "_b", tmp)
|
||||
|
||||
proc loadLSTMCellAdam(z: var ZipArchive; prefix, tmp: string): LSTMCellAdam =
|
||||
result.wCombined = loadAdamVarFromZip(z, prefix & "_wc", tmp)
|
||||
result.bCombined = loadAdamVarFromZip(z, prefix & "_bc", tmp)
|
||||
|
||||
proc loadActorAdam(z: var ZipArchive; prefix, tmp: string): ActorAdam =
|
||||
result.fc1 = loadLinearAdam(z, prefix & "_fc1", tmp)
|
||||
result.lstm = loadLSTMCellAdam(z, prefix & "_lstm", tmp)
|
||||
result.fc2 = loadLinearAdam(z, prefix & "_fc2", tmp)
|
||||
result.muHead = loadLinearAdam(z, prefix & "_mu", tmp)
|
||||
result.logStdHead = loadLinearAdam(z, prefix & "_logstd", tmp)
|
||||
|
||||
proc loadCriticAdam(z: var ZipArchive; prefix, tmp: string): CriticAdam =
|
||||
result.fc1 = loadLinearAdam(z, prefix & "_fc1", tmp)
|
||||
result.lstm = loadLSTMCellAdam(z, prefix & "_lstm", tmp)
|
||||
result.fc2 = loadLinearAdam(z, prefix & "_fc2", tmp)
|
||||
result.fc3 = loadLinearAdam(z, prefix & "_fc3", tmp)
|
||||
|
||||
type
|
||||
WeightCheckpoint* = object
|
||||
actor*: ActorNet
|
||||
critic1*: CriticNet
|
||||
critic2*: CriticNet
|
||||
targetCritic1*: CriticNet
|
||||
targetCritic2*: CriticNet
|
||||
alpha*: float32
|
||||
adam*: SACAdamStates ## initialized=false if not present in zip
|
||||
|
||||
proc loadCheckpoint*(path: string): WeightCheckpoint =
|
||||
## Load all tensors from `path` (.zip). Raises IOError if file not found.
|
||||
## Adam states loaded only if present; result.adam.initialized reflects this.
|
||||
if not fileExists(path):
|
||||
raise newException(IOError, "checkpoint not found: " & path)
|
||||
let tmp = tmpDir()
|
||||
createDir(tmp)
|
||||
try:
|
||||
var z: ZipArchive
|
||||
if not z.open(path, fmRead):
|
||||
raise newException(IOError, "cannot open zip: " & path)
|
||||
|
||||
result.actor = loadActorNet(z, "actor", tmp)
|
||||
result.critic1 = loadCriticNet(z, "c1", tmp)
|
||||
result.critic2 = loadCriticNet(z, "c2", tmp)
|
||||
result.targetCritic1 = loadCriticNet(z, "tc1", tmp)
|
||||
result.targetCritic2 = loadCriticNet(z, "tc2", tmp)
|
||||
|
||||
z.extractFile("alpha.npy", tmp / "alpha.npy")
|
||||
let alphaTensor = read_npy[float32](tmp / "alpha.npy")
|
||||
result.alpha = alphaTensor[0]
|
||||
|
||||
# Adam states — optional
|
||||
var hasAdam = false
|
||||
for f in z.walkFiles:
|
||||
if f.startsWith("adam_"):
|
||||
hasAdam = true
|
||||
break
|
||||
if hasAdam:
|
||||
result.adam.actor = loadActorAdam(z, "adam_actor", tmp)
|
||||
result.adam.critic1 = loadCriticAdam(z, "adam_c1", tmp)
|
||||
result.adam.critic2 = loadCriticAdam(z, "adam_c2", tmp)
|
||||
result.adam.alpha = loadAdamVarFromZip(z, "adam_alpha", tmp)
|
||||
# t counters
|
||||
z.extractFile("adam_t.txt", tmp / "adam_t.txt")
|
||||
let ts = readFile(tmp / "adam_t.txt").strip().splitLines()
|
||||
if ts.len >= 4:
|
||||
let tActor = parseInt(ts[0])
|
||||
let tCritic1 = parseInt(ts[1])
|
||||
let tCritic2 = parseInt(ts[2])
|
||||
let tAlpha = parseInt(ts[3])
|
||||
# propagate t to all Adam vars
|
||||
template setT(v: var AdamVar; tval: int) = v.t = tval
|
||||
setT(result.adam.actor.fc1.w, tActor)
|
||||
setT(result.adam.actor.fc1.b, tActor)
|
||||
setT(result.adam.actor.lstm.wCombined,tActor)
|
||||
setT(result.adam.actor.lstm.bCombined,tActor)
|
||||
setT(result.adam.actor.fc2.w, tActor)
|
||||
setT(result.adam.actor.fc2.b, tActor)
|
||||
setT(result.adam.actor.muHead.w, tActor)
|
||||
setT(result.adam.actor.muHead.b, tActor)
|
||||
setT(result.adam.actor.logStdHead.w, tActor)
|
||||
setT(result.adam.actor.logStdHead.b, tActor)
|
||||
setT(result.adam.critic1.fc1.w, tCritic1)
|
||||
setT(result.adam.critic1.fc1.b, tCritic1)
|
||||
setT(result.adam.critic1.lstm.wCombined,tCritic1)
|
||||
setT(result.adam.critic1.lstm.bCombined,tCritic1)
|
||||
setT(result.adam.critic1.fc2.w, tCritic1)
|
||||
setT(result.adam.critic1.fc2.b, tCritic1)
|
||||
setT(result.adam.critic1.fc3.w, tCritic1)
|
||||
setT(result.adam.critic1.fc3.b, tCritic1)
|
||||
setT(result.adam.critic2.fc1.w, tCritic2)
|
||||
setT(result.adam.critic2.fc1.b, tCritic2)
|
||||
setT(result.adam.critic2.lstm.wCombined,tCritic2)
|
||||
setT(result.adam.critic2.lstm.bCombined,tCritic2)
|
||||
setT(result.adam.critic2.fc2.w, tCritic2)
|
||||
setT(result.adam.critic2.fc2.b, tCritic2)
|
||||
setT(result.adam.critic2.fc3.w, tCritic2)
|
||||
setT(result.adam.critic2.fc3.b, tCritic2)
|
||||
setT(result.adam.alpha, tAlpha)
|
||||
result.adam.initialized = true
|
||||
|
||||
z.close()
|
||||
finally:
|
||||
removeDir(tmp)
|
||||
@@ -0,0 +1,141 @@
|
||||
import unittest
|
||||
import arraymancer
|
||||
import zip/zipfiles
|
||||
import std/[os, math, strutils]
|
||||
import SAC_LSTM_Bot/network
|
||||
import SAC_LSTM_Bot/weights
|
||||
|
||||
const
|
||||
STATE_DIM = 20
|
||||
ACTION_DIM = 4
|
||||
|
||||
proc tensorsEqual(a, b: Tensor[float32]; tol: float32 = 1e-6'f32): bool =
|
||||
if a.shape != b.shape: return false
|
||||
for i in 0 ..< a.size:
|
||||
if abs(a.unsafe_raw_offset[i] - b.unsafe_raw_offset[i]) > tol: return false
|
||||
true
|
||||
|
||||
suite "saveWeights / loadCheckpoint":
|
||||
|
||||
setup:
|
||||
let actor = initActorNet(STATE_DIM)
|
||||
let c1 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let c2 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let tc1 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let tc2 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let alpha = 0.2'f32
|
||||
let zipPath = getTempDir() / "test_weights_latest.zip"
|
||||
|
||||
teardown:
|
||||
if fileExists(zipPath): removeFile(zipPath)
|
||||
|
||||
test "creates a valid zip file":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
check fileExists(zipPath)
|
||||
var z: ZipArchive
|
||||
check z.open(zipPath, fmRead)
|
||||
var count = 0
|
||||
for f in z.walkFiles: inc count
|
||||
z.close()
|
||||
check count > 0
|
||||
|
||||
test "all entries are .npy files (+ alpha.npy)":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
var z: ZipArchive
|
||||
discard z.open(zipPath, fmRead)
|
||||
var allNpy = true
|
||||
for f in z.walkFiles:
|
||||
if not f.endsWith(".npy"): allNpy = false
|
||||
z.close()
|
||||
check allNpy
|
||||
|
||||
test "round-trip: actor weights preserved":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check tensorsEqual(actor.fc1.w, ck.actor.fc1.w)
|
||||
check tensorsEqual(actor.fc1.b, ck.actor.fc1.b)
|
||||
check tensorsEqual(actor.lstm.wCombined, ck.actor.lstm.wCombined)
|
||||
check tensorsEqual(actor.lstm.bCombined, ck.actor.lstm.bCombined)
|
||||
check tensorsEqual(actor.fc2.w, ck.actor.fc2.w)
|
||||
check tensorsEqual(actor.muHead.w, ck.actor.muHead.w)
|
||||
check tensorsEqual(actor.logStdHead.w, ck.actor.logStdHead.w)
|
||||
|
||||
test "round-trip: critic weights preserved":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check tensorsEqual(c1.fc1.w, ck.critic1.fc1.w)
|
||||
check tensorsEqual(c1.lstm.wCombined, ck.critic1.lstm.wCombined)
|
||||
check tensorsEqual(c1.fc3.w, ck.critic1.fc3.w)
|
||||
check tensorsEqual(tc1.fc1.w, ck.targetCritic1.fc1.w)
|
||||
check tensorsEqual(tc2.fc1.w, ck.targetCritic2.fc1.w)
|
||||
|
||||
test "round-trip: alpha preserved":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check abs(ck.alpha - alpha) < 1e-6'f32
|
||||
|
||||
test "round-trip: hiddenDim reconstructed":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check ck.actor.hiddenDim == actor.hiddenDim
|
||||
check ck.critic1.hiddenDim == c1.hiddenDim
|
||||
|
||||
test "adam not present → initialized=false":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check not ck.adam.initialized
|
||||
|
||||
suite "saveCheckpoint (with Adam)":
|
||||
|
||||
setup:
|
||||
let actor = initActorNet(STATE_DIM)
|
||||
let c1 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let c2 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let tc1 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let tc2 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let alpha = 0.1'f32
|
||||
var adam = initSACAdamStates(actor, c1, c2)
|
||||
# put some non-zero values in Adam state
|
||||
adam.actor.fc1.w.m[0, 0] = 0.5'f32
|
||||
adam.actor.fc1.w.t = 42
|
||||
adam.critic1.fc1.w.t = 7
|
||||
adam.alpha.t = 99
|
||||
let zipPath = getTempDir() / "test_checkpoint.zip"
|
||||
|
||||
teardown:
|
||||
if fileExists(zipPath): removeFile(zipPath)
|
||||
|
||||
test "round-trip: Adam m tensor":
|
||||
saveCheckpoint(zipPath, actor, c1, c2, tc1, tc2, alpha, adam)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check ck.adam.initialized
|
||||
check abs(ck.adam.actor.fc1.w.m[0, 0] - 0.5'f32) < 1e-6'f32
|
||||
|
||||
test "round-trip: Adam t counters":
|
||||
saveCheckpoint(zipPath, actor, c1, c2, tc1, tc2, alpha, adam)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check ck.adam.actor.fc1.w.t == 42
|
||||
check ck.adam.critic1.fc1.w.t == 7
|
||||
check ck.adam.alpha.t == 99
|
||||
|
||||
suite "Atomic save":
|
||||
|
||||
test "temp file is cleaned up after successful save":
|
||||
let zipPath = getTempDir() / "test_atomic.zip"
|
||||
let tmpZip = zipPath & ".tmp"
|
||||
let actor = initActorNet(STATE_DIM)
|
||||
let c = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
saveWeights(zipPath, actor, c, c, c, c, 0.2'f32)
|
||||
check fileExists(zipPath)
|
||||
check not fileExists(tmpZip)
|
||||
removeFile(zipPath)
|
||||
|
||||
suite "Error handling":
|
||||
|
||||
test "loadCheckpoint missing file → IOError":
|
||||
var raised = false
|
||||
try:
|
||||
discard loadCheckpoint(getTempDir() / "nonexistent_xxxxxx.zip")
|
||||
except IOError:
|
||||
raised = true
|
||||
check raised
|
||||
Reference in New Issue
Block a user