From 54b8139b1144ef27a5cba88451eb69dba20c8e15 Mon Sep 17 00:00:00 2001 From: Davide Cappellini Date: Thu, 20 Aug 2026 23:57:13 +0200 Subject: [PATCH] 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 --- SAC_LSTM_Bot/config.nims | 3 + SAC_LSTM_Bot/src/SAC_LSTM_Bot/weights.nim | 352 ++++++++++++++++++++++ SAC_LSTM_Bot/tests/test_weights.nim | 141 +++++++++ 3 files changed, 496 insertions(+) create mode 100644 SAC_LSTM_Bot/src/SAC_LSTM_Bot/weights.nim create mode 100644 SAC_LSTM_Bot/tests/test_weights.nim diff --git a/SAC_LSTM_Bot/config.nims b/SAC_LSTM_Bot/config.nims index 678d11a..29e2072 100644 --- a/SAC_LSTM_Bot/config.nims +++ b/SAC_LSTM_Bot/config.nims @@ -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")): diff --git a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/weights.nim b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/weights.nim new file mode 100644 index 0000000..b882ee3 --- /dev/null +++ b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/weights.nim @@ -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) diff --git a/SAC_LSTM_Bot/tests/test_weights.nim b/SAC_LSTM_Bot/tests/test_weights.nim new file mode 100644 index 0000000..025b2c5 --- /dev/null +++ b/SAC_LSTM_Bot/tests/test_weights.nim @@ -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