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
165 changes: 153 additions & 12 deletions LeanBandits/Bandit.lean
Original file line number Diff line number Diff line change
Expand Up @@ -15,11 +15,16 @@ open MeasureTheory ProbabilityTheory Filter Real Finset

open scoped ENNReal NNReal

/-- Measurable equivalence between `Iic 0 → α` and `α`. -/
def MeasurableEquiv.piIicZero (α : Type*) [MeasurableSpace α] :
(Iic 0 → α) ≃ᵐ α :=
have : Unique (Iic 0) := by simp only [mem_Iic, nonpos_iff_eq_zero]; exact Unique.subtypeEq 0
MeasurableEquiv.funUnique _ _
instance : Unique (Iic 0) := by simp only [mem_Iic, nonpos_iff_eq_zero]; exact Unique.subtypeEq 0

lemma coe_default_Iic_zero : ((default : Iic 0) : ℕ) = 0 := by
calc _ = ((⟨0, by simp⟩ : Iic 0) : ℕ) := by congr; exact (Unique.eq_default _).symm
_ = _ := by simp

/-- Measurable equivalence between `Iic 0 → X i` and `X 0`. -/
def MeasurableEquiv.piIicZero (X : ℕ → Type*) [∀ n, MeasurableSpace (X n)] :
((i : Iic 0) → X i) ≃ᵐ X 0 :=
(MeasurableEquiv.piUnique _).trans (coe_default_Iic_zero.symm ▸ MeasurableEquiv.refl _)

namespace Bandits

Expand All @@ -39,6 +44,14 @@ structure Algorithm (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] wher
instance (alg : Algorithm α R) (n : ℕ) : IsMarkovKernel (alg.policy n) := alg.h_policy n
instance (alg : Algorithm α R) : IsProbabilityMeasure alg.p0 := alg.hp0

/-- A deterministic algorithm. -/
noncomputable
def detAlgorithm (nextArm : (n : ℕ) → (Iic n → α × R) → α) (h_next : ∀ n, Measurable (nextArm n))
(arm0 : α) :
Algorithm α R where
policy n := Kernel.deterministic (nextArm n) (h_next n)
p0 := Measure.dirac arm0

namespace Bandit

/-- Kernel describing the distribution of the next arm-reward pair given the history up to `n`. -/
Expand Down Expand Up @@ -71,7 +84,7 @@ deriving IsMarkovKernel
/-- Measure on the sequence of arms pulled and rewards observed generated by the bandit. -/
noncomputable
def trajMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : Measure (ℕ → α × R) :=
(traj alg ν 0) ∘ₘ ((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero _).symm)
(traj alg ν 0) ∘ₘ ((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero (fun _ ↦ α × R)).symm)

