b509195ee9
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
100 lines
3.2 KiB
Nim
100 lines
3.2 KiB
Nim
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")
|