Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
134 changes: 102 additions & 32 deletions ArkLib/Data/Fin/Sigma.lean
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,22 @@ theorem embedSum_succ_zero {n : Fin (m + 1) → ℕ} {j : Fin (n 0)} :
theorem embedSum_succ_succ {n : Fin (m + 1) → ℕ} {i : Fin m} (j : Fin (n i.succ)) :
embedSum (i.succ) j = Fin.natAdd _ (embedSum i j) := rfl

/-- The underlying value of `embedSum i j` is the sum of `n` over all indices before `i`,
plus the value of `j`. -/
theorem val_embedSum {m : ℕ} {n : Fin m → ℕ} (i : Fin m) (j : Fin (n i)) :
(embedSum i j).val = (∑ i' : Fin i.val, n (castLE i.isLt.le i')) + j.val := by
induction m with
| zero => exact Fin.elim0 i
| succ m ih =>
induction i using Fin.cases with
| zero => simp
| succ i =>
have key : (∑ i' : Fin i.succ.val, n (castLE i.succ.isLt.le i'))
= n 0 + ∑ i' : Fin i.val, n ((castLE i.isLt.le i').succ) :=
Fin.sum_univ_succ (fun i' : Fin (i.val + 1) => n (castLE i.succ.isLt.le i'))
rw [embedSum_succ_succ, val_natAdd, ih (n := fun i => n i.succ) i j]
omega

/-- Split a vector sum index `k : Fin (vsum n)` into nested indices `(i : Fin m) × Fin (n i)`.
This converts from indexing into the vector sum back to nested indexing, inverse of `embedSum`. -/
def splitSum {m : ℕ} {n : Fin m → ℕ} (k : Fin (vsum n)) : (i : Fin m) × Fin (n i) := match m with
Expand Down Expand Up @@ -253,16 +269,6 @@ theorem vflatten_one {n : Fin 1 → ℕ} {v : (i : Fin 1) → Fin (n i) → α}
theorem vflatten_two_eq_append {n : Fin 2 → ℕ} {v : (i : Fin 2) → Fin (n i) → α} :
vflatten v = vappend (v 0) (v 1) := rfl

theorem vflatten_eq_vappend_last {m : ℕ} {n : Fin (m + 1) → ℕ}
{v : (i : Fin (m + 1)) → Fin (n i) → α} :
vflatten v =
vappend (vflatten (fun i => v i.castSucc)) (v (last _)) ∘ Fin.cast vsum_castSucc := by
induction m with
| zero => ext i; simp
| succ m ih =>
rw [vflatten_succ, ih, vflatten_succ]
sorry

