//! 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, /// 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 { 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> { 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 = 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()); } }