import JaxLean.Core.Tensor import Mathlib.Algebra.BigOperators.Group.Finset.Basic import Mathlib.Data.ZMod.Basic import Mathlib.Tacticnamespace JaxLean.Tensoropen scoped BigOperators

Pull back a vector along a permutation of its coordinates.

def permute {R : Type} {n : ℕ} (p : Equiv.Perm (Fin n)) (x : Tensor R [n]) : Tensor R [n] := fun i => x (p i.1, ())

Pointwise unary operations commute with every coordinate permutation.

theorem permute_map {R S : Type} {n : ℕ} (p : Equiv.Perm (Fin n)) (f : R → S) (x : Tensor R [n]) : permute p (map f x) = map f (permute p x) := rfl

Both operands of a pointwise binary operation transform together.

theorem permute_map₂ {R S T : Type} {n : ℕ} (p : Equiv.Perm (Fin n)) (f : R → S → T) (x : Tensor R [n]) (y : Tensor S [n]) : permute p (map₂ f x y) = map₂ f (permute p x) (permute p y) := rfltheorem permute_sum {R : Type} [AddCommMonoid R] {n : ℕ} (p : Equiv.Perm (Fin n)) (x : Tensor R [n]) : (∑ i : Fin n, permute p x (i, ())) = ∑ i : Fin n, x (i, ()) := Equiv.sum_comp p (fun i => x (i, ()))

Positive shifts follow jnp.roll: output i reads input i - k.

def cyclicShift (n : ℕ) (k : ℤ) : Equiv.Perm (Fin (n + 1)) := Equiv.addRight (-(k : ZMod (n + 1)))theorem cyclicShift_commute (n : ℕ) (a b : ℤ) (i : Fin (n + 1)) : cyclicShift n a (cyclicShift n b i) = cyclicShift n b (cyclicShift n a i) := n:ℕa:ℤb:ℤi:Fin (n + 1)⊢ (cyclicShift n a) ((cyclicShift n b) i) = (cyclicShift n b) ((cyclicShift n a) i) All goals completed! 🐙

A local flux difference along any permutation, with any pointwise flux.

def fluxStep {R : Type} [Ring R] {n : ℕ} (p : Equiv.Perm (Fin n)) (flux : R → R) (x : Tensor R [n]) : Tensor R [n] := fun i => x i - (flux (x i) - flux (permute p x i))

A stencil commutes with every permutation that commutes with its neighbor map.

theorem fluxStep_equivariant {R : Type} [Ring R] {n : ℕ} (p q : Equiv.Perm (Fin n)) (commute : ∀ i, p (q i) = q (p i)) (flux : R → R) (x : Tensor R [n]) : fluxStep p flux (permute q x) = permute q (fluxStep p flux x) := R:Typeinst✝:Ring Rn:ℕp:Equiv.Perm (Fin n)q:Equiv.Perm (Fin n)commute:∀ (i : Fin n), p (q i) = q (p i)flux:R → Rx:Tensor R [n]⊢ fluxStep p flux (permute q x) = permute q (fluxStep p flux x) R:Typeinst✝:Ring Rn:ℕp:Equiv.Perm (Fin n)q:Equiv.Perm (Fin n)commute:∀ (i : Fin n), p (q i) = q (p i)flux:R → Rx:Tensor R [n]i:Fin nu:Index []⊢ fluxStep p flux (permute q x) (i, u) = permute q (fluxStep p flux x) (i, u) R:Typeinst✝:Ring Rn:ℕp:Equiv.Perm (Fin n)q:Equiv.Perm (Fin n)commute:∀ (i : Fin n), p (q i) = q (p i)flux:R → Rx:Tensor R [n]i:Fin n⊢ fluxStep p flux (permute q x) (i, PUnit.unit) = permute q (fluxStep p flux x) (i, PUnit.unit) All goals completed! 🐙

Cancellation of a permuted flux proves conservation without linearity.

R:Typeinst✝:Ring Rn:ℕp:Equiv.Perm (Fin n)flux:R → Rx:Tensor R [n]⊢ ∑ x_1, x (x_1, ()) - (∑ x_1, flux (x (x_1, ())) - ∑ i, flux (x (i, ()))) = ∑ x_1, x (x_1, ()) All goals completed! 🐙end JaxLean.Tensor