@[simp]
theorem vflatten_splitSum {m : ℕ} {n : Fin m → ℕ} (v : (k : Fin (vsum n)) → α) (k : Fin (vsum n)) :
vflatten (fun i j => v (embedSum i j)) k = v k :=
Expand All @@ -273,6 +279,33 @@ theorem vflatten_embedSum {m : ℕ} {n : Fin m → ℕ} (v : (i : Fin m) → Fin
(j : Fin (n i)) : vflatten v (embedSum i j) = v i j :=
dflatten_embedSum (motive := fun _ => α) v i j

theorem vflatten_eq_vappend_last {m : ℕ} {n : Fin (m + 1) → ℕ}
{v : (i : Fin (m + 1)) → Fin (n i) → α} :
vflatten v =
vappend (vflatten (fun i => v i.castSucc)) (v (last _)) ∘ Fin.cast vsum_castSucc := by
funext k
have hk : embedSum (splitSum k).1 (splitSum k).2 = k := embedSum_splitSum k
rcases hs : splitSum k with ⟨i, j⟩
rw [hs] at hk
clear hs
subst hk
dsimp only
rw [vflatten_embedSum, Function.comp_apply]
induction i using Fin.lastCases with
| last =>
refine (vappend_right (vflatten fun i => v i.castSucc) (v (last m)) j).symm.trans ?_
congr 1
apply Fin.ext
simp only [Fin.val_natAdd, Fin.val_cast, val_embedSum, Fin.val_last, vsum_eq_univ_sum]
rfl
| cast i =>
refine (vflatten_embedSum (fun i => v i.castSucc) i j).symm.trans ?_
rw [← vappend_left (vflatten fun i => v i.castSucc) (v (last m)) (embedSum i j)]
congr 1
apply Fin.ext
simp only [Fin.val_castAdd, Fin.val_cast, val_embedSum, Fin.val_castSucc]
rfl

/-- Functorial flatten: flattens a nested heterogeneous tuple
`(i : Fin m) → (j : Fin (n i)) → F (α i j)` into a single heterogeneous tuple with type
`(k : Fin (vsum n)) → F (vflatten α k)` where `vflatten` operates on the vector of types `α`.
Expand Down Expand Up @@ -309,13 +342,6 @@ theorem fflatten_two_eq_append {A : Sort u} {F : A → Sort v} {n : Fin 2 →
{v : (i : Fin 2) → (j : Fin (n i)) → F (α i j)} :
fflatten v = fappend (F := F) (v 0) (v 1) := rfl

@[simp]
theorem fflatten_splitSum {A : Sort u} {F : A → Sort v} {m : ℕ} {n : Fin m → ℕ}
{α : (i : Fin (vsum n)) → A}
(v : (k : Fin (vsum n)) → F (α k)) (k : Fin (vsum n)) :
fflatten (fun i j => v (embedSum i j)) k = cast (by simp) (v k) := by
sorry

@[simp]
theorem fflatten_embedSum {A : Sort u} {F : A → Sort v} {m : ℕ} {n : Fin m → ℕ}
{α : (i : Fin m) → (j : Fin (n i)) → A}
Expand All @@ -334,6 +360,18 @@ theorem fflatten_embedSum {A : Sort u} {F : A → Sort v} {m : ℕ} {n : Fin m
erw [fappend_right, ih (fun i => v i.succ) i j, _root_.cast_cast]
rfl

@[simp]
theorem fflatten_splitSum {A : Sort u} {F : A → Sort v} {m : ℕ} {n : Fin m → ℕ}
{α : (i : Fin (vsum n)) → A}
(v : (k : Fin (vsum n)) → F (α k)) (k : Fin (vsum n)) :
fflatten (fun i j => v (embedSum i j)) k = cast (by simp) (v k) := by
have hk : embedSum (splitSum k).1 (splitSum k).2 = k := embedSum_splitSum k
rcases hs : splitSum k with ⟨i, j⟩
rw [hs] at hk
clear hs
subst hk
exact fflatten_embedSum _ i j

/-- Functorial flatten with two arguments: flattens two nested heterogeneous tuple
`(i : Fin m) → (j : Fin (n i)) → F (α i j)` into a single heterogeneous tuple with type
`(k : Fin (vsum n)) → F (vflatten α k)` where `vflatten` operates on the vector of types `α`.
Expand Down Expand Up @@ -375,14 +413,6 @@ theorem fflatten₂_two_eq_append {A : Sort u} {B : Sort v} {F : A → B → Sor
{v : (i : Fin 2) → (j : Fin (n i)) → F (α i j) (β i j)} :
fflatten₂ v = fappend₂ (F := F) (v 0) (v 1) := rfl

@[simp]
theorem fflatten₂_splitSum {A : Sort u} {B : Sort v} {F : A → B → Sort w} {m : ℕ} {n : Fin m → ℕ}
{α : (i : Fin m) → (j : Fin (n i)) → A}
{β : (i : Fin m) → (j : Fin (n i)) → B}
(v : (k : Fin (vsum n)) → F (vflatten α k) (vflatten β k)) (k : Fin (vsum n)) :
fflatten₂ (fun i j => v (embedSum i j)) k = cast (by simp) (v k) := by
sorry

@[simp]
theorem fflatten₂_embedSum {A : Sort u} {B : Sort v} {F : A → B → Sort w} {m : ℕ} {n : Fin m → ℕ}
{α : (i : Fin m) → (j : Fin (n i)) → A}
Expand All @@ -402,6 +432,19 @@ theorem fflatten₂_embedSum {A : Sort u} {B : Sort v} {F : A → B → Sort w}
erw [fappend₂_right, ih (fun i => v i.succ) i j, _root_.cast_cast]
rfl

@[simp]
theorem fflatten₂_splitSum {A : Sort u} {B : Sort v} {F : A → B → Sort w} {m : ℕ} {n : Fin m → ℕ}
{α : (i : Fin m) → (j : Fin (n i)) → A}
{β : (i : Fin m) → (j : Fin (n i)) → B}
(v : (k : Fin (vsum n)) → F (vflatten α k) (vflatten β k)) (k : Fin (vsum n)) :
fflatten₂ (fun i j => v (embedSum i j)) k = cast (by simp) (v k) := by
have hk : embedSum (splitSum k).1 (splitSum k).2 = k := embedSum_splitSum k
rcases hs : splitSum k with ⟨i, j⟩
rw [hs] at hk
clear hs
subst hk
exact fflatten₂_embedSum _ i j

/-- Heterogeneous flatten: flattens a nested heterogeneous tuple
`(i : Fin m) → (j : Fin (n i)) → α i j` into a single heterogeneous tuple with type
`(k : Fin (vsum n)) → vflatten α k` where `vflatten` operates on the vector of types `α`.
Expand Down Expand Up @@ -478,9 +521,30 @@ def ranges {n : ℕ} (a : Fin n → ℕ) : (i : Fin n) → Fin (a i) → ℕ :=
def divSum? {m : ℕ} (n : Fin m → ℕ) (k : ℕ) : Option (Fin m) :=
Fin.find? (fun i => k < ∑ j, n (castLE i.isLt j))

/-- The sum of `n` over the first `i + 1` indices is at most the total sum. -/
theorem partialSum_le_sum {m : ℕ} (n : Fin m → ℕ) (i : Fin m) :
∑ j : Fin (i.val + 1), n (castLE i.isLt j) ≤ ∑ j, n j := by
have h : (i.val + 1) + (m - i.val - 1) = m := by omega
conv_rhs => rw [← Fin.sum_congr' n h, Fin.sum_univ_add]
refine le_of_eq_of_le (Finset.sum_congr rfl fun j _ => ?_) (Nat.le_add_right _ _)
rfl

theorem divSum?_is_some_iff_lt_sum {m : ℕ} {n : Fin m → ℕ} {k : ℕ} :
(divSum? n k).isSome ↔ k < ∑ i, n i := by
sorry
rw [divSum?, Fin.isSome_find?_iff]
constructor
· rintro ⟨i, hi⟩
rw [decide_eq_true_eq] at hi
exact lt_of_lt_of_le hi (partialSum_le_sum n i)
· intro hk
obtain ⟨m', rfl⟩ : ∃ m', m = m' + 1 := by
cases m with
| zero => simp at hk
| succ m' => exact ⟨m', rfl⟩
refine ⟨Fin.last m', ?_⟩
rw [decide_eq_true_eq]
refine lt_of_lt_of_le hk (le_of_eq (Finset.sum_congr rfl fun j _ => ?_))
rfl
-- constructor
-- · intro h
-- simp only [divSum?, Nat.succ_eq_add_one, castLE, isSome_find_iff] at h
Expand All @@ -505,20 +569,26 @@ theorem sum_le_of_divSum?_eq_some {m : ℕ} {n : Fin m → ℕ} {k : Fin (∑ j,
by_cases hi' : 0 = i.val
· rw [← Fin.sum_congr' _ hi']
simp only [Finset.univ_eq_empty, Finset.sum_empty, _root_.zero_le]
· have : (i.val - 1) + 1 = i.val := by omega
rw [← Fin.sum_congr' _ this]
sorry
· have hone : (i.val - 1) + 1 = i.val := by omega
rw [← Fin.sum_congr' _ hone]
have hj : (decide (↑k < ∑ j, n (castLE (Fin.mk (i.val - 1) (by omega) : Fin m).isLt j)))
= false :=
Fin.eq_false_of_find?_eq_some_of_lt hi ⟨i.val - 1, by omega⟩ (by simp [Fin.lt_def]; omega)
rw [decide_eq_false_iff_not, not_lt] at hj
refine le_trans (le_of_eq (Finset.sum_congr rfl fun j _ => ?_)) hj
rfl
-- have := Fin.find_min (Option.mem_def.mp hi) (j := ⟨i.val - 1, by omega⟩) <| Fin.lt_def.mpr
-- (by simp only; omega)
-- exact not_lt.mp this

def modSum {m : ℕ} {n : Fin m → ℕ} (k : Fin (∑ j, n j)) : Fin (n (divSum k)) :=
⟨k - ∑ j, n (Fin.castLE (divSum k).isLt.le j), by
-- sorry
have divSum_mem : divSum k ∈ divSum? n k := by
simp only [divSum, divSum?, Option.mem_def, Option.some_get]
have hk : k < ∑ j, n (Fin.castLE (divSum k).isLt j) := by
sorry --Fin.find_spec _ divSum_mem
have hk : (k : ℕ) < ∑ j, n (Fin.castLE (divSum k).isLt j) := by
have := Fin.eq_true_of_find?_eq_some (p := fun i => decide ((k : ℕ) <
∑ j, n (castLE i.isLt j))) divSum_mem
simpa using this
simp only [Fin.sum_univ_succAbove _ (Fin.last (divSum k)), succAbove_last] at hk
rw [Nat.sub_lt_iff_lt_add' (sum_le_of_divSum?_eq_some divSum_mem)]
rw [add_comm]
Expand Down
Loading