Documentation

LeanMlir.Proofs.Nets.ViT.ViTDropBlock

One ViT block with its two drop sites, per example — forward, VJP, input cotangent #

The *drop* ViT renders scale the attention branch (after the out-dense, before the first skip add) and the MLP branch (after fc2, before the second) by the example's own mask entry (ViTRenderB.vBlockFwdB), and the backward puts the same op on each branch's cotangent while the skip fan-ins read the raw one (ViTRenderB's block backward; dropPath_vjp_is_self). Per example, a site is a scalar or absent (dropScalarOpt, Foundation.Batched.Indexed), so this file states the block at two Option ℝ sites sA sM; the batched ties lift it with batchMapIdx, example n at exampleSite of the masks.

noncomputable def Proofs.ViTTieGB.vitCotHVD (gf : GeluForm) {Np1 D mlpDim : ℕ} (ε : ℝ) (γ2 : Vec D) (Wfc1 : Mat D mlpDim) (Wfc2 : Mat mlpDim D) (sM : Option ℝ) (h : Vec (Np1 * D)) (m1 : Vec (Np1 * mlpDim)) (dyOut : Vec (Np1 * D)) :
Vec (Np1 * D)

Cot at the attention-sublayer output h with the MLP branch's site: dyOut raw on the skip, sM ⊙ dyOut into fc2's backward (vitCotHV at none).

Equations
  • One or more equations did not get rendered due to their size.
Instances For
    noncomputable def Proofs.ViTTieGB.vitCotAttVD (gf : GeluForm) {Np1 D mlpDim : ℕ} (ε : ℝ) (γ2 : Vec D) (Wo : Mat D D) (Wfc1 : Mat D mlpDim) (Wfc2 : Mat mlpDim D) (sA sM : Option ℝ) (h : Vec (Np1 * D)) (m1 : Vec (Np1 * mlpDim)) (dyOut : Vec (Np1 * D)) :
    Vec (Np1 * D)

    Cot at the SDPA output with both sites: Woᵀ of sA ⊙ cotH (vitCotAttV at none none).

    Equations
    • One or more equations did not get rendered due to their size.
    Instances For
      noncomputable def Proofs.ViTTieGB.vitBlockSpelledMHVD (gf : GeluForm) (Np1 heads d mlpDim : ℕ) (ε : ℝ) (γ1 β1 : Vec (heads * d)) (Wq Wk Wv Wo : Mat (heads * d) (heads * d)) (bq bk bv bo γ2 β2 : Vec (heads * d)) (Wfc1 : Mat (heads * d) mlpDim) (bfc1 : Vec mlpDim) (Wfc2 : Mat mlpDim (heads * d)) (bfc2 : Vec (heads * d)) (sA sM : Option ℝ) (X : Mat Np1 (heads * d)) :
      Mat Np1 (heads * d)

      The multi-head block with its two drop sites, spelled as the render emits it — vitBlockSpelledMHV with siteScale sA on the out-projection's output and siteScale sM on fc2's, each before its skip add.

      Equations
      • One or more equations did not get rendered due to their size.
      Instances For
        @[reducible, inline]
        noncomputable abbrev Proofs.BlockParamsV.fwdOD (gf : GeluForm) {Np1 heads d mlpDim : ℕ} (p : BlockParamsV (heads * d) mlpDim) (ε : ℝ) (sA sM : Option ℝ) :
        Vec (Np1 * (heads * d)) → Vec (Np1 * (heads * d))

        The block's forward at its two sites (vitBlockFwdOMHV at none none).

        Equations
        • One or more equations did not get rendered due to their size.
        Instances For
          theorem Proofs.ViTTieGB.fwdOD_none {Np1 heads d mlpDim : ℕ} {gf : GeluForm} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) :

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

          noncomputable def Proofs.ViTTieGB.vitBlockCotInAtMHVD (gf : GeluForm) {Np1 heads d mlpDim : ℕ} (ε : ℝ) (γ1 β1 γ2 β2 : Vec (heads * d)) (Wq Wk Wv Wo : Mat (heads * d) (heads * d)) (bq bk bv bo : Vec (heads * d)) (Wfc1 : Mat (heads * d) mlpDim) (bfc1 : Vec mlpDim) (Wfc2 : Mat mlpDim (heads * d)) (sA sM : Option ℝ) (xin dyOut : Vec (Np1 * (heads * d))) :
          Vec (Np1 * (heads * d))

          The block's input cotangent at its two sites — vitBlockCotInAtMHV's let chain with the saves recomputed at the attention site (h reads siteScale sA of the out-projection), the MLP branch fed sM ⊙ dyOut (vitCotHVD) and the attention branch sA ⊙ cotH; both skip fan-ins raw.

          Equations
          • One or more equations did not get rendered due to their size.
          Instances For
            @[reducible, inline]
            noncomputable abbrev Proofs.BlockParamsV.cotInD (gf : GeluForm) {Np1 heads d mlpDim : ℕ} (p : BlockParamsV (heads * d) mlpDim) (ε : ℝ) (sA sM : Option ℝ) :
            Vec (Np1 * (heads * d)) → Vec (Np1 * (heads * d)) → Vec (Np1 * (heads * d))

            The block's input cotangent at its two sites (vitBlockCotInAtMHV at none none).

            Equations
            Instances For
              theorem Proofs.ViTTieGB.cotInD_none {Np1 heads d mlpDim : ℕ} {gf : GeluForm} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) :

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

              noncomputable def Proofs.ViTTieGB.vitAttnBrF {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) :
              Vec (Np1 * (heads * d)) → Vec (Np1 * (heads * d))

              The attention branch mhsa ∘ LN₁, flat.

              Equations
              • One or more equations did not get rendered due to their size.
              Instances For
                noncomputable def Proofs.ViTTieGB.vitMlpBrF {Np1 heads d mlpDim : ℕ} (gf : GeluForm) (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) :
                Vec (Np1 * (heads * d)) → Vec (Np1 * (heads * d))

                The MLP branch mlp ∘ LN₂, flat.

                Equations
                • One or more equations did not get rendered due to their size.
                Instances For
                  theorem Proofs.ViTTieGB.vitAttnBrF_differentiable {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) :
                  theorem Proofs.ViTTieGB.vitMlpBrF_differentiable {Np1 heads d mlpDim : ℕ} {gf : GeluForm} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) :
                  noncomputable def Proofs.ViTTieGB.vitAttnBrHasVJP {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) :

                  The attention branch's VJP — the branch half of transformerAttnSublayerVHasVJPMat.

                  Equations
                  • One or more equations did not get rendered due to their size.
                  Instances For
                    noncomputable def Proofs.ViTTieGB.vitMlpBrHasVJP {Np1 heads d mlpDim : ℕ} (gf : GeluForm) (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) :
                    HasVJP (vitMlpBrF gf ε p)

                    The MLP branch's VJP — the branch half of transformerMlpSublayerVHasVJPMat.

                    Equations
                    • One or more equations did not get rendered due to their size.
                    Instances For
                      theorem Proofs.ViTTieGB.vitAttnBr_back {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (v w : Vec (Np1 * (heads * d))) :
                      (vitAttnBrHasVJP ε hε p).backward v w = (rowLNVecFlatBack Np1 (heads * d) ε p.γ1 v ∘ mhsaBackFlat p.Wq p.Wk p.Wv p.Wo (fun (r : Fin Np1) => dense p.Wq p.bq (layerNormVec (heads * d) ε p.γ1 p.β1 (Mat.unflatten v r))) (fun (r : Fin Np1) => dense p.Wk p.bk (layerNormVec (heads * d) ε p.γ1 p.β1 (Mat.unflatten v r))) fun (r : Fin Np1) => dense p.Wv p.bv (layerNormVec (heads * d) ε p.γ1 p.β1 (Mat.unflatten v r))) w

                      The attention branch's backward is the render's chain (LN₁-back after the multi-head backward, at the saves of input v): attnSubFlat_tie_v with its skip taken off.

                      theorem Proofs.ViTTieGB.vitMlpBr_back {Np1 heads d mlpDim : ℕ} {gf : GeluForm} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (v w : Vec (Np1 * (heads * d))) :
                      (vitMlpBrHasVJP gf ε hε p).backward v w = (rowLNVecFlatBack Np1 (heads * d) ε p.γ2 v ∘ perRowFlatPR Np1 (heads * d) fun (r : Fin Np1) => dense p.Wfc1.transpose 0 ∘ (diagBack fun (c : Fin mlpDim) => gf.scalarDeriv (dense p.Wfc1 p.bfc1 (layerNormVec (heads * d) ε p.γ2 p.β2 (Mat.unflatten v r)) c)) ∘ dense p.Wfc2.transpose 0) w

                      The MLP branch's backward is the render's chain (LN₂-back after the per-row MLP back, at the saves of input v): mlpSubFlat_tie_v with its skip taken off.

                      noncomputable def Proofs.ViTTieGB.vitAttnSiteF {Np1 heads d mlpDim : ℕ} (sA : Option ℝ) (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) :
                      Vec (Np1 * (heads * d)) → Vec (Np1 * (heads * d))

                      The attention sublayer with its site, flat: v ↦ v + sA ⊙ (mhsa ∘ LN₁) v.

                      Equations
                      Instances For
                        noncomputable def Proofs.ViTTieGB.vitMlpSiteF {Np1 heads d mlpDim : ℕ} (gf : GeluForm) (sM : Option ℝ) (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) :
                        Vec (Np1 * (heads * d)) → Vec (Np1 * (heads * d))

                        The MLP sublayer with its site, flat: v ↦ v + sM ⊙ (mlp ∘ LN₂) v.

                        Equations
                        Instances For
                          theorem Proofs.ViTTieGB.vitAttnSiteF_eq {Np1 heads d mlpDim : ℕ} (sA : Option ℝ) (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (xin : Vec (Np1 * (heads * d))) :
                          vitAttnSiteF sA ε p xin = Mat.flatten fun (r : Fin Np1) (s : Fin (heads * d)) => Mat.unflatten xin r s + siteScale sA (mhsaLayer Np1 heads d p.Wq p.Wk p.Wv p.Wo p.bq p.bk p.bv p.bo (fun (n : Fin Np1) => layerNormVec (heads * d) ε p.γ1 p.β1 (Mat.unflatten xin n)) r s)

                          The attention sublayer's output at its site, as a matrix.

                          theorem Proofs.ViTTieGB.fwdOD_eq_sites {Np1 heads d mlpDim : ℕ} {gf : GeluForm} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (sA sM : Option ℝ) :
                          BlockParamsV.fwdOD gf p ε sA sM = vitMlpSiteF gf sM ε p ∘ vitAttnSiteF sA ε p

                          The spelled block with its sites is the two site residuals, composed.

                          theorem Proofs.ViTTieGB.vitAttnSiteF_differentiable {Np1 heads d mlpDim : ℕ} (sA : Option ℝ) (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) :
                          theorem Proofs.ViTTieGB.vitMlpSiteF_differentiable {Np1 heads d mlpDim : ℕ} {gf : GeluForm} (sM : Option ℝ) (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) :
                          noncomputable def Proofs.ViTTieGB.vitAttnSiteHasVJP {Np1 heads d mlpDim : ℕ} (sA : Option ℝ) (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) :

                          The attention site residual's VJP (siteResHasVJP over the branch).

                          Equations
                          Instances For
                            noncomputable def Proofs.ViTTieGB.vitMlpSiteHasVJP {Np1 heads d mlpDim : ℕ} (gf : GeluForm) (sM : Option ℝ) (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) :
                            HasVJP (vitMlpSiteF gf sM ε p)

                            The MLP site residual's VJP (siteResHasVJP over the branch).

                            Equations
                            Instances For
                              theorem Proofs.ViTTieGB.fwdOD_differentiable {Np1 heads d mlpDim : ℕ} {gf : GeluForm} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (sA sM : Option ℝ) :
                              noncomputable def Proofs.BlockParamsV.fwdODHasVJP (gf : GeluForm) {Np1 heads d mlpDim : ℕ} (p : BlockParamsV (heads * d) mlpDim) (ε : ℝ) (hε : 0 < ε) (sA sM : Option ℝ) :
                              HasVJP (fwdOD gf p ε sA sM)

                              The block's VJP at its two sites: the attention site residual, then the MLP one, each siteResHasVJP over its branch's certified VJP.

                              Equations
                              • One or more equations did not get rendered due to their size.
                              Instances For
                                theorem Proofs.ViTTieGB.fwdODHasVJP_backward {Np1 heads d mlpDim : ℕ} {gf : GeluForm} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (sA sM : Option ℝ) (x dy : Vec (Np1 * (heads * d))) :
                                (BlockParamsV.fwdODHasVJP gf p ε hε sA sM).backward x dy = have c := fun (i : Fin (Np1 * (heads * d))) => dy i + (vitMlpBrHasVJP gf ε hε p).backward (vitAttnSiteF sA ε p x) (dropScalarOpt sM dy) i; fun (i : Fin (Np1 * (heads * d))) => c i + (vitAttnBrHasVJP ε hε p).backward x (dropScalarOpt sA c) i

                                The block VJP's backward, unfolded: the MLP branch reads sM ⊙ dy at the attention output, the attention branch sA ⊙ the resulting skip cotangent; both skips raw.

                                theorem Proofs.ViTTieGB.vitCotXinV_attn {Np1 heads d : ℕ} (ε : ℝ) (γ1 : Vec (heads * d)) (Wq Wk Wv Wo : Mat (heads * d) (heads * d)) (Q K V : Mat Np1 (heads * d)) (xin w c : Vec (Np1 * (heads * d))) :
                                vitCotXinV ε γ1 Wq Wk Wv xin (vitCotDQmh Np1 heads d Q.flatten K.flatten V.flatten (StableHLO.rowDenseBackFlat Np1 (heads * d) (heads * d) Wo w)) (vitCotDKmh Np1 heads d Q.flatten K.flatten V.flatten (StableHLO.rowDenseBackFlat Np1 (heads * d) (heads * d) Wo w)) (vitCotDVmh Np1 heads d Q.flatten K.flatten V.flatten (StableHLO.rowDenseBackFlat Np1 (heads * d) (heads * d) Wo w)) c = fun (i : Fin (Np1 * (heads * d))) => c i + (rowLNVecFlatBack Np1 (heads * d) ε γ1 xin ∘ mhsaBackFlat Wq Wk Wv Wo Q K V) w i

                                The attention half of the chain's algebra: the block-input fan-in of the three dense cotangents the core hands back from Woᵀ w is c plus the attention branch's backward of w (vitCotXin_eq_blockBack's attention step, at a general skip cotangent c).

                                theorem Proofs.ViTTieGB.vitCotHVD_eq {Np1 mlpDim : ℕ} {gf : GeluForm} {D : ℕ} (ε : ℝ) (γ2 : Vec D) (Wfc1 : Mat D mlpDim) (Wfc2 : Mat mlpDim D) (sM : Option ℝ) (H : Mat Np1 D) (m1 : Mat Np1 mlpDim) (dy : Vec (Np1 * D)) :
                                vitCotHVD gf ε γ2 Wfc1 Wfc2 sM H.flatten m1.flatten dy = fun (i : Fin (Np1 * D)) => dy i + (rowLNVecFlatBack Np1 D ε γ2 H.flatten ∘ perRowFlatPR Np1 D fun (r : Fin Np1) => dense Wfc1.transpose 0 ∘ (diagBack fun (c : Fin mlpDim) => gf.scalarDeriv (m1 r c)) ∘ dense Wfc2.transpose 0) (dropScalarOpt sM dy) i

                                The MLP half: vitCotHVD is the raw skip plus the MLP branch's backward of sM ⊙ dy.

                                theorem Proofs.ViTTieGB.cotInD_eq_vjp {Np1 heads d mlpDim : ℕ} {gf : GeluForm} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (sA sM : Option ℝ) (xin dyOut : Vec (Np1 * (heads * d))) :
                                BlockParamsV.cotInD gf p ε sA sM xin dyOut = (BlockParamsV.fwdODHasVJP gf p ε hε sA sM).backward xin dyOut

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