import examples.autodiff.generated.Quadratic
import examples.autodiff.generated.JvpTangent
import examples.autodiff.generated.Gradient
import Mathlib.Analysis.Calculus.Deriv.Mul
import Mathlib.Analysis.Calculus.Deriv.Addnamespace JaxLean.Autodiff
theorem quadratic_hasDerivAt (x : ℝ) :
HasDerivAt (fun y : ℝ => quadratic (Tensor.scalar y) ()) (2 * x + 3) x := x:ℝ⊢ HasDerivAt (fun y => quadratic (Tensor.scalar y) ()) (2 * x + 3) x
x:ℝh:HasDerivAt (fun y => y * y + 3 * y) (1 * x + x * 1 + 3 * 1) x⊢ HasDerivAt (fun y => quadratic (Tensor.scalar y) ()) (2 * x + 3) x
x:ℝh:HasDerivAt (fun y => y * y + 3 * y) (1 * x + x * 1 + 3 * 1) x⊢ (fun y => quadratic (Tensor.scalar y) ()) = fun y => y * y + 3 * yx:ℝh:HasDerivAt (fun y => y * y + 3 * y) (1 * x + x * 1 + 3 * 1) x⊢ 2 * x + 3 = 1 * x + x * 1 + 3 * 1 x:ℝh:HasDerivAt (fun y => y * y + 3 * y) (1 * x + x * 1 + 3 * 1) x⊢ (fun y => quadratic (Tensor.scalar y) ()) = fun y => y * y + 3 * yx:ℝh:HasDerivAt (fun y => y * y + 3 * y) (1 * x + x * 1 + 3 * 1) x⊢ 2 * x + 3 = 1 * x + x * 1 + 3 * 1 x:ℝh:HasDerivAt (fun y => y * y + 3 * y) (1 * x + x * 1 + 3 * 1) x⊢ 2 * x = x + x x:ℝh:HasDerivAt (fun y => y * y + 3 * y) (1 * x + x * 1 + 3 * 1) x⊢ 2 * x = x + x All goals completed! 🐙x:ℝ⊢ gradient (Tensor.scalar x) () = 2 * x + 3
simp [gradient, Tensor.scalar] x:ℝ⊢ 3 + x + x = 2 * x + 3
ring All goals completed! 🐙
theorem jvp_correct (x v : ℝ) :
jvp_tangent (Tensor.scalar x) (Tensor.scalar v) () =
deriv (fun y : ℝ => quadratic (Tensor.scalar y) ()) x * v := by x:ℝv:ℝ⊢ jvp_tangent (Tensor.scalar x) (Tensor.scalar v) () = deriv (fun y => quadratic (Tensor.scalar y) ()) x * v
rw [(quadratic_hasDerivAt x).deriv x:ℝv:ℝ⊢ jvp_tangent (Tensor.scalar x) (Tensor.scalar v) () = (2 * x + 3) * v x:ℝv:ℝ⊢ jvp_tangent (Tensor.scalar x) (Tensor.scalar v) () = (2 * x + 3) * v] x:ℝv:ℝ⊢ jvp_tangent (Tensor.scalar x) (Tensor.scalar v) () = (2 * x + 3) * v
simp [jvp_tangent, Tensor.scalar] x:ℝv:ℝ⊢ v * x + x * v + 3 * v = (2 * x + 3) * v
ring All goals completed! 🐙 theorem certified_gradient_correct (x : ℝ) :
Jaxpr.Program.eval (.cons (Tensor.scalar x) .nil) gradient_ir () =
deriv (fun y : ℝ =>
Jaxpr.Program.eval (.cons (Tensor.scalar y) .nil) quadratic_ir ()) x := by x:ℝ⊢ Jaxpr.Program.eval (Jaxpr.Env.cons (Tensor.scalar x) Jaxpr.Env.nil) gradient_ir () =
deriv (fun y => Jaxpr.Program.eval (Jaxpr.Env.cons (Tensor.scalar y) Jaxpr.Env.nil) quadratic_ir ()) x
simp only [gradient_translation_correct, quadratic_translation_correct] x:ℝ⊢ gradient (Tensor.scalar x) () = deriv (fun y => quadratic (Tensor.scalar y) ()) x
exact gradient_correct x All goals completed! 🐙 theorem certified_jvp_correct (x v : ℝ) :
Jaxpr.Program.eval (.cons (Tensor.scalar x) (.cons (Tensor.scalar v) .nil))
jvp_tangent_ir () =
deriv (fun y : ℝ =>
Jaxpr.Program.eval (.cons (Tensor.scalar y) .nil) quadratic_ir ()) x * v := by x:ℝv:ℝ⊢ Jaxpr.Program.eval (Jaxpr.Env.cons (Tensor.scalar x) (Jaxpr.Env.cons (Tensor.scalar v) Jaxpr.Env.nil)) jvp_tangent_ir
() =
deriv (fun y => Jaxpr.Program.eval (Jaxpr.Env.cons (Tensor.scalar y) Jaxpr.Env.nil) quadratic_ir ()) x * v
simp only [jvp_tangent_translation_correct, quadratic_translation_correct] x:ℝv:ℝ⊢ jvp_tangent (Tensor.scalar x) (Tensor.scalar v) () = deriv (fun y => quadratic (Tensor.scalar y) ()) x * v
exact jvp_correct x v All goals completed! 🐙end JaxLean.Autodiff