Documentation

LeanMlir.Proofs.Foundation.Batched.Indexed

The indexed lift — a different per-example map at every example #

batchMap N f runs ONE map on every example, and every lemma the whole-net ties of the batch-separable nets (ViT, ConvNeXt) thread is about that shape: batchSlice_batchMap, batchMapHasVJPAt, HasGradAt.param_batchMap_through, batchShard_batchMap. A drop-path site breaks it: example n's block scales its branch by example n's own mask entry, so the block is a different map at every example. batchMapIdx N f takes the family f : Fin N → Vec a → Vec b, and this file restates each of those lemmas for it.

batchMap N f is batchMapIdx N (fun _ => f) and batchMapAux N f aux is batchMapAuxIdx N (fun _ => f) aux, both by rfl (batchMap_eq_batchMapIdx, batchMapAux_eq_batchMapAuxIdx), so a statement made at the indexed lift holds at the uniform one with no extra step.

One example's site. What the family varies by is a drop site read at one example: dropScalarOpt s scales a vector by s's scalar, or is id at none, and batchSlice_dropPathOpt says the batched site (dropPathOpt, Foundation.DropSites) is it at every example's own mask entry. siteResHasVJP is the residual with a site on its branch, v ↦ v + s ⊙ br v: its backward is dy + br.back (s ⊙ dy) by definition — the skip reads the raw cotangent, the branch the dropped one (eBack's rule, dropPath_vjp_is_self).

noncomputable def Proofs.StableHLO.batchMapIdx (N : ℕ) {a b : ℕ} (f : Fin N → Vec a → Vec b) :
Vec (N * a) → Vec (N * b)

Per-example block-apply, a different map at every example. Example n of the batch (the finProdFinEquiv block {(n, ·)}) is mapped by f n.

