Documentation

LeanMlir.Proofs.Foundation.ParamGradNodes

ParamGradNodes — each batched parameter gradient node is a loss derivative #

GradNodesB states each emitted parameter gradient node as its layer's parameter Jacobian contracted with an arbitrary output cotangent. Here the cotangent is the gradient of a scalar G at the layer's output (HasGradAt), and the node becomes ∂G/∂θ with the layer's parameter varied: one lemma per node kind, shared by every net.

The BatchNorm γ/β nodes are stated in the transposed [C, N·H·W] layout at the reassocB index; their lemmas re-sum the Jacobian over that permutation (bnLA_perm). A conv or depthwise bias that the render reads with the β op (EfficientNet-B0's) is biasBeta_eq_pdiv: the bias enters as a channel broadcast (*_bias_split), so the channel sum the β node computes is its derivative.

theorem Proofs.GradNodeB.cInB_eq_batchMapBackward {N ic oc h w kH kW : ℕ} (W : Kernel4 oc ic kH kW) (b : Vec oc) (x : Vec (N * (ic * h * w))) (dy : Vec (N * (oc * h * w))) :

cInB — the emitted conv input-cotangent — is the batched conv VJP's backward, at any saved input.

theorem Proofs.GradNodeB.cStridedInB_eq_batchMapBackward {N ic oc h w kH kW : ℕ} (W : Kernel4 oc ic kH kW) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (dy : Vec (N * (oc * h * w))) :

cStridedInB — the emitted strided-conv input-cotangent — is the batched strided conv VJP's backward, at any saved input.

theorem Proofs.GradNodeB.dInB_eq_batchMapBackward {N c h w kH kW : ℕ} (W : DepthwiseKernel c kH kW) (b : Vec c) (x dy : Vec (N * (c * h * w))) :

dInB — the emitted depthwise input-cotangent — is the batched depthwise VJP's backward.

theorem Proofs.GradNodeB.dStridedInB_eq_batchMapBackward {N c h w kH kW : ℕ} (W : DepthwiseKernel c kH kW) (b : Vec c) (x : Vec (N * (c * (2 * h) * (2 * w)))) (dy : Vec (N * (c * h * w))) :

dStridedInB — the emitted (symmetric) strided depthwise input-cotangent — is the batched strided depthwise VJP's backward, at any saved input.

theorem Proofs.GradNodeB.hasGradAt_bnBatchLA {N c h w : ℕ} (ε : ℝ) (hε : 0 < ε) (γ β : Vec c) (z : Vec (N * (c * h * w))) {G : Vec (N * (c * h * w)) → Vec 1} {dy : Vec (N * (c * h * w))} (hG : HasGradAt G (StableHLO.bnBatchLA N c h w ε γ β z) dy) :
HasGradAt (fun (z' : Vec (N * (c * h * w))) => G (StableHLO.bnBatchLA N c h w ε γ β z')) z (BackLinks.bnInB N c h w ε γ z dy)

Back through batch BN: the gradient at its input is bnInB of the gradient at its output.

theorem Proofs.GradNodeB.hasGradAt_relu {n : ℕ} (x : Vec n) (hs : ∀ (k : Fin n), x k ≠ 0) {G : Vec n → Vec 1} {dy : Vec n} (hG : HasGradAt G (relu n x) dy) :
HasGradAt (fun (u : Vec n) => G (relu n u)) x (BackLinks.reluMaskB n x dy)

Back through relu, off its kink: reluMaskB.

theorem Proofs.GradNodeB.hasGradAt_bnBackB {N c h w : ℕ} (ε : ℝ) (hε : 0 < ε) (γ β : Vec c) (z : Vec (N * (c * h * w))) {G : Vec (N * (c * h * w)) → Vec 1} {dy : Vec (N * (c * h * w))} (hG : HasGradAt G (StableHLO.bnBatchLA N c h w ε γ β z) dy) :
HasGradAt (fun (z' : Vec (N * (c * h * w))) => G (StableHLO.bnBatchLA N c h w ε γ β z')) z (BackLinks.bnBackB N c h w ε hε γ β z dy)

Back through batch BN, at the certified backward's own spelling bnBackB (EfficientNet-B0's tie threads this form; the ResNets and MobileNets emit bnInB).

