Files
SirRoboGarage/SAC_LSTM_Bot_garage/src/SAC_LSTM_Bot/weights.nim
T

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)