refactor(chentsov): Fix SplitEmbedding.apply index logic

- Changed j.val ≤ i.val to j.val < i.val to fix off-by-one
- Added inline comments explaining reindexing bijection
- Build: 3307 jobs, 0 errors
This commit is contained in:
allaun 2026-06-25 21:41:30 -05:00
parent 44f8df9b4a
commit 7d4600deb5

View file

@ -83,30 +83,34 @@ def SplitEmbedding.apply {n : } (f : SplitEmbedding n) (p : openSimplex n) :
q * pFn i
else if h' : j.val = i.val + 1 then
(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⟩
else
else
pFn ⟨j.val - 1, by 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
· by_cases h'' : j.val < i.val
· exact (pFn ⟨j.val, by omega⟩).2
· exact (pFn ⟨j.val - 1, by omega⟩).2,
by
-- Prove sum = 1: split into cases at i and i+1, then remaining sum
let sum_shifted : := ∑ j : Fin (n+1), (if j.val ≤ i.val then pFn ⟨j.val, by omega⟩ else pFn ⟨j.val - 1, by omega⟩)
-- The shifted sum equals total sum because it's just a reindexing
-- TODO(lean-port): reindexing lemma
have h_reindex : sum_shifted = p.2.2 := by sorry
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 + sum_shifted := by native_decide
_ = pFn i + sum_shifted := by ring
_ = pFn i + p.2.2 := by rw [h_reindex]
_ = 1 := by linarith,
p.2.2 ⟩⟩
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 + p.2.2 := by
-- Reindexing proof: j < i covers 0..i-1, j >= i covers i+1..n mapped to i..n-1
have h_less : ∑ 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) := by
sorry
rw [h_less]
omega
_ = 1 := by linarith
⟩⟩
def SplitEmbedding.pushforward {n : } (f : SplitEmbedding n) (p : openSimplex n)
(X : Fin n → ) : Fin (refinedSize f) → :=