// ppo.mbt — PPO (clipped surrogate) on GridWorld (v0.35.0).
//
// Proximal Policy Optimization with the clipped surrogate objective:
//
//   r_t(θ) = π_θ(a_t | s_t) / π_θ_old(a_t | s_t)
//   L_clip(θ) = E[ min( r_t · A_t, clip(r_t, 1-ε, 1+ε) · A_t ) ]
//
// We use A_t = G_t (full Monte-Carlo return) as the advantage
// (no value baseline; simplest version). The clipped term removes
// gradient signal once the importance ratio leaves [1-ε, 1+ε],
// keeping policy updates conservative.
//
// Gradient of the (negative) loss per transition:
//   g_clip[i, j] = -A_t · effective_ratio · (1{i == a_t} - π_θ(i|s_t)) · x_t[j]
// where effective_ratio = r_t if unclipped, else clip(r_t, 1-ε, 1+ε).
//
// Policy: `LinearSoftmaxPolicy` from reinforce.mbt. Environment:
// `GridWorld` (also from reinforce.mbt).

///|
/// A batch of transitions collected under a "frozen" policy snapshot.
pub(all) struct PpoBatch {
  states : Array[Int]
  actions : Array[Int]
  returns : Array[Float]
  old_probs : Array[Float]  // π_old(a_t | s_t) recorded at collection time
  n_episodes : Int
}

///|
/// Collect a batch of episodes under the current policy. Records
/// `old_probs[a_t | s_t]` for use as the importance-ratio denominator.
pub fn ppo_collect_batch(
  env : GridWorld,
  policy : LinearSoftmaxPolicy,
  n_episodes : Int,
  gamma : Float,
  max_steps : Int,
  seed : UInt64,
) -> PpoBatch {
  let rng = Xoshiro::from_state(seed, seed + 7UL, seed + 13UL, seed + 17UL)
  let states : Array[Int] = []
  let actions : Array[Int] = []
  let returns_arr : Array[Float] = []
  let old_probs : Array[Float] = []
  for _ep in 0.. Unit {
  let n = batch.states.length()
  let n_a = policy.n_actions
  let n_s = policy.n_states
  let lower = 1.0F - clip_eps
  let upper = 1.0F + clip_eps
  for step in 0.. 0: active if r ≤ upper (otherwise clipped to upper).
    // For A = g < 0: active if r ≥ lower (otherwise clipped to lower).
    let effective = if g >= 0.0F {
      if r > upper {
        upper
      } else {
        r
      }
    } else {
      if r < lower {
        lower
      } else {
        r
      }
    }
    // d_logits[i] = (1{i == a} - probs[i]) * (-g * effective)
    // The negative sign comes from converting "maximize L_clip" to
    // "minimize -L_clip". We treat g as the advantage and apply a
    // gradient step on policy weights that moves them in the
    // direction of increased log π for positive g.
    let scale = g * effective
    for i in 0.. Float {
  let mut last_mean = 0.0F
  let mut s = seed
  for _iter in 0..