diff --git a/LeanBandits/Bandit.lean b/LeanBandits/Bandit.lean index f5e6c131..07a3f5d4 100644 --- a/LeanBandits/Bandit.lean +++ b/LeanBandits/Bandit.lean @@ -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 @@ -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`. -/ @@ -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 @@ -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 @@ -130,22 +144,107 @@ 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 : ℕ) : @@ -153,6 +252,17 @@ lemma condDistrib_arm [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace =ᵐ[(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 @@ -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 diff --git a/LeanBandits/ETC.lean b/LeanBandits/ETC.lean index 6af21402..d7defd78 100644 --- a/LeanBandits/ETC.lean +++ b/LeanBandits/ETC.lean @@ -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 diff --git a/LeanBandits/ForMathlib/CondDistrib.lean b/LeanBandits/ForMathlib/CondDistrib.lean index 00ddff42..0e069f4e 100644 --- a/LeanBandits/ForMathlib/CondDistrib.lean +++ b/LeanBandits/ForMathlib/CondDistrib.lean @@ -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 @@ -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 diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 99224810..cb581c7b 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -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 ν @@ -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 ν diff --git a/LeanBandits/UCB.lean b/LeanBandits/UCB.lean index d188a6c3..3ee285ef 100644 --- a/LeanBandits/UCB.lean +++ b/LeanBandits/UCB.lean @@ -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