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.AutodiffUsed `tac1 <;> tac2` where `(tac1; tac2)` would suffice Note: This linter can be disabled with `set_option linter.unnecessarySeqFocus false` 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 Used `tac1 <;> tac2` where `(tac1; tac2)` would suffice Note: This linter can be disabled with `set_option linter.unnecessarySeqFocus false`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 x:ℝ⊢ 3 + x + x = 2 * x + 3 All goals completed! 🐙x:ℝv:ℝ⊢ jvp_tangent (Tensor.scalar x) (Tensor.scalar v) () = (2 * x + 3) * v x:ℝv:ℝ⊢ v * x + x * v + 3 * v = (2 * x + 3) * v 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 := 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 x:ℝ⊢ gradient (Tensor.scalar x) () = deriv (fun y => quadratic (Tensor.scalar y) ()) 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 := 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 x:ℝv:ℝ⊢ jvp_tangent (Tensor.scalar x) (Tensor.scalar v) () = deriv (fun y => quadratic (Tensor.scalar y) ()) x * v All goals completed! 🐙end JaxLean.Autodiff