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.
vitFwdGraphBDrop_slice/vitFwdGraphBDrop_faithful— the graph denotesvitForwardKVDropB, at every depth, every mask and every example;vitForwardKVDrop_ones/vitForwardKVDropB_ones— at all-ones masks (the driver's at eval) the forward isvitForwardKV, per example and lifted, exactly: the keep probability is folded into the mask (Training/DropPath).
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 #
- Huang et al. 2016, Deep Networks with Stochastic Depth. https://arxiv.org/abs/1603.09382
- Touvron et al. 2021, Training data-efficient image transformers & distillation through attention (DeiT). https://arxiv.org/abs/2012.12877
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
At unit scalars the drop block is blockV.
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
The spelled drop block IS blockVDrop (vitBlockSpelledMHV_eq's proof, at the scalars).
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
At unit scalars the tower is vitBodyKV.
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
At unit scalars the forward is vitForwardKV.
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
At the all-ones masks the batched forward is vitForwardKV lifted, exactly — the masks
the driver passes to the forward artifacts at eval.
A batched descriptor token, sliced at example t, is its per-example map at the slice.
matmulFB sliced at example t multiplies example t's two operands
(den_matmulFB_per_example, as a slice).
A drop site, sliced at example t, scales by example t's mask entry — the per-example
content of dropPathB.
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
- Proofs.StableHLO.headsSumGB f = f 0
- Proofs.StableHLO.headsSumGB f = (Proofs.StableHLO.headsSumGB fun (i : Fin (hm1 + 1)) => f i.castSucc).addVB (f (Fin.last (hm1 + 1)))
Instances For
The batched head fold, sliced at example t, is the sum over heads of the slices.
The right-multiplied scale commutes with flattening (scale_flat, operands swapped — the
batched scaleB multiplies on the right).
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
Example t of the batched drop block is the spelled drop block at example t's input and
mask entries.
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
- One or more equations did not get rendered due to their size.
- Proofs.StableHLO.vitBodyGraphBDrop epsStr sStr ε s x✝¹ 0 x_7 x_8 x_9 x✝ = x✝
Instances For
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.
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
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.
The batched ViT forward graph with stochastic depth denotes vitForwardKVDropB — at every
depth, every pair of mask families and every input.