Skip to content
Merged
1 change: 1 addition & 0 deletions LeanBandits.lean
Original file line number Diff line number Diff line change
Expand Up @@ -2,4 +2,5 @@ import LeanBandits.AlgorithmBuilding
import LeanBandits.Bandit
import LeanBandits.ETC
import LeanBandits.Regret
import LeanBandits.RewardByCountMeasure
import LeanBandits.UCB
14 changes: 13 additions & 1 deletion LeanBandits/Bandit.lean
Original file line number Diff line number Diff line change
Expand Up @@ -110,15 +110,27 @@ def hist (n : ℕ) (h : ℕ → α × R) : Iic n → α × R := fun i ↦ h i
@[fun_prop]
lemma measurable_arm (n : ℕ) : Measurable (arm n (α := α) (R := R)) := by unfold arm; fun_prop

@[fun_prop]
lemma measurable_arm_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ arm p.1 p.2) := by
refine measurable_from_prod_countable_right fun n ↦ ?_
simp only
fun_prop

@[fun_prop]
lemma measurable_reward (n : ℕ) : Measurable (reward n (α := α) (R := R)) := by
unfold reward; fun_prop

@[fun_prop]
lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ reward p.1 p.2) := by
refine measurable_from_prod_countable_right fun n ↦ ?_
simp only
fun_prop

@[fun_prop]
lemma measurable_hist (n : ℕ) : Measurable (hist n (α := α) (R := R)) := by unfold hist; fun_prop

/-- Filtration of the bandit process. -/
def ℱ (α : Type*) [MeasurableSpace α] :
def ℱ (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] :
Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) :=
MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R)

Expand Down
87 changes: 80 additions & 7 deletions LeanBandits/Regret.lean
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,8 @@ open scoped ENNReal NNReal

namespace Bandits

variable {α : Type*} {mα : MeasurableSpace α} {ν : Kernel α ℝ} {k : ℕ → α} {t : ℕ} {a : α}
variable {α : Type*} [DecidableEq α] {mα : MeasurableSpace α} {ν : Kernel α ℝ}
{k : ℕ → α} {m n t : ℕ} {a : α} {h : ℕ → α × ℝ}

/-! ### Definitions of regret, gaps, pull counts -/

Expand All @@ -30,13 +31,17 @@ def regret (ν : Kernel α ℝ) (k : ℕ → α) (t : ℕ) : ℝ :=
noncomputable
def gap (ν : Kernel α ℝ) (a : α) : ℝ := (⨆ i, (ν i)[id]) - (ν a)[id]

omit [DecidableEq α] in
lemma gap_nonneg [Fintype α] : 0 ≤ gap ν a := by
rw [gap, sub_nonneg]
exact le_ciSup (f := fun i ↦ (ν i)[id]) (by simp) a

open Classical in
/-- Number of times arm `a` was pulled up to time `t` (excluding `t`). -/
noncomputable def pullCount (k : ℕ → α) (a : α) (t : ℕ) : ℕ := #(filter (fun s ↦ k s = a) (range t))
noncomputable def pullCount [DecidableEq α] (k : ℕ → α) (a : α) (t : ℕ) : ℕ :=
#(filter (fun s ↦ k s = a) (range t))

@[simp]
lemma pullCount_zero (k : ℕ → α) (a : α) : pullCount k a 0 = 0 := by simp [pullCount]

open Classical in
lemma monotone_pullCount (k : ℕ → α) (a : α) : Monotone (pullCount k a) :=
Expand All @@ -46,14 +51,53 @@ lemma pullCount_eq_pullCount_add_one (k : ℕ → α) (t : ℕ) :
pullCount k (k t) (t + 1) = pullCount k (k t) t + 1 := by
simp [pullCount, range_succ, filter_insert]

lemma pullCount_eq_pullCount (k : ℕ → α) (a : α) (t : ℕ) (h : k t ≠ a) :
pullCount k a (t + 1) = pullCount k a t := by
lemma pullCount_eq_pullCount (h : k t ≠ a) : pullCount k a (t + 1) = pullCount k a t := by
simp [pullCount, range_succ, filter_insert, h]

