From 602103febf2b96a6898502a7fea56fd31e6e13f7 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 10 Sep 2025 14:10:17 +0200 Subject: [PATCH] more indep lemmas --- LeanBandits/Bandit.lean | 5 +- LeanBandits/ForMathlib/CondDistrib.lean | 218 ++++++++++++++++++------ LeanBandits/Regret.lean | 37 ++++ LeanBandits/RewardByCountMeasure.lean | 17 +- 4 files changed, 222 insertions(+), 55 deletions(-) diff --git a/LeanBandits/Bandit.lean b/LeanBandits/Bandit.lean index 7f981c3f..f5e6c131 100644 --- a/LeanBandits/Bandit.lean +++ b/LeanBandits/Bandit.lean @@ -4,6 +4,7 @@ Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne, Paulo Rauber -/ import Mathlib +import LeanBandits.ForMathlib.CondDistrib /-! # Bandit @@ -159,10 +160,12 @@ lemma hasLaw_arm_zero [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace 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] +lemma condIndepFun_reward_hist_arm [StandardBorelSpace α] [Nonempty α] + [StandardBorelSpace R] [Nonempty 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 + rw [condIndepFun_iff_condDistrib_prod_ae_eq_prodMkLeft (by fun_prop) (by fun_prop) (by fun_prop)] sorry end MeasureSpace diff --git a/LeanBandits/ForMathlib/CondDistrib.lean b/LeanBandits/ForMathlib/CondDistrib.lean index f1a8b6db..00ddff42 100644 --- a/LeanBandits/ForMathlib/CondDistrib.lean +++ b/LeanBandits/ForMathlib/CondDistrib.lean @@ -13,8 +13,10 @@ 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 γ} +variable {α β γ δ Ω Ω' : Type*} + {m mα : MeasurableSpace α} {μ : Measure α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} + {mδ : MeasurableSpace δ} + [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] [MeasurableSpace Ω'] [StandardBorelSpace Ω'] [Nonempty Ω'] {X : α → β} {Y : α → Ω} {Z : α → Ω'} {T : α → γ} @@ -71,10 +73,96 @@ lemma Kernel.prod_apply_prod {κ : Kernel α β} {η : Kernel α γ} (κ ×ₖ η) 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 : α → δ} +-- fix the lemma in mathlib to allow different types for the functions +theorem CondIndepFun.symm' + [StandardBorelSpace α] {hm : m ≤ mα} [IsFiniteMeasure μ] {f : α → β} {g : α → γ} + (hfg : CondIndepFun m hm f g μ) : + CondIndepFun m hm g f μ := + Kernel.IndepFun.symm hfg + +lemma Kernel.indepFun_const_left {κ : Kernel α β} [IsZeroOrMarkovKernel κ] (c : δ) (X : β → γ) : + IndepFun (fun _ ↦ c) X κ μ := by + rw [IndepFun, MeasurableSpace.comap_const] + exact indep_bot_left _ + +lemma Kernel.indepFun_const_right {κ : Kernel α β} [IsZeroOrMarkovKernel κ] (X : β → γ) (c : δ) : + IndepFun X (fun _ ↦ c) κ μ := + (Kernel.indepFun_const_left c X).symm + +lemma condIndepFun_const_left [StandardBorelSpace α] {hm : m ≤ mα} [IsFiniteMeasure μ] + (c : γ) (X : α → β) : + CondIndepFun m hm (fun _ ↦ c) X μ := + Kernel.indepFun_const_left c X + +lemma condIndepFun_const_right [StandardBorelSpace α] {hm : m ≤ mα} + [IsFiniteMeasure μ] (X : α → β) (c : γ) : + CondIndepFun m hm X (fun _ ↦ c) μ := + Kernel.indepFun_const_right X c + +lemma condIndepFun_of_measurable_left [StandardBorelSpace α] {hm : m ≤ mα} [IsFiniteMeasure μ] + {X : α → β} {Y : α → γ} (hX : Measurable[m] X) (hY : Measurable Y) : + CondIndepFun m hm X Y μ := by + rw [condIndepFun_iff _ hm _ _ (hX.mono hm le_rfl) hY] + rintro _ _ ⟨s, hs, rfl⟩ ⟨t, ht, rfl⟩ + have h_ind_eq ω : (X ⁻¹' s ∩ Y ⁻¹' t).indicator (fun _ ↦ (1 : ℝ)) ω + = (X ⁻¹' s).indicator (fun _ ↦ (1 : ℝ)) ω * (Y ⁻¹' t).indicator (fun _ ↦ (1 : ℝ)) ω := by + simp only [Set.indicator, Set.mem_inter_iff, Set.mem_preimage, mul_ite, mul_one, mul_zero] + split_ifs with h1 h2 h3 h4 h5 + · rfl + · exfalso + rw [Set.mem_inter_iff] at h1 + refine h3 h1.1 + · exfalso + rw [Set.mem_inter_iff] at h1 + exact h2 h1.2 + · exfalso + rw [Set.mem_inter_iff] at h1 + exact h1 ⟨h5, h4⟩ + · rfl + · rfl + calc μ[(X ⁻¹' s ∩ Y ⁻¹' t).indicator fun ω ↦ (1 : ℝ)|m] + _ = μ[fun ω ↦ (X ⁻¹' s).indicator (fun _ ↦ 1) ω * (Y ⁻¹' t).indicator (fun _ ↦ 1) ω|m] := by + simp_rw [← h_ind_eq] + _ =ᵐ[μ] fun ω ↦ (X ⁻¹' s).indicator (fun _ ↦ 1) ω * μ[(Y ⁻¹' t).indicator (fun _ ↦ 1)|m] ω := by + refine condExp_mul_of_stronglyMeasurable_left ?_ ?_ ?_ + · exact (Measurable.indicator (by fun_prop) (hX hs)).stronglyMeasurable + · have : ((X ⁻¹' s).indicator fun x ↦ (1 : ℝ)) * (Y ⁻¹' t).indicator (fun x ↦ 1) + = (X ⁻¹' s ∩ Y ⁻¹' t).indicator (fun _ ↦ (1 : ℝ)) := by ext; simp [h_ind_eq] + rw [this] + rw [integrable_indicator_iff] + · exact (integrable_const (1 : ℝ)).integrableOn + · exact (hm _ (hX hs)).inter (hY ht) + · rw [integrable_indicator_iff] + · exact (integrable_const (1 : ℝ)).integrableOn + · exact hY ht + _ =ᵐ[μ] μ[(X ⁻¹' s).indicator fun ω ↦ 1|m] * μ[(Y ⁻¹' t).indicator fun ω ↦ 1|m] := by + nth_rw 2 [condExp_of_stronglyMeasurable hm] + · rfl + · exact (Measurable.indicator (by fun_prop) (hX hs)).stronglyMeasurable + · rw [integrable_indicator_iff] + · exact (integrable_const (1 : ℝ)).integrableOn + · exact hm _ (hX hs) + +lemma condIndepFun_of_measurable_right [StandardBorelSpace α] {hm : m ≤ mα} [IsFiniteMeasure μ] + {X : α → β} {Y : α → γ} (hX : Measurable X) (hY : Measurable[m] Y) : + CondIndepFun m hm X Y μ := by + refine CondIndepFun.symm' ?_ + exact condIndepFun_of_measurable_left hY hX + +lemma condIndepFun_self_left [StandardBorelSpace α] [IsFiniteMeasure μ] + {X : α → β} {Z : α → δ} (hX : Measurable X) (hZ : Measurable Z) : + CondIndepFun (mδ.comap Z) hZ.comap_le Z X μ := by + refine condIndepFun_of_measurable_left ?_ hX + rw [measurable_iff_comap_le] + +lemma condIndepFun_self_right [StandardBorelSpace α] [IsFiniteMeasure μ] + {X : α → β} {Z : α → δ} (hX : Measurable X) (hZ : Measurable Z) : + CondIndepFun (mδ.comap Z) hZ.comap_le X Z μ := by + refine condIndepFun_of_measurable_right hX ?_ + rw [measurable_iff_comap_le] + +lemma CondIndepFun.prod_right [StandardBorelSpace α] [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 @@ -106,20 +194,26 @@ lemma condDistrib_congr_left {Y' : α → Ω} (hY : Y =ᵐ[μ] Y') : 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 + condDistrib Y X μ =ᵐ[μ.map X] κ := by suffices ∀ᵐ x ∂μ.map (hX.mk X), κ x = condDistrib (hY.mk Y) (hX.mk X) μ x by + symm 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] + rw [condDistrib_congr hY.ae_eq_mk hX.ae_eq_mk, Filter.EventuallyEq] 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_ae_eq_iff_measure_eq_compProd₀ + (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) (κ : Kernel β Ω) [IsFiniteKernel κ] : + (condDistrib Y X μ =ᵐ[μ.map X] κ) ↔ μ.map (fun x => (X x, Y x)) = μ.map X ⊗ₘ κ := by + refine ⟨fun h ↦ ?_, condDistrib_ae_eq_of_measure_eq_compProd₀ hX hY κ⟩ + rw [Measure.compProd_congr h.symm, compProd_map_condDistrib hY] + 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 @@ -133,7 +227,6 @@ lemma condDistrib_const (hX : AEMeasurable X μ) (c : Ω) : 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 @@ -333,51 +426,74 @@ theorem condIndepFun_comap_iff_map_prod_eq_prod_condDistrib_prod_condDistrib · 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 +lemma condIndepFun_iff_condDistrib_prod_ae_eq_prodMkLeft + [StandardBorelSpace α] [StandardBorelSpace β] [Nonempty β] + (hX : Measurable X) (hY : Measurable Y) (hZ : Measurable Z) : + 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 + rw [condDistrib_ae_eq_iff_measure_eq_compProd₀ (μ := μ) (hX.prodMk hZ).aemeasurable + hY.aemeasurable, condIndepFun_comap_iff_map_prod_eq_prod_condDistrib_prod_condDistrib hY hX hZ, + Measure.compProd_eq_comp_prod] + let e : Ω' × Ω × β ≃ᵐ (β × Ω') × Ω := { + toFun := fun p ↦ ((p.2.2, p.1), p.2.1) + invFun := fun p ↦ (p.1.2, p.2, p.1.1) + left_inv p := by simp + right_inv p := by simp + measurable_toFun := by simp only [Equiv.coe_fn_mk]; fun_prop + measurable_invFun := by simp only [Equiv.coe_fn_symm_mk]; fun_prop } + have h_eq : ((condDistrib X Z μ ×ₖ Kernel.id) ×ₖ condDistrib Y Z μ) ∘ₘ μ.map Z + = (Kernel.id ×ₖ Kernel.prodMkLeft β (condDistrib Y Z μ)) ∘ₘ μ.map (fun a ↦ (X a, Z a)) := by + calc ((condDistrib X Z μ ×ₖ Kernel.id) ×ₖ condDistrib Y Z μ) ∘ₘ μ.map Z + _ = (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 + rw [← h_eq] + have h1 : μ.map (fun x ↦ ((X x, Z x), Y x)) = (μ.map (fun a ↦ (Z a , Y a, X a))).map e := by + rw [Measure.map_map (by fun_prop) (by fun_prop)] congr - rw [compProd_map_condDistrib hX.aemeasurable, Measure.swap_comp, - Measure.map_map (by fun_prop) (by fun_prop)] - rfl + have h1_symm : μ.map (fun a ↦ (Z a , Y a, X a)) + = (μ.map (fun x ↦ ((X x, Z x), Y x))).map e.symm := by + rw [h1, Measure.map_map (by fun_prop) (by fun_prop), MeasurableEquiv.symm_comp_self, + Measure.map_id] + have h2 : (condDistrib X Z μ ×ₖ Kernel.id ×ₖ condDistrib Y Z μ) ∘ₘ μ.map Z + = ((Kernel.id ×ₖ (condDistrib Y Z μ ×ₖ condDistrib X Z μ)) ∘ₘ μ.map Z).map e := by + rw [← Measure.deterministic_comp_eq_map e.measurable, Measure.comp_assoc] + sorry + have h2_symm : (Kernel.id ×ₖ (condDistrib Y Z μ ×ₖ condDistrib X Z μ)) ∘ₘ μ.map Z + = ((condDistrib X Z μ ×ₖ Kernel.id ×ₖ condDistrib Y Z μ) ∘ₘ μ.map Z).map e.symm := by + rw [h2, Measure.map_map (by fun_prop) (by fun_prop), MeasurableEquiv.symm_comp_self, + Measure.map_id] + rw [h1, h2] + exact ⟨fun h ↦ by rw [h], fun h ↦ by rw [h1_symm, h1, h2_symm, h2, h]⟩ 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 + symm 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) @@ -439,7 +555,7 @@ lemma cond_of_condIndepFun [StandardBorelSpace α] [StandardBorelSpace β] [None (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 := (condIndepFun_iff_condDistrib_prod_ae_eq_prodMkLeft hX hY hZ).mp 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 diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index 6d5d363c..c28be981 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -43,6 +43,13 @@ noncomputable def pullCount [DecidableEq α] (k : ℕ → α) (a : α) (t : ℕ) @[simp] lemma pullCount_zero (k : ℕ → α) (a : α) : pullCount k a 0 = 0 := by simp [pullCount] +lemma pullCount_one : pullCount k a 1 = if k 0 = a then 1 else 0 := by + simp only [pullCount, range_one] + split_ifs with h + · rw [card_eq_one] + refine ⟨0, by simp [h]⟩ + · simp [h] + open Classical in lemma monotone_pullCount (k : ℕ → α) (a : α) : Monotone (pullCount k a) := fun _ _ _ ↦ card_le_card (filter_subset_filter _ (by simpa)) @@ -110,6 +117,36 @@ lemma stepsUntil_pullCount_eq (k : ℕ → α) (t : ℕ) : simpa [stepsUntil, pullCount_eq_pullCount_add_one] exact fun t' h ↦ Nat.le_of_lt_succ ((monotone_pullCount k (k t)).reflect_lt (h ▸ lt_add_one _)) +/-- If we pull arm `a` at time 0, the first time at which it is pulled once is 0. -/ +lemma stepsUntil_one_of_eq (hka : k 0 = a) : stepsUntil k a 1 = 0 := by + classical + have h_pull : pullCount k a 1 = 1 := by simp [pullCount_one, hka] + have h_le := stepsUntil_pullCount_le k a 0 + simpa [h_pull] using h_le + +lemma stepsUntil_eq_zero_iff : + stepsUntil k a m = 0 ↔ (m = 0 ∧ k 0 ≠ a) ∨ (m = 1 ∧ k 0 = a) := by + classical + refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ + · have h_exists : ∃ s, pullCount k a (s + 1) = m := by + by_contra! h_contra + rw [← stepsUntil_eq_top_iff] at h_contra + simp [h_contra] at h + simp only [stepsUntil_eq_dite, h_exists, ↓reduceDIte, Nat.cast_eq_zero, Nat.find_eq_zero, + zero_add] at h + rw [pullCount_one] at h + by_cases hka : k 0 = a + · simp only [hka, ↓reduceIte] at h + simp [h.symm, hka] + · simp only [hka, ↓reduceIte] at h + simp [h.symm, hka] + · cases h with + | inl h => + rw [h.1, stepsUntil_zero_of_ne h.2] + | inr h => + rw [h.1] + exact stepsUntil_one_of_eq h.2 + lemma arm_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount (arm · h) a (s + 1) = m) : arm (stepsUntil (arm · h) a m).toNat h = a := by classical diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 35db1592..99224810 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -4,7 +4,6 @@ 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` @@ -152,8 +151,20 @@ lemma condIndepFun_reward_stepsUntil_arm [StandardBorelSpace α] [Countable α] -- `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? + by_cases hn : n = 0 + · simp only [hn, CharP.cast_eq_zero] + simp only [stepsUntil_eq_zero_iff, hm, ne_eq, false_and, false_or] + by_cases hm1 : m = 1 + · simp only [hm1, true_and] + have h_indep := condIndepFun_self_right (X := reward 0) (Z := arm 0) + (mβ := inferInstance) (mδ := inferInstance) (μ := Bandit.trajMeasure alg ν) + (by fun_prop) (by fun_prop) + have : {ω : ℕ → α × ℝ | arm 0 ω = a}.indicator (fun x ↦ 1) + = {b | b = a}.indicator (fun _ ↦ 1) ∘ arm 0 := by ext; simp [Set.indicator] + rw [this] + exact h_indep.comp measurable_id (by fun_prop) + · simp only [hm1, false_and, Set.setOf_false, Set.indicator_empty] + exact condIndepFun_const_right (reward 0) 0 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)