From 4de55d6c6c6e96a7b69a0406cf2a98ec5c139e97 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 4 Sep 2025 13:43:28 +0200 Subject: [PATCH 01/16] measurability lemmas --- LeanBandits.lean | 1 + LeanBandits/Bandit.lean | 6 ++ LeanBandits/Regret.lean | 20 +++++- LeanBandits/RewardByCountMeasure.lean | 87 +++++++++++++++++++++++++++ 4 files changed, 111 insertions(+), 3 deletions(-) create mode 100644 LeanBandits/RewardByCountMeasure.lean 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..5d303db7 100644 --- a/LeanBandits/Bandit.lean +++ b/LeanBandits/Bandit.lean @@ -114,6 +114,12 @@ lemma measurable_arm (n : ℕ) : Measurable (arm n (α := α) (R := R)) := by un 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 diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index 27b45389..08e18e55 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -18,6 +18,7 @@ open scoped ENNReal NNReal namespace Bandits variable {α : Type*} {mα : MeasurableSpace α} {ν : Kernel α ℝ} {k : ℕ → α} {t : ℕ} {a : α} + [DecidableEq α] /-! ### Definitions of regret, gaps, pull counts -/ @@ -30,13 +31,14 @@ 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)) open Classical in lemma monotone_pullCount (k : ℕ → α) (a : α) : Monotone (pullCount k a) := @@ -50,6 +52,9 @@ lemma pullCount_eq_pullCount (k : ℕ → α) (a : α) (t : ℕ) (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}) @@ -73,11 +78,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 @@ -116,10 +128,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..b628284f --- /dev/null +++ b/LeanBandits/RewardByCountMeasure.lean @@ -0,0 +1,87 @@ +/- +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 + +/-! # Reward by count measure +-/ + +open MeasureTheory ProbabilityTheory Finset +open scoped ENNReal NNReal + +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 (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] + (a : α) (m : ℕ) : + Measurable (fun k ↦ stepsUntil k a m) := by + sorry + +lemma measurable_stepsUntil'' (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] + (a : α) (m : ℕ) : + Measurable (fun ω : (ℕ → α × ℝ) ↦ stepsUntil (arm · ω) a m) := + (measurable_stepsUntil alg ν a m).comp (by fun_prop) + +lemma measurable_stepsUntil' (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] + (a : α) (m : ℕ) : + Measurable (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ stepsUntil (arm · ω.1) a m) := + (measurable_stepsUntil'' alg ν a m).comp measurable_fst + +lemma measurable_toNat : Measurable (fun n : ℕ∞ ↦ n.toNat) := + measurable_to_countable fun _ ↦ by simp + +omit [DecidableEq α] [MeasurableSingletonClass α] in +@[fun_prop] +lemma _root_.Measurable.toNat {f : α → ℕ∞} (hf : Measurable f) : Measurable (fun a ↦ (f a).toNat) := + measurable_toNat.comp hf + +@[fun_prop] +lemma measurable_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] + (a : α) (m : ℕ) : + Measurable (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ rewardByCount a m ω.1 ω.2) := by + simp_rw [rewardByCount_eq_ite] + refine Measurable.ite ?_ ?_ ?_ + · exact (measurableSet_singleton _).preimage <| measurable_stepsUntil' alg ν 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' alg ν a m).toNat.prodMk (by fun_prop) + exact Measurable.comp (by fun_prop) this + +/-- The reward received at the `m`-th pull of arm `a` has law `ν a`. -/ +lemma hasLaw_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] + (a : α) (m : ℕ) : + HasLaw (fun ω ↦ rewardByCount a m ω.1 ω.2) (ν a) (Bandit.measure alg ν) where + aemeasurable := (measurable_rewardByCount alg ν a m).aemeasurable + map_eq := by + sorry + +lemma iIndepFun_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] : + iIndepFun (fun (p : α × ℕ) ω ↦ rewardByCount p.1 p.2 ω.1 ω.2) (Bandit.measure alg ν) := by + sorry + +lemma identDistrib_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] + (a : α) (n m : ℕ) : + IdentDistrib (fun ω ↦ rewardByCount a n ω.1 ω.2) (fun ω ↦ rewardByCount a m ω.1 ω.2) + (Bandit.measure alg ν) (Bandit.measure alg ν) where + aemeasurable_fst := (measurable_rewardByCount alg ν a n).aemeasurable + aemeasurable_snd := (measurable_rewardByCount alg ν a m).aemeasurable + map_eq := by + sorry + +end Bandits From c94632ac5ae23248d0fc9e6ac04f6e26c06cc98e Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 4 Sep 2025 13:45:06 +0200 Subject: [PATCH 02/16] minor --- LeanBandits/RewardByCountMeasure.lean | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index b628284f..d3c8458e 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -6,7 +6,7 @@ Authors: Rémy Degenne import LeanBandits.Bandit import LeanBandits.Regret -/-! # Reward by count measure +/-! # Laws of `stepsUntil` and `rewardByCount` -/ open MeasureTheory ProbabilityTheory Finset From 14d0ca37fd7673a07447da899c883d3b61a0041c Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 4 Sep 2025 14:08:46 +0200 Subject: [PATCH 03/16] work towards measurability --- LeanBandits/Regret.lean | 11 ++++ LeanBandits/RewardByCountMeasure.lean | 76 ++++++++++++++++++--------- 2 files changed, 62 insertions(+), 25 deletions(-) diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index 08e18e55..a8cf54a0 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -59,6 +59,17 @@ lemma pullCount_eq_sum (k : ℕ → α) (a : α) (t : ℕ) : noncomputable def stepsUntil (k : ℕ → α) (a : α) (m : ℕ) : ℕ∞ := sInf ((↑) '' {s | pullCount k a (s + 1) = m}) +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 + · sorry + · 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] diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index d3c8458e..2a8b57c1 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -16,6 +16,22 @@ namespace Bandits variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [MeasurableSingletonClass α] +lemma measurable_coe_nat_enat : Measurable (fun n : ℕ ↦ (n : ℕ∞)) := by + fun_prop + +omit [DecidableEq α] [MeasurableSingletonClass α] in +@[fun_prop] +lemma Measurable.coe_nat_enat {f : α → ℕ} (hf : Measurable f) : Measurable (fun a ↦ (f a : ℕ∞)) := + measurable_coe_nat_enat.comp hf + +lemma measurable_toNat : Measurable (fun n : ℕ∞ ↦ n.toNat) := + measurable_to_countable fun _ ↦ by simp + +omit [DecidableEq α] [MeasurableSingletonClass α] in +@[fun_prop] +lemma _root_.Measurable.toNat {f : α → ℕ∞} (hf : Measurable f) : Measurable (fun a ↦ (f a).toNat) := + measurable_toNat.comp hf + @[fun_prop] lemma measurable_pullCount (a : α) (t : ℕ) : Measurable (fun k ↦ pullCount k a t) := by simp_rw [pullCount_eq_sum] @@ -25,49 +41,59 @@ lemma measurable_pullCount (a : α) (t : ℕ) : Measurable (fun k ↦ pullCount fun_prop @[fun_prop] -lemma measurable_stepsUntil (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] - (a : α) (m : ℕ) : - Measurable (fun k ↦ stepsUntil k a m) := by - sorry +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'' (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] - (a : α) (m : ℕ) : +lemma measurable_stepsUntil'' (a : α) (m : ℕ) : Measurable (fun ω : (ℕ → α × ℝ) ↦ stepsUntil (arm · ω) a m) := - (measurable_stepsUntil alg ν a m).comp (by fun_prop) + (measurable_stepsUntil a m).comp (by fun_prop) -lemma measurable_stepsUntil' (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] - (a : α) (m : ℕ) : +lemma measurable_stepsUntil' (a : α) (m : ℕ) : Measurable (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ stepsUntil (arm · ω.1) a m) := - (measurable_stepsUntil'' alg ν a m).comp measurable_fst - -lemma measurable_toNat : Measurable (fun n : ℕ∞ ↦ n.toNat) := - measurable_to_countable fun _ ↦ by simp - -omit [DecidableEq α] [MeasurableSingletonClass α] in -@[fun_prop] -lemma _root_.Measurable.toNat {f : α → ℕ∞} (hf : Measurable f) : Measurable (fun a ↦ (f a).toNat) := - measurable_toNat.comp hf + (measurable_stepsUntil'' a m).comp measurable_fst @[fun_prop] -lemma measurable_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] - (a : α) (m : ℕ) : +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' alg ν a m + · 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' alg ν a m).toNat.prodMk (by fun_prop) + (measurable_stepsUntil' a m).toNat.prodMk (by fun_prop) exact Measurable.comp (by fun_prop) this /-- The reward received at the `m`-th pull of arm `a` has law `ν a`. -/ lemma hasLaw_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (a : α) (m : ℕ) : HasLaw (fun ω ↦ rewardByCount a m ω.1 ω.2) (ν a) (Bandit.measure alg ν) where - aemeasurable := (measurable_rewardByCount alg ν a m).aemeasurable map_eq := by sorry @@ -79,8 +105,8 @@ lemma identDistrib_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [ (a : α) (n m : ℕ) : IdentDistrib (fun ω ↦ rewardByCount a n ω.1 ω.2) (fun ω ↦ rewardByCount a m ω.1 ω.2) (Bandit.measure alg ν) (Bandit.measure alg ν) where - aemeasurable_fst := (measurable_rewardByCount alg ν a n).aemeasurable - aemeasurable_snd := (measurable_rewardByCount alg ν a m).aemeasurable + aemeasurable_fst := by fun_prop + aemeasurable_snd := by fun_prop map_eq := by sorry From 7545754103f34de1e0171e710644df54efeb1f98 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 4 Sep 2025 14:11:47 +0200 Subject: [PATCH 04/16] measurability done --- LeanBandits/Regret.lean | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index a8cf54a0..90bb407c 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -64,7 +64,12 @@ lemma stepsUntil_eq_dite (k : ℕ → α) (a : α) (m : ℕ) [Decidable (∃ s, if h : ∃ s, pullCount k a (s + 1) = m then (Nat.find h : ℕ∞) else ⊤ := by unfold stepsUntil split_ifs with h - · sorry + · 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 From e3bb68ce7e3053e15adae621b535f17c916171b7 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 4 Sep 2025 14:18:38 +0200 Subject: [PATCH 05/16] minor --- LeanBandits/RewardByCountMeasure.lean | 13 ++++++------- 1 file changed, 6 insertions(+), 7 deletions(-) diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 2a8b57c1..5ac0843d 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -91,23 +91,22 @@ lemma measurable_rewardByCount (a : α) (m : ℕ) : exact Measurable.comp (by fun_prop) this /-- The reward received at the `m`-th pull of arm `a` has law `ν a`. -/ -lemma hasLaw_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] +lemma hasLaw_rewardByCount {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (m : ℕ) : HasLaw (fun ω ↦ rewardByCount a m ω.1 ω.2) (ν a) (Bandit.measure alg ν) where map_eq := by sorry -lemma iIndepFun_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] : - iIndepFun (fun (p : α × ℕ) ω ↦ rewardByCount p.1 p.2 ω.1 ω.2) (Bandit.measure alg ν) := by - sorry - lemma identDistrib_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (a : α) (n m : ℕ) : 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 - sorry + map_eq := by rw [(hasLaw_rewardByCount a n).map_eq, (hasLaw_rewardByCount a m).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 From c270762e3eb13ece7382e3490cbba84fe5769c5c Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 4 Sep 2025 14:29:47 +0200 Subject: [PATCH 06/16] remove unnecessary lemmas --- LeanBandits/Regret.lean | 2 +- LeanBandits/RewardByCountMeasure.lean | 12 +++--------- 2 files changed, 4 insertions(+), 10 deletions(-) diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index 90bb407c..e07656b5 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -69,7 +69,7 @@ lemma stepsUntil_eq_dite (k : ℕ → α) (a : α) (m : ℕ) [Decidable (∃ s, 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⟩ + 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 diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 5ac0843d..bb3d4b4f 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -16,21 +16,15 @@ namespace Bandits variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [MeasurableSingletonClass α] -lemma measurable_coe_nat_enat : Measurable (fun n : ℕ ↦ (n : ℕ∞)) := by - fun_prop - omit [DecidableEq α] [MeasurableSingletonClass α] in @[fun_prop] -lemma Measurable.coe_nat_enat {f : α → ℕ} (hf : Measurable f) : Measurable (fun a ↦ (f a : ℕ∞)) := - measurable_coe_nat_enat.comp hf - -lemma measurable_toNat : Measurable (fun n : ℕ∞ ↦ n.toNat) := - measurable_to_countable fun _ ↦ by simp +lemma Measurable.coe_nat_enat {f : α → ℕ} (hf : Measurable f) : + Measurable (fun a ↦ (f a : ℕ∞)) := Measurable.comp (by fun_prop) hf omit [DecidableEq α] [MeasurableSingletonClass α] in @[fun_prop] lemma _root_.Measurable.toNat {f : α → ℕ∞} (hf : Measurable f) : Measurable (fun a ↦ (f a).toNat) := - measurable_toNat.comp hf + Measurable.comp (by fun_prop) hf @[fun_prop] lemma measurable_pullCount (a : α) (t : ℕ) : Measurable (fun k ↦ pullCount k a t) := by From 874024593655ac63561b1c33b1ff7e733aaa9b7a Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 4 Sep 2025 15:04:51 +0200 Subject: [PATCH 07/16] add condDistrib lemma --- LeanBandits/Bandit.lean | 2 +- LeanBandits/RewardByCountMeasure.lean | 26 +++++++++++++++++++++++++- 2 files changed, 26 insertions(+), 2 deletions(-) diff --git a/LeanBandits/Bandit.lean b/LeanBandits/Bandit.lean index 5d303db7..c06af810 100644 --- a/LeanBandits/Bandit.lean +++ b/LeanBandits/Bandit.lean @@ -124,7 +124,7 @@ lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ 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/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index bb3d4b4f..90ff8268 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -84,12 +84,36 @@ lemma measurable_rewardByCount (a : α) (m : ℕ) : (measurable_stepsUntil' a m).toNat.prodMk (by fun_prop) exact Measurable.comp (by fun_prop) this +lemma condDistrib_rewardByCount_stepsUntil {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] + (a : α) (m : ℕ) : + 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 {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (m : ℕ) : HasLaw (fun ω ↦ rewardByCount a m ω.1 ω.2) (ν a) (Bandit.measure alg ν) where map_eq := by - sorry + 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 + 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 + sorry + _ = (Kernel.const _ (ν a)) + ∘ₘ ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)) := by + sorry + _ = ν a := by + have : IsProbabilityMeasure + ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)) := + isProbabilityMeasure_map (by fun_prop) + simp lemma identDistrib_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (a : α) (n m : ℕ) : From b9d24ea937c90ce4d5df830aa2460aa7417fdca9 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 4 Sep 2025 15:15:34 +0200 Subject: [PATCH 08/16] progress --- LeanBandits/RewardByCountMeasure.lean | 25 ++++++++++++++++++++++--- 1 file changed, 22 insertions(+), 3 deletions(-) diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 90ff8268..1bcb3716 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -12,6 +12,25 @@ import LeanBandits.Regret open MeasureTheory ProbabilityTheory Finset open scoped ENNReal NNReal +namespace ProbabilityTheory + +variable {α β Ω F : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] + [Nonempty Ω] [NormedAddCommGroup F] {mα : MeasurableSpace α} {μ : Measure α} [IsFiniteMeasure μ] + {X : α → β} {Y : α → Ω} + {mβ : MeasurableSpace β} {s : Set Ω} {t : Set β} {f : β × Ω → F} + +lemma condDistrib_comp_map (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) : + condDistrib Y X μ ∘ₘ (μ.map X) = μ.map Y := by + rw [← Measure.snd_compProd, compProd_map_condDistrib hY] + rw [Measure.snd_map_prodMk₀ hX] + +omit [IsFiniteMeasure μ] in +lemma Measure.comp_congr {κ η : Kernel α β} (h : ∀ᵐ a ∂μ, κ a = η a) : + κ ∘ₘ μ = η ∘ₘ μ := + Measure.bind_congr_right h + +end ProbabilityTheory + namespace Bandits variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [MeasurableSingletonClass α] @@ -105,10 +124,10 @@ lemma hasLaw_rewardByCount {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMark _ = (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 - sorry + rw [condDistrib_comp_map (by fun_prop) (by fun_prop)] _ = (Kernel.const _ (ν a)) - ∘ₘ ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)) := by - sorry + ∘ₘ ((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)) := From 24e4ce3f7191d251fd81f5979fa1cbffc8c66060 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 4 Sep 2025 15:21:46 +0200 Subject: [PATCH 09/16] minor --- LeanBandits/RewardByCountMeasure.lean | 18 ++++++++---------- 1 file changed, 8 insertions(+), 10 deletions(-) diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 1bcb3716..5c61f001 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -12,24 +12,22 @@ import LeanBandits.Regret open MeasureTheory ProbabilityTheory Finset open scoped ENNReal NNReal -namespace ProbabilityTheory +section Aux -variable {α β Ω F : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] - [Nonempty Ω] [NormedAddCommGroup F] {mα : MeasurableSpace α} {μ : Measure α} [IsFiniteMeasure μ] +variable {α β Ω : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] + {mα : MeasurableSpace α} {μ : Measure α} {mβ : MeasurableSpace β} {X : α → β} {Y : α → Ω} - {mβ : MeasurableSpace β} {s : Set Ω} {t : Set β} {f : β × Ω → F} -lemma condDistrib_comp_map (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) : +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] - rw [Measure.snd_map_prodMk₀ hX] + rw [← Measure.snd_compProd, compProd_map_condDistrib hY, Measure.snd_map_prodMk₀ hX] -omit [IsFiniteMeasure μ] in -lemma Measure.comp_congr {κ η : Kernel α β} (h : ∀ᵐ a ∂μ, κ a = η a) : +lemma MeasureTheory.Measure.comp_congr {κ η : Kernel α β} (h : ∀ᵐ a ∂μ, κ a = η a) : κ ∘ₘ μ = η ∘ₘ μ := Measure.bind_congr_right h -end ProbabilityTheory +end Aux namespace Bandits From 465ad0411ea9ec30c8ca66e8cfec86efead96e51 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 4 Sep 2025 16:30:22 +0200 Subject: [PATCH 10/16] add arm_stepsUntil --- LeanBandits/Regret.lean | 28 ++++++++++++++++++++++++++-- 1 file changed, 26 insertions(+), 2 deletions(-) diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index e07656b5..cffa45e3 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -40,6 +40,9 @@ lemma gap_nonneg [Fintype α] : 0 ≤ gap ν a := by 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) := fun _ _ _ ↦ card_le_card (filter_subset_filter _ (by simpa)) @@ -48,7 +51,7 @@ 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) : +lemma pullCount_eq_pullCount {k : ℕ → α} {a : α} {t : ℕ} (h : k t ≠ a) : pullCount k a (t + 1) = pullCount k a t := by simp [pullCount, range_succ, filter_insert, h] @@ -87,6 +90,27 @@ 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 (a : α) (m : ℕ) (h : ℕ → α × ℝ) (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 : ℕ → α → ℝ) : ℝ := @@ -115,7 +139,7 @@ lemma sum_rewardByCount_eq_sum_reward · 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 From 3b4a555c0637d42424702df1c66b11140bb24af8 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 4 Sep 2025 16:46:59 +0200 Subject: [PATCH 11/16] add stepsUntil_zero --- LeanBandits/Regret.lean | 29 +++++++++++++++++++++++------ 1 file changed, 23 insertions(+), 6 deletions(-) diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index cffa45e3..87657de0 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -17,8 +17,8 @@ open scoped ENNReal NNReal namespace Bandits -variable {α : Type*} {mα : MeasurableSpace α} {ν : Kernel α ℝ} {k : ℕ → α} {t : ℕ} {a : α} - [DecidableEq α] +variable {α : Type*} [DecidableEq α] {mα : MeasurableSpace α} {ν : Kernel α ℝ} + {k : ℕ → α} {m n t : ℕ} {a : α} {h : ℕ → α × ℝ} /-! ### Definitions of regret, gaps, pull counts -/ @@ -51,8 +51,7 @@ 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 : ℕ) : @@ -62,6 +61,25 @@ lemma pullCount_eq_sum (k : ℕ → α) (a : α) (t : ℕ) : noncomputable def stepsUntil (k : ℕ → α) (a : α) (m : ℕ) : ℕ∞ := sInf ((↑) '' {s | pullCount k a (s + 1) = m}) +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 + simp only [stepsUntil, sInf_eq_top, Set.mem_image, Set.mem_setOf_eq, forall_exists_index, and_imp, + forall_apply_eq_imp_iff₂, ENat.coe_ne_top, imp_false] + 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 @@ -90,8 +108,7 @@ 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 (a : α) (m : ℕ) (h : ℕ → α × ℝ) (hm : m ≠ 0) - (h_exists : ∃ s, pullCount (arm · h) a (s + 1) = m) : +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] From b34aba1e592d7fec048e797ceefffb3e0df1bafc Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 5 Sep 2025 09:11:18 +0200 Subject: [PATCH 12/16] add special case for pullCount of 0 --- LeanBandits/Regret.lean | 46 ++++++++++++++++++++++++++++++++--------- 1 file changed, 36 insertions(+), 10 deletions(-) diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index 87657de0..f6ba3554 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -38,24 +38,48 @@ lemma gap_nonneg [Fintype α] : 0 ≤ gap ν a := by /-- Number of times arm `a` was pulled up to time `t` (excluding `t`). -/ noncomputable def pullCount [DecidableEq α] (k : ℕ → α) (a : α) (t : ℕ) : ℕ := - #(filter (fun s ↦ k s = a) (range t)) + if t = 0 then 0 else #(filter (fun s ↦ k s = a) (range t)) @[simp] lemma pullCount_zero (k : ℕ → α) (a : α) : pullCount k a 0 = 0 := by simp [pullCount] +@[simp] +lemma pullCount_one (k : ℕ → α) (a : α) : pullCount k a 1 = if k 0 = a then 1 else 0 := by + simp only [pullCount, one_ne_zero, ↓reduceIte, range_one] + split_ifs with h + · suffices ({0} : Finset ℕ).filter (fun s ↦ k s = a) = {0} by simp [this] + ext x + simp only [mem_filter, mem_singleton, and_iff_left_iff_imp] + rintro rfl + exact h + · simp [h] + open Classical in -lemma monotone_pullCount (k : ℕ → α) (a : α) : Monotone (pullCount k a) := - fun _ _ _ ↦ card_le_card (filter_subset_filter _ (by simpa)) +lemma monotone_pullCount (k : ℕ → α) (a : α) : Monotone (pullCount k a) := by + intro m n hmn + by_cases hm : m = 0 + · simp [hm] + have hn : n ≠ 0 := by grind + simp only [pullCount, hm, ↓reduceIte, hn, ge_iff_le] + exact card_le_card (filter_subset_filter _ (by simpa)) 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 (h : k t ≠ a) : pullCount k a (t + 1) = pullCount k a t := by - simp [pullCount, range_succ, filter_insert, h] + cases t with + | zero => simp + | succ n => + rw [pullCount, range_succ, filter_insert] + simp [pullCount] + +lemma pullCount_eq_pullCount (h : k t ≠ a) : pullCount k a (t + 1) = pullCount k a t := by + simp only [pullCount, Nat.add_eq_zero, one_ne_zero, and_false, ↓reduceIte, range_succ, + filter_insert, h, right_eq_ite_iff, card_eq_zero, filter_eq_empty_iff, mem_range] + rintro rfl + simp 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] + pullCount k a t = if t = 0 then 0 else ∑ 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 @@ -164,8 +188,10 @@ lemma sum_pullCount_mul [Fintype α] (k : ℕ → α) (f : α → ℝ) (t : ℕ) classical simp_rw [card_eq_sum_ones] push_cast - simp_rw [sum_mul, one_mul] - exact sum_fiberwise' (range t) k f + by_cases ht : t = 0 + · simp [ht] + · simp only [ht, ↓reduceIte, sum_mul, one_mul] + exact sum_fiberwise' (range t) k f lemma sum_pullCount [Fintype α] : ∑ a, pullCount k a t = t := by suffices ∑ a, pullCount k a t * (1 : ℝ) = t by norm_cast at this; simpa From f7073c70a26dd2e925d180ebe46e99eb397bb180 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 5 Sep 2025 09:12:37 +0200 Subject: [PATCH 13/16] revert --- LeanBandits/Regret.lean | 46 +++++++++-------------------------------- 1 file changed, 10 insertions(+), 36 deletions(-) diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index f6ba3554..87657de0 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -38,48 +38,24 @@ lemma gap_nonneg [Fintype α] : 0 ≤ gap ν a := by /-- Number of times arm `a` was pulled up to time `t` (excluding `t`). -/ noncomputable def pullCount [DecidableEq α] (k : ℕ → α) (a : α) (t : ℕ) : ℕ := - if t = 0 then 0 else #(filter (fun s ↦ k s = a) (range t)) + #(filter (fun s ↦ k s = a) (range t)) @[simp] lemma pullCount_zero (k : ℕ → α) (a : α) : pullCount k a 0 = 0 := by simp [pullCount] -@[simp] -lemma pullCount_one (k : ℕ → α) (a : α) : pullCount k a 1 = if k 0 = a then 1 else 0 := by - simp only [pullCount, one_ne_zero, ↓reduceIte, range_one] - split_ifs with h - · suffices ({0} : Finset ℕ).filter (fun s ↦ k s = a) = {0} by simp [this] - ext x - simp only [mem_filter, mem_singleton, and_iff_left_iff_imp] - rintro rfl - exact h - · simp [h] - open Classical in -lemma monotone_pullCount (k : ℕ → α) (a : α) : Monotone (pullCount k a) := by - intro m n hmn - by_cases hm : m = 0 - · simp [hm] - have hn : n ≠ 0 := by grind - simp only [pullCount, hm, ↓reduceIte, hn, ge_iff_le] - exact card_le_card (filter_subset_filter _ (by simpa)) +lemma monotone_pullCount (k : ℕ → α) (a : α) : Monotone (pullCount k a) := + fun _ _ _ ↦ card_le_card (filter_subset_filter _ (by simpa)) lemma pullCount_eq_pullCount_add_one (k : ℕ → α) (t : ℕ) : pullCount k (k t) (t + 1) = pullCount k (k t) t + 1 := by - cases t with - | zero => simp - | succ n => - rw [pullCount, range_succ, filter_insert] - simp [pullCount] - -lemma pullCount_eq_pullCount (h : k t ≠ a) : pullCount k a (t + 1) = pullCount k a t := by - simp only [pullCount, Nat.add_eq_zero, one_ne_zero, and_false, ↓reduceIte, range_succ, - filter_insert, h, right_eq_ite_iff, card_eq_zero, filter_eq_empty_iff, mem_range] - rintro rfl - simp + simp [pullCount, range_succ, filter_insert] + +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 = if t = 0 then 0 else ∑ s ∈ range t, if k s = a then 1 else 0 := by - simp [pullCount] + 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 @@ -188,10 +164,8 @@ lemma sum_pullCount_mul [Fintype α] (k : ℕ → α) (f : α → ℝ) (t : ℕ) classical simp_rw [card_eq_sum_ones] push_cast - by_cases ht : t = 0 - · simp [ht] - · simp only [ht, ↓reduceIte, sum_mul, one_mul] - exact sum_fiberwise' (range t) k f + simp_rw [sum_mul, one_mul] + exact sum_fiberwise' (range t) k f lemma sum_pullCount [Fintype α] : ∑ a, pullCount k a t = t := by suffices ∑ a, pullCount k a t * (1 : ℝ) = t by norm_cast at this; simpa From 927430b6f83b8db3066112a07dbe3838a12ea865 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 5 Sep 2025 12:32:13 +0200 Subject: [PATCH 14/16] aux lemmas --- LeanBandits/Bandit.lean | 6 +++ LeanBandits/Regret.lean | 6 ++- LeanBandits/RewardByCountMeasure.lean | 61 +++++++++++++++++++++------ 3 files changed, 58 insertions(+), 15 deletions(-) diff --git a/LeanBandits/Bandit.lean b/LeanBandits/Bandit.lean index c06af810..eb25efc4 100644 --- a/LeanBandits/Bandit.lean +++ b/LeanBandits/Bandit.lean @@ -110,6 +110,12 @@ 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 diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index 87657de0..ecd03cb6 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -61,6 +61,9 @@ lemma pullCount_eq_sum (k : ℕ → α) (a : α) (t : ℕ) : 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] @@ -71,8 +74,7 @@ lemma stepsUntil_zero_of_ne (hka : k 0 ≠ a) : stepsUntil k a 0 = 0 := by simp lemma stepsUntil_zero_of_eq (hka : k 0 = a) : stepsUntil k a 0 = ⊤ := by - simp only [stepsUntil, sInf_eq_top, Set.mem_image, Set.mem_setOf_eq, forall_exists_index, and_imp, - forall_apply_eq_imp_iff₂, ENat.coe_ne_top, imp_false] + rw [stepsUntil_eq_top_iff] suffices 0 < pullCount k a 1 by intro n hn refine lt_irrefl 0 ?_ diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 5c61f001..6b01d009 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -14,18 +14,52 @@ open scoped ENNReal NNReal section Aux -variable {α β Ω : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] +variable {α β Ω Ω' : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] {mα : MeasurableSpace α} {μ : Measure α} {mβ : MeasurableSpace β} - {X : α → β} {Y : α → Ω} + [MeasurableSpace Ω'] [StandardBorelSpace Ω'] [Nonempty Ω'] + {X : α → β} {Y : α → Ω} {Z : α → Ω'} + +lemma MeasureTheory.Measure.comp_congr {κ η : Kernel α β} (h : ∀ᵐ a ∂μ, κ a = η a) : + κ ∘ₘ μ = η ∘ₘ μ := + Measure.bind_congr_right h + +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] + calc (Kernel.id ×ₖ Kernel.deterministic X hX) ∘ₘ μ + _ = (Kernel.deterministic (fun ω ↦ (ω, X ω)) (by fun_prop)) ∘ₘ μ := by + rw [Kernel.id, Kernel.deterministic_prod_deterministic] + simp + _ = μ.map (fun ω ↦ (ω, X ω)) := by + rw [Measure.deterministic_comp_eq_map] 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 MeasureTheory.Measure.comp_congr {κ η : Kernel α β} (h : ∀ᵐ a ∂μ, κ a = η a) : - κ ∘ₘ μ = η ∘ₘ μ := - Measure.bind_congr_right h +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] + +lemma ProbabilityTheory.condDistrib_compProd_condDistrib [IsFiniteMeasure μ] + (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) (hZ : AEMeasurable Z μ) : + (condDistrib Y X μ) ⊗ₖ condDistrib Z (fun a ↦ (X a, Y a)) μ + =ᵐ[μ.map X] condDistrib (fun a ↦ (Y a, Z a)) X μ := by + rw [← Kernel.compProd_eq_iff, compProd_map_condDistrib (by fun_prop)] + rw [Measure.compProd_eq_comp_prod] + sorry end Aux @@ -101,23 +135,23 @@ lemma measurable_rewardByCount (a : α) (m : ℕ) : (measurable_stepsUntil' a m).toNat.prodMk (by fun_prop) exact Measurable.comp (by fun_prop) this -lemma condDistrib_rewardByCount_stepsUntil {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] - (a : α) (m : ℕ) : +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 {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] - (a : α) (m : ℕ) : +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 + 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 ν)) @@ -132,13 +166,14 @@ lemma hasLaw_rewardByCount {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMark isProbabilityMeasure_map (by fun_prop) simp -lemma identDistrib_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] - (a : α) (n m : ℕ) : +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).map_eq, (hasLaw_rewardByCount a m).map_eq] + 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 From 7ccbbe17512e3d9ee66b82f034dcab7271e7df67 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 5 Sep 2025 14:08:36 +0200 Subject: [PATCH 15/16] more aux lemmas --- LeanBandits/RewardByCountMeasure.lean | 27 ++++++++++----------------- 1 file changed, 10 insertions(+), 17 deletions(-) diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 6b01d009..1888dead 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -14,8 +14,8 @@ open scoped ENNReal NNReal section Aux -variable {α β Ω Ω' : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] - {mα : MeasurableSpace α} {μ : Measure α} {mβ : MeasurableSpace β} +variable {α β γ Ω Ω' : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] + {mα : MeasurableSpace α} {μ : Measure α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} [MeasurableSpace Ω'] [StandardBorelSpace Ω'] [Nonempty Ω'] {X : α → β} {Y : α → Ω} {Z : α → Ω'} @@ -23,15 +23,16 @@ lemma MeasureTheory.Measure.comp_congr {κ η : Kernel α β} (h : ∀ᵐ 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] - calc (Kernel.id ×ₖ Kernel.deterministic X hX) ∘ₘ μ - _ = (Kernel.deterministic (fun ω ↦ (ω, X ω)) (by fun_prop)) ∘ₘ μ := by - rw [Kernel.id, Kernel.deterministic_prod_deterministic] - simp - _ = μ.map (fun ω ↦ (ω, X ω)) := by - rw [Measure.deterministic_comp_eq_map] + 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 μ) : @@ -53,14 +54,6 @@ lemma ProbabilityTheory.condDistrib_const [IsFiniteMeasure μ] filter_upwards [condDistrib_comp hX (by fun_prop : Measurable (fun _ ↦ c))] with b hb rw [hb] -lemma ProbabilityTheory.condDistrib_compProd_condDistrib [IsFiniteMeasure μ] - (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) (hZ : AEMeasurable Z μ) : - (condDistrib Y X μ) ⊗ₖ condDistrib Z (fun a ↦ (X a, Y a)) μ - =ᵐ[μ.map X] condDistrib (fun a ↦ (Y a, Z a)) X μ := by - rw [← Kernel.compProd_eq_iff, compProd_map_condDistrib (by fun_prop)] - rw [Measure.compProd_eq_comp_prod] - sorry - end Aux namespace Bandits From 0ac44914f02cfb5104e7fa900512f10903d6fe8d Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 5 Sep 2025 16:06:20 +0200 Subject: [PATCH 16/16] move aux lemmas --- LeanBandits/RewardByCountMeasure.lean | 16 +++++++--------- 1 file changed, 7 insertions(+), 9 deletions(-) diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 1888dead..8f948cfa 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -54,22 +54,20 @@ lemma ProbabilityTheory.condDistrib_const [IsFiniteMeasure μ] filter_upwards [condDistrib_comp hX (by fun_prop : Measurable (fun _ ↦ c))] with b hb rw [hb] -end Aux - -namespace Bandits - -variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [MeasurableSingletonClass α] - -omit [DecidableEq α] [MeasurableSingletonClass α] in @[fun_prop] lemma Measurable.coe_nat_enat {f : α → ℕ} (hf : Measurable f) : Measurable (fun a ↦ (f a : ℕ∞)) := Measurable.comp (by fun_prop) hf -omit [DecidableEq α] [MeasurableSingletonClass α] in @[fun_prop] -lemma _root_.Measurable.toNat {f : α → ℕ∞} (hf : Measurable f) : Measurable (fun a ↦ (f a).toNat) := +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]