mirror of
https://github.com/allaunthefox/SilverSight.git
synced 2026-08-20 15:07:29 +00:00
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:
parent
8738d97777
commit
3ef102edd2
1 changed files with 176 additions and 40 deletions
|
|
@ -75,41 +75,119 @@ def SplitEmbedding.refinedSize {n : ℕ} (_ : SplitEmbedding n) : ℕ := n + 1
|
||||||
- states > i are shifted by +1 -/
|
- states > i are shifted by +1 -/
|
||||||
def SplitEmbedding.apply {n : ℕ} (f : SplitEmbedding n) (p : openSimplex n) :
|
def SplitEmbedding.apply {n : ℕ} (f : SplitEmbedding n) (p : openSimplex n) :
|
||||||
openSimplex (refinedSize f) :=
|
openSimplex (refinedSize f) :=
|
||||||
let i := f.splitIdx
|
let i : Fin n := f.splitIdx
|
||||||
let q := f.q
|
let q : ℝ := f.q
|
||||||
let pFn := p.1
|
let pFn : Fin n → ℝ := p.1
|
||||||
⟨fun (j : Fin (n+1)) =>
|
⟨fun (j : Fin (n+1)) =>
|
||||||
if h : j.val = i.val then
|
if h : j.val = i.val then
|
||||||
q * pFn i
|
q * pFn i
|
||||||
else if h' : j.val = i.val + 1 then
|
else if h' : j.val = i.val + 1 then
|
||||||
(1 - q) * pFn i
|
(1 - q) * pFn i
|
||||||
else if h'' : j.val < i.val then
|
else if h'' : j.val < i.val then
|
||||||
pFn ⟨j.val, by omega⟩
|
pFn ⟨j.val, by have := i.isLt; omega⟩
|
||||||
else
|
else
|
||||||
pFn ⟨j.val - 1, by omega⟩,
|
pFn ⟨j.val - 1, by have := i.isLt; omega⟩,
|
||||||
⟨fun j => by
|
⟨fun j => by
|
||||||
by_cases h : j.val = i.val
|
-- beta-reduce (fun j ↦ ...) j before split_ifs can fire
|
||||||
· exact mul_pos q (pFn i).2
|
simp only []
|
||||||
· by_cases h' : j.val = i.val + 1
|
have hjlt : j.val < n + 1 := j.isLt
|
||||||
· exact mul_pos (1 - q) (pFn i).2
|
split_ifs with h h' h''
|
||||||
· by_cases h'' : j.val < i.val
|
· exact mul_pos f.hq_pos (p.2.1 i)
|
||||||
· exact (pFn ⟨j.val, by omega⟩).2
|
· exact mul_pos (by linarith [f.hq_lt_one]) (p.2.1 i)
|
||||||
· exact (pFn ⟨j.val - 1, by omega⟩).2,
|
· exact p.2.1 ⟨j.val, by have := i.isLt; omega⟩
|
||||||
|
· exact p.2.1 ⟨j.val - 1, by have := i.isLt; omega⟩,
|
||||||
by
|
by
|
||||||
-- The sum splits: q*p_i + (1-q)*p_i + sum_{j<i} p_j + sum_{j>i+1} p_{j-1}
|
have hiN_lt : i.val < n + 1 := by have := i.isLt; omega
|
||||||
-- = p_i + sum_{j<i} p_j + sum_{k>i} p_k = p_i + (sum - p_i) = 1
|
have hi1N_lt : i.val + 1 < n + 1 := by have := i.isLt; omega
|
||||||
-- Reindexing: sum_{j>i+1} p_{j-1} = sum_{k>i} p_k by k = j-1
|
let iN : Fin (n+1) := ⟨i.val, hiN_lt⟩
|
||||||
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⟩)
|
let i1N : Fin (n+1) := ⟨i.val + 1, hi1N_lt⟩
|
||||||
= 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
|
-- rfl facts so omega can reason through Fin constructors
|
||||||
_ = 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
|
have hiN_val : iN.val = i.val := rfl
|
||||||
_ = pFn i + ((∑ j : Fin n, pFn j) - pFn i) := by
|
have hi1N_val : i1N.val = i.val + 1 := rfl
|
||||||
-- Σ_{j < i} p_j + Σ_{j > i+1} p_{j-1} = Σ_{j ≠ i} p_j
|
have hi1N_ne_iN : i1N ≠ iN := by
|
||||||
-- where j > i+1 maps to k = j-1 > i, covering indices i+1..n-1
|
intro h; exact absurd (congr_arg Fin.val h) (by simp [hiN_val, hi1N_val]; omega)
|
||||||
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⟩) =
|
have hi1N_mem : i1N ∈ Finset.univ.erase iN :=
|
||||||
∑ j : Fin n, pFn j - pFn i := by
|
Finset.mem_erase.mpr ⟨hi1N_ne_iN, Finset.mem_univ _⟩
|
||||||
sorry
|
let body : Fin (n+1) → ℝ := fun j =>
|
||||||
rw [h_split]
|
if j.val = i.val then q * pFn i
|
||||||
_ = 1 := by omega
|
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)
|
def SplitEmbedding.pushforward {n : ℕ} (f : SplitEmbedding n) (p : openSimplex n)
|
||||||
|
|
@ -488,7 +566,7 @@ lemma metric_at_uniform {N : ℕ} (hN : N ≥ 2) (g : RiemannianMetric N)
|
||||||
from hG0 1 0 (Or.inr rfl),
|
from hG0 1 0 (Or.inr rfl),
|
||||||
show g.toFun (uniformDist 2 (by linarith)) (b (1 : Fin 2)) (b (1 : Fin 2)) = D
|
show g.toFun (uniformDist 2 (by linarith)) (b (1 : Fin 2)) (b (1 : Fin 2)) = D
|
||||||
from hGD 1 (by decide),
|
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)
|
-- Goal: u 1 * (v 1 * D) = D / 2 * (u 0 * v 0 + u 1 * v 1)
|
||||||
have hu0 : u 0 = -u 1 := by
|
have hu0 : u 0 = -u 1 := by
|
||||||
have := hu; simp only [Fin.sum_univ_two] at this; linarith
|
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 →
|
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.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)
|
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
|
-- Normalize: uniformDist N ⋯ = p₀ definitionally (proof irrelevance)
|
||||||
rw [g_sum_left, g_sum_right]
|
show ∑ x : Fin N, u x * ∑ j : Fin N, v j * g.toFun p₀ (b x) (b j) =
|
||||||
-- Partition the sum using b_sum_left/right and g properties
|
D / 2 * ∑ i : Fin N, u i * v i
|
||||||
-- Key: Σ_{x,j: x≠0, j≠0} u_x v_j = u_0 v_0 when Σ u = Σ v = 0
|
let e₀ : Fin N := ⟨0, by omega⟩
|
||||||
-- Proof: u_0 = -Σ_{x≠0} u_x (from Σ u = 0), similarly v_0 = -Σ_{j≠0} v_j
|
-- Inner sum: ∑_j v_j G(b_x, b_j) = D/2 * (v x - v e₀) for x ≠ e₀
|
||||||
-- So u_0 v_0 = (Σ_{x≠0} u_x)(Σ_{j≠0} v_j) = Σ_{x≠0, j≠0} u_x v_j
|
-- Proof: split ∑ via sum_erase_add, ejecting j=x (→ D) and j=e₀ (→ 0),
|
||||||
have h_partition : (∑ x : Fin N, ∑ j : Fin N, u x * v j * g.toFun p₀ (b x) (b j)) =
|
-- leaving ∑_{j≠x,j≠e₀} v j * D/2 = D/2 * (∑_{j≠x,j≠e₀} v j).
|
||||||
D * (∑ i : Fin N, u i * v i) / 2 := by
|
have hinner : ∀ x : Fin N, x ≠ e₀ →
|
||||||
-- Use sum_erase to remove x=0 and j=0 terms, then apply hGOff for off-diagonal
|
∑ j : Fin N, v j * g.toFun p₀ (b x) (b j) = D / 2 * (v x - v e₀) := by
|
||||||
sorry
|
intro x hxe
|
||||||
rw [h_partition]
|
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
|
end UniformMetric
|
||||||
|
|
||||||
-- ============================================================
|
-- ============================================================
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue