// gtrxl_sac_agent.mbt — SAC_GTrXL agent: stochastic recurrent SAC with
// GTrXL block memory for partially-observable continuous-control POMDPs
// (v0.75.0).
//
// Wires:
//   - GTrXLSACActor (stochastic actor, v0.75.0)
//   - twin GTrXLQNetworkContinuous critics (Q-networks, v0.74.0)
//   - 5 Polyak-averaged target networks (actor + 2 critics + 3 targets)
//   - automatic alpha (entropy temperature) via `log_alpha` parameter
//     and `target_entropy = -action_dim` default
//
// Scope of v0.75.0:
//   - SAC_GTrXL struct + constructor (6 distinct seeds)
//   - sac_gtrxl_act (stochastic forward, returns action + block_cache + log_prob)
//   - sac_gtrxl_compute_td_target_seq (twin-critic min, no log-prob needed)
//   - sac_gtrxl_soft_update (Polyak averaging on all 5 targets — each
//     network has MLP w1/b1/w_mean/b_mean/w_log_std/b_log_std + GTrXL
//     block (6 tensors), so 12 tensors per target, 60 total across
//     actor + 2 critics)
//   - sac_gtrxl_update_alpha (gradient step on log_alpha toward target_entropy)
//
// BPTT-driven actor + critic update + SAC soft actor gradient (with
// entropy bonus) are deferred — parallel to Batch F's BPTT follow-up.

///|
/// SAC_GTrXL agent: stochastic actor + twin recurrent critics + 5
/// Polyak-averaged target networks + automatic alpha.
pub(all) struct SAC_GTrXL {
  actor : GTrXLSACActor
  critic1 : GTrXLQNetworkContinuous
  critic2 : GTrXLQNetworkContinuous
  actor_target : GTrXLSACActor
  critic1_target : GTrXLQNetworkContinuous
  critic2_target : GTrXLQNetworkContinuous
  gamma : Float
  mut tau : Float
  d_model : Int
  d_ff : Int
  mut log_alpha : Float
  target_entropy : Float
}

///|
/// Build a fresh SAC_GTrXL. 6 distinct seeds so Polyak targets start
/// far from source networks (they get soft-updated into place).
/// `initial_alpha = 1.0` (log_alpha = 0).
pub fn SAC_GTrXL::new(
  state_dim : Int,
  action_dim : Int,
  d_model : Int,
  d_ff : Int,
  action_low : Float,
  action_high : Float,
  gamma : Float,
  tau : Float,
  seed : UInt64,
) -> SAC_GTrXL {
  let actor : GTrXLSACActor = GTrXLSACActor::new(
    state_dim, action_dim, d_model, d_ff, action_low, action_high, seed,
  )
  let critic1 : GTrXLQNetworkContinuous = GTrXLQNetworkContinuous::new(
    state_dim, action_dim, d_model, d_ff, seed + 1UL,
  )
  let critic2 : GTrXLQNetworkContinuous = GTrXLQNetworkContinuous::new(
    state_dim, action_dim, d_model, d_ff, seed + 2UL,
  )
  let actor_target : GTrXLSACActor = GTrXLSACActor::new(
    state_dim, action_dim, d_model, d_ff, action_low, action_high, seed + 100UL,
  )
  let critic1_target : GTrXLQNetworkContinuous = GTrXLQNetworkContinuous::new(
    state_dim, action_dim, d_model, d_ff, seed + 101UL,
  )
  let critic2_target : GTrXLQNetworkContinuous = GTrXLQNetworkContinuous::new(
    state_dim, action_dim, d_model, d_ff, seed + 102UL,
  )
  {
    actor,
    critic1,
    critic2,
    actor_target,
    critic1_target,
    critic2_target,
    gamma,
    tau,
    d_model,
    d_ff,
    log_alpha : 0.0F,
    target_entropy : -Float::from_int(action_dim),
  }
}

///|
/// Sample a stochastic action for the current observation. Returns
/// `(action, block_cache, log_prob)`. Used at training time (with
/// exploration noise baked into the Gaussian sampling).
pub fn sac_gtrxl_act(
  agent : SAC_GTrXL,
  obs : Array[Float],
  rng : Xoshiro,
) -> (Array[Float], GTrXLTokenCache, Float) {
  let (a, log_prob, cache) = gtrxl_sac_actor_step(agent.actor, obs, rng)
  (a, cache, log_prob)
}

///|
/// Compute the 1-step twin-critic TD target for a T-step sequence.
/// target_t = r_t + γ(1 - d_t) * min(q1_target(next_obs_t, next_act_t),
///                                 q2_target(next_obs_t, next_act_t))
/// `next_act_seq` should be the stochastic action sample from the
/// target actor (caller computes via `gtrxl_sac_select_action` on
/// `actor_target`).
pub fn sac_gtrxl_compute_td_target_seq(
  agent : SAC_GTrXL,
  next_obs_seq : Array[Float],
  next_act_seq : Array[Float],
  reward_seq : Array[Float],
  done_seq : Array[Float],
  seq_len : Int,
) -> Array[Float] {
  let (q1_next, _) = gtrxl_qnetwork_continuous_seq_forward(
    agent.critic1_target, next_obs_seq, next_act_seq, seq_len,
  )
  let (q2_next, _) = gtrxl_qnetwork_continuous_seq_forward(
    agent.critic2_target, next_obs_seq, next_act_seq, seq_len,
  )
  let target : Array[Float] = Array::make(seq_len, 0.0F)
  for t in 0.. Unit {
  let one_minus = 1.0F - tau
  // MLP w1
  for i in 0.. Unit {
  sac_gtrxl_actor_soft_update(agent.actor_target, agent.actor, agent.tau)
  gtrxl_qnetwork_continuous_soft_update(agent.critic1_target, agent.critic1, agent.tau)
  gtrxl_qnetwork_continuous_soft_update(agent.critic2_target, agent.critic2, agent.tau)
}

///|
/// Auto-alpha gradient step. Updates `log_alpha` toward
/// `target_entropy` (typically `-action_dim`). Caller passes
/// `mean_log_prob` = mean of per-timestep log π across the T-step
/// window. Loss: `alpha * (-log_prob - target_entropy)`.
/// Gradient w.r.t. log_alpha: `-(mean_log_prob + target_entropy)`.
pub fn sac_gtrxl_update_alpha(
  agent : SAC_GTrXL,
  mean_log_prob : Float,
  lr : Float,
) -> Unit {
  let target = agent.target_entropy
  // d/d(log_alpha) [alpha * (-log_prob - target)] = -log_prob - target
  let grad_log_alpha = -(mean_log_prob + target)
  // log_alpha <- log_alpha - lr * grad
  agent.log_alpha = agent.log_alpha - lr * grad_log_alpha
}

///|
/// Get the current alpha = exp(log_alpha). Used by callers in
/// actor/critic updates.
pub fn sac_gtrxl_get_alpha(agent : SAC_GTrXL) -> Float {
  expf(agent.log_alpha)
}