feat(SAC_LSTM_Bot): replay buffer module (#45)

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-08-20 23:40:17 +02:00
parent add3e34926
commit cb33551621
2 changed files with 286 additions and 0 deletions
+156
View File
@@ -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"