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,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