diff --git a/LeanBandits.lean b/LeanBandits.lean index ab3c610e..c2443ec0 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -1,6 +1,7 @@ import LeanBandits.AlgorithmBuilding import LeanBandits.Bandit import LeanBandits.ETC +import LeanBandits.ForMathlib.CondDistrib import LeanBandits.Regret import LeanBandits.RewardByCountMeasure import LeanBandits.UCB diff --git a/LeanBandits/Bandit.lean b/LeanBandits/Bandit.lean index eb25efc4..7f981c3f 100644 --- a/LeanBandits/Bandit.lean +++ b/LeanBandits/Bandit.lean @@ -130,7 +130,7 @@ lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ lemma measurable_hist (n : ℕ) : Measurable (hist n (α := α) (R := R)) := by unfold hist; fun_prop /-- Filtration of the bandit process. -/ -def ℱ (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : +protected def filtration (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) := MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R) @@ -152,6 +152,19 @@ lemma condDistrib_arm [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace =ᵐ[(Bandit.trajMeasure alg ν).map (hist n)] alg.policy n := by sorry +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 + map_eq := by + sorry + +/-- The reward at time `n+1` is independent of the history up to time `n` given the arm at `n+1`. -/ +lemma condIndepFun_reward_hist_arm [StandardBorelSpace α] [StandardBorelSpace R] + {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) : + CondIndepFun (MeasurableSpace.comap (arm (n + 1)) inferInstance) + (measurable_arm _).comap_le (reward (n + 1)) (hist n) (Bandit.trajMeasure alg ν) := by + sorry + end MeasureSpace end Bandits diff --git a/LeanBandits/ETC.lean b/LeanBandits/ETC.lean index b805f469..6af21402 100644 --- a/LeanBandits/ETC.lean +++ b/LeanBandits/ETC.lean @@ -43,14 +43,12 @@ def etcAlgorithm (hK : 0 < K) (m : ℕ) : Algorithm (Fin K) ℝ where 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 + 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 lemma ETC.arm_ae_eq_etcNextArm (hK : 0 < K) (m : ℕ) (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] (n : ℕ) : diff --git a/LeanBandits/ForMathlib/CondDistrib.lean b/LeanBandits/ForMathlib/CondDistrib.lean new file mode 100644 index 00000000..f1a8b6db --- /dev/null +++ b/LeanBandits/ForMathlib/CondDistrib.lean @@ -0,0 +1,457 @@ +/- +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.Independence.Basic +import Mathlib.Probability.Independence.Conditional +import Mathlib.Probability.Kernel.Composition.Lemmas +import Mathlib.Probability.Kernel.CompProdEqIff +import Mathlib.Probability.Kernel.Condexp + + +open MeasureTheory ProbabilityTheory Finset +open scoped ENNReal NNReal + +variable {α β γ Ω Ω' : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] + {mα : MeasurableSpace α} {μ : Measure α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} + [MeasurableSpace Ω'] [StandardBorelSpace Ω'] [Nonempty Ω'] + {X : α → β} {Y : α → Ω} {Z : α → Ω'} {T : α → γ} + +@[fun_prop] +lemma Measurable.coe_nat_enat {f : α → ℕ} (hf : Measurable f) : + Measurable (fun a ↦ (f a : ℕ∞)) := Measurable.comp (by fun_prop) hf + +@[fun_prop] +lemma Measurable.toNat {f : α → ℕ∞} (hf : Measurable f) : Measurable (fun a ↦ (f a).toNat) := + Measurable.comp (by fun_prop) hf + +namespace MeasureTheory.Measure + +lemma comp_congr {κ η : Kernel α β} (h : ∀ᵐ a ∂μ, κ a = η a) : + κ ∘ₘ μ = η ∘ₘ μ := + bind_congr_right h + +lemma copy_comp_map (hX : AEMeasurable X μ) : + Kernel.copy β ∘ₘ (μ.map X) = μ.map (fun a ↦ (X a, X a)) := by + rw [Kernel.copy, deterministic_comp_eq_map, AEMeasurable.map_map_of_aemeasurable (by fun_prop) hX] + congr + +lemma compProd_deterministic [SFinite μ] (hX : Measurable X) : + μ ⊗ₘ Kernel.deterministic X hX = μ.map (fun a ↦ (a, X a)) := by + rw [compProd_eq_comp_prod, Kernel.id, Kernel.deterministic_prod_deterministic, + deterministic_comp_eq_map] + rfl + +lemma trim_comap_apply (hX : Measurable X) {s : Set β} (hs : MeasurableSet s) : + μ.trim hX.comap_le (X ⁻¹' s) = μ.map X s := by + rw [trim_measurableSet_eq, Measure.map_apply (by fun_prop) hs] + exact ⟨s, hs, rfl⟩ + +lemma ext_prod₃ {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {μ ν : Measure (α × β × γ)} [IsFiniteMeasure μ] [IsFiniteMeasure ν] + (h : ∀ {s : Set α} {t : Set β} {u : Set γ} (hs : MeasurableSet s) (ht : MeasurableSet t) + (hu : MeasurableSet u), μ (s ×ˢ t ×ˢ u) = ν (s ×ˢ t ×ˢ u)) : + μ = ν := by + sorry + +lemma ext_prod₃_iff {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {μ ν : Measure (α × β × γ)} [IsFiniteMeasure μ] [IsFiniteMeasure ν] : + μ = ν ↔ (∀ {s : Set α} {t : Set β} {u : Set γ}, + MeasurableSet s → MeasurableSet t → MeasurableSet u → + μ (s ×ˢ t ×ˢ u) = ν (s ×ˢ t ×ˢ u)) := + ⟨fun h s t u hs ht hu ↦ by rw [h], Measure.ext_prod₃⟩ + +end MeasureTheory.Measure + +namespace ProbabilityTheory + +lemma Kernel.prod_apply_prod {κ : Kernel α β} {η : Kernel α γ} + [IsSFiniteKernel κ] [IsSFiniteKernel η] {s : Set β} {t : Set γ} {a : α} : + (κ ×ₖ η) a (s ×ˢ t) = (κ a s) * (η a t) := by + rw [Kernel.prod_apply, Measure.prod_prod] + +lemma CondIndepFun.prod_right + {α β γ δ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} + {mδ : MeasurableSpace δ} [StandardBorelSpace α] + {μ : Measure α} [IsFiniteMeasure μ] {X : α → β} {Y : α → γ} {Z : α → δ} + (hX : Measurable X) (hY : Measurable Y) (hZ : Measurable Z) + (h : CondIndepFun (mδ.comap Z) hZ.comap_le X Y μ) : + CondIndepFun (mδ.comap Z) hZ.comap_le X (fun ω ↦ (Y ω, Z ω)) μ := by + sorry + +section CondDistrib + +variable [IsFiniteMeasure μ] + +lemma condDistrib_comp_map (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) : + condDistrib Y X μ ∘ₘ (μ.map X) = μ.map Y := by + rw [← Measure.snd_compProd, compProd_map_condDistrib hY, Measure.snd_map_prodMk₀ hX] + +lemma condDistrib_congr {X' : α → β} {Y' : α → Ω} (hY : Y =ᵐ[μ] Y') (hX : X =ᵐ[μ] X') : + condDistrib Y X μ = condDistrib Y' X' μ := by + rw [condDistrib, condDistrib] + congr 1 + rw [Measure.map_congr] + filter_upwards [hX, hY] with a ha hb using by rw [ha, hb] + +lemma condDistrib_congr_right {X' : α → β} (hX : X =ᵐ[μ] X') : + condDistrib Y X μ = condDistrib Y X' μ := + condDistrib_congr (by rfl) hX + +lemma condDistrib_congr_left {Y' : α → Ω} (hY : Y =ᵐ[μ] Y') : + condDistrib Y X μ = condDistrib Y' X μ := + condDistrib_congr hY (by rfl) + +lemma condDistrib_ae_eq_of_measure_eq_compProd₀ + (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) (κ : Kernel β Ω) [IsFiniteKernel κ] + (hκ : μ.map (fun x => (X x, Y x)) = μ.map X ⊗ₘ κ) : + ∀ᵐ x ∂μ.map X, κ x = condDistrib Y X μ x := by + suffices ∀ᵐ x ∂μ.map (hX.mk X), κ x = condDistrib (hY.mk Y) (hX.mk X) μ x by + rw [Measure.map_congr hX.ae_eq_mk] + convert this using 3 with b + rw [condDistrib_congr hY.ae_eq_mk hX.ae_eq_mk] + refine condDistrib_ae_eq_of_measure_eq_compProd (μ := μ) hX.measurable_mk hY.measurable_mk κ + ((Eq.trans ?_ hκ).trans ?_) + · refine Measure.map_congr ?_ + filter_upwards [hX.ae_eq_mk, hY.ae_eq_mk] with a haX haY using by rw [haX, haY] + · rw [Measure.map_congr hX.ae_eq_mk] + +lemma condDistrib_comp (hX : AEMeasurable X μ) {f : β → Ω} (hf : Measurable f) : + condDistrib (f ∘ X) X μ =ᵐ[μ.map X] Kernel.deterministic f hf := by + symm + refine condDistrib_ae_eq_of_measure_eq_compProd₀ hX (by fun_prop) _ ?_ + rw [Measure.compProd_deterministic, AEMeasurable.map_map_of_aemeasurable (by fun_prop) hX] + rfl + +lemma condDistrib_const (hX : AEMeasurable X μ) (c : Ω) : + condDistrib (fun _ ↦ c) X μ =ᵐ[μ.map X] Kernel.deterministic (fun _ ↦ c) (by fun_prop) := by + have : (fun _ : α ↦ c) = (fun _ : β ↦ c) ∘ X := rfl + conv_lhs => rw [this] + filter_upwards [condDistrib_comp hX (by fun_prop : Measurable (fun _ ↦ c))] with b hb + rw [hb] + +lemma condDistrib_of_indepFun (h : IndepFun X Y μ) (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) : + condDistrib Y X μ =ᵐ[μ.map X] Kernel.const β (μ.map Y) := by + symm + refine condDistrib_ae_eq_of_measure_eq_compProd₀ (μ := μ) hX hY _ ?_ + simp only [Measure.compProd_const] + exact (indepFun_iff_map_prod_eq_prod_map_map hX hY).mp h + +lemma indepFun_iff_condDistrib_eq_const (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) : + IndepFun X Y μ ↔ condDistrib Y X μ =ᵐ[μ.map X] Kernel.const β (μ.map Y) := by + refine ⟨fun h ↦ condDistrib_of_indepFun h hX hY, fun h ↦ ?_⟩ + rw [indepFun_iff_map_prod_eq_prod_map_map hX hY, ← compProd_map_condDistrib hY, + Measure.compProd_congr h] + simp + +-- todo: use this to refactor `indepFun_iff_map_prod_eq_prod_map_map` +theorem Kernel.indepFun_iff_map_prod_eq_prod_map_map {Ω' α β γ : Type*} + {mΩ' : MeasurableSpace Ω'} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {X : α → β} {T : α → γ} + {μ : Measure Ω'} [IsFiniteMeasure μ] + {κ : Kernel Ω' α} [IsFiniteKernel κ] + -- TODO: relax this to CountableOrCountablyGenerated once it is fixed + [StandardBorelSpace β] [StandardBorelSpace γ] + (hf : Measurable X) (hg : Measurable T) : + IndepFun X T κ μ ↔ κ.map (fun ω ↦ (X ω, T ω)) =ᵐ[μ] ((κ.map X) ×ₖ (κ.map T)) := by + classical + rw [indepFun_iff_measure_inter_preimage_eq_mul] + refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ + · rw [← Kernel.compProd_eq_iff] + have : (μ ⊗ₘ κ.map fun ω ↦ (X ω, T ω)) = μ ⊗ₘ (κ.map X ×ₖ κ.map T) + ↔ ∀ {u : Set Ω'} {s : Set β} {t : Set γ}, + MeasurableSet u → MeasurableSet s → MeasurableSet t → + (μ ⊗ₘ κ.map (fun ω ↦ (X ω, T ω))) (u ×ˢ s ×ˢ t) + = (μ ⊗ₘ (κ.map X ×ₖ κ.map T)) (u ×ˢ s ×ˢ t) := by + refine ⟨fun h ↦ by simp [h], fun h ↦ ?_⟩ + exact Measure.ext_prod₃ h + rw [this] + intro u s t hu hs ht + rw [Measure.compProd_apply (hu.prod (hs.prod ht)), + Measure.compProd_apply (hu.prod (hs.prod ht))] + refine lintegral_congr_ae ?_ + have h_set_eq ω : Prod.mk ω ⁻¹' u ×ˢ s ×ˢ t = if ω ∈ u then s ×ˢ t else ∅ := by ext; simp + simp_rw [h_set_eq] + filter_upwards [h s t hs ht] with ω hω + by_cases hωu : ω ∈ u + swap; · simp [hωu] + simp only [hωu, ↓reduceIte] + rw [Kernel.map_apply _ (by fun_prop), Measure.map_apply (by fun_prop) (hs.prod ht)] + rw [Set.mk_preimage_prod, hω, Kernel.prod_apply_prod, Kernel.map_apply' _ (by fun_prop), + Kernel.map_apply' _ (by fun_prop)] + exacts [ht, hs] + · intro s t hs ht + filter_upwards [h] with ω hω + calc (κ ω) (X ⁻¹' s ∩ T ⁻¹' t) + _ = (κ.map (fun ω ↦ (X ω, T ω))) ω (s ×ˢ t) := by + rw [← Kernel.deterministic_comp_eq_map, ← deterministic_prod_deterministic hf hg, + Kernel.comp_apply, Measure.bind_apply (hs.prod ht) (by fun_prop)] + simp_rw [Kernel.prod_apply_prod, Kernel.deterministic_apply' hf _ hs, + Kernel.deterministic_apply' hg _ ht] + calc (κ ω) (X ⁻¹' s ∩ T ⁻¹' t) + _ = ∫⁻ a, (X ⁻¹' s ∩ T ⁻¹' t).indicator (fun x ↦ 1) a ∂κ ω := by + simp [lintegral_indicator ((hf hs).inter (hg ht))] + _ = ∫⁻ a, (X ⁻¹' s).indicator (fun x ↦ 1) a * (T ⁻¹' t).indicator (fun x ↦ 1) a ∂κ ω := by + congr with a + simp only [Set.indicator_apply, Set.mem_inter_iff, Set.mem_preimage, mul_ite, mul_one, + mul_zero] + by_cases has : X a ∈ s <;> simp [has] + _ = ∫⁻ a, s.indicator (fun x ↦ 1) (X a) * t.indicator (fun x ↦ 1) (T a) ∂κ ω := rfl + _ = ((κ.map X) ×ₖ (κ.map T)) ω (s ×ˢ t) := by rw [hω] + _ = (κ ω) (X ⁻¹' s) * (κ ω) (T ⁻¹' t) := by + rw [Kernel.prod_apply_prod, Kernel.map_apply' _ (by fun_prop), + Kernel.map_apply' _ (by fun_prop)] + exacts [ht, hs] + +lemma Kernel.indepFun_iff_compProd_map_prod_eq_compProd_prod_map_map {Ω' α β γ : Type*} + {mΩ' : MeasurableSpace Ω'} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {X : α → β} {T : α → γ} + {μ : Measure Ω'} [IsFiniteMeasure μ] + {κ : Kernel Ω' α} [IsFiniteKernel κ] + -- TODO: relax this to CountableOrCountablyGenerated once it is fixed + [StandardBorelSpace β] [StandardBorelSpace γ] + (hf : Measurable X) (hg : Measurable T) : + IndepFun X T κ μ ↔ (μ ⊗ₘ κ.map fun ω ↦ (X ω, T ω)) = μ ⊗ₘ (κ.map X ×ₖ κ.map T) := by + rw [Kernel.indepFun_iff_map_prod_eq_prod_map_map hf hg, Kernel.compProd_eq_iff] + +theorem condIndepFun_iff_map_prod_eq_prod_map_map {α : Type*} {m mα : MeasurableSpace α} + [StandardBorelSpace α] {X : α → β} {T : α → γ} + {hm : m ≤ mα} {μ : Measure α} [IsFiniteMeasure μ] + -- TODO: relax this to CountableOrCountablyGenerated once it is fixed + [StandardBorelSpace β] [StandardBorelSpace γ] + (hX : Measurable X) (hT : Measurable T) : + CondIndepFun m hm X T μ + ↔ (condExpKernel μ m).map (fun ω ↦ (X ω, T ω)) + =ᵐ[μ.trim hm] (((condExpKernel μ m).map X) ×ₖ ((condExpKernel μ m).map T)) := + Kernel.indepFun_iff_map_prod_eq_prod_map_map hX hT + +lemma condIndepFun_iff_map_prod_eq_prod_comp_trim + {α : Type*} {m mα : MeasurableSpace α} [StandardBorelSpace α] {X : α → β} {T : α → γ} + {hm : m ≤ mα} {μ : Measure α} [IsFiniteMeasure μ] + -- TODO: relax this to CountableOrCountablyGenerated once it is fixed + [StandardBorelSpace β] [StandardBorelSpace γ] + (hX : Measurable X) (hT : Measurable T) : + CondIndepFun m hm X T μ + ↔ @Measure.map _ _ _ (m.prod _) (fun ω ↦ (ω, X ω, T ω)) μ + = (Kernel.id ×ₖ ((condExpKernel μ m).map X ×ₖ (condExpKernel μ m).map T)) ∘ₘ μ.trim hm := by + unfold CondIndepFun + rw [Kernel.indepFun_iff_compProd_map_prod_eq_compProd_prod_map_map hX hT] + congr! + · calc (μ.trim hm ⊗ₘ (condExpKernel μ m).map fun ω ↦ (X ω, T ω)) + _ = (Kernel.id ∥ₖ Kernel.deterministic (fun ω ↦ (X ω, T ω)) (by fun_prop)) + ∘ₘ (μ.trim hm ⊗ₘ (condExpKernel μ m)) := by + rw [Measure.compProd_eq_parallelComp_comp_copy_comp, ← Kernel.deterministic_comp_eq_map, + ← Kernel.parallelComp_id_left_comp_parallelComp, Measure.comp_assoc, Kernel.comp_assoc, + Kernel.parallelComp_comp_copy, ← Measure.comp_assoc, Measure.compProd_eq_comp_prod] + _ = (Kernel.id ∥ₖ Kernel.deterministic (fun ω ↦ (X ω, T ω)) (by fun_prop)) + ∘ₘ (@Measure.map _ _ mα (m.prod mα) (fun ω ↦ (ω, ω)) μ) := by + congr + exact compProd_trim_condExpKernel hm + _ = _ := by + rw [← Measure.deterministic_comp_eq_map, Measure.comp_assoc, + ← Kernel.deterministic_prod_deterministic (g := fun ω ↦ ω), + Kernel.parallelComp_comp_prod, Kernel.deterministic_comp_deterministic, Kernel.id_comp, + Kernel.deterministic_prod_deterministic, Measure.deterministic_comp_eq_map] + · rfl + · exact Measurable.mono measurable_id le_rfl hm + · fun_prop + · rw [Measure.compProd_eq_comp_prod] + +lemma condDistrib_apply_ae_eq_condExpKernel_map + {α : Type*} {mα : MeasurableSpace α} [StandardBorelSpace α] + [StandardBorelSpace β] [Nonempty β] + {X : α → β} {T : α → γ} {μ : Measure α} [IsFiniteMeasure μ] + (hX : Measurable X) (hT : Measurable T) {s : Set β} (hs : MeasurableSet s) : + (fun a ↦ condDistrib X T μ (T a) s) + =ᵐ[μ] fun a ↦ (condExpKernel μ (MeasurableSpace.comap T inferInstance)).map X a s := by + have hT_meas {s : Set γ} (hs : MeasurableSet s) : + MeasurableSet[MeasurableSpace.comap T inferInstance] (T ⁻¹' s) := by + rw [MeasurableSpace.measurableSet_comap] + exact ⟨s, hs, rfl⟩ + have h1 := condDistrib_ae_eq_condExp hT hX (μ := μ) hs + simp_rw [Kernel.map_apply _ hX, Measure.map_apply hX hs] + have h2 := condExpKernel_ae_eq_condExp hT.comap_le (μ := μ) (hX hs) + filter_upwards [h1, h2] with a ha₁ ha₂ + rw [Measure.real] at ha₁ ha₂ + rw [← ENNReal.toReal_eq_toReal (by simp) (by simp), ha₁, ha₂] + +omit [Nonempty Ω'] in +theorem condIndepFun_comap_iff_map_prod_eq_prod_condDistrib_prod_condDistrib + {α : Type*} {mα : MeasurableSpace α} [StandardBorelSpace α] + {X : α → β} {T : α → γ} {Z : α → Ω'} {μ : Measure α} [IsFiniteMeasure μ] + [StandardBorelSpace β] [StandardBorelSpace γ] [Nonempty β] [Nonempty γ] + (hX : Measurable X) (hT : Measurable T) (hZ : Measurable Z) : + CondIndepFun _ hZ.comap_le X T μ + ↔ μ.map (fun ω ↦ (Z ω, X ω, T ω)) + = (Kernel.id ×ₖ (condDistrib X Z μ ×ₖ condDistrib T Z μ)) ∘ₘ μ.map Z := by + rw [condIndepFun_iff_map_prod_eq_prod_comp_trim hX hT] + simp_rw [Measure.ext_prod₃_iff] + have hZ_meas {s : Set Ω'} (hs : MeasurableSet s) : + MeasurableSet[MeasurableSpace.comap Z inferInstance] (Z ⁻¹' s) := by + rw [MeasurableSpace.measurableSet_comap] + exact ⟨s, hs, rfl⟩ + have h_left {s : Set Ω'} {t : Set β} {u : Set γ} (hs : MeasurableSet s) (ht : MeasurableSet t) + (hu : MeasurableSet u) : + (μ.map (fun ω ↦ (Z ω, X ω, T ω))) (s ×ˢ t ×ˢ u) + = (@Measure.map _ _ _ ((MeasurableSpace.comap Z inferInstance).prod inferInstance) + (fun ω ↦ (ω, X ω, T ω)) μ) ((Z ⁻¹' s) ×ˢ t ×ˢ u) := by + rw [Measure.map_apply (by fun_prop) (hs.prod (ht.prod hu)), + Measure.map_apply _ ((hZ_meas hs).prod (ht.prod hu))] + · simp [Set.mk_preimage_prod] + · refine Measurable.prodMk ?_ (by fun_prop) + exact Measurable.mono measurable_id le_rfl hZ.comap_le + have h_right {s : Set Ω'} {t : Set β} {u : Set γ} (hs : MeasurableSet s) (ht : MeasurableSet t) + (hu : MeasurableSet u) : + ((Kernel.id ×ₖ (condDistrib X Z μ ×ₖ condDistrib T Z μ)) ∘ₘ μ.map Z) (s ×ˢ t ×ˢ u) + = ((Kernel.id ×ₖ + ((condExpKernel μ (MeasurableSpace.comap Z inferInstance)).map X ×ₖ + (condExpKernel μ (MeasurableSpace.comap Z inferInstance)).map T)) ∘ₘ + μ.trim hZ.comap_le) ((Z ⁻¹' s) ×ˢ t ×ˢ u) := by + rw [Measure.bind_apply ((hZ_meas hs).prod (ht.prod hu)) (by fun_prop), + Measure.bind_apply (hs.prod (ht.prod hu)) (by fun_prop), lintegral_map ?_ (by fun_prop), + lintegral_trim] + rotate_left + · exact Kernel.measurable_coe _ ((hZ_meas hs).prod (ht.prod hu)) + · exact Kernel.measurable_coe _ (hs.prod (ht.prod hu)) + refine lintegral_congr_ae ?_ + filter_upwards [condDistrib_apply_ae_eq_condExpKernel_map hX hZ ht, + condDistrib_apply_ae_eq_condExpKernel_map hT hZ hu] with a haX haT + simp_rw [Kernel.prod_apply_prod] + simp only [Kernel.id_apply, Measure.dirac_apply] + rw [@Measure.dirac_apply' _ (MeasurableSpace.comap Z inferInstance) _ _ (hZ_meas hs)] + congr + refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ + · intro s t u hs ht hu + specialize h (s := Z ⁻¹' s) (hZ_meas hs) ht hu + convert h + · exact h_left hs ht hu + · exact h_right hs ht hu + · rintro _ t u ⟨s, hs, rfl⟩ ht hu + specialize h hs ht hu + convert h + · exact (h_left hs ht hu).symm + · exact (h_right hs ht hu).symm + +-- todo: should be an iff +lemma condDistrib_prod_of_condIndepFun [StandardBorelSpace α] [StandardBorelSpace β] [Nonempty β] + (hX : Measurable X) (hY : Measurable Y) (hZ : Measurable Z) + (h : CondIndepFun (MeasurableSpace.comap Z inferInstance) hZ.comap_le Y X μ) : + condDistrib Y (fun ω ↦ (X ω, Z ω)) μ + =ᵐ[μ.map (fun ω ↦ (X ω, Z ω))] Kernel.prodMkLeft _ (condDistrib Y Z μ) := by + symm + refine condDistrib_ae_eq_of_measure_eq_compProd₀ (μ := μ) (hX.prodMk hZ).aemeasurable + hY.aemeasurable _ ?_ + rw [condIndepFun_comap_iff_map_prod_eq_prod_condDistrib_prod_condDistrib hY hX hZ] at h + rw [Measure.compProd_eq_comp_prod] + calc μ.map (fun x ↦ ((X x, Z x), Y x)) + _ = ((condDistrib X Z μ ×ₖ Kernel.id) ×ₖ condDistrib Y Z μ) ∘ₘ μ.map Z := by + -- up to shuffling, this is the previous lemma + sorry + _ = (Kernel.id ×ₖ Kernel.prodMkLeft β (condDistrib Y Z μ)) ∘ₘ Kernel.swap _ _ + ∘ₘ (μ.map Z ⊗ₘ condDistrib X Z μ) := by + rw [Measure.compProd_eq_comp_prod, Measure.comp_assoc, Measure.comp_assoc] + congr + rw [Kernel.comp_assoc, Kernel.swap_prod] + ext ω : 1 + simp_rw [Kernel.prod_apply] + rw [Kernel.comp_apply, Kernel.prod_apply, Kernel.id_apply, ← Measure.compProd_eq_comp_prod] + ext s hs + rw [Measure.compProd_apply hs, Measure.prod_apply hs] + simp only [Kernel.prodMkLeft_apply] + rw [lintegral_prod, lintegral_prod] + · simp_rw [lintegral_dirac] + · refine Measurable.aemeasurable ?_ + have : Measurable fun a ↦ (Kernel.prodMkLeft _ (condDistrib Y Z μ) a) (Prod.mk a ⁻¹' s) := + Kernel.measurable_kernel_prodMk_left hs + exact this + · refine Measurable.aemeasurable ?_ + have : Measurable fun x ↦ (Kernel.const _ ((condDistrib Y Z μ) ω) x) (Prod.mk x ⁻¹' s) := + Kernel.measurable_kernel_prodMk_left hs + exact this + _ = (Kernel.id ×ₖ Kernel.prodMkLeft β (condDistrib Y Z μ)) ∘ₘ μ.map (fun a ↦ (X a, Z a)) := by + congr + rw [compProd_map_condDistrib hX.aemeasurable, Measure.swap_comp, + Measure.map_map (by fun_prop) (by fun_prop)] + rfl + +lemma condDistrib_fst_prod (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) + (ν : Measure γ) [IsProbabilityMeasure ν] : + condDistrib (fun ω ↦ Y ω.1) (fun ω ↦ X ω.1) (μ.prod ν) =ᵐ[μ.map X] condDistrib Y X μ := by + refine condDistrib_ae_eq_of_measure_eq_compProd₀ (μ := μ) hX hY _ ?_ + have hX_map : (μ.prod ν).map (fun ω ↦ X ω.1) = μ.map X := by + calc (μ.prod ν).map (fun ω ↦ X ω.1) + _ = ((μ.prod ν).map Prod.fst).map X := by + rw [AEMeasurable.map_map_of_aemeasurable ?_ (by fun_prop)] + · rfl + · rw [Measure.map_fst_prod] + exact hX.smul_measure _ + _ = μ.map X := by simp [Measure.map_fst_prod] + rw [← hX_map, compProd_map_condDistrib] + · calc μ.map (fun x ↦ (X x, Y x)) + _ = ((μ.prod ν).map Prod.fst).map (fun a ↦ (X a, Y a)) := by simp [Measure.map_fst_prod] + _ = (μ.prod ν).map (fun a ↦ (X a.1, Y a.1)) := by + rw [AEMeasurable.map_map_of_aemeasurable ?_ (by fun_prop)] + · rfl + · simp only [Measure.map_fst_prod, measure_univ, one_smul] + fun_prop + · fun_prop + +end CondDistrib + +section Cond + +lemma ae_cond_of_forall_mem {μ : Measure α} {s : Set α} + (hs : MeasurableSet s) {p : α → Prop} (h : ∀ x ∈ s, p x) : + ∀ᵐ x ∂μ[|s], p x := Measure.ae_smul_measure (ae_restrict_of_forall_mem hs h) _ + +lemma condDistrib_ae_eq_cond [Countable β] [MeasurableSingletonClass β] + [IsFiniteMeasure μ] + (hX : Measurable X) (hY : Measurable Y) : + condDistrib Y X μ =ᵐ[μ.map X] fun b ↦ (μ[|X ⁻¹' {b}]).map Y := by + rw [Filter.EventuallyEq, ae_iff_of_countable] + intro b hb + ext s hs + rw [condDistrib_apply_of_ne_zero hY, + Measure.map_apply hX (measurableSet_singleton _), Measure.map_apply hY hs, + Measure.map_apply (hX.prodMk hY) ((measurableSet_singleton _).prod hs), + cond_apply (hX (measurableSet_singleton _))] + · congr + · exact hb + +lemma cond_of_indepFun [IsZeroOrProbabilityMeasure μ] (h : IndepFun X T μ) + (hX : Measurable X) (hT : Measurable T) {s : Set β} (hs : MeasurableSet s) + (hμs : μ (X ⁻¹' s) ≠ 0) : + (μ[|X ⁻¹' s]).map T = μ.map T := by + ext t ht + rw [Measure.map_apply (by fun_prop) ht, Measure.map_apply (by fun_prop) ht, cond_apply (hX hs), + IndepSet.measure_inter_eq_mul, ← mul_assoc, ENNReal.inv_mul_cancel, one_mul] + · exact hμs + · simp + · rw [indepFun_iff_indepSet_preimage hX hT] at h + exact h s t hs ht + +lemma cond_of_condIndepFun [StandardBorelSpace α] [StandardBorelSpace β] [Nonempty β] [Countable β] + [Countable Ω'] + [IsZeroOrProbabilityMeasure μ] + (hZ : Measurable Z) + (h : CondIndepFun (MeasurableSpace.comap Z inferInstance) hZ.comap_le Y X μ) + (hX : Measurable X) (hY : Measurable Y) {b : β} {ω : Ω'} + (hμ : μ (X ⁻¹' {b} ∩ Z ⁻¹' {ω}) ≠ 0) : + (μ[|X ⁻¹' {b} ∩ Z ⁻¹' {ω}]).map Y = (μ[|Z ⁻¹' {ω}]).map Y := by + have h := condDistrib_prod_of_condIndepFun hX hY hZ h + have h_left := condDistrib_ae_eq_cond (hX.prodMk hZ) hY (μ := μ) + have h_right := condDistrib_ae_eq_cond hZ hY (μ := μ) + rw [Filter.EventuallyEq, ae_iff_of_countable] at h h_left h_right + specialize h (b, ω) + specialize h_left (b, ω) + specialize h_right ω + rw [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] at h h_left h_right + rw [← Set.singleton_prod_singleton, Set.mk_preimage_prod] at h h_left + have hZ_ne : μ (Z ⁻¹' {ω}) ≠ 0 := fun h ↦ hμ (measure_mono_null Set.inter_subset_right h) + rw [← h_right hZ_ne, ← h_left hμ, h hμ] + simp + +end Cond + +end ProbabilityTheory diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index ecd03cb6..6d5d363c 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -130,6 +130,27 @@ lemma arm_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount (arm · h) a (s rwa [← pullCount_eq_pullCount] exact h_ne +lemma arm_eq_of_stepsUntil_eq_coe {ω : ℕ → α × ℝ} (hm : m ≠ 0) + (h : stepsUntil (arm · ω) a m = n) : + arm n ω = a := by + have : n = (stepsUntil (fun x ↦ arm x ω) a m).toNat := by simp [h] + rw [this, arm_stepsUntil hm] + by_contra! h_contra + rw [← stepsUntil_eq_top_iff] at h_contra + simp [h_contra] at h + +lemma stepsUntil_eq_congr {k' : ℕ → α} (h : ∀ i ≤ n, k i = k' i) : + stepsUntil k a m = n ↔ stepsUntil k' a m = n := by + sorry + +lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount k a (s + 1) = m) : + pullCount k a (stepsUntil k a m + 1).toNat = m := by + sorry + +lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount k a (s + 1) = m) : + pullCount k a (stepsUntil k a m).toNat = m - 1 := by + sorry + /-- Reward obtained when pulling arm `a` for the `m`-th time. -/ noncomputable def rewardByCount (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : ℝ := @@ -144,6 +165,14 @@ lemma rewardByCount_eq_ite (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ unfold rewardByCount cases stepsUntil (arm · h) a m <;> simp +lemma rewardByCount_of_stepsUntil_eq_top {ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)} + (h : stepsUntil (arm · ω.1) a m = ⊤) : + rewardByCount a m ω.1 ω.2 = ω.2 m a := by simp [rewardByCount_eq_ite, h] + +lemma rewardByCount_of_stepsUntil_eq_coe {ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)} + (h : stepsUntil (arm · ω.1) a m = n) : + rewardByCount a m ω.1 ω.2 = reward n ω.1 := by simp [rewardByCount_eq_ite, h] + lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : rewardByCount (arm t h) (pullCount (arm · h) (arm t h) t + 1) h z = reward t h := by rw [rewardByCount, ← pullCount_eq_pullCount_add_one, stepsUntil_pullCount_eq] diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 8f948cfa..35db1592 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -4,6 +4,7 @@ Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne -/ import LeanBandits.Bandit +import LeanBandits.ForMathlib.CondDistrib import LeanBandits.Regret /-! # Laws of `stepsUntil` and `rewardByCount` @@ -12,58 +13,6 @@ import LeanBandits.Regret open MeasureTheory ProbabilityTheory Finset open scoped ENNReal NNReal -section Aux - -variable {α β γ Ω Ω' : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] - {mα : MeasurableSpace α} {μ : Measure α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} - [MeasurableSpace Ω'] [StandardBorelSpace Ω'] [Nonempty Ω'] - {X : α → β} {Y : α → Ω} {Z : α → Ω'} - -lemma MeasureTheory.Measure.comp_congr {κ η : Kernel α β} (h : ∀ᵐ a ∂μ, κ a = η a) : - κ ∘ₘ μ = η ∘ₘ μ := - Measure.bind_congr_right h - -lemma MeasureTheory.Measure.copy_comp_map (hX : AEMeasurable X μ) : - Kernel.copy β ∘ₘ (μ.map X) = μ.map (fun a ↦ (X a, X a)) := by - rw [Kernel.copy, deterministic_comp_eq_map, AEMeasurable.map_map_of_aemeasurable (by fun_prop) hX] - congr - -lemma MeasureTheory.Measure.compProd_deterministic [SFinite μ] (hX : Measurable X) : - μ ⊗ₘ (Kernel.deterministic X hX) = μ.map (fun a ↦ (a, X a)) := by - rw [Measure.compProd_eq_comp_prod, Kernel.id, Kernel.deterministic_prod_deterministic, - Measure.deterministic_comp_eq_map] - rfl - -lemma ProbabilityTheory.condDistrib_comp_map [IsFiniteMeasure μ] - (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) : - condDistrib Y X μ ∘ₘ (μ.map X) = μ.map Y := by - rw [← Measure.snd_compProd, compProd_map_condDistrib hY, Measure.snd_map_prodMk₀ hX] - -lemma ProbabilityTheory.condDistrib_comp [IsFiniteMeasure μ] - (hX : AEMeasurable X μ) {f : β → Ω} (hf : Measurable f) : - condDistrib (f ∘ X) X μ =ᵐ[μ.map X] Kernel.deterministic f hf := by - rw [← Kernel.compProd_eq_iff, compProd_map_condDistrib (by fun_prop), - Measure.compProd_deterministic, AEMeasurable.map_map_of_aemeasurable (by fun_prop) hX] - congr - -lemma ProbabilityTheory.condDistrib_const [IsFiniteMeasure μ] - (hX : AEMeasurable X μ) (c : Ω) : - condDistrib (fun _ ↦ c) X μ =ᵐ[μ.map X] Kernel.deterministic (fun _ ↦ c) (by fun_prop) := by - have : (fun _ : α ↦ c) = (fun _ : β ↦ c) ∘ X := rfl - conv_lhs => rw [this] - filter_upwards [condDistrib_comp hX (by fun_prop : Measurable (fun _ ↦ c))] with b hb - rw [hb] - -@[fun_prop] -lemma Measurable.coe_nat_enat {f : α → ℕ} (hf : Measurable f) : - Measurable (fun a ↦ (f a : ℕ∞)) := Measurable.comp (by fun_prop) hf - -@[fun_prop] -lemma Measurable.toNat {f : α → ℕ∞} (hf : Measurable f) : Measurable (fun a ↦ (f a).toNat) := - Measurable.comp (by fun_prop) hf - -end Aux - namespace Bandits variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [MeasurableSingletonClass α] @@ -126,16 +75,199 @@ lemma measurable_rewardByCount (a : α) (m : ℕ) : (measurable_stepsUntil' a m).toNat.prodMk (by fun_prop) exact Measurable.comp (by fun_prop) this -lemma condDistrib_rewardByCount_stepsUntil [StandardBorelSpace α] [Nonempty α] - {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (m : ℕ) (hm : m ≠ 0) : +variable {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] + +omit [DecidableEq α] [MeasurableSingletonClass α] in +lemma hasLaw_Z (a : α) (m : ℕ) : + HasLaw (fun ω ↦ ω.2 m a) (ν a) (Bandit.measure alg ν) where + map_eq := by + calc ((Bandit.trajMeasure alg ν).prod (Bandit.streamMeasure ν)).map (fun ω ↦ ω.2 m a) + _ = (((Bandit.trajMeasure alg ν).prod (Bandit.streamMeasure ν)).map (fun ω ↦ ω.2)).map + (fun ω ↦ ω m a) := by + rw [Measure.map_map (by fun_prop) (by fun_prop)] + rfl + _ = (Bandit.streamMeasure ν).map (fun ω ↦ ω m a) := by simp [Measure.map_snd_prod] + _ = ((Measure.infinitePi fun _ ↦ Measure.infinitePi ν).map (fun ω ↦ ω m)).map + (fun ω ↦ ω a) := by + rw [Bandit.streamMeasure, Measure.map_map (by fun_prop) (by fun_prop)] + rfl + _ = ν a := by simp_rw [(measurePreserving_eval_infinitePi _ _).map_eq] + +/-- Law of `Y` conditioned on the event `s`.-/ +notation "𝓛[" Y " | " s "; " μ "]" => Measure.map Y (μ[|s]) +/-- Law of `Y` conditioned on the event that `X` is in `s`. -/ +notation "𝓛[" Y " | " X " in " s "; " μ "]" => Measure.map Y (μ[|X ⁻¹' s]) +/-- Law of `Y` conditioned on the event that `X` equals `x`. -/ +notation "𝓛[" Y " | " X " ← " x "; " μ "]" => Measure.map Y (μ[|X ⁻¹' {x}]) +/-- Law of `Y` conditioned on `X`. -/ +notation "𝓛[" Y " | " X "; " μ "]" => condDistrib Y X μ + +omit [DecidableEq α] [MeasurableSingletonClass α] in +lemma condDistrib_reward' (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 ν + have h_ra' : 𝓛[reward n | arm n; Bandit.trajMeasure alg ν] + =ᵐ[(Bandit.trajMeasure alg ν).map (arm n)] ν := condDistrib_reward alg ν n + have h_law : μ.map (fun ω ↦ arm n ω.1) = (Bandit.trajMeasure alg ν).map (arm n) := by + calc μ.map (fun ω ↦ arm n ω.1) + _ = (μ.map (fun ω ↦ ω.1)).map (fun ω ↦ arm n ω) := by + rw [Measure.map_map (by fun_prop) (by fun_prop)] + rfl + _ = _ := by unfold μ Bandit.measure; simp [Measure.map_fst_prod] + rw [h_law] + have h_prod : 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1; μ] + =ᵐ[(Bandit.trajMeasure alg ν).map (arm n)] 𝓛[reward n | arm n; Bandit.trajMeasure alg ν] := + condDistrib_fst_prod (by fun_prop) (by fun_prop) _ + filter_upwards [h_ra', h_prod] with ω h_eq h_prod + rw [h_prod, h_eq] + +omit [DecidableEq α] in +lemma reward_cond_arm [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 ν + have h_ra : 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1; μ] =ᵐ[μ.map (fun ω ↦ arm n ω.1)] ν := + condDistrib_reward' n + have h_eq := condDistrib_ae_eq_cond (μ := μ) + (X := fun ω ↦ arm n ω.1) (Y := fun ω ↦ reward n ω.1) (by fun_prop) (by fun_prop) + rw [Filter.EventuallyEq, ae_iff_of_countable] at h_ra h_eq + specialize h_ra a hμa + specialize h_eq a hμa + rw [h_ra] at h_eq + exact h_eq.symm + +lemma condIndepFun_reward_stepsUntil_arm [StandardBorelSpace α] [Countable α] [Nonempty α] + (a : α) (m n : ℕ) (hm : m ≠ 0) : + CondIndepFun (mα.comap (fun ω ↦ arm n ω.1)) ((measurable_arm n).comp measurable_fst).comap_le + (fun ω ↦ reward n ω.1) ({ω | stepsUntil (arm · ω.1) a m = ↑n}.indicator (fun _ ↦ 1)) + (Bandit.measure alg ν) := by + -- first restrict to the `trajMeasure` side + suffices h_indep : + CondIndepFun (mα.comap (arm n)) (measurable_arm n).comap_le + (reward n) ({ω | stepsUntil (arm · ω) a m = ↑n}.indicator (fun _ ↦ 1)) + (Bandit.trajMeasure alg ν) by + sorry + -- Now prove the independence : the indicator of `stepsUntil ... = n` is a function of + -- `hist (n-1)` and `arm n`. + -- It thus suffices to prove the independence of `reward n` and `hist (n-1)` conditionally + -- on `arm n`. + have hn : n ≠ 0 := by + sorry -- assume it? + have h_indep : CondIndepFun (mα.comap (arm n)) (measurable_arm n).comap_le (reward n) + (hist (n - 1)) (Bandit.trajMeasure alg ν) := by + convert condIndepFun_reward_hist_arm (alg := alg) (ν := ν) (n - 1) + <;> rw [Nat.sub_add_cancel (by grind)] + have h_indep' : CondIndepFun (mα.comap (arm n)) (measurable_arm n).comap_le (reward n) + (fun ω ↦ (hist (n - 1) ω, arm n ω)) (Bandit.trajMeasure alg ν) := + h_indep.prod_right (by fun_prop) (by fun_prop) (by fun_prop) + suffices ∃ φ : ((Iic (n - 1) → α × ℝ) × α) → ℕ, Measurable φ ∧ + ({ω : ℕ → α × ℝ | stepsUntil (arm · ω) a m = ↑n}.indicator (fun _ ↦ 1)) + = φ ∘ (fun ω : ℕ → α × ℝ ↦ (hist (n - 1) ω, arm n ω)) by + obtain ⟨φ, hφ_meas, h_eq⟩ := this + rw [h_eq] + exact h_indep'.comp measurable_id hφ_meas + -- it would follow from measurability wrt the sigma-algebra generated by + -- `hist (n-1)` and `arm n`, but we can also give an explicit function + let k : ((Iic (n - 1) → α × ℝ) × α) → (ℕ → α) := fun x i ↦ + if hi : i ∈ Iic (n - 1) then (x.1 ⟨i, hi⟩).1 else if i = n then x.2 else a -- a is arbitrary + let φ : ((Iic (n - 1) → α × ℝ) × α) → ℕ := fun x ↦ if stepsUntil (k x) a m = ↑n then 1 else 0 + classical + have hφ_meas : Measurable φ := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + refine (measurableSet_singleton _).preimage ?_ + refine (measurable_stepsUntil a m).comp ?_ + unfold k + rw [measurable_pi_iff] + intro i + split_ifs <;> fun_prop + refine ⟨φ, hφ_meas, ?_⟩ + ext ω + classical + simp only [Set.indicator_apply, Set.mem_setOf_eq, Function.comp_apply, φ] + congr 1 + rw [stepsUntil_eq_congr] + intro i hin + simp only [arm, mem_Iic, hist, dite_eq_ite, left_eq_ite_iff, not_le, k] + intro hni + have : i = n := by grind + simp [this] + +lemma reward_cond_stepsUntil [StandardBorelSpace α] [Countable α] [Nonempty α] (a : α) (m n : ℕ) + (hm : m ≠ 0) + (hμn : (Bandit.measure alg ν) ((fun ω ↦ stepsUntil (arm · ω.1) a m) ⁻¹' {↑n}) ≠ 0) : + 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ stepsUntil (fun x ↦ arm x ω.1) a m ← (n : ℕ∞); + Bandit.measure alg ν] = ν a := by + let μ := Bandit.measure alg ν + have hμna : + μ ((fun ω ↦ stepsUntil (arm · ω.1) a m) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}) ≠ 0 := by + suffices ((fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ + stepsUntil (arm · ω.1) a m) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}) + = (fun ω ↦ stepsUntil (arm · ω.1) a m) ⁻¹' {↑n} by simpa [this] using hμn + ext ω + simp only [Set.mem_inter_iff, Set.mem_preimage, Set.mem_singleton_iff, and_iff_left_iff_imp] + exact arm_eq_of_stepsUntil_eq_coe hm + have hμa : μ.map (fun ω ↦ arm n ω.1) {a} ≠ 0 := by + rw [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] + refine fun h_zero ↦ hμn (measure_mono_null (fun ω ↦ ?_) h_zero) + simp only [Set.mem_preimage, Set.mem_singleton_iff] + exact arm_eq_of_stepsUntil_eq_coe hm + calc 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ stepsUntil (arm · ω.1) a m ← (n : ℕ∞); μ] + _ = (μ[|(fun ω ↦ stepsUntil (fun x ↦ arm x ω.1) a m) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}]).map + (fun ω ↦ reward n ω.1) := by + congr with ω + simp only [Set.mem_preimage, Set.mem_singleton_iff, Set.mem_inter_iff, iff_self_and] + exact arm_eq_of_stepsUntil_eq_coe hm + _ = (μ[|{ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) | stepsUntil (arm · ω.1) a m = ↑n}.indicator 1 ⁻¹' {1} + ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}]).map (fun ω ↦ reward n ω.1) := by + congr 3 with ω + simp [Set.indicator_apply] + _ = 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1 ← a; μ] := by + rw [cond_of_condIndepFun (by fun_prop)] + · exact condIndepFun_reward_stepsUntil_arm a m n hm + · refine measurable_one.indicator ?_ + exact measurableSet_eq_fun' (by fun_prop) (by fun_prop) + · fun_prop + · convert hμna + ext ω + simp [Set.indicator_apply] + _ = ν a := reward_cond_arm a n hμa + +lemma condDistrib_rewardByCount_stepsUntil [Countable α] [StandardBorelSpace α] [Nonempty α] + (a : α) (m : ℕ) (hm : m ≠ 0) : condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil (arm · ω.1) a m) (Bandit.measure alg ν) =ᵐ[(Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)] Kernel.const _ (ν a) := by - sorry + let μ := Bandit.measure alg ν + refine (condDistrib_ae_eq_cond (μ := μ) + (X := fun ω ↦ stepsUntil (arm · ω.1) a m) (by fun_prop) (by fun_prop)).trans ?_ + rw [Filter.EventuallyEq, ae_iff_of_countable] + intro n hn + simp only [Kernel.const_apply] + cases n with + | top => + rw [Measure.map_congr (g := fun ω ↦ ω.2 m a)] + swap + · refine ae_cond_of_forall_mem ((measurableSet_singleton _).preimage (by fun_prop)) ?_ + simp only [Set.mem_preimage, Set.mem_singleton_iff] + exact fun ω ↦ rewardByCount_of_stepsUntil_eq_top + rw [cond_of_indepFun _ (by fun_prop) (by fun_prop) (measurableSet_singleton _)] + · exact (hasLaw_Z a m).map_eq + · rwa [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] at hn + · exact indepFun_prod (X := fun ω : ℕ → α × ℝ ↦ stepsUntil (arm · ω) a m) + (Y := fun ω : ℕ → α → ℝ ↦ ω m a) (by fun_prop) (by fun_prop) + | coe n => + rw [Measure.map_congr (g := fun ω ↦ reward n ω.1)] + swap + · refine ae_cond_of_forall_mem ((measurableSet_singleton _).preimage (by fun_prop)) ?_ + simp only [Set.mem_preimage, Set.mem_singleton_iff] + exact fun ω ↦ rewardByCount_of_stepsUntil_eq_coe + refine reward_cond_stepsUntil a m n hm ?_ + rwa [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] at hn /-- The reward received at the `m`-th pull of arm `a` has law `ν a`. -/ -lemma hasLaw_rewardByCount [StandardBorelSpace α] [Nonempty α] - {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (m : ℕ) (hm : m ≠ 0): +lemma hasLaw_rewardByCount [Countable α] [StandardBorelSpace α] [Nonempty α] + (a : α) (m : ℕ) (hm : m ≠ 0) : HasLaw (fun ω ↦ rewardByCount a m ω.1 ω.2) (ν a) (Bandit.measure alg ν) where map_eq := by have h_condDistrib : @@ -157,8 +289,7 @@ lemma hasLaw_rewardByCount [StandardBorelSpace α] [Nonempty α] isProbabilityMeasure_map (by fun_prop) simp -lemma identDistrib_rewardByCount [StandardBorelSpace α] [Nonempty α] - {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (n m : ℕ) +lemma identDistrib_rewardByCount [Countable α] [StandardBorelSpace α] [Nonempty α] (a : α) (n m : ℕ) (hn : n ≠ 0) (hm : m ≠ 0) : IdentDistrib (fun ω ↦ rewardByCount a n ω.1 ω.2) (fun ω ↦ rewardByCount a m ω.1 ω.2) (Bandit.measure alg ν) (Bandit.measure alg ν) where diff --git a/blueprint/lean_decls b/blueprint/lean_decls index fc9eaeeb..bfc67f6c 100644 --- a/blueprint/lean_decls +++ b/blueprint/lean_decls @@ -3,17 +3,21 @@ Bandits.Bandit.measure Bandits.arm Bandits.reward Bandits.hist +Bandits.pullCount +Bandits.filtration Bandits.condDistrib_reward +Bandits.hasLaw_arm_zero Bandits.condDistrib_arm Bandits.stepsUntil Bandits.rewardByCount +Bandits.hasLaw_rewardByCount +Bandits.iIndepFun_rewardByCount Bandits.stepsUntil_pullCount_le Bandits.stepsUntil_pullCount_eq Bandits.rewardByCount_pullCount_add_one_eq_reward Bandits.sum_rewardByCount_eq_sum_reward Bandits.regret Bandits.gap -Bandits.pullCount Bandits.regret_eq_sum_pullCount_mul_gap ProbabilityTheory.HasSubgaussianMGF ProbabilityTheory.HasSubgaussianMGF.add_of_indepFun diff --git a/blueprint/src/chapters/bandit.tex b/blueprint/src/chapters/bandit.tex index ff2faeb7..41dab130 100644 --- a/blueprint/src/chapters/bandit.tex +++ b/blueprint/src/chapters/bandit.tex @@ -56,6 +56,14 @@ \section{Algorithm, bandit and probability space} \end{definition} +\begin{definition}[Pull counts]\label{def:pullCount} + \uses{def:armAndReward} + \leanok + \lean{Bandits.pullCount} +For an arm $a \in \mathcal{A}$ and a time $t \in \mathbb{N}$, we denote by $N_{t,a}$ the number of times that arm $a$ has been pulled before time $t$, that is $N_{t,a} = \sum_{s=0}^{t-1} \mathbb{I}\{A_s = a\}$. +\end{definition} + + \begin{remark}[Building vs analyzing algorithms] When we describe an algorithm, we give the data of the policies $\pi_t$, which are functions of the partial history up to time $t$, in $(\mathcal{A} \times \mathcal{R})^{t+1}$. That means that any tool used to define a policy must be a function defined on $(\mathcal{A} \times \mathcal{R})^{t+1}$. @@ -68,6 +76,35 @@ \section{Algorithm, bandit and probability space} \end{remark} +\begin{definition}[Filtration]\label{def:banditFiltration} + \uses{def:armAndReward} + \leanok + \lean{Bandits.filtration} +The filtration $\mathcal{F}_B$ on $\Omega_B$ generated by the history of pulls and rewards is the increasing family of $\sigma$-algebras +$\mathcal{F}_{B,t} = \sigma(H_0, \ldots, H_t)$, for $t \in \mathbb{N}$. +\end{definition} + + +\begin{lemma}\label{lem:adapted_hist} + \uses{def:banditFiltration,def:armAndReward} +Seen as processes defined on $\Omega_B$, the processes $(H_t)_{t \in \mathbb{N}}$, $(A_t)_{t \in \mathbb{N}}$, and $(X_t)_{t \in \mathbb{N}}$ are adapted to the filtration $\mathcal{F}_B$. +\end{lemma} + +\begin{proof} + +\end{proof} + + +\begin{lemma}\label{lem:predictable_pullCount} + \uses{def:banditFiltration,def:pullCount} +Let $a \in \mathcal{A}$. Seen as a process defined on $\Omega_B$, $(N_{t,a})_{t \in \mathbb{N}}$ is predictable with respect to the filtration $\mathcal{F}_B$ (that is, $(N_{t,a})_{t \in \mathbb{N}}$ is adapted to $(\mathcal{F}_{B, t-1})_{t \in \mathbb{N}}$). +\end{lemma} + +\begin{proof} + +\end{proof} + + \begin{lemma}\label{lem:condDistrib_reward} \uses{def:Bandit.measure,def:armAndReward} \leanok @@ -82,6 +119,8 @@ \section{Algorithm, bandit and probability space} \begin{lemma}\label{lem:law_arm_zero} \uses{def:Bandit.measure,def:armAndReward} + \leanok + \lean{Bandits.hasLaw_arm_zero} The law of the arm $A_0$ in the bandit probability space $(\Omega, \mathbb{P})$ is $P_0$. \end{lemma} @@ -128,9 +167,54 @@ \section{Alternative model}\label{sec:alt_model} \end{definition} -\begin{lemma}\label{lem:iid_rewardByCount} +\begin{lemma}\label{lem:isStoppingTime_stepsUntil} + \uses{def:stepsUntil} +$T_{n,a}$ is a stopping time for the filtration generated by the history of pulls and rewards. +\end{lemma} + +\begin{proof} +It is the hitting time of a measurable set by the adapted process $(N_{n+1, a})_{n \in \mathbb{N}}$, hence a stopping time. +\end{proof} + + +\begin{lemma}\label{lem:hasLaw_rewardByCount} + \uses{def:rewardByCount} + \leanok + \lean{Bandits.hasLaw_rewardByCount} +For $n > 0$ and $a \in \mathcal{A}$, the law of $Y_{n,a}$ is $\nu(a)$. +\end{lemma} + +\begin{proof} + \uses{lem:condDistrib_reward} +It suffices to show that for all $t \in \mathbb{N} \cup \{\infty\}$, the law of $Y_{n,a}$ conditioned on $T_{n,a} = t$ is $\nu(a)$. +If $t = \infty$, then +\begin{align*} + \mathcal{L}(Y_{n,a} \mid T_{n,a} = t) + = \mathcal{L}(Z_{n,a} \mid T_{n,a} = t) + = \mathcal{L}(Z_{n,a}) + = \nu(a) +\end{align*} +If $t < \infty$, then +\begin{align*} + \mathcal{L}(Y_{n,a} \mid T_{n,a} = t) + &= \mathcal{L}(X_t \mid T_{n,a} = t) + \\ + &= \mathcal{L}(X_t \mid T_{n,a} = t, A_t = a) + \\ + &= \mathcal{L}(X_t \mid A_t = a) + \\ + &= \nu(a) + \: . +\end{align*} +TODO: explain that chain of equalities. There is independence involved. +\end{proof} + + +\begin{lemma}\label{lem:iIndepFun_rewardByCount} \uses{def:rewardByCount} -The rewards $(Y_{n,a})_{n \in \mathbb{N}}$ are independent and identically distributed random variables, with distribution $\nu(a)$. + \leanok + \lean{Bandits.iIndepFun_rewardByCount} +The rewards $(Y_{n,a})_{n \in \mathbb{N}}$ are independent. \end{lemma} \begin{proof} @@ -213,7 +297,7 @@ \section{Regret and other bandit quantities} \begin{definition}[Regret]\label{def:regret} - \uses{def:armMean} + \uses{def:armMean, def:armAndReward} \leanok \lean{Bandits.regret} The regret $R_T$ of a sequence of arms $A_0, \ldots, A_{T-1}$ after $T$ pulls is the difference between the cumulative reward of always playing the best arm and the cumulative reward of the sequence: @@ -231,14 +315,6 @@ \section{Regret and other bandit quantities} \end{definition} -\begin{definition}\label{def:pullCount} - \uses{def:bandit} - \leanok - \lean{Bandits.pullCount} -For an arm $a \in \mathcal{A}$ and a time $t \in \mathbb{N}$, we denote by $N_{t,a}$ the number of times that arm $a$ has been pulled before time $t$, that is $N_{t,a} = \sum_{s=0}^{t-1} \mathbb{I}\{A_s = a\}$. -\end{definition} - - \begin{lemma}\label{lem:regret_eq_sum_pullCount_mul_gap} \uses{def:regret,def:gap,def:pullCount} \leanok diff --git a/blueprint/src/chapters/etc.tex b/blueprint/src/chapters/etc.tex index 113bf6ed..e885c8cf 100644 --- a/blueprint/src/chapters/etc.tex +++ b/blueprint/src/chapters/etc.tex @@ -40,7 +40,7 @@ \section{Explore-Then-Commit} \end{lemma} \begin{proof} - \uses{lem:iid_rewardByCount, lem:independent_rewardByCount, lem:sum_rewardByCount, thm:hoeffding} + \uses{lem:hasLaw_rewardByCount, lem:iIndepFun_rewardByCount, lem:independent_rewardByCount, lem:sum_rewardByCount, thm:hoeffding} \begin{align*} \mathbb{P}(\hat{A}_m^* = a) &\le \mathbb{P}(\hat{\mu}_a \ge \hat{\mu}_{a^*})