// dqn_n_step.mbt — N-step DQN (Sutton 1988; Mnih et al. 2016 §4.2).
//
// 1-step TD target: y_t = r_t + γ · max_a' Q̂(s_{t+1}, a')
// n-step TD target: y_t = Σ_{k=0}^{n-1} γ^k · r_{t+k} + γ^n · max_a' Q̂(s_{t+n}, a')
//
// With truncation: if any transition r_{t+k} is terminal, the sum stops at k
// and γ^n · max ... is replaced by 0 (no bootstrap from a terminal future).
//
// We combine this with the standard (uniform) `ReplayBuffer` from `dqn.mbt`.
// Pairing with `PrioritizedReplayBuffer` is also possible — the priority can
// be the magnitude of the n-step TD error.
//
// Per-step pipeline (during a rollout):
// 1. nstep.push(s, a, r, s', done)
// 2. If nstep.len() == n OR done:
// extract n-step transition G, h, s_root, s_lookahead
// if (h == n) or (s_lookahead corresponds to truncated terminal):
// replay.push(s_root, a_root, G, s_lookahead, done_flag)
// if done: nstep.reset()
// 3. If replay.len() >= batch_size and ep >= warmup:
// sample mini-batch
// dqn_update_step_n_step(...)
// 4. Periodically copy online → target.
///|
/// N-step variant of `dqn_update_step`. Each row in the batch is an n-step
/// transition:
///
/// (s_root, a_root, G_n, s_lookahead, h_eff, use_bootstrap)
///
/// where:
///
/// * `G_n` — discounted reward sum Σ γ^k r_{root+k}
/// * `s_lookahead` — the s_{root + h_eff} used for bootstrap
/// * `h_eff` — effective horizon (number of rewards summed, 1..=n)
/// * `use_bootstrap` — true if `h_eff == n` and the n-step window is non-terminal
///
/// Loss for a row (with bootstrap):
///
/// δ = G_n + γ^h_eff · max_a' Q̂(s_lookahead, a') - Q(s_root, a_root)
///
/// Without bootstrap (terminal reached inside window):
///
/// δ = G_n - Q(s_root, a_root)
///
/// Returns mean |δ| over the batch.
pub fn dqn_update_step_n_step(
online : LinearQNet,
target : LinearQNet,
states : Array[Int],
actions : Array[Int],
rewards : Array[Float], // G_n
next_states : Array[Int], // s_lookahead
h_effs : Array[Int], // effective horizons
use_bootstraps : 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 nstep = NStepBuffer::new(n_step)
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 _ = dqn_update_step_n_step_via_replay(q_net, target, ss, aa, rr, nss, dd, gamma, lr, n_step)
}
// Sync target.
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
}
}
///|
/// Adapter: replay buffer rows don't carry h_eff or use_bootstrap, but we can
/// reconstruct them when n == h_eff and the row's `done` flag is false. For
/// rows with done == true, h_eff is 1 and bootstrap is disabled.
///
/// This adapter assumes the training pipeline produced rows with h_eff = n_step
/// (no early truncation). The flag `done == true` is interpreted as
/// "do not bootstrap".
fn dqn_update_step_n_step_via_replay(
online : LinearQNet,
target : LinearQNet,
states : Array[Int],
actions : Array[Int],
rewards : Array[Float],
next_states : Array[Int],
dones : Array[Bool],
gamma : Float,
lr : Float,
n_step : Int,
) -> Float {
let n = states.length()
let h_arr : Array[Int] = Array::make(n, n_step)
let use_arr : Array[Bool] = Array::make(n, true)
for k in 0..