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