Documentation

LeanMlir.Proofs.Foundation.BceBatchLoss

The batched BCE-with-logits loss, and its gradient is the emitted cotangent #

BceLossCot identifies the bce := true renders' three-op cotangent ROW BY ROW: at example n it is ∂bceLogits/∂z at that example's logits over the baked N·K. This file states the loss as one function of the flat batch of logits — bceBatchLoss, every example's class-summed BCE-with-logits over N·K, i.e. the mean over B×K — and proves its full gradient is the emitted cotangent, read at the head's N·K index (bceBatchLoss_grad). It is SmoothedBatchLoss's twin for ResNet-50's BCE artifacts, and like bceLossCotGraph_row it needs no hypothesis on the target.

noncomputable def Proofs.bceBatchLoss (N K : ℕ) (t : Vec (N * (1 * K))) (z : Vec (N * K)) :
Vec 1

The batched BCE-with-logits loss: Σ_n bceLogits(tₙ, zₙ) / (N·K), as a function of the flat logits — the mean over B×K that timm's BinaryCrossEntropy computes.

Equations
Instances For
    theorem Proofs.bceBatchLoss_pdiv (N K : ℕ) (t : Vec (N * (1 * K))) (z : Vec (N * K)) (n : Fin N) (j : Fin K) :
    pdiv (bceBatchLoss N K t) z (finProdFinEquiv (n, j)) 0 = pdiv (fun (z' : Vec K) (x : Fin 1) => bceLogits K (targetRow N K t n) z') (logitRow N K z n) j 0 / (↑N * ↑K)

    The batched BCE loss's gradient, entry (n, j): example n's own BCE gradient at its logits, over N·K.

    theorem Proofs.bceBatchLoss_grad (N K : ℕ) (bStr logN ohN : String) (t : Vec (N * (1 * K))) (z : Vec (N * K)) (J : Fin (N * K)) :
    pdiv (bceBatchLoss N K t) z J 0 = BackLinks.unrowB N K (StableHLO.den (bceLossCotGraph N K (↑N * ↑K) bStr logN ohN (BackLinks.rowB N K z) t)) J

    The emitted BCE cotangent is the batched loss's gradient. Read at the head's N·K index (unrowB), the three-op chain at the logits rowB z, with the committed divisor N·K, is ∇ bceBatchLoss at z — for every target.