import JaxLean.Stdlib.Equivariance import examples.noether.generated.NoetherAdvect import examples.noether.generated.NoetherBurgersnamespace JaxLean.NoetherProofsopen Tensor NoetherJaxopen scoped BigOperators

The four-cell periodic neighbor map: 0 ↦ 3, 1 ↦ 0, 2 ↦ 1, 3 ↦ 2.

def neighbor := cyclicShift 3 1

Identify 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, ()) All goals completed! 🐙u:Tensor ℝ [4]lam:Tensor ℝ []⊢ ∑ i, fluxStep neighbor (fun z => lam () * (1 / 2 * z ^ 2)) u (i, ()) = ∑ i, u (i, ()) 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) := 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) u:Tensor ℝ [4]c:Tensor ℝ []k:ℤ⊢ Advect.advect (permute (cyclicShift 3 k) u) c = permute (cyclicShift 3 k) (Advect.advect u c) 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) := 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) u:Tensor ℝ [4]lam:Tensor ℝ []k:ℤ⊢ Burgers.burgers (permute (cyclicShift 3 k) u) lam = permute (cyclicShift 3 k) (Burgers.burgers u lam) All goals completed! 🐙end JaxLean.NoetherProofs