// srnn_demo.mbt — Spiking Recurrent Neural Network training demo.
//
// Time-unrolled SRNN with IF (Integrate-and-Fire) spiking neuron,
// trained via BPTT with surrogate gradient (Zenke & Ganguli 2018).
//
// Architecture:
//   input (n_steps, n_in): Poisson spike train from a static "image"
//   For t in 0..n_steps-1:
//     I[t]    = W_in @ x[t] + W_rec @ s[t-1]   # recurrent synaptic input
//     v[t]    = α * v[t-1] + I[t]              # IF membrane (leaky)
//     s[t]    = H(v[t] - Vt)                  # 0 or 1 spike (hard forward)
//     y[t]    = W_out @ s[t]                 # readout
//   loss = CE(y[n_steps-1], target)            # loss at final timestep
//
// BPTT:
//   d_v[t] = d_s[t] * σ'(v[t] - Vt; β)        # surrogate gradient
//   d_v[t] += α * d_v[t+1] if t < n_steps-1   # temporal backflow
//   d_s[t] = W_out^T @ d_y[t]                 # readout backward
//   d_s[t] += W_rec^T @ d_v[t+1]             # recurrent gradient
//
// Reuses:
//   v0.22.0 surrogate (fast_sigmoid_surrogate for the backward)
//   v0.13.2 cross_entropy_loss (forward + backward via softmax)

// ---- Synthetic dataset ----

///|
/// Poisson-encode a static "image" vector into a (n_steps, n_in) spike
/// train. Spike probability per (t, i) = `min(1, input[i] * max_rate * dt)`.
pub fn srnn_poisson_encode(
  input : Array[Float],
  n_steps : Int,
  dt : Float,
  max_rate : Float,
  rng : Xoshiro,
) -> Array[Float] {
  let n_in = input.length()
  let out : Array[Float] = Array::make(n_steps * n_in, 0.0F)
  for t in 0.. 1.0F { 1.0F } else { raw_p }
      let r = next_f32(rng)
      out[t * n_in + i] = if r < p { 1.0F } else { 0.0F }
    }
  }
  out
}

///|
/// Synthetic 2-class dataset: 16 samples (8 per class).
/// Class 0: x[0] in [0.7, 1.0]. Class 1: x[0] in [0, 0.3].
pub fn srnn_dataset_build(
  seed : UInt64,
) -> (Array[Float], Array[Int]) {
  let n_classes = 2
  let per_class = 8
  let n_total = n_classes * per_class
  let n_in = 2
  let images : Array[Float] = Array::make(n_total * n_in, 0.0F)
  let labels : Array[Int] = Array::make(n_total, 0)
  let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  for c in 0.. SRNN {
  let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let std_in = sqrtf(1.0F / Float::from_int(n_in))
  let std_rec = sqrtf(1.0F / Float::from_int(n_hidden))
  let std_out = sqrtf(1.0F / Float::from_int(n_hidden))
  let w_in : Array[Float] = Array::make(n_in * n_hidden, 0.0F)
  let w_rec : Array[Float] = Array::make(n_hidden * n_hidden, 0.0F)
  let w_out : Array[Float] = Array::make(n_hidden * n_classes, 0.0F)
  let b_in : Array[Float] = Array::make(n_hidden, 0.0F)
  let b_out : Array[Float] = Array::make(n_classes, 0.0F)
  let mut i = 0
  while i < n_in * n_hidden {
    let (z, _) = box_muller(rng)
    w_in[i] = Float::from_double(z) * std_in
    i = i + 1
  }
  let mut j = 0
  while j < n_hidden * n_hidden {
    let (z, _) = box_muller(rng)
    w_rec[j] = Float::from_double(z) * std_rec
    j = j + 1
  }
  let mut k = 0
  while k < n_hidden * n_classes {
    let (z, _) = box_muller(rng)
    w_out[k] = Float::from_double(z) * std_out
    k = k + 1
  }
  let dt : Float = 1.0F
  let alpha : Float = 0.9F
  let v_thresh : Float = 0.5F
  { n_in, n_hidden, n_classes, dt, alpha, v_thresh,
    w_in, w_rec, w_out, b_in, b_out }
}

///|
/// Forward pass. Returns (loss, cache).
pub fn srnn_forward(
  x : Array[Float],
  n_steps : Int,
  target : Int,
  srnn : SRNN,
  beta : Float,
) -> (Float, SRNNCache) {
  let n_in = srnn.n_in
  let n_hidden = srnn.n_hidden
  let n_classes = srnn.n_classes
  let alpha = srnn.alpha
  let v_thresh = srnn.v_thresh

  let v : Array[Float] = Array::make(n_steps * n_hidden, 0.0F)
  let spikes : Array[Float] = Array::make(n_steps * n_hidden, 0.0F)
  let y : Array[Float] = Array::make(n_steps * n_classes, 0.0F)

  // Run timesteps.
  for t in 0.. 0 { (t - 1) * n_hidden } else { 0 }
    for h in 0.. 0 {
        for j in 0.. 0 { v[(t - 1) * n_hidden + h] } else { 0.0F }
      let v_t = alpha * v_prev + acc
      v[t * n_hidden + h] = v_t
      spikes[t * n_hidden + h] = if v_t > v_thresh { 1.0F } else { 0.0F }
    }
    // y[t] = W_out @ s[t] + b_out
    for c in 0.. Array[Float] {
  let mut max_v = arr[start]
  for i in 1.. max_v {
      max_v = arr[start + i]
    }
  }
  let mut sum_exp = 0.0F
  let out : Array[Float] = Array::make(n, 0.0F)
  for i in 0.. (
  Array[Float],  // d_w_in
  Array[Float],  // d_w_rec
  Array[Float],  // d_w_out
  Array[Float],  // d_b_in
  Array[Float],  // d_b_out
) {
  let n_steps = cache.n_steps
  let n_in = cache.n_in
  let n_hidden = cache.n_hidden
  let n_classes = cache.n_classes
  let alpha = cache.alpha
  let v_thresh = cache.v_thresh
  let beta = cache.beta

  // d_y[n_steps-1] = probs - one_hot(target).
  let d_y : Array[Float] = Array::make(n_steps * n_classes, 0.0F)
  let final_off = (n_steps - 1) * n_classes
  for c in 0..= 0 {
    for h in 0.. 0 of (s[t-1], d_v[t]).
  let d_w_rec : Array[Float] = Array::make(n_hidden * n_hidden, 0.0F)
  for t in 1.. Float {
  let dt = srnn.dt
  let max_rate = 50.0F
  let x = srnn_poisson_encode(image, n_steps, dt, max_rate, rng)
  let (loss, cache) = srnn_forward(x, n_steps, label, srnn, beta)
  let (d_w_in, d_w_rec, d_w_out, d_b_in, d_b_out) = srnn_backward(cache, srnn)

  clip_grad_array(d_w_in, 1.0F)
  clip_grad_array(d_w_rec, 1.0F)
  clip_grad_array(d_w_out, 1.0F)
  clip_grad_array(d_b_in, 1.0F)
  clip_grad_array(d_b_out, 1.0F)

  for i in 0.. Float {
  let n_total = labels.length()
  let n_in = srnn.n_in
  let mut sum_loss = 0.0F
  for i in 0..