Documentation

LeanMlir.Proofs.Nets.Small.MlpParamGrad

The MNIST MLP — every parameter gradient node IS the loss's derivative #

mlp_train_step_tied_certified ties each of the six SGD updates to the certified per-layer Jacobian contracted with the cotangent the emitted chain threads to it (mlpCotOut1, mlpCotOut0), and leaves open whether that cotangent is the loss gradient at the hidden layers. mlp_net_lossGrad closes it: at the same cotangents, the un-fused weightGrad / biasGrad node of each layer is the gradient of the loss in that parameter, for any loss L of the logits with gradient g there; mlp_net_lossGrad_CE instantiates it at the softmax cross-entropy the render emits. The fused weightSgd / biasSgd ops are θ − lr· these nodes (SmallParamGrad.weightSgd_eq_grad, SmallParamGrad.biasSgd_eq_grad).

How. The loss read at the logits is pulled back one certified stage at a time (SmallParamGrad.hasGradAt_dense, SmallParamGrad.hasGradAt_relu), each landing on the emitted chain's cotangent; at each layer's output the node lemma (SmallParamGrad.denseW_hasGradAt) turns it into the parameter gradient.

Hypotheses. Both hidden pre-activations off the ReLU kink (the pair mlpHasVJPAt takes). Scope. One example (the emitted module batch-contracts; den is per-example).

def Proofs.MlpFold.MlpNetLossTied {d₀ d₁ d₂ d₃ : ℕ} (aN cotN : String) (W₀ : Mat d₀ d₁) (b₀ : Vec d₁) (W₁ : Mat d₁ d₂) (b₁ : Vec d₂) (W₂ : Mat d₂ d₃) (b₂ : Vec d₃) (x : Vec d₀) (L : Vec d₃ → Vec 1) (g : Vec d₃) :

Every MLP parameter node is the gradient of L in that parameter: the six nodes, each at the cotangent the emitted chain threads to its layer (g at the logits, mlpCotOut1, mlpCotOut0), stated against L of mlpForward with that one parameter varied.

Equations
  • One or more equations did not get rendered due to their size.
Instances For
    theorem Proofs.MlpFold.mlp_net_lossGrad {d₀ d₁ d₂ d₃ : ℕ} (aN cotN : String) (W₀ : Mat d₀ d₁) (b₀ : Vec d₁) (W₁ : Mat d₁ d₂) (b₁ : Vec d₂) (W₂ : Mat d₂ d₃) (b₂ : Vec d₃) (x : Vec d₀) (h₀ : ∀ (k : Fin d₁), dense W₀ b₀ x k ≠ 0) (h₁ : ∀ (k : Fin d₂), dense W₁ b₁ (relu d₁ (dense W₀ b₀ x)) k ≠ 0) {L : Vec d₃ → Vec 1} {g : Vec d₃} (hL : HasGradAt L (mlpForward W₀ b₀ W₁ b₁ W₂ b₂ x) g) :
    MlpNetLossTied aN cotN W₀ b₀ W₁ b₁ W₂ b₂ x L g

    Every MLP parameter node is the gradient of L in that parameter, whenever g is L's gradient at the logits and both hidden pre-activations are off the ReLU kink.

    theorem Proofs.MlpFold.mlp_net_lossGrad_CE {d₀ d₁ d₂ d₃ : ℕ} (aN cotN nlogN ohN : String) (W₀ : Mat d₀ d₁) (b₀ : Vec d₁) (W₁ : Mat d₁ d₂) (b₁ : Vec d₂) (W₂ : Mat d₂ d₃) (b₂ : Vec d₃) (x : Vec d₀) (label : Fin d₃) (h₀ : ∀ (k : Fin d₁), dense W₀ b₀ x k ≠ 0) (h₁ : ∀ (k : Fin d₂), dense W₁ b₁ (relu d₁ (dense W₀ b₀ x)) k ≠ 0) :
    MlpNetLossTied aN cotN W₀ b₀ W₁ b₁ W₂ b₂ x (fun (z : Vec d₃) (x : Fin 1) => crossEntropy d₃ z label) (StableHLO.den ((StableHLO.SHlo.operand nlogN (mlpForward W₀ b₀ W₁ b₁ W₂ b₂ x)).expe.softmaxDiv.sub (StableHLO.SHlo.operand ohN (oneHot d₃ label))))

    The artifact's loss: every node is the gradient of the softmax cross-entropy at label, g the emitted loss cotangent.