// vit_trainer.mbt -- ViT trainer: cross-entropy loss + SGD step (v0.112.0).
//
// Scope of v0.112.0:
//   - vit_cross_entropy: cross-entropy loss from logits + true class.
//   - vit_softmax: stable softmax over logits.
//   - vit_train_step_head: one SGD step on the classification head only
//     (analytic gradient through the softmax+CE+linear chain).
//   - vit_train_step: forward + loss + SGD-on-head.
//
// BPTT through N ViTBlocks (so that the full ViT can be trained end-to-end)
// is deferred to a follow-up batch.
//
// Reference: Dosovitskiy et al. 2020; standard softmax+CE+linear
// gradient is computed analytically here.

///|
/// Stable softmax over logits (subtracts the max for numerical stability).
pub fn vit_softmax(logits : Array[Float]) -> Array[Float] {
  let n = logits.length()
  // Find max for stability.
  let mut m = logits[0]
  for i in 1.. m {
      m = logits[i]
    }
  }
  let exp_logits : Array[Float] = Array::make(n, 0.0F)
  let mut sum_exp = 0.0F
  for i in 0.. Float {
  let probs = vit_softmax(logits)
  // Add a tiny epsilon for numerical safety (avoid log(0)).
  let eps = 1.0e-12F
  let p = if probs[target] > eps { probs[target] } else { eps }
  -logf(p)
}

///|
/// One SGD step on the classification head (Linear d_model -> num_classes)
/// given the cached CLS representation, the logits, and the target.
/// Mutates vit.head_w and vit.head_b in place. Returns the loss before
/// the step.
pub fn vit_train_step_head(
  vit : ViT,
  cls_repr : Array[Float],
  logits : Array[Float],
  target : Int,
  lr : Float,
) -> Float {
  let num_classes = vit.num_classes
  let d_model = vit.d_model
  // Cross-entropy gradient: dL/dlogits[k] = softmax[k] - 1[k == target].
  let probs = vit_softmax(logits)
  let loss = -logf(if probs[target] > 1.0e-12F { probs[target] } else { 1.0e-12F })
  // d_logits[k] = probs[k] - (k == target ? 1 : 0).
  let d_logits : Array[Float] = Array::make(num_classes, 0.0F)
  for k in 0.. head_w: dW[k][i] += d_logits[k] * cls_repr[i]
  // d_logits -> head_b: dB[k] += d_logits[k]
  for k in 0.. Array[Float] {
  let d_model = vit.d_model
  let n_patches = vit.n_patches
  let seq_len = n_patches + 1
  let patch_tokens = patch_embed_image(vit.patch_embed, image)
  let tokens : Array[Float] = Array::make(seq_len * d_model, 0.0F)
  for i in 0.. Float {
  let cls_repr = vit_extract_cls(vit, image)
  // Compute logits from cls_repr through head_w/head_b.
  let num_classes = vit.num_classes
  let d_model = vit.d_model
  let logits : Array[Float] = Array::make(num_classes, 0.0F)
  for k in 0.. Float {
  let n = images.length()
  let mut total = 0.0F
  for e in 0.. 1.0e-12F {
      probs[targets[e]]
    } else {
      1.0e-12F
    }
    total = total + (-logf(p))
  }
  total / Float::from_int(n)
}