// 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)
}