diff --git a/LeanBandits.lean b/LeanBandits.lean index 86dda7b8..f2c8344e 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -16,3 +16,4 @@ import LeanBandits.ForMathlib.Traj import LeanBandits.RewardByCountMeasure import LeanBandits.SequentialLearning.Algorithm import LeanBandits.SequentialLearning.Deterministic +import LeanBandits.SequentialLearning.StationaryEnv diff --git a/LeanBandits/AlgorithmBuilding.lean b/LeanBandits/AlgorithmBuilding.lean index 1e1df9f5..971808dd 100644 --- a/LeanBandits/AlgorithmBuilding.lean +++ b/LeanBandits/AlgorithmBuilding.lean @@ -19,7 +19,8 @@ namespace Bandits variable {α : Type*} [DecidableEq α] [MeasurableSpace α] -/-- Number of pulls of arm `a` up to (and including) time `n`. -/ +/-- Number of pulls of arm `a` up to (and including) time `n`. +This is the number of entries in `h` in which the arm is `a`. -/ noncomputable def pullCount' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := #{s | (h s).1 = a} diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 0b88c9bc..2de56bb2 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -4,6 +4,7 @@ Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne, Paulo Rauber -/ import LeanBandits.SequentialLearning.Deterministic +import LeanBandits.SequentialLearning.StationaryEnv import LeanBandits.ForMathlib.IndepInfinitePi import Mathlib.Probability.IdentDistrib diff --git a/LeanBandits/Bandit/Regret.lean b/LeanBandits/Bandit/Regret.lean index 70ac3c2c..905c82e0 100644 --- a/LeanBandits/Bandit/Regret.lean +++ b/LeanBandits/Bandit/Regret.lean @@ -6,6 +6,7 @@ Authors: Rémy Degenne, Paulo Rauber import LeanBandits.Bandit.Bandit import Mathlib.Data.ENat.Lattice import Mathlib.Order.CompletePartialOrder +import Mathlib.Probability.Martingale.BorelCantelli /-! # Regret @@ -55,6 +56,11 @@ open Classical in lemma monotone_pullCount (a : α) (h : ℕ → α × ℝ) : Monotone (pullCount a · h) := fun _ _ _ ↦ card_le_card (filter_subset_filter _ (by simpa)) +@[mono, gcongr] +lemma pullCount_mono (a : α) {n m : ℕ} (hnm : n ≤ m) (h : ℕ → α × ℝ) : + pullCount a n h ≤ pullCount a m h := + monotone_pullCount a h hnm + lemma pullCount_eq_pullCount_add_one (t : ℕ) (h : ℕ → α × ℝ) : pullCount (arm t h) (t + 1) h = pullCount (arm t h) t h + 1 := by simp [pullCount, range_add_one, filter_insert] @@ -83,6 +89,7 @@ lemma pullCount_congr {h' : ℕ → α × ℝ} (h_eq : ∀ i ≤ n, arm i h = ar rw [Nat.lt_add_one_iff] at hs rw [h_eq s hs] +-- TODO: replace this by leastGE? /-- 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}) @@ -93,6 +100,10 @@ lemma stepsUntil_eq_top_iff : stepsUntil a m h = ⊤ ↔ ∀ s, pullCount a (s + lemma stepsUntil_ne_top (h_exists : ∃ s, pullCount a (s + 1) h = m) : stepsUntil a m h ≠ ⊤ := by simpa [stepsUntil_eq_top_iff] +lemma stepsUntil_eq_leastGE (a : α) (m : ℕ) : + stepsUntil a m = leastGE (fun n h ↦ pullCount a (n + 1) h) m := by + sorry + lemma exists_pullCount_eq (h' : stepsUntil a m h ≠ ⊤) : ∃ s, pullCount a (s + 1) h = m := by by_contra! h_contra @@ -283,7 +294,7 @@ 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 + ∑ 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 @@ -296,38 +307,39 @@ end SumRewards section RewardByCount -/-- Reward obtained when pulling arm `a` for the `m`-th time. -/ +/-- Reward obtained when pulling arm `a` for the `m`-th time. +If it is never pulled `m` times, the reward is given by the second component of `ω`, which in +applications will be indepedent with same law. -/ noncomputable -def rewardByCount (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : ℝ := - match (stepsUntil a m h) with - | ⊤ => z m a - | (n : ℕ) => reward n h - -lemma rewardByCount_eq_ite (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : - rewardByCount a m h z = - if (stepsUntil a m h) = ⊤ then z m a else reward (stepsUntil a m h).toNat h := by +def rewardByCount (a : α) (m : ℕ) (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) : ℝ := + match (stepsUntil a m ω.1) with + | ⊤ => ω.2 m a + | (n : ℕ) => reward n ω.1 + +lemma rewardByCount_eq_ite (a : α) (m : ℕ) (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) : + rewardByCount a m ω = + if (stepsUntil a m ω.1) = ⊤ then ω.2 m a else reward (stepsUntil a m ω.1).toNat ω.1 := by unfold rewardByCount - cases stepsUntil a m h <;> simp + cases stepsUntil a m ω.1 <;> simp lemma rewardByCount_of_stepsUntil_eq_top {ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)} (h : stepsUntil a m ω.1 = ⊤) : - rewardByCount a m ω.1 ω.2 = ω.2 m a := by simp [rewardByCount_eq_ite, h] + rewardByCount a m ω = ω.2 m a := by simp [rewardByCount_eq_ite, h] lemma rewardByCount_of_stepsUntil_eq_coe {ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)} (h : stepsUntil a m ω.1 = n) : - rewardByCount a m ω.1 ω.2 = reward n ω.1 := by simp [rewardByCount_eq_ite, h] + rewardByCount a m ω = reward n ω.1 := by simp [rewardByCount_eq_ite, h] -lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : - rewardByCount (arm t h) (pullCount (arm t h) t h + 1) h z = reward t h := by +lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) : + rewardByCount (arm t ω.1) (pullCount (arm t ω.1) t ω.1 + 1) ω = reward t ω.1 := by rw [rewardByCount, ← pullCount_eq_pullCount_add_one, stepsUntil_pullCount_eq] -lemma sum_rewardByCount_eq_sumRewards - (a : α) (t : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : - ∑ m ∈ Icc 1 (pullCount a t h), rewardByCount a m h z = sumRewards a t h := by +lemma sum_rewardByCount_eq_sumRewards (a : α) (t : ℕ) (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) : + ∑ m ∈ Icc 1 (pullCount a t ω.1), rewardByCount a m ω = sumRewards a t ω.1 := by induction t with | zero => simp [pullCount, sumRewards] | succ t ht => - by_cases hta : arm t h = a + by_cases hta : arm t ω.1 = a · rw [← hta] at ht ⊢ rw [pullCount_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht] unfold sumRewards diff --git a/LeanBandits/BanditAlgorithms/ETC.lean b/LeanBandits/BanditAlgorithms/ETC.lean index d3a70c67..a5b06ae9 100644 --- a/LeanBandits/BanditAlgorithms/ETC.lean +++ b/LeanBandits/BanditAlgorithms/ETC.lean @@ -15,6 +15,8 @@ import LeanBandits.RewardByCountMeasure open MeasureTheory ProbabilityTheory Finset Learning open scoped ENNReal NNReal +section Aux + lemma ae_eq_set_iff {α : Type*} {mα : MeasurableSpace α} {μ : Measure α} {s t : Set α} : s =ᵐ[μ] t ↔ ∀ᵐ a ∂μ, a ∈ s ↔ a ∈ t := by rw [Filter.EventuallyEq] @@ -36,11 +38,55 @@ lemma measurable_sum_of_le {α : Type*} {mα : MeasurableSpace α} refine Measurable.ite ?_ (by fun_prop) (by fun_prop) exact (measurableSet_singleton _).preimage (by fun_prop) +lemma sum_mod_range {K : ℕ} (hK : 0 < K) (a : Fin K) : + (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) = 1 := by + have h_iff (s : ℕ) (hs : s < K) : ⟨s % K, Nat.mod_lt _ hK⟩ = a ↔ s = a := by + simp only [Nat.mod_eq_of_lt hs, Fin.ext_iff] + calc (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) + _ = ∑ s ∈ range K, if s = a then 1 else 0 := sum_congr rfl fun s hs ↦ by grind + _ = _ := by + rw [sum_ite_eq'] + simp + +lemma sum_mod_range_mul {K : ℕ} (hK : 0 < K) (m : ℕ) (a : Fin K) : + (∑ s ∈ range (K * m), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) = m := by + induction m with + | zero => simp + | succ n hn => + calc (∑ s ∈ range (K * (n + 1)), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) + _ = (∑ s ∈ range (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by ring_nf + _ = (∑ s ∈ range (K * n), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) + + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by + rw [sum_range_add_sum_Ico] + grind + _ = n + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by + rw [hn] + _ = n + (∑ s ∈ range K, if ⟨(s + K * n) % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by + congr 1 + let e : ℕ ↪ ℕ := ⟨fun i : ℕ ↦ i + K * n, fun i j hij ↦ by grind⟩ + have : Finset.map e (range K) = Ico (K * n) (K * n + K) := by + ext x + simp only [mem_map, mem_range, Function.Embedding.coeFn_mk, mem_Ico, e] + refine ⟨fun h ↦ by grind, fun h ↦ ?_⟩ + use x - K * n + grind + rw [← this, Finset.sum_map] + congr + _ = n + (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by simp + _ = n + 1 := by rw [sum_mod_range hK] + +end Aux + namespace Bandits variable {K : ℕ} -/-- Arm pulled by the ETC algorithm at time `n + 1`. -/ +section AlgorithmDefinition + +/-- Arm pulled by the ETC algorithm at time `n + 1`. +For `n < K * m - 1`, this is arm `n % K`. +For `n = K * m - 1`, this is the arm with the highest empirical mean after the exploration phase. +For `n ≥ K * m`, this is the same arm as at time `n`. -/ noncomputable def ETC.nextArm (hK : 0 < K) (m n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK @@ -50,6 +96,7 @@ def ETC.nextArm (hK : 0 < K) (m n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K := if hn_eq : n = K * m - 1 then measurableArgmax (empMean' n) h else (h ⟨n, by simp⟩).1 +/-- The next arm pulled by ETC is chosen in a measurable way. -/ @[fun_prop] lemma ETC.measurable_nextArm (hK : 0 < K) (m n : ℕ) : Measurable (nextArm hK m n) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK @@ -59,11 +106,14 @@ lemma ETC.measurable_nextArm (hK : 0 < K) (m n : ℕ) : Measurable (nextArm hK m refine Measurable.ite (by simp) ?_ (by fun_prop) exact measurable_measurableArgmax fun a ↦ by fun_prop -/-- The Explore-Then-Commit algorithm. -/ +/-- The Explore-Then-Commit algorithm: deterministic algorithm that chooses the next arm according +to `ETC.nextArm`. -/ noncomputable def etcAlgorithm (hK : 0 < K) (m : ℕ) : Algorithm (Fin K) ℝ := detAlgorithm (ETC.nextArm hK m) (by fun_prop) ⟨0, hK⟩ +end AlgorithmDefinition + namespace ETC variable {hK : 0 < K} {m : ℕ} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] @@ -80,6 +130,7 @@ lemma arm_ae_eq_etcNextArm (n : ℕ) : have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK exact arm_detAlgorithm_ae_eq n +/-- For `n < K * m`, the arm pulled at time `n` is the arm `n % K`. -/ lemma arm_of_lt {n : ℕ} (hn : n < K * m) : arm n =ᵐ[𝔓t] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := by cases n with @@ -89,6 +140,8 @@ lemma arm_of_lt {n : ℕ} (hn : n < K * m) : rw [hn_eq, nextArm, dif_pos] grind +/-- The arm pulled at time `K * m` is the arm with the highest empirical mean after the exploration +phase. -/ lemma arm_mul (hm : m ≠ 0) : have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK arm (K * m) =ᵐ[𝔓t] fun h ↦ measurableArgmax (empMean' (K * m - 1)) (fun i ↦ h i) := by @@ -100,6 +153,7 @@ lemma arm_mul (hm : m ≠ 0) : rw [hn_eq, nextArm, dif_neg (by simp), dif_pos rfl] exact this ▸ rfl +/-- For `n ≥ K * m`, the arm pulled at time `n + 1` is the same as the arm pulled at time `n`. -/ lemma arm_add_one_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : arm (n + 1) =ᵐ[𝔓t] fun ω ↦ arm n ω := by filter_upwards [arm_ae_eq_etcNextArm n] with ω hn_eq @@ -108,7 +162,9 @@ lemma arm_add_one_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : · 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 =ᵐ[𝔓t] arm (K * m) := by +/-- For `n ≥ K * m`, the arm pulled at time `n` is the same as the arm pulled at time `K * m`. -/ +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 @@ -116,43 +172,7 @@ lemma arm_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : arm n =ᵐ[𝔓t] | base => rfl | succ n hmn h_ind => rw [h_ae n hmn, h_ind] -lemma sum_mod_range {K : ℕ} (hK : 0 < K) (a : Fin K) : - (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) = 1 := by - have h_iff (s : ℕ) (hs : s < K) : ⟨s % K, Nat.mod_lt _ hK⟩ = a ↔ s = a := by - simp only [Nat.mod_eq_of_lt hs, Fin.ext_iff] - calc (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) - _ = ∑ s ∈ range K, if s = a then 1 else 0 := sum_congr rfl fun s hs ↦ by grind - _ = _ := by - rw [sum_ite_eq'] - simp - -lemma sum_mod_range_mul {K : ℕ} (hK : 0 < K) (m : ℕ) (a : Fin K) : - (∑ s ∈ range (K * m), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) = m := by - induction m with - | zero => simp - | succ n hn => - calc (∑ s ∈ range (K * (n + 1)), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) - _ = (∑ s ∈ range (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by ring_nf - _ = (∑ s ∈ range (K * n), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) - + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by - rw [sum_range_add_sum_Ico] - grind - _ = n + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by - rw [hn] - _ = n + (∑ s ∈ range K, if ⟨(s + K * n) % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by - congr 1 - let e : ℕ ↪ ℕ := ⟨fun i : ℕ ↦ i + K * n, fun i j hij ↦ by grind⟩ - have : Finset.map e (range K) = Ico (K * n) (K * n + K) := by - ext x - simp only [mem_map, mem_range, Function.Embedding.coeFn_mk, mem_Ico, e] - refine ⟨fun h ↦ by grind, fun h ↦ ?_⟩ - use x - K * n - grind - rw [← this, Finset.sum_map] - congr - _ = n + (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by simp - _ = n + 1 := by rw [sum_mod_range hK] - +/-- At time `K * m`, the number of pulls of each arm is equal to `m`. -/ lemma pullCount_mul (a : Fin K) : pullCount a (K * m) =ᵐ[𝔓t] fun _ ↦ m := by rw [Filter.EventuallyEq] simp_rw [pullCount_eq_sum] @@ -173,6 +193,8 @@ lemma pullCount_add_one_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m filter_upwards [arm_of_ge hm hn] with ω h_arm congr +/-- For `n ≥ K * m`, the number of pulls of each arm `a` at time `n` is equal to `m` plus +`n - K * m` if arm `a` is the best arm after the exploration phase. -/ lemma pullCount_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : pullCount a n =ᵐ[𝔓t] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by @@ -189,6 +211,8 @@ lemma pullCount_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : congr grind +/-- If at time `K * m` the algorithm chooses arm `a`, then the total reward obtained by pulling +arm `a` is at least the total reward obtained by pulling the best arm. -/ 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 @@ -206,37 +230,16 @@ lemma sumRewards_bestArm_le_of_arm_mul_eq (a : Fin K) (hm : m ≠ 0) : lemma identDistrib_aux (m : ℕ) (a b : Fin K) : IdentDistrib - (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount a s ω.1 ω.2, ∑ s ∈ Icc 1 m, rewardByCount b s ω.1 ω.2)) + (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount a s ω, ∑ s ∈ Icc 1 m, rewardByCount b s ω)) (fun ω ↦ (∑ s ∈ range m, ω.2 s a, ∑ s ∈ range m, ω.2 s b)) 𝔓 𝔓 := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - have h1 (a : Fin K) : - IdentDistrib (fun ω s ↦ rewardByCount a (s + 1) ω.1 ω.2) (fun ω s ↦ ω.2 s a) 𝔓 𝔓 := - identDistrib_rewardByCount_stream a - have h2 (a : Fin K) : IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount a s ω.1 ω.2) - (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓 𝔓 := by - have h_eq (ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ)) : ∑ s ∈ Icc 1 m, rewardByCount a s ω.1 ω.2 - = ∑ s ∈ range m, rewardByCount a (s + 1) ω.1 ω.2 := by - let e : Icc 1 m ≃ range m := - { toFun x := ⟨x - 1, by have h := x.2; simp only [mem_Icc] at h; simp; grind⟩ - invFun x := ⟨x + 1, by - have h := x.2 - simp only [mem_Icc, le_add_iff_nonneg_left, zero_le, true_and, ge_iff_le] - simp only [mem_range] at h - grind⟩ - left_inv x := by have h := x.2; simp only [mem_Icc] at h; grind - right_inv x := by have h := x.2; grind } - rw [← sum_coe_sort (Icc 1 m), ← sum_coe_sort (range m), sum_equiv e] - · simp - · simp only [univ_eq_attach, mem_attach, forall_const, Subtype.forall, mem_Icc, - forall_and_index] - grind - simp_rw [h_eq] - exact IdentDistrib.comp (h1 a) (u := fun p ↦ ∑ s ∈ range m, p s) (by fun_prop) + have h2 (a : Fin K) : IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount a s ω) + (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓 𝔓 := identDistrib_sum_Icc_rewardByCount m a by_cases hab : a = b · simp only [hab] exact (h2 b).comp (u := fun p ↦ (p, p)) (by fun_prop) refine (h2 a).prod (h2 b) ?_ ?_ - · suffices IndepFun (fun ω s ↦ rewardByCount a s ω.1 ω.2) (fun ω s ↦ rewardByCount b s ω.1 ω.2) + · suffices IndepFun (fun ω s ↦ rewardByCount a s ω) (fun ω s ↦ rewardByCount b s ω) 𝔓 by exact this.comp (φ := fun p ↦ ∑ i ∈ Icc 1 m, p i) (ψ := fun p ↦ ∑ j ∈ Icc 1 m, p j) (by fun_prop) (by fun_prop) @@ -246,6 +249,8 @@ lemma identDistrib_aux (m : ℕ) (a b : Fin K) : (by fun_prop) (by fun_prop) exact indepFun_eval_snd_measure _ ν hab +/-- The probability that at time `K * m` the ETC algorithm chooses arm `a` is at most +`exp(- m * Δ_a^2 / 4)`. -/ 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 @@ -271,13 +276,12 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i · rfl · 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 + _ = (𝔓).real {ω | ∑ s ∈ Icc 1 (pullCount (bestArm ν) (K * m) ω.1), rewardByCount (bestArm ν) s ω + ≤ ∑ s ∈ Icc 1 (pullCount a (K * m) ω.1), rewardByCount a s ω} := by 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 + _ = (𝔓).real {ω | ∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω + ≤ ∑ s ∈ Icc 1 m, rewardByCount a s ω} := by simp_rw [measureReal_def] congr 1 refine measure_congr ?_ @@ -291,24 +295,24 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i rw [ha, h_best] · simp only [Set.mem_setOf_eq] let f₁ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ - ∑ s ∈ Icc 1 (pullCount (bestArm ν) (K * m) ω.1), rewardByCount (bestArm ν) s ω.1 ω.2 + ∑ s ∈ Icc 1 (pullCount (bestArm ν) (K * m) ω.1), rewardByCount (bestArm ν) s ω let g₁ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ - ∑ s ∈ Icc 1 (pullCount a (K * m) ω.1), rewardByCount a s ω.1 ω.2 + ∑ s ∈ Icc 1 (pullCount a (K * m) ω.1), rewardByCount a s ω let f₂ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ - ∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω.1 ω.2 - let g₂ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ ∑ s ∈ Icc 1 m, rewardByCount a s ω.1 ω.2 + ∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω + let g₂ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ ∑ s ∈ Icc 1 m, rewardByCount a s ω change MeasurableSet {x | f₁ x ≤ g₁ x ↔ f₂ x ≤ g₂ x} have hf₁ : Measurable f₁ := by refine measurable_sum_of_le (n := K * m + 1) (g := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ pullCount (bestArm ν) (K * m) ω.1) - (f := fun s ω ↦ rewardByCount (bestArm ν) s ω.1 ω.2) (fun ω ↦ ?_) + (f := rewardByCount (bestArm ν)) (fun ω ↦ ?_) (by fun_prop) (by fun_prop) have h_le := pullCount_le (bestArm ν) (K * m) ω.1 grind have hg₁ : Measurable g₁ := by refine measurable_sum_of_le (n := K * m + 1) (g := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ pullCount a (K * m) ω.1) - (f := fun s ω ↦ rewardByCount a s ω.1 ω.2) (fun ω ↦ ?_) (by fun_prop) (by fun_prop) + (f := rewardByCount a) (fun ω ↦ ?_) (by fun_prop) (by fun_prop) have h_le := pullCount_le a (K * m) ω.1 grind refine MeasurableSet.iff ?_ ?_ @@ -317,8 +321,8 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i _ = (𝔓).real {ω | ∑ s ∈ range m, ω.2 s (bestArm ν) ≤ ∑ s ∈ range m, ω.2 s a} := by simp_rw [measureReal_def] congr 1 - have : (𝔓).map (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω.1 ω.2, - ∑ s ∈ Icc 1 m, rewardByCount a s ω.1 ω.2)) + have : (𝔓).map (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω, + ∑ s ∈ Icc 1 m, rewardByCount a s ω)) = (𝔓).map (fun ω ↦ (∑ s ∈ range m, ω.2 s (bestArm ν), ∑ s ∈ range m, ω.2 s a)) := (identDistrib_aux m (bestArm ν) a).map_eq rw [Measure.ext_iff] at this @@ -336,6 +340,7 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i _ ≤ Real.exp (-↑m * gap ν a ^ 2 / 4) := by by_cases ha : a = bestArm ν · simp [ha] + -- Apply a sub-Gaussian concentration inequality refine (HasSubgaussianMGF.measure_sum_le_sum_le' (cX := fun _ ↦ 1) (cY := fun _ ↦ 1) ?_ ?_ ?_ ?_ ?_ ?_).trans_eq ?_ · exact iIndepFun_eval_streamMeasure'' ν (bestArm ν) @@ -359,6 +364,7 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i field_simp ring +/-- Bound on the expectation of the number of pulls of each arm by the ETC algorithm. -/ 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) : 𝔓t[fun ω ↦ (pullCount a n ω : ℝ)] @@ -384,12 +390,7 @@ lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - ( 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 ω - +/-- Regret bound for the ETC algorithm. -/ 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 diff --git a/LeanBandits/BanditAlgorithms/UCB.lean b/LeanBandits/BanditAlgorithms/UCB.lean index 2db28bc2..cbfa1e46 100644 --- a/LeanBandits/BanditAlgorithms/UCB.lean +++ b/LeanBandits/BanditAlgorithms/UCB.lean @@ -3,9 +3,11 @@ 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 LeanBandits.AlgorithmBuilding -import LeanBandits.Bandit.Regret +import LeanBandits.AlgorithmAndRandomVariables import LeanBandits.ForMathlib.MeasurableArgMax +import LeanBandits.ForMathlib.SubGaussian +import LeanBandits.RewardByCountMeasure +import LeanBandits.BanditAlgorithms.ETC /-! # UCB algorithm @@ -18,73 +20,702 @@ open scoped ENNReal NNReal namespace Bandits -variable {α : Type*} {mα : MeasurableSpace α} {ν : Kernel α ℝ} {t : ℕ} {a : α} +variable {K : ℕ} -section Algorithm +-- not used +lemma predictatble_pullCount (a : Fin K) : + Adapted (Bandits.filtration (Fin K) ℝ) (fun n ↦ pullCount a (n + 1)) := by + refine fun n ↦ Measurable.stronglyMeasurable ?_ + simp only + have : pullCount a (n + 1) = (fun h ↦ pullCount' n h a) ∘ (hist n) := by + ext + exact pullCount_add_one_eq_pullCount' + rw [Bandits.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe, this] + exact measurable_comp_comap (hist n) (measurable_pullCount' n a) + +-- not used +lemma isStoppingTime_stepsUntil (a : Fin K) (m : ℕ) : + IsStoppingTime (Bandits.filtration (Fin K) ℝ) (stepsUntil a m) := by + rw [stepsUntil_eq_leastGE] + refine Adapted.isStoppingTime_leastGE _ fun n ↦ ?_ + suffices StronglyMeasurable[Bandits.filtration (Fin K) ℝ n] (pullCount a (n + 1)) by fun_prop + exact predictatble_pullCount a n -variable [Nonempty α] [DecidableEq α] [Finite α] [Encodable α] [MeasurableSingletonClass α] +section Algorithm /-- The exploration bonus of the UCB algorithm, which corresponds to the width of a confidence interval. -/ -noncomputable def ucbWidth' (c : ℝ) (n : ℕ) (h : Iic n → α × ℝ) (a : α) : ℝ := - √(c * log (n + 1) / (pullCount' n h a)) +noncomputable def ucbWidth' (c : ℝ) (n : ℕ) (h : Iic n → Fin K × ℝ) (a : Fin K) : ℝ := + √(c * log (n + 2) / pullCount' n h a) open Classical in /-- Arm pulled by the UCB algorithm at time `n + 1`. -/ noncomputable -def ucbNextArm (c : ℝ) (n : ℕ) (h : Iic n → α × ℝ) : α := +def UCB.nextArm (hK : 0 < K) (c : ℝ) (n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K := + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + if n < K - 1 then ⟨(n + 1) % K, Nat.mod_lt _ hK⟩ else measurableArgmax (fun h a ↦ empMean' n h a + ucbWidth' c n h a) h @[fun_prop] -lemma measurable_ucbNextArm (c : ℝ) (n : ℕ) : Measurable (ucbNextArm c n (α := α)) := by - classical +lemma UCB.measurable_nextArm (hK : 0 < K) (c : ℝ) (n : ℕ) : Measurable (nextArm hK c n) := by + refine Measurable.ite (by simp) (by fun_prop) ?_ + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK refine measurable_measurableArgmax fun a ↦ ?_ unfold ucbWidth' fun_prop /-- The UCB algorithm. -/ noncomputable -def ucbAlgorithm (c : ℝ) : Algorithm α ℝ := - detAlgorithm (ucbNextArm c) (by fun_prop) (Classical.arbitrary α) +def ucbAlgorithm (hK : 0 < K) (c : ℝ) : Algorithm (Fin K) ℝ := + detAlgorithm (UCB.nextArm hK c) (by fun_prop) ⟨0, hK⟩ end Algorithm -variable [Fintype α] [Nonempty α] {c : ℝ} {μ : α → ℝ} {N : α → ℕ} {a : α} +namespace UCB + +variable {hK : 0 < K} {c : ℝ} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] {n : ℕ} {h : ℕ → Fin K × ℝ} /-- The exploration bonus of the UCB algorithm, which corresponds to the width of a confidence interval. -/ -noncomputable def ucbWidth (c : ℝ) (N : α → ℕ) (t : ℕ) (a : α) : ℝ := √(c * log t / N a) +noncomputable def ucbWidth (c : ℝ) (a : Fin K) (n : ℕ) (h : ℕ → Fin K × ℝ) : ℝ := + √(c * log (n + 1) / pullCount a n h) -/-- The arm pulled by the UCB algorithm. -/ -noncomputable -def ucbArm (c : ℝ) (μ : α → ℝ) (N : α → ℕ) (t : ℕ) : α := - (exists_max_image univ (fun a ↦ μ a + ucbWidth c N t a) - (univ_nonempty_iff.mpr inferInstance)).choose - -lemma le_ucb (a : α) : - μ a + ucbWidth c N t a ≤ μ (ucbArm c μ N t) + ucbWidth c N t (ucbArm c μ N t) := - (exists_max_image univ (fun a ↦ μ a + ucbWidth c N t a) - (univ_nonempty_iff.mpr inferInstance)).choose_spec.2 _ (mem_univ a) - -lemma gap_ucbArm_le_two_mul_ucbWidth - (h_best : (ν (bestArm ν))[id] ≤ μ (bestArm ν) + ucbWidth c N t (bestArm ν)) - (h_ucb : μ (ucbArm c μ N t) - ucbWidth c N t (ucbArm c μ N t) ≤ (ν (ucbArm c μ N t))[id]) : - gap ν (ucbArm c μ N t) ≤ 2 * ucbWidth c N t (ucbArm c μ N t) := by +@[fun_prop] +lemma measurable_ucbWidth (c : ℝ) (a : Fin K) : Measurable (ucbWidth c a n) := by + unfold ucbWidth + fun_prop + +lemma ucbWidth_eq_ucbWidth' (c : ℝ) (a : Fin K) (n : ℕ) (h : ℕ → Fin K × ℝ) (hn : n ≠ 0) : + ucbWidth c a n h = ucbWidth' c (n - 1) (fun i ↦ h i) a := by + simp only [ucbWidth, pullCount_eq_pullCount' hn, Nat.cast_nonneg, sqrt_div', ucbWidth'] + congr 4 + norm_cast + grind + +local notation "𝔓t" => Bandit.trajMeasure (ucbAlgorithm hK c) ν +local notation "𝔓" => Bandit.measure (ucbAlgorithm hK c) ν + +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_ucbNextArm (n : ℕ) : + arm (n + 1) =ᵐ[𝔓t] fun h ↦ nextArm hK c n (fun i ↦ h i) := by + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + exact arm_detAlgorithm_ae_eq n + +lemma arm_ae_all_eq : + ∀ᵐ h ∂𝔓t, arm 0 h = ⟨0, hK⟩ ∧ ∀ n, arm (n + 1) h = nextArm hK c n (fun i ↦ h i) := by + rw [eventually_and, ae_all_iff] + exact ⟨arm_zero, arm_ae_eq_ucbNextArm⟩ + +lemma ucbIndex_le_ucbIndex_arm (a : Fin K) (hn : K ≤ n) : + ∀ᵐ h ∂𝔓t, empMean a n h + ucbWidth c a n h ≤ + empMean (arm n h) n h + ucbWidth c (arm n h) n h := by + filter_upwards [arm_ae_eq_ucbNextArm (n - 1)] with h h_arm + have : n - 1 + 1 = n := by grind + have h_not_lt : ¬ n - 1 < K - 1 := by grind + simp only [this, nextArm, h_not_lt, ↓reduceIte] at h_arm + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + simp_rw [h_arm, empMean_eq_empMean' (by grind : n ≠ 0), + ucbWidth_eq_ucbWidth' _ _ _ _ (by grind : n ≠ 0)] + exact isMaxOn_measurableArgmax (fun h a ↦ empMean' (n - 1) h a + ucbWidth' c (n - 1) h a) + (fun i ↦ h i) a + +lemma forall_arm_eq_mod_of_lt : + ∀ᵐ h ∂𝔓t, ∀ n < K, arm n h = ⟨n % K, Nat.mod_lt _ hK⟩ := by + simp_rw [ae_all_iff] + intro n hn + induction n with + | zero => exact arm_zero + | succ n _ => + filter_upwards [arm_ae_eq_ucbNextArm n] with h h_eq + rw [h_eq, nextArm, if_pos] + grind + +lemma forall_ucbIndex_le_ucbIndex_arm (a : Fin K) : + ∀ᵐ h ∂𝔓t, ∀ n, K ≤ n → + empMean a n h + ucbWidth c a n h ≤ empMean (arm n h) n h + ucbWidth c (arm n h) n h := by + simp_rw [ae_all_iff] + exact fun _ ↦ ucbIndex_le_ucbIndex_arm a + +lemma forall_arm_prop : + ∀ᵐ h ∂𝔓t, + (∀ n < K, arm n h = ⟨n % K, Nat.mod_lt _ hK⟩) ∧ + (∀ n, K ≤ n → ∀ a, empMean a n h + ucbWidth c a n h ≤ + empMean (arm n h) n h + ucbWidth c (arm n h) n h) := by + simp only [eventually_and] + constructor + · exact forall_arm_eq_mod_of_lt + · simp_rw [ae_all_iff] + intro n hn a + have h_ae := forall_ucbIndex_le_ucbIndex_arm (ν := ν) (c := c) (hK := hK) a + simp_rw [ae_all_iff] at h_ae + exact h_ae n hn + +lemma pullCount_eq_of_time_eq (a : Fin K) : + ∀ᵐ ω ∂𝔓t, pullCount a K ω = 1 := by + filter_upwards [forall_arm_eq_mod_of_lt] with h h_eq + rw [pullCount_eq_sum] + conv_rhs => rw [← sum_mod_range hK a] + refine Finset.sum_congr rfl fun s hs ↦ ?_ + congr + exact h_eq s (by grind) + +lemma time_gt_of_pullCount_gt_one (a : Fin K) : + ∀ᵐ ω ∂𝔓t, ∀ n, 1 < pullCount a n ω → K < n := by + filter_upwards [pullCount_eq_of_time_eq a] with h h_eq n hn + rw [← h_eq] at hn + by_contra! h_lt + exact hn.not_ge (pullCount_mono _ h_lt _) + +lemma pullCount_pos_of_time_ge : + ∀ᵐ ω ∂𝔓t, ∀ n, K ≤ n → ∀ b : Fin K, 0 < pullCount b n ω := by + have h_ae a := pullCount_eq_of_time_eq (ν := ν) (c := c) (hK := hK) a + rw [← ae_all_iff] at h_ae + filter_upwards [h_ae] with ω hω n hn a + refine Nat.one_pos.trans_le ?_ + rw [← hω a] + exact pullCount_mono _ hn _ + +lemma pullCount_pos_of_pullCount_gt_one (a : Fin K) : + ∀ᵐ ω ∂𝔓t, ∀ n, 1 < pullCount a n ω → ∀ b : Fin K, 0 < pullCount b n ω := by + filter_upwards [time_gt_of_pullCount_gt_one a, pullCount_pos_of_time_ge] with ω h1 h2 n h_gt a + exact h2 n (h1 n h_gt).le a + +omit [IsMarkovKernel ν] in +lemma gap_arm_le_two_mul_ucbWidth [Nonempty (Fin K)] + (h_best : (ν (bestArm ν))[id] ≤ empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h) + (h_arm : empMean (arm n h) n h - ucbWidth c (arm n h) n h ≤ (ν (arm n h))[id]) + (h_le : empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h ≤ + empMean (arm n h) n h + ucbWidth c (arm n h) n h) : + gap ν (arm n h) ≤ 2 * ucbWidth c (arm n h) n h := by rw [gap_eq_bestArm_sub, sub_le_iff_le_add'] calc (ν (bestArm ν))[id] - _ ≤ μ (bestArm ν) + ucbWidth c N t (bestArm ν) := h_best - _ ≤ μ (ucbArm c μ N t) + ucbWidth c N t (ucbArm c μ N t) := le_ucb _ - _ ≤ (ν (ucbArm c μ N t))[id] + 2 * ucbWidth c N t (ucbArm c μ N t) := by + _ ≤ empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h := h_best + _ ≤ empMean (arm n h) n h + ucbWidth c (arm n h) n h := h_le + _ ≤ (ν (arm n h))[id] + 2 * ucbWidth c (arm n h) n h := by rw [two_mul, ← add_assoc] gcongr - rwa [sub_le_iff_le_add] at h_ucb - -lemma N_ucbArm_le - (h_best : (ν (bestArm ν))[id] ≤ μ (bestArm ν) + ucbWidth c N t (bestArm ν)) - (h_ucb : μ (ucbArm c μ N t) - ucbWidth c N t (ucbArm c μ N t) ≤ (ν (ucbArm c μ N t))[id]) : - N (ucbArm c μ N t) ≤ 4 * c * log t / gap ν (ucbArm c μ N t) ^ 2 := by - have h_gap := gap_ucbArm_le_two_mul_ucbWidth h_best h_ucb - rw [ucbWidth] at h_gap - sorry + rwa [sub_le_iff_le_add] at h_arm + +omit [IsMarkovKernel ν] in +lemma pullCount_arm_le [Nonempty (Fin K)] (hc : 0 ≤ c) + (h_best : (ν (bestArm ν))[id] ≤ empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h) + (h_arm : empMean (arm n h) n h - ucbWidth c (arm n h) n h ≤ (ν (arm n h))[id]) + (h_le : empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h ≤ + empMean (arm n h) n h + ucbWidth c (arm n h) n h) + (h_gap_pos : 0 < gap ν (arm n h)) (h_pull_pos : 0 < pullCount (arm n h) n h) : + pullCount (arm n h) n h ≤ 4 * c * log (n + 1) / gap ν (arm n h) ^ 2 := by + have h_gap_le := gap_arm_le_two_mul_ucbWidth h_best h_arm h_le + rw [ucbWidth] at h_gap_le + have h2 : (gap ν (arm n h)) ^ 2 ≤ (2 * √(c * log (n + 1) / pullCount (arm n h) n h)) ^ 2 := by + gcongr + rw [mul_pow, sq_sqrt] at h2 + · have : (2 : ℝ) ^ 2 = 4 := by norm_num + rw [this] at h2 + field_simp at h2 ⊢ + exact h2 + · have : 0 ≤ log (n + 1) := by simp [log_nonneg] + positivity + +lemma todo (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + (hc : 0 ≤ c) (a : Fin K) (n k : ℕ) (hk : k ≠ 0) : + 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} ≤ + 1 / (n + 1) ^ (c / 2) := by + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + have h_meas : MeasurableSet {ω | ω / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := + measurableSet_le (by fun_prop) measurable_const + have h_log_nonneg : 0 ≤ log (n + 1) := log_nonneg (by simp) + calc + 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} + _ = ((𝔓).map (fun ω ↦ ∑ m ∈ Icc 1 k, rewardByCount a m ω)) + {ω | ω / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := by + rw [Measure.map_apply (by fun_prop) h_meas] + rfl + _ = ((𝔓).map (fun ω ↦ ∑ s ∈ range k, ω.2 s a)) + {ω | ω / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := by + rw [IdentDistrib.map_eq (identDistrib_sum_Icc_rewardByCount k a)] + _ = 𝔓 {ω | (∑ s ∈ range k, ω.2 s a) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := by + rw [Measure.map_apply (by fun_prop) h_meas] + rfl + _ = 𝔓 {ω | (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) / k ≤ - √(c * log (n + 1) / k)} := by + congr with ω + field_simp + rw [Finset.sum_sub_distrib] + simp + grind + _ = 𝔓 {ω | (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) ≤ - √(c * k * log (n + 1))} := by + congr with ω + field_simp + congr! 2 + rw [sqrt_div (by positivity), ← mul_div_assoc, mul_comm, mul_div_assoc, div_sqrt, + mul_assoc (k : ℝ), sqrt_mul (x := (k : ℝ)) (by positivity), mul_comm] + _ = Bandit.streamMeasure ν + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ - √(c * k * log (n + 1))} := by + rw [← Bandit.snd_measure (ucbAlgorithm hK c), Measure.snd_apply] + · rfl + · exact measurableSet_le (by fun_prop) (by fun_prop) + _ ≤ ENNReal.ofReal (exp (-(√(c * k * log (n + 1))) ^ 2 / (2 * k * 1))) := by + rw [← ofReal_measureReal] + gcongr + refine (HasSubgaussianMGF.measure_sum_range_le_le_of_iIndepFun (c := 1) ?_ ?_ (by positivity)) + · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) + (fun _ ↦ by fun_prop) + · intro i him + refine (hν a).congr_identDistrib ?_ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ + _ = 1 / (n + 1) ^ (c / 2) := by + rw [sq_sqrt] + swap; · exact mul_nonneg (by positivity) (log_nonneg (by simp)) + field_simp + rw [div_eq_inv_mul, ← mul_assoc, ← Real.log_rpow (by positivity), ← Real.log_inv, + Real.exp_log (by positivity), one_div, ENNReal.ofReal_inv_of_pos (by positivity), + ← ENNReal.ofReal_rpow_of_nonneg (by positivity) (by positivity)] + congr 2 + · norm_cast + · field + +lemma todo' (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + (hc : 0 ≤ c) (a : Fin K) (n k : ℕ) (hk : k ≠ 0) : + 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - √(c * log (n + 1) / k)} ≤ + 1 / (n + 1) ^ (c / 2) := by + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + have h_meas : MeasurableSet {ω |(ν a)[id] ≤ ω / k - √(c * log (n + 1) / k)} := + measurableSet_le (by fun_prop) (by fun_prop) + have h_log_nonneg : 0 ≤ log (n + 1) := log_nonneg (by simp) + calc + 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - √(c * log (n + 1) / k)} + _ = ((𝔓).map (fun ω ↦ ∑ m ∈ Icc 1 k, rewardByCount a m ω)) + {ω | (ν a)[id] ≤ ω / k - √(c * log (n + 1) / k)} := by + rw [Measure.map_apply (by fun_prop) h_meas] + rfl + _ = ((𝔓).map (fun ω ↦ ∑ s ∈ range k, ω.2 s a)) + {ω | (ν a)[id] ≤ ω / k - √(c * log (n + 1) / k)} := by + rw [IdentDistrib.map_eq (identDistrib_sum_Icc_rewardByCount k a)] + _ = 𝔓 {ω | (ν a)[id] ≤ (∑ s ∈ range k, ω.2 s a) / k - √(c * log (n + 1) / k)} := by + rw [Measure.map_apply (by fun_prop) h_meas] + rfl + _ = 𝔓 {ω | √(c * log (n + 1) / k) ≤ (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) / k} := by + congr with ω + field_simp + rw [Finset.sum_sub_distrib] + simp + grind + _ = 𝔓 {ω | √(c * k * log (n + 1)) ≤ (∑ s ∈ range k, (ω.2 s a - (ν a)[id]))} := by + congr with ω + field_simp + congr! 1 + rw [sqrt_div (by positivity), ← mul_div_assoc, mul_comm, mul_div_assoc, div_sqrt, + mul_comm _ (k : ℝ), sqrt_mul (x := (k : ℝ)) (by positivity), mul_comm] + _ = Bandit.streamMeasure ν + {ω | √(c * k * log (n + 1)) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} := by + rw [← Bandit.snd_measure (ucbAlgorithm hK c), Measure.snd_apply] + · rfl + · exact measurableSet_le (by fun_prop) (by fun_prop) + _ ≤ ENNReal.ofReal (exp (-(√(c * k * log (n + 1))) ^ 2 / (2 * k * 1))) := by + rw [← ofReal_measureReal] + gcongr + refine (HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := 1) ?_ ?_ (by positivity)) + · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) + (fun _ ↦ by fun_prop) + · intro i him + refine (hν a).congr_identDistrib ?_ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ + _ = 1 / (n + 1) ^ (c / 2) := by + rw [sq_sqrt] + swap; · exact mul_nonneg (by positivity) (log_nonneg (by simp)) + field_simp + rw [div_eq_inv_mul, ← mul_assoc, ← Real.log_rpow (by positivity), ← Real.log_inv, + Real.exp_log (by positivity), one_div, ENNReal.ofReal_inv_of_pos (by positivity), + ← ENNReal.ofReal_rpow_of_nonneg (by positivity) (by positivity)] + congr 2 + · norm_cast + · field + +lemma prob_ucbIndex_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : + 𝔓t {h | 0 < pullCount a n h ∧ empMean a n h + ucbWidth c a n h ≤ (ν a)[id]} ≤ + 1 / (n + 1) ^ (c / 2 - 1) := by + -- extend the probability space + suffices 𝔓 {ω | 0 < pullCount a n ω.1 ∧ + empMean a n ω.1 + ucbWidth c a n ω.1 ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2 - 1) by + rwa [← Bandit.fst_measure (ucbAlgorithm hK c) ν, Measure.fst_apply] + change MeasurableSet ({h | 0 < pullCount a n h} + ∩ {h | empMean a n h + ucbWidth c a n h ≤ ∫ (x : ℝ), id x ∂ν a}) + refine MeasurableSet.inter ?_ ?_ + · exact measurableSet_lt (by fun_prop) (by fun_prop) + · exact measurableSet_le (by fun_prop) (by fun_prop) + -- express with `rewardByCount` and `pullCount` + unfold empMean ucbWidth + simp_rw [← sum_rewardByCount_eq_sumRewards] + calc + 𝔓 {ω | 0 < pullCount a n ω.1 ∧ + (∑ m ∈ Icc 1 (pullCount a n ω.1), rewardByCount a m ω) / pullCount a n ω.1 + + √(c * log (↑n + 1) / pullCount a n ω.1) ≤ (ν a)[id]} + -- list the possible values of `pullCount a n ω.1` + _ ≤ 𝔓 {ω | ∃ k ≤ n, 0 < k ∧ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + + √(c * log (↑n + 1) / k) ≤ (ν a)[id]} := by + refine measure_mono fun ω hω ↦ ?_ + simp only [Nat.cast_nonneg, sqrt_div', id_eq, Set.mem_setOf_eq] at hω ⊢ + exact ⟨pullCount a n ω.1, pullCount_le _ _ _, hω⟩ + _ = 𝔓 (⋃ k ∈ Icc 1 n, {ω |(∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + + √(c * log (↑n + 1) / k) ≤ (ν a)[id]}) := by + congr 1 + ext ω + simp + grind + -- Union bound over the possible values of `pullCount a n ω.1` + _ ≤ ∑ k ∈ Icc 1 n, + 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + √(c * log (↑n + 1) / k) ≤ (ν a)[id]} := + measure_biUnion_finset_le _ _ + _ ≤ ∑ k ∈ Icc 1 n, (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by + gcongr with k hk + exact todo hν hc a n k (by grind) + _ ≤ (n + 1) * (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by + simp only [one_div, sum_const, Nat.card_Icc, add_tsub_cancel_right, nsmul_eq_mul, mul_one] + rw [div_eq_mul_inv ((n : ℝ≥0∞) + 1)] + gcongr + exact le_self_add + _ = 1 / (n + 1) ^ (c / 2 - 1) := by + simp only [mul_one, one_div] + rw [ENNReal.rpow_sub _ _ (by simp) (by finiteness), ENNReal.rpow_one, div_eq_mul_inv, + ENNReal.div_eq_inv_mul, ENNReal.mul_inv (by simp) (by simp), inv_inv] + +lemma prob_ucbIndex_ge (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : + 𝔓t {h | 0 < pullCount a n h ∧ + (ν a)[id] ≤ empMean a n h - ucbWidth c a n h} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by + -- extend the probability space + suffices 𝔓 {ω | 0 < pullCount a n ω.1 ∧ + (ν a)[id] ≤ empMean a n ω.1 - ucbWidth c a n ω.1} ≤ 1 / (n + 1) ^ (c / 2 - 1) by + rwa [← Bandit.fst_measure (ucbAlgorithm hK c) ν, Measure.fst_apply] + change MeasurableSet ({h | 0 < pullCount a n h} + ∩ {h | (ν a)[id] ≤ empMean a n h - ucbWidth c a n h}) + refine MeasurableSet.inter ?_ ?_ + · exact measurableSet_lt (by fun_prop) (by fun_prop) + · exact measurableSet_le (by fun_prop) (by fun_prop) + -- express with `rewardByCount` and `pullCount` + unfold empMean ucbWidth + simp_rw [← sum_rewardByCount_eq_sumRewards] + calc + 𝔓 {ω | 0 < pullCount a n ω.1 ∧ + (ν a)[id] ≤ (∑ m ∈ Icc 1 (pullCount a n ω.1), rewardByCount a m ω) / pullCount a n ω.1 - + √(c * log (↑n + 1) / pullCount a n ω.1)} + -- list the possible values of `pullCount a n ω.1` + _ ≤ 𝔓 {ω | ∃ k ≤ n, 0 < k ∧ (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - + √(c * log (↑n + 1) / k)} := by + refine measure_mono fun ω hω ↦ ?_ + simp only [Nat.cast_nonneg, sqrt_div', id_eq, Set.mem_setOf_eq] at hω ⊢ + exact ⟨pullCount a n ω.1, pullCount_le _ _ _, hω⟩ + _ = 𝔓 (⋃ k ∈ Icc 1 n, {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - + √(c * log (↑n + 1) / k)}) := by + congr 1 + ext ω + simp + grind + -- Union bound over the possible values of `pullCount a n ω.1` + _ ≤ ∑ k ∈ Icc 1 n, + 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - √(c * log (↑n + 1) / k)} := + measure_biUnion_finset_le _ _ + _ ≤ ∑ k ∈ Icc 1 n, (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by + gcongr with k hk + exact todo' hν hc a n k (by grind) + _ ≤ (n + 1) * (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by + simp only [one_div, sum_const, Nat.card_Icc, add_tsub_cancel_right, nsmul_eq_mul, mul_one] + rw [div_eq_mul_inv ((n : ℝ≥0∞) + 1)] + gcongr + exact le_self_add + _ = 1 / (n + 1) ^ (c / 2 - 1) := by + simp only [mul_one, one_div] + rw [ENNReal.rpow_sub _ _ (by simp) (by finiteness), ENNReal.rpow_one, div_eq_mul_inv, + ENNReal.div_eq_inv_mul, ENNReal.mul_inv (by simp) (by simp), inv_inv] + +lemma probReal_ucbIndex_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : + (𝔓t).real {h | 0 < pullCount a n h ∧ empMean a n h + ucbWidth c a n h ≤ (ν a)[id]} ≤ + 1 / (n + 1) ^ (c / 2 - 1) := by + rw [measureReal_def] + grw [prob_ucbIndex_le hν hc a n] + swap; · finiteness + simp only [one_div, ENNReal.toReal_inv] + rw [← ENNReal.toReal_rpow] + norm_cast + +lemma probReal_ucbIndex_ge (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : + (𝔓t).real {h | 0 < pullCount a n h ∧ + (ν a)[id] ≤ empMean a n h - ucbWidth c a n h} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by + rw [measureReal_def] + grw [prob_ucbIndex_ge hν hc a n] + swap; · finiteness + simp only [one_div, ENNReal.toReal_inv] + rw [← ENNReal.toReal_rpow] + norm_cast + +lemma pullCount_le_add (a : Fin K) (n C : ℕ) (ω : ℕ → Fin K × ℝ) : + pullCount a n ω ≤ C + 1 + + ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω}.indicator 1 s := by + rw [pullCount_eq_sum] + calc ∑ s ∈ range n, if arm s ω = a then 1 else 0 + _ ≤ ∑ s ∈ range n, ({s | arm s ω = a ∧ pullCount a s ω ≤ C}.indicator 1 s + + {s | arm s ω = a ∧ C < pullCount a s ω}.indicator 1 s) := by + gcongr with s hs + simp [Set.indicator_apply] + grind + _ = ∑ s ∈ range n, {s | arm s ω = a ∧ pullCount a s ω ≤ C}.indicator 1 s + + ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω}.indicator 1 s := by + rw [Finset.sum_add_distrib] + _ ≤ C + 1 + ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω}.indicator 1 s := by + gcongr + simp_rw [pullCount_eq_sum] + sorry + +omit [IsMarkovKernel ν] in +lemma pullCount_le_add_three [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (ω : ℕ → Fin K × ℝ) : + pullCount a n ω ≤ C + 1 + + ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω ∧ + (ν (bestArm ν))[id] ≤ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω ∧ + empMean (arm s ω) s ω - ucbWidth c (arm s ω) s ω ≤ (ν (arm s ω))[id]}.indicator 1 s + + ∑ s ∈ range n, + {s | C < pullCount a s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < + (ν (bestArm ν))[id]}.indicator 1 s + + ∑ s ∈ range n, + {s | C < pullCount a s ω ∧ (ν a)[id] < + empMean a s ω - ucbWidth c a s ω}.indicator 1 s := by + refine (pullCount_le_add a n C ω).trans ?_ + simp_rw [add_assoc] + gcongr + simp_rw [← add_assoc] + let A := {s | arm s ω = a ∧ C < pullCount a s ω} + let B := {s | arm s ω = a ∧ C < pullCount a s ω ∧ + (ν (bestArm ν))[id] ≤ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω ∧ + empMean (arm s ω) s ω - ucbWidth c (arm s ω) s ω ≤ (ν (arm s ω))[id]} + let C' := {s | C < pullCount a s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < + (ν (bestArm ν))[id]} + let D := {s | C < pullCount a s ω ∧ (ν a)[id] < + empMean a s ω - ucbWidth c a s ω} + change ∑ s ∈ range n, A.indicator 1 s ≤ + ∑ s ∈ range n, B.indicator 1 s + ∑ s ∈ range n, C'.indicator 1 s + + ∑ s ∈ range n, D.indicator 1 s + have h_union : A ⊆ B ∪ C' ∪ D := by simp [A, B, C', D]; grind + calc + (∑ s ∈ range n, A.indicator 1 s) + _ ≤ (∑ s ∈ range n, (B ∪ C' ∪ D).indicator (fun _ ↦ (1 : ℕ)) s) := by + gcongr with n hn + by_cases h : n ∈ A + · have : n ∈ B ∪ C' ∪ D := h_union h + simp [h, this] + · simp [h] + _ ≤ ∑ s ∈ range n, (B.indicator 1 s + C'.indicator 1 s + D.indicator 1 s) := by + gcongr with s + simp [Set.indicator_apply] + grind + _ = ∑ s ∈ range n, B.indicator 1 s + ∑ s ∈ range n, C'.indicator 1 s + + ∑ s ∈ range n, D.indicator 1 s := by + rw [Finset.sum_add_distrib, Finset.sum_add_distrib] + +lemma pullCount_le_add_three_ae [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (hC : C ≠ 0) : + ∀ᵐ ω ∂𝔓t, + pullCount a n ω ≤ C + 1 + + ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω ∧ + (ν (bestArm ν))[id] ≤ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω ∧ + empMean (arm s ω) s ω - ucbWidth c (arm s ω) s ω ≤ (ν (arm s ω))[id]}.indicator 1 s + + ∑ s ∈ range n, + {s | 0 < pullCount (bestArm ν) s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < + (ν (bestArm ν))[id]}.indicator 1 s + + ∑ s ∈ range n, + {s | 0 < pullCount a s ω ∧ (ν a)[id] < + empMean a s ω - ucbWidth c a s ω}.indicator 1 s := by + filter_upwards [pullCount_pos_of_pullCount_gt_one a] with ω hω + refine (pullCount_le_add_three a n C ω (ν := ν) (c := c)).trans ?_ + gcongr 5 with k hk j k hk j + · gcongr 1 + exact fun h_gt ↦ hω _ (lt_of_le_of_lt (by grind) h_gt) _ + · exact fun h_gt ↦ hω _ (lt_of_le_of_lt (by grind) h_gt) _ + +lemma some_sum_eq_zero [Nonempty (Fin K)] (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 < gap ν a) (n C : ℕ) + (hC : C ≠ 0) (hC' : 4 * c * log (n + 1) / gap ν a ^ 2 ≤ C) : + ∀ᵐ ω ∂𝔓t, + ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω ∧ + (ν (bestArm ν))[id] ≤ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω ∧ + empMean (arm s ω) s ω - ucbWidth c (arm s ω) s ω ≤ (ν (arm s ω))[id]}.indicator 1 s = 0 := by + have h_ae := forall_ucbIndex_le_ucbIndex_arm (bestArm ν) (ν := ν) (c := c) (hK := hK) + have h_gt := time_gt_of_pullCount_gt_one a (ν := ν) (c := c) (hK := hK) + filter_upwards [h_ae, h_gt] with ω h_le h_time_ge + simp only [id_eq, tsub_le_iff_right, sum_eq_zero_iff, mem_range, Set.indicator_apply_eq_zero, + Set.mem_setOf_eq, Pi.one_apply, one_ne_zero, imp_false, not_and, not_le] + intro k hn h_arm hC_lt h_le_best + by_contra! h_le_arm + have h := pullCount_arm_le hc h_le_best (by simpa) ?_ ?_ ?_ + rotate_left + · refine h_le _ ?_ + refine (h_time_ge _ ?_).le + refine lt_of_le_of_lt ?_ hC_lt + grind + · rwa [h_arm] + · rw [h_arm] + exact zero_le'.trans_lt hC_lt + refine lt_irrefl (4 * c * log (n + 1) / gap ν a ^ 2) ?_ + refine hC'.trans_lt (lt_of_lt_of_le ?_ (h.trans ?_)) + · rw [h_arm] + exact mod_cast hC_lt + · rw [h_arm] + gcongr + +lemma pullCount_ae_le_add_two [Nonempty (Fin K)] (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 < gap ν a) + (n C : ℕ) (hC : C ≠ 0) (hC' : 4 * c * log (n + 1) / gap ν a ^ 2 ≤ C) : + ∀ᵐ ω ∂𝔓t, + pullCount a n ω ≤ C + 1 + + ∑ s ∈ range n, + {s | 0 < pullCount (bestArm ν) s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < + (ν (bestArm ν))[id]}.indicator 1 s + + ∑ s ∈ range n, + {s | 0 < pullCount a s ω ∧ (ν a)[id] < + empMean a s ω - ucbWidth c a s ω}.indicator 1 s := by + filter_upwards [some_sum_eq_zero hc a h_gap n C hC hC', + pullCount_le_add_three_ae a n C hC] with ω hω_zero hω_le + refine (hω_le).trans_eq ?_ + rw [hω_zero] + +/-- A sum that appears in the UCB regret upper bound. -/ +noncomputable +def constSum (c : ℝ) (n : ℕ) : ℝ≥0∞ := ∑ s ∈ range n, 1 / ((s : ℝ≥0∞) + 1) ^ (c / 2 - 1) + +lemma constSum_lt_top (c : ℝ) (n : ℕ) : constSum c n < ∞ := by + rw [constSum, ENNReal.sum_lt_top] + intro k hk + simp only [one_div, ENNReal.inv_lt_top] + positivity + +/-- Bound on the expectation of the number of pulls of each arm by the UCB algorithm. -/ +lemma expectation_pullCount_le' (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + (hc : 0 < c) (a : Fin K) (h_gap : 0 < gap ν a) (n : ℕ) : + ∫⁻ ω, pullCount a n ω ∂𝔓t ≤ + ENNReal.ofReal (4 * c * log (n + 1) / gap ν a ^ 2 + 1) + 1 + 2 * constSum c n := by + by_cases hn_zero : n = 0 + · simp [hn_zero] + let C a : ℕ := ⌈4 * c * log (n + 1) / gap ν a ^ 2⌉₊ + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + have h_set_1 b : MeasurableSet {a_1 | 0 < pullCount a b a_1 ∧ + (ν a)[id] < empMean a b a_1 - ucbWidth c a b a_1} := by + change MeasurableSet ({a_1 | 0 < pullCount a b a_1} ∩ + {a_1 | (ν a)[id] < empMean a b a_1 - ucbWidth c a b a_1}) + exact (measurableSet_lt (by fun_prop) (by fun_prop)).inter + (measurableSet_lt (by fun_prop) (by fun_prop)) + have h_set_2 b : MeasurableSet {a | 0 < pullCount (bestArm ν) b a ∧ + empMean (bestArm ν) b a + ucbWidth c (bestArm ν) b a < (ν (bestArm ν))[id]} := by + change MeasurableSet ({a | 0 < pullCount (bestArm ν) b a} ∩ + {a | empMean (bestArm ν) b a + ucbWidth c (bestArm ν) b a < (ν (bestArm ν))[id]}) + exact (measurableSet_lt (by fun_prop) (by fun_prop)).inter + (measurableSet_lt (by fun_prop) (by fun_prop)) + have h_meas_1 b : Measurable fun h ↦ {s | 0 < pullCount a s h ∧ (ν a)[id] < + empMean a s h - ucbWidth c a s h}.indicator (1 : ℕ → ℕ) b := by + simp only [id_eq, Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] + exact Measurable.ite (h_set_1 _) (by fun_prop) (by fun_prop) + have h_meas_2 b : Measurable fun h ↦ {s | 0 < pullCount (bestArm ν) s h ∧ + empMean (bestArm ν) s h + ucbWidth c (bestArm ν) s h < + (ν (bestArm ν))[id]}.indicator (1 : ℕ → ℕ) b := by + simp only [id_eq, Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] + exact Measurable.ite (h_set_2 _) (by fun_prop) (by fun_prop) + calc ∫⁻ ω, pullCount a n ω ∂𝔓t + _ ≤ ∫⁻ ω, C a + 1 + + ∑ s ∈ range n, + {s | 0 < pullCount (bestArm ν) s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < + (ν (bestArm ν))[id]}.indicator (1 : ℕ → ℕ) s + + ∑ s ∈ range n, + {s | 0 < pullCount a s ω ∧ (ν a)[id] < + empMean a s ω - ucbWidth c a s ω}.indicator (1 : ℕ → ℕ) s ∂𝔓t := by + refine lintegral_mono_ae ?_ + have hCa : C a ≠ 0 := by + simp only [ne_eq, Nat.ceil_eq_zero, not_le, C] + have : 0 < log (n + 1) := log_pos (by simp; grind) + positivity + filter_upwards [pullCount_ae_le_add_two hc.le a h_gap n (C a) hCa (Nat.le_ceil _)] with ω hω + simp only [id_eq, Nat.cast_sum] + norm_cast + _ ≤ (C a : ℝ≥0∞) + 1 + + ∑ s ∈ range n, + 𝔓t {ω | 0 < pullCount (bestArm ν) s ω ∧ + empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < (ν (bestArm ν))[id]} + + ∑ s ∈ range n, + 𝔓t {ω | 0 < pullCount a s ω ∧ (ν a)[id] < empMean a s ω - ucbWidth c a s ω} := by + simp only [id_eq, Nat.cast_sum] + rw [lintegral_add_left (by fun_prop), lintegral_add_left (by fun_prop)] + simp only [lintegral_const, measure_univ, mul_one] + rw [lintegral_finset_sum _ (by fun_prop), lintegral_finset_sum _ (by fun_prop)] + gcongr with k hk k hk + · rw [← lintegral_indicator_one] + swap; · exact h_set_2 _ + gcongr with h + simp [Set.indicator_apply] + · rw [← lintegral_indicator_one] + swap; · exact h_set_1 _ + gcongr with h + simp [Set.indicator_apply] + _ ≤ (C a : ℝ≥0∞) + 1 + + ∑ s ∈ range n, 1 / ((s : ℝ≥0∞) + 1) ^ (c / 2 - 1) + + ∑ s ∈ range n, 1 / ((s : ℝ≥0∞) + 1) ^ (c / 2 - 1) := by + gcongr with s hs s hs + · refine (measure_mono ?_).trans (prob_ucbIndex_le hν hc.le (bestArm ν) s) + grind + · refine (measure_mono ?_).trans (prob_ucbIndex_ge hν hc.le a s) + grind + _ ≤ ENNReal.ofReal (4 * c * log (n + 1) / gap ν a ^ 2 + 1) + 1 + 2 * constSum c n := by + rw [two_mul, add_assoc, constSum] + gcongr + simp only [C] + rw [← ENNReal.ofReal_natCast] + refine ENNReal.ofReal_le_ofReal ?_ + refine (Nat.ceil_lt_add_one ?_).le + have : 0 ≤ log (n + 1) := log_nonneg (by simp) + positivity + +/-- Bound on the expectation of the number of pulls of each arm by the UCB algorithm. -/ +lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + (hc : 0 < c) (a : Fin K) (h_gap : 0 < gap ν a) (n : ℕ) : + 𝔓t[fun ω ↦ (pullCount a n ω : ℝ)] ≤ + 4 * c * log (n + 1) / gap ν a ^ 2 + 2 + 2 * (constSum c n).toReal := by + have h := expectation_pullCount_le' hν hc a h_gap n (hK := hK) + simp_rw [← ENNReal.ofReal_natCast] at h + rw [← ofReal_integral_eq_lintegral_ofReal] at h + rotate_left + · exact integrable_pullCount _ _ + · exact ae_of_all _ fun _ ↦ by simp + simp only + have : 0 ≤ log (n + 1) := log_nonneg (by simp) + rw [← ENNReal.ofReal_toReal (a := 2 * constSum c n), ← ENNReal.ofReal_one, ← ENNReal.ofReal_add, + ← ENNReal.ofReal_add, ENNReal.ofReal_le_ofReal_iff] at h + rotate_left + · positivity + · positivity + · simp + · have : constSum c n ≠ ∞ := (constSum_lt_top c n).ne + finiteness + · simp + · have : constSum c n ≠ ∞ := (constSum_lt_top c n).ne + finiteness + refine h.trans_eq ?_ + simp only [ENNReal.toReal_mul, ENNReal.toReal_ofNat, add_left_inj] + ring + +/-- Regret bound for the UCB algorithm. -/ +lemma regret_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 < c) (n : ℕ) : + 𝔓t[regret ν n] ≤ + ∑ a, (4 * c * log (n + 1) / gap ν a + gap ν a * (2 + 2 * (constSum c n).toReal)) := by + simp_rw [regret_eq_sum_pullCount_mul_gap] + rw [integral_finset_sum] + swap; · exact fun i _ ↦ (integrable_pullCount i n).mul_const _ + gcongr with a + rw [integral_mul_const] + by_cases h_gap : gap ν a = 0 + · simp [h_gap] + replace h_gap : 0 < gap ν a := lt_of_le_of_ne gap_nonneg (Ne.symm h_gap) + grw [expectation_pullCount_le hν hc a h_gap n] + refine le_of_eq ?_ + rw [mul_add] + field + +end UCB end Bandits diff --git a/LeanBandits/ForMathlib/CondDistrib.lean b/LeanBandits/ForMathlib/CondDistrib.lean index 284eb3ce..ffc4ae1c 100644 --- a/LeanBandits/ForMathlib/CondDistrib.lean +++ b/LeanBandits/ForMathlib/CondDistrib.lean @@ -265,7 +265,25 @@ lemma condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft Y ⟂ᵢ[Z, hZ; μ] X := by refine condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkRight hX hY hZ ?_ (η := η) rw [← Kernel.compProd_eq_iff, compProd_map_condDistrib (by fun_prop)] at h ⊢ - sorry + have : μ.map (fun a ↦ ((Z a, X a), Y a)) + = (μ.map (fun a ↦ ((X a, Z a), Y a))).map (fun p ↦ ((p.1.2, p.1.1), p.2)) := by + rw [Measure.map_map (by fun_prop) (by fun_prop)] + rfl + rw [this, h] + ext s hs + rw [Measure.map_apply, Measure.compProd_apply, Measure.compProd_apply, lintegral_map, + lintegral_map] + · simp only [Kernel.prodMkLeft_apply, Kernel.prodMkRight_apply] + congr + · exact Kernel.measurable_kernel_prodMk_left hs + · fun_prop + · refine Kernel.measurable_kernel_prodMk_left ?_ + exact hs.preimage (by fun_prop) + · fun_prop + · exact hs + · exact hs.preimage (by fun_prop) + · fun_prop + · exact hs /-- Law of `Y` conditioned on `X`. -/ notation "𝓛[" Y " | " X "; " μ "]" => condDistrib Y X μ diff --git a/LeanBandits/ForMathlib/SubGaussian.lean b/LeanBandits/ForMathlib/SubGaussian.lean index 34f628f3..658dd237 100644 --- a/LeanBandits/ForMathlib/SubGaussian.lean +++ b/LeanBandits/ForMathlib/SubGaussian.lean @@ -5,7 +5,7 @@ Authors: Rémy Degenne -/ import Mathlib.Probability.Moments.SubGaussian -open MeasureTheory +open MeasureTheory Real open scoped ENNReal NNReal namespace ProbabilityTheory @@ -14,6 +14,33 @@ namespace HasSubgaussianMGF variable {Ω : Type*} {m mΩ : MeasurableSpace Ω} {μ : Measure Ω} {X Y : Ω → ℝ} {c cX cY : ℝ≥0} +/-- Chernoff bound on the left tail of a sub-Gaussian random variable. -/ +lemma measure_le_le (h : HasSubgaussianMGF X c μ) {ε : ℝ} (hε : 0 ≤ ε) : + μ.real {ω | X ω ≤ -ε} ≤ exp (-ε ^ 2 / (2 * c)) := by + simp_rw [le_neg (b := ε), ← Pi.neg_apply] + exact h.neg.measure_ge_le hε + +/-- **Hoeffding inequality** for sub-Gaussian random variables. -/ +lemma measure_sum_le_le_of_iIndepFun {ι : Type*} {X : ι → Ω → ℝ} (h_indep : iIndepFun X μ) + {c : ι → ℝ≥0} + {s : Finset ι} (h_subG : ∀ i ∈ s, HasSubgaussianMGF (X i) (c i) μ) {ε : ℝ} (hε : 0 ≤ ε) : + μ.real {ω | ∑ i ∈ s, X i ω ≤ -ε} ≤ exp (-ε ^ 2 / (2 * ∑ i ∈ s, c i)) := by + simp_rw [le_neg (b := ε), ← Finset.sum_neg_distrib, ← Pi.neg_apply (f := X _), + ← Pi.neg_apply (f := X)] + refine measure_sum_ge_le_of_iIndepFun (X := -X) (μ := μ) ?_ ?_ hε + · exact h_indep.comp _ (fun _ ↦ measurable_neg) + · exact fun i hi ↦ (h_subG i hi).neg + +/-- **Hoeffding inequality** for sub-Gaussian random variables. -/ +lemma measure_sum_range_le_le_of_iIndepFun {X : ℕ → Ω → ℝ} (h_indep : iIndepFun X μ) {c : ℝ≥0} + {n : ℕ} (h_subG : ∀ i < n, HasSubgaussianMGF (X i) c μ) {ε : ℝ} (hε : 0 ≤ ε) : + μ.real {ω | ∑ i ∈ Finset.range n, X i ω ≤ -ε} ≤ exp (-ε ^ 2 / (2 * n * c)) := by + simp_rw [le_neg (b := ε), ← Finset.sum_neg_distrib, ← Pi.neg_apply (f := X _), + ← Pi.neg_apply (f := X)] + refine measure_sum_range_ge_le_of_iIndepFun (X := -X) (μ := μ) ?_ ?_ hε + · exact h_indep.comp _ (fun _ ↦ measurable_neg) + · exact fun i hi ↦ (h_subG i hi).neg + section Sum variable {ι ι' : Type*} {X : ι → Ω → ℝ} {cX : ι → ℝ≥0} {s : Finset ι} @@ -26,7 +53,7 @@ lemma measure_sum_le_sum_le [IsFiniteMeasure μ] (h_indep_sum : IndepFun (fun ω ↦ ∑ i ∈ s, X i ω) (fun ω ↦ ∑ j ∈ t, Y j ω) μ) (h_le : ∑ j ∈ t, μ[Y j] ≤ ∑ i ∈ s, μ[X i]) : μ.real {ω | ∑ i ∈ s, X i ω ≤ ∑ j ∈ t, Y j ω} - ≤ Real.exp (- (∑ j ∈ t, μ[Y j] - ∑ i ∈ s, μ[X i]) ^ 2 + ≤ exp (- (∑ j ∈ t, μ[Y j] - ∑ i ∈ s, μ[X i]) ^ 2 / (2 * (∑ i ∈ s, cX i + ∑ j ∈ t, cY j))) := by have hX_int i (his : i ∈ s) : Integrable (X i) μ := by have h_int := (hX_subG i his).integrable diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index effb2b92..612d5937 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -26,6 +26,14 @@ lemma measurable_pullCount (a : α) (t : ℕ) : Measurable (fun h ↦ pullCount exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop +lemma integrable_pullCount {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] + (a : α) (n : ℕ) : + Integrable (fun ω ↦ (pullCount a n ω : ℝ)) (Bandit.trajMeasure alg ν) := 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 ω + @[fun_prop] lemma measurable_sumRewards (a : α) (t : ℕ) : Measurable (sumRewards a t) := by unfold sumRewards @@ -34,6 +42,11 @@ lemma measurable_sumRewards (a : α) (t : ℕ) : Measurable (sumRewards a t) := exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop +@[fun_prop] +lemma measurable_empMean (a : α) (n : ℕ) : Measurable (empMean a n) := by + unfold empMean + fun_prop + @[fun_prop] lemma measurable_stepsUntil (a : α) (m : ℕ) : Measurable (fun h ↦ stepsUntil a m h) := by classical @@ -68,7 +81,7 @@ lemma measurable_stepsUntil' (a : α) (m : ℕ) : @[fun_prop] lemma measurable_rewardByCount (a : α) (m : ℕ) : - Measurable (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ rewardByCount a m ω.1 ω.2) := by + Measurable (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ rewardByCount a m ω) := by simp_rw [rewardByCount_eq_ite] refine Measurable.ite ?_ ?_ ?_ · exact (measurableSet_singleton _).preimage <| measurable_stepsUntil' a m @@ -269,7 +282,7 @@ lemma reward_cond_stepsUntil [StandardBorelSpace α] [Countable α] [Nonempty α lemma condDistrib_rewardByCount_stepsUntil [Countable α] [StandardBorelSpace α] [Nonempty α] (a : α) (m : ℕ) (hm : m ≠ 0) : - condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil a m ω.1) + condDistrib (rewardByCount a m) (fun ω ↦ stepsUntil a m ω.1) (Bandit.measure alg ν) =ᵐ[(Bandit.measure alg ν).map (fun ω ↦ stepsUntil a m ω.1)] Kernel.const _ (ν a) := by let μ := Bandit.measure alg ν @@ -302,15 +315,15 @@ lemma condDistrib_rewardByCount_stepsUntil [Countable α] [StandardBorelSpace α /-- The reward received at the `m`-th pull of arm `a` has law `ν a`. -/ lemma hasLaw_rewardByCount [Countable α] [StandardBorelSpace α] [Nonempty α] (a : α) (m : ℕ) (hm : m ≠ 0) : - HasLaw (fun ω ↦ rewardByCount a m ω.1 ω.2) (ν a) (Bandit.measure alg ν) where + HasLaw (rewardByCount a m) (ν a) (Bandit.measure alg ν) where map_eq := by have h_condDistrib : - condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil a m ω.1) + condDistrib (rewardByCount a m) (fun ω ↦ stepsUntil a m ω.1) (Bandit.measure alg ν) =ᵐ[(Bandit.measure alg ν).map (fun ω ↦ stepsUntil a m ω.1)] Kernel.const _ (ν a) := condDistrib_rewardByCount_stepsUntil a m hm - calc (Bandit.measure alg ν).map (fun ω ↦ rewardByCount a m ω.1 ω.2) - _ = (condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil a m ω.1) + calc (Bandit.measure alg ν).map (rewardByCount a m) + _ = (condDistrib (rewardByCount a m) (fun ω ↦ stepsUntil a m ω.1) (Bandit.measure alg ν)) ∘ₘ ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil a m ω.1)) := by rw [condDistrib_comp_map (by fun_prop) (by fun_prop)] @@ -325,7 +338,7 @@ lemma hasLaw_rewardByCount [Countable α] [StandardBorelSpace α] [Nonempty α] lemma identDistrib_rewardByCount [Countable α] [StandardBorelSpace α] [Nonempty α] (a : α) (n m : ℕ) (hn : n ≠ 0) (hm : m ≠ 0) : - IdentDistrib (fun ω ↦ rewardByCount a n ω.1 ω.2) (fun ω ↦ rewardByCount a m ω.1 ω.2) + IdentDistrib (rewardByCount a n) (rewardByCount a m) (Bandit.measure alg ν) (Bandit.measure alg ν) where aemeasurable_fst := by fun_prop aemeasurable_snd := by fun_prop @@ -333,35 +346,34 @@ lemma identDistrib_rewardByCount [Countable α] [StandardBorelSpace α] [Nonempt lemma identDistrib_rewardByCount_id [Countable α] [StandardBorelSpace α] [Nonempty α] (a : α) (n : ℕ) (hn : n ≠ 0) : - IdentDistrib (fun ω ↦ rewardByCount a n ω.1 ω.2) id (Bandit.measure alg ν) (ν a) where + IdentDistrib (rewardByCount a n) id (Bandit.measure alg ν) (ν a) where aemeasurable_fst := by fun_prop aemeasurable_snd := Measurable.aemeasurable <| by fun_prop map_eq := by rw [(hasLaw_rewardByCount a n hn).map_eq, Measure.map_id] lemma identDistrib_rewardByCount_eval [Countable α] [StandardBorelSpace α] [Nonempty α] (a : α) (n m : ℕ) (hn : n ≠ 0) : - IdentDistrib (fun ω ↦ rewardByCount a n ω.1 ω.2) (fun ω ↦ ω m a) + IdentDistrib (rewardByCount a n) (fun ω ↦ ω m a) (Bandit.measure alg ν) (Bandit.streamMeasure ν) := (identDistrib_rewardByCount_id a n hn).trans (identDistrib_eval_eval_id_streamMeasure ν m a).symm lemma indepFun_rewardByCount_Iic (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (a : α) (n : ℕ) : - (fun ω ↦ rewardByCount a (n + 1) ω.1 ω.2) ⟂ᵢ[Bandit.measure alg ν] - fun ω (i : Iic n) ↦ rewardByCount a i ω.1 ω.2 := by + (rewardByCount a (n + 1)) ⟂ᵢ[Bandit.measure alg ν] fun ω (i : Iic n) ↦ rewardByCount a i ω := by sorry lemma iIndepFun_rewardByCount' (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (a : α) : - iIndepFun (fun n ω ↦ rewardByCount a n ω.1 ω.2) (Bandit.measure alg ν) := by + iIndepFun (rewardByCount a) (Bandit.measure alg ν) := by rw [iIndepFun_nat_iff_forall_indepFun (by fun_prop)] exact indepFun_rewardByCount_Iic alg ν a lemma iIndepFun_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] : - iIndepFun (fun (p : α × ℕ) ω ↦ rewardByCount p.1 p.2 ω.1 ω.2) (Bandit.measure alg ν) := by + iIndepFun (fun (p : α × ℕ) ↦ rewardByCount p.1 p.2) (Bandit.measure alg ν) := by sorry lemma identDistrib_rewardByCount_stream' [Countable α] [StandardBorelSpace α] [Nonempty α] (a : α) : - IdentDistrib (fun ω n ↦ rewardByCount a (n + 1) ω.1 ω.2) (fun ω n ↦ ω n a) + IdentDistrib (fun ω n ↦ rewardByCount a (n + 1) ω) (fun ω n ↦ ω n a) (Bandit.measure alg ν) (Bandit.streamMeasure ν) := by refine IdentDistrib.pi (fun n ↦ ?_) ?_ ?_ · refine identDistrib_rewardByCount_eval a (n + 1) n (by simp) (ν := ν) @@ -369,11 +381,10 @@ lemma identDistrib_rewardByCount_stream' [Countable α] [StandardBorelSpace α] exact iIndepFun.precomp (g := fun n ↦ n + 1) (fun i j hij ↦ by grind) h_indep · exact iIndepFun_eval_streamMeasure'' ν a -lemma identDistrib_rewardByCount_stream [Countable α] [StandardBorelSpace α] [Nonempty α] - (a : α) : - IdentDistrib (fun ω n ↦ rewardByCount a (n + 1) ω.1 ω.2) (fun ω n ↦ ω.2 n a) - (Bandit.measure alg ν) (Bandit.measure alg ν) := by - refine (identDistrib_rewardByCount_stream' a).trans ?_ +omit [DecidableEq α] [MeasurableSingletonClass α] in +lemma identDistrib_eval_streamMeasure_measure (a : α) : + IdentDistrib (fun ω n ↦ ω n a) (fun ω n ↦ ω.2 n a) + (Bandit.streamMeasure ν) (Bandit.measure alg ν) := by refine IdentDistrib.pi (fun n ↦ ?_) ?_ ?_ · rw [← Bandit.snd_measure alg ν, Measure.snd, identDistrib_map_left_iff (by fun_prop) (by fun_prop) @@ -385,9 +396,41 @@ lemma identDistrib_rewardByCount_stream [Countable α] [StandardBorelSpace α] [ rw [← Measure.snd, Bandit.snd_measure] exact iIndepFun_eval_streamMeasure'' ν a +lemma identDistrib_rewardByCount_stream [Countable α] [StandardBorelSpace α] [Nonempty α] + (a : α) : + IdentDistrib (fun ω n ↦ rewardByCount a (n + 1) ω) (fun ω n ↦ ω.2 n a) + (Bandit.measure alg ν) (Bandit.measure alg ν) := + (identDistrib_rewardByCount_stream' a).trans (identDistrib_eval_streamMeasure_measure a) + lemma indepFun_rewardByCount_of_ne {a b : α} (hab : a ≠ b) : - IndepFun (fun ω s ↦ rewardByCount a s ω.1 ω.2) (fun ω s ↦ rewardByCount b s ω.1 ω.2) + IndepFun (fun ω s ↦ rewardByCount a s ω) (fun ω s ↦ rewardByCount b s ω) (Bandit.measure alg ν) := by sorry +lemma identDistrib_sum_Icc_rewardByCount [Nonempty α] [Countable α] (m : ℕ) (a : α) : + IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount a s ω) + (fun ω ↦ ∑ s ∈ range m, ω.2 s a) (Bandit.measure alg ν) (Bandit.measure alg ν) := by + have h1 (a : α) : + IdentDistrib (fun ω s ↦ rewardByCount a (s + 1) ω) (fun ω s ↦ ω.2 s a) + (Bandit.measure alg ν) (Bandit.measure alg ν) := + identDistrib_rewardByCount_stream a + have h_eq (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) : ∑ s ∈ Icc 1 m, rewardByCount a s ω + = ∑ s ∈ range m, rewardByCount a (s + 1) ω := by + let e : Icc 1 m ≃ range m := + { toFun x := ⟨x - 1, by have h := x.2; simp only [mem_Icc] at h; simp; grind⟩ + invFun x := ⟨x + 1, by + have h := x.2 + simp only [mem_Icc, le_add_iff_nonneg_left, zero_le, true_and, ge_iff_le] + simp only [mem_range] at h + grind⟩ + left_inv x := by have h := x.2; simp only [mem_Icc] at h; grind + right_inv x := by have h := x.2; grind } + rw [← sum_coe_sort (Icc 1 m), ← sum_coe_sort (range m), sum_equiv e] + · simp + · simp only [univ_eq_attach, mem_attach, forall_const, Subtype.forall, mem_Icc, + forall_and_index] + grind + simp_rw [h_eq] + exact IdentDistrib.comp (h1 a) (u := fun p ↦ ∑ s ∈ range m, p s) (by fun_prop) + end Bandits diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index 7fb28c02..c9736eaf 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -60,7 +60,7 @@ lemma fst_stepKernel (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : on `ℕ → α × ℝ`, supported on full trajectories that start with the partial one. -/ noncomputable def traj (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : Kernel (Iic n → α × R) (ℕ → α × R) := - ProbabilityTheory.Kernel.traj (X := fun _ ↦ α × R) (stepKernel alg env) n + Kernel.traj (X := fun _ ↦ α × R) (stepKernel alg env) n deriving IsMarkovKernel /-- Measure on the sequence of actions and observations generated by the algorithm/environment. -/ @@ -148,7 +148,7 @@ lemma adapted_step [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace [SecondCountableTopology α] [OpensMeasurableSpace α] [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] [SecondCountableTopology R] [OpensMeasurableSpace R] : - Adapted (Learning.filtration α R) (fun n ↦ step (α := α) (R := R) n) := + Adapted (Learning.filtration α R) (step (α := α) (R := R)) := fun n ↦ (measurable_step_filtration n).stronglyMeasurable lemma measurable_hist_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (hist n) := by @@ -235,45 +235,4 @@ lemma condDistrib_reward_zero [StandardBorelSpace R] [Nonempty R] have h_action := (hasLaw_action_zero alg env).map_eq rwa [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop), h_action] -section stationaryEnv - -/-- A stationary environment, in which the distribution of the next reward depends only on the last -action. -/ -@[simps] -def stationaryEnv (ν : Kernel α R) [IsMarkovKernel ν] : Environment α R where - feedback _ := ν.prodMkLeft _ - ν0 := ν - -variable {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] - -local notation "𝔓" => trajMeasure alg (stationaryEnv ν) - -lemma condDistrib_reward_stationaryEnv [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] (n : ℕ) : - condDistrib (reward n) (action n) 𝔓 =ᵐ[(𝔓).map (action n)] ν := by - cases n with - | zero => - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] - change (𝔓).map (step 0) = (𝔓).map (action 0) ⊗ₘ ν - rw [(hasLaw_action_zero alg (stationaryEnv ν)).map_eq, - (hasLaw_step_zero alg (stationaryEnv ν)).map_eq, stationaryEnv_ν0] - | succ n => - have h_eq := condDistrib_reward alg (stationaryEnv ν) n - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h_eq ⊢ - have : (𝔓).map (action (n + 1)) = ((𝔓).map (fun x ↦ (hist n x, action (n + 1) x))).snd := by - rw [Measure.snd_map_prodMk (by fun_prop)] - simp only [stationaryEnv_feedback] at h_eq - rw [this, ← Measure.snd_prodAssoc_compProd_prodMkLeft, ← h_eq, - Measure.snd_map_prodMk (by fun_prop), Measure.map_map (by fun_prop) (by fun_prop)] - congr - -lemma condIndepFun_reward_hist_action [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] (n : ℕ) : - CondIndepFun (MeasurableSpace.comap (action (n + 1)) inferInstance) - (measurable_action _).comap_le (reward (n + 1)) (hist n) (𝔓) := - condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft - (by fun_prop) (by fun_prop) (by fun_prop) (condDistrib_reward alg (stationaryEnv ν) n) - -end stationaryEnv - end Learning diff --git a/LeanBandits/SequentialLearning/Deterministic.lean b/LeanBandits/SequentialLearning/Deterministic.lean index d9f2bc46..f2240745 100644 --- a/LeanBandits/SequentialLearning/Deterministic.lean +++ b/LeanBandits/SequentialLearning/Deterministic.lean @@ -45,7 +45,8 @@ lemma action_detAlgorithm_ae_eq [StandardBorelSpace α] [Nonempty α] [StandardB ae_eq_of_condDistrib_eq_deterministic (by fun_prop) (by fun_prop) (by fun_prop) (condDistrib_action (detAlgorithm nextaction h_next action0) env n) -example [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] : +lemma action_detAlgorithm_ae_all_eq + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] : ∀ᵐ h ∂𝔓, action 0 h = action0 ∧ ∀ n, action (n + 1) h = nextaction n (hist n h) := by rw [eventually_and, ae_all_iff] exact ⟨action_zero_detAlgorithm, action_detAlgorithm_ae_eq⟩ diff --git a/LeanBandits/SequentialLearning/StationaryEnv.lean b/LeanBandits/SequentialLearning/StationaryEnv.lean new file mode 100644 index 00000000..995b7f48 --- /dev/null +++ b/LeanBandits/SequentialLearning/StationaryEnv.lean @@ -0,0 +1,59 @@ +/- +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, Paulo Rauber +-/ +import LeanBandits.SequentialLearning.Algorithm + +/-! +# Stationary environments +-/ + +open MeasureTheory ProbabilityTheory Filter Real Finset + +open scoped ENNReal NNReal + +namespace Learning + +variable {α R : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} + +/-- A stationary environment, in which the distribution of the next reward depends only on the last +action. -/ +@[simps] +def stationaryEnv (ν : Kernel α R) [IsMarkovKernel ν] : Environment α R where + feedback _ := ν.prodMkLeft _ + ν0 := ν + +variable {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] + +local notation "𝔓" => trajMeasure alg (stationaryEnv ν) + +/-- The conditional distribution of the reward at time `n` given the action at time `n` is `ν`. -/ +lemma condDistrib_reward_stationaryEnv [StandardBorelSpace α] [Nonempty α] + [StandardBorelSpace R] [Nonempty R] (n : ℕ) : + condDistrib (reward n) (action n) 𝔓 =ᵐ[(𝔓).map (action n)] ν := by + cases n with + | zero => + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] + change (𝔓).map (step 0) = (𝔓).map (action 0) ⊗ₘ ν + rw [(hasLaw_action_zero alg (stationaryEnv ν)).map_eq, + (hasLaw_step_zero alg (stationaryEnv ν)).map_eq, stationaryEnv_ν0] + | succ n => + have h_eq := condDistrib_reward alg (stationaryEnv ν) n + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h_eq ⊢ + have : (𝔓).map (action (n + 1)) = ((𝔓).map (fun x ↦ (hist n x, action (n + 1) x))).snd := by + rw [Measure.snd_map_prodMk (by fun_prop)] + simp only [stationaryEnv_feedback] at h_eq + rw [this, ← Measure.snd_prodAssoc_compProd_prodMkLeft, ← h_eq, + Measure.snd_map_prodMk (by fun_prop), Measure.map_map (by fun_prop) (by fun_prop)] + congr + +/-- The reward at time `n + 1` is conditionally independent of the history up to time `n` +given the action at time `n + 1`. -/ +lemma condIndepFun_reward_hist_action [StandardBorelSpace α] [Nonempty α] + [StandardBorelSpace R] [Nonempty R] (n : ℕ) : + reward (n + 1) ⟂ᵢ[action (n + 1), measurable_action _ ; 𝔓] hist n := + condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft + (by fun_prop) (by fun_prop) (by fun_prop) (condDistrib_reward alg (stationaryEnv ν) n) + +end Learning