pulsatrix
Loading...
Searching...
No Matches
pulsatrix::ReplayBuffer Class Reference

Fixed-capacity circular replay buffer of (observation, action, reward, next_observation, done) transitions, with uniform-random-with-replacement batch sampling – the off-policy experience store of Mnih et al. 2015 (DQN), reused unchanged by SAC and every other off-policy learner in this campaign. More...

#include <replay_buffer.hpp>

Public Member Functions

 ReplayBuffer (int64_t capacity, int64_t observation_dim, int64_t action_dim, DeviceBackend *backend, uint32_t seed=42)
 Constructs an empty buffer with all storage pre-allocated and zero-filled.
 
void add (const Tensor &observation, const Tensor &action, float reward, const Tensor &next_observation, bool done)
 Stores one transition at the current circular write index, overwriting the oldest transition once the buffer is full.
 
ReplayBatch sample (int64_t batch_size)
 Draws batch_size transitions uniformly at random, with replacement.
 
int64_t size () const
 Transitions currently stored – rises to capacity(), then stays there.
 
int64_t capacity () const
 Maximum transitions retained before the oldest starts being overwritten.
 
int64_t observation_dim () const
 Width of an observation row, as passed to the constructor.
 
int64_t action_dim () const
 Width of an action row, as passed to the constructor.
 

Detailed Description

Fixed-capacity circular replay buffer of (observation, action, reward, next_observation, done) transitions, with uniform-random-with-replacement batch sampling – the off-policy experience store of Mnih et al. 2015 (DQN), reused unchanged by SAC and every other off-policy learner in this campaign.

Once capacity transitions have been stored, the oldest is overwritten by the next add(): memory is bounded regardless of how long training runs.

Note
Not a Module subclass, and not an Environment/Agent either. It has no parameters, no gradient, no forward/backward and nothing for an LRP rule to explain – forcing it into Module would put a pure-virtual propagate_relevance() on a type for which the concept is undefined, exactly the failure mode module.hpp's charter note rules out. Same disposition as MSELoss and Reparameterize: a plain utility class.
Environment-agnostic by construction. It is built from bare observation_dim / action_dim integers, not from an Environment&, so it neither depends on nor outlives any particular environment; CartPoleEnv is simply one source of the tensors a caller happens to add().
Storage is five pre-allocated (capacity, *)-shaped row-major host blocks written into at a circular index, not a std::vector of per-transition tensors. That is one allocation per field for the buffer's whole lifetime instead of five per stored transition, and it makes sample() a set of row-gathers rather than a loop of tensor copies.
Discrete and continuous actions are stored identically, as action_dim floats – a discrete action is its single-element encoded index, exactly as Environment::step() already accepts it. No buffer-side special-casing.
Host boundary (GPU-native-kernels Mission 7): the store is only ever read and written by the host (the circular write and the LCG-driven gather), so it lives in host memory. add() accepts rows on any device (one device->host copy of each row per call); sample() hands back a ReplayBatch uploaded through the buffer's own backend, so a buffer built with a GPU backend hands out GPU tensors.

Constructor & Destructor Documentation

◆ ReplayBuffer()

pulsatrix::ReplayBuffer::ReplayBuffer ( int64_t  capacity,
int64_t  observation_dim,
int64_t  action_dim,
DeviceBackend *  backend,
uint32_t  seed = 42 
)

Constructs an empty buffer with all storage pre-allocated and zero-filled.

Parameters
capacityMaximum number of transitions retained. Must be >= 1.
observation_dimWidth of an observation row. Must be >= 1.
action_dimWidth of an action row (Environment::action_dim()). Must be >= 1.
backendBackend to allocate through. Not owned; must outlive this buffer.
seedSeed for the internal deterministic LCG used by sample().
Exceptions
std::invalid_argumentif capacity, observation_dim or action_dim is <= 0 – external boundary, the same classification as CartPoleEnv's max_steps check.

Member Function Documentation

◆ action_dim()

int64_t pulsatrix::ReplayBuffer::action_dim ( ) const
inline

Width of an action row, as passed to the constructor.

◆ add()

void pulsatrix::ReplayBuffer::add ( const Tensor &  observation,
const Tensor &  action,
float  reward,
const Tensor &  next_observation,
bool  done 
)

Stores one transition at the current circular write index, overwriting the oldest transition once the buffer is full.

Parameters
observationState before the action, shape (1, observation_dim()).
actionAction taken, shape (1, action_dim()).
rewardScalar reward received.
next_observationState after the action, shape (1, observation_dim()).
doneWhether the episode ended on this transition.
Exceptions
std::invalid_argumentif any of the three tensors has the wrong shape – external boundary: a mismatched-shape Tensor can arrive from any caller, and silently writing it would corrupt neighbouring rows of the storage block.
Note
Host boundary: the three tensors may live on any device; each is copied to the host once (Tensor::to_host_vector()).

◆ capacity()

int64_t pulsatrix::ReplayBuffer::capacity ( ) const
inline

Maximum transitions retained before the oldest starts being overwritten.

◆ observation_dim()

int64_t pulsatrix::ReplayBuffer::observation_dim ( ) const
inline

Width of an observation row, as passed to the constructor.

◆ sample()

ReplayBatch pulsatrix::ReplayBuffer::sample ( int64_t  batch_size)

Draws batch_size transitions uniformly at random, with replacement.

Parameters
batch_sizeNumber of transitions to draw. Must be in [1, size()].
Returns
The sampled transitions, one Tensor per field (see ReplayBatch).
Exceptions
std::invalid_argumentif batch_size <= 0, or if batch_size > size() (which includes every call on an empty buffer). External boundary: sampling more than has actually been stored is a real caller error worth a real message, not something to satisfy with garbage rows from never-written slots.
Note
With replacement, so a batch may legitimately repeat a transition – the standard DQN formulation, and what keeps batch_size == size() valid rather than a degenerate full-buffer permutation.
Indices come from the internal LCG seeded at construction (the same Numerical-Recipes constants CartPoleEnv::reset() and this codebase's tests already use), never from \<random\>, whose engines are implementation-defined: two buffers with the same seed and the same contents sample identically, which is what makes sampling testable at all.

◆ size()

int64_t pulsatrix::ReplayBuffer::size ( ) const
inline

Transitions currently stored – rises to capacity(), then stays there.


The documentation for this class was generated from the following file: