feat(SAC_LSTM_Bot): main bot integration (#48)
This commit is contained in:
@@ -0,0 +1,104 @@
|
||||
## 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"
|
||||
Reference in New Issue
Block a user