Skip to content

Commit 8dc181c

Browse files
committed
Refactor TS.lean (in progress)
1 parent 636c66a commit 8dc181c

2 files changed

Lines changed: 98 additions & 147 deletions

File tree

‎LeanMachineLearning/BanditAlgorithms/TS.lean‎

Lines changed: 3 additions & 145 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,6 @@ Authors: Rémy Degenne, Paulo Rauber
55
-/
66
module
77

8-
public import LeanMachineLearning.Bandit.SumRewards
98
public import LeanMachineLearning.BanditAlgorithms.Uniform
109
public import LeanMachineLearning.SequentialLearning.AlgorithmDensity
1110

@@ -195,149 +194,6 @@ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a,
195194

196195
end 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-
341197
variable {K : ℕ}
342198
variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔]
343199
variable (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

‎LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean‎

Lines changed: 95 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,15 +5,15 @@ Authors: Rémy Degenne, Paulo Rauber
55
-/
66
module
77

8-
public import LeanMachineLearning.Bandit.Regret
8+
public import LeanMachineLearning.Bandit.SumRewards
99
public import LeanMachineLearning.ForMathlib.MeasurableArgMax
10-
public import LeanMachineLearning.SequentialLearning.StationaryEnv
1110

1211
/-! # Bayesian stationary environments -/
1312

1413
@[expose] public section
1514

1615
open MeasureTheory ProbabilityTheory Finset
16+
open scoped ENNReal NNReal
1717

1818
namespace Learning
1919

@@ -232,6 +232,99 @@ lemma ae_IsAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P
232232

233233
end CondDistribIsAlgEnvSeq
234234

235+
section HasSubgaussianMGF
236+
237+
variable {K : ℕ} [Nonempty (Fin K)]
238+
variable {κ : Kernel (𝓔 × Fin K) ℝ} {alg : Algorithm (Fin K) ℝ}
239+
variable {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → ℝ}
240+
variable [IsProbabilityMeasure P]
241+
242+
private lemma sqrt_two_mul_le_abs_sub_of_sqrt_div_le {s μ σ L : ℝ} {k : ℕ} (hk : k ≠ 0)
243+
(h : √(2 * σ * L / k) ≤ |s / k - μ|) : √(2 * k * σ * L) ≤ |s - k * μ| := by
244+
have hkp : (0 : ℝ) < k := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hk)
245+
have hks : s / (k : ℝ) - μ = (s - k * μ) / k := by field_simp
246+
by_cases hc : 0 ≤ 2 * σ * L
247+
· have hkc : (2 * k * σ * L : ℝ) = 2 * σ * L / k * k ^ 2 := by field_simp
248+
calc √(2 * k * σ * L)
249+
_ = √(2 * σ * L / k) * k := by
250+
rw [hkc, Real.sqrt_mul (div_nonneg hc hkp.le), Real.sqrt_sq hkp.le]
251+
_ ≤ |s / k - μ| * k := mul_le_mul_of_nonneg_right h hkp.le
252+
_ = |s - k * μ| := by rw [hks, abs_div, abs_of_pos hkp, div_mul_cancel₀ _ (ne_of_gt hkp)]
253+
· push Not at hc
254+
calc √(2 * k * σ * L)
255+
_ = 0 := Real.sqrt_eq_zero_of_nonpos (by nlinarith)
256+
_ ≤ |s - k * μ| := abs_nonneg _
257+
258+
lemma prob_abs_empMean_sub_actionMean_ge_le [IsMarkovKernel κ]
259+
(h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0}
260+
(hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a)))
261+
{δ : ℝ} (hδ : 0 < δ) (n : ℕ) :
262+
P {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧
263+
√(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤
264+
|empMean A R' a t ω - actionMean κ E a ω|}
265+
≤ ENNReal.ofReal (2 * K * (n - 1) * δ) := by
266+
have := h.measurable_E
267+
have := h.measurable_A
268+
have := h.measurable_R
269+
let S : Set (𝓔 × (ℕ → Fin K × ℝ)) := {(e, τ) | ∃ a, ∃ t < n, pullCount IT.action a t τ ≠ 0 ∧
270+
√(2 * pullCount IT.action a t τ * σ2 * Real.log (1 / δ)) ≤
271+
|sumRewards IT.action IT.reward a t τ - pullCount IT.action a t τ * actionMean κ id a e|}
272+
calc _
273+
_ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by
274+
rw [Measure.map_apply (by fun_prop) (by measurability)]
275+
apply measure_mono
276+
intro ω ⟨t, ht, a, hpc, hle⟩
277+
exact ⟨a, t, ht, hpc,
278+
sqrt_two_mul_le_abs_sub_of_sqrt_div_le hpc (by simpa [empMean] using hle)⟩
279+
_ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by
280+
rw [← compProd_map_condDistrib (by fun_prop)]
281+
_ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) :=
282+
Measure.compProd_apply (by measurability)
283+
_ ≤ ∫⁻ _, ENNReal.ofReal (2 * K * (n - 1) * δ) ∂(P.map E) := by
284+
apply lintegral_mono_ae
285+
rw [h.hasLaw_env.map_eq]
286+
filter_upwards [h.ae_IsAlgEnvSeq] with e he
287+
convert Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 (hs e) he hδ
288+
exact (Fintype.card_fin K).symm
289+
_ = _ := by simp [Measure.map_apply h.measurable_E]
290+
291+
lemma prob_abs_empMean_bestAction_sub_actionMean_ge_le [IsMarkovKernel κ]
292+
(h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2)
293+
(hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a)))
294+
{δ : ℝ} (hδ : 0 < δ) (n : ℕ) :
295+
P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧
296+
√(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω : ℝ)) ≤
297+
|empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω|}
298+
≤ ENNReal.ofReal (2 * (n - 1) * δ) := by
299+
have := h.measurable_E
300+
have := h.measurable_A
301+
have := h.measurable_R
302+
let S : Set (𝓔 × (ℕ → Fin K × ℝ)) :=
303+
{(e, τ) | ∃ t < n, pullCount IT.action (bestAction κ id e) t τ ≠ 0 ∧
304+
√(2 * pullCount IT.action (bestAction κ id e) t τ * σ2 * Real.log (1 / δ)) ≤
305+
|sumRewards IT.action IT.reward (bestAction κ id e) t τ -
306+
pullCount IT.action (bestAction κ id e) t τ * actionMean κ id (bestAction κ id e) e|}
307+
calc _
308+
_ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by
309+
rw [Measure.map_apply (by fun_prop) (by measurability)]
310+
apply measure_mono
311+
intro ω ⟨t, ht, hpc, hle⟩
312+
exact ⟨t, ht, hpc,
313+
sqrt_two_mul_le_abs_sub_of_sqrt_div_le hpc (by simpa [empMean] using hle)⟩
314+
_ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by
315+
rw [← compProd_map_condDistrib (by fun_prop)]
316+
_ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) :=
317+
Measure.compProd_apply (by measurability)
318+
_ ≤ ∫⁻ _, ENNReal.ofReal (2 * (n - 1) * δ) ∂(P.map E) := by
319+
apply lintegral_mono_ae
320+
rw [h.hasLaw_env.map_eq]
321+
filter_upwards [h.ae_IsAlgEnvSeq] with e he
322+
exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le (ν := κ.sectR e) hσ2
323+
(hs e (bestAction κ id e)) he hδ
324+
_ = _ := by simp [Measure.map_apply h.measurable_E]
325+
326+
end HasSubgaussianMGF
327+
235328
end IsBayesAlgEnvSeq
236329

237330
section IsAlgEnvSeq

0 commit comments

Comments
 (0)