diff --git a/LeanBandits.lean b/LeanBandits.lean index 8d2c2a9f..ab3c610e 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -2,4 +2,5 @@ import LeanBandits.AlgorithmBuilding import LeanBandits.Bandit import LeanBandits.ETC import LeanBandits.Regret +import LeanBandits.RewardByCountMeasure import LeanBandits.UCB diff --git a/LeanBandits/Bandit.lean b/LeanBandits/Bandit.lean index db98fee1..eb25efc4 100644 --- a/LeanBandits/Bandit.lean +++ b/LeanBandits/Bandit.lean @@ -110,15 +110,27 @@ def hist (n : ℕ) (h : ℕ → α × R) : Iic n → α × R := fun i ↦ h i @[fun_prop] lemma measurable_arm (n : ℕ) : Measurable (arm n (α := α) (R := R)) := by unfold arm; fun_prop +@[fun_prop] +lemma measurable_arm_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ arm p.1 p.2) := by + refine measurable_from_prod_countable_right fun n ↦ ?_ + simp only + fun_prop + @[fun_prop] lemma measurable_reward (n : ℕ) : Measurable (reward n (α := α) (R := R)) := by unfold reward; fun_prop +@[fun_prop] +lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ reward p.1 p.2) := by + refine measurable_from_prod_countable_right fun n ↦ ?_ + simp only + fun_prop + @[fun_prop] lemma measurable_hist (n : ℕ) : Measurable (hist n (α := α) (R := R)) := by unfold hist; fun_prop /-- Filtration of the bandit process. -/ -def ℱ (α : Type*) [MeasurableSpace α] : +def ℱ (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) := MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R) diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index 27b45389..ecd03cb6 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -17,7 +17,8 @@ open scoped ENNReal NNReal namespace Bandits -variable {α : Type*} {mα : MeasurableSpace α} {ν : Kernel α ℝ} {k : ℕ → α} {t : ℕ} {a : α} +variable {α : Type*} [DecidableEq α] {mα : MeasurableSpace α} {ν : Kernel α ℝ} + {k : ℕ → α} {m n t : ℕ} {a : α} {h : ℕ → α × ℝ} /-! ### Definitions of regret, gaps, pull counts -/ @@ -30,13 +31,17 @@ def regret (ν : Kernel α ℝ) (k : ℕ → α) (t : ℕ) : ℝ := noncomputable def gap (ν : Kernel α ℝ) (a : α) : ℝ := (⨆ i, (ν i)[id]) - (ν a)[id] +omit [DecidableEq α] in lemma gap_nonneg [Fintype α] : 0 ≤ gap ν a := by rw [gap, sub_nonneg] exact le_ciSup (f := fun i ↦ (ν i)[id]) (by simp) a -open Classical in /-- Number of times arm `a` was pulled up to time `t` (excluding `t`). -/ -noncomputable def pullCount (k : ℕ → α) (a : α) (t : ℕ) : ℕ := #(filter (fun s ↦ k s = a) (range t)) +noncomputable def pullCount [DecidableEq α] (k : ℕ → α) (a : α) (t : ℕ) : ℕ := + #(filter (fun s ↦ k s = a) (range t)) + +@[simp] +lemma pullCount_zero (k : ℕ → α) (a : α) : pullCount k a 0 = 0 := by simp [pullCount] open Classical in lemma monotone_pullCount (k : ℕ → α) (a : α) : Monotone (pullCount k a) := @@ -46,14 +51,53 @@ lemma pullCount_eq_pullCount_add_one (k : ℕ → α) (t : ℕ) : pullCount k (k t) (t + 1) = pullCount k (k t) t + 1 := by simp [pullCount, range_succ, filter_insert] -lemma pullCount_eq_pullCount (k : ℕ → α) (a : α) (t : ℕ) (h : k t ≠ a) : - pullCount k a (t + 1) = pullCount k a t := by +lemma pullCount_eq_pullCount (h : k t ≠ a) : pullCount k a (t + 1) = pullCount k a t := by simp [pullCount, range_succ, filter_insert, h] +lemma pullCount_eq_sum (k : ℕ → α) (a : α) (t : ℕ) : + pullCount k a t = ∑ s ∈ range t, if k s = a then 1 else 0 := by simp [pullCount] + /-- Number of steps until arm `a` was pulled exactly `m` times. -/ noncomputable def stepsUntil (k : ℕ → α) (a : α) (m : ℕ) : ℕ∞ := sInf ((↑) '' {s | pullCount k a (s + 1) = m}) +lemma stepsUntil_eq_top_iff : stepsUntil k a m = ⊤ ↔ ∀ s, pullCount k a (s + 1) ≠ m := by + simp [stepsUntil, sInf_eq_top] + +lemma stepsUntil_zero_of_ne (hka : k 0 ≠ a) : stepsUntil k a 0 = 0 := by + unfold stepsUntil + simp_rw [← bot_eq_zero, sInf_eq_bot, bot_eq_zero] + intro n hn + refine ⟨0, ?_, hn⟩ + simp only [Set.mem_image, Set.mem_setOf_eq, Nat.cast_eq_zero, exists_eq_right, zero_add] + rw [← zero_add 1, pullCount_eq_pullCount hka] + simp + +lemma stepsUntil_zero_of_eq (hka : k 0 = a) : stepsUntil k a 0 = ⊤ := by + rw [stepsUntil_eq_top_iff] + suffices 0 < pullCount k a 1 by + intro n hn + refine lt_irrefl 0 ?_ + exact this.trans_le (le_trans (monotone_pullCount _ _ (by omega)) hn.le) + rw [← hka, ← zero_add 1, pullCount_eq_pullCount_add_one] + simp + +lemma stepsUntil_eq_dite (k : ℕ → α) (a : α) (m : ℕ) [Decidable (∃ s, pullCount k a (s + 1) = m)] : + stepsUntil k a m = + if h : ∃ s, pullCount k a (s + 1) = m then (Nat.find h : ℕ∞) else ⊤ := by + unfold stepsUntil + split_ifs with h + · refine le_antisymm ?_ ?_ + · refine sInf_le ?_ + simpa using Nat.find_spec h + · simp only [le_sInf_iff, Set.mem_image, Set.mem_setOf_eq, forall_exists_index, and_imp, + forall_apply_eq_imp_iff₂, Nat.cast_le, Nat.find_le_iff] + exact fun n hn ↦ ⟨n, le_rfl, hn⟩ + · push_neg at h + suffices {s | pullCount k a (s + 1) = m} = ∅ by simp [this] + ext s + simpa using (h s) + lemma stepsUntil_pullCount_le (k : ℕ → α) (a : α) (t : ℕ) : stepsUntil k a (pullCount k a (t + 1)) ≤ t := by rw [stepsUntil] @@ -66,6 +110,26 @@ lemma stepsUntil_pullCount_eq (k : ℕ → α) (t : ℕ) : simpa [stepsUntil, pullCount_eq_pullCount_add_one] exact fun t' h ↦ Nat.le_of_lt_succ ((monotone_pullCount k (k t)).reflect_lt (h ▸ lt_add_one _)) +lemma arm_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount (arm · h) a (s + 1) = m) : + arm (stepsUntil (arm · h) a m).toNat h = a := by + classical + simp only [stepsUntil_eq_dite, h_exists, ↓reduceDIte, ENat.toNat_coe] + have h_spec := Nat.find_spec h_exists + have h_spec' n := Nat.find_min h_exists (m := n) + by_cases h_zero : Nat.find h_exists = 0 + · simp only [h_zero, zero_add, not_lt_zero', IsEmpty.forall_iff, implies_true] at * + by_contra h_ne + rw [← zero_add 1, pullCount_eq_pullCount h_ne] at h_spec + simp only [pullCount_zero] at h_spec + exact hm h_spec.symm + have h_pos : 0 < Nat.find h_exists := Nat.pos_of_ne_zero h_zero + by_contra h_ne + refine h_spec' (Nat.find h_exists - 1) ?_ ?_ + · simp [h_pos] + rw [Nat.sub_add_cancel (by omega)] + rwa [← pullCount_eq_pullCount] + exact h_ne + /-- Reward obtained when pulling arm `a` for the `m`-th time. -/ noncomputable def rewardByCount (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : ℝ := @@ -73,11 +137,18 @@ def rewardByCount (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → | ⊤ => z m a | (n : ℕ) => reward n h +lemma rewardByCount_eq_ite (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : + rewardByCount a m h z = + if (stepsUntil (arm · h) a m) = ⊤ then z m a + else reward (stepsUntil (arm · h) a m).toNat h := by + unfold rewardByCount + cases stepsUntil (arm · h) a m <;> simp + lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : rewardByCount (arm t h) (pullCount (arm · h) (arm t h) t + 1) h z = reward t h := by rw [rewardByCount, ← pullCount_eq_pullCount_add_one, stepsUntil_pullCount_eq] -lemma sum_rewardByCount_eq_sum_reward [DecidableEq α] +lemma sum_rewardByCount_eq_sum_reward (a : α) (t : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : ∑ m ∈ Icc 1 (pullCount (arm · h) a t), rewardByCount a m h z = ∑ s ∈ range t, if (arm s h) = a then (reward s h) else 0 := by @@ -87,7 +158,7 @@ lemma sum_rewardByCount_eq_sum_reward [DecidableEq α] · rw [← hta] at ht ⊢ rw [pullCount_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht] 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] + · rwa [pullCount_eq_pullCount hta, sum_range_succ, if_neg hta, add_zero] lemma sum_pullCount_mul [Fintype α] (k : ℕ → α) (f : α → ℝ) (t : ℕ) : ∑ a, pullCount k a t * f a = ∑ s ∈ range t, f (k s) := by @@ -116,10 +187,12 @@ variable [Fintype α] [Nonempty α] noncomputable def bestArm (ν : Kernel α ℝ) : α := (exists_max_image univ (fun a ↦ (ν a)[id]) (univ_nonempty_iff.mpr inferInstance)).choose +omit [DecidableEq α] in lemma le_bestArm (a : α) : (ν a)[id] ≤ (ν (bestArm ν))[id] := (exists_max_image univ (fun a ↦ (ν a)[id]) (univ_nonempty_iff.mpr inferInstance)).choose_spec.2 _ (mem_univ a) +omit [DecidableEq α] in lemma gap_eq_bestArm_sub : gap ν a = (ν (bestArm ν))[id] - (ν a)[id] := by rw [gap] congr diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean new file mode 100644 index 00000000..8f948cfa --- /dev/null +++ b/LeanBandits/RewardByCountMeasure.lean @@ -0,0 +1,173 @@ +/- +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.Bandit +import LeanBandits.Regret + +/-! # Laws of `stepsUntil` and `rewardByCount` +-/ + +open MeasureTheory ProbabilityTheory Finset +open scoped ENNReal NNReal + +section Aux + +variable {α β γ Ω Ω' : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] + {mα : MeasurableSpace α} {μ : Measure α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} + [MeasurableSpace Ω'] [StandardBorelSpace Ω'] [Nonempty Ω'] + {X : α → β} {Y : α → Ω} {Z : α → Ω'} + +lemma MeasureTheory.Measure.comp_congr {κ η : Kernel α β} (h : ∀ᵐ a ∂μ, κ a = η a) : + κ ∘ₘ μ = η ∘ₘ μ := + Measure.bind_congr_right h + +lemma MeasureTheory.Measure.copy_comp_map (hX : AEMeasurable X μ) : + Kernel.copy β ∘ₘ (μ.map X) = μ.map (fun a ↦ (X a, X a)) := by + rw [Kernel.copy, deterministic_comp_eq_map, AEMeasurable.map_map_of_aemeasurable (by fun_prop) hX] + congr + +lemma MeasureTheory.Measure.compProd_deterministic [SFinite μ] (hX : Measurable X) : + μ ⊗ₘ (Kernel.deterministic X hX) = μ.map (fun a ↦ (a, X a)) := by + rw [Measure.compProd_eq_comp_prod, Kernel.id, Kernel.deterministic_prod_deterministic, + Measure.deterministic_comp_eq_map] + rfl + +lemma ProbabilityTheory.condDistrib_comp_map [IsFiniteMeasure μ] + (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) : + condDistrib Y X μ ∘ₘ (μ.map X) = μ.map Y := by + rw [← Measure.snd_compProd, compProd_map_condDistrib hY, Measure.snd_map_prodMk₀ hX] + +lemma ProbabilityTheory.condDistrib_comp [IsFiniteMeasure μ] + (hX : AEMeasurable X μ) {f : β → Ω} (hf : Measurable f) : + condDistrib (f ∘ X) X μ =ᵐ[μ.map X] Kernel.deterministic f hf := by + rw [← Kernel.compProd_eq_iff, compProd_map_condDistrib (by fun_prop), + Measure.compProd_deterministic, AEMeasurable.map_map_of_aemeasurable (by fun_prop) hX] + congr + +lemma ProbabilityTheory.condDistrib_const [IsFiniteMeasure μ] + (hX : AEMeasurable X μ) (c : Ω) : + condDistrib (fun _ ↦ c) X μ =ᵐ[μ.map X] Kernel.deterministic (fun _ ↦ c) (by fun_prop) := by + have : (fun _ : α ↦ c) = (fun _ : β ↦ c) ∘ X := rfl + conv_lhs => rw [this] + filter_upwards [condDistrib_comp hX (by fun_prop : Measurable (fun _ ↦ c))] with b hb + rw [hb] + +@[fun_prop] +lemma Measurable.coe_nat_enat {f : α → ℕ} (hf : Measurable f) : + Measurable (fun a ↦ (f a : ℕ∞)) := Measurable.comp (by fun_prop) hf + +@[fun_prop] +lemma Measurable.toNat {f : α → ℕ∞} (hf : Measurable f) : Measurable (fun a ↦ (f a).toNat) := + Measurable.comp (by fun_prop) hf + +end Aux + +namespace Bandits + +variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [MeasurableSingletonClass α] + +@[fun_prop] +lemma measurable_pullCount (a : α) (t : ℕ) : Measurable (fun k ↦ pullCount k a t) := by + simp_rw [pullCount_eq_sum] + have h_meas s : Measurable (fun k : ℕ → α ↦ if k s = a then 1 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 k ↦ stepsUntil k a m) := by + classical + have h_union : {k' | ∃ s, pullCount k' a (s + 1) = m} + = ⋃ s : ℕ, {k' | pullCount k' a (s + 1) = m} := by ext; simp + have h_meas_set : MeasurableSet {k' | ∃ s, pullCount k' a (s + 1) = m} := by + rw [h_union] + exact MeasurableSet.iUnion fun s ↦ (measurableSet_singleton _).preimage (by fun_prop) + simp_rw [stepsUntil_eq_dite] + suffices Measurable fun k ↦ if h : k ∈ {k' | ∃ s, pullCount k' a (s + 1) = m} + then (Nat.find h : ℕ∞) else ⊤ by convert this + refine Measurable.dite (s := {k' : ℕ → α | ∃ s, pullCount k' a (s + 1) = m}) + (f := fun x ↦ (Nat.find x.2 : ℕ∞)) (g := fun _ ↦ ⊤) ?_ (by fun_prop) h_meas_set + refine Measurable.coe_nat_enat ?_ + refine measurable_find _ fun k ↦ ?_ + suffices MeasurableSet {x : ℕ → α | pullCount x a (k + 1) = m} by + have : Subtype.val '' + {x : {k' : ℕ → α | ∃ s, pullCount k' a (s + 1) = m} | pullCount x a (k + 1) = m} + = {x : ℕ → α | pullCount x a (k + 1) = m} := by + ext x + simp only [Set.mem_setOf_eq, Set.coe_setOf, Set.mem_image, Subtype.exists, exists_and_left, + exists_prop, exists_eq_right_right, and_iff_left_iff_imp] + exact fun h ↦ ⟨_, h⟩ + refine (MeasurableEmbedding.subtype_coe h_meas_set).measurableSet_image.mp ?_ + rw [this] + exact (measurableSet_singleton _).preimage (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + +lemma measurable_stepsUntil'' (a : α) (m : ℕ) : + Measurable (fun ω : (ℕ → α × ℝ) ↦ stepsUntil (arm · ω) a m) := + (measurable_stepsUntil a m).comp (by fun_prop) + +lemma measurable_stepsUntil' (a : α) (m : ℕ) : + Measurable (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ stepsUntil (arm · ω.1) a m) := + (measurable_stepsUntil'' a m).comp measurable_fst + +@[fun_prop] +lemma measurable_rewardByCount (a : α) (m : ℕ) : + Measurable (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ rewardByCount a m ω.1 ω.2) := by + simp_rw [rewardByCount_eq_ite] + refine Measurable.ite ?_ ?_ ?_ + · exact (measurableSet_singleton _).preimage <| measurable_stepsUntil' a m + · fun_prop + · change Measurable ((fun p : ℕ × (ℕ → α × ℝ) ↦ reward p.1 p.2) + ∘ (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ ((stepsUntil (arm · ω.1) a m).toNat, ω.1))) + have : Measurable fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ + ((stepsUntil (arm · ω.1) a m).toNat, ω.1) := + (measurable_stepsUntil' a m).toNat.prodMk (by fun_prop) + exact Measurable.comp (by fun_prop) this + +lemma condDistrib_rewardByCount_stepsUntil [StandardBorelSpace α] [Nonempty α] + {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (m : ℕ) (hm : m ≠ 0) : + condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil (arm · ω.1) a m) + (Bandit.measure alg ν) + =ᵐ[(Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)] Kernel.const _ (ν a) := by + sorry + +/-- The reward received at the `m`-th pull of arm `a` has law `ν a`. -/ +lemma hasLaw_rewardByCount [StandardBorelSpace α] [Nonempty α] + {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (m : ℕ) (hm : m ≠ 0): + HasLaw (fun ω ↦ rewardByCount a m ω.1 ω.2) (ν a) (Bandit.measure alg ν) where + map_eq := by + have h_condDistrib : + condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil (arm · ω.1) a m) + (Bandit.measure alg ν) + =ᵐ[(Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)] + 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 (arm · ω.1) a m) + (Bandit.measure alg ν)) + ∘ₘ ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)) := by + rw [condDistrib_comp_map (by fun_prop) (by fun_prop)] + _ = (Kernel.const _ (ν a)) + ∘ₘ ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)) := + Measure.comp_congr h_condDistrib + _ = ν a := by + have : IsProbabilityMeasure + ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)) := + isProbabilityMeasure_map (by fun_prop) + simp + +lemma identDistrib_rewardByCount [StandardBorelSpace α] [Nonempty α] + {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (n m : ℕ) + (hn : n ≠ 0) (hm : m ≠ 0) : + IdentDistrib (fun ω ↦ rewardByCount a n ω.1 ω.2) (fun ω ↦ rewardByCount a m ω.1 ω.2) + (Bandit.measure alg ν) (Bandit.measure alg ν) where + aemeasurable_fst := by fun_prop + aemeasurable_snd := by fun_prop + map_eq := by rw [(hasLaw_rewardByCount a n hn).map_eq, (hasLaw_rewardByCount a m hm).map_eq] + +lemma iIndepFun_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] : + iIndepFun (fun (p : α × ℕ) ω ↦ rewardByCount p.1 p.2 ω.1 ω.2) (Bandit.measure alg ν) := by + sorry + +end Bandits