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]>
198 lines
4.8 KiB
Rust
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());
|
|
}
|
|
}
|