// c51.mbt — Categorical DQN (Bellemare et al. 2017) — v0.39.0.
//
// Distributional RL: instead of predicting a scalar Q-value, predict
// a probability distribution over a discrete support of "atoms".
//
//   Support: z_i = V_min + i · Δz,    i ∈ {0, 1, ..., N-1}, N = 51
//   Δz = (V_max - V_min) / (N - 1)
//
// Q-recovery:
//   Q(s, a) = Σ_i p(s, a, i) · z_i
//
// Per-transition projected target distribution (the key step):
//   For atom j of next-state distribution, compute
//     Tz_j = r + γ · (1 − done) · z_j
//   Clamp to [V_min, V_max], then map to support:
//     b_j = (Tz_j − V_min) / Δz
//   Project probability mass from j to atoms ⌊b_j⌋ and ⌈b_j⌉:
//     m[⌊b_j⌋] += p(s', a*, j) · (⌈b_j⌉ − b_j)
//     m[⌈b_j⌉] += p(s', a*, j) · (b_j − ⌊b_j⌋)
//   where a* = argmax_a Σ_i p(s', a, i) · z_i.
//
// Loss: cross-entropy between projected target m and predicted p(s, a).
//   L = −Σ_i m[i] · log p(s, a, i)
// Updates: SGD on the (n_states, n_atoms) logits for the chosen (s, a).

