// transformer_block.mbt — Pre-Norm Transformer block
// (Vaswani 2017 §3.1, Pre-LN variant).
//
// One block, two sub-layers, each Pre-LN + residual:
//
//   x2 = x  +  MHA( LN1(x) )                       # self-attention path
//   y  = x2 +  FFN( LN2(x2) )                      # feed-forward path
//
// where:
//
//   FFN(h) = Linear2( GELU( Linear1(h) ) )
//         = W2 · GELU(W1 · h + b1) + b2
//
// and W1 : (d_model → d_ff), W2 : (d_ff → d_model), d_ff = 4 d_model
// by default.
//
// All sub-modules (LayerNorm, MHA, Linear, GELU) are composed. The
// block has its own parameter container and forward/backward cache.
//
// Float32 throughout. `mask` is an optional additive attention mask
// passed straight through to MHA; pass an empty array to skip.

///|
/// Transformer block parameter container.
pub struct TransformerBlock {
  d_model : Int
  num_heads : Int
  d_ff : Int
  ln_1 : LayerNorm
  ln_2 : LayerNorm
  mha : MultiHeadAttention
  ffn_w1 : LinearParam
  ffn_w2 : LinearParam
}

///|
/// Per-block gradient bundle.
pub struct TransformerBlockGrad {
  ln1_d_gamma : Array[Float]
  ln1_d_beta : Array[Float]
  mha_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]
}

///|
/// Forward / backward cache for a Transformer block.
pub struct TransformerBlockCache {
  seq_len : Int
  // Original input (for residual d_x).
  x : Array[Float]
  // Post-attention residual sum (for residual d_x2 from FFN path).
  x2 : Array[Float]
  // Output of LN1, input to MHA (needed to reconstruct d_x_norm1
  // during backward when only LN1's cache is available).
  x_norm1 : Array[Float]
  // Output of LN2, input to linear1 of FFN (needed by linear_backward).
  x_norm2 : Array[Float]
  // Sub-layer caches.
  ln1_cache : LayerNormCache
  attn_cache : AttnCache
  ln2_cache : LayerNormCache
  // FFN intermediates: pre-GELU hidden state (for GELU backward) +
  // post-GELU hidden state (for linear2 backward).
  ffn_hidden_pre_gelu : Array[Float]
  ffn_hidden_post_gelu : Array[Float]
  // Mask used in MHA (or empty array for none).
  mask : Array[Float]
}

///|
/// Construct a Transformer block. `d_ff` defaults to `4 * d_model`.
pub fn TransformerBlock::new(
  d_model : Int,
  num_heads : Int,
  seed : UInt64,
  d_ff? : Int = -1,
) -> TransformerBlock {
  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 mha = MultiHeadAttention::new(d_model, num_heads, seed + 100UL)
  let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  // FFN linear1: (d_model -> d_ff), linear2: (d_ff -> d_model).
  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, ln_1, ln_2, mha, ffn_w1, ffn_w2 }
}

///|
/// Forward pass. `mask` is forwarded to MHA (pass empty array to skip).
pub fn transformer_block_forward(
  x : Array[Float],
  block : TransformerBlock,
  mask : Array[Float],
) -> (Array[Float], TransformerBlockCache) {
  let seq_len = x.length() / block.d_model
  let d_model = block.d_model

  // 1) Pre-LN1 + self-attention + residual.
  let (x_norm1, ln1_cache) = layer_norm_forward(
    x, seq_len, 1, 1, d_model, block.ln_1,
  )
  let (attn_out, attn_cache) = multi_head_attention_forward(
    x_norm1, block.mha, 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] + attn_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], TransformerBlockGrad) {
  let seq_len = cache.seq_len
  let d_model = block.d_model

  // ---- Step A: residual split on FFN path ----
  // out = x2 + ffn_out   ⇒   d_x2 (FFN residual) = d_output,
  //                          d_ffn_out = d_output.
  // d_x2_from_ln2 will be added below.

  // ---- Step B: FFN path backward ----
  // linear2 backward (input = ffn_hidden_post_gelu).
  let (d_ffn_hidden_post_gelu, d_w2, d_b2) = linear_backward(
    cache.ffn_hidden_post_gelu, d_output, seq_len, block.ffn_w2,
  )
  // GELU backward (input = ffn_hidden_pre_gelu).
  let d_ffn_pre_gelu = gelu_backward(
    cache.ffn_hidden_pre_gelu, d_ffn_hidden_post_gelu,
  )
  // linear1 backward (input = x_norm2).
  let (d_x_norm2, d_w1, d_b1) = linear_backward(
    cache.x_norm2, d_ffn_pre_gelu, seq_len, block.ffn_w1,
  )

  // ---- Step C: LN2 backward (input = x2 from cache) ----
  let (d_x2_from_ln2, d_gamma2, d_beta2) = layer_norm_backward(
    d_x_norm2, cache.ln2_cache, block.ln_2,
  )
  // Total d_x2 = d_output (residual) + d_x2_from_ln2.
  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]
  }

  // ---- Step D: residual split on MHA path ----
  // x2 = x + attn_out   ⇒   d_attn_out = d_x2,
  //                          d_x (residual from MHA) = d_x2.

  // ---- Step E: MHA backward ----
  let (d_x_norm1, mha_grad) = multi_head_attention_backward(
    cache.attn_cache, d_x2, block.mha,
  )

  // ---- Step F: LN1 backward ----
  let (d_x_from_ln1, d_gamma1, d_beta1) = layer_norm_backward(
    d_x_norm1, cache.ln1_cache, block.ln_1,
  )

  // ---- Step G: total d_input ----
  // d_x = d_x2 (residual from MHA path) + d_x_from_ln1.
  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 : TransformerBlockGrad = {
    ln1_d_gamma: d_gamma1,
    ln1_d_beta: d_beta1,
    mha_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)
}