lemma pullCount_eq_sum (k : ℕ → α) (a : α) (t : ℕ) :
pullCount k a t = ∑ s ∈ range t, if k s = a then 1 else 0 := by simp [pullCount]

/-- Number of steps until arm `a` was pulled exactly `m` times. -/
noncomputable
def stepsUntil (k : ℕ → α) (a : α) (m : ℕ) : ℕ∞ := sInf ((↑) '' {s | pullCount k a (s + 1) = m})

lemma stepsUntil_eq_top_iff : stepsUntil k a m = ⊤ ↔ ∀ s, pullCount k a (s + 1) ≠ m := by
simp [stepsUntil, sInf_eq_top]

lemma stepsUntil_zero_of_ne (hka : k 0 ≠ a) : stepsUntil k a 0 = 0 := by
unfold stepsUntil
simp_rw [← bot_eq_zero, sInf_eq_bot, bot_eq_zero]
intro n hn
refine ⟨0, ?_, hn⟩
simp only [Set.mem_image, Set.mem_setOf_eq, Nat.cast_eq_zero, exists_eq_right, zero_add]
rw [← zero_add 1, pullCount_eq_pullCount hka]
simp

lemma stepsUntil_zero_of_eq (hka : k 0 = a) : stepsUntil k a 0 = ⊤ := by
rw [stepsUntil_eq_top_iff]
suffices 0 < pullCount k a 1 by
intro n hn
refine lt_irrefl 0 ?_
exact this.trans_le (le_trans (monotone_pullCount _ _ (by omega)) hn.le)
rw [← hka, ← zero_add 1, pullCount_eq_pullCount_add_one]
simp

lemma stepsUntil_eq_dite (k : ℕ → α) (a : α) (m : ℕ) [Decidable (∃ s, pullCount k a (s + 1) = m)] :
stepsUntil k a m =
if h : ∃ s, pullCount k a (s + 1) = m then (Nat.find h : ℕ∞) else ⊤ := by
unfold stepsUntil
split_ifs with h
· refine le_antisymm ?_ ?_
· refine sInf_le ?_
simpa using Nat.find_spec h
· simp only [le_sInf_iff, Set.mem_image, Set.mem_setOf_eq, forall_exists_index, and_imp,
forall_apply_eq_imp_iff₂, Nat.cast_le, Nat.find_le_iff]
exact fun n hn ↦ ⟨n, le_rfl, hn⟩
· push_neg at h
suffices {s | pullCount k a (s + 1) = m} = ∅ by simp [this]
ext s
simpa using (h s)

lemma stepsUntil_pullCount_le (k : ℕ → α) (a : α) (t : ℕ) :
stepsUntil k a (pullCount k a (t + 1)) ≤ t := by
rw [stepsUntil]
Expand All @@ -66,18 +110,45 @@ lemma stepsUntil_pullCount_eq (k : ℕ → α) (t : ℕ) :
simpa [stepsUntil, pullCount_eq_pullCount_add_one]
exact fun t' h ↦ Nat.le_of_lt_succ ((monotone_pullCount k (k t)).reflect_lt (h ▸ lt_add_one _))

