Documentation

LeanMlir.Proofs.Nets.ViT.ViTParamGrad

ViT-Tiny — every parameter gradient node IS the loss's derivative in that parameter #

vit_net_tiedGB says each of the 200 parameter gradient nodes denotes its layer's parameter Jacobian contracted with the cotangent the emitted backward chain threads to it, the chain's top being the smoothed-loss cotangent. vit_net_lossGrad composes that with the chain: for any loss L of the logits whose gradient at the net's output is g, every node is ∂L/∂θ of the WHOLE net vitNetB with that one parameter varied. vit_net_lossGrad_smoothedCE discharges hL for the label-smoothed loss the artifacts ship (smoothedBatchLossDiv, whose gradient is the softmaxDiv cotangent the render emits).

How. ConvNeXt's shape (ConvNeXtParamGrad.lean): no ViT op couples examples, so the work is per example and lifted once.

Two nodes are stated as in the tie. The classifier bias node biasGradB is the identity on its operand and the batch reduce is emitted text, so its statement is the sum over the batch of the node's per-example slices. The CLS token's node carries the batch sum inside den.

Hypotheses. 0 < ε (the LayerNorms' VJPs); no smoothness hypothesis (GELU has no kink). For the smoothed loss, every example's target sums to one and 0 < nC. Drop-path and the bf16 nodes are outside this statement, as they are outside the tie. Stated at ViT-Tiny's literal dims, as the tie is.

noncomputable def Proofs.ViTTiePoCGB.colSlabApplyH {n heads d_in d_out : ℕ} (g : Fin heads → Mat n d_in → Mat n d_out) :
Mat n (heads * d_in) → Mat n (heads * d_out)

colSlabApply with its own map on each head's slab: output column (h, j) is column j of g h applied to input slab h. The attention core in one of Q, K, V is one, since head h reads the other two projections' slab h.

Equations
  • One or more equations did not get rendered due to their size.
Instances For
    theorem Proofs.ViTTiePoCGB.pdivMat_colIndepH {n heads d_in d_out : ℕ} (g : Fin heads → Mat n d_in → Mat n d_out) (h_g_diff : ∀ (h : Fin heads), Differentiable ℝ fun (v : Vec (n * d_in)) => (g h (Mat.unflatten v)).flatten) (A : Mat n (heads * d_in)) (i : Fin n) (h_j : Fin heads) (j' : Fin d_in) (k : Fin n) (h_l : Fin heads) (j'' : Fin d_out) :
    pdivMat (colSlabApplyH g) A i (finProdFinEquiv (h_j, j')) k (finProdFinEquiv (h_l, j'')) = if h_j = h_l then pdivMat (g h_l) (fun (r' : Fin n) (j_in : Fin d_in) => A r' (finProdFinEquiv (h_l, j_in))) i j' k j'' else 0

    The Jacobian stays block-diagonal across heads — pdivMat_colIndep with a per-slab map: zero unless the input and output slabs agree, and g h's own Jacobian on slab h if they do.

    noncomputable def Proofs.ViTTiePoCGB.colSlabwiseHasVJPMatH {n heads d_in d_out : ℕ} {g : Fin heads → Mat n d_in → Mat n d_out} (hg : (h : Fin heads) → HasVJPMat (g h)) (hg_diff : ∀ (h : Fin heads), Differentiable ℝ fun (v : Vec (n * d_in)) => (g h (Mat.unflatten v)).flatten) :

    Lift per-slab VJPs to colSlabApplyH — colSlabwiseHasVJPMat with a per-slab map: the backward runs slab h's own backward on slab h.

    Equations
    • One or more equations did not get rendered due to their size.
    Instances For
      theorem Proofs.ViTTiePoCGB.colSlabApplyH_flat_differentiable {n heads d_in d_out : ℕ} (g : Fin heads → Mat n d_in → Mat n d_out) (hg_diff : ∀ (h : Fin heads), Differentiable ℝ fun (v : Vec (n * d_in)) => (g h (Mat.unflatten v)).flatten) :
      Differentiable ℝ fun (v : Vec (n * (heads * d_in))) => (colSlabApplyH g (Mat.unflatten v)).flatten

      colSlabApplyH is differentiable, flattened, when every slab's map is.

      theorem Proofs.ViTTiePoCGB.sdpaQ_flat_differentiable (n d : ℕ) (K V : Mat n d) :
      Differentiable ℝ fun (v : Vec (n * d)) => (sdpa n d (Mat.unflatten v) K V).flatten

      Single-head attention is differentiable in Q, flattened (K, V fixed).

      theorem Proofs.ViTTiePoCGB.sdpaK_flat_differentiable (n d : ℕ) (Q V : Mat n d) :
      Differentiable ℝ fun (v : Vec (n * d)) => (sdpa n d Q (Mat.unflatten v) V).flatten

      …in K.

      theorem Proofs.ViTTiePoCGB.sdpaV_flat_differentiable (n d : ℕ) (Q K : Mat n d) :
      Differentiable ℝ fun (v : Vec (n * d)) => (sdpa n d Q K (Mat.unflatten v)).flatten

      …in V (the softmax weights are a constant).

      noncomputable def Proofs.ViTTiePoCGB.attnCore {Np1 heads d : ℕ} (Q K V : Mat Np1 (heads * d)) :
      Mat Np1 (heads * d)

      The multi-head attention core as the render spells it: per head, slice → scaled Q·Kᵀ → row softmax → ·V → pad, summed over heads (blkSaves' att).

      Equations
      • One or more equations did not get rendered due to their size.
      Instances For
        theorem Proofs.ViTTiePoCGB.attnCore_eq_colSlabQ {Np1 heads d : ℕ} (K V : Mat Np1 (heads * d)) :
        (fun (Q : Mat Np1 (heads * d)) => attnCore Q K V) = colSlabApplyH fun (hh : Fin heads) (Q' : Mat Np1 d) => sdpa Np1 d Q' (headSliceMat Np1 heads d hh K) (headSliceMat Np1 heads d hh V)

        The core in Q is a per-slab map: each head's attention on its own Q slab (sum_headPadMat_apply).

        theorem Proofs.ViTTiePoCGB.attnCore_eq_colSlabK {Np1 heads d : ℕ} (Q V : Mat Np1 (heads * d)) :
        (fun (K : Mat Np1 (heads * d)) => attnCore Q K V) = colSlabApplyH fun (hh : Fin heads) (K' : Mat Np1 d) => sdpa Np1 d (headSliceMat Np1 heads d hh Q) K' (headSliceMat Np1 heads d hh V)

        …in K.

        theorem Proofs.ViTTiePoCGB.attnCore_eq_colSlabV {Np1 heads d : ℕ} (Q K : Mat Np1 (heads * d)) :
        (fun (V : Mat Np1 (heads * d)) => attnCore Q K V) = colSlabApplyH fun (hh : Fin heads) (V' : Mat Np1 d) => sdpa Np1 d (headSliceMat Np1 heads d hh Q) (headSliceMat Np1 heads d hh K) V'

        …in V.

        noncomputable def Proofs.ViTTiePoCGB.attnCoreQHasVJPMat {Np1 heads d : ℕ} (K V : Mat Np1 (heads * d)) :
        HasVJPMat fun (Q : Mat Np1 (heads * d)) => attnCore Q K V

        The attention core's VJP in Q, K, V fixed: per head, the certified sdpaBackQ.

        Equations
        • One or more equations did not get rendered due to their size.
        Instances For
          noncomputable def Proofs.ViTTiePoCGB.attnCoreKHasVJPMat {Np1 heads d : ℕ} (Q V : Mat Np1 (heads * d)) :
          HasVJPMat fun (K : Mat Np1 (heads * d)) => attnCore Q K V

          …in K: per head, sdpaBackK.

          Equations
          • One or more equations did not get rendered due to their size.
          Instances For
            noncomputable def Proofs.ViTTiePoCGB.attnCoreVHasVJPMat {Np1 heads d : ℕ} (Q K : Mat Np1 (heads * d)) :
            HasVJPMat fun (V : Mat Np1 (heads * d)) => attnCore Q K V

            …in V: per head, sdpaBackV.

            Equations
            • One or more equations did not get rendered due to their size.
            Instances For
              theorem Proofs.ViTTiePoCGB.attnCoreQ_flat_differentiable {Np1 heads d : ℕ} (K V : Mat Np1 (heads * d)) :
              Differentiable ℝ fun (v : Vec (Np1 * (heads * d))) => (attnCore (Mat.unflatten v) K V).flatten

              The core is differentiable in Q, flattened.

              theorem Proofs.ViTTiePoCGB.attnCoreK_flat_differentiable {Np1 heads d : ℕ} (Q V : Mat Np1 (heads * d)) :
              Differentiable ℝ fun (v : Vec (Np1 * (heads * d))) => (attnCore Q (Mat.unflatten v) V).flatten

              …in K.

              theorem Proofs.ViTTiePoCGB.attnCoreV_flat_differentiable {Np1 heads d : ℕ} (Q K : Mat Np1 (heads * d)) :
              Differentiable ℝ fun (v : Vec (Np1 * (heads * d))) => (attnCore Q K (Mat.unflatten v)).flatten

              …in V.

              theorem Proofs.ViTTiePoCGB.attnCoreQ_backward {Np1 heads d : ℕ} (Q K V : Mat Np1 (heads * d)) (dA : Vec (Np1 * (heads * d))) :

              The core's Q backward IS the tie's coreQFlat — the per-head sdpaBackQ on each slab, concatenated; rfl once the saved Q is unflattened.

              theorem Proofs.ViTTiePoCGB.attnCoreK_backward {Np1 heads d : ℕ} (Q K V : Mat Np1 (heads * d)) (dA : Vec (Np1 * (heads * d))) :

              …K: coreKFlat.

              theorem Proofs.ViTTiePoCGB.attnCoreV_backward {Np1 heads d : ℕ} (Q K V : Mat Np1 (heads * d)) (dA : Vec (Np1 * (heads * d))) :

              …V: coreVFlat.

              theorem Proofs.ViTTiePoCGB.hasGradAt_linLoss_constAdd {m : ℕ} (a x dy : Vec m) :
              HasGradAt (fun (u : Vec m) => linLoss dy fun (i : Fin m) => a i + u i) x dy

              ⟨a + ·, dy⟩ has gradient dy: a residual's other branch held fixed.

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

              The flat MLP sublayer h ↦ h + MLP(LN₂ h).

              Equations
              Instances For
                theorem Proofs.ViTTiePoCGB.vitMlpSubF_differentiable {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) :

                The flat MLP sublayer is differentiable (0 < ε, the LayerNorm).

                theorem Proofs.ViTTiePoCGB.vitCotHV_eq_residual {Np1 heads d mlpDim : ℕ} (ε : ℝ) (γ2 : Vec (heads * d)) (Wfc1 : Mat (heads * d) mlpDim) (Wfc2 : Mat mlpDim (heads * d)) (H : Mat Np1 (heads * d)) (m1 : Mat Np1 mlpDim) (dy : Vec (Np1 * (heads * d))) :
                vitCotHV ε γ2 Wfc1 Wfc2 H.flatten m1.flatten dy = residual (rowLNVecFlatBack Np1 (heads * d) ε γ2 H.flatten ∘ perRowFlatPR Np1 (heads * d) fun (r : Fin Np1) => dense Wfc1.transpose 0 ∘ (diagBack fun (c : Fin mlpDim) => geluScalarDeriv (m1 r c)) ∘ dense Wfc2.transpose 0) dy

                vitCotHV is the MLP sublayer's backward in the chain's spelling (vitCotXin_eq_blockBack's first step).

                theorem Proofs.ViTTiePoCGB.vitMlpSub_hasGradAt {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (H : Mat Np1 (heads * d)) (dy : Vec (Np1 * (heads * d))) :
                HasGradAt (fun (v : Vec (Np1 * (heads * d))) => linLoss dy (vitMlpSubF ε p v)) H.flatten (vitCotHV ε p.γ2 p.Wfc1 p.Wfc2 H.flatten (Mat.flatten fun (r : Fin Np1) => dense p.Wfc1 p.bfc1 (layerNormVec (heads * d) ε p.γ2 p.β2 (H r))) dy)

                The MLP sublayer's input gradient is vitCotHV, at any saved sublayer input H.

                The block's forward, per example, as named Mats — blkSaves' let chain verbatim.

                noncomputable def Proofs.ViTTiePoCGB.vitLn1M {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                Mat Np1 (heads * d)

                LN₁'s output.

                Equations
                Instances For
                  noncomputable def Proofs.ViTTiePoCGB.vitQM {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                  Mat Np1 (heads * d)

                  The Q projection.

                  Equations
                  Instances For
                    noncomputable def Proofs.ViTTiePoCGB.vitKM {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                    Mat Np1 (heads * d)

                    The K projection.

                    Equations
                    Instances For
                      noncomputable def Proofs.ViTTiePoCGB.vitVM {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                      Mat Np1 (heads * d)

                      The V projection.

                      Equations
                      Instances For
                        noncomputable def Proofs.ViTTiePoCGB.vitAttM {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                        Mat Np1 (heads * d)

                        The attention core's output (the out-projection's input).

                        Equations
                        Instances For
                          noncomputable def Proofs.ViTTiePoCGB.vitOM {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                          Mat Np1 (heads * d)

                          The out-projection's output.

                          Equations
                          Instances For
                            noncomputable def Proofs.ViTTiePoCGB.vitHM {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                            Mat Np1 (heads * d)

                            The attention sublayer's output h.

                            Equations
                            Instances For
                              noncomputable def Proofs.ViTTiePoCGB.vitLn2M {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                              Mat Np1 (heads * d)

                              LN₂'s output.

                              Equations
                              Instances For
                                noncomputable def Proofs.ViTTiePoCGB.vitM1M {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                                Mat Np1 mlpDim

                                fc1's output (pre-GELU).

                                Equations
                                Instances For
                                  noncomputable def Proofs.ViTTiePoCGB.vitWoF {Np1 heads d mlpDim : ℕ} (p : BlockParamsV (heads * d) mlpDim) :
                                  Vec (Np1 * (heads * d)) → Vec (Np1 * (heads * d))

                                  The out-projection, flat.

                                  Equations
                                  Instances For
                                    noncomputable def Proofs.ViTTiePoCGB.vitPostO {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                                    Vec (Np1 * (heads * d)) → Vec (Np1 * (heads * d))

                                    The block after the out-projection: the attention residual, then the MLP sublayer.

                                    Equations
                                    Instances For
                                      theorem Proofs.ViTTiePoCGB.vitHM_flat {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                                      (fun (i : Fin (Np1 * (heads * d))) => y i + (vitOM ε p y).flatten i) = (vitHM ε p y).flatten

                                      The attention residual's flat sum is h, flattened.

                                      theorem Proofs.ViTTiePoCGB.vitPostO_hasGradAt {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (y dy : Vec (Np1 * (heads * d))) :
                                      HasGradAt (fun (u : Vec (Np1 * (heads * d))) => linLoss dy (vitPostO ε p y u)) (vitOM ε p y).flatten (cH ε p.γ1 p.β1 p.γ2 p.β2 p.Wq p.Wk p.Wv p.Wo p.bq p.bk p.bv p.bo p.Wfc1 p.bfc1 p.Wfc2 y dy)

                                      The out-projection's output cotangent is cH — the loss after the out-projection.

                                      theorem Proofs.ViTTiePoCGB.vitPostAtt_hasGradAt {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (y dy : Vec (Np1 * (heads * d))) :
                                      HasGradAt (fun (a : Vec (Np1 * (heads * d))) => linLoss dy (vitPostO ε p y (vitWoF p a))) (vitAttM ε p y).flatten (cAtt ε p.γ1 p.β1 p.γ2 p.β2 p.Wq p.Wk p.Wv p.Wo p.bq p.bk p.bv p.bo p.Wfc1 p.bfc1 p.Wfc2 y dy)

                                      The block after the attention core: out-projection, then vitPostO.

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

                                      The block after the Q projection.

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

                                        The block after the K projection.

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

                                          The block after the V projection.

                                          Equations
                                          • One or more equations did not get rendered due to their size.
                                          Instances For
                                            theorem Proofs.ViTTiePoCGB.vitPostQ_hasGradAt {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (y dy : Vec (Np1 * (heads * d))) :
                                            HasGradAt (fun (u : Vec (Np1 * (heads * d))) => linLoss dy (vitPostQ ε p y u)) (vitQM ε p y).flatten (cQ ε p.γ1 p.β1 p.γ2 p.β2 p.Wq p.Wk p.Wv p.Wo p.bq p.bk p.bv p.bo p.Wfc1 p.bfc1 p.Wfc2 y dy)

                                            The Q projection's output cotangent is cQ: cAtt pulled back through the core in Q.

                                            theorem Proofs.ViTTiePoCGB.vitPostK_hasGradAt {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (y dy : Vec (Np1 * (heads * d))) :
                                            HasGradAt (fun (u : Vec (Np1 * (heads * d))) => linLoss dy (vitPostK ε p y u)) (vitKM ε p y).flatten (cK ε p.γ1 p.β1 p.γ2 p.β2 p.Wq p.Wk p.Wv p.Wo p.bq p.bk p.bv p.bo p.Wfc1 p.bfc1 p.Wfc2 y dy)

                                            The K projection's output cotangent is cK.

                                            theorem Proofs.ViTTiePoCGB.vitPostV_hasGradAt {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (y dy : Vec (Np1 * (heads * d))) :
                                            HasGradAt (fun (u : Vec (Np1 * (heads * d))) => linLoss dy (vitPostV ε p y u)) (vitVM ε p y).flatten (cV ε p.γ1 p.β1 p.γ2 p.β2 p.Wq p.Wk p.Wv p.Wo p.bq p.bk p.bv p.bo p.Wfc1 p.bfc1 p.Wfc2 y dy)

                                            The V projection's output cotangent is cV.

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

                                            The block after LN₁: the multi-head attention layer, then vitPostO.

                                            Equations
                                            Instances For
                                              theorem Proofs.ViTTiePoCGB.vitPostL1_hasGradAt {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (y dy : Vec (Np1 * (heads * d))) :
                                              HasGradAt (fun (u : Vec (Np1 * (heads * d))) => linLoss dy (vitPostL1 ε p y u)) (vitLn1M ε p y).flatten (cLn1 ε p.γ1 p.β1 p.γ2 p.β2 p.Wq p.Wk p.Wv p.Wo p.bq p.bk p.bv p.bo p.Wfc1 p.bfc1 p.Wfc2 y dy)

                                              LN₁'s output cotangent is cLn1: cH pulled back through the certified attention layer (mhsaBackFlat_eq_mhsa_vjp), whose three paths are the Q/K/V fan-in.

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

                                              The block after LN₂: the MLP body, then the residual.

                                              Equations
                                              • One or more equations did not get rendered due to their size.
                                              Instances For
                                                theorem Proofs.ViTTiePoCGB.vitPostL2_hasGradAt {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y dy : Vec (Np1 * (heads * d))) :
                                                HasGradAt (fun (u : Vec (Np1 * (heads * d))) => linLoss dy (vitPostL2 ε p y u)) (vitLn2M ε p y).flatten (cLn2 ε p.γ1 p.β1 p.γ2 p.β2 p.Wq p.Wk p.Wv p.Wo p.bq p.bk p.bv p.bo p.Wfc1 p.bfc1 p.Wfc2 y dy)

                                                LN₂'s output cotangent is cLn2: the MLP body's certified backward (transformerMlp_back_flat_eq_perRowFlatPR).

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

                                                The block after fc1: GELU, fc2, then the residual.

                                                Equations
                                                • One or more equations did not get rendered due to their size.
                                                Instances For
                                                  theorem Proofs.ViTTiePoCGB.vitPostF1_hasGradAt {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y dy : Vec (Np1 * (heads * d))) :
                                                  HasGradAt (fun (u : Vec (Np1 * mlpDim)) => linLoss dy (vitPostF1 ε p y u)) (vitM1M ε p y).flatten (cM1 ε p.γ1 p.β1 p.γ2 p.β2 p.Wq p.Wk p.Wv p.Wo p.bq p.bk p.bv p.bo p.Wfc1 p.bfc1 p.Wfc2 y dy)

                                                  fc1's output cotangent is cM1: through fc2 and the GELU.

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

                                                  The block after fc2: the residual.

                                                  Equations
                                                  Instances For

                                                    The block forward, read after each node #

                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_eq_postO {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                                                    p.fwdO ε y = vitPostO ε p y (vitOM ε p y).flatten

                                                    The block forward is vitPostO at the out-projection's output.

                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_eq_postF2 {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                                                    p.fwdO ε y = vitPostF2 ε p y (Mat.flatten fun (r : Fin Np1) => dense p.Wfc2 p.bfc2 (gelu mlpDim (vitM1M ε p y r)))

                                                    The block forward is vitPostF2 at fc2's output (rfl).

                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_γ1 {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (θ : Vec (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := θ, β1 := p.β1, Wq := p.Wq, Wk := p.Wk, Wv := p.Wv, Wo := p.Wo, bq := p.bq, bk := p.bk, bv := p.bv, bo := p.bo, γ2 := p.γ2, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostL1 ε p y (Mat.flatten fun (r : Fin Np1) => layerNormVec (heads * d) ε θ p.β1 (Mat.unflatten y r))

                                                    The block with γ1 varied is vitPostL1 after LN₁ at that γ1. The fifteen lemmas below say the same for each other parameter: the node's op at the varied parameter, between the block's prefix and the rest of the block (vitPost*).

                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_β1 {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (θ : Vec (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := θ, Wq := p.Wq, Wk := p.Wk, Wv := p.Wv, Wo := p.Wo, bq := p.bq, bk := p.bk, bv := p.bv, bo := p.bo, γ2 := p.γ2, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostL1 ε p y (Mat.flatten fun (r : Fin Np1) => layerNormVec (heads * d) ε p.γ1 θ (Mat.unflatten y r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_Wq {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (W : Mat (heads * d) (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := W, Wk := p.Wk, Wv := p.Wv, Wo := p.Wo, bq := p.bq, bk := p.bk, bv := p.bv, bo := p.bo, γ2 := p.γ2, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostQ ε p y (Mat.flatten fun (r : Fin Np1) => dense W p.bq (Mat.unflatten (vitLn1M ε p y).flatten r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_bq {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (θ : Vec (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := p.Wq, Wk := p.Wk, Wv := p.Wv, Wo := p.Wo, bq := θ, bk := p.bk, bv := p.bv, bo := p.bo, γ2 := p.γ2, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostQ ε p y (Mat.flatten fun (r : Fin Np1) => dense p.Wq θ (Mat.unflatten (vitLn1M ε p y).flatten r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_Wk {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (W : Mat (heads * d) (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := p.Wq, Wk := W, Wv := p.Wv, Wo := p.Wo, bq := p.bq, bk := p.bk, bv := p.bv, bo := p.bo, γ2 := p.γ2, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostK ε p y (Mat.flatten fun (r : Fin Np1) => dense W p.bk (Mat.unflatten (vitLn1M ε p y).flatten r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_bk {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (θ : Vec (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := p.Wq, Wk := p.Wk, Wv := p.Wv, Wo := p.Wo, bq := p.bq, bk := θ, bv := p.bv, bo := p.bo, γ2 := p.γ2, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostK ε p y (Mat.flatten fun (r : Fin Np1) => dense p.Wk θ (Mat.unflatten (vitLn1M ε p y).flatten r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_Wv {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (W : Mat (heads * d) (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := p.Wq, Wk := p.Wk, Wv := W, Wo := p.Wo, bq := p.bq, bk := p.bk, bv := p.bv, bo := p.bo, γ2 := p.γ2, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostV ε p y (Mat.flatten fun (r : Fin Np1) => dense W p.bv (Mat.unflatten (vitLn1M ε p y).flatten r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_bv {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (θ : Vec (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := p.Wq, Wk := p.Wk, Wv := p.Wv, Wo := p.Wo, bq := p.bq, bk := p.bk, bv := θ, bo := p.bo, γ2 := p.γ2, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostV ε p y (Mat.flatten fun (r : Fin Np1) => dense p.Wv θ (Mat.unflatten (vitLn1M ε p y).flatten r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_Wo {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (W : Mat (heads * d) (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := p.Wq, Wk := p.Wk, Wv := p.Wv, Wo := W, bq := p.bq, bk := p.bk, bv := p.bv, bo := p.bo, γ2 := p.γ2, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostO ε p y (Mat.flatten fun (r : Fin Np1) => dense W p.bo (Mat.unflatten (vitAttM ε p y).flatten r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_bo {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (θ : Vec (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := p.Wq, Wk := p.Wk, Wv := p.Wv, Wo := p.Wo, bq := p.bq, bk := p.bk, bv := p.bv, bo := θ, γ2 := p.γ2, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostO ε p y (Mat.flatten fun (r : Fin Np1) => dense p.Wo θ (Mat.unflatten (vitAttM ε p y).flatten r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_γ2 {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (θ : Vec (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := p.Wq, Wk := p.Wk, Wv := p.Wv, Wo := p.Wo, bq := p.bq, bk := p.bk, bv := p.bv, bo := p.bo, γ2 := θ, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostL2 ε p y (Mat.flatten fun (r : Fin Np1) => layerNormVec (heads * d) ε θ p.β2 (Mat.unflatten (vitHM ε p y).flatten r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_β2 {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (θ : Vec (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := p.Wq, Wk := p.Wk, Wv := p.Wv, Wo := p.Wo, bq := p.bq, bk := p.bk, bv := p.bv, bo := p.bo, γ2 := p.γ2, β2 := θ, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostL2 ε p y (Mat.flatten fun (r : Fin Np1) => layerNormVec (heads * d) ε p.γ2 θ (Mat.unflatten (vitHM ε p y).flatten r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_Wfc1 {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (W : Mat (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := p.Wq, Wk := p.Wk, Wv := p.Wv, Wo := p.Wo, bq := p.bq, bk := p.bk, bv := p.bv, bo := p.bo, γ2 := p.γ2, β2 := p.β2, Wfc1 := W, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostF1 ε p y (Mat.flatten fun (r : Fin Np1) => dense W p.bfc1 (Mat.unflatten (vitLn2M ε p y).flatten r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_bfc1 {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (θ : Vec mlpDim) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := p.Wq, Wk := p.Wk, Wv := p.Wv, Wo := p.Wo, bq := p.bq, bk := p.bk, bv := p.bv, bo := p.bo, γ2 := p.γ2, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := θ, Wfc2 := p.Wfc2, bfc2 := p.bfc2 }.fwdO ε y = vitPostF1 ε p y (Mat.flatten fun (r : Fin Np1) => dense p.Wfc1 θ (Mat.unflatten (vitLn2M ε p y).flatten r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_Wfc2 {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (W : Mat mlpDim (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := p.Wq, Wk := p.Wk, Wv := p.Wv, Wo := p.Wo, bq := p.bq, bk := p.bk, bv := p.bv, bo := p.bo, γ2 := p.γ2, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := W, bfc2 := p.bfc2 }.fwdO ε y = vitPostF2 ε p y (Mat.flatten fun (r : Fin Np1) => dense W p.bfc2 (Mat.unflatten (Mat.flatten fun (r' : Fin Np1) => gelu mlpDim (vitM1M ε p y r')) r))
                                                    theorem Proofs.ViTTiePoCGB.vit_fwd_bfc2 {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (θ : Vec (heads * d)) (y : Vec (Np1 * (heads * d))) :
                                                    { γ1 := p.γ1, β1 := p.β1, Wq := p.Wq, Wk := p.Wk, Wv := p.Wv, Wo := p.Wo, bq := p.bq, bk := p.bk, bv := p.bv, bo := p.bo, γ2 := p.γ2, β2 := p.β2, Wfc1 := p.Wfc1, bfc1 := p.bfc1, Wfc2 := p.Wfc2, bfc2 := θ }.fwdO ε y = vitPostF2 ε p y (Mat.flatten fun (r : Fin Np1) => dense p.Wfc2 θ (Mat.unflatten (Mat.flatten fun (r' : Fin Np1) => gelu mlpDim (vitM1M ε p y r')) r))

                                                    Differentiability #

                                                    theorem Proofs.ViTTiePoCGB.rowDense_weight_differentiable {tk a c : ℕ} (b : Vec c) (x : Vec (tk * a)) :
                                                    Differentiable ℝ fun (θ : Vec (a * c)) => Mat.flatten fun (r : Fin tk) => dense (Mat.unflatten θ) b (Mat.unflatten x r)

                                                    The per-token dense is differentiable in its weight.

                                                    theorem Proofs.ViTTiePoCGB.rowDense_bias_differentiable {tk a c : ℕ} (W : Mat a c) (x : Vec (tk * a)) :
                                                    Differentiable ℝ fun (θ : Vec c) => Mat.flatten fun (r : Fin tk) => dense W θ (Mat.unflatten x r)

                                                    …in its bias.

                                                    theorem Proofs.ViTTiePoCGB.rowVecLN_gamma_differentiable {tk D : ℕ} (ε : ℝ) (β : Vec D) (x : Vec (tk * D)) :
                                                    Differentiable ℝ fun (θ : Vec D) => Mat.flatten fun (r : Fin tk) => layerNormVec D ε θ β (Mat.unflatten x r)

                                                    The per-token vector LayerNorm is differentiable in γ.

                                                    theorem Proofs.ViTTiePoCGB.rowVecLN_beta_differentiable {tk D : ℕ} (ε : ℝ) (γ : Vec D) (x : Vec (tk * D)) :
                                                    Differentiable ℝ fun (θ : Vec D) => Mat.flatten fun (r : Fin tk) => layerNormVec D ε γ θ (Mat.unflatten x r)

                                                    …in β.

                                                    theorem Proofs.ViTTiePoCGB.vitPostO_differentiable {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :

                                                    Each vitPost* is differentiable (0 < ε where the MLP sublayer's LN₂ is inside).

                                                    theorem Proofs.ViTTiePoCGB.vitWoF_differentiable {Np1 heads d mlpDim : ℕ} (p : BlockParamsV (heads * d) mlpDim) :
                                                    theorem Proofs.ViTTiePoCGB.vitPostL1_differentiable {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                                                    theorem Proofs.ViTTiePoCGB.vitPostQ_differentiable {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                                                    theorem Proofs.ViTTiePoCGB.vitPostK_differentiable {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                                                    theorem Proofs.ViTTiePoCGB.vitPostV_differentiable {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                                                    theorem Proofs.ViTTiePoCGB.vitPostL2_differentiable {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                                                    theorem Proofs.ViTTiePoCGB.vitPostF1_differentiable {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                                                    theorem Proofs.ViTTiePoCGB.vitPostF2_differentiable {Np1 heads d mlpDim : ℕ} (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (y : Vec (Np1 * (heads * d))) :
                                                    theorem Proofs.ViTTiePoCGB.lb_batchMap_congr {N a q : ℕ} (Lb : Vec (N * q) → Vec 1) (X : Vec (N * a)) {f g : Vec a → Vec q} (h : ∀ (y : Vec a), f y = g y) :

                                                    Lb of a batched stage, at two pointwise-equal stage maps.

                                                    def Proofs.ViTTiePoCGB.vitBlockLossTiedGB (N : ℕ) {Np1 heads d mlpDim : ℕ} (xN epsStr cotN : String) (ε : ℝ) (p : BlockParamsV (heads * d) mlpDim) (xin : Vec (N * (Np1 * (heads * d)))) (Φ : BlockParamsV (heads * d) mlpDim → Vec 1) (dyOut : Vec (N * (Np1 * (heads * d)))) :

                                                    ViT block, every parameter node a loss derivative — the sixteen nodes vitBlockTiedGB ties, at the tie's batched activations and cotangents, Φ the loss at the block's output as a function of the block's record.

                                                    Equations
                                                    • One or more equations did not get rendered due to their size.
                                                    Instances For
                                                      theorem Proofs.ViTTiePoCGB.vit_block_lossTiedGB (N : ℕ) {Np1 heads d mlpDim : ℕ} (xN epsStr cotN : String) (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (xin : Vec (N * (Np1 * (heads * d)))) {Lb : Vec (N * (Np1 * (heads * d))) → Vec 1} {dyOut : Vec (N * (Np1 * (heads * d)))} (hLb : HasGradAt Lb (StableHLO.batchMap N (p.fwdO ε) xin) dyOut) {Φ : BlockParamsV (heads * d) mlpDim → Vec 1} (hΦ : ∀ (p' : BlockParamsV (heads * d) mlpDim), Φ p' = Lb (StableHLO.batchMap N (p'.fwdO ε) xin)) :
                                                      vitBlockLossTiedGB N xN epsStr cotN ε p xin Φ dyOut
                                                      noncomputable def Proofs.ViTTiePoCGB.vitHeadO {nC : ℕ} (ε : ℝ) (γF βF : Vec 192) (Wcls : Mat 192 nC) (bcls : Vec nC) :
                                                      Vec (197 * 192) → Vec nC

                                                      The head per example: the final vector LN on every token, then the CLS-slice classifier — vitHeadHasVJP's map.

                                                      Equations
                                                      • One or more equations did not get rendered due to their size.
                                                      Instances For
                                                        def Proofs.ViTTiePoCGB.vitHeadLossTiedGB (N : ℕ) {nC : ℕ} (xN aN epsStr cotN : String) (ε : ℝ) (γF βF : Vec 192) (Wcls : Mat 192 nC) (bcls : Vec nC) (b12out : Vec (N * (197 * 192))) (Φ : Vec 192 → Vec 192 → Mat 192 nC → Vec nC → Vec 1) (g : Vec (N * nC)) :

                                                        Head, every parameter node a loss derivative — the final LN's two nodes (vitFinalLNTiedGB) and the classifier's two (vitHeadTiedGB). The classifier bias node is the identity on its operand (the batch reduce is emitted text), so its statement is the batch sum of the node's slices.

                                                        Equations
                                                        • One or more equations did not get rendered due to their size.
                                                        Instances For
                                                          theorem Proofs.ViTTiePoCGB.vit_head_lossTiedGB (N : ℕ) {nC : ℕ} (xN aN epsStr cotN : String) (ε : ℝ) (γF βF : Vec 192) (Wcls : Mat 192 nC) (bcls : Vec nC) (b12out : Vec (N * (197 * 192))) {L : Vec (N * nC) → Vec 1} {g : Vec (N * nC)} (hL : HasGradAt L (StableHLO.batchMap N (vitHeadO ε γF βF Wcls bcls) b12out) g) {Φ : Vec 192 → Vec 192 → Mat 192 nC → Vec nC → Vec 1} (hΦ : ∀ (a b : Vec 192) (W : Mat 192 nC) (bb : Vec nC), Φ a b W bb = L (StableHLO.batchMap N (vitHeadO ε a b W bb) b12out)) :
                                                          vitHeadLossTiedGB N xN aN epsStr cotN ε γF βF Wcls bcls b12out Φ g
                                                          theorem Proofs.ViTTiePoCGB.patchEmbedFlat_weight_differentiable (bc cls : Vec 192) (pos : Mat 197 192) (x : Vec (3 * 224 * 224)) :
                                                          Differentiable ℝ fun (θ : Vec (192 * 3 * 16 * 16)) => patchEmbedFlat 3 224 224 16 196 192 (Kernel4.unflatten θ) bc cls pos x

                                                          The patch embedding is differentiable in its conv weight.

                                                          theorem Proofs.ViTTiePoCGB.patchEmbedFlat_bias_differentiable (Wc : Kernel4 192 3 16 16) (cls : Vec 192) (pos : Mat 197 192) (x : Vec (3 * 224 * 224)) :
                                                          Differentiable ℝ fun (θ : Vec 192) => patchEmbedFlat 3 224 224 16 196 192 Wc θ cls pos x

                                                          …in its conv bias.

                                                          theorem Proofs.ViTTiePoCGB.patchEmbedFlat_cls_differentiable (Wc : Kernel4 192 3 16 16) (bc : Vec 192) (pos : Mat 197 192) (x : Vec (3 * 224 * 224)) :
                                                          Differentiable ℝ fun (θ : Vec 192) => patchEmbedFlat 3 224 224 16 196 192 Wc bc θ pos x

                                                          …in the CLS token.

                                                          theorem Proofs.ViTTiePoCGB.patchEmbedFlat_pos_differentiable (Wc : Kernel4 192 3 16 16) (bc cls : Vec 192) (x : Vec (3 * 224 * 224)) :
                                                          Differentiable ℝ fun (θ : Vec (197 * 192)) => patchEmbedFlat 3 224 224 16 196 192 Wc bc cls (Mat.unflatten θ) x

                                                          …in the position embedding.

                                                          def Proofs.ViTTiePoCGB.vitEmbedLossTiedGB (N : ℕ) (xN cotN : String) (Wc : Kernel4 192 3 16 16) (bc cls : Vec 192) (pos : Mat 197 192) (img : Vec (N * (3 * 224 * 224))) (Φ : Kernel4 192 3 16 16 → Vec 192 → Vec 192 → Mat 197 192 → Vec 1) (dyEmbed : Vec (N * (197 * 192))) :

                                                          Patch embedding, every parameter node a loss derivative — the four nodes vitEmbedTiedGB ties (the CLS token's with the batch sum inside den).

                                                          Equations
                                                          • One or more equations did not get rendered due to their size.
                                                          Instances For
                                                            theorem Proofs.ViTTiePoCGB.vit_embed_lossTiedGB (N : ℕ) (xN cotN : String) (Wc : Kernel4 192 3 16 16) (bc cls : Vec 192) (pos : Mat 197 192) (img : Vec (N * (3 * 224 * 224))) {Lb : Vec (N * (197 * 192)) → Vec 1} {dyEmbed : Vec (N * (197 * 192))} (hLb : HasGradAt Lb (StableHLO.batchMap N (patchEmbedFlat 3 224 224 16 196 192 Wc bc cls pos) img) dyEmbed) {Φ : Kernel4 192 3 16 16 → Vec 192 → Vec 192 → Mat 197 192 → Vec 1} (hΦ : ∀ (W : Kernel4 192 3 16 16) (b c : Vec 192) (q : Mat 197 192), Φ W b c q = Lb (StableHLO.batchMap N (patchEmbedFlat 3 224 224 16 196 192 W b c q) img)) :
                                                            vitEmbedLossTiedGB N xN cotN Wc bc cls pos img Φ dyEmbed
                                                            theorem Proofs.ViTTiePoCGB.vitBlkB_hasGradAt_comp (N : ℕ) {Np1 heads d mlpDim : ℕ} (ε : ℝ) (hε : 0 < ε) (p : BlockParamsV (heads * d) mlpDim) (X : Vec (N * (Np1 * (heads * d)))) {G : Vec (N * (Np1 * (heads * d))) → Vec 1} {dY : Vec (N * (Np1 * (heads * d)))} (hG : HasGradAt G (StableHLO.batchMap N (p.fwdO ε) X) dY) :
                                                            HasGradAt (fun (y : Vec (N * (Np1 * (heads * d)))) => G (StableHLO.batchMap N (p.fwdO ε) y)) X (StableHLO.batchMapAux N (p.cotIn ε) X dY)

                                                            Pull the loss gradient back through a batched block's certified VJP: the cotangent is the tie's batchMapAux N (p.cotIn ε) (vitBlockCotInB_eq_vjp).

                                                            theorem Proofs.ViTTiePoCGB.vitHeadB_hasGradAt_comp (N : ℕ) {nC : ℕ} (ε : ℝ) (hε : 0 < ε) (γF βF : Vec 192) (Wcls : Mat 192 nC) (bcls : Vec nC) (X : Vec (N * (197 * 192))) {L : Vec (N * nC) → Vec 1} {g : Vec (N * nC)} (hL : HasGradAt L (StableHLO.batchMap N (vitHeadO ε γF βF Wcls bcls) X) g) :
                                                            HasGradAt (fun (y : Vec (N * (197 * 192))) => L (StableHLO.batchMap N (vitHeadO ε γF βF Wcls bcls) y)) X (StableHLO.batchMapAux N (vitCotB2outV 196 192 nC ε γF Wcls) X g)

                                                            …and through the batched head (vitCotB2outB_eq_vjp).

                                                            noncomputable def Proofs.ViTTiePoCGB.vitNetB (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                            Vec (N * nC)

                                                            ViT-Tiny, batched: the tie's forward, stage by stage — batchMap N of the patch embedding, each block, then of the head.

                                                            Equations
                                                            • One or more equations did not get rendered due to their size.
                                                            Instances For
                                                              noncomputable def Proofs.ViTTiePoCGB.vitPreE (N : ℕ) {nC : ℕ} (w : ViTTiePoC.ViTTieWeights nC) :
                                                              Vec (N * (3 * 224 * 224)) → Vec (N * (197 * 192))

                                                              The patch embedding's output — block b1's input (the tie's ib1).

                                                              Equations
                                                              Instances For
                                                                noncomputable def Proofs.ViTTiePoCGB.vitPreB1 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                Vec (N * (3 * 224 * 224)) → Vec (N * (197 * 192))

                                                                Block b1's output.

                                                                Equations
                                                                Instances For
                                                                  noncomputable def Proofs.ViTTiePoCGB.vitPreB2 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                  Vec (N * (3 * 224 * 224)) → Vec (N * (197 * 192))

                                                                  Block b2's output.

                                                                  Equations
                                                                  Instances For
                                                                    noncomputable def Proofs.ViTTiePoCGB.vitPreB3 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                    Vec (N * (3 * 224 * 224)) → Vec (N * (197 * 192))

                                                                    Block b3's output.

                                                                    Equations
                                                                    Instances For
                                                                      noncomputable def Proofs.ViTTiePoCGB.vitPreB4 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                      Vec (N * (3 * 224 * 224)) → Vec (N * (197 * 192))

                                                                      Block b4's output.

                                                                      Equations
                                                                      Instances For
                                                                        noncomputable def Proofs.ViTTiePoCGB.vitPreB5 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                        Vec (N * (3 * 224 * 224)) → Vec (N * (197 * 192))

                                                                        Block b5's output.

                                                                        Equations
                                                                        Instances For
                                                                          noncomputable def Proofs.ViTTiePoCGB.vitPreB6 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                          Vec (N * (3 * 224 * 224)) → Vec (N * (197 * 192))

                                                                          Block b6's output.

                                                                          Equations
                                                                          Instances For
                                                                            noncomputable def Proofs.ViTTiePoCGB.vitPreB7 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                            Vec (N * (3 * 224 * 224)) → Vec (N * (197 * 192))

                                                                            Block b7's output.

                                                                            Equations
                                                                            Instances For
                                                                              noncomputable def Proofs.ViTTiePoCGB.vitPreB8 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                              Vec (N * (3 * 224 * 224)) → Vec (N * (197 * 192))

                                                                              Block b8's output.

                                                                              Equations
                                                                              Instances For
                                                                                noncomputable def Proofs.ViTTiePoCGB.vitPreB9 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                Vec (N * (3 * 224 * 224)) → Vec (N * (197 * 192))

                                                                                Block b9's output.

                                                                                Equations
                                                                                Instances For
                                                                                  noncomputable def Proofs.ViTTiePoCGB.vitPreB10 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                  Vec (N * (3 * 224 * 224)) → Vec (N * (197 * 192))

                                                                                  Block b10's output.

                                                                                  Equations
                                                                                  Instances For
                                                                                    noncomputable def Proofs.ViTTiePoCGB.vitPreB11 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                    Vec (N * (3 * 224 * 224)) → Vec (N * (197 * 192))

                                                                                    Block b11's output.

                                                                                    Equations
                                                                                    Instances For
                                                                                      noncomputable def Proofs.ViTTiePoCGB.vitPreB12 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                      Vec (N * (3 * 224 * 224)) → Vec (N * (197 * 192))

                                                                                      Block b12's output.

                                                                                      Equations
                                                                                      Instances For
                                                                                        theorem Proofs.ViTTiePoCGB.vitPreE_apply (N : ℕ) {nC : ℕ} (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                        vitPreE N w img = StableHLO.batchMap N (patchEmbedFlat 3 224 224 16 196 192 w.Wc w.bc w.cls w.pos) img
                                                                                        theorem Proofs.ViTTiePoCGB.vitPreB1_apply (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                        vitPreB1 N ε w img = StableHLO.batchMap N (w.b1.fwdO ε) (vitPreE N w img)
                                                                                        theorem Proofs.ViTTiePoCGB.vitPreB2_apply (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                        vitPreB2 N ε w img = StableHLO.batchMap N (w.b2.fwdO ε) (vitPreB1 N ε w img)
                                                                                        theorem Proofs.ViTTiePoCGB.vitPreB3_apply (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                        vitPreB3 N ε w img = StableHLO.batchMap N (w.b3.fwdO ε) (vitPreB2 N ε w img)
                                                                                        theorem Proofs.ViTTiePoCGB.vitPreB4_apply (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                        vitPreB4 N ε w img = StableHLO.batchMap N (w.b4.fwdO ε) (vitPreB3 N ε w img)
                                                                                        theorem Proofs.ViTTiePoCGB.vitPreB5_apply (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                        vitPreB5 N ε w img = StableHLO.batchMap N (w.b5.fwdO ε) (vitPreB4 N ε w img)
                                                                                        theorem Proofs.ViTTiePoCGB.vitPreB6_apply (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                        vitPreB6 N ε w img = StableHLO.batchMap N (w.b6.fwdO ε) (vitPreB5 N ε w img)
                                                                                        theorem Proofs.ViTTiePoCGB.vitPreB7_apply (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                        vitPreB7 N ε w img = StableHLO.batchMap N (w.b7.fwdO ε) (vitPreB6 N ε w img)
                                                                                        theorem Proofs.ViTTiePoCGB.vitPreB8_apply (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                        vitPreB8 N ε w img = StableHLO.batchMap N (w.b8.fwdO ε) (vitPreB7 N ε w img)
                                                                                        theorem Proofs.ViTTiePoCGB.vitPreB9_apply (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                        vitPreB9 N ε w img = StableHLO.batchMap N (w.b9.fwdO ε) (vitPreB8 N ε w img)
                                                                                        theorem Proofs.ViTTiePoCGB.vitPreB10_apply (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                        vitPreB10 N ε w img = StableHLO.batchMap N (w.b10.fwdO ε) (vitPreB9 N ε w img)
                                                                                        theorem Proofs.ViTTiePoCGB.vitPreB11_apply (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                        vitPreB11 N ε w img = StableHLO.batchMap N (w.b11.fwdO ε) (vitPreB10 N ε w img)
                                                                                        theorem Proofs.ViTTiePoCGB.vitPreB12_apply (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                        vitPreB12 N ε w img = StableHLO.batchMap N (w.b12.fwdO ε) (vitPreB11 N ε w img)
                                                                                        noncomputable def Proofs.ViTTiePoCGB.vitSufB12 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                        Vec (N * (197 * 192)) → Vec (N * nC)

                                                                                        The net after block b12 — the head.

                                                                                        Equations
                                                                                        Instances For
                                                                                          noncomputable def Proofs.ViTTiePoCGB.vitSufB11 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                          Vec (N * (197 * 192)) → Vec (N * nC)

                                                                                          The net after block b11: block b12, then the rest.

                                                                                          Equations
                                                                                          Instances For
                                                                                            noncomputable def Proofs.ViTTiePoCGB.vitSufB10 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                            Vec (N * (197 * 192)) → Vec (N * nC)

                                                                                            The net after block b10: block b11, then the rest.

                                                                                            Equations
                                                                                            Instances For
                                                                                              noncomputable def Proofs.ViTTiePoCGB.vitSufB9 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                              Vec (N * (197 * 192)) → Vec (N * nC)

                                                                                              The net after block b9: block b10, then the rest.

                                                                                              Equations
                                                                                              Instances For
                                                                                                noncomputable def Proofs.ViTTiePoCGB.vitSufB8 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                                Vec (N * (197 * 192)) → Vec (N * nC)

                                                                                                The net after block b8: block b9, then the rest.

                                                                                                Equations
                                                                                                Instances For
                                                                                                  noncomputable def Proofs.ViTTiePoCGB.vitSufB7 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                                  Vec (N * (197 * 192)) → Vec (N * nC)

                                                                                                  The net after block b7: block b8, then the rest.

                                                                                                  Equations
                                                                                                  Instances For
                                                                                                    noncomputable def Proofs.ViTTiePoCGB.vitSufB6 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                                    Vec (N * (197 * 192)) → Vec (N * nC)

                                                                                                    The net after block b6: block b7, then the rest.

                                                                                                    Equations
                                                                                                    Instances For
                                                                                                      noncomputable def Proofs.ViTTiePoCGB.vitSufB5 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                                      Vec (N * (197 * 192)) → Vec (N * nC)

                                                                                                      The net after block b5: block b6, then the rest.

                                                                                                      Equations
                                                                                                      Instances For
                                                                                                        noncomputable def Proofs.ViTTiePoCGB.vitSufB4 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                                        Vec (N * (197 * 192)) → Vec (N * nC)

                                                                                                        The net after block b4: block b5, then the rest.

                                                                                                        Equations
                                                                                                        Instances For
                                                                                                          noncomputable def Proofs.ViTTiePoCGB.vitSufB3 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                                          Vec (N * (197 * 192)) → Vec (N * nC)

                                                                                                          The net after block b3: block b4, then the rest.

                                                                                                          Equations
                                                                                                          Instances For
                                                                                                            noncomputable def Proofs.ViTTiePoCGB.vitSufB2 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                                            Vec (N * (197 * 192)) → Vec (N * nC)

                                                                                                            The net after block b2: block b3, then the rest.

                                                                                                            Equations
                                                                                                            Instances For
                                                                                                              noncomputable def Proofs.ViTTiePoCGB.vitSufB1 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                                              Vec (N * (197 * 192)) → Vec (N * nC)

                                                                                                              The net after block b1: block b2, then the rest.

                                                                                                              Equations
                                                                                                              Instances For
                                                                                                                noncomputable def Proofs.ViTTiePoCGB.vitSufE (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) :
                                                                                                                Vec (N * (197 * 192)) → Vec (N * nC)

                                                                                                                The net after the patch embedding: block b1, then the rest.

                                                                                                                Equations
                                                                                                                Instances For
                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_embed (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (W : Kernel4 192 3 16 16) (b c : Vec 192) (q : Mat 197 192) :
                                                                                                                  vitNetB N ε { Wc := W, bc := b, cls := c, pos := q, b1 := w.b1, b2 := w.b2, b3 := w.b3, b4 := w.b4, b5 := w.b5, b6 := w.b6, b7 := w.b7, b8 := w.b8, b9 := w.b9, b10 := w.b10, b11 := w.b11, b12 := w.b12, γF := w.γF, βF := w.βF, Wcls := w.Wcls, bcls := w.bcls } img = vitSufE N ε w (StableHLO.batchMap N (patchEmbedFlat 3 224 224 16 196 192 W b c q) img)

                                                                                                                  The net with the patch embedding varied is the suffix after it at the varied embedding.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_b1 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (p : BlockParamsV 192 768) :
                                                                                                                  vitNetB N ε { Wc := w.Wc, bc := w.bc, cls := w.cls, pos := w.pos, b1 := p, b2 := w.b2, b3 := w.b3, b4 := w.b4, b5 := w.b5, b6 := w.b6, b7 := w.b7, b8 := w.b8, b9 := w.b9, b10 := w.b10, b11 := w.b11, b12 := w.b12, γF := w.γF, βF := w.βF, Wcls := w.Wcls, bcls := w.bcls } img = vitSufB1 N ε w (StableHLO.batchMap N (p.fwdO ε) (vitPreE N w img))

                                                                                                                  The net with block b1's weights varied is the suffix after it at the varied block.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_b2 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (p : BlockParamsV 192 768) :
                                                                                                                  vitNetB N ε { Wc := w.Wc, bc := w.bc, cls := w.cls, pos := w.pos, b1 := w.b1, b2 := p, b3 := w.b3, b4 := w.b4, b5 := w.b5, b6 := w.b6, b7 := w.b7, b8 := w.b8, b9 := w.b9, b10 := w.b10, b11 := w.b11, b12 := w.b12, γF := w.γF, βF := w.βF, Wcls := w.Wcls, bcls := w.bcls } img = vitSufB2 N ε w (StableHLO.batchMap N (p.fwdO ε) (vitPreB1 N ε w img))

                                                                                                                  The net with block b2's weights varied is the suffix after it at the varied block.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_b3 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (p : BlockParamsV 192 768) :
                                                                                                                  vitNetB N ε { Wc := w.Wc, bc := w.bc, cls := w.cls, pos := w.pos, b1 := w.b1, b2 := w.b2, b3 := p, b4 := w.b4, b5 := w.b5, b6 := w.b6, b7 := w.b7, b8 := w.b8, b9 := w.b9, b10 := w.b10, b11 := w.b11, b12 := w.b12, γF := w.γF, βF := w.βF, Wcls := w.Wcls, bcls := w.bcls } img = vitSufB3 N ε w (StableHLO.batchMap N (p.fwdO ε) (vitPreB2 N ε w img))

                                                                                                                  The net with block b3's weights varied is the suffix after it at the varied block.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_b4 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (p : BlockParamsV 192 768) :
                                                                                                                  vitNetB N ε { Wc := w.Wc, bc := w.bc, cls := w.cls, pos := w.pos, b1 := w.b1, b2 := w.b2, b3 := w.b3, b4 := p, b5 := w.b5, b6 := w.b6, b7 := w.b7, b8 := w.b8, b9 := w.b9, b10 := w.b10, b11 := w.b11, b12 := w.b12, γF := w.γF, βF := w.βF, Wcls := w.Wcls, bcls := w.bcls } img = vitSufB4 N ε w (StableHLO.batchMap N (p.fwdO ε) (vitPreB3 N ε w img))

                                                                                                                  The net with block b4's weights varied is the suffix after it at the varied block.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_b5 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (p : BlockParamsV 192 768) :
                                                                                                                  vitNetB N ε { Wc := w.Wc, bc := w.bc, cls := w.cls, pos := w.pos, b1 := w.b1, b2 := w.b2, b3 := w.b3, b4 := w.b4, b5 := p, b6 := w.b6, b7 := w.b7, b8 := w.b8, b9 := w.b9, b10 := w.b10, b11 := w.b11, b12 := w.b12, γF := w.γF, βF := w.βF, Wcls := w.Wcls, bcls := w.bcls } img = vitSufB5 N ε w (StableHLO.batchMap N (p.fwdO ε) (vitPreB4 N ε w img))

                                                                                                                  The net with block b5's weights varied is the suffix after it at the varied block.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_b6 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (p : BlockParamsV 192 768) :
                                                                                                                  vitNetB N ε { Wc := w.Wc, bc := w.bc, cls := w.cls, pos := w.pos, b1 := w.b1, b2 := w.b2, b3 := w.b3, b4 := w.b4, b5 := w.b5, b6 := p, b7 := w.b7, b8 := w.b8, b9 := w.b9, b10 := w.b10, b11 := w.b11, b12 := w.b12, γF := w.γF, βF := w.βF, Wcls := w.Wcls, bcls := w.bcls } img = vitSufB6 N ε w (StableHLO.batchMap N (p.fwdO ε) (vitPreB5 N ε w img))

                                                                                                                  The net with block b6's weights varied is the suffix after it at the varied block.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_b7 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (p : BlockParamsV 192 768) :
                                                                                                                  vitNetB N ε { Wc := w.Wc, bc := w.bc, cls := w.cls, pos := w.pos, b1 := w.b1, b2 := w.b2, b3 := w.b3, b4 := w.b4, b5 := w.b5, b6 := w.b6, b7 := p, b8 := w.b8, b9 := w.b9, b10 := w.b10, b11 := w.b11, b12 := w.b12, γF := w.γF, βF := w.βF, Wcls := w.Wcls, bcls := w.bcls } img = vitSufB7 N ε w (StableHLO.batchMap N (p.fwdO ε) (vitPreB6 N ε w img))

                                                                                                                  The net with block b7's weights varied is the suffix after it at the varied block.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_b8 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (p : BlockParamsV 192 768) :
                                                                                                                  vitNetB N ε { Wc := w.Wc, bc := w.bc, cls := w.cls, pos := w.pos, b1 := w.b1, b2 := w.b2, b3 := w.b3, b4 := w.b4, b5 := w.b5, b6 := w.b6, b7 := w.b7, b8 := p, b9 := w.b9, b10 := w.b10, b11 := w.b11, b12 := w.b12, γF := w.γF, βF := w.βF, Wcls := w.Wcls, bcls := w.bcls } img = vitSufB8 N ε w (StableHLO.batchMap N (p.fwdO ε) (vitPreB7 N ε w img))

                                                                                                                  The net with block b8's weights varied is the suffix after it at the varied block.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_b9 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (p : BlockParamsV 192 768) :
                                                                                                                  vitNetB N ε { Wc := w.Wc, bc := w.bc, cls := w.cls, pos := w.pos, b1 := w.b1, b2 := w.b2, b3 := w.b3, b4 := w.b4, b5 := w.b5, b6 := w.b6, b7 := w.b7, b8 := w.b8, b9 := p, b10 := w.b10, b11 := w.b11, b12 := w.b12, γF := w.γF, βF := w.βF, Wcls := w.Wcls, bcls := w.bcls } img = vitSufB9 N ε w (StableHLO.batchMap N (p.fwdO ε) (vitPreB8 N ε w img))

                                                                                                                  The net with block b9's weights varied is the suffix after it at the varied block.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_b10 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (p : BlockParamsV 192 768) :
                                                                                                                  vitNetB N ε { Wc := w.Wc, bc := w.bc, cls := w.cls, pos := w.pos, b1 := w.b1, b2 := w.b2, b3 := w.b3, b4 := w.b4, b5 := w.b5, b6 := w.b6, b7 := w.b7, b8 := w.b8, b9 := w.b9, b10 := p, b11 := w.b11, b12 := w.b12, γF := w.γF, βF := w.βF, Wcls := w.Wcls, bcls := w.bcls } img = vitSufB10 N ε w (StableHLO.batchMap N (p.fwdO ε) (vitPreB9 N ε w img))

                                                                                                                  The net with block b10's weights varied is the suffix after it at the varied block.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_b11 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (p : BlockParamsV 192 768) :
                                                                                                                  vitNetB N ε { Wc := w.Wc, bc := w.bc, cls := w.cls, pos := w.pos, b1 := w.b1, b2 := w.b2, b3 := w.b3, b4 := w.b4, b5 := w.b5, b6 := w.b6, b7 := w.b7, b8 := w.b8, b9 := w.b9, b10 := w.b10, b11 := p, b12 := w.b12, γF := w.γF, βF := w.βF, Wcls := w.Wcls, bcls := w.bcls } img = vitSufB11 N ε w (StableHLO.batchMap N (p.fwdO ε) (vitPreB10 N ε w img))

                                                                                                                  The net with block b11's weights varied is the suffix after it at the varied block.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_b12 (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (p : BlockParamsV 192 768) :
                                                                                                                  vitNetB N ε { Wc := w.Wc, bc := w.bc, cls := w.cls, pos := w.pos, b1 := w.b1, b2 := w.b2, b3 := w.b3, b4 := w.b4, b5 := w.b5, b6 := w.b6, b7 := w.b7, b8 := w.b8, b9 := w.b9, b10 := w.b10, b11 := w.b11, b12 := p, γF := w.γF, βF := w.βF, Wcls := w.Wcls, bcls := w.bcls } img = vitSufB12 N ε w (StableHLO.batchMap N (p.fwdO ε) (vitPreB11 N ε w img))

                                                                                                                  The net with block b12's weights varied is the suffix after it at the varied block.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_factor_head (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (a b : Vec 192) (W : Mat 192 nC) (bb : Vec nC) :
                                                                                                                  vitNetB N ε { Wc := w.Wc, bc := w.bc, cls := w.cls, pos := w.pos, b1 := w.b1, b2 := w.b2, b3 := w.b3, b4 := w.b4, b5 := w.b5, b6 := w.b6, b7 := w.b7, b8 := w.b8, b9 := w.b9, b10 := w.b10, b11 := w.b11, b12 := w.b12, γF := a, βF := b, Wcls := W, bcls := bb } img = StableHLO.batchMap N (vitHeadO ε a b W bb) (vitPreB12 N ε w img)

                                                                                                                  The net with the head varied is the head at the varied parameters.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_forward_eq_head (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                                                  vitNetB N ε w img = StableHLO.batchMap N (vitHeadO ε w.γF w.βF w.Wcls w.bcls) (vitPreB12 N ε w img)

                                                                                                                  The net's output is the head at block b12's output.

                                                                                                                  theorem Proofs.ViTTiePoCGB.vit_logitsB_eq (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) :
                                                                                                                  StableHLO.batchMap N (dense w.Wcls w.bcls) (StableHLO.batchMap N (StableHLO.clsSliceFlat 196 192) (StableHLO.batchMap N (fun (b : Vec ((196 + 1) * 192)) => Mat.flatten fun (r : Fin (196 + 1)) => layerNormVec 192 ε w.γF w.βF (Mat.unflatten b r)) (vitPreB12 N ε w img))) = vitNetB N ε w img

                                                                                                                  The logits the tie's loss cotangent reads are vitNetB's. The tie spells the head as three batched ops (batchMap_comp).

                                                                                                                  def Proofs.ViTTiePoCGB.ViTNetLossTiedGB (xN aN epsStr cotN : String) (N : ℕ) {nC : ℕ} (ε : ℝ) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (L : Vec (N * nC) → Vec 1) (g : Vec (N * nC)) :

                                                                                                                  Every ViT-Tiny parameter gradient node is the derivative of L in that parameter, for a loss L of the logits and g the cotangent the chain starts from: the 200 nodes vit_net_tiedGB ties, each at the cotangent the tie threads to it from g, stated against L of vitNetB with that one parameter varied.

                                                                                                                  Equations
                                                                                                                  • One or more equations did not get rendered due to their size.
                                                                                                                  Instances For
                                                                                                                    theorem Proofs.ViTTiePoCGB.vit_net_lossGrad (xN aN epsStr cotN : String) (N : ℕ) {nC : ℕ} (ε : ℝ) (hε : 0 < ε) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) {L : Vec (N * nC) → Vec 1} {g : Vec (N * nC)} (hL : HasGradAt L (vitNetB N ε w img) g) :
                                                                                                                    ViTNetLossTiedGB xN aN epsStr cotN N ε w img L g

                                                                                                                    Every ViT-Tiny parameter gradient node is the derivative of the loss in that parameter. For any loss L of the logits with gradient g at the net's output, each of the 200 nodes vit_net_tiedGB ties — at the same cotangent — is ∂L/∂θ of the WHOLE net, vitNetB with that one parameter varied (a patch-embedding field, a block's record w.bk := p with one slot changed, or a head field).

                                                                                                                    Hypothesis: 0 < ε, the LayerNorms' (the tie itself needs none). The loss enters only through hL; vit_net_lossGrad_smoothedCE discharges it for the loss the artifacts ship.

                                                                                                                    theorem Proofs.ViTTiePoCGB.vit_net_lossGrad_smoothedCE (xN aN epsStr cotN aStr negAK bStr logN ohN : String) (N : ℕ) {nC : ℕ} (hK : 0 < nC) (ε α B : ℝ) (hε : 0 < ε) (w : ViTTiePoC.ViTTieWeights nC) (img : Vec (N * (3 * 224 * 224))) (t : Vec (N * nC)) (ht : ∀ (n : Fin N), ∑ k : Fin nC, StableHLO.batchSlice N nC t n k = 1) :
                                                                                                                    ViTNetLossTiedGB xN aN epsStr cotN N ε w img (smoothedBatchLossDiv N nC α B t) (StableHLO.den (smoothedLossCotGraphDiv N nC α B aStr negAK bStr logN ohN (vitNetB N ε w img) t))

                                                                                                                    The loss the artifacts ship: every node is the derivative of the batched label-smoothed cross-entropy smoothedBatchLossDiv, g the softmaxDiv cotangent the render emits — the tie's own g, whose logits are vitNetB N ε w img (vit_logitsB_eq).