// vit_block.mbt — ViTBlock: transformer block for Vision Transformer
// (v0.110.0).
//
// A standard ViT block (pre-norm variant):
//   x' = x + MHA(LayerNorm(x))
//   y  = x' + MLP(LayerNorm(x'))
// where MLP is a 2-layer FFN with GELU activation:
//   MLP(z) = W_2 · GELU(W_1 · z + b_1) + b_2
//
// Scope of v0.110.0:
//   - ViTBlock struct (MHA + LayerNorm params + 2-layer MLP weights)
//   - vit_block_forward: full pre-norm block pass
//   - vit_block_apply_layernorm: per-token LayerNorm helper
//   - vit_block_mlp: standalone 2-layer FFN with GELU
//
// Reference: Dosovitskiy et al. 2020; standard ViT block follows the
// pre-norm Transformer block from "On Layer Normalization in the
// Transformer Architecture".

///|
/// ViTBlock: MHA + LayerNorm + 2-layer MLP. All weights are Linear
/// projections without bias on the MHA (uses the existing
/// MultiHeadAttention struct from multi_head_attention.mbt).
pub struct ViTBlock {
  d_model : Int
  num_heads : Int
  // Pre-norm parameters (per-token LayerNorm).
  ln1_gamma : Array[Float]
  ln1_beta : Array[Float]
  ln2_gamma : Array[Float]
  ln2_beta : Array[Float]
  // MHA.
  mha : MultiHeadAttention
  // MLP (2-layer FFN with GELU).
  mlp_w1 : Array[Array[Float]]
  mlp_b1 : Array[Float]
  mlp_w2 : Array[Array[Float]]
  mlp_b2 : Array[Float]
  mlp_hidden : Int
}

///|
/// Build a fresh ViTBlock.
pub fn ViTBlock::new(
  d_model : Int,
  num_heads : Int,
  mlp_hidden : Int,
  seed : UInt64,
) -> ViTBlock {
  let ln1_gamma : Array[Float] = Array::make(d_model, 1.0F)
  let ln1_beta : Array[Float] = Array::make(d_model, 0.0F)
  let ln2_gamma : Array[Float] = Array::make(d_model, 1.0F)
  let ln2_beta : Array[Float] = Array::make(d_model, 0.0F)
  let mha = MultiHeadAttention::new(d_model, num_heads, seed)
  let rng1 = Xoshiro::from_state(seed + 4UL, seed + 5UL, seed + 6UL, seed + 7UL)
  let std1 = sqrtf(2.0F / Float::from_int(d_model))
  let mlp_w1 = xavier_normal(mlp_hidden, d_model, std1, rng1)
  let mlp_b1 : Array[Float] = Array::make(mlp_hidden, 0.0F)
  let rng2 = Xoshiro::from_state(seed + 8UL, seed + 9UL, seed + 10UL, seed + 11UL)
  let std2 = sqrtf(2.0F / Float::from_int(mlp_hidden))
  let mlp_w2 = xavier_normal(d_model, mlp_hidden, std2, rng2)
  let mlp_b2 : Array[Float] = Array::make(d_model, 0.0F)
  { d_model, num_heads, ln1_gamma, ln1_beta, ln2_gamma, ln2_beta, mha, mlp_w1, mlp_b1, mlp_w2, mlp_b2, mlp_hidden }
}

///|
/// Per-token LayerNorm (one mean/variance per token row). Operates on
/// flat row-major [seq_len x d_model]. `gamma` and `beta` are length
/// d_model.
pub fn vit_block_apply_layernorm(
  x : Array[Float],
  seq_len : Int,
  d_model : Int,
  gamma : Array[Float],
  beta : Array[Float],
  eps : Float,
) -> Array[Float] {
  let out : Array[Float] = Array::make(seq_len * d_model, 0.0F)
  for s in 0.. Array[Float] {
  let hidden : Array[Float] = Array::make(
    seq_len * block.mlp_hidden, 0.0F,
  )
  let out : Array[Float] = Array::make(
    seq_len * block.d_model, 0.0F,
  )
  for s in 0.. Array[Float] {
  let d_model = block.d_model
  // 1. Pre-norm + MHA + residual.
  let norm1 = vit_block_apply_layernorm(
    x, seq_len, d_model, block.ln1_gamma, block.ln1_beta, 1.0e-5F,
  )
  // Build an empty mask (no masking).
  let empty_mask : Array[Float] = Array::make(0, 0.0F)
  let (attn_out, _cache) = multi_head_attention_forward(
    norm1, block.mha, empty_mask,
  )
  let residual1 : Array[Float] = Array::make(
    seq_len * d_model, 0.0F,
  )
  let total_elems = seq_len * d_model
  for i in 0..