54b8139b11
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>
142 lines
4.7 KiB
Nim
142 lines
4.7 KiB
Nim
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
|