Documentation

LeanMlir.Proofs.Foundation.SmoothedBatchLoss

The batched label-smoothed loss, and its gradient is the emitted cotangent #

SmoothedLossCot identifies the emitted loss cotangent ROW BY ROW: at example n it is (1/B)·∂softCE/∂z at that example's logits. This file states the loss as one function of the flat batch of logits — smoothedBatchLoss, the 1/B-weighted sum of every example's soft-target cross-entropy against its smoothed target — and proves its full gradient is the emitted cotangent, read at the head's N·K index (smoothedBatchLoss_grad). That is the hg a parameter-level statement asks for (HasGradAt at the logits, ParamGrad).

def Proofs.logitRow (N K : ℕ) (z : Vec (N * K)) (n : Fin N) :
Vec K

Example n's logits out of the flat N·K batch.

Equations
Instances For
    noncomputable def Proofs.targetRow (N K : ℕ) (t : Vec (N * (1 * K))) (n : Fin N) :
    Vec K

    Example n's target out of the N·(1·K) graph input %onehot.

    Equations
    Instances For
      noncomputable def Proofs.smoothedBatchLoss (N K : ℕ) (α B : ℝ) (t : Vec (N * (1 * K))) (z : Vec (N * K)) :
      Vec 1

      The batched label-smoothed loss: Σ_n softCE(smooth(tₙ), zₙ) / B, as a function of the flat logits.

      Equations
      Instances For
        theorem Proofs.logitRow_differentiable (N K : ℕ) (n : Fin N) :
        Differentiable ℝ fun (z : Vec (N * K)) => logitRow N K z n
        theorem Proofs.smoothedBatchLoss_differentiable (N K : ℕ) (α B : ℝ) (t : Vec (N * (1 * K))) :
        theorem Proofs.rowSumLoss_pdiv (N K : ℕ) (ℓ : Fin N → Vec K → ℝ) (hℓ : ∀ (m : Fin N), Differentiable ℝ fun (r : Vec K) (x : Fin 1) => ℓ m r) (z : Vec (N * K)) (n : Fin N) (j : Fin K) :
        pdiv (fun (z' : Vec (N * K)) (x : Fin 1) => ∑ m : Fin N, ℓ m (logitRow N K z' m)) z (finProdFinEquiv (n, j)) 0 = pdiv (fun (r : Vec K) (x : Fin 1) => ℓ n r) (logitRow N K z n) j 0

        A loss summed over the rows has each row's own gradient: for Σ_m ℓ_m(zₘ), the partial at (n, j) is ∂ℓₙ/∂z_j at example n's logits — no other example contributes. Shared by every batched loss the renders emit (smoothedBatchLoss, bceBatchLoss).

        theorem Proofs.smoothedBatchLoss_pdiv (N K : ℕ) (α B : ℝ) (t : Vec (N * (1 * K))) (z : Vec (N * K)) (n : Fin N) (j : Fin K) :
        pdiv (smoothedBatchLoss N K α B t) z (finProdFinEquiv (n, j)) 0 = pdiv (fun (z' : Vec K) (x : Fin 1) => softCE K (smoothTarget K α (targetRow N K t n)) z') (logitRow N K z n) j 0 / B

        The batched loss's gradient, entry (n, j): example n's own soft-CE gradient at its logits, over B — no other example contributes.

        theorem Proofs.smoothedBatchLoss_grad (N K : ℕ) (hK : 0 < K) (α B : ℝ) (aStr negAK bStr logN ohN : String) (t : Vec (N * (1 * K))) (z : Vec (N * K)) (ht : ∀ (n : Fin N), ∑ k : Fin K, targetRow N K t n k = 1) (J : Fin (N * K)) :
        pdiv (smoothedBatchLoss N K α B t) z J 0 = BackLinks.unrowB N K (StableHLO.den (smoothedLossCotGraph N K α B aStr negAK bStr logN ohN (BackLinks.rowB N K z) t)) J

        The emitted cotangent is the batched loss's gradient. Read at the head's N·K index (unrowB), the six-op chain at the logits rowB z is ∇ smoothedBatchLoss at z, whenever every example's target sums to 1.

        noncomputable def Proofs.smoothedBatchLossDiv (N K : ℕ) (α B : ℝ) (t z : Vec (N * K)) :
        Vec 1

        The batched label-smoothed loss at the plain N·K index, the target read per example with no row index — the loss whose gradient smoothedLossCotGraphDiv emits (ConvNeXt, ViT).

        Equations
        Instances For
          theorem Proofs.smoothedBatchLossDiv_grad (N K : ℕ) (hK : 0 < K) (α B : ℝ) (aStr negAK bStr logN ohN : String) (t z : Vec (N * K)) (ht : ∀ (n : Fin N), ∑ k : Fin K, StableHLO.batchSlice N K t n k = 1) (J : Fin (N * K)) :
          pdiv (smoothedBatchLossDiv N K α B t) z J 0 = StableHLO.den (smoothedLossCotGraphDiv N K α B aStr negAK bStr logN ohN z t) J

          The emitted softmaxDiv cotangent is the batched loss's gradient, entry by entry, whenever every example's target sums to 1 — smoothedBatchLoss_grad at the plain N·K index.