Documentation

LeanMlir.Proofs.Nets.ViT.ViTFwdGraph

ViT forward graph — ch10 close Item A (planning/archive/vit_close.md) #

Two halves, both living here because StableHLO.lean cannot import Attention.lean (Attention is the proof capstone; StableHLO is the token layer — the ViT den helpers there are local re-spellings, tied back to the proven Attention forms in THIS file):

  1. vitForward2 — the representative distinct-param 2-block ViT forward at the Vec level: classifier ∘ finalLN ∘ block₂ ∘ block₁ ∘ patchEmbed. The proven transformerTower/vit_full share ONE param tuple across blocks; a train step needs distinct per-block params, so the 2-block forward is composed here from transformerBlock directly (the tower proof does exactly this composition — with shared params). vitForward2_has_vjp is the whole-net VJP: vjp_comp chains patchEmbed_flat_has_vjp, two bridged transformerBlock_has_vjp_mat witnesses, the bridged per-token final-LN, and classifier_flat_has_vjp. UNCONDITIONAL except 0 < ε (all-smooth — softmax/GELU/LN have no kinks).

  2. vitFwdGraph — the typed SHlo forward graph over the ch10 token vocabulary (patchEmbedF/lnRowF/denseRowF/matmulF/transposeF/ scaleF/softmaxRowF/geluF/addV/clsSliceF), heads = 1 (SDPA = three matmuls + a row-softmax — the representative granularity trade, like ConvNeXt's 2-block/1×1-stem). vitFwdGraph_faithful: its denotation IS vitForward2 at heads := 1 — the ViT analogue of convNextFwdGraph_faithful.

noncomputable def Proofs.vitForward2 (ic H W patchSize N mlpDim heads d_head nClasses : ) (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)) (ε γ1₁ β1₁ : ) (Wq₁ Wk₁ Wv₁ Wo₁ : Mat (heads * d_head) (heads * d_head)) (bq₁ bk₁ bv₁ bo₁ : Vec (heads * d_head)) (γ2₁ β2₁ : ) (Wfc1₁ : Mat (heads * d_head) mlpDim) (bfc1₁ : Vec mlpDim) (Wfc2₁ : Mat mlpDim (heads * d_head)) (bfc2₁ : Vec (heads * d_head)) (γ1₂ β1₂ : ) (Wq₂ Wk₂ Wv₂ Wo₂ : Mat (heads * d_head) (heads * d_head)) (bq₂ bk₂ bv₂ bo₂ : Vec (heads * d_head)) (γ2₂ β2₂ : ) (Wfc1₂ : Mat (heads * d_head) mlpDim) (bfc1₂ : Vec mlpDim) (Wfc2₂ : Mat mlpDim (heads * d_head)) (bfc2₂ : Vec (heads * d_head)) (γF βF : ) (Wcls : Mat (heads * d_head) nClasses) (bcls : Vec nClasses) :
Vec (ic * H * W)Vec nClasses

Distinct-param 2-block ViT forward (the ch10 representative):

patchEmbed (stride-P conv + CLS + pos) → block₁ → block₂ → final-LN (per-token, scalar γ/β) → CLS slice → dense head

Generic dims; the two transformerBlocks carry distinct parameter sets (…₁/…₂) — beyond the shared-param transformerTower witness, composed from the same proven block VJP. One shared LN ε across all five LN sites (the proof convention, as in vit_full).

Equations
  • One or more equations did not get rendered due to their size.
