chore: rename libs→common_libs, all bot dirs to _garage suffix, fix all path refs
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -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