Documentation

LeanMlir.Proofs.Nets.ConvNeXt.ConvNeXt

ConvNeXt #

A representative ConvNeXt block and a small end-to-end ConvNeXt VJP in flattened Vec space — the ConvNeXt analogue of cnn_has_vjp_at.

ConvNeXt is the "ResNet, modernized" architecture: it keeps the residual skeleton but swaps in the ViT-era ingredients — a large-kernel (7×7) depthwise conv for spatial mixing, LayerNorm instead of BatchNorm, GELU instead of ReLU, an inverted-bottleneck MLP (1×1 expand → GELU → 1×1 project), and a learnable per-channel layer scale. Crucially, the block has no non-smooth activation (GELU is smooth everywhere) and no max-pool, so — unlike cnn_has_vjp_at — the entire VJP composes with no kink hypotheses: the only side conditions are the LayerNorm positivity arguments 0 < ε.

This file contributes three genuinely new pieces:

LayerNorm representation caveat #

The layerNormForward reused here is the proof's Vec→Vec LayerNorm with scalar γ, β that normalizes over the whole flattened vector, whereas true ConvNeXt LayerNorm is per-spatial-position over the channel axis (LayerNorm over NCHW's C). This is the same representation simplification the audit flagged for the LN family; a faithful channel-LN-over-NCHW lift is a follow-up. Every other piece (depthwise 7×7, 1×1 convs, GELU, layer scale, GAP, dense) is exact.

noncomputable def Proofs.layerScale {n : } (γ x : Vec n) :
Vec n

Layer scale — per-channel learnable elementwise multiply by γ. layerScale γ x i = γ i * x i. A diagonal linear map.