Equations
Instances For
    noncomputable def Proofs.StableHLO.batchMapAuxIdx (N : ℕ) {s a b : ℕ} (f : Fin N → Vec s → Vec a → Vec b) (aux : Vec (N * s)) :
    Vec (N * a) → Vec (N * b)

    batchMapAux with a different map at every example: example n is handed its own slice of aux and mapped by f n.

    Equations
    • One or more equations did not get rendered due to their size.
    Instances For
      theorem Proofs.StableHLO.batchMap_eq_batchMapIdx (N : ℕ) {a b : ℕ} (f : Vec a → Vec b) :
      batchMap N f = batchMapIdx N fun (x : Fin N) => f

      The uniform lift is the indexed one at a constant family.

      theorem Proofs.StableHLO.batchMapAux_eq_batchMapAuxIdx (N : ℕ) {s a b : ℕ} (f : Vec s → Vec a → Vec b) (aux : Vec (N * s)) :
      batchMapAux N f aux = batchMapAuxIdx N (fun (x : Fin N) => f) aux

      …and so is the uniform auxiliary lift.

      theorem Proofs.StableHLO.batchMapAuxIdx_eq_batchMapIdx (N : ℕ) {s a b : ℕ} (f : Fin N → Vec s → Vec a → Vec b) (aux : Vec (N * s)) :
      batchMapAuxIdx N f aux = batchMapIdx N fun (n : Fin N) => f n (batchSlice N s aux n)

      The auxiliary lift is the indexed one with each example's slice of aux applied.

      theorem Proofs.StableHLO.batchSlice_batchMapIdx {N a b : ℕ} (f : Fin N → Vec a → Vec b) (x : Vec (N * a)) (n : Fin N) :
      batchSlice N b (batchMapIdx N f x) n = f n (batchSlice N a x n)

      batchSlice of a batchMapIdx is example n's map at the slice.

      theorem Proofs.StableHLO.batchSlice_batchMapAuxIdx {N s a b : ℕ} (f : Fin N → Vec s → Vec a → Vec b) (aux : Vec (N * s)) (x : Vec (N * a)) (n : Fin N) :
      batchSlice N b (batchMapAuxIdx N f aux x) n = f n (batchSlice N s aux n) (batchSlice N a x n)

      batchSlice of a batchMapAuxIdx is example n's map at the two slices.

      theorem Proofs.StableHLO.batchMapIdx_comp (B : ℕ) {a b c : ℕ} (f : Fin B → Vec a → Vec b) (g : Fin B → Vec b → Vec c) :
      (batchMapIdx B fun (n : Fin B) => g n ∘ f n) = batchMapIdx B g ∘ batchMapIdx B f

      batchMapIdx distributes over composition, example by example.

      theorem Proofs.batchMapIdx_eq_rowwiseFlat {N a b : ℕ} (f : Fin N → Vec a → Vec b) :
      StableHLO.batchMapIdx N f = fun (v : Vec (N * a)) => Mat.flatten fun (r : Fin N) => f r (Mat.unflatten v r)

      batchMapIdx N f is the flattened row-wise application of the family.

      theorem Proofs.batchMapIdx_differentiableAt {N a b : ℕ} (f : Fin N → Vec a → Vec b) (v : Vec (N * a)) (hf : ∀ (r : Fin N), DifferentiableAt ℝ (f r) (Mat.unflatten v r)) :

      batchMapIdx N f is differentiable at v when each f r is at row r.

      theorem Proofs.batchMapIdx_differentiable {N a b : ℕ} (f : Fin N → Vec a → Vec b) (hf : ∀ (n : Fin N), Differentiable ℝ (f n)) :

      batchMapIdx N f is differentiable when every f n is.

      theorem Proofs.pdiv_batchMapIdx_at {N a b : ℕ} (f : Fin N → Vec a → Vec b) (v : Vec (N * a)) (hf_diff : ∀ (r : Fin N), DifferentiableAt ℝ (f r) (Mat.unflatten v r)) (idx : Fin (N * a)) (jdx : Fin (N * b)) :

      batchMapIdx's Jacobian is block-diagonal across the batch, at a point: entry (idx, jdx) vanishes unless the two indices name the same example, and is that example's own map's entry otherwise.

      noncomputable def Proofs.batchMapIdxHasVJPAt {N a b : ℕ} (f : Fin N → Vec a → Vec b) (v : Vec (N * a)) (hf : (r : Fin N) → HasVJPAt (f r) (Mat.unflatten v r)) (hf_diff : ∀ (r : Fin N), DifferentiableAt ℝ (f r) (Mat.unflatten v r)) :

      batchMapIdx N f's VJP at a point: each example's row runs its own map's backward. Built field by field (as batchMapHasVJPAt) so .backward reduces.

      Equations
      • One or more equations did not get rendered due to their size.
      Instances For
        noncomputable def Proofs.batchMapIdxHasVJP {N a b : ℕ} (f : Fin N → Vec a → Vec b) (hf : (n : Fin N) → HasVJP (f n)) (hf_diff : ∀ (n : Fin N), Differentiable ℝ (f n)) :

        The global VJP of batchMapIdx N f, every f n globally certified.

        Equations
        • One or more equations did not get rendered due to their size.
        Instances For
          theorem Proofs.batchMapAuxIdx_eq_batchMapIdxHasVJPAt {N a b : ℕ} (f : Fin N → Vec a → Vec b) (g : Fin N → Vec a → Vec b → Vec a) (v : Vec (N * a)) (hf : (r : Fin N) → HasVJPAt (f r) (Mat.unflatten v r)) (hf_diff : ∀ (r : Fin N), DifferentiableAt ℝ (f r) (Mat.unflatten v r)) (hg : ∀ (r : Fin N), g r (Mat.unflatten v r) = (hf r).backward) :

          A batched backward tie from the per-example ones, at an indexed family. If g n is example n's certified backward, batchMapAuxIdx N g v IS the lifted witness's backward.

          theorem Proofs.batchMapIdx_param_differentiableAt {P N a q : ℕ} (per : Fin N → Vec P → Vec a → Vec q) (r : Vec (N * a)) (θ : Vec P) (hper : ∀ (n : Fin N) (y : Vec a), DifferentiableAt ℝ (fun (θ' : Vec P) => per n θ' y) θ) :
          DifferentiableAt ℝ (fun (θ' : Vec P) => StableHLO.batchMapIdx N (fun (n : Fin N) => per n θ') r) θ

          The batched parameterised op θ ↦ batchMapIdx N (fun n => per n θ) r is differentiable when each example's map is differentiable in the parameter.

          theorem Proofs.HasGradAt.param_batchMapIdx {P N a q : ℕ} {G : Vec (N * q) → Vec 1} (per : Fin N → Vec P → Vec a → Vec q) (r : Vec (N * a)) {θ : Vec P} {dy : Vec (N * q)} (hG : HasGradAt G (StableHLO.batchMapIdx N (fun (n : Fin N) => per n θ) r) dy) (hper : ∀ (n : Fin N) (y : Vec a), DifferentiableAt ℝ (fun (θ' : Vec P) => per n θ' y) θ) :
          HasGradAt (fun (θ' : Vec P) => G (StableHLO.batchMapIdx N (fun (n : Fin N) => per n θ') r)) θ fun (i : Fin P) => ∑ n : Fin N, ∑ j : Fin q, pdiv (fun (θ' : Vec P) => per n θ' (StableHLO.batchSlice N a r n)) θ i j * StableHLO.batchSlice N q dy n j

          HasGradAt.param_batchMap at an indexed family: the Jacobian split by example, example n differentiated through its own map.

          theorem Proofs.HasGradAt.param_batchMapIdx_through {P N a b m q : ℕ} {G : Vec (N * q) → Vec 1} (pre : Fin N → Vec a → Vec b) (per : Vec P → Vec b → Vec m) (post : Fin N → Vec a → Vec m → Vec q) (cot : Fin N → Vec a → Vec q → Vec m) (X : Vec (N * a)) {θ : Vec P} {dY : Vec (N * q)} (hG : HasGradAt G (StableHLO.batchMapIdx N (fun (n : Fin N) (y : Vec a) => post n y (per θ (pre n y))) X) dY) (hper : ∀ (y : Vec b), DifferentiableAt ℝ (fun (θ' : Vec P) => per θ' y) θ) (hpost : ∀ (n : Fin N) (y : Vec a), Differentiable ℝ (post n y)) (hcot : ∀ (n : Fin N) (y : Vec a) (dy : Vec q), HasGradAt (fun (u : Vec m) => linLoss dy (post n y u)) (per θ (pre n y)) (cot n y dy)) (A : Vec (N * b)) (COT : Vec (N * m)) (hA : ∀ (n : Fin N), StableHLO.batchSlice N b A n = pre n (StableHLO.batchSlice N a X n)) (hC : ∀ (n : Fin N), StableHLO.batchSlice N m COT n = cot n (StableHLO.batchSlice N a X n) (StableHLO.batchSlice N q dY n)) :
          HasGradAt (fun (θ' : Vec P) => G (StableHLO.batchMapIdx N (fun (n : Fin N) (y : Vec a) => post n y (per θ' (pre n y))) X)) θ fun (i : Fin P) => ∑ n : Fin N, ∑ j : Fin m, pdiv (fun (θ' : Vec P) => per θ' (StableHLO.batchSlice N b A n)) θ i j * StableHLO.batchSlice N m COT n j

          HasGradAt.param_batchMap_through at an indexed family. Example n runs y ↦ post n y (per θ (pre n y)): the shared parameterised op between example n's own prefix and suffix (a drop scale on either side is example n's mask entry). If, per example, the loss ⟨post n y ·, dy⟩ has gradient cot n y dy at the op's output, the batched node Σ_n Σ_j ∂per/∂θ · cotₙ — at any saved activation A and cotangent COT whose slices are pre n yₙ and cot n yₙ dyₙ — is the gradient in θ of the whole batched loss.

          noncomputable def Proofs.siteScale :
          Option ℝ → ℝ → ℝ

          One entry through a drop site that may be absent: none passes it, some a scales it.

          Equations
          Instances For
            @[simp]
            @[simp]
            theorem Proofs.siteScale_some (a x : ℝ) :
            siteScale (some a) x = a * x
            noncomputable def Proofs.dropScalarOpt {k : ℕ} (s : Option ℝ) (v : Vec k) :
            Vec k

            One example's drop scale at a site that may be absent, entry by entry: none is the identity, some a scales every entry by a — example n's reading of dropPathOpt (batchSlice_dropPathOpt). Pointwise, so it reads the same on a flat vector and on a row of its matrix.

            Equations
            Instances For
              @[simp]
              theorem Proofs.dropScalarOpt_some {k : ℕ} (a : ℝ) :
              dropScalarOpt (some a) = fun (v : Vec k) (i : Fin k) => a * v i
              noncomputable def Proofs.dropScalarOptHasVJP {k : ℕ} (s : Option ℝ) :

              The site's VJP is the site itself, stated as the backward field (as dropPathOptHasVJP) so it unfolds at a symbolic site.

              Equations
              Instances For

                The site is linear in what flows through it.

                def Proofs.exampleSite {N : ℕ} (sd : Option (Vec N)) (n : Fin N) :

                Example n's site of a per-example mask that may be absent: its entry at n, or none. Named, not spelled sd.map fun v => v n at every use: a statement that spells the lambda twice gets two hygienic binders, and extract_lets then stops merging the two let chains (planning/droppath_tie.md §2).

                Equations
                Instances For
                  @[simp]
                  @[simp]
                  theorem Proofs.exampleSite_some {N : ℕ} (s : Vec N) (n : Fin N) :
                  exampleSite (some s) n = some (s n)
                  theorem Proofs.batchSlice_dropPathOpt {N k : ℕ} (sd : Option (Vec N)) (x : Vec (N * k)) (n : Fin N) :

                  The batched site, read at one example, is that example's scalar site at its mask entry.

                  noncomputable def Proofs.siteResHasVJP {k : ℕ} (s : Option ℝ) (br : Vec k → Vec k) (hd : Differentiable ℝ br) (hv : HasVJP br) :
                  HasVJP fun (v : Vec k) (i : Fin k) => v i + dropScalarOpt s (br v) i

                  A residual with a drop site on its branch, v ↦ v + s ⊙ br v: the skip's backward is the raw cotangent, the branch's is br's at the dropped one (siteResHasVJP_backward, rfl).

                  Equations
                  • One or more equations did not get rendered due to their size.
                  Instances For
                    theorem Proofs.siteResHasVJP_backward {k : ℕ} (s : Option ℝ) (br : Vec k → Vec k) (hd : Differentiable ℝ br) (hv : HasVJP br) (x dy : Vec k) :
                    (siteResHasVJP s br hd hv).backward x dy = fun (i : Fin k) => dy i + hv.backward x (dropScalarOpt s dy) i
                    theorem Proofs.siteRes_differentiable {k : ℕ} (s : Option ℝ) (br : Vec k → Vec k) (hd : Differentiable ℝ br) :
                    Differentiable ℝ fun (v : Vec k) (i : Fin k) => v i + dropScalarOpt s (br v) i
                    theorem Proofs.StableHLO.batchShard_batchMapIdx {R N a b : ℕ} (f : Fin (R * N) → Vec a → Vec b) (X : Vec (R * N * a)) (r : Fin R) :
                    batchShard R N b (batchMapIdx (R * N) f X) r = batchMapIdx N (fun (n : Fin N) => f (finProdFinEquiv (r, n))) (batchShard R N a X r)

                    batchMapIdx commutes with sharding: shard r's family is the global family at the global indices of its examples.

                    theorem Proofs.StableHLO.batchShard_batchMapAuxIdx {R N s a b : ℕ} (f : Fin (R * N) → Vec s → Vec a → Vec b) (aux : Vec (R * N * s)) (X : Vec (R * N * a)) (r : Fin R) :
                    batchShard R N b (batchMapAuxIdx (R * N) f aux X) r = batchMapAuxIdx N (fun (n : Fin N) => f (finProdFinEquiv (r, n))) (batchShard R N s aux r) (batchShard R N a X r)

                    …and so does batchMapAuxIdx.

                    theorem Proofs.batchMapIdx_smul {N a b : ℕ} (f : Fin N → Vec a → Vec b) (hf : ∀ (n : Fin N), IsHomog (f n)) :

                    An indexed lift of homogeneous maps is homogeneous.

                    theorem Proofs.batchMapAuxIdx_smul {N t a b : ℕ} (f : Fin N → Vec t → Vec a → Vec b) (hf : ∀ (n : Fin N) (x : Vec t), IsHomog (f n x)) (aux : Vec (N * t)) :

                    …and so is an indexed auxiliary lift, homogeneous in its last argument.