Documentation

LeanMlir.Proofs.Nets.Small.CnnParamGrad

The MNIST CNN — every parameter gradient node IS the loss's derivative, at real MNIST #

cnn_train_step_tied_certified ties each of the ten SGD updates to the certified per-layer Jacobian contracted with the cotangent the emitted chain threads to it, and leaves open whether that cotangent is the loss gradient below the output. cnn_net_lossGrad closes it: the un-fused *Grad node of each layer, at the chain cotangent, is the gradient of the loss in that parameter (HasGradAt), for any loss L of the logits with gradient g there; cnn_net_lossGrad_CE instantiates it at the softmax cross-entropy the render emits. The fused *Sgd ops are θ − lr· these nodes (SmallParamGrad.convWeightSgd_eq_grad and its peers).

The pool's clause is stated for the parameters, not the image. On real MNIST almost every image has a 2×2 window whose positive maximum sits at two cells, because the two cells read identical input (a constant background patch): the net has no derivative in its input there, and the pool none in its activation. But two such cells are the same function of the conv weights (CnnPoolTwin, implied by identical two-layer receptive fields, cnnPoolTwin_of_convPatchEq2), so along any parameter the pooled ReLU is the gather at a fixed selection (SmallParamGrad.maxPool_relu_eventuallyEq_sel), and the loss IS differentiable in the parameters. CnnLossSmoothAt allows exactly those ties. The probe scripts/probes/mnist_pool_twin_probe.py checks this clause on the MNIST test set.

Which cotangent. At a tied window the pool's backward must pick one cell. The rendered select_and_scatter (select = GE) does: it routes each window's cotangent to the window's first maximal cell, the gather's adjoint. The capstone is stated at any selection σ naming a maximum of every window (cnnChainCotW2Sel, SmallParamGrad.PoolSelDom). The step tie's cnnChainCotW2 reads the pool backward as maxPoolBackDenote, which routes to the first argmax (maxPool2Argmax), so it IS cnnChainCotW2Sel at that selection, at every point (cnnChainCotW2_eq_sel); SmallParamGrad.poolSelDom_argmax discharges the selection clause there.

How. The loss read at the logits is pulled back through the dense head (SmallParamGrad.hasGradAt_dense, SmallParamGrad.hasGradAt_relu), through the pool as the gather at σ (SmallParamGrad.hasGradAt_gatherRelu; at the point itself the pool IS that gather), then through the convs (SmallParamGrad.hasGradAt_conv). Each conv node is the gather model's parameter gradient, moved to the real net by the germ (HasGradAt.congr_of_eventuallyEq).

Hypotheses. Odd kernels (the rendered conv backward is the conv VJP there), and CnnLossSmoothAt: every ReLU off its kink, and every pool window dead or tied only between twins. Scope. One example (the emitted module batch-contracts; den is per-example).

noncomputable def Proofs.CnnFold.cnnPoolPre {ic h w c kH kW : ℕ} (W₁ : Kernel4 c ic kH kW) (b₁ : Vec c) (W₂ : Kernel4 c c kH kW) (b₂ : Vec c) (x : Vec (ic * (2 * h) * (2 * w))) :
Vec (c * (2 * h) * (2 * w))

The conv2 pre-activation, the input of the pool's ReLU.

