import JaxLean.Verification.Syntax import JaxLean.Verification.Certificate -- Generated by jaxlean from JAX 0.8.0. Edit the source, not this file. -- Real arithmetic abstraction; no claim of IEEE-754 equivalence. import JaxLean.Core.RealOpsopen JaxLeanopen scoped BigOperatorsset_option linter.unusedVariables falsenamespace JaxLean.Blogdef add_then_scale {R : Type} [Field R] (x0 : Tensor R [3]) (x1 : Tensor R [3]) : Tensor R [3] := -- add let v0 : Tensor R [3] := Tensor.map₂ (fun a0 a1 => a0 + a1) x0 x1 -- mul let v1 : Tensor R [3] := Tensor.map (fun a0 => a0 * ((Tensor.scalar (2 : R)) ())) v0 v1end JaxLean.Blog-- IMPORTED IR: the Python importer is trusted to encode the original Jaxpr. namespace JaxLean.Blogdef add_then_scale_ir : Jaxpr.Program [(.real, [3]), (.real, [3])] (.real, [3]) := jaxpr% (a : (.real, [3]), b : (.real, [3])) { c : (.real, [3]) := .add a b (t := [3]); d : (.real, [3]) := .mul c (.literal (2) 1 (c:Jaxpr.Atom [(Jaxpr.DType.real, [3]), (Jaxpr.DType.real, [3]), (Jaxpr.DType.real, [3])] (Jaxpr.DType.real, [3]) := Jaxpr.Atom.var Jaxpr.Var.here⊢ 0 < 1 All goals completed! 🐙)) (t := [3]); return d }-- Relative to Jaxpr.Program.eval's real-arithmetic semantics, for every input. set_option linter.unusedSimpArgs false in theorem add_then_scale_translation_correct (x0 : Tensor ℝ [3]) (x1 : Tensor ℝ [3]) : Jaxpr.Program.eval (.cons x0 (.cons x1 .nil)) _root_.JaxLean.Blog.add_then_scale_ir = _root_.JaxLean.Blog.add_then_scale (R := ℝ) x0 x1 := x0:Tensor ℝ [3]x1:Tensor ℝ [3]⊢ Jaxpr.Program.eval (Jaxpr.Env.cons x0 (Jaxpr.Env.cons x1 Jaxpr.Env.nil)) add_then_scale_ir = add_then_scale x0 x1 x0:Tensor ℝ [3]x1:Tensor ℝ [3]i:Index (Jaxpr.DType.real, [3]).2⊢ Jaxpr.Program.eval (Jaxpr.Env.cons x0 (Jaxpr.Env.cons x1 Jaxpr.Env.nil)) add_then_scale_ir i = add_then_scale x0 x1 i x0:Tensor ℝ [3]x1:Tensor ℝ [3]j0:Fin 3⊢ Jaxpr.Program.eval (Jaxpr.Env.cons x0 (Jaxpr.Env.cons x1 Jaxpr.Env.nil)) add_then_scale_ir (j0, PUnit.unit) = add_then_scale x0 x1 (j0, PUnit.unit) All goals completed! 🐙end JaxLean.Blog