// multi_head_attention.mbt — Standard multi-head self-attention
// (Vaswani et al. 2017).
//
// Self-attention forward (single sequence, no batch):
//
//   x  : (seq_len, d_model)         row-major flat
//   q  = x @ W_q                    (seq_len, d_model)
//   k  = x @ W_k                    (seq_len, d_model)
//   v  = x @ W_v                    (seq_len, d_model)
//
//   split heads: q' has shape (num_heads, seq_len, d_k) where
//                d_k = d_model / num_heads.
//                Flat offset: h * seq_len * d_k + s * d_k + k
//                Maps to x-flat: x[s * d_model + h * d_k + k].
//
//   scores = q' @ k'^T / sqrt(d_k)  shape (num_heads, seq_len, seq_len)
//   if mask:  scores += mask        (additive; -inf blocks)
//   weights = softmax(scores, axis=-1)
//   out_pre = weights @ v'           shape (num_heads, seq_len, d_k)
//
//   merge heads: out_pre' has shape (seq_len, d_model)
//                Flat offset: s * d_model + h * d_k + k.
//   y  = out_pre' @ W_o              (seq_len, d_model)
//
// Bias is added inside each Linear projection.
//
// Initialisation: Xavier-normal for all 4 weight matrices
//   N(0, sqrt(2 / (d_in + d_out)))
//   For square d_in = d_out = d_model, std = sqrt(1/d_model).
//   Biases initialised to zero.
//
// Backward returns (d_x, MHAGrad). The cache retains x / q / k / v /
// scores / weights / out_pre for the reverse pass.

///|
/// Multi-head self-attention parameter container.
pub struct MultiHeadAttention {
  d_model : Int
  num_heads : Int
  d_k : Int
  w_q : LinearParam
  w_k : LinearParam
  w_v : LinearParam
  w_o : LinearParam
}

///|
/// Gradient bundle for an MHA module: d_weight / d_bias for each of
/// the four projections.
pub struct MHAGrad {
  d_w_q : Array[Float]
  d_b_q : Array[Float]
  d_w_k : Array[Float]
  d_b_k : Array[Float]
  d_w_v : Array[Float]
  d_b_v : Array[Float]
  d_w_o : Array[Float]
  d_b_o : Array[Float]
}

///|
/// Forward / backward cache for an MHA forward pass.
pub struct AttnCache {
  seq_len : Int
  // Forward input.
  x : Array[Float]
  // Post-projection q / k / v (each shape seq_len × d_model).
  q : Array[Float]
  k : Array[Float]
  v : Array[Float]
  // Pre-softmax scores, shape num_heads × seq_len × seq_len.
  scores : Array[Float]
  // Post-softmax attention weights.
  weights : Array[Float]
  // Output of weights @ v, pre-merge shape (seq_len × d_model,
  // flat).  Filled by merging heads on the fly.
  out_pre : Array[Float]
  // Optional mask (same shape as scores). Zero when no mask.
  mask : Array[Float]
  // 1 / sqrt(d_k) precomputed.
  scale : Float
}

// ---- initialisation helpers ---------------------------------------

///|
/// Fill an Array[Float] with Xavier-normal draws:
/// N(0, sqrt(2 / (fan_in + fan_out))). Uses Float32 Box-Muller.
fn xavier_normal_init(
  rng : Xoshiro,
  n : Int,
  fan_in : Int,
  fan_out : Int,
) -> Array[Float] {
  let out : Array[Float] = Array::make(n, 0.0F)
  let std = sqrtf(2.0F / Float::from_int(fan_in + fan_out))
  let mut i = 0
  while i < n {
    let (z1, _) = box_muller(rng)
    out[i] = Float::from_double(z1) * std
    i = i + 1
  }
  out
}

///|
/// Build a square LinearParam with Xavier-normal init + zero bias.
fn linear_xavier(
  rng : Xoshiro,
  d : Int,
) -> LinearParam {
  let weight = xavier_normal_init(rng, d * d, d, d)
  let bias : Array[Float] = Array::make(d, 0.0F)
  LinearParam::new(weight, bias, d, d)
}

// ---- public API ---------------------------------------------------

///|
/// Construct a fresh MHA. `seed` initialises the four weight matrices
/// deterministically (xoshiro256++ from_state with derived seeds).
pub fn MultiHeadAttention::new(
  d_model : Int,
  num_heads : Int,
  seed : UInt64,
) -> MultiHeadAttention {
  // Sanity: d_model must be divisible by num_heads.
  if d_model % num_heads != 0 {
    abort("MultiHeadAttention::new: d_model \{d_model} not divisible by num_heads \{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, w_q, w_k, w_v, w_o }
}

///|
/// Self-attention forward pass.
///
/// `x`         : length `seq_len * d_model`, row-major
/// `mask~`     : optional additive mask shape `num_heads * seq_len *
///               seq_len` (pass an empty array to skip).
/// returns     : `(out, cache)` where `out` is length
///               `seq_len * d_model`.
pub fn multi_head_attention_forward(
  x : Array[Float],
  mha : MultiHeadAttention,
  mask : Array[Float],
) -> (Array[Float], AttnCache) {
  let seq_len = x.length() / mha.d_model
  let d_model = mha.d_model
  let num_heads = mha.num_heads
  let d_k = mha.d_k
  let scale = 1.0F / sqrtf(Float::from_int(d_k))

  // 1) Q / K / V projections (each shape seq_len × d_model).
  let q = linear_forward(x, seq_len, mha.w_q)
  let k = linear_forward(x, seq_len, mha.w_k)
  let v = linear_forward(x, seq_len, mha.w_v)

  // 2) Scaled dot-product scores per head.
  //    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 * seq_len, 0.0F)
  let use_mask = mask.length() > 0
  for h in 0.. max_v {
          max_v = scores[s_row + t]
        }
        t = t + 1
      }
      let mut sum_exp = 0.0F
      t = 0
      while t < seq_len {
        let e = expf(scores[s_row + t] - max_v)
        weights[s_row + t] = e
        sum_exp = sum_exp + e
        t = t + 1
      }
      t = 0
      while t < seq_len {
        weights[s_row + t] = weights[s_row + t] / sum_exp
        t = t + 1
      }
    }
  }

  // 4) Merge heads + apply W_o in one go: for each (s, h, kk),
  //    compute out_pre[s, h*d_k + kk] = sum_t weights[h, s, t] *
  //    v[t, h*d_k + kk].  Then linear_forward through W_o.
  let out_pre : Array[Float] = Array::make(seq_len * d_model, 0.0F)
  for h in 0.. (Array[Float], MHAGrad) {
  let seq_len = cache.seq_len
  let d_model = mha.d_model
  let num_heads = mha.num_heads
  let d_k = mha.d_k
  let scale = cache.scale

  // 1) d_out_pre, d_W_o, d_b_o from d_output through W_o.
  let (d_out_pre, d_w_o, d_b_o) = linear_backward(
    cache.out_pre, d_output, seq_len, mha.w_o,
  )

  // Reshape d_out_pre to per-head layout (num_heads, seq_len, d_k).
  // d_out_h[h, s, kk] = d_out_pre[s * d_model + h * d_k + kk]

  // 2) d_weights[h, s, t] and d_v contribution.
  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.. weights.
  //    d_scores[h, s, t] = weights[h, s, t] * (d_weights[h, s, t]
  //        - sum_{t'} weights[h, s, t'] * d_weights[h, s, t'])
  let d_scores : Array[Float] = Array::make(
    num_heads * seq_len * seq_len,
    0.0F,
  )
  for h in 0..