Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions LeanBandits.lean
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import LeanBandits.Bandit.RewardByCountMeasure
import LeanBandits.Bandit.SumRewards
import LeanBandits.BanditAlgorithms.AuxSums
import LeanBandits.BanditAlgorithms.ETC
import LeanBandits.BanditAlgorithms.RoundRobin
import LeanBandits.BanditAlgorithms.UCB
import LeanBandits.ForMathlib.CondDistrib
import LeanBandits.ForMathlib.CondIndepFun
Expand Down
47 changes: 23 additions & 24 deletions LeanBandits/BanditAlgorithms/ETC.lean
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,8 @@ Released under Apache 2.0 license as described in the file LICENSE.
Authors: Rémy Degenne
-/
import LeanBandits.Bandit.SumRewards
import LeanBandits.BanditAlgorithms.AuxSums
import LeanBandits.BanditAlgorithms.RoundRobin
import LeanBandits.ForMathlib.MeasurableArgMax
import LeanBandits.SequentialLearning.Deterministic

/-! # The Explore-Then-Commit Algorithm

Expand Down Expand Up @@ -60,13 +59,28 @@ variable {hK : 0 < K} {m : ℕ} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν]
{A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ}
{σ2 : ℝ≥0}

/-- Until round `K * m - 1`, the ETC algorithm behaves like the Round-Robin algorithm. -/
lemma isAlgEnvSeqUntil_roundRobinAlgorithm [Nonempty (Fin K)]
(h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) :
IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P (K * m - 1) where
measurable_A := h.measurable_A
measurable_R := h.measurable_R
hasLaw_action_zero := h.hasLaw_action_zero
hasCondDistrib_reward_zero := h.hasCondDistrib_reward_zero
hasCondDistrib_action n hn := by
convert h.hasCondDistrib_action n using 1
simp only [roundRobinAlgorithm, detAlgorithm_policy, etcAlgorithm]
congr 1 with h
unfold ETC.nextArm RoundRobin.nextArm
simp [hn]
hasCondDistrib_reward n _ := h.hasCondDistrib_reward n

section AlgorithmBehavior

lemma arm_zero [Nonempty (Fin K)]
(h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) :
A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ := by
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
exact h.action_zero_detAlgorithm
A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ :=
RoundRobin.arm_zero ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono zero_le')

lemma arm_ae_eq_etcNextArm [Nonempty (Fin K)]
(h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (n : ℕ) :
Expand All @@ -77,13 +91,8 @@ lemma arm_ae_eq_etcNextArm [Nonempty (Fin K)]
/-- For `n < K * m`, the arm pulled at time `n` is the arm `n % K`. -/
lemma arm_of_lt [Nonempty (Fin K)]
(h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) {n : ℕ} (hn : n < K * m) :
A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := by
cases n with
| zero => exact arm_zero h
| succ n =>
filter_upwards [arm_ae_eq_etcNextArm h n] with h hn_eq
rw [hn_eq, nextArm, dif_pos]
grind
A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ :=
RoundRobin.arm_ae_eq n ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono (by grind))

