a07e5305f5
Levers 3, 4, 1 of the #57 sign-off (execution order 3->4->1), tracked in #59. - Lever 3 (#59): one JSONL line per trainPass in training_metrics.jsonl with exactly the scalars sacUpdate already exposes (SACMetrics: critic/actor/alpha losses + alpha, averaged per pass) plus epoch, buffer size (replay_buffer.len), cumulative steps and drained count. No trainer change needed. - Lever 4 (#59): sendTrainingMsg drops all training input while SACLSTM_EVAL_MODE=1 (existing #49 harness mechanism) — eval battles can neither pollute the replay buffer nor trigger gradient updates; one-time stderr notice at bot init. - Lever 1 (#59): sac_train.sh evaluates every SAC_EVAL_OPPONENTS entry per cycle (results carry opponent name in eval_log.jsonl); best-gating now uses a composite = mean over opponents of the last-5-evals moving average per opponent. best_score.txt format change: float composite replaces the single-opponent integer win rate semantics (retired). - Tests: metricsLine JSONL scalars + eval-mode suppression asserts. Refs: #59, #57
151 lines
6.4 KiB
Nim
151 lines
6.4 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"
|
|
|
|
# ── 5. Lever 3 (#59): metricsLine emits exactly the exposed trainer scalars ───
|
|
block:
|
|
let line = metricsLine(1787394115.123, 42, 500, 20, 20,
|
|
SACMetrics(criticLoss: 0.5'f32, actorLoss: -1.5'f32,
|
|
alphaLoss: 0.25'f32, alpha: 2.0'f32))
|
|
let j = parseJson(line) # throws on malformed JSONL
|
|
assert j["steps"].getInt() == 42 and j["buffer_size"].getInt() == 500
|
|
assert j["drained"].getInt() == 20 and j["grad_steps"].getInt() == 20
|
|
assert abs(j["critic_loss"].getFloat() - 0.5) < 1e-3
|
|
assert abs(j["actor_loss"].getFloat() + 1.5) < 1e-3
|
|
assert abs(j["alpha_loss"].getFloat() - 0.25) < 1e-3
|
|
assert abs(j["alpha"].getFloat() - 2.0) < 1e-3
|
|
assert j["epoch"].getFloat() > 1e9
|
|
echo "PASS metricsLine JSONL scalars"
|
|
|
|
# ── 6. Lever 4 (#59): SACLSTM_EVAL_MODE=1 suppresses training input ───────────
|
|
block:
|
|
putEnv("SACLSTM_EVAL_MODE", "1")
|
|
# Gate fires before any channel traffic: false = dropped, nothing enqueued.
|
|
assert not sendTrainingMsg(TrainingMsg(kind: tmkTransition)),
|
|
"eval mode must drop transitions"
|
|
assert not sendTrainingMsg(TrainingMsg(kind: tmkNewBattle, enemyId: 7)),
|
|
"eval mode must drop NewBattle (no buffer clears from eval)"
|
|
delEnv("SACLSTM_EVAL_MODE")
|
|
echo "PASS eval-mode training-input suppression"
|