// sattention.mbt — Consolidated Spiking Attention (SAttention) module.
//
// Three related attention variants sharing the same fast-sigmoid
// forward + surrogate backward pattern (Zenke & Ganguli 2018):
//
//   - SpikingMultiHeadAttention (multi-head self-attention, MHA-style)
//   - SpikingSelfAttention (single-head self-attention; just
//                            num_heads=1 convenience over the above)
//   - SpikingCrossAttention (cross-attention: separate Q and K/V inputs)
//
// All three replace softmax along the key axis with
// `fast_sigmoid_forward` from `surrogate.mbt` (v0.22.0) and use
// `fast_sigmoid_surrogate` as the backward gradient.
//
// Backward-compat: `SpikingAttention` is typealiased to
// `SpikingMultiHeadAttention`, so existing calls in `spiking_attention.mbt`
// keep working.

// ============================================================================
// Multi-head self-attention
// ============================================================================

///|
/// Spiking multi-head self-attention parameter container.
pub struct SpikingMultiHeadAttention {
  d_model : Int
  num_heads : Int
  d_k : Int
  beta : Float
  w_q : LinearParam
  w_k : LinearParam
  w_v : LinearParam
  w_o : LinearParam
}

///|
/// Forward / backward cache for multi-head self-attention.
pub struct SpikingMultiHeadAttnCache {
  seq_len : Int
  x : Array[Float]
  q : Array[Float]
  k : Array[Float]
  v : Array[Float]
  scores : Array[Float]
  weights : Array[Float]
  out_pre : Array[Float]
  mask : Array[Float]
  scale : Float
  beta : Float
}

///|
/// Backward-compat struct (was the v0.24.0 name for multi-head
/// self-attention). Structurally identical to `SpikingMultiHeadAttention`.
/// New code should use `SpikingMultiHeadAttention` directly.
pub struct SpikingAttention {
  d_model : Int
  num_heads : Int
  d_k : Int
  beta : Float
  w_q : LinearParam
  w_k : LinearParam
  w_v : LinearParam
  w_o : LinearParam
}

///|
/// Backward-compat cache (was `SpikingAttnCache` in v0.24.0).
pub struct SpikingAttnCache {
  seq_len : Int
  x : Array[Float]
  q : Array[Float]
  k : Array[Float]
  v : Array[Float]
  scores : Array[Float]
  weights : Array[Float]
  out_pre : Array[Float]
  mask : Array[Float]
  scale : Float
  beta : Float
}

///|
/// Construct a SpikingMultiHeadAttention with Xavier-normal init.
pub fn SpikingMultiHeadAttention::new(
  d_model : Int,
  num_heads : Int,
  beta : Float,
  seed : UInt64,
) -> SpikingMultiHeadAttention {
  if d_model % num_heads != 0 {
    abort("SpikingMultiHeadAttention::new: d_model not divisible by num_heads")
  }
  let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let w_q = linear_xavier(rng, d_model)
  let w_k = linear_xavier(rng, d_model)
  let w_v = linear_xavier(rng, d_model)
  let w_o = linear_xavier(rng, d_model)
  { d_model, num_heads, d_k: d_model / num_heads, beta, w_q, w_k, w_v, w_o }
}

///|
/// Convenience: build a single-head spiking self-attention by
/// calling `SpikingMultiHeadAttention::new(d_model, 1, beta, seed)`.
/// No separate struct — single-head is just multi-head with
/// num_heads=1.
pub fn spiking_self_attention_new(
  d_model : Int,
  beta : Float,
  seed : UInt64,
) -> SpikingMultiHeadAttention {
  SpikingMultiHeadAttention::new(d_model, 1, beta, seed)
}

