// spikeformer_demo.mbt — Mini-SpikeFormer training demo.
//
// Demonstrates end-to-end training of a small SpikeFormer-style
// spiking transformer on synthetic 28×28 grayscale images.
//
//   Input: (N, 1, 28, 28) grayscale images, pixel ∈ {0, 1}
//   Patch embed: split into 7×7 = 49 patches of 4×4 pixels; linear
//                project each patch (16-dim) to d_model via Linear.
//   Add positional embedding (learnable, 50 positions for 49 patches
//                + 1 class token).
//   Prepend class token (learnable).
//   SpikingTransformerBlock (Pre-Norm, LN → SpikingAttn → + → LN →
//                FFN(GELU) → +).
//   Readout: take class token (first position after block).
//   Classifier: Linear(d_model, n_classes).
//   Loss: cross-entropy on (logits, target).
//
// Reuses existing modules:
//   v0.22.0 surrogate (fast_sigmoid_forward / _surrogate inside
//                       SpikingMultiHeadAttention from v0.26.0)
//   v0.26.0 SpikingMultiHeadAttention
//   v0.24.1 SpikingTransformerBlock
//   v0.11.1 Linear (forward + backward)
//   v0.23.2 PositionalEmbedding
//   v0.13.1 softmax / log_softmax
//   v0.13.2 cross_entropy
//   v0.15.0 SGD with momentum (via inline sgd_step_arrays)

///|
/// Synthetic 28×28 dataset: 10 classes × 10 samples = 100 images.
/// Each class has a distinct geometric pattern (line, half-frame,
/// cross, etc.) + ~5% pixel flip noise.
pub fn spikeformer_dataset_build(
  seed : UInt64,
) -> (Array[Float], Array[Int]) {
  let n_classes = 10
  let per_class = 10
  let n_total = n_classes * per_class  // 100
  let h = 28
  let w = 28
  let images : Array[Float] = Array::make(n_total * h * w, 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.. v = if y < 14 { 1.0F } else { 0.0F }  // top half
            1 => v = if y >= 14 { 1.0F } else { 0.0F }  // bottom half
            2 => v = if x < 14 { 1.0F } else { 0.0F }  // left half
            3 => v = if x >= 14 { 1.0F } else { 0.0F }  // right half
            4 => v = if x == y || x + y == 27 { 1.0F } else { 0.0F }  // X
            5 => v = if y == 14 || x == 14 { 1.0F } else { 0.0F }  // cross
            6 => v = if x >= 7 && x <= 20 && y >= 7 && y <= 20 { 1.0F } else { 0.0F }  // center block
            7 => v = if y < 7 || y >= 21 { 1.0F } else { 0.0F }  // top+bottom band
            8 => v = if x < 7 || x >= 21 { 1.0F } else { 0.0F }  // left+right band
            _ => v = if x + y < 14 || x + y > 41 { 1.0F } else { 0.0F }  // border ring
          }
          // 5% noise flip.
          let r = next_f32(rng)
          if r < 0.05F {
            v = if v > 0.5F { 0.0F } else { 1.0F }
          }
          images[idx * h * w + y * w + x] = v
        }
      }
    }
  }
  (images, labels)
}

///|
/// One-hot label vector (length n_classes).
pub fn spikeformer_one_hot(label : Int, n_classes : Int) -> Array[Float] {
  let v : Array[Float] = Array::make(n_classes, 0.0F)
  v[label] = 1.0F
  v
}

///|
/// Numerically-stable softmax over a 1D vector.
pub fn spikeformer_softmax(x : Array[Float]) -> Array[Float] {
  let n = x.length()
  let mut max_v = x[0]
  for i in 1.. max_v {
      max_v = x[i]
    }
  }
  let mut sum_exp = 0.0F
  let out : Array[Float] = Array::make(n, 0.0F)
  for i in 0.. Float {
  -logf(probs[target])
}

