// 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
}