// 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..