diff --git a/SAC_LSTM_Bot/src/SAC_LSTM_Bot/replay_buffer.nim b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/replay_buffer.nim new file mode 100644 index 0000000..7672deb --- /dev/null +++ b/SAC_LSTM_Bot/src/SAC_LSTM_Bot/replay_buffer.nim @@ -0,0 +1,130 @@ +## replay_buffer.nim — sequential ring buffer for off-policy SAC+LSTM training. +## +## Stores transitions and samples contiguous sequences for recurrent training. +## Sequences NEVER cross battle boundaries (done=true). +## +## Config env vars: +## SACLSTM_BUFFER_CAPACITY (default: 500_000) +## SACLSTM_BURN_IN (default: 8) +## SACLSTM_TRAIN_WINDOW (default: 16) + +import arraymancer +import std/[os, strutils, random] + +# ── Config ──────────────────────────────────────────────────────────────────── + +proc getBufferCapacity*(): int = + parseInt(getEnv("SACLSTM_BUFFER_CAPACITY", "500000")) + +proc getBurnIn*(): int = + parseInt(getEnv("SACLSTM_BURN_IN", "8")) + +proc getTrainWindow*(): int = + parseInt(getEnv("SACLSTM_TRAIN_WINDOW", "16")) + +# ── Types ───────────────────────────────────────────────────────────────────── + +type + Transition* = object + state*: Tensor[float32] # [stateDim] + action*: Tensor[float32] # [actionDim] + reward*: float32 + nextState*: Tensor[float32] # [stateDim] + done*: bool # true = battle end + + Sequence* = object + burnIn*: seq[Transition] # first burnIn steps (for LSTM warm-up) + train*: seq[Transition] # next trainWindow steps (for gradient computation) + + ReplayBuffer* = object + ## Ring buffer. `head` is the next write position. `count` tracks fill level. + transitions: seq[Transition] + capacity: int + stateDim: int + actionDim: int + head: int # next write index + count: int # number of valid transitions stored + burnIn: int + trainWindow: int + +# ── Construction ────────────────────────────────────────────────────────────── + +proc newReplayBuffer*(capacity, stateDim, actionDim: int; + burnIn = getBurnIn(); + trainWindow = getTrainWindow()): ReplayBuffer = + result.capacity = capacity + result.stateDim = stateDim + result.actionDim = actionDim + result.burnIn = burnIn + result.trainWindow = trainWindow + result.head = 0 + result.count = 0 + result.transitions = newSeq[Transition](capacity) + +# ── Core operations ─────────────────────────────────────────────────────────── + +proc add*(buf: var ReplayBuffer; t: Transition) = + buf.transitions[buf.head] = t + buf.head = (buf.head + 1) mod buf.capacity + if buf.count < buf.capacity: + inc buf.count + +proc len*(buf: ReplayBuffer): int = buf.count + +proc canSample*(buf: ReplayBuffer): bool = + buf.count >= buf.burnIn + buf.trainWindow + +# ── Sampling ────────────────────────────────────────────────────────────────── + +proc sampleSequences*(buf: ReplayBuffer; batchSize: int): seq[Sequence] = + ## Sample `batchSize` contiguous sequences of length burnIn+trainWindow. + ## Sequences never cross a done=true boundary and never wrap the ring buffer. + ## + ## Returns fewer than batchSize sequences if not enough valid starts exist. + ## Returns empty seq if canSample is false. + if not buf.canSample: return @[] + + let seqLen = buf.burnIn + buf.trainWindow + let oldest = if buf.count < buf.capacity: 0 + else: buf.head # oldest valid index when full + + # Build valid starting indices. + # ponytail: O(count) scan per sample call; upgrade to an indexed set of + # boundary positions if count reaches hundreds of thousands and profiling shows + # this is a bottleneck. + var validStarts: seq[int] + for i in 0 ..< buf.count - seqLen + 1: + # Absolute ring-buffer index for the i-th oldest transition + let startIdx = (oldest + i) mod buf.capacity + # Check: the sequence [startIdx .. startIdx+seqLen-2] must not contain done=true + # (a done at position k means the battle ended there; the next transition is + # from a new battle, so the sequence would cross a boundary). + # Also, the sequence must not wrap around the ring buffer. + let endIdx = startIdx + seqLen - 1 # exclusive of wrap check + if endIdx >= buf.capacity: + # Sequence wraps the ring buffer — invalid starting point. + continue + var crosses = false + for j in 0 ..< seqLen - 1: + if buf.transitions[startIdx + j].done: + crosses = true + break + if not crosses: + validStarts.add(startIdx) + + if validStarts.len == 0: return @[] + + result = newSeq[Sequence](min(batchSize, validStarts.len)) + # Sample with replacement if batchSize > validStarts.len, else sample without. + # ponytail: sampling with replacement for simplicity; shuffle+take for + # without-replacement if the caller needs it. + for i in 0 ..< result.len: + let startIdx = validStarts[rand(validStarts.len - 1)] + var s: Sequence + s.burnIn = newSeq[Transition](buf.burnIn) + s.train = newSeq[Transition](buf.trainWindow) + for j in 0 ..< buf.burnIn: + s.burnIn[j] = buf.transitions[startIdx + j] + for j in 0 ..< buf.trainWindow: + s.train[j] = buf.transitions[startIdx + buf.burnIn + j] + result[i] = s diff --git a/SAC_LSTM_Bot/tests/test_replay_buffer.nim b/SAC_LSTM_Bot/tests/test_replay_buffer.nim new file mode 100644 index 0000000..28aff46 --- /dev/null +++ b/SAC_LSTM_Bot/tests/test_replay_buffer.nim @@ -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"