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>
|
| | 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.
|
| |
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.
◆ 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
-
| capacity | Maximum number of transitions retained. Must be >= 1. |
| observation_dim | Width of an observation row. Must be >= 1. |
| action_dim | Width of an action row (Environment::action_dim()). Must be >= 1. |
| backend | Backend to allocate through. Not owned; must outlive this buffer. |
| seed | Seed for the internal deterministic LCG used by sample(). |
- Exceptions
-
| std::invalid_argument | if capacity, observation_dim or action_dim is <= 0 – external boundary, the same classification as CartPoleEnv's max_steps check. |
◆ 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
-
| observation | State before the action, shape (1, observation_dim()). |
| action | Action taken, shape (1, action_dim()). |
| reward | Scalar reward received. |
| next_observation | State after the action, shape (1, observation_dim()). |
| done | Whether the episode ended on this transition. |
- Exceptions
-
| std::invalid_argument | if 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_size | Number of transitions to draw. Must be in [1, size()]. |
- Returns
- The sampled transitions, one Tensor per field (see ReplayBatch).
- Exceptions
-
| std::invalid_argument | if 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: