Skip to content
3 changes: 3 additions & 0 deletions LeanBandits.lean
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,9 @@ import LeanBandits.AlgorithmBuilding
import LeanBandits.Bandit
import LeanBandits.ETC
import LeanBandits.ForMathlib.CondDistrib
import LeanBandits.ForMathlib.KernelCompositionLemmas
import LeanBandits.ForMathlib.KernelCompositionParallelComp
import LeanBandits.ForMathlib.Traj
import LeanBandits.Regret
import LeanBandits.RewardByCountMeasure
import LeanBandits.UCB
84 changes: 7 additions & 77 deletions LeanBandits/Bandit.lean
Original file line number Diff line number Diff line change
Expand Up @@ -5,27 +5,17 @@ Authors: Rémy Degenne, Paulo Rauber
-/
import Mathlib
import LeanBandits.ForMathlib.CondDistrib
import LeanBandits.ForMathlib.KernelCompositionLemmas
import LeanBandits.ForMathlib.Traj

/-!
# Bandit

-/

open MeasureTheory ProbabilityTheory Filter Real Finset

open scoped ENNReal NNReal

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

variable {α R : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R}
Expand Down Expand Up @@ -84,22 +74,15 @@ 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 (fun _ ↦ α × R)).symm)
Kernel.trajMeasure (alg.p0 ⊗ₘ ν) (stepKernel alg ν)
deriving IsProbabilityMeasure

