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).
Example n's logits out of the flat N·K batch.
Equations
- Proofs.logitRow N K z n k = z (finProdFinEquiv (n, k))
Instances For
Example n's target out of the N·(1·K) graph input %onehot.
Equations
- Proofs.targetRow N K t n = Proofs.Mat.unflatten (Proofs.StableHLO.batchSlice N (1 * K) t n) 0
Instances For
The batched label-smoothed loss: Σ_n softCE(smooth(tₙ), zₙ) / B, as a function of the
flat logits.
Equations
- Proofs.smoothedBatchLoss N K α B t z x✝ = ∑ n : Fin N, Proofs.softCE K (Proofs.smoothTarget K α (Proofs.targetRow N K t n)) (Proofs.logitRow N K z n) / B
Instances For
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).
The batched loss's gradient, entry (n, j): example n's own soft-CE gradient at its
logits, over B — no other example contributes.
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.
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
- Proofs.smoothedBatchLossDiv N K α B t z x✝ = ∑ n : Fin N, Proofs.softCE K (Proofs.smoothTarget K α (Proofs.StableHLO.batchSlice N K t n)) (Proofs.logitRow N K z n) / B
Instances For
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.