// rainbow_lite.mbt — Rainbow lite (v0.37.2).
//
// Rainbow (Hessel et al. 2017) is a combination of 6 DQN improvements:
// - Double DQN (decoupled action selection/evaluation)
// - Dueling (V(s) + A(s,a) - mean_a A(s,a))
// - Prioritized replay (sample proportional to |δ|^α)
// - Multi-step (n-step returns instead of 1-step)
// - Distributional (full Q distribution; we skip for simplicity)
// - Noisy Nets (parameter noise; we skip for simplicity)
//
// This lite version combines: Double + Dueling + Prioritized +
// Multi-step (the four most impactful components).
///|
/// Prioritized replay buffer. Each transition has an associated
/// priority (|TD error|^α + ε). Sampling is done with probability
/// p_i ∝ priority_i. Importance-sampling weights correct for the
/// bias.
pub struct PrioritizedBuffer {
capacity : Int
states : Array[Int]
actions : Array[Int]
rewards : Array[Float]
next_states : Array[Int]
dones : Array[Bool]
/// Per-step priority (|δ|^α + ε). Higher → more likely to be sampled.
priorities : Array[Float]
/// Running sum of priorities (for cumulative sampling).
alpha : Float
/// Importance-sampling exponent (β → 0 = no correction, β = 1 = full).
beta : Float
/// Epsilon to keep priority > 0 even after perfect prediction.
epsilon : Float
mut size : Int
mut cursor : Int
mut total_priority : Float
}
///|
pub fn PrioritizedBuffer::new(
capacity : Int,
alpha : Float,
beta : Float,
epsilon : Float,
) -> PrioritizedBuffer {
{
capacity,
states: Array::make(capacity, 0),
actions: Array::make(capacity, 0),
rewards: Array::make(capacity, 0.0F),
next_states: Array::make(capacity, 0),
dones: Array::make(capacity, false),
priorities: Array::make(capacity, 1.0F),
alpha,
beta,
epsilon,
size: 0,
cursor: 0,
total_priority: 0.0F,
}
}
///|
/// Push a new transition with its initial priority. The caller
/// should later call `update_priority(i, new_priority)` to refresh
/// after computing TD errors.
pub fn PrioritizedBuffer::push(
self : PrioritizedBuffer,
s : Int,
a : Int,
r : Float,
s_next : Int,
done : Bool,
) -> Unit {
let i = self.cursor
self.states[i] = s
self.actions[i] = a
self.rewards[i] = r
self.next_states[i] = s_next
self.dones[i] = done
self.priorities[i] = 1.0F
self.total_priority = self.total_priority + 1.0F
self.cursor = (i + 1) % self.capacity
if self.size < self.capacity {
self.size = self.size + 1
}
}
///|
pub fn PrioritizedBuffer::len(self : PrioritizedBuffer) -> Int {
self.size
}
///|
/// Update priority for transition at index `i`.
pub fn PrioritizedBuffer::update_priority(
self : PrioritizedBuffer,
i : Int,
abs_delta : Float,
) -> Unit {
let new_p = pow_fast(abs_delta, self.alpha) + self.epsilon
let old_p = self.priorities[i]
self.priorities[i] = new_p
self.total_priority = self.total_priority + new_p - old_p
}
///|
/// Approximate x^y via expf(y * logf(x)) for x ≥ 0. (Used for
/// priority^(alpha).)
pub fn pow_fast(x : Float, y : Float) -> Float {
if x <= 0.0F {
return 0.0F
}
expf(y * logf(x))
}
///|
/// Sample a mini-batch with probability ∝ priority^alpha. Returns
/// indices, transitions, and importance-sampling weights.
pub fn PrioritizedBuffer::sample(
self : PrioritizedBuffer,
batch_size : Int,
rng : Xoshiro,
) -> (Array[Int], Array[Int], Array[Int], Array[Float], Array[Int], Array[Bool], Array[Float]) {
let indices : Array[Int] = Array::make(batch_size, 0)
let states : Array[Int] = Array::make(batch_size, 0)
let actions : Array[Int] = Array::make(batch_size, 0)
let rewards : Array[Float] = Array::make(batch_size, 0.0F)
let next_states : Array[Int] = Array::make(batch_size, 0)
let dones : Array[Bool] = Array::make(batch_size, false)
let is_weights : Array[Float] = Array::make(batch_size, 0.0F)
if self.size == 0 || self.total_priority <= 0.0F {
return (indices, states, actions, rewards, next_states, dones, is_weights)
}
for k in 0..= target {
found_idx = i
break
}
i = i + 1
}
indices[k] = found_idx
states[k] = self.states[found_idx]
actions[k] = self.actions[found_idx]
rewards[k] = self.rewards[found_idx]
next_states[k] = self.next_states[found_idx]
dones[k] = self.dones[found_idx]
// Importance-sampling weight: w_i = (N · P(i))^(-β), normalised by max.
// P(i) = priority_i / total_priority
let p_i = self.priorities[found_idx] / self.total_priority
let n_over_p = Float::from_int(self.size) / p_i
let w_raw = pow_fast(n_over_p, -self.beta)
is_weights[k] = w_raw
}
// Normalise IS weights by max.
let mut max_w = is_weights[0]
for k in 1.. max_w {
max_w = is_weights[k]
}
}
if max_w > 0.0F {
for k in 0.. Float {
// Bootstrap: γ^n_steps · max_Q
let mut g = q_next_max
let mut i = 0
while i < n_steps {
g = g * gamma
i = i + 1
}
// Accumulate γ^k · r[t+k] for k = 0..n_steps-1
let mut discount = 1.0F
let mut k = 0
while k < n_steps {
let idx = t + k
if idx >= rewards.length() {
break
}
if dones[idx] {
// Truncate cleanly: zero g, stop.
g = 0.0F
break
}
g = g + rewards[idx] * discount
discount = discount * gamma
k = k + 1
}
g
}