Skip to content
import JaxLean.Stdlib
import examples.tensor_puzzles.proofs.TensorPuzzleProofs
import examples.tensor_puzzles.generated.LoopSum
import examples.tensor_puzzles.generated.LoopOuter
import examples.tensor_puzzles.generated.LoopFlip
import examples.tensor_puzzles.generated.LoopFlattennamespace JaxLean.PuzzleJaxopen scoped BigOperatorsunita:Tensor ℝ [4]⊢ ∑ i, a (i, ()) = loop_sum a PUnit.unit
change (∑ i : Fin 4, a (i, ())) =
(((0 + a (0, ())) + a (1, ())) + a (2, ())) + a (3, ())unita:Tensor ℝ [4]⊢ ∑ i, a (i, ()) = 0 + a (0, ()) + a (1, ()) + a (2, ()) + a (3, ())
simp [Fin.sum_univ_succ]unita:Tensor ℝ [4]⊢ a (0, ()) + (a (1, ()) + (a (2, ()) + a (3, ()))) = a (0, ()) + a (1, ()) + a (2, ()) + a (3, ())
ring_nfAll goals completed! 🐙theorem outer_matches_loop (a : Tensor ℝ [2]) (b : Tensor ℝ [3]) :
puzzle_outer a b = loop_outer a b := bya:Tensor ℝ [2]b:Tensor ℝ [3]⊢ puzzle_outer a b = loop_outer a b
funext ia:Tensor ℝ [2]b:Tensor ℝ [3]i:Index [2, 3]⊢ puzzle_outer a b i = loop_outer a b i
rcases i with ⟨i, j, ⟨⟩⟩a:Tensor ℝ [2]b:Tensor ℝ [3]i:Fin 2j:Fin 3⊢ puzzle_outer a b (i, j, PUnit.unit) = loop_outer a b (i, j, PUnit.unit)
fin_cases i«0»a:Tensor ℝ [2]b:Tensor ℝ [3]j:Fin 3⊢ puzzle_outer a b ((fun i => i) ⟨0, ⋯⟩, j, PUnit.unit) = loop_outer a b ((fun i => i) ⟨0, ⋯⟩, j, PUnit.unit)«1»a:Tensor ℝ [2]b:Tensor ℝ [3]j:Fin 3⊢ puzzle_outer a b ((fun i => i) ⟨1, ⋯⟩, j, PUnit.unit) = loop_outer a b ((fun i => i) ⟨1, ⋯⟩, j, PUnit.unit) <;>«0»a:Tensor ℝ [2]b:Tensor ℝ [3]j:Fin 3⊢ puzzle_outer a b ((fun i => i) ⟨0, ⋯⟩, j, PUnit.unit) = loop_outer a b ((fun i => i) ⟨0, ⋯⟩, j, PUnit.unit)«1»a:Tensor ℝ [2]b:Tensor ℝ [3]j:Fin 3⊢ puzzle_outer a b ((fun i => i) ⟨1, ⋯⟩, j, PUnit.unit) = loop_outer a b ((fun i => i) ⟨1, ⋯⟩, j, PUnit.unit) fin_cases j«1».«0»a:Tensor ℝ [2]b:Tensor ℝ [3]⊢ puzzle_outer a b ((fun i => i) ⟨1, ⋯⟩, (fun i => i) ⟨0, ⋯⟩, PUnit.unit) =
loop_outer a b ((fun i => i) ⟨1, ⋯⟩, (fun i => i) ⟨0, ⋯⟩, PUnit.unit)«1».«1»a:Tensor ℝ [2]b:Tensor ℝ [3]⊢ puzzle_outer a b ((fun i => i) ⟨1, ⋯⟩, (fun i => i) ⟨1, ⋯⟩, PUnit.unit) =
loop_outer a b ((fun i => i) ⟨1, ⋯⟩, (fun i => i) ⟨1, ⋯⟩, PUnit.unit)«1».«2»a:Tensor ℝ [2]b:Tensor ℝ [3]⊢ puzzle_outer a b ((fun i => i) ⟨1, ⋯⟩, (fun i => i) ⟨2, ⋯⟩, PUnit.unit) =
loop_outer a b ((fun i => i) ⟨1, ⋯⟩, (fun i => i) ⟨2, ⋯⟩, PUnit.unit) <;>«0».«0»a:Tensor ℝ [2]b:Tensor ℝ [3]⊢ puzzle_outer a b ((fun i => i) ⟨0, ⋯⟩, (fun i => i) ⟨0, ⋯⟩, PUnit.unit) =
loop_outer a b ((fun i => i) ⟨0, ⋯⟩, (fun i => i) ⟨0, ⋯⟩, PUnit.unit)«0».«1»a:Tensor ℝ [2]b:Tensor ℝ [3]⊢ puzzle_outer a b ((fun i => i) ⟨0, ⋯⟩, (fun i => i) ⟨1, ⋯⟩, PUnit.unit) =
loop_outer a b ((fun i => i) ⟨0, ⋯⟩, (fun i => i) ⟨1, ⋯⟩, PUnit.unit)«0».«2»a:Tensor ℝ [2]b:Tensor ℝ [3]⊢ puzzle_outer a b ((fun i => i) ⟨0, ⋯⟩, (fun i => i) ⟨2, ⋯⟩, PUnit.unit) =
loop_outer a b ((fun i => i) ⟨0, ⋯⟩, (fun i => i) ⟨2, ⋯⟩, PUnit.unit)«1».«0»a:Tensor ℝ [2]b:Tensor ℝ [3]⊢ puzzle_outer a b ((fun i => i) ⟨1, ⋯⟩, (fun i => i) ⟨0, ⋯⟩, PUnit.unit) =
loop_outer a b ((fun i => i) ⟨1, ⋯⟩, (fun i => i) ⟨0, ⋯⟩, PUnit.unit)«1».«1»a:Tensor ℝ [2]b:Tensor ℝ [3]⊢ puzzle_outer a b ((fun i => i) ⟨1, ⋯⟩, (fun i => i) ⟨1, ⋯⟩, PUnit.unit) =
loop_outer a b ((fun i => i) ⟨1, ⋯⟩, (fun i => i) ⟨1, ⋯⟩, PUnit.unit)«1».«2»a:Tensor ℝ [2]b:Tensor ℝ [3]⊢ puzzle_outer a b ((fun i => i) ⟨1, ⋯⟩, (fun i => i) ⟨2, ⋯⟩, PUnit.unit) =
loop_outer a b ((fun i => i) ⟨1, ⋯⟩, (fun i => i) ⟨2, ⋯⟩, PUnit.unit) rflAll goals completed! 🐙theorem flip_matches_loop (a : Tensor ℝ [4]) : flip a = loop_flip a := bya:Tensor ℝ [4]⊢ flip a = loop_flip a
funext ia:Tensor ℝ [4]i:Index [4]⊢ flip a i = loop_flip a i
rcases i with ⟨i, ⟨⟩⟩a:Tensor ℝ [4]i:Fin 4⊢ flip a (i, PUnit.unit) = loop_flip a (i, PUnit.unit)
fin_cases i«0»a:Tensor ℝ [4]⊢ flip a ((fun i => i) ⟨0, ⋯⟩, PUnit.unit) = loop_flip a ((fun i => i) ⟨0, ⋯⟩, PUnit.unit)«1»a:Tensor ℝ [4]⊢ flip a ((fun i => i) ⟨1, ⋯⟩, PUnit.unit) = loop_flip a ((fun i => i) ⟨1, ⋯⟩, PUnit.unit)«2»a:Tensor ℝ [4]⊢ flip a ((fun i => i) ⟨2, ⋯⟩, PUnit.unit) = loop_flip a ((fun i => i) ⟨2, ⋯⟩, PUnit.unit)«3»a:Tensor ℝ [4]⊢ flip a ((fun i => i) ⟨3, ⋯⟩, PUnit.unit) = loop_flip a ((fun i => i) ⟨3, ⋯⟩, PUnit.unit) <;>«0»a:Tensor ℝ [4]⊢ flip a ((fun i => i) ⟨0, ⋯⟩, PUnit.unit) = loop_flip a ((fun i => i) ⟨0, ⋯⟩, PUnit.unit)«1»a:Tensor ℝ [4]⊢ flip a ((fun i => i) ⟨1, ⋯⟩, PUnit.unit) = loop_flip a ((fun i => i) ⟨1, ⋯⟩, PUnit.unit)«2»a:Tensor ℝ [4]⊢ flip a ((fun i => i) ⟨2, ⋯⟩, PUnit.unit) = loop_flip a ((fun i => i) ⟨2, ⋯⟩, PUnit.unit)«3»a:Tensor ℝ [4]⊢ flip a ((fun i => i) ⟨3, ⋯⟩, PUnit.unit) = loop_flip a ((fun i => i) ⟨3, ⋯⟩, PUnit.unit) rflAll goals completed! 🐙
theorem flatten_matches_loop (a : Tensor ℝ [2, 3]) :
puzzle_flatten a = loop_flatten a := bya:Tensor ℝ [2, 3]⊢ puzzle_flatten a = loop_flatten a
rw [flatten_matches_pseudocode,a:Tensor ℝ [2, 3]⊢ Pseudocode.flatten a = loop_flatten aa:Tensor ℝ [2, 3]⊢ (fun i => a (i.1.divNat, i.1.modNat, ())) = loop_flatten a Pseudocode.flatten_eqa:Tensor ℝ [2, 3]⊢ (fun i => a (i.1.divNat, i.1.modNat, ())) = loop_flatten aa:Tensor ℝ [2, 3]⊢ (fun i => a (i.1.divNat, i.1.modNat, ())) = loop_flatten a]a:Tensor ℝ [2, 3]⊢ (fun i => a (i.1.divNat, i.1.modNat, ())) = loop_flatten a
funext ia:Tensor ℝ [2, 3]i:Index [2 * 3]⊢ a (i.1.divNat, i.1.modNat, ()) = loop_flatten a i
rcases i with ⟨i, ⟨⟩⟩a:Tensor ℝ [2, 3]i:Fin (2 * 3)⊢ a ((i, PUnit.unit).1.divNat, (i, PUnit.unit).1.modNat, ()) = loop_flatten a (i, PUnit.unit)
fin_cases i«0»a:Tensor ℝ [2, 3]⊢ a (((fun i => i) ⟨0, ⋯⟩, PUnit.unit).1.divNat, ((fun i => i) ⟨0, ⋯⟩, PUnit.unit).1.modNat, ()) =
loop_flatten a ((fun i => i) ⟨0, ⋯⟩, PUnit.unit)«1»a:Tensor ℝ [2, 3]⊢ a (((fun i => i) ⟨1, ⋯⟩, PUnit.unit).1.divNat, ((fun i => i) ⟨1, ⋯⟩, PUnit.unit).1.modNat, ()) =
loop_flatten a ((fun i => i) ⟨1, ⋯⟩, PUnit.unit)«2»a:Tensor ℝ [2, 3]⊢ a (((fun i => i) ⟨2, ⋯⟩, PUnit.unit).1.divNat, ((fun i => i) ⟨2, ⋯⟩, PUnit.unit).1.modNat, ()) =
loop_flatten a ((fun i => i) ⟨2, ⋯⟩, PUnit.unit)«3»a:Tensor ℝ [2, 3]⊢ a (((fun i => i) ⟨3, ⋯⟩, PUnit.unit).1.divNat, ((fun i => i) ⟨3, ⋯⟩, PUnit.unit).1.modNat, ()) =
loop_flatten a ((fun i => i) ⟨3, ⋯⟩, PUnit.unit)«4»a:Tensor ℝ [2, 3]⊢ a (((fun i => i) ⟨4, ⋯⟩, PUnit.unit).1.divNat, ((fun i => i) ⟨4, ⋯⟩, PUnit.unit).1.modNat, ()) =
loop_flatten a ((fun i => i) ⟨4, ⋯⟩, PUnit.unit)«5»a:Tensor ℝ [2, 3]⊢ a (((fun i => i) ⟨5, ⋯⟩, PUnit.unit).1.divNat, ((fun i => i) ⟨5, ⋯⟩, PUnit.unit).1.modNat, ()) =
loop_flatten a ((fun i => i) ⟨5, ⋯⟩, PUnit.unit) <;>«0»a:Tensor ℝ [2, 3]⊢ a (((fun i => i) ⟨0, ⋯⟩, PUnit.unit).1.divNat, ((fun i => i) ⟨0, ⋯⟩, PUnit.unit).1.modNat, ()) =
loop_flatten a ((fun i => i) ⟨0, ⋯⟩, PUnit.unit)«1»a:Tensor ℝ [2, 3]⊢ a (((fun i => i) ⟨1, ⋯⟩, PUnit.unit).1.divNat, ((fun i => i) ⟨1, ⋯⟩, PUnit.unit).1.modNat, ()) =
loop_flatten a ((fun i => i) ⟨1, ⋯⟩, PUnit.unit)«2»a:Tensor ℝ [2, 3]⊢ a (((fun i => i) ⟨2, ⋯⟩, PUnit.unit).1.divNat, ((fun i => i) ⟨2, ⋯⟩, PUnit.unit).1.modNat, ()) =
loop_flatten a ((fun i => i) ⟨2, ⋯⟩, PUnit.unit)«3»a:Tensor ℝ [2, 3]⊢ a (((fun i => i) ⟨3, ⋯⟩, PUnit.unit).1.divNat, ((fun i => i) ⟨3, ⋯⟩, PUnit.unit).1.modNat, ()) =
loop_flatten a ((fun i => i) ⟨3, ⋯⟩, PUnit.unit)«4»a:Tensor ℝ [2, 3]⊢ a (((fun i => i) ⟨4, ⋯⟩, PUnit.unit).1.divNat, ((fun i => i) ⟨4, ⋯⟩, PUnit.unit).1.modNat, ()) =
loop_flatten a ((fun i => i) ⟨4, ⋯⟩, PUnit.unit)«5»a:Tensor ℝ [2, 3]⊢ a (((fun i => i) ⟨5, ⋯⟩, PUnit.unit).1.divNat, ((fun i => i) ⟨5, ⋯⟩, PUnit.unit).1.modNat, ()) =
loop_flatten a ((fun i => i) ⟨5, ⋯⟩, PUnit.unit) rflAll goals completed! 🐙
theorem certified_sum_loop (a : Tensor ℝ [4]) :
Jaxpr.Program.eval (.cons a .nil) sum_ir =
Jaxpr.Program.eval (.cons a .nil) loop_sum_ir := bya:Tensor ℝ [4]⊢ Jaxpr.Program.eval (Jaxpr.Env.cons a Jaxpr.Env.nil) sum_ir =
Jaxpr.Program.eval (Jaxpr.Env.cons a Jaxpr.Env.nil) loop_sum_ir
rw [sum_translation_correct,a:Tensor ℝ [4]⊢ sum a = Jaxpr.Program.eval (Jaxpr.Env.cons a Jaxpr.Env.nil) loop_sum_irAll goals completed! 🐙 loop_sum_translation_correct,a:Tensor ℝ [4]⊢ sum a = loop_sum aAll goals completed! 🐙 sum_matches_loopa:Tensor ℝ [4]⊢ loop_sum a = loop_sum aAll goals completed! 🐙]All goals completed! 🐙
theorem certified_outer_loop (a : Tensor ℝ [2]) (b : Tensor ℝ [3]) :
Jaxpr.Program.eval (.cons a (.cons b .nil)) puzzle_outer_ir =
Jaxpr.Program.eval (.cons a (.cons b .nil)) loop_outer_ir := bya:Tensor ℝ [2]b:Tensor ℝ [3]⊢ Jaxpr.Program.eval (Jaxpr.Env.cons a (Jaxpr.Env.cons b Jaxpr.Env.nil)) puzzle_outer_ir =
Jaxpr.Program.eval (Jaxpr.Env.cons a (Jaxpr.Env.cons b Jaxpr.Env.nil)) loop_outer_ir
rw [puzzle_outer_translation_correct,a:Tensor ℝ [2]b:Tensor ℝ [3]⊢ puzzle_outer a b = Jaxpr.Program.eval (Jaxpr.Env.cons a (Jaxpr.Env.cons b Jaxpr.Env.nil)) loop_outer_irAll goals completed! 🐙 loop_outer_translation_correct,a:Tensor ℝ [2]b:Tensor ℝ [3]⊢ puzzle_outer a b = loop_outer a bAll goals completed! 🐙 outer_matches_loopa:Tensor ℝ [2]b:Tensor ℝ [3]⊢ loop_outer a b = loop_outer a bAll goals completed! 🐙]All goals completed! 🐙
theorem certified_flip_loop (a : Tensor ℝ [4]) :
Jaxpr.Program.eval (.cons a .nil) flip_ir =
Jaxpr.Program.eval (.cons a .nil) loop_flip_ir := bya:Tensor ℝ [4]⊢ Jaxpr.Program.eval (Jaxpr.Env.cons a Jaxpr.Env.nil) flip_ir =
Jaxpr.Program.eval (Jaxpr.Env.cons a Jaxpr.Env.nil) loop_flip_ir
rw [flip_translation_correct,a:Tensor ℝ [4]⊢ flip a = Jaxpr.Program.eval (Jaxpr.Env.cons a Jaxpr.Env.nil) loop_flip_irAll goals completed! 🐙 loop_flip_translation_correct,a:Tensor ℝ [4]⊢ flip a = loop_flip aAll goals completed! 🐙 flip_matches_loopa:Tensor ℝ [4]⊢ loop_flip a = loop_flip aAll goals completed! 🐙]All goals completed! 🐙
theorem certified_flatten_loop (a : Tensor ℝ [2, 3]) :
Jaxpr.Program.eval (.cons a .nil) puzzle_flatten_ir =
Jaxpr.Program.eval (.cons a .nil) loop_flatten_ir := bya:Tensor ℝ [2, 3]⊢ Jaxpr.Program.eval (Jaxpr.Env.cons a Jaxpr.Env.nil) puzzle_flatten_ir =
Jaxpr.Program.eval (Jaxpr.Env.cons a Jaxpr.Env.nil) loop_flatten_ir
rw [puzzle_flatten_translation_correct,a:Tensor ℝ [2, 3]⊢ puzzle_flatten a = Jaxpr.Program.eval (Jaxpr.Env.cons a Jaxpr.Env.nil) loop_flatten_irAll goals completed! 🐙 loop_flatten_translation_correct,a:Tensor ℝ [2, 3]⊢ puzzle_flatten a = loop_flatten aAll goals completed! 🐙 flatten_matches_loopa:Tensor ℝ [2, 3]⊢ loop_flatten a = loop_flatten aAll goals completed! 🐙]All goals completed! 🐙end JaxLean.PuzzleJax