/-- The arm pulled at time `K * m` is the arm with the highest empirical mean after the exploration
phase. -/
Expand Down Expand Up @@ -125,18 +134,8 @@ lemma arm_of_ge [Nonempty (Fin K)]
/-- At time `K * m`, the number of pulls of each arm is equal to `m`. -/
lemma pullCount_mul [Nonempty (Fin K)]
(h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (a : Fin K) :
pullCount A a (K * m) =ᵐ[P] fun _ ↦ m := by
rw [Filter.EventuallyEq]
simp_rw [pullCount_eq_sum]
have h_arm (n : range (K * m)) : A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ :=
arm_of_lt h (mem_range.mp n.2)
simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_arm
filter_upwards [h_arm] with ω h_arm
have h_arm' {i : ℕ} (hi : i ∈ range (K * m)) : A i ω = ⟨i % K, Nat.mod_lt _ hK⟩ := h_arm ⟨i, hi⟩
calc (∑ s ∈ range (K * m), if A s ω = a then 1 else 0)
_ = (∑ s ∈ range (K * m), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) :=
sum_congr rfl fun s hs ↦ by rw [h_arm' hs]
_ = m := sum_mod_range_mul hK m a
pullCount A a (K * m) =ᵐ[P] fun _ ↦ m :=
RoundRobin.pullCount_mul m (isAlgEnvSeqUntil_roundRobinAlgorithm h) a

lemma pullCount_add_one_of_ge [Nonempty (Fin K)]
(h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P)
Expand Down
117 changes: 117 additions & 0 deletions LeanBandits/BanditAlgorithms/RoundRobin.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,117 @@
/-
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.BanditAlgorithms.AuxSums
import LeanBandits.SequentialLearning.Deterministic
import LeanBandits.SequentialLearning.FiniteActions
import LeanBandits.SequentialLearning.StationaryEnv

/-! # Round-Robin algorithm

That algorithm pulls each arm in a round-robin fashion.

-/

open MeasureTheory ProbabilityTheory Finset Learning
open scoped ENNReal NNReal

namespace Bandits

variable {K : ℕ}

section AlgorithmDefinition

/-- Arm pulled by the Round-Robin algorithm at time `n + 1`. This is arm `n % K`. -/
noncomputable
def RoundRobin.nextArm (hK : 0 < K) (n : ℕ) : Fin K := ⟨(n + 1) % K, Nat.mod_lt _ hK⟩

/-- The Round-Robin algorithm: deterministic algorithm that chooses the next arm according
to `RoundRobin.nextArm`. -/
noncomputable
def roundRobinAlgorithm (hK : 0 < K) : Algorithm (Fin K) ℝ :=
detAlgorithm (fun n _ ↦ RoundRobin.nextArm hK n) (by fun_prop) ⟨0, hK⟩

end AlgorithmDefinition

namespace RoundRobin

variable {hK : 0 < K} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν]
{Ω : Type*} {mΩ : MeasurableSpace Ω}
{P : Measure Ω} [IsProbabilityMeasure P]
{A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ}

lemma arm_zero [Nonempty (Fin K)]
(h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P 0) :
A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ := by
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
exact h.action_zero_detAlgorithm

lemma arm_ae_eq_roundRobinNextArm [Nonempty (Fin K)] (n : ℕ)
(h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P (n + 1)) :
A (n + 1) =ᵐ[P] fun _ ↦ nextArm hK n :=
h.action_detAlgorithm_ae_eq (by grind)

/-- The arm pulled at time `n` is the arm `n % K`. -/
lemma arm_ae_eq [Nonempty (Fin K)] (n : ℕ)
(h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P n) :
A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := by
cases n with
| zero => exact arm_zero h
| succ n =>
filter_upwards [arm_ae_eq_roundRobinNextArm n h] with h hn_eq
rw [hn_eq, nextArm]

/-- At time `K * m`, the number of pulls of each arm is equal to `m`. -/
lemma pullCount_mul [Nonempty (Fin K)] (m : ℕ)
(h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P (K * m - 1))
(a : Fin K) :
pullCount A a (K * m) =ᵐ[P] fun _ ↦ m := by
rw [Filter.EventuallyEq]
simp_rw [pullCount_eq_sum]
have h_arm (n : range (K * m)) : A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ :=
arm_ae_eq n (h.mono (by have := n.2; simp only [mem_range] at this; grind))
simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_arm
filter_upwards [h_arm] with ω h_arm
have h_arm' {i : ℕ} (hi : i ∈ range (K * m)) : A i ω = ⟨i % K, Nat.mod_lt _ hK⟩ := h_arm ⟨i, hi⟩
calc (∑ s ∈ range (K * m), if A s ω = a then 1 else 0)
_ = (∑ s ∈ range (K * m), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) :=
sum_congr rfl fun s hs ↦ by rw [h_arm' hs]
_ = m := sum_mod_range_mul hK m a

lemma pullCount_eq_one [Nonempty (Fin K)]
(h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P (K - 1))
(a : Fin K) :
pullCount A a K =ᵐ[P] fun _ ↦ 1 := by
suffices pullCount A a (K * 1) =ᵐ[P] fun _ ↦ 1 by simpa using this
refine pullCount_mul 1 (P := P) (ν := ν) (R := R) (hK := hK) ?_ a
simpa

lemma time_gt_of_pullCount_gt_one [Nonempty (Fin K)]
(h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P (K - 1)) (a : Fin K) :
∀ᵐ ω ∂P, ∀ n, 1 < pullCount A a n ω → K < n := by
filter_upwards [pullCount_eq_one h a] with h h_eq n hn
rw [← h_eq] at hn
by_contra! h_lt
exact hn.not_ge (pullCount_mono _ h_lt _)

lemma pullCount_pos_of_time_ge [Nonempty (Fin K)]
(h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P (K - 1)) :
∀ᵐ ω ∂P, ∀ n, K ≤ n → ∀ b : Fin K, 0 < pullCount A b n ω := by
have h_ae a := pullCount_eq_one h a
simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_ae
filter_upwards [h_ae] with ω hω n hn a
refine Nat.one_pos.trans_le ?_
rw [← hω a]
exact pullCount_mono _ hn _

lemma pullCount_pos_of_pullCount_gt_one [Nonempty (Fin K)]
(h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P (K - 1)) (a : Fin K) :
∀ᵐ ω ∂P, ∀ n, 1 < pullCount A a n ω → ∀ b : Fin K, 0 < pullCount A b n ω := by
filter_upwards [time_gt_of_pullCount_gt_one h a, pullCount_pos_of_time_ge h] with ω h1 h2 n h_gt a
exact h2 n (h1 n h_gt).le a

end RoundRobin

end Bandits
56 changes: 23 additions & 33 deletions LeanBandits/BanditAlgorithms/UCB.lean
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,8 @@ Released under Apache 2.0 license as described in the file LICENSE.
Authors: Rémy Degenne
-/
import LeanBandits.Bandit.SumRewards
import LeanBandits.BanditAlgorithms.AuxSums
import LeanBandits.BanditAlgorithms.RoundRobin
import LeanBandits.ForMathlib.MeasurableArgMax
import LeanBandits.SequentialLearning.Deterministic

/-!
# UCB algorithm
Expand Down Expand Up @@ -59,6 +58,22 @@ variable {hK : 0 < K} {c : ℝ} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν]
{A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ}
{σ2 : ℝ≥0} {n : ℕ} {ω : Ω}

/-- Until round `K - 1`, the UCB algorithm behaves like the Round-Robin algorithm. -/
lemma isAlgEnvSeqUntil_roundRobinAlgorithm [Nonempty (Fin K)]
(h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) :
IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P (K - 1) where
measurable_A := h.measurable_A
measurable_R := h.measurable_R
hasLaw_action_zero := h.hasLaw_action_zero
hasCondDistrib_reward_zero := h.hasCondDistrib_reward_zero
hasCondDistrib_action n hn := by
convert h.hasCondDistrib_action n using 1
simp only [roundRobinAlgorithm, detAlgorithm_policy, ucbAlgorithm]
congr 1 with h
unfold UCB.nextArm RoundRobin.nextArm
simp [hn]
hasCondDistrib_reward n _ := h.hasCondDistrib_reward n

section AlgorithmBehavior

/-- The exploration bonus of the UCB algorithm, which corresponds to the width of
Expand All @@ -82,9 +97,8 @@ lemma ucbWidth_eq_ucbWidth' (c : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) (hn : n

lemma arm_zero [Nonempty (Fin K)]
(h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) :
A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ := by
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
exact h.action_zero_detAlgorithm
A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ :=
RoundRobin.arm_zero ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono zero_le')

lemma arm_ae_eq_ucbNextArm [Nonempty (Fin K)]
(h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (n : ℕ) :
Expand Down Expand Up @@ -147,39 +161,15 @@ lemma forall_arm_prop [Nonempty (Fin K)]
simp_rw [ae_all_iff] at h_ae
exact h_ae n hn

lemma pullCount_eq_of_time_eq [Nonempty (Fin K)]
(h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) :
∀ᵐ ω ∂P, pullCount A a K ω = 1 := by
filter_upwards [forall_arm_eq_mod_of_lt h] with h h_eq
rw [pullCount_eq_sum]
conv_rhs => rw [← sum_mod_range hK a]
refine Finset.sum_congr rfl fun s hs ↦ ?_
congr
exact h_eq s (by grind)

lemma time_gt_of_pullCount_gt_one [Nonempty (Fin K)]
(h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) :
∀ᵐ ω ∂P, ∀ n, 1 < pullCount A a n ω → K < n := by
filter_upwards [pullCount_eq_of_time_eq h a] with h h_eq n hn
rw [← h_eq] at hn
by_contra! h_lt
exact hn.not_ge (pullCount_mono _ h_lt _)

lemma pullCount_pos_of_time_ge [Nonempty (Fin K)]
(h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) :
∀ᵐ ω ∂P, ∀ n, K ≤ n → ∀ b : Fin K, 0 < pullCount A b n ω := by
have h_ae a := pullCount_eq_of_time_eq h a
rw [← ae_all_iff] at h_ae
filter_upwards [h_ae] with ω hω n hn a
refine Nat.one_pos.trans_le ?_
rw [← hω a]
exact pullCount_mono _ hn _
∀ᵐ ω ∂P, ∀ n, 1 < pullCount A a n ω → K < n :=
RoundRobin.time_gt_of_pullCount_gt_one (isAlgEnvSeqUntil_roundRobinAlgorithm h) a

lemma pullCount_pos_of_pullCount_gt_one [Nonempty (Fin K)]
(h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) :
∀ᵐ ω ∂P, ∀ n, 1 < pullCount A a n ω → ∀ b : Fin K, 0 < pullCount A b n ω := by
filter_upwards [time_gt_of_pullCount_gt_one h a, pullCount_pos_of_time_ge h] with ω h1 h2 n h_gt a
exact h2 n (h1 n h_gt).le a
∀ᵐ ω ∂P, ∀ n, 1 < pullCount A a n ω → ∀ b : Fin K, 0 < pullCount A b n ω :=
RoundRobin.pullCount_pos_of_pullCount_gt_one (isAlgEnvSeqUntil_roundRobinAlgorithm h) a

end AlgorithmBehavior

Expand Down
30 changes: 25 additions & 5 deletions LeanBandits/ForMathlib/Traj.lean
Original file line number Diff line number Diff line change
Expand Up @@ -59,14 +59,16 @@ lemma MeasurableEquiv.coe_prodCongr {α β γ δ : Type*}
lemma MeasurableEquiv.coe_refl {α : Type*} {mα : MeasurableSpace α} :
(MeasurableEquiv.refl α : α → α) = id := rfl

theorem hasLaw_Iic_of_forall_hasCondDistrib [∀ n, StandardBorelSpace (X n)] [∀ n, Nonempty (X n)]
{Y : (n : ℕ) → Ω → X n} (h0 : HasLaw (Y 0) μ₀ P)
(h_condDistrib : ∀ n, HasCondDistrib (Y (n + 1)) (fun ω ↦ fun i : Iic n ↦ Y i ω) (κ n) P)
(n : ℕ) :
theorem hasLaw_Iic_of_forall_hasCondDistrib' [∀ n, StandardBorelSpace (X n)] [∀ n, Nonempty (X n)]
{Y : (n : ℕ) → Ω → X n} (h0 : HasLaw (Y 0) μ₀ P) {N n : ℕ}
(h_condDistrib : ∀ n < N, HasCondDistrib (Y (n + 1)) (fun ω ↦ fun i : Iic n ↦ Y i ω) (κ n) P)
(hn : n ≤ N) :
HasLaw (fun ω (i : Iic n) ↦ Y i ω)
((partialTraj κ 0 n) ∘ₘ (μ₀.map (MeasurableEquiv.piUnique _).symm)) P := by
revert hn
induction n with
| zero =>
intro _
simp only [piUnique_symm_apply, partialTraj_self, Measure.id_comp]
rw [← h0.map_eq, AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)]
constructor
Expand All @@ -84,7 +86,9 @@ theorem hasLaw_Iic_of_forall_hasCondDistrib [∀ n, StandardBorelSpace (X n)] [
rw [Unique.eq_default i]
simp [coe_default_Iic_zero]
| succ n hn =>
specialize h_condDistrib n
intro hn_le
specialize h_condDistrib n (by grind)
specialize hn (by grind)
have h_law := hn.prod_of_hasCondDistrib h_condDistrib
have : (fun ω (i : Iic (n + 1)) ↦ Y i ω) =
(MeasurableEquiv.IicSuccProd X n).symm ∘
Expand All @@ -109,13 +113,29 @@ theorem hasLaw_Iic_of_forall_hasCondDistrib [∀ n, StandardBorelSpace (X n)] [
congr
simp [MeasurableEquiv.coe_refl]

theorem hasLaw_Iic_of_forall_hasCondDistrib [∀ n, StandardBorelSpace (X n)] [∀ n, Nonempty (X n)]
{Y : (n : ℕ) → Ω → X n} (h0 : HasLaw (Y 0) μ₀ P)
(h_condDistrib : ∀ n, HasCondDistrib (Y (n + 1)) (fun ω ↦ fun i : Iic n ↦ Y i ω) (κ n) P)
(n : ℕ) :
HasLaw (fun ω (i : Iic n) ↦ Y i ω)
((partialTraj κ 0 n) ∘ₘ (μ₀.map (MeasurableEquiv.piUnique _).symm)) P := by
exact hasLaw_Iic_of_forall_hasCondDistrib' (N := n) h0 (fun n _ ↦ h_condDistrib n) le_rfl

omit [IsProbabilityMeasure μ₀] in
lemma trajMeasure_map_frestrictLe (n : ℕ) :
(trajMeasure μ₀ κ).map (frestrictLe n) =
(partialTraj κ 0 n) ∘ₘ (μ₀.map (MeasurableEquiv.piUnique _).symm) := by
rw [trajMeasure, ← Measure.deterministic_comp_eq_map (by fun_prop), Measure.comp_assoc,
Kernel.deterministic_comp_eq_map, traj_map_frestrictLe]

theorem eq_trajMeasure_map_frestrictLe [∀ n, StandardBorelSpace (X n)] [∀ n, Nonempty (X n)]
{Y : (n : ℕ) → Ω → X n}
(h0 : HasLaw (Y 0) μ₀ P) {N : ℕ}
(h_condDistrib : ∀ n < N, HasCondDistrib (Y (n + 1)) (fun ω ↦ fun i : Iic n ↦ Y i ω) (κ n) P) :
P.map (fun ω (n : Iic N) ↦ Y n ω) = (trajMeasure μ₀ κ).map (frestrictLe N) := by
rw [(hasLaw_Iic_of_forall_hasCondDistrib' h0 h_condDistrib le_rfl).map_eq,
trajMeasure_map_frestrictLe]

-- todo: switch to `HasLaw`
/-- Uniqueness of `trajMeasure`. -/
theorem eq_trajMeasure [∀ n, StandardBorelSpace (X n)] [∀ n, Nonempty (X n)]
Expand Down
Loading