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
2 changes: 2 additions & 0 deletions LeanBandits.lean
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
import LeanBandits.AlgorithmBuilding
import LeanBandits.Bandit
import LeanBandits.ETC
import LeanBandits.Regret
import LeanBandits.UCB
132 changes: 132 additions & 0 deletions LeanBandits/AlgorithmBuilding.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,132 @@
/-
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

/-! # Tools to build bandit algorithms

-/

open MeasureTheory ProbabilityTheory Finset
open scoped ENNReal NNReal

section MeasurableArgmax -- copied from PR #27579 (and changed from argmin to argmax)

lemma measurable_encode {α : Type*} {_ : MeasurableSpace α} [Encodable α]
[MeasurableSingletonClass α] :
Measurable (Encodable.encode (α := α)) := by
refine measurable_to_nat fun a ↦ ?_
have : Encodable.encode ⁻¹' {Encodable.encode a} = {a} := by ext; simp
rw [this]
exact measurableSet_singleton _

lemma measurableEmbedding_encode (α : Type*) {_ : MeasurableSpace α} [Encodable α]
[MeasurableSingletonClass α] :
MeasurableEmbedding (Encodable.encode (α := α)) where
injective := Encodable.encode_injective
measurable := measurable_encode
measurableSet_image' _ _ := .of_discrete

section Finite

variable {𝓧 𝓨 α : Type*} {m𝓧 : MeasurableSpace 𝓧} {m𝓨 : MeasurableSpace 𝓨}
{mα : MeasurableSpace α} [TopologicalSpace α] [LinearOrder α]
[OpensMeasurableSpace α] [OrderClosedTopology α] [SecondCountableTopology α]

lemma measurableSet_isMax [Countable 𝓨]
{f : 𝓧 → 𝓨 → α} (hf : ∀ y, Measurable (fun x ↦ f x y)) (y : 𝓨) :
MeasurableSet {x | ∀ z, f x z ≤ f x y} := by
rw [show {x | ∀ y', f x y' ≤ f x y} = ⋂ y', {x | f x y' ≤ f x y} by ext; simp]
exact MeasurableSet.iInter fun z ↦ measurableSet_le (by fun_prop) (by fun_prop)

lemma exists_isMaxOn' {α : Type*} [LinearOrder α]
[Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] (f : 𝓧 → 𝓨 → α) (x : 𝓧) :
∃ n : ℕ, ∃ y, n = Encodable.encode y ∧ ∀ z, f x z ≤ f x y := by
obtain ⟨y, h⟩ := Finite.exists_max (f x)
exact ⟨Encodable.encode y, y, rfl, h⟩

