Skip to content

Commit b702a65

Browse files
committed
Refactor TS.lean (in progress)
1 parent 35c01e6 commit b702a65

2 files changed

Lines changed: 85 additions & 139 deletions

File tree

‎LeanBandits/BanditAlgorithms/TS.lean‎

Lines changed: 72 additions & 139 deletions
Original file line numberDiff line numberDiff line change
@@ -114,132 +114,68 @@ private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1
114114
have : √n * √n = n := Real.mul_self_sqrt (by positivity)
115115
nlinarith
116116

117-
/-- Helper for `sum_ucb_sub_mean_le`. -/
118-
private lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) :
119-
∑ s ∈ range n, f (pullCount A (A s ω) s ω) =
120-
∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω), f j := by
121-
induction n with
122-
| zero => simp
123-
| succ n ih =>
124-
have hf : f (pullCount A (A n ω) n ω) =
125-
∑ a, if A n ω = a then f (pullCount A a n ω) else 0 := by simp
126-
simp_rw [sum_range_succ, ih, hf, ← sum_add_distrib, pullCount_add_one]
127-
congr 1 with a
128-
split_ifs
129-
· simp [sum_range_succ]
130-
· simp
131-
132-
lemma sum_ucb_sub_mean_le {lo hi : ℝ} {μ : Fin K → ℝ}
133-
(hm : ∀ a, μ a ∈ Set.Icc lo hi)
134-
(hlo : lo ≤ hi) (σ2 δ : ℝ) (n : ℕ) (ω : Ω)
135-
(hconc : ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 →
136-
|empMean A R' a s ω - μ a|
137-
< √(2 * σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ))) :
138-
∑ s ∈ range n, (ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω))
139-
≤ (hi - lo) * ↑K + 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by
140-
-- Split range n into first-pull (pc=0) and non-first-pull (pc≠0) sets
141-
set S0 := (range n).filter (fun s => pullCount A (A s ω) s ω = 0)
142-
set S1 := (range n).filter (fun s => pullCount A (A s ω) s ω ≠ 0)
143-
have hpart : range n = S0 ∪ S1 := (Finset.filter_union_filter_not_eq _ _).symm
144-
have hdisj : Disjoint S0 S1 := Finset.disjoint_filter_filter_not _ _ _
145-
conv_lhs => rw [hpart]
146-
rw [Finset.sum_union hdisj]
147-
-- We bound ∑_{S0} and ∑_{S1} separately, then combine
148-
suffices h_S0 : ∑ s ∈ S0, (ucb A R' lo hi σ2 δ (A s ω) s ω -
149-
μ (A s ω)) ≤ (hi - lo) * ↑K by
150-
suffices h_S1 : ∑ s ∈ S1, (ucb A R' lo hi σ2 δ (A s ω) s ω -
151-
μ (A s ω))
152-
≤ 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) by
153-
have := Finset.sum_union hdisj (f := fun s =>
154-
ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω))
155-
rw [← hpart] at this; linarith
156-
-- Bound ∑_{S1}: each term ≤ 2√(2σ2c/pc) = 2√(2σ2c/max(1,pc)), so ≤ full sum
157-
calc ∑ s ∈ S1, (ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω))
158-
≤ ∑ s ∈ S1,
159-
2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) :=
160-
sum_le_sum fun s hs => by
161-
have hpc : pullCount A (A s ω) s ω ≠ 0 := (Finset.mem_filter.mp hs).2
117+
lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, μ a ∈ Set.Icc l u) (hi : l ≤ u)
118+
(hc : ∀ s < n, pullCount A (A s ω) s ω ≠ 0 → |empMean A R' (A s ω) s ω - μ (A s ω)|
119+
< √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω))) :
120+
∑ s ∈ range n, (ucb A R' l u σ2 δ (A s ω) s ω - μ (A s ω))
121+
≤ (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by
122+
let S₀ := {s ∈ range n | pullCount A (A s ω) s ω = 0}
123+
let S₁ := {s ∈ range n | pullCount A (A s ω) s ω ≠ 0}
124+
have hu : S₀ ∪ S₁ = range n := filter_union_filter_not_eq _ _
125+
have hd : Disjoint S₀ S₁ := disjoint_filter_filter_not _ _ _
126+
rw [← hu, sum_union hd]
127+
gcongr
128+
· calc ∑ s ∈ S₀, (ucb A R' l u σ2 δ (A s ω) s ω - μ (A s ω))
129+
≤ ∑ s ∈ S₀, (u - l) :=
130+
have (s : ℕ) : ucb A R' l u σ2 δ (A s ω) s ω ∈ Set.Icc l u := ucb_mem_Icc hi
131+
sum_le_sum (by grind)
132+
_ = ∑ s ∈ range n, if pullCount A (A s ω) s ω = 0 then (u - l) else 0 := by
133+
rw [sum_filter]
134+
_ = ∑ a, ∑ j ∈ range (pullCount A a n ω), if j = 0 then (u - l) else 0 :=
135+
sum_comp_pullCount (fun j => if j = 0 then (u - l) else 0) n ω
136+
_ ≤ ∑ a, (u - l) := by
137+
gcongr
138+
rw [sum_ite_eq']
139+
grind
140+
_ = (u - l) * K := by
141+
rw [Fin.sum_const, nsmul_eq_mul, mul_comm]
142+
· calc ∑ s ∈ S₁, (ucb A R' l u σ2 δ (A s ω) s ω - μ (A s ω))
143+
≤ ∑ s ∈ S₁, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) := by
144+
gcongr with s hs
162145
unfold ucb
163146
grind
164-
_ ≤ ∑ s ∈ range n,
165-
2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) :=
166-
Finset.sum_le_sum_of_subset_of_nonneg
167-
(Finset.filter_subset _ _) fun s _ _ => by positivity
168-
_ ≤ 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by
169-
set c := Real.log (1 / δ)
170-
by_cases hc : 0 ≤ 2 * σ2 * c
171-
· open Real in
172-
calc ∑ s ∈ range n, 2 * √(2 * σ2 * c / ↑(pullCount A (A s ω) s ω))
173-
= ∑ s ∈ range n, √(8 * σ2 * c) *
174-
(1 / √(↑(pullCount A (A s ω) s ω) : ℝ)) :=
175-
sum_congr rfl fun s _ => by
176-
rw [show (8 : ℝ) * σ2 * c = (2 : ℝ) ^ 2 * (2 * σ2 * c) from by ring]
177-
rw [sqrt_mul (by positivity : (0:ℝ) ≤ 2 ^ 2),
178-
sqrt_sq (by norm_num : (0:ℝ) ≤ 2)]
179-
rw [sqrt_div (by linarith : 0 ≤ 2 * σ2 * c)]; ring
180-
_ = √(8 * σ2 * c) * ∑ s ∈ range n,
181-
(1 / √(↑(pullCount A (A s ω) s ω) : ℝ)) := by
182-
rw [mul_sum]
183-
_ = √(8 * σ2 * c) * ∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω),
184-
(1 / √(↑j : ℝ)) := by
185-
congr 1; exact sum_comp_pullCount (fun j => 1 / √(↑j : ℝ)) n ω
186-
_ ≤ √(8 * σ2 * c) * ∑ a : Fin K, (2 * √↑(pullCount A a n ω)) := by
187-
gcongr with a
188-
by_cases ha : pullCount A a n ω = 0
189-
· simp [ha]
190-
· have := sum_inv_sqrt_le (Nat.pos_of_ne_zero ha)
191-
rw [sum_range_succ] at this
192-
linarith [div_nonneg zero_le_one
193-
(Real.sqrt_nonneg (↑(pullCount A a n ω) : ℝ))]
194-
_ = √(8 * σ2 * c) * (2 * ∑ a : Fin K, √↑(pullCount A a n ω)) := by
195-
simp only [mul_sum]
196-
_ ≤ √(8 * σ2 * c) * (2 * √(↑K * ↑n)) := by
197-
gcongr
198-
calc ∑ a : Fin K, √↑(pullCount A a n ω)
199-
≤ √(↑(Finset.univ.card) * ∑ a, ↑(pullCount A a n ω)) :=
200-
sum_sqrt_le Finset.univ fun a => by positivity
201-
_ = √(↑K * ↑n) := by
202-
congr 1; rw [Finset.card_fin]; congr 1
203-
have h := sum_pullCount (A := A) (t := n) (ω := ω)
204-
exact_mod_cast h
205-
_ = 2 * √(8 * σ2 * c) * √(↑K * ↑n) := by ring
206-
· have h0 : ∀ s ∈ range n,
207-
2 * √(2 * σ2 * c / ↑(pullCount A (A s ω) s ω)) = 0 :=
208-
fun s _ => by
209-
open Real in
210-
have : 2 * σ2 * c / ↑(pullCount A (A s ω) s ω) ≤ 0 :=
211-
div_nonpos_of_nonpos_of_nonneg (by linarith) (Nat.cast_nonneg _)
212-
simp [sqrt_eq_zero'.mpr this]
213-
rw [sum_congr rfl h0]; simp only [sum_const_zero]; positivity
214-
-- Bound ∑_{S0}: each term = hi - μ ≤ hi - lo, and #S0 ≤ K
215-
have hterm_S0 : ∀ s ∈ S0, ucb A R' lo hi σ2 δ (A s ω) s ω -
216-
μ (A s ω) ≤ hi - lo := fun s hs => by
217-
have hpc : pullCount A (A s ω) s ω = 0 := (Finset.mem_filter.mp hs).2
218-
simp only [ucb, hpc, ↓reduceIte]
219-
linarith [(hm (A s ω)).1]
220-
have h_card_S0 : #S0 ≤ K := by
221-
calc #S0 ≤ #(Finset.univ : Finset (Fin K)) :=
222-
Finset.card_le_card_of_injOn (fun s => A s ω)
223-
(fun _ _ => Finset.mem_coe.mpr (Finset.mem_univ _)) (by
224-
intro s₁ hs₁ s₂ hs₂ heq
225-
have hpc₁ := (Finset.mem_filter.mp (Finset.mem_coe.mp hs₁)).2
226-
have hpc₂ := (Finset.mem_filter.mp (Finset.mem_coe.mp hs₂)).2
227-
by_contra h_ne
228-
rcases lt_or_gt_of_ne h_ne with h_lt | h_lt
229-
· have : s₁ ∈ (range s₂).filter (fun i => A i ω = A s₂ ω) := by
230-
simp [mem_range.mpr h_lt, heq]
231-
exact absurd hpc₂ (show pullCount A (A s₂ ω) s₂ ω ≠ 0 from
232-
Finset.card_ne_zero_of_mem this)
233-
· have : s₂ ∈ (range s₁).filter (fun i => A i ω = A s₁ ω) := by
234-
simp [mem_range.mpr h_lt, ← heq]
235-
exact absurd hpc₁ (show pullCount A (A s₁ ω) s₁ ω ≠ 0 from
236-
Finset.card_ne_zero_of_mem this))
237-
_ = K := Finset.card_fin K
238-
calc ∑ s ∈ S0, (ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω))
239-
≤ ∑ _s ∈ S0, (hi - lo) := sum_le_sum hterm_S0
240-
_ = #S0 * (hi - lo) := by rw [sum_const, nsmul_eq_mul]
241-
_ ≤ ↑K * (hi - lo) := by gcongr; linarith
242-
_ = (hi - lo) * ↑K := by ring
147+
_ ≤ ∑ s ∈ range n, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) :=
148+
sum_le_sum_of_subset_of_nonneg (filter_subset _ _) (fun _ _ _ => by positivity)
149+
_ = 2 * √(2 * σ2 * Real.log (1 / δ)) * ∑ s ∈ range n, (1 / √(pullCount A (A s ω) s ω)) := by
150+
rw [mul_sum]
151+
congr with s
152+
rw [Real.sqrt_div' _ (by positivity)]
153+
ring
154+
_ = 2 * √(2 * σ2 * Real.log (1 / δ)) *
155+
∑ a, ∑ j ∈ range (pullCount A a n ω), (1 / √j) := by
156+
rw [sum_comp_pullCount (fun j => 1 / √j)]
157+
_ ≤ 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * ∑ a, √(pullCount A a n ω)) := by
158+
rw [mul_sum _ _ 2]
159+
gcongr with a
160+
by_cases ha : pullCount A a n ω = 0
161+
· simp [ha]
162+
· have hi := sum_inv_sqrt_le (Nat.pos_of_ne_zero ha)
163+
rw [sum_range_succ] at hi
164+
have : 0 ≤ 1 / √(pullCount A a n ω) := by positivity
165+
linarith
166+
_ ≤ 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * √(K * ∑ a, (pullCount A a n ω))) := by
167+
gcongr
168+
have h := sum_sqrt_le Finset.univ (fun a => Nat.cast_nonneg (pullCount A a n ω))
169+
rw [Finset.card_fin] at h
170+
exact_mod_cast h
171+
_ = 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * √(K * n)) := by
172+
congr
173+
exact sum_pullCount (ω := ω)
174+
_ = 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by
175+
ring_nf
176+
rw [← Real.sqrt_mul' _ (by positivity)]
177+
ring_nf
178+
243179

244180
end UCB
245181

@@ -670,7 +606,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp
670606
(n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) :
671607
P[IsBayesAlgEnvSeq.regret κ E A n]
672608
≤ (u - l) * ↑K + 2 * (↑K + 1) * (u - l) * n ^ 2 * δ +
673-
2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by
609+
4 * √(2 * ↑σ2 * Real.log (1 / δ) * ↑K * ↑n) := by
674610
have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _)
675611
have hlo : l ≤ u := h1.trans h2
676612
let bestArm := IsBayesAlgEnvSeq.bestAction κ E
@@ -824,10 +760,10 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp
824760
linarith
825761
have h_second_Eδ : ∀ ω ∈ Eδ,
826762
∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω)
827-
≤ (u - l) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by
763+
≤ (u - l) * ↑K + 4 * √(2 * ↑σ2 * Real.log (1 / δ) * ↑K * ↑n) := by
828764
intro ω hω
829-
exact sum_ucb_sub_mean_le (μ := fun a => armMean a ω)
830-
(hm (E ω)) hlo (↑σ2) δ n ω hω
765+
exact sum_ucb_sub_mean_le (fun a ↦ armMean a ω) (hm (E ω)) hlo
766+
(fun s hs hpc => hω s hs (A s ω) hpc)
831767
have h_prob : P Eδᶜ ≤ ENNReal.ofReal (2 * K * n * δ) := by
832768
have : Eδᶜ = {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧
833769
√(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤
@@ -878,7 +814,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp
878814
(armMean (bestArm ω) ω - uc (bestArm ω) s ω)
879815
set f2 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n,
880816
(uc (A s ω) s ω - armMean (A s ω) ω)
881-
set B := (u - l) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n)
817+
set B := (u - l) * ↑K + 4 * √(2 * ↑σ2 * Real.log (1 / δ) * ↑K * ↑n)
882818
have h1g : ∫ ω in Fδ, f1 ω ∂P ≤ 0 :=
883819
setIntegral_nonpos hFδ_meas fun ω hω ↦ h_first_Fδ ω hω
884820
have h1b : ∫ ω in Fδᶜ, f1 ω ∂P ≤ ↑n * (u - l) * P.real Fδᶜ := by
@@ -954,19 +890,16 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω]
954890
rw [one_div_one_div, Real.log_pow]; norm_cast
955891
calc P[IsBayesAlgEnvSeq.regret κ E A t]
956892
≤ (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * ↑t ^ 2 * (1 / (↑t) ^ 2)
957-
+ 2 * √(8 * ↑σ2 * Real.log (1 / (1 / (↑t) ^ 2))) * √(↑K * ↑t) :=
893+
+ 4 * √(2 * ↑σ2 * Real.log (1 / (1 / (↑t) ^ 2)) * ↑K * ↑t) :=
958894
bayesRegret_le_of_delta (hK := hK) (E := E) (A := A) (R' := R') (Q := Q)
959895
(κ := κ) (P := P) h hσ2 hs hm t (1 / (↑t) ^ 2) hδ hδ1
960-
_ = (3 * ↑K + 2) * (hi - lo) + 8 * (√(↑σ2 * Real.log ↑t) * √(↑K * ↑t)) := by
961-
rw [h_first, h_log,
962-
show (8 : ℝ) * ↑σ2 * (2 * Real.log ↑t) = 4 ^ 2 * (↑σ2 * Real.log ↑t) by ring,
963-
Real.sqrt_mul (by positivity : (0 : ℝ) ≤ 4 ^ 2),
964-
Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 4)]
965-
ring
966896
_ = (3 * ↑K + 2) * (hi - lo) + 8 * √(↑σ2 * ↑K * ↑t * Real.log ↑t) := by
967-
rw [← Real.sqrt_mul (by positivity :
968-
0 ≤ ↑σ2 * Real.log ↑t)]
969-
congr 1; ring_nf
897+
rw [h_first, h_log]; congr 1
898+
rw [show (2 : ℝ) * ↑σ2 * (2 * Real.log ↑t) * ↑K * ↑t =
899+
(2 : ℝ) ^ 2 * (↑σ2 * ↑K * ↑t * Real.log ↑t) from by ring,
900+
Real.sqrt_mul (by positivity : (0 : ℝ) ≤ 2 ^ 2),
901+
Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 2)]
902+
ring
970903

971904
end TS
972905

‎LeanBandits/SequentialLearning/FiniteActions.lean‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -729,6 +729,19 @@ lemma sum_pullCount [Fintype α] {ω : Ω} : ∑ a, pullCount A a t ω = t := by
729729
rw [sum_pullCount_mul]
730730
simp
731731

732+
lemma sum_comp_pullCount [Fintype α] [AddCommMonoid R] (f : ℕ → R) (t : ℕ) (ω : Ω) :
733+
∑ s ∈ range t, f (pullCount A (A s ω) s ω) = ∑ a, ∑ j ∈ range (pullCount A a t ω), f j := by
734+
induction t with
735+
| zero => simp
736+
| succ n ih =>
737+
have hf : f (pullCount A (A n ω) n ω) =
738+
∑ a, if A n ω = a then f (pullCount A a n ω) else 0 := by simp
739+
simp_rw [sum_range_succ, ih, hf, ← sum_add_distrib, pullCount_add_one]
740+
congr 1 with a
741+
split_ifs
742+
· simp [sum_range_succ]
743+
· simp
744+
732745
section SumRewards
733746

734747
/-- Sum of rewards obtained when pulling action `a` up to time `t` (exclusive). -/

0 commit comments

Comments
 (0)