Skip to content

Commit 0c9a8b8

Browse files
Add the pointwise division autograd proof
Add the proof-side division tape node and its Fréchet-derivative correctness theorem under the required nonzero-denominator condition.
1 parent 46658fd commit 0c9a8b8

1 file changed

Lines changed: 150 additions & 0 deletions

File tree

NN/Proofs/Autograd/Tape/Nodes/Arithmetic.lean

Lines changed: 150 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -455,6 +455,156 @@ def squareFderiv {Γ : List Shape} {s : Shape} (idx : Idx Γ s) :
455455
NodeFDerivCorrect (square (Γ := Γ) (s := s) idx) :=
456456
mulFderiv (Γ := Γ) (s := s) idx idx
457457

458+
/--
459+
Pointwise division of two same-shaped context entries.
460+
461+
The forward pass is the coordinatewise quotient `a / b`; the reverse pass is the ordinary quotient
462+
rule `∂(a/b)/∂a = b⁻¹` and `∂(a/b)/∂b = -a · (b²)⁻¹`. This matches the runtime/CUDA `div` node
463+
(whose backward returns `dLdy / b` for the numerator and `-(dLdy · a / b²)` for the denominator).
464+
465+
The `correct_inner` adjoint identity holds unconditionally under the totalized inverse (`0⁻¹ = 0`);
466+
the derivative in `divFderivAt`, by contrast, is the genuine Fréchet derivative only where the
467+
denominator is nonzero, which is the standard mathematical domain of the quotient rule.
468+
-/
469+
def div {Γ : List Shape} {s : Shape} (a b : Idx Γ s) : Node Γ s :=
470+
let n : Nat := Spec.Shape.size s
471+
Node.ofVec (Γ := Γ) (τ := s)
472+
(f := fun x => vecOfFun (n := n) fun i =>
473+
CtxVec.get (Γ := Γ) (s := s) a x i / CtxVec.get (Γ := Γ) (s := s) b x i)
474+
(jvp := fun x dx => vecOfFun (n := n) fun i =>
475+
CtxVec.get (Γ := Γ) (s := s) a dx i * (CtxVec.get (Γ := Γ) (s := s) b x i)⁻¹ -
476+
CtxVec.get (Γ := Γ) (s := s) a x i * CtxVec.get (Γ := Γ) (s := s) b dx i *
477+
((CtxVec.get (Γ := Γ) (s := s) b x i) ^ 2)⁻¹)
478+
(vjp := fun x δ =>
479+
CtxVec.single (Γ := Γ) (s := s) a
480+
(vecOfFun (n := n) fun i => δ i * (CtxVec.get (Γ := Γ) (s := s) b x i)⁻¹) +
481+
CtxVec.single (Γ := Γ) (s := s) b
482+
(vecOfFun (n := n) fun i =>
483+
δ i * (-(CtxVec.get (Γ := Γ) (s := s) a x i) *
484+
((CtxVec.get (Γ := Γ) (s := s) b x i) ^ 2)⁻¹)))
485+
(correct_inner := by
486+
intro x dx δ
487+
classical
488+
rw [inner_add_right, CtxVec.inner_get_single, CtxVec.inner_get_single]
489+
simp only [inner_eq_sum_mul, vecOfFun_apply]
490+
rw [← Finset.sum_add_distrib]
491+
refine Finset.sum_congr rfl (fun i _ => ?_)
492+
ring)
493+
494+
/--
495+
Pointwise `NodeFDerivCorrectAt` for `div` under the assumption that every denominator coordinate is
496+
nonzero. The coordinate derivative is assembled from the product rule applied to `a · b⁻¹`, using
497+
`hasDerivAt_inv` on the denominator projection.
498+
-/
499+
def divFderivAt {Γ : List Shape} {s : Shape} (a b : Idx Γ s) (xV : CtxVec Γ)
500+
(hb : ∀ i : Fin (Spec.Shape.size s), CtxVec.get (Γ := Γ) (s := s) b xV i ≠ 0) :
501+
NodeFDerivCorrectAt (div (Γ := Γ) (s := s) a b) xV := by
502+
classical
503+
let n : Nat := Spec.Shape.size s
504+
refine
505+
{ deriv :=
506+
(euclideanEquiv n).symm.toContinuousLinearMap.comp <|
507+
ContinuousLinearMap.pi (fun i : Fin n =>
508+
(CtxVec.get (Γ := Γ) (s := s) a xV i) •
509+
(-(((CtxVec.get (Γ := Γ) (s := s) b xV i) ^ 2)⁻¹) •
510+
((evalCLM (n := n) i).comp (CtxVec.getCLM (Γ := Γ) (s := s) b))) +
511+
((CtxVec.get (Γ := Γ) (s := s) b xV i)⁻¹) •
512+
((evalCLM (n := n) i).comp (CtxVec.getCLM (Γ := Γ) (s := s) a)))
513+
hasFDerivAt := ?_
514+
jvp_eq := ?_ }
515+
· classical
516+
let aFun : CtxVec Γ → Vec n := fun x => CtxVec.get (Γ := Γ) (s := s) a x
517+
let bFun : CtxVec Γ → Vec n := fun x => CtxVec.get (Γ := Γ) (s := s) b x
518+
have hcoord :
519+
∀ i : Fin n,
520+
HasFDerivAt
521+
((fun x : CtxVec Γ => CtxVec.get (Γ := Γ) (s := s) a x i) *
522+
(fun x : CtxVec Γ => (CtxVec.get (Γ := Γ) (s := s) b x i)⁻¹))
523+
((aFun xV i) •
524+
(-(((CtxVec.get (Γ := Γ) (s := s) b xV i) ^ 2)⁻¹) •
525+
((evalCLM (n := n) i).comp (CtxVec.getCLM (Γ := Γ) (s := s) b))) +
526+
((CtxVec.get (Γ := Γ) (s := s) b xV i)⁻¹) •
527+
((evalCLM (n := n) i).comp (CtxVec.getCLM (Γ := Γ) (s := s) a))) xV := by
528+
intro i
529+
let aCLM : CtxVec Γ →L[ℝ] ℝ :=
530+
(evalCLM (n := n) i).comp (CtxVec.getCLM (Γ := Γ) (s := s) a)
531+
let bCLM : CtxVec Γ →L[ℝ] ℝ :=
532+
(evalCLM (n := n) i).comp (CtxVec.getCLM (Γ := Γ) (s := s) b)
533+
have ha : HasFDerivAt (fun x : CtxVec Γ => CtxVec.get (Γ := Γ) (s := s) a x i) aCLM xV := by
534+
have h0 : HasFDerivAt (fun x : CtxVec Γ => aCLM x) aCLM xV := aCLM.hasFDerivAt (x := xV)
535+
have hEq :
536+
(fun x : CtxVec Γ => CtxVec.get (Γ := Γ) (s := s) a x i) = (fun x : CtxVec Γ => aCLM x) := by
537+
funext x
538+
simp [aCLM, ContinuousLinearMap.comp_apply, evalCLM_apply]
539+
exact (congrArg (fun v : Vec n => v.ofLp i) (CtxVec.getCLM_apply (Γ := Γ) (s := s) a x)).symm
540+
exact h0.congr_of_eventuallyEq hEq.eventuallyEq
541+
have hbder : HasFDerivAt (fun x : CtxVec Γ => CtxVec.get (Γ := Γ) (s := s) b x i) bCLM xV := by
542+
have h0 : HasFDerivAt (fun x : CtxVec Γ => bCLM x) bCLM xV := bCLM.hasFDerivAt (x := xV)
543+
have hEq :
544+
(fun x : CtxVec Γ => CtxVec.get (Γ := Γ) (s := s) b x i) = (fun x : CtxVec Γ => bCLM x) := by
545+
funext x
546+
simp [bCLM, ContinuousLinearMap.comp_apply, evalCLM_apply]
547+
exact (congrArg (fun v : Vec n => v.ofLp i) (CtxVec.getCLM_apply (Γ := Γ) (s := s) b x)).symm
548+
exact h0.congr_of_eventuallyEq hEq.eventuallyEq
549+
have hinv :
550+
HasFDerivAt (fun x : CtxVec Γ => (CtxVec.get (Γ := Γ) (s := s) b x i)⁻¹)
551+
((-((CtxVec.get (Γ := Γ) (s := s) b xV i) ^ 2)⁻¹) • bCLM) xV :=
552+
(hasDerivAt_inv (hb i)).comp_hasFDerivAt xV hbder
553+
have hmul := ha.mul hinv
554+
simpa [aFun, aCLM, bCLM] using hmul
555+
have hpi :
556+
HasFDerivAt
557+
(fun x : CtxVec Γ => fun i : Fin n =>
558+
((fun y : CtxVec Γ => CtxVec.get (Γ := Γ) (s := s) a y i) *
559+
(fun y : CtxVec Γ => (CtxVec.get (Γ := Γ) (s := s) b y i)⁻¹)) x)
560+
(ContinuousLinearMap.pi (fun i : Fin n =>
561+
(aFun xV i) •
562+
(-(((CtxVec.get (Γ := Γ) (s := s) b xV i) ^ 2)⁻¹) •
563+
((evalCLM (n := n) i).comp (CtxVec.getCLM (Γ := Γ) (s := s) b))) +
564+
((CtxVec.get (Γ := Γ) (s := s) b xV i)⁻¹) •
565+
((evalCLM (n := n) i).comp (CtxVec.getCLM (Γ := Γ) (s := s) a)))) xV := by
566+
refine (hasFDerivAt_pi (𝕜 := ℝ)
567+
(φ := fun i : Fin n =>
568+
(fun x : CtxVec Γ => CtxVec.get (Γ := Γ) (s := s) a x i) *
569+
(fun x : CtxVec Γ => (CtxVec.get (Γ := Γ) (s := s) b x i)⁻¹))
570+
(φ' := fun i : Fin n =>
571+
(aFun xV i) •
572+
(-(((CtxVec.get (Γ := Γ) (s := s) b xV i) ^ 2)⁻¹) •
573+
((evalCLM (n := n) i).comp (CtxVec.getCLM (Γ := Γ) (s := s) b))) +
574+
((CtxVec.get (Γ := Γ) (s := s) b xV i)⁻¹) •
575+
((evalCLM (n := n) i).comp (CtxVec.getCLM (Γ := Γ) (s := s) a)))
576+
(x := xV)).2 ?_
577+
intro i
578+
simpa using hcoord i
579+
have he' :
580+
HasFDerivAt (fun g : Fin n → ℝ => (euclideanEquiv n).symm g)
581+
((euclideanEquiv n).symm.toContinuousLinearMap)
582+
(fun i : Fin n =>
583+
((fun y : CtxVec Γ => CtxVec.get (Γ := Γ) (s := s) a y i) *
584+
(fun y : CtxVec Γ => (CtxVec.get (Γ := Γ) (s := s) b y i)⁻¹)) xV) :=
585+
(ContinuousLinearMap.hasFDerivAt (euclideanEquiv n).symm.toContinuousLinearMap)
586+
have hcomp := he'.comp xV hpi
587+
have hEq :
588+
(Node.forwardVec (Γ := Γ) (τ := s) (div (Γ := Γ) (s := s) a b))
589+
= (fun x : CtxVec Γ =>
590+
(euclideanEquiv n).symm (fun i : Fin n =>
591+
((fun y : CtxVec Γ => CtxVec.get (Γ := Γ) (s := s) a y i) *
592+
(fun y : CtxVec Γ => (CtxVec.get (Γ := Γ) (s := s) b y i)⁻¹)) x)) := by
593+
funext x
594+
ext i
595+
simp [div, vecOfFun, Node.forwardVec_ofVec, div_eq_mul_inv, euclideanEquiv]
596+
rw [hEq]
597+
exact hcomp
598+
· intro dxV
599+
ext i
600+
have hga : (CtxVec.getCLM (Γ := Γ) (s := s) a) dxV = CtxVec.get (Γ := Γ) (s := s) a dxV :=
601+
CtxVec.getCLM_apply (Γ := Γ) (s := s) a dxV
602+
have hgb : (CtxVec.getCLM (Γ := Γ) (s := s) b) dxV = CtxVec.get (Γ := Γ) (s := s) b dxV :=
603+
CtxVec.getCLM_apply (Γ := Γ) (s := s) b dxV
604+
simp [div, Node.jvpVec_ofVec, vecOfFun, ContinuousLinearMap.comp_apply, evalCLM_apply,
605+
hga, hgb]
606+
ring
607+
458608
end TapeNodes
459609

460610
end

0 commit comments

Comments
 (0)