import JaxLean.Stdlib.Equivariance
import examples.noether.generated.NoetherAdvect
import examples.noether.generated.NoetherBurgersnamespace JaxLean.NoetherProofsopen Tensor NoetherJaxopen scoped BigOperatorsThe four-cell periodic neighbor map: 0 ↦ 3, 1 ↦ 0, 2 ↦ 1, 3 ↦ 2.
def neighbor := cyclicShift 3 1Identify the actual slice/concatenate implementation at the roll boundary.
theorem advect_roll_spec (u : Tensor ℝ [4]) :
Advect.fn__roll_static u = permute neighbor u := u:Tensor ℝ [4]⊢ Advect.fn__roll_static u = permute neighbor u
u:Tensor ℝ [4]i:Fin 4v:Index []⊢ Advect.fn__roll_static u (i, v) = permute neighbor u (i, v)
u:Tensor ℝ [4]i:Fin 4⊢ Advect.fn__roll_static u (i, PUnit.unit) = permute neighbor u (i, PUnit.unit)
u:Tensor ℝ [4]⊢ Advect.fn__roll_static u ((fun i => i) ⟨0, ⋯⟩, PUnit.unit) = permute neighbor u ((fun i => i) ⟨0, ⋯⟩, PUnit.unit)u:Tensor ℝ [4]⊢ Advect.fn__roll_static u ((fun i => i) ⟨1, ⋯⟩, PUnit.unit) = permute neighbor u ((fun i => i) ⟨1, ⋯⟩, PUnit.unit)u:Tensor ℝ [4]⊢ Advect.fn__roll_static u ((fun i => i) ⟨2, ⋯⟩, PUnit.unit) = permute neighbor u ((fun i => i) ⟨2, ⋯⟩, PUnit.unit)u:Tensor ℝ [4]⊢ Advect.fn__roll_static u ((fun i => i) ⟨3, ⋯⟩, PUnit.unit) = permute neighbor u ((fun i => i) ⟨3, ⋯⟩, PUnit.unit) u:Tensor ℝ [4]⊢ Advect.fn__roll_static u ((fun i => i) ⟨0, ⋯⟩, PUnit.unit) = permute neighbor u ((fun i => i) ⟨0, ⋯⟩, PUnit.unit)u:Tensor ℝ [4]⊢ Advect.fn__roll_static u ((fun i => i) ⟨1, ⋯⟩, PUnit.unit) = permute neighbor u ((fun i => i) ⟨1, ⋯⟩, PUnit.unit)u:Tensor ℝ [4]⊢ Advect.fn__roll_static u ((fun i => i) ⟨2, ⋯⟩, PUnit.unit) = permute neighbor u ((fun i => i) ⟨2, ⋯⟩, PUnit.unit)u:Tensor ℝ [4]⊢ Advect.fn__roll_static u ((fun i => i) ⟨3, ⋯⟩, PUnit.unit) = permute neighbor u ((fun i => i) ⟨3, ⋯⟩, PUnit.unit) All goals completed! 🐙theorem burgers_roll_spec (u : Tensor ℝ [4]) :
Burgers.fn__roll_static u = permute neighbor u := u:Tensor ℝ [4]⊢ Burgers.fn__roll_static u = permute neighbor u
All goals completed! 🐙Readable specification of the original Python function.
theorem advect_spec (u : Tensor ℝ [4]) (c : Tensor ℝ []) :
Advect.advect u c = fluxStep neighbor (fun z => c () * z) u := u:Tensor ℝ [4]c:Tensor ℝ []⊢ Advect.advect u c = fluxStep neighbor (fun z => c () * z) u
u:Tensor ℝ [4]c:Tensor ℝ []⊢ map₂ (fun a0 a1 => a0 - a1) u
(map₂ (fun a0 a1 => a0 - a1) (map (fun a0 => c () * a0) u) (permute neighbor (map (fun a0 => c () * a0) u))) =
fluxStep neighbor (fun z => c () * z) u
All goals completed! 🐙theorem burgers_spec (u : Tensor ℝ [4]) (lam : Tensor ℝ []) :
Burgers.burgers u lam = fluxStep neighbor (fun z => lam () * (1 / 2 * z ^ 2)) u := u:Tensor ℝ [4]lam:Tensor ℝ []⊢ Burgers.burgers u lam = fluxStep neighbor (fun z => lam () * (1 / 2 * z ^ 2)) u
u:Tensor ℝ [4]lam:Tensor ℝ []⊢ map₂ (fun a0 a1 => a0 - a1) u
(map (fun a0 => lam () * a0)
(map₂ (fun a0 a1 => a0 - a1) (map (fun a0 => scalar (1 / 2) () * a0) (map (fun x => x ^ 2) u))
(permute neighbor (map (fun a0 => scalar (1 / 2) () * a0) (map (fun x => x ^ 2) u))))) =
fluxStep neighbor (fun z => lam () * (1 / 2 * z ^ 2)) u
u:Tensor ℝ [4]lam:Tensor ℝ []i:Index [4]⊢ map₂ (fun a0 a1 => a0 - a1) u
(map (fun a0 => lam () * a0)
(map₂ (fun a0 a1 => a0 - a1) (map (fun a0 => scalar (1 / 2) () * a0) (map (fun x => x ^ 2) u))
(permute neighbor (map (fun a0 => scalar (1 / 2) () * a0) (map (fun x => x ^ 2) u)))))
i =
fluxStep neighbor (fun z => lam () * (1 / 2 * z ^ 2)) u i
u:Tensor ℝ [4]lam:Tensor ℝ []i:Index [4]⊢ u i - lam () * (1 / 2 * u i ^ 2 - 1 / 2 * u (neighbor i.1, ()) ^ 2) =
u i - (lam () * (1 / 2 * u i ^ 2) - lam () * (1 / 2 * u (neighbor i.1, ()) ^ 2))
All goals completed! 🐙Every cyclic shift, not just the one-cell shift.
theorem advect_equivariant (u : Tensor ℝ [4]) (c : Tensor ℝ []) (k : ℤ) :
Advect.advect (permute (cyclicShift 3 k) u) c =
permute (cyclicShift 3 k) (Advect.advect u c) := u:Tensor ℝ [4]c:Tensor ℝ []k:ℤ⊢ Advect.advect (permute (cyclicShift 3 k) u) c = permute (cyclicShift 3 k) (Advect.advect u c)
u:Tensor ℝ [4]c:Tensor ℝ []k:ℤ⊢ fluxStep neighbor (fun z => c () * z) (permute (cyclicShift 3 k) u) =
permute (cyclicShift 3 k) (fluxStep neighbor (fun z => c () * z) u)
All goals completed! 🐙theorem burgers_equivariant (u : Tensor ℝ [4]) (lam : Tensor ℝ []) (k : ℤ) :
Burgers.burgers (permute (cyclicShift 3 k) u) lam =
permute (cyclicShift 3 k) (Burgers.burgers u lam) := u:Tensor ℝ [4]lam:Tensor ℝ []k:ℤ⊢ Burgers.burgers (permute (cyclicShift 3 k) u) lam = permute (cyclicShift 3 k) (Burgers.burgers u lam)
u:Tensor ℝ [4]lam:Tensor ℝ []k:ℤ⊢ fluxStep neighbor (fun z => lam () * (1 / 2 * z ^ 2)) (permute (cyclicShift 3 k) u) =
permute (cyclicShift 3 k) (fluxStep neighbor (fun z => lam () * (1 / 2 * z ^ 2)) u)
All goals completed! 🐙u:Tensor ℝ [4]c:Tensor ℝ []⊢ ∑ i, fluxStep neighbor (fun z => c () * z) u (i, ()) = ∑ i, u (i, ())
exact fluxStep_sum _ _ _ All goals completed! 🐙
theorem burgers_conserves_sum (u : Tensor ℝ [4]) (lam : Tensor ℝ []) :
(∑ i : Fin 4, Burgers.burgers u lam (i, ())) = ∑ i : Fin 4, u (i, ()) := by u:Tensor ℝ [4]lam:Tensor ℝ []⊢ ∑ i, Burgers.burgers u lam (i, ()) = ∑ i, u (i, ())
rw [burgers_spec u:Tensor ℝ [4]lam:Tensor ℝ []⊢ ∑ i, fluxStep neighbor (fun z => lam () * (1 / 2 * z ^ 2)) u (i, ()) = ∑ i, u (i, ()) u:Tensor ℝ [4]lam:Tensor ℝ []⊢ ∑ i, fluxStep neighbor (fun z => lam () * (1 / 2 * z ^ 2)) u (i, ()) = ∑ i, u (i, ())] u:Tensor ℝ [4]lam:Tensor ℝ []⊢ ∑ i, fluxStep neighbor (fun z => lam () * (1 / 2 * z ^ 2)) u (i, ()) = ∑ i, u (i, ())
exact fluxStep_sum _ _ _ All goals completed! 🐙Equivariance attached directly to the imported Jaxpr semantics.
theorem advect_certificate (u : Tensor ℝ [4]) (c : Tensor ℝ []) (k : ℤ) :
Jaxpr.Program.eval (.cons (permute (cyclicShift 3 k) u) (.cons c .nil))
Advect.advect_ir =
permute (cyclicShift 3 k)
(Jaxpr.Program.eval (.cons u (.cons c .nil)) Advect.advect_ir) := by u:Tensor ℝ [4]c:Tensor ℝ []k:ℤ⊢ Jaxpr.Program.eval (Jaxpr.Env.cons (permute (cyclicShift 3 k) u) (Jaxpr.Env.cons c Jaxpr.Env.nil)) Advect.advect_ir =
permute (cyclicShift 3 k) (Jaxpr.Program.eval (Jaxpr.Env.cons u (Jaxpr.Env.cons c Jaxpr.Env.nil)) Advect.advect_ir)
simp only [Advect.advect_translation_correct] u:Tensor ℝ [4]c:Tensor ℝ []k:ℤ⊢ Advect.advect (permute (cyclicShift 3 k) u) c = permute (cyclicShift 3 k) (Advect.advect u c)
exact advect_equivariant u c k All goals completed! 🐙theorem burgers_certificate (u : Tensor ℝ [4]) (lam : Tensor ℝ []) (k : ℤ) :
Jaxpr.Program.eval (.cons (permute (cyclicShift 3 k) u) (.cons lam .nil))
Burgers.burgers_ir =
permute (cyclicShift 3 k)
(Jaxpr.Program.eval (.cons u (.cons lam .nil)) Burgers.burgers_ir) := by u:Tensor ℝ [4]lam:Tensor ℝ []k:ℤ⊢ Jaxpr.Program.eval (Jaxpr.Env.cons (permute (cyclicShift 3 k) u) (Jaxpr.Env.cons lam Jaxpr.Env.nil))
Burgers.burgers_ir =
permute (cyclicShift 3 k)
(Jaxpr.Program.eval (Jaxpr.Env.cons u (Jaxpr.Env.cons lam Jaxpr.Env.nil)) Burgers.burgers_ir)
simp only [Burgers.burgers_translation_correct] u:Tensor ℝ [4]lam:Tensor ℝ []k:ℤ⊢ Burgers.burgers (permute (cyclicShift 3 k) u) lam = permute (cyclicShift 3 k) (Burgers.burgers u lam)
exact burgers_equivariant u lam k All goals completed! 🐙end JaxLean.NoetherProofs