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