Files
rustytorch/demos/rtx-embodied-demo/src/replay_buffer.rs
T
osobhandClaude Opus 4.6 02d382d5f6 style: apply rustfmt across all crates and demos
Consistent formatting pass: line wrapping, import sorting, trailing
whitespace removal, let-chain indentation, merged derive attributes,
and unsafe block reformatting.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-04-12 07:01:58 -07:00

198 lines
4.8 KiB
Rust

//! Replay buffer for storing and sampling experience.
use embodied_shared::Transition;
/// Replay buffer for experience storage.
#[derive(Debug)]
pub struct ReplayBuffer {
/// Buffer storage.
buffer: Vec<Transition>,
/// Maximum capacity.
capacity: usize,
/// Current write position.
position: usize,
/// RNG state for sampling.
rng_state: u64,
}
impl ReplayBuffer {
/// Create a new replay buffer.
pub fn new(capacity: usize) -> Self {
Self {
buffer: Vec::with_capacity(capacity),
capacity,
position: 0,
rng_state: 42,
}
}
/// Add a transition to the buffer.
pub fn add(&mut self, transition: Transition) {
if self.buffer.len() < self.capacity {
self.buffer.push(transition);
} else {
self.buffer[self.position] = transition;
}
self.position = (self.position + 1) % self.capacity;
}
/// Get buffer length.
#[must_use]
pub fn len(&self) -> usize {
self.buffer.len()
}
/// Check if buffer is empty.
#[must_use]
pub fn is_empty(&self) -> bool {
self.buffer.is_empty()
}
/// Sample a batch of transitions.
pub fn sample(&mut self, batch_size: usize) -> Vec<Transition> {
let batch_size = batch_size.min(self.buffer.len());
let mut batch = Vec::with_capacity(batch_size);
for _ in 0..batch_size {
let idx = self.random_index();
batch.push(self.buffer[idx].clone());
}
batch
}
/// Sample sequences of transitions.
pub fn sample_sequence(
&mut self,
batch_size: usize,
sequence_length: usize,
) -> Vec<Vec<Transition>> {
let batch_size = batch_size.min(self.buffer.len() / sequence_length.max(1));
let mut batch = Vec::with_capacity(batch_size);
// For simplicity, sample random starting points
for _ in 0..batch_size {
let max_start = (self.buffer.len() - sequence_length).max(0);
let start_idx = if max_start > 0 {
self.random_index() % (max_start + 1)
} else {
0
};
let end_idx = (start_idx + sequence_length).min(self.buffer.len());
let sequence: Vec<Transition> = self.buffer[start_idx..end_idx].to_vec();
batch.push(sequence);
}
batch
}
/// Get a random index.
fn random_index(&mut self) -> usize {
self.rng_state = self
.rng_state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(self.rng_state >> 11) as usize % self.buffer.len().max(1)
}
/// Clear the buffer.
pub fn clear(&mut self) {
self.buffer.clear();
self.position = 0;
}
/// Get capacity.
#[must_use]
pub fn capacity(&self) -> usize {
self.capacity
}
}
#[cfg(test)]
mod tests {
use super::*;
use embodied_shared::sample_transition;
#[test]
fn test_buffer_creation() {
let buffer = ReplayBuffer::new(1000);
assert_eq!(buffer.capacity(), 1000);
assert!(buffer.is_empty());
}
#[test]
fn test_add_transition() {
let mut buffer = ReplayBuffer::new(100);
let transition = sample_transition();
buffer.add(transition);
assert_eq!(buffer.len(), 1);
}
#[test]
fn test_add_overflow() {
let mut buffer = ReplayBuffer::new(10);
for _ in 0..15 {
buffer.add(sample_transition());
}
// Buffer should not exceed capacity
assert_eq!(buffer.len(), 10);
}
#[test]
fn test_sample() {
let mut buffer = ReplayBuffer::new(100);
for _ in 0..50 {
buffer.add(sample_transition());
}
let batch = buffer.sample(16);
assert_eq!(batch.len(), 16);
}
#[test]
fn test_sample_larger_than_buffer() {
let mut buffer = ReplayBuffer::new(100);
for _ in 0..5 {
buffer.add(sample_transition());
}
let batch = buffer.sample(16);
assert_eq!(batch.len(), 5); // Capped at buffer size
}
#[test]
fn test_sample_sequence() {
let mut buffer = ReplayBuffer::new(100);
for _ in 0..50 {
buffer.add(sample_transition());
}
let batch = buffer.sample_sequence(4, 10);
assert_eq!(batch.len(), 4);
for seq in &batch {
assert!(seq.len() <= 10);
}
}
#[test]
fn test_clear() {
let mut buffer = ReplayBuffer::new(100);
for _ in 0..10 {
buffer.add(sample_transition());
}
assert_eq!(buffer.len(), 10);
buffer.clear();
assert!(buffer.is_empty());
}
}