///|
/// One transition in an RL rollout.
pub struct Transition[S, A] {
state : S
action : A
reward : Double
next_state : S
done : Bool
}
///|
pub fn[S, A] Transition::new(
state : S,
action : A,
reward : Double,
next_state : S,
done : Bool,
) -> Transition[S, A] {
{ state, action, reward, next_state, done }
}
///|
pub fn[S, A] Transition::terminal(
state : S,
action : A,
reward : Double,
next_state : S,
) -> Transition[S, A] {
Transition::new(state, action, reward, next_state, true)
}
///|
pub fn[S, A] Transition::non_terminal(
state : S,
action : A,
reward : Double,
next_state : S,
) -> Transition[S, A] {
Transition::new(state, action, reward, next_state, false)
}
///|
/// A compact episode container for ordered transitions.
pub struct Episode[S, A] {
transitions : Array[Transition[S, A]]
terminated : Bool
truncated : Bool
total_reward : Double
}
///|
pub fn[S, A] Episode::new() -> Episode[S, A] {
let transitions : Array[Transition[S, A]] = []
{ transitions, terminated: false, truncated: false, total_reward: 0.0 }
}
///|
pub fn[S, A] Episode::from_transitions(
transitions : Array[Transition[S, A]],
) -> Episode[S, A] {
let mut episode = Episode::new()
for transition in transitions {
episode = episode.push(transition)
}
episode
}
///|
pub fn[S, A] Episode::push(
self : Episode[S, A],
transition : Transition[S, A],
) -> Episode[S, A] {
let transitions = self.transitions
transitions.push(transition)
{
transitions,
terminated: if transition.done {
true
} else {
self.terminated
},
truncated: self.truncated,
total_reward: self.total_reward + transition.reward,
}
}
///|
pub fn[S, A] Episode::mark_terminated(self : Episode[S, A]) -> Episode[S, A] {
{ ..self, terminated: true }
}
///|
pub fn[S, A] Episode::mark_truncated(self : Episode[S, A]) -> Episode[S, A] {
{ ..self, truncated: true }
}
///|
pub fn[S, A] Episode::len(self : Episode[S, A]) -> Int {
self.transitions.length()
}
///|
pub fn[S, A] Episode::is_empty(self : Episode[S, A]) -> Bool {
self.transitions.length() == 0
}
///|
pub fn[S, A] Episode::is_terminated(self : Episode[S, A]) -> Bool {
self.terminated
}
///|
pub fn[S, A] Episode::is_truncated(self : Episode[S, A]) -> Bool {
self.truncated
}
///|
pub fn[S, A] Episode::total_reward(self : Episode[S, A]) -> Double {
self.total_reward
}
///|
pub fn[S, A] Episode::transitions(
self : Episode[S, A],
) -> Array[Transition[S, A]] {
self.transitions
}
///|
pub fn[S, A] Episode::first_transition(
self : Episode[S, A],
) -> Transition[S, A]? {
if self.transitions.length() == 0 {
None
} else {
Some(self.transitions[0])
}
}
///|
pub fn[S, A] Episode::last_transition(
self : Episode[S, A],
) -> Transition[S, A]? {
let len = self.transitions.length()
if len == 0 {
None
} else {
Some(self.transitions[len - 1])
}
}
///|
/// Return discounted returns in rollout order.
pub fn[S, A] Episode::discounted_returns(
self : Episode[S, A],
gamma : Double,
) -> Array[Double] {
let len = self.transitions.length()
let returns : Array[Double] = []
let mut i = 0
while i < len {
returns.push(0.0)
i = i + 1
}
let mut running = 0.0
let mut idx = len
while idx > 0 {
idx = idx - 1
running = self.transitions[idx].reward + gamma * running
returns[idx] = running
}
returns
}
///|
/// Return N-step transitions.
/// For each transition i, the reward is accumulated for N steps (or until done),
/// and the next_state is the state after N steps (or the terminal state).
pub fn[S, A] Episode::n_step_transitions(
self : Episode[S, A],
n : Int,
gamma : Double,
) -> Array[Transition[S, A]] {
if n <= 1 {
return self.transitions
}
let len = self.transitions.length()
let result : Array[Transition[S, A]] = []
let mut i = 0
while i < len {
let start_transition = self.transitions[i]
let mut accumulated_reward = 0.0
let mut current_gamma = 1.0
let mut next_state = start_transition.next_state
let mut done = start_transition.done
let mut step = 0
while step < n && i + step < len {
let t = self.transitions[i + step]
accumulated_reward = accumulated_reward + current_gamma * t.reward
next_state = t.next_state
done = t.done
current_gamma = current_gamma * gamma
if done {
break
}
step = step + 1
}
result.push(
Transition::new(
start_transition.state,
start_transition.action,
accumulated_reward,
next_state,
done,
),
)
i = i + 1
}
result
}