Equations
Instances For

    layerScale γ is differentiable everywhere (diagonal linear).

    theorem Proofs.pdiv_layerScale {n : } (γ x : Vec n) (i j : Fin n) :
    pdiv (layerScale γ) x i j = if i = j then γ i else 0

    Jacobian of layerScale∂(γ_j x_j)/∂x_i = γ_i δ_{ij}.

    noncomputable def Proofs.layerScale_has_vjp {n : } (γ : Vec n) :

    Layer scale VJP: back(x, dy)_i = γ i * dy i.

    Equations
    Instances For
      theorem Proofs.layerScale_has_vjp_correct {n : } (γ x dy : Vec n) (i : Fin n) :
      (layerScale_has_vjp γ).backward x dy i = j : Fin n, pdiv (layerScale γ) x i j * dy j
      noncomputable def Proofs.convNextBlockBody {c cExp h w kH kW : } (Wdw : DepthwiseKernel c kH kW) (bdw : Vec c) (εn γn βn : ) (Wex : Kernel4 cExp c 1 1) (bex : Vec cExp) (Wpr : Kernel4 c cExp 1 1) (bpr : Vec c) (γls : Vec (c * h * w)) :
      Vec (c * h * w)Vec (c * h * w)

      ConvNeXt block body (Vec→Vec, no skip):

      layerScale γ ∘ project(1×1) ∘ gelu ∘ expand(1×1) ∘ layerNorm ∘ depthwise(7×7)

      Channel/spatial dims are generic Nat params; the depthwise kernel is c × kH × kW (the 7×7 in ConvNeXt), the expand conv lifts c → cExp channels (the usual 4× inverted-bottleneck), gelu is applied on the expanded activation, the project conv brings cExp → c back, and the per-channel layer scale closes the block.

      LN representation caveat. The layerNormForward used here is the proof's Vec→Vec LayerNorm with scalar γ_n, β_n that normalizes over the whole flattened c·h·w vector, whereas true ConvNeXt LN is per-spatial-position over the channel axis (LayerNorm over NCHW's C). This is the same representation simplification the audit flagged for the LN family; a faithful channel-LN-over-NCHW is a follow-up. Every other piece is exact. Because gelu is smooth everywhere, LN is smooth given ε>0, and conv/layerScale are linear, the whole body is differentiable everywhere — no ReLU-style kink hypotheses needed.

      Equations
      • One or more equations did not get rendered due to their size.
      Instances For
        theorem Proofs.convNextBlockBody_differentiable {c cExp h w kH kW : } (Wdw : DepthwiseKernel c kH kW) (bdw : Vec c) (εn : ) (hεn : 0 < εn) (γn βn : ) (Wex : Kernel4 cExp c 1 1) (bex : Vec cExp) (Wpr : Kernel4 c cExp 1 1) (bpr : Vec c) (γls : Vec (c * h * w)) :
        Differentiable (convNextBlockBody Wdw bdw εn γn βn Wex bex Wpr bpr γls)

        The block body is differentiable everywhere (composition of everywhere-differentiable maps).

        noncomputable def Proofs.convNextBlockBody_has_vjp {c cExp h w kH kW : } (Wdw : DepthwiseKernel c kH kW) (bdw : Vec c) (εn : ) (hεn : 0 < εn) (γn βn : ) (Wex : Kernel4 cExp c 1 1) (bex : Vec cExp) (Wpr : Kernel4 c cExp 1 1) (bpr : Vec c) (γls : Vec (c * h * w)) :
        HasVJP (convNextBlockBody Wdw bdw εn γn βn Wex bex Wpr bpr γls)

        ConvNeXt block body VJP (global) — built by chaining the everywhere-differentiable piece VJPs through vjp_comp. Needs only 0 < εn (the LayerNorm positivity); no kink hypotheses since gelu is smooth and the rest are linear. Because the body is differentiable everywhere, the VJP is global (HasVJP), not pointwise.

        Equations
        • One or more equations did not get rendered due to their size.
        Instances For
          noncomputable def Proofs.convNextBlockBody_has_vjp_at {c cExp h w kH kW : } (Wdw : DepthwiseKernel c kH kW) (bdw : Vec c) (εn : ) (hεn : 0 < εn) (γn βn : ) (Wex : Kernel4 cExp c 1 1) (bex : Vec cExp) (Wpr : Kernel4 c cExp 1 1) (bpr : Vec c) (γls v : Vec (c * h * w)) :
          HasVJPAt (convNextBlockBody Wdw bdw εn γn βn Wex bex Wpr bpr γls) v

          ConvNeXt block body VJP at a point — the global witness restricted to a point. Kept for downstream _at consumers.

          Equations
          Instances For
            noncomputable def Proofs.convNextBlock {c cExp h w kH kW : } (Wdw : DepthwiseKernel c kH kW) (bdw : Vec c) (εn γn βn : ) (Wex : Kernel4 cExp c 1 1) (bex : Vec cExp) (Wpr : Kernel4 c cExp 1 1) (bpr : Vec c) (γls : Vec (c * h * w)) :
            Vec (c * h * w)Vec (c * h * w)

            Full ConvNeXt block = residual (block body). ConvNeXt uses an identity skip (no projection, no post-add activation), so this is the plain residual of the block body.

            Equations
            Instances For
              theorem Proofs.convNextBlock_differentiable {c cExp h w kH kW : } (Wdw : DepthwiseKernel c kH kW) (bdw : Vec c) (εn : ) (hεn : 0 < εn) (γn βn : ) (Wex : Kernel4 cExp c 1 1) (bex : Vec cExp) (Wpr : Kernel4 c cExp 1 1) (bpr : Vec c) (γls : Vec (c * h * w)) :
              Differentiable (convNextBlock Wdw bdw εn γn βn Wex bex Wpr bpr γls)

              The full ConvNeXt block is differentiable everywhere (residual of an everywhere-differentiable body).

              noncomputable def Proofs.convNextBlock_has_vjp {c cExp h w kH kW : } (Wdw : DepthwiseKernel c kH kW) (bdw : Vec c) (εn : ) (hεn : 0 < εn) (γn βn : ) (Wex : Kernel4 cExp c 1 1) (bex : Vec cExp) (Wpr : Kernel4 c cExp 1 1) (bpr : Vec c) (γls : Vec (c * h * w)) :
              HasVJP (convNextBlock Wdw bdw εn γn βn Wex bex Wpr bpr γls)

              ConvNeXt block VJP (global)residual_has_vjp on top of the block-body VJP. Needs only 0 < εn. Global since the body is everywhere-differentiable and the skip is the identity.

              Equations
              • One or more equations did not get rendered due to their size.
              Instances For
                noncomputable def Proofs.convNextBlock_has_vjp_at {c cExp h w kH kW : } (Wdw : DepthwiseKernel c kH kW) (bdw : Vec c) (εn : ) (hεn : 0 < εn) (γn βn : ) (Wex : Kernel4 cExp c 1 1) (bex : Vec cExp) (Wpr : Kernel4 c cExp 1 1) (bpr : Vec c) (γls v : Vec (c * h * w)) :
                HasVJPAt (convNextBlock Wdw bdw εn γn βn Wex bex Wpr bpr γls) v

                ConvNeXt block VJP at a point — the global witness restricted to a point. Kept for downstream _at consumers.

                Equations
                Instances For
                  noncomputable def Proofs.convNextForward {ic c cExp h w kH kW nClasses : } (Wst : Kernel4 c ic 1 1) (bst : Vec c) (εst γst βst : ) (Wdw₁ : DepthwiseKernel c kH kW) (bdw₁ : Vec c) (εn₁ γn₁ βn₁ : ) (Wex₁ : Kernel4 cExp c 1 1) (bex₁ : Vec cExp) (Wpr₁ : Kernel4 c cExp 1 1) (bpr₁ : Vec c) (γls₁ : Vec (c * h * w)) (Wdw₂ : DepthwiseKernel c kH kW) (bdw₂ : Vec c) (εn₂ γn₂ βn₂ : ) (Wex₂ : Kernel4 cExp c 1 1) (bex₂ : Vec cExp) (Wpr₂ : Kernel4 c cExp 1 1) (bpr₂ : Vec c) (γls₂ : Vec (c * h * w)) (εhd γhd βhd : ) (Wd : Mat c nClasses) (bd : Vec nClasses) :
                  Vec (ic * h * w)Vec nClasses

                  Forward ConvNeXt (representative, fixed block count = 2):

                  stem-patchify(1×1 conv) → stem-LN → block₁ → block₂ → globalAvgPool → head-LN → dense

                  Generic channel/spatial dims; icc patchify, two identity-skip ConvNeXt blocks at c, GAP to Vec c, a final LN over the pooled Vec c, and a Mat c nClasses linear head. Same LN representation caveat as convNextBlockBody applies to the stem-LN and head-LN.

                  Equations
                  • One or more equations did not get rendered due to their size.
                  Instances For
                    noncomputable def Proofs.convnext_has_vjp {ic c cExp h w kH kW nClasses : } (Wst : Kernel4 c ic 1 1) (bst : Vec c) (εst γst βst : ) (hεst : 0 < εst) (Wdw₁ : DepthwiseKernel c kH kW) (bdw₁ : Vec c) (εn₁ γn₁ βn₁ : ) (hεn₁ : 0 < εn₁) (Wex₁ : Kernel4 cExp c 1 1) (bex₁ : Vec cExp) (Wpr₁ : Kernel4 c cExp 1 1) (bpr₁ : Vec c) (γls₁ : Vec (c * h * w)) (Wdw₂ : DepthwiseKernel c kH kW) (bdw₂ : Vec c) (εn₂ γn₂ βn₂ : ) (hεn₂ : 0 < εn₂) (Wex₂ : Kernel4 cExp c 1 1) (bex₂ : Vec cExp) (Wpr₂ : Kernel4 c cExp 1 1) (bpr₂ : Vec c) (γls₂ : Vec (c * h * w)) (εhd γhd βhd : ) (hεhd : 0 < εhd) (Wd : Mat c nClasses) (bd : Vec nClasses) :
                    HasVJP (convNextForward Wst bst εst γst βst Wdw₁ bdw₁ εn₁ γn₁ βn₁ Wex₁ bex₁ Wpr₁ bpr₁ γls₁ Wdw₂ bdw₂ εn₂ γn₂ βn₂ Wex₂ bex₂ Wpr₂ bpr₂ γls₂ εhd γhd βhd Wd bd)

                    End-to-end ConvNeXt VJP (global). Everything is smooth, so the only hypotheses are the four LayerNorm positivity conditions (0 < εst, εn₁, εn₂, εhd) — no ReLU/maxpool kink conditions, unlike cnn_has_vjp_at. Chained entirely through the global vjp_comp, so the VJP holds at every input, not just a fixed point — putting ConvNeXt alongside vit_full_has_vjp as an unconditional whole-network VJP.

                    Equations
                    • One or more equations did not get rendered due to their size.
                    Instances For
                      noncomputable def Proofs.convnext_has_vjp_at {ic c cExp h w kH kW nClasses : } (Wst : Kernel4 c ic 1 1) (bst : Vec c) (εst γst βst : ) (hεst : 0 < εst) (Wdw₁ : DepthwiseKernel c kH kW) (bdw₁ : Vec c) (εn₁ γn₁ βn₁ : ) (hεn₁ : 0 < εn₁) (Wex₁ : Kernel4 cExp c 1 1) (bex₁ : Vec cExp) (Wpr₁ : Kernel4 c cExp 1 1) (bpr₁ : Vec c) (γls₁ : Vec (c * h * w)) (Wdw₂ : DepthwiseKernel c kH kW) (bdw₂ : Vec c) (εn₂ γn₂ βn₂ : ) (hεn₂ : 0 < εn₂) (Wex₂ : Kernel4 cExp c 1 1) (bex₂ : Vec cExp) (Wpr₂ : Kernel4 c cExp 1 1) (bpr₂ : Vec c) (γls₂ : Vec (c * h * w)) (εhd γhd βhd : ) (hεhd : 0 < εhd) (Wd : Mat c nClasses) (bd : Vec nClasses) (x : Vec (ic * h * w)) :
                      HasVJPAt (convNextForward Wst bst εst γst βst Wdw₁ bdw₁ εn₁ γn₁ βn₁ Wex₁ bex₁ Wpr₁ bpr₁ γls₁ Wdw₂ bdw₂ εn₂ γn₂ βn₂ Wex₂ bex₂ Wpr₂ bpr₂ γls₂ εhd γhd βhd Wd bd) x

                      End-to-end ConvNeXt VJP at a point — the global witness restricted to a point. Kept for downstream _at consumers and the comparator.

                      Equations
                      • One or more equations did not get rendered due to their size.
                      Instances For
                        theorem Proofs.convnext_has_vjp_correct {ic c cExp h w kH kW nClasses : } (Wst : Kernel4 c ic 1 1) (bst : Vec c) (εst γst βst : ) (hεst : 0 < εst) (Wdw₁ : DepthwiseKernel c kH kW) (bdw₁ : Vec c) (εn₁ γn₁ βn₁ : ) (hεn₁ : 0 < εn₁) (Wex₁ : Kernel4 cExp c 1 1) (bex₁ : Vec cExp) (Wpr₁ : Kernel4 c cExp 1 1) (bpr₁ : Vec c) (γls₁ : Vec (c * h * w)) (Wdw₂ : DepthwiseKernel c kH kW) (bdw₂ : Vec c) (εn₂ γn₂ βn₂ : ) (hεn₂ : 0 < εn₂) (Wex₂ : Kernel4 cExp c 1 1) (bex₂ : Vec cExp) (Wpr₂ : Kernel4 c cExp 1 1) (bpr₂ : Vec c) (γls₂ : Vec (c * h * w)) (εhd γhd βhd : ) (hεhd : 0 < εhd) (Wd : Mat c nClasses) (bd : Vec nClasses) (x : Vec (ic * h * w)) (dy : Vec nClasses) (i : Fin (ic * h * w)) :
                        (convnext_has_vjp Wst bst εst γst βst hεst Wdw₁ bdw₁ εn₁ γn₁ βn₁ hεn₁ Wex₁ bex₁ Wpr₁ bpr₁ γls₁ Wdw₂ bdw₂ εn₂ γn₂ βn₂ hεn₂ Wex₂ bex₂ Wpr₂ bpr₂ γls₂ εhd γhd βhd hεhd Wd bd).backward x dy i = j : Fin nClasses, pdiv (convNextForward Wst bst εst γst βst Wdw₁ bdw₁ εn₁ γn₁ βn₁ Wex₁ bex₁ Wpr₁ bpr₁ γls₁ Wdw₂ bdw₂ εn₂ γn₂ βn₂ Wex₂ bex₂ Wpr₂ bpr₂ γls₂ εhd γhd βhd Wd bd) x i j * dy j

                        Public correctness theorem for convnext_has_vjp (global) — the end-to-end ConvNeXt's backward equals the pdiv-contracted Jacobian (Jacobian-transpose applied to the cotangent), at every input x. The unconditional ConvNeXt analogue of vit_full_has_vjp_correct.

                        theorem Proofs.convnext_has_vjp_at_correct {ic c cExp h w kH kW nClasses : } (Wst : Kernel4 c ic 1 1) (bst : Vec c) (εst γst βst : ) (hεst : 0 < εst) (Wdw₁ : DepthwiseKernel c kH kW) (bdw₁ : Vec c) (εn₁ γn₁ βn₁ : ) (hεn₁ : 0 < εn₁) (Wex₁ : Kernel4 cExp c 1 1) (bex₁ : Vec cExp) (Wpr₁ : Kernel4 c cExp 1 1) (bpr₁ : Vec c) (γls₁ : Vec (c * h * w)) (Wdw₂ : DepthwiseKernel c kH kW) (bdw₂ : Vec c) (εn₂ γn₂ βn₂ : ) (hεn₂ : 0 < εn₂) (Wex₂ : Kernel4 cExp c 1 1) (bex₂ : Vec cExp) (Wpr₂ : Kernel4 c cExp 1 1) (bpr₂ : Vec c) (γls₂ : Vec (c * h * w)) (εhd γhd βhd : ) (hεhd : 0 < εhd) (Wd : Mat c nClasses) (bd : Vec nClasses) (x : Vec (ic * h * w)) (dy : Vec nClasses) (i : Fin (ic * h * w)) :
                        (convnext_has_vjp_at Wst bst εst γst βst hεst Wdw₁ bdw₁ εn₁ γn₁ βn₁ hεn₁ Wex₁ bex₁ Wpr₁ bpr₁ γls₁ Wdw₂ bdw₂ εn₂ γn₂ βn₂ hεn₂ Wex₂ bex₂ Wpr₂ bpr₂ γls₂ εhd γhd βhd hεhd Wd bd x).backward dy i = j : Fin nClasses, pdiv (convNextForward Wst bst εst γst βst Wdw₁ bdw₁ εn₁ γn₁ βn₁ Wex₁ bex₁ Wpr₁ bpr₁ γls₁ Wdw₂ bdw₂ εn₂ γn₂ βn₂ Wex₂ bex₂ Wpr₂ bpr₂ γls₂ εhd γhd βhd Wd bd) x i j * dy j

                        Public correctness theorem for convnext_has_vjp_at — exposes the witness's .correct field: the end-to-end ConvNeXt's backward equals the pdiv-contracted Jacobian (Jacobian-transpose applied to the cotangent). Analogue of cnn_has_vjp_at_correct.