// sac_gru_actor.mbt — Recurrent SAC (Soft Actor-Critic) stochastic
// actor for partially-observable continuous-control POMDPs (v0.68.0).
//
// Architecture (parallel to v0.58.0 GRUDeterministicPolicy but with
// SAC-style stochastic outputs):
//
//   state_t -> MLP w1 -> ReLU -> GRU cell (over time) -> two MLP
//   branches (mean + log_std) -> tanh squash + Gaussian sampling.
//
//   x_proj_t = ReLU(w1 · state_t + b1)             ∈ R^{hidden}
//   h_t = GRU_cell(x_proj_t, h_{t-1})
//   mean_pre_t = w_mean · h_t + b_mean             ∈ R^{action_dim}
//   log_std_pre_t = w_log_std · h_t + b_log_std     ∈ R^{action_dim}
//   mean_t = squash_to_action_range(tanh(mean_pre_t))
//   a_raw_t = mean_t + exp(log_std_t) * eps_t       (eps ~ N(0, 1))
//   a_t = tanh(a_raw_t) * scale + mid              ∈ R^{action_dim}
//   log π_t = log N(a_raw_t; mean_t, log_std_t) - sum log(1 - tanh²(a_raw_t))
//
// log_std is clamped to [-20, 2] for numerical stability
// (Haarnoja 2018 SAC reference).
//
// Reference: Haarnoja et al. 2018 "Soft Actor-Critic" + Haarnoja 2019
// SAC for discrete action extension to POMDP via recurrent actor.

///|
/// Recurrent SAC stochastic actor. Two output branches (mean, log_std)
/// from the same GRU hidden state. `mean_t` is squashed to
/// [action_low, action_high]; `log_std_t` is clamped to [-20, 2].
pub struct GRUSACActor {
  state_dim : Int
  action_dim : Int
  gru_hidden : Int
  mlp_w1 : Array[Array[Float]]
  mlp_b1 : Array[Float]
  gru : GruCellParam
  mlp_w_mean : Array[Array[Float]]
  mlp_b_mean : Array[Float]
  mlp_w_log_std : Array[Array[Float]]
  mlp_b_log_std : Array[Float]
  action_low : Float
  action_high : Float
  log_std_min : Float
  log_std_max : Float
}

///|
/// Build a fresh `GRUSACActor`. MLP weights use Xavier-normal init
/// (He-style for ReLU); GRU cell uses its own init. log_std biases
/// are init to 0 (so std starts at exp(0)=1).
pub fn GRUSACActor::new(
  state_dim : Int,
  action_dim : Int,
  gru_hidden : Int,
  action_low : Float,
  action_high : Float,
  seed : UInt64,
) -> GRUSACActor {
  let rng1 = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let std1 = sqrtf(2.0F / Float::from_int(state_dim))
  let mlp_w1 = xavier_normal(gru_hidden, state_dim, std1, rng1)
  let mlp_b1 : Array[Float] = Array::make(gru_hidden, 0.0F)
  let gru = GruCellParam::new(gru_hidden, gru_hidden, seed + 4UL)
  let rng2 = Xoshiro::from_state(seed + 5UL, seed + 6UL, seed + 7UL, seed + 8UL)
  let std2 = sqrtf(2.0F / Float::from_int(gru_hidden))
  let mlp_w_mean = xavier_normal(action_dim, gru_hidden, std2, rng2)
  let mlp_b_mean : Array[Float] = Array::make(action_dim, 0.0F)
  let rng3 = Xoshiro::from_state(seed + 9UL, seed + 10UL, seed + 11UL, seed + 12UL)
  let mlp_w_log_std = xavier_normal(action_dim, gru_hidden, std2, rng3)
  let mlp_b_log_std : Array[Float] = Array::make(action_dim, 0.0F)
  {
    state_dim,
    action_dim,
    gru_hidden,
    mlp_w1,
    mlp_b1,
    gru,
    mlp_w_mean,
    mlp_b_mean,
    mlp_w_log_std,
    mlp_b_log_std,
    action_low,
    action_high,
    log_std_min : -20.0F,
    log_std_max : 2.0F,
  }
}

///|
/// Single-step stochastic actor forward. Returns
/// `(action, log_prob, new_hidden)` where:
///   - `action` is the squashed, action_range-scaled sample
///   - `log_prob` is the per-timestep log π(a_t | s_t) (Float scalar = sum of dims)
///   - `new_hidden` is the post-GRU hidden state
pub fn sac_gru_actor_step(
  policy : GRUSACActor,
  state : Array[Float],
  hidden : Array[Float],
  rng : Xoshiro,
) -> (Array[Float], Float, Array[Float]) {
  let x_proj_pre = matvec(policy.mlp_w1, policy.mlp_b1, state)
  let x_proj = relu_forward(x_proj_pre)
  let (hidden_next, _cache) = gru_cell_forward(x_proj, hidden, policy.gru)
  let mean_pre = matvec(policy.mlp_w_mean, policy.mlp_b_mean, hidden_next)
  let log_std_pre = matvec(policy.mlp_w_log_std, policy.mlp_b_log_std, hidden_next)
  let _ = _cache
  // Clamp log_std for numerical safety.
  let log_std : Array[Float] = Array::make(policy.action_dim, 0.0F)
  for i in 0.. policy.log_std_max {
      log_std[i] = policy.log_std_max
    } else {
      log_std[i] = v
    }
  }
  let scale = (policy.action_high - policy.action_low) * 0.5F
  let mid = (policy.action_high + policy.action_low) * 0.5F
  let action : Array[Float] = Array::make(policy.action_dim, 0.0F)
  let mean_tanh : Array[Float] = Array::make(policy.action_dim, 0.0F)
  for i in 0..