///|
/// Multi-head self-attention forward. `mask` is the same additive
/// mask as standard MHA (pass empty array to skip).
pub fn spiking_multi_head_attention_forward(
  x : Array[Float],
  sa : SpikingMultiHeadAttention,
  mask : Array[Float],
) -> (Array[Float], SpikingMultiHeadAttnCache) {
  let seq_len = x.length() / sa.d_model
  let d_model = sa.d_model
  let num_heads = sa.num_heads
  let d_k = sa.d_k
  let scale = 1.0F / sqrtf(Float::from_int(d_k))
  let beta = sa.beta

  let q = linear_forward(x, seq_len, sa.w_q)
  let k = linear_forward(x, seq_len, sa.w_k)
  let v = linear_forward(x, seq_len, sa.w_v)

  let scores : Array[Float] = Array::make(num_heads * seq_len * seq_len, 0.0F)
  let use_mask = mask.length() > 0
  for h in 0.. (Array[Float], MHAGrad) {
  let seq_len = cache.seq_len
  let d_model = sa.d_model
  let num_heads = sa.num_heads
  let d_k = sa.d_k
  let scale = cache.scale
  let beta = cache.beta

  let (d_out_pre, d_w_o, d_b_o) = linear_backward(
    cache.out_pre, d_output, seq_len, sa.w_o,
  )

  let weights = cache.weights
  let v = cache.v
  let d_weights : Array[Float] = Array::make(
    num_heads * seq_len * seq_len, 0.0F,
  )
  let d_v : Array[Float] = Array::make(seq_len * d_model, 0.0F)
  for h in 0.. Float {
  arr[i]
}

// ============================================================================
// Cross-attention (Q from one input, K/V from another)
// ============================================================================

///|
/// Spiking cross-attention parameter container. Q is computed from
/// `x_q`, K/V are computed from a separate `x_kv` input.
pub struct SpikingCrossAttention {
  d_model : Int
  num_heads : Int
  d_k : Int
  beta : Float
  w_q : LinearParam
  w_k : LinearParam
  w_v : LinearParam
  w_o : LinearParam
}

///|
/// Forward / backward cache for cross-attention.
pub struct SpikingCrossAttnCache {
  seq_len_q : Int
  seq_len_kv : Int
  x_q : Array[Float]
  x_kv : Array[Float]
  q : Array[Float]
  k : Array[Float]
  v : Array[Float]
  scores : Array[Float]  // (num_heads * seq_len_q * seq_len_kv)
  weights : Array[Float]
  out_pre : Array[Float]  // (seq_len_q * d_model)
  mask : Array[Float]
  scale : Float
  beta : Float
}

///|
/// Construct a SpikingCrossAttention with Xavier-normal init.
pub fn SpikingCrossAttention::new(
  d_model : Int,
  num_heads : Int,
  beta : Float,
  seed : UInt64,
) -> SpikingCrossAttention {
  if d_model % num_heads != 0 {
    abort("SpikingCrossAttention::new: d_model not divisible by num_heads")
  }
  let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let w_q = linear_xavier(rng, d_model)
  let w_k = linear_xavier(rng, d_model)
  let w_v = linear_xavier(rng, d_model)
  let w_o = linear_xavier(rng, d_model)
  { d_model, num_heads, d_k: d_model / num_heads, beta, w_q, w_k, w_v, w_o }
}

///|
/// Cross-attention forward.
///
/// `x_q`  : (seq_len_q, d_model)
/// `x_kv` : (seq_len_kv, d_model)
/// `mask` : (num_heads, seq_len_q, seq_len_kv) additive, or empty.
/// `returns` : out of shape (seq_len_q, d_model).
pub fn spiking_cross_attention_forward(
  x_q : Array[Float],
  x_kv : Array[Float],
  sa : SpikingCrossAttention,
  mask : Array[Float],
) -> (Array[Float], SpikingCrossAttnCache) {
  let seq_len_q = x_q.length() / sa.d_model
  let seq_len_kv = x_kv.length() / sa.d_model
  let d_model = sa.d_model
  let num_heads = sa.num_heads
  let d_k = sa.d_k
  let scale = 1.0F / sqrtf(Float::from_int(d_k))
  let beta = sa.beta

  let q = linear_forward(x_q, seq_len_q, sa.w_q)
  let k = linear_forward(x_kv, seq_len_kv, sa.w_k)
  let v = linear_forward(x_kv, seq_len_kv, sa.w_v)

  // scores[h, s, t] = sum_k q[s, h*d_k+k] * k[t, h*d_k+k] * scale
  let scores : Array[Float] = Array::make(
    num_heads * seq_len_q * seq_len_kv, 0.0F,
  )
  let use_mask = mask.length() > 0
  for h in 0.. (Array[Float], Array[Float], MHAGrad) {
  let seq_len_q = cache.seq_len_q
  let seq_len_kv = cache.seq_len_kv
  let d_model = sa.d_model
  let num_heads = sa.num_heads
  let d_k = sa.d_k
  let scale = cache.scale
  let beta = cache.beta

  let (d_out_pre, d_w_o, d_b_o) = linear_backward(
    cache.out_pre, d_output, seq_len_q, sa.w_o,
  )

  let weights = cache.weights
  let v = cache.v
  let d_weights : Array[Float] = Array::make(
    num_heads * seq_len_q * seq_len_kv, 0.0F,
  )
  let d_v : Array[Float] = Array::make(seq_len_kv * d_model, 0.0F)
  for h in 0..