/-- Measure of an infinite stream of rewards from each arm. -/
noncomputable
Expand All @@ -82,7 +95,8 @@ deriving IsProbabilityMeasure
instance (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
IsProbabilityMeasure (trajMeasure alg ν) := by
rw [trajMeasure]
have : IsProbabilityMeasure ((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero _).symm) :=
have : IsProbabilityMeasure
((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero (fun _ ↦ α × R)).symm) :=
isProbabilityMeasure_map <| by fun_prop
infer_instance

Expand Down Expand Up @@ -130,29 +144,125 @@ lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦
@[fun_prop]
lemma measurable_hist (n : ℕ) : Measurable (hist n (α := α) (R := R)) := by unfold hist; fun_prop

lemma hist_eq_frestrictLe :
hist = Preorder.frestrictLe («π» := fun _ ↦ α × R) := by
ext n h i : 3
simp [hist, Preorder.frestrictLe]

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

section Traj

open Kernel Preorder

variable {X : ℕ → Type*} [∀ n, MeasurableSpace (X n)]
{κ : (n : ℕ) → Kernel ((i : { x // x ∈ Iic n }) → X i) (X (n + 1))} [∀ n, IsMarkovKernel (κ n)]

lemma Measure.compProd_map {X Y Z : Type*} {mX : MeasurableSpace X} {mY : MeasurableSpace Y}
{mZ : MeasurableSpace Z} {μ : Measure X} {κ : Kernel X Y} [SFinite μ] [IsSFiniteKernel κ]
{f : Y → Z} (hf : Measurable f) :
μ ⊗ₘ (κ.map f) = (μ ⊗ₘ κ).map (Prod.map id f) := by
calc μ ⊗ₘ (κ.map f)
_ = (Kernel.id ∥ₖ Kernel.deterministic f hf) ∘ₘ (Kernel.id ×ₖ κ) ∘ₘ μ := by
rw [Measure.comp_assoc, Kernel.parallelComp_comp_prod, Measure.compProd_eq_comp_prod,
Kernel.id_comp, Kernel.deterministic_comp_eq_map]
_ = (Kernel.id ∥ₖ Kernel.deterministic f hf) ∘ₘ (μ ⊗ₘ κ) := by rw [Measure.compProd_eq_comp_prod]
_ = (μ ⊗ₘ κ).map (Prod.map id f) := by
rw [Kernel.id, Kernel.deterministic_parallelComp_deterministic,
Measure.deterministic_comp_eq_map]

lemma partialTraj_compProd_eq_traj_map_frestrictLe (a : ℕ) (x₀ : (i : Iic 0) → X i) :
(partialTraj κ 0 a x₀) ⊗ₘ (κ a) =
(traj κ 0 x₀).map (fun x ↦ (frestrictLe a x, x (a + 1))) := by
have h1 := partialTraj_compProd_traj (κ := κ) (zero_le a) x₀
have h2 : (fun x : Π n, X n ↦ (frestrictLe a x, x (a + 1))) =
(Prod.map id (fun x ↦ x (a + 1))) ∘ (fun x ↦ (frestrictLe a x, x)) := by ext <;> simp
rw [h2, ← Measure.map_map (by fun_prop) (by fun_prop), ← h1, ← Measure.compProd_map (by fun_prop)]
congr
have : (fun x : Π n, X n ↦ x (a + 1)) =
(fun x : Π i : Iic (a + 1), X i ↦ x ⟨a+1, by simp⟩) ∘ (frestrictLe (a + 1)) := by ext; simp
rw [this, map_comp_right _ (by fun_prop) (by fun_prop), traj_map_frestrictLe,
partialTraj_succ_self, ← map_comp_right _ (by fun_prop) (by fun_prop)]
have : (fun x : Π i : Iic (a + 1), X i ↦ x ⟨a+1, by simp⟩) ∘ IicProdIoc a (a + 1)
= (MeasurableEquiv.piSingleton a).symm ∘ Prod.snd := by
ext; simp [_root_.IicProdIoc, MeasurableEquiv.piSingleton]
rw [this, map_comp_right _ (by fun_prop) (by fun_prop), ← snd_eq, snd_prod,
← map_comp_right _ (by fun_prop) (by fun_prop)]
simp

lemma traj_cond_lemma1 {a : ℕ} (μ₀ : Measure ((i : Iic 0) → X i)) [IsFiniteMeasure μ₀] :
(traj κ 0 ∘ₘ μ₀).map (fun x ↦ (frestrictLe a x, x (a + 1)))
= (traj κ 0 ∘ₘ μ₀).map (frestrictLe a) ⊗ₘ κ a := by
rw [Measure.compProd_eq_comp_prod, Measure.map_comp _ _ (by fun_prop),
Measure.map_comp _ _ (by fun_prop), Measure.comp_assoc, traj_map_frestrictLe]
congr
ext x₀ : 1
rw [ProbabilityTheory.Kernel.comp_apply, ← Measure.compProd_eq_comp_prod]
symm
rw [Kernel.map_apply _ (by fun_prop)]
exact partialTraj_compProd_eq_traj_map_frestrictLe a x₀

lemma condDistrib_lemma (μ₀ : Measure ((i : Iic 0) → X i)) [IsFiniteMeasure μ₀] (a : ℕ)
[Nonempty (X (a + 1))] [StandardBorelSpace (X (a + 1))] :
condDistrib (fun x ↦ x (a + 1)) (frestrictLe a) (traj κ 0 ∘ₘ μ₀)
=ᵐ[(traj κ 0 ∘ₘ μ₀).map (frestrictLe a)] κ a := by
symm
exact condDistrib_ae_eq_of_measure_eq_compProd (by fun_prop) (by fun_prop) _ (traj_cond_lemma1 μ₀)

lemma traj_zero_map_eval_zero :
(Kernel.traj κ 0).map (fun h ↦ h 0)
= Kernel.deterministic (MeasurableEquiv.piIicZero X)
(MeasurableEquiv.piIicZero X).measurable := by
suffices (Kernel.traj κ 0).map (fun h ↦ h 0) = (Kernel.partialTraj κ 0 0).map
(MeasurableEquiv.piIicZero X) by
rwa [Kernel.partialTraj_zero,
Kernel.deterministic_map _ (MeasurableEquiv.piIicZero X).measurable] at this
rw [← Kernel.traj_map_frestrictLe, ← Kernel.map_comp_right _ (by fun_prop) (by fun_prop)]
congr with h
sorry

end Traj

lemma condDistrib_arm_reward [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
(alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
condDistrib (fun h ↦ (arm (n + 1) h, reward (n + 1) h)) (hist n) (Bandit.trajMeasure alg ν)
=ᵐ[(Bandit.trajMeasure alg ν).map (hist n)] Bandit.stepKernel alg ν n := by
sorry
=ᵐ[(Bandit.trajMeasure alg ν).map (hist n)] Bandit.stepKernel alg ν n :=
condDistrib_lemma (X := fun _ ↦ α × R)
((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero (fun _ ↦ α × R)).symm)
(κ := Bandit.stepKernel alg ν) n

lemma condDistrib_reward [StandardBorelSpace R] [Nonempty R] (alg : Algorithm α R) (ν : Kernel α R)
[IsMarkovKernel ν] (n : ℕ) :
lemma condDistrib_reward [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
(alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
condDistrib (reward n) (arm n) (Bandit.trajMeasure alg ν)
=ᵐ[(Bandit.trajMeasure alg ν).map (arm n)] ν := by
sorry
cases n with
| zero => sorry
| succ n =>
have h_ar := condDistrib_arm_reward alg ν n
have h_prod := condDistrib_prod_left (X := arm (n + 1)) (Y := reward (n + 1))
(T := hist n) (μ := Bandit.trajMeasure alg ν) (by fun_prop) (by fun_prop) (by fun_prop)
sorry

lemma condDistrib_arm [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
(alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
condDistrib (arm (n + 1)) (hist n) (Bandit.trajMeasure alg ν)
=ᵐ[(Bandit.trajMeasure alg ν).map (hist n)] alg.policy n := by
sorry

lemma hasLaw_step_zero
(alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
HasLaw (fun h : ℕ → α × R ↦ h 0) (alg.p0 ⊗ₘ ν) (Bandit.trajMeasure alg ν) where
aemeasurable := Measurable.aemeasurable (by fun_prop)
map_eq := by
simp only [Bandit.trajMeasure]
rw [← Measure.deterministic_comp_eq_map (by fun_prop), Measure.comp_assoc,
Kernel.deterministic_comp_eq_map, Bandit.traj, traj_zero_map_eval_zero,
Measure.deterministic_comp_eq_map, Measure.map_map (by fun_prop) (by fun_prop)]
simp

lemma hasLaw_arm_zero [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
(alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
HasLaw (arm 0) alg.p0 (Bandit.trajMeasure alg ν) where
Expand All @@ -168,6 +278,37 @@ lemma condIndepFun_reward_hist_arm [StandardBorelSpace α] [Nonempty α]
rw [condIndepFun_iff_condDistrib_prod_ae_eq_prodMkLeft (by fun_prop) (by fun_prop) (by fun_prop)]
sorry

section DetAlgorithm

variable [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
{nextArm : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextArm n)}
{arm0 : α} {ν : Kernel α R} [IsMarkovKernel ν]

lemma HasLaw_arm_zero_detAlgorithm :
HasLaw (arm 0) (Measure.dirac arm0)
(Bandit.trajMeasure (detAlgorithm nextArm h_next arm0) ν) where
map_eq := (hasLaw_arm_zero _ _).map_eq

lemma arm_zero_detAlgorithm :
arm 0 =ᵐ[Bandit.trajMeasure (detAlgorithm nextArm h_next arm0) ν] fun _ ↦ arm0 := by
have h_eq : ∀ᵐ x ∂((Bandit.trajMeasure (detAlgorithm nextArm h_next arm0) ν).map (arm 0)), x
= arm0 := by
rw [(hasLaw_arm_zero _ _).map_eq]
simp [detAlgorithm]
exact ae_of_ae_map (by fun_prop) h_eq

lemma arm_detAlgorithm_ae_eq (n : ℕ) :
arm (n + 1) =ᵐ[Bandit.trajMeasure (detAlgorithm nextArm h_next arm0) ν]
fun h ↦ nextArm n (fun i ↦ h i) := by
sorry

example : ∀ᵐ h ∂(Bandit.trajMeasure (detAlgorithm nextArm h_next arm0) ν),
arm 0 h = arm0 ∧ ∀ n, arm (n + 1) h = nextArm n (fun i ↦ h i) := by
rw [eventually_and, ae_all_iff]
exact ⟨arm_zero_detAlgorithm, arm_detAlgorithm_ae_eq⟩

end DetAlgorithm

end MeasureSpace

end Bandits
15 changes: 6 additions & 9 deletions LeanBandits/ETC.lean
Original file line number Diff line number Diff line change
Expand Up @@ -38,22 +38,19 @@ lemma measurable_etcNextArm (hK : 0 < K) (m n : ℕ) : Measurable (etcNextArm hK

/-- 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⟩
def etcAlgorithm (hK : 0 < K) (m : ℕ) : Algorithm (Fin K) ℝ :=
detAlgorithm (etcNextArm hK m) (by fun_prop) ⟨0, hK⟩

lemma ETC.arm_zero (hK : 0 < K) (m : ℕ) (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] :
arm 0 =ᵐ[Bandit.trajMeasure (etcAlgorithm hK m) ν] fun _ ↦ ⟨0, hK⟩ := by
have h_eq : ∀ᵐ x ∂((Bandit.trajMeasure (etcAlgorithm hK m) ν).map (arm 0)), x = ⟨0, hK⟩ := by
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
rw [(hasLaw_arm_zero _ _).map_eq]
simp [etcAlgorithm]
exact ae_of_ae_map (by fun_prop) h_eq
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
exact arm_zero_detAlgorithm

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
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
exact arm_detAlgorithm_ae_eq n

end Bandits
28 changes: 28 additions & 0 deletions LeanBandits/ForMathlib/CondDistrib.lean
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,14 @@ end MeasureTheory.Measure

namespace ProbabilityTheory

lemma Kernel.deterministic_parallelComp_deterministic
{f : α → γ} {g : β → δ} (hf : Measurable f) (hg : Measurable g) :
(deterministic f hf) ∥ₖ (deterministic g hg)
= deterministic (Prod.map f g) (hf.prodMap hg) := by
ext x : 1
rw [parallelComp_apply, deterministic_apply, deterministic_apply, deterministic_apply, Prod.map,
Measure.dirac_prod_dirac]

lemma Kernel.prod_apply_prod {κ : Kernel α β} {η : Kernel α γ}
[IsSFiniteKernel κ] [IsSFiniteKernel η] {s : Set β} {t : Set γ} {a : α} :
(κ ×ₖ η) a (s ×ˢ t) = (κ a s) * (η a t) := by
Expand Down Expand Up @@ -513,6 +521,26 @@ lemma condDistrib_fst_prod (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ)
fun_prop
· fun_prop

lemma Measure.compProd_assoc {μ : Measure α} {κ : Kernel α β} {η : Kernel (α × β) γ}
[SFinite μ] [IsSFiniteKernel κ] [IsSFiniteKernel η] :
(μ ⊗ₘ κ) ⊗ₘ η = (μ ⊗ₘ (κ ⊗ₖ η)).map MeasurableEquiv.prodAssoc.symm := by
sorry

lemma Measure.compProd_assoc' {μ : Measure α} {κ : Kernel α β} {η : Kernel (α × β) γ}
[SFinite μ] [IsSFiniteKernel κ] [IsSFiniteKernel η] :
μ ⊗ₘ (κ ⊗ₖ η) = ((μ ⊗ₘ κ) ⊗ₘ η).map MeasurableEquiv.prodAssoc := by
simp [Measure.compProd_assoc]

lemma condDistrib_prod_left [StandardBorelSpace β] [Nonempty β]
(hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) (hT : AEMeasurable T μ) :
condDistrib (fun ω ↦ (X ω, Y ω)) T μ
=ᵐ[μ.map T] condDistrib X T μ ⊗ₖ condDistrib Y (fun ω ↦ (T ω, X ω)) μ := by
refine condDistrib_ae_eq_of_measure_eq_compProd₀ (μ := μ) hT (by fun_prop)
(condDistrib X T μ ⊗ₖ condDistrib Y (fun ω ↦ (T ω, X ω)) μ) ?_
rw [Measure.compProd_assoc', compProd_map_condDistrib hX, compProd_map_condDistrib hY,
AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)]
rfl

end CondDistrib

section Cond
Expand Down
4 changes: 2 additions & 2 deletions LeanBandits/RewardByCountMeasure.lean
Original file line number Diff line number Diff line change
Expand Up @@ -102,7 +102,7 @@ notation "𝓛[" Y " | " X " ← " x "; " μ "]" => Measure.map Y (μ[|X ⁻¹'
notation "𝓛[" Y " | " X "; " μ "]" => condDistrib Y X μ

omit [DecidableEq α] [MeasurableSingletonClass α] in
lemma condDistrib_reward' (n : ℕ) :
lemma condDistrib_reward' [StandardBorelSpace α] [Nonempty α] (n : ℕ) :
𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1; Bandit.measure alg ν]
=ᵐ[(Bandit.measure alg ν).map (fun ω ↦ arm n ω.1)] ν := by
let μ := Bandit.measure alg ν
Expand All @@ -122,7 +122,7 @@ lemma condDistrib_reward' (n : ℕ) :
rw [h_prod, h_eq]

omit [DecidableEq α] in
lemma reward_cond_arm [Countable α] (a : α) (n : ℕ)
lemma reward_cond_arm [StandardBorelSpace α] [Nonempty α] [Countable α] (a : α) (n : ℕ)
(hμa : (Bandit.measure alg ν).map (fun ω ↦ arm n ω.1) {a} ≠ 0) :
𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1 ← a; Bandit.measure alg ν] = ν a := by
let μ := Bandit.measure alg ν
Expand Down
5 changes: 2 additions & 3 deletions LeanBandits/UCB.lean
Original file line number Diff line number Diff line change
Expand Up @@ -43,9 +43,8 @@ lemma measurable_ucbNextArm (c : ℝ) (n : ℕ) : Measurable (ucbNextArm c n (α

/-- The UCB algorithm. -/
noncomputable
def ucbAlgorithm (c : ℝ) : Algorithm α ℝ where
policy n := Kernel.deterministic (ucbNextArm c n) (by fun_prop)
p0 := Measure.dirac (Classical.arbitrary α)
def ucbAlgorithm (c : ℝ) : Algorithm α ℝ :=
detAlgorithm (ucbNextArm c) (by fun_prop) (Classical.arbitrary α)

end Algorithm

Expand Down