Documentation

LeanMlir.Proofs.Nets.ViT.ViTFwdDrop

ViT with stochastic depth — the batched forward graph and its faithfulness #

ViTDepthK states the ViT forward one example at a time (vitFwdGraphKMHV_faithful). The *drop* artifacts — vit_drop_fwd, vitin_drop_fwd, vitsin_drop_fwd and the forward half of every *drop* train step, the book's vitin_emadp128x4wxclipdropbf16 among them — add stochastic depth: two %dp<i> mask inputs per block, a per-example scale dropPath on the attention branch (after the out-dense, before the first skip add) and on the MLP branch (after fc2, before the second), as ViTRenderB.vBlockFwdB emits them. Block i's attention site is %dp<2i> and its MLP site %dp<2i+1> (the render's vitSiteIdx).

Why this graph is batched. A drop mask is per EXAMPLE, so no per-example node can carry it: in the per-example graph a node denotes one example and the batch is lifted outside the AST, which is why the render that writes these artifacts is the batched one. So the statement here is at the batched index B, over the batched tokens the render emits (.batchOp of the row forms, .matmulFB, .scaleB, .addVB, .dropPathB), and it says what the per-example graph says one level up: example t of the batched graph's output is the per-example forward at example t's input, with example t's mask entries as its drop scalars.

The graph uses the render's SSA names (%wConv, b<i>_, %gF, %Wc, …). Like vitFwdGraphKMHV, it is not tied to the artifact text: the render names each shared intermediate once (LN1's output feeds Q, K and V), where a graph term repeats the subterm. These artifacts are f32 where they are forwards; the bf16 train steps' forward differs by the matmul roundings, and their backward through the drop sites is outside this statement, as it is outside ViTStepTieGB.

References #

noncomputable def Proofs.blockVDrop (Np1 heads d_head mlpDim : ℕ) (ε a m : ℝ) (p : BlockParamsV (heads * d_head) mlpDim) :
Mat Np1 (heads * d_head) → Mat Np1 (heads * d_head)

One vector-LN block with its two drop scalars: transformerBlockV with the attention branch scaled by a and the MLP branch by m before their skip adds.

Equations
  • One or more equations did not get rendered due to their size.
