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.
batchSlice_batchMapIdx/batchSlice_batchMapAuxIdx— examplenof the lift isf nat examplen.batchMapIdxHasVJPAt/batchMapIdxHasVJP— the Jacobian is block-diagonal, withf n's own block at examplen(pdiv_batchMapIdx_at, frompdivMat_rowIndep_perRow_at, which already allows a different map on every row);batchMapAuxIdx_eq_batchMapIdxHasVJPAtties a batched backward chain to it.HasGradAt.param_batchMapIdx_through— a parameter op shared by every example, between an indexed prefix and an indexed suffix: the batched node is the parameter's loss gradient.batchShard_batchMapIdx/_batchMapAuxIdx— shardr's family is the global one atfinProdFinEquiv (r, ·).
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).
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
- Proofs.StableHLO.batchMapIdx N f x idx = f (finProdFinEquiv.symm idx).1 (fun (i : Fin a) => x (finProdFinEquiv ((finProdFinEquiv.symm idx).1, i))) (finProdFinEquiv.symm idx).2
Instances For
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
The uniform lift is the indexed one at a constant family.
batchSlice of a batchMapIdx is example n's map at the slice.
batchMapIdx distributes over composition, example by example.
batchMapIdx N f is the flattened row-wise application of the family.
batchMapIdx N f is differentiable at v when each f r is at row r.
batchMapIdx N f is differentiable when every f n is.
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.
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
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
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.
The batched parameterised op θ ↦ batchMapIdx N (fun n => per n θ) r is differentiable when
each example's map is differentiable in the parameter.
HasGradAt.param_batchMap at an indexed family: the Jacobian split by example, example n
differentiated through its own map.
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.
One entry through a drop site that may be absent: none passes it, some a scales it.
Equations
- Proofs.siteScale none = fun (x : ℝ) => x
- Proofs.siteScale (some a) = fun (x : ℝ) => a * x
Instances For
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
- Proofs.dropScalarOpt s v i = Proofs.siteScale s (v i)
Instances For
The site's VJP is the site itself, stated as the backward field (as dropPathOptHasVJP) so
it unfolds at a symbolic site.
Equations
- Proofs.dropScalarOptHasVJP s = { backward := fun (x dy : Proofs.Vec k) => Proofs.dropScalarOpt s dy, correct := ⋯ }
Instances For
The site is linear in what flows through it.
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
- Proofs.exampleSite sd n = Option.map (fun (v : Proofs.Vec N) => v n) sd
Instances For
The batched site, read at one example, is that example's scalar site at its mask entry.
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
batchMapIdx commutes with sharding: shard r's family is the global family at the
global indices of its examples.
…and so does batchMapAuxIdx.
An indexed lift of homogeneous maps is homogeneous.