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