Files
SirRoboGarage/SAC_LSTM_Bot/tests/test_weights.nim
T
SirStone 54b8139b11 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>
2026-08-20 23:57:13 +02:00

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