From 9a760eb4d9b2e736e046f2506c120b4f85487d6b Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 26 Sep 2025 16:28:14 +0200 Subject: [PATCH 01/11] etc work --- LeanBandits/ETC.lean | 55 ++++++++++++++++++++++++++++++++++++++--- LeanBandits/Regret.lean | 6 +++++ 2 files changed, 58 insertions(+), 3 deletions(-) diff --git a/LeanBandits/ETC.lean b/LeanBandits/ETC.lean index e7f950ae..b38738d0 100644 --- a/LeanBandits/ETC.lean +++ b/LeanBandits/ETC.lean @@ -58,14 +58,63 @@ lemma arm_ae_eq_etcNextArm (n : ℕ) : have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK exact arm_detAlgorithm_ae_eq n -lemma pullCount_mul (a : Fin K) : - pullCount a (K * m) =ᵐ[𝔓b] fun _ ↦ m := by +lemma arm_of_lt {n : ℕ} (hn : n < K * m) : + arm n =ᵐ[𝔓b] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := by + cases n with + | zero => exact arm_zero + | succ n => + filter_upwards [arm_ae_eq_etcNextArm n] with h hn_eq + rw [hn_eq, nextArm, dif_pos] + grind + +lemma arm_mul : + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + arm (K * m) =ᵐ[𝔓b] fun h ↦ measurableArgmax (empMean' (K*m+1)) (fun i ↦ h i) := by + have : K * m = (K * m - 1) + 1 := by sorry + rw [this] + filter_upwards [arm_ae_eq_etcNextArm (K * m - 1)] with h hn_eq + rw [hn_eq, nextArm, dif_neg (by simp), dif_pos rfl] sorry +lemma arm_of_ge {n : ℕ} (hn : K * m ≤ n) : arm n =ᵐ[𝔓b] arm (K * m) := by + sorry + +lemma pullCount_mul (a : Fin K) : pullCount a (K * m) =ᵐ[𝔓b] fun _ ↦ m := by + rw [Filter.EventuallyEq] + simp_rw [pullCount_eq_sum] + have h_arm (n : range (K * m)) : arm n =ᵐ[𝔓b] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := + arm_of_lt (mem_range.mp n.2) + simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_arm + filter_upwards [h_arm] with ω h_arm + have h_arm' {i : ℕ} (hi : i ∈ range (K * m)) : arm i ω = ⟨i % K, Nat.mod_lt _ hK⟩ := h_arm ⟨i, hi⟩ + calc (∑ s ∈ range (K * m), if arm s ω = a then 1 else 0) + _ = (∑ s ∈ range (K * m), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := + sum_congr rfl fun s hs ↦ by rw [h_arm' hs] + _ = m := by + sorry + +lemma pullCount_add_one_of_ge (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : + pullCount a (n + 1) + =ᵐ[𝔓b] fun ω ↦ pullCount a n ω + {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by + simp_rw [Filter.EventuallyEq, pullCount_add_one] + filter_upwards [arm_of_ge hn] with ω h_arm + congr + lemma pullCount_of_ge (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : pullCount a n =ᵐ[𝔓b] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by - sorry + have h_ae n : K * m ≤ n → pullCount a (n + 1) + =ᵐ[𝔓b] fun ω ↦ pullCount a n ω + {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := + pullCount_add_one_of_ge a + simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_ae + have h_ae_Km : pullCount a (K * m) =ᵐ[𝔓b] fun _ ↦ m := pullCount_mul a + filter_upwards [h_ae_Km, h_ae] with ω h_Km h_ae + induction n, hn using Nat.le_induction with + | base => simp [h_Km] + | succ n hmn h_ind => + rw [h_ae n hmn, h_ind, add_assoc, ← add_one_mul] + congr + grind lemma prob_arm_mul_eq_le (a : Fin K) : (𝔓b).real {ω | arm (K * m) ω = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index db6de8fc..aaf1c1f1 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -61,6 +61,12 @@ lemma pullCount_eq_pullCount_add_one (t : ℕ) (h : ℕ → α × ℝ) : lemma pullCount_eq_pullCount (ha : arm t h ≠ a) : pullCount a (t + 1) h = pullCount a t h := by simp [pullCount, range_succ, filter_insert, ha] +lemma pullCount_add_one : + pullCount a (t + 1) h = pullCount a t h + if arm t h = a then 1 else 0 := by + split_ifs with h + · rw [← h, pullCount_eq_pullCount_add_one] + · rw [pullCount_eq_pullCount h, add_zero] + lemma pullCount_eq_sum (a : α) (t : ℕ) (h : ℕ → α × ℝ) : pullCount a t h = ∑ s ∈ range t, if arm s h = a then 1 else 0 := by simp [pullCount] From b0c439c9a52c1820cd8489cab3cf739a60810392 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 26 Sep 2025 16:53:47 +0200 Subject: [PATCH 02/11] add subgaussian facts --- LeanBandits/ETC.lean | 19 ++++++--- LeanBandits/ForMathlib/SubGaussian.lean | 52 +++++++++++++++++++++++++ 2 files changed, 65 insertions(+), 6 deletions(-) create mode 100644 LeanBandits/ForMathlib/SubGaussian.lean diff --git a/LeanBandits/ETC.lean b/LeanBandits/ETC.lean index b38738d0..7ec9d9d9 100644 --- a/LeanBandits/ETC.lean +++ b/LeanBandits/ETC.lean @@ -5,6 +5,7 @@ Authors: Rémy Degenne -/ import Mathlib.Probability.Moments.SubGaussian import LeanBandits.AlgorithmBuilding +import LeanBandits.ForMathlib.SubGaussian import LeanBandits.Regret /-! # The Explore-Then-Commit Algorithm @@ -116,7 +117,7 @@ lemma pullCount_of_ge (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : congr grind -lemma prob_arm_mul_eq_le (a : Fin K) : +lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) : (𝔓b).real {ω | arm (K * m) ω = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK -- extend the probability space to include the stream of independent rewards @@ -156,7 +157,11 @@ lemma prob_arm_mul_eq_le (a : Fin K) : sorry sorry · intro i him - sorry + rw [← one_add_one_eq_two] + refine HasSubgaussianMGF.sub_of_indepFun ?_ ?_ ?_ + · sorry + · sorry + · sorry · have : 0 ≤ gap ν a := gap_nonneg positivity · congr 1 @@ -166,7 +171,8 @@ lemma prob_arm_mul_eq_le (a : Fin K) : not_false_eq_true, pow_eq_zero_iff, Nat.cast_eq_zero] norm_num -lemma expectation_pullCount_le (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : +lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : 𝔓b[fun ω ↦ (pullCount a n ω : ℝ)] ≤ m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by have : (fun ω ↦ (pullCount a n ω : ℝ)) @@ -188,10 +194,11 @@ lemma expectation_pullCount_le (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : simp rw [integral_indicator_const, smul_eq_mul, mul_one] · rw [← neg_mul] - exact prob_arm_mul_eq_le a + exact prob_arm_mul_eq_le hν a · exact (measurableSet_singleton _).preimage (by fun_prop) -lemma regret_le (n : ℕ) (hn : K * m ≤ n) : +lemma regret_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + (n : ℕ) (hn : K * m ≤ n) : 𝔓b[regret ν n] ≤ ∑ a, gap ν a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4)) := by simp_rw [regret_eq_sum_pullCount_mul_gap] rw [integral_finset_sum] @@ -202,7 +209,7 @@ lemma regret_le (n : ℕ) (hn : K * m ≤ n) : rw [mul_comm (gap _ _), integral_mul_const] gcongr · exact gap_nonneg - · exact expectation_pullCount_le a hn + · exact expectation_pullCount_le hν a hn end ETC diff --git a/LeanBandits/ForMathlib/SubGaussian.lean b/LeanBandits/ForMathlib/SubGaussian.lean new file mode 100644 index 00000000..de535e60 --- /dev/null +++ b/LeanBandits/ForMathlib/SubGaussian.lean @@ -0,0 +1,52 @@ +/- +Copyright (c) 2025 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +import Mathlib.Probability.Moments.SubGaussian + +open MeasureTheory +open scoped ENNReal NNReal + +namespace ProbabilityTheory + +namespace Kernel.HasSubgaussianMGF + +variable {Ω Ω' : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} + {ν : Measure Ω'} {κ : Kernel Ω' Ω} {X : Ω → ℝ} {c : ℝ≥0} + +protected lemma const_mul (h : HasSubgaussianMGF X c κ ν) (r : ℝ) : + HasSubgaussianMGF (fun ω ↦ r * X ω) (⟨r ^ 2, sq_nonneg r⟩ * c) κ ν where + integrable_exp_mul t := by + simp_rw [← mul_assoc] + exact h.integrable_exp_mul (t * r) + mgf_le := by + filter_upwards [h.mgf_le] with ω hω t + specialize hω (t * r) + rw [mgf] at hω ⊢ + simp_rw [← mul_assoc] + refine hω.trans_eq ?_ + congr 1 + simp only [NNReal.coe_mul, NNReal.coe_mk] + ring + +end Kernel.HasSubgaussianMGF + +namespace HasSubgaussianMGF + +variable {Ω : Type*} {m mΩ : MeasurableSpace Ω} {μ : Measure Ω} {X : Ω → ℝ} {c : ℝ≥0} + +protected lemma const_mul (h : HasSubgaussianMGF X c μ) (r : ℝ) : + HasSubgaussianMGF (fun ω ↦ r * X ω) (⟨r ^ 2, sq_nonneg r⟩ * c) μ := by + rw [HasSubgaussianMGF_iff_kernel] at h ⊢ + exact Kernel.HasSubgaussianMGF.const_mul h r + +lemma sub_of_indepFun {Y : Ω → ℝ} {cX cY : ℝ≥0} (hX : HasSubgaussianMGF X cX μ) + (hY : HasSubgaussianMGF Y cY μ) (hindep : IndepFun X Y μ) : + HasSubgaussianMGF (fun ω ↦ X ω - Y ω) (cX + cY) μ := by + simp_rw [sub_eq_add_neg] + exact hX.add_of_indepFun hY.neg hindep.neg_right + +end HasSubgaussianMGF + +end ProbabilityTheory From 89c8e10882b8885391ed3efb12bae85db3455b9b Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 26 Sep 2025 16:54:44 +0200 Subject: [PATCH 03/11] mk_all --- LeanBandits.lean | 1 + 1 file changed, 1 insertion(+) diff --git a/LeanBandits.lean b/LeanBandits.lean index a5d8fef3..050443eb 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -6,6 +6,7 @@ import LeanBandits.ForMathlib.CondDistrib import LeanBandits.ForMathlib.KernelCompositionLemmas import LeanBandits.ForMathlib.KernelCompositionParallelComp import LeanBandits.ForMathlib.KernelSub +import LeanBandits.ForMathlib.SubGaussian import LeanBandits.ForMathlib.Traj import LeanBandits.Regret import LeanBandits.RewardByCountMeasure From ae73a0e1f3e3258590525f5aeb890a80eeab0d8c Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 26 Sep 2025 16:58:15 +0200 Subject: [PATCH 04/11] add mgf lemma --- LeanBandits/ForMathlib/SubGaussian.lean | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/LeanBandits/ForMathlib/SubGaussian.lean b/LeanBandits/ForMathlib/SubGaussian.lean index de535e60..5291541a 100644 --- a/LeanBandits/ForMathlib/SubGaussian.lean +++ b/LeanBandits/ForMathlib/SubGaussian.lean @@ -10,6 +10,11 @@ open scoped ENNReal NNReal namespace ProbabilityTheory +theorem mgf_const_mul {Ω : Type*} {m : MeasurableSpace Ω} {X : Ω → ℝ} {μ : Measure Ω} + {t : ℝ} (α : ℝ) : mgf (fun ω ↦ α * X ω) μ t = mgf X μ (α * t) := by + rw [← mgf_smul_left] + rfl + namespace Kernel.HasSubgaussianMGF variable {Ω Ω' : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} @@ -23,8 +28,7 @@ protected lemma const_mul (h : HasSubgaussianMGF X c κ ν) (r : ℝ) : mgf_le := by filter_upwards [h.mgf_le] with ω hω t specialize hω (t * r) - rw [mgf] at hω ⊢ - simp_rw [← mul_assoc] + rw [mgf_const_mul, mul_comm] refine hω.trans_eq ?_ congr 1 simp only [NNReal.coe_mul, NNReal.coe_mk] From 3b9d28c90734bced441e423b948fd9d2efe52bdb Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sat, 27 Sep 2025 13:55:58 +0200 Subject: [PATCH 05/11] arms pulled by etc --- LeanBandits/ETC.lean | 45 +++++++++++++++++++++++++++++--------------- 1 file changed, 30 insertions(+), 15 deletions(-) diff --git a/LeanBandits/ETC.lean b/LeanBandits/ETC.lean index 7ec9d9d9..6804c386 100644 --- a/LeanBandits/ETC.lean +++ b/LeanBandits/ETC.lean @@ -27,7 +27,7 @@ def ETC.nextArm (hK : 0 < K) (m n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K := ⟨(n + 1) % K, Nat.mod_lt _ hK⟩ -- for `n = 0` we have pulled arm 0 already, and we pull arm 1 else if hn_eq : n = K * m - 1 then measurableArgmax (empMean' n) h - else (h ⟨n - 1, by simp⟩).1 + else (h ⟨n, by simp⟩).1 @[fun_prop] lemma ETC.measurable_nextArm (hK : 0 < K) (m n : ℕ) : Measurable (nextArm hK m n) := by @@ -68,17 +68,32 @@ lemma arm_of_lt {n : ℕ} (hn : n < K * m) : rw [hn_eq, nextArm, dif_pos] grind -lemma arm_mul : +lemma arm_mul (hm : m ≠ 0) : have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - arm (K * m) =ᵐ[𝔓b] fun h ↦ measurableArgmax (empMean' (K*m+1)) (fun i ↦ h i) := by - have : K * m = (K * m - 1) + 1 := by sorry + arm (K * m) =ᵐ[𝔓b] fun h ↦ measurableArgmax (empMean' (K * m - 1)) (fun i ↦ h i) := by + have : K * m = (K * m - 1) + 1 := by + have : 0 < K * m := Nat.mul_pos hK hm.bot_lt + grind rw [this] filter_upwards [arm_ae_eq_etcNextArm (K * m - 1)] with h hn_eq rw [hn_eq, nextArm, dif_neg (by simp), dif_pos rfl] - sorry + exact this ▸ rfl + +lemma arm_add_one_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : + arm (n + 1) =ᵐ[𝔓b] fun ω ↦ arm n ω := by + filter_upwards [arm_ae_eq_etcNextArm n] with ω hn_eq + rw [hn_eq, nextArm, dif_neg (by grind), dif_neg] + · rfl + · have : 0 < K * m := Nat.mul_pos hK hm.bot_lt + grind -lemma arm_of_ge {n : ℕ} (hn : K * m ≤ n) : arm n =ᵐ[𝔓b] arm (K * m) := by - sorry +lemma arm_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : arm n =ᵐ[𝔓b] arm (K * m) := by + have h_ae n : K * m ≤ n → arm (n + 1) =ᵐ[𝔓b] fun ω ↦ arm n ω := arm_add_one_of_ge hm + simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_ae + filter_upwards [h_ae] with ω h_ae + induction n, hn using Nat.le_induction with + | base => rfl + | succ n hmn h_ind => rw [h_ae n hmn, h_ind] lemma pullCount_mul (a : Fin K) : pullCount a (K * m) =ᵐ[𝔓b] fun _ ↦ m := by rw [Filter.EventuallyEq] @@ -94,19 +109,19 @@ lemma pullCount_mul (a : Fin K) : pullCount a (K * m) =ᵐ[𝔓b] fun _ ↦ m := _ = m := by sorry -lemma pullCount_add_one_of_ge (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : +lemma pullCount_add_one_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : pullCount a (n + 1) =ᵐ[𝔓b] fun ω ↦ pullCount a n ω + {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by simp_rw [Filter.EventuallyEq, pullCount_add_one] - filter_upwards [arm_of_ge hn] with ω h_arm + filter_upwards [arm_of_ge hm hn] with ω h_arm congr -lemma pullCount_of_ge (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : +lemma pullCount_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : pullCount a n =ᵐ[𝔓b] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by have h_ae n : K * m ≤ n → pullCount a (n + 1) =ᵐ[𝔓b] fun ω ↦ pullCount a n ω + {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := - pullCount_add_one_of_ge a + pullCount_add_one_of_ge a hm simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_ae have h_ae_Km : pullCount a (K * m) =ᵐ[𝔓b] fun _ ↦ m := pullCount_mul a filter_upwards [h_ae_Km, h_ae] with ω h_Km h_ae @@ -172,12 +187,12 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i norm_num lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) - (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : + (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : 𝔓b[fun ω ↦ (pullCount a n ω : ℝ)] ≤ m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by have : (fun ω ↦ (pullCount a n ω : ℝ)) =ᵐ[𝔓b] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by - filter_upwards [pullCount_of_ge a hn] with ω h + filter_upwards [pullCount_of_ge a hm hn] with ω h simp only [h, Set.indicator_apply, Set.mem_setOf_eq, mul_ite, mul_one, mul_zero, Nat.cast_add, Nat.cast_ite, CharP.cast_eq_zero, add_right_inj] norm_cast @@ -197,7 +212,7 @@ lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - ( exact prob_arm_mul_eq_le hν a · exact (measurableSet_singleton _).preimage (by fun_prop) -lemma regret_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma regret_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hm : m ≠ 0) (n : ℕ) (hn : K * m ≤ n) : 𝔓b[regret ν n] ≤ ∑ a, gap ν a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4)) := by simp_rw [regret_eq_sum_pullCount_mul_gap] @@ -209,7 +224,7 @@ lemma regret_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν rw [mul_comm (gap _ _), integral_mul_const] gcongr · exact gap_nonneg - · exact expectation_pullCount_le hν a hn + · exact expectation_pullCount_le hν a hm hn end ETC From b575079bd005c4fb98a7e5bda1211f0343322d7a Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sat, 27 Sep 2025 13:57:48 +0200 Subject: [PATCH 06/11] change notation for trajMeasure --- LeanBandits/ETC.lean | 40 ++++++++++++++++++++-------------------- 1 file changed, 20 insertions(+), 20 deletions(-) diff --git a/LeanBandits/ETC.lean b/LeanBandits/ETC.lean index 6804c386..391671ea 100644 --- a/LeanBandits/ETC.lean +++ b/LeanBandits/ETC.lean @@ -47,20 +47,20 @@ namespace ETC variable {hK : 0 < K} {m : ℕ} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] -local notation "𝔓b" => Bandit.trajMeasure (etcAlgorithm hK m) ν +local notation "𝔓t" => Bandit.trajMeasure (etcAlgorithm hK m) ν local notation "𝔓" => Bandit.measure (etcAlgorithm hK m) ν -lemma arm_zero : arm 0 =ᵐ[𝔓b] fun _ ↦ ⟨0, hK⟩ := by +lemma arm_zero : arm 0 =ᵐ[𝔓t] fun _ ↦ ⟨0, hK⟩ := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK exact arm_zero_detAlgorithm lemma arm_ae_eq_etcNextArm (n : ℕ) : - arm (n + 1) =ᵐ[𝔓b] fun h ↦ nextArm hK m n (fun i ↦ h i) := by + arm (n + 1) =ᵐ[𝔓t] fun h ↦ nextArm hK m n (fun i ↦ h i) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK exact arm_detAlgorithm_ae_eq n lemma arm_of_lt {n : ℕ} (hn : n < K * m) : - arm n =ᵐ[𝔓b] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := by + arm n =ᵐ[𝔓t] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := by cases n with | zero => exact arm_zero | succ n => @@ -70,7 +70,7 @@ lemma arm_of_lt {n : ℕ} (hn : n < K * m) : lemma arm_mul (hm : m ≠ 0) : have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - arm (K * m) =ᵐ[𝔓b] fun h ↦ measurableArgmax (empMean' (K * m - 1)) (fun i ↦ h i) := by + arm (K * m) =ᵐ[𝔓t] fun h ↦ measurableArgmax (empMean' (K * m - 1)) (fun i ↦ h i) := by have : K * m = (K * m - 1) + 1 := by have : 0 < K * m := Nat.mul_pos hK hm.bot_lt grind @@ -80,25 +80,25 @@ lemma arm_mul (hm : m ≠ 0) : exact this ▸ rfl lemma arm_add_one_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : - arm (n + 1) =ᵐ[𝔓b] fun ω ↦ arm n ω := by + arm (n + 1) =ᵐ[𝔓t] fun ω ↦ arm n ω := by filter_upwards [arm_ae_eq_etcNextArm n] with ω hn_eq rw [hn_eq, nextArm, dif_neg (by grind), dif_neg] · rfl · have : 0 < K * m := Nat.mul_pos hK hm.bot_lt grind -lemma arm_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : arm n =ᵐ[𝔓b] arm (K * m) := by - have h_ae n : K * m ≤ n → arm (n + 1) =ᵐ[𝔓b] fun ω ↦ arm n ω := arm_add_one_of_ge hm +lemma arm_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : arm n =ᵐ[𝔓t] arm (K * m) := by + have h_ae n : K * m ≤ n → arm (n + 1) =ᵐ[𝔓t] fun ω ↦ arm n ω := arm_add_one_of_ge hm simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_ae filter_upwards [h_ae] with ω h_ae induction n, hn using Nat.le_induction with | base => rfl | succ n hmn h_ind => rw [h_ae n hmn, h_ind] -lemma pullCount_mul (a : Fin K) : pullCount a (K * m) =ᵐ[𝔓b] fun _ ↦ m := by +lemma pullCount_mul (a : Fin K) : pullCount a (K * m) =ᵐ[𝔓t] fun _ ↦ m := by rw [Filter.EventuallyEq] simp_rw [pullCount_eq_sum] - have h_arm (n : range (K * m)) : arm n =ᵐ[𝔓b] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := + have h_arm (n : range (K * m)) : arm n =ᵐ[𝔓t] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := arm_of_lt (mem_range.mp n.2) simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_arm filter_upwards [h_arm] with ω h_arm @@ -111,19 +111,19 @@ lemma pullCount_mul (a : Fin K) : pullCount a (K * m) =ᵐ[𝔓b] fun _ ↦ m := lemma pullCount_add_one_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : pullCount a (n + 1) - =ᵐ[𝔓b] fun ω ↦ pullCount a n ω + {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by + =ᵐ[𝔓t] fun ω ↦ pullCount a n ω + {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by simp_rw [Filter.EventuallyEq, pullCount_add_one] filter_upwards [arm_of_ge hm hn] with ω h_arm congr lemma pullCount_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : pullCount a n - =ᵐ[𝔓b] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by + =ᵐ[𝔓t] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by have h_ae n : K * m ≤ n → pullCount a (n + 1) - =ᵐ[𝔓b] fun ω ↦ pullCount a n ω + {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := + =ᵐ[𝔓t] fun ω ↦ pullCount a n ω + {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := pullCount_add_one_of_ge a hm simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_ae - have h_ae_Km : pullCount a (K * m) =ᵐ[𝔓b] fun _ ↦ m := pullCount_mul a + have h_ae_Km : pullCount a (K * m) =ᵐ[𝔓t] fun _ ↦ m := pullCount_mul a filter_upwards [h_ae_Km, h_ae] with ω h_Km h_ae induction n, hn using Nat.le_induction with | base => simp [h_Km] @@ -133,13 +133,13 @@ lemma pullCount_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : grind lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) : - (𝔓b).real {ω | arm (K * m) ω = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by + (𝔓t).real {ω | arm (K * m) ω = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK -- extend the probability space to include the stream of independent rewards suffices (𝔓).real {ω | arm (K * m) ω.1 = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) by - suffices (𝔓b).real {ω | arm (K * m) ω = a} = (𝔓).real {ω | arm (K * m) ω.1 = a} by + suffices (𝔓t).real {ω | arm (K * m) ω = a} = (𝔓).real {ω | arm (K * m) ω.1 = a} by rwa [this] - calc (𝔓b).real {ω | arm (K * m) ω = a} + calc (𝔓t).real {ω | arm (K * m) ω = a} _ = ((𝔓).fst).real {ω | arm (K * m) ω = a} := by simp _ = (𝔓).real {ω | arm (K * m) ω.1 = a} := by rw [Measure.fst, map_measureReal_apply (by fun_prop)] @@ -188,10 +188,10 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : - 𝔓b[fun ω ↦ (pullCount a n ω : ℝ)] + 𝔓t[fun ω ↦ (pullCount a n ω : ℝ)] ≤ m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by have : (fun ω ↦ (pullCount a n ω : ℝ)) - =ᵐ[𝔓b] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by + =ᵐ[𝔓t] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by filter_upwards [pullCount_of_ge a hm hn] with ω h simp only [h, Set.indicator_apply, Set.mem_setOf_eq, mul_ite, mul_one, mul_zero, Nat.cast_add, Nat.cast_ite, CharP.cast_eq_zero, add_right_inj] @@ -214,7 +214,7 @@ lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - ( lemma regret_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hm : m ≠ 0) (n : ℕ) (hn : K * m ≤ n) : - 𝔓b[regret ν n] ≤ ∑ a, gap ν a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4)) := by + 𝔓t[regret ν n] ≤ ∑ a, gap ν a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4)) := by simp_rw [regret_eq_sum_pullCount_mul_gap] rw [integral_finset_sum] swap From 0db7b411d29c78f3510750c25fa4056cd5a6cda1 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sat, 27 Sep 2025 14:44:24 +0200 Subject: [PATCH 07/11] add sumRewards def and lemmas --- LeanBandits/ETC.lean | 63 ++++++++++++++++++++++++++++++- LeanBandits/Regret.lean | 30 ++++++++++++--- blueprint/lean_decls | 2 +- blueprint/src/chapters/bandit.tex | 2 +- 4 files changed, 88 insertions(+), 9 deletions(-) diff --git a/LeanBandits/ETC.lean b/LeanBandits/ETC.lean index 391671ea..7f0c59d6 100644 --- a/LeanBandits/ETC.lean +++ b/LeanBandits/ETC.lean @@ -132,6 +132,66 @@ lemma pullCount_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : congr grind +lemma pullCount_add_one_eq_pullCount' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} : + pullCount a (n + 1) h = pullCount' n (fun i ↦ h i) a := by + rw [pullCount_eq_sum, pullCount'_eq_sum] + unfold arm + rw [Finset.sum_coe_sort (f := fun s ↦ if (h s).1 = a then 1 else 0) (Iic n)] + congr with m + simp only [mem_range, mem_Iic] + grind + +lemma pullCount_eq_pullCount' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} (hn : n ≠ 0) : + pullCount a n h = pullCount' (n - 1) (fun i ↦ h i) a := by + cases n with + | zero => exact absurd rfl hn + | succ n => + rw [pullCount_add_one_eq_pullCount'] + have : n + 1 - 1 = n := by simp + exact this ▸ rfl + +lemma sumRewards_add_one_eq_sumRewards' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} : + sumRewards a (n + 1) h = sumRewards' n (fun i ↦ h i) a := by + unfold sumRewards sumRewards' arm reward + rw [Finset.sum_coe_sort (f := fun s ↦ if (h s).1 = a then (h s).2 else 0) (Iic n)] + congr with m + simp only [mem_range, mem_Iic] + grind + +lemma sumRewards_eq_sumRewards' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} (hn : n ≠ 0) : + sumRewards a n h = sumRewards' (n - 1) (fun i ↦ h i) a := by + cases n with + | zero => exact absurd rfl hn + | succ n => + rw [sumRewards_add_one_eq_sumRewards'] + have : n + 1 - 1 = n := by simp + exact this ▸ rfl + +lemma empMean_add_one_eq_empMean' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} : + empMean a (n + 1) h = empMean' n (fun i ↦ h i) a := by + unfold empMean empMean' + rw [sumRewards_add_one_eq_sumRewards', pullCount_add_one_eq_pullCount'] + +lemma empMean_eq_empMean' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} (hn : n ≠ 0) : + empMean a n h = empMean' (n - 1) (fun i ↦ h i) a := by + unfold empMean empMean' + rw [sumRewards_eq_sumRewards' hn, pullCount_eq_pullCount' hn] + +lemma sumRewards_bestArm_le_of_arm_mul_eq (a : Fin K) (hm : m ≠ 0) : + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + ∀ᵐ h ∂𝔓t, arm (K * m) h = a → sumRewards (bestArm ν) (K * m) h ≤ sumRewards a (K * m) h := by + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + filter_upwards [arm_mul hm, pullCount_mul a, pullCount_mul (bestArm ν)] with h h_arm ha h_best + h_eq + have h_max := isMaxOn_measurableArgmax (empMean' (K * m - 1)) (fun i ↦ h i) (bestArm ν) + rw [← h_arm, h_eq] at h_max + rw [sumRewards_eq_pullCount_mul_empMean, sumRewards_eq_pullCount_mul_empMean, ha, h_best] + · gcongr + have : 0 < K * m := Nat.mul_pos hK hm.bot_lt + rwa [empMean_eq_empMean' this.ne', empMean_eq_empMean' this.ne'] + · simp [ha, hm] + · simp [h_best, hm] + lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) : (𝔓t).real {ω | arm (K * m) ω = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK @@ -146,8 +206,7 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i · rfl · exact (measurableSet_singleton _).preimage (by fun_prop) calc (𝔓).real {ω | arm (K * m) ω.1 = a} - _ ≤ (𝔓).real {ω | ∑ s ∈ range (K * m), (if (arm s ω.1) = bestArm ν then (reward s ω.1) else 0) - ≤ ∑ s ∈ range (K * m), if (arm s ω.1) = a then (reward s ω.1) else 0} := by + _ ≤ (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} := by sorry _ = (𝔓).real {ω | ∑ s ∈ Icc 1 (pullCount (bestArm ν) (K * m) ω.1), rewardByCount (bestArm ν) s ω.1 ω.2 diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index aaf1c1f1..cfde6e41 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -196,6 +196,23 @@ lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) pullCount a (stepsUntil a m h).toNat h = m - 1 := by sorry +section SumRewards + +/-- Sum of rewards obtained when pulling arm `a` up to time `t` (exclusive). -/ +def sumRewards (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := + ∑ s ∈ range t, if (arm s h) = a then (reward s h) else 0 + +/-- Empirical mean reward obtained when pulling arm `a` up to time `t` (exclusive). -/ +noncomputable +def empMean (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := sumRewards a t h / pullCount a t h + +lemma sumRewards_eq_pullCount_mul_empMean (h_pull : pullCount a t h ≠ 0) : + sumRewards a t h = pullCount a t h * empMean a t h := by unfold empMean; field_simp + +end SumRewards + +section RewardByCount + /-- Reward obtained when pulling arm `a` for the `m`-th time. -/ noncomputable def rewardByCount (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : ℝ := @@ -221,17 +238,18 @@ lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (h : ℕ → α × ℝ rewardByCount (arm t h) (pullCount (arm t h) t h + 1) h z = reward t h := by rw [rewardByCount, ← pullCount_eq_pullCount_add_one, stepsUntil_pullCount_eq] -lemma sum_rewardByCount_eq_sum_reward +lemma sum_rewardByCount_eq_sumRewards (a : α) (t : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : - ∑ m ∈ Icc 1 (pullCount a t h), rewardByCount a m h z = - ∑ s ∈ range t, if (arm s h) = a then (reward s h) else 0 := by + ∑ m ∈ Icc 1 (pullCount a t h), rewardByCount a m h z = sumRewards a t h := by induction' t with t ht - · simp [pullCount] + · simp [pullCount, sumRewards] by_cases hta : arm t h = a · rw [← hta] at ht ⊢ rw [pullCount_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht] + unfold sumRewards rw [sum_range_succ, if_pos rfl, rewardByCount_pullCount_add_one_eq_reward] - · rwa [pullCount_eq_pullCount hta, sum_range_succ, if_neg hta, add_zero] + · unfold sumRewards + rwa [pullCount_eq_pullCount hta, sum_range_succ, if_neg hta, add_zero] lemma sum_pullCount_mul [Fintype α] (h : ℕ → α × ℝ) (f : α → ℝ) (t : ℕ) : ∑ a, pullCount a t h * f a = ∑ s ∈ range t, f (arm s h) := by @@ -252,6 +270,8 @@ lemma regret_eq_sum_pullCount_mul_gap [Fintype α] : simp_rw [sum_pullCount_mul, regret, gap, sum_sub_distrib] simp +end RewardByCount + section BestArm variable [Fintype α] [Nonempty α] diff --git a/blueprint/lean_decls b/blueprint/lean_decls index fa693245..d545e1de 100644 --- a/blueprint/lean_decls +++ b/blueprint/lean_decls @@ -23,7 +23,7 @@ Bandits.iIndepFun_rewardByCount Bandits.stepsUntil_pullCount_le Bandits.stepsUntil_pullCount_eq Bandits.rewardByCount_pullCount_add_one_eq_reward -Bandits.sum_rewardByCount_eq_sum_reward +Bandits.sum_rewardByCount_eq_sumRewards Bandits.regret Bandits.gap Bandits.regret_eq_sum_pullCount_mul_gap diff --git a/blueprint/src/chapters/bandit.tex b/blueprint/src/chapters/bandit.tex index abb30068..f767409a 100644 --- a/blueprint/src/chapters/bandit.tex +++ b/blueprint/src/chapters/bandit.tex @@ -267,7 +267,7 @@ \section{Alternative model}\label{sec:alt_model} \begin{lemma}\label{lem:sum_rewardByCount} \uses{def:rewardByCount,def:pullCount} \leanok - \lean{Bandits.sum_rewardByCount_eq_sum_reward} + \lean{Bandits.sum_rewardByCount_eq_sumRewards} \begin{align*} \sum_{n=1}^{N_{t, a}} Y_{n, a} = \sum_{s=0}^{t-1} \mathbb{I}\{A_s = a\} X_s \: . From 24645a0bc324f6129c9b3371a5bf7a65671ac036 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sat, 27 Sep 2025 15:18:24 +0200 Subject: [PATCH 08/11] integrable pullCount --- LeanBandits/ETC.lean | 49 +++++++++++++++++---------- LeanBandits/Regret.lean | 3 ++ LeanBandits/RewardByCountMeasure.lean | 8 +++++ 3 files changed, 43 insertions(+), 17 deletions(-) diff --git a/LeanBandits/ETC.lean b/LeanBandits/ETC.lean index 7f0c59d6..c15a4d22 100644 --- a/LeanBandits/ETC.lean +++ b/LeanBandits/ETC.lean @@ -7,6 +7,7 @@ import Mathlib.Probability.Moments.SubGaussian import LeanBandits.AlgorithmBuilding import LeanBandits.ForMathlib.SubGaussian import LeanBandits.Regret +import LeanBandits.RewardByCountMeasure /-! # The Explore-Then-Commit Algorithm @@ -192,26 +193,36 @@ lemma sumRewards_bestArm_le_of_arm_mul_eq (a : Fin K) (hm : m ≠ 0) : · simp [ha, hm] · simp [h_best, hm] -lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) : +lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) + (hm : m ≠ 0) : (𝔓t).real {ω | arm (K * m) ω = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + have h_pos : 0 < K * m := Nat.mul_pos hK hm.bot_lt + have h_le : (𝔓t).real {ω | arm (K * m) ω = a} + ≤ (𝔓t).real {ω | sumRewards (bestArm ν) (K * m) ω ≤ sumRewards a (K * m) ω} := by + simp_rw [measureReal_def] + gcongr 1 + · simp + refine measure_mono_ae ?_ + exact sumRewards_bestArm_le_of_arm_mul_eq a hm + refine h_le.trans ?_ -- extend the probability space to include the stream of independent rewards - suffices (𝔓).real {ω | arm (K * m) ω.1 = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) by - suffices (𝔓t).real {ω | arm (K * m) ω = a} = (𝔓).real {ω | arm (K * m) ω.1 = a} by - rwa [this] - calc (𝔓t).real {ω | arm (K * m) ω = a} - _ = ((𝔓).fst).real {ω | arm (K * m) ω = a} := by simp - _ = (𝔓).real {ω | arm (K * m) ω.1 = a} := by + suffices (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} + ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) by + suffices (𝔓t).real {ω | sumRewards (bestArm ν) (K * m) ω ≤ sumRewards a (K * m) ω} + = (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} by rwa [this] + calc (𝔓t).real {ω | sumRewards (bestArm ν) (K * m) ω ≤ sumRewards a (K * m) ω} + _ = ((𝔓).fst).real {ω | sumRewards (bestArm ν) (K * m) ω ≤ sumRewards a (K * m) ω} := by simp + _ = (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} := by rw [Measure.fst, map_measureReal_apply (by fun_prop)] · rfl - · exact (measurableSet_singleton _).preimage (by fun_prop) - calc (𝔓).real {ω | arm (K * m) ω.1 = a} - _ ≤ (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} := by - sorry + · exact measurableSet_le (by fun_prop) (by fun_prop) + calc (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} _ = (𝔓).real {ω | ∑ s ∈ Icc 1 (pullCount (bestArm ν) (K * m) ω.1), rewardByCount (bestArm ν) s ω.1 ω.2 ≤ ∑ s ∈ Icc 1 (pullCount a (K * m) ω.1), rewardByCount a s ω.1 ω.2} := by - sorry + congr with ω + congr! 1 <;> rw [sum_rewardByCount_eq_sumRewards] _ = (𝔓).real {ω | ∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω.1 ω.2 ≤ ∑ s ∈ Icc 1 m, rewardByCount a s ω.1 ω.2} := by sorry @@ -242,7 +253,7 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i field_simp simp_rw [mul_assoc] simp only [NNReal.coe_ofNat, neg_inj, mul_eq_mul_left_iff, ne_eq, OfNat.ofNat_ne_zero, - not_false_eq_true, pow_eq_zero_iff, Nat.cast_eq_zero] + not_false_eq_true, pow_eq_zero_iff] norm_num lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) @@ -268,17 +279,21 @@ lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - ( simp rw [integral_indicator_const, smul_eq_mul, mul_one] · rw [← neg_mul] - exact prob_arm_mul_eq_le hν a + exact prob_arm_mul_eq_le hν a hm · exact (measurableSet_singleton _).preimage (by fun_prop) +lemma integrable_pullCount (a : Fin K) (n : ℕ) : Integrable (fun ω ↦ (pullCount a n ω : ℝ)) 𝔓t := by + refine integrable_of_le_of_le (g₁ := 0) (g₂ := fun _ ↦ n) (by fun_prop) + (ae_of_all _ fun ω ↦ by simp) (ae_of_all _ fun ω ↦ ?_) (integrable_const _) (integrable_const _) + simp only [Nat.cast_le] + exact pullCount_le a n ω + lemma regret_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hm : m ≠ 0) (n : ℕ) (hn : K * m ≤ n) : 𝔓t[regret ν n] ≤ ∑ a, gap ν a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4)) := by simp_rw [regret_eq_sum_pullCount_mul_gap] rw [integral_finset_sum] - swap - · refine fun i _ ↦ Integrable.mul_const ?_ _ - sorry + swap; · exact fun i _ ↦ (integrable_pullCount i n).mul_const _ gcongr with a rw [mul_comm (gap _ _), integral_mul_const] gcongr diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index cfde6e41..ac4b8cb7 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -70,6 +70,9 @@ lemma pullCount_add_one : lemma pullCount_eq_sum (a : α) (t : ℕ) (h : ℕ → α × ℝ) : pullCount a t h = ∑ s ∈ range t, if arm s h = a then 1 else 0 := by simp [pullCount] +lemma pullCount_le (a : α) (t : ℕ) (h : ℕ → α × ℝ) : pullCount a t h ≤ t := + (card_filter_le _ _).trans_eq (by simp) + /-- Number of steps until arm `a` was pulled exactly `m` times. -/ noncomputable def stepsUntil (a : α) (m : ℕ) (h : ℕ → α × ℝ) : ℕ∞ := sInf ((↑) '' {s | pullCount a (s + 1) h = m}) diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 8acb3b6c..2ca511ba 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -24,6 +24,14 @@ lemma measurable_pullCount (a : α) (t : ℕ) : Measurable (fun h ↦ pullCount exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop +@[fun_prop] +lemma measurable_sumRewards (a : α) (t : ℕ) : Measurable (sumRewards a t) := by + unfold sumRewards + have h_meas s : Measurable (fun h : ℕ → α × ℝ ↦ if arm s h = a then reward s h else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + fun_prop + @[fun_prop] lemma measurable_stepsUntil (a : α) (m : ℕ) : Measurable (fun h ↦ stepsUntil a m h) := by classical From 4182bafe6d0afc203033891c6848c38eb47a81ad Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sat, 27 Sep 2025 15:57:01 +0200 Subject: [PATCH 09/11] progress --- LeanBandits/ETC.lean | 31 ++++++++++++++++++++++++++----- 1 file changed, 26 insertions(+), 5 deletions(-) diff --git a/LeanBandits/ETC.lean b/LeanBandits/ETC.lean index c15a4d22..d1445638 100644 --- a/LeanBandits/ETC.lean +++ b/LeanBandits/ETC.lean @@ -193,6 +193,12 @@ lemma sumRewards_bestArm_le_of_arm_mul_eq (a : Fin K) (hm : m ≠ 0) : · simp [ha, hm] · simp [h_best, hm] +lemma ae_eq_set_iff {α : Type*} {mα : MeasurableSpace α} {μ : Measure α} {s t : Set α} : + s =ᵐ[μ] t ↔ ∀ᵐ a ∂μ, a ∈ s ↔ a ∈ t := by + rw [Filter.EventuallyEq] + simp only [eq_iff_iff] + congr! + lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) (hm : m ≠ 0) : (𝔓t).real {ω | arm (K * m) ω = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by @@ -225,7 +231,19 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i congr! 1 <;> rw [sum_rewardByCount_eq_sumRewards] _ = (𝔓).real {ω | ∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω.1 ω.2 ≤ ∑ s ∈ Icc 1 m, rewardByCount a s ω.1 ω.2} := by - sorry + simp_rw [measureReal_def] + congr 1 + refine measure_congr ?_ + have ha := pullCount_mul a (hK := hK) (ν := ν) (m := m) + have h_best := pullCount_mul (bestArm ν) (hK := hK) (ν := ν) (m := m) + rw [ae_eq_set_iff] + change ∀ᵐ ω ∂((𝔓t).prod _), _ + rw [Measure.ae_prod_iff_ae_ae] + · filter_upwards [ha, h_best] with ω ha h_best + refine ae_of_all _ fun ω' ↦ ?_ + rw [ha, h_best] + · simp only [Set.mem_setOf_eq] + sorry _ = (𝔓).real {ω | ∑ s ∈ range m, ω.2 s (bestArm ν) ≤ ∑ s ∈ range m, ω.2 s a} := by sorry _ = (𝔓).real {ω | m * gap ν a @@ -234,13 +252,16 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i simp only [gap_eq_bestArm_sub, id_eq, sum_sub_distrib, sum_const, card_range, nsmul_eq_mul] ring_nf simp + _ = (Bandit.streamMeasure ν).real {ω | m * gap ν a + ≤ ∑ s ∈ range m, ((ω s a - (ν a)[id]) - (ω s (bestArm ν) - (ν (bestArm ν))[id]))} := by + have : Bandit.streamMeasure ν = (𝔓).map Prod.snd := by rw [← Measure.snd, Bandit.snd_measure] + rw [this, measureReal_def, measureReal_def, Measure.map_apply (by fun_prop)] + · rfl + · exact measurableSet_le (by fun_prop) (by fun_prop) _ ≤ Real.exp (-↑m * gap ν a ^ 2 / 4) := by refine (HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := 2) (ε := m * gap ν a) ?_ ?_ ?_).trans_eq ?_ - · suffices iIndepFun (fun s ω ↦ ω s a - (ν a)[id] - (ω s (bestArm ν) - (ν (bestArm ν))[id])) - (Bandit.streamMeasure ν) by - sorry - sorry + · sorry · intro i him rw [← one_add_one_eq_two] refine HasSubgaussianMGF.sub_of_indepFun ?_ ?_ ?_ From 4c1b493c5d359faf7c10383075f7533e7598b6e3 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sun, 28 Sep 2025 17:01:12 +0200 Subject: [PATCH 10/11] add subGaussian lemmas --- LeanBandits/Bandit.lean | 27 +++++++++++++++++++ LeanBandits/ETC.lean | 8 ++++-- LeanBandits/ForMathlib/SubGaussian.lean | 35 +++++++++++++++++++++++++ 3 files changed, 68 insertions(+), 2 deletions(-) diff --git a/LeanBandits/Bandit.lean b/LeanBandits/Bandit.lean index 18c12767..e27198f3 100644 --- a/LeanBandits/Bandit.lean +++ b/LeanBandits/Bandit.lean @@ -74,6 +74,33 @@ lemma snd_measure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] end Bandit +section StreamMeasure + +lemma _root_.hasLaw_eval_infinitePi {ι : Type*} {X : ι → Type*} {mX : ∀ i, MeasurableSpace (X i)} + (μ : (i : ι) → Measure (X i)) [hμ : ∀ i, IsProbabilityMeasure (μ i)] (i : ι) : + HasLaw (Function.eval i) (μ i) (Measure.infinitePi μ) where + aemeasurable := Measurable.aemeasurable (by fun_prop) + map_eq := by exact (measurePreserving_eval_infinitePi μ i).map_eq + +lemma hasLaw_eval_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + HasLaw (fun h : ℕ → α → R ↦ h n) (Measure.infinitePi ν) (Bandit.streamMeasure ν) := + hasLaw_eval_infinitePi (fun _ ↦ Measure.infinitePi ν) n + +lemma hasLaw_eval_eval_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) (a : α) : + HasLaw (fun h : ℕ → α → R ↦ h n a) (ν a) (Bandit.streamMeasure ν) := + (hasLaw_eval_infinitePi ν a).comp (hasLaw_eval_streamMeasure ν n) + +lemma identDistrib_eval_eval_id_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) (a : α) : + IdentDistrib (fun h : ℕ → α → R ↦ h n a) id (Bandit.streamMeasure ν) (ν a) where + aemeasurable_fst := Measurable.aemeasurable (by fun_prop) + aemeasurable_snd := Measurable.aemeasurable (by fun_prop) + map_eq := by + rw [← (hasLaw_eval_eval_streamMeasure ν n a).map_eq, + Measure.map_map (by fun_prop) (by fun_prop)] + simp + +end StreamMeasure + /-- `arm n` is the arm pulled at time `n`. This is a random variable on the measurable space `ℕ → α × ℝ`. -/ def arm (n : ℕ) (h : ℕ → α × R) : α := (h n).1 diff --git a/LeanBandits/ETC.lean b/LeanBandits/ETC.lean index d1445638..32f5b4ac 100644 --- a/LeanBandits/ETC.lean +++ b/LeanBandits/ETC.lean @@ -265,8 +265,12 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i · intro i him rw [← one_add_one_eq_two] refine HasSubgaussianMGF.sub_of_indepFun ?_ ?_ ?_ - · sorry - · sorry + · refine (hν a).congr_identDistrib ?_ + refine IdentDistrib.sub_const ?_ _ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm + · refine (hν (bestArm ν)).congr_identDistrib ?_ + refine IdentDistrib.sub_const ?_ _ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm · sorry · have : 0 ≤ gap ν a := gap_nonneg positivity diff --git a/LeanBandits/ForMathlib/SubGaussian.lean b/LeanBandits/ForMathlib/SubGaussian.lean index 5291541a..ca7e09e3 100644 --- a/LeanBandits/ForMathlib/SubGaussian.lean +++ b/LeanBandits/ForMathlib/SubGaussian.lean @@ -20,6 +20,21 @@ namespace Kernel.HasSubgaussianMGF variable {Ω Ω' : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} {ν : Measure Ω'} {κ : Kernel Ω' Ω} {X : Ω → ℝ} {c : ℝ≥0} +lemma id_map_iff (hX : Measurable X) : + HasSubgaussianMGF X c κ ν ↔ HasSubgaussianMGF id c (κ.map X) ν := by + refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ + · constructor + · intro t + rw [← Kernel.deterministic_comp_eq_map hX, ← Measure.comp_assoc, + Measure.deterministic_comp_eq_map] + rw [integrable_map_measure (by fun_prop) hX.aemeasurable] + exact h.integrable_exp_mul t + · simp_rw [Kernel.map_apply _ hX, mgf_id_map hX.aemeasurable] + exact h.mgf_le + · have : X = id ∘ X := rfl + rw [this] + exact .of_map hX h + protected lemma const_mul (h : HasSubgaussianMGF X c κ ν) (r : ℝ) : HasSubgaussianMGF (fun ω ↦ r * X ω) (⟨r ^ 2, sq_nonneg r⟩ * c) κ ν where integrable_exp_mul t := by @@ -40,6 +55,26 @@ namespace HasSubgaussianMGF variable {Ω : Type*} {m mΩ : MeasurableSpace Ω} {μ : Measure Ω} {X : Ω → ℝ} {c : ℝ≥0} +lemma id_map_iff (hX : AEMeasurable X μ) : + HasSubgaussianMGF X c μ ↔ HasSubgaussianMGF id c (μ.map X) := by + refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ + · constructor + · intro t + rw [integrable_map_measure (by fun_prop) hX] + exact h.integrable_exp_mul t + · intro t + rw [mgf_id_map hX] + exact h.mgf_le t + · have : X = id ∘ X := rfl + rw [this] + exact .of_map hX h + +lemma congr_identDistrib {Ω' : Type*} {mΩ' : MeasurableSpace Ω'} {μ' : Measure Ω'} + {Y : Ω' → ℝ} (hX : HasSubgaussianMGF X c μ) (hXY : IdentDistrib X Y μ μ') : + HasSubgaussianMGF Y c μ' := by + rw [id_map_iff hXY.aemeasurable_fst] at hX + rwa [id_map_iff hXY.aemeasurable_snd, ← hXY.map_eq] + protected lemma const_mul (h : HasSubgaussianMGF X c μ) (r : ℝ) : HasSubgaussianMGF (fun ω ↦ r * X ω) (⟨r ^ 2, sq_nonneg r⟩ * c) μ := by rw [HasSubgaussianMGF_iff_kernel] at h ⊢ From f7df0f26bc1529277fc337b91209f81b13bda9cb Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sun, 28 Sep 2025 17:07:49 +0200 Subject: [PATCH 11/11] work --- LeanBandits/ETC.lean | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/LeanBandits/ETC.lean b/LeanBandits/ETC.lean index 32f5b4ac..46f4b5b7 100644 --- a/LeanBandits/ETC.lean +++ b/LeanBandits/ETC.lean @@ -261,17 +261,20 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i _ ≤ Real.exp (-↑m * gap ν a ^ 2 / 4) := by refine (HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := 2) (ε := m * gap ν a) ?_ ?_ ?_).trans_eq ?_ - · sorry + · suffices iIndepFun (fun s ω ↦ ω s a - ω s (bestArm ν)) (Bandit.streamMeasure ν) by + sorry + sorry · intro i him rw [← one_add_one_eq_two] refine HasSubgaussianMGF.sub_of_indepFun ?_ ?_ ?_ · refine (hν a).congr_identDistrib ?_ - refine IdentDistrib.sub_const ?_ _ - exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ · refine (hν (bestArm ν)).congr_identDistrib ?_ - refine IdentDistrib.sub_const ?_ _ - exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm - · sorry + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ + · suffices IndepFun (fun ω ↦ ω i a) (fun ω ↦ ω i (bestArm ν)) (Bandit.streamMeasure ν) by + exact this.comp (φ := fun x ↦ x - (ν a)[id]) (ψ := fun x ↦ x - (ν (bestArm ν))[id]) + (by fun_prop) (by fun_prop) + sorry · have : 0 ≤ gap ν a := gap_nonneg positivity · congr 1