refactor(chentsov): Prove SplitEmbedding.apply sum via Finset bijection

- Used Finset.sum_nbij' for reindexing bijection between erased indices
- Eliminated apply-sum sorry (was blocking the file)
- Build: 3307 jobs, 0 errors, 3 sorries remaining
This commit is contained in:
allaun 2026-06-25 23:25:48 -05:00
parent 8738d97777
commit 3ef102edd2

View file

@ -75,42 +75,120 @@ def SplitEmbedding.refinedSize {n : } (_ : SplitEmbedding n) : := n + 1
- states > i are shifted by +1 -/
def SplitEmbedding.apply {n : } (f : SplitEmbedding n) (p : openSimplex n) :
openSimplex (refinedSize f) :=
let i := f.splitIdx
let q := f.q
let pFn := p.1
let i : Fin n := f.splitIdx
let q : := f.q
let pFn : Fin n → := p.1
⟨fun (j : Fin (n+1)) =>
if h : j.val = i.val then
q * pFn i
else if h' : j.val = i.val + 1 then
(1 - q) * pFn i
else if h'' : j.val < i.val then
pFn ⟨j.val, by omega⟩
pFn ⟨j.val, by have := i.isLt; omega⟩
else
pFn ⟨j.val - 1, by omega⟩,
pFn ⟨j.val - 1, by have := i.isLt; omega⟩,
⟨fun j => by
by_cases h : j.val = i.val
· exact mul_pos q (pFn i).2
· by_cases h' : j.val = i.val + 1
· exact mul_pos (1 - q) (pFn i).2
· by_cases h'' : j.val < i.val
· exact (pFn ⟨j.val, by omega⟩).2
· exact (pFn ⟨j.val - 1, by omega⟩).2,
-- beta-reduce (fun j ↦ ...) j before split_ifs can fire
simp only []
have hjlt : j.val < n + 1 := j.isLt
split_ifs with h h' h''
· exact mul_pos f.hq_pos (p.2.1 i)
· exact mul_pos (by linarith [f.hq_lt_one]) (p.2.1 i)
· exact p.2.1 ⟨j.val, by have := i.isLt; omega⟩
· exact p.2.1 ⟨j.val - 1, by have := i.isLt; omega⟩,
by
-- The sum splits: q*p_i + (1-q)*p_i + sum_{j<i} p_j + sum_{j>i+1} p_{j-1}
-- = p_i + sum_{j<i} p_j + sum_{k>i} p_k = p_i + (sum - p_i) = 1
-- Reindexing: sum_{j>i+1} p_{j-1} = sum_{k>i} p_k by k = j-1
calc ∑ j : Fin (n+1), (if j.val = i.val then q * pFn i else if j.val = i.val + 1 then (1 - q) * pFn i else if j.val < i.val then pFn ⟨j.val, by omega⟩ else pFn ⟨j.val - 1, by omega⟩)
= q * pFn i + (1 - q) * pFn i + ∑ j : Fin (n+1), (if j.val < i.val then pFn ⟨j.val, by omega⟩ else pFn ⟨j.val - 1, by omega⟩) := by native_decide
_ = pFn i + ∑ j : Fin (n+1), (if j.val < i.val then pFn ⟨j.val, by omega⟩ else pFn ⟨j.val - 1, by omega⟩) := by ring
_ = pFn i + ((∑ j : Fin n, pFn j) - pFn i) := by
-- Σ_{j < i} p_j + Σ_{j > i+1} p_{j-1} = Σ_{j ≠ i} p_j
-- where j > i+1 maps to k = j-1 > i, covering indices i+1..n-1
have h_split : ∑ j : Fin (n+1), (if j.val < i.val then pFn ⟨j.val, by omega⟩ else pFn ⟨j.val - 1, by omega⟩) =
∑ j : Fin n, pFn j - pFn i := by
sorry
rw [h_split]
_ = 1 := by omega
⟩⟩
have hiN_lt : i.val < n + 1 := by have := i.isLt; omega
have hi1N_lt : i.val + 1 < n + 1 := by have := i.isLt; omega
let iN : Fin (n+1) := ⟨i.val, hiN_lt⟩
let i1N : Fin (n+1) := ⟨i.val + 1, hi1N_lt⟩
-- rfl facts so omega can reason through Fin constructors
have hiN_val : iN.val = i.val := rfl
have hi1N_val : i1N.val = i.val + 1 := rfl
have hi1N_ne_iN : i1N ≠ iN := by
intro h; exact absurd (congr_arg Fin.val h) (by simp [hiN_val, hi1N_val]; omega)
have hi1N_mem : i1N ∈ Finset.univ.erase iN :=
Finset.mem_erase.mpr ⟨hi1N_ne_iN, Finset.mem_univ _⟩
let body : Fin (n+1) → := fun j =>
if j.val = i.val then q * pFn i
else if j.val = i.val + 1 then (1 - q) * pFn i
else if j.val < i.val then pFn ⟨j.val, by have := i.isLt; omega⟩
else pFn ⟨j.val - 1, by have := i.isLt; omega⟩
show ∑ j : Fin (n+1), body j = 1
have hbody_iN : body iN = q * pFn i := by dsimp only [body, iN]; simp
have hbody_i1N : body i1N = (1 - q) * pFn i := by
dsimp only [body, i1N]; simp [show i.val + 1 ≠ i.val from by omega]
have hea1 : ∑ j ∈ Finset.univ.erase iN, body j + body iN = ∑ j : Fin (n+1), body j :=
Finset.sum_erase_add Finset.univ body (Finset.mem_univ iN)
have hea2 : ∑ j ∈ (Finset.univ.erase iN).erase i1N, body j + body i1N =
∑ j ∈ Finset.univ.erase iN, body j :=
Finset.sum_erase_add (Finset.univ.erase iN) body hi1N_mem
have hpsum_erase : ∑ k ∈ Finset.univ.erase i, pFn k = 1 - pFn i := by
linarith [Finset.sum_erase_add Finset.univ pFn (Finset.mem_univ i), p.2.2]
have hrest : ∑ j ∈ (Finset.univ.erase iN).erase i1N, body j =
∑ k ∈ Finset.univ.erase i, pFn k :=
Finset.sum_nbij'
(fun j => if j.val < i.val then (⟨j.val, by have := i.isLt; omega⟩ : Fin n)
else ⟨j.val - 1, by have := i.isLt; omega⟩)
(fun k => if k.val < i.val then (⟨k.val, by have := i.isLt; omega⟩ : Fin (n+1))
else ⟨k.val + 1, by have := k.isLt; omega⟩)
-- forward image ∈ erase i
(fun j hj => by
simp only [Finset.mem_erase, Finset.mem_univ, and_true] at hj
-- extract numeric ne conditions via congr_arg Fin.val
have hj1 : j.val ≠ i.val + 1 :=
fun h => hj.1 (Fin.ext (by rw [hi1N_val]; exact h))
have hj2 : j.val ≠ i.val :=
fun h => hj.2 (Fin.ext (by rw [hiN_val]; exact h))
simp only [Finset.mem_erase, Finset.mem_univ, and_true]
split_ifs with h
· exact fun heq => hj2 (congr_arg Fin.val heq)
· exact fun heq => absurd (congr_arg Fin.val heq) (by omega))
-- backward image ∈ rest
(fun k hk => by
simp only [Finset.mem_erase, Finset.mem_univ, and_true] at hk
have hkne : k.val ≠ i.val := fun h => hk (Fin.ext h)
simp only [Finset.mem_erase, Finset.mem_univ, and_true]
constructor
· split_ifs with h
· exact fun heq => absurd (congr_arg Fin.val heq) (by rw [hi1N_val]; omega)
· exact fun heq => absurd (congr_arg Fin.val heq) (by rw [hi1N_val]; omega)
· split_ifs with h
· exact fun heq => absurd (congr_arg Fin.val heq) (by rw [hiN_val]; omega)
· exact fun heq => absurd (congr_arg Fin.val heq) (by rw [hiN_val]; omega))
-- left inverse: ψ(φ(j)) = j
(fun j hj => by
simp only [Finset.mem_erase, Finset.mem_univ, and_true] at hj
have hj1 : j.val ≠ i.val + 1 :=
fun h => hj.1 (Fin.ext (by rw [hi1N_val]; exact h))
have hj2 : j.val ≠ i.val :=
fun h => hj.2 (Fin.ext (by rw [hiN_val]; exact h))
split_ifs with h1 h2
· exact Fin.ext rfl
· omega
· omega
· exact Fin.ext (by omega))
-- right inverse: φ(ψ(k)) = k
(fun k hk => by
simp only [Finset.mem_erase, Finset.mem_univ, and_true] at hk
have hkne : k.val ≠ i.val := fun h => hk (Fin.ext h)
split_ifs with h1 h2
· exact Fin.ext rfl
· omega
· omega
· exact Fin.ext (by omega))
-- body(j) = pFn(φ(j))
(fun j hj => by
simp only [Finset.mem_erase, Finset.mem_univ, and_true] at hj
have hj1 : j.val ≠ i.val + 1 :=
fun h => hj.1 (Fin.ext (by rw [hi1N_val]; exact h))
have hj2 : j.val ≠ i.val :=
fun h => hj.2 (Fin.ext (by rw [hiN_val]; exact h))
dsimp only [body]
simp only [if_neg hj2, if_neg hj1]
split_ifs <;> rfl)
linarith [hea1, hea2, hbody_iN, hbody_i1N, hrest, hpsum_erase,
show q * pFn i + (1 - q) * pFn i = pFn i from by ring]
⟩⟩
def SplitEmbedding.pushforward {n : } (f : SplitEmbedding n) (p : openSimplex n)
(X : Fin n → ) : Fin (refinedSize f) → :=
@ -488,7 +566,7 @@ lemma metric_at_uniform {N : } (hN : N ≥ 2) (g : RiemannianMetric N)
from hG0 1 0 (Or.inr rfl),
show g.toFun (uniformDist 2 (by linarith)) (b (1 : Fin 2)) (b (1 : Fin 2)) = D
from hGD 1 (by decide),
mul_zero, zero_mul, add_zero, zero_add]
mul_zero, add_zero, zero_add]
-- Goal: u 1 * (v 1 * D) = D / 2 * (u 0 * v 0 + u 1 * v 1)
have hu0 : u 0 = -u 1 := by
have := hu; simp only [Fin.sum_univ_two] at this; linarith
@ -499,18 +577,76 @@ lemma metric_at_uniform {N : } (hN : N ≥ 2) (g : RiemannianMetric N)
have hGOff : ∀ x j : Fin N, x.val ≠ 0 → j.val ≠ 0 → x ≠ j →
g.toFun p₀ (b x) (b j) = D / 2 := fun x j hx hj hxj =>
g_offdiag_half hN3 g h_perm x j hx hj (Fin.val_ne_iff.mpr hxj)
-- Expand g(u,v) = Σ_{x,j} u_x v_j g(b_x, b_j) using bilinearity
rw [g_sum_left, g_sum_right]
-- Partition the sum using b_sum_left/right and g properties
-- Key: Σ_{x,j: x≠0, j≠0} u_x v_j = u_0 v_0 when Σ u = Σ v = 0
-- Proof: u_0 = -Σ_{x≠0} u_x (from Σ u = 0), similarly v_0 = -Σ_{j≠0} v_j
-- So u_0 v_0 = (Σ_{x≠0} u_x)(Σ_{j≠0} v_j) = Σ_{x≠0, j≠0} u_x v_j
have h_partition : (∑ x : Fin N, ∑ j : Fin N, u x * v j * g.toFun p₀ (b x) (b j)) =
D * (∑ i : Fin N, u i * v i) / 2 := by
-- Use sum_erase to remove x=0 and j=0 terms, then apply hGOff for off-diagonal
sorry
rw [h_partition]
-- Normalize: uniformDist N ⋯ = p₀ definitionally (proof irrelevance)
show ∑ x : Fin N, u x * ∑ j : Fin N, v j * g.toFun p₀ (b x) (b j) =
D / 2 * ∑ i : Fin N, u i * v i
let e₀ : Fin N := ⟨0, by omega⟩
-- Inner sum: ∑_j v_j G(b_x, b_j) = D/2 * (v x - v e₀) for x ≠ e₀
-- Proof: split ∑ via sum_erase_add, ejecting j=x (→ D) and j=e₀ (→ 0),
-- leaving ∑_{j≠x,j≠e₀} v j * D/2 = D/2 * (∑_{j≠x,j≠e₀} v j).
have hinner : ∀ x : Fin N, x ≠ e₀ →
∑ j : Fin N, v j * g.toFun p₀ (b x) (b j) = D / 2 * (v x - v e₀) := by
intro x hxe
have hxval : x.val ≠ 0 := fun h => hxe (Fin.ext h)
have hmem_e0 : e₀ ∈ Finset.univ.erase x :=
Finset.mem_erase.mpr ⟨hxe.symm, Finset.mem_univ _⟩
-- Partial sums of v over the erased sets
have hv_x : ∑ j ∈ Finset.univ.erase x, v j = -v x :=
by linarith [Finset.sum_erase_add Finset.univ v (Finset.mem_univ x), hv]
have hv_xe : ∑ j ∈ (Finset.univ.erase x).erase e₀, v j = -v x - v e₀ :=
by linarith [Finset.sum_erase_add (Finset.univ.erase x) v hmem_e0, hv_x]
-- G(b_x, b_j) = D/2 for all j ≠ x, j ≠ e₀
have hoff : ∀ j ∈ (Finset.univ.erase x).erase e₀,
v j * g.toFun p₀ (b x) (b j) = v j * (D / 2) := fun j hj => by
simp only [Finset.mem_erase, Finset.mem_univ, and_true] at hj
rw [hGOff x j hxval (fun h => hj.1 (Fin.ext h)) (Ne.symm hj.2)]
-- Reconstruct total sum by splitting out x and e₀
-- Explicit types force beta-reduction of the lambda applications
have h1 : ∑ j ∈ Finset.univ.erase x, v j * g.toFun p₀ (b x) (b j) +
v x * g.toFun p₀ (b x) (b x) = ∑ j : Fin N, v j * g.toFun p₀ (b x) (b j) :=
Finset.sum_erase_add Finset.univ _ (Finset.mem_univ x)
have h2 : ∑ j ∈ (Finset.univ.erase x).erase e₀, v j * g.toFun p₀ (b x) (b j) +
v e₀ * g.toFun p₀ (b x) (b e₀) =
∑ j ∈ Finset.univ.erase x, v j * g.toFun p₀ (b x) (b j) :=
Finset.sum_erase_add (Finset.univ.erase x) _ hmem_e0
rw [hG0 x e₀ (Or.inr rfl), mul_zero, add_zero] at h2
rw [hGD x hxval] at h1
calc ∑ j : Fin N, v j * g.toFun p₀ (b x) (b j)
= ∑ j ∈ (Finset.univ.erase x).erase e₀, v j * g.toFun p₀ (b x) (b j) +
v x * D := by linarith [h1, h2]
_ = ∑ j ∈ (Finset.univ.erase x).erase e₀, v j * (D / 2) + v x * D := by
rw [Finset.sum_congr rfl hoff]
_ = D / 2 * (-v x - v e₀) + v x * D := by
rw [← Finset.sum_mul, hv_xe, mul_comm]
_ = D / 2 * (v x - v e₀) := by ring
-- x = e₀ row is zero
have he0_zero : u e₀ * ∑ j : Fin N, v j * g.toFun p₀ (b e₀) (b j) = 0 := by
suffices h : ∑ j : Fin N, v j * g.toFun p₀ (b e₀) (b j) = 0 by simp [h]
apply Finset.sum_eq_zero; intro j _; rw [hG0 e₀ j (Or.inl rfl)]; ring
-- Split outer sum: e₀ term is 0, remaining terms use hinner
have houter_split : ∑ x : Fin N, u x * ∑ j, v j * g.toFun p₀ (b x) (b j) =
∑ x ∈ Finset.univ.erase e₀, u x * (D / 2 * (v x - v e₀)) := by
have := Finset.sum_erase_add Finset.univ (fun x => u x * ∑ j, v j * g.toFun p₀ (b x) (b j))
(Finset.mem_univ e₀)
simp only [he0_zero] at this
rw [← this]; simp only [add_zero]
apply Finset.sum_congr rfl
intro x hx; rw [hinner x (Finset.mem_erase.mp hx).1]
rw [houter_split]
-- ∑_{x≠e₀} u x * (D/2*(v x - v e₀)) = D/2 * ∑_i u_i v_i
have hue0_sum : ∑ x ∈ Finset.univ.erase e₀, u x = -u e₀ := by
linarith [Finset.sum_erase_add Finset.univ u (Finset.mem_univ e₀), hu]
have hprod_split : ∑ i : Fin N, u i * v i =
u e₀ * v e₀ + ∑ x ∈ Finset.univ.erase e₀, u x * v x := by
linarith [Finset.sum_erase_add Finset.univ (fun x => u x * v x) (Finset.mem_univ e₀)]
-- Expand LHS using ring: u x * (D/2*(v x - v e₀)) = D/2*(u x*v x) - D/2*v e₀*(u x)
have hexpand : ∑ x ∈ Finset.univ.erase e₀, u x * (D / 2 * (v x - v e₀)) =
D / 2 * ∑ x ∈ Finset.univ.erase e₀, u x * v x -
D / 2 * v e₀ * ∑ x ∈ Finset.univ.erase e₀, u x := by
simp_rw [show ∀ x : Fin N, u x * (D / 2 * (v x - v e₀)) =
D / 2 * (u x * v x) - D / 2 * v e₀ * u x from fun x => by ring]
rw [Finset.sum_sub_distrib, ← Finset.mul_sum, ← Finset.mul_sum]
rw [hexpand, hue0_sum, hprod_split]; ring
end UniformMetric
-- ============================================================