// spiking_transformer_block.mbt — Pre-Norm Spiking Transformer block.
//
// Like TransformerBlock but uses `SpikingMultiHeadAttention` (v0.26.0) instead
// of standard `MultiHeadAttention`. The softmax along the key axis is
// replaced by `fast_sigmoid_forward`, and the backward uses the
// `fast_sigmoid_surrogate` (v0.22.0) for BPTT-compatible gradients.
//
// x2 = x + SpikingMultiHeadAttention( LN1(x) ) # spiking self-attention path
// y = x2 + FFN( LN2(x2) ) # feed-forward path
//
// where FFN(h) = Linear2(GELU(Linear1(h))) and d_ff = 4 d_model.
///|
/// Spiking transformer block parameter container.
pub struct SpikingTransformerBlock {
d_model : Int
num_heads : Int
d_ff : Int
beta : Float
ln_1 : LayerNorm
ln_2 : LayerNorm
sa : SpikingMultiHeadAttention
ffn_w1 : LinearParam
ffn_w2 : LinearParam
}
///|
/// Forward / backward cache.
pub struct SpikingTransformerBlockCache {
seq_len : Int
x : Array[Float]
x2 : Array[Float]
x_norm1 : Array[Float]
x_norm2 : Array[Float]
ln1_cache : LayerNormCache
sa_cache : SpikingMultiHeadAttnCache
ln2_cache : LayerNormCache
ffn_hidden_pre_gelu : Array[Float]
ffn_hidden_post_gelu : Array[Float]
mask : Array[Float]
}
///|
/// Gradient bundle (reuses MHAGrad for the spiking attention
/// projection gradients since the Linear layer structure is identical).
pub struct SpikingTransformerBlockGrad {
ln1_d_gamma : Array[Float]
ln1_d_beta : Array[Float]
sa_grad : MHAGrad
ln2_d_gamma : Array[Float]
ln2_d_beta : Array[Float]
ffn_d_w1 : Array[Float]
ffn_d_b1 : Array[Float]
ffn_d_w2 : Array[Float]
ffn_d_b2 : Array[Float]
}
///|
/// Construct a SpikingTransformerBlock. `d_ff` defaults to 4·d_model.
pub fn SpikingTransformerBlock::new(
d_model : Int,
num_heads : Int,
beta : Float,
seed : UInt64,
d_ff? : Int = -1,
) -> SpikingTransformerBlock {
let actual_d_ff = if d_ff < 0 { 4 * d_model } else { d_ff }
let ln_1 = LayerNorm::new(1, 1, d_model)
let ln_2 = LayerNorm::new(1, 1, d_model)
let sa = SpikingMultiHeadAttention::new(d_model, num_heads, beta, seed + 100UL)
let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
let w1_weight = xavier_normal_init(rng, d_model * actual_d_ff, d_model, actual_d_ff)
let w1_bias : Array[Float] = Array::make(actual_d_ff, 0.0F)
let ffn_w1 = LinearParam::new(w1_weight, w1_bias, d_model, actual_d_ff)
let w2_weight = xavier_normal_init(rng, actual_d_ff * d_model, actual_d_ff, d_model)
let w2_bias : Array[Float] = Array::make(d_model, 0.0F)
let ffn_w2 = LinearParam::new(w2_weight, w2_bias, actual_d_ff, d_model)
{ d_model, num_heads, d_ff: actual_d_ff, beta, ln_1, ln_2, sa, ffn_w1, ffn_w2 }
}
///|
/// Forward pass.
pub fn spiking_transformer_block_forward(
x : Array[Float],
block : SpikingTransformerBlock,
mask : Array[Float],
) -> (Array[Float], SpikingTransformerBlockCache) {
let seq_len = x.length() / block.d_model
let d_model = block.d_model
// 1) Pre-LN1 + spiking self-attention + residual.
let (x_norm1, ln1_cache) = layer_norm_forward(
x, seq_len, 1, 1, d_model, block.ln_1,
)
let (sa_out, sa_cache) = spiking_attention_forward(
x_norm1, block.sa, mask,
)
let x2 : Array[Float] = Array::make(seq_len * d_model, 0.0F)
for i in 0..<(seq_len * d_model) {
x2[i] = x[i] + sa_out[i]
}
// 2) Pre-LN2 + FFN + residual.
let (x_norm2, ln2_cache) = layer_norm_forward(
x2, seq_len, 1, 1, d_model, block.ln_2,
)
let pre_gelu = linear_forward(x_norm2, seq_len, block.ffn_w1)
let ffn_hidden_pre_gelu = Array::make(pre_gelu.length(), 0.0F)
for i in 0.. (Array[Float], SpikingTransformerBlockGrad) {
let seq_len = cache.seq_len
let d_model = block.d_model
// ---- FFN residual split ----
// ---- FFN path backward ----
let (d_ffn_hidden_post_gelu, d_w2, d_b2) = linear_backward(
cache.ffn_hidden_post_gelu, d_output, seq_len, block.ffn_w2,
)
let d_ffn_pre_gelu = gelu_backward(
cache.ffn_hidden_pre_gelu, d_ffn_hidden_post_gelu,
)
let (d_x_norm2, d_w1, d_b1) = linear_backward(
cache.x_norm2, d_ffn_pre_gelu, seq_len, block.ffn_w1,
)
// LN2 backward.
let (d_x2_from_ln2, d_gamma2, d_beta2) = layer_norm_backward(
d_x_norm2, cache.ln2_cache, block.ln_2,
)
let d_x2 : Array[Float] = Array::make(seq_len * d_model, 0.0F)
for i in 0..<(seq_len * d_model) {
d_x2[i] = d_output[i] + d_x2_from_ln2[i]
}
// ---- Spiking attention backward ----
let (d_x_norm1, sa_grad) = spiking_attention_backward(
cache.sa_cache, d_x2, block.sa,
)
let (d_x_from_ln1, d_gamma1, d_beta1) = layer_norm_backward(
d_x_norm1, cache.ln1_cache, block.ln_1,
)
let d_x : Array[Float] = Array::make(seq_len * d_model, 0.0F)
for i in 0..<(seq_len * d_model) {
d_x[i] = d_x2[i] + d_x_from_ln1[i]
}
let grads : SpikingTransformerBlockGrad = {
ln1_d_gamma: d_gamma1,
ln1_d_beta: d_beta1,
sa_grad,
ln2_d_gamma: d_gamma2,
ln2_d_beta: d_beta2,
ffn_d_w1: d_w1,
ffn_d_b1: d_b1,
ffn_d_w2: d_w2,
ffn_d_b2: d_b2,
}
(d_x, grads)
}