///|
/// A fixed-size rollout window. `mask[i]` is false for padding positions.
pub struct SequenceSample[S, A] {
start : Int
transitions : Array[Transition[S, A]?]
mask : Array[Bool]
}
///|
pub fn[S, A] SequenceSample::start(self : SequenceSample[S, A]) -> Int {
self.start
}
///|
pub fn[S, A] SequenceSample::length(self : SequenceSample[S, A]) -> Int {
self.transitions.length()
}
///|
pub fn[S, A] SequenceSample::transitions(
self : SequenceSample[S, A],
) -> Array[Transition[S, A]?] {
self.transitions
}
///|
pub fn[S, A] SequenceSample::mask(self : SequenceSample[S, A]) -> Array[Bool] {
self.mask
}
///|
pub fn[S, A] SequenceSample::valid_length(self : SequenceSample[S, A]) -> Int {
let mut count = 0
for valid in self.mask {
if valid {
count = count + 1
}
}
count
}
///|
pub fn[S, A] SequenceSample::is_padded(self : SequenceSample[S, A]) -> Bool {
self.valid_length() < self.length()
}
///|
pub fn[S, A] ReplayBuffer::sequence(
self : ReplayBuffer[S, A],
start : Int,
sequence_length : Int,
pad : Bool,
) -> SequenceSample[S, A] {
let transitions : Array[Transition[S, A]?] = []
let mask : Array[Bool] = []
if sequence_length <= 0 {
return { start, transitions, mask }
}
let mut offset = 0
while offset < sequence_length {
match self.get(start + offset) {
Some(transition) => {
transitions.push(Some(transition))
mask.push(true)
}
None =>
if pad {
transitions.push(None)
mask.push(false)
}
}
offset = offset + 1
}
{ start, transitions, mask }
}
///|
pub fn[S, A] ReplayBuffer::valid_sequence_starts(
self : ReplayBuffer[S, A],
sequence_length : Int,
allow_padding : Bool,
) -> Array[Int] {
let starts : Array[Int] = []
if sequence_length <= 0 || self.len == 0 {
return starts
}
let last_start = if allow_padding {
self.len - 1
} else {
self.len - sequence_length
}
let mut start = 0
while start <= last_start {
starts.push(start)
start = start + 1
}
starts
}
///|
pub fn[S, A] ReplayBuffer::sample_sequences(
self : ReplayBuffer[S, A],
batch_size : Int,
sequence_length : Int,
seed : Int,
allow_padding : Bool,
) -> Array[SequenceSample[S, A]] {
let result : Array[SequenceSample[S, A]] = []
if batch_size <= 0 {
return result
}
let starts = self.valid_sequence_starts(sequence_length, allow_padding)
if starts.length() == 0 {
return result
}
let rng = ReplayRng::new(seed)
let mut i = 0
while i < batch_size {
let start = starts[rng.next_index(starts.length())]
result.push(self.sequence(start, sequence_length, allow_padding))
i = i + 1
}
result
}
///|
pub fn[S, A] ReplayBuffer::all_sequences(
self : ReplayBuffer[S, A],
sequence_length : Int,
) -> Array[SequenceSample[S, A]] {
let result : Array[SequenceSample[S, A]] = []
for start in self.valid_sequence_starts(sequence_length, false) {
result.push(self.sequence(start, sequence_length, false))
}
result
}
///|
pub fn[S, A] SequenceSample::rewards(
self : SequenceSample[S, A],
) -> Array[Double] {
let result : Array[Double] = []
for transition in self.transitions {
match transition {
Some(value) => result.push(value.reward)
None => result.push(0.0)
}
}
result
}
///|
pub fn[S, A] SequenceSample::terminals(
self : SequenceSample[S, A],
) -> Array[Bool] {
let result : Array[Bool] = []
for transition in self.transitions {
match transition {
Some(value) => result.push(value.done)
None => result.push(false)
}
}
result
}