theorem Proofs.GradNodeB.hasGradAt_swish {n : ℕ} (x : Vec n) {G : Vec n → Vec 1} {dy : Vec n} (hG : HasGradAt G (swish n x) dy) :
HasGradAt (fun (u : Vec n) => G (swish n u)) x (BackLinks.swBackB n x dy)

Back through swish: swBackB. Swish has no kink, so there is no smoothness hypothesis.

theorem Proofs.GradNodeB.hasGradAt_conv {N ic oc h w kH kW : ℕ} (W : Kernel4 oc ic kH kW) (b : Vec oc) (x : Vec (N * (ic * h * w))) {G : Vec (N * (oc * h * w)) → Vec 1} {dy : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (flatConv W b) x) dy) :
HasGradAt (fun (y : Vec (N * (ic * h * w))) => G (StableHLO.batchMap N (flatConv W b) y)) x (BackLinks.cInB N W b dy)

Back through a batched conv: cInB.

theorem Proofs.GradNodeB.hasGradAt_convStrided {N ic oc h w kH kW : ℕ} (W : Kernel4 oc ic kH kW) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) {G : Vec (N * (oc * h * w)) → Vec 1} {dy : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (flatConvStride2 W b) x) dy) :
HasGradAt (fun (y : Vec (N * (ic * (2 * h) * (2 * w)))) => G (StableHLO.batchMap N (flatConvStride2 W b) y)) x (BackLinks.cStridedInB N W b dy)

Back through a batched symmetric strided conv: cStridedInB.

theorem Proofs.GradNodeB.hasGradAt_depthwise {N c h w kH kW : ℕ} (W : DepthwiseKernel c kH kW) (b : Vec c) (x : Vec (N * (c * h * w))) {G : Vec (N * (c * h * w)) → Vec 1} {dy : Vec (N * (c * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (depthwiseFlat W b) x) dy) :
HasGradAt (fun (y : Vec (N * (c * h * w))) => G (StableHLO.batchMap N (depthwiseFlat W b) y)) x (BackLinks.dInB N W b dy)

Back through a batched depthwise: dInB.

theorem Proofs.GradNodeB.hasGradAt_depthwiseStrided {N c h w kH kW : ℕ} (W : DepthwiseKernel c kH kW) (b : Vec c) (x : Vec (N * (c * (2 * h) * (2 * w)))) {G : Vec (N * (c * h * w)) → Vec 1} {dy : Vec (N * (c * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (depthwiseStride2Flat W b) x) dy) :
HasGradAt (fun (y : Vec (N * (c * (2 * h) * (2 * w)))) => G (StableHLO.batchMap N (depthwiseStride2Flat W b) y)) x (BackLinks.dStridedInB N W b dy)

Back through a batched symmetric strided depthwise: dStridedInB.

theorem Proofs.GradNodeB.convW_eq_pdiv {N ic oc h w kH kW : ℕ} (xN cotN : String) (b : Vec oc) (x : Vec (N * (ic * h * w))) (W : Kernel4 oc ic kH kW) {G : Vec (N * (oc * h * w)) → Vec 1} {cot : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (flatConv W b) x) cot) (idx : Fin (oc * ic * kH * kW)) :
StableHLO.den (StableHLO.SHlo.convWeightGradB xN b x W (StableHLO.SHlo.operand cotN cot)) idx = pdiv (fun (θ : Vec (oc * ic * kH * kW)) => G (StableHLO.batchMap N (flatConv (Kernel4.unflatten θ) b) x)) W.flatten idx 0

Conv weight node = ∂G/∂W.

theorem Proofs.GradNodeB.convB_eq_pdiv {N ic oc h w kH kW : ℕ} (cotN : String) (W : Kernel4 oc ic kH kW) (x : Vec (N * (ic * h * w))) (b : Vec oc) {G : Vec (N * (oc * h * w)) → Vec 1} {cot : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (flatConv W b) x) cot) (o : Fin oc) :
StableHLO.den (StableHLO.SHlo.convBiasGradB W x b (StableHLO.SHlo.operand cotN cot)) o = pdiv (fun (θ : Vec oc) => G (StableHLO.batchMap N (flatConv W θ) x)) b o 0

Conv bias node = ∂G/∂b.

