Documentation

LeanMlir.Proofs.Foundation.ParamGrad

ParamGrad — the loss gradient in a parameter, from the gradient at its op's output #

A net's train-step tie says each parameter gradient node is its layer's parameter Jacobian contracted with the cotangent the backward chain threads there, and its *_eq_vjp lemmas say the chain's cotangents are certified VJP backwards. This file is the calculus that composes the two into a derivative of the loss:

addConstHasVJPAt / constAddHasVJPAt are the VJP a residual needs when a parameter inside one branch varies: the other branch is a constant. For a net whose every op is batch-separable, HasGradAt.pdiv_param_batchMap_through does the work per example against linLoss dy.

noncomputable def Proofs.addConstHasVJPAt {m n : ℕ} (f : Vec m → Vec n) (c : Vec n) (x : Vec m) (hf : DifferentiableAt ℝ f x) (hv : HasVJPAt f x) :
HasVJPAt (fun (u : Vec m) (i : Fin n) => f u i + c i) x

Adding a constant keeps the VJP. u ↦ f u + c has f's backward: a residual block's skip is a constant once the parameter being varied sits inside the body.

Equations
Instances For
    theorem Proofs.addConstHasVJPAt_backward {m n : ℕ} (f : Vec m → Vec n) (c : Vec n) (x : Vec m) (hf : DifferentiableAt ℝ f x) (hv : HasVJPAt f x) (dy : Vec n) :
    (addConstHasVJPAt f c x hf hv).backward dy = hv.backward dy
    noncomputable def Proofs.constAddHasVJPAt {m n : ℕ} (c : Vec n) (f : Vec m → Vec n) (x : Vec m) (hf : DifferentiableAt ℝ f x) (hv : HasVJPAt f x) :
    HasVJPAt (fun (u : Vec m) (i : Fin n) => c i + f u i) x

    addConstHasVJPAt with the constant on the left: u ↦ c + f u.

    Equations
    Instances For
      theorem Proofs.batchMap_param_differentiableAt {P N a q : ℕ} (per : Vec P → Vec a → Vec q) (r : Vec (N * a)) (θ : Vec P) (hper : ∀ (y : Vec a), DifferentiableAt ℝ (fun (θ' : Vec P) => per θ' y) θ) :
      DifferentiableAt ℝ (fun (θ' : Vec P) => StableHLO.batchMap N (per θ') r) θ

      The batched parameterised op θ ↦ batchMap N (per θ) r is differentiable when each example's map is differentiable in the parameter.

      def Proofs.HasGradAt {m : ℕ} (G : Vec m → Vec 1) (x dy : Vec m) :

      G : Vec m → Vec 1 has gradient dy at x: differentiable there, and each partial is dy's entry. The loss, read as a function of any activation of the net, is such a G; the backward chain's cotangent at that activation is its dy.

      Equations
      Instances For
        theorem Proofs.HasGradAt.comp {m n : ℕ} {G : Vec n → Vec 1} {f : Vec m → Vec n} {x : Vec m} {dy : Vec n} (hG : HasGradAt G (f x) dy) (hf : DifferentiableAt ℝ f x) (vf : HasVJPAt f x) :
        HasGradAt (fun (y : Vec m) => G (f y)) x (vf.backward dy)

        Gradients pull back through a certified VJP: if G has gradient dy at f x, then G ∘ f has gradient f's backward of dy at x.

        theorem Proofs.HasGradAt.comp_global {m n : ℕ} {G : Vec n → Vec 1} {f : Vec m → Vec n} {x : Vec m} {dy : Vec n} (hG : HasGradAt G (f x) dy) (hf : Differentiable ℝ f) (vf : HasVJP f) :
        HasGradAt (fun (y : Vec m) => G (f y)) x (vf.backward x dy)

        HasGradAt.comp through a global VJP, the cotangent spelled vf.backward x dy. At a large certified VJP the two spellings are definitionally equal but the unifier reaches the equality by unfolding the witness; stated once here, the equality is checked at a variable.

        theorem Proofs.HasGradAt.of_eq {m : ℕ} {G : Vec m → Vec 1} {x dy dy' : Vec m} (hG : HasGradAt G x dy) (h : dy = dy') :
        HasGradAt G x dy'

        Restate a gradient at an equal cotangent.

        theorem Proofs.HasGradAt.congr_point {m : ℕ} {G : Vec m → Vec 1} {x x' dy : Vec m} (h : x = x') (hG : HasGradAt G x dy) :
        HasGradAt G x' dy

        Restate a gradient at an equal point.

        theorem Proofs.HasGradAt.pdiv_param {P m : ℕ} {G : Vec m → Vec 1} {layer : Vec P → Vec m} {θ : Vec P} {dy : Vec m} (hG : HasGradAt G (layer θ) dy) (hl : DifferentiableAt ℝ layer θ) (i : Fin P) :
        pdiv (fun (θ' : Vec P) => G (layer θ')) θ i 0 = ∑ j : Fin m, pdiv layer θ i j * dy j

        A parameter's loss derivative, from the gradient at its op's output.

        theorem Proofs.HasGradAt.pdiv_param_batchMap {P N a q : ℕ} {G : Vec (N * q) → Vec 1} (per : Vec P → Vec a → Vec q) (r : Vec (N * a)) {θ : Vec P} {dy : Vec (N * q)} (hG : HasGradAt G (StableHLO.batchMap N (per θ) r) dy) (hper : ∀ (y : Vec a), DifferentiableAt ℝ (fun (θ' : Vec P) => per θ' y) θ) (i : Fin P) :
        pdiv (fun (θ' : Vec P) => G (StableHLO.batchMap N (per θ') r)) θ i 0 = ∑ n : Fin N, ∑ j : Fin q, pdiv (fun (θ' : Vec P) => per θ' (StableHLO.batchSlice N a r n)) θ i j * StableHLO.batchSlice N q dy n j

        …at a batched op: θ ↦ batchMap N (per θ) r, the Jacobian split by example — the Σ_n Σ_j every batched parameter gradient node denotes.

        noncomputable def Proofs.linLoss {m : ℕ} (dy : Vec m) :
        Vec m → Vec 1

        The linear functional u ↦ ⟨u, dy⟩: its gradient is dy everywhere. Read per example, it turns "the chain's cotangent contracts a stage Jacobian" into a HasGradAt statement.

        Equations
        Instances For
          theorem Proofs.hasGradAt_linLoss {m : ℕ} (dy x : Vec m) :
          HasGradAt (linLoss dy) x dy
          theorem Proofs.batchSlice_batchMapAux {N s a b : ℕ} (f : Vec s → Vec a → Vec b) (aux : Vec (N * s)) (x : Vec (N * a)) (n : Fin N) :

          batchSlice of a batchMapAux is the per-example map at the two slices.

          theorem Proofs.HasGradAt.pdiv_param_batchMap_through {P N a b m q : ℕ} {G : Vec (N * q) → Vec 1} (pre : Vec a → Vec b) (per : Vec P → Vec b → Vec m) (post : Vec a → Vec m → Vec q) (cot : Vec a → Vec q → Vec m) (X : Vec (N * a)) {θ : Vec P} {dY : Vec (N * q)} (hG : HasGradAt G (StableHLO.batchMap N (fun (y : Vec a) => post y (per θ (pre y))) X) dY) (hper : ∀ (y : Vec b), DifferentiableAt ℝ (fun (θ' : Vec P) => per θ' y) θ) (hpost : ∀ (y : Vec a), Differentiable ℝ (post y)) (hcot : ∀ (y : Vec a) (dy : Vec q), HasGradAt (fun (u : Vec m) => linLoss dy (post y u)) (per θ (pre y)) (cot y dy)) (A : Vec (N * b)) (COT : Vec (N * m)) (hA : ∀ (n : Fin N), StableHLO.batchSlice N b A n = pre (StableHLO.batchSlice N a X n)) (hC : ∀ (n : Fin N), StableHLO.batchSlice N m COT n = cot (StableHLO.batchSlice N a X n) (StableHLO.batchSlice N q dY n)) (i : Fin P) :
          ∑ n : Fin N, ∑ j : Fin m, pdiv (fun (θ' : Vec P) => per θ' (StableHLO.batchSlice N b A n)) θ i j * StableHLO.batchSlice N m COT n j = pdiv (fun (θ' : Vec P) => G (StableHLO.batchMap N (fun (y : Vec a) => post y (per θ' (pre y))) X)) θ i 0

          A parameter inside a per-example stage, lifted over the batch. Each example runs y ↦ post y (per θ (pre y)): the stage per θ at its input pre y, then the rest of the block post y. If, per example, the loss ⟨post y ·, dy⟩ has gradient cot y dy at the stage output, then the batched node Σ_n Σ_j ∂per/∂θ · cotₙ — at any saved activation A and cotangent COT whose slices are pre yₙ and cot yₙ dyₙ — is ∂G/∂θ of the whole batched block.