Documentation

LeanMlir.Proofs.Nets.Small.Cifar8BnParamGrad

The 8-conv CIFAR CNN with per-channel BN — every gradient node IS the loss's derivative #

cifar8Bn_train_step_tiedG states the 38 un-fused gradient nodes the packed cifar8w_bn_* arms emit, each at the cotangent the chain threads to it. cifar8Bn_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; cifar8Bn_net_lossGrad_CE instantiates it at the softmax cross-entropy the render emits.

The net is the BN-free 8-conv net (Cifar8TieG.cifar8_net_lossGrad) with a per-example, per-channel BN between each conv and its ReLU, so each pool's pre-activation is a BN output. Two cells equal at every conv weight stay equal through BN (one affine map per channel), so the twin relations (Cifar8BnPoolTwin1 … Cifar8BnPoolTwin4, cells equal at every weight upstream of the pool, BN γ/β included) and the selection routing carry over unchanged. Per node kind, BN adds bnGamma_hasGradAt, bnBeta_hasGradAt and the input pull-back hasGradAt_bnPC.

Hypotheses. Odd kernels, every BN ε > 0 (Cifar8BnPos), every ReLU off its kink, every pool window dead or tied only between twins, each selection naming a maximum of every window (Cifar8BnLossSmoothAt). Scope. One example (the emitted module batch-contracts; den is per-example; BN normalises each channel over the example's own spatial cells).

noncomputable def Proofs.Cifar8BnTieG.bnPCHasVJPAt (oc h w : ℕ) (ε : ℝ) (hε : 0 < ε) (γ β : Vec oc) (v : Vec (oc * h * w)) :
HasVJPAt (bnPerChannelTensor3 oc h w ε γ β) v

Per-channel BN's VJP at a point: the renderable backward bnPerChannelTensor3GradInput.

Equations
Instances For
    theorem Proofs.Cifar8BnTieG.hasGradAt_bnPC {oc h w : ℕ} (ε : ℝ) (hε : 0 < ε) (γ β : Vec oc) (v : Vec (oc * h * w)) {G : Vec (oc * h * w) → Vec 1} {dy : Vec (oc * h * w)} (hG : HasGradAt G (bnPerChannelTensor3 oc h w ε γ β v) dy) :
    HasGradAt (fun (y : Vec (oc * h * w)) => G (bnPerChannelTensor3 oc h w ε γ β y)) v (bnPerChannelTensor3GradInput oc h w ε γ v dy)

    Through per-channel BN: the backward is bnPerChannelTensor3GradInput, the BN-back the render emits.

    theorem Proofs.Cifar8BnTieG.bnGamma_hasGradAt {oc h w : ℕ} (vN epsStr cotN : String) (ε : ℝ) (γ β : Vec oc) (v : Vec (oc * h * w)) {G : Vec (oc * h * w) → Vec 1} {c : Vec (oc * h * w)} (hG : HasGradAt G (bnPerChannelTensor3 oc h w ε γ β v) c) :
    HasGradAt (fun (θ : Vec oc) => G (bnPerChannelTensor3 oc h w ε θ β v)) γ (StableHLO.den (StableHLO.SHlo.bnGammaGrad vN epsStr ε v (StableHLO.SHlo.operand cotN c)))

    BN γ node = ∇_γ G.

    theorem Proofs.Cifar8BnTieG.bnBeta_hasGradAt {oc h w : ℕ} (cotN : String) (ε : ℝ) (γ β : Vec oc) (v : Vec (oc * h * w)) {G : Vec (oc * h * w) → Vec 1} {c : Vec (oc * h * w)} (hG : HasGradAt G (bnPerChannelTensor3 oc h w ε γ β v) c) :
    HasGradAt (fun (θ : Vec oc) => G (bnPerChannelTensor3 oc h w ε γ θ v)) β (StableHLO.den (StableHLO.SHlo.operand cotN c).bnBetaGrad)

    BN β node = ∇_β G.

    theorem Proofs.Cifar8BnTieG.bnPC_gamma_continuous {oc h w : ℕ} (ε : ℝ) (β : Vec oc) (v : Vec (oc * h * w)) :
    Continuous fun (θ : Vec oc) => bnPerChannelTensor3 oc h w ε θ β v
    theorem Proofs.Cifar8BnTieG.bnPC_beta_continuous {oc h w : ℕ} (ε : ℝ) (γ : Vec oc) (v : Vec (oc * h * w)) :
    Continuous fun (θ : Vec oc) => bnPerChannelTensor3 oc h w ε γ θ v
    noncomputable def Proofs.Cifar8BnTieG.cifar8BnUp {c c' H W kH kW : ℕ} (Wc : Kernel4 c' c kH kW) (bc : Vec c') (εc : ℝ) (γc βc : Vec c') (Wd : Kernel4 c' c' kH kW) (bd : Vec c') (εd : ℝ) (γd βd : 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, BN, ReLU, conv, BN.

    Equations
    • One or more equations did not get rendered due to their size.
    Instances For
      theorem Proofs.Cifar8BnTieG.cifar8BnUp_continuous {c c' H W kH kW : ℕ} (Wc : Kernel4 c' c kH kW) (bc : Vec c') (εc : ℝ) (hεc : 0 < εc) (γc βc : Vec c') (Wd : Kernel4 c' c' kH kW) (bd : Vec c') (εd : ℝ) (hεd : 0 < εd) (γd βd : Vec c') :
      Continuous (cifar8BnUp Wc bc εc γc βc Wd bd εd γd βd)
      noncomputable def Proofs.Cifar8BnTieG.cifar8BnPre1 {ic c1 h w kH kW : ℕ} (W₁ : Kernel4 c1 ic kH kW) (b₁ : Vec c1) (ε₁ : ℝ) (γ₁ β₁ : Vec c1) (W₂ : Kernel4 c1 c1 kH kW) (b₂ : Vec c1) (ε₂ : ℝ) (γ₂ β₂ : Vec c1) (x : Vec (ic * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) :
      Vec (c1 * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))

      The first pool's pre-activation (BN₂'s output).

      Equations
      • One or more equations did not get rendered due to their size.
      Instances For
        noncomputable def Proofs.Cifar8BnTieG.cifar8BnPre2 {ic c1 c2 h w kH kW : ℕ} (W₁ : Kernel4 c1 ic kH kW) (b₁ : Vec c1) (ε₁ : ℝ) (γ₁ β₁ : Vec c1) (W₂ : Kernel4 c1 c1 kH kW) (b₂ : Vec c1) (ε₂ : ℝ) (γ₂ β₂ : Vec c1) (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (ε₃ : ℝ) (γ₃ β₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) (ε₄ : ℝ) (γ₄ β₄ : 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))))

        Pool 2's pre-activation (BN₄'s output).

        Equations
        • One or more equations did not get rendered due to their size.
        Instances For
          noncomputable def Proofs.Cifar8BnTieG.cifar8BnPre3 {ic c1 c2 c3 h w kH kW : ℕ} (W₁ : Kernel4 c1 ic kH kW) (b₁ : Vec c1) (ε₁ : ℝ) (γ₁ β₁ : Vec c1) (W₂ : Kernel4 c1 c1 kH kW) (b₂ : Vec c1) (ε₂ : ℝ) (γ₂ β₂ : Vec c1) (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (ε₃ : ℝ) (γ₃ β₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) (ε₄ : ℝ) (γ₄ β₄ : Vec c2) (W₅ : Kernel4 c3 c2 kH kW) (b₅ : Vec c3) (ε₅ : ℝ) (γ₅ β₅ : Vec c3) (W₆ : Kernel4 c3 c3 kH kW) (b₆ : Vec c3) (ε₆ : ℝ) (γ₆ β₆ : Vec c3) (x : Vec (ic * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) :
          Vec (c3 * (2 * (2 * h)) * (2 * (2 * w)))

          Pool 3's pre-activation (BN₆'s output).

          Equations
          • One or more equations did not get rendered due to their size.
          Instances For
            noncomputable def Proofs.Cifar8BnTieG.cifar8BnPre4 {ic c1 c2 c3 c4 h w kH kW : ℕ} (W₁ : Kernel4 c1 ic kH kW) (b₁ : Vec c1) (ε₁ : ℝ) (γ₁ β₁ : Vec c1) (W₂ : Kernel4 c1 c1 kH kW) (b₂ : Vec c1) (ε₂ : ℝ) (γ₂ β₂ : Vec c1) (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (ε₃ : ℝ) (γ₃ β₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) (ε₄ : ℝ) (γ₄ β₄ : Vec c2) (W₅ : Kernel4 c3 c2 kH kW) (b₅ : Vec c3) (ε₅ : ℝ) (γ₅ β₅ : Vec c3) (W₆ : Kernel4 c3 c3 kH kW) (b₆ : Vec c3) (ε₆ : ℝ) (γ₆ β₆ : Vec c3) (W₇ : Kernel4 c4 c3 kH kW) (b₇ : Vec c4) (ε₇ : ℝ) (γ₇ β₇ : Vec c4) (W₈ : Kernel4 c4 c4 kH kW) (b₈ : Vec c4) (ε₈ : ℝ) (γ₈ β₈ : Vec c4) (x : Vec (ic * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) :
            Vec (c4 * (2 * h) * (2 * w))

            Pool 4's pre-activation (BN₈'s output).

            Equations
            • One or more equations did not get rendered due to their size.
            Instances For
              def Proofs.Cifar8BnTieG.Cifar8BnPoolTwin1 {ic h w : ℕ} (c1 kH kW : ℕ) (ε₁ ε₂ : ℝ) (x : Vec (ic * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) (p q : Fin (2 * (2 * (2 * (2 * h)))) × Fin (2 * (2 * (2 * (2 * w))))) :

              Twins of pool 1: equal at every weight (conv and BN γ/β) upstream of it, in every channel.

              Equations
              • One or more equations did not get rendered due to their size.
              Instances For
                def Proofs.Cifar8BnTieG.Cifar8BnPoolTwin2 {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 pool 2: equal at every weight (conv and BN γ/β) upstream of it, in every channel.

                Equations
                • One or more equations did not get rendered due to their size.
                Instances For
                  def Proofs.Cifar8BnTieG.Cifar8BnPoolTwin3 {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 pool 3: equal at every weight (conv and BN γ/β) upstream of it, in every channel.

                  Equations
                  • One or more equations did not get rendered due to their size.
                  Instances For
                    def Proofs.Cifar8BnTieG.Cifar8BnPoolTwin4 {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 pool 4: equal at every weight (conv and BN γ/β) upstream of it, in every channel.

                    Equations
                    • One or more equations did not get rendered due to their size.
                    Instances For
                      structure Proofs.Cifar8BnTieG.Cifar8BnPos (ε₁ ε₂ ε₃ ε₄ ε₅ ε₆ ε₇ ε₈ : ℝ) :

                      Every BN ε is positive: BN is then differentiable everywhere.

                      • h1 : 0 < ε₁
                      • h2 : 0 < ε₂
                      • h3 : 0 < ε₃
                      • h4 : 0 < ε₄
                      • h5 : 0 < ε₅
                      • h6 : 0 < ε₆
                      • h7 : 0 < ε₇
                      • h8 : 0 < ε₈
                      Instances For
                        structure Proofs.Cifar8BnTieG.Cifar8BnLossSmoothAt {ic c1 c2 c3 c4 h w d1 kH kW : ℕ} (W₁ : Kernel4 c1 ic kH kW) (b₁ : Vec c1) (ε₁ : ℝ) (γ₁ β₁ : Vec c1) (W₂ : Kernel4 c1 c1 kH kW) (b₂ : Vec c1) (ε₂ : ℝ) (γ₂ β₂ : Vec c1) (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (ε₃ : ℝ) (γ₃ β₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) (ε₄ : ℝ) (γ₄ β₄ : Vec c2) (W₅ : Kernel4 c3 c2 kH kW) (b₅ : Vec c3) (ε₅ : ℝ) (γ₅ β₅ : Vec c3) (W₆ : Kernel4 c3 c3 kH kW) (b₆ : Vec c3) (ε₆ : ℝ) (γ₆ β₆ : Vec c3) (W₇ : Kernel4 c4 c3 kH kW) (b₇ : Vec c4) (ε₇ : ℝ) (γ₇ β₇ : Vec c4) (W₈ : Kernel4 c4 c4 kH kW) (b₈ : Vec c4) (ε₈ : ℝ) (γ₈ β₈ : 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 (at the BN outputs and the dense head); every window of each pool dead or tied only between that pool's twins; each selection naming a maximum of every window.

                        • z1 (k : Fin (c1 * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) : bnPerChannelTensor3 c1 (2 * (2 * (2 * (2 * h)))) (2 * (2 * (2 * (2 * w)))) ε₁ γ₁ β₁ (flatConv W₁ b₁ x) k ≠ 0
                        • z2 (k : Fin (c1 * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w)))))) : cifar8BnPre1 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ x k ≠ 0
                        • pool1 : SmallParamGrad.MaxPool2SmoothUpTo (Cifar8BnPoolTwin1 c1 kH kW ε₁ ε₂ x) (Tensor3.unflatten (cifar8BnPre1 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ x))
                        • sel1 : SmallParamGrad.PoolSelDom σ₁ (relu (c1 * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w))))) (cifar8BnPre1 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ x))
                        • z3 (k : Fin (c2 * (2 * (2 * (2 * h))) * (2 * (2 * (2 * w))))) : bnPerChannelTensor3 c2 (2 * (2 * (2 * h))) (2 * (2 * (2 * w))) ε₃ γ₃ β₃ (flatConv W₃ b₃ (maxPoolFlat c1 (2 * (2 * (2 * h))) (2 * (2 * (2 * w))) (relu (c1 * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w))))) (cifar8BnPre1 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ x)))) k ≠ 0
                        • z4 (k : Fin (c2 * (2 * (2 * (2 * h))) * (2 * (2 * (2 * w))))) : cifar8BnPre2 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ x k ≠ 0
                        • pool2 : SmallParamGrad.MaxPool2SmoothUpTo (Cifar8BnPoolTwin2 c1 c2 kH kW ε₁ ε₂ ε₃ ε₄ x) (Tensor3.unflatten (cifar8BnPre2 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ x))
                        • sel2 : SmallParamGrad.PoolSelDom σ₂ (relu (c2 * (2 * (2 * (2 * h))) * (2 * (2 * (2 * w)))) (cifar8BnPre2 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ x))
                        • z5 (k : Fin (c3 * (2 * (2 * h)) * (2 * (2 * w)))) : bnPerChannelTensor3 c3 (2 * (2 * h)) (2 * (2 * w)) ε₅ γ₅ β₅ (flatConv W₅ b₅ (maxPoolFlat c2 (2 * (2 * h)) (2 * (2 * w)) (relu (c2 * (2 * (2 * (2 * h))) * (2 * (2 * (2 * w)))) (cifar8BnPre2 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ x)))) k ≠ 0
                        • z6 (k : Fin (c3 * (2 * (2 * h)) * (2 * (2 * w)))) : cifar8BnPre3 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ x k ≠ 0
                        • pool3 : SmallParamGrad.MaxPool2SmoothUpTo (Cifar8BnPoolTwin3 c1 c2 c3 kH kW ε₁ ε₂ ε₃ ε₄ ε₅ ε₆ x) (Tensor3.unflatten (cifar8BnPre3 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ x))
                        • sel3 : SmallParamGrad.PoolSelDom σ₃ (relu (c3 * (2 * (2 * h)) * (2 * (2 * w))) (cifar8BnPre3 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ x))
                        • z7 (k : Fin (c4 * (2 * h) * (2 * w))) : bnPerChannelTensor3 c4 (2 * h) (2 * w) ε₇ γ₇ β₇ (flatConv W₇ b₇ (maxPoolFlat c3 (2 * h) (2 * w) (relu (c3 * (2 * (2 * h)) * (2 * (2 * w))) (cifar8BnPre3 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ x)))) k ≠ 0
                        • z8 (k : Fin (c4 * (2 * h) * (2 * w))) : cifar8BnPre4 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ W₇ b₇ ε₇ γ₇ β₇ W₈ b₈ ε₈ γ₈ β₈ x k ≠ 0
                        • pool4 : SmallParamGrad.MaxPool2SmoothUpTo (Cifar8BnPoolTwin4 c1 c2 c3 c4 kH kW ε₁ ε₂ ε₃ ε₄ ε₅ ε₆ ε₇ ε₈ x) (Tensor3.unflatten (cifar8BnPre4 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ W₇ b₇ ε₇ γ₇ β₇ W₈ b₈ ε₈ γ₈ β₈ x))
                        • sel4 : SmallParamGrad.PoolSelDom σ₄ (relu (c4 * (2 * h) * (2 * w)) (cifar8BnPre4 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ W₇ b₇ ε₇ γ₇ β₇ W₈ b₈ ε₈ γ₈ β₈ x))
                        • z9 (k : Fin d1) : dense W₉ b₉ (maxPoolFlat c4 h w (relu (c4 * (2 * h) * (2 * w)) (cifar8BnPre4 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ W₇ b₇ ε₇ γ₇ β₇ W₈ b₈ ε₈ γ₈ β₈ x))) k ≠ 0
                        • za (k : Fin d1) : dense Wa ba (relu d1 (dense W₉ b₉ (maxPoolFlat c4 h w (relu (c4 * (2 * h) * (2 * w)) (cifar8BnPre4 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ W₇ b₇ ε₇ γ₇ β₇ W₈ b₈ ε₈ γ₈ β₈ x))))) k ≠ 0
                        Instances For
                          def Proofs.Cifar8BnTieG.Cifar8BnNetLossTied {ic c1 c2 c3 c4 h w d1 nClasses kH kW : ℕ} (xN vN epsStr cotN : String) (W₁ : Kernel4 c1 ic kH kW) (b₁ : Vec c1) (ε₁ : ℝ) (γ₁ β₁ : Vec c1) (W₂ : Kernel4 c1 c1 kH kW) (b₂ : Vec c1) (ε₂ : ℝ) (γ₂ β₂ : Vec c1) (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (ε₃ : ℝ) (γ₃ β₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) (ε₄ : ℝ) (γ₄ β₄ : Vec c2) (W₅ : Kernel4 c3 c2 kH kW) (b₅ : Vec c3) (ε₅ : ℝ) (γ₅ β₅ : Vec c3) (W₆ : Kernel4 c3 c3 kH kW) (b₆ : Vec c3) (ε₆ : ℝ) (γ₆ β₆ : Vec c3) (W₇ : Kernel4 c4 c3 kH kW) (b₇ : Vec c4) (ε₇ : ℝ) (γ₇ β₇ : Vec c4) (W₈ : Kernel4 c4 c4 kH kW) (b₈ : Vec c4) (ε₈ : ℝ) (γ₈ β₈ : 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-bn gradient node is the gradient of L in that parameter: the 38 un-fused nodes cifar8Bn_train_step_tiedG states, each at the cotangent the chain threads to its layer (each pool routed at its selection), stated against L of cifarCnnBn8Forward with that one parameter varied (F is L of the forward at the given weights, the εs fixed).

                          Equations
                          • One or more equations did not get rendered due to their size.
                          Instances For
                            theorem Proofs.Cifar8BnTieG.cifar8Bn_net_lossGrad {ic c1 c2 c3 c4 h w d1 nClasses kH kW : ℕ} (xN vN epsStr 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) (ε₁ : ℝ) (γ₁ β₁ : Vec c1) (W₂ : Kernel4 c1 c1 kH kW) (b₂ : Vec c1) (ε₂ : ℝ) (γ₂ β₂ : Vec c1) (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (ε₃ : ℝ) (γ₃ β₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) (ε₄ : ℝ) (γ₄ β₄ : Vec c2) (W₅ : Kernel4 c3 c2 kH kW) (b₅ : Vec c3) (ε₅ : ℝ) (γ₅ β₅ : Vec c3) (W₆ : Kernel4 c3 c3 kH kW) (b₆ : Vec c3) (ε₆ : ℝ) (γ₆ β₆ : Vec c3) (W₇ : Kernel4 c4 c3 kH kW) (b₇ : Vec c4) (ε₇ : ℝ) (γ₇ β₇ : Vec c4) (W₈ : Kernel4 c4 c4 kH kW) (b₈ : Vec c4) (ε₈ : ℝ) (γ₈ β₈ : 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) (hq : Cifar8BnPos ε₁ ε₂ ε₃ ε₄ ε₅ ε₆ ε₇ ε₈) (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 : Cifar8BnLossSmoothAt 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 (cifarCnnBn8Forward W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ W₇ b₇ ε₇ γ₇ β₇ W₈ b₈ ε₈ γ₈ β₈ W₉ b₉ Wa ba Wb bb x) g) :
                            Cifar8BnNetLossTied xN vN epsStr 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-bn gradient node is the gradient of L in that parameter, whenever g is L's gradient at the logits.

                            Hypotheses: odd kernels, every BN ε positive (Cifar8BnPos), and Cifar8BnLossSmoothAt — 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.Cifar8BnTieG.cifar8Bn_net_lossGrad_CE {ic c1 c2 c3 c4 h w d1 nClasses kH kW : ℕ} (xN vN epsStr 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) (ε₁ : ℝ) (γ₁ β₁ : Vec c1) (W₂ : Kernel4 c1 c1 kH kW) (b₂ : Vec c1) (ε₂ : ℝ) (γ₂ β₂ : Vec c1) (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (ε₃ : ℝ) (γ₃ β₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) (ε₄ : ℝ) (γ₄ β₄ : Vec c2) (W₅ : Kernel4 c3 c2 kH kW) (b₅ : Vec c3) (ε₅ : ℝ) (γ₅ β₅ : Vec c3) (W₆ : Kernel4 c3 c3 kH kW) (b₆ : Vec c3) (ε₆ : ℝ) (γ₆ β₆ : Vec c3) (W₇ : Kernel4 c4 c3 kH kW) (b₇ : Vec c4) (ε₇ : ℝ) (γ₇ β₇ : Vec c4) (W₈ : Kernel4 c4 c4 kH kW) (b₈ : Vec c4) (ε₈ : ℝ) (γ₈ β₈ : 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) (hq : Cifar8BnPos ε₁ ε₂ ε₃ ε₄ ε₅ ε₆ ε₇ ε₈) (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 : Cifar8BnLossSmoothAt W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ W₇ b₇ ε₇ γ₇ β₇ W₈ b₈ ε₈ γ₈ β₈ W₉ b₉ Wa ba x σ₁ σ₂ σ₃ σ₄) (label : Fin nClasses) :
                            Cifar8BnNetLossTied xN vN epsStr 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 (cifarCnnBn8Forward 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.