Documentation

LeanMlir.Proofs.Nets.Small.CifarParamGrad

The CIFAR CNN — every parameter gradient node IS the loss's derivative, up to pool twins #

cifar_train_step_tied_certified ties each of the fourteen SGD updates to the certified per-layer Jacobian contracted with the cotangent the emitted chain threads to it. cifar_net_lossGrad states that the un-fused *Grad node of each layer, at the chain cotangent, is the gradient of the loss in that parameter, for any loss L of the logits with gradient g there; cifar_net_lossGrad_CE instantiates it at the softmax cross-entropy the render emits.

The two pools are handled as the MNIST CNN's one (CnnFold.cnn_net_lossGrad): each pool's clause allows ties between twins, cells equal at every weight upstream of that pool (CnnFold.CnnPoolTwin for the first, CifarPoolTwin2 for the second), and each pool's backward routes a window's cotangent to the one cell a selection names (CnnFold.cnnChainCotW2Sel, cifarChainCotW2Sel), as the rendered select_and_scatter does. At the first argmax of each window, the cell that op picks, the step tie's own chains are these (CnnFold.cnnChainCotW2_eq_sel, cifarChainCotW2_eq_sel). A parameter of the first stage sees both pools move; its germ rewrites the outer pool first, at the true pre-activation, then the inner one.

Hypotheses. Odd kernels, every ReLU off its kink, every pool window dead or tied only between twins, each selection naming a maximum of every window (CifarLossSmoothAt). Scope. One example (the emitted module batch-contracts; den is per-example).

noncomputable def Proofs.CifarFold.cifarUp2 {c1 c2 h w kH kW : ℕ} (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) (z : Vec (c1 * (2 * (2 * h)) * (2 * (2 * w)))) :
Vec (c2 * (2 * h) * (2 * w))

From the first pool's pre-activation to the second's: ReLU, pool, conv₃, ReLU, conv₄.

Equations
  • One or more equations did not get rendered due to their size.
