@@ -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
244180end 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
971904end TS
972905
0 commit comments