Instances For
    theorem Proofs.blockVDrop_one (Np1 heads d_head mlpDim : ℕ) (ε : ℝ) (p : BlockParamsV (heads * d_head) mlpDim) :
    blockVDrop Np1 heads d_head mlpDim ε 1 1 p = blockV Np1 heads d_head mlpDim ε p

    At unit scalars the drop block is blockV.

    noncomputable def Proofs.vitBlockSpelledMHVDrop (Np1 heads d mlpDim : ℕ) (ε a mk : ℝ) (p : BlockParamsV (heads * d) mlpDim) (X : Mat Np1 (heads * d)) :
    Mat Np1 (heads * d)

    The drop block spelled as the graph emits it — vitBlockSpelledMHV with the two scalars on the branches.

    Equations
    • One or more equations did not get rendered due to their size.
    Instances For
      theorem Proofs.vitBlockSpelledMHVDrop_eq (Np1 heads d mlpDim : ℕ) (ε a mk : ℝ) (p : BlockParamsV (heads * d) mlpDim) (X : Mat Np1 (heads * d)) :
      vitBlockSpelledMHVDrop Np1 heads d mlpDim ε a mk p X = blockVDrop Np1 heads d mlpDim ε a mk p X

      The spelled drop block IS blockVDrop (vitBlockSpelledMHV_eq's proof, at the scalars).

      noncomputable def Proofs.vitBodyKVDrop (Np1 heads d_head mlpDim : ℕ) (ε : ℝ) (k : ℕ) :
      (Fin k → BlockParamsV (heads * d_head) mlpDim) → (Fin k → ℝ × ℝ) → Mat Np1 (heads * d_head) → Mat Np1 (heads * d_head)

      The depth-k tower with drop scalars — vitBodyKV with block i at sd i (attention, MLP).

      Equations
      • One or more equations did not get rendered due to their size.
      • Proofs.vitBodyKVDrop Np1 heads d_head mlpDim ε 0 x_3 x_4 = fun (A : Proofs.Mat Np1 (heads * d_head)) => A
      Instances For
        theorem Proofs.vitBodyKVDrop_ones (Np1 heads d_head mlpDim : ℕ) (ε : ℝ) (k : ℕ) (ps : Fin k → BlockParamsV (heads * d_head) mlpDim) :
        (vitBodyKVDrop Np1 heads d_head mlpDim ε k ps fun (x : Fin k) => (1, 1)) = vitBodyKV Np1 heads d_head mlpDim ε k ps

        At unit scalars the tower is vitBodyKV.

        noncomputable def Proofs.vitForwardKVDrop (ic H W patchSize N mlpDim heads d_head nClasses k : ℕ) (W_conv : Kernel4 (heads * d_head) ic patchSize patchSize) (b_conv cls_token : Vec (heads * d_head)) (pos_embed : Mat (N + 1) (heads * d_head)) (ε : ℝ) (ps : Fin k → BlockParamsV (heads * d_head) mlpDim) (sd : Fin k → ℝ × ℝ) (γF βF : Vec (heads * d_head)) (Wcls : Mat (heads * d_head) nClasses) (bcls : Vec nClasses) :
        Vec (ic * H * W) → Vec nClasses

        The depth-k ViT forward with drop scalars, one example: vitForwardKV with the tower at sd.

        Equations
        • One or more equations did not get rendered due to their size.
        Instances For
          theorem Proofs.vitForwardKVDrop_ones (ic H W patchSize N mlpDim heads d_head nClasses k : ℕ) (W_conv : Kernel4 (heads * d_head) ic patchSize patchSize) (b_conv cls_token : Vec (heads * d_head)) (pos_embed : Mat (N + 1) (heads * d_head)) (ε : ℝ) (ps : Fin k → BlockParamsV (heads * d_head) mlpDim) (γF βF : Vec (heads * d_head)) (Wcls : Mat (heads * d_head) nClasses) (bcls : Vec nClasses) :
          vitForwardKVDrop ic H W patchSize N mlpDim heads d_head nClasses k W_conv b_conv cls_token pos_embed ε ps (fun (x : Fin k) => (1, 1)) γF βF Wcls bcls = vitForwardKV ic H W patchSize N mlpDim heads d_head nClasses k W_conv b_conv cls_token pos_embed ε ps γF βF Wcls bcls

          At unit scalars the forward is vitForwardKV.

          noncomputable def Proofs.vitForwardKVDropB (B ic H W patchSize N mlpDim heads d_head nClasses k : ℕ) (W_conv : Kernel4 (heads * d_head) ic patchSize patchSize) (b_conv cls_token : Vec (heads * d_head)) (pos_embed : Mat (N + 1) (heads * d_head)) (ε : ℝ) (ps : Fin k → BlockParamsV (heads * d_head) mlpDim) (sdA sdM : Fin k → Vec B) (γF βF : Vec (heads * d_head)) (Wcls : Mat (heads * d_head) nClasses) (bcls : Vec nClasses) :
          Vec (B * (ic * H * W)) → Vec (B * nClasses)

          The batched ViT forward with stochastic depth: example t is vitForwardKVDrop at example t's input, with (sdA i t, sdM i t) as block i's drop scalars. sdA i / sdM i are the render's per-example masks %dp<2i> / %dp<2i+1>.

          Equations
          • One or more equations did not get rendered due to their size.
          Instances For
            theorem Proofs.vitForwardKVDropB_ones (B ic H W patchSize N mlpDim heads d_head nClasses k : ℕ) (W_conv : Kernel4 (heads * d_head) ic patchSize patchSize) (b_conv cls_token : Vec (heads * d_head)) (pos_embed : Mat (N + 1) (heads * d_head)) (ε : ℝ) (ps : Fin k → BlockParamsV (heads * d_head) mlpDim) (γF βF : Vec (heads * d_head)) (Wcls : Mat (heads * d_head) nClasses) (bcls : Vec nClasses) :
            vitForwardKVDropB B ic H W patchSize N mlpDim heads d_head nClasses k W_conv b_conv cls_token pos_embed ε ps (fun (x : Fin k) (x_1 : Fin B) => 1) (fun (x : Fin k) (x_1 : Fin B) => 1) γF βF Wcls bcls = StableHLO.batchMap B (vitForwardKV ic H W patchSize N mlpDim heads d_head nClasses k W_conv b_conv cls_token pos_embed ε ps γF βF Wcls bcls)

            At the all-ones masks the batched forward is vitForwardKV lifted, exactly — the masks the driver passes to the forward artifacts at eval.

            theorem Proofs.StableHLO.batchSlice_den_batchOp {B a b : ℕ} (op : BatchableOp a b) (e : SHlo (B * a)) (t : Fin B) :
            batchSlice B b (den (SHlo.batchOp op e)) t = denOp op (batchSlice B a (den e) t)

            A batched descriptor token, sliced at example t, is its per-example map at the slice.

            theorem Proofs.StableHLO.batchSlice_den_matmulFB {B m k n : ℕ} (a : SHlo (B * (m * k))) (b : SHlo (B * (k * n))) (t : Fin B) :
            batchSlice B (m * n) (den (a.matmulFB b)) t = matMulFlat m k n (batchSlice B (m * k) (den a) t) (batchSlice B (k * n) (den b) t)

            matmulFB sliced at example t multiplies example t's two operands (den_matmulFB_per_example, as a slice).

            theorem Proofs.StableHLO.batchSlice_den_scaleB {B n : ℕ} (sS : String) (s : ℝ) (e : SHlo (B * n)) (t : Fin B) :
            batchSlice B n (den (SHlo.scaleB sS s e)) t = fun (i : Fin n) => batchSlice B n (den e) t i * s
            theorem Proofs.StableHLO.batchSlice_den_addVB {B n : ℕ} (a b : SHlo (B * n)) (t : Fin B) :
            batchSlice B n (den (a.addVB b)) t = fun (i : Fin n) => batchSlice B n (den a) t i + batchSlice B n (den b) t i
            theorem Proofs.StableHLO.batchSlice_den_dropPathB {B n : ℕ} (mN : String) (s : Vec B) (e : SHlo (B * n)) (t : Fin B) :
            batchSlice B n (den (SHlo.dropPathB mN s e)) t = fun (i : Fin n) => s t * batchSlice B n (den e) t i

            A drop site, sliced at example t, scales by example t's mask entry — the per-example content of dropPathB.

            def Proofs.StableHLO.headsSumGB {B n hm1 : ℕ} :
            (Fin (hm1 + 1) → SHlo (B * n)) → SHlo (B * n)

            Left-assoc addVB fold of one batched graph per head — headsSumG at the batched index, in the render's order (acc := pd₀, then addVB acc pd_h).

            Equations
            Instances For
              theorem Proofs.StableHLO.batchSlice_den_headsSumGB {B n hm1 : ℕ} (f : Fin (hm1 + 1) → SHlo (B * n)) (t : Fin B) :
              batchSlice B n (den (headsSumGB f)) t = fun (j : Fin n) => ∑ h : Fin (hm1 + 1), batchSlice B n (den (f h)) t j

              The batched head fold, sliced at example t, is the sum over heads of the slices.

              theorem Proofs.StableHLO.scale_flat_right {m n : ℕ} (s : ℝ) (A : Mat m n) :
              (fun (i : Fin (m * n)) => A.flatten i * s) = Mat.flatten fun (r : Fin m) (c : Fin n) => s * A r c

              The right-multiplied scale commutes with flattening (scale_flat, operands swapped — the batched scaleB multiplies on the right).

              theorem Proofs.StableHLO.scale_flat_pt {m n : ℕ} (s : ℝ) (A : Mat m n) (j : Fin (m * n)) :
              s * A.flatten j = Mat.flatten (fun (r : Fin m) (c : Fin n) => s * A r c) j

              A drop site's per-example scale commutes with flattening, pointwise.

              def Proofs.StableHLO.vitBlockGraphBDrop {B Np1 hm1 d mlpDim : ℕ} (pfx epsStr sStr mA mM : String) (ε s : ℝ) (p : BlockParamsV ((hm1 + 1) * d) mlpDim) (a m : Vec B) (x : SHlo (B * (Np1 * ((hm1 + 1) * d)))) :
              SHlo (B * (Np1 * ((hm1 + 1) * d)))

              One batched ViT block with its two drop sites, node for node ViTRenderB.vBlockFwdB: vector-LN 1 (lnRow → rowScale → rowBias), Q/K/V denseRow, per head headSlice → transpose → matmulFB → scaleB → softmaxRow → matmulFB → headPad, summed by headsSumGB; out denseRow, dropPathB at mA, skip addVB; vector-LN 2, fc1, GELU, fc2, dropPathB at mM, skip addVB.

              Equations
              • One or more equations did not get rendered due to their size.
              Instances For
                theorem Proofs.StableHLO.vitBlockGraphBDrop_slice {B Np1 hm1 d mlpDim : ℕ} (pfx epsStr sStr mA mM : String) (ε : ℝ) (p : BlockParamsV ((hm1 + 1) * d) mlpDim) (a m : Vec B) (e : SHlo (B * (Np1 * ((hm1 + 1) * d)))) (t : Fin B) (A : Mat Np1 ((hm1 + 1) * d)) (hA : batchSlice B (Np1 * ((hm1 + 1) * d)) (den e) t = A.flatten) :
                batchSlice B (Np1 * ((hm1 + 1) * d)) (den (vitBlockGraphBDrop pfx epsStr sStr mA mM ε (sdpaScale d) p a m e)) t = (vitBlockSpelledMHVDrop Np1 (hm1 + 1) d mlpDim ε (a t) (m t) p A).flatten

                Example t of the batched drop block is the spelled drop block at example t's input and mask entries.

                def Proofs.StableHLO.vitBodyGraphBDrop {B Np1 hm1 d mlpDim : ℕ} (epsStr sStr : String) (ε s : ℝ) (base k : ℕ) :
                (Fin k → BlockParamsV ((hm1 + 1) * d) mlpDim) → (Fin k → Vec B) → (Fin k → Vec B) → SHlo (B * (Np1 * ((hm1 + 1) * d))) → SHlo (B * (Np1 * ((hm1 + 1) * d)))

                The batched depth-k tower with its drop sites — block base + i carries the prefix b{base+i}_ and reads %dp<2(base+i)> / %dp<2(base+i)+1>, as ViTRenderB.vitFwd12B names them (vitSiteIdx).

                Equations
                Instances For
                  theorem Proofs.StableHLO.vitBodyGraphBDrop_slice {B Np1 hm1 d mlpDim : ℕ} (epsStr sStr : String) (ε : ℝ) (base k : ℕ) (ps : Fin k → BlockParamsV ((hm1 + 1) * d) mlpDim) (sdA sdM : Fin k → Vec B) (e : SHlo (B * (Np1 * ((hm1 + 1) * d)))) (t : Fin B) (A : Mat Np1 ((hm1 + 1) * d)) :
                  batchSlice B (Np1 * ((hm1 + 1) * d)) (den e) t = A.flatten → batchSlice B (Np1 * ((hm1 + 1) * d)) (den (vitBodyGraphBDrop epsStr sStr ε (sdpaScale d) base k ps sdA sdM e)) t = (vitBodyKVDrop Np1 (hm1 + 1) d mlpDim ε k ps (fun (i : Fin k) => (sdA i t, sdM i t)) A).flatten

                  Example t of the batched tower is the drop tower at example t's input and mask entries — by induction on k, one vitBlockGraphBDrop_slice per block.

                  def Proofs.StableHLO.vitFwdGraphBDrop {B ic H W P N hm1 d mlpDim nClasses : ℕ} (epsStr sStr : String) (ε s : ℝ) (Wc : Kernel4 ((hm1 + 1) * d) ic P P) (bc cls : Vec ((hm1 + 1) * d)) (pos : Mat (N + 1) ((hm1 + 1) * d)) (k : ℕ) (ps : Fin k → BlockParamsV ((hm1 + 1) * d) mlpDim) (sdA sdM : Fin k → Vec B) (γF βF : Vec ((hm1 + 1) * d)) (Wcls : Mat ((hm1 + 1) * d) nClasses) (bcls : Vec nClasses) (x : Vec (B * (ic * H * W))) :
                  SHlo (B * nClasses)

                  The batched ViT forward graph with stochastic depth — the typed form of ViTRenderB.vitFwd12B … (sd := true) at depth k: batched patch embed over %x, the drop tower, final vector-LN, CLS slice, dense head.

                  Equations
                  • One or more equations did not get rendered due to their size.
                  Instances For
                    theorem Proofs.StableHLO.vitFwdGraphBDrop_slice {B ic H W patchSize N hm1 d mlpDim nClasses : ℕ} (epsStr sStr : String) (Wc : Kernel4 ((hm1 + 1) * d) ic patchSize patchSize) (bc cls : Vec ((hm1 + 1) * d)) (pos : Mat (N + 1) ((hm1 + 1) * d)) (ε : ℝ) (k : ℕ) (ps : Fin k → BlockParamsV ((hm1 + 1) * d) mlpDim) (sdA sdM : Fin k → Vec B) (γF βF : Vec ((hm1 + 1) * d)) (Wcls : Mat ((hm1 + 1) * d) nClasses) (bcls : Vec nClasses) (x : Vec (B * (ic * H * W))) (t : Fin B) :
                    batchSlice B nClasses (den (vitFwdGraphBDrop epsStr sStr ε (sdpaScale d) Wc bc cls pos k ps sdA sdM γF βF Wcls bcls x)) t = vitForwardKVDrop ic H W patchSize N mlpDim (hm1 + 1) d nClasses k Wc bc cls pos ε ps (fun (i : Fin k) => (sdA i t, sdM i t)) γF βF Wcls bcls (batchSlice B (ic * H * W) x t)

                    Example t of the batched drop graph is the per-example drop forward at example t's input, with example t's mask entries — for every depth k.

                    theorem Proofs.StableHLO.vitFwdGraphBDrop_faithful {B ic H W patchSize N hm1 d mlpDim nClasses : ℕ} (epsStr sStr : String) (Wc : Kernel4 ((hm1 + 1) * d) ic patchSize patchSize) (bc cls : Vec ((hm1 + 1) * d)) (pos : Mat (N + 1) ((hm1 + 1) * d)) (ε : ℝ) (k : ℕ) (ps : Fin k → BlockParamsV ((hm1 + 1) * d) mlpDim) (sdA sdM : Fin k → Vec B) (γF βF : Vec ((hm1 + 1) * d)) (Wcls : Mat ((hm1 + 1) * d) nClasses) (bcls : Vec nClasses) (x : Vec (B * (ic * H * W))) :
                    den (vitFwdGraphBDrop epsStr sStr ε (sdpaScale d) Wc bc cls pos k ps sdA sdM γF βF Wcls bcls x) = vitForwardKVDropB B ic H W patchSize N mlpDim (hm1 + 1) d nClasses k Wc bc cls pos ε ps sdA sdM γF βF Wcls bcls x

                    The batched ViT forward graph with stochastic depth denotes vitForwardKVDropB — at every depth, every pair of mask families and every input.