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>
This commit is contained in:
@@ -1,6 +1,9 @@
|
|||||||
# Static-link OpenBLAS for portable deployment
|
# Static-link OpenBLAS for portable deployment
|
||||||
# ponytail: adjust path per machine, or use pkg-config
|
# ponytail: adjust path per machine, or use pkg-config
|
||||||
switch("passL", "-L/nix/store/v07svn2y92bvzjl51aj7c9ca1cwg7rw7-openblas-0.3.32/lib -lopenblas")
|
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")
|
switch("threads", "on")
|
||||||
# begin Nimble config (version 2)
|
# begin Nimble config (version 2)
|
||||||
when withDir(thisDir(), system.fileExists("nimble.paths")):
|
when withDir(thisDir(), system.fileExists("nimble.paths")):
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -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