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