theorem Proofs.GradNodeB.flatConvStride2_weight_differentiable {ic oc h w kH kW : ℕ} (b : Vec oc) (y : Vec (ic * (2 * h) * (2 * w))) :
Differentiable ℝ fun (θ : Vec (oc * ic * kH * kW)) => flatConvStride2 (Kernel4.unflatten θ) b y
theorem Proofs.GradNodeB.flatConvStride2_bias_differentiable {ic oc h w kH kW : ℕ} (W : Kernel4 oc ic kH kW) (y : Vec (ic * (2 * h) * (2 * w))) :
Differentiable ℝ fun (θ : Vec oc) => flatConvStride2 W θ y
theorem Proofs.GradNodeB.convStridedW_eq_pdiv {N ic oc h w kH kW : ℕ} (xN cotN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) {G : Vec (N * (oc * h * w)) → Vec 1} {cot : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (flatConvStride2 W b) x) cot) (idx : Fin (oc * ic * kH * kW)) :

Stride-2 (symmetric) conv weight node = ∂G/∂W.

theorem Proofs.GradNodeB.convStridedB_eq_pdiv {N ic oc h w kH kW : ℕ} (cotN : String) (W : Kernel4 oc ic kH kW) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (b : Vec oc) {G : Vec (N * (oc * h * w)) → Vec 1} {cot : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (flatConvStride2 W b) x) cot) (o : Fin oc) :

Stride-2 (symmetric) conv bias node = ∂G/∂b.

theorem Proofs.GradNodeB.convStridedXlaW_eq_pdiv {N ic oc h w kH kW : ℕ} (xN cotN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) {G : Vec (N * (oc * h * w)) → Vec 1} {cot : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (flatConvStride2Xla W b) x) cot) (idx : Fin (oc * ic * kH * kW)) :

XLA-SAME strided conv weight node = ∂G/∂W (MobileNetV2's and B0's stem).

theorem Proofs.GradNodeB.convStridedXlaB_eq_pdiv {N ic oc h w kH kW : ℕ} (cotN : String) (W : Kernel4 oc ic kH kW) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (b : Vec oc) {G : Vec (N * (oc * h * w)) → Vec 1} {cot : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (flatConvStride2Xla W b) x) cot) (o : Fin oc) :

XLA-SAME strided conv bias node = ∂G/∂b.

