Merge branch 'worktree-agent-a3ca3066' (ticket #45 replay buffer)
This commit is contained in:
@@ -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
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user