// gru_policy.mbt — REINFORCE with GRU policy on POMDP corridor (v0.34.0).
//
// Same setup as `lstm_policy.mbt` but with a GRU cell instead of
// LSTM. GRU has no cell state — only the hidden state evolves:
//   z_t = σ(W_z · [h_{t-1}; x_t] + b_z)
//   r_t = σ(W_r · [h_{t-1}; x_t] + b_r)
//   n_t = tanh(W_n · [r_t ⊙ h_{t-1}; x_t] + b_n)
//   h_t = (1 - z_t) ⊙ h_{t-1} + z_t ⊙ n_t
//
// The output projection W_out · h_t + b_out is identical to the
// LSTM policy. REINFORCE updates accumulate per-step d_logits
// gradients (using the TD-style ∇log π gradient), then backprop
// through the GRU via `gru_cell_backward`.

///|
/// GRU policy: cell + output projection.
pub struct GruPolicy {
  cell : GruCellParam
  w_out : Array[Array[Float]]  // n_actions × d_h
  b_out : Array[Float]
}

///|
/// Build a GRU policy. d_x = 2 + n_cells, d_h = hidden_dim,
/// n_actions = 2 (left/right).
pub fn GruPolicy::new(
  n_cells : Int,
  d_h : Int,
  seed : UInt64,
) -> GruPolicy {
  let d_x = 2 + n_cells
  let cell = GruCellParam::new(d_x, d_h, seed)
  let n_actions = 2
  let std = sqrtf(1.0F / Float::from_int(d_h))
  let rng = Xoshiro::from_state(seed + 1UL, seed + 2UL, seed + 3UL, seed + 4UL)
  let w_out = xavier_normal(n_actions, d_h, std, rng)
  let b_out : Array[Float] = Array::make(n_actions, 0.0F)
  { cell, w_out, b_out }
}

///|
/// Forward cache for the GRU policy.
pub struct GruPolicyCache {
  cell_caches : Array[GruCellCache]
  probs : Array[Array[Float]]
  hs : Array[Array[Float]]
}

///|
/// Run the GRU policy forward. Returns (action_seq, cache).
pub fn gru_policy_forward(
  policy : GruPolicy,
  obs_seq : Array[Array[Float]],
  h0 : Array[Float],
  rng : Xoshiro,
) -> (Array[Int], GruPolicyCache) {
  let n = obs_seq.length()
  let cell_caches : Array[GruCellCache] = []
  let probs : Array[Array[Float]] = []
  let hs : Array[Array[Float]] = []
  let actions : Array[Int] = []
  let mut cur_h = h0
  for t in 0.. max_l {
        max_l = logits[i]
      }
    }
    let ps : Array[Float] = Array::make(n_actions, 0.0F)
    let mut sum_e = 0.0F
    for i in 0.. Unit {
  let actions_n = actions.length()
  let returns_n = returns.length()
  let n = if actions_n < returns_n { actions_n } else { returns_n }
  let n_actions = policy.w_out.length()
  let d_h = policy.cell.d_h
  let d_w_out : Array[Array[Float]] = Array::make(n_actions, [])
  for i in 0.. Unit {
  for i in 0.. Float {
  let rng = Xoshiro::from_state(seed, seed + 7UL, seed + 13UL, seed + 17UL)
  let mut total_return = 0.0F
  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..