import JaxLean.Verification.Certificate import JaxLean.Core.RealOps-- Generated by jaxlean from JAX 0.8.0. Edit the source, not this file. -- Real arithmetic abstraction; no claim of IEEE-754 equivalence. open JaxLeanopen scoped BigOperatorsset_option linter.unusedVariables falsenamespace JaxLean.PuzzleJaxdef puzzle_outer {R : Type} [Field R] (a : Tensor R [2]) (b : Tensor R [3]) : Tensor R [2, 3] := -- code.py:20 (puzzle_outer) -- broadcast_in_dim let broadcast_in_dim_result : Tensor R [2, 1] := Tensor.reindex (s := [2]) (fun i => (i.1, ())) a -- code.py:20 (puzzle_outer) -- broadcast_in_dim let broadcast_in_dim_result_2 : Tensor R [1, 3] := Tensor.reindex (s := [3]) (fun i => (i.2.1, ())) b -- code.py:20 (puzzle_outer) -- mul let mul_result : Tensor R [2, 3] := Tensor.map₂ (fun a0 a1 => a0 * a1) (fun i => (broadcast_in_dim_result (i.1, 0, ()))) (Tensor.broadcastFirst 2 broadcast_in_dim_result_2) mul_resultend JaxLean.PuzzleJax-- IMPORTED IR: the Python importer is trusted to encode the original Jaxpr. namespace JaxLean.PuzzleJaxdef puzzle_outer_ir : Jaxpr.Program [(.real, [2]), (.real, [3])] (.real, [2, 3]) := .bind (.broadcast_in_dim [2, 1] [0] (.var .here)) <| .bind (.broadcast_in_dim [1, 3] [1] (.var (.there (.there .here)))) <| .bind (.mul (.var (.there .here)) (.var .here) (t := [2, 3])) <| .ret (.var .here)-- Relative to Jaxpr.Program.eval's real-arithmetic semantics, for every input. set_option linter.unusedSimpArgs false in theorem puzzle_outer_translation_correct (a : Tensor ℝ [2]) (b : Tensor ℝ [3]) : Jaxpr.Program.eval (.cons a (.cons b .nil)) _root_.JaxLean.PuzzleJax.puzzle_outer_ir = _root_.JaxLean.PuzzleJax.puzzle_outer (R := ℝ) a b := a:Tensor ℝ [2]b:Tensor ℝ [3]⊢ Jaxpr.Program.eval (Jaxpr.Env.cons a (Jaxpr.Env.cons b Jaxpr.Env.nil)) puzzle_outer_ir = puzzle_outer a b a:Tensor ℝ [2]b:Tensor ℝ [3]i:Index (Jaxpr.DType.real, [2, 3]).2⊢ Jaxpr.Program.eval (Jaxpr.Env.cons a (Jaxpr.Env.cons b Jaxpr.Env.nil)) puzzle_outer_ir i = puzzle_outer a b i a:Tensor ℝ [2]b:Tensor ℝ [3]j0:Fin 2j1:Fin 3⊢ Jaxpr.Program.eval (Jaxpr.Env.cons a (Jaxpr.Env.cons b Jaxpr.Env.nil)) puzzle_outer_ir (j0, j1, PUnit.unit) = puzzle_outer a b (j0, j1, PUnit.unit) All goals completed! 🐙end JaxLean.PuzzleJax-- Generated by jaxlean from JAX 0.8.0. Edit the source, not this file. -- Real arithmetic abstraction; no claim of IEEE-754 equivalence. open JaxLeanopen scoped BigOperatorsset_option linter.unusedVariables falsenamespace JaxLean.PuzzleJaxdef puzzle_flatten {R : Type} [Field R] (a : Tensor R [2, 3]) : Tensor R [6] := -- code.py:30 (puzzle_flatten) -- reshape let reshape_result : Tensor R [6] := Tensor.reshape (t := [6]) (R:Typeinst✝:Field Ra:Tensor R [2, 3]⊢ [2, 3].prod = [6].prod All goals completed! 🐙) (a) reshape_resultend JaxLean.PuzzleJax-- IMPORTED IR: the Python importer is trusted to encode the original Jaxpr. namespace JaxLean.PuzzleJaxdef puzzle_flatten_ir : Jaxpr.Program [(.real, [2, 3])] (.real, [6]) := .bind (.reshape (s := [2, 3]) (t := [6]) (fun i => (((Index.equivFin [2, 3]).symm (Fin.cast (i:Index [6]⊢ [6].prod = [2, 3].prod All goals completed! 🐙) (Index.equivFin [6] i))).1, ((Index.equivFin [2, 3]).symm (Fin.cast (i:Index [6]⊢ [6].prod = [2, 3].prod All goals completed! 🐙) (Index.equivFin [6] i))).2.1, ())) (.var .here)) <| .ret (.var .here)-- Relative to Jaxpr.Program.eval's real-arithmetic semantics, for every input. set_option linter.unusedSimpArgs false in theorem puzzle_flatten_translation_correct (a : Tensor ℝ [2, 3]) : Jaxpr.Program.eval (.cons a .nil) _root_.JaxLean.PuzzleJax.puzzle_flatten_ir = _root_.JaxLean.PuzzleJax.puzzle_flatten (R := ℝ) a := a:Tensor ℝ [2, 3]⊢ Jaxpr.Program.eval (Jaxpr.Env.cons a Jaxpr.Env.nil) puzzle_flatten_ir = puzzle_flatten a a:Tensor ℝ [2, 3]i:Index (Jaxpr.DType.real, [6]).2⊢ Jaxpr.Program.eval (Jaxpr.Env.cons a Jaxpr.Env.nil) puzzle_flatten_ir i = puzzle_flatten a i a:Tensor ℝ [2, 3]j0:Fin 6⊢ Jaxpr.Program.eval (Jaxpr.Env.cons a Jaxpr.Env.nil) puzzle_flatten_ir (j0, PUnit.unit) = puzzle_flatten a (j0, PUnit.unit) All goals completed! 🐙end JaxLean.PuzzleJax-- Generated by jaxlean from JAX 0.8.0. Edit the source, not this file. -- Real arithmetic abstraction; no claim of IEEE-754 equivalence. open JaxLeanopen scoped BigOperatorsset_option linter.unusedVariables falsenamespace JaxLean.PuzzleJaxdef outer_flatten {R : Type} [Field R] [LinearOrder R] (a : Tensor R [2]) (b : Tensor R [3]) : Tensor R [6] := -- code.py:34 (outer_flatten) -- call puzzle_outer let call_puzzle_outer_result : Tensor R [2, 3] := _root_.JaxLean.PuzzleJax.puzzle_outer (R := R) (a) (b) -- code.py:34 (outer_flatten) -- call puzzle_flatten let call_puzzle_flatten_result : Tensor R [6] := _root_.JaxLean.PuzzleJax.puzzle_flatten (R := R) (call_puzzle_outer_result) call_puzzle_flatten_resultend JaxLean.PuzzleJax-- IMPORTED IR: the Python importer is trusted to encode the original Jaxpr. namespace JaxLean.PuzzleJaxdef outer_flatten_ir : Jaxpr.Program [(.real, [2]), (.real, [3])] (.real, [6]) := .call puzzle_outer_ir (.cons (.var .here) (.cons (.var (.there .here)) .nil)) <| .call puzzle_flatten_ir (.cons (.var .here) .nil) <| .ret (.var .here)-- Relative to Jaxpr.Program.eval's real-arithmetic semantics, for every input. set_option linter.unusedSimpArgs false in theorem outer_flatten_translation_correct (a : Tensor ℝ [2]) (b : Tensor ℝ [3]) : Jaxpr.Program.eval (.cons a (.cons b .nil)) _root_.JaxLean.PuzzleJax.outer_flatten_ir = _root_.JaxLean.PuzzleJax.outer_flatten (R := ℝ) a b := a:Tensor ℝ [2]b:Tensor ℝ [3]⊢ Jaxpr.Program.eval (Jaxpr.Env.cons a (Jaxpr.Env.cons b Jaxpr.Env.nil)) outer_flatten_ir = outer_flatten a b a:Tensor ℝ [2]b:Tensor ℝ [3]i:Index (Jaxpr.DType.real, [6]).2⊢ Jaxpr.Program.eval (Jaxpr.Env.cons a (Jaxpr.Env.cons b Jaxpr.Env.nil)) outer_flatten_ir i = outer_flatten a b i a:Tensor ℝ [2]b:Tensor ℝ [3]j0:Fin 6⊢ Jaxpr.Program.eval (Jaxpr.Env.cons a (Jaxpr.Env.cons b Jaxpr.Env.nil)) outer_flatten_ir (j0, PUnit.unit) = outer_flatten a b (j0, PUnit.unit) All goals completed! 🐙end JaxLean.PuzzleJax