use rlkit::replay_buffer::ReplayBuffer;
use rlkit::types::{Sample, Status, Reward};
use rlkit::Action;
#[test]
fn test_new_buffer() {
let buffer: ReplayBuffer<f32, f32> = ReplayBuffer::new(100);
assert_eq!(buffer.len(), 0);
assert!(buffer.is_empty());
}
#[test]
fn test_push_and_len() {
let mut buffer: ReplayBuffer<f32, f32> = ReplayBuffer::new(100);
let sample1 = Sample::<f32, f32> {
state: Status::new(vec![1.0, 2.0, 3.0], vec![10.0, 10.0, 10.0]),
action: Action::new(vec![0.0], vec![10.0]),
reward: Reward(1.0),
next_state: Status::new(vec![4.0, 5.0, 6.0], vec![10.0, 10.0, 10.0]),
done: false,
};
let sample2 = Sample::<f32, f32> {
state: Status::new(vec![7.0, 8.0, 9.0], vec![10.0, 10.0, 10.0]),
action: Action::new(vec![1.0], vec![10.0]),
reward: Reward(2.0),
next_state: Status::new(vec![10.0, 11.0, 12.0], vec![10.0, 10.0, 10.0]),
done: true,
};
buffer.push(sample1.clone());
assert_eq!(buffer.len(), 1);
assert!(!buffer.is_empty());
buffer.push(sample2.clone());
assert_eq!(buffer.len(), 2);
}
#[test]
fn test_buffer_capacity() {
let mut buffer: ReplayBuffer<f32, f32> = ReplayBuffer::new(3);
for i in 0..5 {
let sample = Sample::<f32, f32> {
state: Status::new(vec![i as f32], vec![10.0]),
action: Action::new(vec![i as f32], vec![10.0]),
reward: Reward(i as f32),
next_state: Status::new(vec![(i + 1) as f32], vec![10.0]),
done: i == 4,
};
buffer.push(sample);
}
assert_eq!(buffer.len(), 3);
}
#[test]
fn test_sample() {
let mut buffer: ReplayBuffer<f32, f32> = ReplayBuffer::new(100);
for i in 0..10 {
let sample = Sample::<f32, f32> {
state: Status::new(vec![i as f32], vec![10.0]),
action: Action::new(vec![i as f32], vec![10.0]),
reward: Reward(i as f32),
next_state: Status::new(vec![(i + 1) as f32], vec![10.0]),
done: i == 9,
};
buffer.push(sample);
}
let empty_batch = buffer.sample(0);
assert!(empty_batch.is_empty());
let batch_size = 5;
let samples = buffer.sample(batch_size);
assert_eq!(samples.len(), batch_size);
for sample in samples {
assert!(!sample.state.as_slice().is_empty());
assert!(!sample.next_state.as_slice().is_empty());
}
}
#[test]
fn test_sample_empty_buffer() {
let buffer: ReplayBuffer<f32, f32> = ReplayBuffer::new(100);
let samples = buffer.sample(5);
assert!(samples.is_empty());
}
#[test]
fn test_sample_cloning() {
let mut buffer: ReplayBuffer<f32, f32> = ReplayBuffer::new(100);
let original_sample = Sample::<f32, f32> {
state: Status::new(vec![1.0, 2.0, 3.0], vec![10.0, 10.0, 10.0]),
action: Action::new(vec![0.0], vec![10.0]),
reward: Reward(1.0),
next_state: Status::new(vec![4.0, 5.0, 6.0], vec![10.0, 10.0, 10.0]),
done: false,
};
buffer.push(original_sample.clone());
let mut samples = buffer.sample(1);
let mut sampled_sample = samples.pop().unwrap();
sampled_sample.reward = Reward(100.0);
let samples_after_modification = buffer.sample(1);
assert_eq!(samples_after_modification[0].reward, original_sample.reward);
}