From add3e34926638bd4fc792a3bcd719fac55fcecff Mon Sep 17 00:00:00 2001 From: Davide Cappellini Date: Thu, 20 Aug 2026 23:38:17 +0200 Subject: [PATCH] feat(SAC_LSTM_Bot): LSTM network module (#41) Co-Authored-By: Claude Sonnet 4.6 --- SAC_LSTM_Bot/src/SAC_LSTM_Bot/network.nim | 161 ++++++++++++++++++++++ SAC_LSTM_Bot/tests/test_network.nim | 99 +++++++++++++ 2 files changed, 260 insertions(+) create mode 100644 SAC_LSTM_Bot/src/SAC_LSTM_Bot/network.nim create mode 100644 SAC_LSTM_Bot/tests/test_network.nim diff --git a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/network.nim b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/network.nim new file mode 100644 index 0000000..c1b3096 --- /dev/null +++ b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/network.nim @@ -0,0 +1,161 @@ +## network.nim — LSTM-based Actor and dual Critic for SAC-v2. +## No autograd; inference only. Manual LSTM cell from scratch. + +import arraymancer +import std/[math, random, os, strutils] + +# ── Configuration ───────────────────────────────────────────────────────────── + +proc getHiddenSize*(): int = + let s = getEnv("SACLSTM_HIDDEN_SIZE", "256") + result = parseInt(s) + +proc isEvalMode*(): bool = + getEnv("SACLSTM_EVAL_MODE", "0") == "1" + +# ── Types ───────────────────────────────────────────────────────────────────── + +type + Linear* = object + w*, b*: Tensor[float32] # w: [out, in], b: [out] + + LSTMCell* = object + ## Combined weight matrix Wi|Wf|Wg|Wo stacked: [4*hidden, input+hidden] + ## Combined bias stacked: [4*hidden] + wCombined*: Tensor[float32] + bCombined*: Tensor[float32] + hiddenDim*: int + + ActorNet* = object + fc1*: Linear + lstm*: LSTMCell + fc2*: Linear + muHead*: Linear + logStdHead*: Linear + hiddenDim*: int + + CriticNet* = object + fc1*: Linear + lstm*: LSTMCell + fc2*: Linear + fc3*: Linear + hiddenDim*: int + + LSTMState* = tuple[h, c: Tensor[float32]] # each [hiddenDim] + +# ── Init helpers ────────────────────────────────────────────────────────────── + +proc initLinear*(inDim, outDim: int; scale: float32): Linear = + result.w = randomNormalTensor[float32]([outDim, inDim]) *. scale + result.b = zeros[float32](outDim) + +proc initLinearHe*(inDim, outDim: int): Linear = + initLinear(inDim, outDim, sqrt(2.0'f32 / inDim.float32)) + +proc initLinearOut*(inDim, outDim: int): Linear = + initLinear(inDim, outDim, sqrt(1.0'f32 / inDim.float32)) + +proc initLSTMCell*(inputDim, hiddenDim: int): LSTMCell = + result.hiddenDim = hiddenDim + let fanIn = (inputDim + hiddenDim).float32 + let scale = sqrt(1.0'f32 / fanIn) + result.wCombined = randomNormalTensor[float32]([4 * hiddenDim, inputDim + hiddenDim]) *. scale + result.bCombined = zeros[float32](4 * hiddenDim) + +proc zeroState*(hiddenDim: int): LSTMState = + result = (h: zeros[float32](hiddenDim), c: zeros[float32](hiddenDim)) + +proc initActorNet*(stateDim: int): ActorNet = + let h = getHiddenSize() + result.hiddenDim = h + result.fc1 = initLinearHe(stateDim, h) + result.lstm = initLSTMCell(h, h) + result.fc2 = initLinearHe(h, 128) + result.muHead = initLinearOut(128, 4) + result.logStdHead = initLinearOut(128, 4) + +proc initCriticNet*(stateDim, actionDim: int): CriticNet = + let h = getHiddenSize() + result.hiddenDim = h + result.fc1 = initLinearHe(stateDim + actionDim, h) + result.lstm = initLSTMCell(h, h) + result.fc2 = initLinearHe(h, 128) + result.fc3 = initLinearOut(128, 1) + +# ── Forward helpers ─────────────────────────────────────────────────────────── + +proc linear*(l: Linear; x: Tensor[float32]): Tensor[float32] = + l.w * x + l.b + +proc relu*(x: Tensor[float32]): Tensor[float32] = + x.map(proc(v: float32): float32 = max(0.0'f32, v)) + +proc sigmoid*(x: Tensor[float32]): Tensor[float32] = + x.map(proc(v: float32): float32 = 1.0'f32 / (1.0'f32 + exp(-v))) + +proc tanhT*(x: Tensor[float32]): Tensor[float32] = + x.map(proc(v: float32): float32 = tanh(v)) + +proc lstmStep*(cell: LSTMCell; x, h, c: Tensor[float32]): LSTMState = + ## x: [inputDim], h/c: [hiddenDim] → h', c': [hiddenDim] + let xh = concat(x, h, axis = 0) # [inputDim + hiddenDim] + let gates = cell.wCombined * xh + cell.bCombined # [4*hidden] + let hd = cell.hiddenDim + let iGate = sigmoid(gates[0 ..< hd]) + let fGate = sigmoid(gates[hd ..< 2*hd]) + let gGate = tanhT(gates[2*hd ..< 3*hd]) + let oGate = sigmoid(gates[3*hd ..< 4*hd]) + let cPrime = fGate *. c + iGate *. gGate + let hPrime = oGate *. tanhT(cPrime) + result = (h: hPrime, c: cPrime) + +# ── Actor forward ───────────────────────────────────────────────────────────── + +const + LOG_STD_MIN = -5.0'f32 + LOG_STD_MAX = 2.0'f32 + LOG_PROB_EPS = 1e-6'f32 + +proc actorForward*(net: ActorNet; state: Tensor[float32]; lstm: LSTMState; + deterministic = false): + tuple[actions: Tensor[float32]; logProb: float32; lstm: LSTMState] = + ## state: [stateDim], lstm: (h,c) each [hiddenDim] + ## Returns actions [4], scalar logProb, updated (h',c'). + let h1 = relu(net.fc1.linear(state)) + let lstmOut = lstmStep(net.lstm, h1, lstm.h, lstm.c) + let h2 = relu(net.fc2.linear(lstmOut.h)) + let mu = net.muHead.linear(h2) + let logStdRaw = net.logStdHead.linear(h2) + let logStd = logStdRaw.map(proc(v: float32): float32 = clamp(v, LOG_STD_MIN, LOG_STD_MAX)) + + if deterministic or isEvalMode(): + let actions = tanhT(mu) + return (actions: actions, logProb: 0.0'f32, lstm: lstmOut) + + # Reparameterization: z = mu + std * eps, action = tanh(z) + let std = logStd.map(proc(v: float32): float32 = exp(v)) + var actions = newTensor[float32](4) + var logProb = 0.0'f32 + let twoPiLog = 0.5'f32 * ln(2.0'f32 * PI.float32) + for i in 0 ..< 4: + let eps = gauss(0.0'f64, 1.0'f64).float32 + let z = mu[i] + std[i] * eps + actions[i] = tanh(z) + # log N(z | mu, std) - log(1 - tanh²(z) + eps) + let diff = (z - mu[i]) / std[i] + let logNorm = -0.5'f32 * diff * diff - ln(std[i]) - twoPiLog + let tanhCorr = ln(1.0'f32 - actions[i] * actions[i] + LOG_PROB_EPS) + logProb += logNorm - tanhCorr + + result = (actions: actions, logProb: logProb, lstm: lstmOut) + +# ── Critic forward ──────────────────────────────────────────────────────────── + +proc criticForward*(net: CriticNet; stateAction: Tensor[float32]; lstm: LSTMState): + tuple[q: float32; lstm: LSTMState] = + ## stateAction: [stateDim + actionDim], lstm: (h,c) each [hiddenDim] + let h1 = relu(net.fc1.linear(stateAction)) + let lstmOut = lstmStep(net.lstm, h1, lstm.h, lstm.c) + let h2 = relu(net.fc2.linear(lstmOut.h)) + let q = net.fc3.linear(h2) + result = (q: q[0], lstm: lstmOut) diff --git a/SAC_LSTM_Bot/tests/test_network.nim b/SAC_LSTM_Bot/tests/test_network.nim new file mode 100644 index 0000000..b62b71f --- /dev/null +++ b/SAC_LSTM_Bot/tests/test_network.nim @@ -0,0 +1,99 @@ +import unittest +import arraymancer +import std/[math, os] +import SAC_LSTM_Bot/network + +const + STATE_DIM = 20 + ACTION_DIM = 4 + +suite "ActorNet": + setup: + let actor = initActorNet(STATE_DIM) + let state = randomNormalTensor[float32](STATE_DIM) + let ls0 = zeroState(actor.hiddenDim) + + test "forward output shapes": + let (actions, lp, ls1) = actorForward(actor, state, ls0) + check actions.shape == [4] + check ls1.h.shape == [actor.hiddenDim] + check ls1.c.shape == [actor.hiddenDim] + check not lp.isNaN + + test "actions clamped in (-1, 1) — tanh output": + let (actions, _, _) = actorForward(actor, state, ls0) + for i in 0 ..< 4: + check actions[i] > -1.0'f32 + check actions[i] < 1.0'f32 + + test "hidden state propagates (h/c change after step)": + let (_, _, ls1) = actorForward(actor, state, ls0) + # h' should differ from zero init for non-trivial input + var hChanged = false + for i in 0 ..< actor.hiddenDim: + if abs(ls1.h[i] - ls0.h[i]) > 1e-7'f32: + hChanged = true + break + check hChanged + + test "deterministic mode: same input → same output": + let (a1, _, _) = actorForward(actor, state, ls0, deterministic = true) + let (a2, _, _) = actorForward(actor, state, ls0, deterministic = true) + for i in 0 ..< 4: + check abs(a1[i] - a2[i]) < 1e-7'f32 + + test "stochastic mode: outputs may differ (sampling)": + # Run many times; at least one pair should differ + let (a1, _, _) = actorForward(actor, state, ls0) + let (a2, _, _) = actorForward(actor, state, ls0) + var anyDiff = false + for i in 0 ..< 4: + if abs(a1[i] - a2[i]) > 1e-7'f32: + anyDiff = true + break + check anyDiff + +suite "CriticNet": + setup: + let critic = initCriticNet(STATE_DIM, ACTION_DIM) + let state = randomNormalTensor[float32](STATE_DIM) + let actions = randomNormalTensor[float32](ACTION_DIM) + let stateAct = concat(state, actions, axis = 0) + let ls0 = zeroState(critic.hiddenDim) + + test "forward output shape": + let (q, ls1) = criticForward(critic, stateAct, ls0) + check ls1.h.shape == [critic.hiddenDim] + check ls1.c.shape == [critic.hiddenDim] + # q is scalar float — just ensure it doesn't NaN + check not q.isNaN + + test "hidden state propagates": + let (_, ls1) = criticForward(critic, stateAct, ls0) + var hChanged = false + for i in 0 ..< critic.hiddenDim: + if abs(ls1.h[i] - ls0.h[i]) > 1e-7'f32: + hChanged = true + break + check hChanged + +suite "Dual critics": + test "two independent critics produce different Q values": + let c1 = initCriticNet(STATE_DIM, ACTION_DIM) + let c2 = initCriticNet(STATE_DIM, ACTION_DIM) + let sa = randomNormalTensor[float32](STATE_DIM + ACTION_DIM) + let ls = zeroState(c1.hiddenDim) + let (q1, _) = criticForward(c1, sa, ls) + let (q2, _) = criticForward(c2, sa, ls) + check abs(q1 - q2) > 1e-7'f32 + +suite "Hidden size config": + test "hidden size 128 works": + putEnv("SACLSTM_HIDDEN_SIZE", "128") + let actor = initActorNet(STATE_DIM) + let state = randomNormalTensor[float32](STATE_DIM) + let ls0 = zeroState(128) + let (actions, _, ls1) = actorForward(actor, state, ls0) + check actions.shape == [4] + check ls1.h.shape == [128] + putEnv("SACLSTM_HIDDEN_SIZE", "256")