// vit.mbt -- ViT: Vision Transformer full model (v0.111.0).
//
// A standard ViT (Dosovitskiy et al. 2020):
//   1. Split image into patches and linearly embed each patch (via
//      PatchEmbedding).
//   2. Prepend a learnable [CLS] token.
//   3. Add learnable positional embeddings (one per patch + one for CLS).
//   4. Pass the resulting sequence through N ViTBlocks.
//   5. Read out the CLS token's representation and feed it through a
//      classification head (Linear d_model -> num_classes).
//
// Scope of v0.111.0:
//   - ViT struct (PatchEmbedding + CLS + positional + N blocks + head).
//   - ViT::new (constructor).
//   - vit_forward: image [C*H*W] -> logits [num_classes].
//   - vit_classify: argmax of logits -> Int.
//   - vit_num_params: count of learnable scalars.
//
// Reference: Dosovitskiy et al. 2020 "An Image is Worth 16x16 Words:
// Transformers for Image Recognition at Scale".

///|
/// ViT: full Vision Transformer. The CLS token is prepended to the
/// patch tokens, and positional embeddings are added in-place before
/// the stack of ViTBlocks.
pub struct ViT {
  patch_embed : PatchEmbedding
  cls_token : Array[Float]
  pos_embed : Array[Float]
  blocks : Array[ViTBlock]
  head_w : Array[Array[Float]]
  head_b : Array[Float]
  num_classes : Int
  d_model : Int
  n_patches : Int
}

///|
/// Build a fresh ViT with random init for all learnable parameters.
/// `patch_size`, `d_model`, `num_heads`, `mlp_hidden`, `num_blocks`,
/// `num_classes` are the canonical ViT hyperparameters.
pub fn ViT::new(
  c : Int,
  h : Int,
  w : Int,
  p : Int,
  d_model : Int,
  num_heads : Int,
  mlp_hidden : Int,
  num_blocks : Int,
  num_classes : Int,
  seed : UInt64,
) -> ViT {
  let patch_embed = PatchEmbedding::new(c, h, w, p, d_model, seed)
  let cls_token : Array[Float] = Array::make(d_model, 0.0F)
  let rng_cls = Xoshiro::from_state(
    seed + 100UL, seed + 101UL, seed + 102UL, seed + 103UL,
  )
  let std_cls = sqrtf(2.0F / Float::from_int(d_model))
  for i in 0.. Array[Float] {
  let d_model = vit.d_model
  let n_patches = vit.n_patches
  let seq_len = n_patches + 1
  // 1. Patch embedding.
  let patch_tokens = patch_embed_image(vit.patch_embed, image)
  // 2. Prepend CLS token.
  let tokens : Array[Float] = Array::make(seq_len * d_model, 0.0F)
  for i in 0.. Int {
  let logits = vit_forward(vit, image)
  let mut best_idx = 0
  let mut best_val = logits[0]
  for k in 1.. best_val {
      best_val = logits[k]
      best_idx = k
    }
  }
  best_idx
}

///|
/// Count of learnable scalars in the model. Useful for capacity
/// reporting and as a sanity check on training-step magnitude.
pub fn vit_num_params(vit : ViT) -> Int {
  let mut total = 0
  // PatchEmbedding: d_model * C * P * P + d_model (bias).
  let pe = vit.patch_embed
  total = total + pe.proj.length() * pe.proj[0].length()
  total = total + pe.bias.length()
  // CLS + positional.
  total = total + vit.cls_token.length()
  total = total + vit.pos_embed.length()
  // Per ViTBlock: 4*d_model (LN) + MHA internals + 2*mlp_hidden*d_model (W1) +
  // mlp_hidden (b1) + 2*d_model*mlp_hidden (W2) + d_model (b2).
  for b in 0..