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:
HasGradAt G x dy— a scalarGhas gradientdyatx. The loss read at any activation of the net is such aG, and the chain's cotangent there is itsdy.HasGradAt.comp— gradients pull back through a certified VJP:G ∘ fhas gradientf.backward dy. Applied stage by stage, it walks the loss gradient down the net.HasGradAt.pdiv_param/pdiv_param_batchMap— with the gradient at a parameterised op's output known, the loss derivative in the parameter is the op's parameter Jacobian against it: theΣ_n Σ_jevery batched gradient node denotes.
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.
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
- Proofs.addConstHasVJPAt f c x hf hv = { backward := fun (dy : Proofs.Vec n) => hv.backward dy, correct := ⋯ }
Instances For
addConstHasVJPAt with the constant on the left: u ↦ c + f u.
Equations
- Proofs.constAddHasVJPAt c f x hf hv = { backward := fun (dy : Proofs.Vec n) => hv.backward dy, correct := ⋯ }
Instances For
The batched parameterised op θ ↦ batchMap N (per θ) r is differentiable when each
example's map is differentiable in the parameter.
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
- Proofs.HasGradAt G x dy = (DifferentiableAt ℝ G x ∧ ∀ (j : Fin m), Proofs.pdiv G x j 0 = dy j)
Instances For
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.
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.
…at a batched op: θ ↦ batchMap N (per θ) r, the Jacobian split by example — the
Σ_n Σ_j every batched parameter gradient node denotes.
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
- Proofs.linLoss dy u x✝ = ∑ k : Fin m, u k * dy k
Instances For
batchSlice of a batchMapAux is the per-example map at the two slices.
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.