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