Files
SirRoboGarage/SAC_LSTM_Bot/tests/test_integration.nim
T

105 lines
4.1 KiB
Nim

## Tests for integration.nim (#48) — assert-based, no framework.
## Covers: TrainingMsg channel round-trip (plain arrays through a channel),
## drain-then-train NewBattle/Shutdown handling, trainPass safety below canSample.
import arraymancer except Linear
import std/locks
import SAC_LSTM_Bot/integration
import SAC_LSTM_Bot/state # STATE_DIM
import SAC_LSTM_Bot/actions # ACTION_DIM
import SAC_LSTM_Bot/training # initSACTrainer
import SAC_LSTM_Bot/replay_buffer
# ── 1. TrainingMsg round-trips through a channel with arrays intact ──────────
block:
var ch: Channel[TrainingMsg]
ch.open(4)
var msg = TrainingMsg(kind: tmkTransition)
for i in 0 ..< STATE_DIM:
msg.state[i] = float32(i) * 0.5'f32
msg.nextState[i] = float32(i) * 2.0'f32
for i in 0 ..< ACTION_DIM:
msg.action[i] = float32(i) - 2.0'f32
msg.reward = -1.25'f32
msg.done = true
assert ch.trySend(msg)
assert ch.trySend(TrainingMsg(kind: tmkNewBattle, enemyId: 4242))
assert ch.trySend(TrainingMsg(kind: tmkShutdown))
let r1 = ch.recv()
assert r1.kind == tmkTransition, "first msg is a transition"
for i in 0 ..< STATE_DIM:
assert r1.state[i] == float32(i) * 0.5'f32, "state round-trip at " & $i
assert r1.nextState[i] == float32(i) * 2.0'f32, "nextState round-trip at " & $i
for i in 0 ..< ACTION_DIM:
assert r1.action[i] == float32(i) - 2.0'f32, "action round-trip at " & $i
assert r1.reward == -1.25'f32 and r1.done
let r2 = ch.recv()
assert r2.kind == tmkNewBattle and r2.enemyId == 4242
let r3 = ch.recv()
assert r3.kind == tmkShutdown
ch.close()
# closed + empty -> tryRecv reports no data (Nim 2.2: recv would block forever)
let (ok4, _) = ch.tryRecv()
assert not ok4, "closed channel must report dataAvailable=false"
echo "PASS TrainingMsg channel round-trip"
# ── 2. NewBattle clears only when the opponent changes; Shutdown stops ────────
block:
var st: TrainState
st.trainer = initSACTrainer(STATE_DIM, ACTION_DIM)
st.buf = newReplayBuffer(64, STATE_DIM, ACTION_DIM, burnIn = 2, trainWindow = 3)
assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 7))
assert st.lastEnemyId == 7
for i in 0 ..< 5:
var m = TrainingMsg(kind: tmkTransition)
m.reward = float32(i)
assert handleTrainingMsg(st, m)
assert st.buf.len == 5, "transitions stored"
# Same opponent -> buffer kept (Q12a).
assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 7))
assert st.buf.len == 5, "same opponent must NOT clear"
# Opponent changed -> clear.
assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 8))
assert st.buf.len == 0, "opponent change must clear"
assert st.lastEnemyId == 8
# Shutdown stops the caller's loop.
assert not handleTrainingMsg(st, TrainingMsg(kind: tmkShutdown))
echo "PASS drain-then-train NewBattle/Shutdown"
# ── 3. trainPass is a safe no-op below canSample (no steps, no publish) ───────
block:
var st: TrainState
st.trainer = initSACTrainer(STATE_DIM, ACTION_DIM)
st.buf = newReplayBuffer(64, STATE_DIM, ACTION_DIM, burnIn = 2, trainWindow = 3)
for i in 0 ..< 4:
assert handleTrainingMsg(st, TrainingMsg(kind: tmkTransition))
trainPass(st, 4) # 4 < burnIn+trainWindow = 5
assert st.stepCount == 0, "no gradient steps below canSample"
echo "PASS trainPass no-op below canSample"
# ── 4. Flat snapshot layout: pack/unpack round-trips weights exactly ──────────
block:
let t0 = initSACTrainer(STATE_DIM, ACTION_DIM)
let fs = packFull(t0)
assert fs.data.len == actorSize(t0.actor.hiddenDim) + 4 * criticSize(t0.actor.hiddenDim) + 1
let (a, c1, c2, tc1, tc2, alpha) = unpackFull(fs)
assert a.hiddenDim == t0.actor.hiddenDim
assert c1.fc3.b.shape[0] == 1
let fw = a.muHead.w.flatten()
let fw0 = t0.actor.muHead.w.flatten()
for i in 0 ..< fw.size:
assert fw[i] == fw0[i], "actor mu weights round-trip"
for i in 0 ..< c2.lstm.bCombined.size:
assert c2.lstm.bCombined[i] == t0.critic2.lstm.bCombined[i], "critic lstm bias round-trip"
assert alpha == t0.alpha()
discard tc1
echo "PASS flat snapshot pack/unpack round-trip"