///|
/// A lightweight offline-RL dataset made of complete episodes.
pub struct TrajectoryDataset[S, A] {
episodes : Array[Episode[S, A]]
mut transition_count : Int
}
///|
pub fn[S, A] TrajectoryDataset::new() -> TrajectoryDataset[S, A] {
{ episodes: [], transition_count: 0 }
}
///|
pub fn[S, A] TrajectoryDataset::from_episodes(
episodes : Array[Episode[S, A]],
) -> TrajectoryDataset[S, A] {
let dataset = TrajectoryDataset::new()
for episode in episodes {
let _ = dataset.add(episode)
}
dataset
}
///|
pub fn[S, A] TrajectoryDataset::add(
self : TrajectoryDataset[S, A],
episode : Episode[S, A],
) -> Bool {
if episode.validate().length() != 0 {
return false
}
self.episodes.push(episode)
self.transition_count = self.transition_count + episode.len()
true
}
///|
pub fn[S, A] TrajectoryDataset::add_unchecked(
self : TrajectoryDataset[S, A],
episode : Episode[S, A],
) -> Unit {
self.episodes.push(episode)
self.transition_count = self.transition_count + episode.len()
}
///|
pub fn[S, A] TrajectoryDataset::episode_count(
self : TrajectoryDataset[S, A],
) -> Int {
self.episodes.length()
}
///|
pub fn[S, A] TrajectoryDataset::transition_count(
self : TrajectoryDataset[S, A],
) -> Int {
self.transition_count
}
///|
pub fn[S, A] TrajectoryDataset::is_empty(
self : TrajectoryDataset[S, A],
) -> Bool {
self.episodes.length() == 0
}
///|
pub fn[S, A] TrajectoryDataset::episodes(
self : TrajectoryDataset[S, A],
) -> Array[Episode[S, A]] {
self.episodes
}
///|
pub fn[S, A] TrajectoryDataset::get(
self : TrajectoryDataset[S, A],
index : Int,
) -> Episode[S, A]? {
if index < 0 || index >= self.episodes.length() {
None
} else {
Some(self.episodes[index])
}
}
///|
pub fn[S, A] TrajectoryDataset::flatten(
self : TrajectoryDataset[S, A],
) -> Array[Transition[S, A]] {
let result : Array[Transition[S, A]] = []
for episode in self.episodes {
for transition in episode.transitions() {
result.push(transition)
}
}
result
}
///|
pub fn[S, A] TrajectoryDataset::reward_stats(
self : TrajectoryDataset[S, A],
) -> RewardStats {
reward_stats_from_transitions(self.flatten())
}
///|
pub fn[S, A] TrajectoryDataset::episode_lengths(
self : TrajectoryDataset[S, A],
) -> Array[Int] {
let result : Array[Int] = []
for episode in self.episodes {
result.push(episode.len())
}
result
}
///|
pub fn[S, A] TrajectoryDataset::sample_episodes(
self : TrajectoryDataset[S, A],
count : Int,
seed : Int,
) -> Array[Episode[S, A]] {
let result : Array[Episode[S, A]] = []
if count <= 0 || self.episodes.length() == 0 {
return result
}
let rng = ReplayRng::new(seed)
let mut i = 0
while i < count {
result.push(self.episodes[rng.next_index(self.episodes.length())])
i = i + 1
}
result
}
///|
pub fn[S, A] TrajectoryDataset::split_at(
self : TrajectoryDataset[S, A],
first_count : Int,
) -> (TrajectoryDataset[S, A], TrajectoryDataset[S, A]) {
let pivot = if first_count < 0 {
0
} else if first_count > self.episodes.length() {
self.episodes.length()
} else {
first_count
}
let left = TrajectoryDataset::new()
let right = TrajectoryDataset::new()
let mut i = 0
while i < self.episodes.length() {
if i < pivot {
left.add_unchecked(self.episodes[i])
} else {
right.add_unchecked(self.episodes[i])
}
i = i + 1
}
(left, right)
}
///|
pub fn[S, A] TrajectoryDataset::terminal_episode_count(
self : TrajectoryDataset[S, A],
) -> Int {
let mut count = 0
for episode in self.episodes {
if episode.is_terminated() {
count = count + 1
}
}
count
}
///|
pub fn[S, A] TrajectoryDataset::truncated_episode_count(
self : TrajectoryDataset[S, A],
) -> Int {
let mut count = 0
for episode in self.episodes {
if episode.is_truncated() {
count = count + 1
}
}
count
}
///|
pub fn[S, A] TrajectoryDataset::discounted_returns(
self : TrajectoryDataset[S, A],
gamma : Double,
) -> Array[Double] {
let result : Array[Double] = []
for episode in self.episodes {
for value in episode.discounted_returns(gamma) {
result.push(value)
}
}
result
}