// noisy_dqn.mbt — NoisyDQN agent (v0.43.2).
//
// Fortunato et al. 2018 §3.4: replaces ε-greedy action selection
// from `dqn.mbt::train_dqn` with parametric noise added to the
// NoisyLinear weights. Two NoisyLinear layers, no extra
// exploration hyperparameter — the noise scales σ_w, σ_b are
// themselves learnable parameters updated by gradient descent.
//
// Action selection (training): `argmax noisy_q_forward`.
// Action selection (eval):    `argmax noisy_q_forward_eval` (μ-only).
//
// Gradient update on a mini-batch:
//   δ = r + γ · max_a' Q̂(s', a'; θ_target) · (1 - done) - Q(s, a; θ_online)
//   d_μ_W[a, j]   -= lr · 2 · δ · x[j]
//   d_σ_W[a, j]   -= lr · 2 · δ · x[j] · ε_W[a, j]
//   d_μ_b[a]      -= lr · 2 · δ
//   d_σ_b[a]      -= lr · 2 · δ · ε_b[a]
// For unchosen actions: zero.

///|
/// ε-free action selection. Just greedy over the noisy forward —
/// the noise drives exploration implicitly (paper §3.2).
pub fn noisy_dqn_select_action(
  net : NoisyQNet,
  state : Int,
  rng : Xoshiro,
) -> Int {
  let x = rl_one_hot(state, net.n_states)
  noisy_q_argmax(net, x, rng)
}

///|
/// ε-free action selection at evaluation (μ weights only).
pub fn noisy_dqn_select_action_eval(net : NoisyQNet, state : Int) -> Int {
  let x = rl_one_hot(state, net.n_states)
  noisy_q_argmax_eval(net, x)
}

///|
/// Deep-copy μ + σ weights from `src` to `dst`. Used to sync online
/// → target after every `sync_every` episodes.
pub fn noisy_qnet_copy(dst : NoisyQNet, src : NoisyQNet) -> Unit {
  let src_l1_w = src.layer1.weight_mu
  let dst_l1_w = dst.layer1.weight_mu
  for k in 0.. Float {
  let n = states.length()
  let mut total_abs_delta = 0.0F
  for k in 0..= 0.0F {
      1.0F
    } else {
      -1.0F
    }
    online.layer2.bias_sigma[a] = online.layer2.bias_sigma[a] +
      scale * bias_sigma_proxy * 0.01F
    // Weights: x_norm1 (post-ReLU) needed. Re-run layer1 forward
    // (with same RNG draw that produced q above; in practice the
    // RNG has already been advanced so this re-evaluates to fresh
    // noise — for a per-sample update this is a small approximation,
    // but matches paper §3.4 noise reuse).
    let _ = x
    // Instead of recomputing layer1, use the Jacobian approximation
    // from the chain rule: ∂Q[a]/∂W_l2[a,j] = h1[j]. We approximate
    // h1 ≈ ReLU(layer1_weight_mu @ x) (μ weights, no noise — a
    // noise-free proxy that ignores the forward-time noise on the
    // hidden activations). This is the "linearised" approximation.
    let h1_approx : Array[Float] = Array::make(online.hidden, 0.0F)
    for o in 0.. 0.0F { acc } else { 0.0F }
    }
    // We need ε_w2[a, j] for the σ update. Reconstruct from the
    // outer product of ε_i (size hidden) and ε_j (size n_actions):
    // we don't have them stored, so σ update is approximate using
    // sign(weight_mu) as a noise proxy. This is a simplification —
    // full noisy-DQN SGD would replay the forward and capture ε.
    // We DO update μ correctly.
    for j in 0..= 0.0F {
        1.0F
      } else {
        -1.0F
      }
      online.layer2.weight_sigma[w_idx] = online.layer2.weight_sigma[w_idx] +
        scale * h_j * sigma_proxy * 0.01F
    }
    // ---- Update layer1.weight_mu / bias_mu via chain rule ----
    // Backprop through layer2.weight_mu[a, :].
    for o in 0..= 0.0F {
          1.0F
        } else {
          -1.0F
        }
        online.layer1.weight_sigma[w_l1_idx] = online.layer1.weight_sigma[w_l1_idx] +
          grad_o * x[i] * sigma_proxy * 0.01F
      }
      online.layer1.bias_mu[o] = online.layer1.bias_mu[o] + grad_o
      online.layer1.bias_sigma[o] = online.layer1.bias_sigma[o] +
        grad_o * 0.01F
    }
  }
  total_abs_delta / Float::from_int(n)
}

///|
/// Max of an array (helper for target forward).
fn noisy_q_max(q : Array[Float]) -> Float {
  let mut m = q[0]
  for i in 1.. m {
      m = q[i]
    }
  }
  m
}

///|
/// Train NoisyDQN for `n_episodes` episodes. No ε-greedy schedule
/// — exploration is purely from the noise added in `noisy_q_forward`.
/// Returns the mean undiscounted episode return over the last 10
/// episodes (eval-mode argmax, μ weights).
pub fn train_noisy_dqn(
  env : GridWorld,
  online : NoisyQNet,
  target : NoisyQNet,
  n_episodes : Int,
  gamma : Float,
  lr : Float,
  max_steps : Int,
  buffer_capacity : Int,
  warmup_episodes : Int,
  sync_every : Int,
  batch_size : Int,
  seed : UInt64,
) -> 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 _ = noisy_dqn_update_step(
        online, target, ss, aa, rr, nss, dd, gamma, lr, rng,
      )
    }
    if (ep + 1) % sync_every == 0 {
      noisy_qnet_copy(target, online)
    }
  }
  if return_count > 0 {
    total_return / Float::from_int(return_count)
  } else {
    0.0F
  }
}

///|
/// Evaluate a trained NoisyQNet over `n_episodes` greedy episodes
/// (μ weights, no noise). Returns mean episode return.
pub fn eval_noisy_dqn(
  env : GridWorld,
  online : NoisyQNet,
  n_episodes : Int,
  max_steps : Int,
) -> Float {
  let mut total = 0.0F
  for _ in 0..