Documentation

LeanMlir.Proofs.Foundation.GradNodesBAt

The tie predicates at either precision — ConvWTiedBAt bf16, ConvWSyncAt bf16, … #

GradNodesB states each parameter gradient node against the certified Σ_n gradient (ConvWTiedB, …), ParamGradNodes makes it a loss derivative (convW_hasGradAt, …) and SyncKit ties the all-reduced collective to the batch-R·N node (ConvWSync, …) — all at the f32 constructor. The ImageNet renders the book reports train from are bf16, and a bf16 render emits the *GradBBf16 kind at every conv weight gradient. This file states the same three things with the node chosen by the renderers' own switch (StableHLO.PrecisionSwitch: convWeightGradBAt bf16 id … is the f32 kind at false, the bf16 kind at true), and proves each for either value from the f32 lemma and Bf16Erasure (den_convWeightGradBAt_id): at the identity rounding the bf16 node denotes what the f32 node does, so a tie stated on the switch reads the bf16 artifact's text over ℝ exactly as the f32 tie reads the f32 artifact's.

The kinds here are the ones the ResNet and MobileNet / EfficientNet ties emit — conv, convStrided, convStridedXla (the TF-origin stems), depthwise, depthwiseStrided (B0's and MobileNetV4's downsampling depthwise) and depthwiseStridedXla (MobileNetV2's); the remaining kinds (stride-4, row-dense, patch-embed) follow the same pattern as ConvNeXt's and ViT's ties move onto the switch (planning/bf16_tie.md §3.3). At false each predicate is its f32 original by rfl (convWTiedBAt_false, …), except DepthwiseStridedXlaWTiedBAt, whose f32 form MobileNetV2StepTieB stated inline rather than as a GradNodesB predicate.

Nothing here is about the size of the rounding: rnd is id throughout, as in the renderers.

def Proofs.GradNodeB.ConvWTiedBAt (bf16 : Bool) (N h w : ℕ) {ic oc kH kW : ℕ} (xN cotN : String) (b : Vec oc) (x : Vec (N * (ic * h * w))) (W : Kernel4 oc ic kH kW) (cot : Vec (N * (oc * h * w))) :

ConvWTiedB with the node chosen by the switch: the conv weight gradient node, at either precision, denotes the certified batched Σ_n gradient.

Equations
  • One or more equations did not get rendered due to their size.
Instances For
    theorem Proofs.GradNodeB.convWTiedBAt_false (N h w : ℕ) {ic oc kH kW : ℕ} (xN cotN : String) (b : Vec oc) (x : Vec (N * (ic * h * w))) (W : Kernel4 oc ic kH kW) (cot : Vec (N * (oc * h * w))) :
    ConvWTiedBAt false N h w xN cotN b x W cot = ConvWTiedB N h w xN cotN b x W cot
    theorem Proofs.GradNodeB.convWTiedBAt_holds (bf16 : Bool) {N h w ic oc kH kW : ℕ} {xN cotN : String} {b : Vec oc} {x : Vec (N * (ic * h * w))} {W : Kernel4 oc ic kH kW} {cot : Vec (N * (oc * h * w))} :
    ConvWTiedBAt bf16 N h w xN cotN b x W cot
    def Proofs.GradNodeB.ConvStridedWTiedBAt (bf16 : Bool) (N h w : ℕ) {ic oc kH kW : ℕ} (xN cotN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (cot : Vec (N * (oc * h * w))) :

    ConvStridedWTiedB on the switch (symmetric padding, ResNet's stride-2 sites).

    Equations
    • One or more equations did not get rendered due to their size.
    Instances For
      theorem Proofs.GradNodeB.convStridedWTiedBAt_false (N h w : ℕ) {ic oc kH kW : ℕ} (xN cotN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (cot : Vec (N * (oc * h * w))) :
      ConvStridedWTiedBAt false N h w xN cotN b x W cot = ConvStridedWTiedB N h w xN cotN b x W cot
      theorem Proofs.GradNodeB.convStridedWTiedBAt_holds (bf16 : Bool) {N h w ic oc kH kW : ℕ} {xN cotN : String} {b : Vec oc} {x : Vec (N * (ic * (2 * h) * (2 * w)))} {W : Kernel4 oc ic kH kW} {cot : Vec (N * (oc * h * w))} :
      ConvStridedWTiedBAt bf16 N h w xN cotN b x W cot
      def Proofs.GradNodeB.ConvStridedXlaWTiedBAt (bf16 : Bool) (N h w : ℕ) {ic oc kH kW : ℕ} (xN cotN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (cot : Vec (N * (oc * h * w))) :

      ConvStridedXlaWTiedB on the switch (XLA-SAME padding: MobileNetV2's and EfficientNet-B0's stems). Same type and emitted shape as ConvStridedWTiedBAt; only the certificate differs.

      Equations
      • One or more equations did not get rendered due to their size.
      Instances For
        theorem Proofs.GradNodeB.convStridedXlaWTiedBAt_false (N h w : ℕ) {ic oc kH kW : ℕ} (xN cotN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (cot : Vec (N * (oc * h * w))) :
        ConvStridedXlaWTiedBAt false N h w xN cotN b x W cot = ConvStridedXlaWTiedB N h w xN cotN b x W cot
        theorem Proofs.GradNodeB.convStridedXlaWTiedBAt_holds (bf16 : Bool) {N h w ic oc kH kW : ℕ} {xN cotN : String} {b : Vec oc} {x : Vec (N * (ic * (2 * h) * (2 * w)))} {W : Kernel4 oc ic kH kW} {cot : Vec (N * (oc * h * w))} :
        ConvStridedXlaWTiedBAt bf16 N h w xN cotN b x W cot
        def Proofs.GradNodeB.DepthwiseWTiedBAt (bf16 : Bool) (N h w : ℕ) {c kH kW : ℕ} (xN cotN : String) (b : Vec c) (x : Vec (N * (c * h * w))) (W : DepthwiseKernel c kH kW) (cot : Vec (N * (c * h * w))) :

        DepthwiseWTiedB on the switch (the stride-1 depthwise of every inverted-residual block).

        Equations
        • One or more equations did not get rendered due to their size.
        Instances For
          theorem Proofs.GradNodeB.depthwiseWTiedBAt_false (N h w : ℕ) {c kH kW : ℕ} (xN cotN : String) (b : Vec c) (x : Vec (N * (c * h * w))) (W : DepthwiseKernel c kH kW) (cot : Vec (N * (c * h * w))) :
          DepthwiseWTiedBAt false N h w xN cotN b x W cot = DepthwiseWTiedB N h w xN cotN b x W cot
          theorem Proofs.GradNodeB.depthwiseWTiedBAt_holds (bf16 : Bool) {N h w c kH kW : ℕ} {xN cotN : String} {b : Vec c} {x : Vec (N * (c * h * w))} {W : DepthwiseKernel c kH kW} {cot : Vec (N * (c * h * w))} :
          DepthwiseWTiedBAt bf16 N h w xN cotN b x W cot
          def Proofs.GradNodeB.DepthwiseStridedWTiedBAt (bf16 : Bool) (N h w : ℕ) {c kH kW : ℕ} (xN cotN : String) (b : Vec c) (x : Vec (N * (c * (2 * h) * (2 * w)))) (W : DepthwiseKernel c kH kW) (cot : Vec (N * (c * h * w))) :

          DepthwiseStridedWTiedB on the switch (symmetric padding: EfficientNet-B0's and MobileNetV4's downsampling depthwise).

          Equations
          • One or more equations did not get rendered due to their size.
          Instances For
            theorem Proofs.GradNodeB.depthwiseStridedWTiedBAt_false (N h w : ℕ) {c kH kW : ℕ} (xN cotN : String) (b : Vec c) (x : Vec (N * (c * (2 * h) * (2 * w)))) (W : DepthwiseKernel c kH kW) (cot : Vec (N * (c * h * w))) :
            DepthwiseStridedWTiedBAt false N h w xN cotN b x W cot = DepthwiseStridedWTiedB N h w xN cotN b x W cot
            theorem Proofs.GradNodeB.depthwiseStridedWTiedBAt_holds (bf16 : Bool) {N h w c kH kW : ℕ} {xN cotN : String} {b : Vec c} {x : Vec (N * (c * (2 * h) * (2 * w)))} {W : DepthwiseKernel c kH kW} {cot : Vec (N * (c * h * w))} :
            DepthwiseStridedWTiedBAt bf16 N h w xN cotN b x W cot
            def Proofs.GradNodeB.DepthwiseStridedXlaWTiedBAt (bf16 : Bool) (N h w : ℕ) {c kH kW : ℕ} (xN cotN : String) (b : Vec c) (x : Vec (N * (c * (2 * h) * (2 * w)))) (W : DepthwiseKernel c kH kW) (cot : Vec (N * (c * h * w))) :

            The XLA-SAME strided depthwise weight node on the switch (MobileNetV2's four stride-2 depthwises, b2 / b4 / b7 / b14). GradNodesB has no f32 predicate for this kind — MobileNetV2StepTieB stated the node inline — so false here IS that statement, and the proof is depthwiseStridedXlaWGradB_den under the erasure. Not B0's symmetric DepthwiseStridedWTiedBAt: identical types, different certificates.

            Equations
            • One or more equations did not get rendered due to their size.
            Instances For
              theorem Proofs.GradNodeB.depthwiseStridedXlaWTiedBAt_holds (bf16 : Bool) {N h w c kH kW : ℕ} {xN cotN : String} {b : Vec c} {x : Vec (N * (c * (2 * h) * (2 * w)))} {W : DepthwiseKernel c kH kW} {cot : Vec (N * (c * h * w))} :
              DepthwiseStridedXlaWTiedBAt bf16 N h w xN cotN b x W cot
              theorem Proofs.GradNodeB.convWAt_hasGradAt (bf16 : Bool) {N ic oc h w kH kW : ℕ} (xN cotN : String) (b : Vec oc) (x : Vec (N * (ic * h * w))) (W : Kernel4 oc ic kH kW) {G : Vec (N * (oc * h * w)) → Vec 1} {cot : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (flatConv W b) x) cot) :
              HasGradAt (fun (θ : Vec (oc * ic * kH * kW)) => G (StableHLO.batchMap N (flatConv (Kernel4.unflatten θ) b) x)) W.flatten (StableHLO.den (StableHLO.SHlo.convWeightGradBAt bf16 id xN b x W (StableHLO.SHlo.operand cotN cot)))

              convW_hasGradAt on the switch: the conv weight gradient node, at either precision, is the gradient of G after the conv in the weight.

              theorem Proofs.GradNodeB.convStridedWAt_hasGradAt (bf16 : Bool) {N ic oc h w kH kW : ℕ} (xN cotN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) {G : Vec (N * (oc * h * w)) → Vec 1} {cot : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (flatConvStride2 W b) x) cot) :

              convStridedW_hasGradAt on the switch.

              theorem Proofs.GradNodeB.convStridedXlaWAt_hasGradAt (bf16 : Bool) {N ic oc h w kH kW : ℕ} (xN cotN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) {G : Vec (N * (oc * h * w)) → Vec 1} {cot : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (flatConvStride2Xla W b) x) cot) :

              convStridedXlaW_hasGradAt on the switch.

              theorem Proofs.GradNodeB.depthwiseWAt_hasGradAt (bf16 : Bool) {N c h w kH kW : ℕ} (xN cotN : String) (b : Vec c) (x : Vec (N * (c * h * w))) (W : DepthwiseKernel c kH kW) {G : Vec (N * (c * h * w)) → Vec 1} {cot : Vec (N * (c * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (depthwiseFlat W b) x) cot) :

              depthwiseW_hasGradAt on the switch.

              theorem Proofs.GradNodeB.depthwiseStridedWAt_hasGradAt (bf16 : Bool) {N c h w kH kW : ℕ} (xN cotN : String) (b : Vec c) (x : Vec (N * (c * (2 * h) * (2 * w)))) (W : DepthwiseKernel c kH kW) {G : Vec (N * (c * h * w)) → Vec 1} {cot : Vec (N * (c * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (depthwiseStride2Flat W b) x) cot) :

              depthwiseStridedW_hasGradAt on the switch.

              theorem Proofs.GradNodeB.depthwiseStridedXlaWAt_hasGradAt (bf16 : Bool) {N c h w kH kW : ℕ} (xN cotN : String) (b : Vec c) (x : Vec (N * (c * (2 * h) * (2 * w)))) (W : DepthwiseKernel c kH kW) {G : Vec (N * (c * h * w)) → Vec 1} {cot : Vec (N * (c * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (depthwiseStride2FlatXla W b) x) cot) :

              depthwiseStridedXlaW_hasGradAt on the switch.

              def Proofs.SyncKit.ConvWSyncAt (bf16 : Bool) (R : ℕ) (hR : 0 < R) (N h w : ℕ) {ic oc kH kW : ℕ} (t xN cotN : String) (b : Vec oc) (X : Vec (R * N * (ic * h * w))) (W : Kernel4 oc ic kH kW) (cots : Fin R → Vec (N * (oc * h * w))) (COT : Vec (R * N * (oc * h * w))) :

              ConvWSync on the switch: the all-reduced mean of the replicas' conv weight gradient nodes, at either precision, is the batch-R·N node of the same kind.

              Equations
              • One or more equations did not get rendered due to their size.
              Instances For
                theorem Proofs.SyncKit.convWSyncAt_false (R : ℕ) (hR : 0 < R) (N h w : ℕ) {ic oc kH kW : ℕ} (t xN cotN : String) (b : Vec oc) (X : Vec (R * N * (ic * h * w))) (W : Kernel4 oc ic kH kW) (cots : Fin R → Vec (N * (oc * h * w))) (COT : Vec (R * N * (oc * h * w))) :
                ConvWSyncAt false R hR N h w t xN cotN b X W cots COT = ConvWSync R hR N h w t xN cotN b X W cots COT
                theorem Proofs.SyncKit.convWSyncAt_of_scaled (bf16 : Bool) (R : ℕ) (hR : 0 < R) (N h w : ℕ) {ic oc kH kW : ℕ} (t xN cotN : String) (b : Vec oc) (X : Vec (R * N * (ic * h * w))) (W : Kernel4 oc ic kH kW) (cots : Fin R → Vec (N * (oc * h * w))) (COT : Vec (R * N * (oc * h * w))) (hc : ∀ (r : Fin R), cots r = batchShard R N (oc * h * w) (fun (i : Fin (R * N * (oc * h * w))) => ↑R * COT i) r) :
                ConvWSyncAt bf16 R hR N h w t xN cotN b X W cots COT
                def Proofs.SyncKit.ConvStridedWSyncAt (bf16 : Bool) (R : ℕ) (hR : 0 < R) (N h w : ℕ) {ic oc kH kW : ℕ} (t xN cotN : String) (b : Vec oc) (X : Vec (R * N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (cots : Fin R → Vec (N * (oc * h * w))) (COT : Vec (R * N * (oc * h * w))) :

                ConvStridedWSync on the switch.

                Equations
                • One or more equations did not get rendered due to their size.
                Instances For
                  theorem Proofs.SyncKit.convStridedWSyncAt_false (R : ℕ) (hR : 0 < R) (N h w : ℕ) {ic oc kH kW : ℕ} (t xN cotN : String) (b : Vec oc) (X : Vec (R * N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (cots : Fin R → Vec (N * (oc * h * w))) (COT : Vec (R * N * (oc * h * w))) :
                  ConvStridedWSyncAt false R hR N h w t xN cotN b X W cots COT = ConvStridedWSync R hR N h w t xN cotN b X W cots COT
                  theorem Proofs.SyncKit.convStridedWSyncAt_of_scaled (bf16 : Bool) (R : ℕ) (hR : 0 < R) (N h w : ℕ) {ic oc kH kW : ℕ} (t xN cotN : String) (b : Vec oc) (X : Vec (R * N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (cots : Fin R → Vec (N * (oc * h * w))) (COT : Vec (R * N * (oc * h * w))) (hc : ∀ (r : Fin R), cots r = batchShard R N (oc * h * w) (fun (i : Fin (R * N * (oc * h * w))) => ↑R * COT i) r) :
                  ConvStridedWSyncAt bf16 R hR N h w t xN cotN b X W cots COT
                  def Proofs.SyncKit.ConvStridedXlaWSyncAt (bf16 : Bool) (R : ℕ) (hR : 0 < R) (N h w : ℕ) {ic oc kH kW : ℕ} (t xN cotN : String) (b : Vec oc) (X : Vec (R * N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (cots : Fin R → Vec (N * (oc * h * w))) (COT : Vec (R * N * (oc * h * w))) :

                  ConvStridedXlaWSync on the switch (the TF-origin stems).

                  Equations
                  • One or more equations did not get rendered due to their size.
                  Instances For
                    theorem Proofs.SyncKit.convStridedXlaWSyncAt_false (R : ℕ) (hR : 0 < R) (N h w : ℕ) {ic oc kH kW : ℕ} (t xN cotN : String) (b : Vec oc) (X : Vec (R * N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (cots : Fin R → Vec (N * (oc * h * w))) (COT : Vec (R * N * (oc * h * w))) :
                    ConvStridedXlaWSyncAt false R hR N h w t xN cotN b X W cots COT = ConvStridedXlaWSync R hR N h w t xN cotN b X W cots COT
                    theorem Proofs.SyncKit.convStridedXlaWSyncAt_of_scaled (bf16 : Bool) (R : ℕ) (hR : 0 < R) (N h w : ℕ) {ic oc kH kW : ℕ} (t xN cotN : String) (b : Vec oc) (X : Vec (R * N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (cots : Fin R → Vec (N * (oc * h * w))) (COT : Vec (R * N * (oc * h * w))) (hc : ∀ (r : Fin R), cots r = batchShard R N (oc * h * w) (fun (i : Fin (R * N * (oc * h * w))) => ↑R * COT i) r) :
                    ConvStridedXlaWSyncAt bf16 R hR N h w t xN cotN b X W cots COT
                    def Proofs.SyncKit.DepthwiseWSyncAt (bf16 : Bool) (R : ℕ) (hR : 0 < R) (N h w : ℕ) {c kH kW : ℕ} (t xN cotN : String) (b : Vec c) (X : Vec (R * N * (c * h * w))) (W : DepthwiseKernel c kH kW) (cots : Fin R → Vec (N * (c * h * w))) (COT : Vec (R * N * (c * h * w))) :

                    DepthwiseWSync on the switch: the all-reduced mean of the replicas' depthwise weight gradient nodes, at either precision, is the batch-R·N node of the same kind.

                    Equations
                    • One or more equations did not get rendered due to their size.
                    Instances For
                      theorem Proofs.SyncKit.depthwiseWSyncAt_false (R : ℕ) (hR : 0 < R) (N h w : ℕ) {c kH kW : ℕ} (t xN cotN : String) (b : Vec c) (X : Vec (R * N * (c * h * w))) (W : DepthwiseKernel c kH kW) (cots : Fin R → Vec (N * (c * h * w))) (COT : Vec (R * N * (c * h * w))) :
                      DepthwiseWSyncAt false R hR N h w t xN cotN b X W cots COT = DepthwiseWSync R hR N h w t xN cotN b X W cots COT
                      theorem Proofs.SyncKit.depthwiseWSyncAt_of_scaled (bf16 : Bool) (R : ℕ) (hR : 0 < R) (N h w : ℕ) {c kH kW : ℕ} (t xN cotN : String) (b : Vec c) (X : Vec (R * N * (c * h * w))) (W : DepthwiseKernel c kH kW) (cots : Fin R → Vec (N * (c * h * w))) (COT : Vec (R * N * (c * h * w))) (hc : ∀ (r : Fin R), cots r = batchShard R N (c * h * w) (fun (i : Fin (R * N * (c * h * w))) => ↑R * COT i) r) :
                      DepthwiseWSyncAt bf16 R hR N h w t xN cotN b X W cots COT
                      def Proofs.SyncKit.DepthwiseStridedWSyncAt (bf16 : Bool) (R : ℕ) (hR : 0 < R) (N h w : ℕ) {c kH kW : ℕ} (t xN cotN : String) (b : Vec c) (X : Vec (R * N * (c * (2 * h) * (2 * w)))) (W : DepthwiseKernel c kH kW) (cots : Fin R → Vec (N * (c * h * w))) (COT : Vec (R * N * (c * h * w))) :

                      DepthwiseStridedWSync on the switch (symmetric padding).

                      Equations
                      • One or more equations did not get rendered due to their size.
                      Instances For
                        theorem Proofs.SyncKit.depthwiseStridedWSyncAt_false (R : ℕ) (hR : 0 < R) (N h w : ℕ) {c kH kW : ℕ} (t xN cotN : String) (b : Vec c) (X : Vec (R * N * (c * (2 * h) * (2 * w)))) (W : DepthwiseKernel c kH kW) (cots : Fin R → Vec (N * (c * h * w))) (COT : Vec (R * N * (c * h * w))) :
                        DepthwiseStridedWSyncAt false R hR N h w t xN cotN b X W cots COT = DepthwiseStridedWSync R hR N h w t xN cotN b X W cots COT
                        theorem Proofs.SyncKit.depthwiseStridedWSyncAt_of_scaled (bf16 : Bool) (R : ℕ) (hR : 0 < R) (N h w : ℕ) {c kH kW : ℕ} (t xN cotN : String) (b : Vec c) (X : Vec (R * N * (c * (2 * h) * (2 * w)))) (W : DepthwiseKernel c kH kW) (cots : Fin R → Vec (N * (c * h * w))) (COT : Vec (R * N * (c * h * w))) (hc : ∀ (r : Fin R), cots r = batchShard R N (c * h * w) (fun (i : Fin (R * N * (c * h * w))) => ↑R * COT i) r) :
                        DepthwiseStridedWSyncAt bf16 R hR N h w t xN cotN b X W cots COT