Documentation

LeanMlir.Proofs.Nets.Small.Cifar8ParamGrad

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

cifar8_train_step_tiedG states the 22 un-fused gradient nodes the packed cifar8w_* arms emit, each at the cotangent the chain threads to it. cifar8_net_lossGrad states that each node, at the chain cotangent, is the gradient of the loss in that parameter, for any loss L of the logits with gradient g there; cifar8_net_lossGrad_CE instantiates it at the softmax cross-entropy the render emits.

The four pools are handled as the 2-stage net's two (CifarFold.cifar_net_lossGrad): each pool's clause allows ties between twins, cells equal at every weight upstream of that pool (CnnFold.CnnPoolTwin, Cifar8PoolTwin2, Cifar8PoolTwin3, Cifar8PoolTwin4), and each pool's backward routes a window's cotangent to the one cell a selection names, as the rendered select_and_scatter does. A stage-s parameter sees pools s…4 move; its germ rewrites them outermost first, at the true pre-activations (germ4 … germ1 in the proof). cifar8Up is the step between two pools' pre-activations.

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 (Cifar8LossSmoothAt). Scope. One example (the emitted module batch-contracts; den is per-example).

noncomputable def Proofs.Cifar8TieG.cifar8Up {c c' H W kH kW : ℕ} (Wc : Kernel4 c' c kH kW) (bc : Vec c') (Wd : Kernel4 c' c' kH kW) (bd : Vec c') (z : Vec (c * (2 * H) * (2 * W))) :
Vec (c' * H * W)

From one pool's pre-activation to the next pool's: ReLU, pool, conv, ReLU, conv.

Equations
Instances For
    theorem Proofs.Cifar8TieG.cifar8Up_continuous {c c' H W kH kW : ℕ} (Wc : Kernel4 c' c kH kW) (bc : Vec c') (Wd : Kernel4 c' c' kH kW) (bd : Vec c') :
    Continuous (cifar8Up Wc bc Wd bd)
    noncomputable def Proofs.Cifar8TieG.cifar8Pre2 {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 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) :
    Vec (c2 * (2 * (2 * (2 * h))) * (2 * (2 * (2 * w))))

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

    Equations
    Instances For
      noncomputable def Proofs.Cifar8TieG.cifar8Pre3 {ic c1 c2 c3 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) (W₅ : Kernel4 c3 c2 kH kW) (b₅ : Vec c3) (W₆ : Kernel4 c3 c3 kH kW) (b₆ : Vec c3) (x : Vec (ic * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) :
      Vec (c3 * (2 * (2 * h)) * (2 * (2 * w)))

      The third pool's pre-activation (conv₆'s output).

      Equations
      Instances For
        noncomputable def Proofs.Cifar8TieG.cifar8Pre4 {ic c1 c2 c3 c4 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) (W₅ : Kernel4 c3 c2 kH kW) (b₅ : Vec c3) (W₆ : Kernel4 c3 c3 kH kW) (b₆ : Vec c3) (W₇ : Kernel4 c4 c3 kH kW) (b₇ : Vec c4) (W₈ : Kernel4 c4 c4 kH kW) (b₈ : Vec c4) (x : Vec (ic * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) :
        Vec (c4 * (2 * h) * (2 * w))

        The fourth pool's pre-activation (conv₈'s output).

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

          Twins of the second pool: equal at every weight of convs 1–4, in every channel.

          Equations
          • One or more equations did not get rendered due to their size.
          Instances For
            def Proofs.Cifar8TieG.Cifar8PoolTwin3 {ic h w : ℕ} (c1 c2 c3 kH kW : ℕ) (x : Vec (ic * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) (p q : Fin (2 * (2 * h)) × Fin (2 * (2 * w))) :

            Twins of the third pool: equal at every weight of convs 1–6, in every channel.

            Equations
            • One or more equations did not get rendered due to their size.
            Instances For
              def Proofs.Cifar8TieG.Cifar8PoolTwin4 {ic h w : ℕ} (c1 c2 c3 c4 kH kW : ℕ) (x : Vec (ic * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) (p q : Fin (2 * h) × Fin (2 * w)) :

              Twins of the fourth pool: equal at every weight of the eight convs, in every channel.

              Equations
              • One or more equations did not get rendered due to their size.
              Instances For
                structure Proofs.Cifar8TieG.Cifar8LossSmoothAt {ic c1 c2 c3 c4 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₅ : Kernel4 c3 c2 kH kW) (b₅ : Vec c3) (W₆ : Kernel4 c3 c3 kH kW) (b₆ : Vec c3) (W₇ : Kernel4 c4 c3 kH kW) (b₇ : Vec c4) (W₈ : Kernel4 c4 c4 kH kW) (b₈ : Vec c4) (W₉ : Mat (c4 * h * w) d1) (b₉ : Vec d1) (Wa : Mat d1 d1) (ba : Vec d1) (x : Vec (ic * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) (σ₁ : Fin c1 → Fin (2 * (2 * (2 * h))) → Fin (2 * (2 * (2 * w))) → Fin 2 × Fin 2) (σ₂ : Fin c2 → Fin (2 * (2 * h)) → Fin (2 * (2 * w)) → Fin 2 × Fin 2) (σ₃ : Fin c3 → Fin (2 * h) → Fin (2 * w) → Fin 2 × Fin 2) (σ₄ : Fin c4 → 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.Cifar8TieG.Cifar8NetLossTied {ic c1 c2 c3 c4 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₅ : Kernel4 c3 c2 kH kW) (b₅ : Vec c3) (W₆ : Kernel4 c3 c3 kH kW) (b₆ : Vec c3) (W₇ : Kernel4 c4 c3 kH kW) (b₇ : Vec c4) (W₈ : Kernel4 c4 c4 kH kW) (b₈ : Vec c4) (W₉ : Mat (c4 * h * w) d1) (b₉ : Vec d1) (Wa : Mat d1 d1) (ba : Vec d1) (Wb : Mat d1 nClasses) (bb : Vec nClasses) (x : Vec (ic * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) (σ₁ : Fin c1 → Fin (2 * (2 * (2 * h))) → Fin (2 * (2 * (2 * w))) → Fin 2 × Fin 2) (σ₂ : Fin c2 → Fin (2 * (2 * h)) → Fin (2 * (2 * w)) → Fin 2 × Fin 2) (σ₃ : Fin c3 → Fin (2 * h) → Fin (2 * w) → Fin 2 × Fin 2) (σ₄ : Fin c4 → Fin h → Fin w → Fin 2 × Fin 2) (L : Vec nClasses → Vec 1) (g : Vec nClasses) :

                  Every cifar8 gradient node is the gradient of L in that parameter: the 22 un-fused nodes cifar8_train_step_tiedG states, each at the cotangent the chain threads to its layer (each pool routed at its selection), stated against L of cifarCnn8Forward with that one parameter varied.

                  Equations
                  • One or more equations did not get rendered due to their size.
                  Instances For
                    theorem Proofs.Cifar8TieG.cifar8_net_lossGrad {ic c1 c2 c3 c4 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₅ : Kernel4 c3 c2 kH kW) (b₅ : Vec c3) (W₆ : Kernel4 c3 c3 kH kW) (b₆ : Vec c3) (W₇ : Kernel4 c4 c3 kH kW) (b₇ : Vec c4) (W₈ : Kernel4 c4 c4 kH kW) (b₈ : Vec c4) (W₉ : Mat (c4 * h * w) d1) (b₉ : Vec d1) (Wa : Mat d1 d1) (ba : Vec d1) (Wb : Mat d1 nClasses) (bb : Vec nClasses) (x : Vec (ic * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) (σ₁ : Fin c1 → Fin (2 * (2 * (2 * h))) → Fin (2 * (2 * (2 * w))) → Fin 2 × Fin 2) (σ₂ : Fin c2 → Fin (2 * (2 * h)) → Fin (2 * (2 * w)) → Fin 2 × Fin 2) (σ₃ : Fin c3 → Fin (2 * h) → Fin (2 * w) → Fin 2 × Fin 2) (σ₄ : Fin c4 → Fin h → Fin w → Fin 2 × Fin 2) (hx : Cifar8LossSmoothAt W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ W₆ b₆ W₇ b₇ W₈ b₈ W₉ b₉ Wa ba x σ₁ σ₂ σ₃ σ₄) {L : Vec nClasses → Vec 1} {g : Vec nClasses} (hL : HasGradAt L (cifarCnn8Forward W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ W₆ b₆ W₇ b₇ W₈ b₈ W₉ b₉ Wa ba Wb bb x) g) :
                    Cifar8NetLossTied xN cotN W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ W₆ b₆ W₇ b₇ W₈ b₈ W₉ b₉ Wa ba Wb bb x σ₁ σ₂ σ₃ σ₄ L g

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

                    Hypotheses: odd kernels, and Cifar8LossSmoothAt — 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.Cifar8TieG.cifar8_net_lossGrad_CE {ic c1 c2 c3 c4 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₅ : Kernel4 c3 c2 kH kW) (b₅ : Vec c3) (W₆ : Kernel4 c3 c3 kH kW) (b₆ : Vec c3) (W₇ : Kernel4 c4 c3 kH kW) (b₇ : Vec c4) (W₈ : Kernel4 c4 c4 kH kW) (b₈ : Vec c4) (W₉ : Mat (c4 * h * w) d1) (b₉ : Vec d1) (Wa : Mat d1 d1) (ba : Vec d1) (Wb : Mat d1 nClasses) (bb : Vec nClasses) (x : Vec (ic * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) (σ₁ : Fin c1 → Fin (2 * (2 * (2 * h))) → Fin (2 * (2 * (2 * w))) → Fin 2 × Fin 2) (σ₂ : Fin c2 → Fin (2 * (2 * h)) → Fin (2 * (2 * w)) → Fin 2 × Fin 2) (σ₃ : Fin c3 → Fin (2 * h) → Fin (2 * w) → Fin 2 × Fin 2) (σ₄ : Fin c4 → Fin h → Fin w → Fin 2 × Fin 2) (hx : Cifar8LossSmoothAt W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ W₆ b₆ W₇ b₇ W₈ b₈ W₉ b₉ Wa ba x σ₁ σ₂ σ₃ σ₄) (label : Fin nClasses) :
                    Cifar8NetLossTied xN cotN W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ W₆ b₆ W₇ b₇ W₈ b₈ W₉ b₉ Wa ba Wb bb x σ₁ σ₂ σ₃ σ₄ (fun (z : Vec nClasses) (x : Fin 1) => crossEntropy nClasses z label) (StableHLO.den ((StableHLO.SHlo.operand nlogN (cifarCnn8Forward W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ W₆ b₆ W₇ b₇ W₈ b₈ W₉ b₉ Wa ba Wb bb 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.