Files
SirRoboGarage/SAC_LSTM_Bot_garage/tests/test_network.nim
T

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")