/-- Measure of an infinite stream of rewards from each arm. -/
noncomputable
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 (fun _ ↦ α × R)).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
Expand Down Expand Up @@ -171,57 +154,6 @@ 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)
Expand All @@ -240,9 +172,7 @@ lemma condDistrib_arm_reward [StandardBorelSpace α] [Nonempty α] [StandardBore
(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 :=
condDistrib_lemma (X := fun _ ↦ α × R)
((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero (fun _ ↦ α × R)).symm)
(κ := Bandit.stepKernel alg ν) n
Kernel.condDistrib_trajMeasure_ae_eq_kernel

lemma condDistrib_reward [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
(alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
Expand All @@ -267,9 +197,9 @@ lemma hasLaw_step_zero
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]
simp only [Bandit.trajMeasure, Kernel.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,
Kernel.deterministic_comp_eq_map, traj_zero_map_eval_zero,
Measure.deterministic_comp_eq_map, Measure.map_map (by fun_prop) (by fun_prop)]
simp

Expand Down
10 changes: 1 addition & 9 deletions LeanBandits/ForMathlib/CondDistrib.lean
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ import Mathlib.Probability.Independence.Conditional
import Mathlib.Probability.Kernel.Composition.Lemmas
import Mathlib.Probability.Kernel.CompProdEqIff
import Mathlib.Probability.Kernel.Condexp

import LeanBandits.ForMathlib.KernelCompositionParallelComp

open MeasureTheory ProbabilityTheory Finset
open scoped ENNReal NNReal
Expand Down Expand Up @@ -68,14 +68,6 @@ 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
21 changes: 21 additions & 0 deletions LeanBandits/ForMathlib/KernelCompositionLemmas.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,21 @@
import Mathlib.Probability.Kernel.Composition.Lemmas
import LeanBandits.ForMathlib.KernelCompositionParallelComp

open MeasureTheory ProbabilityTheory
open scoped ENNReal

variable {α β γ δ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β}
{mγ : MeasurableSpace γ} {mδ : MeasurableSpace δ}
{μ : Measure α} {ν : Measure β} {κ : Kernel α β}

-- PR: https://github.com/leanprover-community/mathlib4/pull/29555
lemma MeasureTheory.Measure.compProd_map [SFinite μ] [IsSFiniteKernel κ]
{f : β → γ} (hf : Measurable f) :
μ ⊗ₘ (κ.map f) = (μ ⊗ₘ κ).map (Prod.map id f) := by
calc μ ⊗ₘ (κ.map f)
_ = (Kernel.id ∥ₖ Kernel.deterministic f hf) ∘ₘ (Kernel.id ×ₖ κ) ∘ₘ μ := by
rw [comp_assoc, Kernel.parallelComp_comp_prod, compProd_eq_comp_prod,
Kernel.id_comp, Kernel.deterministic_comp_eq_map]
_ = (Kernel.id ∥ₖ Kernel.deterministic f hf) ∘ₘ (μ ⊗ₘ κ) := by rw [compProd_eq_comp_prod]
_ = (μ ⊗ₘ κ).map (Prod.map id f) := by
rw [Kernel.id, Kernel.deterministic_parallelComp_deterministic, deterministic_comp_eq_map]
26 changes: 26 additions & 0 deletions LeanBandits/ForMathlib/KernelCompositionParallelComp.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
import Mathlib.Probability.Kernel.Composition.ParallelComp

open MeasureTheory
open scoped ENNReal

namespace ProbabilityTheory.Kernel

variable {α β γ δ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β}
{mγ : MeasurableSpace γ} {mδ : MeasurableSpace δ}
{κ : Kernel α β} {η : Kernel γ δ} {x : α × γ}

-- PR: https://github.com/leanprover-community/mathlib4/pull/29555
lemma parallelComp_apply_prod [IsSFiniteKernel κ] [IsSFiniteKernel η] (s : Set β) (t : Set δ) :
(κ ∥ₖ η) x (s ×ˢ t) = (κ x.1 s) * (η x.2 t) := by
rw [parallelComp_apply, Measure.prod_prod]

-- PR: https://github.com/leanprover-community/mathlib4/pull/29555
lemma 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]

end ProbabilityTheory.Kernel
87 changes: 87 additions & 0 deletions LeanBandits/ForMathlib/Traj.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
import Mathlib.Probability.Kernel.IonescuTulcea.Traj
import Mathlib.Probability.Kernel.CondDistrib
import LeanBandits.ForMathlib.CondDistrib
import LeanBandits.ForMathlib.KernelCompositionLemmas

open Filter Finset Function MeasurableEquiv MeasurableSpace MeasureTheory Preorder ProbabilityTheory

variable {X : ℕ → Type*} [∀ n, MeasurableSpace (X n)]
variable {κ : (n : ℕ) → Kernel (Π i : Iic n, X i) (X (n + 1))} [∀ n, IsMarkovKernel (κ n)]
variable {μ₀ : Measure (X 0)} [IsProbabilityMeasure μ₀]

section MeasurableEquiv

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 _)

end MeasurableEquiv

namespace ProbabilityTheory.Kernel

-- Probability/Kernel/IonescuTulcea/Traj.lean
/-- Distribution of the infinite trajectory given the distribution of `X 0`. -/
noncomputable
def trajMeasure (μ₀ : Measure (X 0)) (κ : (n : ℕ) → Kernel (Π i : Iic n, X i) (X (n + 1)))
[∀ n, IsMarkovKernel (κ n)] :
Measure (Π n, X n) :=
(traj κ 0) ∘ₘ (μ₀.map (MeasurableEquiv.piIicZero _).symm)

-- Probability/Kernel/IonescuTulcea/Traj.lean
instance : IsProbabilityMeasure (trajMeasure μ₀ κ) := by
rw [trajMeasure]
have : IsProbabilityMeasure (μ₀.map (MeasurableEquiv.piIicZero _).symm) :=
isProbabilityMeasure_map <| by fun_prop
infer_instance

-- Probability/Kernel/IonescuTulcea/Traj.lean
lemma traj_map_eq_kernel {a : ℕ} : (traj κ a).map (fun x ↦ x (a + 1)) = κ a := by
set f : (Π n, X n) → X (a + 1) := fun x ↦ x (a + 1)
set g : (Π n : Iic (a + 1), X n) → X (a + 1) := fun x ↦ x ⟨a + 1, by simp⟩
have hf : f = g ∘ (frestrictLe (a + 1)) := by rfl
have hp : g ∘ IicProdIoc a (a + 1) = (piSingleton a).symm ∘ Prod.snd := by
ext
simp [g, _root_.IicProdIoc, piSingleton]
rw [hf, map_comp_right, traj_map_frestrictLe, partialTraj_succ_self, ← map_comp_right, hp,
map_comp_right, ← snd_eq, snd_prod, ← map_comp_right]
all_goals measurability

-- Probability/Kernel/IonescuTulcea/Traj.lean
lemma partialTraj_compProd_kernel_eq_traj_map {a : ℕ} {x₀ : Π n : Iic 0, X n} :
(partialTraj κ 0 a x₀) ⊗ₘ (κ a) = (traj κ 0 x₀).map (fun x ↦ (frestrictLe a x, x (a + 1))) := by
set f := fun x ↦ (frestrictLe a x, x (a + 1))
set g := fun x ↦ (frestrictLe a x, x)
have hf : f = (Prod.map id (fun x ↦ x (a + 1))) ∘ g := rfl
rw [hf, ← Measure.map_map, ← partialTraj_compProd_traj, ← MeasureTheory.Measure.compProd_map,
traj_map_eq_kernel]
all_goals measurability

-- (Extract kernel lemmas from rewrites?) Probability/Kernel/IonescuTulcea/Traj.lean
lemma trajMeasure_map_frestrictLe_compProd_kernel_eq_trajMeasure_map {a : ℕ} :
(trajMeasure μ₀ κ).map (frestrictLe a) ⊗ₘ κ a =
(trajMeasure μ₀ κ).map (fun x ↦ (frestrictLe a x, x (a + 1))) := by
rw [Measure.compProd_eq_comp_prod, trajMeasure, Measure.map_comp, traj_map_frestrictLe,
Measure.comp_assoc, Measure.map_comp]
any_goals fun_prop
congr
ext1 x₀
rw [comp_apply, ← Measure.compProd_eq_comp_prod, map_apply,
partialTraj_compProd_kernel_eq_traj_map]
fun_prop

-- Probability/Kernel/IonescuTulcea/Traj.lean
lemma condDistrib_trajMeasure_ae_eq_kernel {a : ℕ}
[StandardBorelSpace (X (a + 1))] [Nonempty (X (a + 1))] :
condDistrib (fun x ↦ x (a + 1)) (frestrictLe a) (trajMeasure μ₀ κ)
=ᵐ[(trajMeasure μ₀ κ).map (frestrictLe a)] κ a := by
apply condDistrib_ae_eq_of_measure_eq_compProd₀ (by measurability) (by measurability)
exact trajMeasure_map_frestrictLe_compProd_kernel_eq_trajMeasure_map.symm

end ProbabilityTheory.Kernel