feat(SAC_LSTM_Bot): LSTM network module (#41)
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -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)
|
||||||
@@ -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")
|
||||||
Reference in New Issue
Block a user