/-- A measurable argmax function. -/
noncomputable
def measurableArgmax [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨]
(f : 𝓧 → 𝓨 → α)
[∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y]
(x : 𝓧) :
𝓨 :=
(measurableEmbedding_encode 𝓨).invFun (Nat.find (exists_isMaxOn' f x))

lemma measurable_measurableArgmax [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨]
{f : 𝓧 → 𝓨 → α}
[∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y]
(hf : ∀ y, Measurable (fun x ↦ f x y)) :
Measurable (measurableArgmax f) := by
refine (MeasurableEmbedding.measurable_invFun (measurableEmbedding_encode 𝓨)).comp ?_
refine measurable_find _ fun n ↦ ?_
have : {x | ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y}
= ⋃ y, ({x | n = Encodable.encode y} ∩ {x | ∀ z, f x z ≤ f x y}) := by ext; simp
rw [this]
refine MeasurableSet.iUnion fun y ↦ (MeasurableSet.inter (by simp) ?_)
exact measurableSet_isMax (by fun_prop) y

lemma isMaxOn_measurableArgmax {α : Type*} [LinearOrder α]
[Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨]
(f : 𝓧 → 𝓨 → α)
[∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y]
(x : 𝓧) (z : 𝓨) :
f x z ≤ f x (measurableArgmax f x) := by
obtain ⟨y, h_eq, h_le⟩ := Nat.find_spec (exists_isMaxOn' f x)
refine le_trans (h_le z) (le_of_eq ?_)
rw [measurableArgmax, h_eq,
MeasurableEmbedding.leftInverse_invFun (measurableEmbedding_encode 𝓨) y]

end Finite
end MeasurableArgmax

namespace Bandits

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

/-- Number of pulls of arm `a` up to (and including) time `n`. -/
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']

@[fun_prop]
lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) :
Measurable (fun h ↦ 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' [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 Bandits
127 changes: 67 additions & 60 deletions LeanBandits/Bandit.lean
Original file line number Diff line number Diff line change
@@ -1,7 +1,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
Authors: Rémy Degenne, Paulo Rauber
-/
import Mathlib

Expand All @@ -22,115 +22,122 @@ def MeasurableEquiv.piIicZero (α : Type*) [MeasurableSpace α] :

namespace Bandits

variable {α : Type*} {mα : MeasurableSpace α}
variable {α R : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R}

section MeasureSpace

/-- A bandit interaction between an agent described by a policy and an environment given by
reward distributions. -/
structure Bandit (α : Type*) [MeasurableSpace α] where
/-- Conditional distribution of the rewards given the arm pulled. -/
ν : Kernel α ℝ
hν : IsMarkovKernel ν
/-- A stochastic, sequential algorithm. -/
structure Algorithm (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] where
/-- Policy or sampling rule: distribution of the next pull. -/
policy : (n : ℕ) → Kernel (Iic n → α × ℝ) α
h_policy n : IsMarkovKernel (policy n)
policy : (n : ℕ) → Kernel (Iic n → α × R) α
[h_policy : ∀ n, IsMarkovKernel (policy n)]
/-- Distribution of the first pull. -/
p0 : Measure α
hp0 : IsProbabilityMeasure p0
[hp0 : IsProbabilityMeasure p0]

instance (b : Bandit α) : IsMarkovKernel b.ν := b.hν
instance (b : Bandit α) (n : ℕ) : IsMarkovKernel (b.policy n) := b.h_policy n
instance (b : Bandit α) : IsProbabilityMeasure b.p0 := b.hp0
instance (alg : Algorithm α R) (n : ℕ) : IsMarkovKernel (alg.policy n) := alg.h_policy n
instance (alg : Algorithm α R) : IsProbabilityMeasure alg.p0 := alg.hp0

namespace Bandit

/-- Kernel describing the distribution of the next arm-reward pair given the history up to `n`. -/
noncomputable
def stepKernel (b : Bandit α) (n : ℕ) : Kernel (Iic n → α × ℝ) (α × ℝ) :=
(b.policy n) ⊗ₖ b.ν.prodMkLeft (Iic n → α × ℝ)
def stepKernel (alg : Algorithm α R) (ν : Kernel α R) (n : ℕ) : Kernel (Iic n → α × R) (α × R) :=
(alg.policy n) ⊗ₖ ν.prodMkLeft (Iic n → α × R)

instance (b : Bandit α) (n : ℕ) : IsMarkovKernel (b.stepKernel n) := by
instance (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
IsMarkovKernel (stepKernel alg ν n) := by
rw [stepKernel]
infer_instance

@[simp]
lemma fst_stepKernel (b : Bandit α) (n : ℕ) : (b.stepKernel n).fst = b.policy n := by
lemma fst_stepKernel (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
(stepKernel alg ν n).fst = alg.policy n := by
rw [stepKernel, Kernel.fst_compProd]

@[simp]
lemma snd_stepKernel (b : Bandit α) (n : ℕ) : (b.stepKernel n).snd = b.ν ∘ₖ b.policy n := by
lemma snd_stepKernel (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
(stepKernel alg ν n).snd = ν ∘ₖ alg.policy n := by
rw [stepKernel, Kernel.snd_compProd_prodMkLeft]

/-- Kernel sending a partial trajectory of the bandit interaction `Iic n → α × ℝ` to a measure
on `ℕ → α × ℝ`, supported on full trajectories that start with the partial one. -/
noncomputable def traj (b : Bandit α) (n : ℕ) : Kernel (Iic n → α × ℝ) (ℕ → α × ℝ) :=
ProbabilityTheory.Kernel.traj (X := fun _ ↦ α × ℝ) b.stepKernel n

instance (b : Bandit α) (n : ℕ) : IsMarkovKernel (b.traj n) := by
rw [traj]
infer_instance
noncomputable def traj (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
Kernel (Iic n → α × R) (ℕ → α × R) :=
ProbabilityTheory.Kernel.traj (X := fun _ ↦ α × R) (stepKernel alg ν) n
deriving IsMarkovKernel

/-- Measure on the sequence of arms pulled and rewards observed generated by the bandit. -/
noncomputable
def trajMeasure (b : Bandit α) : Measure (ℕ → α × ℝ) :=
(b.traj 0) ∘ₘ ((b.p0 ⊗ₘ b.ν).map (MeasurableEquiv.piIicZero _).symm)
def trajMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : Measure (ℕ → α × R) :=
(traj alg ν 0) ∘ₘ ((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero _).symm)

/-- Measure of an infinite stream of rewards from each arm. -/
noncomputable
def streamMeasure (b : Bandit α) : Measure (ℕ → α → ℝ) :=
Measure.infinitePi fun _ ↦ Measure.infinitePi b.ν
def streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] : Measure (ℕ → α → R) :=
Measure.infinitePi fun _ ↦ Measure.infinitePi ν
deriving IsProbabilityMeasure

instance (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
IsProbabilityMeasure (trajMeasure alg ν) := by
rw [trajMeasure]
have : IsProbabilityMeasure ((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero _).symm) :=
isProbabilityMeasure_map <| by fun_prop
infer_instance

/-- Joint distribution of the sequence of arm pulled and rewards, and a stream of independent
rewards from all arms. -/
noncomputable
def measure (b : Bandit α) : Measure ((ℕ → α × ℝ) × (ℕ → α → ℝ)) :=
(b.trajMeasure).prod (b.streamMeasure)

instance (b : Bandit α) : IsProbabilityMeasure b.trajMeasure := by
rw [Bandit.trajMeasure]
have : IsProbabilityMeasure ((b.p0 ⊗ₘ b.ν).map (MeasurableEquiv.piIicZero _).symm) :=
isProbabilityMeasure_map <| by fun_prop
infer_instance

instance (b : Bandit α) : IsProbabilityMeasure b.streamMeasure := by
rw [streamMeasure]
infer_instance

instance (b : Bandit α) : IsProbabilityMeasure b.measure := by
rw [measure]
infer_instance
def measure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
Measure ((ℕ → α × R) × (ℕ → α → R)) :=
(trajMeasure alg ν).prod (streamMeasure ν)
deriving IsProbabilityMeasure

end Bandit

/-- `arm n` is the arm pulled at time `n`. This is a random variable on the measurable space
`ℕ → α × ℝ`. -/
def arm (n : ℕ) (h : ℕ → α × ℝ) : α := (h n).1
def arm (n : ℕ) (h : ℕ → α × R) : α := (h n).1

/-- `reward n` is the reward at time `n`. This is a random variable on the measurable space
`ℕ → α × ℝ`. -/
def reward (n : ℕ) (h : ℕ → α × ℝ) : ℝ := (h n).2
`ℕ → α × R`. -/
def reward (n : ℕ) (h : ℕ → α × R) : R := (h n).2

/-- `hist n` is the history up to time `n`. This is a random variable on the measurable space
`ℕ → α × ℝ`. -/
def hist (n : ℕ) (h : ℕ → α × ℝ) : Iic n → α × ℝ := fun i ↦ h i
`ℕ → α × R`. -/
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_reward (n : ℕ) : Measurable (reward n (α := α) (R := R)) := by
unfold reward; 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 α] :
Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × ℝ)) :=
MeasureTheory.Filtration.piLE (X := fun _ ↦ α × ℝ)

lemma condDistrib_arm_reward [StandardBorelSpace α] [Nonempty α] (b : Bandit α) (n : ℕ) :
condDistrib (fun h ↦ (arm n h, reward n h)) (hist n) b.trajMeasure = b.stepKernel n := by
Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) :=
MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R)

lemma condDistrib_arm_reward [StandardBorelSpace α] [Nonempty α]
[StandardBorelSpace R] [Nonempty R] (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν]
(n : ℕ) :
condDistrib (fun h ↦ (arm n h, reward n h)) (hist n) (Bandit.trajMeasure alg ν)
= Bandit.stepKernel alg ν n := by
sorry

lemma condDistrib_reward (b : Bandit α) (n : ℕ) :
condDistrib (reward n) (arm n) b.trajMeasure = b.ν := by
lemma condDistrib_reward [StandardBorelSpace R] [Nonempty R] (alg : Algorithm α R)
(ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
condDistrib (reward n) (arm n) (Bandit.trajMeasure alg ν) = ν := by
sorry

lemma condDistrib_arm [StandardBorelSpace α] [Nonempty α] (b : Bandit α) (n : ℕ) :
condDistrib (arm n) (hist n) b.trajMeasure = b.policy n := by
rw [← b.fst_stepKernel, ← condDistrib_arm_reward]
lemma condDistrib_arm [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
(alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
condDistrib (arm n) (hist n) (Bandit.trajMeasure alg ν) = alg.policy n := by
rw [← Bandit.fst_stepKernel alg ν n, ← condDistrib_arm_reward alg ν n]
sorry

end MeasureSpace
Expand Down
61 changes: 61 additions & 0 deletions LeanBandits/ETC.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
/-
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.Moments.SubGaussian
import LeanBandits.AlgorithmBuilding

/-! # The Explore-Then-Commit Algorithm

-/

open MeasureTheory ProbabilityTheory Finset
open scoped ENNReal NNReal

namespace Bandits

variable {K : ℕ}

/-- Arm pulled by the ETC algorithm at time `n + 1`. -/
noncomputable
def etcNextArm (hK : 0 < K) (m n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K :=
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
if hn : n < K * m - 1 then
⟨(n + 1) % K, Nat.mod_lt _ hK⟩ -- for `n = 0` we have pulled arm 0 already, and we pull arm 1
else
if hn_eq : n = K * m - 1 then measurableArgmax (empMean' n) h
else (h ⟨n - 1, by simp⟩).1

@[fun_prop]
lemma measurable_etcNextArm (hK : 0 < K) (m n : ℕ) : Measurable (etcNextArm hK m n) := by
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
unfold etcNextArm
simp only [dite_eq_ite]
refine Measurable.ite (by simp) (by fun_prop) ?_
refine Measurable.ite (by simp) ?_ (by fun_prop)
exact measurable_measurableArgmax fun a ↦ by fun_prop

/-- The Explore-Then-Commit algorithm. -/
noncomputable
def etcAlgorithm (hK : 0 < K) (m : ℕ) : Algorithm (Fin K) ℝ where
policy n := Kernel.deterministic (etcNextArm hK m n) (by fun_prop)
p0 := Measure.dirac ⟨0, hK⟩

lemma ETC.arm_zero (hK : 0 < K) (m : ℕ) (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] :
arm 0 =ᵐ[Bandit.trajMeasure (etcAlgorithm hK m) ν] fun h ↦ ⟨0, hK⟩ := by
suffices h : (Bandit.trajMeasure (etcAlgorithm hK m) ν).map (arm 0) = (etcAlgorithm hK m).p0 by
have h_eq : ∀ᵐ x ∂((Bandit.trajMeasure (etcAlgorithm hK m) ν).map (arm 0)), x = ⟨0, hK⟩ := by
rw [h]
simp [etcAlgorithm]
exact ae_of_ae_map (by fun_prop) h_eq
-- extract lemma
sorry

lemma ETC.arm_ae_eq_etcNextArm (hK : 0 < K) (m : ℕ) (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν]
(n : ℕ) :
arm (n + 1) =ᵐ[(Bandit.trajMeasure (etcAlgorithm hK m) ν)]
fun h ↦ etcNextArm hK m n (fun i ↦ h i) := by
sorry

end Bandits
Loading