chore: rename libs→common_libs, all bot dirs to _garage suffix, fix all path refs
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,2 @@
|
||||
switch("path", "../src")
|
||||
switch("path", "../../libs")
|
||||
@@ -0,0 +1,97 @@
|
||||
import unittest
|
||||
import arraymancer
|
||||
import std/math
|
||||
import SAC_LSTM_Bot/actions
|
||||
|
||||
proc makeOutput(a0, a1, a2, a3: float): Tensor[float32] =
|
||||
result = newTensor[float32](4)
|
||||
result[0] = a0.float32
|
||||
result[1] = a1.float32
|
||||
result[2] = a2.float32
|
||||
result[3] = a3.float32
|
||||
|
||||
suite "mapActions":
|
||||
|
||||
test "ACTION_DIM is 4":
|
||||
check ACTION_DIM == 4
|
||||
|
||||
# Speed-aware turn rate
|
||||
test "turn: output +1 at speed 0 -> +10 degrees":
|
||||
let m = mapActions(makeOutput(1.0, 0.0, 0.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.turnRate - 10.0) < 1e-6
|
||||
|
||||
test "turn: output -1 at speed 0 -> -10 degrees":
|
||||
let m = mapActions(makeOutput(-1.0, 0.0, 0.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.turnRate - (-10.0)) < 1e-6
|
||||
|
||||
test "turn: output +1 at speed 8 -> +4 degrees":
|
||||
let m = mapActions(makeOutput(1.0, 0.0, 0.0, -1.0), 8.0, 0.0)
|
||||
check abs(m.turnRate - 4.0) < 1e-6
|
||||
|
||||
test "turn: output -1 at speed 8 -> -4 degrees":
|
||||
let m = mapActions(makeOutput(-1.0, 0.0, 0.0, -1.0), 8.0, 0.0)
|
||||
check abs(m.turnRate - (-4.0)) < 1e-6
|
||||
|
||||
test "turn: output 0 -> 0 regardless of speed":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, -1.0), 5.0, 0.0)
|
||||
check abs(m.turnRate) < 1e-6
|
||||
|
||||
# Acceleration range
|
||||
test "accel: output -1 -> -2.0":
|
||||
let m = mapActions(makeOutput(0.0, -1.0, 0.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.acceleration - (-2.0)) < 1e-6
|
||||
|
||||
test "accel: output +1 -> +1.0":
|
||||
let m = mapActions(makeOutput(0.0, 1.0, 0.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.acceleration - 1.0) < 1e-6
|
||||
|
||||
test "accel: output 0 -> midpoint -0.5":
|
||||
# value*1.5 - 0.5 at value=0 -> -0.5 (correct midpoint between -2 and +1)
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.acceleration - (-0.5)) < 1e-6
|
||||
|
||||
# Gun turn rate
|
||||
test "gun turn: output +1 -> +20 degrees":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 1.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.gunTurnRate - 20.0) < 1e-6
|
||||
|
||||
test "gun turn: output -1 -> -20 degrees":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, -1.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.gunTurnRate - (-20.0)) < 1e-6
|
||||
|
||||
test "gun turn: output 0 -> 0":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.gunTurnRate) < 1e-6
|
||||
|
||||
# Fire threshold
|
||||
test "fire: output -0.5 -> no fire (firePower == 0)":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, -0.5), 0.0, 0.0)
|
||||
check m.firePower == 0.0
|
||||
|
||||
test "fire: output 0.0 -> no fire (boundary, not positive)":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 0.0), 0.0, 0.0)
|
||||
check m.firePower == 0.0
|
||||
|
||||
test "fire: output +0.5 -> fire with correct power":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 0.5), 0.0, 0.0)
|
||||
# 0.5 * 2.9 + 0.1 = 1.55
|
||||
check abs(m.firePower - 1.55) < 1e-5
|
||||
|
||||
test "fire: output +1.0 -> fire power near 3.0":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 1.0), 0.0, 0.0)
|
||||
check abs(m.firePower - 3.0) < 1e-5
|
||||
|
||||
test "fire: output +1.0 but gunHeat > 0 -> no fire":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 1.0), 0.0, 1.5)
|
||||
check m.firePower == 0.0
|
||||
|
||||
test "fire: output +0.001 (just above 0) -> fires with power near 0.1":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 0.001), 0.0, 0.0)
|
||||
check m.firePower > 0.0
|
||||
check m.firePower < 0.2
|
||||
|
||||
# Speed-aware turn with negative speed (reverse)
|
||||
test "turn: speed -8 (reversing) -> same magnitude as speed +8":
|
||||
let fwd = mapActions(makeOutput(1.0, 0.0, 0.0, -1.0), 8.0, 0.0)
|
||||
let rev = mapActions(makeOutput(1.0, 0.0, 0.0, -1.0), -8.0, 0.0)
|
||||
check abs(fwd.turnRate - rev.turnRate) < 1e-6
|
||||
@@ -0,0 +1,150 @@
|
||||
## 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"
|
||||
@@ -0,0 +1,99 @@
|
||||
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")
|
||||
@@ -0,0 +1,156 @@
|
||||
## Tests for replay_buffer.nim — assert-based, no framework.
|
||||
|
||||
import arraymancer
|
||||
import std/sequtils
|
||||
import SAC_LSTM_Bot/replay_buffer
|
||||
import SAC_LSTM_Bot/state # STATE_DIM
|
||||
|
||||
const
|
||||
S_DIM = STATE_DIM # 35
|
||||
A_DIM = 4
|
||||
|
||||
proc makeTrans(reward: float32; done: bool): Transition =
|
||||
Transition(
|
||||
state: zeros[float32](S_DIM),
|
||||
action: zeros[float32](A_DIM),
|
||||
reward: reward,
|
||||
nextState: zeros[float32](S_DIM),
|
||||
done: done
|
||||
)
|
||||
|
||||
# ── 1. len and canSample on empty buffer ──────────────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(100, S_DIM, A_DIM, burnIn = 8, trainWindow = 16)
|
||||
assert buf.len == 0, "empty len"
|
||||
assert not buf.canSample, "empty canSample"
|
||||
echo "PASS empty buffer"
|
||||
|
||||
# ── 2. len grows, canSample becomes true ──────────────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(100, S_DIM, A_DIM, burnIn = 2, trainWindow = 3)
|
||||
for i in 0 ..< 4:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
assert buf.len == 4
|
||||
assert not buf.canSample, "needs 5 to sample"
|
||||
buf.add(makeTrans(99, false))
|
||||
assert buf.len == 5
|
||||
assert buf.canSample, "5 transitions, seqLen=5 → canSample"
|
||||
echo "PASS len + canSample"
|
||||
|
||||
# ── 3. Ring buffer wraps at capacity ─────────────────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(10, S_DIM, A_DIM, burnIn = 2, trainWindow = 3)
|
||||
for i in 0 ..< 15:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
assert buf.len == 10, "wraps at capacity, len stays 10"
|
||||
echo "PASS wrap at capacity"
|
||||
|
||||
# ── 4. Sampled sequences have correct length ──────────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(200, S_DIM, A_DIM, burnIn = 4, trainWindow = 6)
|
||||
for i in 0 ..< 50:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
let seqs = buf.sampleSequences(8)
|
||||
assert seqs.len > 0, "should have valid starts"
|
||||
for s in seqs:
|
||||
assert s.burnIn.len == 4, "burnIn len"
|
||||
assert s.train.len == 6, "train len"
|
||||
echo "PASS sequence length"
|
||||
|
||||
# ── 5. Burn-in / train split is correct ───────────────────────────────────────
|
||||
block:
|
||||
# Fill with distinct rewards so we can identify positions
|
||||
var buf = newReplayBuffer(200, S_DIM, A_DIM, burnIn = 3, trainWindow = 4)
|
||||
for i in 0 ..< 30:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
let seqs = buf.sampleSequences(1)
|
||||
assert seqs.len == 1
|
||||
let s = seqs[0]
|
||||
# The 4th reward of the sequence (index 3) must match s.train[0].reward
|
||||
# and s.burnIn[2].reward must be s.burnIn[2].reward (just check no overlap)
|
||||
let allRewards = s.burnIn.mapIt(it.reward) & s.train.mapIt(it.reward)
|
||||
# consecutive integer rewards → each must be strictly increasing by 1
|
||||
var ok = true
|
||||
for i in 1 ..< allRewards.len:
|
||||
if allRewards[i] != allRewards[i-1] + 1.0f32:
|
||||
ok = false
|
||||
break
|
||||
assert ok, "burn-in and train must form a contiguous sequence"
|
||||
echo "PASS burn-in/train split"
|
||||
|
||||
# ── 6. Sequences never cross a done=true boundary ─────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(200, S_DIM, A_DIM, burnIn = 2, trainWindow = 3)
|
||||
# Episode 1: transitions 0..4 (done at index 4)
|
||||
for i in 0 ..< 4:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
buf.add(makeTrans(99, true)) # battle end at index 4
|
||||
# Episode 2: transitions 5..14
|
||||
for i in 5 ..< 15:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
|
||||
let seqs = buf.sampleSequences(20)
|
||||
for s in seqs:
|
||||
# No transition in burnIn (except the last) or train (except the last)
|
||||
# may have done=true, since that would mean the next step crosses a boundary.
|
||||
for i in 0 ..< s.burnIn.len - 1:
|
||||
assert not s.burnIn[i].done, "done in middle of burnIn"
|
||||
for i in 0 ..< s.train.len - 1:
|
||||
assert not s.train[i].done, "done in middle of train"
|
||||
# The join between burnIn and train must not cross a done=true
|
||||
if s.burnIn.len > 0:
|
||||
assert not s.burnIn[^1].done, "done at end of burnIn crosses boundary to train"
|
||||
echo "PASS no cross-boundary sequences"
|
||||
|
||||
# ── 7. canSample false when buffer too small ───────────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(100, S_DIM, A_DIM, burnIn = 8, trainWindow = 16)
|
||||
for i in 0 ..< 23: # seqLen = 24; 23 < 24
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
assert not buf.canSample, "23 < 24 seqLen"
|
||||
buf.add(makeTrans(23, false))
|
||||
assert buf.canSample, "24 == seqLen"
|
||||
echo "PASS canSample threshold"
|
||||
|
||||
# ── 8. Empty buffer doesn't crash on sample ────────────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(100, S_DIM, A_DIM, burnIn = 8, trainWindow = 16)
|
||||
let seqs = buf.sampleSequences(4)
|
||||
assert seqs.len == 0, "empty buffer → empty result"
|
||||
echo "PASS empty sample"
|
||||
|
||||
# ── 9. Single episode (no done except at very end) ────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(200, S_DIM, A_DIM, burnIn = 3, trainWindow = 5)
|
||||
for i in 0 ..< 19:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
buf.add(makeTrans(19, true)) # last
|
||||
let seqs = buf.sampleSequences(5)
|
||||
assert seqs.len > 0, "should find valid starts"
|
||||
for s in seqs:
|
||||
assert s.burnIn.len == 3
|
||||
assert s.train.len == 5
|
||||
echo "PASS single episode"
|
||||
|
||||
# ── 10. Multiple short episodes: all boundaries respected ─────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(300, S_DIM, A_DIM, burnIn = 2, trainWindow = 3)
|
||||
# 10 episodes of 5 transitions each (done at end of each episode)
|
||||
var reward = 0'f32
|
||||
for ep in 0 ..< 10:
|
||||
for i in 0 ..< 4:
|
||||
buf.add(makeTrans(reward, false))
|
||||
reward += 1
|
||||
buf.add(makeTrans(reward, true)) # battle end
|
||||
reward += 1
|
||||
|
||||
let seqs = buf.sampleSequences(30)
|
||||
assert seqs.len > 0
|
||||
for s in seqs:
|
||||
# No done in the middle of any sequence
|
||||
let all = s.burnIn & s.train
|
||||
for i in 0 ..< all.len - 1:
|
||||
assert not all[i].done, "boundary crossed in multi-episode test"
|
||||
echo "PASS multiple episodes"
|
||||
|
||||
echo "ALL TESTS PASSED"
|
||||
@@ -0,0 +1,109 @@
|
||||
## Assert-based tests for rewards.nim.
|
||||
## Run: nim c -r tests/test_rewards.nim
|
||||
|
||||
import std/[math, strformat]
|
||||
import SAC_LSTM_Bot/rewards
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
# ── computeReward ─────────────────────────────────────────────────────────────
|
||||
|
||||
block damageInflicted:
|
||||
# p=1: 1.25 * (6*1 - 2) = 5 (lever 2 aggression mult, #59)
|
||||
check abs(computeReward(damageInflicted = 1.0) - 5.0) < 1e-9, "p=1 damage = +5"
|
||||
# p=3: 1.25 * (6*3 - 2) = 20
|
||||
check abs(computeReward(damageInflicted = 3.0) - 20.0) < 1e-9, "p=3 damage = +20"
|
||||
# low-power spam stays unprofitable: 1.25*(6*0.1-2) < 0
|
||||
check computeReward(damageInflicted = 0.1) < 0.0, "p=0.1 spam still negative"
|
||||
|
||||
block damageReceived:
|
||||
# p_e=1: -(6*1 - 2) = -4 (unchanged by lever 2)
|
||||
check abs(computeReward(damageReceived = 1.0) - (-4.0)) < 1e-9, "p_e=1 received = -4"
|
||||
# p_e=3: -(6*3 - 2) = -16
|
||||
check abs(computeReward(damageReceived = 3.0) - (-16.0)) < 1e-9, "p_e=3 received = -16"
|
||||
|
||||
block hitBonus:
|
||||
# flat +0.5 per landed shot: p=1 hit -> 5.0 + 0.5
|
||||
check abs(computeReward(damageInflicted = 1.0, hitCount = 1) - 5.5) < 1e-9,
|
||||
"p=1 hit = +5.5"
|
||||
# two hits in one step: 1.25*(6*2-2) + 2*0.5 = 12.5 + 1.0 = 13.5
|
||||
check abs(computeReward(damageInflicted = 2.0, hitCount = 2) - 13.5) < 1e-9,
|
||||
"two hits = +13.5"
|
||||
|
||||
block ramTaken:
|
||||
# flat per victim collision (#59)
|
||||
check abs(computeReward(ramTakenCount = 1) - (-3.0)) < 1e-9, "ram taken x1 = -3"
|
||||
check abs(computeReward(ramTakenCount = 2) - (-6.0)) < 1e-9, "ram taken x2 = -6"
|
||||
|
||||
block chargeDeterrent:
|
||||
# zero-damage case at half threshold depth: -2 * (1 - 0.06/0.12) = -1
|
||||
let rHalf = computeReward(enemyDistFrac = 0.06)
|
||||
check abs(rHalf - (-1.0)) < 1e-9, "charge at frac 0.06 = -1"
|
||||
# at zero distance: full ceiling
|
||||
check abs(computeReward(enemyDistFrac = 0.0) - (-2.0)) < 1e-9, "charge at frac 0 = -2"
|
||||
# at/beyond threshold and no-contact sentinel: no penalty
|
||||
check abs(computeReward(enemyDistFrac = 0.12)) < 1e-9, "at threshold = 0"
|
||||
check abs(computeReward(enemyDistFrac = 0.5)) < 1e-9, "beyond threshold = 0"
|
||||
check abs(computeReward(enemyDistFrac = 2.0)) < 1e-9, "no-contact sentinel = 0"
|
||||
# suppressed while dealing damage that step (fighting back at close range is fine)
|
||||
let rFight = computeReward(damageInflicted = 1.0, enemyDistFrac = 0.06)
|
||||
check abs(rFight - 5.0) < 1e-9, "dealing damage cancels charge penalty"
|
||||
|
||||
block wallHit:
|
||||
check abs(computeReward(wallHitTicks = 1) - (-5.0)) < 1e-9, "1 wall tick = -5"
|
||||
|
||||
block wastedShot:
|
||||
# p=2: -0.1 * 2 = -0.2
|
||||
check abs(computeReward(wastedShotPower = 2.0) - (-0.2)) < 1e-9, "missed p=2 = -0.2"
|
||||
|
||||
block winLoss:
|
||||
# terminal terms stay dominant over shaping (#59 scale discipline)
|
||||
check abs(computeReward(win = true) - 20.0) < 1e-9, "win = +20"
|
||||
check abs(computeReward(loss = true) - (-10.0)) < 1e-9, "loss = -10"
|
||||
|
||||
# ── RewardNormalizer cold start ───────────────────────────────────────────────
|
||||
|
||||
block coldStart:
|
||||
var rn: RewardNormalizer
|
||||
# 0 samples
|
||||
let v0 = rn.normalize(99.0)
|
||||
check not isNaN(v0), "0 samples: not NaN"
|
||||
check classify(v0) != fcInf and classify(v0) != fcNegInf, "0 samples: not inf"
|
||||
check abs(v0) < 1e-9, "0 samples: returns 0"
|
||||
# 1 sample (variance undefined)
|
||||
rn.update(5.0)
|
||||
let v1 = rn.normalize(5.0)
|
||||
check not isNaN(v1), "1 sample: not NaN"
|
||||
check classify(v1) != fcInf and classify(v1) != fcNegInf, "1 sample: not inf"
|
||||
check abs(v1) < 1e-9, "1 sample: returns 0"
|
||||
|
||||
# ── Running normalization convergence ─────────────────────────────────────────
|
||||
|
||||
block convergence:
|
||||
var rn: RewardNormalizer
|
||||
# Feed 1000 identical samples of 5.0 — mean=5.0, std=0 → normalizer returns ~0
|
||||
for _ in 0 ..< 1000:
|
||||
rn.update(5.0)
|
||||
let v = rn.normalize(5.0)
|
||||
check not isNaN(v), "convergence: not NaN"
|
||||
check classify(v) != fcInf and classify(v) != fcNegInf, "convergence: not inf"
|
||||
# (5 - 5) / (0 + eps) = 0
|
||||
check abs(v) < 1e-6, "convergence to mean: normalized ≈ 0"
|
||||
|
||||
block knownMeanStd:
|
||||
# Insert samples -1 and +1 repeatedly → mean=0, std=1
|
||||
var rn: RewardNormalizer
|
||||
for _ in 0 ..< 500:
|
||||
rn.update(-1.0)
|
||||
rn.update( 1.0)
|
||||
# normalize(1.0) ≈ (1 - 0) / (1 + eps) ≈ 1
|
||||
let vPos = rn.normalize(1.0)
|
||||
check abs(vPos - 1.0) < 1e-4, &"normalize(+1) ≈ +1, got {vPos}"
|
||||
let vNeg = rn.normalize(-1.0)
|
||||
check abs(vNeg - (-1.0)) < 1e-4, &"normalize(-1) ≈ -1, got {vNeg}"
|
||||
let vMid = rn.normalize(0.0)
|
||||
check abs(vMid) < 1e-4, &"normalize(0) ≈ 0, got {vMid}"
|
||||
|
||||
echo "test_rewards: all passed"
|
||||
@@ -0,0 +1,87 @@
|
||||
## Tests for state.nim — assert-based, no framework.
|
||||
|
||||
import std/math
|
||||
import arraymancer
|
||||
import SAC_LSTM_Bot/state
|
||||
|
||||
proc makeBase(): GameState =
|
||||
result.arenaWidth = 1200.0
|
||||
result.arenaHeight = 800.0
|
||||
result.x = 600.0; result.y = 400.0
|
||||
result.direction = 90.0; result.speed = 4.0
|
||||
result.energy = 50.0
|
||||
result.gunDirection = 90.0; result.gunHeat = 0.5
|
||||
|
||||
proc allInRange(t: Tensor[float32]): bool =
|
||||
for v in t:
|
||||
if v < -1.01f32 or v > 1.01f32: return false
|
||||
true
|
||||
|
||||
proc hasNaN(t: Tensor[float32]): bool =
|
||||
for v in t:
|
||||
if v.float64.isNaN: return true
|
||||
false
|
||||
|
||||
# 1. Correct shape
|
||||
block:
|
||||
let gs = makeBase()
|
||||
let t = buildState(gs)
|
||||
assert t.shape[0] == STATE_DIM, "wrong dim: " & $t.shape[0]
|
||||
echo "PASS shape"
|
||||
|
||||
# 2. All values in [-1, 1] for typical input
|
||||
block:
|
||||
var gs = makeBase()
|
||||
gs.hasContact = true
|
||||
gs.enemy = EnemyData(x: 800.0, y: 300.0, direction: 45.0, speed: 6.0,
|
||||
energy: 80.0, hasFired: true, lastFirePower: 2.0,
|
||||
prevSpeed: 5.0, prevDirection: 40.0, hasPrevScan: true)
|
||||
gs.bulletCount = 1
|
||||
gs.bullets[0] = BulletData(x: 610.0, y: 410.0, power: 1.0)
|
||||
gs.ticksSinceLastScan = 10
|
||||
let t = buildState(gs)
|
||||
assert not hasNaN(t), "NaN in tensor"
|
||||
assert allInRange(t), "value out of [-1,1]"
|
||||
echo "PASS range"
|
||||
|
||||
# 3. Missing scan → valid tensor, no NaN, zeros for enemy fields
|
||||
block:
|
||||
let gs = makeBase() # hasContact = false
|
||||
let t = buildState(gs)
|
||||
assert not hasNaN(t), "NaN with no scan"
|
||||
for i in 7 .. 17:
|
||||
assert t[i] == 0.0f32, "enemy slot " & $i & " should be 0"
|
||||
echo "PASS no-scan zeros"
|
||||
|
||||
# 4. Bullet tracking: 0, 1, 2, 3 bullets
|
||||
block:
|
||||
for n in 0 .. 3:
|
||||
var gs = makeBase()
|
||||
gs.bulletCount = n
|
||||
for i in 0 ..< n:
|
||||
gs.bullets[i] = BulletData(x: float64(600 + i * 10), y: float64(400 + i * 5), power: 1.0)
|
||||
let t = buildState(gs)
|
||||
assert not hasNaN(t), "NaN with " & $n & " bullets"
|
||||
# slots beyond bulletCount must be 0
|
||||
for i in n ..< 3:
|
||||
let base = 22 + i * 4
|
||||
for j in 0 ..< 4:
|
||||
assert t[base + j] == 0.0f32, "bullet slot " & $i & " slot " & $j & " should be 0"
|
||||
echo "PASS bullet tracking 0-3"
|
||||
|
||||
# 5. Scan staleness increments and clamps
|
||||
block:
|
||||
var gs = makeBase()
|
||||
gs.hasContact = true
|
||||
gs.ticksSinceLastScan = 0
|
||||
assert buildState(gs)[34] == 0.0f32, "staleness at 0"
|
||||
gs.ticksSinceLastScan = 15
|
||||
let mid = buildState(gs)[34]
|
||||
assert mid > 0.0f32 and mid < 1.0f32, "staleness mid not in (0,1)"
|
||||
gs.ticksSinceLastScan = 30
|
||||
assert buildState(gs)[34] == 1.0f32, "staleness at 30"
|
||||
gs.ticksSinceLastScan = 60
|
||||
assert buildState(gs)[34] == 1.0f32, "staleness clamp at 60"
|
||||
echo "PASS staleness"
|
||||
|
||||
echo "ALL TESTS PASSED"
|
||||
@@ -0,0 +1,159 @@
|
||||
## test_training.nim — stdlib unittest for SAC training module.
|
||||
import unittest
|
||||
import arraymancer
|
||||
import std/[math, random, os]
|
||||
import SAC_LSTM_Bot/network
|
||||
import SAC_LSTM_Bot/weights
|
||||
import SAC_LSTM_Bot/replay_buffer
|
||||
import SAC_LSTM_Bot/training
|
||||
|
||||
const
|
||||
STATE_DIM = 10
|
||||
ACTION_DIM = 4
|
||||
HIDDEN = 16 # small for speed; override via env not needed in tests
|
||||
|
||||
proc makeTrainer(): SACTrainer =
|
||||
## Small trainer for tests — override hidden size via env before calling.
|
||||
putEnv("SACLSTM_HIDDEN_SIZE", "16")
|
||||
initSACTrainer(STATE_DIM, ACTION_DIM)
|
||||
|
||||
proc makeBuffer(): ReplayBuffer =
|
||||
newReplayBuffer(capacity = 2000, stateDim = STATE_DIM, actionDim = ACTION_DIM,
|
||||
burnIn = 4, trainWindow = 8)
|
||||
|
||||
proc randState(): Tensor[float32] =
|
||||
randomNormalTensor[float32](STATE_DIM)
|
||||
|
||||
proc randAction(): Tensor[float32] =
|
||||
randomTensor[float32](ACTION_DIM, 1.0'f32) *. 2.0'f32 -. 1.0'f32 # uniform [-1,1]
|
||||
|
||||
proc fillBuffer(buf: var ReplayBuffer; n: int; donePeriod = 20) =
|
||||
for i in 0 ..< n:
|
||||
let t = Transition(
|
||||
state: randState(),
|
||||
action: randAction(),
|
||||
reward: rand(-1.0'f32 .. 1.0'f32),
|
||||
nextState: randState(),
|
||||
done: (i mod donePeriod == donePeriod - 1))
|
||||
buf.add(t)
|
||||
|
||||
proc isFinite(x: float32): bool =
|
||||
not (x != x) and x < Inf and x > -Inf # not NaN and not Inf
|
||||
|
||||
suite "SACTrainer — basic update":
|
||||
|
||||
setup:
|
||||
randomize(42)
|
||||
var trainer = makeTrainer()
|
||||
var buf = makeBuffer()
|
||||
fillBuffer(buf, 1000)
|
||||
let seqs = buf.sampleSequences(4)
|
||||
|
||||
test "sampleSequences returns non-empty batch":
|
||||
check seqs.len > 0
|
||||
|
||||
test "sacUpdate returns finite losses":
|
||||
let m = sacUpdate(trainer, seqs)
|
||||
check isFinite(m.criticLoss)
|
||||
check isFinite(m.actorLoss)
|
||||
check isFinite(m.alphaLoss)
|
||||
|
||||
test "alpha stays positive after update":
|
||||
var t2 = trainer
|
||||
discard sacUpdate(t2, seqs)
|
||||
check t2.alpha() > 0.0'f32
|
||||
|
||||
test "critic1 weights change after update":
|
||||
let w_before = trainer.critic1.fc3.w.clone()
|
||||
discard sacUpdate(trainer, seqs)
|
||||
let w_after = trainer.critic1.fc3.w
|
||||
var changed = false
|
||||
for i in 0 ..< w_before.size:
|
||||
if abs(w_before.unsafe_raw_offset[i] - w_after.unsafe_raw_offset[i]) > 1e-10'f32:
|
||||
changed = true
|
||||
break
|
||||
check changed
|
||||
|
||||
test "critic2 weights change after update":
|
||||
let w_before = trainer.critic2.fc3.w.clone()
|
||||
discard sacUpdate(trainer, seqs)
|
||||
let w_after = trainer.critic2.fc3.w
|
||||
var changed = false
|
||||
for i in 0 ..< w_before.size:
|
||||
if abs(w_before.unsafe_raw_offset[i] - w_after.unsafe_raw_offset[i]) > 1e-10'f32:
|
||||
changed = true
|
||||
break
|
||||
check changed
|
||||
|
||||
test "actor weights change after update":
|
||||
let w_before = trainer.actor.fc1.w.clone()
|
||||
discard sacUpdate(trainer, seqs)
|
||||
let w_after = trainer.actor.fc1.w
|
||||
var changed = false
|
||||
for i in 0 ..< w_before.size:
|
||||
if abs(w_before.unsafe_raw_offset[i] - w_after.unsafe_raw_offset[i]) > 1e-10'f32:
|
||||
changed = true
|
||||
break
|
||||
check changed
|
||||
|
||||
test "soft target update: target moves toward critic":
|
||||
## After update, target fc3.w should be closer to critic1.fc3.w than before.
|
||||
let targetBefore = trainer.targetCritic1.fc3.w.clone()
|
||||
let criticW = trainer.critic1.fc3.w.clone()
|
||||
discard sacUpdate(trainer, seqs)
|
||||
let targetAfter = trainer.targetCritic1.fc3.w
|
||||
|
||||
# Distance before: ||targetBefore - criticW||
|
||||
var distBefore = 0.0'f32
|
||||
for i in 0 ..< targetBefore.size:
|
||||
let d = targetBefore.unsafe_raw_offset[i] - criticW.unsafe_raw_offset[i]
|
||||
distBefore += d * d
|
||||
|
||||
# Distance after: ||targetAfter - criticW_new|| (critic changed too, use original for ref)
|
||||
var distAfter = 0.0'f32
|
||||
for i in 0 ..< targetAfter.size:
|
||||
let d = targetAfter.unsafe_raw_offset[i] - criticW.unsafe_raw_offset[i]
|
||||
distAfter += d * d
|
||||
|
||||
# Target moved toward critic (distance decreased).
|
||||
# With tau=0.005, it moves a tiny bit — just check direction.
|
||||
check distAfter <= distBefore + 1e-3'f32 # soft bound; critic also moves
|
||||
|
||||
test "empty sequences is a no-op":
|
||||
let m = sacUpdate(trainer, @[])
|
||||
check m.criticLoss == 0.0'f32
|
||||
check m.actorLoss == 0.0'f32
|
||||
check m.alpha == 0.0'f32
|
||||
|
||||
suite "SACTrainer — done=true terminal transitions":
|
||||
|
||||
test "update with done=true transitions produces finite losses":
|
||||
randomize(7)
|
||||
putEnv("SACLSTM_HIDDEN_SIZE", "16")
|
||||
var trainer = makeTrainer()
|
||||
var buf = makeBuffer()
|
||||
# Fill with short episodes: done every 12 steps (burnIn=4, trainWindow=8 → seqLen=12)
|
||||
fillBuffer(buf, 800, donePeriod = 12)
|
||||
let seqs = buf.sampleSequences(2)
|
||||
if seqs.len > 0:
|
||||
let m = sacUpdate(trainer, seqs)
|
||||
check isFinite(m.criticLoss)
|
||||
check isFinite(m.actorLoss)
|
||||
check isFinite(m.alphaLoss)
|
||||
check m.alpha > 0.0'f32
|
||||
|
||||
suite "SACTrainer — multiple updates":
|
||||
|
||||
test "three sequential updates all stay finite":
|
||||
randomize(99)
|
||||
putEnv("SACLSTM_HIDDEN_SIZE", "16")
|
||||
var trainer = makeTrainer()
|
||||
var buf = makeBuffer()
|
||||
fillBuffer(buf, 1000)
|
||||
for _ in 1..3:
|
||||
let seqs = buf.sampleSequences(4)
|
||||
if seqs.len > 0:
|
||||
let m = sacUpdate(trainer, seqs)
|
||||
check isFinite(m.criticLoss)
|
||||
check isFinite(m.actorLoss)
|
||||
check m.alpha > 0.0'f32
|
||||
@@ -0,0 +1,141 @@
|
||||
import unittest
|
||||
import arraymancer
|
||||
import zip/zipfiles
|
||||
import std/[os, math, strutils]
|
||||
import SAC_LSTM_Bot/network
|
||||
import SAC_LSTM_Bot/weights
|
||||
|
||||
const
|
||||
STATE_DIM = 20
|
||||
ACTION_DIM = 4
|
||||
|
||||
proc tensorsEqual(a, b: Tensor[float32]; tol: float32 = 1e-6'f32): bool =
|
||||
if a.shape != b.shape: return false
|
||||
for i in 0 ..< a.size:
|
||||
if abs(a.unsafe_raw_offset[i] - b.unsafe_raw_offset[i]) > tol: return false
|
||||
true
|
||||
|
||||
suite "saveWeights / loadCheckpoint":
|
||||
|
||||
setup:
|
||||
let actor = initActorNet(STATE_DIM)
|
||||
let c1 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let c2 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let tc1 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let tc2 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let alpha = 0.2'f32
|
||||
let zipPath = getTempDir() / "test_weights_latest.zip"
|
||||
|
||||
teardown:
|
||||
if fileExists(zipPath): removeFile(zipPath)
|
||||
|
||||
test "creates a valid zip file":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
check fileExists(zipPath)
|
||||
var z: ZipArchive
|
||||
check z.open(zipPath, fmRead)
|
||||
var count = 0
|
||||
for f in z.walkFiles: inc count
|
||||
z.close()
|
||||
check count > 0
|
||||
|
||||
test "all entries are .npy files (+ alpha.npy)":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
var z: ZipArchive
|
||||
discard z.open(zipPath, fmRead)
|
||||
var allNpy = true
|
||||
for f in z.walkFiles:
|
||||
if not f.endsWith(".npy"): allNpy = false
|
||||
z.close()
|
||||
check allNpy
|
||||
|
||||
test "round-trip: actor weights preserved":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check tensorsEqual(actor.fc1.w, ck.actor.fc1.w)
|
||||
check tensorsEqual(actor.fc1.b, ck.actor.fc1.b)
|
||||
check tensorsEqual(actor.lstm.wCombined, ck.actor.lstm.wCombined)
|
||||
check tensorsEqual(actor.lstm.bCombined, ck.actor.lstm.bCombined)
|
||||
check tensorsEqual(actor.fc2.w, ck.actor.fc2.w)
|
||||
check tensorsEqual(actor.muHead.w, ck.actor.muHead.w)
|
||||
check tensorsEqual(actor.logStdHead.w, ck.actor.logStdHead.w)
|
||||
|
||||
test "round-trip: critic weights preserved":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check tensorsEqual(c1.fc1.w, ck.critic1.fc1.w)
|
||||
check tensorsEqual(c1.lstm.wCombined, ck.critic1.lstm.wCombined)
|
||||
check tensorsEqual(c1.fc3.w, ck.critic1.fc3.w)
|
||||
check tensorsEqual(tc1.fc1.w, ck.targetCritic1.fc1.w)
|
||||
check tensorsEqual(tc2.fc1.w, ck.targetCritic2.fc1.w)
|
||||
|
||||
test "round-trip: alpha preserved":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check abs(ck.alpha - alpha) < 1e-6'f32
|
||||
|
||||
test "round-trip: hiddenDim reconstructed":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check ck.actor.hiddenDim == actor.hiddenDim
|
||||
check ck.critic1.hiddenDim == c1.hiddenDim
|
||||
|
||||
test "adam not present → initialized=false":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check not ck.adam.initialized
|
||||
|
||||
suite "saveCheckpoint (with Adam)":
|
||||
|
||||
setup:
|
||||
let actor = initActorNet(STATE_DIM)
|
||||
let c1 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let c2 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let tc1 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let tc2 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let alpha = 0.1'f32
|
||||
var adam = initSACAdamStates(actor, c1, c2)
|
||||
# put some non-zero values in Adam state
|
||||
adam.actor.fc1.w.m[0, 0] = 0.5'f32
|
||||
adam.actor.fc1.w.t = 42
|
||||
adam.critic1.fc1.w.t = 7
|
||||
adam.alpha.t = 99
|
||||
let zipPath = getTempDir() / "test_checkpoint.zip"
|
||||
|
||||
teardown:
|
||||
if fileExists(zipPath): removeFile(zipPath)
|
||||
|
||||
test "round-trip: Adam m tensor":
|
||||
saveCheckpoint(zipPath, actor, c1, c2, tc1, tc2, alpha, adam)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check ck.adam.initialized
|
||||
check abs(ck.adam.actor.fc1.w.m[0, 0] - 0.5'f32) < 1e-6'f32
|
||||
|
||||
test "round-trip: Adam t counters":
|
||||
saveCheckpoint(zipPath, actor, c1, c2, tc1, tc2, alpha, adam)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check ck.adam.actor.fc1.w.t == 42
|
||||
check ck.adam.critic1.fc1.w.t == 7
|
||||
check ck.adam.alpha.t == 99
|
||||
|
||||
suite "Atomic save":
|
||||
|
||||
test "temp file is cleaned up after successful save":
|
||||
let zipPath = getTempDir() / "test_atomic.zip"
|
||||
let tmpZip = zipPath & ".tmp"
|
||||
let actor = initActorNet(STATE_DIM)
|
||||
let c = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
saveWeights(zipPath, actor, c, c, c, c, 0.2'f32)
|
||||
check fileExists(zipPath)
|
||||
check not fileExists(tmpZip)
|
||||
removeFile(zipPath)
|
||||
|
||||
suite "Error handling":
|
||||
|
||||
test "loadCheckpoint missing file → IOError":
|
||||
var raised = false
|
||||
try:
|
||||
discard loadCheckpoint(getTempDir() / "nonexistent_xxxxxx.zip")
|
||||
except IOError:
|
||||
raised = true
|
||||
check raised
|
||||
Reference in New Issue
Block a user