///|
/// Convert (N, 1, 28, 28) → (N, 49, 16) patches (NCHW row-major).
pub fn spikeformer_extract_patches(
  images : Array[Float],
  n : Int,
  patch_dim : Int,
  patch_h : Int,
  patch_w : Int,
  img_h : Int,
  img_w : Int,
) -> Array[Float] {
  let n_patches_h = img_h / patch_h
  let n_patches_w = img_w / patch_w
  let n_patches = n_patches_h * n_patches_w
  let out : Array[Float] = Array::make(n * n_patches * patch_dim, 0.0F)
  for batch in 0.. MiniSpikeFormer {
  let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let patch_dim = 16
  let n_patches = 49
  // Patch embedding Linear(16, d_model): Xavier-normal.
  let std_pe = sqrtf(2.0F / Float::from_int(patch_dim + d_model))
  let patch_w : Array[Float] = Array::make(patch_dim * d_model, 0.0F)
  for i in 0..<(patch_dim * d_model) {
    let (z, _) = box_muller(rng)
    patch_w[i] = Float::from_double(z) * std_pe
  }
  let patch_b : Array[Float] = Array::make(d_model, 0.0F)
  // Class token: small random init.
  let class_token : Array[Float] = Array::make(d_model, 0.0F)
  for i in 0.. (Float, MiniSpikeFormerCache) {
  let patch_dim = msf.patch_dim
  let n_patches = msf.n_patches
  let d_model = msf.d_model

  // Extract patches (batch × 49 × 16).
  let patches = spikeformer_extract_patches(
    images, batch, patch_dim, 4, 4, 28, 28,
  )

  // Patch embed: Linear(patch_dim → d_model).
  let patch_emb = Array::make(batch * n_patches * d_model, 0.0F)
  for b in 0.. (
  Array[Float],  // d_patch_w (patch_dim * d_model)
  Array[Float],  // d_patch_b (d_model)
  Array[Float],  // d_class_token (d_model)
  Array[Float],  // d_pos_weight ((n_patches + 1) * d_model)
  SpikingTransformerBlockGrad,  // block grads
  Array[Float],  // d_cls_w (d_model * n_classes)
  Array[Float],  // d_cls_b (n_classes)
) {
  let d_model = msf.d_model
  let n_classes = msf.n_classes
  let n_patches = msf.n_patches
  let seq_len = n_patches + 1
  let patch_dim = msf.patch_dim

  // ---- d_logits = probs - one_hot(target) ----
  let d_logits : Array[Float] = Array::make(n_classes, 0.0F)
  for c in 0.. Float {
  if v > clip {
    clip
  } else if v < -clip {
    -clip
  } else {
    v
  }
}

///|
/// In-place clip an Array[Float] to [-clip, clip] per element.
fn clip_grad_array(arr : Array[Float], clip : Float) -> Unit {
  let n = arr.length()
  for i in 0.. Float {
  let (loss, cache) = mini_spike_former_forward(images, 1, label, msf)
  let (_d_pw, d_pb, d_ct, d_pos, bg, d_cw, d_cb) = mini_spike_former_backward(
    cache, msf,
  )
  let clip = 1.0F
  let d_model = msf.d_model
  let n_classes = msf.n_classes
  let n_patches = msf.n_patches

  // Clip all gradients in place.
  clip_grad_array(d_ct, clip)
  clip_grad_array(d_pos, clip)
  clip_grad_array(d_cw, clip)
  clip_grad_array(d_cb, clip)
  clip_grad_array(d_pb, clip)
  clip_grad_array(bg.ln1_d_gamma, clip)
  clip_grad_array(bg.ln1_d_beta, clip)
  clip_grad_array(bg.ln2_d_gamma, clip)
  clip_grad_array(bg.ln2_d_beta, clip)
  clip_grad_array(bg.ffn_d_w1, clip)
  clip_grad_array(bg.ffn_d_b1, clip)
  clip_grad_array(bg.ffn_d_w2, clip)
  clip_grad_array(bg.ffn_d_b2, clip)
  clip_grad_array(bg.sa_grad.d_w_q, clip)
  clip_grad_array(bg.sa_grad.d_w_k, clip)
  clip_grad_array(bg.sa_grad.d_w_v, clip)
  clip_grad_array(bg.sa_grad.d_w_o, clip)
  clip_grad_array(bg.sa_grad.d_b_q, clip)
  clip_grad_array(bg.sa_grad.d_b_k, clip)
  clip_grad_array(bg.sa_grad.d_b_v, clip)
  clip_grad_array(bg.sa_grad.d_b_o, clip)

  // SGD update for class token.
  for i in 0.. Float {
  let n = labels.length()
  let mut sum_loss = 0.0F
  for i in 0.. Int {
  let n = probs.length()
  let mut best_i = 0
  let mut best_v = probs[0]
  for i in 1.. best_v {
      best_v = probs[i]
      best_i = i
    }
  }
  best_i
}

///|
/// Predict argmax class for a single 28×28 image. Reuses
/// `mini_spike_former_forward` (the label is unused for argmax).
pub fn mini_spike_former_predict(
  images : Array[Float],
  msf : MiniSpikeFormer,
) -> Int {
  let (_l, cache) = mini_spike_former_forward(images, 1, 0, msf)
  ignore(_l)
  spikeformer_argmax(cache.probs)
}

///|
/// Classification accuracy on a list of (image, label) samples. Each
/// image is evaluated independently (batch=1 forward pass). Returns
/// correct / total as a Float; returns 0.0 for an empty label list.
pub fn mini_spike_former_eval_accuracy(
  images : Array[Float],
  labels : Array[Int],
  msf : MiniSpikeFormer,
) -> Float {
  let n = labels.length()
  if n == 0 {
    return 0.0F
  }
  let mut correct = 0
  for i in 0.. (Float, Float) {
  let init_acc = mini_spike_former_eval_accuracy(images, [label], msf)
  for _k in 0..