lemma arm_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount (arm · h) a (s + 1) = m) :
arm (stepsUntil (arm · h) a m).toNat h = a := by
classical
simp only [stepsUntil_eq_dite, h_exists, ↓reduceDIte, ENat.toNat_coe]
have h_spec := Nat.find_spec h_exists
have h_spec' n := Nat.find_min h_exists (m := n)
by_cases h_zero : Nat.find h_exists = 0
· simp only [h_zero, zero_add, not_lt_zero', IsEmpty.forall_iff, implies_true] at *
by_contra h_ne
rw [← zero_add 1, pullCount_eq_pullCount h_ne] at h_spec
simp only [pullCount_zero] at h_spec
exact hm h_spec.symm
have h_pos : 0 < Nat.find h_exists := Nat.pos_of_ne_zero h_zero
by_contra h_ne
refine h_spec' (Nat.find h_exists - 1) ?_ ?_
· simp [h_pos]
rw [Nat.sub_add_cancel (by omega)]
rwa [← pullCount_eq_pullCount]
exact h_ne

/-- Reward obtained when pulling arm `a` for the `m`-th time. -/
noncomputable
def rewardByCount (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : ℝ :=
match (stepsUntil (arm · h) a m) with
| ⊤ => 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
Expand All @@ -87,7 +158,7 @@ lemma sum_rewardByCount_eq_sum_reward [DecidableEq α]
· rw [← hta] at ht ⊢
rw [pullCount_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht]
rw [sum_range_succ, if_pos rfl, rewardByCount_pullCount_add_one_eq_reward]
· rwa [pullCount_eq_pullCount _ _ _ hta, sum_range_succ, if_neg hta, add_zero]
· rwa [pullCount_eq_pullCount hta, sum_range_succ, if_neg hta, add_zero]

lemma sum_pullCount_mul [Fintype α] (k : ℕ → α) (f : α → ℝ) (t : ℕ) :
∑ a, pullCount k a t * f a = ∑ s ∈ range t, f (k s) := by
Expand Down Expand Up @@ -116,10 +187,12 @@ variable [Fintype α] [Nonempty α]
noncomputable def bestArm (ν : Kernel α ℝ) : α :=
(exists_max_image univ (fun a ↦ (ν a)[id]) (univ_nonempty_iff.mpr inferInstance)).choose

omit [DecidableEq α] in
lemma le_bestArm (a : α) : (ν a)[id] ≤ (ν (bestArm ν))[id] :=
(exists_max_image univ (fun a ↦ (ν a)[id])
(univ_nonempty_iff.mpr inferInstance)).choose_spec.2 _ (mem_univ a)

omit [DecidableEq α] in
lemma gap_eq_bestArm_sub : gap ν a = (ν (bestArm ν))[id] - (ν a)[id] := by
rw [gap]
congr
Expand Down
173 changes: 173 additions & 0 deletions LeanBandits/RewardByCountMeasure.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,173 @@
/-
Copyright (c) 2025 Rémy Degenne. All rights reserved.
Released under Apache 2.0 license as described in the file LICENSE.
Authors: Rémy Degenne
-/
import LeanBandits.Bandit
import LeanBandits.Regret

/-! # Laws of `stepsUntil` and `rewardByCount`
-/

open MeasureTheory ProbabilityTheory Finset
open scoped ENNReal NNReal

section Aux

variable {α β γ Ω Ω' : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω]
{mα : MeasurableSpace α} {μ : Measure α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ}
[MeasurableSpace Ω'] [StandardBorelSpace Ω'] [Nonempty Ω']
{X : α → β} {Y : α → Ω} {Z : α → Ω'}

lemma MeasureTheory.Measure.comp_congr {κ η : Kernel α β} (h : ∀ᵐ a ∂μ, κ a = η a) :
κ ∘ₘ μ = η ∘ₘ μ :=
Measure.bind_congr_right h

lemma MeasureTheory.Measure.copy_comp_map (hX : AEMeasurable X μ) :
Kernel.copy β ∘ₘ (μ.map X) = μ.map (fun a ↦ (X a, X a)) := by
rw [Kernel.copy, deterministic_comp_eq_map, AEMeasurable.map_map_of_aemeasurable (by fun_prop) hX]
congr

lemma MeasureTheory.Measure.compProd_deterministic [SFinite μ] (hX : Measurable X) :
μ ⊗ₘ (Kernel.deterministic X hX) = μ.map (fun a ↦ (a, X a)) := by
rw [Measure.compProd_eq_comp_prod, Kernel.id, Kernel.deterministic_prod_deterministic,
Measure.deterministic_comp_eq_map]
rfl

lemma ProbabilityTheory.condDistrib_comp_map [IsFiniteMeasure μ]
(hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) :
condDistrib Y X μ ∘ₘ (μ.map X) = μ.map Y := by
rw [← Measure.snd_compProd, compProd_map_condDistrib hY, Measure.snd_map_prodMk₀ hX]

lemma ProbabilityTheory.condDistrib_comp [IsFiniteMeasure μ]
(hX : AEMeasurable X μ) {f : β → Ω} (hf : Measurable f) :
condDistrib (f ∘ X) X μ =ᵐ[μ.map X] Kernel.deterministic f hf := by
rw [← Kernel.compProd_eq_iff, compProd_map_condDistrib (by fun_prop),
Measure.compProd_deterministic, AEMeasurable.map_map_of_aemeasurable (by fun_prop) hX]
congr

lemma ProbabilityTheory.condDistrib_const [IsFiniteMeasure μ]
(hX : AEMeasurable X μ) (c : Ω) :
condDistrib (fun _ ↦ c) X μ =ᵐ[μ.map X] Kernel.deterministic (fun _ ↦ c) (by fun_prop) := by
have : (fun _ : α ↦ c) = (fun _ : β ↦ c) ∘ X := rfl
conv_lhs => rw [this]
filter_upwards [condDistrib_comp hX (by fun_prop : Measurable (fun _ ↦ c))] with b hb
rw [hb]

@[fun_prop]
lemma Measurable.coe_nat_enat {f : α → ℕ} (hf : Measurable f) :
Measurable (fun a ↦ (f a : ℕ∞)) := Measurable.comp (by fun_prop) hf

@[fun_prop]
lemma Measurable.toNat {f : α → ℕ∞} (hf : Measurable f) : Measurable (fun a ↦ (f a).toNat) :=
Measurable.comp (by fun_prop) hf

end Aux

namespace Bandits

variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [MeasurableSingletonClass α]

@[fun_prop]
lemma measurable_pullCount (a : α) (t : ℕ) : Measurable (fun k ↦ pullCount k a t) := by
simp_rw [pullCount_eq_sum]
have h_meas s : Measurable (fun k : ℕ → α ↦ if k s = a then 1 else 0) := by
refine Measurable.ite ?_ (by fun_prop) (by fun_prop)
exact (measurableSet_singleton _).preimage (by fun_prop)
fun_prop

@[fun_prop]
lemma measurable_stepsUntil (a : α) (m : ℕ) : Measurable (fun k ↦ stepsUntil k a m) := by
classical
have h_union : {k' | ∃ s, pullCount k' a (s + 1) = m}
= ⋃ s : ℕ, {k' | pullCount k' a (s + 1) = m} := by ext; simp
have h_meas_set : MeasurableSet {k' | ∃ s, pullCount k' a (s + 1) = m} := by
rw [h_union]
exact MeasurableSet.iUnion fun s ↦ (measurableSet_singleton _).preimage (by fun_prop)
simp_rw [stepsUntil_eq_dite]
suffices Measurable fun k ↦ if h : k ∈ {k' | ∃ s, pullCount k' a (s + 1) = m}
then (Nat.find h : ℕ∞) else ⊤ by convert this
refine Measurable.dite (s := {k' : ℕ → α | ∃ s, pullCount k' a (s + 1) = m})
(f := fun x ↦ (Nat.find x.2 : ℕ∞)) (g := fun _ ↦ ⊤) ?_ (by fun_prop) h_meas_set
refine Measurable.coe_nat_enat ?_
refine measurable_find _ fun k ↦ ?_
suffices MeasurableSet {x : ℕ → α | pullCount x a (k + 1) = m} by
have : Subtype.val ''
{x : {k' : ℕ → α | ∃ s, pullCount k' a (s + 1) = m} | pullCount x a (k + 1) = m}
= {x : ℕ → α | pullCount x a (k + 1) = m} := by
ext x
simp only [Set.mem_setOf_eq, Set.coe_setOf, Set.mem_image, Subtype.exists, exists_and_left,
exists_prop, exists_eq_right_right, and_iff_left_iff_imp]
exact fun h ↦ ⟨_, h⟩
refine (MeasurableEmbedding.subtype_coe h_meas_set).measurableSet_image.mp ?_
rw [this]
exact (measurableSet_singleton _).preimage (by fun_prop)
exact (measurableSet_singleton _).preimage (by fun_prop)

lemma measurable_stepsUntil'' (a : α) (m : ℕ) :
Measurable (fun ω : (ℕ → α × ℝ) ↦ stepsUntil (arm · ω) a m) :=
(measurable_stepsUntil a m).comp (by fun_prop)

lemma measurable_stepsUntil' (a : α) (m : ℕ) :
Measurable (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ stepsUntil (arm · ω.1) a m) :=
(measurable_stepsUntil'' a m).comp measurable_fst

@[fun_prop]
lemma measurable_rewardByCount (a : α) (m : ℕ) :
Measurable (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ rewardByCount a m ω.1 ω.2) := by
simp_rw [rewardByCount_eq_ite]
refine Measurable.ite ?_ ?_ ?_
· exact (measurableSet_singleton _).preimage <| measurable_stepsUntil' a m
· fun_prop
· change Measurable ((fun p : ℕ × (ℕ → α × ℝ) ↦ reward p.1 p.2)
∘ (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ ((stepsUntil (arm · ω.1) a m).toNat, ω.1)))
have : Measurable fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦
((stepsUntil (arm · ω.1) a m).toNat, ω.1) :=
(measurable_stepsUntil' a m).toNat.prodMk (by fun_prop)
exact Measurable.comp (by fun_prop) this

lemma condDistrib_rewardByCount_stepsUntil [StandardBorelSpace α] [Nonempty α]
{alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (m : ℕ) (hm : m ≠ 0) :
condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil (arm · ω.1) a m)
(Bandit.measure alg ν)
=ᵐ[(Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)] Kernel.const _ (ν a) := by
sorry

/-- The reward received at the `m`-th pull of arm `a` has law `ν a`. -/
lemma hasLaw_rewardByCount [StandardBorelSpace α] [Nonempty α]
{alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (m : ℕ) (hm : m ≠ 0):
HasLaw (fun ω ↦ rewardByCount a m ω.1 ω.2) (ν a) (Bandit.measure alg ν) where
map_eq := by
have h_condDistrib :
condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil (arm · ω.1) a m)
(Bandit.measure alg ν)
=ᵐ[(Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)]
Kernel.const _ (ν a) := condDistrib_rewardByCount_stepsUntil a m hm
calc (Bandit.measure alg ν).map (fun ω ↦ rewardByCount a m ω.1 ω.2)
_ = (condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil (arm · ω.1) a m)
(Bandit.measure alg ν))
∘ₘ ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)) := by
rw [condDistrib_comp_map (by fun_prop) (by fun_prop)]
_ = (Kernel.const _ (ν a))
∘ₘ ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)) :=
Measure.comp_congr h_condDistrib
_ = ν a := by
have : IsProbabilityMeasure
((Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)) :=
isProbabilityMeasure_map (by fun_prop)
simp

lemma identDistrib_rewardByCount [StandardBorelSpace α] [Nonempty α]
{alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (n m : ℕ)
(hn : n ≠ 0) (hm : m ≠ 0) :
IdentDistrib (fun ω ↦ rewardByCount a n ω.1 ω.2) (fun ω ↦ rewardByCount a m ω.1 ω.2)
(Bandit.measure alg ν) (Bandit.measure alg ν) where
aemeasurable_fst := by fun_prop
aemeasurable_snd := by fun_prop
map_eq := by rw [(hasLaw_rewardByCount a n hn).map_eq, (hasLaw_rewardByCount a m hm).map_eq]

lemma iIndepFun_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] :
iIndepFun (fun (p : α × ℕ) ω ↦ rewardByCount p.1 p.2 ω.1 ω.2) (Bandit.measure alg ν) := by
sorry

end Bandits