## 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)