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