Files
SirRoboGarage/SAC_LSTM_Bot/tests/test_integration.nim
T
SirStone df256b4d3e feat(SAC_LSTM_Bot): training harness (#49)
sac_train.sh orchestrates chunked self-play via tools/training_runner/
RunTraining.java: weighted opponent sampling per chunk, deterministic
eval (SACLSTM_EVAL_MODE=1) every N chunks with win-rate tracking, best
checkpoint (weights/sac_best.zip) by eval score, crash-restart loop on
the runner's liveness detection.

Supporting changes:
- integration.nim: opponentKey() keys the NewBattle buffer-clear rule on
  getBotName(id) with numeric-id fallback (#49 Q14 follow-up);
  bumpRoundCounter() emits the per-round liveness signal.
- SAC_LSTM_Bot.nim: onRoundEnded -> bumpRoundCounter().
- RunTraining.java: BOT_NAME env parameterizes result matching
  (default PPO_Bot, unchanged behavior for PPO).
- Launch packaging: root SAC_LSTM_Bot.json + .sh for the booter;
  src json name aligned to 'SAC_LSTM_Bot' so self-reported identity
  matches the booted identity (mismatch = runner connect timeout).
2026-08-21 21:51:27 +02:00

125 lines
5.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, json, os]
import tankroyale_botapi # updateBotNames: seed the vendored id->name table
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 NAME changes; Shutdown stops ───
block:
# Seed the API's id->name table (v1.0.1) for name-based identity (#49).
updateBotNames(parseJson(
"""{"bots":[{"id":7,"name":"Corners"},{"id":8,"name":"Crazy"}]}"""))
assert opponentKey(7) == "Corners", "known id resolves to name"
assert opponentKey(99) == "99", "unknown id falls back to numeric string"
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.lastEnemyKey == "Corners"
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.lastEnemyKey == "Crazy"
# Nameless window (pre-BotListUpdate) is a distinct key.
assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 99))
assert st.lastEnemyKey == "99" and st.buf.len == 0
# Shutdown stops the caller's loop.
assert not handleTrainingMsg(st, TrainingMsg(kind: tmkShutdown))
echo "PASS name-keyed NewBattle clear / Shutdown"
# ── 2b. bumpRoundCounter increments per call, cold-starts at 1 ────────────────
block:
let tmp = getTempDir() / "sac_test_weights_" & $getCurrentProcessId()
putEnv("SACLSTM_WEIGHTS_PATH", tmp / "sac_latest.zip")
bumpRoundCounter()
bumpRoundCounter()
assert readFile(tmp / "round_counter.txt") == "2", "counter increments per round"
echo "PASS bumpRoundCounter"
# ── 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"