// ppo_kl.mbt — PPO with KL penalty (v0.35.2).
//
// KL-penalized PPO objective (without explicit clipping):
//
//   L_KL(θ) = E[ r_t(θ) · A_t ] - β · KL(π_θ(·|s_t) || π_θ_old(·|s_t))
//
// where r_t(θ) = π_θ(a_t|s_t) / π_θ_old(a_t|s_t) and the per-state
// KL is the standard forward KL.
//
// Gradient of -β·KL w.r.t. θ:
//   ∇_θ(-β·KL) = -β · Σ_a (log(π_θ(a|s) / π_old(a|s)) + 1) · ∇_θ π_θ(a|s)
//
// Combined with the surrogate gradient:
//   ∇_θ(-L_KL) = -A_t · ∇_θ r_t(θ)  -  β · (∇_θ KL term above)
//
// For LinearSoftmaxPolicy:
//   ∇_θ π_θ(a|s) = (1{i==a} - π_θ(i|s)) · x[j]  for W[i,j]

///|
/// Batch storing per-step (state, action, advantage) plus the full
/// per-state old distribution (for the KL term).
pub(all) struct PpoKlBatch {
  states : Array[Int]
  actions : Array[Int]
  advantages : Array[Float]
  old_probs : Array[Float]  // π_old(a_t|s_t)
  old_dist : Array[Array[Float]]  // π_old(·|s_t), full distribution per state
}

///|
/// Collect a batch with GAE-computed advantages. Uses a value net
/// for advantage computation but does NOT update it (this variant
/// focuses on the KL penalty term; assumes V is held fixed or
/// pre-trained).
pub fn ppo_kl_collect_batch(
  env : GridWorld,
  policy : LinearSoftmaxPolicy,
  value_net : LinearValueNet,
  n_episodes : Int,
  gamma : Float,
  max_steps : Int,
  seed : UInt64,
) -> PpoKlBatch {
  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 old_dist : Array[Array[Float]] = []
  for _ep in 0..= 0 {
      running_r = ep_rewards[i] + gamma * running_r
      rets[i] = running_r
      i = i - 1
    }
    // GAE advantages with λ=1 (full Monte-Carlo = G_t - V(s_t)).
    for t in 0.. Float {
  let n = pi_cur.length()
  let mut kl = 0.0F
  for i in 0.. 0.0F && pi_old[i] > 0.0F {
      kl = kl + pi_cur[i] * logf(pi_cur[i] / pi_old[i])
    }
  }
  kl
}

///|
/// One PPO-KL update step. Returns the mean per-state KL after the
/// update (for monitoring). Modifies `policy.w` in place.
pub fn ppo_kl_update_step(
  policy : LinearSoftmaxPolicy,
  batch : PpoKlBatch,
  kl_beta : Float,
  lr : Float,
) -> Float {
  let n = batch.states.length()
  let n_a = policy.n_actions
  let n_s = policy.n_states
  let mut total_kl = 0.0F
  for step in 0.. 0.0F && old_dist[k] > 0.0F {
          logf(probs[k] / old_dist[k])
        } else {
          0.0F
        }
        grad_row = grad_row - lr * kl_beta * (lp + 1.0F) * (indicator_k - probs[k])
      }
      let row = policy.w[i]
      for j in 0.. 0 {
    total_kl / Float::from_int(n)
  } else {
    0.0F
  }
}

///|
/// Train PPO-KL for `n_iters`. Returns the last mean episode return.
pub fn train_ppo_kl(
  env : GridWorld,
  policy : LinearSoftmaxPolicy,
  value_net : LinearValueNet,
  n_iters : Int,
  n_episodes_per_iter : Int,
  gamma : Float,
  kl_beta : 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
}