// ppo_gae.mbt — PPO with GAE (generalized advantage estimation) (v0.35.1).
//
// Full PPO with value baseline:
// - Critic: V(s; w_v) = `LinearValueNet`
// - TD residual: δ_t = r_t + γ · V(s_{t+1}) · (1 - done) - V(s_t)
// - GAE: A_t = Σ_{k >= 0} (γλ)^k · δ_{t+k}
// - PPO clipped surrogate: L_clip = E[ min(r·A, clip(r,1-ε,1+ε)·A) ]
// - Critic update: w_v -= α_v · ∇_w_v (V(s) - R)^2 / N
//
// Reuses:
// - `LinearSoftmaxPolicy` from reinforce.mbt
// - `LinearValueNet` from actor_critic.mbt
// - `GridWorld` from reinforce.mbt
///|
/// Episode record augmented with V(s_t), V(s_{t+1}), done flags for
/// GAE computation. Mirrors `EpisodeWithValues` from actor_critic.mbt.
pub(all) struct PpoGaeEpisode {
states : Array[Int]
actions : Array[Int]
rewards : Array[Float]
values : Array[Float]
next_values : Array[Float]
dones : Array[Bool]
}
///|
/// Collect a batch of episodes recording V at every step.
pub fn ppo_gae_collect_batch(
env : GridWorld,
policy : LinearSoftmaxPolicy,
value_net : LinearValueNet,
n_episodes : Int,
max_steps : Int,
seed : UInt64,
) -> Array[PpoGaeEpisode] {
let rng = Xoshiro::from_state(seed, seed + 7UL, seed + 13UL, seed + 17UL)
let episodes : Array[PpoGaeEpisode] = []
for _ep in 0.. (Array[Float], Array[Float]) {
let t = episode.states.length()
let advs : Array[Float] = Array::make(t, 0.0F)
let rets : Array[Float] = Array::make(t, 0.0F)
// Compute returns R_t via backwards accumulation.
let mut running_r = 0.0F
let mut i = t - 1
while i >= 0 {
running_r = episode.rewards[i] + gamma * running_r
rets[i] = running_r
i = i - 1
}
// Compute GAE advantages via backwards accumulation.
let mut running_adv = 0.0F
let mut k = t - 1
while k >= 0 {
let v = episode.values[k]
let v_next = episode.next_values[k]
let not_done = if episode.dones[k] { 0.0F } else { 1.0F }
let delta = episode.rewards[k] + gamma * v_next * not_done - v
running_adv = delta + gamma * gae_lambda * not_done * running_adv
advs[k] = running_adv
k = k - 1
}
(advs, rets)
}
///|
/// Apply one PPO update step using GAE advantages. Also updates the
/// value network via simple MSE regression on the returns.
///
/// `clip_eps` is the PPO clipping epsilon (typically 0.2).
pub fn ppo_gae_update(
policy : LinearSoftmaxPolicy,
value_net : LinearValueNet,
episodes : Array[PpoGaeEpisode],
gamma : Float,
gae_lambda : Float,
clip_eps : Float,
lr_policy : Float,
lr_value : Float,
) -> Unit {
let lower = 1.0F - clip_eps
let upper = 1.0F + clip_eps
let n_a = policy.n_actions
let n_s = policy.n_states
for ep in episodes {
let (advs, rets) = gae_advantages_returns(ep, gamma, gae_lambda)
let t = ep.states.length()
// First pass: critic MSE update.
for step in 0.. Float {
let mut last_mean = 0.0F
let mut s = seed
for _iter in 0.. 0 { sum / Float::from_int(count) } else { 0.0F }
s = s + 1UL
}
last_mean
}