Documentation

LeanMlir.Proofs.Nets.ConvNeXt.ConvNeXtDropBlock

One ConvNeXt block with its drop site, per example — forward, VJP, input cotangent #

The *drop* ConvNeXt renders put one stochastic-depth site per block on the residual branch, between LayerScale and the skip add (ConvNeXtRenderB's block forward), and the backward puts the same op on the block-output cotangent before the WHOLE branch reads it — LayerScale γ's node included — while the skip fan-in reads the raw one (ConvNeXtRenderB.bwdBlockB; dropPath_vjp_is_self). Per example the site is a scalar or absent (dropScalarOpt, Foundation.Batched.Indexed).

theorem Proofs.CnxTieGB.cnxBlockCotInChAt_eq_vjp {gf : GeluForm} {c cExp h w : ℕ} (ε : ℝ) (hε : 0 < ε) (Wdw : DepthwiseKernel c 7 7) (bdw ng nbt : Vec c) (Wex : Kernel4 cExp c 1 1) (bex : Vec cExp) (Wpr : Kernel4 c cExp 1 1) (bpr lg : Vec c) (xin dyOut : Vec (c * h * w)) :
CnxTie.cnxBlockCotInChAt gf ε Wdw bdw ng nbt Wex bex Wpr bpr lg xin dyOut = (cnxBlockChWHasVJP gf { Wdw := Wdw, bdw := bdw, εn := ε, γn := ng, βn := nbt, Wex := Wex, bex := bex, Wpr := Wpr, bpr := bpr, γls := lg } hε).backward xin dyOut

A ConvNeXt block's input cotangent is its certified VJP's backward. The chain's per-op pieces (depthwiseFlatHasVJP, the 1×1 conv2dHasVJP3s, the GELU mask, layer scale, chanLNTensor3Back) are rewritten into cnxBlockChBack_eq_vjp's form, which ties the block.

@[reducible, inline]
abbrev Proofs.CnxTie.CnxTieBlk.toCh {c cExp : ℕ} (p : CnxTieBlk c cExp) (h w : ℕ) (ε : ℝ) :
CnxBlockParamsCh c cExp h w 7 7

The block's weights as ConvNeXtFullT's record, at the shared ε.

Equations
  • p.toCh h w ε = { Wdw := p.aW, bdw := p.aB, εn := ε, γn := p.nG, βn := p.nB, Wex := p.eW, bex := p.eB, Wpr := p.pW, bpr := p.pB, γls := p.sL }
Instances For
    @[reducible, inline]
    noncomputable abbrev Proofs.CnxTie.CnxTieBlk.bodyF {c cExp : ℕ} (gf : GeluForm) (p : CnxTieBlk c cExp) {h w : ℕ} (ε : ℝ) :
    Vec (c * h * w) → Vec (c * h * w)

    The block's residual branch, flat: depthwise → channel LN → expand → GELU → project → LayerScale (cnxBodyWith at the record).

    Equations
    Instances For
      @[reducible, inline]
      noncomputable abbrev Proofs.CnxTie.CnxTieBlk.fwdOD {c cExp : ℕ} (gf : GeluForm) (p : CnxTieBlk c cExp) {h w : ℕ} (ε : ℝ) (s : Option ℝ) :
      Vec (c * h * w) → Vec (c * h * w)

      The block forward at its drop site: the branch through siteScale s, then the skip.

      Equations
      Instances For
        theorem Proofs.CnxTieGB.fwdOD_none {c cExp : ℕ} {gf : GeluForm} {h w : ℕ} (ε : ℝ) (p : CnxTie.CnxTieBlk c cExp) :

        With no site rendered the block is the drop-free one.

        noncomputable def Proofs.CnxTieGB.cnxBlockCotInChAtD (gf : GeluForm) {c cExp h w : ℕ} (ε : ℝ) (Wdw : DepthwiseKernel c 7 7) (bdw ng nbt : Vec c) (Wex : Kernel4 cExp c 1 1) (bex : Vec cExp) (Wpr : Kernel4 c cExp 1 1) (bpr lg : Vec c) (s : Option ℝ) (xin dyOut : Vec (c * h * w)) :
        Vec (c * h * w)

        The block's input cotangent at its drop site — cnxBlockCotInChAt's let chain with the branch fed s ⊙ dyOut and the skip fan-in the raw dyOut.

        Equations
        • One or more equations did not get rendered due to their size.
        Instances For
          @[reducible, inline]
          noncomputable abbrev Proofs.CnxTie.CnxTieBlk.cotInD {c cExp : ℕ} (gf : GeluForm) (p : CnxTieBlk c cExp) {h w : ℕ} (ε : ℝ) (s : Option ℝ) :
          Vec (c * h * w) → Vec (c * h * w) → Vec (c * h * w)

          The block's input cotangent at its drop site (cnxBlockCotInChAt at none).

          Equations
          Instances For
            theorem Proofs.CnxTieGB.cotInD_none {c cExp : ℕ} {gf : GeluForm} {h w : ℕ} (ε : ℝ) (p : CnxTie.CnxTieBlk c cExp) :

            With no site rendered the chain is the drop-free one.

            theorem Proofs.CnxTieGB.cotInD_eq_cotIn {c cExp : ℕ} {gf : GeluForm} {h w : ℕ} (ε : ℝ) (p : CnxTie.CnxTieBlk c cExp) (s : Option ℝ) (xin dy : Vec (c * h * w)) :
            CnxTie.CnxTieBlk.cotInD gf p ε s xin dy = fun (i : Fin (c * h * w)) => p.cotIn gf ε xin (dropScalarOpt s dy) i - dropScalarOpt s dy i + dy i

            The chain at the site is the drop-free chain at the dropped cotangent, with the skip's s ⊙ dy swapped back for dy.

            theorem Proofs.CnxTieGB.bodyF_differentiable {c cExp : ℕ} {gf : GeluForm} {h w : ℕ} (ε : ℝ) (hε : 0 < ε) (p : CnxTie.CnxTieBlk c cExp) :
            noncomputable def Proofs.CnxTieGB.bodyFHasVJP {c cExp : ℕ} (gf : GeluForm) {h w : ℕ} (ε : ℝ) (hε : 0 < ε) (p : CnxTie.CnxTieBlk c cExp) :

            The branch's VJP — the one cnxBlockChWHasVJP puts under its residual.

            Equations
            Instances For
              theorem Proofs.CnxTieGB.bodyF_back {c cExp : ℕ} {gf : GeluForm} {h w : ℕ} (ε : ℝ) (hε : 0 < ε) (p : CnxTie.CnxTieBlk c cExp) (x v : Vec (c * h * w)) :
              (bodyFHasVJP gf ε hε p).backward x v = fun (i : Fin (c * h * w)) => p.cotIn gf ε x v i - v i

              The branch's backward is the render's chain without its skip (cnxBlockCotInChAt_eq_vjp minus the residual's identity).

              theorem Proofs.CnxTieGB.fwdOD_differentiable {c cExp : ℕ} {gf : GeluForm} {h w : ℕ} (ε : ℝ) (hε : 0 < ε) (p : CnxTie.CnxTieBlk c cExp) (s : Option ℝ) :
              noncomputable def Proofs.CnxTie.CnxTieBlk.fwdODHasVJP {c cExp : ℕ} (gf : GeluForm) (p : CnxTieBlk c cExp) {h w : ℕ} (ε : ℝ) (hε : 0 < ε) (s : Option ℝ) :
              HasVJP (fwdOD gf p ε s)

              The block's VJP at its drop site: the residual over the dropped branch; its backward is body.back x (s ⊙ dy) + dy by definition.

              Equations
              • One or more equations did not get rendered due to their size.
              Instances For
                theorem Proofs.CnxTieGB.cotInD_eq_vjp {c cExp : ℕ} {gf : GeluForm} {h w : ℕ} (ε : ℝ) (hε : 0 < ε) (p : CnxTie.CnxTieBlk c cExp) (s : Option ℝ) (xin dy : Vec (c * h * w)) :
                CnxTie.CnxTieBlk.cotInD gf p ε s xin dy = (CnxTie.CnxTieBlk.fwdODHasVJP gf p ε hε s).backward xin dy

                The chain's block-input cotangent at the site is the block VJP's backward.