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