Instances For
    theorem Proofs.CifarFold.cifarUp2_continuous {c1 c2 h w kH kW : ℕ} (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) :
    Continuous (cifarUp2 W₃ b₃ W₄ b₄)
    noncomputable def Proofs.CifarFold.cifarPre2 {ic c1 c2 h w kH kW : ℕ} (W₁ : Kernel4 c1 ic kH kW) (b₁ : Vec c1) (W₂ : Kernel4 c1 c1 kH kW) (b₂ : Vec c1) (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) (x : Vec (ic * (2 * (2 * h)) * (2 * (2 * w)))) :
    Vec (c2 * (2 * h) * (2 * w))

    The second pool's pre-activation (conv₄'s output).

    Equations
    Instances For
      def Proofs.CifarFold.CifarPoolTwin2 {ic h w : ℕ} (c1 c2 kH kW : ℕ) (x : Vec (ic * (2 * (2 * h)) * (2 * (2 * w)))) (p q : Fin (2 * h) × Fin (2 * w)) :

      Two cells of the second pool's input are twins: equal at every weight of the four convs, in every channel.

      Equations
      • One or more equations did not get rendered due to their size.
      Instances For
        noncomputable def Proofs.CifarFold.cifarChainCotW2Sel {c1 c2 h w kH kW : ℕ} (σ₁ : Fin c1 → Fin (2 * h) → Fin (2 * w) → Fin 2 × Fin 2) (W₃ : Kernel4 c2 c1 kH kW) (hc2 : Vec (c1 * (2 * (2 * h)) * (2 * (2 * w)))) (cotW3 : Vec (c2 * (2 * h) * (2 * w))) :
        Vec (c1 * (2 * (2 * h)) * (2 * (2 * w)))

        The conv₂-output cotangent at a first-pool selection σ₁: cifarChainCotW2 with the pool backward routing each window's cotangent to the ONE cell σ₁ names, then the ReLU mask.

        Equations
        • One or more equations did not get rendered due to their size.
        Instances For
          theorem Proofs.CifarFold.cifarChainCotW2_eq_sel {c1 c2 h w kH kW : ℕ} (W₃ : Kernel4 c2 c1 kH kW) (ac2 : Tensor3 c1 (2 * (2 * h)) (2 * (2 * w))) (hc2 : Vec (c1 * (2 * (2 * h)) * (2 * (2 * w)))) (cotW3 : Vec (c2 * (2 * h) * (2 * w))) :
          cifarChainCotW2 W₃ ac2 hc2 cotW3 = cifarChainCotW2Sel (maxPool2Argmax ac2) W₃ hc2 cotW3

          The step tie's conv₂ cotangent is the capstone's at the first argmax of the first pool's windows (maxPool2Argmax), the cell the rendered select_and_scatter picks, at every point.

          structure Proofs.CifarFold.CifarLossSmoothAt {ic c1 c2 h w d1 kH kW : ℕ} (W₁ : Kernel4 c1 ic kH kW) (b₁ : Vec c1) (W₂ : Kernel4 c1 c1 kH kW) (b₂ : Vec c1) (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) (W₅ : Mat (c2 * h * w) d1) (b₅ : Vec d1) (W₆ : Mat d1 d1) (b₆ : Vec d1) (x : Vec (ic * (2 * (2 * h)) * (2 * (2 * w)))) (σ₁ : Fin c1 → Fin (2 * h) → Fin (2 * w) → Fin 2 × Fin 2) (σ₂ : Fin c2 → Fin h → Fin w → Fin 2 × Fin 2) :

          The smooth-point bundle the loss gradient needs. Every ReLU off its kink; every window of each pool dead or tied only between that pool's twins; each selection naming a maximum of every window.

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

            Every CIFAR-CNN parameter node is the gradient of L in that parameter: the fourteen un-fused nodes, each at the cotangent the chain threads to its layer (each pool routed at its selection), stated against L of cifarCnnForward with that one parameter varied.

            Equations
            • One or more equations did not get rendered due to their size.
            Instances For
              theorem Proofs.CifarFold.cifar_net_lossGrad {ic c1 c2 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 c1 ic kH kW) (b₁ : Vec c1) (W₂ : Kernel4 c1 c1 kH kW) (b₂ : Vec c1) (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) (W₅ : Mat (c2 * h * w) d1) (b₅ : Vec d1) (W₆ : Mat d1 d1) (b₆ : Vec d1) (W₇ : Mat d1 nClasses) (b₇ : Vec nClasses) (x : Vec (ic * (2 * (2 * h)) * (2 * (2 * w)))) (σ₁ : Fin c1 → Fin (2 * h) → Fin (2 * w) → Fin 2 × Fin 2) (σ₂ : Fin c2 → Fin h → Fin w → Fin 2 × Fin 2) (hx : CifarLossSmoothAt W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ W₆ b₆ x σ₁ σ₂) {L : Vec nClasses → Vec 1} {g : Vec nClasses} (hL : HasGradAt L (cifarCnnForward W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ W₆ b₆ W₇ b₇ x) g) :
              CifarNetLossTied xN cotN W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ W₆ b₆ W₇ b₇ x σ₁ σ₂ L g

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

              Hypotheses: odd kernels, and CifarLossSmoothAt — every ReLU off its kink, every window of each pool dead or tied only between cells that are the same function of the weights upstream of it, each selection naming a maximum of every window.

              theorem Proofs.CifarFold.cifar_net_lossGrad_CE {ic c1 c2 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 c1 ic kH kW) (b₁ : Vec c1) (W₂ : Kernel4 c1 c1 kH kW) (b₂ : Vec c1) (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) (W₅ : Mat (c2 * h * w) d1) (b₅ : Vec d1) (W₆ : Mat d1 d1) (b₆ : Vec d1) (W₇ : Mat d1 nClasses) (b₇ : Vec nClasses) (x : Vec (ic * (2 * (2 * h)) * (2 * (2 * w)))) (σ₁ : Fin c1 → Fin (2 * h) → Fin (2 * w) → Fin 2 × Fin 2) (σ₂ : Fin c2 → Fin h → Fin w → Fin 2 × Fin 2) (hx : CifarLossSmoothAt W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ W₆ b₆ x σ₁ σ₂) (label : Fin nClasses) :
              CifarNetLossTied xN cotN W₁ b₁ W₂ b₂ 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 (cifarCnnForward W₁ b₁ W₂ b₂ 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.