// 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)
}