From dd690dfafa762c40a7ca6050a8e084ae42cf164f Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Mon, 29 Dec 2025 13:49:31 +0100 Subject: [PATCH 1/3] reorganize --- LeanBandits.lean | 2 - LeanBandits/AlgorithmAndRandomVariables.lean | 66 --- LeanBandits/AlgorithmBuilding.lean | 36 +- LeanBandits/Bandit/Regret.lean | 322 +------------ LeanBandits/BanditAlgorithms/ETC.lean | 6 +- LeanBandits/BanditAlgorithms/UCB.lean | 9 +- LeanBandits/ForMathlib/IdentDistrib.lean | 44 -- LeanBandits/RewardByCountMeasure.lean | 29 +- .../SequentialLearning/FiniteActions.lean | 432 ++++++++++++++++++ 9 files changed, 462 insertions(+), 484 deletions(-) delete mode 100644 LeanBandits/AlgorithmAndRandomVariables.lean delete mode 100644 LeanBandits/ForMathlib/IdentDistrib.lean create mode 100644 LeanBandits/SequentialLearning/FiniteActions.lean diff --git a/LeanBandits.lean b/LeanBandits.lean index f2c8344e..511518a0 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -1,11 +1,9 @@ -import LeanBandits.AlgorithmAndRandomVariables import LeanBandits.AlgorithmBuilding import LeanBandits.Bandit.Bandit import LeanBandits.Bandit.Regret import LeanBandits.BanditAlgorithms.ETC import LeanBandits.BanditAlgorithms.UCB import LeanBandits.ForMathlib.CondDistrib -import LeanBandits.ForMathlib.IdentDistrib import LeanBandits.ForMathlib.IndepFun import LeanBandits.ForMathlib.IndepInfinitePi import LeanBandits.ForMathlib.KernelSub diff --git a/LeanBandits/AlgorithmAndRandomVariables.lean b/LeanBandits/AlgorithmAndRandomVariables.lean deleted file mode 100644 index 04c305ff..00000000 --- a/LeanBandits/AlgorithmAndRandomVariables.lean +++ /dev/null @@ -1,66 +0,0 @@ -/- -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.Regret -import LeanBandits.AlgorithmBuilding - -/-! -# Equalities between definitions of random variables used in bandit algorithms - --/ - -open MeasureTheory ProbabilityTheory Finset -open scoped ENNReal NNReal - -namespace Bandits - -variable {K : ℕ} (hK : 0 < K) - -lemma pullCount_add_one_eq_pullCount' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} : - pullCount a (n + 1) h = pullCount' n (fun i ↦ h i) a := by - rw [pullCount_eq_sum, pullCount'_eq_sum] - unfold arm - rw [Finset.sum_coe_sort (f := fun s ↦ if (h s).1 = a then 1 else 0) (Iic n)] - congr with m - simp only [mem_range, mem_Iic] - grind - -lemma pullCount_eq_pullCount' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} (hn : n ≠ 0) : - pullCount a n h = pullCount' (n - 1) (fun i ↦ h i) a := by - cases n with - | zero => exact absurd rfl hn - | succ n => - rw [pullCount_add_one_eq_pullCount'] - have : n + 1 - 1 = n := by simp - exact this ▸ rfl - -lemma sumRewards_add_one_eq_sumRewards' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} : - sumRewards a (n + 1) h = sumRewards' n (fun i ↦ h i) a := by - unfold sumRewards sumRewards' arm reward - rw [Finset.sum_coe_sort (f := fun s ↦ if (h s).1 = a then (h s).2 else 0) (Iic n)] - congr with m - simp only [mem_range, mem_Iic] - grind - -lemma sumRewards_eq_sumRewards' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} (hn : n ≠ 0) : - sumRewards a n h = sumRewards' (n - 1) (fun i ↦ h i) a := by - cases n with - | zero => exact absurd rfl hn - | succ n => - rw [sumRewards_add_one_eq_sumRewards'] - have : n + 1 - 1 = n := by simp - exact this ▸ rfl - -lemma empMean_add_one_eq_empMean' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} : - empMean a (n + 1) h = empMean' n (fun i ↦ h i) a := by - unfold empMean empMean' - rw [sumRewards_add_one_eq_sumRewards', pullCount_add_one_eq_pullCount'] - -lemma empMean_eq_empMean' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} (hn : n ≠ 0) : - empMean a n h = empMean' (n - 1) (fun i ↦ h i) a := by - unfold empMean empMean' - rw [sumRewards_eq_sumRewards' hn, pullCount_eq_pullCount' hn] - -end Bandits diff --git a/LeanBandits/AlgorithmBuilding.lean b/LeanBandits/AlgorithmBuilding.lean index 971808dd..616ee345 100644 --- a/LeanBandits/AlgorithmBuilding.lean +++ b/LeanBandits/AlgorithmBuilding.lean @@ -3,44 +3,22 @@ Copyright (c) 2025 Rémy Degenne. All rights reserved. Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne -/ -import Mathlib.Analysis.Normed.Ring.Basic -import Mathlib.MeasureTheory.Constructions.BorelSpace.Basic -import Mathlib.Topology.Compactness.Lindelof -import Mathlib.Topology.Metrizable.Basic +import LeanBandits.SequentialLearning.FiniteActions /-! # Tools to build bandit algorithms -/ -open MeasureTheory Finset +open MeasureTheory Finset Learning open scoped ENNReal NNReal namespace Bandits -variable {α : Type*} [DecidableEq α] [MeasurableSpace α] - -/-- 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} - -/-- Sum of rewards of arm `a` up to (and including) time `n`. -/ -noncomputable -def sumRewards' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := - ∑ s, if (h s).1 = a then (h s).2 else 0 - -/-- Empirical mean of arm `a` at time `n`. -/ -noncomputable -def empMean' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := - (sumRewards' n h a) / (pullCount' n h a) - -omit [MeasurableSpace α] in -lemma pullCount'_eq_sum (n : ℕ) (h : Iic n → α × ℝ) (a : α) : - pullCount' n h a = ∑ s : Iic n, if (h s).1 = a then 1 else 0 := by simp [pullCount'] +variable {α : Type*} [DecidableEq α] [MeasurableSpace α] [MeasurableSingletonClass α] @[fun_prop] -lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : - Measurable (fun h ↦ pullCount' n h a) := by +lemma measurable_pullCount' (n : ℕ) (a : α) : + Measurable (fun h : Iic n → α × ℝ ↦ pullCount' n h a) := by simp_rw [pullCount'_eq_sum] have h_meas s : Measurable (fun (h : Iic n → α × ℝ) ↦ if (h s).1 = a then 1 else 0) := by refine Measurable.ite ?_ (by fun_prop) (by fun_prop) @@ -48,7 +26,7 @@ lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : fun_prop @[fun_prop] -lemma measurable_sumRewards' [MeasurableSingletonClass α] (n : ℕ) (a : α) : +lemma measurable_sumRewards' (n : ℕ) (a : α) : Measurable (fun h ↦ sumRewards' n h a) := by simp_rw [sumRewards'] have h_meas s : Measurable (fun (h : Iic n → α × ℝ) ↦ if (h s).1 = a then (h s).2 else 0) := by @@ -57,7 +35,7 @@ lemma measurable_sumRewards' [MeasurableSingletonClass α] (n : ℕ) (a : α) : fun_prop @[fun_prop] -lemma measurable_empMean' [MeasurableSingletonClass α] (n : ℕ) (a : α) : +lemma measurable_empMean' (n : ℕ) (a : α) : Measurable (fun h ↦ empMean' n h a) := by unfold empMean' fun_prop diff --git a/LeanBandits/Bandit/Regret.lean b/LeanBandits/Bandit/Regret.lean index 905c82e0..f13772e6 100644 --- a/LeanBandits/Bandit/Regret.lean +++ b/LeanBandits/Bandit/Regret.lean @@ -7,13 +7,14 @@ import LeanBandits.Bandit.Bandit import Mathlib.Data.ENat.Lattice import Mathlib.Order.CompletePartialOrder import Mathlib.Probability.Martingale.BorelCantelli +import LeanBandits.SequentialLearning.FiniteActions /-! # Regret -/ -open MeasureTheory ProbabilityTheory Filter Real Finset +open MeasureTheory ProbabilityTheory Filter Real Finset Learning open scoped ENNReal NNReal @@ -38,333 +39,20 @@ lemma gap_nonneg [Fintype α] : 0 ≤ gap ν a := by rw [gap, sub_nonneg] exact le_ciSup (f := fun i ↦ (ν i)[id]) (by simp) a -/-- Number of times arm `a` was pulled up to time `t` (excluding `t`). -/ -noncomputable def pullCount [DecidableEq α] (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℕ := - #(filter (fun s ↦ arm s h = a) (range t)) - -@[simp] -lemma pullCount_zero (a : α) (h : ℕ → α × ℝ) : pullCount a 0 h = 0 := by simp [pullCount] - -lemma pullCount_one : pullCount a 1 h = if arm 0 h = a then 1 else 0 := by - simp only [pullCount, range_one] - split_ifs with h - · rw [card_eq_one] - refine ⟨0, by simp [h]⟩ - · simp [h] - -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] - -lemma pullCount_eq_pullCount (ha : arm t h ≠ a) : pullCount a (t + 1) h = pullCount a t h := by - simp [pullCount, range_add_one, filter_insert, ha] - -lemma pullCount_add_one : - pullCount a (t + 1) h = pullCount a t h + if arm t h = a then 1 else 0 := by - split_ifs with h - · rw [← h, pullCount_eq_pullCount_add_one] - · rw [pullCount_eq_pullCount h, add_zero] - -lemma pullCount_eq_sum (a : α) (t : ℕ) (h : ℕ → α × ℝ) : - pullCount a t h = ∑ s ∈ range t, if arm s h = a then 1 else 0 := by simp [pullCount] - -lemma pullCount_le (a : α) (t : ℕ) (h : ℕ → α × ℝ) : pullCount a t h ≤ t := - (card_filter_le _ _).trans_eq (by simp) - -lemma pullCount_congr {h' : ℕ → α × ℝ} (h_eq : ∀ i ≤ n, arm i h = arm i h') : - pullCount a (n + 1) h = pullCount a (n + 1) h' := by - unfold pullCount - congr 1 with s - simp only [mem_filter, mem_range, and_congr_right_iff] - intro hs - 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}) - -lemma stepsUntil_eq_top_iff : stepsUntil a m h = ⊤ ↔ ∀ s, pullCount a (s + 1) h ≠ m := by - simp [stepsUntil, sInf_eq_top] - -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 - rw [← stepsUntil_eq_top_iff] at h_contra - simp [h_contra] at h' - -lemma stepsUntil_zero_of_ne (hka : arm 0 h ≠ a) : stepsUntil a 0 h = 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 : arm 0 h = a) : stepsUntil a 0 h = ⊤ := by - rw [stepsUntil_eq_top_iff] - suffices 0 < pullCount a 1 h 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 (a : α) (m : ℕ) (h : ℕ → α × ℝ) - [Decidable (∃ s, pullCount a (s + 1) h = m)] : - stepsUntil a m h = - if h : ∃ s, pullCount a (s + 1) h = 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 a (s + 1) h = m} = ∅ by simp [this] - ext s - simpa using (h' s) - -lemma stepsUntil_pullCount_le (h : ℕ → α × ℝ) (a : α) (t : ℕ) : - stepsUntil a (pullCount a (t + 1) h) h ≤ t := by - rw [stepsUntil] - exact csInf_le (OrderBot.bddBelow _) ⟨t, rfl, rfl⟩ - -lemma stepsUntil_pullCount_eq (h : ℕ → α × ℝ) (t : ℕ) : - stepsUntil (arm t h) (pullCount (arm t h) (t + 1) h) h = t := by - apply le_antisymm (stepsUntil_pullCount_le h (arm t h) t) - suffices ∀ t', pullCount (arm t h) (t' + 1) h = pullCount (arm t h) t h + 1 → t ≤ t' by - simpa [stepsUntil, pullCount_eq_pullCount_add_one] - exact fun t' h' ↦ Nat.le_of_lt_succ ((monotone_pullCount (arm t h) h).reflect_lt - (h' ▸ lt_add_one _)) - -/-- If we pull arm `a` at time 0, the first time at which it is pulled once is 0. -/ -lemma stepsUntil_one_of_eq (hka : arm 0 h = a) : stepsUntil a 1 h = 0 := by - classical - have h_pull : pullCount a 1 h = 1 := by simp [pullCount_one, hka] - have h_le := stepsUntil_pullCount_le h a 0 - simpa [h_pull] using h_le - -lemma stepsUntil_eq_zero_iff : - stepsUntil a m h = 0 ↔ (m = 0 ∧ arm 0 h ≠ a) ∨ (m = 1 ∧ arm 0 h = a) := by - classical - refine ⟨fun h' ↦ ?_, fun h' ↦ ?_⟩ - · have h_exists : ∃ s, pullCount a (s + 1) h = m := exists_pullCount_eq (by simp [h']) - simp only [stepsUntil_eq_dite, h_exists, ↓reduceDIte, Nat.cast_eq_zero, Nat.find_eq_zero, - zero_add] at h' - rw [pullCount_one] at h' - by_cases hka : arm 0 h = a - · simp only [hka, ↓reduceIte] at h' - simp [h'.symm, hka] - · simp only [hka, ↓reduceIte] at h' - simp [h'.symm, hka] - · cases h' with - | inl h => - rw [h.1, stepsUntil_zero_of_ne h.2] - | inr h => - rw [h.1] - exact stepsUntil_one_of_eq h.2 - lemma arm_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) h = m) : arm (stepsUntil a m h).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 + exact action_stepsUntil hm h_exists lemma arm_eq_of_stepsUntil_eq_coe {ω : ℕ → α × ℝ} (hm : m ≠ 0) (h : stepsUntil a m ω = n) : arm n ω = a := by - have : n = (stepsUntil a m ω).toNat := by simp [h] - rw [this, arm_stepsUntil hm] - exact exists_pullCount_eq (by simp [h]) - -lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount a (s + 1) h = m) : - pullCount a (stepsUntil a m h + 1).toNat h = m := by - classical - have h_eq := stepsUntil_eq_dite a m h - simp only [h_exists, ↓reduceDIte] at h_eq - have h' := Nat.find_spec h_exists - rw [h_eq] - rw [ENat.toNat_add (by simp) (by simp)] - simp only [ENat.toNat_coe, ENat.toNat_one] - exact h' - -lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) h = m) : - pullCount a (stepsUntil a m h).toNat h = m - 1 := by - have h_arm := arm_eq_of_stepsUntil_eq_coe (n := (stepsUntil a m h).toNat) (a := a) (ω := h) hm ?_ - swap; · symm; simpa [stepsUntil_eq_top_iff] - have h_add_one := pullCount_stepsUntil_add_one h_exists - nth_rw 1 [← h_arm] at h_add_one - rw [ENat.toNat_add ?_ (by simp), ENat.toNat_one, pullCount_eq_pullCount_add_one] at h_add_one - swap; · simpa [stepsUntil_eq_top_iff] - grind - -lemma pullCount_lt_of_le_stepsUntil (a : α) {n m : ℕ} (h : ℕ → α × ℝ) - (h_exists : ∃ s, pullCount a (s + 1) h = m) (hn : n < stepsUntil a m h) : - pullCount a (n + 1) h < m := by - classical - have h_eq := stepsUntil_eq_dite a m h - simp only [h_exists, ↓reduceDIte] at h_eq - rw [← ENat.coe_toNat (stepsUntil_ne_top h_exists)] at hn - refine lt_of_le_of_ne ?_ ?_ - · calc pullCount a (n + 1) h - _ ≤ pullCount a (stepsUntil a m h + 1).toNat h := by - refine monotone_pullCount a h ?_ - rw [ENat.toNat_add (stepsUntil_ne_top h_exists) (by simp)] - simp only [ENat.toNat_one, add_le_add_iff_right] - exact mod_cast hn.le - _ = m := pullCount_stepsUntil_add_one h_exists - · refine Nat.find_min h_exists (m := n) ?_ - suffices n < (stepsUntil a m h).toNat by - rwa [h_eq, ENat.toNat_coe] at this - exact mod_cast hn - -lemma pullCount_eq_of_stepsUntil_eq_coe {ω : ℕ → α × ℝ} (hm : m ≠ 0) - (h : stepsUntil a m ω = n) : - pullCount a n ω = m - 1 := by - have : n = (stepsUntil a m ω).toNat := by simp [h] - rw [this, pullCount_stepsUntil hm] - exact exists_pullCount_eq (by simp [h]) - -lemma pullCount_add_one_eq_of_stepsUntil_eq_coe {ω : ℕ → α × ℝ} - (h : stepsUntil a m ω = n) : - pullCount a (n + 1) ω = m := by - have : n + 1 = (stepsUntil a m ω + 1).toNat := by - rw [ENat.toNat_add (by simp [h]) (by simp)]; simp [h] - rw [this, pullCount_stepsUntil_add_one] - exact exists_pullCount_eq (by simp [h]) - -lemma stepsUntil_eq_iff {ω : ℕ → α × ℝ} (n : ℕ) : - stepsUntil a m ω = n ↔ - pullCount a (n + 1) ω = m ∧ (∀ k < n, pullCount a (k + 1) ω < m) := by - refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ - · have h_exists : ∃ s, pullCount a (s + 1) ω = m := exists_pullCount_eq (by simp [h]) - refine ⟨pullCount_add_one_eq_of_stepsUntil_eq_coe h, fun k hk ↦ ?_⟩ - exact pullCount_lt_of_le_stepsUntil a ω h_exists (by rw [h]; exact mod_cast hk) - · classical - rw [stepsUntil_eq_dite a m ω, dif_pos ⟨n, h.1⟩] - simp only [Nat.cast_inj] - rw [Nat.find_eq_iff] - exact ⟨h.1, fun k hk ↦ (h.2 k hk).ne⟩ - -lemma stepsUntil_eq_congr {h' : ℕ → α × ℝ} (h_eq : ∀ i ≤ n, arm i h = arm i h') : - stepsUntil a m h = n ↔ stepsUntil a m h' = n := by - simp_rw [stepsUntil_eq_iff n] - congr! 1 - · rw [pullCount_congr h_eq] - · congr! 3 with k hk - rw [pullCount_congr] - grind - -section SumRewards - -/-- Sum of rewards obtained when pulling arm `a` up to time `t` (exclusive). -/ -def sumRewards (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := - ∑ s ∈ range t, if arm s h = a then reward s h else 0 - -/-- Empirical mean reward obtained when pulling arm `a` up to time `t` (exclusive). -/ -noncomputable -def empMean (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := sumRewards a t h / pullCount a t h - -lemma sumRewards_eq_pullCount_mul_empMean (h_pull : pullCount a t h ≠ 0) : - sumRewards a t h = pullCount a t h * empMean a t h := by unfold empMean; field_simp - -end SumRewards + exact action_eq_of_stepsUntil_eq_coe hm h section RewardByCount -/-- 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 : ℕ) (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) : ℝ := - 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 ω.1 <;> simp - -lemma rewardByCount_of_stepsUntil_eq_top {ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)} - (h : stepsUntil a m ω.1 = ⊤) : - 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 ω = reward n ω.1 := by simp [rewardByCount_eq_ite, h] - -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 : ℕ) (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) : - ∑ 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 ω.1 = a - · rw [← hta] at ht ⊢ - rw [pullCount_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht] - unfold sumRewards - rw [sum_range_succ, if_pos rfl, rewardByCount_pullCount_add_one_eq_reward] - · unfold sumRewards - rwa [pullCount_eq_pullCount hta, sum_range_succ, if_neg hta, add_zero] - -lemma sum_pullCount_mul [Fintype α] (h : ℕ → α × ℝ) (f : α → ℝ) (t : ℕ) : - ∑ a, pullCount a t h * f a = ∑ s ∈ range t, f (arm s h) := by - unfold pullCount - classical - simp_rw [card_eq_sum_ones] - push_cast - simp_rw [sum_mul, one_mul] - exact sum_fiberwise' (range t) (arm · h) f - -lemma sum_pullCount [Fintype α] : ∑ a, pullCount a t h = t := by - suffices ∑ a, pullCount a t h * (1 : ℝ) = t by norm_cast at this; simpa - rw [sum_pullCount_mul] - simp - lemma regret_eq_sum_pullCount_mul_gap [Fintype α] : regret ν t h = ∑ a, pullCount a t h * gap ν a := by - simp_rw [sum_pullCount_mul, regret, gap, sum_sub_distrib] - simp + simp [sum_pullCount_mul, regret, gap, sum_sub_distrib, arm, action] end RewardByCount diff --git a/LeanBandits/BanditAlgorithms/ETC.lean b/LeanBandits/BanditAlgorithms/ETC.lean index a5b06ae9..4fd327c5 100644 --- a/LeanBandits/BanditAlgorithms/ETC.lean +++ b/LeanBandits/BanditAlgorithms/ETC.lean @@ -3,7 +3,7 @@ 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.AlgorithmAndRandomVariables +import LeanBandits.AlgorithmBuilding import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.ForMathlib.SubGaussian import LeanBandits.RewardByCountMeasure @@ -191,7 +191,7 @@ lemma pullCount_add_one_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m =ᵐ[𝔓t] fun ω ↦ pullCount a n ω + {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by simp_rw [Filter.EventuallyEq, pullCount_add_one] filter_upwards [arm_of_ge hm hn] with ω h_arm - congr + congr 3 /-- 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. -/ @@ -238,7 +238,7 @@ lemma identDistrib_aux (m : ℕ) (a b : Fin K) : 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) ?_ ?_ + refine (h2 a).prodMk (h2 b) ?_ ?_ · 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) diff --git a/LeanBandits/BanditAlgorithms/UCB.lean b/LeanBandits/BanditAlgorithms/UCB.lean index 11b28924..675dcd89 100644 --- a/LeanBandits/BanditAlgorithms/UCB.lean +++ b/LeanBandits/BanditAlgorithms/UCB.lean @@ -3,7 +3,6 @@ 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.AlgorithmAndRandomVariables import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.ForMathlib.SubGaussian import LeanBandits.RewardByCountMeasure @@ -23,11 +22,11 @@ namespace Bandits variable {K : ℕ} -- not used -lemma predictatble_pullCount (a : Fin K) : +lemma predictable_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 + have : pullCount a (n + 1) = (fun h : Iic n → Fin K × ℝ ↦ 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] @@ -39,7 +38,7 @@ lemma isStoppingTime_stepsUntil (a : Fin K) (m : ℕ) : 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 + exact predictable_pullCount a n section Algorithm @@ -472,7 +471,7 @@ lemma pullCount_le_add (a : Fin K) (n C : ℕ) (ω : ℕ → Fin K × ℝ) : pullCount a n ω := by rw [pullCount_eq_sum] gcongr with s hs - simp only [Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] + simp only [Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply, arm, action] grind induction n with | zero => simp diff --git a/LeanBandits/ForMathlib/IdentDistrib.lean b/LeanBandits/ForMathlib/IdentDistrib.lean deleted file mode 100644 index 0e542109..00000000 --- a/LeanBandits/ForMathlib/IdentDistrib.lean +++ /dev/null @@ -1,44 +0,0 @@ -/- -Copyright (c) 2025 Rémy Degenne. All rights reserved. -Released under Apache 2.0 license as described in the file LICENSE. -Authors: Rémy Degenne --/ -import Mathlib.Probability.IdentDistrib -import Mathlib.Probability.Independence.InfinitePi - -open MeasureTheory ProbabilityTheory Finset -open scoped ENNReal NNReal - -namespace ProbabilityTheory - -variable {Ω Ω' ι : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} - {μ : Measure Ω} {ν : Measure Ω'} {X Y : Ω → ℝ} {Z W : Ω' → ℝ} - -lemma IdentDistrib.prod [IsFiniteMeasure μ] [IsFiniteMeasure ν] - (hXZ : IdentDistrib X Z μ ν) (hYW : IdentDistrib Y W μ ν) - (hXY : IndepFun X Y μ) (hZW : IndepFun Z W ν) : - IdentDistrib (fun ω ↦ (X ω, Y ω)) (fun ω' ↦ (Z ω', W ω')) μ ν where - aemeasurable_fst := hXZ.aemeasurable_fst.prodMk hYW.aemeasurable_fst - aemeasurable_snd := hXZ.aemeasurable_snd.prodMk hYW.aemeasurable_snd - map_eq := by - rw [(indepFun_iff_map_prod_eq_prod_map_map hXZ.aemeasurable_fst hYW.aemeasurable_fst).mp hXY, - (indepFun_iff_map_prod_eq_prod_map_map hXZ.aemeasurable_snd hYW.aemeasurable_snd).mp hZW, - hXZ.map_eq, hYW.map_eq] - -lemma IdentDistrib.pi [Countable ι] [IsProbabilityMeasure μ] [IsProbabilityMeasure ν] - {X : (i : ι) → Ω → ℝ} {Y : (i : ι) → Ω' → ℝ} - (h : ∀ i, IdentDistrib (X i) (Y i) μ ν) (hX_ind : iIndepFun X μ) (hY_ind : iIndepFun Y ν) : - IdentDistrib (fun ω ↦ (X · ω)) (fun ω ↦ (Y · ω)) μ ν where - aemeasurable_fst := by - rw [aemeasurable_pi_iff] - exact fun i ↦ (h i).aemeasurable_fst - aemeasurable_snd := by - rw [aemeasurable_pi_iff] - exact fun i ↦ (h i).aemeasurable_snd - map_eq := by - rw [(iIndepFun_iff_map_fun_eq_infinitePi_map₀' (fun i ↦ (h i).aemeasurable_fst)).mp hX_ind, - (iIndepFun_iff_map_fun_eq_infinitePi_map₀' (fun i ↦ (h i).aemeasurable_snd)).mp hY_ind] - congr with i - rw [(h i).map_eq] - -end ProbabilityTheory diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index cc8ded25..66826b53 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -5,8 +5,8 @@ Authors: Rémy Degenne -/ import LeanBandits.Bandit.Bandit import LeanBandits.Bandit.Regret -import LeanBandits.ForMathlib.IdentDistrib import LeanBandits.ForMathlib.IndepFun +import Mathlib.Probability.IdentDistribIndep /-! # Laws of `stepsUntil` and `rewardByCount` -/ @@ -18,14 +18,6 @@ namespace Bandits variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [MeasurableSingletonClass α] -@[fun_prop] -lemma measurable_pullCount (a : α) (t : ℕ) : Measurable (fun h ↦ pullCount a t h) := by - simp_rw [pullCount_eq_sum] - have h_meas s : Measurable (fun h : ℕ → α × ℝ ↦ if arm s h = a then 1 else 0) := by - refine Measurable.ite ?_ (by fun_prop) (by fun_prop) - 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 @@ -48,11 +40,12 @@ lemma measurable_empMean (a : α) (n : ℕ) : Measurable (empMean a n) := by fun_prop @[fun_prop] -lemma measurable_stepsUntil (a : α) (m : ℕ) : Measurable (fun h ↦ stepsUntil a m h) := by +lemma measurable_stepsUntil (a : α) (m : ℕ) : + Measurable (fun h : ℕ → α × ℝ ↦ stepsUntil a m h) := by classical - have h_union : {h' | ∃ s, pullCount a (s + 1) h' = m} + have h_union : {h' : ℕ → α × ℝ | ∃ s, pullCount a (s + 1) h' = m} = ⋃ s : ℕ, {h' | pullCount a (s + 1) h' = m} := by ext; simp - have h_meas_set : MeasurableSet {h' | ∃ s, pullCount a (s + 1) h' = m} := by + have h_meas_set : MeasurableSet {h' : ℕ → α × ℝ | ∃ s, pullCount a (s + 1) h' = m} := by rw [h_union] exact MeasurableSet.iUnion fun s ↦ (measurableSet_singleton _).preimage (by fun_prop) simp_rw [stepsUntil_eq_dite] @@ -63,8 +56,8 @@ lemma measurable_stepsUntil (a : α) (m : ℕ) : Measurable (fun h ↦ stepsUnti refine Measurable.coe_nat_enat ?_ refine measurable_find _ fun k ↦ ?_ suffices MeasurableSet {x : ℕ → α × ℝ | pullCount a (k + 1) x = m} by - have : Subtype.val '' - {x : {k' : ℕ → α × ℝ | ∃ s, pullCount a (s + 1) k' = m} | pullCount a (k + 1) x = m} + have : Subtype.val '' {x : {k' : ℕ → α × ℝ | + ∃ s, pullCount a (s + 1) k' = m} | pullCount a (k + 1) (x : ℕ → α × ℝ) = m} = {x : ℕ → α × ℝ | pullCount a (k + 1) x = m} := by ext x simp only [Set.mem_setOf_eq, Set.coe_setOf, Set.mem_image, Subtype.exists, exists_and_left, @@ -175,11 +168,11 @@ lemma measurable_comap_indicator_stepsUntil_eq (a : α) (m n : ℕ) : congr 1 rw [stepsUntil_eq_congr] intro i hin - simp only [arm, mem_Iic, hist, dite_eq_ite, k] + simp only [arm, mem_Iic, hist, dite_eq_ite, k, action] grind lemma measurable_indicator_stepsUntil_eq (a : α) (m n : ℕ) : - Measurable ({ω | stepsUntil a m ω = ↑n}.indicator fun _ ↦ 1) := by + Measurable ({ω : ℕ → α × ℝ | stepsUntil a m ω = ↑n}.indicator fun _ ↦ 1) := by refine (measurable_comap_indicator_stepsUntil_eq a m n).mono ?_ le_rfl refine Measurable.comap_le ?_ fun_prop @@ -207,8 +200,8 @@ lemma condIndepFun_reward_stepsUntil_arm' [StandardBorelSpace α] [Countable α] have h_indep := condIndepFun_self_right (X := reward 0) (Z := arm 0) (mβ := inferInstance) (mβ' := inferInstance) (μ := 𝔓t) (by fun_prop) (by fun_prop) - have : {ω : ℕ → α × ℝ | arm 0 ω = a}.indicator (fun x ↦ 1) - = {b | b = a}.indicator (fun _ ↦ 1) ∘ arm 0 := by ext; simp [Set.indicator] + have : {ω : ℕ → α × ℝ | action 0 ω = a}.indicator (fun x ↦ 1) + = {b | b = a}.indicator (fun _ ↦ 1) ∘ action 0 := by ext; simp [Set.indicator] rw [this] exact h_indep.comp measurable_id (by fun_prop) · simp only [hm1, false_and, Set.setOf_false, Set.indicator_empty] diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean new file mode 100644 index 00000000..73aefdc9 --- /dev/null +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -0,0 +1,432 @@ +/- +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 +import Mathlib.Order.CompletePartialOrder +import Mathlib.Probability.Martingale.BorelCantelli + +/-! +# Bookkeeping definitions for finite action space sequential learning problems + +If the number of actions is finite, it makes sense to define the number of times each action was +chosen, the time at which an action was chosen for the nth time, the value of the reward at that +time, the sum of rewards obtained for each action, the empirical mean reward for each action, etc. + +For each definition that take as arguments a time `t : ℕ`, a history `h : ℕ → α × R`, and possibly +other parameters, we put the time and history at the end in this order, so that the definition can +be seen as a stochastic process indexed by time `t` on the measurable space `ℕ → α × R`. + +-/ + +open MeasureTheory Finset + +namespace Learning + +variable {α R : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} [DecidableEq α] + {a : α} {m n t : ℕ} {h : ℕ → α × R} + +/-- Number of times action `a` was chosen up to time `t` (excluding `t`). -/ +noncomputable +def pullCount (a : α) (t : ℕ) (h : ℕ → α × R) : ℕ := + #(filter (fun s ↦ action s h = a) (range t)) + +/-- 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 → α × R) (a : α) := #{s | (h s).1 = a} + +@[simp] +lemma pullCount_zero (a : α) (h : ℕ → α × R) : pullCount a 0 h = 0 := by simp [pullCount] + +lemma pullCount_one : pullCount a 1 h = if action 0 h = a then 1 else 0 := by + simp only [pullCount, range_one] + split_ifs with h + · rw [card_eq_one] + refine ⟨0, by simp [h]⟩ + · simp [h] + +lemma monotone_pullCount (a : α) (h : ℕ → α × R) : 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 : ℕ → α × R) : + pullCount a n h ≤ pullCount a m h := + monotone_pullCount a h hnm + +lemma pullCount_action_eq_pullCount_add_one (t : ℕ) (h : ℕ → α × R) : + pullCount (action t h) (t + 1) h = pullCount (action t h) t h + 1 := by + simp [pullCount, range_add_one, filter_insert] + +lemma pullCount_eq_pullCount_of_action_ne (ha : action t h ≠ a) : + pullCount a (t + 1) h = pullCount a t h := by + simp [pullCount, range_add_one, filter_insert, ha] + +lemma pullCount_add_one : + pullCount a (t + 1) h = pullCount a t h + if action t h = a then 1 else 0 := by + split_ifs with h + · rw [← h, pullCount_action_eq_pullCount_add_one] + · rw [pullCount_eq_pullCount_of_action_ne h, add_zero] + +lemma pullCount_eq_sum (a : α) (t : ℕ) (h : ℕ → α × R) : + pullCount a t h = ∑ s ∈ range t, if action s h = a then 1 else 0 := by simp [pullCount] + +lemma pullCount'_eq_sum (n : ℕ) (h : Iic n → α × R) (a : α) : + pullCount' n h a = ∑ s : Iic n, if (h s).1 = a then 1 else 0 := by simp [pullCount'] + +lemma pullCount_add_one_eq_pullCount' {n : ℕ} {h : ℕ → α × R} : + pullCount a (n + 1) h = pullCount' n (fun i ↦ h i) a := by + rw [pullCount_eq_sum, pullCount'_eq_sum] + unfold action + rw [Finset.sum_coe_sort (f := fun s ↦ if (h s).1 = a then 1 else 0) (Iic n)] + congr with m + simp only [mem_range, mem_Iic] + grind + +lemma pullCount_eq_pullCount' {n : ℕ} {h : ℕ → α × R} (hn : n ≠ 0) : + pullCount a n h = pullCount' (n - 1) (fun i ↦ h i) a := by + cases n with + | zero => exact absurd rfl hn + | succ n => + rw [pullCount_add_one_eq_pullCount'] + have : n + 1 - 1 = n := by simp + exact this ▸ rfl + +lemma pullCount_le (a : α) (t : ℕ) (h : ℕ → α × R) : pullCount a t h ≤ t := + (card_filter_le _ _).trans_eq (by simp) + +lemma pullCount_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h = action i h') : + pullCount a (n + 1) h = pullCount a (n + 1) h' := by + unfold pullCount + congr 1 with s + simp only [mem_filter, mem_range, and_congr_right_iff] + intro hs + rw [Nat.lt_add_one_iff] at hs + rw [h_eq s hs] + +@[fun_prop] +lemma measurable_pullCount [MeasurableSingletonClass α] (a : α) (t : ℕ) : + Measurable (fun h : ℕ → α × R ↦ pullCount a t h) := by + simp_rw [pullCount_eq_sum] + have h_meas s : Measurable (fun h : ℕ → α × R ↦ if action s h = a then 1 else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + fun_prop + +-- TODO: replace this by leastGE +/-- Number of steps until action `a` was pulled exactly `m` times. -/ +noncomputable +def stepsUntil (a : α) (m : ℕ) (h : ℕ → α × R) : ℕ∞ := sInf ((↑) '' {s | pullCount a (s + 1) h = m}) + +lemma stepsUntil_eq_top_iff : stepsUntil a m h = ⊤ ↔ ∀ s, pullCount a (s + 1) h ≠ m := by + simp [stepsUntil, sInf_eq_top] + +lemma stepsUntil_ne_top (h_exists : ∃ s, pullCount a (s + 1) h = m) : stepsUntil a m h ≠ ⊤ := by + simpa [stepsUntil_eq_top_iff] + +-- todo: this is in ℝ because of the limited def of leastGE +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 + rw [← stepsUntil_eq_top_iff] at h_contra + simp [h_contra] at h' + +lemma stepsUntil_zero_of_ne (hka : action 0 h ≠ a) : stepsUntil a 0 h = 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_of_action_ne hka] + simp + +lemma stepsUntil_zero_of_eq (hka : action 0 h = a) : stepsUntil a 0 h = ⊤ := by + rw [stepsUntil_eq_top_iff] + suffices 0 < pullCount a 1 h 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_action_eq_pullCount_add_one] + simp + +lemma stepsUntil_eq_dite (a : α) (m : ℕ) (h : ℕ → α × R) + [Decidable (∃ s, pullCount a (s + 1) h = m)] : + stepsUntil a m h = + if h : ∃ s, pullCount a (s + 1) h = 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 a (s + 1) h = m} = ∅ by simp [this] + ext s + simpa using (h' s) + +lemma stepsUntil_pullCount_le (h : ℕ → α × R) (a : α) (t : ℕ) : + stepsUntil a (pullCount a (t + 1) h) h ≤ t := by + rw [stepsUntil] + exact csInf_le (OrderBot.bddBelow _) ⟨t, rfl, rfl⟩ + +lemma stepsUntil_pullCount_eq (h : ℕ → α × R) (t : ℕ) : + stepsUntil (action t h) (pullCount (action t h) (t + 1) h) h = t := by + apply le_antisymm (stepsUntil_pullCount_le h (action t h) t) + suffices ∀ t', pullCount (action t h) (t' + 1) h = pullCount (action t h) t h + 1 → t ≤ t' by + simpa [stepsUntil, pullCount_action_eq_pullCount_add_one] + exact fun t' h' ↦ Nat.le_of_lt_succ ((monotone_pullCount (action t h) h).reflect_lt + (h' ▸ lt_add_one _)) + +/-- If we pull action `a` at time 0, the first time at which it is pulled once is 0. -/ +lemma stepsUntil_one_of_eq (hka : action 0 h = a) : stepsUntil a 1 h = 0 := by + classical + have h_pull : pullCount a 1 h = 1 := by simp [pullCount_one, hka] + have h_le := stepsUntil_pullCount_le h a 0 + simpa [h_pull] using h_le + +lemma stepsUntil_eq_zero_iff : + stepsUntil a m h = 0 ↔ (m = 0 ∧ action 0 h ≠ a) ∨ (m = 1 ∧ action 0 h = a) := by + classical + refine ⟨fun h' ↦ ?_, fun h' ↦ ?_⟩ + · have h_exists : ∃ s, pullCount a (s + 1) h = m := exists_pullCount_eq (by simp [h']) + simp only [stepsUntil_eq_dite, h_exists, ↓reduceDIte, Nat.cast_eq_zero, Nat.find_eq_zero, + zero_add] at h' + rw [pullCount_one] at h' + by_cases hka : action 0 h = a + · simp only [hka, ↓reduceIte] at h' + simp [h'.symm, hka] + · simp only [hka, ↓reduceIte] at h' + simp [h'.symm, hka] + · cases h' with + | inl h => + rw [h.1, stepsUntil_zero_of_ne h.2] + | inr h => + rw [h.1] + exact stepsUntil_one_of_eq h.2 + +lemma action_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) h = m) : + action (stepsUntil a m h).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_of_action_ne 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_of_action_ne] + exact h_ne + +lemma action_eq_of_stepsUntil_eq_coe {ω : ℕ → α × R} (hm : m ≠ 0) + (h : stepsUntil a m ω = n) : + action n ω = a := by + have : n = (stepsUntil a m ω).toNat := by simp [h] + rw [this, action_stepsUntil hm] + exact exists_pullCount_eq (by simp [h]) + +lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount a (s + 1) h = m) : + pullCount a (stepsUntil a m h + 1).toNat h = m := by + classical + have h_eq := stepsUntil_eq_dite a m h + simp only [h_exists, ↓reduceDIte] at h_eq + have h' := Nat.find_spec h_exists + rw [h_eq] + rw [ENat.toNat_add (by simp) (by simp)] + simp only [ENat.toNat_coe, ENat.toNat_one] + exact h' + +lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) h = m) : + pullCount a (stepsUntil a m h).toNat h = m - 1 := by + have h_action := action_eq_of_stepsUntil_eq_coe (n := (stepsUntil a m h).toNat) (a := a) (ω := h) + hm ?_ + swap; · symm; simpa [stepsUntil_eq_top_iff] + have h_add_one := pullCount_stepsUntil_add_one h_exists + nth_rw 1 [← h_action] at h_add_one + rw [ENat.toNat_add ?_ (by simp), ENat.toNat_one, pullCount_action_eq_pullCount_add_one] + at h_add_one + swap; · simpa [stepsUntil_eq_top_iff] + grind + +lemma pullCount_lt_of_le_stepsUntil (a : α) {n m : ℕ} (h : ℕ → α × R) + (h_exists : ∃ s, pullCount a (s + 1) h = m) (hn : n < stepsUntil a m h) : + pullCount a (n + 1) h < m := by + classical + have h_eq := stepsUntil_eq_dite a m h + simp only [h_exists, ↓reduceDIte] at h_eq + rw [← ENat.coe_toNat (stepsUntil_ne_top h_exists)] at hn + refine lt_of_le_of_ne ?_ ?_ + · calc pullCount a (n + 1) h + _ ≤ pullCount a (stepsUntil a m h + 1).toNat h := by + refine monotone_pullCount a h ?_ + rw [ENat.toNat_add (stepsUntil_ne_top h_exists) (by simp)] + simp only [ENat.toNat_one, add_le_add_iff_right] + exact mod_cast hn.le + _ = m := pullCount_stepsUntil_add_one h_exists + · refine Nat.find_min h_exists (m := n) ?_ + suffices n < (stepsUntil a m h).toNat by + rwa [h_eq, ENat.toNat_coe] at this + exact mod_cast hn + +lemma pullCount_eq_of_stepsUntil_eq_coe {ω : ℕ → α × R} (hm : m ≠ 0) + (h : stepsUntil a m ω = n) : + pullCount a n ω = m - 1 := by + have : n = (stepsUntil a m ω).toNat := by simp [h] + rw [this, pullCount_stepsUntil hm] + exact exists_pullCount_eq (by simp [h]) + +lemma pullCount_add_one_eq_of_stepsUntil_eq_coe {ω : ℕ → α × R} + (h : stepsUntil a m ω = n) : + pullCount a (n + 1) ω = m := by + have : n + 1 = (stepsUntil a m ω + 1).toNat := by + rw [ENat.toNat_add (by simp [h]) (by simp)]; simp [h] + rw [this, pullCount_stepsUntil_add_one] + exact exists_pullCount_eq (by simp [h]) + +lemma stepsUntil_eq_iff {ω : ℕ → α × R} (n : ℕ) : + stepsUntil a m ω = n ↔ + pullCount a (n + 1) ω = m ∧ (∀ k < n, pullCount a (k + 1) ω < m) := by + refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ + · have h_exists : ∃ s, pullCount a (s + 1) ω = m := exists_pullCount_eq (by simp [h]) + refine ⟨pullCount_add_one_eq_of_stepsUntil_eq_coe h, fun k hk ↦ ?_⟩ + exact pullCount_lt_of_le_stepsUntil a ω h_exists (by rw [h]; exact mod_cast hk) + · classical + rw [stepsUntil_eq_dite a m ω, dif_pos ⟨n, h.1⟩] + simp only [Nat.cast_inj] + rw [Nat.find_eq_iff] + exact ⟨h.1, fun k hk ↦ (h.2 k hk).ne⟩ + +lemma stepsUntil_eq_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h = action i h') : + stepsUntil a m h = n ↔ stepsUntil a m h' = n := by + simp_rw [stepsUntil_eq_iff n] + congr! 1 + · rw [pullCount_congr h_eq] + · congr! 3 with k hk + rw [pullCount_congr] + grind + +section RewardByCount + +/-- Reward obtained when pulling action `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 : ℕ) (ω : (ℕ → α × R) × (ℕ → α → R)) : R := + match (stepsUntil a m ω.1) with + | ⊤ => ω.2 m a + | (n : ℕ) => reward n ω.1 + +lemma rewardByCount_eq_ite (a : α) (m : ℕ) (ω : (ℕ → α × R) × (ℕ → α → R)) : + 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 ω.1 <;> simp + +lemma rewardByCount_of_stepsUntil_eq_top {ω : (ℕ → α × R) × (ℕ → α → R)} + (h : stepsUntil a m ω.1 = ⊤) : + rewardByCount a m ω = ω.2 m a := by simp [rewardByCount_eq_ite, h] + +lemma rewardByCount_of_stepsUntil_eq_coe {ω : (ℕ → α × R) × (ℕ → α → R)} + (h : stepsUntil a m ω.1 = n) : + rewardByCount a m ω = reward n ω.1 := by simp [rewardByCount_eq_ite, h] + +lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (ω : (ℕ → α × R) × (ℕ → α → R)) : + rewardByCount (action t ω.1) (pullCount (action t ω.1) t ω.1 + 1) ω = reward t ω.1 := by + rw [rewardByCount, ← pullCount_action_eq_pullCount_add_one, stepsUntil_pullCount_eq] + +end RewardByCount + +lemma sum_pullCount_mul [Fintype α] [Semiring R] (h : ℕ → α × R) (f : α → R) (t : ℕ) : + ∑ a, pullCount a t h * f a = ∑ s ∈ range t, f (action s h) := by + unfold pullCount + classical + simp_rw [card_eq_sum_ones] + push_cast + simp_rw [sum_mul, one_mul] + exact sum_fiberwise' (range t) (action · h) f + +-- todo: only in ℝ for now +lemma sum_pullCount [Fintype α] {h : ℕ → α × ℝ} : ∑ a, pullCount a t h = t := by + suffices ∑ a, pullCount a t h * (1 : ℝ) = t by norm_cast at this; simpa + rw [sum_pullCount_mul] + simp + +section SumRewards + +/-- Sum of rewards obtained when pulling action `a` up to time `t` (exclusive). -/ +def sumRewards (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := + ∑ s ∈ range t, if action s h = a then reward s h else 0 + +/-- Sum of rewards of arm `a` up to (and including) time `n`. -/ +noncomputable +def sumRewards' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := + ∑ s, if (h s).1 = a then (h s).2 else 0 + +/-- Empirical mean reward obtained when pulling action `a` up to time `t` (exclusive). -/ +noncomputable +def empMean (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := sumRewards a t h / pullCount a t h + +/-- Empirical mean of arm `a` at time `n`. -/ +noncomputable +def empMean' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := + (sumRewards' n h a) / (pullCount' n h a) + +lemma sumRewards_eq_pullCount_mul_empMean {h : ℕ → α × ℝ} (h_pull : pullCount a t h ≠ 0) : + sumRewards a t h = pullCount a t h * empMean a t h := by unfold empMean; field_simp + +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 : action t ω.1 = a + · rw [← hta] at ht ⊢ + rw [pullCount_action_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht] + unfold sumRewards + rw [sum_range_succ, if_pos rfl, rewardByCount_pullCount_add_one_eq_reward] + · unfold sumRewards + rwa [pullCount_eq_pullCount_of_action_ne hta, sum_range_succ, if_neg hta, add_zero] + +lemma sumRewards_add_one_eq_sumRewards' {n : ℕ} {h : ℕ → α × ℝ} : + sumRewards a (n + 1) h = sumRewards' n (fun i ↦ h i) a := by + unfold sumRewards sumRewards' action Learning.reward + rw [Finset.sum_coe_sort (f := fun s ↦ if (h s).1 = a then (h s).2 else 0) (Iic n)] + congr with m + simp only [mem_range, mem_Iic] + grind + +lemma sumRewards_eq_sumRewards' {n : ℕ} {h : ℕ → α × ℝ} (hn : n ≠ 0) : + sumRewards a n h = sumRewards' (n - 1) (fun i ↦ h i) a := by + cases n with + | zero => exact absurd rfl hn + | succ n => + rw [sumRewards_add_one_eq_sumRewards'] + have : n + 1 - 1 = n := by simp + exact this ▸ rfl + +lemma empMean_add_one_eq_empMean' {n : ℕ} {h : ℕ → α × ℝ} : + empMean a (n + 1) h = empMean' n (fun i ↦ h i) a := by + unfold empMean empMean' + rw [sumRewards_add_one_eq_sumRewards', pullCount_add_one_eq_pullCount'] + +lemma empMean_eq_empMean' {n : ℕ} {h : ℕ → α × ℝ} (hn : n ≠ 0) : + empMean a n h = empMean' (n - 1) (fun i ↦ h i) a := by + unfold empMean empMean' + rw [sumRewards_eq_sumRewards' hn, pullCount_eq_pullCount' hn] + +end SumRewards + +end Learning From a2482ae65f278ee19cff73c4b2cb495b1a077e9e Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Mon, 29 Dec 2025 14:30:10 +0100 Subject: [PATCH 2/3] reorganize --- LeanBandits.lean | 1 - LeanBandits/AlgorithmBuilding.lean | 43 ------------------- LeanBandits/BanditAlgorithms/ETC.lean | 2 +- .../SequentialLearning/FiniteActions.lean | 24 +++++++++++ blueprint/lean_decls | 14 +++--- blueprint/src/chapters/bandit.tex | 14 +++--- 6 files changed, 39 insertions(+), 59 deletions(-) delete mode 100644 LeanBandits/AlgorithmBuilding.lean diff --git a/LeanBandits.lean b/LeanBandits.lean index 511518a0..4f77e7a2 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -1,4 +1,3 @@ -import LeanBandits.AlgorithmBuilding import LeanBandits.Bandit.Bandit import LeanBandits.Bandit.Regret import LeanBandits.BanditAlgorithms.ETC diff --git a/LeanBandits/AlgorithmBuilding.lean b/LeanBandits/AlgorithmBuilding.lean deleted file mode 100644 index 616ee345..00000000 --- a/LeanBandits/AlgorithmBuilding.lean +++ /dev/null @@ -1,43 +0,0 @@ -/- -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.SequentialLearning.FiniteActions - -/-! # Tools to build bandit algorithms - --/ - -open MeasureTheory Finset Learning -open scoped ENNReal NNReal - -namespace Bandits - -variable {α : Type*} [DecidableEq α] [MeasurableSpace α] [MeasurableSingletonClass α] - -@[fun_prop] -lemma measurable_pullCount' (n : ℕ) (a : α) : - Measurable (fun h : Iic n → α × ℝ ↦ pullCount' n h a) := by - simp_rw [pullCount'_eq_sum] - have h_meas s : Measurable (fun (h : Iic n → α × ℝ) ↦ if (h s).1 = 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_sumRewards' (n : ℕ) (a : α) : - Measurable (fun h ↦ sumRewards' n h a) := by - simp_rw [sumRewards'] - have h_meas s : Measurable (fun (h : Iic n → α × ℝ) ↦ if (h s).1 = a then (h s).2 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_empMean' (n : ℕ) (a : α) : - Measurable (fun h ↦ empMean' n h a) := by - unfold empMean' - fun_prop - -end Bandits diff --git a/LeanBandits/BanditAlgorithms/ETC.lean b/LeanBandits/BanditAlgorithms/ETC.lean index 4fd327c5..fb2187fe 100644 --- a/LeanBandits/BanditAlgorithms/ETC.lean +++ b/LeanBandits/BanditAlgorithms/ETC.lean @@ -3,7 +3,7 @@ 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.SequentialLearning.FiniteActions import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.ForMathlib.SubGaussian import LeanBandits.RewardByCountMeasure diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index 73aefdc9..b79e4a19 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -114,6 +114,15 @@ lemma measurable_pullCount [MeasurableSingletonClass α] (a : α) (t : ℕ) : exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop +@[fun_prop] +lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : + Measurable (fun h : Iic n → α × R ↦ pullCount' n h a) := by + simp_rw [pullCount'_eq_sum] + have h_meas s : Measurable (fun (h : Iic n → α × R) ↦ if (h s).1 = a then 1 else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + fun_prop + -- TODO: replace this by leastGE /-- Number of steps until action `a` was pulled exactly `m` times. -/ noncomputable @@ -427,6 +436,21 @@ lemma empMean_eq_empMean' {n : ℕ} {h : ℕ → α × ℝ} (hn : n ≠ 0) : unfold empMean empMean' rw [sumRewards_eq_sumRewards' hn, pullCount_eq_pullCount' hn] +@[fun_prop] +lemma measurable_sumRewards' [MeasurableSingletonClass α] (n : ℕ) (a : α) : + Measurable (fun h ↦ sumRewards' n h a) := by + simp_rw [sumRewards'] + have h_meas s : Measurable (fun (h : Iic n → α × ℝ) ↦ if (h s).1 = a then (h s).2 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_empMean' [MeasurableSingletonClass α] (n : ℕ) (a : α) : + Measurable (fun h ↦ empMean' n h a) := by + unfold empMean' + fun_prop + end SumRewards end Learning diff --git a/blueprint/lean_decls b/blueprint/lean_decls index c92a2d0a..ab9b326e 100644 --- a/blueprint/lean_decls +++ b/blueprint/lean_decls @@ -11,19 +11,19 @@ Bandits.Bandit.measure Bandits.arm Bandits.reward Bandits.hist -Bandits.pullCount +Learning.pullCount Bandits.filtration Bandits.condDistrib_reward Bandits.hasLaw_arm_zero Bandits.condDistrib_arm -Bandits.stepsUntil -Bandits.rewardByCount +Learning.stepsUntil +Learning.rewardByCount Bandits.hasLaw_rewardByCount Bandits.iIndepFun_rewardByCount -Bandits.stepsUntil_pullCount_le -Bandits.stepsUntil_pullCount_eq -Bandits.rewardByCount_pullCount_add_one_eq_reward -Bandits.sum_rewardByCount_eq_sumRewards +Learning.stepsUntil_pullCount_le +Learning.stepsUntil_pullCount_eq +Learning.rewardByCount_pullCount_add_one_eq_reward +Learning.sum_rewardByCount_eq_sumRewards Bandits.regret Bandits.gap Bandits.regret_eq_sum_pullCount_mul_gap diff --git a/blueprint/src/chapters/bandit.tex b/blueprint/src/chapters/bandit.tex index f767409a..43f4b164 100644 --- a/blueprint/src/chapters/bandit.tex +++ b/blueprint/src/chapters/bandit.tex @@ -50,7 +50,7 @@ \section{Algorithm, bandit and probability space} \begin{definition}[Pull counts]\label{def:pullCount} \uses{def:armAndReward} \leanok - \lean{Bandits.pullCount} + \lean{Learning.pullCount} For an arm $a \in \mathcal{A}$ and a time $t \in \mathbb{N}$, we denote by $N_{t,a}$ the number of times that arm $a$ has been pulled before time $t$, that is $N_{t,a} = \sum_{s=0}^{t-1} \mathbb{I}\{A_s = a\}$. \end{definition} @@ -145,7 +145,7 @@ \section{Alternative model}\label{sec:alt_model} \begin{definition}\label{def:stepsUntil} \uses{def:pullCount} \leanok - \lean{Bandits.stepsUntil} + \lean{Learning.stepsUntil} For an arm $a \in \mathcal{A}$ and a time $n \in \mathbb{N}$, we denote by $T_{n,a}$ the time at which arm $a$ was pulled for the $n$-th time, that is $T_{n,a} = \min\{s \in \mathbb{N} \mid N_{s+1,a} = n\}$. Note that $T_{n, a}$ can be infinite if the arm is not pulled $n$ times. \end{definition} @@ -154,7 +154,7 @@ \section{Alternative model}\label{sec:alt_model} \begin{definition}\label{def:rewardByCount} \uses{def:stepsUntil} \leanok - \lean{Bandits.rewardByCount} + \lean{Learning.rewardByCount} For $a \in \mathcal{A}$ and $n \in \mathbb{N}$, let $Z_{n,a} \sim \nu(a)$, independent of the bandit interaction and other $Z_{m,b}$. In our probability space $\Omega$, we can take for $Z_{n,a}$ the function $\omega \mapsto \omega_{2,n,a}$. We define $Y_{n, a} = X_{T_{n,a}} \mathbb{I}\{T_{n, a} < \infty\} + Z_{n,a} \mathbb{I}\{T_{n, a} = \infty\}$, the reward received when pulling arm $a$ for the $n$-th time if that time is finite, and equal to $Z_{n,a}$ otherwise. @@ -229,7 +229,7 @@ \section{Alternative model}\label{sec:alt_model} \begin{lemma}\label{lem:stepsUntil_pullCount_le} \uses{def:stepsUntil,def:pullCount} \leanok - \lean{Bandits.stepsUntil_pullCount_le} + \lean{Learning.stepsUntil_pullCount_le} $T_{N_{t+1, a}, a} \le t < \infty$ for all $t \in \mathbb{N}$ and $a \in \mathcal{A}$. \end{lemma} @@ -241,7 +241,7 @@ \section{Alternative model}\label{sec:alt_model} \begin{lemma}\label{lem:stepsUntil_pullCount_eq} \uses{def:stepsUntil,def:pullCount} \leanok - \lean{Bandits.stepsUntil_pullCount_eq} + \lean{Learning.stepsUntil_pullCount_eq} $T_{N_{t+1, A_t}, A_t} = t$ for all $t \in \mathbb{N}$. \end{lemma} @@ -253,7 +253,7 @@ \section{Alternative model}\label{sec:alt_model} \begin{lemma}\label{lem:rewardByCount_pullCount} \uses{def:rewardByCount,def:pullCount} \leanok - \lean{Bandits.rewardByCount_pullCount_add_one_eq_reward} + \lean{Learning.rewardByCount_pullCount_add_one_eq_reward} $Y_{N_{t+1, A_t}, A_t} = X_t$ for all $t \in \mathbb{N}$ and $a \in \mathcal{A}$. \end{lemma} @@ -267,7 +267,7 @@ \section{Alternative model}\label{sec:alt_model} \begin{lemma}\label{lem:sum_rewardByCount} \uses{def:rewardByCount,def:pullCount} \leanok - \lean{Bandits.sum_rewardByCount_eq_sumRewards} + \lean{Learning.sum_rewardByCount_eq_sumRewards} \begin{align*} \sum_{n=1}^{N_{t, a}} Y_{n, a} = \sum_{s=0}^{t-1} \mathbb{I}\{A_s = a\} X_s \: . From 4f2e1d7dedd25f9207b6a25c2d7313fe11f57077 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Mon, 29 Dec 2025 14:55:33 +0100 Subject: [PATCH 3/3] reorganize --- LeanBandits.lean | 1 + LeanBandits/Bandit/Bandit.lean | 2 +- LeanBandits/Bandit/Regret.lean | 3 - LeanBandits/BanditAlgorithms/ETC.lean | 1 - LeanBandits/BanditAlgorithms/UCB.lean | 22 ----- LeanBandits/RewardByCountMeasure.lean | 60 ------------- .../SequentialLearning/FiniteActions.lean | 90 ++++++++++++++++++- 7 files changed, 91 insertions(+), 88 deletions(-) diff --git a/LeanBandits.lean b/LeanBandits.lean index 4f77e7a2..aaea891f 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -13,4 +13,5 @@ import LeanBandits.ForMathlib.Traj import LeanBandits.RewardByCountMeasure import LeanBandits.SequentialLearning.Algorithm import LeanBandits.SequentialLearning.Deterministic +import LeanBandits.SequentialLearning.FiniteActions import LeanBandits.SequentialLearning.StationaryEnv diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 2de56bb2..c31e3bb7 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -3,9 +3,9 @@ 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.ForMathlib.IndepInfinitePi 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 f13772e6..acabf6c8 100644 --- a/LeanBandits/Bandit/Regret.lean +++ b/LeanBandits/Bandit/Regret.lean @@ -4,9 +4,6 @@ Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne, Paulo Rauber -/ import LeanBandits.Bandit.Bandit -import Mathlib.Data.ENat.Lattice -import Mathlib.Order.CompletePartialOrder -import Mathlib.Probability.Martingale.BorelCantelli import LeanBandits.SequentialLearning.FiniteActions /-! diff --git a/LeanBandits/BanditAlgorithms/ETC.lean b/LeanBandits/BanditAlgorithms/ETC.lean index fb2187fe..048cb41c 100644 --- a/LeanBandits/BanditAlgorithms/ETC.lean +++ b/LeanBandits/BanditAlgorithms/ETC.lean @@ -3,7 +3,6 @@ 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.SequentialLearning.FiniteActions import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.ForMathlib.SubGaussian import LeanBandits.RewardByCountMeasure diff --git a/LeanBandits/BanditAlgorithms/UCB.lean b/LeanBandits/BanditAlgorithms/UCB.lean index 675dcd89..0fd5462b 100644 --- a/LeanBandits/BanditAlgorithms/UCB.lean +++ b/LeanBandits/BanditAlgorithms/UCB.lean @@ -3,9 +3,6 @@ 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.ForMathlib.MeasurableArgMax -import LeanBandits.ForMathlib.SubGaussian -import LeanBandits.RewardByCountMeasure import LeanBandits.BanditAlgorithms.ETC /-! @@ -21,25 +18,6 @@ namespace Bandits variable {K : ℕ} --- not used -lemma predictable_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 : Iic n → Fin K × ℝ ↦ 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 predictable_pullCount a n - section Algorithm /-- The exploration bonus of the UCB algorithm, which corresponds to the width of diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 66826b53..bd6782db 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -3,7 +3,6 @@ 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.Bandit import LeanBandits.Bandit.Regret import LeanBandits.ForMathlib.IndepFun import Mathlib.Probability.IdentDistribIndep @@ -26,65 +25,6 @@ lemma integrable_pullCount {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMark simp only [Nat.cast_le] exact pullCount_le a n ω -@[fun_prop] -lemma measurable_sumRewards (a : α) (t : ℕ) : Measurable (sumRewards a t) := by - unfold sumRewards - have h_meas s : Measurable (fun h : ℕ → α × ℝ ↦ if arm s h = a then reward s h else 0) := by - refine Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact (measurableSet_singleton _).preimage (by fun_prop) - fun_prop - -@[fun_prop] -lemma measurable_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 - have h_union : {h' : ℕ → α × ℝ | ∃ s, pullCount a (s + 1) h' = m} - = ⋃ s : ℕ, {h' | pullCount a (s + 1) h' = m} := by ext; simp - have h_meas_set : MeasurableSet {h' : ℕ → α × ℝ | ∃ s, pullCount a (s + 1) h' = 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 a (s + 1) k' = m} - then (Nat.find h : ℕ∞) else ⊤ by convert this - refine Measurable.dite (s := {k' : ℕ → α × ℝ | ∃ s, pullCount a (s + 1) k' = 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 a (k + 1) x = m} by - have : Subtype.val '' {x : {k' : ℕ → α × ℝ | - ∃ s, pullCount a (s + 1) k' = m} | pullCount a (k + 1) (x : ℕ → α × ℝ) = m} - = {x : ℕ → α × ℝ | pullCount a (k + 1) x = 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 a m ω.1) := - (measurable_stepsUntil a m).comp measurable_fst - -@[fun_prop] -lemma measurable_rewardByCount (a : α) (m : ℕ) : - Measurable (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ rewardByCount a m ω) := 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 a m ω.1).toNat, ω.1))) - have : Measurable fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ ((stepsUntil a m ω.1).toNat, ω.1) := - (measurable_stepsUntil' a m).toNat.prodMk (by fun_prop) - exact Measurable.comp (by fun_prop) this - variable {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] omit [DecidableEq α] [MeasurableSingletonClass α] in diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index b79e4a19..0dc602bb 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -38,7 +38,9 @@ noncomputable def pullCount' (n : ℕ) (h : Iic n → α × R) (a : α) := #{s | (h s).1 = a} @[simp] -lemma pullCount_zero (a : α) (h : ℕ → α × R) : pullCount a 0 h = 0 := by simp [pullCount] +lemma pullCount_zero (a : α) : pullCount a 0 (R := R) = 0 := by ext; simp [pullCount] + +lemma pullCount_zero_apply (a : α) (h : ℕ → α × R) : pullCount a 0 h = 0 := by simp lemma pullCount_one : pullCount a 1 h = if action 0 h = a then 1 else 0 := by simp only [pullCount, range_one] @@ -123,6 +125,23 @@ lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop +lemma adapted_pullCount_add_one [MeasurableSingletonClass α] (a : α) : + Adapted (Learning.filtration α R) (fun n ↦ pullCount a (n + 1)) := by + refine fun n ↦ Measurable.stronglyMeasurable ?_ + simp only + have : pullCount a (n + 1) = (fun h : Iic n → α × R ↦ pullCount' n h a) ∘ (hist n) := by + ext + exact pullCount_add_one_eq_pullCount' + rw [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe, this] + exact measurable_comp_comap (hist n) (measurable_pullCount' n a) + +lemma isPredictable_pullCount [MeasurableSingletonClass α] (a : α) : + IsPredictable (Learning.filtration α R) (pullCount a) := by + rw [isPredictable_iff_measurable_add_one] + refine ⟨?_, fun n ↦ (adapted_pullCount_add_one a n).measurable⟩ + simp only [pullCount_zero] + fun_prop + -- TODO: replace this by leastGE /-- Number of steps until action `a` was pulled exactly `m` times. -/ noncomputable @@ -327,6 +346,47 @@ lemma stepsUntil_eq_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h rw [pullCount_congr] grind +lemma isStoppingTime_stepsUntil [MeasurableSingletonClass α] (a : α) (m : ℕ) : + IsStoppingTime (Learning.filtration α ℝ) (stepsUntil a m) := by + rw [stepsUntil_eq_leastGE] + refine Adapted.isStoppingTime_leastGE _ fun n ↦ ?_ + suffices StronglyMeasurable[Learning.filtration α ℝ n] (pullCount a (n + 1)) by fun_prop + exact adapted_pullCount_add_one a n + +-- todo: get this from the stopping time property? +@[fun_prop] +lemma measurable_stepsUntil [MeasurableSingletonClass α] (a : α) (m : ℕ) : + Measurable (fun h : ℕ → α × R ↦ stepsUntil a m h) := by + classical + have h_union : {h' : ℕ → α × R | ∃ s, pullCount a (s + 1) h' = m} + = ⋃ s : ℕ, {h' | pullCount a (s + 1) h' = m} := by ext; simp + have h_meas_set : MeasurableSet {h' : ℕ → α × R | ∃ s, pullCount a (s + 1) h' = 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 a (s + 1) k' = m} + then (Nat.find h : ℕ∞) else ⊤ by convert this + refine Measurable.dite (s := {k' : ℕ → α × R | ∃ s, pullCount a (s + 1) k' = 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 : ℕ → α × R | pullCount a (k + 1) x = m} by + have : Subtype.val '' {x : {k' : ℕ → α × R | + ∃ s, pullCount a (s + 1) k' = m} | pullCount a (k + 1) (x : ℕ → α × R) = m} + = {x : ℕ → α × R | pullCount a (k + 1) x = 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' [MeasurableSingletonClass α] (a : α) (m : ℕ) : + Measurable (fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ stepsUntil a m ω.1) := + (measurable_stepsUntil a m).comp measurable_fst + section RewardByCount /-- Reward obtained when pulling action `a` for the `m`-th time. @@ -356,6 +416,19 @@ lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (ω : (ℕ → α × R rewardByCount (action t ω.1) (pullCount (action t ω.1) t ω.1 + 1) ω = reward t ω.1 := by rw [rewardByCount, ← pullCount_action_eq_pullCount_add_one, stepsUntil_pullCount_eq] +@[fun_prop] +lemma measurable_rewardByCount [MeasurableSingletonClass α] (a : α) (m : ℕ) : + Measurable (fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ rewardByCount a m ω) := by + simp_rw [rewardByCount_eq_ite] + refine Measurable.ite ?_ ?_ ?_ + · exact (measurableSet_singleton _).preimage <| measurable_stepsUntil' a m + · fun_prop + · change Measurable ((fun p : ℕ × (ℕ → α × R) ↦ reward p.1 p.2) + ∘ (fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ ((stepsUntil a m ω.1).toNat, ω.1))) + have : Measurable fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ ((stepsUntil a m ω.1).toNat, ω.1) := + (measurable_stepsUntil' a m).toNat.prodMk (by fun_prop) + exact Measurable.comp (by fun_prop) this + end RewardByCount lemma sum_pullCount_mul [Fintype α] [Semiring R] (h : ℕ → α × R) (f : α → R) (t : ℕ) : @@ -436,6 +509,21 @@ lemma empMean_eq_empMean' {n : ℕ} {h : ℕ → α × ℝ} (hn : n ≠ 0) : unfold empMean empMean' rw [sumRewards_eq_sumRewards' hn, pullCount_eq_pullCount' hn] +@[fun_prop] +lemma measurable_sumRewards [MeasurableSingletonClass α] (a : α) (t : ℕ) : + Measurable (sumRewards a t) := by + unfold sumRewards + have h_meas s : Measurable (fun h : ℕ → α × ℝ ↦ if action s h = a then reward s h else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + fun_prop + +@[fun_prop] +lemma measurable_empMean [MeasurableSingletonClass α] (a : α) (n : ℕ) : + Measurable (empMean a n) := by + unfold empMean + fun_prop + @[fun_prop] lemma measurable_sumRewards' [MeasurableSingletonClass α] (n : ℕ) (a : α) : Measurable (fun h ↦ sumRewards' n h a) := by