@@ -5,7 +5,6 @@ Authors: Rémy Degenne, Paulo Rauber
55-/
66module
77
8- public import LeanMachineLearning.Bandit.SumRewards
98public import LeanMachineLearning.BanditAlgorithms.Uniform
109public import LeanMachineLearning.SequentialLearning.AlgorithmDensity
1110
@@ -195,149 +194,6 @@ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a,
195194
196195end UCB
197196
198- end TS
199-
200- end Bandits
201-
202- namespace Learning.IsBayesAlgEnvSeq
203-
204- variable {K : ℕ} [Nonempty (Fin K)]
205- variable {𝓔 Ω : Type *} [MeasurableSpace 𝓔] [MeasurableSpace Ω]
206- variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] {alg : Algorithm (Fin K) ℝ}
207- variable {E : Ω → 𝓔} {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → ℝ}
208- variable {P : Measure Ω} [IsProbabilityMeasure P]
209-
210- lemma prob_abs_sumRewards_sub_pullCount_mul_actionMean_ge_le
211- (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0 }
212- (hσ2 : 0 < σ2 ) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a)))
213- {δ : ℝ} (hδ : 0 < δ) (n : ℕ) :
214- P {ω | ∃ a, ∃ t < n, pullCount A a t ω ≠ 0 ∧
215- √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤
216- |sumRewards A R' a t ω - pullCount A a t ω * actionMean κ E a ω|}
217- ≤ ENNReal.ofReal (2 * K * (n - 1 ) * δ) := by
218- have := h.measurable_E
219- have := h.measurable_A
220- have := h.measurable_R
221- let B e := {τ | ∃ a, ∃ t < n, pullCount IT.action a t τ ≠ 0 ∧
222- √(2 * pullCount IT.action a t τ * σ2 * Real.log (1 / δ)) ≤
223- |sumRewards IT.action IT.reward a t τ - pullCount IT.action a t τ * actionMean κ id a e|}
224- calc P ((fun ω ↦ (E ω, trajectory A R' ω)) ⁻¹' {(e, τ) : 𝓔 × (ℕ → Fin K × ℝ) | τ ∈ B e})
225- _ = (P.map (fun ω ↦ (E ω, trajectory A R' ω))) {(e, τ) : 𝓔 × (ℕ → Fin K × ℝ) | τ ∈ B e} :=
226- (Measure.map_apply (by fun_prop) (by measurability)).symm
227- _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) {(e, τ) : 𝓔 × _ | τ ∈ B e} := by
228- rw [← compProd_map_condDistrib (by fun_prop)]
229- _ = ∫⁻ e, (condDistrib (trajectory A R') E P e) (B e) ∂(P.map E) := by
230- rw [Measure.compProd_apply (by measurability)]
231- rfl
232- _ ≤ ∫⁻ e, ENNReal.ofReal (2 * (Fintype.card (Fin K)) * (n - 1 ) * δ) ∂(P.map E) := by
233- apply lintegral_mono_ae
234- rw [h.hasLaw_env.map_eq]
235- filter_upwards [h.ae_IsAlgEnvSeq] with e he
236- exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 (hs e) he hδ
237- _ = ENNReal.ofReal (2 * K * (n - 1 ) * δ) := by
238- simp [lintegral_const, Measure.map_apply h.measurable_E]
239-
240- lemma prob_abs_sumRewards_bestAction_sub_pullCount_mul_actionMean_ge_le
241- (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0 } (hσ2 : 0 < σ2 )
242- (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a)))
243- {δ : ℝ} (hδ : 0 < δ) (n : ℕ) :
244- P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧
245- √(2 * pullCount A (bestAction κ E ω) t ω * σ2 * Real.log (1 / δ)) ≤
246- |sumRewards A R' (bestAction κ E ω) t ω -
247- pullCount A (bestAction κ E ω) t ω * actionMean κ E (bestAction κ E ω) ω|}
248- ≤ ENNReal.ofReal (2 * (n - 1 ) * δ) := by
249- have := h.measurable_E
250- have := h.measurable_A
251- have := h.measurable_R
252- let B e := {τ | ∃ t < n, pullCount IT.action (bestAction κ id e) t τ ≠ 0 ∧
253- √(2 * pullCount IT.action (bestAction κ id e) t τ * σ2 * Real.log (1 / δ)) ≤
254- |sumRewards IT.action IT.reward (bestAction κ id e) t τ -
255- pullCount IT.action (bestAction κ id e) t τ * actionMean κ id (bestAction κ id e) e|}
256- calc P ((fun ω ↦ (E ω, trajectory A R' ω)) ⁻¹' {(e, τ) : 𝓔 × (ℕ → Fin K × ℝ) | τ ∈ B e})
257- _ = (P.map (fun ω ↦ (E ω, trajectory A R' ω))) {(e, τ) : 𝓔 × (ℕ → Fin K × ℝ) | τ ∈ B e} :=
258- (Measure.map_apply (by fun_prop) (by measurability)).symm
259- _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) {(e, τ) : 𝓔 × _ | τ ∈ B e} := by
260- rw [← compProd_map_condDistrib (by fun_prop)]
261- _ = ∫⁻ e, (condDistrib (trajectory A R') E P e) (B e) ∂(P.map E) := by
262- rw [Measure.compProd_apply (by measurability)]
263- rfl
264- _ ≤ ∫⁻ e, ENNReal.ofReal (2 * (n - 1 ) * δ) ∂(P.map E) := by
265- apply lintegral_mono_ae
266- rw [h.hasLaw_env.map_eq]
267- filter_upwards [h.ae_IsAlgEnvSeq] with e he
268- exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le (ν := κ.sectR e)
269- hσ2 (hs e (bestAction κ id e)) he hδ
270- _ = ENNReal.ofReal (2 * (n - 1 ) * δ) := by
271- simp [lintegral_const, Measure.map_apply h.measurable_E]
272-
273- omit [Nonempty (Fin K)] [MeasurableSpace 𝓔] [MeasurableSpace Ω] in
274- private lemma abs_sumRewards_sub_pullCount_mul_ge {a : Fin K} {n : ℕ} {ω : Ω}
275- {μ σ2 δ : ℝ} (hpc : pullCount A a n ω ≠ 0 )
276- (h : √(2 * σ2 * Real.log (1 / δ) / pullCount A a n ω) ≤
277- |empMean A R' a n ω - μ|) :
278- √(2 * pullCount A a n ω * σ2 * Real.log (1 / δ)) ≤
279- |sumRewards A R' a n ω - pullCount A a n ω * μ| := by
280- have hk : (0 : ℝ) < pullCount A a n ω := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hpc)
281- by_cases hc : 0 ≤ 2 * σ2 * Real.log (1 / δ)
282- · calc
283- _ = √(2 * σ2 * Real.log (1 / δ) / pullCount A a n ω) * pullCount A a n ω := by
284- have : 2 * pullCount A a n ω * σ2 * Real.log (1 / δ) =
285- 2 * σ2 * Real.log (1 / δ) / pullCount A a n ω * pullCount A a n ω ^ 2 := by
286- field_simp
287- rw [this, Real.sqrt_mul (div_nonneg hc hk.le), Real.sqrt_sq hk.le]
288- _ ≤ |sumRewards A R' a n ω / pullCount A a n ω - μ| * pullCount A a n ω :=
289- mul_le_mul_of_nonneg_right h hk.le
290- _ = |sumRewards A R' a n ω - pullCount A a n ω * μ| := by
291- have : sumRewards A R' a n ω / ↑(pullCount A a n ω) - μ =
292- (sumRewards A R' a n ω - pullCount A a n ω * μ) / pullCount A a n ω := by
293- field_simp
294- rw [this, abs_div, abs_of_pos hk, div_mul_cancel₀ _ (ne_of_gt hk)]
295- · calc
296- _ = 0 := Real.sqrt_eq_zero_of_nonpos (by push Not at hc; nlinarith)
297- _ ≤ |sumRewards A R' a n ω - pullCount A a n ω * μ| := abs_nonneg _
298-
299- lemma prob_abs_empMean_sub_actionMean_ge_le
300- (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0 }
301- (hσ2 : 0 < σ2 ) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a)))
302- {δ : ℝ} (hδ : 0 < δ) (n : ℕ) :
303- P {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧
304- √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤
305- |empMean A R' a t ω - actionMean κ E a ω|}
306- ≤ ENNReal.ofReal (2 * K * (n - 1 ) * δ) :=
307- calc
308- _ ≤ P {ω | ∃ a, ∃ t < n, pullCount A a t ω ≠ 0 ∧
309- √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤
310- |sumRewards A R' a t ω - pullCount A a t ω * actionMean κ E a ω|} := by
311- apply measure_mono
312- intro ω ⟨t, ht, a, hpc, hle⟩
313- exact ⟨a, t, ht, hpc, abs_sumRewards_sub_pullCount_mul_ge hpc hle⟩
314- _ ≤ _ := h.prob_abs_sumRewards_sub_pullCount_mul_actionMean_ge_le hσ2 hs hδ n
315-
316- lemma prob_abs_empMean_bestAction_sub_actionMean_ge_le
317- (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0 } (hσ2 : 0 < σ2 )
318- (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a)))
319- {δ : ℝ} (hδ : 0 < δ) (n : ℕ) :
320- P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧
321- √(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω : ℝ)) ≤
322- |empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω|}
323- ≤ ENNReal.ofReal (2 * (n - 1 ) * δ) :=
324- calc
325- _ ≤ P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧
326- √(2 * pullCount A (bestAction κ E ω) t ω * σ2 * Real.log (1 / δ)) ≤
327- |sumRewards A R' (bestAction κ E ω) t ω -
328- pullCount A (bestAction κ E ω) t ω *
329- actionMean κ E (bestAction κ E ω) ω|} := by
330- apply measure_mono
331- intro ω ⟨t, ht, hpc, hle⟩
332- exact ⟨t, ht, hpc, abs_sumRewards_sub_pullCount_mul_ge hpc hle⟩
333- _ ≤ _ :=
334- h.prob_abs_sumRewards_bestAction_sub_pullCount_mul_actionMean_ge_le
335- hσ2 hs hδ n
336-
337- end Learning.IsBayesAlgEnvSeq
338-
339- namespace Bandits.TS
340-
341197variable {K : ℕ}
342198variable {𝓔 : Type *} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔]
343199variable (hK : 0 < K)
@@ -670,4 +526,6 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω]
670526 Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 2 )]
671527 ring
672528
673- end Bandits.TS
529+ end TS
530+
531+ end Bandits
0 commit comments