// sac.mbt — Soft Actor-Critic (v0.38.0).
//
// SAC (Haarnoja et al. 2018) for discrete actions on GridWorld.
// Stochastic softmax policy + twin Q-networks + entropy bonus.
//
// Key components:
//   - SoftmaxPolicy (linear)         — π(a|s) = softmax(W · x)
//   - Q-net (linear)                  — Q(s, a) = W_q · x + b_q
//   - Twin Q-nets + target Q-nets (Polyak averaged)
//   - ReplayBuffer (reused from dqn.mbt)
//
// Soft Bellman target (discrete):
//   Q_target(s, a) = r + γ · Σ_{a'} π(a'|s') · [Q̂_min(s', a') - α · log π(a'|s')]
// where Q̂_min = min(Q̂1, Q̂2).
//
// Actor loss: maximise Σ_a π(a|s) · α · log π(a|s) - Q(s, a)
// (equivalently minimise the negative).
//
// Critic loss: MSE(Q(s,a), Q_target(s,a)) for both Q1 and Q2.

///|
/// SAC policy + critic + target critics + temperature α.
pub struct Sac {
  policy : LinearSoftmaxPolicy
  q1 : LinearQNet
  q2 : LinearQNet
  q1_target : LinearQNet
  q2_target : LinearQNet
  /// Entropy temperature. Higher = more exploration.
  alpha : Float
}

///|
pub fn Sac::new(
  n_states : Int,
  n_actions : Int,
  alpha : Float,
  seed : UInt64,
) -> Sac {
  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)
  // Sync target = online initially.
  qnet_copy(q1_target, q1)
  qnet_copy(q2_target, q2)
  { policy, q1, q2, q1_target, q2_target, alpha }
}

///|
/// Sample action from the softmax policy.
pub fn sac_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
}

///|
/// Compute the SAC soft Bellman target for a single transition.
/// Q_target(s, a) = r + γ · (1 - done) · Σ_a' π(a'|s') · [Q̂_min(s', a') - α · log π(a'|s')]
pub fn sac_soft_target(
  sac : Sac,
  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 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 - sac.alpha * lp)
  }
  reward + gamma * expected
}

///|
/// Update the SAC critics for one batch. Returns mean squared TD error.
pub fn sac_critic_update(
  sac : Sac,
  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
  for k in 0.. 0.0F { logf(pi[a]) } else { -20.0F }
        h_contrib = h_contrib + (lp + 1.0F) * (indicator - pi[a])
      }
      // Apply descent on Q - α·H (negative sign from log-of-policy gradient).
      let grad_row = lr * (q_contrib + sac.alpha * h_contrib)
      let row = sac.policy.w[i]
      for j in 0.. 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_critic_update(sac, ss, aa, rr, nss, dd, gamma, lr_critic)
      let _ = sac_actor_update(sac, ss, lr_actor)
      sac_soft_update(sac, tau)
    }
  }
  if return_count > 0 {
    total_return / Float::from_int(return_count)
  } else {
    0.0F
  }
}