Documentation

LeanMlir.Proofs.Nets.Small.Cifar8StepTieG

The cifar8 step tie at its UN-FUSED gradient nodes — the packed cifar8w_* arms #

Cifar8Tie.cifar8_train_step_tied_certified ties the fused-SGD cifar8_train_step.mlir, which no trainer runs. The packed wide no-BN arms (cifar8w_{sgd,mom,adam}_train_step.mlir, from cifar8AdamTrainStepText) emit the same forward and backward chain feeding *Grad ops to a separate optimizer. This file states those nodes, each at the same chain cotangent: all 22 parameter tensors, via GradNode (Foundation/SgdNodes.lean). The optimizer update is outside the statement.

Scope (as the fused tie) #

theorem Proofs.Cifar8TieG.cifar8_train_step_tiedG {ic c1 c2 c3 c4 h w d1 nClasses kH kW : ℕ} (xN cotN : String) (W₁ : Kernel4 c1 ic kH kW) (b₁ : Vec c1) (W₂ : Kernel4 c1 c1 kH kW) (b₂ : Vec c1) (W₃ : Kernel4 c2 c1 kH kW) (b₃ : Vec c2) (W₄ : Kernel4 c2 c2 kH kW) (b₄ : Vec c2) (W₅ : Kernel4 c3 c2 kH kW) (b₅ : Vec c3) (W₆ : Kernel4 c3 c3 kH kW) (b₆ : Vec c3) (W₇ : Kernel4 c4 c3 kH kW) (b₇ : Vec c4) (W₈ : Kernel4 c4 c4 kH kW) (b₈ : Vec c4) (W₉ : Mat (c4 * h * w) d1) (b₉ : Vec d1) (Wa : Mat d1 d1) (ba : Vec d1) (Wb : Mat d1 nClasses) (bb : Vec nClasses) (x : Tensor3 ic (2 * (2 * (2 * (2 * h)))) (2 * (2 * (2 * (2 * w))))) (label : Fin nClasses) :
have xv := x.flatten; have cc1 := flatConv W₁ b₁ xv; have r1 := relu (c1 * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w))))) cc1; have r1t := Tensor3.unflatten r1; have cc2 := flatConv W₂ b₂ r1; have r2 := relu (c1 * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w))))) cc2; have r2t := Tensor3.unflatten r2; have zp1 := maxPoolFlat c1 (2 * (2 * (2 * h))) (2 * (2 * (2 * w))) r2; have zp1t := Tensor3.unflatten zp1; have cc3 := flatConv W₃ b₃ zp1; have r3 := relu (c2 * (2 * (2 * (2 * h))) * (2 * (2 * (2 * w)))) cc3; have r3t := Tensor3.unflatten r3; have cc4 := flatConv W₄ b₄ r3; have r4 := relu (c2 * (2 * (2 * (2 * h))) * (2 * (2 * (2 * w)))) cc4; have r4t := Tensor3.unflatten r4; have zp2 := maxPoolFlat c2 (2 * (2 * h)) (2 * (2 * w)) r4; have zp2t := Tensor3.unflatten zp2; have cc5 := flatConv W₅ b₅ zp2; have r5 := relu (c3 * (2 * (2 * h)) * (2 * (2 * w))) cc5; have r5t := Tensor3.unflatten r5; have cc6 := flatConv W₆ b₆ r5; have r6 := relu (c3 * (2 * (2 * h)) * (2 * (2 * w))) cc6; have r6t := Tensor3.unflatten r6; have zp3 := maxPoolFlat c3 (2 * h) (2 * w) r6; have zp3t := Tensor3.unflatten zp3; have cc7 := flatConv W₇ b₇ zp3; have r7 := relu (c4 * (2 * h) * (2 * w)) cc7; have r7t := Tensor3.unflatten r7; have cc8 := flatConv W₈ b₈ r7; have r8 := relu (c4 * (2 * h) * (2 * w)) cc8; have r8t := Tensor3.unflatten r8; have zp4 := maxPoolFlat c4 h w r8; have h9 := dense W₉ b₉ zp4; have ha := dense Wa ba (relu d1 h9); have g := fun (k : Fin nClasses) => softmax nClasses (cifarCnn8Forward W₁ b₁ W₂ b₂ W₃ b₃ W₄ b₄ W₅ b₅ W₆ b₆ W₇ b₇ W₈ b₈ W₉ b₉ Wa ba Wb bb xv) k - oneHot nClasses label k; have cotC8 := cnnChainCotW2 W₉ Wa Wb h9 ha r8t cc8 g; have cotC7 := cnnChainCotW1 W₈ cc7 cotC8; have cotC6 := CifarFold.cifarChainCotW2 W₇ r6t cc6 cotC7; have cotC5 := cnnChainCotW1 W₆ cc5 cotC6; have cotC4 := CifarFold.cifarChainCotW2 W₅ r4t cc4 cotC5; have cotC3 := cnnChainCotW1 W₄ cc3 cotC4; have cotC2 := CifarFold.cifarChainCotW2 W₃ r2t cc2 cotC3; have cotC1 := cnnChainCotW1 W₂ cc1 cotC2; GradNode.ConvWGradTied xN cotN b₁ x W₁ cotC1 ∧ GradNode.ConvBGradTied cotN W₁ x b₁ cotC1 ∧ GradNode.ConvWGradTied xN cotN b₂ r1t W₂ cotC2 ∧ GradNode.ConvBGradTied cotN W₂ r1t b₂ cotC2 ∧ GradNode.ConvWGradTied xN cotN b₃ zp1t W₃ cotC3 ∧ GradNode.ConvBGradTied cotN W₃ zp1t b₃ cotC3 ∧ GradNode.ConvWGradTied xN cotN b₄ r3t W₄ cotC4 ∧ ⋯ ∧ ⋯

Whole cifar8 train step at its gradient nodes. All 22 parameter tensors (8 conv W+b, the dense head W₉,b₉,Wa,ba,Wb,bb), at the real cifar8 forward: each emitted *Grad node denotes the certified per-layer Jacobian contracted with the rendered backward-chain cotangent, driven by the composed softmax-CE cotangent g — the fused tie's chain, node for node.