b509195ee9
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
353 lines
15 KiB
Nim
353 lines
15 KiB
Nim
## 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)
|