// dsac.mbt — Distributional SAC (Ma et al. 2021) for discrete-action control.
//
// SAC (Haarnoja 2018) learns a Q-function and a softmax policy with an entropy
// bonus. Standard SAC uses a *scalar* Q-value; DSAC replaces the critic head
// with a *Gaussian* distribution Q(s, a) ~ N(μ(s, a), σ²(s, a)). The target
// distribution is the SAC soft-Bellman target plus noise (Gaussian mixture of
// the min-of-two critics). The critic loss is a Gaussian negative log-likelihood
// (NLL) on the sampled target scalar.
//
// Reference: Ma, Wang, Chen, Liu, Wei, "Distributional Soft Actor-Critic for
// Off-policy Reinforcement Learning", 2021.
//
// Key differences vs. `sac.mbt`:
//   - `LinearGaussianQNet` outputs μ AND σ (one Float each) per action.
//   - Critic loss: NLL of target scalar under N(μ, σ²). σ is updated explicitly
//     (not just through the gradient of an MSE loss).
//   - Actor loss: same as SAC (entropy-augmented expected return under π).

///|
/// Linear Q-network with Gaussian heads. Each (state, action) pair has its own
/// μ and σ. Stored as two parallel weight matrices + bias vectors.
pub struct LinearGaussianQNet {
  n_states : Int
  n_actions : Int
  // Layout: w_mu[a][s] = μ-coefficient for action a on state dim s.
  w_mu : Array[Array[Float]]
  b_mu : Array[Float]
  w_sigma : Array[Array[Float]]
  b_sigma : Array[Float]
}

///|
/// Construct a Gaussian critic. Initial μ is small-random, σ is initialised
/// to 1.0 + small jitter (in log-space for stability under direct σ updates).
pub fn LinearGaussianQNet::new(
  n_states : Int,
  n_actions : Int,
  seed : UInt64,
) -> LinearGaussianQNet {
  let rng = Xoshiro::from_state(seed, seed + 7UL, seed + 13UL, seed + 17UL)
  let bound = 1.0F / Float::from_int(n_states).to_double().sqrt().to_float()
  let w_mu : Array[Array[Float]] = []
  let b_mu : Array[Float] = []
  let w_sigma : Array[Array[Float]] = []
  let b_sigma : Array[Float] = []
  for a in 0.. (Array[Float], Array[Float]) {
  let mus : Array[Float] = Array::make(net.n_actions, 0.0F)
  let sigmas : Array[Float] = Array::make(net.n_actions, 1.0F)
  for a in 0.. 10.0F {
      10.0F
    } else {
      exp_ls
    }
  }
  (mus, sigmas)
}

///|
/// Compute Gaussian NLL: -log p(y | μ, σ) = 0.5·((y-μ)/σ)² + log σ + 0.5·log(2π).
/// Returns a scalar Float.
pub fn gaussian_nll(y : Float, mu : Float, sigma : Float) -> Float {
  let diff = y - mu
  let z = diff / sigma
  let half_z2 = 0.5F * z * z
  let log_sigma = Float::from_double(@math.ln(sigma.to_double()))
  half_z2 + log_sigma + 0.9189385F // 0.5·log(2π)
}

///|
/// Compute the DSAC soft-Bellman target scalar for a transition. Returns a
/// single sample from the target distribution (a Gaussian centred on the SAC
/// soft-target of the min-of-two critics with a stochastic perturbation).
pub fn dsac_soft_target(
  q1_target : LinearGaussianQNet,
  q2_target : LinearGaussianQNet,
  policy : LinearSoftmaxPolicy,
  alpha : Float,
  next_state : Int,
  reward : Float,
  done : Bool,
  gamma : Float,
  rng : Xoshiro,
) -> Float {
  if done {
    return reward
  }
  let x_next = rl_one_hot(next_state, policy.n_states)
  let (_l, pi) = policy_forward(policy, x_next)
  let (mu1, sig1) = gaussian_q_forward(q1_target, x_next)
  let (mu2, sig2) = gaussian_q_forward(q2_target, x_next)
  let n_a = policy.n_actions
  let mut expected = 0.0F
  for i in 0.. DSac {
  let policy = LinearSoftmaxPolicy::new(n_states, n_actions, seed)
  let q1 = LinearGaussianQNet::new(n_states, n_actions, seed + 1UL)
  let q2 = LinearGaussianQNet::new(n_states, n_actions, seed + 2UL)
  let q1_target = LinearGaussianQNet::new(n_states, n_actions, seed + 3UL)
  let q2_target = LinearGaussianQNet::new(n_states, n_actions, seed + 4UL)
  gaussian_qnet_copy(q1_target, q1)
  gaussian_qnet_copy(q2_target, q2)
  { policy, q1, q2, q1_target, q2_target, alpha }
}

///|
/// Copy weights from `src` to `dst` (all four: w_mu, b_mu, w_sigma, b_sigma).
pub fn gaussian_qnet_copy(
  dst : LinearGaussianQNet,
  src : LinearGaussianQNet,
) -> Unit {
  for a in 0.. Float {
  let n = states.length()
  let mut total_nll = 0.0F
  for k in 0.. Unit {
  for a in 0.. Int {
  sac_sample_action(agent.policy, state, rng)
}