// lstm_ppo.mbt — Recurrent PPO with LSTM policy (v0.35.3).
//
// Recurrent PPO on the POMDP corridor. The forward pass is the
// `LstmPolicy` from `lstm_policy.mbt` (v0.33.2); the update step
// applies PPO-style clipped surrogate with the importance ratio
// computed from cached π_old and the (re-derived) current π.
//
// Note: this is a single-pass online PPO — the cached probs serve
// as π_old, and the current π is the policy at update time. The
// importance ratio deviates from 1 only because updating `w_out`
// shifts the policy between when we cached probs and when we
// recompute them. In practice with 1 epoch and small lr, ratio
// stays close to 1.

///|
/// Per-step transition record for recurrent PPO. Includes both the
/// cached π_old (at batch collection) and the recomputed π_cur
/// (at update time, after policy weights have shifted).
pub(all) struct RecurrentPpoBatch {
  states : Array[Int]
  actions : Array[Int]
  advantages : Array[Float]
  old_probs : Array[Float]
  cur_probs : Array[Float]
  cur_dist : Array[Array[Float]]
  old_dist : Array[Array[Float]]
}

///|
/// Collect a batch from the LSTM policy on the corridor POMDP.
/// Uses Monte-Carlo returns as advantages (no value baseline).
pub fn lstm_ppo_collect_batch(
  env_n_cells : Int,
  policy : LstmPolicy,
  n_episodes : Int,
  gamma : Float,
  max_steps : Int,
  seed : UInt64,
) -> RecurrentPpoBatch {
  let rng = Xoshiro::from_state(seed, seed + 7UL, seed + 13UL, seed + 17UL)
  let states : Array[Int] = []
  let actions : Array[Int] = []
  let advantages : Array[Float] = []
  let old_probs : Array[Float] = []
  let cur_probs : Array[Float] = []
  let cur_dist : Array[Array[Float]] = []
  let old_dist : Array[Array[Float]] = []
  let start = env_n_cells / 2
  for _ep in 0.. 0.0 { 1 } else { 0 }
    let env : CorridorEnv = CorridorEnv::{
      n_cells: env_n_cells,
      goal_side: side,
      step_penalty: -0.1F,
      goal_reward: 1.0F,
      max_steps,
    }
    let obs_seq : Array[Array[Float]] = []
    for _t in 0.. Float {
  let lower = 1.0F - clip_eps
  let upper = 1.0F + clip_eps
  let n_actions = policy.w_out.length()
  let d_h = policy.cell.d_h
  let n = batch.states.length()
  let mut total_kl = 0.0F
  for t in 0..= 0.0F {
      if r > upper {
        upper
      } else {
        r
      }
    } else {
      if r < lower {
        lower
      } else {
        r
      }
    }
    let scale = advantage * effective
    for i in 0.. 0.0F && old_d[i] > 0.0F {
        kl = kl + cur_d[i] * logf(cur_d[i] / old_d[i])
      }
    }
    total_kl = total_kl + kl
  }
  if n > 0 {
    total_kl / Float::from_int(n)
  } else {
    0.0F
  }
}

///|
/// Train recurrent PPO for `n_iters`. Returns the last mean
/// per-step advantage.
pub fn train_lstm_ppo(
  env_n_cells : Int,
  policy : LstmPolicy,
  n_iters : Int,
  n_episodes_per_iter : Int,
  gamma : Float,
  clip_eps : Float,
  lr : Float,
  max_steps : Int,
  seed : UInt64,
) -> 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
}