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
@@ -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