use rand::{seq::SliceRandom, thread_rng};
use crate::{ds::RingBuffer, env::Environment};
use super::{Exp, ExpBatch};
pub struct ReplayMemory<E: Environment> {
memory: RingBuffer<Exp<E>>,
pub batch_size: usize,
}
impl<E: Environment> ReplayMemory<E> {
pub fn new(capacity: usize, batch_size: usize) -> Self {
Self {
memory: RingBuffer::<Exp<E>>::new(capacity),
batch_size,
}
}
pub fn push(&mut self, exp: Exp<E>) {
self.memory.push(exp);
}
pub fn sample(&self) -> Option<Vec<&Exp<E>>> {
if self.batch_size <= self.memory.len() {
Some(
self.memory
.view()
.choose_multiple(&mut thread_rng(), self.batch_size)
.collect(),
)
} else {
None
}
}
pub fn sample_zipped(&self) -> Option<ExpBatch<E>> {
if self.batch_size <= self.memory.len() {
let experiences = self
.memory
.view()
.choose_multiple(&mut thread_rng(), self.batch_size)
.cloned();
let batch = ExpBatch::from_iter(experiences, self.batch_size);
Some(batch)
} else {
None
}
}
}
#[cfg(test)]
mod tests {
use crate::env::tests::MockEnv;
use super::*;
fn create_mock_exp_vec() -> Vec<Exp<MockEnv>> {
(0..4)
.map(|i| Exp {
state: i,
action: i + 1,
next_state: Some(i + 1),
reward: 1.0,
})
.collect()
}
#[test]
fn replay_memory_functional() {
let experiences = create_mock_exp_vec();
let mut memory = ReplayMemory::new(4, 2);
assert!(
memory.sample().is_none(),
"sample none when too few experiences"
);
assert!(
memory.sample_zipped().is_none(),
"sample_zipped none when too few experiences"
);
for exp in experiences {
memory.push(exp);
}
assert!(
memory.sample().is_some_and(|b| b.len() == 2),
"sample works"
);
assert!(
memory.sample_zipped().is_some_and(|b| b.states.len() == 2),
"sample_zipped works"
);
}
}