Equations
Instances For
    def Proofs.CnnFold.CnnPoolTwin {ic h w : ℕ} (c kH kW : ℕ) (x : Vec (ic * (2 * h) * (2 * w))) (p q : Fin (2 * h) × Fin (2 * w)) :

    Two pool-input cells are twins: they are equal at every conv weight, in every channel. At such a pair the pool can tie at a positive maximum, and the tie is the same at every weight, which is why the parameter gradient survives it.

    Equations
    • One or more equations did not get rendered due to their size.
    Instances For
      theorem Proofs.CnnFold.cnnPoolTwin_of_convPatchEq2 {ic h w c kH kW : ℕ} {x : Vec (ic * (2 * h) * (2 * w))} {p q : Fin (2 * h) × Fin (2 * w)} (hpq : ConvPatchEq2 kH kW (Tensor3.unflatten x) p q) :
      CnnPoolTwin c kH kW x p q

      Cells with identical two-layer receptive fields in the image are twins.

      noncomputable def Proofs.CnnFold.cnnChainCotW2Sel {c h w d1 nClasses : ℕ} (σ : Fin c → Fin h → Fin w → Fin 2 × Fin 2) (W₃ : Mat (c * h * w) d1) (W₄ : Mat d1 d1) (W₅ : Mat d1 nClasses) (h3 h4 : Vec d1) (hc2 : Vec (c * (2 * h) * (2 * w))) (dy : Vec nClasses) :
      Vec (c * (2 * h) * (2 * w))

      The conv2-output cotangent at a pool selection σ: cnnChainCotW2 with the pool backward routing each window's cotangent to the ONE cell σ names (selScatter), as the rendered select_and_scatter does, then the ReLU mask.

      Equations
      • One or more equations did not get rendered due to their size.
      Instances For
        theorem Proofs.CnnFold.cnnChainCotW2_eq_sel {c h w d1 nClasses : ℕ} (W₃ : Mat (c * h * w) d1) (W₄ : Mat d1 d1) (W₅ : Mat d1 nClasses) (h3 h4 : Vec d1) (ac2 : Tensor3 c (2 * h) (2 * w)) (hc2 : Vec (c * (2 * h) * (2 * w))) (dy : Vec nClasses) :
        cnnChainCotW2 W₃ W₄ W₅ h3 h4 ac2 hc2 dy = cnnChainCotW2Sel (maxPool2Argmax ac2) W₃ W₄ W₅ h3 h4 hc2 dy

        The step tie's conv₂ cotangent is the capstone's at the first argmax. The rendered pool backward routes each window to its first maximum (maxPool2Argmax), so cnnChainCotW2 is cnnChainCotW2Sel at that selection, at every point, ties included.

        structure Proofs.CnnFold.CnnLossSmoothAt {ic c h w d1 kH kW : ℕ} (W₁ : Kernel4 c ic kH kW) (b₁ : Vec c) (W₂ : Kernel4 c c kH kW) (b₂ : Vec c) (W₃ : Mat (c * h * w) d1) (b₃ : Vec d1) (W₄ : Mat d1 d1) (b₄ : Vec d1) (x : Vec (ic * (2 * h) * (2 * w))) (σ : Fin c → Fin h → Fin w → Fin 2 × Fin 2) :

        The smooth-point bundle the loss gradient needs. Every ReLU off its kink; every pool window dead or tied only between twins (CnnPoolTwin); and σ names a maximum of every window.

        Instances For
          def Proofs.CnnFold.CnnNetLossTied {ic c h w d1 nClasses kH kW : ℕ} (xN cotN : String) (W₁ : Kernel4 c ic kH kW) (b₁ : Vec c) (W₂ : Kernel4 c c kH kW) (b₂ : Vec c) (W₃ : Mat (c * h * w) d1) (b₃ : Vec d1) (W₄ : Mat d1 d1) (b₄ : Vec d1) (W₅ : Mat d1 nClasses) (b₅ : Vec nClasses) (x : Vec (ic * (2 * h) * (2 * w))) (σ : Fin c → Fin h → Fin w → Fin 2 × Fin 2) (L : Vec nClasses → Vec 1) (g : Vec nClasses) :

          Every MNIST-CNN parameter node is the gradient of L in that parameter: the ten un-fused nodes, each at the cotangent the chain threads to its layer (the head's mlpCotOut1 / mlpCotOut0, the pool routed at σ, the conv backward), stated against L of mnistCnnNoBnForward with that one parameter varied.

          Equations
          • One or more equations did not get rendered due to their size.
          Instances For
            theorem Proofs.CnnFold.cnn_net_lossGrad {ic c h w d1 nClasses kH kW : ℕ} (xN cotN : String) (hkH : 2 * ((kH - 1) / 2) + 1 = kH) (hkW : 2 * ((kW - 1) / 2) + 1 = kW) (W₁ : Kernel4 c ic kH kW) (b₁ : Vec c) (W₂ : Kernel4 c c kH kW) (b₂ : Vec c) (W₃ : Mat (c * h * w) d1) (b₃ : Vec d1) (W₄ : Mat d1 d1) (b₄ : Vec d1) (W₅ : Mat d1 nClasses) (b₅ : Vec nClasses) (x : Vec (ic * (2 * h) * (2 * w))) (σ : Fin c → Fin h → Fin w → Fin 2 × Fin 2) (hx : CnnLossSmoothAt W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ x σ) {L : Vec nClasses → Vec 1} {g : Vec nClasses} (hL : HasGradAt L (mnistCnnNoBnForward W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ x) g) :
            CnnNetLossTied xN cotN W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ x σ L g

            Every MNIST-CNN parameter node is the gradient of L in that parameter, whenever g is L's gradient at the logits.

            Hypotheses: odd kernels, every ReLU off its kink, and every pool window dead or tied only between cells that are the same function of the conv weights (CnnLossSmoothAt), with σ naming a maximum of every window.

            theorem Proofs.CnnFold.cnn_net_lossGrad_CE {ic c h w d1 nClasses kH kW : ℕ} (xN cotN nlogN ohN : String) (hkH : 2 * ((kH - 1) / 2) + 1 = kH) (hkW : 2 * ((kW - 1) / 2) + 1 = kW) (W₁ : Kernel4 c ic kH kW) (b₁ : Vec c) (W₂ : Kernel4 c c kH kW) (b₂ : Vec c) (W₃ : Mat (c * h * w) d1) (b₃ : Vec d1) (W₄ : Mat d1 d1) (b₄ : Vec d1) (W₅ : Mat d1 nClasses) (b₅ : Vec nClasses) (x : Vec (ic * (2 * h) * (2 * w))) (σ : Fin c → Fin h → Fin w → Fin 2 × Fin 2) (hx : CnnLossSmoothAt W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ x σ) (label : Fin nClasses) :
            CnnNetLossTied xN cotN W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ x σ (fun (z : Vec nClasses) (x : Fin 1) => crossEntropy nClasses z label) (StableHLO.den ((StableHLO.SHlo.operand nlogN (mnistCnnNoBnForward W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ x)).expe.softmaxDiv.sub (StableHLO.SHlo.operand ohN (oneHot nClasses label))))

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