///|
pub struct ReplaySample[S, A] {
index : Int
transition : Transition[S, A]
}
///|
pub struct ReplayBuffer[S, A] {
capacity : Int
mut len : Int
mut head : Int
data : Array[Transition[S, A]?]
}
///|
pub fn[S, A] ReplayBuffer::new(capacity : Int) -> ReplayBuffer[S, A] {
let data : Array[Transition[S, A]?] = []
let mut i = 0
while i < capacity {
data.push(None)
i = i + 1
}
{ capacity, len: 0, head: 0, data }
}
///|
pub fn[S, A] ReplayBuffer::capacity(self : ReplayBuffer[S, A]) -> Int {
self.capacity
}
///|
pub fn[S, A] ReplayBuffer::len(self : ReplayBuffer[S, A]) -> Int {
self.len
}
///|
pub fn[S, A] ReplayBuffer::is_empty(self : ReplayBuffer[S, A]) -> Bool {
self.len == 0
}
///|
pub fn[S, A] ReplayBuffer::is_full(self : ReplayBuffer[S, A]) -> Bool {
self.capacity > 0 && self.len == self.capacity
}
///|
pub fn[S, A] ReplayBuffer::clear(
self : ReplayBuffer[S, A],
) -> ReplayBuffer[S, A] {
let mut i = 0
while i < self.data.length() {
self.data[i] = None
i = i + 1
}
{ ..self, len: 0, head: 0 }
}
///|
pub fn[S, A] ReplayBuffer::get(
self : ReplayBuffer[S, A],
logical_index : Int,
) -> Transition[S, A]? {
if logical_index < 0 || logical_index >= self.len || self.capacity == 0 {
None
} else {
let slot = (self.head + logical_index) % self.capacity
self.data[slot]
}
}
///|
pub fn[S, A] ReplayBuffer::oldest(
self : ReplayBuffer[S, A],
) -> Transition[S, A]? {
self.get(0)
}
///|
pub fn[S, A] ReplayBuffer::newest(
self : ReplayBuffer[S, A],
) -> Transition[S, A]? {
if self.len == 0 {
None
} else {
self.get(self.len - 1)
}
}
///|
pub fn[S, A] ReplayBuffer::push(
self : ReplayBuffer[S, A],
transition : Transition[S, A],
) -> Int {
if self.capacity == 0 {
return 0
}
if self.len < self.capacity {
self.data[self.len] = Some(transition)
self.len = self.len + 1
} else {
self.data[self.head] = Some(transition)
self.head = (self.head + 1) % self.capacity
}
self.len
}
///|
pub fn[S, A] ReplayBuffer::extend_episode(
self : ReplayBuffer[S, A],
episode : Episode[S, A],
) -> Int {
for transition in episode.transitions() {
let _ = self.push(transition)
}
self.len
}
///|
pub fn[S, A] ReplayBuffer::to_array(
self : ReplayBuffer[S, A],
) -> Array[Transition[S, A]] {
let result : Array[Transition[S, A]] = []
let mut i = 0
while i < self.len {
match self.get(i) {
Some(transition) => result.push(transition)
None => ()
}
i = i + 1
}
result
}
///|
pub fn[S, A] ReplayBuffer::sample_batch(
self : ReplayBuffer[S, A],
batch_size : Int,
seed : Int,
) -> Array[ReplaySample[S, A]] {
let target = if batch_size < self.len { batch_size } else { self.len }
let samples : Array[ReplaySample[S, A]] = []
if target == 0 {
return samples
}
let indices = shuffle_prefix(make_sequential_indices(self.len), target, seed)
let mut i = 0
while i < target {
let idx = indices[i]
match self.get(idx) {
Some(transition) => samples.push({ index: idx, transition })
None => ()
}
i = i + 1
}
samples
}