## 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"