// sac_auto_alpha.mbt — SAC with auto-tuned entropy temperature (v0.38.1).
//
// Extension of v0.38.0 SAC where α is itself a learnable parameter.
//
// Algorithm (Haarnoja et al. 2018, Appendix):
//   target_entropy H̄ ∈ ℝ (typically negative; e.g. −|A| for discrete actions)
//   α parameterised as α = exp(log_alpha) ∈ (0, ∞)
//
//   For each batch of states:
//     neg_entropy_s = Σ_a π(a|s) · log π(a|s)   (≤ 0)
//     mean_neg_entropy = E_s[neg_entropy_s]
//     delta = mean_neg_entropy − target_entropy
//     log_alpha ← log_alpha + alpha_lr · delta
//
//   Sign interpretation:
//     delta < 0 → policy is "too random" (mean_neg_entropy below target) → shrink α
//     delta > 0 → policy is "too deterministic" → grow α
//     delta = 0 → equilibrium
//
// The actor and critic use α = exp(log_alpha) in place of the fixed
// `sac.alpha` from v0.38.0.

///|
/// SAC agent with auto-tuned entropy temperature.
pub struct SacAutoAlpha {
  policy : LinearSoftmaxPolicy
  q1 : LinearQNet
  q2 : LinearQNet
  q1_target : LinearQNet
  q2_target : LinearQNet
  /// log of entropy temperature; α = exp(log_alpha).
  mut log_alpha : Float
  /// Target for E[Σ_a π(a|s) log π(a|s)]. Negative scalar.
  /// Typical choice: −|A| for discrete action spaces.
  target_entropy : Float
  /// Learning rate for the log_alpha update.
  alpha_lr : Float
}

///|
pub fn SacAutoAlpha::new(
  n_states : Int,
  n_actions : Int,
  log_alpha_init : Float,
  target_entropy : Float,
  alpha_lr : Float,
  seed : UInt64,
) -> SacAutoAlpha {
  let policy = LinearSoftmaxPolicy::new(n_states, n_actions, seed)
  let q1 = LinearQNet::new(n_states, n_actions, seed + 1UL)
  let q2 = LinearQNet::new(n_states, n_actions, seed + 2UL)
  let q1_target = LinearQNet::new(n_states, n_actions, seed + 3UL)
  let q2_target = LinearQNet::new(n_states, n_actions, seed + 4UL)
  qnet_copy(q1_target, q1)
  qnet_copy(q2_target, q2)
  { policy, q1, q2, q1_target, q2_target, log_alpha: log_alpha_init, target_entropy, alpha_lr }
}

///|
/// Current α = exp(log_alpha).
pub fn sac_auto_alpha_get(sac : SacAutoAlpha) -> Float {
  expf(sac.log_alpha)
}

///|
/// Sample action from the softmax policy.
pub fn sac_auto_alpha_sample_action(
  policy : LinearSoftmaxPolicy,
  state : Int,
  rng : Xoshiro,
) -> Int {
  let x = rl_one_hot(state, policy.n_states)
  let (_l, probs) = policy_forward(policy, x)
  let (a, _) = sample_categorical(probs, rng)
  a
}

///|
/// Soft Bellman target (discrete). Identical to v0.38.0 except α is
/// computed from `log_alpha` instead of the fixed `alpha` field.
pub fn sac_auto_alpha_soft_target(
  sac : SacAutoAlpha,
  next_state : Int,
  reward : Float,
  done : Bool,
  gamma : Float,
) -> Float {
  if done {
    return reward
  }
  let x_next = rl_one_hot(next_state, sac.policy.n_states)
  let (_l, pi) = policy_forward(sac.policy, x_next)
  let q1_t = q_forward(sac.q1_target, x_next)
  let q2_t = q_forward(sac.q2_target, x_next)
  let alpha = sac_auto_alpha_get(sac)
  let mut expected = 0.0F
  let n_a = sac.policy.n_actions
  for i in 0.. 0.0F { logf(pi[i]) } else { -20.0F }
    expected = expected + pi[i] * (q_min - alpha * lp)
  }
  reward + gamma * expected
}

///|
/// Critic MSE update (same as v0.38.0).
pub fn sac_auto_alpha_critic_update(
  sac : SacAutoAlpha,
  states : Array[Int],
  actions : Array[Int],
  rewards : Array[Float],
  next_states : Array[Int],
  dones : Array[Bool],
  gamma : Float,
  lr : Float,
) -> Float {
  let n = states.length()
  let mut total_loss = 0.0F
  for k in 0.. Unit {
  let n = states.length()
  let n_a = sac.policy.n_actions
  let n_s = sac.policy.n_states
  let alpha = sac_auto_alpha_get(sac)
  for k in 0.. 0.0F { logf(pi[a]) } else { -20.0F }
        h_contrib = h_contrib + (lp + 1.0F) * (indicator - pi[a])
      }
      let grad_row = lr * (q_contrib + alpha * h_contrib)
      let row = sac.policy.w[i]
      for j in 0.. 0  →  policy too deterministic  →  α grows
///   delta < 0  →  policy too random          →  α shrinks
pub fn sac_auto_alpha_update_alpha(
  sac : SacAutoAlpha,
  states : Array[Int],
) -> Float {
  let n = states.length()
  let n_a = sac.policy.n_actions
  let n_s = sac.policy.n_states
  let mut sum_neg_ent = 0.0F
  for k in 0.. 0.0F { logf(pi[a]) } else { -20.0F }
      neg_ent = neg_ent + pi[a] * lp
    }
    sum_neg_ent = sum_neg_ent + neg_ent
  }
  let mean_neg_ent = sum_neg_ent / Float::from_int(n)
  let delta = mean_neg_ent - sac.target_entropy
  sac.log_alpha = sac.log_alpha + sac.alpha_lr * delta
  delta
}

///|
/// Polyak soft target update (identical to v0.38.0).
pub fn sac_auto_alpha_soft_update(sac : SacAutoAlpha, tau : Float) -> Unit {
  for i in 0.. Float {
  let rng = Xoshiro::from_state(seed, seed + 7UL, seed + 13UL, seed + 17UL)
  let buffer = ReplayBuffer::new(buffer_capacity)
  let mut total_return = 0.0F
  let mut return_count = 0
  for ep in 0..= warmup_episodes && buffer.len() >= batch_size {
      let (ss, aa, rr, nss, dd) = buffer.sample(batch_size, rng)
      let _ = sac_auto_alpha_critic_update(sac, ss, aa, rr, nss, dd, gamma, lr_critic)
      sac_auto_alpha_actor_update(sac, ss, lr_actor)
      let _ = sac_auto_alpha_update_alpha(sac, ss)
      sac_auto_alpha_soft_update(sac, tau)
    }
  }
  if return_count > 0 {
    total_return / Float::from_int(return_count)
  } else {
    0.0F
  }
}