///|
/// Xavier-normal init for a 3D weight tensor (rows × cols × depth).
/// All entries are independent N(0, std²).
fn xavier_normal_3d(
  rows : Int,
  cols : Int,
  depth : Int,
  std : Float,
  rng : Xoshiro,
) -> Array[Array[Array[Float]]] {
  let w : Array[Array[Array[Float]]] = Array::make(rows, [])
  for i in 0.. C51Net {
  let std = sqrtf(1.0F / Float::from_int(n_states * n_atoms))
  let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let w = xavier_normal_3d(n_actions, n_states, n_atoms, std, rng)
  let b : Array[Array[Float]] = Array::make(n_actions, [])
  for a in 0.. Array[Array[Float]] {
  let n_a = net.n_actions
  let n_i = net.n_atoms
  let logits : Array[Array[Float]] = Array::make(n_a, [])
  for a in 0..= 0 && state < net.n_states {
      for i in 0.. Array[Float] {
  let n = logits.length()
  let mut max_l = logits[0]
  for i in 1.. max_l {
      max_l = logits[i]
    }
  }
  let mut sum_e = 0.0F
  let exps : Array[Float] = Array::make(n, 0.0F)
  for i in 0.. Array[Array[Float]] {
  let logits = c51_net_logits(net, state)
  let probs : Array[Array[Float]] = Array::make(net.n_actions, [])
  for a in 0.. Float {
  let n = probs.length()
  let mut s = 0.0F
  for i in 0.. Array[Float] {
  let support : Array[Float] = Array::make(n_atoms, 0.0F)
  if n_atoms > 1 {
    let delta_z = (v_max - v_min) / Float::from_int(n_atoms - 1)
    for i in 0.. C51 {
  let net = C51Net::new(n_states, n_actions, n_atoms, v_min, v_max, seed)
  let net_target = C51Net::new(
    n_states, n_actions, n_atoms, v_min, v_max, seed + 100UL,
  )
  c51_copy(net_target, net)
  let support = c51_support(v_min, v_max, n_atoms)
  let delta_z = if n_atoms > 1 {
    (v_max - v_min) / Float::from_int(n_atoms - 1)
  } else {
    1.0F
  }
  { net, net_target, gamma, lr, n_atoms, v_min, v_max, support, delta_z }
}

///|
/// Compute the projected target distribution for one transition.
/// Returns m[i] (n_atoms) — a probability distribution over atoms.
///
/// The "greedy" next action is a* = argmax_a Σ_i p(s', a, i) · z_i
/// (Double DQN-style online-net evaluation is not implemented here —
/// we use the target net for both selection and evaluation, which is
/// the original Bellemare 2017 formulation).
pub fn c51_projected_target(
  c51 : C51,
  reward : Float,
  next_state : Int,
  done : Bool,
) -> Array[Float] {
  let n = c51.n_atoms
  let m : Array[Float] = Array::make(n, 0.0F)
  if done {
    // Target = reward projected onto a single atom (delta distribution).
    let tz = if reward < c51.v_min {
      c51.v_min
    } else if reward > c51.v_max {
      c51.v_max
    } else {
      reward
    }
    let b = (tz - c51.v_min) / c51.delta_z
    let lo = Float::to_int(b)
    let lo_clamped = if lo < 0 {
      0
    } else if lo >= n - 1 {
      n - 2
    } else {
      lo
    }
    let hi = lo_clamped + 1
    let frac = b - Float::from_int(lo_clamped)
    m[lo_clamped] = 1.0F - frac
    m[hi] = frac
    return m
  }
  // 1) Greedy next-action a* via target net.
  let probs_next = c51_forward(c51.net_target, next_state)
  let mut best_a = 0
  let mut best_q = c51_q_value(probs_next[0], c51.support)
  for a in 1.. best_q {
      best_q = q
      best_a = a
    }
  }
  // 2) For each atom j, project Tz_j = r + γ · z_j into support.
  let p_next = probs_next[best_a]
  for j in 0.. c51.v_max {
      c51.v_max
    } else {
      tz
    }
    let b = (tz_clamped - c51.v_min) / c51.delta_z
    let lo = Float::to_int(b)
    let lo_clamped = if lo < 0 {
      0
    } else if lo >= n - 1 {
      n - 2
    } else {
      lo
    }
    let hi = lo_clamped + 1
    let frac = b - Float::from_int(lo_clamped)
    m[lo_clamped] = m[lo_clamped] + p_next[j] * (1.0F - frac)
    m[hi] = m[hi] + p_next[j] * frac
  }
  m
}

///|
/// Greedy action for state `s` based on the target net.
/// a* = argmax_a Σ_i p(s, a, i) · z_i.
pub fn c51_greedy_action(c51 : C51, state : Int) -> Int {
  let probs = c51_forward(c51.net_target, state)
  let mut best_a = 0
  let mut best_q = c51_q_value(probs[0], c51.support)
  for a in 1.. best_q {
      best_q = q
      best_a = a
    }
  }
  best_a
}

///|
/// One C51 update step: cross-entropy loss + SGD on (s, a) logits.
///
/// For each transition in the batch:
///   m = projected_target(r, s', done)        // n_atoms
///   p = softmax(logits(s)[a])                 // n_atoms
///   L = -Σ_i m[i] · log p[i]                // cross-entropy
///   d_logits[a][i] = (p[i] - m[i])            // gradient of CE wrt logits
///
/// Returns mean cross-entropy loss.
pub fn c51_update(
  c51 : C51,
  states : Array[Int],
  actions : Array[Int],
  rewards : Array[Float],
  next_states : Array[Int],
  dones : Array[Bool],
) -> Float {
  let n = states.length()
  let mut total_loss = 0.0F
  for k in 0.. 0.0F { logf(p[i]) } else { -20.0F }
      loss = loss - m[i] * lp
    }
    total_loss = total_loss + loss
    // Gradient: d_logits[i] = p[i] - m[i]  (CE derivative w.r.t. logit).
    let g = 2.0F * c51.lr
    for i in 0..= 0 && s < c51.net.n_states {
        c51.net.w[a][s][i] = c51.net.w[a][s][i] - d_logit
        c51.net.b[a][i] = c51.net.b[a][i] - d_logit
      }
    }
  }
  total_loss / Float::from_int(n)
}

///|
/// Polyak soft update of target net from online net.
pub fn c51_soft_update(c51 : C51, tau : Float) -> Unit {
  let n_a = c51.net.n_actions
  let n_s = c51.net.n_states
  let n_i = c51.net.n_atoms
  for a in 0.. Unit {
  let n_a = source.n_actions
  let n_s = source.n_states
  let n_i = source.n_atoms
  for a in 0..