// double_dqn.mbt — Double DQN (v0.37.0).
//
// Double DQN (van Hasselt et al. 2016) decouples action selection
// from action evaluation to reduce Q-value overestimation bias:
//
// Vanilla: target = r + γ · max_a' Q̂(s', a'; θ_target)
// Double: a* = argmax_a' Q(s', a'; θ_online) (online net)
// target = r + γ · Q̂(s', a*; θ_target) (target net)
//
// The two nets have the same architecture but may be at different
// optimisation stages. Using the online net to pick the action
// (less biased than the target net, which has accumulated over-
// estimation error) and the target net to evaluate it gives a
// lower-variance TD target.
//
// Reuses:
// - `LinearQNet`, `ReplayBuffer`, `qnet_copy`, `q_forward`,
// `q_argmax`, `q_max` from dqn.mbt
///|
/// Double DQN update step. Same signature as `dqn_update_step`
/// but uses the online net for argmax selection.
///
/// Returns mean |δ| for monitoring.
pub fn double_dqn_update_step(
online : LinearQNet,
target : LinearQNet,
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_abs_delta = 0.0F
for k 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 _ = double_dqn_update_step(q_net, target, ss, aa, rr, nss, dd, gamma, lr)
}
if (ep + 1) % sync_every == 0 {
qnet_copy(target, q_net)
}
}
if return_count > 0 {
total_return / Float::from_int(return_count)
} else {
0.0F
}
}
///|
/// Compute the absolute difference between Double DQN and vanilla
/// DQN target values for a batch. Useful as a sanity check that the
/// decoupled rule produces different gradients than the coupled rule
/// (which it should whenever online != target).
pub fn double_dqn_vs_vanilla_diff(
online : LinearQNet,
target : LinearQNet,
next_states : Array[Int],
gamma : Float,
) -> Float {
let n = next_states.length()
let mut total_diff = 0.0F
for k in 0..