feat(SAC_LSTM_Bot): LSTM network module (#41)

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-08-20 23:38:17 +02:00
parent f130bf1254
commit add3e34926
2 changed files with 260 additions and 0 deletions
+161
View File
@@ -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)