Instances For
    noncomputable def Proofs.vitForward2_has_vjp (ic H W patchSize N mlpDim heads d_head nClasses : ) (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)) (ε : ) ( : 0 < ε) (γ1₁ β1₁ : ) (Wq₁ Wk₁ Wv₁ Wo₁ : Mat (heads * d_head) (heads * d_head)) (bq₁ bk₁ bv₁ bo₁ : Vec (heads * d_head)) (γ2₁ β2₁ : ) (Wfc1₁ : Mat (heads * d_head) mlpDim) (bfc1₁ : Vec mlpDim) (Wfc2₁ : Mat mlpDim (heads * d_head)) (bfc2₁ : Vec (heads * d_head)) (γ1₂ β1₂ : ) (Wq₂ Wk₂ Wv₂ Wo₂ : Mat (heads * d_head) (heads * d_head)) (bq₂ bk₂ bv₂ bo₂ : Vec (heads * d_head)) (γ2₂ β2₂ : ) (Wfc1₂ : Mat (heads * d_head) mlpDim) (bfc1₂ : Vec mlpDim) (Wfc2₂ : Mat mlpDim (heads * d_head)) (bfc2₂ : Vec (heads * d_head)) (γF βF : ) (Wcls : Mat (heads * d_head) nClasses) (bcls : Vec nClasses) :
    HasVJP (vitForward2 ic H W patchSize N mlpDim heads d_head nClasses W_conv b_conv cls_token pos_embed ε γ1₁ β1₁ Wq₁ Wk₁ Wv₁ Wo₁ bq₁ bk₁ bv₁ bo₁ γ2₁ β2₁ Wfc1₁ bfc1₁ Wfc2₁ bfc2₁ γ1₂ β1₂ Wq₂ Wk₂ Wv₂ Wo₂ bq₂ bk₂ bv₂ bo₂ γ2₂ β2₂ Wfc1₂ bfc1₂ Wfc2₂ bfc2₂ γF βF Wcls bcls)

    Whole-net VJP for the distinct-param 2-block ViT (global). All-smooth, so the only hypothesis is the LayerNorm positivity 0 < ε — joins vit_full_has_vjp/convnext_has_vjp as an unconditional whole-network VJP holding at every input. Four vjp_comp steps glueing patchEmbed_flat_has_vjp, two bridged distinct-param transformerBlock_has_vjp_mat witnesses, the bridged per-token final-LN, and classifier_flat_has_vjp.

    Equations
    • One or more equations did not get rendered due to their size.
    Instances For
      theorem Proofs.vitForward2_has_vjp_correct (ic H W patchSize N mlpDim heads d_head nClasses : ) (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)) (ε : ) ( : 0 < ε) (γ1₁ β1₁ : ) (Wq₁ Wk₁ Wv₁ Wo₁ : Mat (heads * d_head) (heads * d_head)) (bq₁ bk₁ bv₁ bo₁ : Vec (heads * d_head)) (γ2₁ β2₁ : ) (Wfc1₁ : Mat (heads * d_head) mlpDim) (bfc1₁ : Vec mlpDim) (Wfc2₁ : Mat mlpDim (heads * d_head)) (bfc2₁ : Vec (heads * d_head)) (γ1₂ β1₂ : ) (Wq₂ Wk₂ Wv₂ Wo₂ : Mat (heads * d_head) (heads * d_head)) (bq₂ bk₂ bv₂ bo₂ : Vec (heads * d_head)) (γ2₂ β2₂ : ) (Wfc1₂ : Mat (heads * d_head) mlpDim) (bfc1₂ : Vec mlpDim) (Wfc2₂ : Mat mlpDim (heads * d_head)) (bfc2₂ : Vec (heads * d_head)) (γF βF : ) (Wcls : Mat (heads * d_head) nClasses) (bcls : Vec nClasses) (x : Vec (ic * H * W)) (dy : Vec nClasses) (i : Fin (ic * H * W)) :
      (vitForward2_has_vjp ic H W patchSize N mlpDim heads d_head nClasses W_conv b_conv cls_token pos_embed ε γ1₁ β1₁ Wq₁ Wk₁ Wv₁ Wo₁ bq₁ bk₁ bv₁ bo₁ γ2₁ β2₁ Wfc1₁ bfc1₁ Wfc2₁ bfc2₁ γ1₂ β1₂ Wq₂ Wk₂ Wv₂ Wo₂ bq₂ bk₂ bv₂ bo₂ γ2₂ β2₂ Wfc1₂ bfc1₂ Wfc2₂ bfc2₂ γF βF Wcls bcls).backward x dy i = j : Fin nClasses, pdiv (vitForward2 ic H W patchSize N mlpDim heads d_head nClasses W_conv b_conv cls_token pos_embed ε γ1₁ β1₁ Wq₁ Wk₁ Wv₁ Wo₁ bq₁ bk₁ bv₁ bo₁ γ2₁ β2₁ Wfc1₁ bfc1₁ Wfc2₁ bfc2₁ γ1₂ β1₂ Wq₂ Wk₂ Wv₂ Wo₂ bq₂ bk₂ bv₂ bo₂ γ2₂ β2₂ Wfc1₂ bfc1₂ Wfc2₂ bfc2₂ γF βF Wcls bcls) x i j * dy j

      Public correctness theorem for vitForward2_has_vjp — the distinct-param 2-block ViT's backward equals the pdiv-contracted Jacobian (Jacobian-transpose applied to the cotangent), at every input. The ch10 analogue of convnext_has_vjp_correct.

      theorem Proofs.mhsa_layer_one_head (Np1 d : ) (Wq Wk Wv Wo : Mat (1 * d) (1 * d)) (bq bk bv bo : Vec (1 * d)) (X : Mat Np1 (1 * d)) :
      mhsa_layer Np1 1 d Wq Wk Wv Wo bq bk bv bo X = fun (n : Fin Np1) => dense Wo bo ((rowSoftmax fun (i j : Fin Np1) => sdpa_scale d * Mat.mul (fun (r : Fin Np1) (c : Fin (1 * d)) => dense Wq bq (X r) c) (Mat.transpose fun (r : Fin Np1) (c : Fin (1 * d)) => dense Wk bk (X r) c) i j).mul (fun (r : Fin Np1) (c : Fin (1 * d)) => dense Wv bv (X r) c) n)

      MHSA at heads = 1 is three matmuls + a row-softmax. The per-head slice/concat plumbing of mhsa_layer collapses (the head axis is Fin 1), leaving exactly the ch10 graph spelling: Q/K/V per-token dense → Q·Kᵀ·1/√d → row-softmax → P·V → output dense. This is the load-bearing tie for vitFwdGraph_faithful.

      noncomputable def Proofs.vitBlockSpelled (Np1 d mlpDim : ) (ε γ1 β1 : ) (Wq Wk Wv Wo : Mat (1 * d) (1 * d)) (bq bk bv bo : Vec (1 * d)) (γ2 β2 : ) (Wfc1 : Mat (1 * d) mlpDim) (bfc1 : Vec mlpDim) (Wfc2 : Mat mlpDim (1 * d)) (bfc2 : Vec (1 * d)) (X : Mat Np1 (1 * d)) :
      Mat Np1 (1 * d)

      The ch10 spelled pre-norm transformer block at heads = 1 (Mat level) — the exact op sequence vitBlockGraph denotes: LN₁ → Q/K/V per-token dense → Q·Kᵀ·1/√d → row-softmax → P·V → output dense → +res → LN₂ → fc1 → GELU → fc2 → +res. Equals transformerBlock at one head (vitBlockSpelled_eq).

      Equations
      • One or more equations did not get rendered due to their size.
      Instances For
        theorem Proofs.vitBlockSpelled_eq (Np1 d mlpDim : ) (ε γ1 β1 : ) (Wq Wk Wv Wo : Mat (1 * d) (1 * d)) (bq bk bv bo : Vec (1 * d)) (γ2 β2 : ) (Wfc1 : Mat (1 * d) mlpDim) (bfc1 : Vec mlpDim) (Wfc2 : Mat mlpDim (1 * d)) (bfc2 : Vec (1 * d)) (X : Mat Np1 (1 * d)) :
        vitBlockSpelled Np1 d mlpDim ε γ1 β1 Wq Wk Wv Wo bq bk bv bo γ2 β2 Wfc1 bfc1 Wfc2 bfc2 X = transformerBlock Np1 1 d mlpDim ε γ1 β1 Wq Wk Wv Wo bq bk bv bo γ2 β2 Wfc1 bfc1 Wfc2 bfc2 X

        The spelled block IS transformerBlock at one head. The sublayer/residual structure matches definitionally once mhsa_layer_one_head collapses the per-head plumbing.

        def Proofs.StableHLO.vitBlockGraph {Np1 D mlpDim : } (pfx epsStr sStr : String) (ε s γ1 β1 : ) (Wq Wk Wv Wo : Mat D D) (bq bk bv bo : Vec D) (γ2 β2 : ) (Wfc1 : Mat D mlpDim) (bfc1 : Vec mlpDim) (Wfc2 : Mat mlpDim D) (bfc2 : Vec D) (x : SHlo (Np1 * D)) :
        SHlo (Np1 * D)

        One spelled pre-norm transformer block over the ch10 tokens (heads = 1): lnRowF → Q/K/V denseRowFmatmulF(Q, transposeF K) → scaleFsoftmaxRowFmatmulF(P, V) → output denseRowFaddV residual → lnRowF → fc1 → geluF → fc2 → addV residual. Generic D; the faithfulness theorem instantiates D := 1 * d, s := sdpa_scale d.

        Equations
        • One or more equations did not get rendered due to their size.
        Instances For

          Flat ↔ Mat commutation bridges #

          Each ch10 den helper applied to a Mat.flatten is the flatten of the corresponding Mat-level op (the Mat.unflatten_flatten round-trip cancels); the pointwise ops (scaleF/geluF/addV) commute with flattening definitionally. Public — ViTChainClose reuses them to tie the matmul-spelled SDPA backward to the proven closed forms.

          theorem Proofs.StableHLO.rowLNFlat_flat {m n : } (ε γ β : ) (A : Mat m n) :
          rowLNFlat m n ε γ β A.flatten = Mat.flatten fun (r : Fin m) => layerNormForward n ε γ β (A r)
          theorem Proofs.StableHLO.rowDenseFlat_flat {N a c : } (W : Mat a c) (b : Vec c) (A : Mat N a) :
          rowDenseFlat N a c W b A.flatten = Mat.flatten fun (r : Fin N) => dense W b (A r)
          theorem Proofs.StableHLO.matMulFlat_flat {m k n : } (A : Mat m k) (B : Mat k n) :
          theorem Proofs.StableHLO.scale_flat {m n : } (s : ) (A : Mat m n) :
          (fun (i : Fin (m * n)) => s * A.flatten i) = Mat.flatten fun (r : Fin m) (c : Fin n) => s * A r c
          theorem Proofs.StableHLO.gelu_flat {m n : } (A : Mat m n) :
          gelu (m * n) A.flatten = Mat.flatten fun (r : Fin m) => gelu n (A r)
          theorem Proofs.StableHLO.add_flat_pt {m n : } (A B : Mat m n) (j : Fin (m * n)) :
          A.flatten j + B.flatten j = Mat.flatten (fun (r : Fin m) (s : Fin n) => A r s + B r s) j
          def Proofs.StableHLO.vitFwdGraph {ic H W P N D mlpDim nClasses : } (epsStr sStr : String) (ε s : ) (Wc : Kernel4 D ic P P) (bc cls : Vec D) (pos : Mat (N + 1) D) (γ1₁ β1₁ : ) (Wq₁ Wk₁ Wv₁ Wo₁ : Mat D D) (bq₁ bk₁ bv₁ bo₁ : Vec D) (γ2₁ β2₁ : ) (Wfc1₁ : Mat D mlpDim) (bfc1₁ : Vec mlpDim) (Wfc2₁ : Mat mlpDim D) (bfc2₁ : Vec D) (γ1₂ β1₂ : ) (Wq₂ Wk₂ Wv₂ Wo₂ : Mat D D) (bq₂ bk₂ bv₂ bo₂ : Vec D) (γ2₂ β2₂ : ) (Wfc1₂ : Mat D mlpDim) (bfc1₂ : Vec mlpDim) (Wfc2₂ : Mat mlpDim D) (bfc2₂ : Vec D) (γF βF : ) (Wcls : Mat D nClasses) (bcls : Vec nClasses) (x : Vec (ic * H * W)) :
          SHlo nClasses

          Whole ViT forward graph (the ch10 representative, peer of convNextFwdGraph): patch embed (stride-P conv + CLS + pos-embed) → 2 spelled transformer blocks (distinct params) → final per-token LN → CLS slice → dense head. Generic D/s; faithful at D := 1 * d, s := sdpa_scale d (heads = 1).

          Equations
          • One or more equations did not get rendered due to their size.
          Instances For
            theorem Proofs.StableHLO.vitFwdGraph_faithful (ic H W patchSize N d mlpDim nClasses : ) (epsStr sStr : String) (Wc : Kernel4 (1 * d) ic patchSize patchSize) (bc cls : Vec (1 * d)) (pos : Mat (N + 1) (1 * d)) (ε γ1₁ β1₁ : ) (Wq₁ Wk₁ Wv₁ Wo₁ : Mat (1 * d) (1 * d)) (bq₁ bk₁ bv₁ bo₁ : Vec (1 * d)) (γ2₁ β2₁ : ) (Wfc1₁ : Mat (1 * d) mlpDim) (bfc1₁ : Vec mlpDim) (Wfc2₁ : Mat mlpDim (1 * d)) (bfc2₁ : Vec (1 * d)) (γ1₂ β1₂ : ) (Wq₂ Wk₂ Wv₂ Wo₂ : Mat (1 * d) (1 * d)) (bq₂ bk₂ bv₂ bo₂ : Vec (1 * d)) (γ2₂ β2₂ : ) (Wfc1₂ : Mat (1 * d) mlpDim) (bfc1₂ : Vec mlpDim) (Wfc2₂ : Mat mlpDim (1 * d)) (bfc2₂ : Vec (1 * d)) (γF βF : ) (Wcls : Mat (1 * d) nClasses) (bcls : Vec nClasses) (x : Vec (ic * H * W)) :
            den (vitFwdGraph epsStr sStr ε (sdpa_scale d) Wc bc cls pos γ1₁ β1₁ Wq₁ Wk₁ Wv₁ Wo₁ bq₁ bk₁ bv₁ bo₁ γ2₁ β2₁ Wfc1₁ bfc1₁ Wfc2₁ bfc2₁ γ1₂ β1₂ Wq₂ Wk₂ Wv₂ Wo₂ bq₂ bk₂ bv₂ bo₂ γ2₂ β2₂ Wfc1₂ bfc1₂ Wfc2₂ bfc2₂ γF βF Wcls bcls x) = vitForward2 ic H W patchSize N mlpDim 1 d nClasses Wc bc cls pos ε γ1₁ β1₁ Wq₁ Wk₁ Wv₁ Wo₁ bq₁ bk₁ bv₁ bo₁ γ2₁ β2₁ Wfc1₁ bfc1₁ Wfc2₁ bfc2₁ γ1₂ β1₂ Wq₂ Wk₂ Wv₂ Wo₂ bq₂ bk₂ bv₂ bo₂ γ2₂ β2₂ Wfc1₂ bfc1₂ Wfc2₂ bfc2₂ γF βF Wcls bcls x

            ViT forward faithfulness — the ch10 close's Item A apex: the representative forward graph denotes the proven distinct-param 2-block vitForward2 at one head (heads := 1, D := 1 * d, s := sdpa_scale d). The ViT analogue of convNextFwdGraph_faithful: per-block vitBlockGraph_den_aux + vitBlockSpelled_eq (mhsa_layer_one_head under the hood), then the patch-embed / CLS-slice den helpers are the proven Attention forms verbatim.