theorem Proofs.GradNodeB.depthwiseW_eq_pdiv {N c h w kH kW : ℕ} (xN cotN : String) (b : Vec c) (x : Vec (N * (c * h * w))) (W : DepthwiseKernel c kH kW) {G : Vec (N * (c * h * w)) → Vec 1} {cot : Vec (N * (c * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (depthwiseFlat W b) x) cot) (idx : Fin (c * kH * kW)) :

Depthwise weight node = ∂G/∂W.

theorem Proofs.GradNodeB.depthwiseB_eq_pdiv {N c h w kH kW : ℕ} (cotN : String) (W : DepthwiseKernel c kH kW) (x : Vec (N * (c * h * w))) (b : Vec c) {G : Vec (N * (c * h * w)) → Vec 1} {cot : Vec (N * (c * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (depthwiseFlat W b) x) cot) (o : Fin c) :

Depthwise bias node = ∂G/∂b.

theorem Proofs.GradNodeB.depthwiseStridedW_eq_pdiv {N c h w kH kW : ℕ} (xN cotN : String) (b : Vec c) (x : Vec (N * (c * (2 * h) * (2 * w)))) (W : DepthwiseKernel c kH kW) {G : Vec (N * (c * h * w)) → Vec 1} {cot : Vec (N * (c * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (depthwiseStride2Flat W b) x) cot) (idx : Fin (c * kH * kW)) :

Symmetric strided depthwise weight node = ∂G/∂W (MobileNetV4's dw_mid at its three downsampling rows).

theorem Proofs.GradNodeB.depthwiseStridedXlaW_eq_pdiv {N c h w kH kW : ℕ} (xN cotN : String) (b : Vec c) (x : Vec (N * (c * (2 * h) * (2 * w)))) (W : DepthwiseKernel c kH kW) {G : Vec (N * (c * h * w)) → Vec 1} {cot : Vec (N * (c * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (depthwiseStride2FlatXla W b) x) cot) (idx : Fin (c * kH * kW)) :

XLA-SAME strided depthwise weight node = ∂G/∂W.

theorem Proofs.GradNodeB.depthwiseStridedXlaB_eq_pdiv {N c h w kH kW : ℕ} (cotN : String) (W : DepthwiseKernel c kH kW) (x : Vec (N * (c * (2 * h) * (2 * w)))) (b : Vec c) {G : Vec (N * (c * h * w)) → Vec 1} {cot : Vec (N * (c * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (depthwiseStride2FlatXla W b) x) cot) (o : Fin c) :

XLA-SAME strided depthwise bias node = ∂G/∂b.

theorem Proofs.GradNodeB.denseW_eq_pdiv {N a c : ℕ} (xN cotN : String) (x : Vec (N * a)) (W : Mat a c) (b : Vec c) {G : Vec (N * c) → Vec 1} {cot : Vec (N * c)} (hG : HasGradAt G (StableHLO.batchMap N (dense W b) x) cot) (i : Fin a) (j : Fin c) :

Dense weight node = ∂G/∂W.

theorem Proofs.GradNodeB.dense_bias_differentiable {a c : ℕ} (W : Mat a c) (x : Vec a) :
Differentiable ℝ fun (b' : Vec c) => dense W b' x
theorem Proofs.GradNodeB.denseB_eq_pdiv {N a c : ℕ} (cotN : String) (W : Mat a c) (x₀ : Vec a) (x : Vec (N * a)) (b : Vec c) {G : Vec (N * c) → Vec 1} {cot : Vec (N * c)} (hG : HasGradAt G (StableHLO.batchMap N (dense W b) x) cot) (j : Fin c) :
StableHLO.den (StableHLO.SHlo.operand cotN cot).denseBiasGradB j = pdiv (fun (θ : Vec c) => G (StableHLO.batchMap N (dense W θ) x)) b j 0

Dense bias node = ∂G/∂b. The node's statement carries one activation x₀ for every example; the bias Jacobian is the identity whatever the activation, so any x₀ serves.

noncomputable def Proofs.GradNodeB.bnLAPerm (N oc h w : ℕ) :
Fin (N * (oc * h * w)) ≃ Fin (oc * (N * (h * w)))

The permutation bnBatchLA reads its per-channel core through: network index J ↦ the [C, N·H·W] cell bnchwBackIdx (J at the mul_assoc cast).

Equations
  • One or more equations did not get rendered due to their size.
Instances For
    theorem Proofs.GradNodeB.bnBatchLA_apply_perm (N oc h w : ℕ) (ε : ℝ) (γ β : Vec oc) (v : Vec (N * (oc * h * w))) (J : Fin (N * (oc * h * w))) :
    StableHLO.bnBatchLA N oc h w ε γ β v J = bnPerChannelFlat oc (N * (h * w)) ε γ β (bnchwFwd N oc h w (BackLinks.reassocB N oc h w v)) ((bnLAPerm N oc h w) J)

    bnBatchLA at a network index IS the per-channel core at the permuted cell.

    theorem Proofs.GradNodeB.bnGamma_eq_pdiv {N oc h w : ℕ} (vN epsStr cotN : String) (ε : ℝ) (γ β : Vec oc) (v : Vec (N * (oc * h * w))) {G : Vec (N * (oc * h * w)) → Vec 1} {cot : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.bnBatchLA N oc h w ε γ β v) cot) (c : Fin oc) :
    StableHLO.den (StableHLO.SHlo.bnGammaGradB vN epsStr ε (BackLinks.reassocB N oc h w v) (StableHLO.SHlo.operand cotN (BackLinks.reassocB N oc h w cot))) c = pdiv (fun (θ : Vec oc) => G (StableHLO.bnBatchLA N oc h w ε θ β v)) γ c 0

    BatchNorm γ node = ∂G/∂γ, at the reassocB index the render's node reads.

    theorem Proofs.GradNodeB.bnBeta_eq_pdiv {N oc h w : ℕ} (cotN : String) (ε : ℝ) (γ β : Vec oc) (v : Vec (N * (oc * h * w))) {G : Vec (N * (oc * h * w)) → Vec 1} {cot : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.bnBatchLA N oc h w ε γ β v) cot) (c : Fin oc) :
    StableHLO.den (StableHLO.SHlo.operand cotN (BackLinks.reassocB N oc h w cot)).bnBetaGradB c = pdiv (fun (θ : Vec oc) => G (StableHLO.bnBatchLA N oc h w ε γ θ v)) β c 0

    BatchNorm β node = ∂G/∂β, at the reassocB index.

    theorem Proofs.GradNodeB.pdiv_bias_of_split {a oc h w : ℕ} (per : Vec oc → Vec a → Vec (oc * h * w)) (hsplit : ∀ (θ : Vec oc) (y : Vec a), per θ y = fun (k : Fin (oc * h * w)) => per 0 y k + broadcastFlat oc h w θ k) (y : Vec a) (b : Vec oc) (o : Fin oc) (j : Fin (oc * h * w)) :
    pdiv (fun (θ : Vec oc) => per θ y) b o j = if o = flatChannel oc h w j then 1 else 0

    A channel-broadcast bias's Jacobian is the channel indicator.

    theorem Proofs.GradNodeB.biasBeta_eq_pdiv {N a oc h w : ℕ} (cotN : String) (per : Vec oc → Vec a → Vec (oc * h * w)) (hsplit : ∀ (θ : Vec oc) (y : Vec a), per θ y = fun (k : Fin (oc * h * w)) => per 0 y k + broadcastFlat oc h w θ k) (x : Vec (N * a)) (b : Vec oc) {G : Vec (N * (oc * h * w)) → Vec 1} {cot : Vec (N * (oc * h * w))} (hG : HasGradAt G (StableHLO.batchMap N (per b) x) cot) (o : Fin oc) :
    StableHLO.den (StableHLO.SHlo.operand cotN (BackLinks.reassocB N oc h w cot)).bnBetaGradB o = pdiv (fun (θ : Vec oc) => G (StableHLO.batchMap N (per θ) x)) b o 0

    A channel-broadcast bias, read by the β node, = ∂G/∂b. For any per-example op whose bias enters as per 0 y + broadcast b, the emitted bnBetaGradB on the op's output cotangent is the loss derivative in the bias.

    theorem Proofs.GradNodeB.flatConv_bias_split {ic oc h w kH kW : ℕ} (W : Kernel4 oc ic kH kW) (θ : Vec oc) (y : Vec (ic * h * w)) :
    flatConv W θ y = fun (k : Fin (oc * h * w)) => flatConv W 0 y k + broadcastFlat oc h w θ k

    A conv's bias is a channel broadcast.

    theorem Proofs.GradNodeB.depthwiseFlat_bias_split {c h w kH kW : ℕ} (W : DepthwiseKernel c kH kW) (θ : Vec c) (y : Vec (c * h * w)) :
    depthwiseFlat W θ y = fun (k : Fin (c * h * w)) => depthwiseFlat W 0 y k + broadcastFlat c h w θ k

    A depthwise conv's bias is a channel broadcast.

    theorem Proofs.GradNodeB.depthwiseStride2Flat_bias_split {c h w kH kW : ℕ} (W : DepthwiseKernel c kH kW) (θ : Vec c) (y : Vec (c * (2 * h) * (2 * w))) :
    depthwiseStride2Flat W θ y = fun (k : Fin (c * h * w)) => depthwiseStride2Flat W 0 y k + broadcastFlat c h w θ k

    A symmetric strided depthwise conv's bias is a channel broadcast: decimation keeps channels.

    theorem Proofs.GradNodeB.flatConvStride2Xla_bias_split {ic oc h w kH kW : ℕ} (W : Kernel4 oc ic kH kW) (θ : Vec oc) (y : Vec (ic * (2 * h) * (2 * w))) :
    flatConvStride2Xla W θ y = fun (k : Fin (oc * h * w)) => flatConvStride2Xla W 0 y k + broadcastFlat oc h w θ k

    An XLA-SAME strided conv's bias is a channel broadcast.

    theorem Proofs.GradNodeB.flatConvStride4_bias_split {ic oc h w kH kW : ℕ} (W : Kernel4 oc ic kH kW) (θ : Vec oc) (y : Vec (ic * (2 * (2 * h)) * (2 * (2 * w)))) :
    flatConvStride4 W θ y = fun (k : Fin (oc * h * w)) => flatConvStride4 W 0 y k + broadcastFlat oc h w θ k

    The stride-4 patchify conv's bias is a channel broadcast: both decimations keep channels.