Documentation

LeanMlir.Proofs.Nets.Small.SmallParamGrad

SmallParamGrad — the per-example kit for the chapter nets' loss gradients #

The seven ImageNet nets state their parameter gradients through ParamGradNodes, at the batched *GradB nodes. The chapter nets (linear, MLP, MNIST CNN, the CIFAR CNNs) run one example at a time and emit the per-example GradNode ops (SgdNodes). This file is their kit:

Which cells are twins is per net: each net file names them (cells equal at every value of the weights upstream of the pool) and discharges maxPool_relu_eventuallyEq_sel along each parameter.

theorem Proofs.SmallParamGrad.hasGradAt_crossEntropy {n : ℕ} (label : Fin n) (z : Vec n) :
HasGradAt (fun (z' : Vec n) (x : Fin 1) => crossEntropy n z' label) z fun (j : Fin n) => softmax n z j - oneHot n label j

Softmax cross-entropy at a hard label has gradient softmax − onehot in the logits.

theorem Proofs.SmallParamGrad.convW_hasGradAt {ic oc h w kH kW : ℕ} (xN cotN : String) (b : Vec oc) (x : Vec (ic * h * w)) (W : Kernel4 oc ic kH kW) {G : Vec (oc * h * w) → Vec 1} {c : Vec (oc * h * w)} (hG : HasGradAt G (flatConv W b x) c) :

Conv weight node = ∇_W G.

theorem Proofs.SmallParamGrad.convB_hasGradAt {ic oc h w kH kW : ℕ} (cotN : String) (W : Kernel4 oc ic kH kW) (x : Vec (ic * h * w)) (b : Vec oc) {G : Vec (oc * h * w) → Vec 1} {c : Vec (oc * h * w)} (hG : HasGradAt G (flatConv W b x) c) :

Conv bias node = ∇_b G.

theorem Proofs.SmallParamGrad.denseW_hasGradAt {m n : ℕ} (aN cotN : String) (a : Vec m) (W : Mat m n) (b : Vec n) {G : Vec n → Vec 1} {c : Vec n} (hG : HasGradAt G (dense W b a) c) :

Dense weight node = ∇_W G.

theorem Proofs.SmallParamGrad.denseB_hasGradAt {m n : ℕ} (cotN : String) (W : Mat m n) (a : Vec m) (b : Vec n) {G : Vec n → Vec 1} {c : Vec n} (hG : HasGradAt G (dense W b a) c) :
HasGradAt (fun (θ : Vec n) => G (dense W θ a)) b (StableHLO.den (StableHLO.SHlo.operand cotN c).biasGrad)

Dense bias node = ∇_b G.

The fused *Sgd op is θ − lr· its un-fused *Grad peer, at any cotangent: the SGD renders (the linear, MLP, MNIST-CNN and CIFAR arms) step by exactly the node the lemmas above identify.

theorem Proofs.SmallParamGrad.convWeightSgd_eq_grad {ic oc h w kH kW : ℕ} (xN wN lrStr cotN : String) (b : Vec oc) (x : Tensor3 ic h w) (W : Kernel4 oc ic kH kW) (c : Vec (oc * h * w)) (lr : ℝ) (idx : Fin (oc * ic * kH * kW)) :
theorem Proofs.SmallParamGrad.convBiasSgd_eq_grad {ic oc h w kH kW : ℕ} (bN lrStr cotN : String) (W : Kernel4 oc ic kH kW) (x : Tensor3 ic h w) (b : Vec oc) (c : Vec (oc * h * w)) (lr : ℝ) (o : Fin oc) :
theorem Proofs.SmallParamGrad.weightSgd_eq_grad {m n : ℕ} (aN wN lrStr cotN : String) (a : Vec m) (W : Mat m n) (b c : Vec n) (lr : ℝ) (i : Fin m) (j : Fin n) :
theorem Proofs.SmallParamGrad.biasSgd_eq_grad {m n : ℕ} (bN lrStr cotN : String) (W : Mat m n) (a : Vec m) (b c : Vec n) (lr : ℝ) (i : Fin n) :
theorem Proofs.SmallParamGrad.hasGradAt_dense {m n : ℕ} (W : Mat m n) (b : Vec n) (u : Vec m) {G : Vec n → Vec 1} {dy : Vec n} (hG : HasGradAt G (dense W b u) dy) :
HasGradAt (fun (y : Vec m) => G (dense W b y)) u ((IR.emitDenseBack W).denote dy)

Through a dense layer: the backward is W · dy (emitDenseBack).

theorem Proofs.SmallParamGrad.hasGradAt_relu {n : ℕ} (z : Vec n) (hz : ∀ (k : Fin n), z k ≠ 0) {G : Vec n → Vec 1} {dy : Vec n} (hG : HasGradAt G (relu n z) dy) :
HasGradAt (fun (y : Vec n) => G (relu n y)) z ((IR.emitReluBack z).denote dy)

Through a ReLU off its kink: the backward is the mask relu'(z) ⊙ dy (emitReluBack).

theorem Proofs.SmallParamGrad.hasGradAt_conv {ic oc h w kH kW : ℕ} (hkH : 2 * ((kH - 1) / 2) + 1 = kH) (hkW : 2 * ((kW - 1) / 2) + 1 = kW) (W : Kernel4 oc ic kH kW) (b : Vec oc) (v : Vec (ic * h * w)) {G : Vec (oc * h * w) → Vec 1} {dy : Vec (oc * h * w)} (hG : HasGradAt G (flatConv W b v) dy) :
HasGradAt (fun (y : Vec (ic * h * w)) => G (flatConv W b y)) v ((IR.Back3.conv W IR.Back3.cot).flatDenote dy)

Through a stride-1 conv with odd kernels: the backward is the rendered reversed-kernel conv (Back3.conv, conv_flatten_bridge).

def Proofs.SmallParamGrad.poolSelIdx {c h w : ℕ} (σ : Fin c → Fin h → Fin w → Fin 2 × Fin 2) (k : Fin (c * h * w)) :
Fin (c * (2 * h) * (2 * w))

The flat index of the cell σ selects in pooled entry k's window.

Equations
  • One or more equations did not get rendered due to their size.
Instances For
    theorem Proofs.SmallParamGrad.poolSelIdx_t3Idx {c h w : ℕ} (σ : Fin c → Fin h → Fin w → Fin 2 × Fin 2) (ci : Fin c) (ho : Fin h) (wo : Fin w) :
    poolSelIdx σ (t3Idx ci ho wo) = t3Idx ci (winRowInv ho (σ ci ho wo).1) (winColInv wo (σ ci ho wo).2)
    theorem Proofs.SmallParamGrad.poolGatherFlat_eq_sel {c h w : ℕ} (σ : Fin c → Fin h → Fin w → Fin 2 × Fin 2) (u : Vec (c * (2 * h) * (2 * w))) :
    poolGatherFlat σ u = fun (k : Fin (c * h * w)) => u (poolSelIdx σ k)

    poolGatherFlat is the reindex along poolSelIdx.

    def Proofs.SmallParamGrad.PoolSelDom {c h w : ℕ} (σ : Fin c → Fin h → Fin w → Fin 2 × Fin 2) (u : Vec (c * (2 * h) * (2 * w))) :

    The selection names a maximum of every window of u. The rendered select_and_scatter (select = GE) makes such a choice; the canonical argmax is one (poolSelDom_argmax).

    Equations
    • One or more equations did not get rendered due to their size.
    Instances For
      @[reducible, inline]
      abbrev Proofs.SmallParamGrad.MaxPool2SmoothUpTo {c h w : ℕ} (T : Fin (2 * h) × Fin (2 * w) → Fin (2 * h) × Fin (2 * w) → Prop) (x : Tensor3 c (2 * h) (2 * w)) :

      Smooth, dead, or tied only between twins at the 2×2 windows (WindowSmoothUpTo), on the pool's PRE-activation: a window is dead when its cells are all ≤ 0. The pre-activation margin MaxPool2MarginQUpTo δ T implies it at any δ ≥ 0 (windowSmoothUpTo_of_margin).

      Equations
      Instances For

        The 2×2 pool is continuous (a max of coordinates), so a pre-activation computed through earlier pools moves continuously with the parameters.

        def Proofs.SmallParamGrad.selScatter {m n : ℕ} (σ : Fin n → Fin m) (dy : Vec n) :
        Vec m

        The scatter along a selection: each pooled cotangent lands on the one cell it read.

        Equations
        Instances For
          theorem Proofs.SmallParamGrad.hasGradAt_gatherRelu {m n : ℕ} (σ : Fin n → Fin m) (z : Vec m) (hz : ∀ (k : Fin m), z k ≠ 0) {G : Vec n → Vec 1} {dy : Vec n} (hG : HasGradAt G (fun (k : Fin n) => relu m z (σ k)) dy) :
          HasGradAt (fun (y : Vec m) => G fun (k : Fin n) => relu m y (σ k)) z ((IR.emitReluBack z).denote (selScatter σ dy))

          Through ReLU then a gather (y ↦ relu y ∘ σ), off the ReLU kinks: the backward is the scatter along σ, then the ReLU mask. No pool hypothesis: the gather is linear.

          theorem Proofs.SmallParamGrad.maxPool_relu_eventuallyEq_sel {P c h w : ℕ} (Z : Vec P → Vec (c * (2 * h) * (2 * w))) (σ : Fin c → Fin h → Fin w → Fin 2 × Fin 2) (T : Fin (2 * h) × Fin (2 * w) → Fin (2 * h) × Fin (2 * w) → Prop) (hT : ∀ (θ : Vec P) (ci : Fin c) (p q : Fin (2 * h) × Fin (2 * w)), T p q → Z θ (t3Idx ci p.1 p.2) = Z θ (t3Idx ci q.1 q.2)) (θ₀ : Vec P) (hZc : ContinuousAt Z θ₀) (hz : ∀ (k : Fin (c * (2 * h) * (2 * w))), Z θ₀ k ≠ 0) (hs : MaxPool2SmoothUpTo T (Tensor3.unflatten (Z θ₀))) (hσ : PoolSelDom σ (relu (c * (2 * h) * (2 * w)) (Z θ₀))) :
          ∀ᶠ (θ : Vec P) in nhds θ₀, maxPoolFlat c h w (relu (c * (2 * h) * (2 * w)) (Z θ)) = fun (k : Fin (c * h * w)) => relu (c * (2 * h) * (2 * w)) (Z θ) (poolSelIdx σ k)

          Along a parameter, ReLU then the 2×2 pool is the gather at a fixed selection. Z θ is the pool's pre-activation as the parameter moves; at θ₀ it has no zero entry, every window is dead or tied only between T-twins, σ names a maximum of every window of its ReLU, and T-twins are equal at EVERY θ. Then near θ₀ the pooled ReLU reads each window at σ: a dead window stays negative, a strict maximum stays strict, and a twin stays tied with it.

          The step ties' pool backward is the scatter at the first argmax. The Back3 maxpool node (maxPoolBackDenote, the den of the rendered maxPoolBack) routes each window's cotangent to maxPool2Argmax's cell, the window's first maximum, so through the flatten it is selScatter along that selection, at every point, ties included. With poolSelDom_argmax, the chain the step ties read is the loss-gradient chain at σ = maxPool2Argmax.