@@ -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+
458608end TapeNodes
459609
460610end
0 commit comments