From dc318ace28975c5b8a49dc01a91a1c0273d4569a Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 2 Jan 2026 22:57:52 +0100 Subject: [PATCH 01/30] partial independence proof --- LeanBandits.lean | 1 + LeanBandits/Bandit/Bandit.lean | 18 ++ LeanBandits/ForMathlib/CondDistrib.lean | 15 -- LeanBandits/ForMathlib/HasCondDistrib.lean | 26 +++ LeanBandits/RewardByCountMeasure.lean | 187 +++++++++++++++++- .../SequentialLearning/FiniteActions.lean | 14 ++ blueprint/lean_decls | 6 +- 7 files changed, 241 insertions(+), 26 deletions(-) create mode 100644 LeanBandits/ForMathlib/HasCondDistrib.lean diff --git a/LeanBandits.lean b/LeanBandits.lean index aaea891f..10b0f463 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -3,6 +3,7 @@ import LeanBandits.Bandit.Regret import LeanBandits.BanditAlgorithms.ETC import LeanBandits.BanditAlgorithms.UCB import LeanBandits.ForMathlib.CondDistrib +import LeanBandits.ForMathlib.HasCondDistrib import LeanBandits.ForMathlib.IndepFun import LeanBandits.ForMathlib.IndepInfinitePi import LeanBandits.ForMathlib.KernelSub diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 5e047baa..d379032d 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -239,6 +239,24 @@ lemma condIndepFun_reward_hist_arm [StandardBorelSpace α] [Nonempty α] (measurable_arm _).comap_le (reward (n + 1)) (hist n) (Bandit.trajMeasure alg ν) := Learning.condIndepFun_reward_hist_action n +lemma condIndepFun_reward_hist_arm_arm [StandardBorelSpace α] [Countable α] [Nonempty α] + [StandardBorelSpace R] [Nonempty R] + {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) : + reward (n + 1) ⟂ᵢ[arm (n + 1), measurable_arm (n + 1); Bandit.trajMeasure alg ν] + (fun ω ↦ (hist n ω, arm (n + 1) ω)) := by + have h_indep : reward (n + 1) ⟂ᵢ[arm (n + 1), measurable_arm (n + 1); Bandit.trajMeasure alg ν] + hist n := by + convert condIndepFun_reward_hist_arm (alg := alg) (ν := ν) n + exact h_indep.prod_right (by fun_prop) (by fun_prop) (by fun_prop) + +lemma condIndepFun_reward_hist_arm_arm' [StandardBorelSpace α] [Countable α] [Nonempty α] + [StandardBorelSpace R] [Nonempty R] + {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) (hn : n ≠ 0) : + reward n ⟂ᵢ[arm n, measurable_arm n; Bandit.trajMeasure alg ν] + (fun ω ↦ (hist (n - 1) ω, arm n ω)) := by + have := condIndepFun_reward_hist_arm_arm (alg := alg) (ν := ν) (n - 1) + grind + end Laws section DetAlgorithm diff --git a/LeanBandits/ForMathlib/CondDistrib.lean b/LeanBandits/ForMathlib/CondDistrib.lean index 9aa67375..101b4c93 100644 --- a/LeanBandits/ForMathlib/CondDistrib.lean +++ b/LeanBandits/ForMathlib/CondDistrib.lean @@ -33,14 +33,6 @@ lemma Measurable.toNat {f : α → ℕ∞} (hf : Measurable f) : Measurable (fun namespace MeasureTheory.Measure -lemma trim_eq_map {hm : m ≤ mα} : μ.trim hm = @Measure.map _ _ mα m id μ := by - refine @Measure.ext _ m _ _ fun s hs ↦ ?_ - rw [trim_measurableSet_eq _ hs, Measure.map_apply _ hs] - · simp - · intro t ht - simp only [Set.preimage_id_eq, id_eq] - exact hm _ ht - 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] @@ -52,13 +44,6 @@ namespace ProbabilityTheory section IndepFun --- 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.of_prod_right {ε Ω : Type*} {mΩ : MeasurableSpace Ω} {mε : MeasurableSpace ε} {μ : Measure Ω} {κ : Kernel Ω α} {X : α → β} {Y : α → γ} {T : α → ε} (h : IndepFun X (fun ω ↦ (Y ω, T ω)) κ μ) : diff --git a/LeanBandits/ForMathlib/HasCondDistrib.lean b/LeanBandits/ForMathlib/HasCondDistrib.lean new file mode 100644 index 00000000..e26dacfb --- /dev/null +++ b/LeanBandits/ForMathlib/HasCondDistrib.lean @@ -0,0 +1,26 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne, Paulo Rauber +-/ +import LeanBandits.ForMathlib.CondDistrib + +/-! +# A predicate for having a specified conditional distribution +-/ + +open MeasureTheory + +namespace ProbabilityTheory + +variable {α β Ω : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mΩ : MeasurableSpace Ω} [StandardBorelSpace Ω] [Nonempty Ω] + {μ : Measure α} {X : α → β} {Y : α → Ω} {κ : Kernel β Ω} + +structure HasCondDistrib (Y : α → Ω) (X : α → β) (κ : Kernel β Ω) + (μ : Measure α) [IsFiniteMeasure μ] : Prop where + aemeasurable_fst : AEMeasurable Y μ + aemeasurable_snd : AEMeasurable X μ + condDistrib_eq : condDistrib Y X μ =ᵐ[μ.map X] κ + +end ProbabilityTheory diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index baac2523..250d2ba9 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -142,16 +142,13 @@ lemma condIndepFun_reward_stepsUntil_arm' [StandardBorelSpace α] [Countable α] 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 : reward n ⟂ᵢ[arm n, measurable_arm n; 𝔓t] hist (n - 1) := by - convert condIndepFun_reward_hist_arm (alg := alg) (ν := ν) (n - 1) - <;> rw [Nat.sub_add_cancel (by grind)] - have h_indep' : reward n ⟂ᵢ[arm n, measurable_arm n; 𝔓t] fun ω ↦ (hist (n - 1) ω, arm n ω) := - h_indep.prod_right (by fun_prop) (by fun_prop) (by fun_prop) + have h_indep : reward n ⟂ᵢ[arm n, measurable_arm n; 𝔓t] fun ω ↦ (hist (n - 1) ω, arm n ω) := + condIndepFun_reward_hist_arm_arm' (alg := alg) (ν := ν) n (by grind) obtain ⟨φ, hφ_meas, h_eq⟩ : ∃ φ : ((Iic (n - 1) → α × ℝ) × α) → ℕ, Measurable φ ∧ {ω | stepsUntil a m ω = ↑n}.indicator (fun _ ↦ 1) = φ ∘ (fun ω ↦ (hist (n - 1) ω, arm n ω)) := (measurable_comap_indicator_stepsUntil_eq a m n).exists_eq_measurable_comp rw [h_eq] - exact h_indep'.comp measurable_id hφ_meas + exact h_indep.comp measurable_id hφ_meas lemma condIndepFun_reward_stepsUntil_arm [StandardBorelSpace α] [Countable α] [Nonempty α] (a : α) (m n : ℕ) (hm : m ≠ 0) : @@ -281,10 +278,184 @@ lemma iIndepFun_rewardByCount' (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [Is rw [iIndepFun_nat_iff_forall_indepFun (by fun_prop)] exact indepFun_rewardByCount_Iic alg ν a -lemma iIndepFun_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] : - iIndepFun (fun (p : α × ℕ) ↦ rewardByCount p.1 p.2) 𝔓 := by +def E' (I : Finset (α × ℕ)) (S : Finset I) : Set (ℕ → α × ℝ) := + {ω | (∀ i ∈ S, stepsUntil i.1.1 (i.1.2 + 1) ω < ⊤) ∧ + (∀ j ∉ S, stepsUntil j.1.1 (j.1.2 + 1) ω = ⊤)} + +lemma measurableSet_E' [Countable α] [Nonempty α] (I : Finset (α × ℕ)) (S : Finset I) : + MeasurableSet (E' I S) := by + have h_eq : E' I S + = (⋂ i ∈ S, {ω | stepsUntil i.1.1 (i.1.2 + 1) ω ≠ ⊤}) ∩ + (⋂ j ∉ S, {ω | stepsUntil j.1.1 (j.1.2 + 1) ω = ⊤}) := by ext; simp [E', lt_top_iff_ne_top] + rw [h_eq] + refine MeasurableSet.inter ?_ ?_ + · refine MeasurableSet.iInter fun i ↦ MeasurableSet.iInter fun hi ↦ ?_ + exact (measurableSet_singleton _).compl.preimage (by fun_prop) + · refine MeasurableSet.iInter fun j ↦ MeasurableSet.iInter fun hj ↦ ?_ + exact (measurableSet_singleton _).preimage (by fun_prop) + +def E (I : Finset (α × ℕ)) (S : Finset I) : Set ((ℕ → α × ℝ) × (ℕ → α → ℝ)) := + {ω | (∀ i ∈ S, stepsUntil i.1.1 (i.1.2 + 1) ω.1 < ⊤) ∧ + (∀ j ∉ S, stepsUntil j.1.1 (j.1.2 + 1) ω.1 = ⊤)} + +lemma measurableSet_E [Countable α] [Nonempty α] (I : Finset (α × ℕ)) (S : Finset I) : + MeasurableSet (E I S) := by + have : E I S = Prod.fst ⁻¹' (E' I S) := by ext; simp [E, E'] + rw [this] + exact measurable_fst (measurableSet_E' I S) + +lemma iIndepFun_rewardByCount.extracted_1 [Countable α] [Nonempty α] + (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (I : Finset (α × ℕ)) + {B : α × ℕ → Set ℝ} (hB : ∀ i ∈ I, MeasurableSet (B i)) (S : Finset I) : + 𝔓 (E I S ∩ ⋂ i ∈ S, (fun ω ↦ reward (stepsUntil i.1.1 (i.1.2 + 1) ω.1).toNat ω.1) ⁻¹' B i) = + 𝔓 (E I S) * ∏ i ∈ S, (ν i.1.1) (B i) := by sorry +lemma iIndepFun_rewardByCount.extracted_2 [Countable α] [Nonempty α] + (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (I : Finset (α × ℕ)) + {B : α × ℕ → Set ℝ} (hB : ∀ i ∈ I, MeasurableSet (B i)) (S : Finset I) : + Bandit.measure alg ν (⋂ j ∉ S, (fun ω ↦ ω.2 (j.1.2 + 1) j.1.1) ⁻¹' B j) = + ∏ j ∉ S, (ν j.1.1) (B j) := by + have h_indep : iIndepFun (fun (i : I) ω ↦ ω.2 (i.1.2 + 1) i.1.1) (Bandit.measure alg ν) := by + suffices iIndepFun (fun (i : I) ω ↦ ω (i.1.2 + 1) i.1.1) (Bandit.streamMeasure ν) by + sorry + sorry + rw [iIndepFun_iff_measure_inter_preimage_eq_mul] at h_indep + specialize h_indep Sᶜ (sets := fun i ↦ B i) (fun i hi ↦ hB i i.2) + simp only [mem_compl] at h_indep ⊢ + rw [h_indep] + congr with i + rw [← Measure.map_apply (by fun_prop) (hB i i.2)] + congr + exact (hasLaw_Z i.1.1 (i.1.2 + 1)).map_eq + +lemma iIndepFun_rewardByCount.extracted_3 [Countable α] [Nonempty α] + (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [inst_4 : IsMarkovKernel ν] (I : Finset (α × ℕ)) + {B : α × ℕ → Set ℝ} (hB : ∀ i ∈ I, MeasurableSet (B i)) (S : Finset I) : + IndepSet (E I S ∩ + ⋂ i ∈ S, (fun ω ↦ reward (stepsUntil i.1.1 (i.1.2 + 1) ω.1).toNat ω.1) ⁻¹' B i) + (⋂ j ∉ S, (fun ω ↦ ω.2 (j.1.2 + 1) j.1.1) ⁻¹' B j) (Bandit.measure alg ν) := by + let A := E I S ∩ ⋂ i ∈ S, (fun ω ↦ reward (stepsUntil i.1.1 (i.1.2 + 1) ω.1).toNat ω.1) ⁻¹' B i + let C := ⋂ j ∉ S, (fun (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) ↦ ω.2 (j.1.2 + 1) j.1.1) ⁻¹' B j + let A' := E' I S ∩ ⋂ i ∈ S, (fun ω ↦ reward (stepsUntil i.1.1 (i.1.2 + 1) ω).toNat ω) ⁻¹' B i + have hA : A = Prod.fst ⁻¹' A' := by + ext ω + simp [A, A', E, E'] + let C' := ⋂ j ∉ S, (fun (ω : ℕ → α → ℝ) ↦ ω (j.1.2 + 1) j.1.1) ⁻¹' B j + have hC : C = Prod.snd ⁻¹' C' := by + ext ω + simp [C, C'] + have hAC : A ∩ C = A' ×ˢ C' := by rw [hA, hC]; ext; simp + have hA'_meas : MeasurableSet A' := by + refine MeasurableSet.inter ?_ (MeasurableSet.iInter fun i ↦ MeasurableSet.iInter fun hi ↦ ?_) + · exact measurableSet_E' I S + · exact (hB i.1 i.2).preimage (by fun_prop) + have hC'_meas : MeasurableSet C' := by + refine MeasurableSet.iInter fun j ↦ MeasurableSet.iInter fun hj ↦ ?_ + exact (hB j.1 j.2).preimage (by fun_prop) + change IndepSet A C (Bandit.measure alg ν) + rw [indepSet_iff_measure_inter_eq_mul (μ := Bandit.measure alg ν)] + rotate_left + · rw [hA] + exact measurable_fst hA'_meas + · rw [hC] + exact measurable_snd hC'_meas + rw [hAC, hA, hC, Bandit.measure, ← Measure.fst_apply, ← Measure.snd_apply] + · simp + · exact hC'_meas + · exact hA'_meas + +lemma iIndepFun_rewardByCount [Countable α] [Nonempty α] + (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] : + iIndepFun (fun (p : α × ℕ) ↦ rewardByCount p.1 (p.2 + 1)) (Bandit.measure alg ν) := by + rw [iIndepFun_iff_measure_inter_preimage_eq_mul] + intro I B hB + suffices Bandit.measure alg ν (⋂ i ∈ I, rewardByCount i.1 (i.2 + 1) ⁻¹' B i) = + ∏ i ∈ I, ν i.1 (B i) by + rw [this] + refine Finset.prod_congr rfl fun i hi ↦ ?_ + rw [← Measure.map_apply (by fun_prop) (hB i hi)] + congr + exact (hasLaw_rewardByCount i.1 (i.2 + 1) (by simp)).map_eq.symm + have hE_disj (S T : Finset I) (hST : S ≠ T) : Disjoint (E I S) (E I T) := by + rw [Set.disjoint_iff_forall_ne] + simp only [Subtype.forall, Prod.forall, Set.mem_setOf_eq, ne_eq, and_imp, Prod.mk.injEq, + not_and, E] + grind + have hE_union : ⋃ (S : Finset I), E I S = Set.univ := by + ext ω + simp only [Subtype.forall, Prod.forall, Set.mem_iUnion, Set.mem_setOf_eq, Set.mem_univ, + iff_true, E] + use Finset.univ.filter (fun i : I ↦ stepsUntil i.1.1 (i.1.2 + 1) ω.1 < ⊤) + simp + have : ⋂ i ∈ I, rewardByCount i.1 (i.2 + 1) ⁻¹' B i = + ⋃ (S : Finset I), E I S ∩ (⋂ i ∈ I, rewardByCount i.1 (i.2 + 1) ⁻¹' B i) := by + rw [← Set.iUnion_inter, hE_union, Set.univ_inter] + rw [this, measure_iUnion] + rotate_left + · intro S T hST + simp only [Function.onFun] + exact Disjoint.inter_left _ (Disjoint.inter_right _ (hE_disj S T hST)) + · refine fun S ↦ (measurableSet_E I S).inter ?_ + refine MeasurableSet.iInter fun i ↦ MeasurableSet.iInter fun hi ↦ ?_ + exact (hB i hi).preimage (by fun_prop) + suffices ∀ (S : Finset I), + Bandit.measure alg ν (E I S ∩ ⋂ i ∈ I, rewardByCount i.1 (i.2 + 1) ⁻¹' B i) = + Bandit.measure alg ν (E I S) * ∏ i ∈ I, ν i.1 (B i) by + simp_rw [this] + rw [ENNReal.tsum_mul_right, ← measure_iUnion hE_disj (measurableSet_E I), hE_union, + measure_univ, one_mul] + intro S + have h_eq : E I S ∩ ⋂ i ∈ I, rewardByCount i.1 (i.2 + 1) ⁻¹' B i + = E I S ∩ (⋂ i ∈ S, (fun ω ↦ reward (stepsUntil i.1.1 (i.1.2 + 1) ω.1).toNat ω.1) ⁻¹' B i) ∩ + (⋂ j ∉ S, (fun ω ↦ ω.2 (j.1.2 + 1) j.1.1) ⁻¹' B j) := by + ext ω + rw [Set.inter_assoc] + simp only [Set.mem_inter_iff, and_congr_right_iff] + intro hω + simp only [Subtype.forall, Prod.forall, Set.mem_setOf_eq, E] at hω + conv_rhs => rw [Set.iInter_subtype, Set.iInter_subtype] + rw [← Set.mem_inter_iff, ← Set.iInter_inter_distrib] + simp_rw [← Set.iInter_inter_distrib] + simp only [Set.mem_iInter, Set.mem_preimage, Prod.forall, Set.mem_inter_iff] + constructor + · intro h_all a i hai + constructor + · intro haiS + convert h_all a i hai + replace hω := hω.1 a i hai haiS + rw [rewardByCount_of_stepsUntil_ne_top hω.ne] + rfl + · intro haiS + convert h_all a i hai + replace hω := hω.2 a i hai haiS + rw [rewardByCount_of_stepsUntil_eq_top hω] + · intro h a i hai + specialize h a i hai + by_cases haiS : ⟨⟨a, i⟩, hai⟩ ∈ S + · convert h.1 haiS + replace hω := hω.1 a i hai haiS + rw [rewardByCount_of_stepsUntil_ne_top hω.ne] + rfl + · convert h.2 haiS + replace hω := hω.2 a i hai haiS + rw [rewardByCount_of_stepsUntil_eq_top hω] + rw [h_eq, IndepSet.measure_inter_eq_mul] + swap; · exact iIndepFun_rewardByCount.extracted_3 alg ν I hB S + rw [iIndepFun_rewardByCount.extracted_1 alg ν I hB S, + iIndepFun_rewardByCount.extracted_2 alg ν I hB S, mul_assoc] + congr + rw [Finset.prod_mul_prod_compl, Finset.prod_subtype I] + simp + +lemma identDistrib_rewardByCount_stream_all [Countable α] [StandardBorelSpace α] [Nonempty α] : + IdentDistrib (fun ω (p : α × ℕ) ↦ rewardByCount p.1 (p.2 + 1) ω) + (fun ω p ↦ ω p.2 p.1) 𝔓 (Bandit.streamMeasure ν) := by + refine IdentDistrib.pi (fun p ↦ ?_) ?_ ?_ + · refine identDistrib_rewardByCount_eval p.1 (p.2 + 1) p.2 (by simp) (ν := ν) + · exact iIndepFun_rewardByCount alg ν + · sorry + lemma identDistrib_rewardByCount_stream' [Countable α] [StandardBorelSpace α] [Nonempty α] (a : α) : IdentDistrib (fun ω n ↦ rewardByCount a (n + 1) ω) (fun ω n ↦ ω n a) diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index e4f6d6b5..6153e550 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -461,10 +461,24 @@ lemma rewardByCount_of_stepsUntil_eq_top {ω : (ℕ → α × R) × (ℕ → α (h : stepsUntil a m ω.1 = ⊤) : rewardByCount a m ω = ω.2 m a := by simp [rewardByCount_eq_ite, h] +lemma rewardByCount_of_stepsUntil_ne_top {ω : (ℕ → α × R) × (ℕ → α → R)} + (h : stepsUntil a m ω.1 ≠ ⊤) : + rewardByCount a m ω = reward (stepsUntil a m ω.1).toNat ω.1 := by simp [rewardByCount_eq_ite, h] + lemma rewardByCount_of_stepsUntil_eq_coe {ω : (ℕ → α × R) × (ℕ → α → R)} (h : stepsUntil a m ω.1 = n) : rewardByCount a m ω = reward n ω.1 := by simp [rewardByCount_eq_ite, h] +/-- The value at 0 does not matter (it would be the "zeroth" reward). +It should be considered a junk value. -/ +@[simp] +lemma rewardByCount_zero (a : α) (ω : (ℕ → α × R) × (ℕ → α → R)) : + rewardByCount a 0 ω = if action 0 ω.1 = a then ω.2 0 a else reward 0 ω.1 := by + rw [rewardByCount_eq_ite] + by_cases ha : action 0 ω.1 = a + · simp [ha, stepsUntil_zero_of_eq] + · simp [stepsUntil_zero_of_ne, ha] + lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (ω : (ℕ → α × R) × (ℕ → α → R)) : rewardByCount (action t ω.1) (pullCount (action t ω.1) t ω.1 + 1) ω = reward t ω.1 := by rw [rewardByCount, ← pullCount_action_eq_pullCount_add_one, stepsUntil_pullCount_eq] diff --git a/blueprint/lean_decls b/blueprint/lean_decls index a677429e..6dace1d1 100644 --- a/blueprint/lean_decls +++ b/blueprint/lean_decls @@ -45,13 +45,11 @@ Learning.sum_rewardByCount_eq_sumRewards Bandits.Bandit.trajMeasure Bandits.Bandit.measure Bandits.measurable_comap_indicator_stepsUntil_eq -ProbabilityTheory.CondIndepFun.prod_right Bandits.condIndepFun_reward_stepsUntil_arm Bandits.reward_cond_stepsUntil ProbabilityTheory.condDistrib_ae_eq_cond Bandits.condDistrib_rewardByCount_stepsUntil Bandits.hasLaw_rewardByCount -ProbabilityTheory.iIndepFun_nat_iff_forall_indepFun Bandits.iIndepFun_rewardByCount' Bandits.identDistrib_rewardByCount_stream Bandits.identDistrib_sum_Icc_rewardByCount @@ -83,4 +81,6 @@ Bandits.UCB.pullCount_le_add_three Bandits.UCB.pullCount_le_add_three_ae Bandits.UCB.some_sum_eq_zero Bandits.UCB.expectation_pullCount_le -Bandits.UCB.regret_le \ No newline at end of file +Bandits.UCB.regret_le +ProbabilityTheory.CondIndepFun.prod_right +ProbabilityTheory.iIndepFun_nat_iff_forall_indepFun \ No newline at end of file From be27b5651f205e5d68dcee0f8332a0696eecfc7e Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sat, 3 Jan 2026 11:58:36 +0100 Subject: [PATCH 02/30] define new filtration --- LeanBandits/RewardByCountMeasure.lean | 40 ------- LeanBandits/SequentialLearning/Algorithm.lean | 108 ++++++++++++++++-- .../SequentialLearning/FiniteActions.lean | 93 ++++++++++++++- blueprint/lean_decls | 2 +- blueprint/src/chapters/bandit.tex | 2 +- 5 files changed, 189 insertions(+), 56 deletions(-) diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 250d2ba9..e017dd70 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -81,46 +81,6 @@ lemma reward_cond_arm [StandardBorelSpace α] [Nonempty α] [Countable α] (a : rw [h_ra] at h_eq exact h_eq.symm --- after the Mathlib stopping time refactor, we will be able to prove that stepsUntil is a --- stopping time -lemma measurable_comap_indicator_stepsUntil_eq (a : α) (m n : ℕ) : - Measurable[MeasurableSpace.comap (fun ω : ℕ → α × ℝ ↦ (hist (n-1) ω, arm n ω)) inferInstance] - ({ω | stepsUntil a m ω = ↑n}.indicator fun _ ↦ 1) := by - let k : ((Iic (n - 1) → α × ℝ) × α) → (ℕ → α × ℝ) := fun x i ↦ - if hi : i ∈ Iic (n - 1) then (x.1 ⟨i, hi⟩) else if i = n then (x.2, 0) else (a, 0) - have hk : Measurable k := by - unfold k - rw [measurable_pi_iff] - intro i - split_ifs <;> fun_prop - let φ : ((Iic (n - 1) → α × ℝ) × α) → ℕ := fun x ↦ if stepsUntil a m (k x) = ↑n then 1 else 0 - have hφ : Measurable φ := - Measurable.ite ((measurableSet_singleton _).preimage (by fun_prop)) (by fun_prop) (by fun_prop) - suffices {ω | stepsUntil a m ω = ↑n}.indicator (fun x ↦ 1) - = φ ∘ fun ω ↦ (hist (n - 1) ω, arm n ω) from this ▸ measurable_comp_comap _ hφ - 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, k, action] - grind - -lemma measurable_indicator_stepsUntil_eq (a : α) (m n : ℕ) : - Measurable ({ω : ℕ → α × ℝ | stepsUntil a m ω = ↑n}.indicator fun _ ↦ 1) := by - refine (measurable_comap_indicator_stepsUntil_eq a m n).mono ?_ le_rfl - refine Measurable.comap_le ?_ - fun_prop - -lemma measurableSet_stepsUntil_eq (a : α) (m n : ℕ) : - MeasurableSet[MeasurableSpace.comap (fun ω : ℕ → α × ℝ ↦ (hist (n-1) ω, arm n ω)) inferInstance] - {ω : ℕ → α × ℝ | stepsUntil a m ω = ↑n} := by - let mProd := MeasurableSpace.comap (fun ω : ℕ → α × ℝ ↦ (hist (n-1) ω, arm n ω)) inferInstance - suffices Measurable[mProd] ({ω | stepsUntil a m ω = ↑n}.indicator fun x ↦ 1) by - rwa [measurable_indicator_const_iff] at this - exact measurable_comap_indicator_stepsUntil_eq a m n - lemma condIndepFun_reward_stepsUntil_arm' [StandardBorelSpace α] [Countable α] [Nonempty α] (a : α) (m n : ℕ) (hm : m ≠ 0) : reward n ⟂ᵢ[arm n, measurable_arm n; 𝔓t] {ω | stepsUntil a m ω = ↑n}.indicator (fun _ ↦ 1) := by diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index c9736eaf..418bc5cb 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -130,6 +130,10 @@ protected def filtration (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) := MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R) +lemma filtration_eq_comap (n : ℕ) : + Learning.filtration α R n = MeasurableSpace.comap (hist n) inferInstance := by + simp [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe] + lemma step_eq_eval_comp_hist (n : ℕ) : step (α := α) (R := R) n = (fun x ↦ x ⟨n, by simp⟩) ∘ (hist n) := rfl @@ -140,8 +144,7 @@ lemma reward_eq_eval_comp_hist (n : ℕ) : reward (α := α) (R := R) n = (fun x ↦ (x ⟨n, by simp⟩).2) ∘ (hist n) := rfl lemma measurable_step_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (step n) := by - simp only [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe] - rw [step_eq_eval_comp_hist] + rw [filtration_eq_comap, step_eq_eval_comp_hist] exact measurable_comp_comap _ (by fun_prop) lemma adapted_step [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] @@ -152,8 +155,7 @@ lemma adapted_step [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace fun n ↦ (measurable_step_filtration n).stronglyMeasurable lemma measurable_hist_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (hist n) := by - simp [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe, - measurable_iff_comap_le] + simp [filtration_eq_comap, measurable_iff_comap_le] lemma adapted_hist [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] [SecondCountableTopology α] [OpensMeasurableSpace α] @@ -163,8 +165,7 @@ lemma adapted_hist [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace fun n ↦ (measurable_hist_filtration n).stronglyMeasurable lemma measurable_action_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (action n) := by - simp only [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe] - rw [action_eq_eval_comp_hist] + rw [filtration_eq_comap, action_eq_eval_comp_hist] exact measurable_comp_comap _ (by fun_prop) lemma adapted_action [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] @@ -173,8 +174,7 @@ lemma adapted_action [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpa fun n ↦ (measurable_action_filtration n).stronglyMeasurable lemma measurable_reward_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (reward n) := by - simp only [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe] - rw [reward_eq_eval_comp_hist] + rw [filtration_eq_comap, reward_eq_eval_comp_hist] exact measurable_comp_comap _ (by fun_prop) lemma adapted_reward [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] @@ -182,6 +182,96 @@ lemma adapted_reward [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpac Adapted (Learning.filtration α R) reward := fun n ↦ (measurable_reward_filtration n).stronglyMeasurable +section FiltrationAction + +/-- Filtration generated by the history at time `n-1` together with the action at time `n`. -/ +def filtrationAction (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : + Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) where + seq n := if n = 0 then MeasurableSpace.comap (action 0) inferInstance + else Learning.filtration α R (n - 1) ⊔ MeasurableSpace.comap (action n) inferInstance + mono' n m hnm := by + simp only + by_cases hn : n = 0 + · by_cases hm : m = 0 + · simp [hn, hm] + · simp only [hn, ↓reduceIte, hm] + refine le_sup_of_le_left ?_ + rw [← measurable_iff_comap_le] + suffices Measurable[Learning.filtration α R 0] (action 0) from + this.mono ((Learning.filtration α R).mono zero_le') le_rfl + exact measurable_action_filtration 0 + have hm : m ≠ 0 := by grind + simp only [hn, hm, ↓reduceIte] + have hnm' : n - 1 ≤ m - 1 := by grind + simp only [sup_le_iff] + constructor + · refine le_sup_of_le_left ?_ + exact (Learning.filtration α R).mono hnm' + · rcases eq_or_lt_of_le hnm with rfl | hlt + · exact le_sup_of_le_right le_rfl + refine le_sup_of_le_left ?_ + rw [← measurable_iff_comap_le] + have h_le : n ≤ m - 1 := by grind + suffices Measurable[Learning.filtration α R n] (action n) from + this.mono ((Learning.filtration α R).mono h_le) le_rfl + exact measurable_action_filtration n + le' n := by + by_cases hn : n = 0 + · simp only [hn, ↓reduceIte] + rw [← measurable_iff_comap_le] + fun_prop + simp only [hn, ↓reduceIte, sup_le_iff] + constructor + · exact (Learning.filtration α R).le _ + · rw [← measurable_iff_comap_le] + fun_prop + +lemma filtrationAction_zero_eq_comap : + filtrationAction α R 0 = MeasurableSpace.comap (action 0) inferInstance := by + simp [filtrationAction] + +lemma filtrationAction_eq_comap (n : ℕ) (hn : n ≠ 0) : + filtrationAction α R n = + MeasurableSpace.comap (fun ω ↦ (hist (n - 1) ω, action n ω)) inferInstance := by + simp only [filtrationAction, filtration_eq_comap, ← MeasurableSpace.comap_prodMk, hn, ↓reduceIte] + rfl + +lemma filtration_le_filtrationAction_add_one (n : ℕ) : + Learning.filtration α R n ≤ filtrationAction α R (n + 1) := le_sup_of_le_left le_rfl + +lemma filtration_le_filtrationAction {m n : ℕ} (h : n < m) : + Learning.filtration α R n ≤ filtrationAction α R m := by + have h' : n + 1 ≤ m := by grind + exact (filtration_le_filtrationAction_add_one n).trans ((filtrationAction α R).mono h') + +lemma filtrationAction_le_filtration_self (n : ℕ) : + filtrationAction α R n ≤ Learning.filtration α R n := by + by_cases hn : n = 0 + · simp only [hn, filtrationAction_zero_eq_comap] + rw [← measurable_iff_comap_le] + exact measurable_action_filtration 0 + simp only [filtrationAction, hn, ↓reduceIte, sup_le_iff] + constructor + · exact (Learning.filtration α R).mono (by grind) + · rw [← measurable_iff_comap_le] + exact measurable_action_filtration _ + +lemma filtrationAction_le_filtration {m n : ℕ} (h : m ≤ n) : + filtrationAction α R m ≤ Learning.filtration α R n := + (filtrationAction_le_filtration_self m).trans ((Learning.filtration α R).mono h) + +lemma measurable_action_filtrationAction (n : ℕ) : + Measurable[filtrationAction α R n] (action n) := by + simp only [filtrationAction] + rw [measurable_iff_comap_le] + split_ifs with hn + · simp [hn] + · exact le_sup_of_le_right le_rfl + +end FiltrationAction + +section Laws + lemma condDistrib_step [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : condDistrib (step (n + 1)) (hist n) (trajMeasure alg env) @@ -235,4 +325,6 @@ lemma condDistrib_reward_zero [StandardBorelSpace R] [Nonempty R] have h_action := (hasLaw_action_zero alg env).map_eq rwa [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop), h_action] +end Laws + end Learning diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index 6153e550..cbd1fbfa 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -397,6 +397,8 @@ lemma stepsUntil_eq_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h rw [pullCount_congr] grind +section Measurability + lemma isStoppingTime_stepsUntil [MeasurableSingletonClass α] (a : α) (hm : m ≠ 0) : IsStoppingTime (Learning.filtration α ℝ) (stepsUntil a m) := by rw [stepsUntil_eq_leastGE _ hm] @@ -438,6 +440,80 @@ lemma measurable_stepsUntil' [MeasurableSingletonClass α] (a : α) (m : ℕ) : Measurable (fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ stepsUntil a m ω.1) := (measurable_stepsUntil a m).comp measurable_fst +lemma measurable_comap_indicator_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] + (a : α) (m n : ℕ) : + Measurable[MeasurableSpace.comap (fun ω : ℕ → α × R ↦ (hist (n-1) ω, action n ω)) inferInstance] + ({ω | stepsUntil a m ω = ↑n}.indicator fun _ ↦ 1) := by + let r₀ : R := Nonempty.some inferInstance + let k : ((Iic (n - 1) → α × R) × α) → (ℕ → α × R) := fun x i ↦ + if hi : i ∈ Iic (n - 1) then (x.1 ⟨i, hi⟩) else if i = n then (x.2, r₀) else (a, r₀) + have hk : Measurable k := by + unfold k + rw [measurable_pi_iff] + intro i + split_ifs <;> fun_prop + let φ : ((Iic (n - 1) → α × R) × α) → ℕ := fun x ↦ if stepsUntil a m (k x) = ↑n then 1 else 0 + have hφ : Measurable φ := + Measurable.ite ((measurableSet_singleton _).preimage (by fun_prop)) (by fun_prop) (by fun_prop) + suffices {ω | stepsUntil a m ω = ↑n}.indicator (fun x ↦ 1) + = φ ∘ fun ω ↦ (hist (n - 1) ω, action n ω) from this ▸ measurable_comp_comap _ hφ + 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 [action, mem_Iic, hist, dite_eq_ite, k, action] + grind + +lemma measurable_indicator_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] + (a : α) (m n : ℕ) : + Measurable ({ω : ℕ → α × R | stepsUntil a m ω = ↑n}.indicator fun _ ↦ 1) := by + refine (measurable_comap_indicator_stepsUntil_eq (mR := mR) a m n).mono ?_ le_rfl + refine Measurable.comap_le ?_ + fun_prop + +lemma measurableSet_stepsUntil_eq_zero [Nonempty R] [MeasurableSingletonClass α] (a : α) (m : ℕ) : + MeasurableSet[MeasurableSpace.comap (action 0) inferInstance] + {ω : ℕ → α × R | stepsUntil a m ω = 0} := by + simp only [stepsUntil_eq_zero_iff (a := a) (m := m), ne_eq] + by_cases hm : m = 0 + · simp only [hm, true_and, zero_ne_one, false_and, or_false] + refine (measurableSet_singleton _).compl.preimage ?_ + rw [measurable_iff_comap_le] + by_cases hm1 : m = 1 + swap; · simp [hm, hm1] + simp only [hm1, one_ne_zero, false_and, true_and, false_or] + refine (measurableSet_singleton _).preimage ?_ + rw [measurable_iff_comap_le] + +lemma measurableSet_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] (a : α) (m n : ℕ) : + MeasurableSet[MeasurableSpace.comap (fun ω : ℕ → α × R ↦ (hist (n-1) ω, action n ω)) + inferInstance] + {ω : ℕ → α × R | stepsUntil a m ω = ↑n} := by + let mProd := MeasurableSpace.comap (fun ω : ℕ → α × R ↦ (hist (n-1) ω, action n ω)) inferInstance + suffices Measurable[mProd] ({ω | stepsUntil a m ω = ↑n}.indicator fun x ↦ 1) by + rwa [measurable_indicator_const_iff] at this + exact measurable_comap_indicator_stepsUntil_eq a m n + +/-- `stepsUntil a m` is a stopping time with respect to the filtration `filtrationAction`. -/ +theorem isStoppingTime_stepsUntil_filtrationAction [Nonempty R] [MeasurableSingletonClass α] + (a : α) (m : ℕ) : + IsStoppingTime (filtrationAction α R) (stepsUntil a m) := by + refine isStoppingTime_of_measurableSet_eq fun n ↦ ?_ + by_cases hn : n = 0 + · simp only [hn, filtrationAction_zero_eq_comap, WithTop.coe_zero] + exact measurableSet_stepsUntil_eq_zero a m + · rw [filtrationAction_eq_comap _ hn] + exact measurableSet_stepsUntil_eq a m n + +-- /-- Sigma-algebra generated by the stopping time `stepsUntil a m`. -/ +-- def stepsUntilMeasurableSpace [Nonempty R] [MeasurableSingletonClass α] (a : α) (m : ℕ) : +-- MeasurableSpace (ℕ → α × R) := +-- (isStoppingTime_stepsUntil_filtrationAction a m (mR := mR)).measurableSpace + +end Measurability + end StepsUntil section RewardByCount @@ -451,22 +527,27 @@ def rewardByCount (a : α) (m : ℕ) (ω : (ℕ → α × R) × (ℕ → α → | ⊤ => ω.2 m a | (n : ℕ) => reward n ω.1 +variable {ω : (ℕ → α × R) × (ℕ → α → R)} + lemma rewardByCount_eq_ite (a : α) (m : ℕ) (ω : (ℕ → α × R) × (ℕ → α → R)) : rewardByCount a m ω = if (stepsUntil a m ω.1) = ⊤ then ω.2 m a else reward (stepsUntil a m ω.1).toNat ω.1 := by unfold rewardByCount cases stepsUntil a m ω.1 <;> simp -lemma rewardByCount_of_stepsUntil_eq_top {ω : (ℕ → α × R) × (ℕ → α → R)} - (h : stepsUntil a m ω.1 = ⊤) : +lemma rewardByCount_of_stepsUntil_eq_top (h : stepsUntil a m ω.1 = ⊤) : rewardByCount a m ω = ω.2 m a := by simp [rewardByCount_eq_ite, h] -lemma rewardByCount_of_stepsUntil_ne_top {ω : (ℕ → α × R) × (ℕ → α → R)} - (h : stepsUntil a m ω.1 ≠ ⊤) : +lemma rewardByCount_of_stepsUntil_ne_top (h : stepsUntil a m ω.1 ≠ ⊤) : rewardByCount a m ω = reward (stepsUntil a m ω.1).toNat ω.1 := by simp [rewardByCount_eq_ite, h] -lemma rewardByCount_of_stepsUntil_eq_coe {ω : (ℕ → α × R) × (ℕ → α → R)} - (h : stepsUntil a m ω.1 = n) : +lemma rewardByCount_eq_stoppedValue (h : stepsUntil a m ω.1 ≠ ⊤) : + rewardByCount a m ω = stoppedValue reward (stepsUntil a m) ω.1 := by + rw [rewardByCount_of_stepsUntil_ne_top h, stoppedValue] + lift stepsUntil a m ω.1 to ℕ using h with n + simp + +lemma rewardByCount_of_stepsUntil_eq_coe (h : stepsUntil a m ω.1 = n) : rewardByCount a m ω = reward n ω.1 := by simp [rewardByCount_eq_ite, h] /-- The value at 0 does not matter (it would be the "zeroth" reward). diff --git a/blueprint/lean_decls b/blueprint/lean_decls index 6dace1d1..0c5255e1 100644 --- a/blueprint/lean_decls +++ b/blueprint/lean_decls @@ -44,7 +44,7 @@ Learning.empMean Learning.sum_rewardByCount_eq_sumRewards Bandits.Bandit.trajMeasure Bandits.Bandit.measure -Bandits.measurable_comap_indicator_stepsUntil_eq +Learning.measurable_comap_indicator_stepsUntil_eq Bandits.condIndepFun_reward_stepsUntil_arm Bandits.reward_cond_stepsUntil ProbabilityTheory.condDistrib_ae_eq_cond diff --git a/blueprint/src/chapters/bandit.tex b/blueprint/src/chapters/bandit.tex index a3f3af6a..72d9e31f 100644 --- a/blueprint/src/chapters/bandit.tex +++ b/blueprint/src/chapters/bandit.tex @@ -51,7 +51,7 @@ \section{Alternative models: rewards indexed by time or pull count}\label{sec:al \begin{lemma}\label{lem:measurable_comap_indicator_stepsUntil_eq} \uses{def:stepsUntil} \leanok - \lean{Bandits.measurable_comap_indicator_stepsUntil_eq} + \lean{Learning.measurable_comap_indicator_stepsUntil_eq} The function $\mathbb{I}\{T_{n,a} = t\} : \Omega \to \{0, 1\}$ is measurable with respect to the sigma-algebra generated by $(H_{t-1}, A_t)$. \end{lemma} From a00150d25a07c393f85521ac368372798872b229 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Mon, 5 Jan 2026 11:02:34 +0100 Subject: [PATCH 03/30] lots of stuff --- LeanBandits.lean | 1 + LeanBandits/Bandit/Bandit.lean | 209 +++++ LeanBandits/ForMathlib/HasCondDistrib.lean | 49 +- LeanBandits/RewardByCountMeasure.lean | 84 +- LeanBandits/SequentialLearning/Algorithm.lean | 201 ++++- LeanBandits/SequentialLearning/Draft.lean | 772 ++++++++++++++++++ .../SequentialLearning/FiniteActions.lean | 17 + 7 files changed, 1272 insertions(+), 61 deletions(-) create mode 100644 LeanBandits/SequentialLearning/Draft.lean diff --git a/LeanBandits.lean b/LeanBandits.lean index 10b0f463..d6909492 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -14,5 +14,6 @@ import LeanBandits.ForMathlib.Traj import LeanBandits.RewardByCountMeasure import LeanBandits.SequentialLearning.Algorithm import LeanBandits.SequentialLearning.Deterministic +import LeanBandits.SequentialLearning.Draft import LeanBandits.SequentialLearning.FiniteActions import LeanBandits.SequentialLearning.StationaryEnv diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index d379032d..5d7dd348 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -6,7 +6,9 @@ Authors: Rémy Degenne, Paulo Rauber import LeanBandits.ForMathlib.IndepInfinitePi import LeanBandits.SequentialLearning.Deterministic import LeanBandits.SequentialLearning.StationaryEnv +import LeanBandits.SequentialLearning.FiniteActions import Mathlib.Probability.IdentDistrib +import Mathlib.MeasureTheory.Constructions.UnitInterval /-! # Bandit @@ -286,6 +288,213 @@ example [StandardBorelSpace α] [Nonempty α] end DetAlgorithm +section ArrayModel + +open unitInterval + +section Aux + +-- from Mathlib PR #30112 +theorem representation {α β : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + [Nonempty β] [StandardBorelSpace β] + (κ : Kernel α β) [IsMarkovKernel κ] : + ∃ (f : α → I → β), Measurable (Function.uncurry f) ∧ ∀ a, volume.map (f a) = κ a := sorry + +theorem representation_measure {β : Type*} {mβ : MeasurableSpace β} + [Nonempty β] [StandardBorelSpace β] + (μ : Measure β) [IsProbabilityMeasure μ] : + ∃ (f : I → β), Measurable f ∧ volume.map f = μ := by + obtain ⟨f, hf_meas, hf_map⟩ := representation (Kernel.const Unit μ) + specialize hf_map ⟨⟩ + exact ⟨f ⟨⟩, by fun_prop, by simpa⟩ + +end Aux + +variable (α R) in +def probSpace : Type _ := (ℕ → I) × (ℕ → α → R) + +instance {α R : Type*} [MeasurableSpace R] : MeasurableSpace (probSpace α R) := + inferInstanceAs (MeasurableSpace ((ℕ → I) × (ℕ → α → R))) + +noncomputable +def arrayMeasure (ν : Kernel α R) : Measure (probSpace α R) := + (Measure.infinitePi fun _ ↦ volume).prod (Bandit.streamMeasure ν) + +instance (ν : Kernel α R) [IsMarkovKernel ν] : IsProbabilityMeasure (arrayMeasure ν) := + Measure.prod.instIsProbabilityMeasure _ _ + +variable [Nonempty α] [StandardBorelSpace α] + +noncomputable +def initAlgFunction (alg : Algorithm α R) : I → α := + (representation_measure alg.p0).choose + +lemma initAlgFunction_map (alg : Algorithm α R) : volume.map (initAlgFunction alg) = alg.p0 := + (representation_measure alg.p0).choose_spec.2 + +@[fun_prop] +lemma measurable_initAlgFunction (alg : Algorithm α R) : + Measurable (initAlgFunction alg) := (representation_measure alg.p0).choose_spec.1 + +noncomputable +def algFunction (alg : Algorithm α R) (n : ℕ) : + (Iic n → α × R) → I → α := + (representation (alg.policy n)).choose + +lemma algFunction_map (alg : Algorithm α R) (n : ℕ) (h : Iic n → α × R) : + volume.map (algFunction alg n h) = alg.policy n h := + (representation (alg.policy n)).choose_spec.2 h + +@[fun_prop] +lemma measurable_algFunction (alg : Algorithm α R) (n : ℕ) : + Measurable (Function.uncurry (algFunction alg n)) := + (representation (alg.policy n)).choose_spec.1 + +noncomputable +def altHist [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) : (n : ℕ) → Iic n → α × R +| 0 => fun _ ↦ (initAlgFunction alg (ω.1 0), ω.2 0 (initAlgFunction alg (ω.1 0))) +| n + 1 => + let hn : Iic n → α × R := altHist alg ω n + let a : α := algFunction alg n hn (ω.1 (n + 1)) + fun i ↦ if hin : i ≤ n then hn ⟨i, by simp [hin]⟩ else (a, ω.2 (pullCount' n hn a + 1) a) + +@[simp] +lemma altHist_zero [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) : + altHist alg ω 0 = fun _ ↦ (initAlgFunction alg (ω.1 0), ω.2 0 (initAlgFunction alg (ω.1 0))) := + rfl + +lemma altHist_add_one [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) (n : ℕ) : + let a : α := algFunction alg n (altHist alg ω n) (ω.1 (n + 1)) + altHist alg ω (n + 1) = + fun (i : Iic (n + 1)) ↦ if hin : i ≤ n then altHist alg ω n ⟨i, by simp [hin]⟩ + else (a, ω.2 (pullCount' n (altHist alg ω n) a + 1) a) := + rfl + +lemma altHist_eq [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) (n : ℕ) : + altHist alg ω n = fun i : Iic n ↦ altHist alg ω i ⟨i.1, by simp⟩ := by + induction n with + | zero => + ext i : 1 + simp only [altHist] + sorry + | succ n hn => + ext i : 1 + by_cases hin : i ≤ n + · rw [altHist_add_one] + simp only [hin, ↓reduceDIte] + rw [funext_iff] at hn + simp_rw [hn] + · grind + +@[fun_prop] +lemma measurable_altHist [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : + Measurable (fun ω ↦ altHist alg ω n) := by + induction n with + | zero => + simp_rw [altHist_zero, measurable_pi_iff] + refine fun _ ↦ Measurable.prodMk (by fun_prop) ?_ + sorry + | succ n hn => + refine measurable_pi_iff.mpr fun i ↦ ?_ + by_cases hin : i ≤ n + · simp only [altHist, hin, ↓reduceDIte] + rw [measurable_pi_iff] at hn + exact hn ⟨i.1, by simp [hin]⟩ + · simp only [altHist, hin, ↓reduceDIte] + refine Measurable.prodMk (by fun_prop) ?_ + sorry + +noncomputable +def altArm [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : α := + (altHist alg ω n ⟨n, by simp⟩).1 + +lemma altArm_zero [DecidableEq α] (alg : Algorithm α R) : + altArm alg 0 = fun ω ↦ initAlgFunction alg (ω.1 0) := by + ext + simp [altArm, altHist_zero] + +@[fun_prop] +lemma measurable_altArm [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : + Measurable (altArm alg n) := by unfold altArm; fun_prop + +noncomputable +def altReward [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : R := + (altHist alg ω n ⟨n, by simp⟩).2 + +lemma altReward_zero [DecidableEq α] (alg : Algorithm α R) : + altReward alg 0 = fun ω ↦ ω.2 0 (altArm alg 0 ω) := by + ext + simp [altReward, altHist_zero, altArm_zero] + +@[fun_prop] +lemma measurable_altReward [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : + Measurable (altReward alg n) := by unfold altReward; fun_prop + +variable [DecidableEq α] + +lemma hasLaw_altArm_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : + HasLaw (altArm alg 0) alg.p0 (arrayMeasure ν) where + map_eq := by + calc (arrayMeasure ν).map (fun ω ↦ initAlgFunction alg (ω.1 0)) + _ = ((arrayMeasure ν).fst.map (Function.eval 0)).map (initAlgFunction alg) := by + rw [Measure.fst, Measure.map_map (by fun_prop) (by fun_prop), + Measure.map_map (by fun_prop) (by fun_prop)] + rfl + _ = (volume : Measure I).map (initAlgFunction alg) := by + simp only [arrayMeasure, Measure.fst_prod] + rw [(measurePreserving_eval_infinitePi (fun _ ↦ volume) 0).map_eq] + _ = alg.p0 := initAlgFunction_map alg + +variable [StandardBorelSpace R] [Nonempty R] + +lemma hasCondDistrib_altReward_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : + HasCondDistrib (altReward alg 0) (altArm alg 0) (stationaryEnv ν).ν0 (arrayMeasure ν) where + condDistrib_eq := by + simp only [stationaryEnv_ν0, (hasLaw_altArm_zero alg ν).map_eq, altReward_zero] + sorry + +lemma hasCondDistrib_altStep' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] + (n : ℕ) : + HasCondDistrib (altHist alg · (n + 1) ⟨n + 1, by simp⟩) (altHist alg · n) + (Bandit.stepKernel alg ν n) (arrayMeasure ν) where + condDistrib_eq := by + simp only [Bandit.stepKernel, stepKernel, stationaryEnv_feedback] + sorry + +lemma hasCondDistrib_altStep (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] + (n : ℕ) : + HasCondDistrib (fun ω ↦ (altArm alg (n + 1) ω, altReward alg (n + 1) ω)) + (fun ω (i : Iic n) ↦ (altArm alg i ω, altReward alg i ω)) + (Bandit.stepKernel alg ν n) (arrayMeasure ν) := by + convert hasCondDistrib_altStep' alg ν n with ω i + · simp only [altArm] + rw [altHist_eq _ _ n] + · simp only [altReward] + rw [altHist_eq _ _ n] + +lemma hasCondDistrib_altArm (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + HasCondDistrib (altArm alg (n + 1)) (fun ω (i : Iic n) ↦ (altArm alg i ω, altReward alg i ω)) + (alg.policy n) (arrayMeasure ν) := by + convert HasCondDistrib.fst (hasCondDistrib_altStep alg ν n) + simp + +lemma hasCondDistrib_altReward (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] + (n : ℕ) : + HasCondDistrib (altReward alg (n + 1)) + (fun ω ↦ (fun (i : Iic n) ↦ (altArm alg i ω, altReward alg i ω), altArm alg (n + 1) ω)) + ((stationaryEnv ν).feedback n) (arrayMeasure ν) := by + simp only [stationaryEnv_feedback] + sorry + +lemma isAlgEnvInteraction_arrayMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : + IsAlgEnvInteraction (altArm alg) (altReward alg) alg (stationaryEnv ν) (arrayMeasure ν) where + hasLaw_action_zero := hasLaw_altArm_zero alg ν + hasCondDistrib_reward_zero := hasCondDistrib_altReward_zero alg ν + hasCondDistrib_action := hasCondDistrib_altArm alg ν + hasCondDistrib_reward := hasCondDistrib_altReward alg ν + +end ArrayModel + end MeasureSpace end Bandits diff --git a/LeanBandits/ForMathlib/HasCondDistrib.lean b/LeanBandits/ForMathlib/HasCondDistrib.lean index e26dacfb..e6d9f06d 100644 --- a/LeanBandits/ForMathlib/HasCondDistrib.lean +++ b/LeanBandits/ForMathlib/HasCondDistrib.lean @@ -13,14 +13,57 @@ open MeasureTheory namespace ProbabilityTheory -variable {α β Ω : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} +variable {α β γ Ω Ω' : Type*} + {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} {mΩ : MeasurableSpace Ω} [StandardBorelSpace Ω] [Nonempty Ω] + {mΩ' : MeasurableSpace Ω'} [StandardBorelSpace Ω'] [Nonempty Ω'] {μ : Measure α} {X : α → β} {Y : α → Ω} {κ : Kernel β Ω} structure HasCondDistrib (Y : α → Ω) (X : α → β) (κ : Kernel β Ω) (μ : Measure α) [IsFiniteMeasure μ] : Prop where - aemeasurable_fst : AEMeasurable Y μ - aemeasurable_snd : AEMeasurable X μ + aemeasurable_fst : AEMeasurable Y μ := by fun_prop + aemeasurable_snd : AEMeasurable X μ := by fun_prop condDistrib_eq : condDistrib Y X μ =ᵐ[μ.map X] κ +lemma hasCondDistrib_fst_prod {Y : α → Ω} {X : α → β} + {κ : Kernel β Ω} + {μ : Measure α} [IsFiniteMeasure μ] {ν : Measure γ} [IsProbabilityMeasure ν] + (h : HasCondDistrib Y X κ μ) : + HasCondDistrib (fun ω ↦ Y ω.1) (fun ω ↦ X ω.1) κ (μ.prod ν) where + aemeasurable_fst := by have := h.aemeasurable_fst; fun_prop + aemeasurable_snd := by have := h.aemeasurable_snd; fun_prop + condDistrib_eq := by + have : ((μ.prod ν).map (fun ω ↦ X ω.1)) = μ.map X := by + conv_rhs => rw [← Measure.fst_prod (μ := μ) (ν := ν), Measure.fst] + rw [AEMeasurable.map_map_of_aemeasurable _ (by fun_prop)] + · rfl + · have := h.aemeasurable_snd + simpa + rw [this] + exact (condDistrib_fst_prod X h.aemeasurable_fst ν).trans h.condDistrib_eq + +lemma HasCondDistrib.comp [IsFiniteMeasure μ] + (h : HasCondDistrib Y X κ μ) {f : Ω → Ω'} (hf : Measurable f) : + HasCondDistrib (fun ω ↦ f (Y ω)) X (κ.map f) μ where + aemeasurable_fst := by have := h.aemeasurable_fst; fun_prop + aemeasurable_snd := by have := h.aemeasurable_snd; fun_prop + condDistrib_eq := by + have h_comp := condDistrib_comp X (Y := Y) (f := f) (mβ := mβ) h.aemeasurable_fst hf + refine h_comp.trans ?_ + have h' := h.condDistrib_eq + filter_upwards [h'] with ω hω + rw [Kernel.map_apply _ hf, hω, Kernel.map_apply _ hf] + +lemma HasCondDistrib.fst {Y : α → Ω × Ω'} {κ : Kernel β (Ω × Ω')} [IsFiniteMeasure μ] + (h : HasCondDistrib Y X κ μ) : + HasCondDistrib (fun ω ↦ (Y ω).1) X κ.fst μ := by + rw [Kernel.fst_eq] + exact HasCondDistrib.comp h measurable_fst + +lemma HasCondDistrib.snd {Y : α → Ω × Ω'} {κ : Kernel β (Ω × Ω')} [IsFiniteMeasure μ] + (h : HasCondDistrib Y X κ μ) : + HasCondDistrib (fun ω ↦ (Y ω).2) X κ.snd μ := by + rw [Kernel.snd_eq] + exact HasCondDistrib.comp h measurable_snd + end ProbabilityTheory diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index e017dd70..1558fffb 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -13,6 +13,44 @@ import Mathlib.Probability.IdentDistribIndep open MeasureTheory ProbabilityTheory Finset Learning open scoped ENNReal NNReal +section Aux -- todo: move + +namespace ProbabilityTheory + +variable {α β γ δ γ' δ' : Type*} + {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} + {mδ : MeasurableSpace δ} {mγ' : MeasurableSpace γ'} {mδ' : MeasurableSpace δ'} + [StandardBorelSpace α] + [StandardBorelSpace δ'] [Nonempty δ'] [StandardBorelSpace γ'] [Nonempty γ'] + {μ : Measure α} [IsFiniteMeasure μ] + {X : α → β} {hX : Measurable X} {Y : α → γ} {Z : α → δ} {Y' : α → γ'} {Z' : α → δ'} + +lemma CondIndepFun.of_measurable (h_indep : Y ⟂ᵢ[X, hX; μ] Z) + (hY_meas : Measurable[mγ.comap Y] Y') (hZ_meas : Measurable[mδ.comap Z] Z') : + Y' ⟂ᵢ[X, hX; μ] Z' := by + obtain ⟨φ, hφ_meas, h_eqY⟩ : ∃ φ, Measurable φ ∧ Y' = φ ∘ Y := hY_meas.exists_eq_measurable_comp + obtain ⟨ψ, hψ_meas, h_eqZ⟩ : ∃ ψ, Measurable ψ ∧ Z' = ψ ∘ Z := hZ_meas.exists_eq_measurable_comp + rw [h_eqY, h_eqZ] + exact h_indep.comp hφ_meas hψ_meas + +lemma CondIndepFun.of_measurable_left + (h_indep : Y ⟂ᵢ[X, hX; μ] Z) (hY_meas : Measurable[mγ.comap Y] Y') : + Y' ⟂ᵢ[X, hX; μ] Z := by + obtain ⟨φ, hφ_meas, h_eqY⟩ : ∃ φ, Measurable φ ∧ Y' = φ ∘ Y := hY_meas.exists_eq_measurable_comp + rw [h_eqY] + exact h_indep.comp hφ_meas measurable_id + +lemma CondIndepFun.of_measurable_right + (h_indep : Y ⟂ᵢ[X, hX; μ] Z) (hZ_meas : Measurable[mδ.comap Z] Z') : + Y ⟂ᵢ[X, hX; μ] Z' := by + obtain ⟨ψ, hψ_meas, h_eqZ⟩ : ∃ ψ, Measurable ψ ∧ Z' = ψ ∘ Z := hZ_meas.exists_eq_measurable_comp + rw [h_eqZ] + exact h_indep.comp measurable_id hψ_meas + +end ProbabilityTheory + +end Aux + namespace Bandits variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [MeasurableSingletonClass α] @@ -82,41 +120,29 @@ lemma reward_cond_arm [StandardBorelSpace α] [Nonempty α] [Countable α] (a : exact h_eq.symm lemma condIndepFun_reward_stepsUntil_arm' [StandardBorelSpace α] [Countable α] [Nonempty α] - (a : α) (m n : ℕ) (hm : m ≠ 0) : + (a : α) (m n : ℕ) : reward n ⟂ᵢ[arm n, measurable_arm n; 𝔓t] {ω | stepsUntil a m ω = ↑n}.indicator (fun _ ↦ 1) := by - -- 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 + -- the indicator of `stepsUntil ... = n` is a function of `hist (n-1)` and `arm n`. + -- It thus suffices to use the independence of `reward n` and `hist (n-1)` conditionally -- on `arm n`. 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) (μ := 𝔓t) - (by fun_prop) (by fun_prop) - have : {ω : ℕ → α × ℝ | action 0 ω = a}.indicator (fun x ↦ 1) - = {b | b = a}.indicator (fun _ ↦ 1) ∘ action 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 : reward n ⟂ᵢ[arm n, measurable_arm n; 𝔓t] fun ω ↦ (hist (n - 1) ω, arm n ω) := - condIndepFun_reward_hist_arm_arm' (alg := alg) (ν := ν) n (by grind) - obtain ⟨φ, hφ_meas, h_eq⟩ : ∃ φ : ((Iic (n - 1) → α × ℝ) × α) → ℕ, Measurable φ ∧ - {ω | stepsUntil a m ω = ↑n}.indicator (fun _ ↦ 1) = φ ∘ (fun ω ↦ (hist (n - 1) ω, arm n ω)) := - (measurable_comap_indicator_stepsUntil_eq a m n).exists_eq_measurable_comp - rw [h_eq] - exact h_indep.comp measurable_id hφ_meas + · have h_indep : reward 0 ⟂ᵢ[arm 0, measurable_arm 0; 𝔓t] arm 0 := + condIndepFun_self_right (by fun_prop) (by fun_prop) + simp only [hn, CharP.cast_eq_zero] + refine h_indep.of_measurable_right (hX := measurable_arm 0) ?_ + exact measurable_comap_indicator_stepsUntil_eq_zero a m + · have h_indep : reward n ⟂ᵢ[arm n, measurable_arm n; 𝔓t] fun ω ↦ (hist (n - 1) ω, arm n ω) := + condIndepFun_reward_hist_arm_arm' (alg := alg) (ν := ν) n (by grind) + refine h_indep.of_measurable_right (hX := measurable_arm n) ?_ + exact measurable_comap_indicator_stepsUntil_eq a m n lemma condIndepFun_reward_stepsUntil_arm [StandardBorelSpace α] [Countable α] [Nonempty α] - (a : α) (m n : ℕ) (hm : m ≠ 0) : + (a : α) (m n : ℕ) : CondIndepFun (mα.comap (fun ω ↦ arm n ω.1)) ((measurable_arm n).comp measurable_fst).comap_le (fun ω ↦ reward n ω.1) ({ω | stepsUntil a m ω.1 = ↑n}.indicator (fun _ ↦ 1)) 𝔓 := condIndepFun_fst_prod (ν := Bandit.streamMeasure ν) (measurable_indicator_stepsUntil_eq a m n) (by fun_prop) (by fun_prop) - (condIndepFun_reward_stepsUntil_arm' a m n hm) + (condIndepFun_reward_stepsUntil_arm' a m n) lemma reward_cond_stepsUntil [StandardBorelSpace α] [Countable α] [Nonempty α] (a : α) (m n : ℕ) (hm : m ≠ 0) (hμn : 𝔓 ((fun ω ↦ stepsUntil a m ω.1) ⁻¹' {↑n}) ≠ 0) : @@ -149,7 +175,7 @@ lemma reward_cond_stepsUntil [StandardBorelSpace α] [Countable α] [Nonempty α rw [and_comm] _ = 𝓛[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 + · exact condIndepFun_reward_stepsUntil_arm a m n · refine measurable_one.indicator ?_ exact measurableSet_eq_fun (by fun_prop) (by fun_prop) · fun_prop @@ -159,7 +185,9 @@ lemma reward_cond_stepsUntil [StandardBorelSpace α] [Countable α] [Nonempty α simp [Set.indicator_apply] _ = ν a := reward_cond_arm a n hμa -lemma condDistrib_rewardByCount_stepsUntil [Countable α] [StandardBorelSpace α] [Nonempty α] +/-- The conditional distribution of the reward received at the `m`-th pull of arm `a` +given the time at which number of pulls is `m` is the constant kernel with value `ν a`. -/ +theorem condDistrib_rewardByCount_stepsUntil [Countable α] [StandardBorelSpace α] [Nonempty α] (a : α) (m : ℕ) (hm : m ≠ 0) : condDistrib (rewardByCount a m) (fun ω ↦ stepsUntil a m ω.1) 𝔓 =ᵐ[(𝔓).map (fun ω ↦ stepsUntil a m ω.1)] Kernel.const _ (ν a) := by diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index 418bc5cb..335d4729 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -3,6 +3,7 @@ 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, Paulo Rauber -/ +import LeanBandits.ForMathlib.HasCondDistrib import LeanBandits.ForMathlib.Measurable import LeanBandits.ForMathlib.Traj import Mathlib.Probability.HasLaw @@ -17,7 +18,7 @@ open scoped ENNReal NNReal namespace Learning -variable {α R : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} +variable {α R Ω : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} {mΩ : MeasurableSpace Ω} /-- A stochastic, sequential algorithm. -/ structure Algorithm (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] where @@ -56,6 +57,124 @@ lemma fst_stepKernel (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : (stepKernel alg env n).fst = alg.policy n := by rw [stepKernel, Kernel.fst_compProd] +section IsAlgEnvInteraction + +variable {A : ℕ → Ω → α} {R' : ℕ → Ω → R} {alg : Algorithm α R} {env : Environment α R} + {P : Measure Ω} [IsFiniteMeasure P] + +structure IsAlgEnvInteraction + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (alg : Algorithm α R) (env : Environment α R) + (P : Measure Ω) [IsFiniteMeasure P] : Prop where + measurable_A n : Measurable (A n) := by fun_prop + measurable_R n : Measurable (R' n) := by fun_prop + hasLaw_action_zero : HasLaw (fun ω ↦ (A 0 ω)) alg.p0 P + hasCondDistrib_reward_zero : HasCondDistrib (R' 0) (A 0) env.ν0 P + hasCondDistrib_action n : + HasCondDistrib (A (n + 1)) (fun ω ↦ fun (i : Iic n) ↦ (A i ω, R' i ω)) (alg.policy n) P + hasCondDistrib_reward n : + HasCondDistrib (R' (n + 1)) (fun ω ↦ (fun (i : Iic n) ↦ (A i ω, R' i ω), A (n + 1) ω)) + (env.feedback n) P + +def IsAlgEnvInteraction.step (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (n : ℕ) (ω : Ω) : α × R := + (A n ω, R' n ω) + +@[fun_prop] +lemma IsAlgEnvInteraction.measurable_step (n : ℕ) (hA : Measurable (A n)) + (hR' : Measurable (R' n)) : + Measurable (IsAlgEnvInteraction.step A R' n) := by + unfold IsAlgEnvInteraction.step + fun_prop + +def IsAlgEnvInteraction.hist (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (n : ℕ) (ω : Ω) : Iic n → α × R := + fun i ↦ (A i ω, R' i ω) + +@[fun_prop] +lemma IsAlgEnvInteraction.measurable_hist (hA : ∀ n, Measurable (A n)) + (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : + Measurable (IsAlgEnvInteraction.hist A R' n) := by + unfold IsAlgEnvInteraction.hist + fun_prop + +def IsAlgEnvInteraction.filtration (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) : + Filtration ℕ mΩ where + seq i := MeasurableSpace.comap (hist A R' i) inferInstance + mono' i j hij := by + simp only + rw [← measurable_iff_comap_le] + have : hist A R' i = (fun h k ↦ h ⟨k.1, by grind⟩) ∘ hist A R' j := rfl + rw [this] + exact measurable_comp_comap _ (by fun_prop) + le' i := by + rw [← measurable_iff_comap_le] + exact measurable_hist hA hR' i + +lemma IsAlgEnvInteraction.measurable_action_filtration + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : + Measurable[IsAlgEnvInteraction.filtration hA hR' n] (A n) := by + have : A n = (fun h ↦ (h ⟨n, by simp⟩).1) ∘ (hist A R' n) := by + ext ω : 1 + simp [IsAlgEnvInteraction.hist] + rw [this] + exact measurable_comp_comap _ (by fun_prop) + +/-- Filtration generated by the history at time `n-1` together with the action at time `n`. -/ +def IsAlgEnvInteraction.filtrationAction + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) : + Filtration ℕ mΩ where + seq n := if n = 0 then MeasurableSpace.comap (A 0) inferInstance + else IsAlgEnvInteraction.filtration hA hR' (n - 1) ⊔ MeasurableSpace.comap (A n) inferInstance + mono' n m hnm := by + simp only + by_cases hn : n = 0 + · by_cases hm : m = 0 + · simp [hn, hm] + · simp only [hn, ↓reduceIte, hm] + refine le_sup_of_le_left ?_ + rw [← measurable_iff_comap_le] + suffices Measurable[IsAlgEnvInteraction.filtration hA hR' 0] (A 0) from + this.mono ((IsAlgEnvInteraction.filtration hA hR').mono zero_le') le_rfl + exact measurable_action_filtration hA hR' 0 + have hm : m ≠ 0 := by grind + simp only [hn, hm, ↓reduceIte] + have hnm' : n - 1 ≤ m - 1 := by grind + simp only [sup_le_iff] + constructor + · refine le_sup_of_le_left ?_ + exact (IsAlgEnvInteraction.filtration hA hR').mono hnm' + · rcases eq_or_lt_of_le hnm with rfl | hlt + · exact le_sup_of_le_right le_rfl + refine le_sup_of_le_left ?_ + rw [← measurable_iff_comap_le] + have h_le : n ≤ m - 1 := by grind + suffices Measurable[IsAlgEnvInteraction.filtration hA hR' n] (A n) from + this.mono ((IsAlgEnvInteraction.filtration hA hR').mono h_le) le_rfl + exact measurable_action_filtration hA hR' n + le' n := by + by_cases hn : n = 0 + · simp only [hn, ↓reduceIte] + rw [← measurable_iff_comap_le] + fun_prop + simp only [hn, ↓reduceIte, sup_le_iff] + constructor + · exact (IsAlgEnvInteraction.filtration hA hR').le _ + · rw [← measurable_iff_comap_le] + fun_prop + +lemma IsAlgEnvInteraction.filtrationAction_zero_eq_comap + {hA : ∀ n, Measurable (A n)} {hR' : ∀ n, Measurable (R' n)} : + filtrationAction hA hR' 0 = MeasurableSpace.comap (A 0) inferInstance := by + simp [filtrationAction] + +lemma IsAlgEnvInteraction.filtrationAction_eq_comap + {hA : ∀ n, Measurable (A n)} {hR' : ∀ n, Measurable (R' n)} (n : ℕ) (hn : n ≠ 0) : + filtrationAction hA hR' n = + MeasurableSpace.comap (fun ω ↦ (hist A R' (n - 1) ω, A n ω)) inferInstance := by + simp only [filtrationAction, filtration, ← MeasurableSpace.comap_prodMk, hn, ↓reduceIte] + rfl + +end IsAlgEnvInteraction + /-- Kernel sending a partial trajectory of the bandit interaction `Iic n → α × ℝ` to a measure on `ℕ → α × ℝ`, supported on full trajectories that start with the partial one. -/ noncomputable def traj (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : @@ -272,33 +391,6 @@ end FiltrationAction section Laws -lemma condDistrib_step [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : - condDistrib (step (n + 1)) (hist n) (trajMeasure alg env) - =ᵐ[(trajMeasure alg env).map (hist n)] stepKernel alg env n := - Kernel.condDistrib_trajMeasure - -lemma condDistrib_action [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : - condDistrib (action (n + 1)) (hist n) (trajMeasure alg env) - =ᵐ[(trajMeasure alg env).map (hist n)] alg.policy n := by - rw [← fst_comp_step] - refine (condDistrib_comp _ (by fun_prop) (by fun_prop)).trans ?_ - filter_upwards [condDistrib_step alg env n] with h h_eq - rw [Kernel.map_apply _ (by fun_prop), h_eq, ← Kernel.map_apply _ (by fun_prop), ← Kernel.fst_eq, - fst_stepKernel] - -lemma condDistrib_reward [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : - condDistrib (reward (n + 1)) (fun ω ↦ (hist n ω, action (n + 1) ω)) (trajMeasure alg env) - =ᵐ[(trajMeasure alg env).map (fun ω ↦ (hist n ω, action (n + 1) ω))] env.feedback n := by - have h_step := condDistrib_step alg env n - have h_action := condDistrib_action alg env n - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h_step h_action ⊢ - rw [h_action, ← Measure.compProd_assoc, ← stepKernel, ← h_step, - Measure.map_map (by fun_prop) (by fun_prop)] - rfl - lemma hasLaw_step_zero (alg : Algorithm α R) (env : Environment α R) : HasLaw (step 0) (alg.p0 ⊗ₘ env.ν0) (trajMeasure alg env) where aemeasurable := Measurable.aemeasurable (by fun_prop) @@ -317,14 +409,63 @@ lemma hasLaw_action_zero (alg : Algorithm α R) (env : Environment α R) : rw [← fst_comp_step, ← Measure.map_map (by fun_prop) (by fun_prop), (hasLaw_step_zero alg env).map_eq, ← Measure.fst, Measure.fst_compProd] -lemma condDistrib_reward_zero [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) (env : Environment α R) : +variable [StandardBorelSpace R] [Nonempty R] + +lemma condDistrib_reward_zero (alg : Algorithm α R) (env : Environment α R) : condDistrib (reward 0) (action 0) (trajMeasure alg env) =ᵐ[(trajMeasure alg env).map (action 0)] env.ν0 := by have h_step := (hasLaw_step_zero alg env).map_eq have h_action := (hasLaw_action_zero alg env).map_eq rwa [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop), h_action] +variable [StandardBorelSpace α] [Nonempty α] + +lemma condDistrib_step (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : + condDistrib (step (n + 1)) (hist n) (trajMeasure alg env) + =ᵐ[(trajMeasure alg env).map (hist n)] stepKernel alg env n := + Kernel.condDistrib_trajMeasure + +lemma condDistrib_action (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : + condDistrib (action (n + 1)) (hist n) (trajMeasure alg env) + =ᵐ[(trajMeasure alg env).map (hist n)] alg.policy n := by + rw [← fst_comp_step] + refine (condDistrib_comp _ (by fun_prop) (by fun_prop)).trans ?_ + filter_upwards [condDistrib_step alg env n] with h h_eq + rw [Kernel.map_apply _ (by fun_prop), h_eq, ← Kernel.map_apply _ (by fun_prop), ← Kernel.fst_eq, + fst_stepKernel] + +lemma condDistrib_reward (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : + condDistrib (reward (n + 1)) (fun ω ↦ (hist n ω, action (n + 1) ω)) (trajMeasure alg env) + =ᵐ[(trajMeasure alg env).map (fun ω ↦ (hist n ω, action (n + 1) ω))] env.feedback n := by + have h_step := condDistrib_step alg env n + have h_action := condDistrib_action alg env n + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h_step h_action ⊢ + rw [h_action, ← Measure.compProd_assoc, ← stepKernel, ← h_step, + Measure.map_map (by fun_prop) (by fun_prop)] + rfl + +lemma isAlgEnvInteraction_trajMeasure (alg : Algorithm α R) (env : Environment α R) : + IsAlgEnvInteraction action reward alg env (trajMeasure alg env) where + hasLaw_action_zero := hasLaw_action_zero alg env + hasCondDistrib_reward_zero := ⟨by fun_prop, by fun_prop, condDistrib_reward_zero alg env⟩ + hasCondDistrib_action n := ⟨by fun_prop, by fun_prop, condDistrib_action alg env n⟩ + hasCondDistrib_reward n := ⟨by fun_prop, by fun_prop, condDistrib_reward alg env n⟩ + end Laws +section ModelEquivalence + +variable {Ω Ω' : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + {alg : Algorithm α R} {env : Environment α R} + {P : Measure Ω} [IsFiniteMeasure P] {P' : Measure Ω'} [IsFiniteMeasure P'] + {A₁ : ℕ → Ω → α} {R₁ : ℕ → Ω → R} {A₂ : ℕ → Ω' → α} {R₂ : ℕ → Ω' → R} + +theorem isAlgEnvInteraction_unique (h1 : IsAlgEnvInteraction A₁ R₁ alg env P) + (h2 : IsAlgEnvInteraction A₂ R₂ alg env P') : + P.map (fun ω n ↦ (A₁ n ω, R₁ n ω)) = P'.map (fun ω n ↦ (A₂ n ω, R₂ n ω)) := by + sorry + +end ModelEquivalence + end Learning diff --git a/LeanBandits/SequentialLearning/Draft.lean b/LeanBandits/SequentialLearning/Draft.lean new file mode 100644 index 00000000..5d3a5c71 --- /dev/null +++ b/LeanBandits/SequentialLearning/Draft.lean @@ -0,0 +1,772 @@ +/- +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, Paulo Rauber +-/ +import LeanBandits.SequentialLearning.Algorithm +import Mathlib.Order.CompletePartialOrder +import Mathlib.Probability.Martingale.BorelCantelli + +/-! +# Bookkeeping definitions for finite action space sequential learning problems + +If the number of actions is finite, it makes sense to define the number of times each action was +chosen, the time at which an action was chosen for the nth time, the value of the reward at that +time, the sum of rewards obtained for each action, the empirical mean reward for each action, etc. + +For each definition that take as arguments a time `t : ℕ`, a history `h : ℕ → α × R`, and possibly +other parameters, we put the time and history at the end in this order, so that the definition can +be seen as a stochastic process indexed by time `t` on the measurable space `ℕ → α × R`. + +-/ + +open MeasureTheory Finset Learning + +namespace LearningDraft + +variable {α R Ω : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} {mΩ : MeasurableSpace Ω} + [DecidableEq α] + {alg : Algorithm α R} {env : Environment α R} + {P : Measure Ω} [IsProbabilityMeasure P] + {A : ℕ → Ω → α} {R' : ℕ → Ω → R} + {a : α} {m n t : ℕ} {ω : Ω} + +section PullCount + +/-- Number of times action `a` was chosen up to time `t` (excluding `t`). -/ +noncomputable +def pullCount (A : ℕ → Ω → α) (a : α) (t : ℕ) (ω : Ω) : ℕ := + #(filter (fun s ↦ A s ω = a) (range t)) + +/-- Number of pulls of arm `a` up to (and including) time `n`. +This is the number of entries in `h` in which the arm is `a`. -/ +noncomputable +def pullCount' (n : ℕ) (h : Iic n → α × R) (a : α) := #{s | (h s).1 = a} + +@[simp] +lemma pullCount_zero (a : α) : pullCount A a 0 = 0 := by ext; simp [pullCount] + +lemma pullCount_zero_apply (a : α) (ω : Ω) : pullCount A a 0 ω = 0 := by simp + +lemma pullCount_one : pullCount A a 1 ω = if A 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] + +lemma monotone_pullCount (a : α) (ω : Ω) : Monotone (pullCount A a · ω) := + fun _ _ _ ↦ card_le_card (filter_subset_filter _ (by simpa)) + +@[mono, gcongr] +lemma pullCount_mono (a : α) {n m : ℕ} (hnm : n ≤ m) (ω : Ω) : + pullCount A a n ω ≤ pullCount A a m ω := + monotone_pullCount a ω hnm + +lemma pullCount_action_eq_pullCount_add_one (t : ℕ) (ω : Ω) : + pullCount A (A t ω) (t + 1) ω = pullCount A (A t ω) t ω + 1 := by + simp [pullCount, range_add_one, filter_insert] + +lemma pullCount_eq_pullCount_of_action_ne (ha : A t ω ≠ a) : + pullCount A a (t + 1) ω = pullCount A a t ω := by + simp [pullCount, range_add_one, filter_insert, ha] + +lemma pullCount_add_one : + pullCount A a (t + 1) ω = pullCount A a t ω + if A t ω = a then 1 else 0 := by + split_ifs with h + · rw [← h, pullCount_action_eq_pullCount_add_one] + · rw [pullCount_eq_pullCount_of_action_ne h, add_zero] + +lemma pullCount_eq_sum (a : α) (t : ℕ) (ω : Ω) : + pullCount A a t ω = ∑ s ∈ range t, if A s ω = a then 1 else 0 := by simp [pullCount] + +lemma pullCount'_eq_sum (n : ℕ) (h : Iic n → α × R) (a : α) : + pullCount' n h a = ∑ s : Iic n, if (h s).1 = a then 1 else 0 := by simp [pullCount'] + +lemma pullCount_add_one_eq_pullCount' {n : ℕ} {ω : Ω} : + pullCount A a (n + 1) ω = pullCount' n (fun i ↦ (A i ω, R' i ω)) a := by + rw [pullCount_eq_sum, pullCount'_eq_sum] + rw [Finset.sum_coe_sort (f := fun s ↦ if A s ω = a then 1 else 0) (Iic n)] + congr with m + simp only [mem_range, mem_Iic] + grind + +lemma pullCount_eq_pullCount' {n : ℕ} {ω : Ω} (hn : n ≠ 0) : + pullCount A a n ω = pullCount' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by + cases n with + | zero => exact absurd rfl hn + | succ n => + rw [pullCount_add_one_eq_pullCount' (R' := R')] + have : n + 1 - 1 = n := by simp + exact this ▸ rfl + +lemma pullCount_le (a : α) (t : ℕ) (ω : Ω) : pullCount A a t ω ≤ t := + (card_filter_le _ _).trans_eq (by simp) + +lemma pullCount_congr {ω' : Ω} (h_eq : ∀ i ≤ n, A i ω = A i ω') : + pullCount A a (n + 1) ω = pullCount A a (n + 1) ω' := by + unfold pullCount + congr 1 with s + simp only [mem_filter, mem_range, and_congr_right_iff] + intro hs + rw [Nat.lt_add_one_iff] at hs + rw [h_eq s hs] + +lemma pullCount_lt_of_forall_ne (h_lt : ∀ s, pullCount A a (s + 1) ω ≠ t) (ht : t ≠ 0) : + pullCount A a n ω < t := by + induction n with + | zero => simpa using ht.bot_lt + | succ n hn => + specialize h_lt n + rw [pullCount_add_one] at h_lt ⊢ + grind + +lemma exists_pullCount_eq_of_le (hnm : t ≤ pullCount A a (n + 1) ω) (ht : t ≠ 0) : + ∃ s, pullCount A a (s + 1) ω = t := by + by_contra! h_contra + refine lt_irrefl (pullCount A a (n + 1) ω) ?_ + refine lt_of_lt_of_le ?_ hnm + exact pullCount_lt_of_forall_ne h_contra ht + +section Measurability + +@[fun_prop] +lemma measurable_pullCount [MeasurableSingletonClass α] (hA : ∀ n, Measurable (A n)) + (a : α) (t : ℕ) : + Measurable (fun ω : Ω ↦ pullCount A a t ω) := by + simp_rw [pullCount_eq_sum] + have h_meas s : Measurable (fun ω : Ω ↦ if A s ω = a then 1 else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + fun_prop + +@[fun_prop] +lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : + Measurable (fun h : Iic n → α × R ↦ pullCount' n h a) := by + simp_rw [pullCount'_eq_sum] + have h_meas s : Measurable (fun (h : Iic n → α × R) ↦ if (h s).1 = a then 1 else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + fun_prop + +lemma adapted_pullCount_add_one' [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (n : ℕ) : + Measurable[IsAlgEnvInteraction.filtration hA hR' n] (pullCount A a (n + 1)) := by + have : pullCount A a (n + 1) = (fun h : Iic n → α × R ↦ pullCount' n h a) ∘ + (IsAlgEnvInteraction.hist A R' n) := by + ext + exact pullCount_add_one_eq_pullCount' + rw [IsAlgEnvInteraction.filtration, this] + exact measurable_comp_comap _ (measurable_pullCount' n a) + +lemma adapted_pullCount_add_one [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) : + Adapted (IsAlgEnvInteraction.filtration hA hR') (fun n ↦ pullCount A a (n + 1)) := + fun n ↦ Measurable.stronglyMeasurable <| adapted_pullCount_add_one' hA hR' a n + +lemma isPredictable_pullCount [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) : + IsPredictable (IsAlgEnvInteraction.filtration hA hR') (pullCount A a) := by + rw [isPredictable_iff_measurable_add_one] + refine ⟨?_, fun n ↦ (adapted_pullCount_add_one hA hR' a n).measurable⟩ + simp only [pullCount_zero] + fun_prop + +end Measurability + +end PullCount + +section StepsUntil + +-- TODO: replace this by leastGE, once leastGE is generalized +/-- Number of steps until action `a` was pulled exactly `m` times. -/ +noncomputable +def stepsUntil (A : ℕ → Ω → α) (a : α) (m : ℕ) (ω : Ω) : ℕ∞ := + sInf ((↑) '' {s | pullCount A a (s + 1) ω = m}) + +lemma stepsUntil_eq_top_iff : stepsUntil A a m ω = ⊤ ↔ ∀ s, pullCount A a (s + 1) ω ≠ m := by + simp [stepsUntil, sInf_eq_top] + +lemma stepsUntil_ne_top (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : stepsUntil A a m ω ≠ ⊤ := by + simpa [stepsUntil_eq_top_iff] + +lemma exists_pullCount_eq (h' : stepsUntil A a m ω ≠ ⊤) : + ∃ s, pullCount A a (s + 1) ω = m := by + by_contra! h_contra + rw [← stepsUntil_eq_top_iff] at h_contra + simp [h_contra] at h' + +lemma stepsUntil_zero_of_ne (hka : A 0 ω ≠ a) : stepsUntil A a 0 ω = 0 := by + unfold stepsUntil + simp_rw [← bot_eq_zero, sInf_eq_bot, bot_eq_zero] + intro n hn + refine ⟨0, ?_, hn⟩ + simp only [Set.mem_image, Set.mem_setOf_eq, Nat.cast_eq_zero, exists_eq_right, zero_add] + rw [← zero_add 1, pullCount_eq_pullCount_of_action_ne hka] + simp + +lemma stepsUntil_zero_of_eq (hka : A 0 ω = a) : stepsUntil A a 0 ω = ⊤ := by + rw [stepsUntil_eq_top_iff] + suffices 0 < pullCount A a 1 ω by + intro n hn + refine lt_irrefl 0 ?_ + exact this.trans_le (le_trans (monotone_pullCount _ _ (by omega)) hn.le) + rw [← hka, ← zero_add 1, pullCount_action_eq_pullCount_add_one] + simp + +lemma stepsUntil_eq_dite (a : α) (m : ℕ) (ω : Ω) + [Decidable (∃ s, pullCount A a (s + 1) ω = m)] : + stepsUntil A a m ω = + if h : ∃ s, pullCount A a (s + 1) ω = m then (Nat.find h : ℕ∞) else ⊤ := by + unfold stepsUntil + split_ifs with h' + · refine le_antisymm ?_ ?_ + · refine sInf_le ?_ + simpa using Nat.find_spec h' + · simp only [le_sInf_iff, Set.mem_image, Set.mem_setOf_eq, forall_exists_index, and_imp, + forall_apply_eq_imp_iff₂, Nat.cast_le, Nat.find_le_iff] + exact fun n hn ↦ ⟨n, le_rfl, hn⟩ + · push_neg at h' + suffices {s | pullCount A a (s + 1) ω = m} = ∅ by simp [this] + ext s + simpa using (h' s) + +-- todo: this is in ℝ because of the limited def of leastGE +lemma stepsUntil_eq_leastGE (a : α) (hm : m ≠ 0) : + stepsUntil A a m = leastGE (fun n (ω : Ω) ↦ pullCount A a (n + 1) ω) m := by + classical + ext ω + rw [stepsUntil_eq_dite] + unfold leastGE hittingAfter + simp only [zero_le, Set.mem_Ici, Nat.cast_le, true_and, ENat.some_eq_coe] + have h_iff : (∃ s, pullCount A a (s + 1) ω = m) ↔ (∃ s, m ≤ pullCount A a (s + 1) ω) := by + refine ⟨fun ⟨s, hs⟩ ↦ ⟨s, hs.ge⟩, fun ⟨s, hs⟩ ↦ ?_⟩ + exact exists_pullCount_eq_of_le hs hm + by_cases h_exists : ∃ s, m ≤ pullCount A a (s + 1) ω + swap; · simp_rw [h_iff]; simp [h_exists] + rw [if_pos h_exists, dif_pos] + swap; · rwa [h_iff] + norm_cast + rw [Nat.find_eq_iff] + constructor + · apply le_antisymm + · by_contra! h_contra + obtain ⟨s, hs⟩ : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq_of_le h_contra.le hm + rw [← hs] at h_contra + refine h_contra.not_ge ?_ + gcongr + exact csInf_le (by simp) (by simp) + · exact Nat.sInf_mem (s := {j | m ≤ pullCount A a (j + 1) ω}) h_exists + · intro n hn h_contra + refine hn.not_ge ?_ + exact csInf_le (by simp) (by simp [h_contra]) + +lemma stepsUntil_pullCount_le (ω : Ω) (a : α) (t : ℕ) : + stepsUntil A a (pullCount A a (t + 1) ω) ω ≤ t := by + rw [stepsUntil] + exact csInf_le (OrderBot.bddBelow _) ⟨t, rfl, rfl⟩ + +lemma stepsUntil_pullCount_eq (ω : Ω) (t : ℕ) : + stepsUntil A (A t ω) (pullCount A (A t ω) (t + 1) ω) ω = t := by + apply le_antisymm (stepsUntil_pullCount_le ω (A t ω) t) + suffices ∀ t', pullCount A (A t ω) (t' + 1) ω = pullCount A (A t ω) t ω + 1 → t ≤ t' by + simpa [stepsUntil, pullCount_action_eq_pullCount_add_one] + exact fun t' h' ↦ Nat.le_of_lt_succ ((monotone_pullCount (A t ω) ω).reflect_lt + (h' ▸ lt_add_one _)) + +/-- If we pull action `a` at time 0, the first time at which it is pulled once is 0. -/ +lemma stepsUntil_one_of_eq (hka : A 0 ω = a) : stepsUntil A a 1 ω = 0 := by + classical + have h_pull : pullCount A a 1 ω = 1 := by simp [pullCount_one, hka] + have h_le := stepsUntil_pullCount_le (A := A) ω a 0 + simpa [h_pull] using h_le + +lemma stepsUntil_eq_zero_iff : + stepsUntil A a m ω = 0 ↔ (m = 0 ∧ A 0 ω ≠ a) ∨ (m = 1 ∧ A 0 ω = a) := by + classical + refine ⟨fun h' ↦ ?_, fun h' ↦ ?_⟩ + · have h_exists : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq (by simp [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 : A 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 action_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : + A (stepsUntil A a m ω).toNat ω = a := by + classical + simp only [stepsUntil_eq_dite, h_exists, ↓reduceDIte, ENat.toNat_coe] + have h_spec := Nat.find_spec h_exists + have h_spec' n := Nat.find_min h_exists (m := n) + by_cases h_zero : Nat.find h_exists = 0 + · simp only [h_zero, zero_add, not_lt_zero', IsEmpty.forall_iff, implies_true] at * + by_contra h_ne + rw [← zero_add 1, pullCount_eq_pullCount_of_action_ne h_ne] at h_spec + simp only [pullCount_zero] at h_spec + exact hm h_spec.symm + have h_pos : 0 < Nat.find h_exists := Nat.pos_of_ne_zero h_zero + by_contra h_ne + refine h_spec' (Nat.find h_exists - 1) ?_ ?_ + · simp [h_pos] + rw [Nat.sub_add_cancel (by omega)] + rwa [← pullCount_eq_pullCount_of_action_ne] + exact h_ne + +lemma action_eq_of_stepsUntil_eq_coe (hm : m ≠ 0) (h : stepsUntil A a m ω = n) : + A n ω = a := by + have : n = (stepsUntil A a m ω).toNat := by simp [h] + rw [this] + have h_exists : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq (by simp [h]) + exact action_stepsUntil hm h_exists + +lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : + pullCount A a (stepsUntil A a m ω + 1).toNat ω = m := by + classical + have h_eq := stepsUntil_eq_dite (A := A) a m ω + simp only [h_exists, ↓reduceDIte] at h_eq + have h' := Nat.find_spec h_exists + rw [h_eq] + rw [ENat.toNat_add (by simp) (by simp)] + simp only [ENat.toNat_coe, ENat.toNat_one] + exact h' + +lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : + pullCount A a (stepsUntil A a m ω).toNat ω = m - 1 := by + have h_action := action_eq_of_stepsUntil_eq_coe (A := A) (n := (stepsUntil A a m ω).toNat) + (a := a) (ω := ω) hm ?_ + swap; · symm; simpa [stepsUntil_eq_top_iff] + have h_add_one := pullCount_stepsUntil_add_one h_exists + nth_rw 1 [← h_action] at h_add_one + rw [ENat.toNat_add ?_ (by simp), ENat.toNat_one, pullCount_action_eq_pullCount_add_one] + at h_add_one + swap; · simpa [stepsUntil_eq_top_iff] + grind + +lemma pullCount_lt_of_le_stepsUntil (a : α) {n m : ℕ} (ω : Ω) + (h_exists : ∃ s, pullCount A a (s + 1) ω = m) (hn : n < stepsUntil A a m ω) : + pullCount A a (n + 1) ω < m := by + classical + have h_eq := stepsUntil_eq_dite (A := A) a m ω + simp only [h_exists, ↓reduceDIte] at h_eq + rw [← ENat.coe_toNat (stepsUntil_ne_top h_exists)] at hn + refine lt_of_le_of_ne ?_ ?_ + · calc pullCount A a (n + 1) ω + _ ≤ pullCount A a (stepsUntil A a m ω + 1).toNat ω := by + refine monotone_pullCount a ω ?_ + rw [ENat.toNat_add (stepsUntil_ne_top h_exists) (by simp)] + simp only [ENat.toNat_one, add_le_add_iff_right] + exact mod_cast hn.le + _ = m := pullCount_stepsUntil_add_one h_exists + · refine Nat.find_min h_exists (m := n) ?_ + suffices n < (stepsUntil A a m ω).toNat by + rwa [h_eq, ENat.toNat_coe] at this + exact mod_cast hn + +lemma pullCount_eq_of_stepsUntil_eq_coe {ω : Ω} (hm : m ≠ 0) + (h : stepsUntil A a m ω = n) : + pullCount A a n ω = m - 1 := by + have : n = (stepsUntil A a m ω).toNat := by simp [h] + rw [this, pullCount_stepsUntil hm] + exact exists_pullCount_eq (by simp [h]) + +lemma pullCount_add_one_eq_of_stepsUntil_eq_coe {ω : Ω} + (h : stepsUntil A a m ω = n) : + pullCount A a (n + 1) ω = m := by + have : n + 1 = (stepsUntil A a m ω + 1).toNat := by + rw [ENat.toNat_add (by simp [h]) (by simp)]; simp [h] + rw [this, pullCount_stepsUntil_add_one] + exact exists_pullCount_eq (by simp [h]) + +lemma stepsUntil_eq_iff {ω : Ω} (n : ℕ) : + stepsUntil A a m ω = n ↔ + pullCount A a (n + 1) ω = m ∧ (∀ k < n, pullCount A a (k + 1) ω < m) := by + refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ + · have h_exists : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq (by simp [h]) + refine ⟨pullCount_add_one_eq_of_stepsUntil_eq_coe h, fun k hk ↦ ?_⟩ + exact pullCount_lt_of_le_stepsUntil a ω h_exists (by rw [h]; exact mod_cast hk) + · classical + rw [stepsUntil_eq_dite a m ω, dif_pos ⟨n, h.1⟩] + simp only [Nat.cast_inj] + rw [Nat.find_eq_iff] + exact ⟨h.1, fun k hk ↦ (h.2 k hk).ne⟩ + +lemma stepsUntil_eq_iff' {ω : Ω} (hm : m ≠ 0) (n : ℕ) : + stepsUntil A a m ω = n ↔ A n ω = a ∧ pullCount A a n ω = m - 1 := by + by_cases hn : n = 0 + · simp [hn, stepsUntil_eq_zero_iff, hm] + grind + rw [stepsUntil_eq_iff n] + refine ⟨fun ⟨h1, h2⟩ ↦ ⟨?_, ?_⟩, fun ⟨h1, h2⟩ ↦ ⟨?_, fun k hk ↦ ?_⟩⟩ + · rw [pullCount_add_one] at h1 + specialize h2 (n - 1) (by lia) + grind + · rw [pullCount_add_one] at h1 + specialize h2 (n - 1) (by lia) + grind + · rw [pullCount_add_one, h1, h2] + grind + · rw [Nat.lt_iff_le_pred (by grind)] + rw [← h2] + refine monotone_pullCount a ω ?_ + grind + +lemma stepsUntil_eq_congr {ω' : Ω} (h_eq : ∀ i ≤ n, A i ω = A i ω') : + stepsUntil A a m ω = n ↔ stepsUntil A a m ω' = n := by + simp_rw [stepsUntil_eq_iff n] + congr! 1 + · rw [pullCount_congr h_eq] + · congr! 3 with k hk + rw [pullCount_congr] + grind + +section Measurability + +lemma isStoppingTime_stepsUntil [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (hm : m ≠ 0) : + IsStoppingTime (IsAlgEnvInteraction.filtration hA hR') (stepsUntil A a m) := by + rw [stepsUntil_eq_leastGE _ hm] + refine Adapted.isStoppingTime_leastGE _ fun n ↦ ?_ + suffices StronglyMeasurable[IsAlgEnvInteraction.filtration hA hR' n] (pullCount A a (n + 1)) by + fun_prop + exact adapted_pullCount_add_one hA hR' a n + +-- todo: get this from the stopping time property? +@[fun_prop] +lemma measurable_stepsUntil [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (a : α) (m : ℕ) : + Measurable (stepsUntil A a m) := by + classical + have h_union : {h' : Ω | ∃ s, pullCount A a (s + 1) h' = m} + = ⋃ s : ℕ, {h' | pullCount A a (s + 1) h' = m} := by ext; simp + have h_meas_set : MeasurableSet {h' : Ω | ∃ s, pullCount A a (s + 1) h' = m} := by + rw [h_union] + refine MeasurableSet.iUnion fun s ↦ (measurableSet_singleton _).preimage ?_ + exact measurable_pullCount hA a (s + 1) + --simp_rw [stepsUntil_eq_dite] + suffices Measurable fun k ↦ if h : k ∈ {k' | ∃ s, pullCount A a (s + 1) k' = m} + then (Nat.find h : ℕ∞) else ⊤ by + convert this with ω + rw [stepsUntil_eq_dite a m ω] + rfl + refine Measurable.dite (s := {k' : Ω | ∃ s, pullCount A a (s + 1) k' = m}) + (f := fun x ↦ (Nat.find x.2 : ℕ∞)) (g := fun _ ↦ ⊤) ?_ (by fun_prop) h_meas_set + refine Measurable.coe_nat_enat ?_ + refine measurable_find _ fun k ↦ ?_ + suffices MeasurableSet {x : Ω | pullCount A a (k + 1) x = m} by + have : Subtype.val '' {x : {k' : Ω | + ∃ s, pullCount A a (s + 1) k' = m} | pullCount A a (k + 1) (x : Ω) = m} + = {x : Ω | pullCount A a (k + 1) x = m} := by + ext x + simp only [Set.mem_setOf_eq, Set.coe_setOf, Set.mem_image, Subtype.exists, exists_and_left, + exists_prop, exists_eq_right_right, and_iff_left_iff_imp] + exact fun h ↦ ⟨_, h⟩ + refine (MeasurableEmbedding.subtype_coe h_meas_set).measurableSet_image.mp ?_ + rw [this] + exact (measurableSet_singleton _).preimage (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + +lemma measurable_stepsUntil' [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (a : α) (m : ℕ) : + Measurable (fun ω : Ω × (ℕ → α → R) ↦ stepsUntil A a m ω.1) := + (measurable_stepsUntil hA a m).comp measurable_fst + +lemma measurable_comap_indicator_stepsUntil_eq [MeasurableSingletonClass α] [Nonempty R] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : + Measurable[MeasurableSpace.comap + (fun ω : Ω ↦ (IsAlgEnvInteraction.hist A R' (n-1) ω, A n ω)) inferInstance] + ({ω | stepsUntil A a m ω = ↑n}.indicator fun _ ↦ 1) := by + by_cases hm : m = 0 + · simp only [hm] + by_cases hn : n = 0 + · simp only [hn, CharP.cast_eq_zero, stepsUntil_eq_zero_iff, ne_eq, true_and, zero_ne_one, + false_and, or_false] + refine Measurable.indicator measurable_const ?_ + refine (measurableSet_singleton _).compl.preimage ?_ + rw [measurable_iff_comap_le] + rw [Prod.instMeasurableSpace, MeasurableSpace.comap_prodMk] + exact le_sup_of_le_right le_rfl + · have : {ω | stepsUntil A a 0 ω = n} = ∅ := by + ext ω + by_cases ha : A 0 ω = a + · simp [stepsUntil_zero_of_eq ha] + · simp only [Set.mem_setOf_eq, stepsUntil_zero_of_ne ha, Set.mem_empty_iff_false, + iff_false] + norm_cast + exact Ne.symm hn + simp [this] + simp_rw [stepsUntil_eq_iff' hm] + refine Measurable.indicator measurable_const ?_ + refine ((measurableSet_singleton _).preimage ?_).inter ((measurableSet_singleton _).preimage ?_) + · rw [measurable_iff_comap_le] + rw [Prod.instMeasurableSpace, MeasurableSpace.comap_prodMk] + exact le_sup_of_le_right le_rfl + · rw [measurable_iff_comap_le] + rw [Prod.instMeasurableSpace, MeasurableSpace.comap_prodMk] + refine le_sup_of_le_left ?_ + rw [← measurable_iff_comap_le] + by_cases hn : n = 0 + · simp only [hn, pullCount_zero] + exact measurable_const + have h_meas := adapted_pullCount_add_one' hA hR' a (n - 1) + rwa [Nat.sub_add_cancel (by lia)] at h_meas + +lemma measurable_indicator_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : + Measurable ({ω : Ω | stepsUntil A a m ω = ↑n}.indicator fun _ ↦ 1) := by + refine (measurable_comap_indicator_stepsUntil_eq hA hR' a m n).mono ?_ le_rfl + refine Measurable.comap_le ?_ + fun_prop + +lemma measurableSet_stepsUntil_eq_zero [MeasurableSingletonClass α] (a : α) (m : ℕ) : + MeasurableSet[MeasurableSpace.comap (A 0) inferInstance] + {ω : Ω | stepsUntil A a m ω = 0} := by + simp only [stepsUntil_eq_zero_iff (a := a) (m := m), ne_eq] + by_cases hm : m = 0 + · simp only [hm, true_and, zero_ne_one, false_and, or_false] + refine (measurableSet_singleton _).compl.preimage ?_ + rw [measurable_iff_comap_le] + by_cases hm1 : m = 1 + swap; · simp [hm, hm1] + simp only [hm1, one_ne_zero, false_and, true_and, false_or] + refine (measurableSet_singleton _).preimage ?_ + rw [measurable_iff_comap_le] + +lemma measurable_comap_indicator_stepsUntil_eq_zero [MeasurableSingletonClass α] (a : α) (m : ℕ) : + Measurable[MeasurableSpace.comap (A 0) inferInstance] + ({ω | stepsUntil A a m ω = 0}.indicator fun _ ↦ 1) := by + rw [measurable_indicator_const_iff] + exact measurableSet_stepsUntil_eq_zero a m + +lemma measurableSet_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : + MeasurableSet[MeasurableSpace.comap (fun ω : Ω ↦ (IsAlgEnvInteraction.hist A R' (n-1) ω, A n ω)) + inferInstance] + {ω : Ω | stepsUntil A a m ω = ↑n} := by + let mProd := MeasurableSpace.comap + (fun ω : Ω ↦ (IsAlgEnvInteraction.hist A R' (n-1) ω, A n ω)) inferInstance + suffices Measurable[mProd] ({ω | stepsUntil A a m ω = ↑n}.indicator fun x ↦ 1) by + rwa [measurable_indicator_const_iff] at this + exact measurable_comap_indicator_stepsUntil_eq hA hR' a m n + +/-- `stepsUntil a m` is a stopping time with respect to the filtration `filtrationAction`. -/ +theorem isStoppingTime_stepsUntil_filtrationAction [Nonempty R] [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m : ℕ) : + IsStoppingTime (IsAlgEnvInteraction.filtrationAction hA hR') (stepsUntil A a m) := by + refine isStoppingTime_of_measurableSet_eq fun n ↦ ?_ + by_cases hn : n = 0 + · simp only [hn, IsAlgEnvInteraction.filtrationAction_zero_eq_comap, WithTop.coe_zero] + exact measurableSet_stepsUntil_eq_zero a m + · rw [IsAlgEnvInteraction.filtrationAction_eq_comap _ hn] + exact measurableSet_stepsUntil_eq hA hR' a m n + +-- /-- Sigma-algebra generated by the stopping time `stepsUntil a m`. -/ +-- def stepsUntilMeasurableSpace [Nonempty R] [MeasurableSingletonClass α] (a : α) (m : ℕ) : +-- MeasurableSpace (ℕ → α × R) := +-- (isStoppingTime_stepsUntil_filtrationAction a m (mR := mR)).measurableSpace + +end Measurability + +end StepsUntil + +section RewardByCount + +/-- Reward obtained when pulling action `a` for the `m`-th time. +If it is never pulled `m` times, the reward is given by the second component of `ω`, which in +applications will be indepedent with same law. -/ +noncomputable +def rewardByCount (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (a : α) (m : ℕ) (ω : Ω × (ℕ → α → R)) : R := + match (stepsUntil A a m ω.1) with + | ⊤ => ω.2 m a + | (n : ℕ) => R' n ω.1 + +variable {ω : Ω × (ℕ → α → R)} + +lemma rewardByCount_eq_ite (a : α) (m : ℕ) (ω : Ω × (ℕ → α → R)) : + rewardByCount A R' a m ω = + if (stepsUntil A a m ω.1) = ⊤ then ω.2 m a else R' (stepsUntil A a m ω.1).toNat ω.1 := by + unfold rewardByCount + cases stepsUntil A a m ω.1 <;> simp + +lemma rewardByCount_eq_add [AddMonoid R] (a : α) (m : ℕ) : + rewardByCount A R' a m = + {ω : Ω × (ℕ → α → R) | stepsUntil A a m ω.1 ≠ ⊤}.indicator + (fun ω ↦ R' (stepsUntil A a m ω.1).toNat ω.1) + + {ω | stepsUntil A a m ω.1 = ⊤}.indicator (fun ω ↦ ω.2 m a) := by + ext ω + simp only [rewardByCount_eq_ite, ne_eq, Pi.add_apply, Set.indicator_apply, Set.mem_setOf_eq, + ite_not] + grind + +lemma rewardByCount_of_stepsUntil_eq_top (h : stepsUntil A a m ω.1 = ⊤) : + rewardByCount A R' a m ω = ω.2 m a := by simp [rewardByCount_eq_ite, h] + +lemma rewardByCount_of_stepsUntil_ne_top (h : stepsUntil A a m ω.1 ≠ ⊤) : + rewardByCount A R' a m ω = R' (stepsUntil A a m ω.1).toNat ω.1 := by + simp [rewardByCount_eq_ite, h] + +lemma rewardByCount_eq_stoppedValue (h : stepsUntil A a m ω.1 ≠ ⊤) : + rewardByCount A R' a m ω = stoppedValue R' (stepsUntil A a m) ω.1 := by + rw [rewardByCount_of_stepsUntil_ne_top h, stoppedValue] + lift stepsUntil A a m ω.1 to ℕ using h with n + simp + +lemma rewardByCount_of_stepsUntil_eq_coe (h : stepsUntil A a m ω.1 = n) : + rewardByCount A R' a m ω = R' n ω.1 := by simp [rewardByCount_eq_ite, h] + +/-- The value at 0 does not matter (it would be the "zeroth" reward). +It should be considered a junk value. -/ +@[simp] +lemma rewardByCount_zero (a : α) (ω : Ω × (ℕ → α → R)) : + rewardByCount A R' a 0 ω = if A 0 ω.1 = a then ω.2 0 a else R' 0 ω.1 := by + rw [rewardByCount_eq_ite] + by_cases ha : A 0 ω.1 = a + · simp [ha, stepsUntil_zero_of_eq] + · simp [stepsUntil_zero_of_ne, ha] + +lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (ω : Ω × (ℕ → α → R)) : + rewardByCount A R' (A t ω.1) (pullCount A (A t ω.1) t ω.1 + 1) ω = R' t ω.1 := by + rw [rewardByCount, ← pullCount_action_eq_pullCount_add_one, stepsUntil_pullCount_eq] + +@[fun_prop] +lemma measurable_rewardByCount [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m : ℕ) : + Measurable (fun ω : Ω × (ℕ → α → R) ↦ rewardByCount A R' a m ω) := by + simp_rw [rewardByCount_eq_ite] + refine Measurable.ite ?_ ?_ ?_ + · exact (measurableSet_singleton _).preimage <| measurable_stepsUntil' hA a m + · fun_prop + · change Measurable ((fun p : ℕ × Ω ↦ R' p.1 p.2) + ∘ (fun ω : Ω × (ℕ → α → R) ↦ ((stepsUntil A a m ω.1).toNat, ω.1))) + have : Measurable fun ω : Ω × (ℕ → α → R) ↦ ((stepsUntil A a m ω.1).toNat, ω.1) := + (measurable_stepsUntil' hA a m).toNat.prodMk (by fun_prop) + refine Measurable.comp ?_ this + refine measurable_from_prod_countable_right fun n ↦ ?_ + simp only + fun_prop + +end RewardByCount + +lemma sum_pullCount_mul [Fintype α] [Semiring R] (ω : Ω) (f : α → R) (t : ℕ) : + ∑ a, pullCount A a t ω * f a = ∑ s ∈ range t, f (A s ω) := by + unfold pullCount + classical + simp_rw [card_eq_sum_ones] + push_cast + simp_rw [sum_mul, one_mul] + exact sum_fiberwise' (range t) (A · ω) f + +-- todo: only in ℝ for now +lemma sum_pullCount [Fintype α] {ω : Ω} : ∑ a, pullCount A a t ω = t := by + suffices ∑ a, pullCount A a t ω * (1 : ℝ) = t by norm_cast at this; simpa + rw [sum_pullCount_mul] + simp + +section SumRewards + +/-- Sum of rewards obtained when pulling action `a` up to time `t` (exclusive). -/ +def sumRewards (A : ℕ → Ω → α) (R' : ℕ → Ω → ℝ) (a : α) (t : ℕ) (ω : Ω) : ℝ := + ∑ s ∈ range t, if A s ω = a then R' s ω else 0 + +/-- Sum of rewards of arm `a` up to (and including) time `n`. -/ +noncomputable +def sumRewards' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := + ∑ s, if (h s).1 = a then (h s).2 else 0 + +/-- Empirical mean reward obtained when pulling action `a` up to time `t` (exclusive). -/ +noncomputable +def empMean (A : ℕ → Ω → α) (R' : ℕ → Ω → ℝ) (a : α) (t : ℕ) (ω : Ω) : ℝ := + sumRewards A R' a t ω / pullCount A a t ω + +/-- Empirical mean of arm `a` at time `n`. -/ +noncomputable +def empMean' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := + (sumRewards' n h a) / (pullCount' n h a) + +lemma sumRewards_eq_pullCount_mul_empMean {R' : ℕ → Ω → ℝ} {ω : Ω} + (h_pull : pullCount A a t ω ≠ 0) : + sumRewards A R' a t ω = pullCount A a t ω * empMean A R' a t ω := by unfold empMean; field_simp + +lemma sum_rewardByCount_eq_sumRewards {R' : ℕ → Ω → ℝ} (a : α) (t : ℕ) (ω : Ω × (ℕ → α → ℝ)) : + ∑ m ∈ Icc 1 (pullCount A a t ω.1), rewardByCount A R' a m ω = sumRewards A R' a t ω.1 := by + induction t with + | zero => simp [pullCount, sumRewards] + | succ t ht => + by_cases hta : A t ω.1 = a + · rw [← hta] at ht ⊢ + rw [pullCount_action_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht] + unfold sumRewards + rw [sum_range_succ, if_pos rfl, rewardByCount_pullCount_add_one_eq_reward] + · unfold sumRewards + rwa [pullCount_eq_pullCount_of_action_ne hta, sum_range_succ, if_neg hta, add_zero] + +lemma sumRewards_add_one_eq_sumRewards' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} : + sumRewards A R' a (n + 1) ω = sumRewards' n (fun i ↦ (A i ω, R' i ω)) a := by + unfold sumRewards sumRewards' + rw [Finset.sum_coe_sort (f := fun s ↦ if A s ω = a then R' s ω else 0) (Iic n)] + congr with m + simp only [mem_range, mem_Iic] + grind + +lemma sumRewards_eq_sumRewards' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (hn : n ≠ 0) : + sumRewards A R' a n ω = sumRewards' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by + cases n with + | zero => exact absurd rfl hn + | succ n => + rw [sumRewards_add_one_eq_sumRewards'] + have : n + 1 - 1 = n := by simp + exact this ▸ rfl + +lemma empMean_add_one_eq_empMean' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} : + empMean A R' a (n + 1) ω = empMean' n (fun i ↦ (A i ω, R' i ω)) a := by + unfold empMean empMean' + rw [sumRewards_add_one_eq_sumRewards', pullCount_add_one_eq_pullCount'] + +lemma empMean_eq_empMean' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (hn : n ≠ 0) : + empMean A R' a n ω = empMean' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by + unfold empMean empMean' + rw [sumRewards_eq_sumRewards' hn, pullCount_eq_pullCount' hn] + +@[fun_prop] +lemma measurable_sumRewards [MeasurableSingletonClass α] {R' : ℕ → Ω → ℝ} + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (t : ℕ) : + Measurable (sumRewards A R' a t) := by + unfold sumRewards + have h_meas s : Measurable (fun h : Ω ↦ if A s h = a then R' s h else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + fun_prop + +@[fun_prop] +lemma measurable_empMean [MeasurableSingletonClass α] {R' : ℕ → Ω → ℝ} (hA : ∀ n, Measurable (A n)) + (hR' : ∀ n, Measurable (R' n)) (a : α) (n : ℕ) : + Measurable (empMean A R' a n) := by + unfold empMean + fun_prop + +@[fun_prop] +lemma measurable_sumRewards' [MeasurableSingletonClass α] (n : ℕ) (a : α) : + Measurable (fun h ↦ sumRewards' n h a) := by + simp_rw [sumRewards'] + have h_meas s : Measurable (fun (h : Iic n → α × ℝ) ↦ if (h s).1 = a then (h s).2 else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + fun_prop + +@[fun_prop] +lemma measurable_empMean' [MeasurableSingletonClass α] (n : ℕ) (a : α) : + Measurable (fun h ↦ empMean' n h a) := by + unfold empMean' + fun_prop + +end SumRewards + +end LearningDraft diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index cbd1fbfa..a4e5ff01 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -487,6 +487,13 @@ lemma measurableSet_stepsUntil_eq_zero [Nonempty R] [MeasurableSingletonClass α refine (measurableSet_singleton _).preimage ?_ rw [measurable_iff_comap_le] +lemma measurable_comap_indicator_stepsUntil_eq_zero [Nonempty R] [MeasurableSingletonClass α] + (a : α) (m : ℕ) : + Measurable[MeasurableSpace.comap (action 0 (R := R)) inferInstance] + ({ω | stepsUntil a m ω = 0}.indicator fun _ ↦ 1) := by + rw [measurable_indicator_const_iff] + exact measurableSet_stepsUntil_eq_zero a m + lemma measurableSet_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] (a : α) (m n : ℕ) : MeasurableSet[MeasurableSpace.comap (fun ω : ℕ → α × R ↦ (hist (n-1) ω, action n ω)) inferInstance] @@ -535,6 +542,16 @@ lemma rewardByCount_eq_ite (a : α) (m : ℕ) (ω : (ℕ → α × R) × (ℕ unfold rewardByCount cases stepsUntil a m ω.1 <;> simp +lemma rewardByCount_eq_add [AddMonoid R] (a : α) (m : ℕ) : + rewardByCount a m = + {ω : (ℕ → α × R) × (ℕ → α → R) | stepsUntil a m ω.1 ≠ ⊤}.indicator + (fun ω ↦ reward (stepsUntil a m ω.1).toNat ω.1) + + {ω | stepsUntil a m ω.1 = ⊤}.indicator (fun ω ↦ ω.2 m a) := by + ext ω + simp only [rewardByCount_eq_ite, ne_eq, Pi.add_apply, Set.indicator_apply, Set.mem_setOf_eq, + ite_not] + grind + lemma rewardByCount_of_stepsUntil_eq_top (h : stepsUntil a m ω.1 = ⊤) : rewardByCount a m ω = ω.2 m a := by simp [rewardByCount_eq_ite, h] From ace4f950fdf740fdc2465794be1eee99e280e5cb Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Tue, 6 Jan 2026 11:59:06 +0100 Subject: [PATCH 04/30] update to IsAlgEnvSeq --- LeanBandits.lean | 1 - LeanBandits/Bandit/Bandit.lean | 186 +---- LeanBandits/Bandit/Regret.lean | 31 +- LeanBandits/BanditAlgorithms/ETC.lean | 249 +++--- LeanBandits/BanditAlgorithms/UCB.lean | 508 +++++++----- LeanBandits/ForMathlib/HasCondDistrib.lean | 35 + LeanBandits/ForMathlib/Traj.lean | 17 +- LeanBandits/RewardByCountMeasure.lean | 466 ++++------- LeanBandits/SequentialLearning/Algorithm.lean | 212 +++-- .../SequentialLearning/Deterministic.lean | 59 +- LeanBandits/SequentialLearning/Draft.lean | 772 ------------------ .../SequentialLearning/FiniteActions.lean | 555 +++++++------ .../SequentialLearning/StationaryEnv.lean | 100 ++- 13 files changed, 1269 insertions(+), 1922 deletions(-) delete mode 100644 LeanBandits/SequentialLearning/Draft.lean diff --git a/LeanBandits.lean b/LeanBandits.lean index d6909492..10b0f463 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -14,6 +14,5 @@ import LeanBandits.ForMathlib.Traj import LeanBandits.RewardByCountMeasure import LeanBandits.SequentialLearning.Algorithm import LeanBandits.SequentialLearning.Deterministic -import LeanBandits.SequentialLearning.Draft import LeanBandits.SequentialLearning.FiniteActions import LeanBandits.SequentialLearning.StationaryEnv diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 5d7dd348..1788b187 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -26,7 +26,8 @@ section MeasureSpace namespace Bandit -/-- Kernel describing the distribution of the next arm-reward pair given the history up to `n`. -/ +/-- Kernel describing the distribution of the next action-reward pair given the history up to +time `n`. -/ noncomputable def stepKernel (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : Kernel (Iic n → α × R) (α × R) := @@ -43,13 +44,13 @@ lemma snd_stepKernel (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel (stepKernel alg ν n).snd = ν ∘ₖ alg.policy n := by rw [stepKernel, Learning.stepKernel, stationaryEnv_feedback, Kernel.snd_compProd_prodMkLeft] -/-- Measure on the sequence of arms pulled and rewards observed generated by the bandit. -/ +/-- Measure on the sequence of actions pulled and rewards observed generated by the bandit. -/ noncomputable def trajMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : Measure (ℕ → α × R) := Learning.trajMeasure alg (stationaryEnv ν) deriving IsProbabilityMeasure -/-- Measure of an infinite stream of rewards from each arm. -/ +/-- Measure of an infinite stream of rewards from each action. -/ noncomputable def streamMeasure (ν : Kernel α R) : Measure (ℕ → α → R) := Measure.infinitePi fun _ ↦ Measure.infinitePi ν @@ -58,8 +59,8 @@ instance (ν : Kernel α R) [IsMarkovKernel ν] : IsProbabilityMeasure (streamMe unfold streamMeasure infer_instance -/-- Joint distribution of the sequence of arm pulled and rewards, and a stream of independent -rewards from all arms. -/ +/-- Joint distribution of the sequence of action pulled and rewards, and a stream of independent +rewards from all actions. -/ noncomputable def measure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : Measure ((ℕ → α × R) × (ℕ → α → R)) := @@ -158,133 +159,31 @@ lemma indepFun_eval_snd_measure (alg : Algorithm α R) (ν : Kernel α R) [IsMar end StreamMeasure -/-- `arm n` is the arm pulled at time `n`. This is a random variable on the measurable space -`ℕ → α × ℝ`. -/ -def arm (n : ℕ) (h : ℕ → α × R) : α := (h n).1 - -/-- `reward n` is the reward at time `n`. This is a random variable on the measurable space -`ℕ → α × R`. -/ -def reward (n : ℕ) (h : ℕ → α × R) : R := (h n).2 - -/-- `hist n` is the history up to time `n`. This is a random variable on the measurable space -`ℕ → α × R`. -/ -def hist (n : ℕ) (h : ℕ → α × R) : Iic n → α × R := fun i ↦ h i - -@[fun_prop] -lemma measurable_arm (n : ℕ) : Measurable (arm n (α := α) (R := R)) := measurable_action n - -@[fun_prop] -lemma measurable_arm_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ arm p.1 p.2) := - measurable_action_prod - -@[fun_prop] -lemma measurable_reward (n : ℕ) : Measurable (reward n (α := α) (R := R)) := - Learning.measurable_reward n - -@[fun_prop] -lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ reward p.1 p.2) := - Learning.measurable_reward_prod - -@[fun_prop] -lemma measurable_hist (n : ℕ) : Measurable (hist n (α := α) (R := R)) := - Learning.measurable_hist n - -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 Laws - -lemma hasLaw_step_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : - HasLaw (fun h : ℕ → α × R ↦ h 0) (alg.p0 ⊗ₘ ν) (Bandit.trajMeasure alg ν) := - Learning.hasLaw_step_zero alg (stationaryEnv ν) - -lemma hasLaw_arm_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : - HasLaw (arm 0) alg.p0 (Bandit.trajMeasure alg ν) := - Learning.hasLaw_action_zero alg (stationaryEnv ν) - -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 := - Learning.condDistrib_step alg (stationaryEnv ν) n - -lemma condDistrib_reward' [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - condDistrib (reward (n + 1)) (fun ω ↦ (hist n ω, arm (n + 1) ω)) (Bandit.trajMeasure alg ν) - =ᵐ[(Bandit.trajMeasure alg ν).map (fun ω ↦ (hist n ω, arm (n + 1) ω))] ν.prodMkLeft _ := - Learning.condDistrib_reward alg (stationaryEnv ν) 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)] ν := - Learning.condDistrib_reward_stationaryEnv n - -lemma condDistrib_arm [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - condDistrib (arm (n + 1)) (hist n) (Bandit.trajMeasure alg ν) - =ᵐ[(Bandit.trajMeasure alg ν).map (hist n)] alg.policy n := - Learning.condDistrib_action alg (stationaryEnv ν) n - -/-- 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 α] [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 ν) := - Learning.condIndepFun_reward_hist_action n - -lemma condIndepFun_reward_hist_arm_arm [StandardBorelSpace α] [Countable α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] - {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) : - reward (n + 1) ⟂ᵢ[arm (n + 1), measurable_arm (n + 1); Bandit.trajMeasure alg ν] - (fun ω ↦ (hist n ω, arm (n + 1) ω)) := by - have h_indep : reward (n + 1) ⟂ᵢ[arm (n + 1), measurable_arm (n + 1); Bandit.trajMeasure alg ν] - hist n := by - convert condIndepFun_reward_hist_arm (alg := alg) (ν := ν) n - exact h_indep.prod_right (by fun_prop) (by fun_prop) (by fun_prop) - -lemma condIndepFun_reward_hist_arm_arm' [StandardBorelSpace α] [Countable α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] - {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) (hn : n ≠ 0) : - reward n ⟂ᵢ[arm n, measurable_arm n; Bandit.trajMeasure alg ν] - (fun ω ↦ (hist (n - 1) ω, arm n ω)) := by - have := condIndepFun_reward_hist_arm_arm (alg := alg) (ν := ν) (n - 1) - grind - -end Laws - section DetAlgorithm -variable {nextArm : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextArm n)} - {arm0 : α} {ν : Kernel α R} [IsMarkovKernel ν] +variable {nextaction : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextaction n)} + {action0 : α} {ν : Kernel α R} [IsMarkovKernel ν] -local notation "𝔓t" => Bandit.trajMeasure (detAlgorithm nextArm h_next arm0) ν +local notation "𝔓t" => Bandit.trajMeasure (detAlgorithm nextaction h_next action0) ν -lemma HasLaw_arm_zero_detAlgorithm : HasLaw (arm 0) (Measure.dirac arm0) 𝔓t where - map_eq := (hasLaw_arm_zero _ _).map_eq +lemma HasLaw_action_zero_detAlgorithm : HasLaw (IT.action 0) (Measure.dirac action0) 𝔓t where + map_eq := (IT.hasLaw_action_zero _ _).map_eq -lemma arm_zero_detAlgorithm [MeasurableSingletonClass α] : - arm 0 =ᵐ[𝔓t] fun _ ↦ arm0 := - Learning.action_zero_detAlgorithm +lemma action_zero_detAlgorithm [MeasurableSingletonClass α] : + IT.action 0 =ᵐ[𝔓t] fun _ ↦ action0 := + IT.action_zero_detAlgorithm -lemma arm_detAlgorithm_ae_eq [StandardBorelSpace α] [Nonempty α] +lemma action_detAlgorithm_ae_eq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] (n : ℕ) : - arm (n + 1) =ᵐ[𝔓t] fun h ↦ nextArm n (fun i ↦ h i) := - Learning.action_detAlgorithm_ae_eq n + IT.action (n + 1) =ᵐ[𝔓t] fun h ↦ nextaction n (fun i ↦ h i) := + IT.action_detAlgorithm_ae_eq n example [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] : - ∀ᵐ h ∂(𝔓t), arm 0 h = arm0 ∧ ∀ n, arm (n + 1) h = nextArm n (fun i ↦ h i) := by + ∀ᵐ h ∂(𝔓t), IT.action 0 h = action0 ∧ + ∀ n, IT.action (n + 1) h = nextaction n (fun i ↦ h i) := by rw [eventually_and, ae_all_iff] - exact ⟨arm_zero_detAlgorithm, arm_detAlgorithm_ae_eq⟩ + exact ⟨action_zero_detAlgorithm, action_detAlgorithm_ae_eq⟩ end DetAlgorithm @@ -405,26 +304,26 @@ lemma measurable_altHist [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : sorry noncomputable -def altArm [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : α := +def altaction [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : α := (altHist alg ω n ⟨n, by simp⟩).1 -lemma altArm_zero [DecidableEq α] (alg : Algorithm α R) : - altArm alg 0 = fun ω ↦ initAlgFunction alg (ω.1 0) := by +lemma altaction_zero [DecidableEq α] (alg : Algorithm α R) : + altaction alg 0 = fun ω ↦ initAlgFunction alg (ω.1 0) := by ext - simp [altArm, altHist_zero] + simp [altaction, altHist_zero] @[fun_prop] -lemma measurable_altArm [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : - Measurable (altArm alg n) := by unfold altArm; fun_prop +lemma measurable_altaction [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : + Measurable (altaction alg n) := by unfold altaction; fun_prop noncomputable def altReward [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : R := (altHist alg ω n ⟨n, by simp⟩).2 lemma altReward_zero [DecidableEq α] (alg : Algorithm α R) : - altReward alg 0 = fun ω ↦ ω.2 0 (altArm alg 0 ω) := by + altReward alg 0 = fun ω ↦ ω.2 0 (altaction alg 0 ω) := by ext - simp [altReward, altHist_zero, altArm_zero] + simp [altReward, altHist_zero, altaction_zero] @[fun_prop] lemma measurable_altReward [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : @@ -432,8 +331,8 @@ lemma measurable_altReward [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : variable [DecidableEq α] -lemma hasLaw_altArm_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : - HasLaw (altArm alg 0) alg.p0 (arrayMeasure ν) where +lemma hasLaw_altaction_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : + HasLaw (altaction alg 0) alg.p0 (arrayMeasure ν) where map_eq := by calc (arrayMeasure ν).map (fun ω ↦ initAlgFunction alg (ω.1 0)) _ = ((arrayMeasure ν).fst.map (Function.eval 0)).map (initAlgFunction alg) := by @@ -448,9 +347,9 @@ lemma hasLaw_altArm_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKern variable [StandardBorelSpace R] [Nonempty R] lemma hasCondDistrib_altReward_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : - HasCondDistrib (altReward alg 0) (altArm alg 0) (stationaryEnv ν).ν0 (arrayMeasure ν) where + HasCondDistrib (altReward alg 0) (altaction alg 0) (stationaryEnv ν).ν0 (arrayMeasure ν) where condDistrib_eq := by - simp only [stationaryEnv_ν0, (hasLaw_altArm_zero alg ν).map_eq, altReward_zero] + simp only [stationaryEnv_ν0, (hasLaw_altaction_zero alg ν).map_eq, altReward_zero] sorry lemma hasCondDistrib_altStep' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] @@ -463,17 +362,18 @@ lemma hasCondDistrib_altStep' (alg : Algorithm α R) (ν : Kernel α R) [IsMarko lemma hasCondDistrib_altStep (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - HasCondDistrib (fun ω ↦ (altArm alg (n + 1) ω, altReward alg (n + 1) ω)) - (fun ω (i : Iic n) ↦ (altArm alg i ω, altReward alg i ω)) + HasCondDistrib (fun ω ↦ (altaction alg (n + 1) ω, altReward alg (n + 1) ω)) + (fun ω (i : Iic n) ↦ (altaction alg i ω, altReward alg i ω)) (Bandit.stepKernel alg ν n) (arrayMeasure ν) := by convert hasCondDistrib_altStep' alg ν n with ω i - · simp only [altArm] + · simp only [altaction] rw [altHist_eq _ _ n] · simp only [altReward] rw [altHist_eq _ _ n] -lemma hasCondDistrib_altArm (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - HasCondDistrib (altArm alg (n + 1)) (fun ω (i : Iic n) ↦ (altArm alg i ω, altReward alg i ω)) +lemma hasCondDistrib_altaction (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + HasCondDistrib (altaction alg (n + 1)) + (fun ω (i : Iic n) ↦ (altaction alg i ω, altReward alg i ω)) (alg.policy n) (arrayMeasure ν) := by convert HasCondDistrib.fst (hasCondDistrib_altStep alg ν n) simp @@ -481,16 +381,16 @@ lemma hasCondDistrib_altArm (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovK lemma hasCondDistrib_altReward (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : HasCondDistrib (altReward alg (n + 1)) - (fun ω ↦ (fun (i : Iic n) ↦ (altArm alg i ω, altReward alg i ω), altArm alg (n + 1) ω)) + (fun ω ↦ (fun (i : Iic n) ↦ (altaction alg i ω, altReward alg i ω), altaction alg (n + 1) ω)) ((stationaryEnv ν).feedback n) (arrayMeasure ν) := by simp only [stationaryEnv_feedback] sorry -lemma isAlgEnvInteraction_arrayMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : - IsAlgEnvInteraction (altArm alg) (altReward alg) alg (stationaryEnv ν) (arrayMeasure ν) where - hasLaw_action_zero := hasLaw_altArm_zero alg ν +lemma isAlgEnvSeq_arrayMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : + IsAlgEnvSeq (altaction alg) (altReward alg) alg (stationaryEnv ν) (arrayMeasure ν) where + hasLaw_action_zero := hasLaw_altaction_zero alg ν hasCondDistrib_reward_zero := hasCondDistrib_altReward_zero alg ν - hasCondDistrib_action := hasCondDistrib_altArm alg ν + hasCondDistrib_action := hasCondDistrib_altaction alg ν hasCondDistrib_reward := hasCondDistrib_altReward alg ν end ArrayModel diff --git a/LeanBandits/Bandit/Regret.lean b/LeanBandits/Bandit/Regret.lean index acabf6c8..3f35f89f 100644 --- a/LeanBandits/Bandit/Regret.lean +++ b/LeanBandits/Bandit/Regret.lean @@ -17,17 +17,19 @@ open scoped ENNReal NNReal namespace Bandits -variable {α : Type*} [DecidableEq α] {mα : MeasurableSpace α} {ν : Kernel α ℝ} - {h : ℕ → α × ℝ} {m n t : ℕ} {a : α} +variable {α Ω : Type*} [DecidableEq α] {mα : MeasurableSpace α} {mΩ : MeasurableSpace Ω} + {ν : Kernel α ℝ} + {A : ℕ → Ω → α} {R : ℕ → Ω → ℝ} + {ω : Ω} {m n t : ℕ} {a : α} /-! ### Definitions of regret, gaps, pull counts -/ /-- Regret of a sequence of pulls `k : ℕ → α` at time `t` for the reward kernel `ν ; Kernel α ℝ`. -/ noncomputable -def regret (ν : Kernel α ℝ) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := - t * (⨆ a, (ν a)[id]) - ∑ s ∈ range t, (ν (arm s h))[id] +def regret (ν : Kernel α ℝ) (A : ℕ → Ω → α) (t : ℕ) (ω : Ω) : ℝ := + t * (⨆ a, (ν a)[id]) - ∑ s ∈ range t, (ν (A s ω))[id] -/-- Gap of an arm `a`: difference between the highest mean of the arms and the mean of `a`. -/ +/-- Gap of an action `a`: difference between the highest mean of the actions and the mean of `a`. -/ noncomputable def gap (ν : Kernel α ℝ) (a : α) : ℝ := (⨆ i, (ν i)[id]) - (ν a)[id] @@ -36,28 +38,19 @@ lemma gap_nonneg [Fintype α] : 0 ≤ gap ν a := by rw [gap, sub_nonneg] exact le_ciSup (f := fun i ↦ (ν i)[id]) (by simp) a -lemma arm_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) h = m) : - arm (stepsUntil a m h).toNat h = a := by - exact action_stepsUntil hm h_exists - -lemma arm_eq_of_stepsUntil_eq_coe {ω : ℕ → α × ℝ} (hm : m ≠ 0) - (h : stepsUntil a m ω = n) : - arm n ω = a := by - exact action_eq_of_stepsUntil_eq_coe hm h - section RewardByCount lemma regret_eq_sum_pullCount_mul_gap [Fintype α] : - regret ν t h = ∑ a, pullCount a t h * gap ν a := by - simp [sum_pullCount_mul, regret, gap, sum_sub_distrib, arm, action] + regret ν A t ω = ∑ a, pullCount A a t ω * gap ν a := by + simp [sum_pullCount_mul, regret, gap, sum_sub_distrib] end RewardByCount -section BestArm +section bestArm variable [Fintype α] [Nonempty α] -/-- Arm with the highest mean. -/ +/-- action with the highest mean. -/ noncomputable def bestArm (ν : Kernel α ℝ) : α := (exists_max_image univ (fun a ↦ (ν a)[id]) (univ_nonempty_iff.mpr inferInstance)).choose @@ -78,6 +71,6 @@ omit [DecidableEq α] in lemma gap_bestArm : gap ν (bestArm ν) = 0 := by rw [gap_eq_bestArm_sub, sub_self] -end BestArm +end bestArm end Bandits diff --git a/LeanBandits/BanditAlgorithms/ETC.lean b/LeanBandits/BanditAlgorithms/ETC.lean index 048cb41c..49c3d90c 100644 --- a/LeanBandits/BanditAlgorithms/ETC.lean +++ b/LeanBandits/BanditAlgorithms/ETC.lean @@ -116,55 +116,66 @@ end AlgorithmDefinition namespace ETC variable {hK : 0 < K} {m : ℕ} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] + {Ω : Type*} {mΩ : MeasurableSpace Ω} + {P : Measure Ω} [IsProbabilityMeasure P] + {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} -local notation "𝔓t" => Bandit.trajMeasure (etcAlgorithm hK m) ν -local notation "𝔓" => Bandit.measure (etcAlgorithm hK m) ν +local notation "𝔓" => P.prod (Bandit.streamMeasure ν) -lemma arm_zero : arm 0 =ᵐ[𝔓t] fun _ ↦ ⟨0, hK⟩ := by +lemma arm_zero [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) : + A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - exact arm_zero_detAlgorithm + exact h.action_zero_detAlgorithm -lemma arm_ae_eq_etcNextArm (n : ℕ) : - arm (n + 1) =ᵐ[𝔓t] fun h ↦ nextArm hK m n (fun i ↦ h i) := by +lemma arm_ae_eq_etcNextArm [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (n : ℕ) : + A (n + 1) =ᵐ[P] fun ω ↦ nextArm hK m n (IsAlgEnvSeq.hist A R n ω) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - exact arm_detAlgorithm_ae_eq n + exact h.action_detAlgorithm_ae_eq n /-- For `n < K * m`, the arm pulled at time `n` is the arm `n % K`. -/ -lemma arm_of_lt {n : ℕ} (hn : n < K * m) : - arm n =ᵐ[𝔓t] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := by +lemma arm_of_lt [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) {n : ℕ} (hn : n < K * m) : + A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := by cases n with - | zero => exact arm_zero + | zero => exact arm_zero h | succ n => - filter_upwards [arm_ae_eq_etcNextArm n] with h hn_eq + filter_upwards [arm_ae_eq_etcNextArm h n] with h hn_eq rw [hn_eq, nextArm, dif_pos] grind /-- The arm pulled at time `K * m` is the arm with the highest empirical mean after the exploration phase. -/ -lemma arm_mul (hm : m ≠ 0) : - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - arm (K * m) =ᵐ[𝔓t] fun h ↦ measurableArgmax (empMean' (K * m - 1)) (fun i ↦ h i) := by +lemma arm_mul [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (hm : m ≠ 0) : + A (K * m) =ᵐ[P] fun ω ↦ measurableArgmax (empMean' (K * m - 1)) + (IsAlgEnvSeq.hist A R (K * m - 1) ω) := by have : K * m = (K * m - 1) + 1 := by have : 0 < K * m := Nat.mul_pos hK hm.bot_lt grind rw [this] - filter_upwards [arm_ae_eq_etcNextArm (K * m - 1)] with h hn_eq + filter_upwards [arm_ae_eq_etcNextArm h (K * m - 1)] with ω hn_eq rw [hn_eq, nextArm, dif_neg (by simp), dif_pos rfl] exact this ▸ rfl /-- For `n ≥ K * m`, the arm pulled at time `n + 1` is the same as the arm pulled at time `n`. -/ -lemma arm_add_one_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : - arm (n + 1) =ᵐ[𝔓t] fun ω ↦ arm n ω := by - filter_upwards [arm_ae_eq_etcNextArm n] with ω hn_eq +lemma arm_add_one_of_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : + A (n + 1) =ᵐ[P] fun ω ↦ A n ω := by + filter_upwards [arm_ae_eq_etcNextArm h n] with ω hn_eq rw [hn_eq, nextArm, dif_neg (by grind), dif_neg] · rfl · have : 0 < K * m := Nat.mul_pos hK hm.bot_lt grind /-- For `n ≥ K * m`, the arm pulled at time `n` is the same as the arm pulled at time `K * m`. -/ -lemma arm_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : - arm n =ᵐ[𝔓t] arm (K * m) := by - have h_ae n : K * m ≤ n → arm (n + 1) =ᵐ[𝔓t] fun ω ↦ arm n ω := arm_add_one_of_ge hm +lemma arm_of_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : + A n =ᵐ[P] A (K * m) := by + have h_ae n : K * m ≤ n → A (n + 1) =ᵐ[P] fun ω ↦ A n ω := arm_add_one_of_ge h hm simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_ae filter_upwards [h_ae] with ω h_ae induction n, hn using Nat.le_induction with @@ -172,36 +183,42 @@ lemma arm_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : | succ n hmn h_ind => rw [h_ae n hmn, h_ind] /-- At time `K * m`, the number of pulls of each arm is equal to `m`. -/ -lemma pullCount_mul (a : Fin K) : pullCount a (K * m) =ᵐ[𝔓t] fun _ ↦ m := by +lemma pullCount_mul [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (a : Fin K) : + pullCount A a (K * m) =ᵐ[P] fun _ ↦ m := by rw [Filter.EventuallyEq] simp_rw [pullCount_eq_sum] - have h_arm (n : range (K * m)) : arm n =ᵐ[𝔓t] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := - arm_of_lt (mem_range.mp n.2) + have h_arm (n : range (K * m)) : A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := + arm_of_lt h (mem_range.mp n.2) simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_arm filter_upwards [h_arm] with ω h_arm - have h_arm' {i : ℕ} (hi : i ∈ range (K * m)) : arm i ω = ⟨i % K, Nat.mod_lt _ hK⟩ := h_arm ⟨i, hi⟩ - calc (∑ s ∈ range (K * m), if arm s ω = a then 1 else 0) + have h_arm' {i : ℕ} (hi : i ∈ range (K * m)) : A i ω = ⟨i % K, Nat.mod_lt _ hK⟩ := h_arm ⟨i, hi⟩ + calc (∑ s ∈ range (K * m), if A s ω = a then 1 else 0) _ = (∑ s ∈ range (K * m), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := sum_congr rfl fun s hs ↦ by rw [h_arm' hs] _ = m := sum_mod_range_mul hK m a -lemma pullCount_add_one_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : - pullCount a (n + 1) - =ᵐ[𝔓t] fun ω ↦ pullCount a n ω + {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by +lemma pullCount_add_one_of_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : + pullCount A a (n + 1) + =ᵐ[P] fun ω ↦ pullCount A a n ω + {ω' | A (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by simp_rw [Filter.EventuallyEq, pullCount_add_one] - filter_upwards [arm_of_ge hm hn] with ω h_arm + filter_upwards [arm_of_ge h hm hn] with ω h_arm congr 3 /-- For `n ≥ K * m`, the number of pulls of each arm `a` at time `n` is equal to `m` plus `n - K * m` if arm `a` is the best arm after the exploration phase. -/ -lemma pullCount_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : - pullCount a n - =ᵐ[𝔓t] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by - have h_ae n : K * m ≤ n → pullCount a (n + 1) - =ᵐ[𝔓t] fun ω ↦ pullCount a n ω + {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := - pullCount_add_one_of_ge a hm +lemma pullCount_of_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : + pullCount A a n + =ᵐ[P] fun ω ↦ m + (n - K * m) * {ω' | A (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by + have h_ae n : K * m ≤ n → pullCount A a (n + 1) + =ᵐ[P] fun ω ↦ pullCount A a n ω + {ω' | A (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := + pullCount_add_one_of_ge h a hm simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_ae - have h_ae_Km : pullCount a (K * m) =ᵐ[𝔓t] fun _ ↦ m := pullCount_mul a + have h_ae_Km : pullCount A a (K * m) =ᵐ[P] fun _ ↦ m := pullCount_mul h a filter_upwards [h_ae_Km, h_ae] with ω h_Km h_ae induction n, hn using Nat.le_induction with | base => simp [h_Km] @@ -212,13 +229,14 @@ lemma pullCount_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : /-- If at time `K * m` the algorithm chooses arm `a`, then the total reward obtained by pulling arm `a` is at least the total reward obtained by pulling the best arm. -/ -lemma sumRewards_bestArm_le_of_arm_mul_eq (a : Fin K) (hm : m ≠ 0) : - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - ∀ᵐ h ∂𝔓t, arm (K * m) h = a → sumRewards (bestArm ν) (K * m) h ≤ sumRewards a (K * m) h := by - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - filter_upwards [arm_mul hm, pullCount_mul a, pullCount_mul (bestArm ν)] with h h_arm ha h_best - h_eq - have h_max := isMaxOn_measurableArgmax (empMean' (K * m - 1)) (fun i ↦ h i) (bestArm ν) +lemma sumRewards_bestArm_le_of_arm_mul_eq [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (a : Fin K) (hm : m ≠ 0) : + ∀ᵐ h ∂P, A (K * m) h = a → sumRewards A R (bestArm ν) (K * m) h ≤ + sumRewards A R a (K * m) h := by + filter_upwards [arm_mul h hm, pullCount_mul h a, pullCount_mul h (bestArm ν)] + with h h_arm ha h_best h_eq + have h_max := isMaxOn_measurableArgmax (empMean' (K * m - 1)) (IsAlgEnvSeq.hist A R (K * m - 1) h) + (bestArm ν) rw [← h_arm, h_eq] at h_max rw [sumRewards_eq_pullCount_mul_empMean, sumRewards_eq_pullCount_mul_empMean, ha, h_best] · gcongr @@ -227,108 +245,121 @@ lemma sumRewards_bestArm_le_of_arm_mul_eq (a : Fin K) (hm : m ≠ 0) : · simp [ha, hm] · simp [h_best, hm] -lemma identDistrib_aux (m : ℕ) (a b : Fin K) : +variable [StandardBorelSpace Ω] [Nonempty Ω] + +lemma identDistrib_aux [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (a b : Fin K) : IdentDistrib - (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount a s ω, ∑ s ∈ Icc 1 m, rewardByCount b s ω)) - (fun ω ↦ (∑ s ∈ range m, ω.2 s a, ∑ s ∈ range m, ω.2 s b)) 𝔓 𝔓 := by - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - have h2 (a : Fin K) : IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount a s ω) - (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓 𝔓 := identDistrib_sum_Icc_rewardByCount m a + (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount A R a s ω, ∑ s ∈ Icc 1 m, rewardByCount A R b s ω)) + (fun ω ↦ (∑ s ∈ range m, ω.2 s a, ∑ s ∈ range m, ω.2 s b)) + 𝔓 (Bandit.measure (etcAlgorithm hK m) ν) := by + have h2 (a : Fin K) : IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount A R a s ω) + (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓 (Bandit.measure (etcAlgorithm hK m) ν) := + identDistrib_sum_Icc_rewardByCount h m a by_cases hab : a = b · simp only [hab] exact (h2 b).comp (u := fun p ↦ (p, p)) (by fun_prop) refine (h2 a).prodMk (h2 b) ?_ ?_ - · suffices IndepFun (fun ω s ↦ rewardByCount a s ω) (fun ω s ↦ rewardByCount b s ω) + · suffices IndepFun (fun ω s ↦ rewardByCount A R a s ω) (fun ω s ↦ rewardByCount A R b s ω) 𝔓 by exact this.comp (φ := fun p ↦ ∑ i ∈ Icc 1 m, p i) (ψ := fun p ↦ ∑ j ∈ Icc 1 m, p j) (by fun_prop) (by fun_prop) - exact indepFun_rewardByCount_of_ne hab - · suffices IndepFun (fun ω s ↦ ω.2 s a) (fun ω s ↦ ω.2 s b) 𝔓 by + exact indepFun_rewardByCount_of_ne h hab + · suffices IndepFun (fun ω s ↦ ω.2 s a) (fun ω s ↦ ω.2 s b) + (Bandit.measure (etcAlgorithm hK m) ν) by exact this.comp (φ := fun p ↦ ∑ i ∈ range m, p i) (ψ := fun p ↦ ∑ j ∈ range m, p j) (by fun_prop) (by fun_prop) exact indepFun_eval_snd_measure _ ν hab /-- The probability that at time `K * m` the ETC algorithm chooses arm `a` is at most `exp(- m * Δ_a^2 / 4)`. -/ -lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) +lemma prob_arm_mul_eq_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) (hm : m ≠ 0) : - (𝔓t).real {ω | arm (K * m) ω = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + P.real {ω | A (K * m) ω = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by + have hA := h.measurable_A + have hR := h.measurable_R have h_pos : 0 < K * m := Nat.mul_pos hK hm.bot_lt - have h_le : (𝔓t).real {ω | arm (K * m) ω = a} - ≤ (𝔓t).real {ω | sumRewards (bestArm ν) (K * m) ω ≤ sumRewards a (K * m) ω} := by + have h_le : P.real {ω | A (K * m) ω = a} + ≤ P.real {ω | sumRewards A R (bestArm ν) (K * m) ω ≤ sumRewards A R a (K * m) ω} := by simp_rw [measureReal_def] gcongr 1 · simp refine measure_mono_ae ?_ - exact sumRewards_bestArm_le_of_arm_mul_eq a hm + exact sumRewards_bestArm_le_of_arm_mul_eq h a hm refine h_le.trans ?_ -- extend the probability space to include the stream of independent rewards - suffices (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} + suffices (𝔓).real {ω | sumRewards A R (bestArm ν) (K * m) ω.1 ≤ sumRewards A R a (K * m) ω.1} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) by - suffices (𝔓t).real {ω | sumRewards (bestArm ν) (K * m) ω ≤ sumRewards a (K * m) ω} - = (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} by rwa [this] - calc (𝔓t).real {ω | sumRewards (bestArm ν) (K * m) ω ≤ sumRewards a (K * m) ω} - _ = ((𝔓).fst).real {ω | sumRewards (bestArm ν) (K * m) ω ≤ sumRewards a (K * m) ω} := by simp - _ = (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} := by + suffices P.real {ω | sumRewards A R (bestArm ν) (K * m) ω ≤ sumRewards A R a (K * m) ω} + = (𝔓).real {ω | sumRewards A R (bestArm ν) (K * m) ω.1 ≤ sumRewards A R a (K * m) ω.1} by + rwa [this] + calc P.real {ω | sumRewards A R (bestArm ν) (K * m) ω ≤ sumRewards A R a (K * m) ω} + _ = ((𝔓).fst).real {ω | sumRewards A R (bestArm ν) (K * m) ω ≤ sumRewards A R a (K * m) ω} := by + simp + _ = (𝔓).real {ω | sumRewards A R (bestArm ν) (K * m) ω.1 ≤ sumRewards A R a (K * m) ω.1} := by rw [Measure.fst, map_measureReal_apply (by fun_prop)] · rfl - · exact measurableSet_le (by fun_prop) (by fun_prop) - calc (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} - _ = (𝔓).real {ω | ∑ s ∈ Icc 1 (pullCount (bestArm ν) (K * m) ω.1), rewardByCount (bestArm ν) s ω - ≤ ∑ s ∈ Icc 1 (pullCount a (K * m) ω.1), rewardByCount a s ω} := by + · have h_meas := measurable_sumRewards h.measurable_A h.measurable_R + exact measurableSet_le (by fun_prop) (by fun_prop) + calc (𝔓).real {ω | sumRewards A R (bestArm ν) (K * m) ω.1 ≤ sumRewards A R a (K * m) ω.1} + _ = (𝔓).real {ω | ∑ s ∈ Icc 1 (pullCount A (bestArm ν) (K * m) ω.1), + rewardByCount A R (bestArm ν) s ω + ≤ ∑ s ∈ Icc 1 (pullCount A a (K * m) ω.1), rewardByCount A R a s ω} := by congr with ω congr! 1 <;> rw [sum_rewardByCount_eq_sumRewards] - _ = (𝔓).real {ω | ∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω - ≤ ∑ s ∈ Icc 1 m, rewardByCount a s ω} := by + _ = (𝔓).real {ω | ∑ s ∈ Icc 1 m, rewardByCount A R (bestArm ν) s ω + ≤ ∑ s ∈ Icc 1 m, rewardByCount A R a s ω} := by simp_rw [measureReal_def] congr 1 refine measure_congr ?_ - have ha := pullCount_mul a (hK := hK) (ν := ν) (m := m) - have h_best := pullCount_mul (bestArm ν) (hK := hK) (ν := ν) (m := m) - rw [ae_eq_set_iff] - change ∀ᵐ ω ∂((𝔓t).prod _), _ - rw [Measure.ae_prod_iff_ae_ae] + have ha := pullCount_mul h a (hK := hK) (ν := ν) (m := m) + have h_best := pullCount_mul h (bestArm ν) (hK := hK) (ν := ν) (m := m) + rw [ae_eq_set_iff, Measure.ae_prod_iff_ae_ae] · filter_upwards [ha, h_best] with ω ha h_best refine ae_of_all _ fun ω' ↦ ?_ rw [ha, h_best] · simp only [Set.mem_setOf_eq] - let f₁ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ - ∑ s ∈ Icc 1 (pullCount (bestArm ν) (K * m) ω.1), rewardByCount (bestArm ν) s ω - let g₁ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ - ∑ s ∈ Icc 1 (pullCount a (K * m) ω.1), rewardByCount a s ω - let f₂ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ - ∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω - let g₂ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ ∑ s ∈ Icc 1 m, rewardByCount a s ω + let f₁ := fun ω : Ω × (ℕ → Fin K → ℝ) ↦ + ∑ s ∈ Icc 1 (pullCount A (bestArm ν) (K * m) ω.1), rewardByCount A R (bestArm ν) s ω + let g₁ := fun ω : Ω × (ℕ → Fin K → ℝ) ↦ + ∑ s ∈ Icc 1 (pullCount A a (K * m) ω.1), rewardByCount A R a s ω + let f₂ := fun ω : Ω × (ℕ → Fin K → ℝ) ↦ + ∑ s ∈ Icc 1 m, rewardByCount A R (bestArm ν) s ω + let g₂ := fun ω : Ω × (ℕ → Fin K → ℝ) ↦ ∑ s ∈ Icc 1 m, rewardByCount A R a s ω change MeasurableSet {x | f₁ x ≤ g₁ x ↔ f₂ x ≤ g₂ x} have hf₁ : Measurable f₁ := by refine measurable_sum_of_le (n := K * m + 1) - (g := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ pullCount (bestArm ν) (K * m) ω.1) - (f := rewardByCount (bestArm ν)) (fun ω ↦ ?_) + (g := fun ω : Ω × (ℕ → Fin K → ℝ) ↦ pullCount A (bestArm ν) (K * m) ω.1) + (f := rewardByCount A R (bestArm ν)) (fun ω ↦ ?_) (by fun_prop) (by fun_prop) - have h_le := pullCount_le (bestArm ν) (K * m) ω.1 + have h_le := pullCount_le (A := A) (bestArm ν) (K * m) ω.1 grind have hg₁ : Measurable g₁ := by refine measurable_sum_of_le (n := K * m + 1) - (g := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ pullCount a (K * m) ω.1) - (f := rewardByCount a) (fun ω ↦ ?_) (by fun_prop) (by fun_prop) - have h_le := pullCount_le a (K * m) ω.1 + (g := fun ω : Ω × (ℕ → Fin K → ℝ) ↦ pullCount A a (K * m) ω.1) + (f := rewardByCount A R a) (fun ω ↦ ?_) (by fun_prop) (by fun_prop) + have h_le := pullCount_le (A := A) a (K * m) ω.1 grind refine MeasurableSet.iff ?_ ?_ · exact measurableSet_le (by fun_prop) (by fun_prop) · exact measurableSet_le (by fun_prop) (by fun_prop) - _ = (𝔓).real {ω | ∑ s ∈ range m, ω.2 s (bestArm ν) ≤ ∑ s ∈ range m, ω.2 s a} := by + _ = (Bandit.measure (etcAlgorithm hK m) ν).real + {ω | ∑ s ∈ range m, ω.2 s (bestArm ν) ≤ ∑ s ∈ range m, ω.2 s a} := by simp_rw [measureReal_def] congr 1 - have : (𝔓).map (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω, - ∑ s ∈ Icc 1 m, rewardByCount a s ω)) - = (𝔓).map (fun ω ↦ (∑ s ∈ range m, ω.2 s (bestArm ν), ∑ s ∈ range m, ω.2 s a)) := - (identDistrib_aux m (bestArm ν) a).map_eq + have : (𝔓).map (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount A R (bestArm ν) s ω, + ∑ s ∈ Icc 1 m, rewardByCount A R a s ω)) + = (Bandit.measure (etcAlgorithm hK m) ν).map + (fun ω ↦ (∑ s ∈ range m, ω.2 s (bestArm ν), ∑ s ∈ range m, ω.2 s a)) := + (identDistrib_aux h (bestArm ν) a).map_eq rw [Measure.ext_iff] at this have h_meas : MeasurableSet {x : ℝ × ℝ | x.1 ≤ x.2} := measurableSet_le (by fun_prop) (by fun_prop) specialize this {x | x.1 ≤ x.2} h_meas - rwa [Measure.map_apply (by fun_prop) h_meas, Measure.map_apply (by fun_prop) h_meas] at this + rw [Measure.map_apply (by fun_prop) h_meas, Measure.map_apply (by fun_prop) h_meas] at this + convert this _ = (Bandit.streamMeasure ν).real {ω | ∑ s ∈ range m, ω s (bestArm ν) ≤ ∑ s ∈ range m, ω s a} := by simp_rw [measureReal_def] @@ -364,13 +395,17 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i ring /-- Bound on the expectation of the number of pulls of each arm by the ETC algorithm. -/ -lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma expectation_pullCount_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : - 𝔓t[fun ω ↦ (pullCount a n ω : ℝ)] + P[fun ω ↦ (pullCount A a n ω : ℝ)] ≤ m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by - have : (fun ω ↦ (pullCount a n ω : ℝ)) - =ᵐ[𝔓t] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by - filter_upwards [pullCount_of_ge a hm hn] with ω h + have hA := h.measurable_A + have hR := h.measurable_R + have : (fun ω ↦ (pullCount A a n ω : ℝ)) + =ᵐ[P] fun ω ↦ m + (n - K * m) * {ω' | A (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by + filter_upwards [pullCount_of_ge h a hm hn] with ω h simp only [h, Set.indicator_apply, Set.mem_setOf_eq, mul_ite, mul_one, mul_zero, Nat.cast_add, Nat.cast_ite, CharP.cast_eq_zero, add_right_inj] norm_cast @@ -386,21 +421,25 @@ lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - ( simp rw [integral_indicator_const, smul_eq_mul, mul_one] · rw [← neg_mul] - exact prob_arm_mul_eq_le hν a hm + exact prob_arm_mul_eq_le h hν a hm · exact (measurableSet_singleton _).preimage (by fun_prop) /-- Regret bound for the ETC algorithm. -/ -lemma regret_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hm : m ≠ 0) +lemma regret_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hm : m ≠ 0) (n : ℕ) (hn : K * m ≤ n) : - 𝔓t[regret ν n] ≤ ∑ a, gap ν a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4)) := by + P[regret ν A n] ≤ + ∑ a, gap ν a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4)) := by + have hA := h.measurable_A simp_rw [regret_eq_sum_pullCount_mul_gap] rw [integral_finset_sum] - swap; · exact fun i _ ↦ (integrable_pullCount i n).mul_const _ + swap; · exact fun i _ ↦ (integrable_pullCount hA i n).mul_const _ gcongr with a rw [mul_comm (gap _ _), integral_mul_const] gcongr · exact gap_nonneg - · exact expectation_pullCount_le hν a hm hn + · exact expectation_pullCount_le h hν a hm hn end ETC diff --git a/LeanBandits/BanditAlgorithms/UCB.lean b/LeanBandits/BanditAlgorithms/UCB.lean index 0fd5462b..74973c59 100644 --- a/LeanBandits/BanditAlgorithms/UCB.lean +++ b/LeanBandits/BanditAlgorithms/UCB.lean @@ -50,143 +50,161 @@ end Algorithm namespace UCB -variable {hK : 0 < K} {c : ℝ} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] {n : ℕ} {h : ℕ → Fin K × ℝ} +variable {hK : 0 < K} {c : ℝ} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] + {Ω : Type*} {mΩ : MeasurableSpace Ω} + {P : Measure Ω} [IsProbabilityMeasure P] + {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} + {n : ℕ} {ω : Ω} /-- The exploration bonus of the UCB algorithm, which corresponds to the width of a confidence interval. -/ -noncomputable def ucbWidth (c : ℝ) (a : Fin K) (n : ℕ) (h : ℕ → Fin K × ℝ) : ℝ := - √(c * log (n + 1) / pullCount a n h) +noncomputable def ucbWidth (A : ℕ → Ω → Fin K) (c : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ := + √(c * log (n + 1) / pullCount A a n ω) @[fun_prop] -lemma measurable_ucbWidth (c : ℝ) (a : Fin K) : Measurable (ucbWidth c a n) := by +lemma measurable_ucbWidth (hA : ∀ n, Measurable (A n)) (c : ℝ) (a : Fin K) : + Measurable (ucbWidth A c a n) := by unfold ucbWidth fun_prop -lemma ucbWidth_eq_ucbWidth' (c : ℝ) (a : Fin K) (n : ℕ) (h : ℕ → Fin K × ℝ) (hn : n ≠ 0) : - ucbWidth c a n h = ucbWidth' c (n - 1) (fun i ↦ h i) a := by - simp only [ucbWidth, pullCount_eq_pullCount' hn, Nat.cast_nonneg, sqrt_div', ucbWidth'] +lemma ucbWidth_eq_ucbWidth' (c : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) (hn : n ≠ 0) : + ucbWidth A c a n ω = ucbWidth' c (n - 1) (IsAlgEnvSeq.hist A R (n - 1) ω) a := by + simp only [ucbWidth, pullCount_eq_pullCount' (A := A) (R' := R) hn, Nat.cast_nonneg, sqrt_div', + ucbWidth'] congr 4 norm_cast grind -local notation "𝔓t" => Bandit.trajMeasure (ucbAlgorithm hK c) ν -local notation "𝔓" => Bandit.measure (ucbAlgorithm hK c) ν +local notation "𝔓" => P.prod (Bandit.streamMeasure ν) -lemma arm_zero : arm 0 =ᵐ[𝔓t] fun _ ↦ ⟨0, hK⟩ := by +lemma arm_zero [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) : + A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - exact arm_zero_detAlgorithm + exact h.action_zero_detAlgorithm -lemma arm_ae_eq_ucbNextArm (n : ℕ) : - arm (n + 1) =ᵐ[𝔓t] fun h ↦ nextArm hK c n (fun i ↦ h i) := by +lemma arm_ae_eq_ucbNextArm [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (n : ℕ) : + A (n + 1) =ᵐ[P] fun ω ↦ nextArm hK c n (IsAlgEnvSeq.hist A R n ω) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - exact arm_detAlgorithm_ae_eq n + exact h.action_detAlgorithm_ae_eq n -lemma arm_ae_all_eq : - ∀ᵐ h ∂𝔓t, arm 0 h = ⟨0, hK⟩ ∧ ∀ n, arm (n + 1) h = nextArm hK c n (fun i ↦ h i) := by +lemma arm_ae_all_eq [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) : + ∀ᵐ h ∂P, A 0 h = ⟨0, hK⟩ ∧ ∀ n, A (n + 1) h = nextArm hK c n (IsAlgEnvSeq.hist A R n h) := by rw [eventually_and, ae_all_iff] - exact ⟨arm_zero, arm_ae_eq_ucbNextArm⟩ + exact ⟨arm_zero h, arm_ae_eq_ucbNextArm h⟩ -lemma ucbIndex_le_ucbIndex_arm (a : Fin K) (hn : K ≤ n) : - ∀ᵐ h ∂𝔓t, empMean a n h + ucbWidth c a n h ≤ - empMean (arm n h) n h + ucbWidth c (arm n h) n h := by - filter_upwards [arm_ae_eq_ucbNextArm (n - 1)] with h h_arm +lemma ucbIndex_le_ucbIndex_arm [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) (hn : K ≤ n) : + ∀ᵐ h ∂P, empMean A R a n h + ucbWidth A c a n h ≤ + empMean A R (A n h) n h + ucbWidth A c (A n h) n h := by + filter_upwards [arm_ae_eq_ucbNextArm h (n - 1)] with h h_arm have : n - 1 + 1 = n := by grind have h_not_lt : ¬ n - 1 < K - 1 := by grind simp only [this, nextArm, h_not_lt, ↓reduceIte] at h_arm have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK simp_rw [h_arm, empMean_eq_empMean' (by grind : n ≠ 0), - ucbWidth_eq_ucbWidth' _ _ _ _ (by grind : n ≠ 0)] + ucbWidth_eq_ucbWidth' (A := A) (R := R) _ _ _ _ (by grind : n ≠ 0)] exact isMaxOn_measurableArgmax (fun h a ↦ empMean' (n - 1) h a + ucbWidth' c (n - 1) h a) - (fun i ↦ h i) a + (IsAlgEnvSeq.hist A R (n - 1) h) a -lemma forall_arm_eq_mod_of_lt : - ∀ᵐ h ∂𝔓t, ∀ n < K, arm n h = ⟨n % K, Nat.mod_lt _ hK⟩ := by +lemma forall_arm_eq_mod_of_lt [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) : + ∀ᵐ h ∂P, ∀ n < K, A n h = ⟨n % K, Nat.mod_lt _ hK⟩ := by simp_rw [ae_all_iff] intro n hn induction n with - | zero => exact arm_zero + | zero => exact arm_zero h | succ n _ => - filter_upwards [arm_ae_eq_ucbNextArm n] with h h_eq + filter_upwards [arm_ae_eq_ucbNextArm h n] with h h_eq rw [h_eq, nextArm, if_pos] grind -lemma forall_ucbIndex_le_ucbIndex_arm (a : Fin K) : - ∀ᵐ h ∂𝔓t, ∀ n, K ≤ n → - empMean a n h + ucbWidth c a n h ≤ empMean (arm n h) n h + ucbWidth c (arm n h) n h := by +lemma forall_ucbIndex_le_ucbIndex_arm [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) : + ∀ᵐ h ∂P, ∀ n, K ≤ n → + empMean A R a n h + ucbWidth A c a n h ≤ + empMean A R (A n h) n h + ucbWidth A c (A n h) n h := by simp_rw [ae_all_iff] - exact fun _ ↦ ucbIndex_le_ucbIndex_arm a - -lemma forall_arm_prop : - ∀ᵐ h ∂𝔓t, - (∀ n < K, arm n h = ⟨n % K, Nat.mod_lt _ hK⟩) ∧ - (∀ n, K ≤ n → ∀ a, empMean a n h + ucbWidth c a n h ≤ - empMean (arm n h) n h + ucbWidth c (arm n h) n h) := by + exact fun _ ↦ ucbIndex_le_ucbIndex_arm h a + +lemma forall_arm_prop [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) : + ∀ᵐ h ∂P, + (∀ n < K, A n h = ⟨n % K, Nat.mod_lt _ hK⟩) ∧ + (∀ n, K ≤ n → ∀ a, empMean A R a n h + ucbWidth A c a n h ≤ + empMean A R (A n h) n h + ucbWidth A c (A n h) n h) := by simp only [eventually_and] constructor - · exact forall_arm_eq_mod_of_lt + · exact forall_arm_eq_mod_of_lt h · simp_rw [ae_all_iff] intro n hn a - have h_ae := forall_ucbIndex_le_ucbIndex_arm (ν := ν) (c := c) (hK := hK) a + have h_ae := forall_ucbIndex_le_ucbIndex_arm h a simp_rw [ae_all_iff] at h_ae exact h_ae n hn -lemma pullCount_eq_of_time_eq (a : Fin K) : - ∀ᵐ ω ∂𝔓t, pullCount a K ω = 1 := by - filter_upwards [forall_arm_eq_mod_of_lt] with h h_eq +lemma pullCount_eq_of_time_eq [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) : + ∀ᵐ ω ∂P, pullCount A a K ω = 1 := by + filter_upwards [forall_arm_eq_mod_of_lt h] with h h_eq rw [pullCount_eq_sum] conv_rhs => rw [← sum_mod_range hK a] refine Finset.sum_congr rfl fun s hs ↦ ?_ congr exact h_eq s (by grind) -lemma time_gt_of_pullCount_gt_one (a : Fin K) : - ∀ᵐ ω ∂𝔓t, ∀ n, 1 < pullCount a n ω → K < n := by - filter_upwards [pullCount_eq_of_time_eq a] with h h_eq n hn +lemma time_gt_of_pullCount_gt_one [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) : + ∀ᵐ ω ∂P, ∀ n, 1 < pullCount A a n ω → K < n := by + filter_upwards [pullCount_eq_of_time_eq h a] with h h_eq n hn rw [← h_eq] at hn by_contra! h_lt exact hn.not_ge (pullCount_mono _ h_lt _) -lemma pullCount_pos_of_time_ge : - ∀ᵐ ω ∂𝔓t, ∀ n, K ≤ n → ∀ b : Fin K, 0 < pullCount b n ω := by - have h_ae a := pullCount_eq_of_time_eq (ν := ν) (c := c) (hK := hK) a +lemma pullCount_pos_of_time_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) : + ∀ᵐ ω ∂P, ∀ n, K ≤ n → ∀ b : Fin K, 0 < pullCount A b n ω := by + have h_ae a := pullCount_eq_of_time_eq h a rw [← ae_all_iff] at h_ae filter_upwards [h_ae] with ω hω n hn a refine Nat.one_pos.trans_le ?_ rw [← hω a] exact pullCount_mono _ hn _ -lemma pullCount_pos_of_pullCount_gt_one (a : Fin K) : - ∀ᵐ ω ∂𝔓t, ∀ n, 1 < pullCount a n ω → ∀ b : Fin K, 0 < pullCount b n ω := by - filter_upwards [time_gt_of_pullCount_gt_one a, pullCount_pos_of_time_ge] with ω h1 h2 n h_gt a +lemma pullCount_pos_of_pullCount_gt_one [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) : + ∀ᵐ ω ∂P, ∀ n, 1 < pullCount A a n ω → ∀ b : Fin K, 0 < pullCount A b n ω := by + filter_upwards [time_gt_of_pullCount_gt_one h a, pullCount_pos_of_time_ge h] with ω h1 h2 n h_gt a exact h2 n (h1 n h_gt).le a omit [IsMarkovKernel ν] in lemma gap_arm_le_two_mul_ucbWidth [Nonempty (Fin K)] - (h_best : (ν (bestArm ν))[id] ≤ empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h) - (h_arm : empMean (arm n h) n h - ucbWidth c (arm n h) n h ≤ (ν (arm n h))[id]) - (h_le : empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h ≤ - empMean (arm n h) n h + ucbWidth c (arm n h) n h) : - gap ν (arm n h) ≤ 2 * ucbWidth c (arm n h) n h := by + (h_best : (ν (bestArm ν))[id] ≤ empMean A R (bestArm ν) n ω + ucbWidth A c (bestArm ν) n ω) + (h_arm : empMean A R (A n ω) n ω - ucbWidth A c (A n ω) n ω ≤ (ν (A n ω))[id]) + (h_le : empMean A R (bestArm ν) n ω + ucbWidth A c (bestArm ν) n ω ≤ + empMean A R (A n ω) n ω + ucbWidth A c (A n ω) n ω) : + gap ν (A n ω) ≤ 2 * ucbWidth A c (A n ω) n ω := by rw [gap_eq_bestArm_sub, sub_le_iff_le_add'] calc (ν (bestArm ν))[id] - _ ≤ empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h := h_best - _ ≤ empMean (arm n h) n h + ucbWidth c (arm n h) n h := h_le - _ ≤ (ν (arm n h))[id] + 2 * ucbWidth c (arm n h) n h := by + _ ≤ empMean A R (bestArm ν) n ω + ucbWidth A c (bestArm ν) n ω := h_best + _ ≤ empMean A R (A n ω) n ω + ucbWidth A c (A n ω) n ω := h_le + _ ≤ (ν (A n ω))[id] + 2 * ucbWidth A c (A n ω) n ω := by rw [two_mul, ← add_assoc] gcongr rwa [sub_le_iff_le_add] at h_arm omit [IsMarkovKernel ν] in lemma pullCount_arm_le [Nonempty (Fin K)] (hc : 0 ≤ c) - (h_best : (ν (bestArm ν))[id] ≤ empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h) - (h_arm : empMean (arm n h) n h - ucbWidth c (arm n h) n h ≤ (ν (arm n h))[id]) - (h_le : empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h ≤ - empMean (arm n h) n h + ucbWidth c (arm n h) n h) - (h_gap_pos : 0 < gap ν (arm n h)) (h_pull_pos : 0 < pullCount (arm n h) n h) : - pullCount (arm n h) n h ≤ 4 * c * log (n + 1) / gap ν (arm n h) ^ 2 := by + (h_best : (ν (bestArm ν))[id] ≤ empMean A R (bestArm ν) n ω + ucbWidth A c (bestArm ν) n ω) + (h_arm : empMean A R (A n ω) n ω - ucbWidth A c (A n ω) n ω ≤ (ν (A n ω))[id]) + (h_le : empMean A R (bestArm ν) n ω + ucbWidth A c (bestArm ν) n ω ≤ + empMean A R (A n ω) n ω + ucbWidth A c (A n ω) n ω) + (h_gap_pos : 0 < gap ν (A n ω)) (h_pull_pos : 0 < pullCount A (A n ω) n ω) : + pullCount A (A n ω) n ω ≤ 4 * c * log (n + 1) / gap ν (A n ω) ^ 2 := by have h_gap_le := gap_arm_le_two_mul_ucbWidth h_best h_arm h_le rw [ucbWidth] at h_gap_le - have h2 : (gap ν (arm n h)) ^ 2 ≤ (2 * √(c * log (n + 1) / pullCount (arm n h) n h)) ^ 2 := by + have h2 : (gap ν (A n ω)) ^ 2 ≤ (2 * √(c * log (n + 1) / pullCount A (A n ω) n ω)) ^ 2 := by gcongr rw [mul_pow, sq_sqrt] at h2 · have : (2 : ℝ) ^ 2 = 4 := by norm_num @@ -196,33 +214,41 @@ lemma pullCount_arm_le [Nonempty (Fin K)] (hc : 0 ≤ c) · have : 0 ≤ log (n + 1) := by simp [log_nonneg] positivity -lemma todo (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +variable [StandardBorelSpace Ω] [Nonempty Ω] + +lemma todo [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n k : ℕ) (hk : k ≠ 0) : - 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} ≤ + 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2) := by - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + have hA := h.measurable_A + have hR := h.measurable_R have h_meas : MeasurableSet {ω | ω / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := measurableSet_le (by fun_prop) measurable_const have h_log_nonneg : 0 ≤ log (n + 1) := log_nonneg (by simp) calc - 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} - _ = ((𝔓).map (fun ω ↦ ∑ m ∈ Icc 1 k, rewardByCount a m ω)) + 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} + _ = ((𝔓).map (fun ω ↦ ∑ m ∈ Icc 1 k, rewardByCount A R a m ω)) {ω | ω / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := by rw [Measure.map_apply (by fun_prop) h_meas] rfl - _ = ((𝔓).map (fun ω ↦ ∑ s ∈ range k, ω.2 s a)) + _ = ((Bandit.measure (ucbAlgorithm hK c) ν).map (fun ω ↦ ∑ s ∈ range k, ω.2 s a)) {ω | ω / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := by - rw [IdentDistrib.map_eq (identDistrib_sum_Icc_rewardByCount k a)] - _ = 𝔓 {ω | (∑ s ∈ range k, ω.2 s a) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := by + rw [IdentDistrib.map_eq (identDistrib_sum_Icc_rewardByCount h k a)] + _ = (Bandit.measure (ucbAlgorithm hK c) ν) + {ω | (∑ s ∈ range k, ω.2 s a) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := by rw [Measure.map_apply (by fun_prop) h_meas] rfl - _ = 𝔓 {ω | (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) / k ≤ - √(c * log (n + 1) / k)} := by + _ = (Bandit.measure (ucbAlgorithm hK c) ν) + {ω | (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) / k ≤ - √(c * log (n + 1) / k)} := by congr with ω field_simp rw [Finset.sum_sub_distrib] simp grind - _ = 𝔓 {ω | (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) ≤ - √(c * k * log (n + 1))} := by + _ = (Bandit.measure (ucbAlgorithm hK c) ν) + {ω | (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) ≤ - √(c * k * log (n + 1))} := by congr with ω field_simp congr! 2 @@ -253,33 +279,40 @@ lemma todo (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) · norm_cast · field -lemma todo' (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma todo' [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n k : ℕ) (hk : k ≠ 0) : - 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - √(c * log (n + 1) / k)} ≤ + 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k - √(c * log (n + 1) / k)} ≤ 1 / (n + 1) ^ (c / 2) := by + have hA := h.measurable_A + have hR := h.measurable_R have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK have h_meas : MeasurableSet {ω |(ν a)[id] ≤ ω / k - √(c * log (n + 1) / k)} := measurableSet_le (by fun_prop) (by fun_prop) have h_log_nonneg : 0 ≤ log (n + 1) := log_nonneg (by simp) calc - 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - √(c * log (n + 1) / k)} - _ = ((𝔓).map (fun ω ↦ ∑ m ∈ Icc 1 k, rewardByCount a m ω)) + 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k - √(c * log (n + 1) / k)} + _ = ((𝔓).map (fun ω ↦ ∑ m ∈ Icc 1 k, rewardByCount A R a m ω)) {ω | (ν a)[id] ≤ ω / k - √(c * log (n + 1) / k)} := by rw [Measure.map_apply (by fun_prop) h_meas] rfl - _ = ((𝔓).map (fun ω ↦ ∑ s ∈ range k, ω.2 s a)) + _ = ((Bandit.measure (ucbAlgorithm hK c) ν).map (fun ω ↦ ∑ s ∈ range k, ω.2 s a)) {ω | (ν a)[id] ≤ ω / k - √(c * log (n + 1) / k)} := by - rw [IdentDistrib.map_eq (identDistrib_sum_Icc_rewardByCount k a)] - _ = 𝔓 {ω | (ν a)[id] ≤ (∑ s ∈ range k, ω.2 s a) / k - √(c * log (n + 1) / k)} := by + rw [IdentDistrib.map_eq (identDistrib_sum_Icc_rewardByCount h k a)] + _ = (Bandit.measure (ucbAlgorithm hK c) ν) + {ω | (ν a)[id] ≤ (∑ s ∈ range k, ω.2 s a) / k - √(c * log (n + 1) / k)} := by rw [Measure.map_apply (by fun_prop) h_meas] rfl - _ = 𝔓 {ω | √(c * log (n + 1) / k) ≤ (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) / k} := by + _ = (Bandit.measure (ucbAlgorithm hK c) ν) + {ω | √(c * log (n + 1) / k) ≤ (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) / k} := by congr with ω field_simp rw [Finset.sum_sub_distrib] simp grind - _ = 𝔓 {ω | √(c * k * log (n + 1)) ≤ (∑ s ∈ range k, (ω.2 s a - (ν a)[id]))} := by + _ = (Bandit.measure (ucbAlgorithm hK c) ν) + {ω | √(c * k * log (n + 1)) ≤ (∑ s ∈ range k, (ω.2 s a - (ν a)[id]))} := by congr with ω field_simp congr! 1 @@ -310,16 +343,20 @@ lemma todo' (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) · norm_cast · field -lemma prob_ucbIndex_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma prob_ucbIndex_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : - 𝔓t {h | 0 < pullCount a n h ∧ empMean a n h + ucbWidth c a n h ≤ (ν a)[id]} ≤ + P {h | 0 < pullCount A a n h ∧ empMean A R a n h + ucbWidth A c a n h ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by + have hA := h.measurable_A + have hR := h.measurable_R -- extend the probability space - suffices 𝔓 {ω | 0 < pullCount a n ω.1 ∧ - empMean a n ω.1 + ucbWidth c a n ω.1 ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2 - 1) by - rwa [← Bandit.fst_measure (ucbAlgorithm hK c) ν, Measure.fst_apply] - change MeasurableSet ({h | 0 < pullCount a n h} - ∩ {h | empMean a n h + ucbWidth c a n h ≤ ∫ (x : ℝ), id x ∂ν a}) + suffices 𝔓 {ω | 0 < pullCount A a n ω.1 ∧ + empMean A R a n ω.1 + ucbWidth A c a n ω.1 ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2 - 1) by + rwa [← Measure.fst_prod (μ := P) (ν := Bandit.streamMeasure ν), Measure.fst_apply] + change MeasurableSet ({h | 0 < pullCount A a n h} + ∩ {h | empMean A R a n h + ucbWidth A c a n h ≤ ∫ (x : ℝ), id x ∂ν a}) refine MeasurableSet.inter ?_ ?_ · exact measurableSet_lt (by fun_prop) (by fun_prop) · exact measurableSet_le (by fun_prop) (by fun_prop) @@ -327,16 +364,16 @@ lemma prob_ucbIndex_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id] unfold empMean ucbWidth simp_rw [← sum_rewardByCount_eq_sumRewards] calc - 𝔓 {ω | 0 < pullCount a n ω.1 ∧ - (∑ m ∈ Icc 1 (pullCount a n ω.1), rewardByCount a m ω) / pullCount a n ω.1 + - √(c * log (↑n + 1) / pullCount a n ω.1) ≤ (ν a)[id]} + 𝔓 {ω | 0 < pullCount A a n ω.1 ∧ + (∑ m ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a m ω) / pullCount A a n ω.1 + + √(c * log (↑n + 1) / pullCount A a n ω.1) ≤ (ν a)[id]} -- list the possible values of `pullCount a n ω.1` - _ ≤ 𝔓 {ω | ∃ k ≤ n, 0 < k ∧ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + + _ ≤ 𝔓 {ω | ∃ k ≤ n, 0 < k ∧ (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k + √(c * log (↑n + 1) / k) ≤ (ν a)[id]} := by refine measure_mono fun ω hω ↦ ?_ simp only [Nat.cast_nonneg, sqrt_div', id_eq, Set.mem_setOf_eq] at hω ⊢ - exact ⟨pullCount a n ω.1, pullCount_le _ _ _, hω⟩ - _ = 𝔓 (⋃ k ∈ Icc 1 n, {ω |(∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + + exact ⟨pullCount A a n ω.1, pullCount_le _ _ _, hω⟩ + _ = 𝔓 (⋃ k ∈ Icc 1 n, {ω |(∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k + √(c * log (↑n + 1) / k) ≤ (ν a)[id]}) := by congr 1 ext ω @@ -344,11 +381,11 @@ lemma prob_ucbIndex_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id] grind -- Union bound over the possible values of `pullCount a n ω.1` _ ≤ ∑ k ∈ Icc 1 n, - 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + √(c * log (↑n + 1) / k) ≤ (ν a)[id]} := + 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k + √(c * log (↑n + 1) / k) ≤ (ν a)[id]} := measure_biUnion_finset_le _ _ _ ≤ ∑ k ∈ Icc 1 n, (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by gcongr with k hk - exact todo hν hc a n k (by grind) + exact todo h hν hc a n k (by grind) _ ≤ (n + 1) * (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by simp only [one_div, sum_const, Nat.card_Icc, add_tsub_cancel_right, nsmul_eq_mul, mul_one] rw [div_eq_mul_inv ((n : ℝ≥0∞) + 1)] @@ -359,16 +396,20 @@ lemma prob_ucbIndex_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id] rw [ENNReal.rpow_sub _ _ (by simp) (by finiteness), ENNReal.rpow_one, div_eq_mul_inv, ENNReal.div_eq_inv_mul, ENNReal.mul_inv (by simp) (by simp), inv_inv] -lemma prob_ucbIndex_ge (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma prob_ucbIndex_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : - 𝔓t {h | 0 < pullCount a n h ∧ - (ν a)[id] ≤ empMean a n h - ucbWidth c a n h} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by + P {h | 0 < pullCount A a n h ∧ + (ν a)[id] ≤ empMean A R a n h - ucbWidth A c a n h} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by + have hA := h.measurable_A + have hR := h.measurable_R -- extend the probability space - suffices 𝔓 {ω | 0 < pullCount a n ω.1 ∧ - (ν a)[id] ≤ empMean a n ω.1 - ucbWidth c a n ω.1} ≤ 1 / (n + 1) ^ (c / 2 - 1) by - rwa [← Bandit.fst_measure (ucbAlgorithm hK c) ν, Measure.fst_apply] - change MeasurableSet ({h | 0 < pullCount a n h} - ∩ {h | (ν a)[id] ≤ empMean a n h - ucbWidth c a n h}) + suffices 𝔓 {ω | 0 < pullCount A a n ω.1 ∧ + (ν a)[id] ≤ empMean A R a n ω.1 - ucbWidth A c a n ω.1} ≤ 1 / (n + 1) ^ (c / 2 - 1) by + rwa [← Measure.fst_prod (μ := P) (ν := Bandit.streamMeasure ν), Measure.fst_apply] + change MeasurableSet ({h | 0 < pullCount A a n h} + ∩ {h | (ν a)[id] ≤ empMean A R a n h - ucbWidth A c a n h}) refine MeasurableSet.inter ?_ ?_ · exact measurableSet_lt (by fun_prop) (by fun_prop) · exact measurableSet_le (by fun_prop) (by fun_prop) @@ -376,16 +417,16 @@ lemma prob_ucbIndex_ge (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id] unfold empMean ucbWidth simp_rw [← sum_rewardByCount_eq_sumRewards] calc - 𝔓 {ω | 0 < pullCount a n ω.1 ∧ - (ν a)[id] ≤ (∑ m ∈ Icc 1 (pullCount a n ω.1), rewardByCount a m ω) / pullCount a n ω.1 - - √(c * log (↑n + 1) / pullCount a n ω.1)} - -- list the possible values of `pullCount a n ω.1` - _ ≤ 𝔓 {ω | ∃ k ≤ n, 0 < k ∧ (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - + 𝔓 {ω | 0 < pullCount A a n ω.1 ∧ + (ν a)[id] ≤ (∑ m ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a m ω) / pullCount A a n ω.1 - + √(c * log (↑n + 1) / pullCount A a n ω.1)} + -- list the possible values of `pullCount A a n ω.1` + _ ≤ 𝔓 {ω | ∃ k ≤ n, 0 < k ∧ (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k - √(c * log (↑n + 1) / k)} := by refine measure_mono fun ω hω ↦ ?_ simp only [Nat.cast_nonneg, sqrt_div', id_eq, Set.mem_setOf_eq] at hω ⊢ - exact ⟨pullCount a n ω.1, pullCount_le _ _ _, hω⟩ - _ = 𝔓 (⋃ k ∈ Icc 1 n, {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - + exact ⟨pullCount A a n ω.1, pullCount_le _ _ _, hω⟩ + _ = 𝔓 (⋃ k ∈ Icc 1 n, {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k - √(c * log (↑n + 1) / k)}) := by congr 1 ext ω @@ -393,11 +434,11 @@ lemma prob_ucbIndex_ge (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id] grind -- Union bound over the possible values of `pullCount a n ω.1` _ ≤ ∑ k ∈ Icc 1 n, - 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - √(c * log (↑n + 1) / k)} := + 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k - √(c * log (↑n + 1) / k)} := measure_biUnion_finset_le _ _ _ ≤ ∑ k ∈ Icc 1 n, (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by gcongr with k hk - exact todo' hν hc a n k (by grind) + exact todo' h hν hc a n k (by grind) _ ≤ (n + 1) * (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by simp only [one_div, sum_const, Nat.card_Icc, add_tsub_cancel_right, nsmul_eq_mul, mul_one] rw [div_eq_mul_inv ((n : ℝ≥0∞) + 1)] @@ -408,55 +449,60 @@ lemma prob_ucbIndex_ge (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id] rw [ENNReal.rpow_sub _ _ (by simp) (by finiteness), ENNReal.rpow_one, div_eq_mul_inv, ENNReal.div_eq_inv_mul, ENNReal.mul_inv (by simp) (by simp), inv_inv] -lemma probReal_ucbIndex_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma probReal_ucbIndex_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : - (𝔓t).real {h | 0 < pullCount a n h ∧ empMean a n h + ucbWidth c a n h ≤ (ν a)[id]} ≤ + P.real {h | 0 < pullCount A a n h ∧ empMean A R a n h + ucbWidth A c a n h ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by rw [measureReal_def] - grw [prob_ucbIndex_le hν hc a n] + grw [prob_ucbIndex_le h hν hc a n] swap; · finiteness simp only [one_div, ENNReal.toReal_inv] rw [← ENNReal.toReal_rpow] norm_cast -lemma probReal_ucbIndex_ge (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma probReal_ucbIndex_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : - (𝔓t).real {h | 0 < pullCount a n h ∧ - (ν a)[id] ≤ empMean a n h - ucbWidth c a n h} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by + P.real {h | 0 < pullCount A a n h ∧ + (ν a)[id] ≤ empMean A R a n h - ucbWidth A c a n h} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by rw [measureReal_def] - grw [prob_ucbIndex_ge hν hc a n] + grw [prob_ucbIndex_ge h hν hc a n] swap; · finiteness simp only [one_div, ENNReal.toReal_inv] rw [← ENNReal.toReal_rpow] norm_cast -lemma pullCount_le_add (a : Fin K) (n C : ℕ) (ω : ℕ → Fin K × ℝ) : - pullCount a n ω ≤ C + 1 + - ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω}.indicator 1 s := by +omit [Nonempty Ω] in +lemma pullCount_le_add [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (ω : Ω) : + pullCount A a n ω ≤ C + 1 + + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s := by rw [pullCount_eq_sum] - calc ∑ s ∈ range n, if arm s ω = a then 1 else 0 - _ ≤ ∑ s ∈ range n, ({s | arm s ω = a ∧ pullCount a s ω ≤ C}.indicator 1 s + - {s | arm s ω = a ∧ C < pullCount a s ω}.indicator 1 s) := by + calc ∑ s ∈ range n, if A s ω = a then 1 else 0 + _ ≤ ∑ s ∈ range n, ({s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s + + {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s) := by gcongr with s hs simp [Set.indicator_apply] grind - _ = ∑ s ∈ range n, {s | arm s ω = a ∧ pullCount a s ω ≤ C}.indicator 1 s + - ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω}.indicator 1 s := by + _ = ∑ s ∈ range n, {s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s + + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s := by rw [Finset.sum_add_distrib] - _ ≤ C + 1 + ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω}.indicator 1 s := by + _ ≤ C + 1 + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s := by gcongr - have h_le n : ∑ s ∈ range n, {s | arm s ω = a ∧ pullCount a s ω ≤ C}.indicator 1 s ≤ - pullCount a n ω := by + have h_le n : ∑ s ∈ range n, {s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s ≤ + pullCount A a n ω := by rw [pullCount_eq_sum] gcongr with s hs - simp only [Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply, arm, action] + simp only [Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] grind induction n with | zero => simp | succ n hn => rw [Finset.sum_range_succ] - rcases le_or_gt (pullCount a n ω) C with h_pc | h_pc - · have hn' : ∑ s ∈ range n, {s | arm s ω = a ∧ pullCount a s ω ≤ C}.indicator 1 s ≤ C := + rcases le_or_gt (pullCount A a n ω) C with h_pc | h_pc + · have hn' : ∑ s ∈ range n, {s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s ≤ C := (h_le n).trans h_pc grw [hn'] gcongr @@ -465,39 +511,39 @@ lemma pullCount_le_add (a : Fin K) (n C : ℕ) (ω : ℕ → Fin K × ℝ) : · refine le_trans ?_ hn simp [h_pc] -omit [IsMarkovKernel ν] in -lemma pullCount_le_add_three [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (ω : ℕ → Fin K × ℝ) : - pullCount a n ω ≤ C + 1 + - ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω ∧ - (ν (bestArm ν))[id] ≤ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω ∧ - empMean (arm s ω) s ω - ucbWidth c (arm s ω) s ω ≤ (ν (arm s ω))[id]}.indicator 1 s + +omit [StandardBorelSpace Ω] [Nonempty Ω] [IsMarkovKernel ν] in +lemma pullCount_le_add_three [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (ω : Ω) : + pullCount A a n ω ≤ C + 1 + + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω ∧ + (ν (bestArm ν))[id] ≤ empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω ∧ + empMean A R (A s ω) s ω - ucbWidth A c (A s ω) s ω ≤ (ν (A s ω))[id]}.indicator 1 s + ∑ s ∈ range n, - {s | C < pullCount a s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < + {s | C < pullCount A a s ω ∧ empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω < (ν (bestArm ν))[id]}.indicator 1 s + ∑ s ∈ range n, - {s | C < pullCount a s ω ∧ (ν a)[id] < - empMean a s ω - ucbWidth c a s ω}.indicator 1 s := by + {s | C < pullCount A a s ω ∧ (ν a)[id] < + empMean A R a s ω - ucbWidth A c a s ω}.indicator 1 s := by refine (pullCount_le_add a n C ω).trans ?_ simp_rw [add_assoc] gcongr simp_rw [← add_assoc] - let A := {s | arm s ω = a ∧ C < pullCount a s ω} - let B := {s | arm s ω = a ∧ C < pullCount a s ω ∧ - (ν (bestArm ν))[id] ≤ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω ∧ - empMean (arm s ω) s ω - ucbWidth c (arm s ω) s ω ≤ (ν (arm s ω))[id]} - let C' := {s | C < pullCount a s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < - (ν (bestArm ν))[id]} - let D := {s | C < pullCount a s ω ∧ (ν a)[id] < - empMean a s ω - ucbWidth c a s ω} - change ∑ s ∈ range n, A.indicator 1 s ≤ + let A' := {s | A s ω = a ∧ C < pullCount A a s ω} + let B := {s | A s ω = a ∧ C < pullCount A a s ω ∧ + (ν (bestArm ν))[id] ≤ empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω ∧ + empMean A R (A s ω) s ω - ucbWidth A c (A s ω) s ω ≤ (ν (A s ω))[id]} + let C' := {s | C < pullCount A a s ω ∧ + empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω < (ν (bestArm ν))[id]} + let D := {s | C < pullCount A a s ω ∧ (ν a)[id] < + empMean A R a s ω - ucbWidth A c a s ω} + change ∑ s ∈ range n, A'.indicator 1 s ≤ ∑ s ∈ range n, B.indicator 1 s + ∑ s ∈ range n, C'.indicator 1 s + ∑ s ∈ range n, D.indicator 1 s - have h_union : A ⊆ B ∪ C' ∪ D := by simp [A, B, C', D]; grind + have h_union : A' ⊆ B ∪ C' ∪ D := by simp [A', B, C', D]; grind calc - (∑ s ∈ range n, A.indicator 1 s) + (∑ s ∈ range n, A'.indicator 1 s) _ ≤ (∑ s ∈ range n, (B ∪ C' ∪ D).indicator (fun _ ↦ (1 : ℕ)) s) := by gcongr with n hn - by_cases h : n ∈ A + by_cases h : n ∈ A' · have : n ∈ B ∪ C' ∪ D := h_union h simp [h, this] · simp [h] @@ -509,33 +555,40 @@ lemma pullCount_le_add_three [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (ω : ∑ s ∈ range n, D.indicator 1 s := by rw [Finset.sum_add_distrib, Finset.sum_add_distrib] -lemma pullCount_le_add_three_ae [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (hC : C ≠ 0) : - ∀ᵐ ω ∂𝔓t, - pullCount a n ω ≤ C + 1 + - ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω ∧ - (ν (bestArm ν))[id] ≤ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω ∧ - empMean (arm s ω) s ω - ucbWidth c (arm s ω) s ω ≤ (ν (arm s ω))[id]}.indicator 1 s + +omit [StandardBorelSpace Ω] [Nonempty Ω] in +lemma pullCount_le_add_three_ae [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (a : Fin K) (n C : ℕ) (hC : C ≠ 0) : + ∀ᵐ ω ∂P, + pullCount A a n ω ≤ C + 1 + + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω ∧ + (ν (bestArm ν))[id] ≤ empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω ∧ + empMean A R (A s ω) s ω - ucbWidth A c (A s ω) s ω ≤ (ν (A s ω))[id]}.indicator 1 s + ∑ s ∈ range n, - {s | 0 < pullCount (bestArm ν) s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < - (ν (bestArm ν))[id]}.indicator 1 s + + {s | 0 < pullCount A (bestArm ν) s ω ∧ + empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω < + (ν (bestArm ν))[id]}.indicator 1 s + ∑ s ∈ range n, - {s | 0 < pullCount a s ω ∧ (ν a)[id] < - empMean a s ω - ucbWidth c a s ω}.indicator 1 s := by - filter_upwards [pullCount_pos_of_pullCount_gt_one a] with ω hω - refine (pullCount_le_add_three a n C ω (ν := ν) (c := c)).trans ?_ + {s | 0 < pullCount A a s ω ∧ (ν a)[id] < + empMean A R a s ω - ucbWidth A c a s ω}.indicator 1 s := by + filter_upwards [pullCount_pos_of_pullCount_gt_one h a] with ω hω + refine (pullCount_le_add_three (R := R) a n C ω (ν := ν) (c := c)).trans ?_ gcongr 5 with k hk j k hk j · gcongr 1 exact fun h_gt ↦ hω _ (lt_of_le_of_lt (by grind) h_gt) _ · exact fun h_gt ↦ hω _ (lt_of_le_of_lt (by grind) h_gt) _ -lemma some_sum_eq_zero [Nonempty (Fin K)] (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 < gap ν a) (n C : ℕ) +omit [StandardBorelSpace Ω] [Nonempty Ω] in +lemma some_sum_eq_zero [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 < gap ν a) (n C : ℕ) (hC : C ≠ 0) (hC' : 4 * c * log (n + 1) / gap ν a ^ 2 ≤ C) : - ∀ᵐ ω ∂𝔓t, - ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω ∧ - (ν (bestArm ν))[id] ≤ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω ∧ - empMean (arm s ω) s ω - ucbWidth c (arm s ω) s ω ≤ (ν (arm s ω))[id]}.indicator 1 s = 0 := by - have h_ae := forall_ucbIndex_le_ucbIndex_arm (bestArm ν) (ν := ν) (c := c) (hK := hK) - have h_gt := time_gt_of_pullCount_gt_one a (ν := ν) (c := c) (hK := hK) + ∀ᵐ ω ∂P, + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω ∧ + (ν (bestArm ν))[id] ≤ empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω ∧ + empMean A R (A s ω) s ω - ucbWidth A c (A s ω) s ω ≤ (ν (A s ω))[id]}.indicator 1 s = 0 := by + have h_ae := forall_ucbIndex_le_ucbIndex_arm h (bestArm ν) (ν := ν) (c := c) (hK := hK) + have h_gt := time_gt_of_pullCount_gt_one h a (ν := ν) (c := c) (hK := hK) filter_upwards [h_ae, h_gt] with ω h_le h_time_ge simp only [id_eq, tsub_le_iff_right, sum_eq_zero_iff, mem_range, Set.indicator_apply_eq_zero, Set.mem_setOf_eq, Pi.one_apply, one_ne_zero, imp_false, not_and, not_le] @@ -557,18 +610,22 @@ lemma some_sum_eq_zero [Nonempty (Fin K)] (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 · rw [h_arm] gcongr -lemma pullCount_ae_le_add_two [Nonempty (Fin K)] (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 < gap ν a) +omit [StandardBorelSpace Ω] [Nonempty Ω] in +lemma pullCount_ae_le_add_two [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 < gap ν a) (n C : ℕ) (hC : C ≠ 0) (hC' : 4 * c * log (n + 1) / gap ν a ^ 2 ≤ C) : - ∀ᵐ ω ∂𝔓t, - pullCount a n ω ≤ C + 1 + + ∀ᵐ ω ∂P, + pullCount A a n ω ≤ C + 1 + ∑ s ∈ range n, - {s | 0 < pullCount (bestArm ν) s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < - (ν (bestArm ν))[id]}.indicator 1 s + + {s | 0 < pullCount A (bestArm ν) s ω ∧ + empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω < + (ν (bestArm ν))[id]}.indicator 1 s + ∑ s ∈ range n, - {s | 0 < pullCount a s ω ∧ (ν a)[id] < - empMean a s ω - ucbWidth c a s ω}.indicator 1 s := by - filter_upwards [some_sum_eq_zero hc a h_gap n C hC hC', - pullCount_le_add_three_ae a n C hC] with ω hω_zero hω_le + {s | 0 < pullCount A a s ω ∧ (ν a)[id] < + empMean A R a s ω - ucbWidth A c a s ω}.indicator 1 s := by + filter_upwards [some_sum_eq_zero h hc a h_gap n C hC hC', + pullCount_le_add_three_ae h a n C hC] with ω hω_zero hω_le refine (hω_le).trans_eq ?_ rw [hω_zero] @@ -583,57 +640,62 @@ lemma constSum_lt_top (c : ℝ) (n : ℕ) : constSum c n < ∞ := by positivity /-- Bound on the expectation of the number of pulls of each arm by the UCB algorithm. -/ -lemma expectation_pullCount_le' (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma expectation_pullCount_le' [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 < c) (a : Fin K) (h_gap : 0 < gap ν a) (n : ℕ) : - ∫⁻ ω, pullCount a n ω ∂𝔓t ≤ + ∫⁻ ω, pullCount A a n ω ∂P ≤ ENNReal.ofReal (4 * c * log (n + 1) / gap ν a ^ 2 + 1) + 1 + 2 * constSum c n := by + have hA := h.measurable_A + have hR := h.measurable_R by_cases hn_zero : n = 0 · simp [hn_zero] let C a : ℕ := ⌈4 * c * log (n + 1) / gap ν a ^ 2⌉₊ have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - have h_set_1 b : MeasurableSet {a_1 | 0 < pullCount a b a_1 ∧ - (ν a)[id] < empMean a b a_1 - ucbWidth c a b a_1} := by - change MeasurableSet ({a_1 | 0 < pullCount a b a_1} ∩ - {a_1 | (ν a)[id] < empMean a b a_1 - ucbWidth c a b a_1}) + have h_set_1 b : MeasurableSet {a_1 | 0 < pullCount A a b a_1 ∧ + (ν a)[id] < empMean A R a b a_1 - ucbWidth A c a b a_1} := by + change MeasurableSet ({a_1 | 0 < pullCount A a b a_1} ∩ + {a_1 | (ν a)[id] < empMean A R a b a_1 - ucbWidth A c a b a_1}) exact (measurableSet_lt (by fun_prop) (by fun_prop)).inter (measurableSet_lt (by fun_prop) (by fun_prop)) - have h_set_2 b : MeasurableSet {a | 0 < pullCount (bestArm ν) b a ∧ - empMean (bestArm ν) b a + ucbWidth c (bestArm ν) b a < (ν (bestArm ν))[id]} := by - change MeasurableSet ({a | 0 < pullCount (bestArm ν) b a} ∩ - {a | empMean (bestArm ν) b a + ucbWidth c (bestArm ν) b a < (ν (bestArm ν))[id]}) + have h_set_2 b : MeasurableSet {a | 0 < pullCount A (bestArm ν) b a ∧ + empMean A R (bestArm ν) b a + ucbWidth A c (bestArm ν) b a < (ν (bestArm ν))[id]} := by + change MeasurableSet ({a | 0 < pullCount A (bestArm ν) b a} ∩ + {a | empMean A R (bestArm ν) b a + ucbWidth A c (bestArm ν) b a < (ν (bestArm ν))[id]}) exact (measurableSet_lt (by fun_prop) (by fun_prop)).inter (measurableSet_lt (by fun_prop) (by fun_prop)) - have h_meas_1 b : Measurable fun h ↦ {s | 0 < pullCount a s h ∧ (ν a)[id] < - empMean a s h - ucbWidth c a s h}.indicator (1 : ℕ → ℕ) b := by + have h_meas_1 b : Measurable fun h ↦ {s | 0 < pullCount A a s h ∧ (ν a)[id] < + empMean A R a s h - ucbWidth A c a s h}.indicator (1 : ℕ → ℕ) b := by simp only [id_eq, Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] exact Measurable.ite (h_set_1 _) (by fun_prop) (by fun_prop) - have h_meas_2 b : Measurable fun h ↦ {s | 0 < pullCount (bestArm ν) s h ∧ - empMean (bestArm ν) s h + ucbWidth c (bestArm ν) s h < + have h_meas_2 b : Measurable fun h ↦ {s | 0 < pullCount A (bestArm ν) s h ∧ + empMean A R (bestArm ν) s h + ucbWidth A c (bestArm ν) s h < (ν (bestArm ν))[id]}.indicator (1 : ℕ → ℕ) b := by simp only [id_eq, Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] exact Measurable.ite (h_set_2 _) (by fun_prop) (by fun_prop) - calc ∫⁻ ω, pullCount a n ω ∂𝔓t + calc ∫⁻ ω, pullCount A a n ω ∂P _ ≤ ∫⁻ ω, C a + 1 + ∑ s ∈ range n, - {s | 0 < pullCount (bestArm ν) s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < - (ν (bestArm ν))[id]}.indicator (1 : ℕ → ℕ) s + + {s | 0 < pullCount A (bestArm ν) s ω ∧ + empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω < + (ν (bestArm ν))[id]}.indicator (1 : ℕ → ℕ) s + ∑ s ∈ range n, - {s | 0 < pullCount a s ω ∧ (ν a)[id] < - empMean a s ω - ucbWidth c a s ω}.indicator (1 : ℕ → ℕ) s ∂𝔓t := by + {s | 0 < pullCount A a s ω ∧ (ν a)[id] < + empMean A R a s ω - ucbWidth A c a s ω}.indicator (1 : ℕ → ℕ) s ∂P := by refine lintegral_mono_ae ?_ have hCa : C a ≠ 0 := by simp only [ne_eq, Nat.ceil_eq_zero, not_le, C] have : 0 < log (n + 1) := log_pos (by simp; grind) positivity - filter_upwards [pullCount_ae_le_add_two hc.le a h_gap n (C a) hCa (Nat.le_ceil _)] with ω hω + filter_upwards [pullCount_ae_le_add_two h hc.le a h_gap n (C a) hCa (Nat.le_ceil _)] with ω hω simp only [id_eq, Nat.cast_sum] norm_cast _ ≤ (C a : ℝ≥0∞) + 1 + ∑ s ∈ range n, - 𝔓t {ω | 0 < pullCount (bestArm ν) s ω ∧ - empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < (ν (bestArm ν))[id]} + + P {ω | 0 < pullCount A (bestArm ν) s ω ∧ + empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω < (ν (bestArm ν))[id]} + ∑ s ∈ range n, - 𝔓t {ω | 0 < pullCount a s ω ∧ (ν a)[id] < empMean a s ω - ucbWidth c a s ω} := by + P {ω | 0 < pullCount A a s ω ∧ (ν a)[id] < empMean A R a s ω - ucbWidth A c a s ω} := by simp only [id_eq, Nat.cast_sum] rw [lintegral_add_left (by fun_prop), lintegral_add_left (by fun_prop)] simp only [lintegral_const, measure_univ, mul_one] @@ -651,9 +713,9 @@ lemma expectation_pullCount_le' (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - ( ∑ s ∈ range n, 1 / ((s : ℝ≥0∞) + 1) ^ (c / 2 - 1) + ∑ s ∈ range n, 1 / ((s : ℝ≥0∞) + 1) ^ (c / 2 - 1) := by gcongr with s hs s hs - · refine (measure_mono ?_).trans (prob_ucbIndex_le hν hc.le (bestArm ν) s) + · refine (measure_mono ?_).trans (prob_ucbIndex_le h hν hc.le (bestArm ν) s) grind - · refine (measure_mono ?_).trans (prob_ucbIndex_ge hν hc.le a s) + · refine (measure_mono ?_).trans (prob_ucbIndex_ge h hν hc.le a s) grind _ ≤ ENNReal.ofReal (4 * c * log (n + 1) / gap ν a ^ 2 + 1) + 1 + 2 * constSum c n := by rw [two_mul, add_assoc, constSum] @@ -666,15 +728,18 @@ lemma expectation_pullCount_le' (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - ( positivity /-- Bound on the expectation of the number of pulls of each arm by the UCB algorithm. -/ -lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma expectation_pullCount_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 < c) (a : Fin K) (h_gap : 0 < gap ν a) (n : ℕ) : - 𝔓t[fun ω ↦ (pullCount a n ω : ℝ)] ≤ + P[fun ω ↦ (pullCount A a n ω : ℝ)] ≤ 4 * c * log (n + 1) / gap ν a ^ 2 + 2 + 2 * (constSum c n).toReal := by - have h := expectation_pullCount_le' hν hc a h_gap n (hK := hK) + have hA := h.measurable_A + have h := expectation_pullCount_le' h hν hc a h_gap n (hK := hK) simp_rw [← ENNReal.ofReal_natCast] at h rw [← ofReal_integral_eq_lintegral_ofReal] at h rotate_left - · exact integrable_pullCount _ _ + · exact integrable_pullCount hA _ _ · exact ae_of_all _ fun _ ↦ by simp simp only have : 0 ≤ log (n + 1) := log_nonneg (by simp) @@ -694,18 +759,21 @@ lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - ( ring /-- Regret bound for the UCB algorithm. -/ -lemma regret_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 < c) (n : ℕ) : - 𝔓t[regret ν n] ≤ +lemma regret_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 < c) (n : ℕ) : + P[regret ν A n] ≤ ∑ a, (4 * c * log (n + 1) / gap ν a + gap ν a * (2 + 2 * (constSum c n).toReal)) := by + have hA := h.measurable_A simp_rw [regret_eq_sum_pullCount_mul_gap] rw [integral_finset_sum] - swap; · exact fun i _ ↦ (integrable_pullCount i n).mul_const _ + swap; · exact fun i _ ↦ (integrable_pullCount hA i n).mul_const _ gcongr with a rw [integral_mul_const] by_cases h_gap : gap ν a = 0 · simp [h_gap] replace h_gap : 0 < gap ν a := lt_of_le_of_ne gap_nonneg (Ne.symm h_gap) - grw [expectation_pullCount_le hν hc a h_gap n] + grw [expectation_pullCount_le h hν hc a h_gap n] refine le_of_eq ?_ rw [mul_add] field diff --git a/LeanBandits/ForMathlib/HasCondDistrib.lean b/LeanBandits/ForMathlib/HasCondDistrib.lean index e6d9f06d..2f76fe2e 100644 --- a/LeanBandits/ForMathlib/HasCondDistrib.lean +++ b/LeanBandits/ForMathlib/HasCondDistrib.lean @@ -4,6 +4,7 @@ Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne, Paulo Rauber -/ import LeanBandits.ForMathlib.CondDistrib +import Mathlib.Probability.HasLaw /-! # A predicate for having a specified conditional distribution @@ -66,4 +67,38 @@ lemma HasCondDistrib.snd {Y : α → Ω × Ω'} {κ : Kernel β (Ω × Ω')} [Is rw [Kernel.snd_eq] exact HasCondDistrib.comp h measurable_snd +lemma HasLaw.prod_of_hasCondDistrib {P : Measure β} [IsFiniteMeasure μ] [IsSFiniteKernel κ] + (h1 : HasLaw X P μ) (h2 : HasCondDistrib Y X κ μ) : + HasLaw (fun ω ↦ (X ω, Y ω)) (P ⊗ₘ κ) μ := by + have hX := h1.aemeasurable + have hY := h2.aemeasurable_fst + refine ⟨by fun_prop, ?_⟩ + rw [← compProd_map_condDistrib (by fun_prop), h1.map_eq] + refine Measure.compProd_congr ?_ + rw [← h1.map_eq] + exact h2.condDistrib_eq + +lemma HasCondDistrib.prod [IsFiniteMeasure μ] [IsFiniteKernel κ] + {Z : α → Ω'} {η : Kernel (β × Ω) Ω'} [IsFiniteKernel η] + (h1 : HasCondDistrib Y X κ μ) (h2 : HasCondDistrib Z (fun ω ↦ (X ω, Y ω)) η μ) : + HasCondDistrib (fun ω ↦ (Y ω, Z ω)) X (κ ⊗ₖ η) μ := by + have hX := h1.aemeasurable_snd + have hY := h1.aemeasurable_fst + have hZ := h2.aemeasurable_fst + refine ⟨by fun_prop, by fun_prop, ?_⟩ + have h_condDistrib_Y := h1.condDistrib_eq + have h_condDistrib_Z := h2.condDistrib_eq + have h_prod := condDistrib_prod_left hY hZ hX + have h_prod' : 𝓛[fun ω ↦ (Y ω, Z ω) | X; μ] =ᵐ[μ.map X] (κ ⊗ₖ 𝓛[Z | fun ω ↦ (X ω, Y ω); μ]) := by + filter_upwards [h_condDistrib_Y, h_prod] with ω hω₁ hω₂ + rw [hω₂] + ext s hs + rw [Kernel.compProd_apply hs, Kernel.compProd_apply hs] + simp [hω₁] + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] + at h_condDistrib_Z h_condDistrib_Y ⊢ + rw [← Measure.compProd_assoc', ← h_condDistrib_Y, ← h_condDistrib_Z, + AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + rfl + end ProbabilityTheory diff --git a/LeanBandits/ForMathlib/Traj.lean b/LeanBandits/ForMathlib/Traj.lean index 8df53a22..f4e84f1c 100644 --- a/LeanBandits/ForMathlib/Traj.lean +++ b/LeanBandits/ForMathlib/Traj.lean @@ -4,9 +4,10 @@ import LeanBandits.ForMathlib.CondDistrib 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 μ₀] +variable {Ω : Type*} {mΩ : MeasurableSpace Ω} {P : Measure Ω} [IsFiniteMeasure P] + {X : ℕ → Type*} [∀ n, MeasurableSpace (X n)] + {κ : (n : ℕ) → Kernel (Π i : Iic n, X i) (X (n + 1))} [∀ n, IsMarkovKernel (κ n)] + {μ₀ : Measure (X 0)} [IsProbabilityMeasure μ₀] section MeasurableEquiv @@ -27,4 +28,14 @@ lemma traj_zero_map_eval_zero : rw [← Kernel.traj_map_frestrictLe, ← Kernel.map_comp_right _ (by fun_prop) (by fun_prop)] rfl +/-- Uniqueness of `trajMeasure`. -/ +theorem eq_trajMeasure [∀ n, StandardBorelSpace (X n)] [∀ n, Nonempty (X n)] + {Y : (n : ℕ) → Ω → X n} (hY_meas : ∀ n, Measurable (Y n)) + (h0 : P.map (Y 0) = μ₀) + (h_condDistrib : ∀ n, + condDistrib (Y (n + 1)) (fun ω ↦ fun i : Iic n ↦ Y i ω) P + =ᵐ[P.map (fun ω ↦ fun i : Iic n ↦ Y i ω)] κ n) : + P.map (fun ω n ↦ Y n ω) = trajMeasure μ₀ κ := by + sorry + end ProbabilityTheory.Kernel diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 1558fffb..980b01c7 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -53,24 +53,29 @@ end Aux namespace Bandits -variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [MeasurableSingletonClass α] - -lemma integrable_pullCount {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] - (a : α) (n : ℕ) : - Integrable (fun ω ↦ (pullCount a n ω : ℝ)) (Bandit.trajMeasure alg ν) := by +variable {α Ω : Type*} {mα : MeasurableSpace α} {mΩ : MeasurableSpace Ω} [DecidableEq α] + [StandardBorelSpace α] [Nonempty α] + {A : ℕ → Ω → α} {R : ℕ → Ω → ℝ} {P : Measure Ω} [IsProbabilityMeasure P] + {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] + {h_inter : IsAlgEnvSeq A R alg (stationaryEnv ν) P} + +omit [StandardBorelSpace α] [Nonempty α] in +lemma integrable_pullCount [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (a : α) (n : ℕ) : + Integrable (fun ω ↦ (pullCount A a n ω : ℝ)) P := by refine integrable_of_le_of_le (g₁ := 0) (g₂ := fun _ ↦ n) (by fun_prop) (ae_of_all _ fun ω ↦ by simp) (ae_of_all _ fun ω ↦ ?_) (integrable_const _) (integrable_const _) simp only [Nat.cast_le] exact pullCount_le a n ω -variable {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] +local notation "𝔓'" => P.prod (Bandit.streamMeasure ν) -omit [DecidableEq α] [MeasurableSingletonClass α] in +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in lemma hasLaw_Z (a : α) (m : ℕ) : - HasLaw (fun ω ↦ ω.2 m a) (ν a) (Bandit.measure alg ν) where + HasLaw (fun ω ↦ ω.2 m a) (ν a) 𝔓' where map_eq := by - calc (Bandit.measure alg ν).map (fun ω ↦ ω.2 m a) - _ = ((Bandit.measure alg ν).snd).map (fun ω ↦ ω m a) := by + calc (𝔓').map (fun ω ↦ ω.2 m a) + _ = ((𝔓').snd).map (fun ω ↦ ω m a) := by rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)] rfl _ = (Bandit.streamMeasure ν).map (fun ω ↦ ω m a) := by simp @@ -90,92 +95,105 @@ notation "𝓛[" Y " | " X " ← " x "; " μ "]" => Measure.map Y (μ[|X ⁻¹' local notation "𝔓t" => Bandit.trajMeasure alg ν local notation "𝔓" => Bandit.measure alg ν -omit [DecidableEq α] [MeasurableSingletonClass α] in -lemma condDistrib_reward'' [StandardBorelSpace α] [Nonempty α] (n : ℕ) : - 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1; 𝔓] - =ᵐ[(𝔓).map (fun ω ↦ arm n ω.1)] ν := by - have h_ra' : 𝓛[reward n | arm n; 𝔓t] =ᵐ[(𝔓t).map (arm n)] ν := condDistrib_reward alg ν n - have h_law : (𝔓).map (fun ω ↦ arm n ω.1) = (𝔓t).map (arm n) := by - rw [← Bandit.fst_measure, Measure.fst, Measure.map_map (by fun_prop) (by fun_prop)] - rfl +omit [DecidableEq α] in +lemma condDistrib_reward'' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (n : ℕ) : + 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1; 𝔓'] =ᵐ[(𝔓').map (fun ω ↦ A n ω.1)] ν := by + have hA := h.measurable_A + have hR := h.measurable_R + have h_ra' : 𝓛[R n | A n; P] =ᵐ[P.map (A n)] ν := h.condDistrib_reward_stationaryEnv n + have h_law : (𝔓').map (fun ω ↦ A n ω.1) = P.map (A n) := by + change ((𝔓').map (A n ∘ Prod.fst)) = _ + rw [← Measure.map_map (by fun_prop) (by fun_prop), ← Measure.fst, Measure.fst_prod] rw [h_law] - have h_prod : 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1; 𝔓] - =ᵐ[(𝔓t).map (arm n)] 𝓛[reward n | arm n; 𝔓t] := + have h_prod : 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1; 𝔓'] + =ᵐ[P.map (A n)] 𝓛[R n | A n; P] := condDistrib_fst_prod _ (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 [StandardBorelSpace α] [Nonempty α] [Countable α] (a : α) (n : ℕ) - (hμa : (𝔓).map (fun ω ↦ arm n ω.1) {a} ≠ 0) : - 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1 ← a; 𝔓] = ν a := by - 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) +lemma reward_cond_action [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n : ℕ) + (hμa : (𝔓').map (fun ω ↦ A n ω.1) {a} ≠ 0) : + 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1 ← a; 𝔓'] = ν a := by + have hA := h.measurable_A + have hR := h.measurable_R + have h_ra : 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1; 𝔓'] =ᵐ[(𝔓').map (fun ω ↦ A n ω.1)] ν := + condDistrib_reward'' h n + have h_eq := condDistrib_ae_eq_cond (μ := 𝔓') + (X := fun ω ↦ A n ω.1) (Y := fun ω ↦ R 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 : ℕ) : - reward n ⟂ᵢ[arm n, measurable_arm n; 𝔓t] {ω | stepsUntil a m ω = ↑n}.indicator (fun _ ↦ 1) := by - -- the indicator of `stepsUntil ... = n` is a function of `hist (n-1)` and `arm n`. +lemma condIndepFun_reward_stepsUntil_action' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m n : ℕ) : + R n ⟂ᵢ[A n, h.measurable_A n; P] {ω | stepsUntil A a m ω = ↑n}.indicator (fun _ ↦ 1) := by + -- the indicator of `stepsUntil ... = n` is a function of `hist (n-1)` and `action n`. -- It thus suffices to use the independence of `reward n` and `hist (n-1)` conditionally - -- on `arm n`. + -- on `action n`. + have hA := h.measurable_A + have hR := h.measurable_R by_cases hn : n = 0 - · have h_indep : reward 0 ⟂ᵢ[arm 0, measurable_arm 0; 𝔓t] arm 0 := + · have h_indep : R 0 ⟂ᵢ[A 0, hA 0; P] A 0 := condIndepFun_self_right (by fun_prop) (by fun_prop) simp only [hn, CharP.cast_eq_zero] - refine h_indep.of_measurable_right (hX := measurable_arm 0) ?_ + refine h_indep.of_measurable_right (hX := hA 0) ?_ exact measurable_comap_indicator_stepsUntil_eq_zero a m - · have h_indep : reward n ⟂ᵢ[arm n, measurable_arm n; 𝔓t] fun ω ↦ (hist (n - 1) ω, arm n ω) := - condIndepFun_reward_hist_arm_arm' (alg := alg) (ν := ν) n (by grind) - refine h_indep.of_measurable_right (hX := measurable_arm n) ?_ - exact measurable_comap_indicator_stepsUntil_eq a m n + · have h_indep : R n ⟂ᵢ[A n, hA n; P] fun ω ↦ (IsAlgEnvSeq.hist A R (n - 1) ω, A n ω) := + IsAlgEnvSeq.condIndepFun_reward_hist_action_action' h n (by grind) + refine h_indep.of_measurable_right (hX := hA n) ?_ + exact measurable_comap_indicator_stepsUntil_eq hA hR a m n -lemma condIndepFun_reward_stepsUntil_arm [StandardBorelSpace α] [Countable α] [Nonempty α] +lemma condIndepFun_reward_stepsUntil_action [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m n : ℕ) : - CondIndepFun (mα.comap (fun ω ↦ arm n ω.1)) ((measurable_arm n).comp measurable_fst).comap_le - (fun ω ↦ reward n ω.1) ({ω | stepsUntil a m ω.1 = ↑n}.indicator (fun _ ↦ 1)) 𝔓 := - condIndepFun_fst_prod (ν := Bandit.streamMeasure ν) - (measurable_indicator_stepsUntil_eq a m n) (by fun_prop) (by fun_prop) - (condIndepFun_reward_stepsUntil_arm' a m n) - -lemma reward_cond_stepsUntil [StandardBorelSpace α] [Countable α] [Nonempty α] (a : α) (m n : ℕ) - (hm : m ≠ 0) (hμn : 𝔓 ((fun ω ↦ stepsUntil a m ω.1) ⁻¹' {↑n}) ≠ 0) : - 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ stepsUntil a m ω.1 ← ↑n; 𝔓] = ν a := by + CondIndepFun (mα.comap (fun ω ↦ A n ω.1)) ((h.measurable_A n).comp measurable_fst).comap_le + (fun ω ↦ R n ω.1) ({ω | stepsUntil A a m ω.1 = ↑n}.indicator (fun _ ↦ 1)) 𝔓' := by + have hA := h.measurable_A + have hR := h.measurable_R + exact condIndepFun_fst_prod (ν := Bandit.streamMeasure ν) + (measurable_indicator_stepsUntil_eq hA hR a m n) (by fun_prop) (by fun_prop) + (condIndepFun_reward_stepsUntil_action' h a m n) + +lemma reward_cond_stepsUntil [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m n : ℕ) + (hm : m ≠ 0) (hμn : 𝔓' ((fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n}) ≠ 0) : + 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ stepsUntil A a m ω.1 ← ↑n; 𝔓'] = ν a := by + have hA := h.measurable_A + have hR := h.measurable_R have hμna : - 𝔓 ((fun ω ↦ stepsUntil a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}) ≠ 0 := by - suffices ((fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ - stepsUntil a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}) - = (fun ω ↦ stepsUntil a m ω.1) ⁻¹' {↑n} by simpa [this] using hμn + 𝔓' ((fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ A n ω.1) ⁻¹' {a}) ≠ 0 := by + suffices ((fun ω : Ω × (ℕ → α → ℝ) ↦ + stepsUntil A a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ A n ω.1) ⁻¹' {a}) + = (fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑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 + exact action_eq_of_stepsUntil_eq_coe hm + have hμa : (𝔓').map (fun ω ↦ A 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 a m ω.1 ← (n : ℕ∞); 𝔓] - _ = (𝔓[|(fun ω ↦ stepsUntil a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}]).map - (fun ω ↦ reward n ω.1) := by + exact action_eq_of_stepsUntil_eq_coe hm + calc 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ stepsUntil A a m ω.1 ← (n : ℕ∞); 𝔓'] + _ = (𝔓'[|(fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ A n ω.1) ⁻¹' {a}]).map + (fun ω ↦ R 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 - _ = (𝔓[|(fun ω ↦ arm n ω.1) ⁻¹' {a} - ∩ {ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) | stepsUntil a m ω.1 = ↑n}.indicator 1 ⁻¹' {1} ]).map - (fun ω ↦ reward n ω.1) := by + exact action_eq_of_stepsUntil_eq_coe hm + _ = (𝔓'[|(fun ω ↦ A n ω.1) ⁻¹' {a} + ∩ {ω : Ω × (ℕ → α → ℝ) | stepsUntil A a m ω.1 = ↑n}.indicator 1 ⁻¹' {1} ]).map + (fun ω ↦ R n ω.1) := by congr 2 with ω simp only [Set.mem_inter_iff, Set.mem_preimage, Set.mem_singleton_iff, Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply, ite_eq_left_iff, zero_ne_one, imp_false, Decidable.not_not] rw [and_comm] - _ = 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1 ← a; 𝔓] := by + _ = 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1 ← a; 𝔓'] := by rw [cond_of_condIndepFun (by fun_prop)] - · exact condIndepFun_reward_stepsUntil_arm a m n + · exact condIndepFun_reward_stepsUntil_action h a m n · refine measurable_one.indicator ?_ exact measurableSet_eq_fun (by fun_prop) (by fun_prop) · fun_prop @@ -183,16 +201,18 @@ lemma reward_cond_stepsUntil [StandardBorelSpace α] [Countable α] [Nonempty α rw [Set.inter_comm] congr 1 with ω simp [Set.indicator_apply] - _ = ν a := reward_cond_arm a n hμa + _ = ν a := reward_cond_action h a n hμa -/-- The conditional distribution of the reward received at the `m`-th pull of arm `a` +/-- The conditional distribution of the reward received at the `m`-th pull of action `a` given the time at which number of pulls is `m` is the constant kernel with value `ν a`. -/ -theorem condDistrib_rewardByCount_stepsUntil [Countable α] [StandardBorelSpace α] [Nonempty α] - (a : α) (m : ℕ) (hm : m ≠ 0) : - condDistrib (rewardByCount a m) (fun ω ↦ stepsUntil a m ω.1) 𝔓 - =ᵐ[(𝔓).map (fun ω ↦ stepsUntil a m ω.1)] Kernel.const _ (ν a) := by - refine (condDistrib_ae_eq_cond (μ := 𝔓) - (X := fun ω ↦ stepsUntil a m ω.1) (by fun_prop) (by fun_prop)).trans ?_ +theorem condDistrib_rewardByCount_stepsUntil [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m : ℕ) (hm : m ≠ 0) : + condDistrib (rewardByCount A R a m) (fun ω ↦ stepsUntil A a m ω.1) 𝔓' + =ᵐ[(𝔓').map (fun ω ↦ stepsUntil A a m ω.1)] Kernel.const _ (ν a) := by + have hA := h.measurable_A + have hR := h.measurable_R + refine (condDistrib_ae_eq_cond (μ := 𝔓') + (X := fun ω ↦ stepsUntil A a m ω.1) (by fun_prop) (by fun_prop)).trans ?_ rw [Filter.EventuallyEq, ae_iff_of_countable] intro n hn simp only [Kernel.const_apply] @@ -206,255 +226,95 @@ theorem condDistrib_rewardByCount_stepsUntil [Countable α] [StandardBorelSpace 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 a m ω) + · exact indepFun_prod (X := fun ω : Ω ↦ stepsUntil A a m ω) (Y := fun ω : ℕ → α → ℝ ↦ ω m a) (by fun_prop) (by fun_prop) | coe n => - rw [Measure.map_congr (g := fun ω ↦ reward n ω.1)] + rw [Measure.map_congr (g := fun ω ↦ R 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 ?_ + refine reward_cond_stepsUntil h 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 [Countable α] [StandardBorelSpace α] [Nonempty α] - (a : α) (m : ℕ) (hm : m ≠ 0) : - HasLaw (rewardByCount a m) (ν a) 𝔓 where +/-- The reward received at the `m`-th pull of action `a` has law `ν a`. -/ +lemma hasLaw_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m : ℕ) (hm : m ≠ 0) : + HasLaw (rewardByCount A R a m) (ν a) 𝔓' where + aemeasurable := (measurable_rewardByCount h.measurable_A h.measurable_R a m).aemeasurable map_eq := by + have hA := h.measurable_A + have hR := h.measurable_R have h_condDistrib : - condDistrib (rewardByCount a m) (fun ω ↦ stepsUntil a m ω.1) 𝔓 - =ᵐ[(𝔓).map (fun ω ↦ stepsUntil a m ω.1)] - Kernel.const _ (ν a) := condDistrib_rewardByCount_stepsUntil a m hm - calc (𝔓).map (rewardByCount a m) - _ = (condDistrib (rewardByCount a m) (fun ω ↦ stepsUntil a m ω.1) 𝔓) - ∘ₘ ((𝔓).map (fun ω ↦ stepsUntil a m ω.1)) := by + condDistrib (rewardByCount A R a m) (fun ω ↦ stepsUntil A a m ω.1) 𝔓' + =ᵐ[(𝔓').map (fun ω ↦ stepsUntil A a m ω.1)] + Kernel.const _ (ν a) := condDistrib_rewardByCount_stepsUntil h a m hm + calc (𝔓').map (rewardByCount A R a m) + _ = (condDistrib (rewardByCount A R a m) (fun ω ↦ stepsUntil A a m ω.1) 𝔓') + ∘ₘ ((𝔓').map (fun ω ↦ stepsUntil A a m ω.1)) := by rw [condDistrib_comp_map (by fun_prop) (by fun_prop)] - _ = (Kernel.const _ (ν a)) ∘ₘ ((𝔓).map (fun ω ↦ stepsUntil a m ω.1)) := + _ = (Kernel.const _ (ν a)) ∘ₘ ((𝔓').map (fun ω ↦ stepsUntil A a m ω.1)) := Measure.comp_congr h_condDistrib _ = ν a := by - have : IsProbabilityMeasure ((𝔓).map (fun ω ↦ stepsUntil a m ω.1)) := + have : IsProbabilityMeasure ((𝔓').map (fun ω ↦ stepsUntil A a m ω.1)) := Measure.isProbabilityMeasure_map (by fun_prop) simp -lemma identDistrib_rewardByCount [Countable α] [StandardBorelSpace α] [Nonempty α] (a : α) (n m : ℕ) +lemma identDistrib_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n m : ℕ) (hn : n ≠ 0) (hm : m ≠ 0) : - IdentDistrib (rewardByCount a n) (rewardByCount a m) 𝔓 𝔓 where - aemeasurable_fst := by fun_prop - aemeasurable_snd := by fun_prop - map_eq := by rw [(hasLaw_rewardByCount a n hn).map_eq, (hasLaw_rewardByCount a m hm).map_eq] - -lemma identDistrib_rewardByCount_id [Countable α] [StandardBorelSpace α] [Nonempty α] - (a : α) (n : ℕ) (hn : n ≠ 0) : - IdentDistrib (rewardByCount a n) id 𝔓 (ν a) where - aemeasurable_fst := by fun_prop + IdentDistrib (rewardByCount A R a n) (rewardByCount A R a m) 𝔓' 𝔓' where + aemeasurable_fst := (measurable_rewardByCount h.measurable_A h.measurable_R a n).aemeasurable + aemeasurable_snd := (measurable_rewardByCount h.measurable_A h.measurable_R a m).aemeasurable + map_eq := by rw [(hasLaw_rewardByCount h a n hn).map_eq, (hasLaw_rewardByCount h a m hm).map_eq] + +lemma identDistrib_rewardByCount_id [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n : ℕ) (hn : n ≠ 0) : + IdentDistrib (rewardByCount A R a n) id 𝔓' (ν a) where + aemeasurable_fst := (measurable_rewardByCount h.measurable_A h.measurable_R a n).aemeasurable aemeasurable_snd := Measurable.aemeasurable <| by fun_prop - map_eq := by rw [(hasLaw_rewardByCount a n hn).map_eq, Measure.map_id] + map_eq := by rw [(hasLaw_rewardByCount h a n hn).map_eq, Measure.map_id] -lemma identDistrib_rewardByCount_eval [Countable α] [StandardBorelSpace α] [Nonempty α] - (a : α) (n m : ℕ) (hn : n ≠ 0) : - IdentDistrib (rewardByCount a n) (fun ω ↦ ω m a) 𝔓 (Bandit.streamMeasure ν) := - (identDistrib_rewardByCount_id a n hn).trans (identDistrib_eval_eval_id_streamMeasure ν m a).symm +lemma identDistrib_rewardByCount_eval [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n m : ℕ) (hn : n ≠ 0) : + IdentDistrib (rewardByCount A R a n) (fun ω ↦ ω m a) 𝔓' (Bandit.streamMeasure ν) := + (identDistrib_rewardByCount_id h a n hn).trans + (identDistrib_eval_eval_id_streamMeasure ν m a).symm -lemma indepFun_rewardByCount_Iic (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (a : α) +lemma indepFun_rewardByCount_Iic [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n : ℕ) : - (rewardByCount a (n + 1)) ⟂ᵢ[𝔓] fun ω (i : Iic n) ↦ rewardByCount a i ω := by + (rewardByCount A R a (n + 1)) ⟂ᵢ[𝔓'] fun ω (i : Iic n) ↦ rewardByCount A R a i ω := by sorry -lemma iIndepFun_rewardByCount' (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (a : α) : - iIndepFun (rewardByCount a) (Bandit.measure alg ν) := by +lemma iIndepFun_rewardByCount' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) : + iIndepFun (rewardByCount A R a) 𝔓' := by + have hA := h.measurable_A + have hR := h.measurable_R rw [iIndepFun_nat_iff_forall_indepFun (by fun_prop)] - exact indepFun_rewardByCount_Iic alg ν a - -def E' (I : Finset (α × ℕ)) (S : Finset I) : Set (ℕ → α × ℝ) := - {ω | (∀ i ∈ S, stepsUntil i.1.1 (i.1.2 + 1) ω < ⊤) ∧ - (∀ j ∉ S, stepsUntil j.1.1 (j.1.2 + 1) ω = ⊤)} - -lemma measurableSet_E' [Countable α] [Nonempty α] (I : Finset (α × ℕ)) (S : Finset I) : - MeasurableSet (E' I S) := by - have h_eq : E' I S - = (⋂ i ∈ S, {ω | stepsUntil i.1.1 (i.1.2 + 1) ω ≠ ⊤}) ∩ - (⋂ j ∉ S, {ω | stepsUntil j.1.1 (j.1.2 + 1) ω = ⊤}) := by ext; simp [E', lt_top_iff_ne_top] - rw [h_eq] - refine MeasurableSet.inter ?_ ?_ - · refine MeasurableSet.iInter fun i ↦ MeasurableSet.iInter fun hi ↦ ?_ - exact (measurableSet_singleton _).compl.preimage (by fun_prop) - · refine MeasurableSet.iInter fun j ↦ MeasurableSet.iInter fun hj ↦ ?_ - exact (measurableSet_singleton _).preimage (by fun_prop) - -def E (I : Finset (α × ℕ)) (S : Finset I) : Set ((ℕ → α × ℝ) × (ℕ → α → ℝ)) := - {ω | (∀ i ∈ S, stepsUntil i.1.1 (i.1.2 + 1) ω.1 < ⊤) ∧ - (∀ j ∉ S, stepsUntil j.1.1 (j.1.2 + 1) ω.1 = ⊤)} - -lemma measurableSet_E [Countable α] [Nonempty α] (I : Finset (α × ℕ)) (S : Finset I) : - MeasurableSet (E I S) := by - have : E I S = Prod.fst ⁻¹' (E' I S) := by ext; simp [E, E'] - rw [this] - exact measurable_fst (measurableSet_E' I S) - -lemma iIndepFun_rewardByCount.extracted_1 [Countable α] [Nonempty α] - (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (I : Finset (α × ℕ)) - {B : α × ℕ → Set ℝ} (hB : ∀ i ∈ I, MeasurableSet (B i)) (S : Finset I) : - 𝔓 (E I S ∩ ⋂ i ∈ S, (fun ω ↦ reward (stepsUntil i.1.1 (i.1.2 + 1) ω.1).toNat ω.1) ⁻¹' B i) = - 𝔓 (E I S) * ∏ i ∈ S, (ν i.1.1) (B i) := by - sorry + exact indepFun_rewardByCount_Iic h a -lemma iIndepFun_rewardByCount.extracted_2 [Countable α] [Nonempty α] - (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (I : Finset (α × ℕ)) - {B : α × ℕ → Set ℝ} (hB : ∀ i ∈ I, MeasurableSet (B i)) (S : Finset I) : - Bandit.measure alg ν (⋂ j ∉ S, (fun ω ↦ ω.2 (j.1.2 + 1) j.1.1) ⁻¹' B j) = - ∏ j ∉ S, (ν j.1.1) (B j) := by - have h_indep : iIndepFun (fun (i : I) ω ↦ ω.2 (i.1.2 + 1) i.1.1) (Bandit.measure alg ν) := by - suffices iIndepFun (fun (i : I) ω ↦ ω (i.1.2 + 1) i.1.1) (Bandit.streamMeasure ν) by - sorry - sorry - rw [iIndepFun_iff_measure_inter_preimage_eq_mul] at h_indep - specialize h_indep Sᶜ (sets := fun i ↦ B i) (fun i hi ↦ hB i i.2) - simp only [mem_compl] at h_indep ⊢ - rw [h_indep] - congr with i - rw [← Measure.map_apply (by fun_prop) (hB i i.2)] - congr - exact (hasLaw_Z i.1.1 (i.1.2 + 1)).map_eq - -lemma iIndepFun_rewardByCount.extracted_3 [Countable α] [Nonempty α] - (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [inst_4 : IsMarkovKernel ν] (I : Finset (α × ℕ)) - {B : α × ℕ → Set ℝ} (hB : ∀ i ∈ I, MeasurableSet (B i)) (S : Finset I) : - IndepSet (E I S ∩ - ⋂ i ∈ S, (fun ω ↦ reward (stepsUntil i.1.1 (i.1.2 + 1) ω.1).toNat ω.1) ⁻¹' B i) - (⋂ j ∉ S, (fun ω ↦ ω.2 (j.1.2 + 1) j.1.1) ⁻¹' B j) (Bandit.measure alg ν) := by - let A := E I S ∩ ⋂ i ∈ S, (fun ω ↦ reward (stepsUntil i.1.1 (i.1.2 + 1) ω.1).toNat ω.1) ⁻¹' B i - let C := ⋂ j ∉ S, (fun (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) ↦ ω.2 (j.1.2 + 1) j.1.1) ⁻¹' B j - let A' := E' I S ∩ ⋂ i ∈ S, (fun ω ↦ reward (stepsUntil i.1.1 (i.1.2 + 1) ω).toNat ω) ⁻¹' B i - have hA : A = Prod.fst ⁻¹' A' := by - ext ω - simp [A, A', E, E'] - let C' := ⋂ j ∉ S, (fun (ω : ℕ → α → ℝ) ↦ ω (j.1.2 + 1) j.1.1) ⁻¹' B j - have hC : C = Prod.snd ⁻¹' C' := by - ext ω - simp [C, C'] - have hAC : A ∩ C = A' ×ˢ C' := by rw [hA, hC]; ext; simp - have hA'_meas : MeasurableSet A' := by - refine MeasurableSet.inter ?_ (MeasurableSet.iInter fun i ↦ MeasurableSet.iInter fun hi ↦ ?_) - · exact measurableSet_E' I S - · exact (hB i.1 i.2).preimage (by fun_prop) - have hC'_meas : MeasurableSet C' := by - refine MeasurableSet.iInter fun j ↦ MeasurableSet.iInter fun hj ↦ ?_ - exact (hB j.1 j.2).preimage (by fun_prop) - change IndepSet A C (Bandit.measure alg ν) - rw [indepSet_iff_measure_inter_eq_mul (μ := Bandit.measure alg ν)] - rotate_left - · rw [hA] - exact measurable_fst hA'_meas - · rw [hC] - exact measurable_snd hC'_meas - rw [hAC, hA, hC, Bandit.measure, ← Measure.fst_apply, ← Measure.snd_apply] - · simp - · exact hC'_meas - · exact hA'_meas - -lemma iIndepFun_rewardByCount [Countable α] [Nonempty α] - (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] : - iIndepFun (fun (p : α × ℕ) ↦ rewardByCount p.1 (p.2 + 1)) (Bandit.measure alg ν) := by - rw [iIndepFun_iff_measure_inter_preimage_eq_mul] - intro I B hB - suffices Bandit.measure alg ν (⋂ i ∈ I, rewardByCount i.1 (i.2 + 1) ⁻¹' B i) = - ∏ i ∈ I, ν i.1 (B i) by - rw [this] - refine Finset.prod_congr rfl fun i hi ↦ ?_ - rw [← Measure.map_apply (by fun_prop) (hB i hi)] - congr - exact (hasLaw_rewardByCount i.1 (i.2 + 1) (by simp)).map_eq.symm - have hE_disj (S T : Finset I) (hST : S ≠ T) : Disjoint (E I S) (E I T) := by - rw [Set.disjoint_iff_forall_ne] - simp only [Subtype.forall, Prod.forall, Set.mem_setOf_eq, ne_eq, and_imp, Prod.mk.injEq, - not_and, E] - grind - have hE_union : ⋃ (S : Finset I), E I S = Set.univ := by - ext ω - simp only [Subtype.forall, Prod.forall, Set.mem_iUnion, Set.mem_setOf_eq, Set.mem_univ, - iff_true, E] - use Finset.univ.filter (fun i : I ↦ stepsUntil i.1.1 (i.1.2 + 1) ω.1 < ⊤) - simp - have : ⋂ i ∈ I, rewardByCount i.1 (i.2 + 1) ⁻¹' B i = - ⋃ (S : Finset I), E I S ∩ (⋂ i ∈ I, rewardByCount i.1 (i.2 + 1) ⁻¹' B i) := by - rw [← Set.iUnion_inter, hE_union, Set.univ_inter] - rw [this, measure_iUnion] - rotate_left - · intro S T hST - simp only [Function.onFun] - exact Disjoint.inter_left _ (Disjoint.inter_right _ (hE_disj S T hST)) - · refine fun S ↦ (measurableSet_E I S).inter ?_ - refine MeasurableSet.iInter fun i ↦ MeasurableSet.iInter fun hi ↦ ?_ - exact (hB i hi).preimage (by fun_prop) - suffices ∀ (S : Finset I), - Bandit.measure alg ν (E I S ∩ ⋂ i ∈ I, rewardByCount i.1 (i.2 + 1) ⁻¹' B i) = - Bandit.measure alg ν (E I S) * ∏ i ∈ I, ν i.1 (B i) by - simp_rw [this] - rw [ENNReal.tsum_mul_right, ← measure_iUnion hE_disj (measurableSet_E I), hE_union, - measure_univ, one_mul] - intro S - have h_eq : E I S ∩ ⋂ i ∈ I, rewardByCount i.1 (i.2 + 1) ⁻¹' B i - = E I S ∩ (⋂ i ∈ S, (fun ω ↦ reward (stepsUntil i.1.1 (i.1.2 + 1) ω.1).toNat ω.1) ⁻¹' B i) ∩ - (⋂ j ∉ S, (fun ω ↦ ω.2 (j.1.2 + 1) j.1.1) ⁻¹' B j) := by - ext ω - rw [Set.inter_assoc] - simp only [Set.mem_inter_iff, and_congr_right_iff] - intro hω - simp only [Subtype.forall, Prod.forall, Set.mem_setOf_eq, E] at hω - conv_rhs => rw [Set.iInter_subtype, Set.iInter_subtype] - rw [← Set.mem_inter_iff, ← Set.iInter_inter_distrib] - simp_rw [← Set.iInter_inter_distrib] - simp only [Set.mem_iInter, Set.mem_preimage, Prod.forall, Set.mem_inter_iff] - constructor - · intro h_all a i hai - constructor - · intro haiS - convert h_all a i hai - replace hω := hω.1 a i hai haiS - rw [rewardByCount_of_stepsUntil_ne_top hω.ne] - rfl - · intro haiS - convert h_all a i hai - replace hω := hω.2 a i hai haiS - rw [rewardByCount_of_stepsUntil_eq_top hω] - · intro h a i hai - specialize h a i hai - by_cases haiS : ⟨⟨a, i⟩, hai⟩ ∈ S - · convert h.1 haiS - replace hω := hω.1 a i hai haiS - rw [rewardByCount_of_stepsUntil_ne_top hω.ne] - rfl - · convert h.2 haiS - replace hω := hω.2 a i hai haiS - rw [rewardByCount_of_stepsUntil_eq_top hω] - rw [h_eq, IndepSet.measure_inter_eq_mul] - swap; · exact iIndepFun_rewardByCount.extracted_3 alg ν I hB S - rw [iIndepFun_rewardByCount.extracted_1 alg ν I hB S, - iIndepFun_rewardByCount.extracted_2 alg ν I hB S, mul_assoc] - congr - rw [Finset.prod_mul_prod_compl, Finset.prod_subtype I] - simp - -lemma identDistrib_rewardByCount_stream_all [Countable α] [StandardBorelSpace α] [Nonempty α] : - IdentDistrib (fun ω (p : α × ℕ) ↦ rewardByCount p.1 (p.2 + 1) ω) - (fun ω p ↦ ω p.2 p.1) 𝔓 (Bandit.streamMeasure ν) := by +lemma identDistrib_rewardByCount_stream_all [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) : + IdentDistrib (fun ω (p : α × ℕ) ↦ rewardByCount A R p.1 (p.2 + 1) ω) + (fun ω p ↦ ω p.2 p.1) 𝔓' (Bandit.streamMeasure ν) := by refine IdentDistrib.pi (fun p ↦ ?_) ?_ ?_ - · refine identDistrib_rewardByCount_eval p.1 (p.2 + 1) p.2 (by simp) (ν := ν) - · exact iIndepFun_rewardByCount alg ν + · refine identDistrib_rewardByCount_eval h p.1 (p.2 + 1) p.2 (by simp) (ν := ν) + · sorry · sorry -lemma identDistrib_rewardByCount_stream' [Countable α] [StandardBorelSpace α] [Nonempty α] - (a : α) : - IdentDistrib (fun ω n ↦ rewardByCount a (n + 1) ω) (fun ω n ↦ ω n a) - 𝔓 (Bandit.streamMeasure ν) := by +lemma identDistrib_rewardByCount_stream' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) : + IdentDistrib (fun ω n ↦ rewardByCount A R a (n + 1) ω) (fun ω n ↦ ω n a) + 𝔓' (Bandit.streamMeasure ν) := by refine IdentDistrib.pi (fun n ↦ ?_) ?_ ?_ - · refine identDistrib_rewardByCount_eval a (n + 1) n (by simp) (ν := ν) - · have h_indep := iIndepFun_rewardByCount' alg ν a + · refine identDistrib_rewardByCount_eval h a (n + 1) n (by simp) (ν := ν) + · have h_indep := iIndepFun_rewardByCount' h a exact iIndepFun.precomp (g := fun n ↦ n + 1) (fun i j hij ↦ by grind) h_indep · exact iIndepFun_eval_streamMeasure'' ν a -omit [DecidableEq α] [MeasurableSingletonClass α] in +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in lemma identDistrib_eval_streamMeasure_measure (a : α) : IdentDistrib (fun ω n ↦ ω n a) (fun ω n ↦ ω.2 n a) (Bandit.streamMeasure ν) 𝔓 := by @@ -469,23 +329,25 @@ lemma identDistrib_eval_streamMeasure_measure (a : α) : rw [← Measure.snd, Bandit.snd_measure] exact iIndepFun_eval_streamMeasure'' ν a -lemma identDistrib_rewardByCount_stream [Countable α] [StandardBorelSpace α] [Nonempty α] - (a : α) : - IdentDistrib (fun ω n ↦ rewardByCount a (n + 1) ω) (fun ω n ↦ ω.2 n a) 𝔓 𝔓 := - (identDistrib_rewardByCount_stream' a).trans (identDistrib_eval_streamMeasure_measure a) +lemma identDistrib_rewardByCount_stream [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) : + IdentDistrib (fun ω n ↦ rewardByCount A R a (n + 1) ω) (fun ω n ↦ ω.2 n a) 𝔓' 𝔓 := + (identDistrib_rewardByCount_stream' h a).trans (identDistrib_eval_streamMeasure_measure a) -lemma indepFun_rewardByCount_of_ne {a b : α} (hab : a ≠ b) : - IndepFun (fun ω s ↦ rewardByCount a s ω) (fun ω s ↦ rewardByCount b s ω) 𝔓 := by +lemma indepFun_rewardByCount_of_ne [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {a b : α} (hab : a ≠ b) : + IndepFun (fun ω s ↦ rewardByCount A R a s ω) (fun ω s ↦ rewardByCount A R b s ω) 𝔓' := by sorry -lemma identDistrib_sum_Icc_rewardByCount [Nonempty α] [Countable α] (m : ℕ) (a : α) : - IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount a s ω) - (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓 𝔓 := by +lemma identDistrib_sum_Icc_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (m : ℕ) (a : α) : + IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount A R a s ω) + (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓' 𝔓 := by have h1 (a : α) : - IdentDistrib (fun ω s ↦ rewardByCount a (s + 1) ω) (fun ω s ↦ ω.2 s a) 𝔓 𝔓 := - identDistrib_rewardByCount_stream a - have h_eq (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) : ∑ s ∈ Icc 1 m, rewardByCount a s ω - = ∑ s ∈ range m, rewardByCount a (s + 1) ω := by + IdentDistrib (fun ω s ↦ rewardByCount A R a (s + 1) ω) (fun ω s ↦ ω.2 s a) 𝔓' 𝔓 := + identDistrib_rewardByCount_stream h a + have h_eq (ω : Ω × (ℕ → α → ℝ)) : ∑ s ∈ Icc 1 m, rewardByCount A R a s ω + = ∑ s ∈ range m, rewardByCount A R a (s + 1) ω := by let e : Icc 1 m ≃ range m := { toFun x := ⟨x - 1, by have h := x.2; simp only [mem_Icc] at h; simp; grind⟩ invFun x := ⟨x + 1, by diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index 335d4729..a9d9b183 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -57,12 +57,41 @@ lemma fst_stepKernel (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : (stepKernel alg env n).fst = alg.policy n := by rw [stepKernel, Kernel.fst_compProd] -section IsAlgEnvInteraction +section IsAlgEnvSeq variable {A : ℕ → Ω → α} {R' : ℕ → Ω → R} {alg : Algorithm α R} {env : Environment α R} {P : Measure Ω} [IsFiniteMeasure P] -structure IsAlgEnvInteraction +def IsAlgEnvSeq.step (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (n : ℕ) (ω : Ω) : α × R := + (A n ω, R' n ω) + +@[fun_prop] +lemma IsAlgEnvSeq.measurable_step (n : ℕ) (hA : Measurable (A n)) + (hR' : Measurable (R' n)) : + Measurable (IsAlgEnvSeq.step A R' n) := by + unfold IsAlgEnvSeq.step + fun_prop + +def IsAlgEnvSeq.hist (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (n : ℕ) (ω : Ω) : Iic n → α × R := + fun i ↦ (A i ω, R' i ω) + +@[fun_prop] +lemma IsAlgEnvSeq.measurable_hist (hA : ∀ n, Measurable (A n)) + (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : + Measurable (IsAlgEnvSeq.hist A R' n) := by + unfold IsAlgEnvSeq.hist + fun_prop + +lemma IsAlgEnvSeq.eval_comp_hist (n : ℕ) : + (fun x ↦ x ⟨n, by simp⟩) ∘ (hist A R' n) = step A R' n := rfl + +lemma IsAlgEnvSeq.fst_eval_comp_hist (n : ℕ) : + (fun x ↦ (x ⟨n, by simp⟩).1) ∘ (hist A R' n) = A n := rfl + +lemma IsAlgEnvSeq.snd_eval_comp_hist (n : ℕ) : + (fun x ↦ (x ⟨n, by simp⟩).2) ∘ (hist A R' n) = R' n := rfl + +structure IsAlgEnvSeq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (alg : Algorithm α R) (env : Environment α R) (P : Measure Ω) [IsFiniteMeasure P] : Prop where @@ -71,32 +100,24 @@ structure IsAlgEnvInteraction hasLaw_action_zero : HasLaw (fun ω ↦ (A 0 ω)) alg.p0 P hasCondDistrib_reward_zero : HasCondDistrib (R' 0) (A 0) env.ν0 P hasCondDistrib_action n : - HasCondDistrib (A (n + 1)) (fun ω ↦ fun (i : Iic n) ↦ (A i ω, R' i ω)) (alg.policy n) P + HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) (alg.policy n) P hasCondDistrib_reward n : - HasCondDistrib (R' (n + 1)) (fun ω ↦ (fun (i : Iic n) ↦ (A i ω, R' i ω), A (n + 1) ω)) + HasCondDistrib (R' (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) (env.feedback n) P -def IsAlgEnvInteraction.step (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (n : ℕ) (ω : Ω) : α × R := - (A n ω, R' n ω) - -@[fun_prop] -lemma IsAlgEnvInteraction.measurable_step (n : ℕ) (hA : Measurable (A n)) - (hR' : Measurable (R' n)) : - Measurable (IsAlgEnvInteraction.step A R' n) := by - unfold IsAlgEnvInteraction.step - fun_prop - -def IsAlgEnvInteraction.hist (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (n : ℕ) (ω : Ω) : Iic n → α × R := - fun i ↦ (A i ω, R' i ω) +lemma IsAlgEnvSeq.hasLaw_step_zero + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + (h : IsAlgEnvSeq A R' alg env P) : + HasLaw (step A R' 0) (alg.p0 ⊗ₘ env.ν0) P := + HasLaw.prod_of_hasCondDistrib h.hasLaw_action_zero h.hasCondDistrib_reward_zero -@[fun_prop] -lemma IsAlgEnvInteraction.measurable_hist (hA : ∀ n, Measurable (A n)) - (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : - Measurable (IsAlgEnvInteraction.hist A R' n) := by - unfold IsAlgEnvInteraction.hist - fun_prop +lemma IsAlgEnvSeq.hasCondDistrib_step + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + (h : IsAlgEnvSeq A R' alg env P) (n : ℕ) : + HasCondDistrib (step A R' (n + 1)) (hist A R' n) (stepKernel alg env n) P := + HasCondDistrib.prod (h.hasCondDistrib_action n) (h.hasCondDistrib_reward n) -def IsAlgEnvInteraction.filtration (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) : +def IsAlgEnvSeq.filtration (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) : Filtration ℕ mΩ where seq i := MeasurableSpace.comap (hist A R' i) inferInstance mono' i j hij := by @@ -109,21 +130,21 @@ def IsAlgEnvInteraction.filtration (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, rw [← measurable_iff_comap_le] exact measurable_hist hA hR' i -lemma IsAlgEnvInteraction.measurable_action_filtration +lemma IsAlgEnvSeq.measurable_action_filtration (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : - Measurable[IsAlgEnvInteraction.filtration hA hR' n] (A n) := by + Measurable[IsAlgEnvSeq.filtration hA hR' n] (A n) := by have : A n = (fun h ↦ (h ⟨n, by simp⟩).1) ∘ (hist A R' n) := by ext ω : 1 - simp [IsAlgEnvInteraction.hist] + simp [IsAlgEnvSeq.hist] rw [this] exact measurable_comp_comap _ (by fun_prop) /-- Filtration generated by the history at time `n-1` together with the action at time `n`. -/ -def IsAlgEnvInteraction.filtrationAction +def IsAlgEnvSeq.filtrationAction (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) : Filtration ℕ mΩ where seq n := if n = 0 then MeasurableSpace.comap (A 0) inferInstance - else IsAlgEnvInteraction.filtration hA hR' (n - 1) ⊔ MeasurableSpace.comap (A n) inferInstance + else IsAlgEnvSeq.filtration hA hR' (n - 1) ⊔ MeasurableSpace.comap (A n) inferInstance mono' n m hnm := by simp only by_cases hn : n = 0 @@ -132,8 +153,8 @@ def IsAlgEnvInteraction.filtrationAction · simp only [hn, ↓reduceIte, hm] refine le_sup_of_le_left ?_ rw [← measurable_iff_comap_le] - suffices Measurable[IsAlgEnvInteraction.filtration hA hR' 0] (A 0) from - this.mono ((IsAlgEnvInteraction.filtration hA hR').mono zero_le') le_rfl + suffices Measurable[IsAlgEnvSeq.filtration hA hR' 0] (A 0) from + this.mono ((IsAlgEnvSeq.filtration hA hR').mono zero_le') le_rfl exact measurable_action_filtration hA hR' 0 have hm : m ≠ 0 := by grind simp only [hn, hm, ↓reduceIte] @@ -141,14 +162,14 @@ def IsAlgEnvInteraction.filtrationAction simp only [sup_le_iff] constructor · refine le_sup_of_le_left ?_ - exact (IsAlgEnvInteraction.filtration hA hR').mono hnm' + exact (IsAlgEnvSeq.filtration hA hR').mono hnm' · rcases eq_or_lt_of_le hnm with rfl | hlt · exact le_sup_of_le_right le_rfl refine le_sup_of_le_left ?_ rw [← measurable_iff_comap_le] have h_le : n ≤ m - 1 := by grind - suffices Measurable[IsAlgEnvInteraction.filtration hA hR' n] (A n) from - this.mono ((IsAlgEnvInteraction.filtration hA hR').mono h_le) le_rfl + suffices Measurable[IsAlgEnvSeq.filtration hA hR' n] (A n) from + this.mono ((IsAlgEnvSeq.filtration hA hR').mono h_le) le_rfl exact measurable_action_filtration hA hR' n le' n := by by_cases hn : n = 0 @@ -157,25 +178,25 @@ def IsAlgEnvInteraction.filtrationAction fun_prop simp only [hn, ↓reduceIte, sup_le_iff] constructor - · exact (IsAlgEnvInteraction.filtration hA hR').le _ + · exact (IsAlgEnvSeq.filtration hA hR').le _ · rw [← measurable_iff_comap_le] fun_prop -lemma IsAlgEnvInteraction.filtrationAction_zero_eq_comap +lemma IsAlgEnvSeq.filtrationAction_zero_eq_comap {hA : ∀ n, Measurable (A n)} {hR' : ∀ n, Measurable (R' n)} : filtrationAction hA hR' 0 = MeasurableSpace.comap (A 0) inferInstance := by simp [filtrationAction] -lemma IsAlgEnvInteraction.filtrationAction_eq_comap +lemma IsAlgEnvSeq.filtrationAction_eq_comap {hA : ∀ n, Measurable (A n)} {hR' : ∀ n, Measurable (R' n)} (n : ℕ) (hn : n ≠ 0) : filtrationAction hA hR' n = MeasurableSpace.comap (fun ω ↦ (hist A R' (n - 1) ω, A n ω)) inferInstance := by simp only [filtrationAction, filtration, ← MeasurableSpace.comap_prodMk, hn, ↓reduceIte] rfl -end IsAlgEnvInteraction +end IsAlgEnvSeq -/-- Kernel sending a partial trajectory of the bandit interaction `Iic n → α × ℝ` to a measure +/-- Kernel sending a partial trajectory of the bandit Seq `Iic n → α × ℝ` to a measure on `ℕ → α × ℝ`, supported on full trajectories that start with the partial one. -/ noncomputable def traj (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : Kernel (Iic n → α × R) (ℕ → α × R) := @@ -189,6 +210,36 @@ def trajMeasure (alg : Algorithm α R) (env : Environment α R) : Kernel.trajMeasure (alg.p0 ⊗ₘ env.ν0) (stepKernel alg env) deriving IsProbabilityMeasure +section ModelEquivalence + +variable {Ω Ω' : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + {alg : Algorithm α R} {env : Environment α R} + {P : Measure Ω} [IsProbabilityMeasure P] {P' : Measure Ω'} [IsProbabilityMeasure P'] + {A₁ : ℕ → Ω → α} {R₁ : ℕ → Ω → R} {A₂ : ℕ → Ω' → α} {R₂ : ℕ → Ω' → R} + +theorem eq_trajMeasure_of_isAlgEnvSeq (h : IsAlgEnvSeq A₁ R₁ alg env P) : + P.map (fun ω n ↦ (A₁ n ω, R₁ n ω)) = trajMeasure alg env := by + rw [trajMeasure] + have h := Kernel.eq_trajMeasure (Y := fun n ω ↦ (A₁ n ω, R₁ n ω)) (P := P) + (μ₀ := alg.p0 ⊗ₘ env.ν0) (κ := stepKernel alg env) (fun n ↦ ?_) ?_ (fun n ↦ ?_) + · exact h + · have hA := h.measurable_A n + have hR := h.measurable_R n + fun_prop + · simp only + exact h.hasLaw_step_zero.map_eq + · exact (h.hasCondDistrib_step n).condDistrib_eq + +theorem isAlgEnvSeq_unique (h1 : IsAlgEnvSeq A₁ R₁ alg env P) + (h2 : IsAlgEnvSeq A₂ R₂ alg env P') : + P.map (fun ω n ↦ (A₁ n ω, R₁ n ω)) = P'.map (fun ω n ↦ (A₂ n ω, R₂ n ω)) := by + rw [eq_trajMeasure_of_isAlgEnvSeq h1, eq_trajMeasure_of_isAlgEnvSeq h2] + +end ModelEquivalence + +namespace IT + /-- Action and reward at step `n`. -/ def step (n : ℕ) (h : ℕ → α × R) : α × R := h n @@ -211,30 +262,24 @@ lemma measurable_step (n : ℕ) : Measurable (step n (α := α) (R := R)) := by unfold step; fun_prop @[fun_prop] -lemma measurable_step_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ step p.1 p.2) := by - refine measurable_from_prod_countable_right fun n ↦ ?_ - simp only - fun_prop +lemma measurable_step_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ step p.1 p.2) := + measurable_from_prod_countable_right fun n ↦ (by simp only; fun_prop) @[fun_prop] lemma measurable_action (n : ℕ) : Measurable (action n (α := α) (R := R)) := by unfold action; fun_prop @[fun_prop] -lemma measurable_action_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ action p.1 p.2) := by - refine measurable_from_prod_countable_right fun n ↦ ?_ - simp only - fun_prop +lemma measurable_action_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ action p.1 p.2) := + measurable_from_prod_countable_right fun n ↦ (by simp only; fun_prop) @[fun_prop] lemma measurable_reward (n : ℕ) : Measurable (reward n (α := α) (R := R)) := by unfold reward; fun_prop @[fun_prop] -lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ reward p.1 p.2) := by - refine measurable_from_prod_countable_right fun n ↦ ?_ - simp only - fun_prop +lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ reward p.1 p.2) := + measurable_from_prod_countable_right fun n ↦ (by simp only; fun_prop) @[fun_prop] lemma measurable_hist (n : ℕ) : Measurable (hist n (α := α) (R := R)) := by unfold hist; fun_prop @@ -244,14 +289,14 @@ lemma hist_eq_frestrictLe : ext n h i : 3 simp [hist, Preorder.frestrictLe] -/-- Filtration of the algorithm interaction. -/ +/-- Filtration of the algorithm Seq. -/ protected def filtration (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) := MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R) lemma filtration_eq_comap (n : ℕ) : - Learning.filtration α R n = MeasurableSpace.comap (hist n) inferInstance := by - simp [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe] + IT.filtration α R n = MeasurableSpace.comap (hist n) inferInstance := by + simp [IT.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe] lemma step_eq_eval_comp_hist (n : ℕ) : step (α := α) (R := R) n = (fun x ↦ x ⟨n, by simp⟩) ∘ (hist n) := rfl @@ -262,7 +307,7 @@ lemma action_eq_eval_comp_hist (n : ℕ) : lemma reward_eq_eval_comp_hist (n : ℕ) : reward (α := α) (R := R) n = (fun x ↦ (x ⟨n, by simp⟩).2) ∘ (hist n) := rfl -lemma measurable_step_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (step n) := by +lemma measurable_step_filtration (n : ℕ) : Measurable[IT.filtration α R n] (step n) := by rw [filtration_eq_comap, step_eq_eval_comp_hist] exact measurable_comp_comap _ (by fun_prop) @@ -270,35 +315,35 @@ lemma adapted_step [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace [SecondCountableTopology α] [OpensMeasurableSpace α] [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] [SecondCountableTopology R] [OpensMeasurableSpace R] : - Adapted (Learning.filtration α R) (step (α := α) (R := R)) := + Adapted (IT.filtration α R) (step (α := α) (R := R)) := fun n ↦ (measurable_step_filtration n).stronglyMeasurable -lemma measurable_hist_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (hist n) := by +lemma measurable_hist_filtration (n : ℕ) : Measurable[IT.filtration α R n] (hist n) := by simp [filtration_eq_comap, measurable_iff_comap_le] lemma adapted_hist [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] [SecondCountableTopology α] [OpensMeasurableSpace α] [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] [SecondCountableTopology R] [OpensMeasurableSpace R] : - Adapted (Learning.filtration α R) hist := + Adapted (IT.filtration α R) hist := fun n ↦ (measurable_hist_filtration n).stronglyMeasurable -lemma measurable_action_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (action n) := by +lemma measurable_action_filtration (n : ℕ) : Measurable[IT.filtration α R n] (action n) := by rw [filtration_eq_comap, action_eq_eval_comp_hist] exact measurable_comp_comap _ (by fun_prop) lemma adapted_action [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] [SecondCountableTopology α] [OpensMeasurableSpace α] : - Adapted (Learning.filtration α R) action := + Adapted (IT.filtration α R) action := fun n ↦ (measurable_action_filtration n).stronglyMeasurable -lemma measurable_reward_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (reward n) := by +lemma measurable_reward_filtration (n : ℕ) : Measurable[IT.filtration α R n] (reward n) := by rw [filtration_eq_comap, reward_eq_eval_comp_hist] exact measurable_comp_comap _ (by fun_prop) lemma adapted_reward [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] [SecondCountableTopology R] [OpensMeasurableSpace R] : - Adapted (Learning.filtration α R) reward := + Adapted (IT.filtration α R) reward := fun n ↦ (measurable_reward_filtration n).stronglyMeasurable section FiltrationAction @@ -307,7 +352,7 @@ section FiltrationAction def filtrationAction (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) where seq n := if n = 0 then MeasurableSpace.comap (action 0) inferInstance - else Learning.filtration α R (n - 1) ⊔ MeasurableSpace.comap (action n) inferInstance + else IT.filtration α R (n - 1) ⊔ MeasurableSpace.comap (action n) inferInstance mono' n m hnm := by simp only by_cases hn : n = 0 @@ -316,8 +361,8 @@ def filtrationAction (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : · simp only [hn, ↓reduceIte, hm] refine le_sup_of_le_left ?_ rw [← measurable_iff_comap_le] - suffices Measurable[Learning.filtration α R 0] (action 0) from - this.mono ((Learning.filtration α R).mono zero_le') le_rfl + suffices Measurable[IT.filtration α R 0] (action 0) from + this.mono ((IT.filtration α R).mono zero_le') le_rfl exact measurable_action_filtration 0 have hm : m ≠ 0 := by grind simp only [hn, hm, ↓reduceIte] @@ -325,14 +370,14 @@ def filtrationAction (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : simp only [sup_le_iff] constructor · refine le_sup_of_le_left ?_ - exact (Learning.filtration α R).mono hnm' + exact (IT.filtration α R).mono hnm' · rcases eq_or_lt_of_le hnm with rfl | hlt · exact le_sup_of_le_right le_rfl refine le_sup_of_le_left ?_ rw [← measurable_iff_comap_le] have h_le : n ≤ m - 1 := by grind - suffices Measurable[Learning.filtration α R n] (action n) from - this.mono ((Learning.filtration α R).mono h_le) le_rfl + suffices Measurable[IT.filtration α R n] (action n) from + this.mono ((IT.filtration α R).mono h_le) le_rfl exact measurable_action_filtration n le' n := by by_cases hn : n = 0 @@ -341,7 +386,7 @@ def filtrationAction (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : fun_prop simp only [hn, ↓reduceIte, sup_le_iff] constructor - · exact (Learning.filtration α R).le _ + · exact (IT.filtration α R).le _ · rw [← measurable_iff_comap_le] fun_prop @@ -356,28 +401,28 @@ lemma filtrationAction_eq_comap (n : ℕ) (hn : n ≠ 0) : rfl lemma filtration_le_filtrationAction_add_one (n : ℕ) : - Learning.filtration α R n ≤ filtrationAction α R (n + 1) := le_sup_of_le_left le_rfl + IT.filtration α R n ≤ filtrationAction α R (n + 1) := le_sup_of_le_left le_rfl lemma filtration_le_filtrationAction {m n : ℕ} (h : n < m) : - Learning.filtration α R n ≤ filtrationAction α R m := by + IT.filtration α R n ≤ filtrationAction α R m := by have h' : n + 1 ≤ m := by grind exact (filtration_le_filtrationAction_add_one n).trans ((filtrationAction α R).mono h') lemma filtrationAction_le_filtration_self (n : ℕ) : - filtrationAction α R n ≤ Learning.filtration α R n := by + filtrationAction α R n ≤ IT.filtration α R n := by by_cases hn : n = 0 · simp only [hn, filtrationAction_zero_eq_comap] rw [← measurable_iff_comap_le] exact measurable_action_filtration 0 simp only [filtrationAction, hn, ↓reduceIte, sup_le_iff] constructor - · exact (Learning.filtration α R).mono (by grind) + · exact (IT.filtration α R).mono (by grind) · rw [← measurable_iff_comap_le] exact measurable_action_filtration _ lemma filtrationAction_le_filtration {m n : ℕ} (h : m ≤ n) : - filtrationAction α R m ≤ Learning.filtration α R n := - (filtrationAction_le_filtration_self m).trans ((Learning.filtration α R).mono h) + filtrationAction α R m ≤ IT.filtration α R n := + (filtrationAction_le_filtration_self m).trans ((IT.filtration α R).mono h) lemma measurable_action_filtrationAction (n : ℕ) : Measurable[filtrationAction α R n] (action n) := by @@ -444,8 +489,8 @@ lemma condDistrib_reward (alg : Algorithm α R) (env : Environment α R) (n : Measure.map_map (by fun_prop) (by fun_prop)] rfl -lemma isAlgEnvInteraction_trajMeasure (alg : Algorithm α R) (env : Environment α R) : - IsAlgEnvInteraction action reward alg env (trajMeasure alg env) where +lemma isAlgEnvSeq_trajMeasure (alg : Algorithm α R) (env : Environment α R) : + IsAlgEnvSeq action reward alg env (trajMeasure alg env) where hasLaw_action_zero := hasLaw_action_zero alg env hasCondDistrib_reward_zero := ⟨by fun_prop, by fun_prop, condDistrib_reward_zero alg env⟩ hasCondDistrib_action n := ⟨by fun_prop, by fun_prop, condDistrib_action alg env n⟩ @@ -453,19 +498,6 @@ lemma isAlgEnvInteraction_trajMeasure (alg : Algorithm α R) (env : Environment end Laws -section ModelEquivalence - -variable {Ω Ω' : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} - [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - {alg : Algorithm α R} {env : Environment α R} - {P : Measure Ω} [IsFiniteMeasure P] {P' : Measure Ω'} [IsFiniteMeasure P'] - {A₁ : ℕ → Ω → α} {R₁ : ℕ → Ω → R} {A₂ : ℕ → Ω' → α} {R₂ : ℕ → Ω' → R} - -theorem isAlgEnvInteraction_unique (h1 : IsAlgEnvInteraction A₁ R₁ alg env P) - (h2 : IsAlgEnvInteraction A₂ R₂ alg env P') : - P.map (fun ω n ↦ (A₁ n ω, R₁ n ω)) = P'.map (fun ω n ↦ (A₂ n ω, R₂ n ω)) := by - sorry - -end ModelEquivalence +end IT end Learning diff --git a/LeanBandits/SequentialLearning/Deterministic.lean b/LeanBandits/SequentialLearning/Deterministic.lean index f2240745..2c1ba945 100644 --- a/LeanBandits/SequentialLearning/Deterministic.lean +++ b/LeanBandits/SequentialLearning/Deterministic.lean @@ -29,26 +29,69 @@ def detAlgorithm (nextaction : (n : ℕ) → (Iic n → α × R) → α) variable {nextaction : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextaction n)} {action0 : α} {env : Environment α R} +section IsAlgEnvSeq + +variable {Ω : Type*} {mΩ : MeasurableSpace Ω} + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] + {P : Measure Ω} [IsProbabilityMeasure P] {A : ℕ → Ω → α} {R' : ℕ → Ω → R} + +lemma IsAlgEnvSeq.HasLaw_action_zero_detAlgorithm + (h : IsAlgEnvSeq A R' (detAlgorithm nextaction h_next action0) env P) : + HasLaw (A 0) (Measure.dirac action0) P where + aemeasurable := have hA := h.measurable_A; by fun_prop + map_eq := (hasLaw_action_zero h).map_eq + +lemma IsAlgEnvSeq.action_zero_detAlgorithm + (h : IsAlgEnvSeq A R' (detAlgorithm nextaction h_next action0) env P) : + A 0 =ᵐ[P] fun _ ↦ action0 := by + have h_eq : ∀ᵐ x ∂(P.map (A 0)), x = action0 := by + rw [(hasLaw_action_zero h).map_eq] + simp [detAlgorithm] + have hA := h.measurable_A + exact ae_of_ae_map (by fun_prop) h_eq + +lemma IsAlgEnvSeq.action_detAlgorithm_ae_eq + (h : IsAlgEnvSeq A R' (detAlgorithm nextaction h_next action0) env P) (n : ℕ) : + A (n + 1) =ᵐ[P] fun ω ↦ nextaction n (hist A R' n ω) := by + have hA := h.measurable_A + have hR' := h.measurable_R + exact ae_eq_of_condDistrib_eq_deterministic (by fun_prop) (by fun_prop) (by fun_prop) + (h.hasCondDistrib_action n).condDistrib_eq + +lemma IsAlgEnvSeq.action_detAlgorithm_ae_all_eq + (h : IsAlgEnvSeq A R' (detAlgorithm nextaction h_next action0) env P) : + ∀ᵐ ω ∂P, A 0 ω = action0 ∧ ∀ n, A (n + 1) ω = nextaction n (hist A R' n ω) := by + rw [eventually_and, ae_all_iff] + exact ⟨action_zero_detAlgorithm h, action_detAlgorithm_ae_eq h⟩ + +end IsAlgEnvSeq + +namespace IT + local notation "𝔓" => trajMeasure (detAlgorithm nextaction h_next action0) env -lemma HasLaw_action_zero_detAlgorithm : HasLaw (action 0) (Measure.dirac action0) 𝔓 where - map_eq := (hasLaw_action_zero _ _).map_eq +lemma HasLaw_action_zero_detAlgorithm : HasLaw (IT.action 0) (Measure.dirac action0) 𝔓 where + map_eq := (IT.hasLaw_action_zero _ _).map_eq -lemma action_zero_detAlgorithm [MeasurableSingletonClass α] : action 0 =ᵐ[𝔓] fun _ ↦ action0 := by - have h_eq : ∀ᵐ x ∂((𝔓).map (action 0)), x = action0 := by - rw [(hasLaw_action_zero _ _).map_eq] +lemma action_zero_detAlgorithm [MeasurableSingletonClass α] : + IT.action 0 =ᵐ[𝔓] fun _ ↦ action0 := by + have h_eq : ∀ᵐ x ∂((𝔓).map (IT.action 0)), x = action0 := by + rw [(IT.hasLaw_action_zero _ _).map_eq] simp [detAlgorithm] exact ae_of_ae_map (by fun_prop) h_eq lemma action_detAlgorithm_ae_eq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] - [Nonempty R] (n : ℕ) : action (n + 1) =ᵐ[𝔓] fun h ↦ nextaction n (hist n h) := + [Nonempty R] (n : ℕ) : IT.action (n + 1) =ᵐ[𝔓] fun h ↦ nextaction n (IT.hist n h) := ae_eq_of_condDistrib_eq_deterministic (by fun_prop) (by fun_prop) (by fun_prop) - (condDistrib_action (detAlgorithm nextaction h_next action0) env n) + (IT.condDistrib_action (detAlgorithm nextaction h_next action0) env n) lemma action_detAlgorithm_ae_all_eq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] : - ∀ᵐ h ∂𝔓, action 0 h = action0 ∧ ∀ n, action (n + 1) h = nextaction n (hist n h) := by + ∀ᵐ h ∂𝔓, IT.action 0 h = action0 ∧ ∀ n, IT.action (n + 1) h = nextaction n (IT.hist n h) := by rw [eventually_and, ae_all_iff] exact ⟨action_zero_detAlgorithm, action_detAlgorithm_ae_eq⟩ +end IT + end Learning diff --git a/LeanBandits/SequentialLearning/Draft.lean b/LeanBandits/SequentialLearning/Draft.lean deleted file mode 100644 index 5d3a5c71..00000000 --- a/LeanBandits/SequentialLearning/Draft.lean +++ /dev/null @@ -1,772 +0,0 @@ -/- -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, Paulo Rauber --/ -import LeanBandits.SequentialLearning.Algorithm -import Mathlib.Order.CompletePartialOrder -import Mathlib.Probability.Martingale.BorelCantelli - -/-! -# Bookkeeping definitions for finite action space sequential learning problems - -If the number of actions is finite, it makes sense to define the number of times each action was -chosen, the time at which an action was chosen for the nth time, the value of the reward at that -time, the sum of rewards obtained for each action, the empirical mean reward for each action, etc. - -For each definition that take as arguments a time `t : ℕ`, a history `h : ℕ → α × R`, and possibly -other parameters, we put the time and history at the end in this order, so that the definition can -be seen as a stochastic process indexed by time `t` on the measurable space `ℕ → α × R`. - --/ - -open MeasureTheory Finset Learning - -namespace LearningDraft - -variable {α R Ω : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} {mΩ : MeasurableSpace Ω} - [DecidableEq α] - {alg : Algorithm α R} {env : Environment α R} - {P : Measure Ω} [IsProbabilityMeasure P] - {A : ℕ → Ω → α} {R' : ℕ → Ω → R} - {a : α} {m n t : ℕ} {ω : Ω} - -section PullCount - -/-- Number of times action `a` was chosen up to time `t` (excluding `t`). -/ -noncomputable -def pullCount (A : ℕ → Ω → α) (a : α) (t : ℕ) (ω : Ω) : ℕ := - #(filter (fun s ↦ A s ω = a) (range t)) - -/-- Number of pulls of arm `a` up to (and including) time `n`. -This is the number of entries in `h` in which the arm is `a`. -/ -noncomputable -def pullCount' (n : ℕ) (h : Iic n → α × R) (a : α) := #{s | (h s).1 = a} - -@[simp] -lemma pullCount_zero (a : α) : pullCount A a 0 = 0 := by ext; simp [pullCount] - -lemma pullCount_zero_apply (a : α) (ω : Ω) : pullCount A a 0 ω = 0 := by simp - -lemma pullCount_one : pullCount A a 1 ω = if A 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] - -lemma monotone_pullCount (a : α) (ω : Ω) : Monotone (pullCount A a · ω) := - fun _ _ _ ↦ card_le_card (filter_subset_filter _ (by simpa)) - -@[mono, gcongr] -lemma pullCount_mono (a : α) {n m : ℕ} (hnm : n ≤ m) (ω : Ω) : - pullCount A a n ω ≤ pullCount A a m ω := - monotone_pullCount a ω hnm - -lemma pullCount_action_eq_pullCount_add_one (t : ℕ) (ω : Ω) : - pullCount A (A t ω) (t + 1) ω = pullCount A (A t ω) t ω + 1 := by - simp [pullCount, range_add_one, filter_insert] - -lemma pullCount_eq_pullCount_of_action_ne (ha : A t ω ≠ a) : - pullCount A a (t + 1) ω = pullCount A a t ω := by - simp [pullCount, range_add_one, filter_insert, ha] - -lemma pullCount_add_one : - pullCount A a (t + 1) ω = pullCount A a t ω + if A t ω = a then 1 else 0 := by - split_ifs with h - · rw [← h, pullCount_action_eq_pullCount_add_one] - · rw [pullCount_eq_pullCount_of_action_ne h, add_zero] - -lemma pullCount_eq_sum (a : α) (t : ℕ) (ω : Ω) : - pullCount A a t ω = ∑ s ∈ range t, if A s ω = a then 1 else 0 := by simp [pullCount] - -lemma pullCount'_eq_sum (n : ℕ) (h : Iic n → α × R) (a : α) : - pullCount' n h a = ∑ s : Iic n, if (h s).1 = a then 1 else 0 := by simp [pullCount'] - -lemma pullCount_add_one_eq_pullCount' {n : ℕ} {ω : Ω} : - pullCount A a (n + 1) ω = pullCount' n (fun i ↦ (A i ω, R' i ω)) a := by - rw [pullCount_eq_sum, pullCount'_eq_sum] - rw [Finset.sum_coe_sort (f := fun s ↦ if A s ω = a then 1 else 0) (Iic n)] - congr with m - simp only [mem_range, mem_Iic] - grind - -lemma pullCount_eq_pullCount' {n : ℕ} {ω : Ω} (hn : n ≠ 0) : - pullCount A a n ω = pullCount' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by - cases n with - | zero => exact absurd rfl hn - | succ n => - rw [pullCount_add_one_eq_pullCount' (R' := R')] - have : n + 1 - 1 = n := by simp - exact this ▸ rfl - -lemma pullCount_le (a : α) (t : ℕ) (ω : Ω) : pullCount A a t ω ≤ t := - (card_filter_le _ _).trans_eq (by simp) - -lemma pullCount_congr {ω' : Ω} (h_eq : ∀ i ≤ n, A i ω = A i ω') : - pullCount A a (n + 1) ω = pullCount A a (n + 1) ω' := by - unfold pullCount - congr 1 with s - simp only [mem_filter, mem_range, and_congr_right_iff] - intro hs - rw [Nat.lt_add_one_iff] at hs - rw [h_eq s hs] - -lemma pullCount_lt_of_forall_ne (h_lt : ∀ s, pullCount A a (s + 1) ω ≠ t) (ht : t ≠ 0) : - pullCount A a n ω < t := by - induction n with - | zero => simpa using ht.bot_lt - | succ n hn => - specialize h_lt n - rw [pullCount_add_one] at h_lt ⊢ - grind - -lemma exists_pullCount_eq_of_le (hnm : t ≤ pullCount A a (n + 1) ω) (ht : t ≠ 0) : - ∃ s, pullCount A a (s + 1) ω = t := by - by_contra! h_contra - refine lt_irrefl (pullCount A a (n + 1) ω) ?_ - refine lt_of_lt_of_le ?_ hnm - exact pullCount_lt_of_forall_ne h_contra ht - -section Measurability - -@[fun_prop] -lemma measurable_pullCount [MeasurableSingletonClass α] (hA : ∀ n, Measurable (A n)) - (a : α) (t : ℕ) : - Measurable (fun ω : Ω ↦ pullCount A a t ω) := by - simp_rw [pullCount_eq_sum] - have h_meas s : Measurable (fun ω : Ω ↦ if A s ω = a then 1 else 0) := by - refine Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact (measurableSet_singleton _).preimage (by fun_prop) - fun_prop - -@[fun_prop] -lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : - Measurable (fun h : Iic n → α × R ↦ pullCount' n h a) := by - simp_rw [pullCount'_eq_sum] - have h_meas s : Measurable (fun (h : Iic n → α × R) ↦ if (h s).1 = a then 1 else 0) := by - refine Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact (measurableSet_singleton _).preimage (by fun_prop) - fun_prop - -lemma adapted_pullCount_add_one' [MeasurableSingletonClass α] - (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (n : ℕ) : - Measurable[IsAlgEnvInteraction.filtration hA hR' n] (pullCount A a (n + 1)) := by - have : pullCount A a (n + 1) = (fun h : Iic n → α × R ↦ pullCount' n h a) ∘ - (IsAlgEnvInteraction.hist A R' n) := by - ext - exact pullCount_add_one_eq_pullCount' - rw [IsAlgEnvInteraction.filtration, this] - exact measurable_comp_comap _ (measurable_pullCount' n a) - -lemma adapted_pullCount_add_one [MeasurableSingletonClass α] - (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) : - Adapted (IsAlgEnvInteraction.filtration hA hR') (fun n ↦ pullCount A a (n + 1)) := - fun n ↦ Measurable.stronglyMeasurable <| adapted_pullCount_add_one' hA hR' a n - -lemma isPredictable_pullCount [MeasurableSingletonClass α] - (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) : - IsPredictable (IsAlgEnvInteraction.filtration hA hR') (pullCount A a) := by - rw [isPredictable_iff_measurable_add_one] - refine ⟨?_, fun n ↦ (adapted_pullCount_add_one hA hR' a n).measurable⟩ - simp only [pullCount_zero] - fun_prop - -end Measurability - -end PullCount - -section StepsUntil - --- TODO: replace this by leastGE, once leastGE is generalized -/-- Number of steps until action `a` was pulled exactly `m` times. -/ -noncomputable -def stepsUntil (A : ℕ → Ω → α) (a : α) (m : ℕ) (ω : Ω) : ℕ∞ := - sInf ((↑) '' {s | pullCount A a (s + 1) ω = m}) - -lemma stepsUntil_eq_top_iff : stepsUntil A a m ω = ⊤ ↔ ∀ s, pullCount A a (s + 1) ω ≠ m := by - simp [stepsUntil, sInf_eq_top] - -lemma stepsUntil_ne_top (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : stepsUntil A a m ω ≠ ⊤ := by - simpa [stepsUntil_eq_top_iff] - -lemma exists_pullCount_eq (h' : stepsUntil A a m ω ≠ ⊤) : - ∃ s, pullCount A a (s + 1) ω = m := by - by_contra! h_contra - rw [← stepsUntil_eq_top_iff] at h_contra - simp [h_contra] at h' - -lemma stepsUntil_zero_of_ne (hka : A 0 ω ≠ a) : stepsUntil A a 0 ω = 0 := by - unfold stepsUntil - simp_rw [← bot_eq_zero, sInf_eq_bot, bot_eq_zero] - intro n hn - refine ⟨0, ?_, hn⟩ - simp only [Set.mem_image, Set.mem_setOf_eq, Nat.cast_eq_zero, exists_eq_right, zero_add] - rw [← zero_add 1, pullCount_eq_pullCount_of_action_ne hka] - simp - -lemma stepsUntil_zero_of_eq (hka : A 0 ω = a) : stepsUntil A a 0 ω = ⊤ := by - rw [stepsUntil_eq_top_iff] - suffices 0 < pullCount A a 1 ω by - intro n hn - refine lt_irrefl 0 ?_ - exact this.trans_le (le_trans (monotone_pullCount _ _ (by omega)) hn.le) - rw [← hka, ← zero_add 1, pullCount_action_eq_pullCount_add_one] - simp - -lemma stepsUntil_eq_dite (a : α) (m : ℕ) (ω : Ω) - [Decidable (∃ s, pullCount A a (s + 1) ω = m)] : - stepsUntil A a m ω = - if h : ∃ s, pullCount A a (s + 1) ω = m then (Nat.find h : ℕ∞) else ⊤ := by - unfold stepsUntil - split_ifs with h' - · refine le_antisymm ?_ ?_ - · refine sInf_le ?_ - simpa using Nat.find_spec h' - · simp only [le_sInf_iff, Set.mem_image, Set.mem_setOf_eq, forall_exists_index, and_imp, - forall_apply_eq_imp_iff₂, Nat.cast_le, Nat.find_le_iff] - exact fun n hn ↦ ⟨n, le_rfl, hn⟩ - · push_neg at h' - suffices {s | pullCount A a (s + 1) ω = m} = ∅ by simp [this] - ext s - simpa using (h' s) - --- todo: this is in ℝ because of the limited def of leastGE -lemma stepsUntil_eq_leastGE (a : α) (hm : m ≠ 0) : - stepsUntil A a m = leastGE (fun n (ω : Ω) ↦ pullCount A a (n + 1) ω) m := by - classical - ext ω - rw [stepsUntil_eq_dite] - unfold leastGE hittingAfter - simp only [zero_le, Set.mem_Ici, Nat.cast_le, true_and, ENat.some_eq_coe] - have h_iff : (∃ s, pullCount A a (s + 1) ω = m) ↔ (∃ s, m ≤ pullCount A a (s + 1) ω) := by - refine ⟨fun ⟨s, hs⟩ ↦ ⟨s, hs.ge⟩, fun ⟨s, hs⟩ ↦ ?_⟩ - exact exists_pullCount_eq_of_le hs hm - by_cases h_exists : ∃ s, m ≤ pullCount A a (s + 1) ω - swap; · simp_rw [h_iff]; simp [h_exists] - rw [if_pos h_exists, dif_pos] - swap; · rwa [h_iff] - norm_cast - rw [Nat.find_eq_iff] - constructor - · apply le_antisymm - · by_contra! h_contra - obtain ⟨s, hs⟩ : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq_of_le h_contra.le hm - rw [← hs] at h_contra - refine h_contra.not_ge ?_ - gcongr - exact csInf_le (by simp) (by simp) - · exact Nat.sInf_mem (s := {j | m ≤ pullCount A a (j + 1) ω}) h_exists - · intro n hn h_contra - refine hn.not_ge ?_ - exact csInf_le (by simp) (by simp [h_contra]) - -lemma stepsUntil_pullCount_le (ω : Ω) (a : α) (t : ℕ) : - stepsUntil A a (pullCount A a (t + 1) ω) ω ≤ t := by - rw [stepsUntil] - exact csInf_le (OrderBot.bddBelow _) ⟨t, rfl, rfl⟩ - -lemma stepsUntil_pullCount_eq (ω : Ω) (t : ℕ) : - stepsUntil A (A t ω) (pullCount A (A t ω) (t + 1) ω) ω = t := by - apply le_antisymm (stepsUntil_pullCount_le ω (A t ω) t) - suffices ∀ t', pullCount A (A t ω) (t' + 1) ω = pullCount A (A t ω) t ω + 1 → t ≤ t' by - simpa [stepsUntil, pullCount_action_eq_pullCount_add_one] - exact fun t' h' ↦ Nat.le_of_lt_succ ((monotone_pullCount (A t ω) ω).reflect_lt - (h' ▸ lt_add_one _)) - -/-- If we pull action `a` at time 0, the first time at which it is pulled once is 0. -/ -lemma stepsUntil_one_of_eq (hka : A 0 ω = a) : stepsUntil A a 1 ω = 0 := by - classical - have h_pull : pullCount A a 1 ω = 1 := by simp [pullCount_one, hka] - have h_le := stepsUntil_pullCount_le (A := A) ω a 0 - simpa [h_pull] using h_le - -lemma stepsUntil_eq_zero_iff : - stepsUntil A a m ω = 0 ↔ (m = 0 ∧ A 0 ω ≠ a) ∨ (m = 1 ∧ A 0 ω = a) := by - classical - refine ⟨fun h' ↦ ?_, fun h' ↦ ?_⟩ - · have h_exists : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq (by simp [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 : A 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 action_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : - A (stepsUntil A a m ω).toNat ω = a := by - classical - simp only [stepsUntil_eq_dite, h_exists, ↓reduceDIte, ENat.toNat_coe] - have h_spec := Nat.find_spec h_exists - have h_spec' n := Nat.find_min h_exists (m := n) - by_cases h_zero : Nat.find h_exists = 0 - · simp only [h_zero, zero_add, not_lt_zero', IsEmpty.forall_iff, implies_true] at * - by_contra h_ne - rw [← zero_add 1, pullCount_eq_pullCount_of_action_ne h_ne] at h_spec - simp only [pullCount_zero] at h_spec - exact hm h_spec.symm - have h_pos : 0 < Nat.find h_exists := Nat.pos_of_ne_zero h_zero - by_contra h_ne - refine h_spec' (Nat.find h_exists - 1) ?_ ?_ - · simp [h_pos] - rw [Nat.sub_add_cancel (by omega)] - rwa [← pullCount_eq_pullCount_of_action_ne] - exact h_ne - -lemma action_eq_of_stepsUntil_eq_coe (hm : m ≠ 0) (h : stepsUntil A a m ω = n) : - A n ω = a := by - have : n = (stepsUntil A a m ω).toNat := by simp [h] - rw [this] - have h_exists : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq (by simp [h]) - exact action_stepsUntil hm h_exists - -lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : - pullCount A a (stepsUntil A a m ω + 1).toNat ω = m := by - classical - have h_eq := stepsUntil_eq_dite (A := A) a m ω - simp only [h_exists, ↓reduceDIte] at h_eq - have h' := Nat.find_spec h_exists - rw [h_eq] - rw [ENat.toNat_add (by simp) (by simp)] - simp only [ENat.toNat_coe, ENat.toNat_one] - exact h' - -lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : - pullCount A a (stepsUntil A a m ω).toNat ω = m - 1 := by - have h_action := action_eq_of_stepsUntil_eq_coe (A := A) (n := (stepsUntil A a m ω).toNat) - (a := a) (ω := ω) hm ?_ - swap; · symm; simpa [stepsUntil_eq_top_iff] - have h_add_one := pullCount_stepsUntil_add_one h_exists - nth_rw 1 [← h_action] at h_add_one - rw [ENat.toNat_add ?_ (by simp), ENat.toNat_one, pullCount_action_eq_pullCount_add_one] - at h_add_one - swap; · simpa [stepsUntil_eq_top_iff] - grind - -lemma pullCount_lt_of_le_stepsUntil (a : α) {n m : ℕ} (ω : Ω) - (h_exists : ∃ s, pullCount A a (s + 1) ω = m) (hn : n < stepsUntil A a m ω) : - pullCount A a (n + 1) ω < m := by - classical - have h_eq := stepsUntil_eq_dite (A := A) a m ω - simp only [h_exists, ↓reduceDIte] at h_eq - rw [← ENat.coe_toNat (stepsUntil_ne_top h_exists)] at hn - refine lt_of_le_of_ne ?_ ?_ - · calc pullCount A a (n + 1) ω - _ ≤ pullCount A a (stepsUntil A a m ω + 1).toNat ω := by - refine monotone_pullCount a ω ?_ - rw [ENat.toNat_add (stepsUntil_ne_top h_exists) (by simp)] - simp only [ENat.toNat_one, add_le_add_iff_right] - exact mod_cast hn.le - _ = m := pullCount_stepsUntil_add_one h_exists - · refine Nat.find_min h_exists (m := n) ?_ - suffices n < (stepsUntil A a m ω).toNat by - rwa [h_eq, ENat.toNat_coe] at this - exact mod_cast hn - -lemma pullCount_eq_of_stepsUntil_eq_coe {ω : Ω} (hm : m ≠ 0) - (h : stepsUntil A a m ω = n) : - pullCount A a n ω = m - 1 := by - have : n = (stepsUntil A a m ω).toNat := by simp [h] - rw [this, pullCount_stepsUntil hm] - exact exists_pullCount_eq (by simp [h]) - -lemma pullCount_add_one_eq_of_stepsUntil_eq_coe {ω : Ω} - (h : stepsUntil A a m ω = n) : - pullCount A a (n + 1) ω = m := by - have : n + 1 = (stepsUntil A a m ω + 1).toNat := by - rw [ENat.toNat_add (by simp [h]) (by simp)]; simp [h] - rw [this, pullCount_stepsUntil_add_one] - exact exists_pullCount_eq (by simp [h]) - -lemma stepsUntil_eq_iff {ω : Ω} (n : ℕ) : - stepsUntil A a m ω = n ↔ - pullCount A a (n + 1) ω = m ∧ (∀ k < n, pullCount A a (k + 1) ω < m) := by - refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ - · have h_exists : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq (by simp [h]) - refine ⟨pullCount_add_one_eq_of_stepsUntil_eq_coe h, fun k hk ↦ ?_⟩ - exact pullCount_lt_of_le_stepsUntil a ω h_exists (by rw [h]; exact mod_cast hk) - · classical - rw [stepsUntil_eq_dite a m ω, dif_pos ⟨n, h.1⟩] - simp only [Nat.cast_inj] - rw [Nat.find_eq_iff] - exact ⟨h.1, fun k hk ↦ (h.2 k hk).ne⟩ - -lemma stepsUntil_eq_iff' {ω : Ω} (hm : m ≠ 0) (n : ℕ) : - stepsUntil A a m ω = n ↔ A n ω = a ∧ pullCount A a n ω = m - 1 := by - by_cases hn : n = 0 - · simp [hn, stepsUntil_eq_zero_iff, hm] - grind - rw [stepsUntil_eq_iff n] - refine ⟨fun ⟨h1, h2⟩ ↦ ⟨?_, ?_⟩, fun ⟨h1, h2⟩ ↦ ⟨?_, fun k hk ↦ ?_⟩⟩ - · rw [pullCount_add_one] at h1 - specialize h2 (n - 1) (by lia) - grind - · rw [pullCount_add_one] at h1 - specialize h2 (n - 1) (by lia) - grind - · rw [pullCount_add_one, h1, h2] - grind - · rw [Nat.lt_iff_le_pred (by grind)] - rw [← h2] - refine monotone_pullCount a ω ?_ - grind - -lemma stepsUntil_eq_congr {ω' : Ω} (h_eq : ∀ i ≤ n, A i ω = A i ω') : - stepsUntil A a m ω = n ↔ stepsUntil A a m ω' = n := by - simp_rw [stepsUntil_eq_iff n] - congr! 1 - · rw [pullCount_congr h_eq] - · congr! 3 with k hk - rw [pullCount_congr] - grind - -section Measurability - -lemma isStoppingTime_stepsUntil [MeasurableSingletonClass α] - (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (hm : m ≠ 0) : - IsStoppingTime (IsAlgEnvInteraction.filtration hA hR') (stepsUntil A a m) := by - rw [stepsUntil_eq_leastGE _ hm] - refine Adapted.isStoppingTime_leastGE _ fun n ↦ ?_ - suffices StronglyMeasurable[IsAlgEnvInteraction.filtration hA hR' n] (pullCount A a (n + 1)) by - fun_prop - exact adapted_pullCount_add_one hA hR' a n - --- todo: get this from the stopping time property? -@[fun_prop] -lemma measurable_stepsUntil [MeasurableSingletonClass α] - (hA : ∀ n, Measurable (A n)) (a : α) (m : ℕ) : - Measurable (stepsUntil A a m) := by - classical - have h_union : {h' : Ω | ∃ s, pullCount A a (s + 1) h' = m} - = ⋃ s : ℕ, {h' | pullCount A a (s + 1) h' = m} := by ext; simp - have h_meas_set : MeasurableSet {h' : Ω | ∃ s, pullCount A a (s + 1) h' = m} := by - rw [h_union] - refine MeasurableSet.iUnion fun s ↦ (measurableSet_singleton _).preimage ?_ - exact measurable_pullCount hA a (s + 1) - --simp_rw [stepsUntil_eq_dite] - suffices Measurable fun k ↦ if h : k ∈ {k' | ∃ s, pullCount A a (s + 1) k' = m} - then (Nat.find h : ℕ∞) else ⊤ by - convert this with ω - rw [stepsUntil_eq_dite a m ω] - rfl - refine Measurable.dite (s := {k' : Ω | ∃ s, pullCount A a (s + 1) k' = m}) - (f := fun x ↦ (Nat.find x.2 : ℕ∞)) (g := fun _ ↦ ⊤) ?_ (by fun_prop) h_meas_set - refine Measurable.coe_nat_enat ?_ - refine measurable_find _ fun k ↦ ?_ - suffices MeasurableSet {x : Ω | pullCount A a (k + 1) x = m} by - have : Subtype.val '' {x : {k' : Ω | - ∃ s, pullCount A a (s + 1) k' = m} | pullCount A a (k + 1) (x : Ω) = m} - = {x : Ω | pullCount A a (k + 1) x = m} := by - ext x - simp only [Set.mem_setOf_eq, Set.coe_setOf, Set.mem_image, Subtype.exists, exists_and_left, - exists_prop, exists_eq_right_right, and_iff_left_iff_imp] - exact fun h ↦ ⟨_, h⟩ - refine (MeasurableEmbedding.subtype_coe h_meas_set).measurableSet_image.mp ?_ - rw [this] - exact (measurableSet_singleton _).preimage (by fun_prop) - exact (measurableSet_singleton _).preimage (by fun_prop) - -lemma measurable_stepsUntil' [MeasurableSingletonClass α] - (hA : ∀ n, Measurable (A n)) (a : α) (m : ℕ) : - Measurable (fun ω : Ω × (ℕ → α → R) ↦ stepsUntil A a m ω.1) := - (measurable_stepsUntil hA a m).comp measurable_fst - -lemma measurable_comap_indicator_stepsUntil_eq [MeasurableSingletonClass α] [Nonempty R] - (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : - Measurable[MeasurableSpace.comap - (fun ω : Ω ↦ (IsAlgEnvInteraction.hist A R' (n-1) ω, A n ω)) inferInstance] - ({ω | stepsUntil A a m ω = ↑n}.indicator fun _ ↦ 1) := by - by_cases hm : m = 0 - · simp only [hm] - by_cases hn : n = 0 - · simp only [hn, CharP.cast_eq_zero, stepsUntil_eq_zero_iff, ne_eq, true_and, zero_ne_one, - false_and, or_false] - refine Measurable.indicator measurable_const ?_ - refine (measurableSet_singleton _).compl.preimage ?_ - rw [measurable_iff_comap_le] - rw [Prod.instMeasurableSpace, MeasurableSpace.comap_prodMk] - exact le_sup_of_le_right le_rfl - · have : {ω | stepsUntil A a 0 ω = n} = ∅ := by - ext ω - by_cases ha : A 0 ω = a - · simp [stepsUntil_zero_of_eq ha] - · simp only [Set.mem_setOf_eq, stepsUntil_zero_of_ne ha, Set.mem_empty_iff_false, - iff_false] - norm_cast - exact Ne.symm hn - simp [this] - simp_rw [stepsUntil_eq_iff' hm] - refine Measurable.indicator measurable_const ?_ - refine ((measurableSet_singleton _).preimage ?_).inter ((measurableSet_singleton _).preimage ?_) - · rw [measurable_iff_comap_le] - rw [Prod.instMeasurableSpace, MeasurableSpace.comap_prodMk] - exact le_sup_of_le_right le_rfl - · rw [measurable_iff_comap_le] - rw [Prod.instMeasurableSpace, MeasurableSpace.comap_prodMk] - refine le_sup_of_le_left ?_ - rw [← measurable_iff_comap_le] - by_cases hn : n = 0 - · simp only [hn, pullCount_zero] - exact measurable_const - have h_meas := adapted_pullCount_add_one' hA hR' a (n - 1) - rwa [Nat.sub_add_cancel (by lia)] at h_meas - -lemma measurable_indicator_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] - (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : - Measurable ({ω : Ω | stepsUntil A a m ω = ↑n}.indicator fun _ ↦ 1) := by - refine (measurable_comap_indicator_stepsUntil_eq hA hR' a m n).mono ?_ le_rfl - refine Measurable.comap_le ?_ - fun_prop - -lemma measurableSet_stepsUntil_eq_zero [MeasurableSingletonClass α] (a : α) (m : ℕ) : - MeasurableSet[MeasurableSpace.comap (A 0) inferInstance] - {ω : Ω | stepsUntil A a m ω = 0} := by - simp only [stepsUntil_eq_zero_iff (a := a) (m := m), ne_eq] - by_cases hm : m = 0 - · simp only [hm, true_and, zero_ne_one, false_and, or_false] - refine (measurableSet_singleton _).compl.preimage ?_ - rw [measurable_iff_comap_le] - by_cases hm1 : m = 1 - swap; · simp [hm, hm1] - simp only [hm1, one_ne_zero, false_and, true_and, false_or] - refine (measurableSet_singleton _).preimage ?_ - rw [measurable_iff_comap_le] - -lemma measurable_comap_indicator_stepsUntil_eq_zero [MeasurableSingletonClass α] (a : α) (m : ℕ) : - Measurable[MeasurableSpace.comap (A 0) inferInstance] - ({ω | stepsUntil A a m ω = 0}.indicator fun _ ↦ 1) := by - rw [measurable_indicator_const_iff] - exact measurableSet_stepsUntil_eq_zero a m - -lemma measurableSet_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] - (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : - MeasurableSet[MeasurableSpace.comap (fun ω : Ω ↦ (IsAlgEnvInteraction.hist A R' (n-1) ω, A n ω)) - inferInstance] - {ω : Ω | stepsUntil A a m ω = ↑n} := by - let mProd := MeasurableSpace.comap - (fun ω : Ω ↦ (IsAlgEnvInteraction.hist A R' (n-1) ω, A n ω)) inferInstance - suffices Measurable[mProd] ({ω | stepsUntil A a m ω = ↑n}.indicator fun x ↦ 1) by - rwa [measurable_indicator_const_iff] at this - exact measurable_comap_indicator_stepsUntil_eq hA hR' a m n - -/-- `stepsUntil a m` is a stopping time with respect to the filtration `filtrationAction`. -/ -theorem isStoppingTime_stepsUntil_filtrationAction [Nonempty R] [MeasurableSingletonClass α] - (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m : ℕ) : - IsStoppingTime (IsAlgEnvInteraction.filtrationAction hA hR') (stepsUntil A a m) := by - refine isStoppingTime_of_measurableSet_eq fun n ↦ ?_ - by_cases hn : n = 0 - · simp only [hn, IsAlgEnvInteraction.filtrationAction_zero_eq_comap, WithTop.coe_zero] - exact measurableSet_stepsUntil_eq_zero a m - · rw [IsAlgEnvInteraction.filtrationAction_eq_comap _ hn] - exact measurableSet_stepsUntil_eq hA hR' a m n - --- /-- Sigma-algebra generated by the stopping time `stepsUntil a m`. -/ --- def stepsUntilMeasurableSpace [Nonempty R] [MeasurableSingletonClass α] (a : α) (m : ℕ) : --- MeasurableSpace (ℕ → α × R) := --- (isStoppingTime_stepsUntil_filtrationAction a m (mR := mR)).measurableSpace - -end Measurability - -end StepsUntil - -section RewardByCount - -/-- Reward obtained when pulling action `a` for the `m`-th time. -If it is never pulled `m` times, the reward is given by the second component of `ω`, which in -applications will be indepedent with same law. -/ -noncomputable -def rewardByCount (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (a : α) (m : ℕ) (ω : Ω × (ℕ → α → R)) : R := - match (stepsUntil A a m ω.1) with - | ⊤ => ω.2 m a - | (n : ℕ) => R' n ω.1 - -variable {ω : Ω × (ℕ → α → R)} - -lemma rewardByCount_eq_ite (a : α) (m : ℕ) (ω : Ω × (ℕ → α → R)) : - rewardByCount A R' a m ω = - if (stepsUntil A a m ω.1) = ⊤ then ω.2 m a else R' (stepsUntil A a m ω.1).toNat ω.1 := by - unfold rewardByCount - cases stepsUntil A a m ω.1 <;> simp - -lemma rewardByCount_eq_add [AddMonoid R] (a : α) (m : ℕ) : - rewardByCount A R' a m = - {ω : Ω × (ℕ → α → R) | stepsUntil A a m ω.1 ≠ ⊤}.indicator - (fun ω ↦ R' (stepsUntil A a m ω.1).toNat ω.1) - + {ω | stepsUntil A a m ω.1 = ⊤}.indicator (fun ω ↦ ω.2 m a) := by - ext ω - simp only [rewardByCount_eq_ite, ne_eq, Pi.add_apply, Set.indicator_apply, Set.mem_setOf_eq, - ite_not] - grind - -lemma rewardByCount_of_stepsUntil_eq_top (h : stepsUntil A a m ω.1 = ⊤) : - rewardByCount A R' a m ω = ω.2 m a := by simp [rewardByCount_eq_ite, h] - -lemma rewardByCount_of_stepsUntil_ne_top (h : stepsUntil A a m ω.1 ≠ ⊤) : - rewardByCount A R' a m ω = R' (stepsUntil A a m ω.1).toNat ω.1 := by - simp [rewardByCount_eq_ite, h] - -lemma rewardByCount_eq_stoppedValue (h : stepsUntil A a m ω.1 ≠ ⊤) : - rewardByCount A R' a m ω = stoppedValue R' (stepsUntil A a m) ω.1 := by - rw [rewardByCount_of_stepsUntil_ne_top h, stoppedValue] - lift stepsUntil A a m ω.1 to ℕ using h with n - simp - -lemma rewardByCount_of_stepsUntil_eq_coe (h : stepsUntil A a m ω.1 = n) : - rewardByCount A R' a m ω = R' n ω.1 := by simp [rewardByCount_eq_ite, h] - -/-- The value at 0 does not matter (it would be the "zeroth" reward). -It should be considered a junk value. -/ -@[simp] -lemma rewardByCount_zero (a : α) (ω : Ω × (ℕ → α → R)) : - rewardByCount A R' a 0 ω = if A 0 ω.1 = a then ω.2 0 a else R' 0 ω.1 := by - rw [rewardByCount_eq_ite] - by_cases ha : A 0 ω.1 = a - · simp [ha, stepsUntil_zero_of_eq] - · simp [stepsUntil_zero_of_ne, ha] - -lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (ω : Ω × (ℕ → α → R)) : - rewardByCount A R' (A t ω.1) (pullCount A (A t ω.1) t ω.1 + 1) ω = R' t ω.1 := by - rw [rewardByCount, ← pullCount_action_eq_pullCount_add_one, stepsUntil_pullCount_eq] - -@[fun_prop] -lemma measurable_rewardByCount [MeasurableSingletonClass α] - (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m : ℕ) : - Measurable (fun ω : Ω × (ℕ → α → R) ↦ rewardByCount A R' a m ω) := by - simp_rw [rewardByCount_eq_ite] - refine Measurable.ite ?_ ?_ ?_ - · exact (measurableSet_singleton _).preimage <| measurable_stepsUntil' hA a m - · fun_prop - · change Measurable ((fun p : ℕ × Ω ↦ R' p.1 p.2) - ∘ (fun ω : Ω × (ℕ → α → R) ↦ ((stepsUntil A a m ω.1).toNat, ω.1))) - have : Measurable fun ω : Ω × (ℕ → α → R) ↦ ((stepsUntil A a m ω.1).toNat, ω.1) := - (measurable_stepsUntil' hA a m).toNat.prodMk (by fun_prop) - refine Measurable.comp ?_ this - refine measurable_from_prod_countable_right fun n ↦ ?_ - simp only - fun_prop - -end RewardByCount - -lemma sum_pullCount_mul [Fintype α] [Semiring R] (ω : Ω) (f : α → R) (t : ℕ) : - ∑ a, pullCount A a t ω * f a = ∑ s ∈ range t, f (A s ω) := by - unfold pullCount - classical - simp_rw [card_eq_sum_ones] - push_cast - simp_rw [sum_mul, one_mul] - exact sum_fiberwise' (range t) (A · ω) f - --- todo: only in ℝ for now -lemma sum_pullCount [Fintype α] {ω : Ω} : ∑ a, pullCount A a t ω = t := by - suffices ∑ a, pullCount A a t ω * (1 : ℝ) = t by norm_cast at this; simpa - rw [sum_pullCount_mul] - simp - -section SumRewards - -/-- Sum of rewards obtained when pulling action `a` up to time `t` (exclusive). -/ -def sumRewards (A : ℕ → Ω → α) (R' : ℕ → Ω → ℝ) (a : α) (t : ℕ) (ω : Ω) : ℝ := - ∑ s ∈ range t, if A s ω = a then R' s ω else 0 - -/-- Sum of rewards of arm `a` up to (and including) time `n`. -/ -noncomputable -def sumRewards' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := - ∑ s, if (h s).1 = a then (h s).2 else 0 - -/-- Empirical mean reward obtained when pulling action `a` up to time `t` (exclusive). -/ -noncomputable -def empMean (A : ℕ → Ω → α) (R' : ℕ → Ω → ℝ) (a : α) (t : ℕ) (ω : Ω) : ℝ := - sumRewards A R' a t ω / pullCount A a t ω - -/-- Empirical mean of arm `a` at time `n`. -/ -noncomputable -def empMean' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := - (sumRewards' n h a) / (pullCount' n h a) - -lemma sumRewards_eq_pullCount_mul_empMean {R' : ℕ → Ω → ℝ} {ω : Ω} - (h_pull : pullCount A a t ω ≠ 0) : - sumRewards A R' a t ω = pullCount A a t ω * empMean A R' a t ω := by unfold empMean; field_simp - -lemma sum_rewardByCount_eq_sumRewards {R' : ℕ → Ω → ℝ} (a : α) (t : ℕ) (ω : Ω × (ℕ → α → ℝ)) : - ∑ m ∈ Icc 1 (pullCount A a t ω.1), rewardByCount A R' a m ω = sumRewards A R' a t ω.1 := by - induction t with - | zero => simp [pullCount, sumRewards] - | succ t ht => - by_cases hta : A t ω.1 = a - · rw [← hta] at ht ⊢ - rw [pullCount_action_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht] - unfold sumRewards - rw [sum_range_succ, if_pos rfl, rewardByCount_pullCount_add_one_eq_reward] - · unfold sumRewards - rwa [pullCount_eq_pullCount_of_action_ne hta, sum_range_succ, if_neg hta, add_zero] - -lemma sumRewards_add_one_eq_sumRewards' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} : - sumRewards A R' a (n + 1) ω = sumRewards' n (fun i ↦ (A i ω, R' i ω)) a := by - unfold sumRewards sumRewards' - rw [Finset.sum_coe_sort (f := fun s ↦ if A s ω = a then R' s ω else 0) (Iic n)] - congr with m - simp only [mem_range, mem_Iic] - grind - -lemma sumRewards_eq_sumRewards' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (hn : n ≠ 0) : - sumRewards A R' a n ω = sumRewards' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by - cases n with - | zero => exact absurd rfl hn - | succ n => - rw [sumRewards_add_one_eq_sumRewards'] - have : n + 1 - 1 = n := by simp - exact this ▸ rfl - -lemma empMean_add_one_eq_empMean' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} : - empMean A R' a (n + 1) ω = empMean' n (fun i ↦ (A i ω, R' i ω)) a := by - unfold empMean empMean' - rw [sumRewards_add_one_eq_sumRewards', pullCount_add_one_eq_pullCount'] - -lemma empMean_eq_empMean' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (hn : n ≠ 0) : - empMean A R' a n ω = empMean' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by - unfold empMean empMean' - rw [sumRewards_eq_sumRewards' hn, pullCount_eq_pullCount' hn] - -@[fun_prop] -lemma measurable_sumRewards [MeasurableSingletonClass α] {R' : ℕ → Ω → ℝ} - (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (t : ℕ) : - Measurable (sumRewards A R' a t) := by - unfold sumRewards - have h_meas s : Measurable (fun h : Ω ↦ if A s h = a then R' s h else 0) := by - refine Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact (measurableSet_singleton _).preimage (by fun_prop) - fun_prop - -@[fun_prop] -lemma measurable_empMean [MeasurableSingletonClass α] {R' : ℕ → Ω → ℝ} (hA : ∀ n, Measurable (A n)) - (hR' : ∀ n, Measurable (R' n)) (a : α) (n : ℕ) : - Measurable (empMean A R' a n) := by - unfold empMean - fun_prop - -@[fun_prop] -lemma measurable_sumRewards' [MeasurableSingletonClass α] (n : ℕ) (a : α) : - Measurable (fun h ↦ sumRewards' n h a) := by - simp_rw [sumRewards'] - have h_meas s : Measurable (fun (h : Iic n → α × ℝ) ↦ if (h s).1 = a then (h s).2 else 0) := by - refine Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact (measurableSet_singleton _).preimage (by fun_prop) - fun_prop - -@[fun_prop] -lemma measurable_empMean' [MeasurableSingletonClass α] (n : ℕ) (a : α) : - Measurable (fun h ↦ empMean' n h a) := by - unfold empMean' - fun_prop - -end SumRewards - -end LearningDraft diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index a4e5ff01..18d57302 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -20,19 +20,23 @@ be seen as a stochastic process indexed by time `t` on the measurable space `ℕ -/ -open MeasureTheory Finset +open MeasureTheory Finset Learning namespace Learning -variable {α R : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} [DecidableEq α] - {a : α} {m n t : ℕ} {h : ℕ → α × R} +variable {α R Ω : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} {mΩ : MeasurableSpace Ω} + [DecidableEq α] + {alg : Algorithm α R} {env : Environment α R} + {P : Measure Ω} [IsProbabilityMeasure P] + {A : ℕ → Ω → α} {R' : ℕ → Ω → R} + {a : α} {m n t : ℕ} {ω : Ω} section PullCount /-- Number of times action `a` was chosen up to time `t` (excluding `t`). -/ noncomputable -def pullCount (a : α) (t : ℕ) (h : ℕ → α × R) : ℕ := - #(filter (fun s ↦ action s h = a) (range t)) +def pullCount (A : ℕ → Ω → α) (a : α) (t : ℕ) (ω : Ω) : ℕ := + #(filter (fun s ↦ A s ω = a) (range t)) /-- Number of pulls of arm `a` up to (and including) time `n`. This is the number of entries in `h` in which the arm is `a`. -/ @@ -40,68 +44,67 @@ noncomputable def pullCount' (n : ℕ) (h : Iic n → α × R) (a : α) := #{s | (h s).1 = a} @[simp] -lemma pullCount_zero (a : α) : pullCount a 0 (R := R) = 0 := by ext; simp [pullCount] +lemma pullCount_zero (a : α) : pullCount A a 0 = 0 := by ext; simp [pullCount] -lemma pullCount_zero_apply (a : α) (h : ℕ → α × R) : pullCount a 0 h = 0 := by simp +lemma pullCount_zero_apply (a : α) (ω : Ω) : pullCount A a 0 ω = 0 := by simp -lemma pullCount_one : pullCount a 1 h = if action 0 h = a then 1 else 0 := by +lemma pullCount_one : pullCount A a 1 ω = if A 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] -lemma monotone_pullCount (a : α) (h : ℕ → α × R) : Monotone (pullCount a · h) := +lemma monotone_pullCount (a : α) (ω : Ω) : Monotone (pullCount A a · ω) := fun _ _ _ ↦ card_le_card (filter_subset_filter _ (by simpa)) @[mono, gcongr] -lemma pullCount_mono (a : α) {n m : ℕ} (hnm : n ≤ m) (h : ℕ → α × R) : - pullCount a n h ≤ pullCount a m h := - monotone_pullCount a h hnm +lemma pullCount_mono (a : α) {n m : ℕ} (hnm : n ≤ m) (ω : Ω) : + pullCount A a n ω ≤ pullCount A a m ω := + monotone_pullCount a ω hnm -lemma pullCount_action_eq_pullCount_add_one (t : ℕ) (h : ℕ → α × R) : - pullCount (action t h) (t + 1) h = pullCount (action t h) t h + 1 := by +lemma pullCount_action_eq_pullCount_add_one (t : ℕ) (ω : Ω) : + pullCount A (A t ω) (t + 1) ω = pullCount A (A t ω) t ω + 1 := by simp [pullCount, range_add_one, filter_insert] -lemma pullCount_eq_pullCount_of_action_ne (ha : action t h ≠ a) : - pullCount a (t + 1) h = pullCount a t h := by +lemma pullCount_eq_pullCount_of_action_ne (ha : A t ω ≠ a) : + pullCount A a (t + 1) ω = pullCount A a t ω := by simp [pullCount, range_add_one, filter_insert, ha] lemma pullCount_add_one : - pullCount a (t + 1) h = pullCount a t h + if action t h = a then 1 else 0 := by + pullCount A a (t + 1) ω = pullCount A a t ω + if A t ω = a then 1 else 0 := by split_ifs with h · rw [← h, pullCount_action_eq_pullCount_add_one] · rw [pullCount_eq_pullCount_of_action_ne h, add_zero] -lemma pullCount_eq_sum (a : α) (t : ℕ) (h : ℕ → α × R) : - pullCount a t h = ∑ s ∈ range t, if action s h = a then 1 else 0 := by simp [pullCount] +lemma pullCount_eq_sum (a : α) (t : ℕ) (ω : Ω) : + pullCount A a t ω = ∑ s ∈ range t, if A s ω = a then 1 else 0 := by simp [pullCount] lemma pullCount'_eq_sum (n : ℕ) (h : Iic n → α × R) (a : α) : pullCount' n h a = ∑ s : Iic n, if (h s).1 = a then 1 else 0 := by simp [pullCount'] -lemma pullCount_add_one_eq_pullCount' {n : ℕ} {h : ℕ → α × R} : - pullCount a (n + 1) h = pullCount' n (fun i ↦ h i) a := by +lemma pullCount_add_one_eq_pullCount' {n : ℕ} {ω : Ω} : + pullCount A a (n + 1) ω = pullCount' n (fun i ↦ (A i ω, R' i ω)) a := by rw [pullCount_eq_sum, pullCount'_eq_sum] - unfold action - rw [Finset.sum_coe_sort (f := fun s ↦ if (h s).1 = a then 1 else 0) (Iic n)] + rw [Finset.sum_coe_sort (f := fun s ↦ if A s ω = a then 1 else 0) (Iic n)] congr with m simp only [mem_range, mem_Iic] grind -lemma pullCount_eq_pullCount' {n : ℕ} {h : ℕ → α × R} (hn : n ≠ 0) : - pullCount a n h = pullCount' (n - 1) (fun i ↦ h i) a := by +lemma pullCount_eq_pullCount' {n : ℕ} {ω : Ω} (hn : n ≠ 0) : + pullCount A a n ω = pullCount' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by cases n with | zero => exact absurd rfl hn | succ n => - rw [pullCount_add_one_eq_pullCount'] + rw [pullCount_add_one_eq_pullCount' (R' := R')] have : n + 1 - 1 = n := by simp exact this ▸ rfl -lemma pullCount_le (a : α) (t : ℕ) (h : ℕ → α × R) : pullCount a t h ≤ t := +lemma pullCount_le (a : α) (t : ℕ) (ω : Ω) : pullCount A a t ω ≤ t := (card_filter_le _ _).trans_eq (by simp) -lemma pullCount_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h = action i h') : - pullCount a (n + 1) h = pullCount a (n + 1) h' := by +lemma pullCount_congr {ω' : Ω} (h_eq : ∀ i ≤ n, A i ω = A i ω') : + pullCount A a (n + 1) ω = pullCount A a (n + 1) ω' := by unfold pullCount congr 1 with s simp only [mem_filter, mem_range, and_congr_right_iff] @@ -109,8 +112,8 @@ lemma pullCount_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h = a rw [Nat.lt_add_one_iff] at hs rw [h_eq s hs] -lemma pullCount_lt_of_forall_ne (h_lt : ∀ s, pullCount a (s + 1) h ≠ t) (ht : t ≠ 0) : - pullCount a n h < t := by +lemma pullCount_lt_of_forall_ne (h_lt : ∀ s, pullCount A a (s + 1) ω ≠ t) (ht : t ≠ 0) : + pullCount A a n ω < t := by induction n with | zero => simpa using ht.bot_lt | succ n hn => @@ -118,20 +121,21 @@ lemma pullCount_lt_of_forall_ne (h_lt : ∀ s, pullCount a (s + 1) h ≠ t) (ht rw [pullCount_add_one] at h_lt ⊢ grind -lemma exists_pullCount_eq_of_le (hnm : t ≤ pullCount a (n + 1) h) (ht : t ≠ 0) : - ∃ s, pullCount a (s + 1) h = t := by +lemma exists_pullCount_eq_of_le (hnm : t ≤ pullCount A a (n + 1) ω) (ht : t ≠ 0) : + ∃ s, pullCount A a (s + 1) ω = t := by by_contra! h_contra - refine lt_irrefl (pullCount a (n + 1) h) ?_ + refine lt_irrefl (pullCount A a (n + 1) ω) ?_ refine lt_of_lt_of_le ?_ hnm exact pullCount_lt_of_forall_ne h_contra ht section Measurability @[fun_prop] -lemma measurable_pullCount [MeasurableSingletonClass α] (a : α) (t : ℕ) : - Measurable (fun h : ℕ → α × R ↦ pullCount a t h) := by +lemma measurable_pullCount [MeasurableSingletonClass α] (hA : ∀ n, Measurable (A n)) + (a : α) (t : ℕ) : + Measurable (fun ω : Ω ↦ pullCount A a t ω) := by simp_rw [pullCount_eq_sum] - have h_meas s : Measurable (fun h : ℕ → α × R ↦ if action s h = a then 1 else 0) := by + have h_meas s : Measurable (fun ω : Ω ↦ if A s ω = a then 1 else 0) := by refine Measurable.ite ?_ (by fun_prop) (by fun_prop) exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop @@ -145,20 +149,26 @@ lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop -lemma adapted_pullCount_add_one [MeasurableSingletonClass α] (a : α) : - Adapted (Learning.filtration α R) (fun n ↦ pullCount a (n + 1)) := by - refine fun n ↦ Measurable.stronglyMeasurable ?_ - simp only - have : pullCount a (n + 1) = (fun h : Iic n → α × R ↦ pullCount' n h a) ∘ (hist n) := by +lemma adapted_pullCount_add_one' [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (n : ℕ) : + Measurable[IsAlgEnvSeq.filtration hA hR' n] (pullCount A a (n + 1)) := by + have : pullCount A a (n + 1) = (fun h : Iic n → α × R ↦ pullCount' n h a) ∘ + (IsAlgEnvSeq.hist A R' n) := by ext exact pullCount_add_one_eq_pullCount' - rw [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe, this] - exact measurable_comp_comap (hist n) (measurable_pullCount' n a) + rw [IsAlgEnvSeq.filtration, this] + exact measurable_comp_comap _ (measurable_pullCount' n a) -lemma isPredictable_pullCount [MeasurableSingletonClass α] (a : α) : - IsPredictable (Learning.filtration α R) (pullCount a) := by +lemma adapted_pullCount_add_one [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) : + Adapted (IsAlgEnvSeq.filtration hA hR') (fun n ↦ pullCount A a (n + 1)) := + fun n ↦ Measurable.stronglyMeasurable <| adapted_pullCount_add_one' hA hR' a n + +lemma isPredictable_pullCount [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) : + IsPredictable (IsAlgEnvSeq.filtration hA hR') (pullCount A a) := by rw [isPredictable_iff_measurable_add_one] - refine ⟨?_, fun n ↦ (adapted_pullCount_add_one a n).measurable⟩ + refine ⟨?_, fun n ↦ (adapted_pullCount_add_one hA hR' a n).measurable⟩ simp only [pullCount_zero] fun_prop @@ -171,21 +181,22 @@ section StepsUntil -- TODO: replace this by leastGE, once leastGE is generalized /-- Number of steps until action `a` was pulled exactly `m` times. -/ noncomputable -def stepsUntil (a : α) (m : ℕ) (h : ℕ → α × R) : ℕ∞ := sInf ((↑) '' {s | pullCount a (s + 1) h = m}) +def stepsUntil (A : ℕ → Ω → α) (a : α) (m : ℕ) (ω : Ω) : ℕ∞ := + sInf ((↑) '' {s | pullCount A a (s + 1) ω = m}) -lemma stepsUntil_eq_top_iff : stepsUntil a m h = ⊤ ↔ ∀ s, pullCount a (s + 1) h ≠ m := by +lemma stepsUntil_eq_top_iff : stepsUntil A a m ω = ⊤ ↔ ∀ s, pullCount A a (s + 1) ω ≠ m := by simp [stepsUntil, sInf_eq_top] -lemma stepsUntil_ne_top (h_exists : ∃ s, pullCount a (s + 1) h = m) : stepsUntil a m h ≠ ⊤ := by +lemma stepsUntil_ne_top (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : stepsUntil A a m ω ≠ ⊤ := by simpa [stepsUntil_eq_top_iff] -lemma exists_pullCount_eq (h' : stepsUntil a m h ≠ ⊤) : - ∃ s, pullCount a (s + 1) h = m := by +lemma exists_pullCount_eq (h' : stepsUntil A a m ω ≠ ⊤) : + ∃ s, pullCount A a (s + 1) ω = m := by by_contra! h_contra rw [← stepsUntil_eq_top_iff] at h_contra simp [h_contra] at h' -lemma stepsUntil_zero_of_ne (hka : action 0 h ≠ a) : stepsUntil a 0 h = 0 := by +lemma stepsUntil_zero_of_ne (hka : A 0 ω ≠ a) : stepsUntil A a 0 ω = 0 := by unfold stepsUntil simp_rw [← bot_eq_zero, sInf_eq_bot, bot_eq_zero] intro n hn @@ -194,19 +205,19 @@ lemma stepsUntil_zero_of_ne (hka : action 0 h ≠ a) : stepsUntil a 0 h = 0 := b rw [← zero_add 1, pullCount_eq_pullCount_of_action_ne hka] simp -lemma stepsUntil_zero_of_eq (hka : action 0 h = a) : stepsUntil a 0 h = ⊤ := by +lemma stepsUntil_zero_of_eq (hka : A 0 ω = a) : stepsUntil A a 0 ω = ⊤ := by rw [stepsUntil_eq_top_iff] - suffices 0 < pullCount a 1 h by + suffices 0 < pullCount A a 1 ω by intro n hn refine lt_irrefl 0 ?_ exact this.trans_le (le_trans (monotone_pullCount _ _ (by omega)) hn.le) rw [← hka, ← zero_add 1, pullCount_action_eq_pullCount_add_one] simp -lemma stepsUntil_eq_dite (a : α) (m : ℕ) (h : ℕ → α × R) - [Decidable (∃ s, pullCount a (s + 1) h = m)] : - stepsUntil a m h = - if h : ∃ s, pullCount a (s + 1) h = m then (Nat.find h : ℕ∞) else ⊤ := by +lemma stepsUntil_eq_dite (a : α) (m : ℕ) (ω : Ω) + [Decidable (∃ s, pullCount A a (s + 1) ω = m)] : + stepsUntil A a m ω = + if h : ∃ s, pullCount A a (s + 1) ω = m then (Nat.find h : ℕ∞) else ⊤ := by unfold stepsUntil split_ifs with h' · refine le_antisymm ?_ ?_ @@ -216,22 +227,22 @@ lemma stepsUntil_eq_dite (a : α) (m : ℕ) (h : ℕ → α × R) forall_apply_eq_imp_iff₂, Nat.cast_le, Nat.find_le_iff] exact fun n hn ↦ ⟨n, le_rfl, hn⟩ · push_neg at h' - suffices {s | pullCount a (s + 1) h = m} = ∅ by simp [this] + suffices {s | pullCount A a (s + 1) ω = m} = ∅ by simp [this] ext s simpa using (h' s) -- todo: this is in ℝ because of the limited def of leastGE lemma stepsUntil_eq_leastGE (a : α) (hm : m ≠ 0) : - stepsUntil a m = leastGE (fun n (h : ℕ → α × ℝ) ↦ pullCount a (n + 1) h) m := by + stepsUntil A a m = leastGE (fun n (ω : Ω) ↦ pullCount A a (n + 1) ω) m := by classical - ext h + ext ω rw [stepsUntil_eq_dite] unfold leastGE hittingAfter simp only [zero_le, Set.mem_Ici, Nat.cast_le, true_and, ENat.some_eq_coe] - have h_iff : (∃ s, pullCount a (s + 1) h = m) ↔ (∃ s, m ≤ pullCount a (s + 1) h) := by + have h_iff : (∃ s, pullCount A a (s + 1) ω = m) ↔ (∃ s, m ≤ pullCount A a (s + 1) ω) := by refine ⟨fun ⟨s, hs⟩ ↦ ⟨s, hs.ge⟩, fun ⟨s, hs⟩ ↦ ?_⟩ exact exists_pullCount_eq_of_le hs hm - by_cases h_exists : ∃ s, m ≤ pullCount a (s + 1) h + by_cases h_exists : ∃ s, m ≤ pullCount A a (s + 1) ω swap; · simp_rw [h_iff]; simp [h_exists] rw [if_pos h_exists, dif_pos] swap; · rwa [h_iff] @@ -240,45 +251,45 @@ lemma stepsUntil_eq_leastGE (a : α) (hm : m ≠ 0) : constructor · apply le_antisymm · by_contra! h_contra - obtain ⟨s, hs⟩ : ∃ s, pullCount a (s + 1) h = m := exists_pullCount_eq_of_le h_contra.le hm + obtain ⟨s, hs⟩ : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq_of_le h_contra.le hm rw [← hs] at h_contra refine h_contra.not_ge ?_ gcongr exact csInf_le (by simp) (by simp) - · exact Nat.sInf_mem (s := {j | m ≤ pullCount a (j + 1) h}) h_exists + · exact Nat.sInf_mem (s := {j | m ≤ pullCount A a (j + 1) ω}) h_exists · intro n hn h_contra refine hn.not_ge ?_ exact csInf_le (by simp) (by simp [h_contra]) -lemma stepsUntil_pullCount_le (h : ℕ → α × R) (a : α) (t : ℕ) : - stepsUntil a (pullCount a (t + 1) h) h ≤ t := by +lemma stepsUntil_pullCount_le (ω : Ω) (a : α) (t : ℕ) : + stepsUntil A a (pullCount A a (t + 1) ω) ω ≤ t := by rw [stepsUntil] exact csInf_le (OrderBot.bddBelow _) ⟨t, rfl, rfl⟩ -lemma stepsUntil_pullCount_eq (h : ℕ → α × R) (t : ℕ) : - stepsUntil (action t h) (pullCount (action t h) (t + 1) h) h = t := by - apply le_antisymm (stepsUntil_pullCount_le h (action t h) t) - suffices ∀ t', pullCount (action t h) (t' + 1) h = pullCount (action t h) t h + 1 → t ≤ t' by +lemma stepsUntil_pullCount_eq (ω : Ω) (t : ℕ) : + stepsUntil A (A t ω) (pullCount A (A t ω) (t + 1) ω) ω = t := by + apply le_antisymm (stepsUntil_pullCount_le ω (A t ω) t) + suffices ∀ t', pullCount A (A t ω) (t' + 1) ω = pullCount A (A t ω) t ω + 1 → t ≤ t' by simpa [stepsUntil, pullCount_action_eq_pullCount_add_one] - exact fun t' h' ↦ Nat.le_of_lt_succ ((monotone_pullCount (action t h) h).reflect_lt + exact fun t' h' ↦ Nat.le_of_lt_succ ((monotone_pullCount (A t ω) ω).reflect_lt (h' ▸ lt_add_one _)) /-- If we pull action `a` at time 0, the first time at which it is pulled once is 0. -/ -lemma stepsUntil_one_of_eq (hka : action 0 h = a) : stepsUntil a 1 h = 0 := by +lemma stepsUntil_one_of_eq (hka : A 0 ω = a) : stepsUntil A a 1 ω = 0 := by classical - have h_pull : pullCount a 1 h = 1 := by simp [pullCount_one, hka] - have h_le := stepsUntil_pullCount_le h a 0 + have h_pull : pullCount A a 1 ω = 1 := by simp [pullCount_one, hka] + have h_le := stepsUntil_pullCount_le (A := A) ω a 0 simpa [h_pull] using h_le lemma stepsUntil_eq_zero_iff : - stepsUntil a m h = 0 ↔ (m = 0 ∧ action 0 h ≠ a) ∨ (m = 1 ∧ action 0 h = a) := by + stepsUntil A a m ω = 0 ↔ (m = 0 ∧ A 0 ω ≠ a) ∨ (m = 1 ∧ A 0 ω = a) := by classical refine ⟨fun h' ↦ ?_, fun h' ↦ ?_⟩ - · have h_exists : ∃ s, pullCount a (s + 1) h = m := exists_pullCount_eq (by simp [h']) + · have h_exists : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq (by simp [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 : action 0 h = a + by_cases hka : A 0 ω = a · simp only [hka, ↓reduceIte] at h' simp [h'.symm, hka] · simp only [hka, ↓reduceIte] at h' @@ -290,8 +301,8 @@ lemma stepsUntil_eq_zero_iff : rw [h.1] exact stepsUntil_one_of_eq h.2 -lemma action_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) h = m) : - action (stepsUntil a m h).toNat h = a := by +lemma action_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : + A (stepsUntil A a m ω).toNat ω = a := by classical simp only [stepsUntil_eq_dite, h_exists, ↓reduceDIte, ENat.toNat_coe] have h_spec := Nat.find_spec h_exists @@ -310,17 +321,17 @@ lemma action_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) h rwa [← pullCount_eq_pullCount_of_action_ne] exact h_ne -lemma action_eq_of_stepsUntil_eq_coe {ω : ℕ → α × R} (hm : m ≠ 0) - (h : stepsUntil a m ω = n) : - action n ω = a := by - have : n = (stepsUntil a m ω).toNat := by simp [h] - rw [this, action_stepsUntil hm] - exact exists_pullCount_eq (by simp [h]) +lemma action_eq_of_stepsUntil_eq_coe (hm : m ≠ 0) (h : stepsUntil A a m ω = n) : + A n ω = a := by + have : n = (stepsUntil A a m ω).toNat := by simp [h] + rw [this] + have h_exists : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq (by simp [h]) + exact action_stepsUntil hm h_exists -lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount a (s + 1) h = m) : - pullCount a (stepsUntil a m h + 1).toNat h = m := by +lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : + pullCount A a (stepsUntil A a m ω + 1).toNat ω = m := by classical - have h_eq := stepsUntil_eq_dite a m h + have h_eq := stepsUntil_eq_dite (A := A) a m ω simp only [h_exists, ↓reduceDIte] at h_eq have h' := Nat.find_spec h_exists rw [h_eq] @@ -328,10 +339,10 @@ lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount a (s + 1) h = m) simp only [ENat.toNat_coe, ENat.toNat_one] exact h' -lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) h = m) : - pullCount a (stepsUntil a m h).toNat h = m - 1 := by - have h_action := action_eq_of_stepsUntil_eq_coe (n := (stepsUntil a m h).toNat) (a := a) (ω := h) - hm ?_ +lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : + pullCount A a (stepsUntil A a m ω).toNat ω = m - 1 := by + have h_action := action_eq_of_stepsUntil_eq_coe (A := A) (n := (stepsUntil A a m ω).toNat) + (a := a) (ω := ω) hm ?_ swap; · symm; simpa [stepsUntil_eq_top_iff] have h_add_one := pullCount_stepsUntil_add_one h_exists nth_rw 1 [← h_action] at h_add_one @@ -340,46 +351,46 @@ lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) swap; · simpa [stepsUntil_eq_top_iff] grind -lemma pullCount_lt_of_le_stepsUntil (a : α) {n m : ℕ} (h : ℕ → α × R) - (h_exists : ∃ s, pullCount a (s + 1) h = m) (hn : n < stepsUntil a m h) : - pullCount a (n + 1) h < m := by +lemma pullCount_lt_of_le_stepsUntil (a : α) {n m : ℕ} (ω : Ω) + (h_exists : ∃ s, pullCount A a (s + 1) ω = m) (hn : n < stepsUntil A a m ω) : + pullCount A a (n + 1) ω < m := by classical - have h_eq := stepsUntil_eq_dite a m h + have h_eq := stepsUntil_eq_dite (A := A) a m ω simp only [h_exists, ↓reduceDIte] at h_eq rw [← ENat.coe_toNat (stepsUntil_ne_top h_exists)] at hn refine lt_of_le_of_ne ?_ ?_ - · calc pullCount a (n + 1) h - _ ≤ pullCount a (stepsUntil a m h + 1).toNat h := by - refine monotone_pullCount a h ?_ + · calc pullCount A a (n + 1) ω + _ ≤ pullCount A a (stepsUntil A a m ω + 1).toNat ω := by + refine monotone_pullCount a ω ?_ rw [ENat.toNat_add (stepsUntil_ne_top h_exists) (by simp)] simp only [ENat.toNat_one, add_le_add_iff_right] exact mod_cast hn.le _ = m := pullCount_stepsUntil_add_one h_exists · refine Nat.find_min h_exists (m := n) ?_ - suffices n < (stepsUntil a m h).toNat by + suffices n < (stepsUntil A a m ω).toNat by rwa [h_eq, ENat.toNat_coe] at this exact mod_cast hn -lemma pullCount_eq_of_stepsUntil_eq_coe {ω : ℕ → α × R} (hm : m ≠ 0) - (h : stepsUntil a m ω = n) : - pullCount a n ω = m - 1 := by - have : n = (stepsUntil a m ω).toNat := by simp [h] +lemma pullCount_eq_of_stepsUntil_eq_coe {ω : Ω} (hm : m ≠ 0) + (h : stepsUntil A a m ω = n) : + pullCount A a n ω = m - 1 := by + have : n = (stepsUntil A a m ω).toNat := by simp [h] rw [this, pullCount_stepsUntil hm] exact exists_pullCount_eq (by simp [h]) -lemma pullCount_add_one_eq_of_stepsUntil_eq_coe {ω : ℕ → α × R} - (h : stepsUntil a m ω = n) : - pullCount a (n + 1) ω = m := by - have : n + 1 = (stepsUntil a m ω + 1).toNat := by +lemma pullCount_add_one_eq_of_stepsUntil_eq_coe {ω : Ω} + (h : stepsUntil A a m ω = n) : + pullCount A a (n + 1) ω = m := by + have : n + 1 = (stepsUntil A a m ω + 1).toNat := by rw [ENat.toNat_add (by simp [h]) (by simp)]; simp [h] rw [this, pullCount_stepsUntil_add_one] exact exists_pullCount_eq (by simp [h]) -lemma stepsUntil_eq_iff {ω : ℕ → α × R} (n : ℕ) : - stepsUntil a m ω = n ↔ - pullCount a (n + 1) ω = m ∧ (∀ k < n, pullCount a (k + 1) ω < m) := by +lemma stepsUntil_eq_iff {ω : Ω} (n : ℕ) : + stepsUntil A a m ω = n ↔ + pullCount A a (n + 1) ω = m ∧ (∀ k < n, pullCount A a (k + 1) ω < m) := by refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ - · have h_exists : ∃ s, pullCount a (s + 1) ω = m := exists_pullCount_eq (by simp [h]) + · have h_exists : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq (by simp [h]) refine ⟨pullCount_add_one_eq_of_stepsUntil_eq_coe h, fun k hk ↦ ?_⟩ exact pullCount_lt_of_le_stepsUntil a ω h_exists (by rw [h]; exact mod_cast hk) · classical @@ -388,8 +399,28 @@ lemma stepsUntil_eq_iff {ω : ℕ → α × R} (n : ℕ) : rw [Nat.find_eq_iff] exact ⟨h.1, fun k hk ↦ (h.2 k hk).ne⟩ -lemma stepsUntil_eq_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h = action i h') : - stepsUntil a m h = n ↔ stepsUntil a m h' = n := by +lemma stepsUntil_eq_iff' {ω : Ω} (hm : m ≠ 0) (n : ℕ) : + stepsUntil A a m ω = n ↔ A n ω = a ∧ pullCount A a n ω = m - 1 := by + by_cases hn : n = 0 + · simp [hn, stepsUntil_eq_zero_iff, hm] + grind + rw [stepsUntil_eq_iff n] + refine ⟨fun ⟨h1, h2⟩ ↦ ⟨?_, ?_⟩, fun ⟨h1, h2⟩ ↦ ⟨?_, fun k hk ↦ ?_⟩⟩ + · rw [pullCount_add_one] at h1 + specialize h2 (n - 1) (by lia) + grind + · rw [pullCount_add_one] at h1 + specialize h2 (n - 1) (by lia) + grind + · rw [pullCount_add_one, h1, h2] + grind + · rw [Nat.lt_iff_le_pred (by grind)] + rw [← h2] + refine monotone_pullCount a ω ?_ + grind + +lemma stepsUntil_eq_congr {ω' : Ω} (h_eq : ∀ i ≤ n, A i ω = A i ω') : + stepsUntil A a m ω = n ↔ stepsUntil A a m ω' = n := by simp_rw [stepsUntil_eq_iff n] congr! 1 · rw [pullCount_congr h_eq] @@ -399,34 +430,41 @@ lemma stepsUntil_eq_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h section Measurability -lemma isStoppingTime_stepsUntil [MeasurableSingletonClass α] (a : α) (hm : m ≠ 0) : - IsStoppingTime (Learning.filtration α ℝ) (stepsUntil a m) := by +lemma isStoppingTime_stepsUntil [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (hm : m ≠ 0) : + IsStoppingTime (IsAlgEnvSeq.filtration hA hR') (stepsUntil A a m) := by rw [stepsUntil_eq_leastGE _ hm] refine Adapted.isStoppingTime_leastGE _ fun n ↦ ?_ - suffices StronglyMeasurable[Learning.filtration α ℝ n] (pullCount a (n + 1)) by fun_prop - exact adapted_pullCount_add_one a n + suffices StronglyMeasurable[IsAlgEnvSeq.filtration hA hR' n] (pullCount A a (n + 1)) by + fun_prop + exact adapted_pullCount_add_one hA hR' a n -- todo: get this from the stopping time property? @[fun_prop] -lemma measurable_stepsUntil [MeasurableSingletonClass α] (a : α) (m : ℕ) : - Measurable (fun h : ℕ → α × R ↦ stepsUntil a m h) := by +lemma measurable_stepsUntil [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (a : α) (m : ℕ) : + Measurable (stepsUntil A a m) := by classical - have h_union : {h' : ℕ → α × R | ∃ s, pullCount a (s + 1) h' = m} - = ⋃ s : ℕ, {h' | pullCount a (s + 1) h' = m} := by ext; simp - have h_meas_set : MeasurableSet {h' : ℕ → α × R | ∃ s, pullCount a (s + 1) h' = m} := by + have h_union : {h' : Ω | ∃ s, pullCount A a (s + 1) h' = m} + = ⋃ s : ℕ, {h' | pullCount A a (s + 1) h' = m} := by ext; simp + have h_meas_set : MeasurableSet {h' : Ω | ∃ s, pullCount A a (s + 1) h' = m} := by rw [h_union] - exact MeasurableSet.iUnion fun s ↦ (measurableSet_singleton _).preimage (by fun_prop) - simp_rw [stepsUntil_eq_dite] - suffices Measurable fun k ↦ if h : k ∈ {k' | ∃ s, pullCount a (s + 1) k' = m} - then (Nat.find h : ℕ∞) else ⊤ by convert this - refine Measurable.dite (s := {k' : ℕ → α × R | ∃ s, pullCount a (s + 1) k' = m}) + refine MeasurableSet.iUnion fun s ↦ (measurableSet_singleton _).preimage ?_ + exact measurable_pullCount hA a (s + 1) + --simp_rw [stepsUntil_eq_dite] + suffices Measurable fun k ↦ if h : k ∈ {k' | ∃ s, pullCount A a (s + 1) k' = m} + then (Nat.find h : ℕ∞) else ⊤ by + convert this with ω + rw [stepsUntil_eq_dite a m ω] + rfl + refine Measurable.dite (s := {k' : Ω | ∃ s, pullCount A a (s + 1) k' = m}) (f := fun x ↦ (Nat.find x.2 : ℕ∞)) (g := fun _ ↦ ⊤) ?_ (by fun_prop) h_meas_set refine Measurable.coe_nat_enat ?_ refine measurable_find _ fun k ↦ ?_ - suffices MeasurableSet {x : ℕ → α × R | pullCount a (k + 1) x = m} by - have : Subtype.val '' {x : {k' : ℕ → α × R | - ∃ s, pullCount a (s + 1) k' = m} | pullCount a (k + 1) (x : ℕ → α × R) = m} - = {x : ℕ → α × R | pullCount a (k + 1) x = m} := by + suffices MeasurableSet {x : Ω | pullCount A a (k + 1) x = m} by + have : Subtype.val '' {x : {k' : Ω | + ∃ s, pullCount A a (s + 1) k' = m} | pullCount A a (k + 1) (x : Ω) = m} + = {x : Ω | pullCount A a (k + 1) x = m} := by ext x simp only [Set.mem_setOf_eq, Set.coe_setOf, Set.mem_image, Subtype.exists, exists_and_left, exists_prop, exists_eq_right_right, and_iff_left_iff_imp] @@ -436,46 +474,61 @@ lemma measurable_stepsUntil [MeasurableSingletonClass α] (a : α) (m : ℕ) : exact (measurableSet_singleton _).preimage (by fun_prop) exact (measurableSet_singleton _).preimage (by fun_prop) -lemma measurable_stepsUntil' [MeasurableSingletonClass α] (a : α) (m : ℕ) : - Measurable (fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ stepsUntil a m ω.1) := - (measurable_stepsUntil a m).comp measurable_fst - -lemma measurable_comap_indicator_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] - (a : α) (m n : ℕ) : - Measurable[MeasurableSpace.comap (fun ω : ℕ → α × R ↦ (hist (n-1) ω, action n ω)) inferInstance] - ({ω | stepsUntil a m ω = ↑n}.indicator fun _ ↦ 1) := by - let r₀ : R := Nonempty.some inferInstance - let k : ((Iic (n - 1) → α × R) × α) → (ℕ → α × R) := fun x i ↦ - if hi : i ∈ Iic (n - 1) then (x.1 ⟨i, hi⟩) else if i = n then (x.2, r₀) else (a, r₀) - have hk : Measurable k := by - unfold k - rw [measurable_pi_iff] - intro i - split_ifs <;> fun_prop - let φ : ((Iic (n - 1) → α × R) × α) → ℕ := fun x ↦ if stepsUntil a m (k x) = ↑n then 1 else 0 - have hφ : Measurable φ := - Measurable.ite ((measurableSet_singleton _).preimage (by fun_prop)) (by fun_prop) (by fun_prop) - suffices {ω | stepsUntil a m ω = ↑n}.indicator (fun x ↦ 1) - = φ ∘ fun ω ↦ (hist (n - 1) ω, action n ω) from this ▸ measurable_comp_comap _ hφ - 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 [action, mem_Iic, hist, dite_eq_ite, k, action] - grind +lemma measurable_stepsUntil' [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (a : α) (m : ℕ) : + Measurable (fun ω : Ω × (ℕ → α → R) ↦ stepsUntil A a m ω.1) := + (measurable_stepsUntil hA a m).comp measurable_fst + +lemma measurable_comap_indicator_stepsUntil_eq [MeasurableSingletonClass α] [Nonempty R] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : + Measurable[MeasurableSpace.comap + (fun ω : Ω ↦ (IsAlgEnvSeq.hist A R' (n-1) ω, A n ω)) inferInstance] + ({ω | stepsUntil A a m ω = ↑n}.indicator fun _ ↦ 1) := by + by_cases hm : m = 0 + · simp only [hm] + by_cases hn : n = 0 + · simp only [hn, CharP.cast_eq_zero, stepsUntil_eq_zero_iff, ne_eq, true_and, zero_ne_one, + false_and, or_false] + refine Measurable.indicator measurable_const ?_ + refine (measurableSet_singleton _).compl.preimage ?_ + rw [measurable_iff_comap_le] + rw [Prod.instMeasurableSpace, MeasurableSpace.comap_prodMk] + exact le_sup_of_le_right le_rfl + · have : {ω | stepsUntil A a 0 ω = n} = ∅ := by + ext ω + by_cases ha : A 0 ω = a + · simp [stepsUntil_zero_of_eq ha] + · simp only [Set.mem_setOf_eq, stepsUntil_zero_of_ne ha, Set.mem_empty_iff_false, + iff_false] + norm_cast + exact Ne.symm hn + simp [this] + simp_rw [stepsUntil_eq_iff' hm] + refine Measurable.indicator measurable_const ?_ + refine ((measurableSet_singleton _).preimage ?_).inter ((measurableSet_singleton _).preimage ?_) + · rw [measurable_iff_comap_le] + rw [Prod.instMeasurableSpace, MeasurableSpace.comap_prodMk] + exact le_sup_of_le_right le_rfl + · rw [measurable_iff_comap_le] + rw [Prod.instMeasurableSpace, MeasurableSpace.comap_prodMk] + refine le_sup_of_le_left ?_ + rw [← measurable_iff_comap_le] + by_cases hn : n = 0 + · simp only [hn, pullCount_zero] + exact measurable_const + have h_meas := adapted_pullCount_add_one' hA hR' a (n - 1) + rwa [Nat.sub_add_cancel (by lia)] at h_meas lemma measurable_indicator_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] - (a : α) (m n : ℕ) : - Measurable ({ω : ℕ → α × R | stepsUntil a m ω = ↑n}.indicator fun _ ↦ 1) := by - refine (measurable_comap_indicator_stepsUntil_eq (mR := mR) a m n).mono ?_ le_rfl + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : + Measurable ({ω : Ω | stepsUntil A a m ω = ↑n}.indicator fun _ ↦ 1) := by + refine (measurable_comap_indicator_stepsUntil_eq hA hR' a m n).mono ?_ le_rfl refine Measurable.comap_le ?_ fun_prop -lemma measurableSet_stepsUntil_eq_zero [Nonempty R] [MeasurableSingletonClass α] (a : α) (m : ℕ) : - MeasurableSet[MeasurableSpace.comap (action 0) inferInstance] - {ω : ℕ → α × R | stepsUntil a m ω = 0} := by +lemma measurableSet_stepsUntil_eq_zero [MeasurableSingletonClass α] (a : α) (m : ℕ) : + MeasurableSet[MeasurableSpace.comap (A 0) inferInstance] + {ω : Ω | stepsUntil A a m ω = 0} := by simp only [stepsUntil_eq_zero_iff (a := a) (m := m), ne_eq] by_cases hm : m = 0 · simp only [hm, true_and, zero_ne_one, false_and, or_false] @@ -487,32 +540,33 @@ lemma measurableSet_stepsUntil_eq_zero [Nonempty R] [MeasurableSingletonClass α refine (measurableSet_singleton _).preimage ?_ rw [measurable_iff_comap_le] -lemma measurable_comap_indicator_stepsUntil_eq_zero [Nonempty R] [MeasurableSingletonClass α] - (a : α) (m : ℕ) : - Measurable[MeasurableSpace.comap (action 0 (R := R)) inferInstance] - ({ω | stepsUntil a m ω = 0}.indicator fun _ ↦ 1) := by +lemma measurable_comap_indicator_stepsUntil_eq_zero [MeasurableSingletonClass α] (a : α) (m : ℕ) : + Measurable[MeasurableSpace.comap (A 0) inferInstance] + ({ω | stepsUntil A a m ω = 0}.indicator fun _ ↦ 1) := by rw [measurable_indicator_const_iff] exact measurableSet_stepsUntil_eq_zero a m -lemma measurableSet_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] (a : α) (m n : ℕ) : - MeasurableSet[MeasurableSpace.comap (fun ω : ℕ → α × R ↦ (hist (n-1) ω, action n ω)) +lemma measurableSet_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : + MeasurableSet[MeasurableSpace.comap (fun ω : Ω ↦ (IsAlgEnvSeq.hist A R' (n-1) ω, A n ω)) inferInstance] - {ω : ℕ → α × R | stepsUntil a m ω = ↑n} := by - let mProd := MeasurableSpace.comap (fun ω : ℕ → α × R ↦ (hist (n-1) ω, action n ω)) inferInstance - suffices Measurable[mProd] ({ω | stepsUntil a m ω = ↑n}.indicator fun x ↦ 1) by + {ω : Ω | stepsUntil A a m ω = ↑n} := by + let mProd := MeasurableSpace.comap + (fun ω : Ω ↦ (IsAlgEnvSeq.hist A R' (n-1) ω, A n ω)) inferInstance + suffices Measurable[mProd] ({ω | stepsUntil A a m ω = ↑n}.indicator fun x ↦ 1) by rwa [measurable_indicator_const_iff] at this - exact measurable_comap_indicator_stepsUntil_eq a m n + exact measurable_comap_indicator_stepsUntil_eq hA hR' a m n /-- `stepsUntil a m` is a stopping time with respect to the filtration `filtrationAction`. -/ theorem isStoppingTime_stepsUntil_filtrationAction [Nonempty R] [MeasurableSingletonClass α] - (a : α) (m : ℕ) : - IsStoppingTime (filtrationAction α R) (stepsUntil a m) := by + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m : ℕ) : + IsStoppingTime (IsAlgEnvSeq.filtrationAction hA hR') (stepsUntil A a m) := by refine isStoppingTime_of_measurableSet_eq fun n ↦ ?_ by_cases hn : n = 0 - · simp only [hn, filtrationAction_zero_eq_comap, WithTop.coe_zero] + · simp only [hn, IsAlgEnvSeq.filtrationAction_zero_eq_comap, WithTop.coe_zero] exact measurableSet_stepsUntil_eq_zero a m - · rw [filtrationAction_eq_comap _ hn] - exact measurableSet_stepsUntil_eq a m n + · rw [IsAlgEnvSeq.filtrationAction_eq_comap _ hn] + exact measurableSet_stepsUntil_eq hA hR' a m n -- /-- Sigma-algebra generated by the stopping time `stepsUntil a m`. -/ -- def stepsUntilMeasurableSpace [Nonempty R] [MeasurableSingletonClass α] (a : α) (m : ℕ) : @@ -529,93 +583,98 @@ section RewardByCount If it is never pulled `m` times, the reward is given by the second component of `ω`, which in applications will be indepedent with same law. -/ noncomputable -def rewardByCount (a : α) (m : ℕ) (ω : (ℕ → α × R) × (ℕ → α → R)) : R := - match (stepsUntil a m ω.1) with +def rewardByCount (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (a : α) (m : ℕ) (ω : Ω × (ℕ → α → R)) : R := + match (stepsUntil A a m ω.1) with | ⊤ => ω.2 m a - | (n : ℕ) => reward n ω.1 + | (n : ℕ) => R' n ω.1 -variable {ω : (ℕ → α × R) × (ℕ → α → R)} +variable {ω : Ω × (ℕ → α → R)} -lemma rewardByCount_eq_ite (a : α) (m : ℕ) (ω : (ℕ → α × R) × (ℕ → α → R)) : - rewardByCount a m ω = - if (stepsUntil a m ω.1) = ⊤ then ω.2 m a else reward (stepsUntil a m ω.1).toNat ω.1 := by +lemma rewardByCount_eq_ite (a : α) (m : ℕ) (ω : Ω × (ℕ → α → R)) : + rewardByCount A R' a m ω = + if (stepsUntil A a m ω.1) = ⊤ then ω.2 m a else R' (stepsUntil A a m ω.1).toNat ω.1 := by unfold rewardByCount - cases stepsUntil a m ω.1 <;> simp + cases stepsUntil A a m ω.1 <;> simp lemma rewardByCount_eq_add [AddMonoid R] (a : α) (m : ℕ) : - rewardByCount a m = - {ω : (ℕ → α × R) × (ℕ → α → R) | stepsUntil a m ω.1 ≠ ⊤}.indicator - (fun ω ↦ reward (stepsUntil a m ω.1).toNat ω.1) - + {ω | stepsUntil a m ω.1 = ⊤}.indicator (fun ω ↦ ω.2 m a) := by + rewardByCount A R' a m = + {ω : Ω × (ℕ → α → R) | stepsUntil A a m ω.1 ≠ ⊤}.indicator + (fun ω ↦ R' (stepsUntil A a m ω.1).toNat ω.1) + + {ω | stepsUntil A a m ω.1 = ⊤}.indicator (fun ω ↦ ω.2 m a) := by ext ω simp only [rewardByCount_eq_ite, ne_eq, Pi.add_apply, Set.indicator_apply, Set.mem_setOf_eq, ite_not] grind -lemma rewardByCount_of_stepsUntil_eq_top (h : stepsUntil a m ω.1 = ⊤) : - rewardByCount a m ω = ω.2 m a := by simp [rewardByCount_eq_ite, h] +lemma rewardByCount_of_stepsUntil_eq_top (h : stepsUntil A a m ω.1 = ⊤) : + rewardByCount A R' a m ω = ω.2 m a := by simp [rewardByCount_eq_ite, h] -lemma rewardByCount_of_stepsUntil_ne_top (h : stepsUntil a m ω.1 ≠ ⊤) : - rewardByCount a m ω = reward (stepsUntil a m ω.1).toNat ω.1 := by simp [rewardByCount_eq_ite, h] +lemma rewardByCount_of_stepsUntil_ne_top (h : stepsUntil A a m ω.1 ≠ ⊤) : + rewardByCount A R' a m ω = R' (stepsUntil A a m ω.1).toNat ω.1 := by + simp [rewardByCount_eq_ite, h] -lemma rewardByCount_eq_stoppedValue (h : stepsUntil a m ω.1 ≠ ⊤) : - rewardByCount a m ω = stoppedValue reward (stepsUntil a m) ω.1 := by +lemma rewardByCount_eq_stoppedValue (h : stepsUntil A a m ω.1 ≠ ⊤) : + rewardByCount A R' a m ω = stoppedValue R' (stepsUntil A a m) ω.1 := by rw [rewardByCount_of_stepsUntil_ne_top h, stoppedValue] - lift stepsUntil a m ω.1 to ℕ using h with n + lift stepsUntil A a m ω.1 to ℕ using h with n simp -lemma rewardByCount_of_stepsUntil_eq_coe (h : stepsUntil a m ω.1 = n) : - rewardByCount a m ω = reward n ω.1 := by simp [rewardByCount_eq_ite, h] +lemma rewardByCount_of_stepsUntil_eq_coe (h : stepsUntil A a m ω.1 = n) : + rewardByCount A R' a m ω = R' n ω.1 := by simp [rewardByCount_eq_ite, h] /-- The value at 0 does not matter (it would be the "zeroth" reward). It should be considered a junk value. -/ @[simp] -lemma rewardByCount_zero (a : α) (ω : (ℕ → α × R) × (ℕ → α → R)) : - rewardByCount a 0 ω = if action 0 ω.1 = a then ω.2 0 a else reward 0 ω.1 := by +lemma rewardByCount_zero (a : α) (ω : Ω × (ℕ → α → R)) : + rewardByCount A R' a 0 ω = if A 0 ω.1 = a then ω.2 0 a else R' 0 ω.1 := by rw [rewardByCount_eq_ite] - by_cases ha : action 0 ω.1 = a + by_cases ha : A 0 ω.1 = a · simp [ha, stepsUntil_zero_of_eq] · simp [stepsUntil_zero_of_ne, ha] -lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (ω : (ℕ → α × R) × (ℕ → α → R)) : - rewardByCount (action t ω.1) (pullCount (action t ω.1) t ω.1 + 1) ω = reward t ω.1 := by +lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (ω : Ω × (ℕ → α → R)) : + rewardByCount A R' (A t ω.1) (pullCount A (A t ω.1) t ω.1 + 1) ω = R' t ω.1 := by rw [rewardByCount, ← pullCount_action_eq_pullCount_add_one, stepsUntil_pullCount_eq] @[fun_prop] -lemma measurable_rewardByCount [MeasurableSingletonClass α] (a : α) (m : ℕ) : - Measurable (fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ rewardByCount a m ω) := by +lemma measurable_rewardByCount [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m : ℕ) : + Measurable (fun ω : Ω × (ℕ → α → R) ↦ rewardByCount A R' a m ω) := by simp_rw [rewardByCount_eq_ite] refine Measurable.ite ?_ ?_ ?_ - · exact (measurableSet_singleton _).preimage <| measurable_stepsUntil' a m + · exact (measurableSet_singleton _).preimage <| measurable_stepsUntil' hA a m · fun_prop - · change Measurable ((fun p : ℕ × (ℕ → α × R) ↦ reward p.1 p.2) - ∘ (fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ ((stepsUntil a m ω.1).toNat, ω.1))) - have : Measurable fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ ((stepsUntil a m ω.1).toNat, ω.1) := - (measurable_stepsUntil' a m).toNat.prodMk (by fun_prop) - exact Measurable.comp (by fun_prop) this + · change Measurable ((fun p : ℕ × Ω ↦ R' p.1 p.2) + ∘ (fun ω : Ω × (ℕ → α → R) ↦ ((stepsUntil A a m ω.1).toNat, ω.1))) + have : Measurable fun ω : Ω × (ℕ → α → R) ↦ ((stepsUntil A a m ω.1).toNat, ω.1) := + (measurable_stepsUntil' hA a m).toNat.prodMk (by fun_prop) + refine Measurable.comp ?_ this + refine measurable_from_prod_countable_right fun n ↦ ?_ + simp only + fun_prop end RewardByCount -lemma sum_pullCount_mul [Fintype α] [Semiring R] (h : ℕ → α × R) (f : α → R) (t : ℕ) : - ∑ a, pullCount a t h * f a = ∑ s ∈ range t, f (action s h) := by +lemma sum_pullCount_mul [Fintype α] [Semiring R] (ω : Ω) (f : α → R) (t : ℕ) : + ∑ a, pullCount A a t ω * f a = ∑ s ∈ range t, f (A s ω) := by unfold pullCount classical simp_rw [card_eq_sum_ones] push_cast simp_rw [sum_mul, one_mul] - exact sum_fiberwise' (range t) (action · h) f + exact sum_fiberwise' (range t) (A · ω) f -- todo: only in ℝ for now -lemma sum_pullCount [Fintype α] {h : ℕ → α × ℝ} : ∑ a, pullCount a t h = t := by - suffices ∑ a, pullCount a t h * (1 : ℝ) = t by norm_cast at this; simpa +lemma sum_pullCount [Fintype α] {ω : Ω} : ∑ a, pullCount A a t ω = t := by + suffices ∑ a, pullCount A a t ω * (1 : ℝ) = t by norm_cast at this; simpa rw [sum_pullCount_mul] simp section SumRewards /-- Sum of rewards obtained when pulling action `a` up to time `t` (exclusive). -/ -def sumRewards (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := - ∑ s ∈ range t, if action s h = a then reward s h else 0 +def sumRewards (A : ℕ → Ω → α) (R' : ℕ → Ω → ℝ) (a : α) (t : ℕ) (ω : Ω) : ℝ := + ∑ s ∈ range t, if A s ω = a then R' s ω else 0 /-- Sum of rewards of arm `a` up to (and including) time `n`. -/ noncomputable @@ -624,22 +683,24 @@ def sumRewards' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := /-- Empirical mean reward obtained when pulling action `a` up to time `t` (exclusive). -/ noncomputable -def empMean (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := sumRewards a t h / pullCount a t h +def empMean (A : ℕ → Ω → α) (R' : ℕ → Ω → ℝ) (a : α) (t : ℕ) (ω : Ω) : ℝ := + sumRewards A R' a t ω / pullCount A a t ω /-- Empirical mean of arm `a` at time `n`. -/ noncomputable def empMean' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := (sumRewards' n h a) / (pullCount' n h a) -lemma sumRewards_eq_pullCount_mul_empMean {h : ℕ → α × ℝ} (h_pull : pullCount a t h ≠ 0) : - sumRewards a t h = pullCount a t h * empMean a t h := by unfold empMean; field_simp +lemma sumRewards_eq_pullCount_mul_empMean {R' : ℕ → Ω → ℝ} {ω : Ω} + (h_pull : pullCount A a t ω ≠ 0) : + sumRewards A R' a t ω = pullCount A a t ω * empMean A R' a t ω := by unfold empMean; field_simp -lemma sum_rewardByCount_eq_sumRewards (a : α) (t : ℕ) (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) : - ∑ m ∈ Icc 1 (pullCount a t ω.1), rewardByCount a m ω = sumRewards a t ω.1 := by +lemma sum_rewardByCount_eq_sumRewards {R' : ℕ → Ω → ℝ} (a : α) (t : ℕ) (ω : Ω × (ℕ → α → ℝ)) : + ∑ m ∈ Icc 1 (pullCount A a t ω.1), rewardByCount A R' a m ω = sumRewards A R' a t ω.1 := by induction t with | zero => simp [pullCount, sumRewards] | succ t ht => - by_cases hta : action t ω.1 = a + by_cases hta : A t ω.1 = a · rw [← hta] at ht ⊢ rw [pullCount_action_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht] unfold sumRewards @@ -647,16 +708,16 @@ lemma sum_rewardByCount_eq_sumRewards (a : α) (t : ℕ) (ω : (ℕ → α × · unfold sumRewards rwa [pullCount_eq_pullCount_of_action_ne hta, sum_range_succ, if_neg hta, add_zero] -lemma sumRewards_add_one_eq_sumRewards' {n : ℕ} {h : ℕ → α × ℝ} : - sumRewards a (n + 1) h = sumRewards' n (fun i ↦ h i) a := by - unfold sumRewards sumRewards' action Learning.reward - rw [Finset.sum_coe_sort (f := fun s ↦ if (h s).1 = a then (h s).2 else 0) (Iic n)] +lemma sumRewards_add_one_eq_sumRewards' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} : + sumRewards A R' a (n + 1) ω = sumRewards' n (fun i ↦ (A i ω, R' i ω)) a := by + unfold sumRewards sumRewards' + rw [Finset.sum_coe_sort (f := fun s ↦ if A s ω = a then R' s ω else 0) (Iic n)] congr with m simp only [mem_range, mem_Iic] grind -lemma sumRewards_eq_sumRewards' {n : ℕ} {h : ℕ → α × ℝ} (hn : n ≠ 0) : - sumRewards a n h = sumRewards' (n - 1) (fun i ↦ h i) a := by +lemma sumRewards_eq_sumRewards' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (hn : n ≠ 0) : + sumRewards A R' a n ω = sumRewards' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by cases n with | zero => exact absurd rfl hn | succ n => @@ -664,28 +725,30 @@ lemma sumRewards_eq_sumRewards' {n : ℕ} {h : ℕ → α × ℝ} (hn : n ≠ 0) have : n + 1 - 1 = n := by simp exact this ▸ rfl -lemma empMean_add_one_eq_empMean' {n : ℕ} {h : ℕ → α × ℝ} : - empMean a (n + 1) h = empMean' n (fun i ↦ h i) a := by +lemma empMean_add_one_eq_empMean' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} : + empMean A R' a (n + 1) ω = empMean' n (fun i ↦ (A i ω, R' i ω)) a := by unfold empMean empMean' rw [sumRewards_add_one_eq_sumRewards', pullCount_add_one_eq_pullCount'] -lemma empMean_eq_empMean' {n : ℕ} {h : ℕ → α × ℝ} (hn : n ≠ 0) : - empMean a n h = empMean' (n - 1) (fun i ↦ h i) a := by +lemma empMean_eq_empMean' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (hn : n ≠ 0) : + empMean A R' a n ω = empMean' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by unfold empMean empMean' rw [sumRewards_eq_sumRewards' hn, pullCount_eq_pullCount' hn] @[fun_prop] -lemma measurable_sumRewards [MeasurableSingletonClass α] (a : α) (t : ℕ) : - Measurable (sumRewards a t) := by +lemma measurable_sumRewards [MeasurableSingletonClass α] {R' : ℕ → Ω → ℝ} + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (t : ℕ) : + Measurable (sumRewards A R' a t) := by unfold sumRewards - have h_meas s : Measurable (fun h : ℕ → α × ℝ ↦ if action s h = a then reward s h else 0) := by + have h_meas s : Measurable (fun h : Ω ↦ if A s h = a then R' s h else 0) := by refine Measurable.ite ?_ (by fun_prop) (by fun_prop) exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop @[fun_prop] -lemma measurable_empMean [MeasurableSingletonClass α] (a : α) (n : ℕ) : - Measurable (empMean a n) := by +lemma measurable_empMean [MeasurableSingletonClass α] {R' : ℕ → Ω → ℝ} (hA : ∀ n, Measurable (A n)) + (hR' : ∀ n, Measurable (R' n)) (a : α) (n : ℕ) : + Measurable (empMean A R' a n) := by unfold empMean fun_prop diff --git a/LeanBandits/SequentialLearning/StationaryEnv.lean b/LeanBandits/SequentialLearning/StationaryEnv.lean index 995b7f48..3560acf8 100644 --- a/LeanBandits/SequentialLearning/StationaryEnv.lean +++ b/LeanBandits/SequentialLearning/StationaryEnv.lean @@ -24,24 +24,81 @@ def stationaryEnv (ν : Kernel α R) [IsMarkovKernel ν] : Environment α R wher feedback _ := ν.prodMkLeft _ ν0 := ν -variable {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] +variable {Ω : Type*} {mΩ : MeasurableSpace Ω} + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] + {P : Measure Ω} [IsProbabilityMeasure P] {A : ℕ → Ω → α} {R' : ℕ → Ω → R} + +section IsAlgEnvSeq + +/-- The conditional distribution of the reward at time `n` given the action at time `n` is `ν`. -/ +lemma IsAlgEnvSeq.condDistrib_reward_stationaryEnv + (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) : + condDistrib (R' n) (A n) P =ᵐ[P.map (A n)] ν := by + have hA := h.measurable_A + have hR' := h.measurable_R + cases n with + | zero => + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] + change P.map (step A R' 0) = P.map (A 0) ⊗ₘ ν + rw [(hasLaw_action_zero h).map_eq, (hasLaw_step_zero h).map_eq, stationaryEnv_ν0] + | succ n => + have h_eq := (h.hasCondDistrib_reward n).condDistrib_eq + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h_eq ⊢ + have : P.map (A (n + 1)) = + (P.map (fun x ↦ (hist A R' n x, A (n + 1) x))).snd := by + rw [Measure.snd_map_prodMk (by fun_prop)] + simp only [stationaryEnv_feedback] at h_eq + rw [this, ← Measure.snd_prodAssoc_compProd_prodMkLeft, ← h_eq, + Measure.snd_map_prodMk (by fun_prop), Measure.map_map (by fun_prop) (by fun_prop)] + congr + +/-- The reward at time `n + 1` is conditionally independent of the history up to time `n` +given the action at time `n + 1`. -/ +lemma IsAlgEnvSeq.condIndepFun_reward_hist_action [StandardBorelSpace Ω] [Nonempty Ω] + (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) : + R' (n + 1) ⟂ᵢ[A (n + 1), h.measurable_A _ ; P] hist A R' n := by + have hA := h.measurable_A + have hR' := h.measurable_R + exact condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft + (by fun_prop) (by fun_prop) (by fun_prop) (h.hasCondDistrib_reward n).condDistrib_eq + +lemma IsAlgEnvSeq.condIndepFun_reward_hist_action_action [StandardBorelSpace Ω] [Nonempty Ω] + (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) : + R' (n + 1) ⟂ᵢ[A (n + 1), h.measurable_A (n + 1); P] + (fun ω ↦ (hist A R' n ω, A (n + 1) ω)) := by + have h_indep : R' (n + 1) ⟂ᵢ[A (n + 1), h.measurable_A (n + 1); P] hist A R' n := by + convert h.condIndepFun_reward_hist_action n + have hA := h.measurable_A + have hR' := h.measurable_R + exact h_indep.prod_right (by fun_prop) (by fun_prop) (by fun_prop) + +lemma IsAlgEnvSeq.condIndepFun_reward_hist_action_action' [StandardBorelSpace Ω] [Nonempty Ω] + (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) (hn : n ≠ 0) : + R' n ⟂ᵢ[A n, h.measurable_A n; P] (fun ω ↦ (hist A R' (n - 1) ω, A n ω)) := by + have := h.condIndepFun_reward_hist_action_action (n - 1) + grind + +end IsAlgEnvSeq + +namespace IT local notation "𝔓" => trajMeasure alg (stationaryEnv ν) /-- The conditional distribution of the reward at time `n` given the action at time `n` is `ν`. -/ -lemma condDistrib_reward_stationaryEnv [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] (n : ℕ) : - condDistrib (reward n) (action n) 𝔓 =ᵐ[(𝔓).map (action n)] ν := by +lemma condDistrib_reward_stationaryEnv (n : ℕ) : + condDistrib (IT.reward n) (IT.action n) 𝔓 =ᵐ[(𝔓).map (IT.action n)] ν := by cases n with | zero => rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] - change (𝔓).map (step 0) = (𝔓).map (action 0) ⊗ₘ ν - rw [(hasLaw_action_zero alg (stationaryEnv ν)).map_eq, - (hasLaw_step_zero alg (stationaryEnv ν)).map_eq, stationaryEnv_ν0] + change (𝔓).map (IT.step 0) = (𝔓).map (IT.action 0) ⊗ₘ ν + rw [(IT.hasLaw_action_zero alg (stationaryEnv ν)).map_eq, + (IT.hasLaw_step_zero alg (stationaryEnv ν)).map_eq, stationaryEnv_ν0] | succ n => - have h_eq := condDistrib_reward alg (stationaryEnv ν) n + have h_eq := IT.condDistrib_reward alg (stationaryEnv ν) n rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h_eq ⊢ - have : (𝔓).map (action (n + 1)) = ((𝔓).map (fun x ↦ (hist n x, action (n + 1) x))).snd := by + have : (𝔓).map (IT.action (n + 1)) = + ((𝔓).map (fun x ↦ (IT.hist n x, IT.action (n + 1) x))).snd := by rw [Measure.snd_map_prodMk (by fun_prop)] simp only [stationaryEnv_feedback] at h_eq rw [this, ← Measure.snd_prodAssoc_compProd_prodMkLeft, ← h_eq, @@ -50,10 +107,27 @@ lemma condDistrib_reward_stationaryEnv [StandardBorelSpace α] [Nonempty α] /-- The reward at time `n + 1` is conditionally independent of the history up to time `n` given the action at time `n + 1`. -/ -lemma condIndepFun_reward_hist_action [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] (n : ℕ) : - reward (n + 1) ⟂ᵢ[action (n + 1), measurable_action _ ; 𝔓] hist n := +lemma condIndepFun_reward_hist_action (n : ℕ) : + IT.reward (n + 1) ⟂ᵢ[IT.action (n + 1), IT.measurable_action _ ; 𝔓] IT.hist n := condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft - (by fun_prop) (by fun_prop) (by fun_prop) (condDistrib_reward alg (stationaryEnv ν) n) + (by fun_prop) (by fun_prop) (by fun_prop) (IT.condDistrib_reward alg (stationaryEnv ν) n) + +lemma condIndepFun_reward_hist_action_action + {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) : + reward (n + 1) ⟂ᵢ[action (n + 1), measurable_action (n + 1); trajMeasure alg (stationaryEnv ν)] + (fun ω ↦ (hist n ω, action (n + 1) ω)) := by + have h_indep : reward (n + 1) ⟂ᵢ[action (n + 1), measurable_action (n + 1); + trajMeasure alg (stationaryEnv ν)] hist n := by + convert condIndepFun_reward_hist_action (alg := alg) (ν := ν) n + exact h_indep.prod_right (by fun_prop) (by fun_prop) (by fun_prop) + +lemma condIndepFun_reward_hist_action_action' + {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) (hn : n ≠ 0) : + reward n ⟂ᵢ[action n, measurable_action n; trajMeasure alg (stationaryEnv ν)] + (fun ω ↦ (hist (n - 1) ω, action n ω)) := by + have := condIndepFun_reward_hist_action_action (alg := alg) (ν := ν) (n - 1) + grind + +end IT end Learning From 9ab6f4429218e8b03799e56b6dcdf43dc5da1bc2 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Tue, 6 Jan 2026 15:40:46 +0100 Subject: [PATCH 05/30] more --- LeanBandits.lean | 3 +- LeanBandits/Bandit/Bandit.lean | 135 +++++---- .../{ => Bandit}/RewardByCountMeasure.lean | 7 +- LeanBandits/Bandit/SumRewards.lean | 83 ++++++ LeanBandits/BanditAlgorithms/ETC.lean | 2 +- LeanBandits/SequentialLearning/Algorithm.lean | 262 ---------------- .../SequentialLearning/Deterministic.lean | 2 +- .../SequentialLearning/FiniteActions.lean | 11 + .../IonescuTulceaSpace.lean | 282 ++++++++++++++++++ .../SequentialLearning/StationaryEnv.lean | 2 +- 10 files changed, 464 insertions(+), 325 deletions(-) rename LeanBandits/{ => Bandit}/RewardByCountMeasure.lean (98%) create mode 100644 LeanBandits/Bandit/SumRewards.lean create mode 100644 LeanBandits/SequentialLearning/IonescuTulceaSpace.lean diff --git a/LeanBandits.lean b/LeanBandits.lean index 10b0f463..4de3f6dd 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -1,5 +1,6 @@ import LeanBandits.Bandit.Bandit import LeanBandits.Bandit.Regret +import LeanBandits.Bandit.RewardByCountMeasure import LeanBandits.BanditAlgorithms.ETC import LeanBandits.BanditAlgorithms.UCB import LeanBandits.ForMathlib.CondDistrib @@ -11,8 +12,8 @@ import LeanBandits.ForMathlib.Measurable import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.ForMathlib.SubGaussian import LeanBandits.ForMathlib.Traj -import LeanBandits.RewardByCountMeasure import LeanBandits.SequentialLearning.Algorithm import LeanBandits.SequentialLearning.Deterministic import LeanBandits.SequentialLearning.FiniteActions +import LeanBandits.SequentialLearning.IonescuTulceaSpace import LeanBandits.SequentialLearning.StationaryEnv diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 1788b187..ea3e8601 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -187,7 +187,7 @@ example [StandardBorelSpace α] [Nonempty α] end DetAlgorithm -section ArrayModel +namespace ArrayModel open unitInterval @@ -250,89 +250,108 @@ lemma measurable_algFunction (alg : Algorithm α R) (n : ℕ) : (representation (alg.policy n)).choose_spec.1 noncomputable -def altHist [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) : (n : ℕ) → Iic n → α × R +def hist [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) : (n : ℕ) → Iic n → α × R | 0 => fun _ ↦ (initAlgFunction alg (ω.1 0), ω.2 0 (initAlgFunction alg (ω.1 0))) | n + 1 => - let hn : Iic n → α × R := altHist alg ω n + let hn : Iic n → α × R := hist alg ω n let a : α := algFunction alg n hn (ω.1 (n + 1)) - fun i ↦ if hin : i ≤ n then hn ⟨i, by simp [hin]⟩ else (a, ω.2 (pullCount' n hn a + 1) a) + fun i ↦ if hin : i ≤ n then hn ⟨i, by simp [hin]⟩ else (a, ω.2 (pullCount' n hn a) a) @[simp] -lemma altHist_zero [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) : - altHist alg ω 0 = fun _ ↦ (initAlgFunction alg (ω.1 0), ω.2 0 (initAlgFunction alg (ω.1 0))) := +lemma hist_zero [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) : + hist alg ω 0 = fun _ ↦ (initAlgFunction alg (ω.1 0), ω.2 0 (initAlgFunction alg (ω.1 0))) := rfl -lemma altHist_add_one [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) (n : ℕ) : - let a : α := algFunction alg n (altHist alg ω n) (ω.1 (n + 1)) - altHist alg ω (n + 1) = - fun (i : Iic (n + 1)) ↦ if hin : i ≤ n then altHist alg ω n ⟨i, by simp [hin]⟩ - else (a, ω.2 (pullCount' n (altHist alg ω n) a + 1) a) := +lemma hist_add_one [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) (n : ℕ) : + let a : α := algFunction alg n (hist alg ω n) (ω.1 (n + 1)) + hist alg ω (n + 1) = + fun (i : Iic (n + 1)) ↦ if hin : i ≤ n then hist alg ω n ⟨i, by simp [hin]⟩ + else (a, ω.2 (pullCount' n (hist alg ω n) a) a) := rfl -lemma altHist_eq [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) (n : ℕ) : - altHist alg ω n = fun i : Iic n ↦ altHist alg ω i ⟨i.1, by simp⟩ := by +lemma hist_eq [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) (n : ℕ) : + hist alg ω n = fun i : Iic n ↦ hist alg ω i ⟨i.1, by simp⟩ := by induction n with | zero => ext i : 1 - simp only [altHist] + simp only [hist] sorry | succ n hn => ext i : 1 by_cases hin : i ≤ n - · rw [altHist_add_one] + · rw [hist_add_one] simp only [hin, ↓reduceDIte] rw [funext_iff] at hn simp_rw [hn] · grind @[fun_prop] -lemma measurable_altHist [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : - Measurable (fun ω ↦ altHist alg ω n) := by +lemma measurable_hist [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : + Measurable (fun ω ↦ hist alg ω n) := by induction n with | zero => - simp_rw [altHist_zero, measurable_pi_iff] + simp_rw [hist_zero, measurable_pi_iff] refine fun _ ↦ Measurable.prodMk (by fun_prop) ?_ sorry | succ n hn => refine measurable_pi_iff.mpr fun i ↦ ?_ by_cases hin : i ≤ n - · simp only [altHist, hin, ↓reduceDIte] + · simp only [hist, hin, ↓reduceDIte] rw [measurable_pi_iff] at hn exact hn ⟨i.1, by simp [hin]⟩ - · simp only [altHist, hin, ↓reduceDIte] + · simp only [hist, hin, ↓reduceDIte] refine Measurable.prodMk (by fun_prop) ?_ sorry noncomputable -def altaction [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : α := - (altHist alg ω n ⟨n, by simp⟩).1 +def action [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : α := + (hist alg ω n ⟨n, by simp⟩).1 -lemma altaction_zero [DecidableEq α] (alg : Algorithm α R) : - altaction alg 0 = fun ω ↦ initAlgFunction alg (ω.1 0) := by +lemma action_zero [DecidableEq α] (alg : Algorithm α R) : + action alg 0 = fun ω ↦ initAlgFunction alg (ω.1 0) := by ext - simp [altaction, altHist_zero] + simp [action, hist_zero] + +lemma action_add_one_eq [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : + action alg (n + 1) = fun ω ↦ algFunction alg n (hist alg ω n) (ω.1 (n + 1)) := by + ext ω + rw [action, hist_add_one] + simp only [add_le_iff_nonpos_right, nonpos_iff_eq_zero, one_ne_zero, ↓reduceDIte] @[fun_prop] -lemma measurable_altaction [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : - Measurable (altaction alg n) := by unfold altaction; fun_prop +lemma measurable_action [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : + Measurable (action alg n) := by unfold action; fun_prop noncomputable -def altReward [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : R := - (altHist alg ω n ⟨n, by simp⟩).2 +def reward [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : R := + (hist alg ω n ⟨n, by simp⟩).2 -lemma altReward_zero [DecidableEq α] (alg : Algorithm α R) : - altReward alg 0 = fun ω ↦ ω.2 0 (altaction alg 0 ω) := by +lemma reward_zero [DecidableEq α] (alg : Algorithm α R) : + reward alg 0 = fun ω ↦ ω.2 0 (action alg 0 ω) := by ext - simp [altReward, altHist_zero, altaction_zero] + simp [reward, hist_zero, action_zero] + +lemma reward_eq [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : + reward alg n = fun ω ↦ ω.2 (pullCount (action alg) (action alg n ω) n ω) (action alg n ω) := by + cases n with + | zero => ext; simp [reward_zero, action_zero] + | succ n => + ext ω + rw [reward, hist_add_one] + simp only [add_le_iff_nonpos_right, nonpos_iff_eq_zero, one_ne_zero, ↓reduceDIte] + rw [action_add_one_eq, pullCount_eq_pullCount' (R' := reward alg) (by simp)] + simp only [Nat.add_one_sub_one] + rw [hist_eq] + rfl @[fun_prop] -lemma measurable_altReward [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : - Measurable (altReward alg n) := by unfold altReward; fun_prop +lemma measurable_reward [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : + Measurable (reward alg n) := by unfold reward; fun_prop variable [DecidableEq α] -lemma hasLaw_altaction_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : - HasLaw (altaction alg 0) alg.p0 (arrayMeasure ν) where +lemma hasLaw_action_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : + HasLaw (action alg 0) alg.p0 (arrayMeasure ν) where map_eq := by calc (arrayMeasure ν).map (fun ω ↦ initAlgFunction alg (ω.1 0)) _ = ((arrayMeasure ν).fst.map (Function.eval 0)).map (initAlgFunction alg) := by @@ -346,15 +365,15 @@ lemma hasLaw_altaction_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovK variable [StandardBorelSpace R] [Nonempty R] -lemma hasCondDistrib_altReward_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : - HasCondDistrib (altReward alg 0) (altaction alg 0) (stationaryEnv ν).ν0 (arrayMeasure ν) where +lemma hasCondDistrib_reward_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : + HasCondDistrib (reward alg 0) (action alg 0) (stationaryEnv ν).ν0 (arrayMeasure ν) where condDistrib_eq := by - simp only [stationaryEnv_ν0, (hasLaw_altaction_zero alg ν).map_eq, altReward_zero] + simp only [stationaryEnv_ν0, (hasLaw_action_zero alg ν).map_eq, reward_zero] sorry lemma hasCondDistrib_altStep' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - HasCondDistrib (altHist alg · (n + 1) ⟨n + 1, by simp⟩) (altHist alg · n) + HasCondDistrib (hist alg · (n + 1) ⟨n + 1, by simp⟩) (hist alg · n) (Bandit.stepKernel alg ν n) (arrayMeasure ν) where condDistrib_eq := by simp only [Bandit.stepKernel, stepKernel, stationaryEnv_feedback] @@ -362,36 +381,36 @@ lemma hasCondDistrib_altStep' (alg : Algorithm α R) (ν : Kernel α R) [IsMarko lemma hasCondDistrib_altStep (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - HasCondDistrib (fun ω ↦ (altaction alg (n + 1) ω, altReward alg (n + 1) ω)) - (fun ω (i : Iic n) ↦ (altaction alg i ω, altReward alg i ω)) + HasCondDistrib (fun ω ↦ (action alg (n + 1) ω, reward alg (n + 1) ω)) + (fun ω (i : Iic n) ↦ (action alg i ω, reward alg i ω)) (Bandit.stepKernel alg ν n) (arrayMeasure ν) := by convert hasCondDistrib_altStep' alg ν n with ω i - · simp only [altaction] - rw [altHist_eq _ _ n] - · simp only [altReward] - rw [altHist_eq _ _ n] - -lemma hasCondDistrib_altaction (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - HasCondDistrib (altaction alg (n + 1)) - (fun ω (i : Iic n) ↦ (altaction alg i ω, altReward alg i ω)) + · simp only [action] + rw [hist_eq _ _ n] + · simp only [reward] + rw [hist_eq _ _ n] + +lemma hasCondDistrib_action (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + HasCondDistrib (action alg (n + 1)) + (fun ω (i : Iic n) ↦ (action alg i ω, reward alg i ω)) (alg.policy n) (arrayMeasure ν) := by convert HasCondDistrib.fst (hasCondDistrib_altStep alg ν n) simp -lemma hasCondDistrib_altReward (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] +lemma hasCondDistrib_reward (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - HasCondDistrib (altReward alg (n + 1)) - (fun ω ↦ (fun (i : Iic n) ↦ (altaction alg i ω, altReward alg i ω), altaction alg (n + 1) ω)) + HasCondDistrib (reward alg (n + 1)) + (fun ω ↦ (fun (i : Iic n) ↦ (action alg i ω, reward alg i ω), action alg (n + 1) ω)) ((stationaryEnv ν).feedback n) (arrayMeasure ν) := by simp only [stationaryEnv_feedback] sorry lemma isAlgEnvSeq_arrayMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : - IsAlgEnvSeq (altaction alg) (altReward alg) alg (stationaryEnv ν) (arrayMeasure ν) where - hasLaw_action_zero := hasLaw_altaction_zero alg ν - hasCondDistrib_reward_zero := hasCondDistrib_altReward_zero alg ν - hasCondDistrib_action := hasCondDistrib_altaction alg ν - hasCondDistrib_reward := hasCondDistrib_altReward alg ν + IsAlgEnvSeq (action alg) (reward alg) alg (stationaryEnv ν) (arrayMeasure ν) where + hasLaw_action_zero := hasLaw_action_zero alg ν + hasCondDistrib_reward_zero := hasCondDistrib_reward_zero alg ν + hasCondDistrib_action := hasCondDistrib_action alg ν + hasCondDistrib_reward := hasCondDistrib_reward alg ν end ArrayModel diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/Bandit/RewardByCountMeasure.lean similarity index 98% rename from LeanBandits/RewardByCountMeasure.lean rename to LeanBandits/Bandit/RewardByCountMeasure.lean index 980b01c7..6a2e7083 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/Bandit/RewardByCountMeasure.lean @@ -295,13 +295,18 @@ lemma iIndepFun_rewardByCount' [StandardBorelSpace Ω] [Nonempty Ω] [Countable rw [iIndepFun_nat_iff_forall_indepFun (by fun_prop)] exact indepFun_rewardByCount_Iic h a +lemma iIndepFun_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) : + iIndepFun (fun (p : α × ℕ) ↦ rewardByCount A R p.1 (p.2 + 1)) 𝔓' := by + sorry + lemma identDistrib_rewardByCount_stream_all [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) : IdentDistrib (fun ω (p : α × ℕ) ↦ rewardByCount A R p.1 (p.2 + 1) ω) (fun ω p ↦ ω p.2 p.1) 𝔓' (Bandit.streamMeasure ν) := by refine IdentDistrib.pi (fun p ↦ ?_) ?_ ?_ · refine identDistrib_rewardByCount_eval h p.1 (p.2 + 1) p.2 (by simp) (ν := ν) - · sorry + · exact iIndepFun_rewardByCount h · sorry lemma identDistrib_rewardByCount_stream' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean new file mode 100644 index 00000000..c47a4bbb --- /dev/null +++ b/LeanBandits/Bandit/SumRewards.lean @@ -0,0 +1,83 @@ +/- +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 LeanBandits.Bandit.Regret + +/-! # Law of the sum of rewards +-/ + +open MeasureTheory ProbabilityTheory Finset Learning +open scoped ENNReal NNReal + +namespace Bandits + +namespace ArrayModel + +variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [StandardBorelSpace α] [Nonempty α] + {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] + +local notation "A" => action alg +local notation "R" => reward alg + +lemma identDistrib_sum_Icc_rewardByCount_pullCount' [Countable α] (a : α) (n : ℕ) : + IdentDistrib (fun ω ↦ ∑ i ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a i ω) + (fun ω ↦ ∑ i ∈ Icc 1 (pullCount A a n ω), ω.2 (i - 1) a) + ((arrayMeasure ν).prod (Bandit.streamMeasure ν)) (arrayMeasure ν) where + aemeasurable_fst := sorry -- the issue is the random bound + aemeasurable_snd := sorry + map_eq := by + by_cases hn : n = 0 + · simp [hn] + have h_eq (i : ℕ) (ω : probSpace α ℝ × (ℕ → α → ℝ)) (hi : i ∈ Icc 1 (pullCount A a n ω.1)) : + rewardByCount A R a i ω = ω.1.2 (i - 1) a := by + rw [rewardByCount_of_stepsUntil_ne_top] + · simp only [reward_eq] + have h_exists : ∃ s, pullCount A a (s + 1) ω.1 = i := + exists_pullCount_eq_of_le (n := n - 1) (by grind) (by grind) + have h_action : A (stepsUntil A a i ω.1).toNat ω.1 = a := + action_stepsUntil («A» := A) (by grind) h_exists + congr! + rw [h_action, pullCount_stepsUntil (by grind) h_exists] + · have : stepsUntil A a (pullCount A a (n + 1) ω.1) ω.1 ≠ ⊤ := by + refine ne_top_of_le_ne_top ?_ (stepsUntil_pullCount_le _ _ _) + simp + refine ne_top_of_le_ne_top this ?_ + refine stepsUntil_mono a ω.1 (by grind) ?_ + simp only [mem_Icc] at hi + refine hi.2.trans ?_ + exact pullCount_mono _ (by grind) _ + have h_sum_eq (ω : probSpace α ℝ × (ℕ → α → ℝ)) : + ∑ i ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a i ω = + ∑ i ∈ Icc 1 (pullCount A a n ω.1), ω.1.2 (i - 1) a := + Finset.sum_congr rfl fun i hi ↦ h_eq i ω hi + simp_rw [h_sum_eq] + conv_rhs => rw [← Measure.fst_prod (μ := arrayMeasure ν) (ν := Bandit.streamMeasure ν), + Measure.fst] + rw [AEMeasurable.map_map_of_aemeasurable _ (by fun_prop)] + · rfl + simp only [Measure.map_fst_prod, measure_univ, one_smul] + sorry + +lemma identDistrib_sum_Icc_rewardByCount_pullCount [Countable α] (a : α) (n : ℕ) : + IdentDistrib (fun ω ↦ ∑ i ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a i ω) + (fun ω ↦ ∑ i ∈ range (pullCount A a n ω), ω.2 i a) + ((arrayMeasure ν).prod (Bandit.streamMeasure ν)) (arrayMeasure ν) := by + convert identDistrib_sum_Icc_rewardByCount_pullCount' a n using 2 with ω + swap; · infer_instance + sorry + +lemma identDistrib_sumRewards_pullCount [Countable α] (a : α) (n : ℕ) : + IdentDistrib (sumRewards A R a n) (fun ω ↦ ∑ i ∈ range (pullCount A a n ω), ω.2 i a) + (arrayMeasure ν) (arrayMeasure ν) := by + suffices IdentDistrib (fun ω ↦ sumRewards A R a n ω.1) + (fun ω ↦ ∑ i ∈ range (pullCount A a n ω), ω.2 i a) + ((arrayMeasure ν).prod (Bandit.streamMeasure ν)) (arrayMeasure ν) by + sorry + simp_rw [← sum_rewardByCount_eq_sumRewards] + exact identDistrib_sum_Icc_rewardByCount_pullCount a n + +end ArrayModel + +end Bandits diff --git a/LeanBandits/BanditAlgorithms/ETC.lean b/LeanBandits/BanditAlgorithms/ETC.lean index 49c3d90c..fe691261 100644 --- a/LeanBandits/BanditAlgorithms/ETC.lean +++ b/LeanBandits/BanditAlgorithms/ETC.lean @@ -3,9 +3,9 @@ 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 LeanBandits.Bandit.RewardByCountMeasure import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.ForMathlib.SubGaussian -import LeanBandits.RewardByCountMeasure /-! # The Explore-Then-Commit Algorithm diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index a9d9b183..5f7e2b92 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -238,266 +238,4 @@ theorem isAlgEnvSeq_unique (h1 : IsAlgEnvSeq A₁ R₁ alg env P) end ModelEquivalence -namespace IT - -/-- Action and reward at step `n`. -/ -def step (n : ℕ) (h : ℕ → α × R) : α × R := h n - -/-- `action n` is the action pulled at time `n`. This is a random variable on the measurable space -`ℕ → α × ℝ`. -/ -def action (n : ℕ) (h : ℕ → α × R) : α := (h n).1 - -/-- `reward n` is the reward at time `n`. This is a random variable on the measurable space -`ℕ → α × R`. -/ -def reward (n : ℕ) (h : ℕ → α × R) : R := (h n).2 - -/-- `hist n` is the history up to time `n`. This is a random variable on the measurable space -`ℕ → α × R`. -/ -def hist (n : ℕ) (h : ℕ → α × R) : Iic n → α × R := fun i ↦ h i - -lemma fst_comp_step (n : ℕ) : Prod.fst ∘ step (α := α) (R := R) n = action n := rfl - -@[fun_prop] -lemma measurable_step (n : ℕ) : Measurable (step n (α := α) (R := R)) := by - unfold step; fun_prop - -@[fun_prop] -lemma measurable_step_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ step p.1 p.2) := - measurable_from_prod_countable_right fun n ↦ (by simp only; fun_prop) - -@[fun_prop] -lemma measurable_action (n : ℕ) : Measurable (action n (α := α) (R := R)) := by - unfold action; fun_prop - -@[fun_prop] -lemma measurable_action_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ action p.1 p.2) := - measurable_from_prod_countable_right fun n ↦ (by simp only; fun_prop) - -@[fun_prop] -lemma measurable_reward (n : ℕ) : Measurable (reward n (α := α) (R := R)) := by - unfold reward; fun_prop - -@[fun_prop] -lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ reward p.1 p.2) := - measurable_from_prod_countable_right fun n ↦ (by simp only; fun_prop) - -@[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 algorithm Seq. -/ -protected def filtration (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : - Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) := - MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R) - -lemma filtration_eq_comap (n : ℕ) : - IT.filtration α R n = MeasurableSpace.comap (hist n) inferInstance := by - simp [IT.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe] - -lemma step_eq_eval_comp_hist (n : ℕ) : - step (α := α) (R := R) n = (fun x ↦ x ⟨n, by simp⟩) ∘ (hist n) := rfl - -lemma action_eq_eval_comp_hist (n : ℕ) : - action (α := α) (R := R) n = (fun x ↦ (x ⟨n, by simp⟩).1) ∘ (hist n) := rfl - -lemma reward_eq_eval_comp_hist (n : ℕ) : - reward (α := α) (R := R) n = (fun x ↦ (x ⟨n, by simp⟩).2) ∘ (hist n) := rfl - -lemma measurable_step_filtration (n : ℕ) : Measurable[IT.filtration α R n] (step n) := by - rw [filtration_eq_comap, step_eq_eval_comp_hist] - exact measurable_comp_comap _ (by fun_prop) - -lemma adapted_step [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] - [SecondCountableTopology α] [OpensMeasurableSpace α] - [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] - [SecondCountableTopology R] [OpensMeasurableSpace R] : - Adapted (IT.filtration α R) (step (α := α) (R := R)) := - fun n ↦ (measurable_step_filtration n).stronglyMeasurable - -lemma measurable_hist_filtration (n : ℕ) : Measurable[IT.filtration α R n] (hist n) := by - simp [filtration_eq_comap, measurable_iff_comap_le] - -lemma adapted_hist [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] - [SecondCountableTopology α] [OpensMeasurableSpace α] - [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] - [SecondCountableTopology R] [OpensMeasurableSpace R] : - Adapted (IT.filtration α R) hist := - fun n ↦ (measurable_hist_filtration n).stronglyMeasurable - -lemma measurable_action_filtration (n : ℕ) : Measurable[IT.filtration α R n] (action n) := by - rw [filtration_eq_comap, action_eq_eval_comp_hist] - exact measurable_comp_comap _ (by fun_prop) - -lemma adapted_action [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] - [SecondCountableTopology α] [OpensMeasurableSpace α] : - Adapted (IT.filtration α R) action := - fun n ↦ (measurable_action_filtration n).stronglyMeasurable - -lemma measurable_reward_filtration (n : ℕ) : Measurable[IT.filtration α R n] (reward n) := by - rw [filtration_eq_comap, reward_eq_eval_comp_hist] - exact measurable_comp_comap _ (by fun_prop) - -lemma adapted_reward [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] - [SecondCountableTopology R] [OpensMeasurableSpace R] : - Adapted (IT.filtration α R) reward := - fun n ↦ (measurable_reward_filtration n).stronglyMeasurable - -section FiltrationAction - -/-- Filtration generated by the history at time `n-1` together with the action at time `n`. -/ -def filtrationAction (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : - Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) where - seq n := if n = 0 then MeasurableSpace.comap (action 0) inferInstance - else IT.filtration α R (n - 1) ⊔ MeasurableSpace.comap (action n) inferInstance - mono' n m hnm := by - simp only - by_cases hn : n = 0 - · by_cases hm : m = 0 - · simp [hn, hm] - · simp only [hn, ↓reduceIte, hm] - refine le_sup_of_le_left ?_ - rw [← measurable_iff_comap_le] - suffices Measurable[IT.filtration α R 0] (action 0) from - this.mono ((IT.filtration α R).mono zero_le') le_rfl - exact measurable_action_filtration 0 - have hm : m ≠ 0 := by grind - simp only [hn, hm, ↓reduceIte] - have hnm' : n - 1 ≤ m - 1 := by grind - simp only [sup_le_iff] - constructor - · refine le_sup_of_le_left ?_ - exact (IT.filtration α R).mono hnm' - · rcases eq_or_lt_of_le hnm with rfl | hlt - · exact le_sup_of_le_right le_rfl - refine le_sup_of_le_left ?_ - rw [← measurable_iff_comap_le] - have h_le : n ≤ m - 1 := by grind - suffices Measurable[IT.filtration α R n] (action n) from - this.mono ((IT.filtration α R).mono h_le) le_rfl - exact measurable_action_filtration n - le' n := by - by_cases hn : n = 0 - · simp only [hn, ↓reduceIte] - rw [← measurable_iff_comap_le] - fun_prop - simp only [hn, ↓reduceIte, sup_le_iff] - constructor - · exact (IT.filtration α R).le _ - · rw [← measurable_iff_comap_le] - fun_prop - -lemma filtrationAction_zero_eq_comap : - filtrationAction α R 0 = MeasurableSpace.comap (action 0) inferInstance := by - simp [filtrationAction] - -lemma filtrationAction_eq_comap (n : ℕ) (hn : n ≠ 0) : - filtrationAction α R n = - MeasurableSpace.comap (fun ω ↦ (hist (n - 1) ω, action n ω)) inferInstance := by - simp only [filtrationAction, filtration_eq_comap, ← MeasurableSpace.comap_prodMk, hn, ↓reduceIte] - rfl - -lemma filtration_le_filtrationAction_add_one (n : ℕ) : - IT.filtration α R n ≤ filtrationAction α R (n + 1) := le_sup_of_le_left le_rfl - -lemma filtration_le_filtrationAction {m n : ℕ} (h : n < m) : - IT.filtration α R n ≤ filtrationAction α R m := by - have h' : n + 1 ≤ m := by grind - exact (filtration_le_filtrationAction_add_one n).trans ((filtrationAction α R).mono h') - -lemma filtrationAction_le_filtration_self (n : ℕ) : - filtrationAction α R n ≤ IT.filtration α R n := by - by_cases hn : n = 0 - · simp only [hn, filtrationAction_zero_eq_comap] - rw [← measurable_iff_comap_le] - exact measurable_action_filtration 0 - simp only [filtrationAction, hn, ↓reduceIte, sup_le_iff] - constructor - · exact (IT.filtration α R).mono (by grind) - · rw [← measurable_iff_comap_le] - exact measurable_action_filtration _ - -lemma filtrationAction_le_filtration {m n : ℕ} (h : m ≤ n) : - filtrationAction α R m ≤ IT.filtration α R n := - (filtrationAction_le_filtration_self m).trans ((IT.filtration α R).mono h) - -lemma measurable_action_filtrationAction (n : ℕ) : - Measurable[filtrationAction α R n] (action n) := by - simp only [filtrationAction] - rw [measurable_iff_comap_le] - split_ifs with hn - · simp [hn] - · exact le_sup_of_le_right le_rfl - -end FiltrationAction - -section Laws - -lemma hasLaw_step_zero (alg : Algorithm α R) (env : Environment α R) : - HasLaw (step 0) (alg.p0 ⊗ₘ env.ν0) (trajMeasure alg env) where - aemeasurable := Measurable.aemeasurable (by fun_prop) - map_eq := by - unfold step - rw [← coe_default_Iic_zero] - simp only [trajMeasure, Kernel.trajMeasure] - rw [← Measure.deterministic_comp_eq_map (by fun_prop), Measure.comp_assoc, - Kernel.deterministic_comp_eq_map, Kernel.traj_zero_map_eval_zero, - Measure.deterministic_comp_eq_map, Measure.map_map (by fun_prop) (by fun_prop)] - exact Measure.map_id - -lemma hasLaw_action_zero (alg : Algorithm α R) (env : Environment α R) : - HasLaw (action 0) alg.p0 (trajMeasure alg env) where - map_eq := by - rw [← fst_comp_step, ← Measure.map_map (by fun_prop) (by fun_prop), - (hasLaw_step_zero alg env).map_eq, ← Measure.fst, Measure.fst_compProd] - -variable [StandardBorelSpace R] [Nonempty R] - -lemma condDistrib_reward_zero (alg : Algorithm α R) (env : Environment α R) : - condDistrib (reward 0) (action 0) (trajMeasure alg env) - =ᵐ[(trajMeasure alg env).map (action 0)] env.ν0 := by - have h_step := (hasLaw_step_zero alg env).map_eq - have h_action := (hasLaw_action_zero alg env).map_eq - rwa [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop), h_action] - -variable [StandardBorelSpace α] [Nonempty α] - -lemma condDistrib_step (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : - condDistrib (step (n + 1)) (hist n) (trajMeasure alg env) - =ᵐ[(trajMeasure alg env).map (hist n)] stepKernel alg env n := - Kernel.condDistrib_trajMeasure - -lemma condDistrib_action (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : - condDistrib (action (n + 1)) (hist n) (trajMeasure alg env) - =ᵐ[(trajMeasure alg env).map (hist n)] alg.policy n := by - rw [← fst_comp_step] - refine (condDistrib_comp _ (by fun_prop) (by fun_prop)).trans ?_ - filter_upwards [condDistrib_step alg env n] with h h_eq - rw [Kernel.map_apply _ (by fun_prop), h_eq, ← Kernel.map_apply _ (by fun_prop), ← Kernel.fst_eq, - fst_stepKernel] - -lemma condDistrib_reward (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : - condDistrib (reward (n + 1)) (fun ω ↦ (hist n ω, action (n + 1) ω)) (trajMeasure alg env) - =ᵐ[(trajMeasure alg env).map (fun ω ↦ (hist n ω, action (n + 1) ω))] env.feedback n := by - have h_step := condDistrib_step alg env n - have h_action := condDistrib_action alg env n - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h_step h_action ⊢ - rw [h_action, ← Measure.compProd_assoc, ← stepKernel, ← h_step, - Measure.map_map (by fun_prop) (by fun_prop)] - rfl - -lemma isAlgEnvSeq_trajMeasure (alg : Algorithm α R) (env : Environment α R) : - IsAlgEnvSeq action reward alg env (trajMeasure alg env) where - hasLaw_action_zero := hasLaw_action_zero alg env - hasCondDistrib_reward_zero := ⟨by fun_prop, by fun_prop, condDistrib_reward_zero alg env⟩ - hasCondDistrib_action n := ⟨by fun_prop, by fun_prop, condDistrib_action alg env n⟩ - hasCondDistrib_reward n := ⟨by fun_prop, by fun_prop, condDistrib_reward alg env n⟩ - -end Laws - -end IT - end Learning diff --git a/LeanBandits/SequentialLearning/Deterministic.lean b/LeanBandits/SequentialLearning/Deterministic.lean index 2c1ba945..1bab62a4 100644 --- a/LeanBandits/SequentialLearning/Deterministic.lean +++ b/LeanBandits/SequentialLearning/Deterministic.lean @@ -3,7 +3,7 @@ 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 LeanBandits.SequentialLearning.Algorithm +import LeanBandits.SequentialLearning.IonescuTulceaSpace /-! # Deterministic algorithms diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index 18d57302..c1b48a8c 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -261,6 +261,17 @@ lemma stepsUntil_eq_leastGE (a : α) (hm : m ≠ 0) : refine hn.not_ge ?_ exact csInf_le (by simp) (by simp [h_contra]) +lemma stepsUntil_mono (a : α) (ω : Ω) {n m : ℕ} (hn : n ≠ 0) (hnm : n ≤ m) : + stepsUntil A a n ω ≤ stepsUntil A a m ω := by + rw [stepsUntil_eq_leastGE a hn, stepsUntil_eq_leastGE a (by lia)] + simp_rw [leastGE] + have h_Ici_subset : Set.Ici (m : ℝ) ⊆ Set.Ici (n : ℝ) := by + intro x hx + simp only [Set.mem_Ici] at hx ⊢ + refine le_trans ?_ hx + exact mod_cast hnm + exact hittingAfter_anti (fun n ω ↦ (pullCount A a (n + 1) ω : ℝ)) 0 h_Ici_subset ω + lemma stepsUntil_pullCount_le (ω : Ω) (a : α) (t : ℕ) : stepsUntil A a (pullCount A a (t + 1) ω) ω ≤ t := by rw [stepsUntil] diff --git a/LeanBandits/SequentialLearning/IonescuTulceaSpace.lean b/LeanBandits/SequentialLearning/IonescuTulceaSpace.lean new file mode 100644 index 00000000..9057372f --- /dev/null +++ b/LeanBandits/SequentialLearning/IonescuTulceaSpace.lean @@ -0,0 +1,282 @@ +/- +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, Paulo Rauber +-/ +import LeanBandits.SequentialLearning.Algorithm + +/-! +# Algorithms +-/ + +open MeasureTheory ProbabilityTheory Filter Real Finset + +open scoped ENNReal NNReal + +namespace Learning + +variable {α R Ω : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} {mΩ : MeasurableSpace Ω} + +namespace IT + +/-- Action and reward at step `n`. -/ +def step (n : ℕ) (h : ℕ → α × R) : α × R := h n + +/-- `action n` is the action pulled at time `n`. This is a random variable on the measurable space +`ℕ → α × ℝ`. -/ +def action (n : ℕ) (h : ℕ → α × R) : α := (h n).1 + +/-- `reward n` is the reward at time `n`. This is a random variable on the measurable space +`ℕ → α × R`. -/ +def reward (n : ℕ) (h : ℕ → α × R) : R := (h n).2 + +/-- `hist n` is the history up to time `n`. This is a random variable on the measurable space +`ℕ → α × R`. -/ +def hist (n : ℕ) (h : ℕ → α × R) : Iic n → α × R := fun i ↦ h i + +lemma fst_comp_step (n : ℕ) : Prod.fst ∘ step (α := α) (R := R) n = action n := rfl + +@[fun_prop] +lemma measurable_step (n : ℕ) : Measurable (step n (α := α) (R := R)) := by + unfold step; fun_prop + +@[fun_prop] +lemma measurable_step_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ step p.1 p.2) := + measurable_from_prod_countable_right fun n ↦ (by simp only; fun_prop) + +@[fun_prop] +lemma measurable_action (n : ℕ) : Measurable (action n (α := α) (R := R)) := by + unfold action; fun_prop + +@[fun_prop] +lemma measurable_action_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ action p.1 p.2) := + measurable_from_prod_countable_right fun n ↦ (by simp only; fun_prop) + +@[fun_prop] +lemma measurable_reward (n : ℕ) : Measurable (reward n (α := α) (R := R)) := by + unfold reward; fun_prop + +@[fun_prop] +lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ reward p.1 p.2) := + measurable_from_prod_countable_right fun n ↦ (by simp only; fun_prop) + +@[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 algorithm Seq. -/ +protected def filtration (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : + Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) := + MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R) + +lemma filtration_eq_comap (n : ℕ) : + IT.filtration α R n = MeasurableSpace.comap (hist n) inferInstance := by + simp [IT.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe] + +lemma step_eq_eval_comp_hist (n : ℕ) : + step (α := α) (R := R) n = (fun x ↦ x ⟨n, by simp⟩) ∘ (hist n) := rfl + +lemma action_eq_eval_comp_hist (n : ℕ) : + action (α := α) (R := R) n = (fun x ↦ (x ⟨n, by simp⟩).1) ∘ (hist n) := rfl + +lemma reward_eq_eval_comp_hist (n : ℕ) : + reward (α := α) (R := R) n = (fun x ↦ (x ⟨n, by simp⟩).2) ∘ (hist n) := rfl + +lemma measurable_step_filtration (n : ℕ) : Measurable[IT.filtration α R n] (step n) := by + rw [filtration_eq_comap, step_eq_eval_comp_hist] + exact measurable_comp_comap _ (by fun_prop) + +lemma adapted_step [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] + [SecondCountableTopology α] [OpensMeasurableSpace α] + [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] + [SecondCountableTopology R] [OpensMeasurableSpace R] : + Adapted (IT.filtration α R) (step (α := α) (R := R)) := + fun n ↦ (measurable_step_filtration n).stronglyMeasurable + +lemma measurable_hist_filtration (n : ℕ) : Measurable[IT.filtration α R n] (hist n) := by + simp [filtration_eq_comap, measurable_iff_comap_le] + +lemma adapted_hist [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] + [SecondCountableTopology α] [OpensMeasurableSpace α] + [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] + [SecondCountableTopology R] [OpensMeasurableSpace R] : + Adapted (IT.filtration α R) hist := + fun n ↦ (measurable_hist_filtration n).stronglyMeasurable + +lemma measurable_action_filtration (n : ℕ) : Measurable[IT.filtration α R n] (action n) := by + rw [filtration_eq_comap, action_eq_eval_comp_hist] + exact measurable_comp_comap _ (by fun_prop) + +lemma adapted_action [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] + [SecondCountableTopology α] [OpensMeasurableSpace α] : + Adapted (IT.filtration α R) action := + fun n ↦ (measurable_action_filtration n).stronglyMeasurable + +lemma measurable_reward_filtration (n : ℕ) : Measurable[IT.filtration α R n] (reward n) := by + rw [filtration_eq_comap, reward_eq_eval_comp_hist] + exact measurable_comp_comap _ (by fun_prop) + +lemma adapted_reward [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] + [SecondCountableTopology R] [OpensMeasurableSpace R] : + Adapted (IT.filtration α R) reward := + fun n ↦ (measurable_reward_filtration n).stronglyMeasurable + +section FiltrationAction + +/-- Filtration generated by the history at time `n-1` together with the action at time `n`. -/ +def filtrationAction (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : + Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) where + seq n := if n = 0 then MeasurableSpace.comap (action 0) inferInstance + else IT.filtration α R (n - 1) ⊔ MeasurableSpace.comap (action n) inferInstance + mono' n m hnm := by + simp only + by_cases hn : n = 0 + · by_cases hm : m = 0 + · simp [hn, hm] + · simp only [hn, ↓reduceIte, hm] + refine le_sup_of_le_left ?_ + rw [← measurable_iff_comap_le] + suffices Measurable[IT.filtration α R 0] (action 0) from + this.mono ((IT.filtration α R).mono zero_le') le_rfl + exact measurable_action_filtration 0 + have hm : m ≠ 0 := by grind + simp only [hn, hm, ↓reduceIte] + have hnm' : n - 1 ≤ m - 1 := by grind + simp only [sup_le_iff] + constructor + · refine le_sup_of_le_left ?_ + exact (IT.filtration α R).mono hnm' + · rcases eq_or_lt_of_le hnm with rfl | hlt + · exact le_sup_of_le_right le_rfl + refine le_sup_of_le_left ?_ + rw [← measurable_iff_comap_le] + have h_le : n ≤ m - 1 := by grind + suffices Measurable[IT.filtration α R n] (action n) from + this.mono ((IT.filtration α R).mono h_le) le_rfl + exact measurable_action_filtration n + le' n := by + by_cases hn : n = 0 + · simp only [hn, ↓reduceIte] + rw [← measurable_iff_comap_le] + fun_prop + simp only [hn, ↓reduceIte, sup_le_iff] + constructor + · exact (IT.filtration α R).le _ + · rw [← measurable_iff_comap_le] + fun_prop + +lemma filtrationAction_zero_eq_comap : + filtrationAction α R 0 = MeasurableSpace.comap (action 0) inferInstance := by + simp [filtrationAction] + +lemma filtrationAction_eq_comap (n : ℕ) (hn : n ≠ 0) : + filtrationAction α R n = + MeasurableSpace.comap (fun ω ↦ (hist (n - 1) ω, action n ω)) inferInstance := by + simp only [filtrationAction, filtration_eq_comap, ← MeasurableSpace.comap_prodMk, hn, ↓reduceIte] + rfl + +lemma filtration_le_filtrationAction_add_one (n : ℕ) : + IT.filtration α R n ≤ filtrationAction α R (n + 1) := le_sup_of_le_left le_rfl + +lemma filtration_le_filtrationAction {m n : ℕ} (h : n < m) : + IT.filtration α R n ≤ filtrationAction α R m := by + have h' : n + 1 ≤ m := by grind + exact (filtration_le_filtrationAction_add_one n).trans ((filtrationAction α R).mono h') + +lemma filtrationAction_le_filtration_self (n : ℕ) : + filtrationAction α R n ≤ IT.filtration α R n := by + by_cases hn : n = 0 + · simp only [hn, filtrationAction_zero_eq_comap] + rw [← measurable_iff_comap_le] + exact measurable_action_filtration 0 + simp only [filtrationAction, hn, ↓reduceIte, sup_le_iff] + constructor + · exact (IT.filtration α R).mono (by grind) + · rw [← measurable_iff_comap_le] + exact measurable_action_filtration _ + +lemma filtrationAction_le_filtration {m n : ℕ} (h : m ≤ n) : + filtrationAction α R m ≤ IT.filtration α R n := + (filtrationAction_le_filtration_self m).trans ((IT.filtration α R).mono h) + +lemma measurable_action_filtrationAction (n : ℕ) : + Measurable[filtrationAction α R n] (action n) := by + simp only [filtrationAction] + rw [measurable_iff_comap_le] + split_ifs with hn + · simp [hn] + · exact le_sup_of_le_right le_rfl + +end FiltrationAction + +section Laws + +lemma hasLaw_step_zero (alg : Algorithm α R) (env : Environment α R) : + HasLaw (step 0) (alg.p0 ⊗ₘ env.ν0) (trajMeasure alg env) where + aemeasurable := Measurable.aemeasurable (by fun_prop) + map_eq := by + unfold step + rw [← coe_default_Iic_zero] + simp only [trajMeasure, Kernel.trajMeasure] + rw [← Measure.deterministic_comp_eq_map (by fun_prop), Measure.comp_assoc, + Kernel.deterministic_comp_eq_map, Kernel.traj_zero_map_eval_zero, + Measure.deterministic_comp_eq_map, Measure.map_map (by fun_prop) (by fun_prop)] + exact Measure.map_id + +lemma hasLaw_action_zero (alg : Algorithm α R) (env : Environment α R) : + HasLaw (action 0) alg.p0 (trajMeasure alg env) where + map_eq := by + rw [← fst_comp_step, ← Measure.map_map (by fun_prop) (by fun_prop), + (hasLaw_step_zero alg env).map_eq, ← Measure.fst, Measure.fst_compProd] + +variable [StandardBorelSpace R] [Nonempty R] + +lemma condDistrib_reward_zero (alg : Algorithm α R) (env : Environment α R) : + condDistrib (reward 0) (action 0) (trajMeasure alg env) + =ᵐ[(trajMeasure alg env).map (action 0)] env.ν0 := by + have h_step := (hasLaw_step_zero alg env).map_eq + have h_action := (hasLaw_action_zero alg env).map_eq + rwa [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop), h_action] + +variable [StandardBorelSpace α] [Nonempty α] + +lemma condDistrib_step (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : + condDistrib (step (n + 1)) (hist n) (trajMeasure alg env) + =ᵐ[(trajMeasure alg env).map (hist n)] stepKernel alg env n := + Kernel.condDistrib_trajMeasure + +lemma condDistrib_action (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : + condDistrib (action (n + 1)) (hist n) (trajMeasure alg env) + =ᵐ[(trajMeasure alg env).map (hist n)] alg.policy n := by + rw [← fst_comp_step] + refine (condDistrib_comp _ (by fun_prop) (by fun_prop)).trans ?_ + filter_upwards [condDistrib_step alg env n] with h h_eq + rw [Kernel.map_apply _ (by fun_prop), h_eq, ← Kernel.map_apply _ (by fun_prop), ← Kernel.fst_eq, + fst_stepKernel] + +lemma condDistrib_reward (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : + condDistrib (reward (n + 1)) (fun ω ↦ (hist n ω, action (n + 1) ω)) (trajMeasure alg env) + =ᵐ[(trajMeasure alg env).map (fun ω ↦ (hist n ω, action (n + 1) ω))] env.feedback n := by + have h_step := condDistrib_step alg env n + have h_action := condDistrib_action alg env n + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h_step h_action ⊢ + rw [h_action, ← Measure.compProd_assoc, ← stepKernel, ← h_step, + Measure.map_map (by fun_prop) (by fun_prop)] + rfl + +lemma isAlgEnvSeq_trajMeasure (alg : Algorithm α R) (env : Environment α R) : + IsAlgEnvSeq action reward alg env (trajMeasure alg env) where + hasLaw_action_zero := hasLaw_action_zero alg env + hasCondDistrib_reward_zero := ⟨by fun_prop, by fun_prop, condDistrib_reward_zero alg env⟩ + hasCondDistrib_action n := ⟨by fun_prop, by fun_prop, condDistrib_action alg env n⟩ + hasCondDistrib_reward n := ⟨by fun_prop, by fun_prop, condDistrib_reward alg env n⟩ + +end Laws + +end IT + +end Learning diff --git a/LeanBandits/SequentialLearning/StationaryEnv.lean b/LeanBandits/SequentialLearning/StationaryEnv.lean index 3560acf8..b5eaffa3 100644 --- a/LeanBandits/SequentialLearning/StationaryEnv.lean +++ b/LeanBandits/SequentialLearning/StationaryEnv.lean @@ -3,7 +3,7 @@ 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, Paulo Rauber -/ -import LeanBandits.SequentialLearning.Algorithm +import LeanBandits.SequentialLearning.IonescuTulceaSpace /-! # Stationary environments From 812c0d1dd5001286b56c4de32e9935ba63a7c5c5 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Tue, 6 Jan 2026 17:30:11 +0100 Subject: [PATCH 06/30] work on sumRewards --- LeanBandits/Bandit/SumRewards.lean | 263 ++++++++++++++++++++++++++--- 1 file changed, 239 insertions(+), 24 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index c47a4bbb..50167be0 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -11,6 +11,34 @@ import LeanBandits.Bandit.Regret open MeasureTheory ProbabilityTheory Finset Learning open scoped ENNReal NNReal +lemma measurable_sum_range_of_le {α : Type*} {mα : MeasurableSpace α} + {f : ℕ → α → ℝ} {g : α → ℕ} {n : ℕ} (hg_le : ∀ a, g a ≤ n) (hf : ∀ i, Measurable (f i)) + (hg : Measurable g) : + Measurable (fun a ↦ ∑ i ∈ range (g a), f i a) := by + have h_eq : (fun a ↦ ∑ i ∈ range (g a), f i a) + = fun a ↦ ∑ i ∈ range (n + 1), if g a = i then ∑ j ∈ range i, f j a else 0 := by + ext ω + rw [sum_ite_eq_of_mem] + grind + rw [h_eq] + refine measurable_sum _ fun n hn ↦ ?_ + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + +lemma measurable_sum_Icc_of_le {α : Type*} {mα : MeasurableSpace α} + {f : ℕ → α → ℝ} {g : α → ℕ} {n : ℕ} (hg_le : ∀ a, g a ≤ n) (hf : ∀ i, Measurable (f i)) + (hg : Measurable g) : + Measurable (fun a ↦ ∑ i ∈ Icc 1 (g a), f i a) := by + have h_eq : (fun a ↦ ∑ i ∈ Icc 1 (g a), f i a) + = fun a ↦ ∑ i ∈ range (n + 1), if g a = i then ∑ j ∈ Icc 1 i, f j a else 0 := by + ext ω + rw [sum_ite_eq_of_mem] + grind + rw [h_eq] + refine measurable_sum _ fun n hn ↦ ?_ + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + namespace Bandits namespace ArrayModel @@ -20,17 +48,28 @@ variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [StandardBorel local notation "A" => action alg local notation "R" => reward alg +local notation "𝔓" => arrayMeasure ν -lemma identDistrib_sum_Icc_rewardByCount_pullCount' [Countable α] (a : α) (n : ℕ) : - IdentDistrib (fun ω ↦ ∑ i ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a i ω) - (fun ω ↦ ∑ i ∈ Icc 1 (pullCount A a n ω), ω.2 (i - 1) a) - ((arrayMeasure ν).prod (Bandit.streamMeasure ν)) (arrayMeasure ν) where - aemeasurable_fst := sorry -- the issue is the random bound - aemeasurable_snd := sorry +lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount' [Countable α] (n : ℕ) : + IdentDistrib (fun ω a ↦ (pullCount A a n ω.1, + ∑ i ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a i ω)) + (fun ω a ↦ (pullCount A a n ω, ∑ i ∈ Icc 1 (pullCount A a n ω), ω.2 (i - 1) a)) + ((𝔓).prod (Bandit.streamMeasure ν)) 𝔓 where + aemeasurable_fst := by + refine Measurable.aemeasurable ?_ + rw [measurable_pi_iff] + refine fun a ↦ Measurable.prod (by fun_prop) ?_ + exact measurable_sum_Icc_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) + aemeasurable_snd := by + refine Measurable.aemeasurable ?_ + rw [measurable_pi_iff] + refine fun a ↦ Measurable.prod (by fun_prop) ?_ + exact measurable_sum_Icc_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) map_eq := by by_cases hn : n = 0 · simp [hn] - have h_eq (i : ℕ) (ω : probSpace α ℝ × (ℕ → α → ℝ)) (hi : i ∈ Icc 1 (pullCount A a n ω.1)) : + have h_eq (a : α) (i : ℕ) (ω : probSpace α ℝ × (ℕ → α → ℝ)) + (hi : i ∈ Icc 1 (pullCount A a n ω.1)) : rewardByCount A R a i ω = ω.1.2 (i - 1) a := by rw [rewardByCount_of_stepsUntil_ne_top] · simp only [reward_eq] @@ -48,36 +87,212 @@ lemma identDistrib_sum_Icc_rewardByCount_pullCount' [Countable α] (a : α) (n : simp only [mem_Icc] at hi refine hi.2.trans ?_ exact pullCount_mono _ (by grind) _ - have h_sum_eq (ω : probSpace α ℝ × (ℕ → α → ℝ)) : + have h_sum_eq (a : α) (ω : probSpace α ℝ × (ℕ → α → ℝ)) : ∑ i ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a i ω = ∑ i ∈ Icc 1 (pullCount A a n ω.1), ω.1.2 (i - 1) a := - Finset.sum_congr rfl fun i hi ↦ h_eq i ω hi + Finset.sum_congr rfl fun i hi ↦ h_eq a i ω hi simp_rw [h_sum_eq] - conv_rhs => rw [← Measure.fst_prod (μ := arrayMeasure ν) (ν := Bandit.streamMeasure ν), + conv_rhs => rw [← Measure.fst_prod (μ := 𝔓) (ν := Bandit.streamMeasure ν), Measure.fst] rw [AEMeasurable.map_map_of_aemeasurable _ (by fun_prop)] · rfl simp only [Measure.map_fst_prod, measure_univ, one_smul] - sorry + refine Measurable.aemeasurable ?_ + rw [measurable_pi_iff] + refine fun a ↦ Measurable.prod (by fun_prop) ?_ + exact measurable_sum_Icc_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) -lemma identDistrib_sum_Icc_rewardByCount_pullCount [Countable α] (a : α) (n : ℕ) : - IdentDistrib (fun ω ↦ ∑ i ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a i ω) - (fun ω ↦ ∑ i ∈ range (pullCount A a n ω), ω.2 i a) - ((arrayMeasure ν).prod (Bandit.streamMeasure ν)) (arrayMeasure ν) := by - convert identDistrib_sum_Icc_rewardByCount_pullCount' a n using 2 with ω - swap; · infer_instance +lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount [Countable α] (n : ℕ) : + IdentDistrib (fun ω a ↦ (pullCount A a n ω.1, + ∑ i ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a i ω)) + (fun ω a ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) + ((𝔓).prod (Bandit.streamMeasure ν)) 𝔓 := by + convert identDistrib_pullCount_prod_sum_Icc_rewardByCount' n using 2 with ω + rotate_left + · infer_instance + · infer_instance + ext a : 1 + congr 1 sorry -lemma identDistrib_sumRewards_pullCount [Countable α] (a : α) (n : ℕ) : - IdentDistrib (sumRewards A R a n) (fun ω ↦ ∑ i ∈ range (pullCount A a n ω), ω.2 i a) - (arrayMeasure ν) (arrayMeasure ν) := by - suffices IdentDistrib (fun ω ↦ sumRewards A R a n ω.1) - (fun ω ↦ ∑ i ∈ range (pullCount A a n ω), ω.2 i a) - ((arrayMeasure ν).prod (Bandit.streamMeasure ν)) (arrayMeasure ν) by +lemma identDistrib_pullCount_prod_sumRewards [Countable α] (n : ℕ) : + IdentDistrib (fun ω a ↦ (pullCount A a n ω, sumRewards A R a n ω)) + (fun ω a ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) 𝔓 𝔓 := by + suffices IdentDistrib (fun ω a ↦ (pullCount A a n ω.1, sumRewards A R a n ω.1)) + (fun ω a ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) + ((𝔓).prod (Bandit.streamMeasure ν)) 𝔓 by sorry simp_rw [← sum_rewardByCount_eq_sumRewards] - exact identDistrib_sum_Icc_rewardByCount_pullCount a n + exact identDistrib_pullCount_prod_sum_Icc_rewardByCount n + +lemma identDistrib_pullCount_prod_sumRewards_arm [Countable α] (a : α) (n : ℕ) : + IdentDistrib (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω)) + (fun ω ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) 𝔓 𝔓 := by + have h1 : (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω)) = + (fun p ↦ p a) ∘ (fun ω a ↦ (pullCount A a n ω, sumRewards A R a n ω)) := rfl + have h2 : (fun ω ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) = + (fun p ↦ p a) ∘ + (fun ω a ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) := rfl + rw [h1, h2] + refine (identDistrib_pullCount_prod_sumRewards n).comp ?_ + fun_prop + +lemma identDistrib_sumRewards [Countable α] (n : ℕ) : + IdentDistrib (fun ω a ↦ sumRewards A R a n ω) + (fun ω a ↦ ∑ i ∈ range (pullCount A a n ω), ω.2 i a) 𝔓 𝔓 := by + have h_ident := identDistrib_pullCount_prod_sumRewards (ν := ν) (alg := alg) n + exact h_ident.comp (u := fun p a ↦ (p a).2) (by fun_prop) + +lemma identDistrib_sumRewards_arm [Countable α] (a : α) (n : ℕ) : + IdentDistrib (sumRewards A R a n) + (fun ω ↦ ∑ i ∈ range (pullCount A a n ω), ω.2 i a) 𝔓 𝔓 := by + have h1 : sumRewards A R a n = (fun p ↦ p a) ∘ (fun ω a ↦ sumRewards A R a n ω) := rfl + have h2 : (fun ω ↦ ∑ i ∈ range (pullCount A a n ω), ω.2 i a) = + (fun p ↦ p a) ∘ (fun ω a ↦ ∑ i ∈ range (pullCount A a n ω), ω.2 i a) := rfl + rw [h1, h2] + refine (identDistrib_sumRewards n).comp ?_ + fun_prop + +lemma todo'' [Countable α] (a : α) (n : ℕ) + {s : Set ℕ} [DecidablePred (· ∈ s)] (hs : MeasurableSet s) {B : Set ℝ} (hB : MeasurableSet B) : + 𝔓 {ω | pullCount A a n ω ∈ s ∧ sumRewards A R a n ω ∈ B} ≤ + ∑ k ∈ (range (n + 1)).filter (· ∈ s), + Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := by + have h_ident := identDistrib_pullCount_prod_sumRewards_arm a n (ν := ν) (alg := alg) + have : 𝔓 {ω | pullCount A a n ω ∈ s ∧ sumRewards A R a n ω ∈ B} = + (𝔓).map (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω)) (s ×ˢ B) := by + rw [Measure.map_apply (by fun_prop) (hs.prod hB), Set.mk_preimage_prod] + rfl + rw [this, h_ident.map_eq, Measure.map_apply ?_ (hs.prod hB)] + swap + · refine Measurable.prod (by fun_prop) ?_ + exact measurable_sum_range_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) + rw [Set.mk_preimage_prod] + calc 𝔓 {ω | pullCount A a n ω ∈ s ∧ ∑ i ∈ range (pullCount A a n ω), ω.2 i a ∈ B} + _ ≤ 𝔓 {ω | ∃ k ≤ n, k ∈ s ∧ ∑ i ∈ range k, ω.2 i a ∈ B} := by + refine measure_mono fun ω hω ↦ ?_ + simp only [Set.mem_setOf_eq] at hω ⊢ + exact ⟨pullCount A a n ω, pullCount_le _ _ _, hω⟩ + _ = 𝔓 (⋃ k ∈ (range (n + 1)).filter (· ∈ s), {ω | ∑ i ∈ range k, ω.2 i a ∈ B}) := by + congr 1 + ext ω + simp + grind + _ ≤ ∑ k ∈ (range (n + 1)).filter (· ∈ s), 𝔓 {ω | ∑ i ∈ range k, ω.2 i a ∈ B} := + measure_biUnion_finset_le _ _ + _ = ∑ k ∈ (range (n + 1)).filter (· ∈ s), + (Bandit.streamMeasure ν) {ω | ∑ i ∈ range k, ω i a ∈ B} := by + congr with k + sorry end ArrayModel +variable {α Ω Ω' : Type*} [DecidableEq α] {mα : MeasurableSpace α} {mΩ : MeasurableSpace Ω} + {mΩ' : MeasurableSpace Ω'} + {P : Measure Ω} [IsProbabilityMeasure P] {P' : Measure Ω'} [IsProbabilityMeasure P'] + {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] + {A : ℕ → Ω → α} {R : ℕ → Ω → ℝ} {A₂ : ℕ → Ω' → α} {R₂ : ℕ → Ω' → ℝ} + {ω : Ω} {m n t : ℕ} {a : α} + +variable [StandardBorelSpace α] [Nonempty α] + +omit [Nonempty α] in +lemma sumRewards_eq_comp : + sumRewards A R a n = + (fun p ↦ ∑ i ∈ range n, if (p i).1 = a then (p i).2 else 0) ∘ (fun ω n ↦ (A n ω, R n ω)) := by + ext + simp [sumRewards] + +omit [Nonempty α] in +lemma pullCount_eq_comp : + pullCount A a n = + (fun p ↦ ∑ i ∈ range n, if (p i).1 = a then 1 else 0) ∘ (fun ω n ↦ (A n ω, R n ω)) := by + ext + simp [pullCount] + +lemma _root_.Learning.IsAlgEnvSeq.law_sumRewards_unique + (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + (h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') : + P.map (sumRewards A R a n) = P'.map (sumRewards A₂ R₂ a n) := by + have hA := h1.measurable_A + have hR := h1.measurable_R + have hA2 := h2.measurable_A + have hR2 := h2.measurable_R + have h_unique := isAlgEnvSeq_unique h1 h2 + rw [sumRewards_eq_comp, sumRewards_eq_comp, ← Measure.map_map, h_unique, Measure.map_map, + ← sumRewards_eq_comp] + · refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + · rw [measurable_pi_iff] + exact fun n ↦ Measurable.prodMk (hA2 n) (hR2 n) + · refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + · rw [measurable_pi_iff] + exact fun n ↦ Measurable.prodMk (hA n) (hR n) + +lemma _root_.Learning.IsAlgEnvSeq.law_pullCount_sumRewards_unique + (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + (h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') : + P.map (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω)) = + P'.map (fun ω ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) := by + have hA := h1.measurable_A + have hR := h1.measurable_R + have hA2 := h2.measurable_A + have hR2 := h2.measurable_R + have h_unique := isAlgEnvSeq_unique h1 h2 + let f := fun p : ℕ → α × ℝ ↦ (∑ i ∈ range n, if (p i).1 = a then 1 else 0, + ∑ i ∈ range n, if (p i).1 = a then (p i).2 else 0) + have hf : Measurable f := by + refine Measurable.prod ?_ ?_ + · simp only [f] + refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + · simp only [f] + refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + have h_eq_comp : (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω)) + = f ∘ (fun ω n ↦ (A n ω, R n ω)) := by + ext ω : 1 + rw [pullCount_eq_comp (R := R), sumRewards_eq_comp] + grind + have h_eq_comp2 : (fun ω ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) + = f ∘ (fun ω n ↦ (A₂ n ω, R₂ n ω)) := by + ext ω : 1 + rw [pullCount_eq_comp (R := R₂), sumRewards_eq_comp] + grind + rw [h_eq_comp, h_eq_comp2, ← Measure.map_map hf, h_unique, Measure.map_map hf, + ← h_eq_comp2] + · rw [measurable_pi_iff] + exact fun n ↦ Measurable.prodMk (hA2 n) (hR2 n) + · rw [measurable_pi_iff] + exact fun n ↦ Measurable.prodMk (hA n) (hR n) + +lemma todo2 [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + {s : Set ℕ} [DecidablePred (· ∈ s)] (hs : MeasurableSet s) {B : Set ℝ} (hB : MeasurableSet B) : + P {ω | pullCount A a n ω ∈ s ∧ sumRewards A R a n ω ∈ B} ≤ + ∑ k ∈ (range (n + 1)).filter (· ∈ s), + Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := by + have hA := h.measurable_A + have hR := h.measurable_R + calc P {ω | pullCount A a n ω ∈ s ∧ sumRewards A R a n ω ∈ B} + _ = (P.map (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω))) (s ×ˢ B) := by + rw [Measure.map_apply (by fun_prop) (hs.prod hB)]; rfl + _ = ((ArrayModel.arrayMeasure ν).map + (fun ω ↦ (pullCount (ArrayModel.action alg) a n ω, + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a n ω))) (s ×ˢ B) := by + rw [h.law_pullCount_sumRewards_unique (ArrayModel.isAlgEnvSeq_arrayMeasure alg ν)] + _ = (ArrayModel.arrayMeasure ν) {ω | pullCount (ArrayModel.action alg) a n ω ∈ s ∧ + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a n ω ∈ B} := by + rw [Measure.map_apply (by fun_prop) (hs.prod hB)]; rfl + _ ≤ ∑ k ∈ (range (n + 1)).filter (· ∈ s), Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := + ArrayModel.todo'' a n hs hB + +lemma todo [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + {B : Set ℝ} (hB : MeasurableSet B) : + P (sumRewards A R a n ⁻¹' B) ≤ + ∑ k ∈ range (n + 1), Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := by + classical + have h_le := todo2 h .univ hB (a := a) (n := n) + simpa using h_le + end Bandits From 1b146e79a4be3c5b097118c831430270a363cabe Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 7 Jan 2026 14:10:38 +0100 Subject: [PATCH 07/30] update --- LeanBandits/Bandit/RewardByCountMeasure.lean | 9 - LeanBandits/Bandit/SumRewards.lean | 404 ++++++++++++++++-- LeanBandits/BanditAlgorithms/ETC.lean | 191 ++------- LeanBandits/BanditAlgorithms/UCB.lean | 266 +++--------- .../SequentialLearning/FiniteActions.lean | 43 ++ 5 files changed, 501 insertions(+), 412 deletions(-) diff --git a/LeanBandits/Bandit/RewardByCountMeasure.lean b/LeanBandits/Bandit/RewardByCountMeasure.lean index 6a2e7083..820f6a37 100644 --- a/LeanBandits/Bandit/RewardByCountMeasure.lean +++ b/LeanBandits/Bandit/RewardByCountMeasure.lean @@ -59,15 +59,6 @@ variable {α Ω : Type*} {mα : MeasurableSpace α} {mΩ : MeasurableSpace Ω} [ {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] {h_inter : IsAlgEnvSeq A R alg (stationaryEnv ν) P} -omit [StandardBorelSpace α] [Nonempty α] in -lemma integrable_pullCount [MeasurableSingletonClass α] - (hA : ∀ n, Measurable (A n)) (a : α) (n : ℕ) : - Integrable (fun ω ↦ (pullCount A a n ω : ℝ)) P := by - refine integrable_of_le_of_le (g₁ := 0) (g₂ := fun _ ↦ n) (by fun_prop) - (ae_of_all _ fun ω ↦ by simp) (ae_of_all _ fun ω ↦ ?_) (integrable_const _) (integrable_const _) - simp only [Nat.cast_le] - exact pullCount_le a n ω - local notation "𝔓'" => P.prod (Bandit.streamMeasure ν) omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 50167be0..4b382d95 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -4,6 +4,7 @@ Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne -/ import LeanBandits.Bandit.Regret +import LeanBandits.ForMathlib.SubGaussian /-! # Law of the sum of rewards -/ @@ -50,7 +51,7 @@ local notation "A" => action alg local notation "R" => reward alg local notation "𝔓" => arrayMeasure ν -lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount' [Countable α] (n : ℕ) : +lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount' (n : ℕ) : IdentDistrib (fun ω a ↦ (pullCount A a n ω.1, ∑ i ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a i ω)) (fun ω a ↦ (pullCount A a n ω, ∑ i ∈ Icc 1 (pullCount A a n ω), ω.2 (i - 1) a)) @@ -102,7 +103,7 @@ lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount' [Countable α] (n : ℕ refine fun a ↦ Measurable.prod (by fun_prop) ?_ exact measurable_sum_Icc_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) -lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount [Countable α] (n : ℕ) : +lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount (n : ℕ) : IdentDistrib (fun ω a ↦ (pullCount A a n ω.1, ∑ i ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a i ω)) (fun ω a ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) @@ -110,22 +111,44 @@ lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount [Countable α] (n : ℕ) convert identDistrib_pullCount_prod_sum_Icc_rewardByCount' n using 2 with ω rotate_left · infer_instance - · infer_instance ext a : 1 congr 1 - sorry + let e : Icc 1 (pullCount A a n ω) ≃ range (pullCount A a n ω) := + { toFun x := ⟨x - 1, by have h := x.2; simp only [mem_Icc] at h; simp; grind⟩ + invFun x := ⟨x + 1, by + have h := x.2 + simp only [mem_Icc, le_add_iff_nonneg_left, zero_le, true_and, ge_iff_le] + simp only [mem_range] at h + grind⟩ + left_inv x := by have h := x.2; simp only [mem_Icc] at h; grind + right_inv x := by have h := x.2; grind } + rw [← sum_coe_sort (Icc 1 (pullCount A a n ω)), ← sum_coe_sort (range (pullCount A a n ω)), + sum_equiv e] + · simp + · simp [e] -lemma identDistrib_pullCount_prod_sumRewards [Countable α] (n : ℕ) : +lemma identDistrib_pullCount_prod_sumRewards (n : ℕ) : IdentDistrib (fun ω a ↦ (pullCount A a n ω, sumRewards A R a n ω)) (fun ω a ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) 𝔓 𝔓 := by suffices IdentDistrib (fun ω a ↦ (pullCount A a n ω.1, sumRewards A R a n ω.1)) (fun ω a ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) ((𝔓).prod (Bandit.streamMeasure ν)) 𝔓 by - sorry + -- todo: missing lemma about IdentDistrib? + constructor + · refine Measurable.aemeasurable ?_ + fun_prop + · refine Measurable.aemeasurable ?_ + rw [measurable_pi_iff] + refine fun a ↦ Measurable.prod (by fun_prop) ?_ + exact measurable_sum_range_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) + have h_eq := this.map_eq + nth_rw 1 [← Measure.fst_prod (μ := 𝔓) (ν := Bandit.streamMeasure ν), Measure.fst, + Measure.map_map (by fun_prop) (by fun_prop)] + exact h_eq simp_rw [← sum_rewardByCount_eq_sumRewards] exact identDistrib_pullCount_prod_sum_Icc_rewardByCount n -lemma identDistrib_pullCount_prod_sumRewards_arm [Countable α] (a : α) (n : ℕ) : +lemma identDistrib_pullCount_prod_sumRewards_arm (a : α) (n : ℕ) : IdentDistrib (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω)) (fun ω ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) 𝔓 𝔓 := by have h1 : (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω)) = @@ -137,13 +160,22 @@ lemma identDistrib_pullCount_prod_sumRewards_arm [Countable α] (a : α) (n : refine (identDistrib_pullCount_prod_sumRewards n).comp ?_ fun_prop -lemma identDistrib_sumRewards [Countable α] (n : ℕ) : +lemma identDistrib_pullCount_prod_sumRewards_two_arms (a b : α) (n : ℕ) : + IdentDistrib (fun ω ↦ (pullCount A a n ω, pullCount A b n ω, + sumRewards A R a n ω, sumRewards A R b n ω)) + (fun ω ↦ (pullCount A a n ω, pullCount A b n ω, + ∑ i ∈ range (pullCount A a n ω), ω.2 i a, + ∑ i ∈ range (pullCount A b n ω), ω.2 i b)) 𝔓 𝔓 := by + have h_ident := identDistrib_pullCount_prod_sumRewards (ν := ν) (alg := alg) n + exact h_ident.comp (u := fun p ↦ ((p a).1, (p b).1, (p a).2, (p b).2)) (by fun_prop) + +lemma identDistrib_sumRewards (n : ℕ) : IdentDistrib (fun ω a ↦ sumRewards A R a n ω) (fun ω a ↦ ∑ i ∈ range (pullCount A a n ω), ω.2 i a) 𝔓 𝔓 := by have h_ident := identDistrib_pullCount_prod_sumRewards (ν := ν) (alg := alg) n exact h_ident.comp (u := fun p a ↦ (p a).2) (by fun_prop) -lemma identDistrib_sumRewards_arm [Countable α] (a : α) (n : ℕ) : +lemma identDistrib_sumRewards_arm (a : α) (n : ℕ) : IdentDistrib (sumRewards A R a n) (fun ω ↦ ∑ i ∈ range (pullCount A a n ω), ω.2 i a) 𝔓 𝔓 := by have h1 : sumRewards A R a n = (fun p ↦ p a) ∘ (fun ω a ↦ sumRewards A R a n ω) := rfl @@ -153,37 +185,114 @@ lemma identDistrib_sumRewards_arm [Countable α] (a : α) (n : ℕ) : refine (identDistrib_sumRewards n).comp ?_ fun_prop -lemma todo'' [Countable α] (a : α) (n : ℕ) - {s : Set ℕ} [DecidablePred (· ∈ s)] (hs : MeasurableSet s) {B : Set ℝ} (hB : MeasurableSet B) : - 𝔓 {ω | pullCount A a n ω ∈ s ∧ sumRewards A R a n ω ∈ B} ≤ - ∑ k ∈ (range (n + 1)).filter (· ∈ s), - Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := by +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in +lemma identDistrib_sum_range_snd (a : α) (k : ℕ) : + IdentDistrib (fun ω ↦ ∑ i ∈ range k, ω.2 i a) (fun ω ↦ ∑ i ∈ range k, ω i a) + 𝔓 (Bandit.streamMeasure ν) where + aemeasurable_fst := by fun_prop + aemeasurable_snd := (measurable_sum _ fun i _ ↦ by fun_prop).aemeasurable + map_eq := by + rw [← Measure.snd_prod (μ := (Measure.infinitePi fun (_ : ℕ) ↦ (volume : Measure unitInterval))) + (ν := Bandit.streamMeasure ν), Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)] + rfl + +lemma prob_pullCount_prod_sumRewards_mem_le (a : α) (n : ℕ) + {s : Set (ℕ × ℝ)} [DecidablePred (· ∈ Prod.fst '' s)] (hs : MeasurableSet s) : + 𝔓 {ω | (pullCount A a n ω, sumRewards A R a n ω) ∈ s} ≤ + ∑ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), + Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := by have h_ident := identDistrib_pullCount_prod_sumRewards_arm a n (ν := ν) (alg := alg) - have : 𝔓 {ω | pullCount A a n ω ∈ s ∧ sumRewards A R a n ω ∈ B} = - (𝔓).map (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω)) (s ×ˢ B) := by - rw [Measure.map_apply (by fun_prop) (hs.prod hB), Set.mk_preimage_prod] + have : 𝔓 {ω | (pullCount A a n ω, sumRewards A R a n ω) ∈ s} = + (𝔓).map (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω)) s := by + rw [Measure.map_apply (by fun_prop) hs] rfl - rw [this, h_ident.map_eq, Measure.map_apply ?_ (hs.prod hB)] + rw [this, h_ident.map_eq, Measure.map_apply ?_ hs] swap · refine Measurable.prod (by fun_prop) ?_ exact measurable_sum_range_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) - rw [Set.mk_preimage_prod] - calc 𝔓 {ω | pullCount A a n ω ∈ s ∧ ∑ i ∈ range (pullCount A a n ω), ω.2 i a ∈ B} - _ ≤ 𝔓 {ω | ∃ k ≤ n, k ∈ s ∧ ∑ i ∈ range k, ω.2 i a ∈ B} := by + calc 𝔓 ((fun ω ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) ⁻¹' s) + _ ≤ 𝔓 {ω | ∃ k ≤ n, (k, ∑ i ∈ range k, ω.2 i a) ∈ s} := by refine measure_mono fun ω hω ↦ ?_ simp only [Set.mem_setOf_eq] at hω ⊢ exact ⟨pullCount A a n ω, pullCount_le _ _ _, hω⟩ - _ = 𝔓 (⋃ k ∈ (range (n + 1)).filter (· ∈ s), {ω | ∑ i ∈ range k, ω.2 i a ∈ B}) := by - congr 1 - ext ω - simp - grind - _ ≤ ∑ k ∈ (range (n + 1)).filter (· ∈ s), 𝔓 {ω | ∑ i ∈ range k, ω.2 i a ∈ B} := - measure_biUnion_finset_le _ _ - _ = ∑ k ∈ (range (n + 1)).filter (· ∈ s), - (Bandit.streamMeasure ν) {ω | ∑ i ∈ range k, ω i a ∈ B} := by + _ = 𝔓 (⋃ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), + {ω | (k, ∑ i ∈ range k, ω.2 i a) ∈ s}) := by congr 1; ext; simp; grind + _ ≤ ∑ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), + 𝔓 {ω | ∑ i ∈ range k, ω.2 i a ∈ Prod.mk k ⁻¹' s} := measure_biUnion_finset_le _ _ + _ = ∑ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), + (Bandit.streamMeasure ν) {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := by congr with k - sorry + have : (𝔓).map (fun ω ↦ ∑ i ∈ range k, ω.2 i a) = + (Bandit.streamMeasure ν).map (fun ω ↦ ∑ i ∈ range k, ω i a) := + (identDistrib_sum_range_snd a k).map_eq + rw [Measure.ext_iff] at this + specialize this (Prod.mk k ⁻¹' s) (hs.preimage (by fun_prop)) + rwa [Measure.map_apply (by fun_prop) (hs.preimage (by fun_prop)), + Measure.map_apply (by fun_prop) (hs.preimage (by fun_prop))] at this + +lemma prob_pullCount_mem_and_sumRewards_mem_le (a : α) (n : ℕ) + {s : Set ℕ} [DecidablePred (· ∈ s)] (hs : MeasurableSet s) {B : Set ℝ} (hB : MeasurableSet B) : + 𝔓 {ω | pullCount A a n ω ∈ s ∧ sumRewards A R a n ω ∈ B} ≤ + ∑ k ∈ (range (n + 1)).filter (· ∈ s), + Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := by + classical + rcases Set.eq_empty_or_nonempty B with h_empty | h_nonempty + · simp [h_empty] + convert prob_pullCount_prod_sumRewards_mem_le a n (hs.prod hB) (ν := ν) (alg := alg) with _ _ k hk + · ext n + have : ∃ x, x ∈ B := h_nonempty + simp [this] + · ext x + simp only [Set.mem_image, Set.mem_prod, Prod.exists, exists_and_right, exists_and_left, + exists_eq_right, mem_filter, mem_range] at hk + simp [hk.2.1] + +lemma prob_sumRewards_le_sumRewards_le [Fintype α] (a : α) (n m₁ m₂ : ℕ) : + (𝔓) {ω | pullCount A (bestArm ν) n ω = m₁ ∧ pullCount A a n ω = m₂ ∧ + sumRewards A R (bestArm ν) n ω ≤ sumRewards A R a n ω} ≤ + Bandit.streamMeasure ν + {ω | ∑ i ∈ range m₁, ω i (bestArm ν) ≤ ∑ i ∈ range m₂, ω i a} := by + have h_ident := identDistrib_pullCount_prod_sumRewards_two_arms (bestArm ν) a n + (ν := ν) (alg := alg) + let s := {p : ℕ × ℕ × ℝ × ℝ | p.1 = m₁ ∧ p.2.1 = m₂ ∧ p.2.2.1 ≤ p.2.2.2} + have hs : MeasurableSet s := by simp only [measurableSet_setOf, s]; fun_prop + calc 𝔓 {ω | pullCount A (bestArm ν) n ω = m₁ ∧ pullCount A a n ω = m₂ ∧ + sumRewards A R (bestArm ν) n ω ≤ sumRewards A R a n ω} + _ = 𝔓 ((fun ω ↦ (pullCount A (bestArm ν) n ω, pullCount A a n ω, + sumRewards A R (bestArm ν) n ω, sumRewards A R a n ω)) ⁻¹' + {p | p.1 = m₁ ∧ p.2.1 = m₂ ∧ p.2.2.1 ≤ p.2.2.2}) := rfl + _ = 𝔓 ((fun ω ↦ (pullCount A (bestArm ν) n ω, pullCount A a n ω, + ∑ i ∈ range (pullCount A (bestArm ν) n ω), ω.2 i (bestArm ν), + ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) ⁻¹' + {p | p.1 = m₁ ∧ p.2.1 = m₂ ∧ p.2.2.1 ≤ p.2.2.2}) := by + rw [← Measure.map_apply (by fun_prop) hs, h_ident.map_eq, + Measure.map_apply _ hs] + refine Measurable.prod (by fun_prop) (Measurable.prod (by fun_prop) ?_) + refine Measurable.prod ?_ ?_ + · exact measurable_sum_range_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) + · exact measurable_sum_range_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) + _ ≤ 𝔓 ((fun ω ↦ (∑ i ∈ range m₁, ω.2 i (bestArm ν), ∑ i ∈ range m₂, ω.2 i a)) ⁻¹' + {p | p.1 ≤ p.2}) := by + refine measure_mono fun ω hω ↦ ?_ + simp only [Set.preimage_setOf_eq, Set.mem_setOf_eq] at hω ⊢ + grind + _ = Bandit.streamMeasure ν + {ω | ∑ i ∈ range m₁, ω i (bestArm ν) ≤ ∑ i ∈ range m₂, ω i a} := by + rw [← Measure.snd_prod (μ := (Measure.infinitePi fun (_ : ℕ) ↦ (volume : Measure unitInterval))) + (ν := Bandit.streamMeasure ν), Measure.snd, Measure.map_apply (by fun_prop)] + · rfl + simp only [measurableSet_setOf] + fun_prop + +lemma probReal_sumRewards_le_sumRewards_le [Fintype α] (a : α) (n m₁ m₂ : ℕ) : + (𝔓).real {ω | pullCount A (bestArm ν) n ω = m₁ ∧ pullCount A a n ω = m₂ ∧ + sumRewards A R (bestArm ν) n ω ≤ sumRewards A R a n ω} ≤ + (Bandit.streamMeasure ν).real + {ω | ∑ i ∈ range m₁, ω i (bestArm ν) ≤ ∑ i ∈ range m₂, ω i a} := by + simp_rw [measureReal_def] + gcongr + · finiteness + · exact prob_sumRewards_le_sumRewards_le a n m₁ m₂ end ArrayModel @@ -210,6 +319,7 @@ lemma pullCount_eq_comp : ext simp [pullCount] +-- todo: write those lemmas with IdentDistrib instead of equality of maps lemma _root_.Learning.IsAlgEnvSeq.law_sumRewards_unique (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') : @@ -230,6 +340,52 @@ lemma _root_.Learning.IsAlgEnvSeq.law_sumRewards_unique · rw [measurable_pi_iff] exact fun n ↦ Measurable.prodMk (hA n) (hR n) +lemma _root_.Learning.IsAlgEnvSeq.law_pullCount_sumRewards_unique' + (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + (h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') : + IdentDistrib (fun ω a ↦ (pullCount A a n ω, sumRewards A R a n ω)) + (fun ω a ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) P P' := by + have hA := h1.measurable_A + have hR := h1.measurable_R + have hA2 := h2.measurable_A + have hR2 := h2.measurable_R + constructor + · refine Measurable.aemeasurable ?_ + rw [measurable_pi_iff] + exact fun a ↦ Measurable.prod (by fun_prop) (measurable_sumRewards hA hR _ _) + · refine Measurable.aemeasurable ?_ + rw [measurable_pi_iff] + exact fun a ↦ Measurable.prod (by fun_prop) (measurable_sumRewards hA2 hR2 _ _) + have h_unique := isAlgEnvSeq_unique h1 h2 + let f := fun (p : ℕ → α × ℝ ) (a : α) ↦ (∑ i ∈ range n, if (p i).1 = a then 1 else 0, + ∑ i ∈ range n, if (p i).1 = a then (p i).2 else 0) + have hf : Measurable f := by + rw [measurable_pi_iff] + intro a + refine Measurable.prod ?_ ?_ + · simp only [f] + refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + · simp only [f] + refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + have h_eq_comp : (fun ω a ↦ (pullCount A a n ω, sumRewards A R a n ω)) + = f ∘ (fun ω n ↦ (A n ω, R n ω)) := by + ext ω a : 2 + rw [pullCount_eq_comp (R := R), sumRewards_eq_comp] + grind + have h_eq_comp2 : (fun ω a ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) + = f ∘ (fun ω n ↦ (A₂ n ω, R₂ n ω)) := by + ext ω a : 2 + rw [pullCount_eq_comp (R := R₂), sumRewards_eq_comp] + grind + rw [h_eq_comp, h_eq_comp2, ← Measure.map_map hf, h_unique, Measure.map_map hf, + ← h_eq_comp2] + · rw [measurable_pi_iff] + exact fun n ↦ Measurable.prodMk (hA2 n) (hR2 n) + · rw [measurable_pi_iff] + exact fun n ↦ Measurable.prodMk (hA n) (hR n) + lemma _root_.Learning.IsAlgEnvSeq.law_pullCount_sumRewards_unique (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') : @@ -267,32 +423,186 @@ lemma _root_.Learning.IsAlgEnvSeq.law_pullCount_sumRewards_unique · rw [measurable_pi_iff] exact fun n ↦ Measurable.prodMk (hA n) (hR n) -lemma todo2 [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) - {s : Set ℕ} [DecidablePred (· ∈ s)] (hs : MeasurableSet s) {B : Set ℝ} (hB : MeasurableSet B) : - P {ω | pullCount A a n ω ∈ s ∧ sumRewards A R a n ω ∈ B} ≤ - ∑ k ∈ (range (n + 1)).filter (· ∈ s), - Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := by +-- this is what we will use for UCB +lemma prob_pullCount_prod_sumRewards_mem_le (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + {s : Set (ℕ × ℝ)} [DecidablePred (· ∈ Prod.fst '' s)] (hs : MeasurableSet s) : + P {ω | (pullCount A a n ω, sumRewards A R a n ω) ∈ s} ≤ + ∑ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), + Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := by have hA := h.measurable_A have hR := h.measurable_R - calc P {ω | pullCount A a n ω ∈ s ∧ sumRewards A R a n ω ∈ B} - _ = (P.map (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω))) (s ×ˢ B) := by - rw [Measure.map_apply (by fun_prop) (hs.prod hB)]; rfl + calc P {ω | (pullCount A a n ω, sumRewards A R a n ω) ∈ s} + _ = (P.map (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω))) s := by + rw [Measure.map_apply (by fun_prop) hs]; rfl _ = ((ArrayModel.arrayMeasure ν).map (fun ω ↦ (pullCount (ArrayModel.action alg) a n ω, - sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a n ω))) (s ×ˢ B) := by + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a n ω))) s := by rw [h.law_pullCount_sumRewards_unique (ArrayModel.isAlgEnvSeq_arrayMeasure alg ν)] - _ = (ArrayModel.arrayMeasure ν) {ω | pullCount (ArrayModel.action alg) a n ω ∈ s ∧ - sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a n ω ∈ B} := by - rw [Measure.map_apply (by fun_prop) (hs.prod hB)]; rfl - _ ≤ ∑ k ∈ (range (n + 1)).filter (· ∈ s), Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := - ArrayModel.todo'' a n hs hB + _ = (ArrayModel.arrayMeasure ν) {ω | (pullCount (ArrayModel.action alg) a n ω, + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a n ω) ∈ s} := by + rw [Measure.map_apply (by fun_prop) hs]; rfl + _ ≤ ∑ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), + Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := + ArrayModel.prob_pullCount_prod_sumRewards_mem_le a n hs + +lemma prob_pullCount_mem_and_sumRewards_mem_le (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + {s : Set ℕ} [DecidablePred (· ∈ s)] (hs : MeasurableSet s) {B : Set ℝ} (hB : MeasurableSet B) : + P {ω | pullCount A a n ω ∈ s ∧ sumRewards A R a n ω ∈ B} ≤ + ∑ k ∈ (range (n + 1)).filter (· ∈ s), + Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := by + classical + rcases Set.eq_empty_or_nonempty B with h_empty | h_nonempty + · simp [h_empty] + convert prob_pullCount_prod_sumRewards_mem_le h (hs.prod hB) (ν := ν) (alg := alg) with _ _ k hk + · ext n + have : ∃ x, x ∈ B := h_nonempty + simp [this] + · ext x + simp only [Set.mem_image, Set.mem_prod, Prod.exists, exists_and_right, exists_and_left, + exists_eq_right, mem_filter, mem_range] at hk + simp [hk.2.1] -lemma todo [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) +lemma todo (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {B : Set ℝ} (hB : MeasurableSet B) : P (sumRewards A R a n ⁻¹' B) ≤ ∑ k ∈ range (n + 1), Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := by classical - have h_le := todo2 h .univ hB (a := a) (n := n) + have h_le := prob_pullCount_mem_and_sumRewards_mem_le h .univ hB (a := a) (n := n) simpa using h_le +lemma prob_pullCount_eq_and_sumRewards_mem_le (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + {m : ℕ} (hm : m ≤ n) {B : Set ℝ} (hB : MeasurableSet B) : + P {ω | pullCount A a n ω = m ∧ sumRewards A R a n ω ∈ B} ≤ + Bandit.streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by + have h_le := prob_pullCount_mem_and_sumRewards_mem_le h (s := {m}) (by simp) hB (a := a) (n := n) + have hm' : m < n + 1 := by lia + simpa [hm'] using h_le + +lemma probReal_sumRewards_le_sumRewards_le [Fintype α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + (a : α) (n m₁ m₂ : ℕ) : + P.real {ω | pullCount A (bestArm ν) n ω = m₁ ∧ pullCount A a n ω = m₂ ∧ + sumRewards A R (bestArm ν) n ω ≤ sumRewards A R a n ω} ≤ + (Bandit.streamMeasure ν).real + {ω | ∑ i ∈ range m₁, ω i (bestArm ν) ≤ ∑ i ∈ range m₂, ω i a} := by + have hA := h.measurable_A + have hR := h.measurable_R + refine le_trans (le_of_eq ?_) + (ArrayModel.probReal_sumRewards_le_sumRewards_le (alg := alg) a n m₁ m₂) + let s := {p : ℕ × ℕ × ℝ × ℝ | p.1 = m₁ ∧ p.2.1 = m₂ ∧ p.2.2.1 ≤ p.2.2.2} + have hs : MeasurableSet s := by simp only [measurableSet_setOf, s]; fun_prop + change P.real ((fun ω ↦ (pullCount A (bestArm ν) n ω, + pullCount A a n ω, sumRewards A R (bestArm ν) n ω, sumRewards A R a n ω)) ⁻¹' s) = + (ArrayModel.arrayMeasure ν).real + ((fun ω ↦ (pullCount (ArrayModel.action alg) (bestArm ν) n ω, + pullCount (ArrayModel.action alg) a n ω, + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) (bestArm ν) n ω, + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a n ω)) ⁻¹' s) + simp_rw [measureReal_def] + congr 1 + rw [← Measure.map_apply ?_ hs, ← Measure.map_apply (by fun_prop) hs] + swap + · refine Measurable.prod (by fun_prop) (Measurable.prod (by fun_prop) ?_) + exact (measurable_sumRewards hA hR _ _).prod (measurable_sumRewards hA hR _ _) + congr 1 + refine IdentDistrib.map_eq ?_ + have h_eq := h.law_pullCount_sumRewards_unique' (ArrayModel.isAlgEnvSeq_arrayMeasure alg ν) + (n := n) + exact h_eq.comp (u := fun p ↦ ((p (bestArm ν)).1, (p a).1, (p (bestArm ν)).2, (p a).2)) + (by fun_prop) + +section Subgaussian + +omit [DecidableEq α] [StandardBorelSpace α] in +lemma probReal_sum_le_sum_streamMeasure [Fintype α] + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : α) (m : ℕ) : + (Bandit.streamMeasure ν).real + {ω | ∑ s ∈ range m, ω s (bestArm ν) ≤ ∑ s ∈ range m, ω s a} ≤ + Real.exp (-↑m * gap ν a ^ 2 / 4) := by + by_cases ha : a = bestArm ν + · simp [ha] + refine (HasSubgaussianMGF.measure_sum_le_sum_le' (cX := fun _ ↦ 1) (cY := fun _ ↦ 1) + ?_ ?_ ?_ ?_ ?_ ?_).trans_eq ?_ + · exact iIndepFun_eval_streamMeasure'' ν (bestArm ν) + · exact iIndepFun_eval_streamMeasure'' ν a + · intro i him + simp_rw [integral_eval_streamMeasure] + refine (hν (bestArm ν)).congr_identDistrib ?_ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ + · intro i him + simp_rw [integral_eval_streamMeasure] + refine (hν a).congr_identDistrib ?_ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ + · exact indepFun_eval_streamMeasure' ν (Ne.symm ha) + · gcongr 1 with i him + simp_rw [integral_eval_streamMeasure] + exact le_bestArm a + · congr 1 + simp_rw [integral_eval_streamMeasure] + simp only [id_eq, sum_const, card_range, nsmul_eq_mul, mul_one, NNReal.coe_natCast, + gap_eq_bestArm_sub, neg_mul] + field_simp + ring + +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in +lemma prob_sum_le_sqrt_log + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) {c : ℝ} (hc : 0 ≤ c) + (a : α) (k : ℕ) (hk : k ≠ 0) : + Bandit.streamMeasure ν + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ - √(c * k * Real.log (n + 1))} ≤ + 1 / (n + 1) ^ (c / 2) := by + calc + Bandit.streamMeasure ν + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ - √(c * k * Real.log (n + 1))} + _ ≤ ENNReal.ofReal (Real.exp (-(√(c * k * Real.log (n + 1))) ^ 2 / (2 * k * 1))) := by + rw [← ofReal_measureReal] + gcongr + refine (HasSubgaussianMGF.measure_sum_range_le_le_of_iIndepFun (c := 1) ?_ ?_ (by positivity)) + · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) + (fun _ ↦ by fun_prop) + · intro i him + refine (hν a).congr_identDistrib ?_ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ + _ = 1 / (n + 1) ^ (c / 2) := by + rw [Real.sq_sqrt] + swap; · exact mul_nonneg (by positivity) (Real.log_nonneg (by simp)) + field_simp + rw [div_eq_inv_mul, ← mul_assoc, ← Real.log_rpow (by positivity), ← Real.log_inv, + Real.exp_log (by positivity), one_div, ENNReal.ofReal_inv_of_pos (by positivity), + ← ENNReal.ofReal_rpow_of_nonneg (by positivity) (by positivity)] + congr 2 + · norm_cast + · field + +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in +lemma prob_sum_ge_sqrt_log + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) {c : ℝ} (hc : 0 ≤ c) + (a : α) (k : ℕ) (hk : k ≠ 0) : + Bandit.streamMeasure ν + {ω | √(c * k * Real.log (n + 1)) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} ≤ + 1 / (n + 1) ^ (c / 2) := by + calc + Bandit.streamMeasure ν + {ω | √(c * k * Real.log (n + 1)) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} + _ ≤ ENNReal.ofReal (Real.exp (-(√(c * k * Real.log (n + 1))) ^ 2 / (2 * k * 1))) := by + rw [← ofReal_measureReal] + gcongr + refine (HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := 1) ?_ ?_ (by positivity)) + · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) + (fun _ ↦ by fun_prop) + · intro i him + refine (hν a).congr_identDistrib ?_ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ + _ = 1 / (n + 1) ^ (c / 2) := by + rw [Real.sq_sqrt] + swap; · exact mul_nonneg (by positivity) (Real.log_nonneg (by simp)) + field_simp + rw [div_eq_inv_mul, ← mul_assoc, ← Real.log_rpow (by positivity), ← Real.log_inv, + Real.exp_log (by positivity), one_div, ENNReal.ofReal_inv_of_pos (by positivity), + ← ENNReal.ofReal_rpow_of_nonneg (by positivity) (by positivity)] + congr 2 + · norm_cast + · field + +end Subgaussian + end Bandits diff --git a/LeanBandits/BanditAlgorithms/ETC.lean b/LeanBandits/BanditAlgorithms/ETC.lean index fe691261..ee84d0d2 100644 --- a/LeanBandits/BanditAlgorithms/ETC.lean +++ b/LeanBandits/BanditAlgorithms/ETC.lean @@ -3,9 +3,8 @@ 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 LeanBandits.Bandit.RewardByCountMeasure +import LeanBandits.Bandit.SumRewards import LeanBandits.ForMathlib.MeasurableArgMax -import LeanBandits.ForMathlib.SubGaussian /-! # The Explore-Then-Commit Algorithm @@ -22,21 +21,6 @@ lemma ae_eq_set_iff {α : Type*} {mα : MeasurableSpace α} {μ : Measure α} {s simp only [eq_iff_iff] congr! ---todo: generalize Icc -lemma measurable_sum_of_le {α : Type*} {mα : MeasurableSpace α} - {f : ℕ → α → ℝ} {g : α → ℕ} {n : ℕ} (hg_le : ∀ a, g a ≤ n) (hf : ∀ i, Measurable (f i)) - (hg : Measurable g) : - Measurable (fun a ↦ ∑ i ∈ Icc 1 (g a), f i a) := by - have h_eq : (fun a ↦ ∑ i ∈ Icc 1 (g a), f i a) - = fun a ↦ ∑ i ∈ range (n + 1), if g a = i then ∑ j ∈ Icc 1 i, f j a else 0 := by - ext ω - rw [sum_ite_eq_of_mem] - grind - rw [h_eq] - refine measurable_sum _ fun n hn ↦ ?_ - refine Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact (measurableSet_singleton _).preimage (by fun_prop) - lemma sum_mod_range {K : ℕ} (hK : 0 < K) (a : Fin K) : (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) = 1 := by have h_iff (s : ℕ) (hs : s < K) : ⟨s % K, Nat.mod_lt _ hK⟩ = a ↔ s = a := by @@ -245,31 +229,46 @@ lemma sumRewards_bestArm_le_of_arm_mul_eq [Nonempty (Fin K)] · simp [ha, hm] · simp [h_best, hm] -variable [StandardBorelSpace Ω] [Nonempty Ω] - -lemma identDistrib_aux [Nonempty (Fin K)] - (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (a b : Fin K) : - IdentDistrib - (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount A R a s ω, ∑ s ∈ Icc 1 m, rewardByCount A R b s ω)) - (fun ω ↦ (∑ s ∈ range m, ω.2 s a, ∑ s ∈ range m, ω.2 s b)) - 𝔓 (Bandit.measure (etcAlgorithm hK m) ν) := by - have h2 (a : Fin K) : IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount A R a s ω) - (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓 (Bandit.measure (etcAlgorithm hK m) ν) := - identDistrib_sum_Icc_rewardByCount h m a - by_cases hab : a = b - · simp only [hab] - exact (h2 b).comp (u := fun p ↦ (p, p)) (by fun_prop) - refine (h2 a).prodMk (h2 b) ?_ ?_ - · suffices IndepFun (fun ω s ↦ rewardByCount A R a s ω) (fun ω s ↦ rewardByCount A R b s ω) - 𝔓 by - exact this.comp (φ := fun p ↦ ∑ i ∈ Icc 1 m, p i) (ψ := fun p ↦ ∑ j ∈ Icc 1 m, p j) - (by fun_prop) (by fun_prop) - exact indepFun_rewardByCount_of_ne h hab - · suffices IndepFun (fun ω s ↦ ω.2 s a) (fun ω s ↦ ω.2 s b) - (Bandit.measure (etcAlgorithm hK m) ν) by - exact this.comp (φ := fun p ↦ ∑ i ∈ range m, p i) (ψ := fun p ↦ ∑ j ∈ range m, p j) - (by fun_prop) (by fun_prop) - exact indepFun_eval_snd_measure _ ν hab +-- lemma identDistrib_aux [Nonempty (Fin K)] +-- (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (a b : Fin K) : +-- IdentDistrib +-- (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount A R a s ω, ∑ s ∈ Icc 1 m, rewardByCount A R b s ω)) +-- (fun ω ↦ (∑ s ∈ range m, ω.2 s a, ∑ s ∈ range m, ω.2 s b)) +-- 𝔓 (Bandit.measure (etcAlgorithm hK m) ν) := by +-- have h2 (a : Fin K) : IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount A R a s ω) +-- (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓 (Bandit.measure (etcAlgorithm hK m) ν) := +-- identDistrib_sum_Icc_rewardByCount h m a +-- by_cases hab : a = b +-- · simp only [hab] +-- exact (h2 b).comp (u := fun p ↦ (p, p)) (by fun_prop) +-- refine (h2 a).prodMk (h2 b) ?_ ?_ +-- · suffices IndepFun (fun ω s ↦ rewardByCount A R a s ω) (fun ω s ↦ rewardByCount A R b s ω) +-- 𝔓 by +-- exact this.comp (φ := fun p ↦ ∑ i ∈ Icc 1 m, p i) (ψ := fun p ↦ ∑ j ∈ Icc 1 m, p j) +-- (by fun_prop) (by fun_prop) +-- exact indepFun_rewardByCount_of_ne h hab +-- · suffices IndepFun (fun ω s ↦ ω.2 s a) (fun ω s ↦ ω.2 s b) +-- (Bandit.measure (etcAlgorithm hK m) ν) by +-- exact this.comp (φ := fun p ↦ ∑ i ∈ range m, p i) (ψ := fun p ↦ ∑ j ∈ range m, p j) +-- (by fun_prop) (by fun_prop) +-- exact indepFun_eval_snd_measure _ ν hab + +lemma probReal_sumRewards_le_sumRewards_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) : + P.real {ω | sumRewards A R (bestArm ν) (K * m) ω ≤ sumRewards A R a (K * m) ω} ≤ + Real.exp (-↑m * gap ν a ^ 2 / 4) := by + have hA := h.measurable_A + have hR := h.measurable_R + have h1 := Bandits.probReal_sumRewards_le_sumRewards_le h a (K * m) m m + have h2 := probReal_sum_le_sum_streamMeasure hν a m + refine le_trans (le_of_eq ?_) (h1.trans h2) + simp_rw [measureReal_def] + congr 1 + refine measure_congr ?_ + rw [ae_eq_set_iff] + filter_upwards [pullCount_mul h a, pullCount_mul h (bestArm ν)] with ω ha h_best + simp [ha, h_best] /-- The probability that at time `K * m` the ETC algorithm chooses arm `a` is at most `exp(- m * Δ_a^2 / 4)`. -/ @@ -278,8 +277,6 @@ lemma prob_arm_mul_eq_le [Nonempty (Fin K)] (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) (hm : m ≠ 0) : P.real {ω | A (K * m) ω = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by - have hA := h.measurable_A - have hR := h.measurable_R have h_pos : 0 < K * m := Nat.mul_pos hK hm.bot_lt have h_le : P.real {ω | A (K * m) ω = a} ≤ P.real {ω | sumRewards A R (bestArm ν) (K * m) ω ≤ sumRewards A R a (K * m) ω} := by @@ -288,111 +285,7 @@ lemma prob_arm_mul_eq_le [Nonempty (Fin K)] · simp refine measure_mono_ae ?_ exact sumRewards_bestArm_le_of_arm_mul_eq h a hm - refine h_le.trans ?_ - -- extend the probability space to include the stream of independent rewards - suffices (𝔓).real {ω | sumRewards A R (bestArm ν) (K * m) ω.1 ≤ sumRewards A R a (K * m) ω.1} - ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) by - suffices P.real {ω | sumRewards A R (bestArm ν) (K * m) ω ≤ sumRewards A R a (K * m) ω} - = (𝔓).real {ω | sumRewards A R (bestArm ν) (K * m) ω.1 ≤ sumRewards A R a (K * m) ω.1} by - rwa [this] - calc P.real {ω | sumRewards A R (bestArm ν) (K * m) ω ≤ sumRewards A R a (K * m) ω} - _ = ((𝔓).fst).real {ω | sumRewards A R (bestArm ν) (K * m) ω ≤ sumRewards A R a (K * m) ω} := by - simp - _ = (𝔓).real {ω | sumRewards A R (bestArm ν) (K * m) ω.1 ≤ sumRewards A R a (K * m) ω.1} := by - rw [Measure.fst, map_measureReal_apply (by fun_prop)] - · rfl - · have h_meas := measurable_sumRewards h.measurable_A h.measurable_R - exact measurableSet_le (by fun_prop) (by fun_prop) - calc (𝔓).real {ω | sumRewards A R (bestArm ν) (K * m) ω.1 ≤ sumRewards A R a (K * m) ω.1} - _ = (𝔓).real {ω | ∑ s ∈ Icc 1 (pullCount A (bestArm ν) (K * m) ω.1), - rewardByCount A R (bestArm ν) s ω - ≤ ∑ s ∈ Icc 1 (pullCount A a (K * m) ω.1), rewardByCount A R a s ω} := by - congr with ω - congr! 1 <;> rw [sum_rewardByCount_eq_sumRewards] - _ = (𝔓).real {ω | ∑ s ∈ Icc 1 m, rewardByCount A R (bestArm ν) s ω - ≤ ∑ s ∈ Icc 1 m, rewardByCount A R a s ω} := by - simp_rw [measureReal_def] - congr 1 - refine measure_congr ?_ - have ha := pullCount_mul h a (hK := hK) (ν := ν) (m := m) - have h_best := pullCount_mul h (bestArm ν) (hK := hK) (ν := ν) (m := m) - rw [ae_eq_set_iff, Measure.ae_prod_iff_ae_ae] - · filter_upwards [ha, h_best] with ω ha h_best - refine ae_of_all _ fun ω' ↦ ?_ - rw [ha, h_best] - · simp only [Set.mem_setOf_eq] - let f₁ := fun ω : Ω × (ℕ → Fin K → ℝ) ↦ - ∑ s ∈ Icc 1 (pullCount A (bestArm ν) (K * m) ω.1), rewardByCount A R (bestArm ν) s ω - let g₁ := fun ω : Ω × (ℕ → Fin K → ℝ) ↦ - ∑ s ∈ Icc 1 (pullCount A a (K * m) ω.1), rewardByCount A R a s ω - let f₂ := fun ω : Ω × (ℕ → Fin K → ℝ) ↦ - ∑ s ∈ Icc 1 m, rewardByCount A R (bestArm ν) s ω - let g₂ := fun ω : Ω × (ℕ → Fin K → ℝ) ↦ ∑ s ∈ Icc 1 m, rewardByCount A R a s ω - change MeasurableSet {x | f₁ x ≤ g₁ x ↔ f₂ x ≤ g₂ x} - have hf₁ : Measurable f₁ := by - refine measurable_sum_of_le (n := K * m + 1) - (g := fun ω : Ω × (ℕ → Fin K → ℝ) ↦ pullCount A (bestArm ν) (K * m) ω.1) - (f := rewardByCount A R (bestArm ν)) (fun ω ↦ ?_) - (by fun_prop) (by fun_prop) - have h_le := pullCount_le (A := A) (bestArm ν) (K * m) ω.1 - grind - have hg₁ : Measurable g₁ := by - refine measurable_sum_of_le (n := K * m + 1) - (g := fun ω : Ω × (ℕ → Fin K → ℝ) ↦ pullCount A a (K * m) ω.1) - (f := rewardByCount A R a) (fun ω ↦ ?_) (by fun_prop) (by fun_prop) - have h_le := pullCount_le (A := A) a (K * m) ω.1 - grind - refine MeasurableSet.iff ?_ ?_ - · exact measurableSet_le (by fun_prop) (by fun_prop) - · exact measurableSet_le (by fun_prop) (by fun_prop) - _ = (Bandit.measure (etcAlgorithm hK m) ν).real - {ω | ∑ s ∈ range m, ω.2 s (bestArm ν) ≤ ∑ s ∈ range m, ω.2 s a} := by - simp_rw [measureReal_def] - congr 1 - have : (𝔓).map (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount A R (bestArm ν) s ω, - ∑ s ∈ Icc 1 m, rewardByCount A R a s ω)) - = (Bandit.measure (etcAlgorithm hK m) ν).map - (fun ω ↦ (∑ s ∈ range m, ω.2 s (bestArm ν), ∑ s ∈ range m, ω.2 s a)) := - (identDistrib_aux h (bestArm ν) a).map_eq - rw [Measure.ext_iff] at this - have h_meas : MeasurableSet {x : ℝ × ℝ | x.1 ≤ x.2} := - measurableSet_le (by fun_prop) (by fun_prop) - specialize this {x | x.1 ≤ x.2} h_meas - rw [Measure.map_apply (by fun_prop) h_meas, Measure.map_apply (by fun_prop) h_meas] at this - convert this - _ = (Bandit.streamMeasure ν).real - {ω | ∑ s ∈ range m, ω s (bestArm ν) ≤ ∑ s ∈ range m, ω s a} := by - simp_rw [measureReal_def] - congr 1 - rw [← Bandit.snd_measure (etcAlgorithm hK m), Measure.snd_apply] - · rfl - · exact measurableSet_le (by fun_prop) (by fun_prop) - _ ≤ Real.exp (-↑m * gap ν a ^ 2 / 4) := by - by_cases ha : a = bestArm ν - · simp [ha] - -- Apply a sub-Gaussian concentration inequality - refine (HasSubgaussianMGF.measure_sum_le_sum_le' (cX := fun _ ↦ 1) (cY := fun _ ↦ 1) - ?_ ?_ ?_ ?_ ?_ ?_).trans_eq ?_ - · exact iIndepFun_eval_streamMeasure'' ν (bestArm ν) - · exact iIndepFun_eval_streamMeasure'' ν a - · intro i him - simp_rw [integral_eval_streamMeasure] - refine (hν (bestArm ν)).congr_identDistrib ?_ - exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ - · intro i him - simp_rw [integral_eval_streamMeasure] - refine (hν a).congr_identDistrib ?_ - exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ - · exact indepFun_eval_streamMeasure' ν (Ne.symm ha) - · gcongr 1 with i him - simp_rw [integral_eval_streamMeasure] - exact le_bestArm a - · congr 1 - simp_rw [integral_eval_streamMeasure] - simp only [id_eq, sum_const, card_range, nsmul_eq_mul, mul_one, NNReal.coe_natCast, - gap_eq_bestArm_sub, neg_mul] - field_simp - ring + exact h_le.trans (probReal_sumRewards_le_sumRewards_le h hν a) /-- Bound on the expectation of the number of pulls of each arm by the ETC algorithm. -/ lemma expectation_pullCount_le [Nonempty (Fin K)] diff --git a/LeanBandits/BanditAlgorithms/UCB.lean b/LeanBandits/BanditAlgorithms/UCB.lean index 74973c59..bc4d13c2 100644 --- a/LeanBandits/BanditAlgorithms/UCB.lean +++ b/LeanBandits/BanditAlgorithms/UCB.lean @@ -214,134 +214,52 @@ lemma pullCount_arm_le [Nonempty (Fin K)] (hc : 0 ≤ c) · have : 0 ≤ log (n + 1) := by simp [log_nonneg] positivity -variable [StandardBorelSpace Ω] [Nonempty Ω] - lemma todo [Nonempty (Fin K)] - (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n k : ℕ) (hk : k ≠ 0) : - 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} ≤ + Bandit.streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2) := by - have hA := h.measurable_A - have hR := h.measurable_R - have h_meas : MeasurableSet {ω | ω / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := - measurableSet_le (by fun_prop) measurable_const have h_log_nonneg : 0 ≤ log (n + 1) := log_nonneg (by simp) - calc - 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} - _ = ((𝔓).map (fun ω ↦ ∑ m ∈ Icc 1 k, rewardByCount A R a m ω)) - {ω | ω / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := by - rw [Measure.map_apply (by fun_prop) h_meas] - rfl - _ = ((Bandit.measure (ucbAlgorithm hK c) ν).map (fun ω ↦ ∑ s ∈ range k, ω.2 s a)) - {ω | ω / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := by - rw [IdentDistrib.map_eq (identDistrib_sum_Icc_rewardByCount h k a)] - _ = (Bandit.measure (ucbAlgorithm hK c) ν) - {ω | (∑ s ∈ range k, ω.2 s a) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := by - rw [Measure.map_apply (by fun_prop) h_meas] - rfl - _ = (Bandit.measure (ucbAlgorithm hK c) ν) - {ω | (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) / k ≤ - √(c * log (n + 1) / k)} := by + calc Bandit.streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} + _ = Bandit.streamMeasure ν + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) / k ≤ - √(c * log (n + 1) / k)} := by congr with ω field_simp rw [Finset.sum_sub_distrib] simp grind - _ = (Bandit.measure (ucbAlgorithm hK c) ν) - {ω | (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) ≤ - √(c * k * log (n + 1))} := by + _ = Bandit.streamMeasure ν + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ - √(c * k * log (n + 1))} := by congr with ω field_simp congr! 2 rw [sqrt_div (by positivity), ← mul_div_assoc, mul_comm, mul_div_assoc, div_sqrt, mul_assoc (k : ℝ), sqrt_mul (x := (k : ℝ)) (by positivity), mul_comm] - _ = Bandit.streamMeasure ν - {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ - √(c * k * log (n + 1))} := by - rw [← Bandit.snd_measure (ucbAlgorithm hK c), Measure.snd_apply] - · rfl - · exact measurableSet_le (by fun_prop) (by fun_prop) - _ ≤ ENNReal.ofReal (exp (-(√(c * k * log (n + 1))) ^ 2 / (2 * k * 1))) := by - rw [← ofReal_measureReal] - gcongr - refine (HasSubgaussianMGF.measure_sum_range_le_le_of_iIndepFun (c := 1) ?_ ?_ (by positivity)) - · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) - (fun _ ↦ by fun_prop) - · intro i him - refine (hν a).congr_identDistrib ?_ - exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ - _ = 1 / (n + 1) ^ (c / 2) := by - rw [sq_sqrt] - swap; · exact mul_nonneg (by positivity) (log_nonneg (by simp)) - field_simp - rw [div_eq_inv_mul, ← mul_assoc, ← Real.log_rpow (by positivity), ← Real.log_inv, - Real.exp_log (by positivity), one_div, ENNReal.ofReal_inv_of_pos (by positivity), - ← ENNReal.ofReal_rpow_of_nonneg (by positivity) (by positivity)] - congr 2 - · norm_cast - · field + _ ≤ 1 / (n + 1) ^ (c / 2) := prob_sum_le_sqrt_log hν hc a k hk lemma todo' [Nonempty (Fin K)] - (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n k : ℕ) (hk : k ≠ 0) : - 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k - √(c * log (n + 1) / k)} ≤ + Bandit.streamMeasure ν + {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - √(c * log (n + 1) / k)} ≤ 1 / (n + 1) ^ (c / 2) := by - have hA := h.measurable_A - have hR := h.measurable_R - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - have h_meas : MeasurableSet {ω |(ν a)[id] ≤ ω / k - √(c * log (n + 1) / k)} := - measurableSet_le (by fun_prop) (by fun_prop) have h_log_nonneg : 0 ≤ log (n + 1) := log_nonneg (by simp) - calc - 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k - √(c * log (n + 1) / k)} - _ = ((𝔓).map (fun ω ↦ ∑ m ∈ Icc 1 k, rewardByCount A R a m ω)) - {ω | (ν a)[id] ≤ ω / k - √(c * log (n + 1) / k)} := by - rw [Measure.map_apply (by fun_prop) h_meas] - rfl - _ = ((Bandit.measure (ucbAlgorithm hK c) ν).map (fun ω ↦ ∑ s ∈ range k, ω.2 s a)) - {ω | (ν a)[id] ≤ ω / k - √(c * log (n + 1) / k)} := by - rw [IdentDistrib.map_eq (identDistrib_sum_Icc_rewardByCount h k a)] - _ = (Bandit.measure (ucbAlgorithm hK c) ν) - {ω | (ν a)[id] ≤ (∑ s ∈ range k, ω.2 s a) / k - √(c * log (n + 1) / k)} := by - rw [Measure.map_apply (by fun_prop) h_meas] - rfl - _ = (Bandit.measure (ucbAlgorithm hK c) ν) - {ω | √(c * log (n + 1) / k) ≤ (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) / k} := by + calc Bandit.streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - √(c * log (n + 1) / k)} + _ = Bandit.streamMeasure ν + {ω | √(c * log (n + 1) / k) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id])) / k} := by congr with ω field_simp rw [Finset.sum_sub_distrib] simp grind - _ = (Bandit.measure (ucbAlgorithm hK c) ν) - {ω | √(c * k * log (n + 1)) ≤ (∑ s ∈ range k, (ω.2 s a - (ν a)[id]))} := by + _ = Bandit.streamMeasure ν + {ω | √(c * k * log (n + 1)) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} := by congr with ω field_simp congr! 1 rw [sqrt_div (by positivity), ← mul_div_assoc, mul_comm, mul_div_assoc, div_sqrt, mul_comm _ (k : ℝ), sqrt_mul (x := (k : ℝ)) (by positivity), mul_comm] - _ = Bandit.streamMeasure ν - {ω | √(c * k * log (n + 1)) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} := by - rw [← Bandit.snd_measure (ucbAlgorithm hK c), Measure.snd_apply] - · rfl - · exact measurableSet_le (by fun_prop) (by fun_prop) - _ ≤ ENNReal.ofReal (exp (-(√(c * k * log (n + 1))) ^ 2 / (2 * k * 1))) := by - rw [← ofReal_measureReal] - gcongr - refine (HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := 1) ?_ ?_ (by positivity)) - · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) - (fun _ ↦ by fun_prop) - · intro i him - refine (hν a).congr_identDistrib ?_ - exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ - _ = 1 / (n + 1) ^ (c / 2) := by - rw [sq_sqrt] - swap; · exact mul_nonneg (by positivity) (log_nonneg (by simp)) - field_simp - rw [div_eq_inv_mul, ← mul_assoc, ← Real.log_rpow (by positivity), ← Real.log_inv, - Real.exp_log (by positivity), one_div, ENNReal.ofReal_inv_of_pos (by positivity), - ← ENNReal.ofReal_rpow_of_nonneg (by positivity) (by positivity)] - congr 2 - · norm_cast - · field + _ ≤ 1 / (n + 1) ^ (c / 2) := prob_sum_ge_sqrt_log hν hc a k hk lemma prob_ucbIndex_le [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) @@ -349,43 +267,30 @@ lemma prob_ucbIndex_le [Nonempty (Fin K)] (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : P {h | 0 < pullCount A a n h ∧ empMean A R a n h + ucbWidth A c a n h ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by - have hA := h.measurable_A - have hR := h.measurable_R - -- extend the probability space - suffices 𝔓 {ω | 0 < pullCount A a n ω.1 ∧ - empMean A R a n ω.1 + ucbWidth A c a n ω.1 ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2 - 1) by - rwa [← Measure.fst_prod (μ := P) (ν := Bandit.streamMeasure ν), Measure.fst_apply] - change MeasurableSet ({h | 0 < pullCount A a n h} - ∩ {h | empMean A R a n h + ucbWidth A c a n h ≤ ∫ (x : ℝ), id x ∂ν a}) - refine MeasurableSet.inter ?_ ?_ - · exact measurableSet_lt (by fun_prop) (by fun_prop) - · exact measurableSet_le (by fun_prop) (by fun_prop) - -- express with `rewardByCount` and `pullCount` - unfold empMean ucbWidth - simp_rw [← sum_rewardByCount_eq_sumRewards] - calc - 𝔓 {ω | 0 < pullCount A a n ω.1 ∧ - (∑ m ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a m ω) / pullCount A a n ω.1 + - √(c * log (↑n + 1) / pullCount A a n ω.1) ≤ (ν a)[id]} - -- list the possible values of `pullCount a n ω.1` - _ ≤ 𝔓 {ω | ∃ k ≤ n, 0 < k ∧ (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k + - √(c * log (↑n + 1) / k) ≤ (ν a)[id]} := by - refine measure_mono fun ω hω ↦ ?_ - simp only [Nat.cast_nonneg, sqrt_div', id_eq, Set.mem_setOf_eq] at hω ⊢ - exact ⟨pullCount A a n ω.1, pullCount_le _ _ _, hω⟩ - _ = 𝔓 (⋃ k ∈ Icc 1 n, {ω |(∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k + - √(c * log (↑n + 1) / k) ≤ (ν a)[id]}) := by - congr 1 - ext ω - simp - grind - -- Union bound over the possible values of `pullCount a n ω.1` + let s : Set (ℕ × ℝ) := {(m, x) | 0 < m ∧ x / m + √(c * log (↑n + 1) / m) ≤ (ν a)[id]} + have hs : MeasurableSet s := by + simp only [Nat.cast_nonneg, sqrt_div', id_eq, measurableSet_setOf, s] + fun_prop + classical + calc P {h | 0 < pullCount A a n h ∧ empMean A R a n h + ucbWidth A c a n h ≤ (ν a)[id]} + _ ≤ ∑ k ∈ range (n + 1) with k ∈ Prod.fst '' s, + (Bandit.streamMeasure ν) {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := + prob_pullCount_prod_sumRewards_mem_le h hs _ ≤ ∑ k ∈ Icc 1 n, - 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k + √(c * log (↑n + 1) / k) ≤ (ν a)[id]} := - measure_biUnion_finset_le _ _ + (Bandit.streamMeasure ν) {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := by + refine Finset.sum_le_sum_of_subset_of_nonneg (fun m ↦ ?_) fun _ _ _ ↦ by positivity + simp [s] + grind + _ = ∑ k ∈ Icc 1 n, + (Bandit.streamMeasure ν) {ω | (∑ i ∈ range k, ω i a) / k + √(c * log (↑n + 1) / k) ≤ + (ν a)[id]} := by + refine Finset.sum_congr rfl fun k hk ↦ ?_ + congr with ω + have hk : 0 < k := by grind + simp [s, hk] _ ≤ ∑ k ∈ Icc 1 n, (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by gcongr with k hk - exact todo h hν hc a n k (by grind) + exact todo hν hc a n k (by grind) _ ≤ (n + 1) * (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by simp only [one_div, sum_const, Nat.card_Icc, add_tsub_cancel_right, nsmul_eq_mul, mul_one] rw [div_eq_mul_inv ((n : ℝ≥0∞) + 1)] @@ -402,43 +307,30 @@ lemma prob_ucbIndex_ge [Nonempty (Fin K)] (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : P {h | 0 < pullCount A a n h ∧ (ν a)[id] ≤ empMean A R a n h - ucbWidth A c a n h} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by - have hA := h.measurable_A - have hR := h.measurable_R - -- extend the probability space - suffices 𝔓 {ω | 0 < pullCount A a n ω.1 ∧ - (ν a)[id] ≤ empMean A R a n ω.1 - ucbWidth A c a n ω.1} ≤ 1 / (n + 1) ^ (c / 2 - 1) by - rwa [← Measure.fst_prod (μ := P) (ν := Bandit.streamMeasure ν), Measure.fst_apply] - change MeasurableSet ({h | 0 < pullCount A a n h} - ∩ {h | (ν a)[id] ≤ empMean A R a n h - ucbWidth A c a n h}) - refine MeasurableSet.inter ?_ ?_ - · exact measurableSet_lt (by fun_prop) (by fun_prop) - · exact measurableSet_le (by fun_prop) (by fun_prop) - -- express with `rewardByCount` and `pullCount` - unfold empMean ucbWidth - simp_rw [← sum_rewardByCount_eq_sumRewards] - calc - 𝔓 {ω | 0 < pullCount A a n ω.1 ∧ - (ν a)[id] ≤ (∑ m ∈ Icc 1 (pullCount A a n ω.1), rewardByCount A R a m ω) / pullCount A a n ω.1 - - √(c * log (↑n + 1) / pullCount A a n ω.1)} - -- list the possible values of `pullCount A a n ω.1` - _ ≤ 𝔓 {ω | ∃ k ≤ n, 0 < k ∧ (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k - - √(c * log (↑n + 1) / k)} := by - refine measure_mono fun ω hω ↦ ?_ - simp only [Nat.cast_nonneg, sqrt_div', id_eq, Set.mem_setOf_eq] at hω ⊢ - exact ⟨pullCount A a n ω.1, pullCount_le _ _ _, hω⟩ - _ = 𝔓 (⋃ k ∈ Icc 1 n, {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k - - √(c * log (↑n + 1) / k)}) := by - congr 1 - ext ω - simp - grind - -- Union bound over the possible values of `pullCount a n ω.1` + let s : Set (ℕ × ℝ) := {(m, x) | 0 < m ∧ (ν a)[id] ≤ x / m - √(c * log (↑n + 1) / m)} + have hs : MeasurableSet s := by + simp only [Nat.cast_nonneg, sqrt_div', id_eq, measurableSet_setOf, s] + fun_prop + classical + calc P {h | 0 < pullCount A a n h ∧ (ν a)[id] ≤ empMean A R a n h - ucbWidth A c a n h} + _ ≤ ∑ k ∈ range (n + 1) with k ∈ Prod.fst '' s, + (Bandit.streamMeasure ν) {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := + prob_pullCount_prod_sumRewards_mem_le h hs _ ≤ ∑ k ∈ Icc 1 n, - 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount A R a m ω) / k - √(c * log (↑n + 1) / k)} := - measure_biUnion_finset_le _ _ + (Bandit.streamMeasure ν) {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := by + refine Finset.sum_le_sum_of_subset_of_nonneg (fun m ↦ ?_) fun _ _ _ ↦ by positivity + simp [s] + grind + _ = ∑ k ∈ Icc 1 n, + (Bandit.streamMeasure ν) + {ω | (ν a)[id] ≤ (∑ i ∈ range k, ω i a) / k - √(c * log (↑n + 1) / k)} := by + refine Finset.sum_congr rfl fun k hk ↦ ?_ + congr with ω + have hk : 0 < k := by grind + simp [s, hk] _ ≤ ∑ k ∈ Icc 1 n, (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by gcongr with k hk - exact todo' h hν hc a n k (by grind) + exact todo' hν hc a n k (by grind) _ ≤ (n + 1) * (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by simp only [one_div, sum_const, Nat.card_Icc, add_tsub_cancel_right, nsmul_eq_mul, mul_one] rw [div_eq_mul_inv ((n : ℝ≥0∞) + 1)] @@ -475,43 +367,7 @@ lemma probReal_ucbIndex_ge [Nonempty (Fin K)] rw [← ENNReal.toReal_rpow] norm_cast -omit [Nonempty Ω] in -lemma pullCount_le_add [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (ω : Ω) : - pullCount A a n ω ≤ C + 1 + - ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s := by - rw [pullCount_eq_sum] - calc ∑ s ∈ range n, if A s ω = a then 1 else 0 - _ ≤ ∑ s ∈ range n, ({s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s + - {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s) := by - gcongr with s hs - simp [Set.indicator_apply] - grind - _ = ∑ s ∈ range n, {s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s + - ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s := by - rw [Finset.sum_add_distrib] - _ ≤ C + 1 + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s := by - gcongr - have h_le n : ∑ s ∈ range n, {s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s ≤ - pullCount A a n ω := by - rw [pullCount_eq_sum] - gcongr with s hs - simp only [Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] - grind - induction n with - | zero => simp - | succ n hn => - rw [Finset.sum_range_succ] - rcases le_or_gt (pullCount A a n ω) C with h_pc | h_pc - · have hn' : ∑ s ∈ range n, {s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s ≤ C := - (h_le n).trans h_pc - grw [hn'] - gcongr - simp only [Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] - grind - · refine le_trans ?_ hn - simp [h_pc] - -omit [StandardBorelSpace Ω] [Nonempty Ω] [IsMarkovKernel ν] in +omit [IsMarkovKernel ν] in lemma pullCount_le_add_three [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (ω : Ω) : pullCount A a n ω ≤ C + 1 + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω ∧ @@ -533,8 +389,7 @@ lemma pullCount_le_add_three [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (ω : Ω empMean A R (A s ω) s ω - ucbWidth A c (A s ω) s ω ≤ (ν (A s ω))[id]} let C' := {s | C < pullCount A a s ω ∧ empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω < (ν (bestArm ν))[id]} - let D := {s | C < pullCount A a s ω ∧ (ν a)[id] < - empMean A R a s ω - ucbWidth A c a s ω} + let D := {s | C < pullCount A a s ω ∧ (ν a)[id] < empMean A R a s ω - ucbWidth A c a s ω} change ∑ s ∈ range n, A'.indicator 1 s ≤ ∑ s ∈ range n, B.indicator 1 s + ∑ s ∈ range n, C'.indicator 1 s + ∑ s ∈ range n, D.indicator 1 s @@ -555,7 +410,6 @@ lemma pullCount_le_add_three [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (ω : Ω ∑ s ∈ range n, D.indicator 1 s := by rw [Finset.sum_add_distrib, Finset.sum_add_distrib] -omit [StandardBorelSpace Ω] [Nonempty Ω] in lemma pullCount_le_add_three_ae [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) (n C : ℕ) (hC : C ≠ 0) : @@ -578,7 +432,6 @@ lemma pullCount_le_add_three_ae [Nonempty (Fin K)] exact fun h_gt ↦ hω _ (lt_of_le_of_lt (by grind) h_gt) _ · exact fun h_gt ↦ hω _ (lt_of_le_of_lt (by grind) h_gt) _ -omit [StandardBorelSpace Ω] [Nonempty Ω] in lemma some_sum_eq_zero [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 < gap ν a) (n C : ℕ) @@ -610,7 +463,6 @@ lemma some_sum_eq_zero [Nonempty (Fin K)] · rw [h_arm] gcongr -omit [StandardBorelSpace Ω] [Nonempty Ω] in lemma pullCount_ae_le_add_two [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 < gap ν a) diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index c1b48a8c..128c0e8b 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -128,6 +128,41 @@ lemma exists_pullCount_eq_of_le (hnm : t ≤ pullCount A a (n + 1) ω) (ht : t refine lt_of_lt_of_le ?_ hnm exact pullCount_lt_of_forall_ne h_contra ht +lemma pullCount_le_add [Nonempty α] (a : α) (n C : ℕ) (ω : Ω) : + pullCount A a n ω ≤ C + 1 + + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s := by + rw [pullCount_eq_sum] + calc ∑ s ∈ range n, if A s ω = a then 1 else 0 + _ ≤ ∑ s ∈ range n, ({s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s + + {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s) := by + gcongr with s hs + simp [Set.indicator_apply] + grind + _ = ∑ s ∈ range n, {s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s + + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s := by + rw [Finset.sum_add_distrib] + _ ≤ C + 1 + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s := by + gcongr + have h_le n : ∑ s ∈ range n, {s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s ≤ + pullCount A a n ω := by + rw [pullCount_eq_sum] + gcongr with s hs + simp only [Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] + grind + induction n with + | zero => simp + | succ n hn => + rw [Finset.sum_range_succ] + rcases le_or_gt (pullCount A a n ω) C with h_pc | h_pc + · have hn' : ∑ s ∈ range n, {s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s ≤ C := + (h_le n).trans h_pc + grw [hn'] + gcongr + simp only [Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] + grind + · refine le_trans ?_ hn + simp [h_pc] + section Measurability @[fun_prop] @@ -172,6 +207,14 @@ lemma isPredictable_pullCount [MeasurableSingletonClass α] simp only [pullCount_zero] fun_prop +lemma integrable_pullCount [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (a : α) (n : ℕ) : + Integrable (fun ω ↦ (pullCount A a n ω : ℝ)) P := by + refine integrable_of_le_of_le (g₁ := 0) (g₂ := fun _ ↦ n) (by fun_prop) + (ae_of_all _ fun ω ↦ by simp) (ae_of_all _ fun ω ↦ ?_) (integrable_const _) (integrable_const _) + simp only [Nat.cast_le] + exact pullCount_le a n ω + end Measurability end PullCount From de5d855e74cee854fb80de062b3f87cd2740a9cb Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 8 Jan 2026 10:02:38 +0100 Subject: [PATCH 08/30] prove uniqueness of trajMeasure --- LeanBandits/Bandit/Bandit.lean | 41 +++++++++- LeanBandits/Bandit/SumRewards.lean | 11 ++- LeanBandits/ForMathlib/HasCondDistrib.lean | 2 + LeanBandits/ForMathlib/Traj.lean | 74 +++++++++++++++++-- LeanBandits/SequentialLearning/Algorithm.lean | 4 +- 5 files changed, 117 insertions(+), 15 deletions(-) diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index ea3e8601..39fcf209 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -365,10 +365,44 @@ lemma hasLaw_action_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKern variable [StandardBorelSpace R] [Nonempty R] -lemma hasCondDistrib_reward_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : +lemma hasCondDistrib_reward_zero [Countable α] + (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : HasCondDistrib (reward alg 0) (action alg 0) (stationaryEnv ν).ν0 (arrayMeasure ν) where condDistrib_eq := by - simp only [stationaryEnv_ν0, (hasLaw_action_zero alg ν).map_eq, reward_zero] + -- simp only [stationaryEnv_ν0] + -- refine condDistrib_ae_eq_of_measure_eq_compProd _ (by fun_prop) ?_ + -- rw [reward_zero] + -- simp only + -- have : (fun x ↦ (action alg 0 x, x.2 0 (action alg 0 x))) = + -- (fun p ↦ (p.2, p.1.2 0 p.2)) ∘ (fun x ↦ (x, action alg 0 x)) := rfl + + refine (condDistrib_ae_eq_cond (by fun_prop) (by fun_prop)).trans ?_ + rw [Filter.EventuallyEq, ae_iff_of_countable] + intro a ha + simp only [stationaryEnv_ν0, reward_zero] + calc ((arrayMeasure ν)[|action alg 0 ⁻¹' {a}]).map (fun ω ↦ ω.2 0 (action alg 0 ω)) + _ = ((arrayMeasure ν)[|action alg 0 ⁻¹' {a}]).map (fun ω ↦ ω.2 0 a) := by + refine Measure.map_congr + (ae_cond_of_forall_mem ((measurableSet_singleton _).preimage (by fun_prop)) ?_) + intro x hx + simp only [Set.mem_preimage, Set.mem_singleton_iff] at hx + simp [hx] + _ = ν a := sorry + +lemma hasLaw_and_hasCondDistrib (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] + (n : ℕ) : + HasLaw (hist alg · n) ((Bandit.trajMeasure alg ν).map (fun ω (i : Iic n) ↦ ω i)) + (arrayMeasure ν) ∧ + HasCondDistrib (hist alg · (n + 1) ⟨n + 1, by simp⟩) (hist alg · n) + (Bandit.stepKernel alg ν n) (arrayMeasure ν) := by + induction n with + | zero => sorry + | succ n hn => + have h1 : HasLaw (fun x ↦ hist alg x (n + 1)) (Measure.map (fun ω i ↦ ω ↑i) + (Bandit.trajMeasure alg ν)) (arrayMeasure ν) := by + have h_law := HasLaw.prod_of_hasCondDistrib hn.1 hn.2 + sorry + refine ⟨h1, ?_⟩ sorry lemma hasCondDistrib_altStep' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] @@ -405,7 +439,8 @@ lemma hasCondDistrib_reward (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovK simp only [stationaryEnv_feedback] sorry -lemma isAlgEnvSeq_arrayMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : +lemma isAlgEnvSeq_arrayMeasure [Countable α] + (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : IsAlgEnvSeq (action alg) (reward alg) alg (stationaryEnv ν) (arrayMeasure ν) where hasLaw_action_zero := hasLaw_action_zero alg ν hasCondDistrib_reward_zero := hasCondDistrib_reward_zero alg ν diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 4b382d95..0e54bbc7 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -424,7 +424,8 @@ lemma _root_.Learning.IsAlgEnvSeq.law_pullCount_sumRewards_unique exact fun n ↦ Measurable.prodMk (hA n) (hR n) -- this is what we will use for UCB -lemma prob_pullCount_prod_sumRewards_mem_le (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) +lemma prob_pullCount_prod_sumRewards_mem_le [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {s : Set (ℕ × ℝ)} [DecidablePred (· ∈ Prod.fst '' s)] (hs : MeasurableSet s) : P {ω | (pullCount A a n ω, sumRewards A R a n ω) ∈ s} ≤ ∑ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), @@ -445,7 +446,8 @@ lemma prob_pullCount_prod_sumRewards_mem_le (h : IsAlgEnvSeq A R alg (stationary Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := ArrayModel.prob_pullCount_prod_sumRewards_mem_le a n hs -lemma prob_pullCount_mem_and_sumRewards_mem_le (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) +lemma prob_pullCount_mem_and_sumRewards_mem_le [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {s : Set ℕ} [DecidablePred (· ∈ s)] (hs : MeasurableSet s) {B : Set ℝ} (hB : MeasurableSet B) : P {ω | pullCount A a n ω ∈ s ∧ sumRewards A R a n ω ∈ B} ≤ ∑ k ∈ (range (n + 1)).filter (· ∈ s), @@ -462,7 +464,7 @@ lemma prob_pullCount_mem_and_sumRewards_mem_le (h : IsAlgEnvSeq A R alg (station exists_eq_right, mem_filter, mem_range] at hk simp [hk.2.1] -lemma todo (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) +lemma todo [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {B : Set ℝ} (hB : MeasurableSet B) : P (sumRewards A R a n ⁻¹' B) ≤ ∑ k ∈ range (n + 1), Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := by @@ -470,7 +472,8 @@ lemma todo (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) have h_le := prob_pullCount_mem_and_sumRewards_mem_le h .univ hB (a := a) (n := n) simpa using h_le -lemma prob_pullCount_eq_and_sumRewards_mem_le (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) +lemma prob_pullCount_eq_and_sumRewards_mem_le [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {m : ℕ} (hm : m ≤ n) {B : Set ℝ} (hB : MeasurableSet B) : P {ω | pullCount A a n ω = m ∧ sumRewards A R a n ω ∈ B} ≤ Bandit.streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by diff --git a/LeanBandits/ForMathlib/HasCondDistrib.lean b/LeanBandits/ForMathlib/HasCondDistrib.lean index 2f76fe2e..fcec9c6f 100644 --- a/LeanBandits/ForMathlib/HasCondDistrib.lean +++ b/LeanBandits/ForMathlib/HasCondDistrib.lean @@ -26,6 +26,8 @@ structure HasCondDistrib (Y : α → Ω) (X : α → β) (κ : Kernel β Ω) aemeasurable_snd : AEMeasurable X μ := by fun_prop condDistrib_eq : condDistrib Y X μ =ᵐ[μ.map X] κ +attribute [fun_prop] HasCondDistrib.aemeasurable_fst HasCondDistrib.aemeasurable_snd + lemma hasCondDistrib_fst_prod {Y : α → Ω} {X : α → β} {κ : Kernel β Ω} {μ : Measure α} [IsFiniteMeasure μ] {ν : Measure γ} [IsProbabilityMeasure ν] diff --git a/LeanBandits/ForMathlib/Traj.lean b/LeanBandits/ForMathlib/Traj.lean index f4e84f1c..bbe7023c 100644 --- a/LeanBandits/ForMathlib/Traj.lean +++ b/LeanBandits/ForMathlib/Traj.lean @@ -1,6 +1,7 @@ import Mathlib.Probability.Kernel.IonescuTulcea.Traj import Mathlib.Probability.Kernel.CondDistrib -import LeanBandits.ForMathlib.CondDistrib +import Mathlib.Probability.Process.FiniteDimensionalLaws +import LeanBandits.ForMathlib.HasCondDistrib open Filter Finset Function MeasurableEquiv MeasurableSpace MeasureTheory Preorder ProbabilityTheory @@ -28,14 +29,75 @@ lemma traj_zero_map_eval_zero : rw [← Kernel.traj_map_frestrictLe, ← Kernel.map_comp_right _ (by fun_prop) (by fun_prop)] rfl +theorem hasLaw_Iic_of_forall_hasCondDistrib [∀ n, StandardBorelSpace (X n)] [∀ n, Nonempty (X n)] + {Y : (n : ℕ) → Ω → X n} (h0 : HasLaw (Y 0) μ₀ P) + (h_condDistrib : ∀ n, HasCondDistrib (Y (n + 1)) (fun ω ↦ fun i : Iic n ↦ Y i ω) (κ n) P) + (n : ℕ) : + HasLaw (fun ω (i : Iic n) ↦ Y i ω) + ((partialTraj κ 0 n) ∘ₘ (μ₀.map (MeasurableEquiv.piUnique _).symm)) P := by + induction n with + | zero => + simp only [piUnique_symm_apply, partialTraj_self, Measure.id_comp] + rw [← h0.map_eq, AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + constructor + · have h_meas := h0.aemeasurable + have : (fun ω (i : Iic 0) ↦ Y i ω) = (MeasurableEquiv.piUnique _).symm ∘ (Y 0) := by + ext ω i + simp only [piUnique_symm_apply, Function.comp_apply] + -- casting hell + sorry + rw [this] + exact AEMeasurable.comp_aemeasurable (by fun_prop) h_meas + · congr + ext ω i + simp only [Function.comp_apply] + -- same goal as above + sorry + | succ n hn => + specialize h_condDistrib n + have h_law := hn.prod_of_hasCondDistrib h_condDistrib + have : (fun ω (i : Iic (n + 1)) ↦ Y i ω) = + ((IicProdIoc n (n + 1)) ∘ (Prod.map id (MeasurableEquiv.piSingleton n))) ∘ + (fun ω ↦ (fun i : Iic n ↦ Y i ω, Y (n + 1) ω)) := by + ext ω i + simp only [Function.comp_apply] + simp only [_root_.IicProdIoc, piSingleton, MeasurableEquiv.coe_mk, Equiv.coe_fn_mk, + Prod.map_apply, id_eq, left_eq_dite_iff, not_le] + intro hi + grind + rw [this] + refine HasLaw.comp ?_ h_law + refine ⟨by fun_prop, ?_⟩ + rw [Measure.compProd_eq_comp_prod, partialTraj_succ_eq_comp (by simp), Measure.comp_assoc, + ← Measure.deterministic_comp_eq_map (by fun_prop), Measure.comp_assoc] + congr 1 + rw [← Kernel.comp_assoc] + congr + rw [Kernel.deterministic_comp_eq_map, partialTraj_succ_self, + Kernel.map_comp_right _ (by fun_prop) (by fun_prop), + ← Kernel.map_prod_map _ _ (by fun_prop) (by fun_prop)] + simp only [map_id] + +omit [IsProbabilityMeasure μ₀] in +lemma trajMeasure_map_frestrictLe (n : ℕ) : + (trajMeasure μ₀ κ).map (frestrictLe n) = + (partialTraj κ 0 n) ∘ₘ (μ₀.map (MeasurableEquiv.piUnique _).symm) := by + rw [trajMeasure, ← Measure.deterministic_comp_eq_map (by fun_prop), Measure.comp_assoc, + Kernel.deterministic_comp_eq_map, traj_map_frestrictLe] + +-- todo: switch to `HasLaw` /-- Uniqueness of `trajMeasure`. -/ theorem eq_trajMeasure [∀ n, StandardBorelSpace (X n)] [∀ n, Nonempty (X n)] {Y : (n : ℕ) → Ω → X n} (hY_meas : ∀ n, Measurable (Y n)) - (h0 : P.map (Y 0) = μ₀) - (h_condDistrib : ∀ n, - condDistrib (Y (n + 1)) (fun ω ↦ fun i : Iic n ↦ Y i ω) P - =ᵐ[P.map (fun ω ↦ fun i : Iic n ↦ Y i ω)] κ n) : + (h0 : HasLaw (Y 0) μ₀ P) + (h_condDistrib : ∀ n, HasCondDistrib (Y (n + 1)) (fun ω ↦ fun i : Iic n ↦ Y i ω) (κ n) P) : P.map (fun ω n ↦ Y n ω) = trajMeasure μ₀ κ := by - sorry + refine IsProjectiveLimit.unique (P := fun (J : Finset ℕ) ↦ P.map (fun ω (i : J) ↦ Y i ω)) ?_ ?_ + · exact isProjectiveLimit_map (by fun_prop) + rw [isProjectiveLimit_nat_iff] + swap; · exact isProjectiveMeasureFamily_map_restrict (by fun_prop) + intro n + rw [(hasLaw_Iic_of_forall_hasCondDistrib h0 h_condDistrib n).map_eq, + trajMeasure_map_frestrictLe] end ProbabilityTheory.Kernel diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index 5f7e2b92..5e965465 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -228,8 +228,8 @@ theorem eq_trajMeasure_of_isAlgEnvSeq (h : IsAlgEnvSeq A₁ R₁ alg env P) : have hR := h.measurable_R n fun_prop · simp only - exact h.hasLaw_step_zero.map_eq - · exact (h.hasCondDistrib_step n).condDistrib_eq + exact h.hasLaw_step_zero + · exact h.hasCondDistrib_step n theorem isAlgEnvSeq_unique (h1 : IsAlgEnvSeq A₁ R₁ alg env P) (h2 : IsAlgEnvSeq A₂ R₂ alg env P') : From a7dd567f80c6b3feb2079b519fdc292f91c8f094 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 8 Jan 2026 10:05:39 +0100 Subject: [PATCH 09/30] finish uniqueness proof --- LeanBandits/ForMathlib/Traj.lean | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/LeanBandits/ForMathlib/Traj.lean b/LeanBandits/ForMathlib/Traj.lean index bbe7023c..9046c137 100644 --- a/LeanBandits/ForMathlib/Traj.lean +++ b/LeanBandits/ForMathlib/Traj.lean @@ -44,15 +44,15 @@ theorem hasLaw_Iic_of_forall_hasCondDistrib [∀ n, StandardBorelSpace (X n)] [ have : (fun ω (i : Iic 0) ↦ Y i ω) = (MeasurableEquiv.piUnique _).symm ∘ (Y 0) := by ext ω i simp only [piUnique_symm_apply, Function.comp_apply] - -- casting hell - sorry + rw [Unique.eq_default i] + simp [uniqueElim_default, coe_default_Iic_zero] rw [this] exact AEMeasurable.comp_aemeasurable (by fun_prop) h_meas · congr ext ω i simp only [Function.comp_apply] - -- same goal as above - sorry + rw [Unique.eq_default i] + simp [uniqueElim_default, coe_default_Iic_zero] | succ n hn => specialize h_condDistrib n have h_law := hn.prod_of_hasCondDistrib h_condDistrib From aed649232bab1d5174e9e44796755d609566e474 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 8 Jan 2026 14:53:08 +0100 Subject: [PATCH 10/30] condDistrib progess --- LeanBandits/Bandit/Bandit.lean | 147 ++++++++++++++++-- LeanBandits/Bandit/SumRewards.lean | 4 +- LeanBandits/ForMathlib/Traj.lean | 53 +++++-- .../SequentialLearning/FiniteActions.lean | 4 + 4 files changed, 181 insertions(+), 27 deletions(-) diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 39fcf209..67f41988 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -275,7 +275,8 @@ lemma hist_eq [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) (n : | zero => ext i : 1 simp only [hist] - sorry + rw [Unique.eq_default i] + simp [coe_default_Iic_zero] | succ n hn => ext i : 1 by_cases hin : i ≤ n @@ -285,14 +286,46 @@ lemma hist_eq [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) (n : simp_rw [hn] · grind +lemma hist_add_one_eq_IicSuccProd' [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) + (n : ℕ) : + let a : α := algFunction alg n (hist alg ω n) (ω.1 (n + 1)) + hist alg ω (n + 1) = + (MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n).symm + (hist alg ω n, (a, ω.2 (pullCount' n (hist alg ω n) a) a)) := by + intro a + rw [hist_add_one] + ext i : 1 + simp only [Kernel.symm_IicSuccProd, MeasurableEquiv.prodCongr, MeasurableEquiv.refl_toEquiv, + MeasurableEquiv.piSingleton, eq_rec_constant, MeasurableEquiv.IicProdIoc, + MeasurableEquiv.trans_apply, MeasurableEquiv.coe_mk, Equiv.prodCongr_apply, Equiv.coe_refl, + Equiv.coe_fn_mk, Prod.map_apply, id_eq] + rfl + +lemma measurable_action_add_one' [DecidableEq α] {alg : Algorithm α R} + (n : ℕ) (h : Measurable (hist alg · n)) : + Measurable (fun x ↦ algFunction alg n (hist alg x n) (x.1 (n + 1))) := by fun_prop + +lemma measurable_pullCount'_action_add_one [DecidableEq α] [Countable α] {alg : Algorithm α R} + (n : ℕ) (h_hist : Measurable (hist alg · n)) : + Measurable (fun x ↦ + pullCount' n (hist alg x n) (algFunction alg n (hist alg x n) (x.1 (n + 1)))) := by + have h_alg_meas : Measurable (fun x ↦ algFunction alg n (hist alg x n) (x.1 (n + 1))) := + measurable_action_add_one' n h_hist + exact (measurable_uncurry_pullCount' (α := α) n).comp (h_hist.prodMk h_alg_meas) + @[fun_prop] -lemma measurable_hist [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : +lemma measurable_hist [DecidableEq α] [Countable α] (alg : Algorithm α R) (n : ℕ) : Measurable (fun ω ↦ hist alg ω n) := by induction n with | zero => simp_rw [hist_zero, measurable_pi_iff] refine fun _ ↦ Measurable.prodMk (by fun_prop) ?_ - sorry + unfold probSpace + change Measurable ((fun x : α × ((ℕ → I) × (ℕ → α → R)) ↦ x.2.2 0 x.1) ∘ + (fun x : (ℕ → I) × (ℕ → α → R) ↦ (initAlgFunction alg (x.1 0), x))) + have : Measurable (fun x : α × ((ℕ → I) × (ℕ → α → R)) ↦ x.2.2 0 x.1) := + measurable_from_prod_countable_right fun p ↦ by simp only; fun_prop + exact Measurable.comp (by fun_prop) (Measurable.prodMk (by fun_prop) (by fun_prop)) | succ n hn => refine measurable_pi_iff.mpr fun i ↦ ?_ by_cases hin : i ≤ n @@ -301,7 +334,14 @@ lemma measurable_hist [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : exact hn ⟨i.1, by simp [hin]⟩ · simp only [hist, hin, ↓reduceDIte] refine Measurable.prodMk (by fun_prop) ?_ - sorry + change Measurable ((fun (x : (ℕ → α → R) × ℕ × α) ↦ x.1 x.2.1 x.2.2) ∘ + (fun x ↦ (x.2, pullCount' n (hist alg x n) (algFunction alg n (hist alg x n) (x.1 (n + 1))), + (algFunction alg n (hist alg x n) (x.1 (n + 1)))))) + have h1 : Measurable (fun (x : (ℕ → α → R) × ℕ × α) ↦ x.1 x.2.1 x.2.2) := + measurable_from_prod_countable_left fun p : ℕ × α ↦ (by simp only; fun_prop) + refine Measurable.comp (by fun_prop) (Measurable.prodMk (by fun_prop) ?_) + refine Measurable.prodMk ?_ (by fun_prop) + exact measurable_pullCount'_action_add_one n hn noncomputable def action [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : α := @@ -319,7 +359,7 @@ lemma action_add_one_eq [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : simp only [add_le_iff_nonpos_right, nonpos_iff_eq_zero, one_ne_zero, ↓reduceDIte] @[fun_prop] -lemma measurable_action [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : +lemma measurable_action [DecidableEq α] [Countable α] (alg : Algorithm α R) (n : ℕ) : Measurable (action alg n) := by unfold action; fun_prop noncomputable @@ -331,6 +371,12 @@ lemma reward_zero [DecidableEq α] (alg : Algorithm α R) : ext simp [reward, hist_zero, action_zero] +lemma reward_add_one [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : + reward alg (n + 1) = + fun ω ↦ ω.2 (pullCount' n (hist alg ω n) (action alg (n + 1) ω)) (action alg (n + 1) ω) := by + ext ω + simp [reward, hist_add_one, action_add_one_eq] + lemma reward_eq [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : reward alg n = fun ω ↦ ω.2 (pullCount (action alg) (action alg n ω) n ω) (action alg n ω) := by cases n with @@ -345,10 +391,17 @@ lemma reward_eq [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : rfl @[fun_prop] -lemma measurable_reward [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : +lemma measurable_reward [DecidableEq α] [Countable α] (alg : Algorithm α R) (n : ℕ) : Measurable (reward alg n) := by unfold reward; fun_prop -variable [DecidableEq α] +lemma hist_add_one_eq_IicSuccProd [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) + (n : ℕ) : + hist alg ω (n + 1) = + (MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n).symm + (hist alg ω n, (action alg (n + 1) ω, reward alg (n + 1) ω)) := by + rw [hist_add_one_eq_IicSuccProd', reward_add_one, action_add_one_eq] + +variable [DecidableEq α] [Countable α] lemma hasLaw_action_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : HasLaw (action alg 0) alg.p0 (arrayMeasure ν) where @@ -389,20 +442,91 @@ lemma hasCondDistrib_reward_zero [Countable α] simp [hx] _ = ν a := sorry +lemma indepFun_fst_add_one_hist (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + IndepFun (fun ω ↦ ω.1 (n + 1)) (hist alg · n) (arrayMeasure ν) := by + sorry + +lemma hasCondDistrib_action' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + HasCondDistrib (action alg (n + 1)) (hist alg · n) (alg.policy n) (arrayMeasure ν) := by + rw [action_add_one_eq] + have h_fun ω := algFunction_map alg n (hist alg ω n) + refine ⟨by fun_prop, by fun_prop, ?_⟩ + refine condDistrib_ae_eq_of_measure_eq_compProd _ (by fun_prop) ?_ + have h_indep : (arrayMeasure ν).map (fun ω ↦ (ω.1 (n + 1), hist alg ω n)) = + (ℙ).prod ((arrayMeasure ν).map (hist alg · n)) := by + have h_indep' := indepFun_fst_add_one_hist alg ν n + rw [indepFun_iff_map_prod_eq_prod_map_map (by fun_prop) (by fun_prop)] at h_indep' + rw [h_indep'] + congr + simp only [arrayMeasure] + calc ((Measure.infinitePi fun x ↦ ℙ).prod (Bandit.streamMeasure ν)).map (fun ω ↦ ω.1 (n + 1)) + _ = (Measure.infinitePi fun x ↦ ℙ).map (Function.eval (n + 1)) := by + nth_rw 2 [← Measure.fst_prod (μ := Measure.infinitePi fun x ↦ ℙ) + (ν := Bandit.streamMeasure ν)] + rw [Measure.fst, Measure.map_map (by fun_prop) (by fun_prop)] + rfl + _ = ℙ := by rw [Measure.infinitePi_map_eval] + have : (fun x ↦ (hist alg x n, algFunction alg n (hist alg x n) (x.1 (n + 1)))) = + (fun p ↦ (p.2, algFunction alg n (p.2) (p.1))) ∘ (fun x ↦ (x.1 (n + 1), hist alg x n)) := rfl + rw [this, ← Measure.map_map (by fun_prop) (by fun_prop), h_indep] + have : (ℙ : Measure I).prod ((arrayMeasure ν).map (hist alg · n)) = + ((Kernel.const _ ℙ) ×ₖ Kernel.id) ∘ₘ ((arrayMeasure ν).map (hist alg · n)) := by + have h := Measure.compProd_const (μ := (arrayMeasure ν).map (hist alg · n)) + (ν := (ℙ : Measure I)) + rw [Measure.compProd_eq_comp_prod] at h + rw [← Measure.prod_swap, ← h, ← Measure.deterministic_comp_eq_map (by fun_prop), + Measure.comp_assoc, ← Kernel.swap, Kernel.swap_prod] + rw [this, ← Measure.deterministic_comp_eq_map (by fun_prop), + ← Measure.deterministic_comp_eq_map (by fun_prop), Measure.compProd_eq_comp_prod, + Measure.comp_assoc, Measure.comp_assoc, Measure.comp_assoc] + congr 2 + ext ω : 1 + simp only [Kernel.deterministic_comp_eq_map, Kernel.comp_deterministic_eq_comap, Kernel.coe_comap, + Function.comp_apply] + rw [Kernel.map_apply _ (by fun_prop), Kernel.prod_apply, Kernel.const_apply, Kernel.id_apply, + Kernel.prod_apply, Kernel.id_apply, ← h_fun] + calc (((ℙ).prod (Measure.dirac (hist alg ω n)))).map (fun p ↦ (p.2, algFunction alg n p.2 p.1)) + _ = (((ℙ).prod (Measure.dirac (hist alg ω n))).map Prod.swap).map + (fun p ↦ (p.1, algFunction alg n p.1 p.2)) := by + rw [Measure.map_map (by fun_prop) (by fun_prop)] + rfl + _ = ((Measure.dirac (hist alg ω n)).prod ℙ).map (fun p ↦ (p.1, algFunction alg n p.1 p.2)) := by + rw [Measure.prod_swap] + _ = (Measure.dirac (hist alg ω n)).prod ((ℙ).map (algFunction alg n (hist alg ω n))) := by + ext s hs + rw [Measure.map_apply (by fun_prop) hs, Measure.prod_apply, lintegral_dirac, Measure.prod_apply, + lintegral_dirac, Measure.map_apply (by fun_prop)] + · congr + · exact hs.preimage (by fun_prop) + · exact hs + · exact hs.preimage (by fun_prop) + +lemma hasCondDistrib_reward' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] + (n : ℕ) : + HasCondDistrib (reward alg (n + 1)) (fun ω ↦ (hist alg ω n, action alg (n + 1) ω)) + (ν.prodMkLeft _) (arrayMeasure ν) := by + rw [reward_eq] + refine ⟨?_, by fun_prop, ?_⟩ + · sorry + refine condDistrib_ae_eq_of_measure_eq_compProd _ ?_ ?_ + · sorry + sorry + lemma hasLaw_and_hasCondDistrib (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - HasLaw (hist alg · n) ((Bandit.trajMeasure alg ν).map (fun ω (i : Iic n) ↦ ω i)) + HasLaw (hist alg · n) ((Bandit.trajMeasure alg ν).map (Preorder.frestrictLe n)) (arrayMeasure ν) ∧ HasCondDistrib (hist alg · (n + 1) ⟨n + 1, by simp⟩) (hist alg · n) (Bandit.stepKernel alg ν n) (arrayMeasure ν) := by induction n with | zero => sorry | succ n hn => - have h1 : HasLaw (fun x ↦ hist alg x (n + 1)) (Measure.map (fun ω i ↦ ω ↑i) - (Bandit.trajMeasure alg ν)) (arrayMeasure ν) := by + have h1 : HasLaw (fun x ↦ hist alg x (n + 1)) + ((Bandit.trajMeasure alg ν).map (Preorder.frestrictLe (n + 1))) (arrayMeasure ν) := by have h_law := HasLaw.prod_of_hasCondDistrib hn.1 hn.2 sorry refine ⟨h1, ?_⟩ + simp_rw [hist_add_one_eq_IicSuccProd _ _ (n + 1)] sorry lemma hasCondDistrib_altStep' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] @@ -439,8 +563,7 @@ lemma hasCondDistrib_reward (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovK simp only [stationaryEnv_feedback] sorry -lemma isAlgEnvSeq_arrayMeasure [Countable α] - (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : +lemma isAlgEnvSeq_arrayMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : IsAlgEnvSeq (action alg) (reward alg) alg (stationaryEnv ν) (arrayMeasure ν) where hasLaw_action_zero := hasLaw_action_zero alg ν hasCondDistrib_reward_zero := hasCondDistrib_reward_zero alg ν diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 0e54bbc7..ce143813 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -44,7 +44,8 @@ namespace Bandits namespace ArrayModel -variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [StandardBorelSpace α] [Nonempty α] +variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [Countable α] + [StandardBorelSpace α] [Nonempty α] {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] local notation "A" => action alg @@ -111,6 +112,7 @@ lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount (n : ℕ) : convert identDistrib_pullCount_prod_sum_Icc_rewardByCount' n using 2 with ω rotate_left · infer_instance + · infer_instance ext a : 1 congr 1 let e : Icc 1 (pullCount A a n ω) ≃ range (pullCount A a n ω) := diff --git a/LeanBandits/ForMathlib/Traj.lean b/LeanBandits/ForMathlib/Traj.lean index 9046c137..93333263 100644 --- a/LeanBandits/ForMathlib/Traj.lean +++ b/LeanBandits/ForMathlib/Traj.lean @@ -29,6 +29,30 @@ lemma traj_zero_map_eval_zero : rw [← Kernel.traj_map_frestrictLe, ← Kernel.map_comp_right _ (by fun_prop) (by fun_prop)] rfl +def _root_.MeasurableEquiv.IicSuccProd (X : ℕ → Type*) [∀ n, MeasurableSpace (X n)] (n : ℕ) : + MeasurableEquiv (Π i : Iic (n + 1), X i) ((Π i : Iic n, X i) × X (n + 1)) := + (MeasurableEquiv.IicProdIoc (Nat.le_succ n)).symm.trans + (MeasurableEquiv.prodCongr (MeasurableEquiv.refl _) (MeasurableEquiv.piSingleton n).symm) + +lemma symm_IicSuccProd (n : ℕ) : + (MeasurableEquiv.IicSuccProd X n).symm = + (MeasurableEquiv.prodCongr (MeasurableEquiv.refl _) (MeasurableEquiv.piSingleton n)).trans + (MeasurableEquiv.IicProdIoc (Nat.le_succ n)) := rfl + +@[simp] +lemma MeasurableEquiv.IicSuccProd_apply (n : ℕ) (h : Π i : Iic (n + 1), X i) : + MeasurableEquiv.IicSuccProd X n h = (fun i : Iic n ↦ h ⟨i.1, by grind⟩, h ⟨n + 1, by simp⟩) := + rfl + +lemma MeasurableEquiv.coe_prodCongr {α β γ δ : Type*} + {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {mδ : MeasurableSpace δ} + (e₁ : MeasurableEquiv α β) (e₂ : MeasurableEquiv γ δ) : + (MeasurableEquiv.prodCongr e₁ e₂ : (α × γ) → (β × δ)) = Prod.map e₁ e₂ := rfl + +lemma MeasurableEquiv.coe_refl {α : Type*} {mα : MeasurableSpace α} : + (MeasurableEquiv.refl α : α → α) = id := rfl + theorem hasLaw_Iic_of_forall_hasCondDistrib [∀ n, StandardBorelSpace (X n)] [∀ n, Nonempty (X n)] {Y : (n : ℕ) → Ω → X n} (h0 : HasLaw (Y 0) μ₀ P) (h_condDistrib : ∀ n, HasCondDistrib (Y (n + 1)) (fun ω ↦ fun i : Iic n ↦ Y i ω) (κ n) P) @@ -45,38 +69,39 @@ theorem hasLaw_Iic_of_forall_hasCondDistrib [∀ n, StandardBorelSpace (X n)] [ ext ω i simp only [piUnique_symm_apply, Function.comp_apply] rw [Unique.eq_default i] - simp [uniqueElim_default, coe_default_Iic_zero] + simp [coe_default_Iic_zero] rw [this] exact AEMeasurable.comp_aemeasurable (by fun_prop) h_meas · congr ext ω i simp only [Function.comp_apply] rw [Unique.eq_default i] - simp [uniqueElim_default, coe_default_Iic_zero] + simp [coe_default_Iic_zero] | succ n hn => specialize h_condDistrib n have h_law := hn.prod_of_hasCondDistrib h_condDistrib have : (fun ω (i : Iic (n + 1)) ↦ Y i ω) = - ((IicProdIoc n (n + 1)) ∘ (Prod.map id (MeasurableEquiv.piSingleton n))) ∘ + (MeasurableEquiv.IicSuccProd X n).symm ∘ (fun ω ↦ (fun i : Iic n ↦ Y i ω, Y (n + 1) ω)) := by - ext ω i - simp only [Function.comp_apply] - simp only [_root_.IicProdIoc, piSingleton, MeasurableEquiv.coe_mk, Equiv.coe_fn_mk, - Prod.map_apply, id_eq, left_eq_dite_iff, not_le] - intro hi - grind + suffices (MeasurableEquiv.IicSuccProd X n) ∘ (fun ω (i : Iic (n + 1)) ↦ Y i ω) = + (fun ω ↦ (fun i : Iic n ↦ Y i ω, Y (n + 1) ω)) by + rw [← this, ← Function.comp_assoc, MeasurableEquiv.symm_comp_self] + simp + ext ω : 1 + simp rw [this] - refine HasLaw.comp ?_ h_law - refine ⟨by fun_prop, ?_⟩ + refine HasLaw.comp ⟨by fun_prop, ?_⟩ h_law rw [Measure.compProd_eq_comp_prod, partialTraj_succ_eq_comp (by simp), Measure.comp_assoc, ← Measure.deterministic_comp_eq_map (by fun_prop), Measure.comp_assoc] congr 1 rw [← Kernel.comp_assoc] congr - rw [Kernel.deterministic_comp_eq_map, partialTraj_succ_self, - Kernel.map_comp_right _ (by fun_prop) (by fun_prop), + rw [Kernel.deterministic_comp_eq_map, partialTraj_succ_self, symm_IicSuccProd] + rw [MeasurableEquiv.coe_trans, MeasurableEquiv.coe_prodCongr] + rw [Kernel.map_comp_right _ (by fun_prop) (by fun_prop), ← Kernel.map_prod_map _ _ (by fun_prop) (by fun_prop)] - simp only [map_id] + congr + simp [MeasurableEquiv.coe_refl] omit [IsProbabilityMeasure μ₀] in lemma trajMeasure_map_frestrictLe (n : ℕ) : diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index 128c0e8b..2291b919 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -184,6 +184,10 @@ lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop +lemma measurable_uncurry_pullCount' [MeasurableSingletonClass α] [Countable α] (n : ℕ) : + Measurable (fun p : (Iic n → α × R) × α ↦ pullCount' n p.1 p.2) := by + refine measurable_from_prod_countable_left fun a ↦ measurable_pullCount' n a + lemma adapted_pullCount_add_one' [MeasurableSingletonClass α] (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (n : ℕ) : Measurable[IsAlgEnvSeq.filtration hA hR' n] (pullCount A a (n + 1)) := by From 5cade77bb3fcbcf2353108e85714027411c08b8a Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 8 Jan 2026 16:31:40 +0100 Subject: [PATCH 11/30] progress --- LeanBandits/Bandit/Bandit.lean | 56 +++++++++++++++---- .../SequentialLearning/FiniteActions.lean | 8 ++- 2 files changed, 50 insertions(+), 14 deletions(-) diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 67f41988..d6a323f2 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -305,7 +305,11 @@ lemma measurable_action_add_one' [DecidableEq α] {alg : Algorithm α R} (n : ℕ) (h : Measurable (hist alg · n)) : Measurable (fun x ↦ algFunction alg n (hist alg x n) (x.1 (n + 1))) := by fun_prop -lemma measurable_pullCount'_action_add_one [DecidableEq α] [Countable α] {alg : Algorithm α R} +instance : MeasurableEq α := by + letI := upgradeStandardBorel α + infer_instance + +lemma measurable_pullCount'_action_add_one [DecidableEq α] {alg : Algorithm α R} (n : ℕ) (h_hist : Measurable (hist alg · n)) : Measurable (fun x ↦ pullCount' n (hist alg x n) (algFunction alg n (hist alg x n) (x.1 (n + 1)))) := by @@ -320,7 +324,6 @@ lemma measurable_hist [DecidableEq α] [Countable α] (alg : Algorithm α R) (n | zero => simp_rw [hist_zero, measurable_pi_iff] refine fun _ ↦ Measurable.prodMk (by fun_prop) ?_ - unfold probSpace change Measurable ((fun x : α × ((ℕ → I) × (ℕ → α → R)) ↦ x.2.2 0 x.1) ∘ (fun x : (ℕ → I) × (ℕ → α → R) ↦ (initAlgFunction alg (x.1 0), x))) have : Measurable (fun x : α × ((ℕ → I) × (ℕ → α → R)) ↦ x.2.2 0 x.1) := @@ -418,17 +421,36 @@ lemma hasLaw_action_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKern variable [StandardBorelSpace R] [Nonempty R] -lemma hasCondDistrib_reward_zero [Countable α] - (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : +omit [Nonempty α] [StandardBorelSpace α] [DecidableEq α] [Countable α] + [StandardBorelSpace R] [Nonempty R] in +lemma indepFun_fst_snd (ν : Kernel α R) [IsMarkovKernel ν] : + IndepFun Prod.fst Prod.snd (arrayMeasure ν) := + indepFun_prod measurable_id measurable_id + +omit [Nonempty α] [StandardBorelSpace α] [DecidableEq α] [Countable α] + [StandardBorelSpace R] [Nonempty R] in +lemma indepFun_fst_zero_snd_zero_action (ν : Kernel α R) [IsMarkovKernel ν] (a : α) : + IndepFun (fun ω ↦ ω.1 0) (fun ω ↦ ω.2 0 a) (arrayMeasure ν) := + indepFun_prod (X := fun ω : ℕ → I ↦ ω 0) (Y := fun ω : ℕ → α → R ↦ ω 0 a) + (by fun_prop) (by fun_prop) + +omit [Nonempty α] [StandardBorelSpace α] [DecidableEq α] [Countable α] + [StandardBorelSpace R] [Nonempty R] in +lemma map_snd_apply_arrayMeasure {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) (a : α) : + (arrayMeasure ν).map (fun ω ↦ ω.2 n a) = ν a := by + calc (arrayMeasure ν).map (fun ω ↦ ω.2 n a) + _ = (arrayMeasure ν).snd.map (fun ω ↦ ω n a) := by + rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)] + rfl + _ = ν a := by + rw [arrayMeasure, Measure.snd_prod, Bandit.streamMeasure] + have : (fun ω ↦ ω n a) = (fun h : α → R ↦ h a) ∘ (fun ω : ℕ → α → R ↦ ω n) := rfl + rw [this, ← Measure.map_map (by fun_prop) (by fun_prop), Measure.infinitePi_map_eval, + Measure.infinitePi_map_eval] + +lemma hasCondDistrib_reward_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : HasCondDistrib (reward alg 0) (action alg 0) (stationaryEnv ν).ν0 (arrayMeasure ν) where condDistrib_eq := by - -- simp only [stationaryEnv_ν0] - -- refine condDistrib_ae_eq_of_measure_eq_compProd _ (by fun_prop) ?_ - -- rw [reward_zero] - -- simp only - -- have : (fun x ↦ (action alg 0 x, x.2 0 (action alg 0 x))) = - -- (fun p ↦ (p.2, p.1.2 0 p.2)) ∘ (fun x ↦ (x, action alg 0 x)) := rfl - refine (condDistrib_ae_eq_cond (by fun_prop) (by fun_prop)).trans ?_ rw [Filter.EventuallyEq, ae_iff_of_countable] intro a ha @@ -440,7 +462,17 @@ lemma hasCondDistrib_reward_zero [Countable α] intro x hx simp only [Set.mem_preimage, Set.mem_singleton_iff] at hx simp [hx] - _ = ν a := sorry + _ = ν a := by + rw [cond_of_indepFun] + · exact map_snd_apply_arrayMeasure 0 a + · have : (fun ω ↦ ω.1 0) ⟂ᵢ[arrayMeasure ν] fun ω ↦ ω.2 0 a := + indepFun_fst_zero_snd_zero_action ν a + rw [action_zero] + exact this.comp (φ := initAlgFunction alg) (by fun_prop) measurable_id + · fun_prop + · fun_prop + · simp + · rwa [Measure.map_apply (by fun_prop) (by simp)] at ha lemma indepFun_fst_add_one_hist (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : IndepFun (fun ω ↦ ω.1 (n + 1)) (hist alg · n) (arrayMeasure ν) := by diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index 2291b919..7c180d4f 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -184,9 +184,13 @@ lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop -lemma measurable_uncurry_pullCount' [MeasurableSingletonClass α] [Countable α] (n : ℕ) : +lemma measurable_uncurry_pullCount' [MeasurableSingletonClass α] [MeasurableEq α] (n : ℕ) : Measurable (fun p : (Iic n → α × R) × α ↦ pullCount' n p.1 p.2) := by - refine measurable_from_prod_countable_left fun a ↦ measurable_pullCount' n a + simp_rw [pullCount'_eq_sum] + have h_meas s : Measurable (fun h : (Iic n → α × R) × α ↦ if (h.1 s).1 = h.2 then 1 else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact measurableSet_eq_fun (by fun_prop) (by fun_prop) + fun_prop lemma adapted_pullCount_add_one' [MeasurableSingletonClass α] (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (n : ℕ) : From d84d16ee6af7e0c75a24fd957d58e843a36d9590 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 8 Jan 2026 16:35:30 +0100 Subject: [PATCH 12/30] delete unused lemmas --- LeanBandits/Bandit/Bandit.lean | 48 +++++++--------------------------- 1 file changed, 9 insertions(+), 39 deletions(-) diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index d6a323f2..98ebe561 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -544,56 +544,26 @@ lemma hasCondDistrib_reward' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkov · sorry sorry -lemma hasLaw_and_hasCondDistrib (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] - (n : ℕ) : - HasLaw (hist alg · n) ((Bandit.trajMeasure alg ν).map (Preorder.frestrictLe n)) - (arrayMeasure ν) ∧ - HasCondDistrib (hist alg · (n + 1) ⟨n + 1, by simp⟩) (hist alg · n) - (Bandit.stepKernel alg ν n) (arrayMeasure ν) := by - induction n with - | zero => sorry - | succ n hn => - have h1 : HasLaw (fun x ↦ hist alg x (n + 1)) - ((Bandit.trajMeasure alg ν).map (Preorder.frestrictLe (n + 1))) (arrayMeasure ν) := by - have h_law := HasLaw.prod_of_hasCondDistrib hn.1 hn.2 - sorry - refine ⟨h1, ?_⟩ - simp_rw [hist_add_one_eq_IicSuccProd _ _ (n + 1)] - sorry - -lemma hasCondDistrib_altStep' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] - (n : ℕ) : - HasCondDistrib (hist alg · (n + 1) ⟨n + 1, by simp⟩) (hist alg · n) - (Bandit.stepKernel alg ν n) (arrayMeasure ν) where - condDistrib_eq := by - simp only [Bandit.stepKernel, stepKernel, stationaryEnv_feedback] - sorry - -lemma hasCondDistrib_altStep (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] - (n : ℕ) : - HasCondDistrib (fun ω ↦ (action alg (n + 1) ω, reward alg (n + 1) ω)) +lemma hasCondDistrib_action (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + HasCondDistrib (action alg (n + 1)) (fun ω (i : Iic n) ↦ (action alg i ω, reward alg i ω)) - (Bandit.stepKernel alg ν n) (arrayMeasure ν) := by - convert hasCondDistrib_altStep' alg ν n with ω i + (alg.policy n) (arrayMeasure ν) := by + convert hasCondDistrib_action' alg ν n with ω i · simp only [action] rw [hist_eq _ _ n] · simp only [reward] rw [hist_eq _ _ n] -lemma hasCondDistrib_action (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - HasCondDistrib (action alg (n + 1)) - (fun ω (i : Iic n) ↦ (action alg i ω, reward alg i ω)) - (alg.policy n) (arrayMeasure ν) := by - convert HasCondDistrib.fst (hasCondDistrib_altStep alg ν n) - simp - lemma hasCondDistrib_reward (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : HasCondDistrib (reward alg (n + 1)) (fun ω ↦ (fun (i : Iic n) ↦ (action alg i ω, reward alg i ω), action alg (n + 1) ω)) ((stationaryEnv ν).feedback n) (arrayMeasure ν) := by - simp only [stationaryEnv_feedback] - sorry + convert hasCondDistrib_reward' alg ν n with ω i + · simp only [action] + rw [hist_eq _ _ n] + · simp only [reward] + rw [hist_eq _ _ n] lemma isAlgEnvSeq_arrayMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : IsAlgEnvSeq (action alg) (reward alg) alg (stationaryEnv ν) (arrayMeasure ν) where From 39ad41cee16dd09cc3cae1a00c21a83426ef16ac Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Tue, 13 Jan 2026 22:04:10 +0100 Subject: [PATCH 13/30] progress --- LeanBandits/Bandit/Bandit.lean | 531 +++++++++++++++++- LeanBandits/Bandit/RewardByCountMeasure.lean | 39 +- LeanBandits/ForMathlib/CondDistrib.lean | 114 ++++ LeanBandits/ForMathlib/CondIndepFun.lean | 71 +++ .../ForMathlib/KernelRepresentation.lean | 146 +++++ .../SequentialLearning/FiniteActions.lean | 16 +- 6 files changed, 858 insertions(+), 59 deletions(-) create mode 100644 LeanBandits/ForMathlib/CondIndepFun.lean create mode 100644 LeanBandits/ForMathlib/KernelRepresentation.lean diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 98ebe561..483e550c 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -3,7 +3,9 @@ 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, Paulo Rauber -/ +import LeanBandits.ForMathlib.CondIndepFun import LeanBandits.ForMathlib.IndepInfinitePi +import LeanBandits.ForMathlib.KernelRepresentation import LeanBandits.SequentialLearning.Deterministic import LeanBandits.SequentialLearning.StationaryEnv import LeanBandits.SequentialLearning.FiniteActions @@ -193,17 +195,11 @@ open unitInterval section Aux --- from Mathlib PR #30112 -theorem representation {α β : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} - [Nonempty β] [StandardBorelSpace β] - (κ : Kernel α β) [IsMarkovKernel κ] : - ∃ (f : α → I → β), Measurable (Function.uncurry f) ∧ ∀ a, volume.map (f a) = κ a := sorry - theorem representation_measure {β : Type*} {mβ : MeasurableSpace β} [Nonempty β] [StandardBorelSpace β] (μ : Measure β) [IsProbabilityMeasure μ] : ∃ (f : I → β), Measurable f ∧ volume.map f = μ := by - obtain ⟨f, hf_meas, hf_map⟩ := representation (Kernel.const Unit μ) + obtain ⟨f, hf_meas, hf_map⟩ := Kernel.representation (Kernel.const Unit μ) specialize hf_map ⟨⟩ exact ⟨f ⟨⟩, by fun_prop, by simpa⟩ @@ -215,6 +211,12 @@ def probSpace : Type _ := (ℕ → I) × (ℕ → α → R) instance {α R : Type*} [MeasurableSpace R] : MeasurableSpace (probSpace α R) := inferInstanceAs (MeasurableSpace ((ℕ → I) × (ℕ → α → R))) +instance {α R : Type*} [MeasurableSpace α] [Countable α] + [MeasurableSpace R] [StandardBorelSpace R] [Nonempty R] : + StandardBorelSpace (probSpace α R) := by + unfold probSpace + infer_instance + noncomputable def arrayMeasure (ν : Kernel α R) : Measure (probSpace α R) := (Measure.infinitePi fun _ ↦ volume).prod (Bandit.streamMeasure ν) @@ -238,16 +240,16 @@ lemma measurable_initAlgFunction (alg : Algorithm α R) : noncomputable def algFunction (alg : Algorithm α R) (n : ℕ) : (Iic n → α × R) → I → α := - (representation (alg.policy n)).choose + (Kernel.representation (alg.policy n)).choose lemma algFunction_map (alg : Algorithm α R) (n : ℕ) (h : Iic n → α × R) : volume.map (algFunction alg n h) = alg.policy n h := - (representation (alg.policy n)).choose_spec.2 h + (Kernel.representation (alg.policy n)).choose_spec.2 h @[fun_prop] lemma measurable_algFunction (alg : Algorithm α R) (n : ℕ) : Measurable (Function.uncurry (algFunction alg n)) := - (representation (alg.policy n)).choose_spec.1 + (Kernel.representation (alg.policy n)).choose_spec.1 noncomputable def hist [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) : (n : ℕ) → Iic n → α × R @@ -449,12 +451,12 @@ lemma map_snd_apply_arrayMeasure {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ Measure.infinitePi_map_eval] lemma hasCondDistrib_reward_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : - HasCondDistrib (reward alg 0) (action alg 0) (stationaryEnv ν).ν0 (arrayMeasure ν) where + HasCondDistrib (reward alg 0) (action alg 0) ν (arrayMeasure ν) where condDistrib_eq := by refine (condDistrib_ae_eq_cond (by fun_prop) (by fun_prop)).trans ?_ rw [Filter.EventuallyEq, ae_iff_of_countable] intro a ha - simp only [stationaryEnv_ν0, reward_zero] + simp only [reward_zero] calc ((arrayMeasure ν)[|action alg 0 ⁻¹' {a}]).map (fun ω ↦ ω.2 0 (action alg 0 ω)) _ = ((arrayMeasure ν)[|action alg 0 ⁻¹' {a}]).map (fun ω ↦ ω.2 0 a) := by refine Measure.map_congr @@ -474,10 +476,91 @@ lemma hasCondDistrib_reward_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMa · simp · rwa [Measure.map_apply (by fun_prop) (by simp)] at ha -lemma indepFun_fst_add_one_hist (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - IndepFun (fun ω ↦ ω.1 (n + 1)) (hist alg · n) (arrayMeasure ν) := by +lemma indepFun_fst_add_one_aux (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + (fun ω ↦ ω.1 (n + 1)) ⟂ᵢ[arrayMeasure ν] (fun ω ↦ (fun (i : Iic n) ↦ ω.1 i, ω.2)) := by + rw [indepFun_iff_map_prod_eq_prod_map_map (by fun_prop) (by fun_prop)] sorry +omit [StandardBorelSpace R] [Nonempty R] in +lemma measurable_hist_todo (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + Measurable[MeasurableSpace.comap (fun ω ↦ (fun (i : Iic n) ↦ ω.1 i, ω.2)) inferInstance] + (hist alg · n) := by + induction n with + | zero => + simp only [hist_zero] + have : (fun (ω : probSpace α R) (i : Iic 0) ↦ + (initAlgFunction alg (ω.1 0), ω.2 0 (initAlgFunction alg (ω.1 0)))) = + (fun (p : (Iic 0 → I) × (ℕ → α → R)) (i : Iic 0) ↦ (initAlgFunction alg (p.1 ⟨0, by simp⟩), + p.2 0 (initAlgFunction alg (p.1 ⟨0, by simp⟩)))) ∘ + (fun (ω : probSpace α R) ↦ (fun (i : Iic 0) ↦ ω.1 i, ω.2)) := rfl + rw [this] + have h_meas : Measurable (fun (p : (Iic 0 → I) × (ℕ → α → R)) (i : Iic 0) ↦ + (initAlgFunction alg (p.1 ⟨0, by simp⟩), + p.2 0 (initAlgFunction alg (p.1 ⟨0, by simp⟩)))) := by + rw [measurable_pi_iff] + intro i + refine Measurable.prodMk (by fun_prop) ?_ + change Measurable ((fun x : (α × (ℕ → α → R)) ↦ x.2 0 x.1) ∘ + (fun x : (Iic 0 → I) × (ℕ → α → R) ↦ (initAlgFunction alg (x.1 ⟨0, by simp⟩), x.2))) + have : Measurable (fun x : (α × (ℕ → α → R)) ↦ x.2 0 x.1) := + measurable_from_prod_countable_right fun p ↦ by simp only; fun_prop + exact this.comp (by fun_prop) + refine Measurable.comp h_meas ?_ + exact Measurable.of_comap_le le_rfl + | succ n hn => + simp_rw [hist_add_one_eq_IicSuccProd] + have h_hist : Measurable[MeasurableSpace.comap (fun ω ↦ (fun (i : Iic (n + 1)) ↦ ω.1 i, ω.2)) + inferInstance] (hist alg · n) := by + rw [measurable_iff_comap_le] at hn ⊢ + refine hn.trans ?_ + rw [← measurable_iff_comap_le] + have : (fun (ω : probSpace α R) ↦ (fun (i : Iic n) ↦ ω.1 i, ω.2)) = + (fun (p : (Iic (n + 1) → I) × (ℕ → α → R)) ↦ (fun (i : Iic n) ↦ p.1 ⟨i, by grind⟩, p.2)) ∘ + (fun (ω : probSpace α R) ↦ (fun (i : Iic (n + 1)) ↦ ω.1 i, ω.2)) := rfl + rw [this] + exact Measurable.comp (by fun_prop) (Measurable.of_comap_le le_rfl) + refine (MeasurableEquiv.measurable _).comp (Measurable.prodMk h_hist ?_) + have h_action : Measurable[MeasurableSpace.comap (fun ω ↦ (fun (i : Iic (n + 1)) ↦ ω.1 i, ω.2)) + inferInstance] (action alg (n + 1)) := by + rw [action_add_one_eq] + have : (fun ω ↦ algFunction alg n (hist alg ω n) (ω.1 (n + 1))) = + (Function.uncurry (algFunction alg n)) ∘ (fun ω ↦ (hist alg ω n, ω.1 (n + 1))) := rfl + rw [this] + refine (measurable_algFunction alg n).comp (h_hist.prodMk ?_) + have : (fun ω : probSpace α R ↦ ω.1 (n + 1)) = + (fun (p : (Iic (n + 1) → I) × (ℕ → α → R)) ↦ p.1 ⟨n + 1, by simp⟩) ∘ + (fun ω ↦ (fun (i : Iic (n + 1)) ↦ ω.1 i, ω.2)) := rfl + rw [this] + exact Measurable.comp (by fun_prop) (Measurable.of_comap_le le_rfl) + refine h_action.prodMk ?_ + rw [reward_add_one] + have : (fun ω ↦ ω.2 (pullCount' n (hist alg ω n) (action alg (n + 1) ω)) + (action alg (n + 1) ω)) = + (fun p : ((Iic (n + 1) → I) × (ℕ → α → R)) × + (Iic n → α × R) × α ↦ p.1.2 (pullCount' n p.2.1 p.2.2) p.2.2) ∘ + (fun ω ↦ ((fun i : Iic (n + 1) ↦ ω.1 i, ω.2), hist alg ω n, action alg (n + 1) ω)) := rfl + rw [this] + have h_meas : Measurable + (fun p : ((Iic (n + 1) → I) × (ℕ → α → R)) × (Iic n → α × R) × α ↦ + p.1.2 (pullCount' n p.2.1 p.2.2) p.2.2) := by + have : (fun p : ((Iic (n + 1) → I) × (ℕ → α → R)) × (Iic n → α × R) × α ↦ + p.1.2 (pullCount' n p.2.1 p.2.2) p.2.2) = + (fun (x : (ℕ → α → R) × ℕ × α) ↦ x.1 x.2.1 x.2.2) ∘ + (fun p : ((Iic (n + 1) → I) × (ℕ → α → R)) × (Iic n → α × R) × α ↦ + (p.1.2, pullCount' n p.2.1 p.2.2, p.2.2)) := rfl + rw [this] + refine Measurable.comp (measurable_from_prod_countable_left (m := inferInstance) fun p ↦ ?_) + ?_ + · simp only; fun_prop + refine Measurable.prodMk (by fun_prop) (Measurable.prodMk ?_ (by fun_prop)) + exact (measurable_uncurry_pullCount' (α := α) (mR := mR) n).comp (by fun_prop) + refine h_meas.comp (Measurable.prodMk ?_ (Measurable.prodMk h_hist h_action)) + exact Measurable.of_comap_le le_rfl + +lemma indepFun_fst_add_one_hist (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + IndepFun (fun ω ↦ ω.1 (n + 1)) (hist alg · n) (arrayMeasure ν) := + (indepFun_fst_add_one_aux alg ν n).of_measurable_right (measurable_hist_todo alg ν n) + lemma hasCondDistrib_action' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : HasCondDistrib (action alg (n + 1)) (hist alg · n) (alg.policy n) (arrayMeasure ν) := by rw [action_add_one_eq] @@ -533,16 +616,424 @@ lemma hasCondDistrib_action' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkov · exact hs · exact hs.preimage (by fun_prop) +omit [Countable α] [StandardBorelSpace R] [Nonempty R] in +lemma hist_congr (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) + {ω ω' : probSpace α R} + (hω1 : ∀ i ≤ n, ω.1 i = ω'.1 i) + (hω2 : ∀ i a, i < pullCount (action alg) a (n + 1) ω → ω.2 i a = ω'.2 i a) : + hist alg ω n = hist alg ω' n := by + induction n with + | zero => + simp only [zero_add, pullCount_one] at hω2 + simp_rw [hist_zero] + ext i : 1 + simp only [le_refl, hω1, Prod.mk.injEq, true_and] + refine hω2 0 _ ?_ + simp [action, hω1] + | succ n hn => + simp_rw [hist_add_one_eq_IicSuccProd] + specialize hn fun i hin ↦ hω1 i (by grind) + have h_hist : hist alg ω n = hist alg ω' n := by + refine hn fun i a hi ↦ hω2 i a (hi.trans_le ?_) + exact pullCount_mono _ (by lia) _ + have h_action : action alg (n + 1) ω = action alg (n + 1) ω' := by + simp_rw [action_add_one_eq] + rw [h_hist, hω1 _ le_rfl] + congr 3 + simp only [reward_add_one, h_hist, h_action] + refine hω2 _ _ ?_ + rw [pullCount_add_one, h_action] + simp only [↓reduceIte] + rw [pullCount_eq_pullCount' (R' := reward alg) (by simp)] + simp only [Nat.add_one_sub_one] + rw [← h_hist, hist_eq] + change pullCount' n (fun i ↦ (action alg i ω, reward alg i ω)) (action alg (n + 1) ω') < + pullCount' n (fun i ↦ (action alg i ω, reward alg i ω)) (action alg (n + 1) ω') + 1 + grind + +-- very bad name +noncomputable +def truePast (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] + (a : α) (n : ℕ) (ω : probSpace α R) : + probSpace α R := + (ω.1, fun i b ↦ if b = a then ω.2 (min i ((pullCount (action alg) a (n + 1) ω) - 1)) a + else ω.2 i b) + +omit [Countable α] [StandardBorelSpace R] [Nonempty R] in +lemma truePast_eq_of_pullCount_eq (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] + (a : α) (n m : ℕ) (ω : probSpace α R) + (h_pc : pullCount (action alg) a (n + 1) ω = m) : + truePast alg ν a n ω = (ω.1, fun i b ↦ if b = a then ω.2 (min i (m - 1)) a else ω.2 i b) := by + simp [truePast, h_pc] + +omit [StandardBorelSpace R] [Nonempty R] in +lemma measurable_hist_truePast (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] + (a : α) (n : ℕ) : + Measurable[MeasurableSpace.comap (truePast alg ν a n) inferInstance] (hist alg · n) := by + have h_eq : (hist alg · n) = (hist alg · n) ∘ (truePast alg ν a n) := by + ext ω : 1 + refine hist_congr alg ν n (fun _ _ ↦ rfl) fun i b hi ↦ ?_ + by_cases hb : b = a + · subst hb + simp only [truePast, ↓reduceIte] + rw [min_eq_left] + grind + · simp [truePast, hb] + rw [h_eq] + refine Measurable.comp ?_ (Measurable.of_comap_le le_rfl) + fun_prop + +omit [StandardBorelSpace R] [Nonempty R] in +lemma measurable_action_add_one_truePast (alg : Algorithm α R) + (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (n : ℕ) : + Measurable[MeasurableSpace.comap (truePast alg ν a n) inferInstance] + (action alg (n + 1)) := by + rw [action_add_one_eq] + change Measurable[MeasurableSpace.comap (truePast alg ν a n) inferInstance] + ((fun p ↦ algFunction alg n p.1 p.2) ∘ (fun ω ↦ (hist alg ω n, ω.1 (n + 1)))) + refine (measurable_algFunction alg n).comp (Measurable.prodMk ?_ ?_) + · exact measurable_hist_truePast alg ν a n + · have : (fun ω ↦ ω.1 (n + 1)) = + (fun (p : probSpace α R) ↦ p.1 (n + 1)) ∘ (truePast alg ν a n) := rfl + rw [this] + exact Measurable.comp (by fun_prop) (Measurable.of_comap_le le_rfl) + +omit [StandardBorelSpace R] [Nonempty R] in +lemma measurable_pullCount_add_one_truePast (alg : Algorithm α R) + (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (n : ℕ) : + Measurable[MeasurableSpace.comap (truePast alg ν a n) inferInstance] + (pullCount (action alg) a (n + 1)) := by + change Measurable[MeasurableSpace.comap (truePast alg ν a n) inferInstance] + (fun ω ↦ pullCount (action alg) a (n + 1) ω) + simp_rw [pullCount_eq_sum] + refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + refine (measurableSet_singleton _).preimage ?_ + have h_meas := measurable_hist_truePast alg ν a n + simp_rw [hist_eq _ _ n, @measurable_pi_iff] at h_meas + exact (h_meas ⟨i, by grind⟩).fst + +lemma indepFun_snd_apply_aux (alg : Algorithm α R) + (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m n : ℕ) : + (fun ω ↦ ω.2 m a) ⟂ᵢ[arrayMeasure ν] + (fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b)) := by + sorry + +omit [Countable α] [StandardBorelSpace R] [Nonempty R] in +lemma stepsUntil_congr_aux (alg : Algorithm α R) + (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m n : ℕ) {ω ω' : probSpace α R} + (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) + (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) + (h_eq : action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m) : + action alg (n + 1) ω' = a ∧ pullCount (action alg) a (n + 1) ω' = m := by + obtain ⟨h_action, h_pc⟩ := h_eq + have h_hist := hist_congr alg ν n (ω := ω) (ω' := ω') (by grind) fun i b hi ↦ ?_ + swap + · rcases eq_or_ne b a with (rfl | hba) + · refine hω2_eq i ?_ + rw [h_pc] at hi + grind + · grind + constructor + · rw [← h_action, action_add_one_eq] + simp [h_hist, hω1] + · simp_rw [← h_pc, pullCount_eq_sum] + refine Finset.sum_congr rfl fun i hi ↦ ?_ + congr 2 + rw [hist_eq _ _ n, hist_eq _ _ n, funext_iff] at h_hist + unfold action + specialize h_hist ⟨i, by grind⟩ + simp only at h_hist + rw [h_hist] + +omit [Countable α] [StandardBorelSpace R] [Nonempty R] in +lemma stepsUntil_congr (alg : Algorithm α R) + (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m n : ℕ) {ω ω' : probSpace α R} + (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) + (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) : + (action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m) ↔ + (action alg (n + 1) ω' = a ∧ pullCount (action alg) a (n + 1) ω' = m) := + ⟨stepsUntil_congr_aux alg ν a m n hω1 hω2_ne hω2_eq, + stepsUntil_congr_aux alg ν a m n (by grind) (by grind) (by grind)⟩ + +omit [Countable α] [StandardBorelSpace R] [Nonempty R] in +lemma stepsUntil_indicator_congr (alg : Algorithm α R) + (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m n : ℕ) {ω ω' : probSpace α R} + (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) + (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) : + {ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}.indicator (fun _ ↦ 1) + ω = + {ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}.indicator + (fun _ ↦ 1) ω' := by + simp only [Set.indicator_apply, Set.mem_setOf_eq] + simp_rw [stepsUntil_congr alg ν a m n hω1 hω2_ne hω2_eq] + +omit [StandardBorelSpace R] [Nonempty R] in +lemma measurable_stepsUntil (alg : Algorithm α R) + (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m n : ℕ) : + Measurable[MeasurableSpace.comap + (fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b)) inferInstance] + (({ω | action alg (n + 1) ω = a ∧ + pullCount (action alg) a (n + 1) ω = m}).indicator (fun _ ↦ 1)) := by + let f := ({ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}).indicator + (fun _ ↦ 1) + have h_eq : f = f ∘ + fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b) := by + ext ω + exact stepsUntil_indicator_congr alg ν a m n (by grind) (by grind) (by grind) + change Measurable[MeasurableSpace.comap + (fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b)) inferInstance] f + rw [h_eq] + refine Measurable.comp ?_ (Measurable.of_comap_le le_rfl) + refine Measurable.indicator (by fun_prop) ?_ + exact MeasurableSet.inter ((measurableSet_singleton _).preimage (by fun_prop)) + ((measurableSet_singleton _).preimage (by fun_prop)) + +lemma indepFun_snd_apply_pullCount_action (alg : Algorithm α R) + (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m n : ℕ) : + (fun ω ↦ ω.2 m a) ⟂ᵢ[arrayMeasure ν] + ({ω | action alg (n + 1) ω = a ∧ + pullCount (action alg) a (n + 1) ω = m}).indicator (fun _ ↦ 1) := + (indepFun_snd_apply_aux alg ν a m n).of_measurable_right (measurable_stepsUntil alg ν a m n) + +omit [StandardBorelSpace R] [Nonempty R] in +@[fun_prop] +lemma measurable_pullCount_action_add_one (alg : Algorithm α R) + (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + Measurable (fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω) := by + change Measurable ((fun p : (probSpace α R) × α ↦ pullCount (action alg) p.2 (n + 1) p.1) ∘ + (fun ω : probSpace α R ↦ (ω, action alg (n + 1) ω))) + exact (measurable_uncurry_pullCount (by fun_prop) _).comp (by fun_prop) + +lemma hasCondDistrib_reward_pullCount_action + (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + HasCondDistrib (reward alg (n + 1)) + (fun ω ↦ (action alg (n + 1) ω, pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) + (ν.prodMkRight _) (arrayMeasure ν) := by + have h_meas : Measurable fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω := by + change Measurable ((fun p : (probSpace α R) × α ↦ pullCount (action alg) p.2 (n + 1) p.1) ∘ + (fun ω : probSpace α R ↦ (ω, action alg (n + 1) ω))) + exact (measurable_uncurry_pullCount (by fun_prop) _).comp (by fun_prop) + refine ⟨by fun_prop, by fun_prop, ?_⟩ + refine (condDistrib_ae_eq_cond + (Measurable.prodMk (by fun_prop) (by fun_prop)) (by fun_prop)).trans ?_ + rw [Filter.EventuallyEq, ae_iff_of_countable] + intro ⟨a, m⟩ ham + simp only [Kernel.prodMkRight_apply] + calc + Measure.map (reward alg (n + 1)) + (arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}] + _ = Measure.map (fun ω ↦ ω.2 m a) + (arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}] := by + rw [reward_eq] + refine Measure.map_congr + (ae_cond_of_forall_mem ((measurableSet_singleton _).preimage (by fun_prop)) (fun x hx ↦ ?_)) + simp only [Set.mem_preimage, Set.mem_singleton_iff, Prod.mk.injEq] at hx + simp only [hx.1] at hx ⊢ + simp [hx.2] + _ = Measure.map (fun ω ↦ ω.2 m a) + (arrayMeasure ν)[|({ω | action alg (n + 1) ω = a ∧ + pullCount (action alg) a (n + 1) ω = m}).indicator 1 ⁻¹' {1}] := by + congr with ω + simp only [Set.mem_preimage, Set.mem_singleton_iff, Prod.mk.injEq, Set.indicator_apply, + Set.mem_setOf_eq, Pi.one_apply, ite_eq_left_iff, not_and, zero_ne_one, imp_false, + Classical.not_imp, Decidable.not_not, and_congr_right_iff] + intro ha + simp [ha] + _ = ν a := by + rw [cond_of_indepFun, map_snd_apply_arrayMeasure m a] + · exact (indepFun_snd_apply_pullCount_action alg ν a m n).symm + · refine Measurable.indicator (by fun_prop) ?_ + exact MeasurableSet.inter ((measurableSet_singleton _).preimage (by fun_prop)) + ((measurableSet_singleton _).preimage (by fun_prop)) + · fun_prop + · simp + · rw [Measure.map_apply (by fun_prop) (by simp)] at ham + convert ham + ext ω + simp only [Set.mem_preimage, Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply, + Set.mem_singleton_iff, ite_eq_left_iff, not_and, zero_ne_one, imp_false, Classical.not_imp, + Decidable.not_not, Prod.mk.injEq, and_congr_right_iff] + intro ha + simp [ha] + +-- lemma hasCondDistrib_reward_hist_action_pullCount' +-- (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (n m : ℕ) : +-- HasCondDistrib (reward alg (n + 1)) +-- (fun ω ↦ (hist alg ω n, {ω' | action alg (n + 1) ω' = a ∧ +-- pullCount (action alg) (action alg (n + 1) ω') (n + 1) ω' = m}.indicator +-- (fun _ ↦ (1 : ℕ)) ω)) +-- (Kernel.const _ (ν a)) (arrayMeasure ν) := by +-- sorry + +omit [StandardBorelSpace R] [Nonempty R] in +lemma reward_ae_eq_cond + (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (n m : ℕ) : + reward alg (n + 1) =ᵐ[(arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}]] + (fun ω ↦ ω.2 m a) := by + rw [reward_eq] + refine ae_cond_of_forall_mem ?_ ?_ + · have : Measurable fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω := by + change Measurable ((fun p : (probSpace α R) × α ↦ pullCount (action alg) p.2 (n + 1) p.1) ∘ + (fun ω : probSpace α R ↦ (ω, action alg (n + 1) ω))) + exact (measurable_uncurry_pullCount (by fun_prop) _).comp (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + intro ω hω + simp only [Set.mem_preimage, Set.mem_singleton_iff, Prod.mk.injEq] at hω + simp only [hω.2] + simp [hω.1] + +lemma indepFun_cond_of_indepFun {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {μ : Measure α} + {X : α → β} {Y : α → γ} (hXY : X ⟂ᵢ[μ] Y) (hY : Measurable Y) {s : Set γ} + (hs : MeasurableSet s) : + X ⟂ᵢ[μ[|Y ⁻¹' s]] Y := by + by_cases h_zero : μ[|Y ⁻¹' s] = 0 + · simp only [h_zero] + -- missing simp lemma : `X ⟂ᵢ[0] Y` + simp [indepFun_iff_measure_inter_preimage_eq_mul] + rw [cond_eq_zero] at h_zero + push_neg at h_zero + rw [indepFun_iff_measure_inter_preimage_eq_mul] at hXY ⊢ + intro u t hu ht + rw [cond_apply (hs.preimage hY), cond_apply (hs.preimage hY), cond_apply (hs.preimage hY)] + have h_eq : Y ⁻¹' s ∩ (X ⁻¹' u ∩ Y ⁻¹' t) = X ⁻¹' u ∩ Y ⁻¹' (s ∩ t) := by grind + have hsu := hXY u s hu hs + have hust := hXY u (s ∩ t) hu (hs.inter ht) + rw [Set.inter_comm] at hsu + rw [hsu, h_eq, hust] + simp_rw [mul_assoc] + congr 1 + rw [← mul_assoc (μ (Y ⁻¹' s)), ENNReal.mul_inv_cancel h_zero.2 h_zero.1, one_mul] + congr + +lemma indepFun_todo {α β γ δ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {mδ : MeasurableSpace δ} [MeasurableSingletonClass δ] {μ : Measure α} + {X : α → β} {Y : α → γ} (hXY : X ⟂ᵢ[μ] Y) (hY : Measurable Y) + {Z : γ → δ} (hZ : Measurable Z) (z : δ) : + X ⟂ᵢ[μ[|(Z ∘ Y) ⁻¹' {z}]] Y := by + have h_preim : (Z ∘ Y) ⁻¹' {z} = Y ⁻¹' (Z ⁻¹' {z}) := by grind + simp_rw [h_preim] + exact indepFun_cond_of_indepFun hXY hY (hZ (measurableSet_singleton z)) + +lemma indepFun_snd_hist_cond (alg : Algorithm α R) + (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (n m : ℕ) : + (fun ω ↦ ω.2 m a) ⟂ᵢ[(arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}]] + (hist alg · n) := by + have h_meas := measurable_hist_truePast alg ν a n + refine IndepFun.of_measurable_right ?_ h_meas + have h_ae_eq : truePast alg ν a n =ᵐ[(arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}]] + (fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b)) := by + refine ae_cond_of_forall_mem ?_ fun x hx ↦ ?_ + · refine (measurableSet_singleton _).preimage ?_ + have h_meas_pc : Measurable fun ω ↦ + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω := by + change Measurable ((fun p : (probSpace α R) × α ↦ pullCount (action alg) p.2 (n + 1) p.1) ∘ + (fun ω : probSpace α R ↦ (ω, action alg (n + 1) ω))) + exact (measurable_uncurry_pullCount (by fun_prop) _).comp (by fun_prop) + fun_prop + simp only [Set.mem_preimage, Set.mem_singleton_iff, Prod.mk.injEq] at hx + simp only [truePast] + congr with i b + by_cases hb : b = a + · simp only [hb, ↓reduceIte] + simp only [hx.1, true_and] at hx + congr! + · simp [hb] + refine IndepFun.congr ?_ EventuallyEq.rfl h_ae_eq.symm + suffices (fun ω ↦ ω.2 m a) ⟂ᵢ[(arrayMeasure ν)[|(({ω | action alg (n + 1) ω = a ∧ + pullCount (action alg) a (n + 1) ω = m}).indicator (fun _ ↦ 1)) ⁻¹' {1}]] + fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b) by + convert this + ext ω + simp only [Set.mem_preimage, Set.mem_singleton_iff, Prod.mk.injEq, Set.indicator_apply, + Set.mem_setOf_eq, ite_eq_left_iff, not_and, zero_ne_one, imp_false, + Classical.not_imp, Decidable.not_not, and_congr_right_iff] + intro ha + simp [ha] + have h_meas := measurable_stepsUntil alg ν a m n + obtain ⟨f, hf, hf_eq⟩ := h_meas.exists_eq_measurable_comp + simp_rw [hf_eq] + refine indepFun_todo (Z := f) (z := 1) ?_ ?_ hf + · exact indepFun_snd_apply_aux alg ν a m n + · refine Measurable.prodMk (by fun_prop) ?_ + simp_rw [measurable_pi_iff] + intro m a + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact MeasurableSet.const _ + +lemma hasCondDistrib_reward_hist_action_pullCount + (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + HasCondDistrib (reward alg (n + 1)) + (fun ω ↦ (hist alg ω n, action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) + ((ν.prodMkRight _).prodMkLeft _) (arrayMeasure ν) := by + have h_meas : Measurable fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω := by + change Measurable ((fun p : (probSpace α R) × α ↦ pullCount (action alg) p.2 (n + 1) p.1) ∘ + (fun ω : probSpace α R ↦ (ω, action alg (n + 1) ω))) + exact (measurable_uncurry_pullCount (by fun_prop) _).comp (by fun_prop) + refine ⟨by fun_prop, by fun_prop, ?_⟩ + refine condDistrib_prod_of_forall_condDistrib_cond (by fun_prop) (by fun_prop) (by fun_prop) _ ?_ + intro (a, m) ham + have h_eq : ((ν.prodMkRight _).prodMkLeft _).comap (fun ω : (Iic n → α × R) ↦ (ω, a, m)) + (by fun_prop) = + Kernel.const _ (ν a) := by ext; simp + rw [h_eq, condDistrib_congr_left (reward_ae_eq_cond alg ν a n m)] + refine (condDistrib_of_indepFun ?_ (by fun_prop) (by fun_prop)).trans (ae_of_all _ fun ω ↦ ?_) + · exact (indepFun_snd_hist_cond alg ν a n m).symm + · simp only [Kernel.const_apply] + have : (fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)} = + ({ω | action alg (n + 1) ω = a ∧ + pullCount (action alg) a (n + 1) ω = m}).indicator 1 ⁻¹' {1} := by + ext ω + simp [Set.indicator_apply] + grind + rw [this, cond_of_indepFun, map_snd_apply_arrayMeasure m a] + · exact (indepFun_snd_apply_pullCount_action alg ν a m n).symm + · refine Measurable.indicator (by fun_prop) ?_ + exact MeasurableSet.inter ((measurableSet_singleton _).preimage (by fun_prop)) + ((measurableSet_singleton _).preimage (by fun_prop)) + · fun_prop + · simp + · convert ham + ext ω + simp only [Set.mem_preimage, Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply, + Set.mem_singleton_iff, ite_eq_left_iff, not_and, zero_ne_one, imp_false, Classical.not_imp, + Decidable.not_not, Prod.mk.injEq, and_congr_right_iff] + intro ha + simp [ha] + +lemma condIndepFun_todo (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + (reward alg (n + 1)) ⟂ᵢ[(fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)), + Measurable.prodMk (by fun_prop) (measurable_pullCount_action_add_one alg ν n); + arrayMeasure ν] + (hist alg · n) := by + have h_cond := hasCondDistrib_reward_hist_action_pullCount alg ν n + refine condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (by fun_prop) (by fun_prop) ?_ + h_cond.condDistrib_eq + exact Measurable.prodMk (by fun_prop) (measurable_pullCount_action_add_one alg ν n) + lemma hasCondDistrib_reward' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : HasCondDistrib (reward alg (n + 1)) (fun ω ↦ (hist alg ω n, action alg (n + 1) ω)) (ν.prodMkLeft _) (arrayMeasure ν) := by - rw [reward_eq] - refine ⟨?_, by fun_prop, ?_⟩ - · sorry - refine condDistrib_ae_eq_of_measure_eq_compProd _ ?_ ?_ - · sorry - sorry + suffices HasCondDistrib (reward alg (n + 1)) + (fun ω ↦ (hist alg ω n, action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) + ((ν.prodMkRight _).prodMkLeft _) (arrayMeasure ν) by + sorry + suffices HasCondDistrib (reward alg (n + 1)) + (fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) + (ν.prodMkRight _) (arrayMeasure ν) by + sorry + exact hasCondDistrib_reward_pullCount_action alg ν n lemma hasCondDistrib_action (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : HasCondDistrib (action alg (n + 1)) diff --git a/LeanBandits/Bandit/RewardByCountMeasure.lean b/LeanBandits/Bandit/RewardByCountMeasure.lean index 820f6a37..78a96464 100644 --- a/LeanBandits/Bandit/RewardByCountMeasure.lean +++ b/LeanBandits/Bandit/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.Regret +import LeanBandits.ForMathlib.CondIndepFun import LeanBandits.ForMathlib.IndepFun import Mathlib.Probability.IdentDistribIndep @@ -13,44 +14,6 @@ import Mathlib.Probability.IdentDistribIndep open MeasureTheory ProbabilityTheory Finset Learning open scoped ENNReal NNReal -section Aux -- todo: move - -namespace ProbabilityTheory - -variable {α β γ δ γ' δ' : Type*} - {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} - {mδ : MeasurableSpace δ} {mγ' : MeasurableSpace γ'} {mδ' : MeasurableSpace δ'} - [StandardBorelSpace α] - [StandardBorelSpace δ'] [Nonempty δ'] [StandardBorelSpace γ'] [Nonempty γ'] - {μ : Measure α} [IsFiniteMeasure μ] - {X : α → β} {hX : Measurable X} {Y : α → γ} {Z : α → δ} {Y' : α → γ'} {Z' : α → δ'} - -lemma CondIndepFun.of_measurable (h_indep : Y ⟂ᵢ[X, hX; μ] Z) - (hY_meas : Measurable[mγ.comap Y] Y') (hZ_meas : Measurable[mδ.comap Z] Z') : - Y' ⟂ᵢ[X, hX; μ] Z' := by - obtain ⟨φ, hφ_meas, h_eqY⟩ : ∃ φ, Measurable φ ∧ Y' = φ ∘ Y := hY_meas.exists_eq_measurable_comp - obtain ⟨ψ, hψ_meas, h_eqZ⟩ : ∃ ψ, Measurable ψ ∧ Z' = ψ ∘ Z := hZ_meas.exists_eq_measurable_comp - rw [h_eqY, h_eqZ] - exact h_indep.comp hφ_meas hψ_meas - -lemma CondIndepFun.of_measurable_left - (h_indep : Y ⟂ᵢ[X, hX; μ] Z) (hY_meas : Measurable[mγ.comap Y] Y') : - Y' ⟂ᵢ[X, hX; μ] Z := by - obtain ⟨φ, hφ_meas, h_eqY⟩ : ∃ φ, Measurable φ ∧ Y' = φ ∘ Y := hY_meas.exists_eq_measurable_comp - rw [h_eqY] - exact h_indep.comp hφ_meas measurable_id - -lemma CondIndepFun.of_measurable_right - (h_indep : Y ⟂ᵢ[X, hX; μ] Z) (hZ_meas : Measurable[mδ.comap Z] Z') : - Y ⟂ᵢ[X, hX; μ] Z' := by - obtain ⟨ψ, hψ_meas, h_eqZ⟩ : ∃ ψ, Measurable ψ ∧ Z' = ψ ∘ Z := hZ_meas.exists_eq_measurable_comp - rw [h_eqZ] - exact h_indep.comp measurable_id hψ_meas - -end ProbabilityTheory - -end Aux - namespace Bandits variable {α Ω : Type*} {mα : MeasurableSpace α} {mΩ : MeasurableSpace Ω} [DecidableEq α] diff --git a/LeanBandits/ForMathlib/CondDistrib.lean b/LeanBandits/ForMathlib/CondDistrib.lean index 101b4c93..b1cde46a 100644 --- a/LeanBandits/ForMathlib/CondDistrib.lean +++ b/LeanBandits/ForMathlib/CondDistrib.lean @@ -392,6 +392,120 @@ lemma condDistrib_ae_eq_cond [Countable β] [MeasurableSingletonClass β] · congr · exact hb +lemma lintegral_cond {μ : Measure α} (s : Set α) (f : α → ℝ≥0∞) : + ∫⁻ x, f x ∂μ[|s] = (μ s)⁻¹ * ∫⁻ (a : α) in s, f a ∂μ := by + unfold cond + simp [lintegral_smul_measure] + +omit [Nonempty Ω'] in +lemma condDistrib_prod_of_forall_condDistrib_cond [Countable Ω'] [IsFiniteMeasure μ] + (hX : Measurable X) (hY : Measurable Y) (hZ : Measurable Z) + (κ : Kernel (β × Ω') Ω) [IsFiniteKernel κ] + (h_cond : ∀ b, μ (Z ⁻¹' {b}) ≠ 0 → condDistrib Y X μ[|Z ⁻¹' {b}] =ᵐ[μ[|Z ⁻¹' {b}].map X] + (κ.comap (fun ω ↦ (ω, b)) (by fun_prop))) : + condDistrib Y (fun ω ↦ (X ω, Z ω)) μ =ᵐ[μ.map (fun ω ↦ (X ω, Z ω))] κ := by + refine condDistrib_ae_eq_of_measure_eq_compProd _ (by fun_prop) ?_ + ext s hs + suffices ∀ b, (Measure.map (fun x ↦ ((X x, Z x), Y x)) μ) (s ∩ {p | p.1.2 = b}) = + (Measure.map (fun ω ↦ (X ω, Z ω)) μ ⊗ₘ κ) (s ∩ {p | p.1.2 = b}) by + have hs_iUnion : s = ⋃ b, s ∩ {p | p.1.2 = b} := by + ext p + simp only [Set.mem_iUnion, Set.mem_inter_iff, Set.mem_setOf_eq] + grind + have h_disj : Pairwise (Function.onFun Disjoint fun b ↦ s ∩ {p | p.1.2 = b}) := by + intro i j hij + simp only [Set.disjoint_iff_inter_eq_empty] + ext + grind + have h_meas (b : Ω') : MeasurableSet (s ∩ {p | p.1.2 = b}) := + hs.inter ((measurableSet_singleton _).preimage (by fun_prop)) + rw [hs_iUnion, measure_iUnion h_disj h_meas, measure_iUnion h_disj h_meas] + congr with b + exact this b + intro b + by_cases hb : μ (Z ⁻¹' {b}) = 0 + · have h_left : (Measure.map (fun x ↦ ((X x, Z x), Y x)) μ) (s ∩ {p | p.1.2 = b}) = 0 := by + suffices (Measure.map (fun x ↦ ((X x, Z x), Y x)) μ) {p | p.1.2 = b} = 0 from + measure_mono_null Set.inter_subset_right this + rw [Measure.map_apply (by fun_prop)] + · simpa + · exact (measurableSet_singleton _).preimage (by fun_prop) + have h_right : (Measure.map (fun ω ↦ (X ω, Z ω)) μ ⊗ₘ κ) (s ∩ {p | p.1.2 = b}) = 0 := by + suffices (Measure.map (fun ω ↦ (X ω, Z ω)) μ ⊗ₘ κ) {p | p.1.2 = b} = 0 from + measure_mono_null Set.inter_subset_right this + rw [Measure.compProd_apply, lintegral_map] + rotate_left + · exact Kernel.measurable_kernel_prodMk_left + ((measurableSet_singleton _).preimage (by fun_prop)) + · fun_prop + · exact (measurableSet_singleton _).preimage (by fun_prop) + simp only [Set.preimage_setOf_eq] + classical + have h_le : ∫⁻ a, (κ (X a, Z a)) {a_1 | Z a = b} ∂μ ≤ + ∫⁻ a, {a' | Z a' = b}.indicator (fun _ ↦ κ.bound) a ∂μ := by + gcongr with a + by_cases hZ : Z a = b + · simp only [hZ, Set.setOf_true, Set.mem_setOf_eq, Set.indicator_of_mem] + exact κ.measure_le_bound _ _ + · simp [hZ] + refine le_antisymm (h_le.trans ?_) zero_le' + rw [lintegral_indicator] + swap; · exact (measurableSet_singleton _).preimage (by fun_prop) + simp only [lintegral_const, MeasurableSet.univ, Measure.restrict_apply, Set.univ_inter, + nonpos_iff_eq_zero, mul_eq_zero] + exact .inr hb + rw [h_left, h_right] + specialize h_cond b hb + rw [condDistrib_ae_eq_iff_measure_eq_compProd] at h_cond + swap; · fun_prop + rw [Measure.ext_iff] at h_cond + have hs' : MeasurableSet {p : β × Ω | ((p.1, b), p.2) ∈ s} := hs.preimage (by fun_prop) + have h1 := h_cond {p | ((p.1, b), p.2) ∈ s} hs' + have h_indicator : Measurable ({ω' | Z ω' = b}.indicator (fun x ↦ 1)) := + Measurable.indicator (by fun_prop) ((measurableSet_singleton _).preimage (by fun_prop)) + rw [Measure.map_apply] at h1 ⊢ + rotate_left + · fun_prop + · exact hs.inter ((measurableSet_singleton _).preimage (by fun_prop)) + · fun_prop + · exact hs' + rw [cond_apply] at h1 + swap; · exact (measurableSet_singleton _).preimage (by fun_prop) + have h1' : μ (Z ⁻¹' {b} ∩ (fun x ↦ (X x, Y x)) ⁻¹' {p | ((p.1, b), p.2) ∈ s}) = + (μ (Z ⁻¹' {b})) * + (Measure.map X μ[|Z ⁻¹' {b}] ⊗ₘ κ.comap (fun ω ↦ (ω, b)) (by fun_prop)) + {p | ((p.1, b), p.2) ∈ s} := by + rw [← h1, ← mul_assoc, ENNReal.mul_inv_cancel hb (by simp), one_mul] + convert h1' + · ext x + simp only [Set.preimage_inter, Set.preimage_setOf_eq, Set.mem_inter_iff, Set.mem_preimage, + Set.mem_setOf_eq] + grind + · rw [Measure.compProd_apply, Measure.compProd_apply, lintegral_map, lintegral_map] + rotate_left + · exact Kernel.measurable_kernel_prodMk_left hs' + · fun_prop + · apply Kernel.measurable_kernel_prodMk_left + exact hs.inter ((measurableSet_singleton _).preimage (by fun_prop)) + · fun_prop + · exact hs' + · exact hs.inter ((measurableSet_singleton _).preimage (by fun_prop)) + rw [lintegral_cond, ← mul_assoc, ENNReal.mul_inv_cancel hb (by simp), one_mul] + simp only [Set.preimage_inter, Set.preimage_setOf_eq, Kernel.coe_comap, Function.comp_apply] + classical + have h_eq : (fun a ↦ κ (X a, Z a) (Prod.mk (X a, Z a) ⁻¹' s ∩ {a_1 | Z a = b})) = + {a | Z a = b}.indicator + (fun a ↦ κ (X a, b) (Prod.mk (X a, b) ⁻¹' s ∩ {a_1 | Z a = b})) := by + ext a + by_cases hZ : Z a = b <;> simp [hZ] + simp_rw [h_eq] + rw [lintegral_indicator] + swap; · exact (measurableSet_singleton _).preimage (by fun_prop) + refine setLIntegral_congr_fun ((measurableSet_singleton _).preimage (by fun_prop)) fun a ha ↦ ?_ + congr 1 with ω + simp only [Set.mem_inter_iff, Set.mem_preimage, Set.mem_setOf_eq, and_iff_left_iff_imp] + grind + 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) : diff --git a/LeanBandits/ForMathlib/CondIndepFun.lean b/LeanBandits/ForMathlib/CondIndepFun.lean new file mode 100644 index 00000000..d361705a --- /dev/null +++ b/LeanBandits/ForMathlib/CondIndepFun.lean @@ -0,0 +1,71 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +import Mathlib.MeasureTheory.Function.FactorsThrough +import Mathlib.Probability.Independence.Basic +import Mathlib.Probability.Independence.Conditional + +/-! # Laws of `stepsUntil` and `rewardByCount` +-/ + +open MeasureTheory ProbabilityTheory Finset +open scoped ENNReal NNReal + +namespace ProbabilityTheory + +variable {α β γ δ γ' δ' : Type*} + {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} + {mδ : MeasurableSpace δ} {mγ' : MeasurableSpace γ'} {mδ' : MeasurableSpace δ'} + [StandardBorelSpace δ'] [Nonempty δ'] [StandardBorelSpace γ'] [Nonempty γ'] + {μ : Measure α} + {X : α → β} {hX : Measurable X} {Y : α → γ} {Z : α → δ} {Y' : α → γ'} {Z' : α → δ'} + +lemma IndepFun.of_measurable (h_indep : Y ⟂ᵢ[μ] Z) + (hY_meas : Measurable[mγ.comap Y] Y') (hZ_meas : Measurable[mδ.comap Z] Z') : + Y' ⟂ᵢ[μ] Z' := by + obtain ⟨φ, hφ_meas, h_eqY⟩ : ∃ φ, Measurable φ ∧ Y' = φ ∘ Y := hY_meas.exists_eq_measurable_comp + obtain ⟨ψ, hψ_meas, h_eqZ⟩ : ∃ ψ, Measurable ψ ∧ Z' = ψ ∘ Z := hZ_meas.exists_eq_measurable_comp + rw [h_eqY, h_eqZ] + exact h_indep.comp hφ_meas hψ_meas + +lemma IndepFun.of_measurable_left + (h_indep : Y ⟂ᵢ[μ] Z) (hY_meas : Measurable[mγ.comap Y] Y') : + Y' ⟂ᵢ[μ] Z := by + obtain ⟨φ, hφ_meas, h_eqY⟩ : ∃ φ, Measurable φ ∧ Y' = φ ∘ Y := hY_meas.exists_eq_measurable_comp + rw [h_eqY] + exact h_indep.comp hφ_meas measurable_id + +lemma IndepFun.of_measurable_right + (h_indep : Y ⟂ᵢ[μ] Z) (hZ_meas : Measurable[mδ.comap Z] Z') : + Y ⟂ᵢ[μ] Z' := by + obtain ⟨ψ, hψ_meas, h_eqZ⟩ : ∃ ψ, Measurable ψ ∧ Z' = ψ ∘ Z := hZ_meas.exists_eq_measurable_comp + rw [h_eqZ] + exact h_indep.comp measurable_id hψ_meas + +variable [StandardBorelSpace α] [IsFiniteMeasure μ] + +lemma CondIndepFun.of_measurable (h_indep : Y ⟂ᵢ[X, hX; μ] Z) + (hY_meas : Measurable[mγ.comap Y] Y') (hZ_meas : Measurable[mδ.comap Z] Z') : + Y' ⟂ᵢ[X, hX; μ] Z' := by + obtain ⟨φ, hφ_meas, h_eqY⟩ : ∃ φ, Measurable φ ∧ Y' = φ ∘ Y := hY_meas.exists_eq_measurable_comp + obtain ⟨ψ, hψ_meas, h_eqZ⟩ : ∃ ψ, Measurable ψ ∧ Z' = ψ ∘ Z := hZ_meas.exists_eq_measurable_comp + rw [h_eqY, h_eqZ] + exact h_indep.comp hφ_meas hψ_meas + +lemma CondIndepFun.of_measurable_left + (h_indep : Y ⟂ᵢ[X, hX; μ] Z) (hY_meas : Measurable[mγ.comap Y] Y') : + Y' ⟂ᵢ[X, hX; μ] Z := by + obtain ⟨φ, hφ_meas, h_eqY⟩ : ∃ φ, Measurable φ ∧ Y' = φ ∘ Y := hY_meas.exists_eq_measurable_comp + rw [h_eqY] + exact h_indep.comp hφ_meas measurable_id + +lemma CondIndepFun.of_measurable_right + (h_indep : Y ⟂ᵢ[X, hX; μ] Z) (hZ_meas : Measurable[mδ.comap Z] Z') : + Y ⟂ᵢ[X, hX; μ] Z' := by + obtain ⟨ψ, hψ_meas, h_eqZ⟩ : ∃ ψ, Measurable ψ ∧ Z' = ψ ∘ Z := hZ_meas.exists_eq_measurable_comp + rw [h_eqZ] + exact h_indep.comp measurable_id hψ_meas + +end ProbabilityTheory diff --git a/LeanBandits/ForMathlib/KernelRepresentation.lean b/LeanBandits/ForMathlib/KernelRepresentation.lean new file mode 100644 index 00000000..68366915 --- /dev/null +++ b/LeanBandits/ForMathlib/KernelRepresentation.lean @@ -0,0 +1,146 @@ +/- +Copyright (c) 2025 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré, Rémy Degenne +-/ + +import Mathlib.Analysis.SpecialFunctions.Sigmoid +import Mathlib.MeasureTheory.Constructions.UnitInterval +import Mathlib.Order.CompletePartialOrder +import Mathlib.Probability.CDF + +-- copied from PR #30112 + +/-! +# Representation of kernels + +This file contains results about isolation of kernels randomness. In particular, it shows that, +when the target space is a standard Borel space, any Markov kernel can be represented as the image +of the uniform measure on `[0,1]` by a deterministic map. It corresponds to Lemma 4.22 in +"Foundations of Modern Probability" by Olav Kallenberg, 2021. + +## Statements + +* `ProbabilityTheory.Kernel.unitInterval_representation`: + for a Markov kernel `κ : Kernel α I`, there exists a jointly measurable function + `f : α → I → I` such that for all `a : α`, `volume.map (f a) = κ a`. + +* `ProbabilityTheory.Kernel.embedding_representation`: + for a measurable embedding `g : β → I` and a Markov kernel `κ : Kernel α β`, + there exists a jointly measurable function `f : α → I → β` such that for all `a : α`, + `volume.map (f a) = κ a`. + +* `ProbabilityTheory.Kernel.representation`: + for a Markov kernel `κ : Kernel α β` with `β` a standard Borel space, + there exists a jointly measurable function `f : α → I → β` such that for all `a : α`, + `volume.map (f a) = κ a`. + This is a consequence of `ProbabilityTheory.Kernel.embedding_representation` and the fact that + any standard Borel space can be embedded in `ℝ`, and then composed with `unitInterval.sigmoid`. +-/ + +open MeasureTheory ProbabilityTheory Set ENNReal unitInterval Filter Topology Function + +namespace ProbabilityTheory.Kernel + +variable {α : Type*} [MeasurableSpace α] + +lemma unitInterval_representation (κ : Kernel α I) [IsMarkovKernel κ] : + ∃ (f : α → I → I), Measurable (uncurry f) ∧ ∀ a, volume.map (f a) = κ a := by + let f := fun s (t : I) ↦ sSup {x | (κ s).real (Icc 0 x) < t} + have measurable_f : Measurable (uncurry f) := by + refine measurable_of_Ioi fun a ↦ ?_ + simp only [preimage, uncurry, mem_Ioi] + have h_monotone s : Monotone (fun x ↦ (κ s).real (Icc 0 x)) := by + intro x y hxy + suffices h : Icc 0 x ⊆ Icc 0 y from measureReal_mono h + exact Icc_subset_Icc_right hxy + have sSup_eq_iUnion_rat : {x : α × I | a < f x.1 x.2} = ⋃ (q : ℚ), ⋃ (hqI : ↑q ∈ I), + ⋃ (_ : a < (q : ℝ)), {e | (κ e.1).real (Icc 0 ⟨q, hqI⟩) < e.2} := by + ext e + simp only [f] + constructor + · intro (he : a < sSup {x | (κ e.1).real (Icc 0 x) < e.2}) + simp_rw [Set.mem_iUnion] + rw [lt_sSup_iff] at he + obtain ⟨y, y_mem, (hy : a.1 < y.1)⟩ := he + obtain ⟨q, hqa, hqy⟩ := exists_rat_btwn hy + have q_in_I : (q : ℝ) ∈ I := ⟨a.2.1.trans hqa.le, hqy.le.trans y.2.2⟩ + refine ⟨q, q_in_I, hqa, ?_⟩ + exact lt_of_lt_of_le' y_mem (h_monotone e.1 hqy.le) + · intro he + simp_all only [lt_sSup_iff, Set.mem_iUnion] + obtain ⟨q, q_in_I, hqa, h⟩ := he + exact ⟨⟨q, q_in_I⟩, h, hqa⟩ + rw [sSup_eq_iUnion_rat] + refine MeasurableSet.iUnion (fun b ↦ MeasurableSet.iUnion + (fun bI ↦ MeasurableSet.iUnion (fun _ ↦ ?_))) + refine measurableSet_lt ?_ measurable_snd.subtype_val + simp_rw [measureReal_def] + have hκ := κ.measurable_coe (s := Icc 0 ⟨b, bI⟩) measurableSet_Icc + fun_prop + refine ⟨f, measurable_f, fun a ↦ (volume.map (f a)).ext_of_Iic (κ a) fun x ↦ ?_⟩ + rw [volume.map_apply measurable_f.of_uncurry_left measurableSet_Iic, preimage] + simp only [mem_Iic] + have Iic_to_Icc : Iic x = Icc 0 x := by ext; simp + rw [Iic_to_Icc] + clear Iic_to_Icc + rw [← ofReal_measureReal (measure_ne_top (κ a) _)] + have κ_in_I : ((κ a).real (Icc 0 x)) ∈ I := ⟨measureReal_nonneg, measureReal_le_one⟩ + rw [← volume_Iic ⟨_, κ_in_I⟩] + congr with ξ + constructor + swap + · intro (hξ : ξ ≤ (κ a).real (Icc 0 x)) + simp only [sSup_le_iff, f] + intro c hc + have le1 := lt_of_le_of_lt' hξ hc + by_contra h + push_neg at h + have le2 : (κ a).real (Icc 0 x) ≤ (κ a).real (Icc 0 c) := by + suffices h : Icc 0 x ⊆ Icc 0 c from measureReal_mono h + refine (Icc_subset_Icc_iff unitInterval.nonneg').mpr ?_ + exact ⟨nonneg', h.le⟩ + linarith + · intro (hξ : f a ξ ≤ x) + change ξ ≤ (κ a).real (Icc 0 x) + by_cases hx : x = 1 + · simp [hx, ← univ_eq_Icc, ξ.2.2] + let g := fun y ↦ (κ a).real (Icc 0 y) + letI nebot : NeBot (𝓝[>] x) := by + refine nhdsGT_neBot_of_exists_gt ?_ + use 1 + exact lt_of_le_of_ne x.2.2 hx + refine le_of_tendsto_of_tendsto (b := 𝓝[>] x) (g := g) continuousWithinAt_const ?_ ?_ + · let h := cdf ((κ a).map Subtype.val) + have h_continuousWithinAt := continuousWithinAt_Ioi_iff_Ici.mpr (h.right_continuous x) + simp_rw [g, ← unitInterval.cdf_eq_real (κ a)] + exact h_continuousWithinAt.comp (Continuous.continuousWithinAt (by fun_prop)) (fun y hy ↦ hy) + · apply eventually_nhdsWithin_of_forall + intro y hy + by_contra h + push_neg at h + simp only [sSup_le_iff, f] at hξ + specialize hξ y h + replace hξ : y.1 ≤ x.1 := hξ + have : y.1 > x.1 := hy + linarith + +lemma embedding_representation {β : Type*} [Nonempty β] [MeasurableSpace β] {g : β → I} + (hg : MeasurableEmbedding g) (κ : Kernel α β) [IsMarkovKernel κ] : + ∃ (f : α → I → β), Measurable (uncurry f) ∧ ∀ a, volume.map (f a) = κ a := by + have hκg : IsMarkovKernel (κ.map g) := Kernel.IsMarkovKernel.map κ hg.measurable + classical + have hg'κ : κ = (κ.map g).map hg.invFun := by + rw [← Kernel.map_comp_right _ hg.measurable (by fun_prop), LeftInverse.id hg.leftInverse_invFun, + Kernel.map_id] + obtain ⟨f', hf', hf'κ⟩ := (κ.map g).unitInterval_representation + refine ⟨fun a u ↦ hg.invFun (f' a u), by fun_prop, fun a ↦ ?_⟩ + rw [hg'κ, Kernel.map_apply _ (by fun_prop), ← hf'κ, Measure.map_map (by fun_prop) (by fun_prop)] + rfl + +theorem representation {β : Type*} [Nonempty β] [MeasurableSpace β] [StandardBorelSpace β] + (κ : Kernel α β) [IsMarkovKernel κ] : + ∃ (f : α → I → β), Measurable (uncurry f) ∧ ∀ a, volume.map (f a) = κ a := + κ.embedding_representation (measurableEmbedding_sigmoid_comp_embeddingReal β) + +end ProbabilityTheory.Kernel diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index 7c180d4f..a3810696 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -100,6 +100,11 @@ lemma pullCount_eq_pullCount' {n : ℕ} {ω : Ω} (hn : n ≠ 0) : have : n + 1 - 1 = n := by simp exact this ▸ rfl +lemma pullCount'_mono {n m : ℕ} (hnm : n ≤ m) : + pullCount' n (fun i ↦ (A i ω, R' i ω)) a ≤ pullCount' m (fun i ↦ (A i ω, R' i ω)) a := by + simp_rw [pullCount'] + sorry + lemma pullCount_le (a : α) (t : ℕ) (ω : Ω) : pullCount A a t ω ≤ t := (card_filter_le _ _).trans_eq (by simp) @@ -175,6 +180,16 @@ lemma measurable_pullCount [MeasurableSingletonClass α] (hA : ∀ n, Measurable exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop +@[fun_prop] +lemma measurable_uncurry_pullCount [MeasurableSingletonClass α] [MeasurableEq α] + (hA : ∀ n, Measurable (A n)) (t : ℕ) : + Measurable (fun p : Ω × α ↦ pullCount A p.2 t p.1) := by + simp_rw [pullCount_eq_sum] + have h_meas s : Measurable (fun h : Ω × α ↦ if A s h.1 = h.2 then 1 else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact measurableSet_eq_fun (by fun_prop) (by fun_prop) + fun_prop + @[fun_prop] lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : Measurable (fun h : Iic n → α × R ↦ pullCount' n h a) := by @@ -513,7 +528,6 @@ lemma measurable_stepsUntil [MeasurableSingletonClass α] rw [h_union] refine MeasurableSet.iUnion fun s ↦ (measurableSet_singleton _).preimage ?_ exact measurable_pullCount hA a (s + 1) - --simp_rw [stepsUntil_eq_dite] suffices Measurable fun k ↦ if h : k ∈ {k' | ∃ s, pullCount A a (s + 1) k' = m} then (Nat.find h : ℕ∞) else ⊤ by convert this with ω From a548ef2c5b7acc817c63c9b72cda07eaad77c23b Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Tue, 13 Jan 2026 22:26:35 +0100 Subject: [PATCH 14/30] getting close --- LeanBandits/Bandit/Bandit.lean | 37 +++++++++++++++++++++++----------- 1 file changed, 25 insertions(+), 12 deletions(-) diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 483e550c..dfc9b304 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -1008,7 +1008,7 @@ lemma hasCondDistrib_reward_hist_action_pullCount intro ha simp [ha] -lemma condIndepFun_todo (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : +lemma condIndepFun_reward_hist (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : (reward alg (n + 1)) ⟂ᵢ[(fun ω ↦ (action alg (n + 1) ω, pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)), Measurable.prodMk (by fun_prop) (measurable_pullCount_action_add_one alg ν n); @@ -1019,20 +1019,33 @@ lemma condIndepFun_todo (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKerne h_cond.condDistrib_eq exact Measurable.prodMk (by fun_prop) (measurable_pullCount_action_add_one alg ν n) -lemma hasCondDistrib_reward' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] - (n : ℕ) : +lemma hasCondDistrib_reward' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : HasCondDistrib (reward alg (n + 1)) (fun ω ↦ (hist alg ω n, action alg (n + 1) ω)) (ν.prodMkLeft _) (arrayMeasure ν) := by - suffices HasCondDistrib (reward alg (n + 1)) - (fun ω ↦ (hist alg ω n, action alg (n + 1) ω, - pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) - ((ν.prodMkRight _).prodMkLeft _) (arrayMeasure ν) by - sorry - suffices HasCondDistrib (reward alg (n + 1)) - (fun ω ↦ (action alg (n + 1) ω, - pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) - (ν.prodMkRight _) (arrayMeasure ν) by + let R := reward alg (n + 1) + let H := (hist alg · n) + let A := action alg (n + 1) + let P := fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω + have hP : Measurable P := measurable_pullCount_action_add_one alg ν n + change HasCondDistrib R (fun ω ↦ (H ω, A ω)) (ν.prodMkLeft _) _ + suffices HasCondDistrib R (fun ω ↦ ((A ω, P ω), H ω)) + ((ν.prodMkRight _).prodMkRight _) (arrayMeasure ν) by + -- use that `P` is measurable wrt `(A, H)` to drop it from the conditioning sorry + suffices HasCondDistrib R (fun ω ↦ (A ω, P ω)) (ν.prodMkRight _) (arrayMeasure ν) by + have h_indep : H ⟂ᵢ[(fun ω ↦ (A ω, P ω)), (by fun_prop); arrayMeasure ν] R := + (condIndepFun_reward_hist alg ν n).symm + have h_condDistrib := this.condDistrib_eq + rw [condIndepFun_iff_condDistrib_prod_ae_eq_prodMkRight (by fun_prop) (by fun_prop) + (by fun_prop)] at h_indep + refine ⟨by fun_prop, by fun_prop, ?_⟩ + refine h_indep.trans ?_ + rw [Filter.EventuallyEq, ae_map_iff] at h_condDistrib ⊢ + · simpa only [Kernel.prodMkRight_apply] + · fun_prop + · exact Kernel.measurableSet_eq _ _ + · fun_prop + · exact Kernel.measurableSet_eq _ _ exact hasCondDistrib_reward_pullCount_action alg ν n lemma hasCondDistrib_action (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : From 830bad08e54ef1ee86469e4fc8090a24316feec4 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 14 Jan 2026 13:06:01 +0100 Subject: [PATCH 15/30] only two indepedence sorry left --- LeanBandits/Bandit/Bandit.lean | 71 +++++++------ LeanBandits/ForMathlib/HasCondDistrib.lean | 111 +++++++++++++++++++++ LeanBandits/ForMathlib/IndepFun.lean | 29 ++++++ 3 files changed, 181 insertions(+), 30 deletions(-) diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index dfc9b304..8b334b57 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/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 LeanBandits.ForMathlib.CondIndepFun +import LeanBandits.ForMathlib.IndepFun import LeanBandits.ForMathlib.IndepInfinitePi import LeanBandits.ForMathlib.KernelRepresentation import LeanBandits.SequentialLearning.Deterministic @@ -885,30 +886,6 @@ lemma reward_ae_eq_cond simp only [hω.2] simp [hω.1] -lemma indepFun_cond_of_indepFun {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} - {mγ : MeasurableSpace γ} {μ : Measure α} - {X : α → β} {Y : α → γ} (hXY : X ⟂ᵢ[μ] Y) (hY : Measurable Y) {s : Set γ} - (hs : MeasurableSet s) : - X ⟂ᵢ[μ[|Y ⁻¹' s]] Y := by - by_cases h_zero : μ[|Y ⁻¹' s] = 0 - · simp only [h_zero] - -- missing simp lemma : `X ⟂ᵢ[0] Y` - simp [indepFun_iff_measure_inter_preimage_eq_mul] - rw [cond_eq_zero] at h_zero - push_neg at h_zero - rw [indepFun_iff_measure_inter_preimage_eq_mul] at hXY ⊢ - intro u t hu ht - rw [cond_apply (hs.preimage hY), cond_apply (hs.preimage hY), cond_apply (hs.preimage hY)] - have h_eq : Y ⁻¹' s ∩ (X ⁻¹' u ∩ Y ⁻¹' t) = X ⁻¹' u ∩ Y ⁻¹' (s ∩ t) := by grind - have hsu := hXY u s hu hs - have hust := hXY u (s ∩ t) hu (hs.inter ht) - rw [Set.inter_comm] at hsu - rw [hsu, h_eq, hust] - simp_rw [mul_assoc] - congr 1 - rw [← mul_assoc (μ (Y ⁻¹' s)), ENNReal.mul_inv_cancel h_zero.2 h_zero.1, one_mul] - congr - lemma indepFun_todo {α β γ δ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} {mδ : MeasurableSpace δ} [MeasurableSingletonClass δ] {μ : Measure α} {X : α → β} {Y : α → γ} (hXY : X ⟂ᵢ[μ] Y) (hY : Measurable Y) @@ -1019,21 +996,55 @@ lemma condIndepFun_reward_hist (alg : Algorithm α R) (ν : Kernel α R) [IsMark h_cond.condDistrib_eq exact Measurable.prodMk (by fun_prop) (measurable_pullCount_action_add_one alg ν n) +omit [Countable α] [StandardBorelSpace R] [Nonempty R] in +lemma measurable_pullCount_action_add_one_hist (alg : Algorithm α R) + (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + Measurable[MeasurableSpace.comap (fun ω ↦ (action alg (n + 1) ω, hist alg ω n)) inferInstance] + (fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω) := by + simp_rw [pullCount_eq_sum] + refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + refine measurableSet_eq_fun ?_ (measurable_comp_comap _ measurable_fst) + simp_rw [hist_eq _ _ n] + unfold action + refine Measurable.fst (mγ := inferInstance) ?_ + have : (hist alg · i ⟨i, by grind⟩) = + (fun ω : α × (Iic n → α × R) ↦ ω.2 ⟨i, by grind⟩) ∘ + (fun ω ↦ (action alg (n + 1) ω, fun i : Iic n ↦ hist alg ω i ⟨i, by grind⟩)) := rfl + rw [this] + exact measurable_comp_comap _ (Measurable.prodMk (by fun_prop) (by fun_prop)) + lemma hasCondDistrib_reward' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : HasCondDistrib (reward alg (n + 1)) (fun ω ↦ (hist alg ω n, action alg (n + 1) ω)) (ν.prodMkLeft _) (arrayMeasure ν) := by - let R := reward alg (n + 1) + let R' := reward alg (n + 1) let H := (hist alg · n) let A := action alg (n + 1) let P := fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω have hP : Measurable P := measurable_pullCount_action_add_one alg ν n - change HasCondDistrib R (fun ω ↦ (H ω, A ω)) (ν.prodMkLeft _) _ - suffices HasCondDistrib R (fun ω ↦ ((A ω, P ω), H ω)) + change HasCondDistrib R' (fun ω ↦ (H ω, A ω)) (ν.prodMkLeft _) _ + suffices HasCondDistrib R' (fun ω ↦ (A ω, H ω)) (ν.prodMkRight _) (arrayMeasure ν) by + have h_eq : (fun ω ↦ (H ω, A ω)) = MeasurableEquiv.prodComm ∘ (fun ω ↦ (A ω, H ω)) := rfl + rw [h_eq] + exact this.comp_right (κ := ν.prodMkRight _) _ + suffices HasCondDistrib R' (fun ω ↦ ((A ω, H ω), P ω)) ((ν.prodMkRight _).prodMkRight _) (arrayMeasure ν) by -- use that `P` is measurable wrt `(A, H)` to drop it from the conditioning - sorry - suffices HasCondDistrib R (fun ω ↦ (A ω, P ω)) (ν.prodMkRight _) (arrayMeasure ν) by - have h_indep : H ⟂ᵢ[(fun ω ↦ (A ω, P ω)), (by fun_prop); arrayMeasure ν] R := + have hP_meas : + Measurable[MeasurableSpace.comap (fun ω ↦ (A ω, H ω)) inferInstance] P := + measurable_pullCount_action_add_one_hist alg ν n + obtain ⟨f, hf_meas, hf_eq⟩ := hP_meas.exists_eq_measurable_comp + simp only [hf_eq, Function.comp_apply] at this + rwa [hasCondDistrib_prod_right_iff _ _ hf_meas] at this + suffices HasCondDistrib R' (fun ω ↦ ((A ω, P ω), H ω)) + ((ν.prodMkRight _).prodMkRight _) (arrayMeasure ν) by + let e : ((α × ℕ) × (Iic n → α × R)) ≃ᵐ ((α × (Iic n → α × R)) × ℕ) := + { toFun := fun x ↦ ((x.1.1, x.2), x.1.2) + invFun := fun x ↦ ((x.1.1, x.2), x.1.2) + measurable_toFun := by fun_prop + measurable_invFun := by fun_prop } + exact this.comp_right e + suffices HasCondDistrib R' (fun ω ↦ (A ω, P ω)) (ν.prodMkRight _) (arrayMeasure ν) by + have h_indep : H ⟂ᵢ[(fun ω ↦ (A ω, P ω)), (by fun_prop); arrayMeasure ν] R' := (condIndepFun_reward_hist alg ν n).symm have h_condDistrib := this.condDistrib_eq rw [condIndepFun_iff_condDistrib_prod_ae_eq_prodMkRight (by fun_prop) (by fun_prop) diff --git a/LeanBandits/ForMathlib/HasCondDistrib.lean b/LeanBandits/ForMathlib/HasCondDistrib.lean index fcec9c6f..9a23c90a 100644 --- a/LeanBandits/ForMathlib/HasCondDistrib.lean +++ b/LeanBandits/ForMathlib/HasCondDistrib.lean @@ -69,6 +69,117 @@ lemma HasCondDistrib.snd {Y : α → Ω × Ω'} {κ : Kernel β (Ω × Ω')} [Is rw [Kernel.snd_eq] exact HasCondDistrib.comp h measurable_snd +lemma HasCondDistrib.comp_right [IsFiniteMeasure μ] [IsFiniteKernel κ] (h : HasCondDistrib Y X κ μ) + (f : β ≃ᵐ γ) : + HasCondDistrib Y (f ∘ X) (κ.comap f.symm (by fun_prop)) μ := by + have hY := h.aemeasurable_fst + have hX := h.aemeasurable_snd + refine ⟨h.aemeasurable_fst, by fun_prop, ?_⟩ + have h_eq := h.condDistrib_eq + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop) _] at h_eq ⊢ + calc μ.map (fun ω ↦ ((f ∘ X) ω, Y ω)) + _ = μ.map ((fun p ↦ (f p.1, p.2)) ∘ fun ω ↦ (X ω, Y ω)) := by congr + _ = (μ.map (fun ω ↦ (X ω, Y ω))).map (fun p ↦ (f p.1, p.2)) := by + rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + _ = (μ.map X ⊗ₘ κ).map (fun p ↦ (f p.1, p.2)) := by rw [h_eq] + _ = μ.map (f ∘ X) ⊗ₘ (κ.comap f.symm (by fun_prop)) := by + -- this is probably very inefficient. + have hX_eq : X = f.symm ∘ (f ∘ X) := by ext; simp + conv_lhs => rw [hX_eq] + rw [← AEMeasurable.map_map_of_aemeasurable, Measure.compProd_eq_comp_prod, + ← Measure.deterministic_comp_eq_map (f := f.symm), ← Measure.deterministic_comp_eq_map] + rotate_left + · fun_prop + · fun_prop + · fun_prop + · fun_prop + rw [← Kernel.comp_deterministic_eq_comap, Measure.compProd_eq_comp_prod] + simp_rw [Measure.comp_assoc] + congr 1 + ext c : 1 + rw [Kernel.comp_apply, Kernel.comp_apply, Kernel.prod_apply, Kernel.comp_apply] + simp only [Kernel.deterministic_apply, Kernel.id_apply, Measure.dirac_bind κ.measurable, + Measure.dirac_bind (Kernel.id ×ₖ κ).measurable, Kernel.prod_apply, + Measure.deterministic_comp_eq_map] + ext s hs + rw [Measure.map_apply (by fun_prop) hs, Measure.prod_apply, Measure.prod_apply, + lintegral_dirac', lintegral_dirac'] + · congr + ext + simp + · exact measurable_measure_prodMk_left hs + · exact measurable_measure_prodMk_left (hs.preimage (by fun_prop)) + · exact hs + · exact hs.preimage (by fun_prop) + +lemma HasCondDistrib.prod_right [IsFiniteMeasure μ] [IsFiniteKernel κ] (h : HasCondDistrib Y X κ μ) + {f : β → γ} (hf : Measurable f) : + HasCondDistrib Y (fun a ↦ (X a, f (X a))) (κ.prodMkRight _) μ := by + have hY := h.aemeasurable_fst + have hX := h.aemeasurable_snd + refine ⟨h.aemeasurable_fst, by fun_prop, ?_⟩ + have h_eq := h.condDistrib_eq + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop) _] at h_eq ⊢ + calc μ.map (fun x ↦ ((X x, f (X x)), Y x)) + _ = (μ.map (fun ω ↦ (X ω, Y ω))).map (fun p ↦ ((p.1, f p.1), p.2)) := by + rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + congr + _ = (μ.map X ⊗ₘ κ).map (fun p ↦ ((p.1, f p.1), p.2)) := by rw [h_eq] + _ = (μ.map X).map (fun a ↦ (a, f a)) ⊗ₘ κ.prodMkRight γ := by + rw [Measure.compProd_eq_comp_prod, Measure.compProd_eq_comp_prod, + ← Measure.deterministic_comp_eq_map (f := fun a ↦ (a, f a)), + ← Measure.deterministic_comp_eq_map, Measure.comp_assoc, Measure.comp_assoc] + swap; · fun_prop + swap; · fun_prop + congr 1 + ext b : 1 + rw [Kernel.comp_apply, Kernel.comp_apply, Kernel.prod_apply, Kernel.deterministic_apply, + Kernel.id_apply, Measure.dirac_bind (Kernel.measurable _), Kernel.prod_apply, + Measure.deterministic_comp_eq_map, Kernel.prodMkRight_apply, Kernel.id_apply] + change Measure.map (Prod.map (fun x ↦ (x, f x)) id) ((Measure.dirac b).prod (κ b)) = + (Measure.dirac (b, f b)).prod (κ b) + rw [← Measure.map_prod_map _ _ (by fun_prop) (by fun_prop), Measure.map_id, + Measure.map_dirac (by fun_prop)] + _ = μ.map (fun a ↦ (X a, f (X a))) ⊗ₘ κ.prodMkRight γ := by + rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + congr + +lemma hasCondDistrib_prod_right_iff [IsFiniteMeasure μ] [IsFiniteKernel κ] (X : α → β) (Y : α → Ω) + {f : β → γ} (hf : Measurable f) : + HasCondDistrib Y (fun a ↦ (X a, f (X a))) (κ.prodMkRight _) μ ↔ HasCondDistrib Y X κ μ := by + refine ⟨fun h ↦ ?_, fun h ↦ h.prod_right hf⟩ + have hX : AEMeasurable X μ := by + have := h.aemeasurable_snd + have h_eq : X = (fun p ↦ p.1) ∘ (fun a ↦ (X a, f (X a))) := by ext; simp + rw [h_eq] + exact Measurable.comp_aemeasurable (by fun_prop) (by fun_prop) + have hY := h.aemeasurable_fst + refine ⟨by fun_prop, by fun_prop, ?_⟩ + have h_eq := h.condDistrib_eq + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop) _] at h_eq ⊢ + calc μ.map (fun x ↦ (X x, Y x)) + _ = (μ.map (fun ω ↦ ((X ω, f (X ω)), Y ω))).map (fun p ↦ (p.1.1, p.2)) := by + rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + congr + _ = (μ.map (fun a ↦ (X a, f (X a))) ⊗ₘ κ.prodMkRight γ).map (fun p ↦ (p.1.1, p.2)) := by rw [h_eq] + _ = ((μ.map X).map (fun a ↦ (a, f a)) ⊗ₘ κ.prodMkRight γ).map (fun p ↦ (p.1.1, p.2)) := by + rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + congr + _ = μ.map X ⊗ₘ κ := by + simp_rw [Measure.compProd_eq_comp_prod, + ← Measure.deterministic_comp_eq_map (f := fun a ↦ (a, f a)) (by fun_prop), + ← Measure.deterministic_comp_eq_map (f := fun p : (β × γ) × Ω ↦ (p.1.1, p.2)) (by fun_prop), + Measure.comp_assoc] + congr 1 + ext b : 1 + rw [Kernel.comp_apply, Kernel.comp_apply, Kernel.prod_apply, Kernel.id_apply, + Kernel.deterministic_apply, Measure.dirac_bind (Kernel.measurable _), + Kernel.prod_apply, Measure.deterministic_comp_eq_map, Kernel.prodMkRight_apply, + Kernel.id_apply] + change Measure.map (Prod.map (fun x ↦ x.1) id) ((Measure.dirac (b, f b)).prod (κ b)) = _ + rw [← Measure.map_prod_map _ _ (by fun_prop) (by fun_prop), Measure.map_id, + Measure.map_dirac (by fun_prop)] + lemma HasLaw.prod_of_hasCondDistrib {P : Measure β} [IsFiniteMeasure μ] [IsSFiniteKernel κ] (h1 : HasLaw X P μ) (h2 : HasCondDistrib Y X κ μ) : HasLaw (fun ω ↦ (X ω, Y ω)) (P ⊗ₘ κ) μ := by diff --git a/LeanBandits/ForMathlib/IndepFun.lean b/LeanBandits/ForMathlib/IndepFun.lean index 15901bd4..d0834089 100644 --- a/LeanBandits/ForMathlib/IndepFun.lean +++ b/LeanBandits/ForMathlib/IndepFun.lean @@ -9,6 +9,35 @@ variable {α Ω Ω' E ι : Type*} [Countable ι] {mα : MeasurableSpace α} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} {mE : MeasurableSpace E} {μ ν : Measure Ω} +@[simp] +lemma indepFun_zero_measure {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} (X : α → β) (Y : α → γ) : + X ⟂ᵢ[(0 : Measure α)] Y := by + simp [indepFun_iff_measure_inter_preimage_eq_mul] + +lemma indepFun_cond_of_indepFun {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {μ : Measure α} + {X : α → β} {Y : α → γ} (hXY : X ⟂ᵢ[μ] Y) (hY : Measurable Y) {s : Set γ} + (hs : MeasurableSet s) : + X ⟂ᵢ[μ[|Y ⁻¹' s]] Y := by + by_cases h_zero : μ[|Y ⁻¹' s] = 0 + · simp [h_zero] + rw [cond_eq_zero] at h_zero + push_neg at h_zero -- `h_zero : μ (Y ⁻¹' s) ≠ ⊤ ∧ μ (Y ⁻¹' s) ≠ 0` + rw [indepFun_iff_measure_inter_preimage_eq_mul] at hXY ⊢ + intro u t hu ht + rw [cond_apply (hs.preimage hY), cond_apply (hs.preimage hY), cond_apply (hs.preimage hY)] + have h_eq : Y ⁻¹' s ∩ (X ⁻¹' u ∩ Y ⁻¹' t) = X ⁻¹' u ∩ Y ⁻¹' (s ∩ t) := by grind + have hsu : μ (X ⁻¹' u ∩ Y ⁻¹' s) = μ (X ⁻¹' u) * μ (Y ⁻¹' s) := hXY u s hu hs + rw [Set.inter_comm] at hsu + have hust : μ (X ⁻¹' u ∩ Y ⁻¹' (s ∩ t)) = μ (X ⁻¹' u) * μ (Y ⁻¹' (s ∩ t)) := + hXY u (s ∩ t) hu (hs.inter ht) + rw [hsu, h_eq, hust] + simp_rw [mul_assoc] + congr 1 + rw [← mul_assoc (μ (Y ⁻¹' s)), ENNReal.mul_inv_cancel h_zero.2 h_zero.1, one_mul] + congr + lemma iIndepFun_nat_iff_forall_indepFun [IsProbabilityMeasure μ] {X : ℕ → Ω → E} (hX : ∀ n, AEMeasurable (X n) μ) : iIndepFun X μ ↔ ∀ n, X (n + 1) ⟂ᵢ[μ] fun ω (i : Iic n) ↦ X i ω := by From adfc69305d2e3bee2c8daad4c9a92cb5d4aff607 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 14 Jan 2026 14:11:17 +0100 Subject: [PATCH 16/30] lint --- LeanBandits/Bandit/Bandit.lean | 141 +++++++----------- LeanBandits/Bandit/RewardByCountMeasure.lean | 20 +-- LeanBandits/BanditAlgorithms/UCB.lean | 6 +- LeanBandits/ForMathlib/HasCondDistrib.lean | 2 + .../ForMathlib/KernelRepresentation.lean | 8 + LeanBandits/ForMathlib/StandardBorel.lean | 18 +++ LeanBandits/ForMathlib/Traj.lean | 2 + LeanBandits/SequentialLearning/Algorithm.lean | 5 + .../SequentialLearning/FiniteActions.lean | 18 +-- .../SequentialLearning/StationaryEnv.lean | 6 +- 10 files changed, 113 insertions(+), 113 deletions(-) create mode 100644 LeanBandits/ForMathlib/StandardBorel.lean diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 8b334b57..0fe71146 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -7,6 +7,7 @@ import LeanBandits.ForMathlib.CondIndepFun import LeanBandits.ForMathlib.IndepFun import LeanBandits.ForMathlib.IndepInfinitePi import LeanBandits.ForMathlib.KernelRepresentation +import LeanBandits.ForMathlib.StandardBorel import LeanBandits.SequentialLearning.Deterministic import LeanBandits.SequentialLearning.StationaryEnv import LeanBandits.SequentialLearning.FiniteActions @@ -181,43 +182,24 @@ lemma action_detAlgorithm_ae_eq [StandardBorelSpace α] [Nonempty α] IT.action (n + 1) =ᵐ[𝔓t] fun h ↦ nextaction n (fun i ↦ h i) := IT.action_detAlgorithm_ae_eq n -example [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] : - ∀ᵐ h ∂(𝔓t), IT.action 0 h = action0 ∧ - ∀ n, IT.action (n + 1) h = nextaction n (fun i ↦ h i) := by - rw [eventually_and, ae_all_iff] - exact ⟨action_zero_detAlgorithm, action_detAlgorithm_ae_eq⟩ - end DetAlgorithm namespace ArrayModel open unitInterval -section Aux - -theorem representation_measure {β : Type*} {mβ : MeasurableSpace β} - [Nonempty β] [StandardBorelSpace β] - (μ : Measure β) [IsProbabilityMeasure μ] : - ∃ (f : I → β), Measurable f ∧ volume.map f = μ := by - obtain ⟨f, hf_meas, hf_map⟩ := Kernel.representation (Kernel.const Unit μ) - specialize hf_map ⟨⟩ - exact ⟨f ⟨⟩, by fun_prop, by simpa⟩ - -end Aux - variable (α R) in +/-- Probability space for the array model of stochastic bandits. -/ def probSpace : Type _ := (ℕ → I) × (ℕ → α → R) instance {α R : Type*} [MeasurableSpace R] : MeasurableSpace (probSpace α R) := inferInstanceAs (MeasurableSpace ((ℕ → I) × (ℕ → α → R))) -instance {α R : Type*} [MeasurableSpace α] [Countable α] - [MeasurableSpace R] [StandardBorelSpace R] [Nonempty R] : - StandardBorelSpace (probSpace α R) := by - unfold probSpace - infer_instance +instance {α R : Type*} [Countable α] [MeasurableSpace R] [StandardBorelSpace R] : + StandardBorelSpace (probSpace α R) := + inferInstanceAs (StandardBorelSpace ((ℕ → I) × (ℕ → α → R))) +/-- Probability measure for the array model of stochastic bandits. -/ noncomputable def arrayMeasure (ν : Kernel α R) : Measure (probSpace α R) := (Measure.infinitePi fun _ ↦ volume).prod (Bandit.streamMeasure ν) @@ -227,6 +209,7 @@ instance (ν : Kernel α R) [IsMarkovKernel ν] : IsProbabilityMeasure (arrayMea variable [Nonempty α] [StandardBorelSpace α] +/-- The initial action is the image of a uniform random variable by this function. -/ noncomputable def initAlgFunction (alg : Algorithm α R) : I → α := (representation_measure alg.p0).choose @@ -238,6 +221,7 @@ lemma initAlgFunction_map (alg : Algorithm α R) : volume.map (initAlgFunction a lemma measurable_initAlgFunction (alg : Algorithm α R) : Measurable (initAlgFunction alg) := (representation_measure alg.p0).choose_spec.1 +/-- The next action is the image of the history and a uniform random variable by this function. -/ noncomputable def algFunction (alg : Algorithm α R) (n : ℕ) : (Iic n → α × R) → I → α := @@ -252,6 +236,7 @@ lemma measurable_algFunction (alg : Algorithm α R) (n : ℕ) : Measurable (Function.uncurry (algFunction alg n)) := (Kernel.representation (alg.policy n)).choose_spec.1 +/-- History of actions and rewards up to time `n` in the array model. -/ noncomputable def hist [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) : (n : ℕ) → Iic n → α × R | 0 => fun _ ↦ (initAlgFunction alg (ω.1 0), ω.2 0 (initAlgFunction alg (ω.1 0))) @@ -269,8 +254,7 @@ lemma hist_add_one [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) let a : α := algFunction alg n (hist alg ω n) (ω.1 (n + 1)) hist alg ω (n + 1) = fun (i : Iic (n + 1)) ↦ if hin : i ≤ n then hist alg ω n ⟨i, by simp [hin]⟩ - else (a, ω.2 (pullCount' n (hist alg ω n) a) a) := - rfl + else (a, ω.2 (pullCount' n (hist alg ω n) a) a) := rfl lemma hist_eq [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) (n : ℕ) : hist alg ω n = fun i : Iic n ↦ hist alg ω i ⟨i.1, by simp⟩ := by @@ -308,10 +292,6 @@ lemma measurable_action_add_one' [DecidableEq α] {alg : Algorithm α R} (n : ℕ) (h : Measurable (hist alg · n)) : Measurable (fun x ↦ algFunction alg n (hist alg x n) (x.1 (n + 1))) := by fun_prop -instance : MeasurableEq α := by - letI := upgradeStandardBorel α - infer_instance - lemma measurable_pullCount'_action_add_one [DecidableEq α] {alg : Algorithm α R} (n : ℕ) (h_hist : Measurable (hist alg · n)) : Measurable (fun x ↦ @@ -349,6 +329,7 @@ lemma measurable_hist [DecidableEq α] [Countable α] (alg : Algorithm α R) (n refine Measurable.prodMk ?_ (by fun_prop) exact measurable_pullCount'_action_add_one n hn +/-- Action taken at time `n` in the array model. -/ noncomputable def action [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : α := (hist alg ω n ⟨n, by simp⟩).1 @@ -368,6 +349,7 @@ lemma action_add_one_eq [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : lemma measurable_action [DecidableEq α] [Countable α] (alg : Algorithm α R) (n : ℕ) : Measurable (action alg n) := by unfold action; fun_prop +/-- Reward received at time `n` in the array model. -/ noncomputable def reward [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : R := (hist alg ω n ⟨n, by simp⟩).2 @@ -477,13 +459,14 @@ lemma hasCondDistrib_reward_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMa · simp · rwa [Measure.map_apply (by fun_prop) (by simp)] at ha +omit [DecidableEq α] in lemma indepFun_fst_add_one_aux (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : (fun ω ↦ ω.1 (n + 1)) ⟂ᵢ[arrayMeasure ν] (fun ω ↦ (fun (i : Iic n) ↦ ω.1 i, ω.2)) := by rw [indepFun_iff_map_prod_eq_prod_map_map (by fun_prop) (by fun_prop)] sorry omit [StandardBorelSpace R] [Nonempty R] in -lemma measurable_hist_todo (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : +lemma measurable_hist_todo (alg : Algorithm α R) (n : ℕ) : Measurable[MeasurableSpace.comap (fun ω ↦ (fun (i : Iic n) ↦ ω.1 i, ω.2)) inferInstance] (hist alg · n) := by induction n with @@ -560,7 +543,7 @@ lemma measurable_hist_todo (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKe lemma indepFun_fst_add_one_hist (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : IndepFun (fun ω ↦ ω.1 (n + 1)) (hist alg · n) (arrayMeasure ν) := - (indepFun_fst_add_one_aux alg ν n).of_measurable_right (measurable_hist_todo alg ν n) + (indepFun_fst_add_one_aux alg ν n).of_measurable_right (measurable_hist_todo alg n) lemma hasCondDistrib_action' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : HasCondDistrib (action alg (n + 1)) (hist alg · n) (alg.policy n) (arrayMeasure ν) := by @@ -618,8 +601,7 @@ lemma hasCondDistrib_action' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkov · exact hs.preimage (by fun_prop) omit [Countable α] [StandardBorelSpace R] [Nonempty R] in -lemma hist_congr (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) - {ω ω' : probSpace α R} +lemma hist_congr (alg : Algorithm α R) (n : ℕ) {ω ω' : probSpace α R} (hω1 : ∀ i ≤ n, ω.1 i = ω'.1 i) (hω2 : ∀ i a, i < pullCount (action alg) a (n + 1) ω → ω.2 i a = ω'.2 i a) : hist alg ω n = hist alg ω' n := by @@ -653,27 +635,28 @@ lemma hist_congr (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] ( grind -- very bad name +/-- All random variables in the space, except for the unseen rewards for action `a` after +time `n`. -/ noncomputable -def truePast (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] - (a : α) (n : ℕ) (ω : probSpace α R) : +def truePast (alg : Algorithm α R) (a : α) (n : ℕ) (ω : probSpace α R) : probSpace α R := (ω.1, fun i b ↦ if b = a then ω.2 (min i ((pullCount (action alg) a (n + 1) ω) - 1)) a else ω.2 i b) omit [Countable α] [StandardBorelSpace R] [Nonempty R] in -lemma truePast_eq_of_pullCount_eq (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] +lemma truePast_eq_of_pullCount_eq (alg : Algorithm α R) (a : α) (n m : ℕ) (ω : probSpace α R) (h_pc : pullCount (action alg) a (n + 1) ω = m) : - truePast alg ν a n ω = (ω.1, fun i b ↦ if b = a then ω.2 (min i (m - 1)) a else ω.2 i b) := by + truePast alg a n ω = (ω.1, fun i b ↦ if b = a then ω.2 (min i (m - 1)) a else ω.2 i b) := by simp [truePast, h_pc] omit [StandardBorelSpace R] [Nonempty R] in -lemma measurable_hist_truePast (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] +lemma measurable_hist_truePast (alg : Algorithm α R) (a : α) (n : ℕ) : - Measurable[MeasurableSpace.comap (truePast alg ν a n) inferInstance] (hist alg · n) := by - have h_eq : (hist alg · n) = (hist alg · n) ∘ (truePast alg ν a n) := by + Measurable[MeasurableSpace.comap (truePast alg a n) inferInstance] (hist alg · n) := by + have h_eq : (hist alg · n) = (hist alg · n) ∘ (truePast alg a n) := by ext ω : 1 - refine hist_congr alg ν n (fun _ _ ↦ rfl) fun i b hi ↦ ?_ + refine hist_congr alg n (fun _ _ ↦ rfl) fun i b hi ↦ ?_ by_cases hb : b = a · subst hb simp only [truePast, ↓reduceIte] @@ -686,30 +669,29 @@ lemma measurable_hist_truePast (alg : Algorithm α R) (ν : Kernel α R) [IsMark omit [StandardBorelSpace R] [Nonempty R] in lemma measurable_action_add_one_truePast (alg : Algorithm α R) - (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (n : ℕ) : - Measurable[MeasurableSpace.comap (truePast alg ν a n) inferInstance] + (a : α) (n : ℕ) : + Measurable[MeasurableSpace.comap (truePast alg a n) inferInstance] (action alg (n + 1)) := by rw [action_add_one_eq] - change Measurable[MeasurableSpace.comap (truePast alg ν a n) inferInstance] + change Measurable[MeasurableSpace.comap (truePast alg a n) inferInstance] ((fun p ↦ algFunction alg n p.1 p.2) ∘ (fun ω ↦ (hist alg ω n, ω.1 (n + 1)))) refine (measurable_algFunction alg n).comp (Measurable.prodMk ?_ ?_) - · exact measurable_hist_truePast alg ν a n + · exact measurable_hist_truePast alg a n · have : (fun ω ↦ ω.1 (n + 1)) = - (fun (p : probSpace α R) ↦ p.1 (n + 1)) ∘ (truePast alg ν a n) := rfl + (fun (p : probSpace α R) ↦ p.1 (n + 1)) ∘ (truePast alg a n) := rfl rw [this] exact Measurable.comp (by fun_prop) (Measurable.of_comap_le le_rfl) omit [StandardBorelSpace R] [Nonempty R] in -lemma measurable_pullCount_add_one_truePast (alg : Algorithm α R) - (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (n : ℕ) : - Measurable[MeasurableSpace.comap (truePast alg ν a n) inferInstance] +lemma measurable_pullCount_add_one_truePast (alg : Algorithm α R) (a : α) (n : ℕ) : + Measurable[MeasurableSpace.comap (truePast alg a n) inferInstance] (pullCount (action alg) a (n + 1)) := by - change Measurable[MeasurableSpace.comap (truePast alg ν a n) inferInstance] + change Measurable[MeasurableSpace.comap (truePast alg a n) inferInstance] (fun ω ↦ pullCount (action alg) a (n + 1) ω) simp_rw [pullCount_eq_sum] refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) refine (measurableSet_singleton _).preimage ?_ - have h_meas := measurable_hist_truePast alg ν a n + have h_meas := measurable_hist_truePast alg a n simp_rw [hist_eq _ _ n, @measurable_pi_iff] at h_meas exact (h_meas ⟨i, by grind⟩).fst @@ -721,13 +703,13 @@ lemma indepFun_snd_apply_aux (alg : Algorithm α R) omit [Countable α] [StandardBorelSpace R] [Nonempty R] in lemma stepsUntil_congr_aux (alg : Algorithm α R) - (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m n : ℕ) {ω ω' : probSpace α R} + (a : α) (m n : ℕ) {ω ω' : probSpace α R} (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) (h_eq : action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m) : action alg (n + 1) ω' = a ∧ pullCount (action alg) a (n + 1) ω' = m := by obtain ⟨h_action, h_pc⟩ := h_eq - have h_hist := hist_congr alg ν n (ω := ω) (ω' := ω') (by grind) fun i b hi ↦ ?_ + have h_hist := hist_congr alg n (ω := ω) (ω' := ω') (by grind) fun i b hi ↦ ?_ swap · rcases eq_or_ne b a with (rfl | hba) · refine hω2_eq i ?_ @@ -747,18 +729,16 @@ lemma stepsUntil_congr_aux (alg : Algorithm α R) rw [h_hist] omit [Countable α] [StandardBorelSpace R] [Nonempty R] in -lemma stepsUntil_congr (alg : Algorithm α R) - (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m n : ℕ) {ω ω' : probSpace α R} +lemma stepsUntil_congr (alg : Algorithm α R) (a : α) (m n : ℕ) {ω ω' : probSpace α R} (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) : (action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m) ↔ (action alg (n + 1) ω' = a ∧ pullCount (action alg) a (n + 1) ω' = m) := - ⟨stepsUntil_congr_aux alg ν a m n hω1 hω2_ne hω2_eq, - stepsUntil_congr_aux alg ν a m n (by grind) (by grind) (by grind)⟩ + ⟨stepsUntil_congr_aux alg a m n hω1 hω2_ne hω2_eq, + stepsUntil_congr_aux alg a m n (by grind) (by grind) (by grind)⟩ omit [Countable α] [StandardBorelSpace R] [Nonempty R] in -lemma stepsUntil_indicator_congr (alg : Algorithm α R) - (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m n : ℕ) {ω ω' : probSpace α R} +lemma stepsUntil_indicator_congr (alg : Algorithm α R) (a : α) (m n : ℕ) {ω ω' : probSpace α R} (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) : {ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}.indicator (fun _ ↦ 1) @@ -766,11 +746,10 @@ lemma stepsUntil_indicator_congr (alg : Algorithm α R) {ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}.indicator (fun _ ↦ 1) ω' := by simp only [Set.indicator_apply, Set.mem_setOf_eq] - simp_rw [stepsUntil_congr alg ν a m n hω1 hω2_ne hω2_eq] + simp_rw [stepsUntil_congr alg a m n hω1 hω2_ne hω2_eq] omit [StandardBorelSpace R] [Nonempty R] in -lemma measurable_stepsUntil (alg : Algorithm α R) - (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m n : ℕ) : +lemma measurable_stepsUntil (alg : Algorithm α R) (a : α) (m n : ℕ) : Measurable[MeasurableSpace.comap (fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b)) inferInstance] (({ω | action alg (n + 1) ω = a ∧ @@ -780,7 +759,7 @@ lemma measurable_stepsUntil (alg : Algorithm α R) have h_eq : f = f ∘ fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b) := by ext ω - exact stepsUntil_indicator_congr alg ν a m n (by grind) (by grind) (by grind) + exact stepsUntil_indicator_congr alg a m n (by grind) (by grind) (by grind) change Measurable[MeasurableSpace.comap (fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b)) inferInstance] f rw [h_eq] @@ -794,12 +773,11 @@ lemma indepFun_snd_apply_pullCount_action (alg : Algorithm α R) (fun ω ↦ ω.2 m a) ⟂ᵢ[arrayMeasure ν] ({ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}).indicator (fun _ ↦ 1) := - (indepFun_snd_apply_aux alg ν a m n).of_measurable_right (measurable_stepsUntil alg ν a m n) + (indepFun_snd_apply_aux alg ν a m n).of_measurable_right (measurable_stepsUntil alg a m n) omit [StandardBorelSpace R] [Nonempty R] in @[fun_prop] -lemma measurable_pullCount_action_add_one (alg : Algorithm α R) - (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : +lemma measurable_pullCount_action_add_one (alg : Algorithm α R) (n : ℕ) : Measurable (fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω) := by change Measurable ((fun p : (probSpace α R) × α ↦ pullCount (action alg) p.2 (n + 1) p.1) ∘ (fun ω : probSpace α R ↦ (ω, action alg (n + 1) ω))) @@ -859,18 +837,8 @@ lemma hasCondDistrib_reward_pullCount_action intro ha simp [ha] --- lemma hasCondDistrib_reward_hist_action_pullCount' --- (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (n m : ℕ) : --- HasCondDistrib (reward alg (n + 1)) --- (fun ω ↦ (hist alg ω n, {ω' | action alg (n + 1) ω' = a ∧ --- pullCount (action alg) (action alg (n + 1) ω') (n + 1) ω' = m}.indicator --- (fun _ ↦ (1 : ℕ)) ω)) --- (Kernel.const _ (ν a)) (arrayMeasure ν) := by --- sorry - omit [StandardBorelSpace R] [Nonempty R] in -lemma reward_ae_eq_cond - (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (n m : ℕ) : +lemma reward_ae_eq_cond (alg : Algorithm α R) (ν : Kernel α R) (a : α) (n m : ℕ) : reward alg (n + 1) =ᵐ[(arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}]] (fun ω ↦ ω.2 m a) := by @@ -900,9 +868,9 @@ lemma indepFun_snd_hist_cond (alg : Algorithm α R) (fun ω ↦ ω.2 m a) ⟂ᵢ[(arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}]] (hist alg · n) := by - have h_meas := measurable_hist_truePast alg ν a n + have h_meas := measurable_hist_truePast alg a n refine IndepFun.of_measurable_right ?_ h_meas - have h_ae_eq : truePast alg ν a n =ᵐ[(arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, + have h_ae_eq : truePast alg a n =ᵐ[(arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}]] (fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b)) := by refine ae_cond_of_forall_mem ?_ fun x hx ↦ ?_ @@ -932,7 +900,7 @@ lemma indepFun_snd_hist_cond (alg : Algorithm α R) Classical.not_imp, Decidable.not_not, and_congr_right_iff] intro ha simp [ha] - have h_meas := measurable_stepsUntil alg ν a m n + have h_meas := measurable_stepsUntil alg a m n obtain ⟨f, hf, hf_eq⟩ := h_meas.exists_eq_measurable_comp simp_rw [hf_eq] refine indepFun_todo (Z := f) (z := 1) ?_ ?_ hf @@ -988,17 +956,16 @@ lemma hasCondDistrib_reward_hist_action_pullCount lemma condIndepFun_reward_hist (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : (reward alg (n + 1)) ⟂ᵢ[(fun ω ↦ (action alg (n + 1) ω, pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)), - Measurable.prodMk (by fun_prop) (measurable_pullCount_action_add_one alg ν n); + Measurable.prodMk (by fun_prop) (measurable_pullCount_action_add_one alg n); arrayMeasure ν] (hist alg · n) := by have h_cond := hasCondDistrib_reward_hist_action_pullCount alg ν n refine condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (by fun_prop) (by fun_prop) ?_ h_cond.condDistrib_eq - exact Measurable.prodMk (by fun_prop) (measurable_pullCount_action_add_one alg ν n) + exact Measurable.prodMk (by fun_prop) (measurable_pullCount_action_add_one alg n) omit [Countable α] [StandardBorelSpace R] [Nonempty R] in -lemma measurable_pullCount_action_add_one_hist (alg : Algorithm α R) - (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : +lemma measurable_pullCount_action_add_one_hist (alg : Algorithm α R) (n : ℕ) : Measurable[MeasurableSpace.comap (fun ω ↦ (action alg (n + 1) ω, hist alg ω n)) inferInstance] (fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω) := by simp_rw [pullCount_eq_sum] @@ -1020,7 +987,7 @@ lemma hasCondDistrib_reward' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkov let H := (hist alg · n) let A := action alg (n + 1) let P := fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω - have hP : Measurable P := measurable_pullCount_action_add_one alg ν n + have hP : Measurable P := measurable_pullCount_action_add_one alg n change HasCondDistrib R' (fun ω ↦ (H ω, A ω)) (ν.prodMkLeft _) _ suffices HasCondDistrib R' (fun ω ↦ (A ω, H ω)) (ν.prodMkRight _) (arrayMeasure ν) by have h_eq : (fun ω ↦ (H ω, A ω)) = MeasurableEquiv.prodComm ∘ (fun ω ↦ (A ω, H ω)) := rfl @@ -1031,7 +998,7 @@ lemma hasCondDistrib_reward' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkov -- use that `P` is measurable wrt `(A, H)` to drop it from the conditioning have hP_meas : Measurable[MeasurableSpace.comap (fun ω ↦ (A ω, H ω)) inferInstance] P := - measurable_pullCount_action_add_one_hist alg ν n + measurable_pullCount_action_add_one_hist alg n obtain ⟨f, hf_meas, hf_eq⟩ := hP_meas.exists_eq_measurable_comp simp only [hf_eq, Function.comp_apply] at this rwa [hasCondDistrib_prod_right_iff _ _ hf_meas] at this diff --git a/LeanBandits/Bandit/RewardByCountMeasure.lean b/LeanBandits/Bandit/RewardByCountMeasure.lean index 78a96464..07b20131 100644 --- a/LeanBandits/Bandit/RewardByCountMeasure.lean +++ b/LeanBandits/Bandit/RewardByCountMeasure.lean @@ -50,7 +50,7 @@ local notation "𝔓t" => Bandit.trajMeasure alg ν local notation "𝔓" => Bandit.measure alg ν omit [DecidableEq α] in -lemma condDistrib_reward'' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +lemma condDistrib_reward'' [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (n : ℕ) : 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1; 𝔓'] =ᵐ[(𝔓').map (fun ω ↦ A n ω.1)] ν := by have hA := h.measurable_A @@ -67,7 +67,7 @@ lemma condDistrib_reward'' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] rw [h_prod, h_eq] omit [DecidableEq α] in -lemma reward_cond_action [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +lemma reward_cond_action [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n : ℕ) (hμa : (𝔓').map (fun ω ↦ A n ω.1) {a} ≠ 0) : 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1 ← a; 𝔓'] = ν a := by @@ -83,7 +83,7 @@ lemma reward_cond_action [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] rw [h_ra] at h_eq exact h_eq.symm -lemma condIndepFun_reward_stepsUntil_action' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +lemma condIndepFun_reward_stepsUntil_action' [StandardBorelSpace Ω] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m n : ℕ) : R n ⟂ᵢ[A n, h.measurable_A n; P] {ω | stepsUntil A a m ω = ↑n}.indicator (fun _ ↦ 1) := by -- the indicator of `stepsUntil ... = n` is a function of `hist (n-1)` and `action n`. @@ -102,7 +102,7 @@ lemma condIndepFun_reward_stepsUntil_action' [StandardBorelSpace Ω] [Nonempty refine h_indep.of_measurable_right (hX := hA n) ?_ exact measurable_comap_indicator_stepsUntil_eq hA hR a m n -lemma condIndepFun_reward_stepsUntil_action [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +lemma condIndepFun_reward_stepsUntil_action [StandardBorelSpace Ω] [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m n : ℕ) : CondIndepFun (mα.comap (fun ω ↦ A n ω.1)) ((h.measurable_A n).comp measurable_fst).comap_le @@ -113,7 +113,7 @@ lemma condIndepFun_reward_stepsUntil_action [StandardBorelSpace Ω] [Nonempty Ω (measurable_indicator_stepsUntil_eq hA hR a m n) (by fun_prop) (by fun_prop) (condIndepFun_reward_stepsUntil_action' h a m n) -lemma reward_cond_stepsUntil [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +lemma reward_cond_stepsUntil [StandardBorelSpace Ω] [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m n : ℕ) (hm : m ≠ 0) (hμn : 𝔓' ((fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n}) ≠ 0) : 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ stepsUntil A a m ω.1 ← ↑n; 𝔓'] = ν a := by @@ -159,7 +159,7 @@ lemma reward_cond_stepsUntil [StandardBorelSpace Ω] [Nonempty Ω] [Countable α /-- The conditional distribution of the reward received at the `m`-th pull of action `a` given the time at which number of pulls is `m` is the constant kernel with value `ν a`. -/ -theorem condDistrib_rewardByCount_stepsUntil [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +theorem condDistrib_rewardByCount_stepsUntil [StandardBorelSpace Ω] [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m : ℕ) (hm : m ≠ 0) : condDistrib (rewardByCount A R a m) (fun ω ↦ stepsUntil A a m ω.1) 𝔓' =ᵐ[(𝔓').map (fun ω ↦ stepsUntil A a m ω.1)] Kernel.const _ (ν a) := by @@ -192,7 +192,7 @@ theorem condDistrib_rewardByCount_stepsUntil [StandardBorelSpace Ω] [Nonempty rwa [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] at hn /-- The reward received at the `m`-th pull of action `a` has law `ν a`. -/ -lemma hasLaw_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +lemma hasLaw_rewardByCount [StandardBorelSpace Ω] [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m : ℕ) (hm : m ≠ 0) : HasLaw (rewardByCount A R a m) (ν a) 𝔓' where aemeasurable := (measurable_rewardByCount h.measurable_A h.measurable_R a m).aemeasurable @@ -214,7 +214,7 @@ lemma hasLaw_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] Measure.isProbabilityMeasure_map (by fun_prop) simp -lemma identDistrib_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +lemma identDistrib_rewardByCount [StandardBorelSpace Ω] [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n m : ℕ) (hn : n ≠ 0) (hm : m ≠ 0) : IdentDistrib (rewardByCount A R a n) (rewardByCount A R a m) 𝔓' 𝔓' where @@ -222,14 +222,14 @@ lemma identDistrib_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countabl aemeasurable_snd := (measurable_rewardByCount h.measurable_A h.measurable_R a m).aemeasurable map_eq := by rw [(hasLaw_rewardByCount h a n hn).map_eq, (hasLaw_rewardByCount h a m hm).map_eq] -lemma identDistrib_rewardByCount_id [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +lemma identDistrib_rewardByCount_id [StandardBorelSpace Ω] [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n : ℕ) (hn : n ≠ 0) : IdentDistrib (rewardByCount A R a n) id 𝔓' (ν a) where aemeasurable_fst := (measurable_rewardByCount h.measurable_A h.measurable_R a n).aemeasurable aemeasurable_snd := Measurable.aemeasurable <| by fun_prop map_eq := by rw [(hasLaw_rewardByCount h a n hn).map_eq, Measure.map_id] -lemma identDistrib_rewardByCount_eval [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +lemma identDistrib_rewardByCount_eval [StandardBorelSpace Ω] [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n m : ℕ) (hn : n ≠ 0) : IdentDistrib (rewardByCount A R a n) (fun ω ↦ ω m a) 𝔓' (Bandit.streamMeasure ν) := (identDistrib_rewardByCount_id h a n hn).trans diff --git a/LeanBandits/BanditAlgorithms/UCB.lean b/LeanBandits/BanditAlgorithms/UCB.lean index bc4d13c2..66ad5f4c 100644 --- a/LeanBandits/BanditAlgorithms/UCB.lean +++ b/LeanBandits/BanditAlgorithms/UCB.lean @@ -214,8 +214,7 @@ lemma pullCount_arm_le [Nonempty (Fin K)] (hc : 0 ≤ c) · have : 0 ≤ log (n + 1) := by simp [log_nonneg] positivity -lemma todo [Nonempty (Fin K)] - (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma todo (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n k : ℕ) (hk : k ≠ 0) : Bandit.streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2) := by @@ -237,8 +236,7 @@ lemma todo [Nonempty (Fin K)] mul_assoc (k : ℝ), sqrt_mul (x := (k : ℝ)) (by positivity), mul_comm] _ ≤ 1 / (n + 1) ^ (c / 2) := prob_sum_le_sqrt_log hν hc a k hk -lemma todo' [Nonempty (Fin K)] - (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma todo' (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n k : ℕ) (hk : k ≠ 0) : Bandit.streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - √(c * log (n + 1) / k)} ≤ diff --git a/LeanBandits/ForMathlib/HasCondDistrib.lean b/LeanBandits/ForMathlib/HasCondDistrib.lean index 9a23c90a..a151124e 100644 --- a/LeanBandits/ForMathlib/HasCondDistrib.lean +++ b/LeanBandits/ForMathlib/HasCondDistrib.lean @@ -20,6 +20,8 @@ variable {α β γ Ω Ω' : Type*} {mΩ' : MeasurableSpace Ω'} [StandardBorelSpace Ω'] [Nonempty Ω'] {μ : Measure α} {X : α → β} {Y : α → Ω} {κ : Kernel β Ω} +/-- Predicate stating that the conditional distribution of `Y` given `X` under the measure `μ` +is equal to the kernel `κ`. -/ structure HasCondDistrib (Y : α → Ω) (X : α → β) (κ : Kernel β Ω) (μ : Measure α) [IsFiniteMeasure μ] : Prop where aemeasurable_fst : AEMeasurable Y μ := by fun_prop diff --git a/LeanBandits/ForMathlib/KernelRepresentation.lean b/LeanBandits/ForMathlib/KernelRepresentation.lean index 68366915..aebf144b 100644 --- a/LeanBandits/ForMathlib/KernelRepresentation.lean +++ b/LeanBandits/ForMathlib/KernelRepresentation.lean @@ -144,3 +144,11 @@ theorem representation {β : Type*} [Nonempty β] [MeasurableSpace β] [Standard κ.embedding_representation (measurableEmbedding_sigmoid_comp_embeddingReal β) end ProbabilityTheory.Kernel + +theorem ProbabilityTheory.representation_measure {β : Type*} {mβ : MeasurableSpace β} + [Nonempty β] [StandardBorelSpace β] + (μ : Measure β) [IsProbabilityMeasure μ] : + ∃ (f : I → β), Measurable f ∧ volume.map f = μ := by + obtain ⟨f, hf_meas, hf_map⟩ := Kernel.representation (Kernel.const Unit μ) + specialize hf_map ⟨⟩ + exact ⟨f ⟨⟩, by fun_prop, by simpa⟩ diff --git a/LeanBandits/ForMathlib/StandardBorel.lean b/LeanBandits/ForMathlib/StandardBorel.lean new file mode 100644 index 00000000..e31407b7 --- /dev/null +++ b/LeanBandits/ForMathlib/StandardBorel.lean @@ -0,0 +1,18 @@ +/- +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.MeasureTheory.Constructions.Polish.Basic + +/-! +# Properties of standard Borel spaces +-/ + +open MeasureTheory + +variable {Ω : Type*} {mΩ : MeasurableSpace Ω} + +instance [StandardBorelSpace Ω] : MeasurableEq Ω := by + letI := upgradeStandardBorel Ω + infer_instance diff --git a/LeanBandits/ForMathlib/Traj.lean b/LeanBandits/ForMathlib/Traj.lean index 93333263..cc0055ec 100644 --- a/LeanBandits/ForMathlib/Traj.lean +++ b/LeanBandits/ForMathlib/Traj.lean @@ -29,6 +29,8 @@ lemma traj_zero_map_eval_zero : rw [← Kernel.traj_map_frestrictLe, ← Kernel.map_comp_right _ (by fun_prop) (by fun_prop)] rfl +/-- Measurable equivalence between a product up to `n + 1` and the pair of the product up to `n` and +the space at `n + 1`. -/ def _root_.MeasurableEquiv.IicSuccProd (X : ℕ → Type*) [∀ n, MeasurableSpace (X n)] (n : ℕ) : MeasurableEquiv (Π i : Iic (n + 1), X i) ((Π i : Iic n, X i) × X (n + 1)) := (MeasurableEquiv.IicProdIoc (Nat.le_succ n)).symm.trans diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index 5e965465..ac1de8c1 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -62,6 +62,7 @@ section IsAlgEnvSeq variable {A : ℕ → Ω → α} {R' : ℕ → Ω → R} {alg : Algorithm α R} {env : Environment α R} {P : Measure Ω} [IsFiniteMeasure P] +/-- Step of the algorithm-environment sequence: the action-reward pair at time `n`. -/ def IsAlgEnvSeq.step (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (n : ℕ) (ω : Ω) : α × R := (A n ω, R' n ω) @@ -72,6 +73,7 @@ lemma IsAlgEnvSeq.measurable_step (n : ℕ) (hA : Measurable (A n)) unfold IsAlgEnvSeq.step fun_prop +/-- History of the algorithm-environment sequence up to time `n`. -/ def IsAlgEnvSeq.hist (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (n : ℕ) (ω : Ω) : Iic n → α × R := fun i ↦ (A i ω, R' i ω) @@ -91,6 +93,8 @@ lemma IsAlgEnvSeq.fst_eval_comp_hist (n : ℕ) : lemma IsAlgEnvSeq.snd_eval_comp_hist (n : ℕ) : (fun x ↦ (x ⟨n, by simp⟩).2) ∘ (hist A R' n) = R' n := rfl +/-- An algorithm-environment sequence: a sequence of actions and rewards generated +by an algorithm interacting with an environment. -/ structure IsAlgEnvSeq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (alg : Algorithm α R) (env : Environment α R) @@ -117,6 +121,7 @@ lemma IsAlgEnvSeq.hasCondDistrib_step HasCondDistrib (step A R' (n + 1)) (hist A R' n) (stepKernel alg env n) P := HasCondDistrib.prod (h.hasCondDistrib_action n) (h.hasCondDistrib_reward n) +/-- Filtration generated by the history up to time `n`. -/ def IsAlgEnvSeq.filtration (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) : Filtration ℕ mΩ where seq i := MeasurableSpace.comap (hist A R' i) inferInstance diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index a3810696..527db1f6 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -102,8 +102,8 @@ lemma pullCount_eq_pullCount' {n : ℕ} {ω : Ω} (hn : n ≠ 0) : lemma pullCount'_mono {n m : ℕ} (hnm : n ≤ m) : pullCount' n (fun i ↦ (A i ω, R' i ω)) a ≤ pullCount' m (fun i ↦ (A i ω, R' i ω)) a := by - simp_rw [pullCount'] - sorry + rw [← pullCount_add_one_eq_pullCount', ← pullCount_add_one_eq_pullCount'] + exact pullCount_mono a (by lia) _ lemma pullCount_le (a : α) (t : ℕ) (ω : Ω) : pullCount A a t ω ≤ t := (card_filter_le _ _).trans_eq (by simp) @@ -133,7 +133,7 @@ lemma exists_pullCount_eq_of_le (hnm : t ≤ pullCount A a (n + 1) ω) (ht : t refine lt_of_lt_of_le ?_ hnm exact pullCount_lt_of_forall_ne h_contra ht -lemma pullCount_le_add [Nonempty α] (a : α) (n C : ℕ) (ω : Ω) : +lemma pullCount_le_add (a : α) (n C : ℕ) (ω : Ω) : pullCount A a n ω ≤ C + 1 + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s := by rw [pullCount_eq_sum] @@ -181,7 +181,7 @@ lemma measurable_pullCount [MeasurableSingletonClass α] (hA : ∀ n, Measurable fun_prop @[fun_prop] -lemma measurable_uncurry_pullCount [MeasurableSingletonClass α] [MeasurableEq α] +lemma measurable_uncurry_pullCount [MeasurableEq α] (hA : ∀ n, Measurable (A n)) (t : ℕ) : Measurable (fun p : Ω × α ↦ pullCount A p.2 t p.1) := by simp_rw [pullCount_eq_sum] @@ -199,7 +199,7 @@ lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop -lemma measurable_uncurry_pullCount' [MeasurableSingletonClass α] [MeasurableEq α] (n : ℕ) : +lemma measurable_uncurry_pullCount' [MeasurableEq α] (n : ℕ) : Measurable (fun p : (Iic n → α × R) × α ↦ pullCount' n p.1 p.2) := by simp_rw [pullCount'_eq_sum] have h_meas s : Measurable (fun h : (Iic n → α × R) × α ↦ if (h.1 s).1 = h.2 then 1 else 0) := by @@ -555,7 +555,7 @@ lemma measurable_stepsUntil' [MeasurableSingletonClass α] Measurable (fun ω : Ω × (ℕ → α → R) ↦ stepsUntil A a m ω.1) := (measurable_stepsUntil hA a m).comp measurable_fst -lemma measurable_comap_indicator_stepsUntil_eq [MeasurableSingletonClass α] [Nonempty R] +lemma measurable_comap_indicator_stepsUntil_eq [MeasurableSingletonClass α] (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : Measurable[MeasurableSpace.comap (fun ω : Ω ↦ (IsAlgEnvSeq.hist A R' (n-1) ω, A n ω)) inferInstance] @@ -595,7 +595,7 @@ lemma measurable_comap_indicator_stepsUntil_eq [MeasurableSingletonClass α] [No have h_meas := adapted_pullCount_add_one' hA hR' a (n - 1) rwa [Nat.sub_add_cancel (by lia)] at h_meas -lemma measurable_indicator_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] +lemma measurable_indicator_stepsUntil_eq [MeasurableSingletonClass α] (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : Measurable ({ω : Ω | stepsUntil A a m ω = ↑n}.indicator fun _ ↦ 1) := by refine (measurable_comap_indicator_stepsUntil_eq hA hR' a m n).mono ?_ le_rfl @@ -622,7 +622,7 @@ lemma measurable_comap_indicator_stepsUntil_eq_zero [MeasurableSingletonClass α rw [measurable_indicator_const_iff] exact measurableSet_stepsUntil_eq_zero a m -lemma measurableSet_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] +lemma measurableSet_stepsUntil_eq [MeasurableSingletonClass α] (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : MeasurableSet[MeasurableSpace.comap (fun ω : Ω ↦ (IsAlgEnvSeq.hist A R' (n-1) ω, A n ω)) inferInstance] @@ -634,7 +634,7 @@ lemma measurableSet_stepsUntil_eq [Nonempty R] [MeasurableSingletonClass α] exact measurable_comap_indicator_stepsUntil_eq hA hR' a m n /-- `stepsUntil a m` is a stopping time with respect to the filtration `filtrationAction`. -/ -theorem isStoppingTime_stepsUntil_filtrationAction [Nonempty R] [MeasurableSingletonClass α] +theorem isStoppingTime_stepsUntil_filtrationAction [MeasurableSingletonClass α] (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m : ℕ) : IsStoppingTime (IsAlgEnvSeq.filtrationAction hA hR') (stepsUntil A a m) := by refine isStoppingTime_of_measurableSet_eq fun n ↦ ?_ diff --git a/LeanBandits/SequentialLearning/StationaryEnv.lean b/LeanBandits/SequentialLearning/StationaryEnv.lean index b5eaffa3..563bb49a 100644 --- a/LeanBandits/SequentialLearning/StationaryEnv.lean +++ b/LeanBandits/SequentialLearning/StationaryEnv.lean @@ -55,7 +55,7 @@ lemma IsAlgEnvSeq.condDistrib_reward_stationaryEnv /-- The reward at time `n + 1` is conditionally independent of the history up to time `n` given the action at time `n + 1`. -/ -lemma IsAlgEnvSeq.condIndepFun_reward_hist_action [StandardBorelSpace Ω] [Nonempty Ω] +lemma IsAlgEnvSeq.condIndepFun_reward_hist_action [StandardBorelSpace Ω] (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) : R' (n + 1) ⟂ᵢ[A (n + 1), h.measurable_A _ ; P] hist A R' n := by have hA := h.measurable_A @@ -63,7 +63,7 @@ lemma IsAlgEnvSeq.condIndepFun_reward_hist_action [StandardBorelSpace Ω] [Nonem exact condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (by fun_prop) (by fun_prop) (by fun_prop) (h.hasCondDistrib_reward n).condDistrib_eq -lemma IsAlgEnvSeq.condIndepFun_reward_hist_action_action [StandardBorelSpace Ω] [Nonempty Ω] +lemma IsAlgEnvSeq.condIndepFun_reward_hist_action_action [StandardBorelSpace Ω] (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) : R' (n + 1) ⟂ᵢ[A (n + 1), h.measurable_A (n + 1); P] (fun ω ↦ (hist A R' n ω, A (n + 1) ω)) := by @@ -73,7 +73,7 @@ lemma IsAlgEnvSeq.condIndepFun_reward_hist_action_action [StandardBorelSpace Ω] have hR' := h.measurable_R exact h_indep.prod_right (by fun_prop) (by fun_prop) (by fun_prop) -lemma IsAlgEnvSeq.condIndepFun_reward_hist_action_action' [StandardBorelSpace Ω] [Nonempty Ω] +lemma IsAlgEnvSeq.condIndepFun_reward_hist_action_action' [StandardBorelSpace Ω] (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) (hn : n ≠ 0) : R' n ⟂ᵢ[A n, h.measurable_A n; P] (fun ω ↦ (hist A R' (n - 1) ω, A n ω)) := by have := h.condIndepFun_reward_hist_action_action (n - 1) From e41825dc909134ef96654d6395d8acd19843fe06 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 14 Jan 2026 14:14:50 +0100 Subject: [PATCH 17/30] lake exe mk_all --- LeanBandits.lean | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/LeanBandits.lean b/LeanBandits.lean index 4de3f6dd..ef2e8c26 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -1,15 +1,19 @@ import LeanBandits.Bandit.Bandit import LeanBandits.Bandit.Regret import LeanBandits.Bandit.RewardByCountMeasure +import LeanBandits.Bandit.SumRewards import LeanBandits.BanditAlgorithms.ETC import LeanBandits.BanditAlgorithms.UCB import LeanBandits.ForMathlib.CondDistrib +import LeanBandits.ForMathlib.CondIndepFun import LeanBandits.ForMathlib.HasCondDistrib import LeanBandits.ForMathlib.IndepFun import LeanBandits.ForMathlib.IndepInfinitePi +import LeanBandits.ForMathlib.KernelRepresentation import LeanBandits.ForMathlib.KernelSub import LeanBandits.ForMathlib.Measurable import LeanBandits.ForMathlib.MeasurableArgMax +import LeanBandits.ForMathlib.StandardBorel import LeanBandits.ForMathlib.SubGaussian import LeanBandits.ForMathlib.Traj import LeanBandits.SequentialLearning.Algorithm From a6af9887bc7c4e0065be2190c6e6ec6bea53b001 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 14 Jan 2026 14:26:15 +0100 Subject: [PATCH 18/30] temporary blueprint fix --- blueprint/lean_decls | 34 ++++++++++++++-------------- blueprint/src/chapters/algorithm.tex | 24 ++++++++++---------- blueprint/src/chapters/bandit.tex | 2 +- 3 files changed, 30 insertions(+), 30 deletions(-) diff --git a/blueprint/lean_decls b/blueprint/lean_decls index 0c5255e1..671d8fec 100644 --- a/blueprint/lean_decls +++ b/blueprint/lean_decls @@ -4,23 +4,23 @@ Learning.detAlgorithm Learning.stationaryEnv ProbabilityTheory.Kernel.traj ProbabilityTheory.Kernel.trajMeasure -Learning.step -Learning.hist -Learning.filtration -Learning.adapted_step -Learning.adapted_hist +Learning.IT.step +Learning.IT.hist +Learning.IT.filtration +Learning.IT.adapted_step +Learning.IT.adapted_hist ProbabilityTheory.Kernel.condDistrib_trajMeasure -Learning.hasLaw_step_zero -Learning.action -Learning.reward -Learning.adapted_action -Learning.adapted_reward -Learning.condDistrib_action -Learning.condDistrib_reward -Learning.hasLaw_action_zero -Learning.condDistrib_reward_zero -Learning.condDistrib_reward_stationaryEnv -Learning.condIndepFun_reward_hist_action +Learning.IsAlgEnvSeq.hasLaw_step_zero +Learning.IT.action +Learning.IT.reward +Learning.IT.adapted_action +Learning.IT.adapted_reward +Learning.IsAlgEnvSeq.hasCondDistrib_action +Learning.IsAlgEnvSeq.hasCondDistrib_reward +Learning.IsAlgEnvSeq.hasLaw_action_zero +Learning.IsAlgEnvSeq.hasCondDistrib_reward_zero +Learning.IsAlgEnvSeq.condDistrib_reward_stationaryEnv +Learning.IsAlgEnvSeq.condIndepFun_reward_hist_action Learning.pullCount Learning.pullCount_zero Learning.pullCount_mono @@ -45,7 +45,7 @@ Learning.sum_rewardByCount_eq_sumRewards Bandits.Bandit.trajMeasure Bandits.Bandit.measure Learning.measurable_comap_indicator_stepsUntil_eq -Bandits.condIndepFun_reward_stepsUntil_arm +Bandits.condIndepFun_reward_stepsUntil_action Bandits.reward_cond_stepsUntil ProbabilityTheory.condDistrib_ae_eq_cond Bandits.condDistrib_rewardByCount_stepsUntil diff --git a/blueprint/src/chapters/algorithm.tex b/blueprint/src/chapters/algorithm.tex index 7b3de6d7..dbbf6e74 100644 --- a/blueprint/src/chapters/algorithm.tex +++ b/blueprint/src/chapters/algorithm.tex @@ -102,7 +102,7 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{definition}[Step and history]\label{def:history} \leanok - \lean{Learning.step, Learning.hist} + \lean{Learning.IT.step, Learning.IT.hist} For $t \in \mathbb{N}$, we denote by $X_t \in \Omega_t$ the random variable describing the time step $t$, and by $H_t \in \prod_{s=0}^t \Omega_s$ the history up to time $t$. Formally, these are measurable functions on $\Omega_{\mathcal{T}}$, defined by $X_t(\omega) = \omega_t$ and $H_t(\omega) = (\omega_1, \ldots, \omega_t)$. \end{definition} @@ -113,7 +113,7 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{definition}[Filtration]\label{def:filtration} \uses{def:history} \leanok - \lean{Learning.filtration} + \lean{Learning.IT.filtration} For $t \in \mathbb{N}$, we denote by $\mathcal{F}_t$ the sigma-algebra generated by the history up to time $t$: $\mathcal{F}_t = \sigma(H_t)$. The family $(\mathcal{F}_t)_{t \in \mathbb{N}}$ is a filtration on $\Omega_{\mathcal{T}}$. \end{definition} @@ -124,7 +124,7 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:adapted_history} \uses{def:history, def:filtration} \leanok - \lean{Learning.adapted_step, Learning.adapted_hist} + \lean{Learning.IT.adapted_step, Learning.IT.adapted_hist} The random variables $X_t$ and $H_t$ are $\mathcal{F}_t$-measurable. Said differently, the processes $(X_t)_{t \in \mathbb{N}}$ and $(H_t)_{t \in \mathbb{N}}$ are adapted to the filtration $(\mathcal{F}_t)_{t \in \mathbb{N}}$. \end{lemma} @@ -154,7 +154,7 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:law_X_zero} \uses{def:history, def:trajMeasure} \leanok - \lean{Learning.hasLaw_step_zero} + \lean{Learning.IsAlgEnvSeq.hasLaw_step_zero} The law of $X_0$ under $P_{\mathcal{T}}$ is $\mu$. \end{lemma} @@ -174,7 +174,7 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{definition}\label{def:actionReward} \uses{def:history} \leanok - \lean{Learning.action, Learning.reward} + \lean{Learning.IT.action, Learning.IT.reward} We write $A_t$ and $R_t$ for the projections of $X_t$ on $\mathcal{A}_t$ and $\mathcal{R}_t$ respectively. $A_t$ is the action taken at time $t$ and $R_t$ is the reward received at time $t$. Formally, $A_t(\omega) = \omega_{t,1}$ and $R_t(\omega) = \omega_{t,2}$ for $\omega = \prod_{t=0}^{+\infty}(\omega_{t,1}, \omega_{t,2}) \in \prod_{t=0}^{+\infty} \mathcal{A}_t \times \mathcal{R}_t$. @@ -184,7 +184,7 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:adapted_action_reward} \uses{def:actionReward, def:filtration} \leanok - \lean{Learning.adapted_action, Learning.adapted_reward} + \lean{Learning.IT.adapted_action, Learning.IT.adapted_reward} The random variables $A_t$ and $R_t$ are $\mathcal{F}_t$-measurable. Said differently, the processes $(A_t)_{t \in \mathbb{N}}$ and $(R_t)_{t \in \mathbb{N}}$ are adapted to the filtration $(\mathcal{F}_t)_{t \in \mathbb{N}}$. \end{lemma} @@ -200,7 +200,7 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:condDistrib_A_add_one} \uses{def:actionReward, def:trajMeasure, def:algorithm, def:environment} \leanok - \lean{Learning.condDistrib_action} + \lean{Learning.IsAlgEnvSeq.hasCondDistrib_action} For any $t \in \mathbb{N}$, the conditional distribution $P_{\mathcal{T}}\left[A_{t+1} \mid H_t\right]$ is $((H_t)_* P_{\mathcal{T}})$-almost surely equal to $\pi_t$. \end{lemma} @@ -214,7 +214,7 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:condDistrib_R_add_one} \uses{def:actionReward, def:trajMeasure, def:algorithm, def:environment} \leanok - \lean{Learning.condDistrib_reward} + \lean{Learning.IsAlgEnvSeq.hasCondDistrib_reward} For any $t \in \mathbb{N}$, the conditional distribution $P_{\mathcal{T}}\left[R_{t+1} \mid H_t, A_{t+1}\right]$ is $((H_t, A_{t+1})_* P_{\mathcal{T}})$-almost surely equal to $\nu_t$. \end{lemma} @@ -233,7 +233,7 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:law_A_zero} \uses{def:actionReward, def:trajMeasure, def:algorithm, def:environment} \leanok - \lean{Learning.hasLaw_action_zero} + \lean{Learning.IsAlgEnvSeq.hasLaw_action_zero} The law of $A_0$ under $P_{\mathcal{T}}$ is $\alpha_0$. \end{lemma} @@ -246,7 +246,7 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:condDistrib_R_zero} \uses{def:actionReward, def:trajMeasure, def:algorithm, def:environment} \leanok - \lean{Learning.condDistrib_reward_zero} + \lean{Learning.IsAlgEnvSeq.hasCondDistrib_reward_zero} The conditional distribution $P_{\mathcal{T}}\left[R_0 \mid A_0\right]$ is $(A_{0*} P_{\mathcal{T}})$-almost surely equal to $\nu'_0$. \end{lemma} @@ -269,7 +269,7 @@ \section{Stationary environment} \begin{lemma}\label{lem:condDistrib_reward_stationaryEnv} \uses{def:actionReward, def:trajMeasure, def:algorithm, def:stationaryEnv} \leanok - \lean{Learning.condDistrib_reward_stationaryEnv} + \lean{Learning.IsAlgEnvSeq.condDistrib_reward_stationaryEnv} In a stationary environment, for any $t \in \mathbb{N}$, the conditional distribution $P_{\mathcal{T}}\left[R_t \mid A_t\right]$ is $(A_{t*} P_{\mathcal{T}})$-almost surely equal to $\nu$. \end{lemma} @@ -282,7 +282,7 @@ \section{Stationary environment} \begin{lemma}\label{lem:condIndepFun_reward_hist_action} \uses{def:actionReward, def:trajMeasure, def:algorithm, def:stationaryEnv} \leanok - \lean{Learning.condIndepFun_reward_hist_action} + \lean{Learning.IsAlgEnvSeq.condIndepFun_reward_hist_action} In a stationary environment, for any $t \in \mathbb{N}$, the reward $R_{t+1}$ is conditionally independent of the history $H_t$ given the action $A_{t+1}$ (more succinctly, $R_{t+1} \ind H_t \mid A_{t+1}$). \end{lemma} diff --git a/blueprint/src/chapters/bandit.tex b/blueprint/src/chapters/bandit.tex index 72d9e31f..e504f4ca 100644 --- a/blueprint/src/chapters/bandit.tex +++ b/blueprint/src/chapters/bandit.tex @@ -63,7 +63,7 @@ \section{Alternative models: rewards indexed by time or pull count}\label{sec:al \begin{lemma}\label{lem:condIndepFun_reward_stepsUntil_arm} \uses{def:stepsUntil, def:actionReward, def:Bandit.measure} \leanok - \lean{Bandits.condIndepFun_reward_stepsUntil_arm} + \lean{Bandits.condIndepFun_reward_stepsUntil_action} For $t > 0$, $R_t \ind \mathbb{I}\{T_{n, a} = t\} \mid A_t$. \end{lemma} From 6647491627020229b13c0de94d063264c3a5ad5a Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 14 Jan 2026 14:57:00 +0100 Subject: [PATCH 19/30] min_imports --- LeanBandits/Bandit/Bandit.lean | 4 +- LeanBandits/Bandit/Regret.lean | 1 - LeanBandits/Bandit/RewardByCountMeasure.lean | 4 +- LeanBandits/Bandit/SumRewards.lean | 1 + LeanBandits/BanditAlgorithms/AuxSums.lean | 47 +++++++++++++++++++ LeanBandits/BanditAlgorithms/ETC.lean | 38 +-------------- LeanBandits/BanditAlgorithms/UCB.lean | 4 +- LeanBandits/ForMathlib/Traj.lean | 8 +++- LeanBandits/SequentialLearning/Algorithm.lean | 2 - .../SequentialLearning/Deterministic.lean | 10 ++-- .../SequentialLearning/StationaryEnv.lean | 10 ++-- 11 files changed, 70 insertions(+), 59 deletions(-) create mode 100644 LeanBandits/BanditAlgorithms/AuxSums.lean diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 0fe71146..7d19b73a 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -9,10 +9,8 @@ import LeanBandits.ForMathlib.IndepInfinitePi import LeanBandits.ForMathlib.KernelRepresentation import LeanBandits.ForMathlib.StandardBorel import LeanBandits.SequentialLearning.Deterministic -import LeanBandits.SequentialLearning.StationaryEnv import LeanBandits.SequentialLearning.FiniteActions -import Mathlib.Probability.IdentDistrib -import Mathlib.MeasureTheory.Constructions.UnitInterval +import LeanBandits.SequentialLearning.StationaryEnv /-! # Bandit diff --git a/LeanBandits/Bandit/Regret.lean b/LeanBandits/Bandit/Regret.lean index 3f35f89f..89adb008 100644 --- a/LeanBandits/Bandit/Regret.lean +++ b/LeanBandits/Bandit/Regret.lean @@ -3,7 +3,6 @@ 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, Paulo Rauber -/ -import LeanBandits.Bandit.Bandit import LeanBandits.SequentialLearning.FiniteActions /-! diff --git a/LeanBandits/Bandit/RewardByCountMeasure.lean b/LeanBandits/Bandit/RewardByCountMeasure.lean index 07b20131..60a2b369 100644 --- a/LeanBandits/Bandit/RewardByCountMeasure.lean +++ b/LeanBandits/Bandit/RewardByCountMeasure.lean @@ -3,9 +3,7 @@ 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 LeanBandits.Bandit.Regret -import LeanBandits.ForMathlib.CondIndepFun -import LeanBandits.ForMathlib.IndepFun +import LeanBandits.Bandit.Bandit import Mathlib.Probability.IdentDistribIndep /-! # Laws of `stepsUntil` and `rewardByCount` diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index ce143813..7ac161c4 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -3,6 +3,7 @@ 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 LeanBandits.Bandit.Bandit import LeanBandits.Bandit.Regret import LeanBandits.ForMathlib.SubGaussian diff --git a/LeanBandits/BanditAlgorithms/AuxSums.lean b/LeanBandits/BanditAlgorithms/AuxSums.lean new file mode 100644 index 00000000..27bdf41f --- /dev/null +++ b/LeanBandits/BanditAlgorithms/AuxSums.lean @@ -0,0 +1,47 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +import Mathlib.Algebra.BigOperators.Intervals +import Mathlib.Algebra.BigOperators.Ring.Finset +import Mathlib.Tactic.Ring.RingNF + +open Finset + +lemma sum_mod_range {K : ℕ} (hK : 0 < K) (a : Fin K) : + (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) = 1 := by + have h_iff (s : ℕ) (hs : s < K) : ⟨s % K, Nat.mod_lt _ hK⟩ = a ↔ s = a := by + simp only [Nat.mod_eq_of_lt hs, Fin.ext_iff] + calc (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) + _ = ∑ s ∈ range K, if s = a then 1 else 0 := sum_congr rfl fun s hs ↦ by grind + _ = _ := by + rw [sum_ite_eq'] + simp + +lemma sum_mod_range_mul {K : ℕ} (hK : 0 < K) (m : ℕ) (a : Fin K) : + (∑ s ∈ range (K * m), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) = m := by + induction m with + | zero => simp + | succ n hn => + calc (∑ s ∈ range (K * (n + 1)), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) + _ = (∑ s ∈ range (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by ring_nf + _ = (∑ s ∈ range (K * n), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) + + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by + rw [sum_range_add_sum_Ico] + grind + _ = n + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by + rw [hn] + _ = n + (∑ s ∈ range K, if ⟨(s + K * n) % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by + congr 1 + let e : ℕ ↪ ℕ := ⟨fun i : ℕ ↦ i + K * n, fun i j hij ↦ by grind⟩ + have : Finset.map e (range K) = Ico (K * n) (K * n + K) := by + ext x + simp only [mem_map, mem_range, Function.Embedding.coeFn_mk, mem_Ico, e] + refine ⟨fun h ↦ by grind, fun h ↦ ?_⟩ + use x - K * n + grind + rw [← this, Finset.sum_map] + congr + _ = n + (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by simp + _ = n + 1 := by rw [sum_mod_range hK] diff --git a/LeanBandits/BanditAlgorithms/ETC.lean b/LeanBandits/BanditAlgorithms/ETC.lean index ee84d0d2..c173d916 100644 --- a/LeanBandits/BanditAlgorithms/ETC.lean +++ b/LeanBandits/BanditAlgorithms/ETC.lean @@ -4,6 +4,7 @@ Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne -/ import LeanBandits.Bandit.SumRewards +import LeanBandits.BanditAlgorithms.AuxSums import LeanBandits.ForMathlib.MeasurableArgMax /-! # The Explore-Then-Commit Algorithm @@ -21,43 +22,6 @@ lemma ae_eq_set_iff {α : Type*} {mα : MeasurableSpace α} {μ : Measure α} {s simp only [eq_iff_iff] congr! -lemma sum_mod_range {K : ℕ} (hK : 0 < K) (a : Fin K) : - (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) = 1 := by - have h_iff (s : ℕ) (hs : s < K) : ⟨s % K, Nat.mod_lt _ hK⟩ = a ↔ s = a := by - simp only [Nat.mod_eq_of_lt hs, Fin.ext_iff] - calc (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) - _ = ∑ s ∈ range K, if s = a then 1 else 0 := sum_congr rfl fun s hs ↦ by grind - _ = _ := by - rw [sum_ite_eq'] - simp - -lemma sum_mod_range_mul {K : ℕ} (hK : 0 < K) (m : ℕ) (a : Fin K) : - (∑ s ∈ range (K * m), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) = m := by - induction m with - | zero => simp - | succ n hn => - calc (∑ s ∈ range (K * (n + 1)), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) - _ = (∑ s ∈ range (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by ring_nf - _ = (∑ s ∈ range (K * n), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) - + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by - rw [sum_range_add_sum_Ico] - grind - _ = n + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by - rw [hn] - _ = n + (∑ s ∈ range K, if ⟨(s + K * n) % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by - congr 1 - let e : ℕ ↪ ℕ := ⟨fun i : ℕ ↦ i + K * n, fun i j hij ↦ by grind⟩ - have : Finset.map e (range K) = Ico (K * n) (K * n + K) := by - ext x - simp only [mem_map, mem_range, Function.Embedding.coeFn_mk, mem_Ico, e] - refine ⟨fun h ↦ by grind, fun h ↦ ?_⟩ - use x - K * n - grind - rw [← this, Finset.sum_map] - congr - _ = n + (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by simp - _ = n + 1 := by rw [sum_mod_range hK] - end Aux namespace Bandits diff --git a/LeanBandits/BanditAlgorithms/UCB.lean b/LeanBandits/BanditAlgorithms/UCB.lean index 66ad5f4c..bd300d8a 100644 --- a/LeanBandits/BanditAlgorithms/UCB.lean +++ b/LeanBandits/BanditAlgorithms/UCB.lean @@ -3,7 +3,9 @@ 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 LeanBandits.BanditAlgorithms.ETC +import LeanBandits.Bandit.SumRewards +import LeanBandits.BanditAlgorithms.AuxSums +import LeanBandits.ForMathlib.MeasurableArgMax /-! # UCB algorithm diff --git a/LeanBandits/ForMathlib/Traj.lean b/LeanBandits/ForMathlib/Traj.lean index cc0055ec..d9fdf7e2 100644 --- a/LeanBandits/ForMathlib/Traj.lean +++ b/LeanBandits/ForMathlib/Traj.lean @@ -1,7 +1,11 @@ +/- +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, Paulo Rauber +-/ +import LeanBandits.ForMathlib.HasCondDistrib import Mathlib.Probability.Kernel.IonescuTulcea.Traj -import Mathlib.Probability.Kernel.CondDistrib import Mathlib.Probability.Process.FiniteDimensionalLaws -import LeanBandits.ForMathlib.HasCondDistrib open Filter Finset Function MeasurableEquiv MeasurableSpace MeasureTheory Preorder ProbabilityTheory diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index ac1de8c1..df3c528b 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -3,10 +3,8 @@ 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, Paulo Rauber -/ -import LeanBandits.ForMathlib.HasCondDistrib import LeanBandits.ForMathlib.Measurable import LeanBandits.ForMathlib.Traj -import Mathlib.Probability.HasLaw /-! # Algorithms diff --git a/LeanBandits/SequentialLearning/Deterministic.lean b/LeanBandits/SequentialLearning/Deterministic.lean index 1bab62a4..8be520da 100644 --- a/LeanBandits/SequentialLearning/Deterministic.lean +++ b/LeanBandits/SequentialLearning/Deterministic.lean @@ -29,20 +29,20 @@ def detAlgorithm (nextaction : (n : ℕ) → (Iic n → α × R) → α) variable {nextaction : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextaction n)} {action0 : α} {env : Environment α R} -section IsAlgEnvSeq +namespace IsAlgEnvSeq variable {Ω : Type*} {mΩ : MeasurableSpace Ω} [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] {P : Measure Ω} [IsProbabilityMeasure P] {A : ℕ → Ω → α} {R' : ℕ → Ω → R} -lemma IsAlgEnvSeq.HasLaw_action_zero_detAlgorithm +lemma HasLaw_action_zero_detAlgorithm (h : IsAlgEnvSeq A R' (detAlgorithm nextaction h_next action0) env P) : HasLaw (A 0) (Measure.dirac action0) P where aemeasurable := have hA := h.measurable_A; by fun_prop map_eq := (hasLaw_action_zero h).map_eq -lemma IsAlgEnvSeq.action_zero_detAlgorithm +lemma action_zero_detAlgorithm (h : IsAlgEnvSeq A R' (detAlgorithm nextaction h_next action0) env P) : A 0 =ᵐ[P] fun _ ↦ action0 := by have h_eq : ∀ᵐ x ∂(P.map (A 0)), x = action0 := by @@ -51,7 +51,7 @@ lemma IsAlgEnvSeq.action_zero_detAlgorithm have hA := h.measurable_A exact ae_of_ae_map (by fun_prop) h_eq -lemma IsAlgEnvSeq.action_detAlgorithm_ae_eq +lemma action_detAlgorithm_ae_eq (h : IsAlgEnvSeq A R' (detAlgorithm nextaction h_next action0) env P) (n : ℕ) : A (n + 1) =ᵐ[P] fun ω ↦ nextaction n (hist A R' n ω) := by have hA := h.measurable_A @@ -59,7 +59,7 @@ lemma IsAlgEnvSeq.action_detAlgorithm_ae_eq exact ae_eq_of_condDistrib_eq_deterministic (by fun_prop) (by fun_prop) (by fun_prop) (h.hasCondDistrib_action n).condDistrib_eq -lemma IsAlgEnvSeq.action_detAlgorithm_ae_all_eq +lemma action_detAlgorithm_ae_all_eq (h : IsAlgEnvSeq A R' (detAlgorithm nextaction h_next action0) env P) : ∀ᵐ ω ∂P, A 0 ω = action0 ∧ ∀ n, A (n + 1) ω = nextaction n (hist A R' n ω) := by rw [eventually_and, ae_all_iff] diff --git a/LeanBandits/SequentialLearning/StationaryEnv.lean b/LeanBandits/SequentialLearning/StationaryEnv.lean index 563bb49a..35d44157 100644 --- a/LeanBandits/SequentialLearning/StationaryEnv.lean +++ b/LeanBandits/SequentialLearning/StationaryEnv.lean @@ -29,10 +29,10 @@ variable {Ω : Type*} {mΩ : MeasurableSpace Ω} {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] {P : Measure Ω} [IsProbabilityMeasure P] {A : ℕ → Ω → α} {R' : ℕ → Ω → R} -section IsAlgEnvSeq +namespace IsAlgEnvSeq /-- The conditional distribution of the reward at time `n` given the action at time `n` is `ν`. -/ -lemma IsAlgEnvSeq.condDistrib_reward_stationaryEnv +lemma condDistrib_reward_stationaryEnv (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) : condDistrib (R' n) (A n) P =ᵐ[P.map (A n)] ν := by have hA := h.measurable_A @@ -55,7 +55,7 @@ lemma IsAlgEnvSeq.condDistrib_reward_stationaryEnv /-- The reward at time `n + 1` is conditionally independent of the history up to time `n` given the action at time `n + 1`. -/ -lemma IsAlgEnvSeq.condIndepFun_reward_hist_action [StandardBorelSpace Ω] +lemma condIndepFun_reward_hist_action [StandardBorelSpace Ω] (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) : R' (n + 1) ⟂ᵢ[A (n + 1), h.measurable_A _ ; P] hist A R' n := by have hA := h.measurable_A @@ -63,7 +63,7 @@ lemma IsAlgEnvSeq.condIndepFun_reward_hist_action [StandardBorelSpace Ω] exact condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (by fun_prop) (by fun_prop) (by fun_prop) (h.hasCondDistrib_reward n).condDistrib_eq -lemma IsAlgEnvSeq.condIndepFun_reward_hist_action_action [StandardBorelSpace Ω] +lemma condIndepFun_reward_hist_action_action [StandardBorelSpace Ω] (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) : R' (n + 1) ⟂ᵢ[A (n + 1), h.measurable_A (n + 1); P] (fun ω ↦ (hist A R' n ω, A (n + 1) ω)) := by @@ -73,7 +73,7 @@ lemma IsAlgEnvSeq.condIndepFun_reward_hist_action_action [StandardBorelSpace Ω] have hR' := h.measurable_R exact h_indep.prod_right (by fun_prop) (by fun_prop) (by fun_prop) -lemma IsAlgEnvSeq.condIndepFun_reward_hist_action_action' [StandardBorelSpace Ω] +lemma condIndepFun_reward_hist_action_action' [StandardBorelSpace Ω] (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) (hn : n ≠ 0) : R' n ⟂ᵢ[A n, h.measurable_A n; P] (fun ω ↦ (hist A R' (n - 1) ω, A n ω)) := by have := h.condIndepFun_reward_hist_action_action (n - 1) From a565bed85076cc521aa32f351f84d6be19df1d20 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 14 Jan 2026 14:57:37 +0100 Subject: [PATCH 20/30] mk_all --- LeanBandits.lean | 1 + 1 file changed, 1 insertion(+) diff --git a/LeanBandits.lean b/LeanBandits.lean index ef2e8c26..c1a8f31c 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -2,6 +2,7 @@ import LeanBandits.Bandit.Bandit import LeanBandits.Bandit.Regret import LeanBandits.Bandit.RewardByCountMeasure import LeanBandits.Bandit.SumRewards +import LeanBandits.BanditAlgorithms.AuxSums import LeanBandits.BanditAlgorithms.ETC import LeanBandits.BanditAlgorithms.UCB import LeanBandits.ForMathlib.CondDistrib From f9591ced76c39567cd27ae61a02dc7067e2f8e17 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 14 Jan 2026 16:13:21 +0100 Subject: [PATCH 21/30] add Claude proof --- LeanBandits/Bandit/Bandit.lean | 21 ----- LeanBandits/ForMathlib/CondDistrib.lean | 79 +++++++++++++++++-- .../SequentialLearning/Deterministic.lean | 30 +++---- .../SequentialLearning/StationaryEnv.lean | 38 +++------ 4 files changed, 97 insertions(+), 71 deletions(-) diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 7d19b73a..0e19def2 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -161,27 +161,6 @@ lemma indepFun_eval_snd_measure (alg : Algorithm α R) (ν : Kernel α R) [IsMar end StreamMeasure -section DetAlgorithm - -variable {nextaction : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextaction n)} - {action0 : α} {ν : Kernel α R} [IsMarkovKernel ν] - -local notation "𝔓t" => Bandit.trajMeasure (detAlgorithm nextaction h_next action0) ν - -lemma HasLaw_action_zero_detAlgorithm : HasLaw (IT.action 0) (Measure.dirac action0) 𝔓t where - map_eq := (IT.hasLaw_action_zero _ _).map_eq - -lemma action_zero_detAlgorithm [MeasurableSingletonClass α] : - IT.action 0 =ᵐ[𝔓t] fun _ ↦ action0 := - IT.action_zero_detAlgorithm - -lemma action_detAlgorithm_ae_eq [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] (n : ℕ) : - IT.action (n + 1) =ᵐ[𝔓t] fun h ↦ nextaction n (fun i ↦ h i) := - IT.action_detAlgorithm_ae_eq n - -end DetAlgorithm - namespace ArrayModel open unitInterval diff --git a/LeanBandits/ForMathlib/CondDistrib.lean b/LeanBandits/ForMathlib/CondDistrib.lean index b1cde46a..bf137f29 100644 --- a/LeanBandits/ForMathlib/CondDistrib.lean +++ b/LeanBandits/ForMathlib/CondDistrib.lean @@ -72,14 +72,6 @@ lemma CondIndepFun.of_prod_left {ε : Type*} {mε : MeasurableSpace ε} X ⟂ᵢ[Z, hZ; μ] Y := Kernel.IndepFun.of_prod_left h -lemma CondIndepFun.prod_right [StandardBorelSpace α] [StandardBorelSpace β] [Nonempty β] - [StandardBorelSpace γ] [Nonempty γ] [StandardBorelSpace δ] [Nonempty δ] [IsFiniteMeasure μ] - {X : α → β} {Y : α → γ} {Z : α → δ} - (hX : Measurable X) (hY : Measurable Y) (hZ : Measurable Z) - (h : X ⟂ᵢ[Z, hZ; μ] Y) : - X ⟂ᵢ[Z, hZ; μ] (fun ω ↦ (Y ω, Z ω)) := by - sorry - end IndepFun section CondDistrib @@ -95,6 +87,77 @@ lemma condDistrib_prod_left [StandardBorelSpace β] [Nonempty β] AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] rfl +lemma condDistrib_prod_self_left [StandardBorelSpace β] [Nonempty β] [StandardBorelSpace γ] + [Nonempty γ] + (hX : AEMeasurable X μ) (hT : AEMeasurable T μ) : + condDistrib (fun ω ↦ (X ω, T ω)) T μ =ᵐ[μ.map T] condDistrib X T μ ×ₖ Kernel.id := by + have h_prod := condDistrib_prod_left hX hT hT (μ := μ) + have h_fst := condDistrib_comp_self (μ := μ) (fun ω ↦ (T ω, X ω)) (f := Prod.fst) (by fun_prop) + rw [(compProd_map_condDistrib hX).symm] at h_fst + have h_fst' := (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_fst + filter_upwards [h_prod, h_fst'] with z hz1 hz2 + rw [hz1] + ext t ht + rw [Kernel.compProd_apply ht] + have hz2' : ∀ᵐ y ∂(condDistrib X T μ z), + condDistrib T (fun ω ↦ (T ω, X ω)) μ (z, y) = Measure.dirac z := by + filter_upwards [hz2] with y hy; convert hy using 2 + calc ∫⁻ y, condDistrib T (fun ω ↦ (T ω, X ω)) μ (z, y) (Prod.mk y ⁻¹' t) ∂condDistrib X T μ z + _ = ∫⁻ y, (Measure.dirac z) (Prod.mk y ⁻¹' t) ∂condDistrib X T μ z := + lintegral_congr_ae (hz2'.mono fun y hy ↦ by simp only [hy]) + _ = ∫⁻ y, (Prod.mk y ⁻¹' t).indicator 1 z ∂condDistrib X T μ z := + lintegral_congr fun y ↦ Measure.dirac_apply' _ (ht.preimage (by fun_prop)) + _ = (condDistrib X T μ z) ((fun y ↦ (y, z)) ⁻¹' t) := by + rw [← lintegral_indicator_one (ht.preimage (by fun_prop : Measurable fun y ↦ (y, z)))] + exact lintegral_congr fun _ ↦ rfl + _ = ((condDistrib X T μ ×ₖ Kernel.id) z) t := by + rw [Kernel.prod_apply, Kernel.id_apply, Measure.prod_apply_symm ht, lintegral_dirac] + +-- proved by Claude, then modified +lemma CondIndepFun.prod_right [StandardBorelSpace α] [StandardBorelSpace β] [Nonempty β] + [StandardBorelSpace γ] [Nonempty γ] [StandardBorelSpace δ] [Nonempty δ] + {X : α → β} {Y : α → γ} {Z : α → δ} + (hX : Measurable X) (hY : Measurable Y) (hZ : Measurable Z) + (h : X ⟂ᵢ[Z, hZ; μ] Y) : + X ⟂ᵢ[Z, hZ; μ] (fun ω ↦ (Y ω, Z ω)) := by + rw [condIndepFun_iff_condDistrib_prod_ae_eq_prodMkRight hY hX hZ, + condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h + rw [condIndepFun_iff_condDistrib_prod_ae_eq_prodMkRight (by fun_prop) hX hZ, + condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] + -- Key: condDistrib (Y, Z) Z μ z = (condDistrib Y Z μ z).map (y ↦ (y, z)) + have h_cond : condDistrib (fun ω ↦ (Y ω, Z ω)) Z μ =ᵐ[μ.map Z] + fun z ↦ (condDistrib Y Z μ z).map (fun y ↦ (y, z)) := by + suffices condDistrib (fun ω ↦ (Y ω, Z ω)) Z μ =ᵐ[μ.map Z] + (condDistrib Y Z μ) ×ₖ Kernel.id by + refine this.trans (ae_of_all _ fun x ↦ ?_) + simp only + rw [Kernel.prod_apply, Kernel.id_apply] + ext s hs + rw [Measure.map_apply (by fun_prop) hs, Measure.prod_apply_symm hs, lintegral_dirac] + exact condDistrib_prod_self_left hY.aemeasurable hZ.aemeasurable + -- Main calculation + calc μ.map (fun x ↦ ((Z x, X x), (Y x, Z x))) + _ = (μ.map (fun x ↦ ((Z x, X x), Y x))).map (fun p ↦ (p.1, (p.2, p.1.1))) := by + rw [Measure.map_map (by fun_prop) (by fun_prop)]; rfl + _ = (μ.map (fun ω ↦ (Z ω, X ω)) ⊗ₘ (condDistrib Y Z μ).prodMkRight β).map + (fun p ↦ (p.1, (p.2, p.1.1))) := by rw [h] + _ = μ.map (fun ω ↦ (Z ω, X ω)) ⊗ₘ (condDistrib (fun ω ↦ (Y ω, Z ω)) Z μ).prodMkRight β := by + ext s hs + rw [Measure.map_apply (by fun_prop) hs, + Measure.compProd_apply (hs.preimage (by fun_prop)), Measure.compProd_apply hs] + have h_cond' : ∀ᵐ p ∂(μ.map (fun ω ↦ (Z ω, X ω))), + condDistrib (fun ω ↦ (Y ω, Z ω)) Z μ p.1 = + (condDistrib Y Z μ p.1).map (fun y ↦ (y, p.1)) := by + have h_fst : (μ.map (fun ω ↦ (Z ω, X ω))).map Prod.fst = μ.map Z := by + rw [Measure.map_map (by fun_prop) (by fun_prop)]; rfl + rw [← h_fst] at h_cond + exact mem_ae_of_mem_ae_map (by fun_prop) h_cond + refine lintegral_congr_ae (h_cond'.mono fun ⟨z, x⟩ hzx ↦ ?_) + simp only [Kernel.prodMkRight_apply, hzx, + Measure.map_apply (by fun_prop : Measurable fun y ↦ (y, z)) + (hs.preimage (by fun_prop : Measurable (Prod.mk (z, x))))] + congr 1 + lemma fst_condDistrib_prod [StandardBorelSpace β] [Nonempty β] (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) (hT : AEMeasurable T μ) : (condDistrib (fun ω ↦ (X ω, Y ω)) T μ).fst =ᵐ[μ.map T] condDistrib X T μ := by diff --git a/LeanBandits/SequentialLearning/Deterministic.lean b/LeanBandits/SequentialLearning/Deterministic.lean index 8be520da..5ab42a05 100644 --- a/LeanBandits/SequentialLearning/Deterministic.lean +++ b/LeanBandits/SequentialLearning/Deterministic.lean @@ -17,16 +17,16 @@ namespace Learning variable {α R : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} -/-- A deterministic algorithm. -/ +/-- A deterministic algorithm, which chooses the action given by the function `nextAction`. -/ @[simps] noncomputable -def detAlgorithm (nextaction : (n : ℕ) → (Iic n → α × R) → α) - (h_next : ∀ n, Measurable (nextaction n)) (action0 : α) : +def detAlgorithm (nextAction : (n : ℕ) → (Iic n → α × R) → α) + (h_next : ∀ n, Measurable (nextAction n)) (action0 : α) : Algorithm α R where - policy n := Kernel.deterministic (nextaction n) (h_next n) + policy n := Kernel.deterministic (nextAction n) (h_next n) p0 := Measure.dirac action0 -variable {nextaction : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextaction n)} +variable {nextAction : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextAction n)} {action0 : α} {env : Environment α R} namespace IsAlgEnvSeq @@ -37,13 +37,13 @@ variable {Ω : Type*} {mΩ : MeasurableSpace Ω} {P : Measure Ω} [IsProbabilityMeasure P] {A : ℕ → Ω → α} {R' : ℕ → Ω → R} lemma HasLaw_action_zero_detAlgorithm - (h : IsAlgEnvSeq A R' (detAlgorithm nextaction h_next action0) env P) : + (h : IsAlgEnvSeq A R' (detAlgorithm nextAction h_next action0) env P) : HasLaw (A 0) (Measure.dirac action0) P where aemeasurable := have hA := h.measurable_A; by fun_prop map_eq := (hasLaw_action_zero h).map_eq lemma action_zero_detAlgorithm - (h : IsAlgEnvSeq A R' (detAlgorithm nextaction h_next action0) env P) : + (h : IsAlgEnvSeq A R' (detAlgorithm nextAction h_next action0) env P) : A 0 =ᵐ[P] fun _ ↦ action0 := by have h_eq : ∀ᵐ x ∂(P.map (A 0)), x = action0 := by rw [(hasLaw_action_zero h).map_eq] @@ -52,16 +52,16 @@ lemma action_zero_detAlgorithm exact ae_of_ae_map (by fun_prop) h_eq lemma action_detAlgorithm_ae_eq - (h : IsAlgEnvSeq A R' (detAlgorithm nextaction h_next action0) env P) (n : ℕ) : - A (n + 1) =ᵐ[P] fun ω ↦ nextaction n (hist A R' n ω) := by + (h : IsAlgEnvSeq A R' (detAlgorithm nextAction h_next action0) env P) (n : ℕ) : + A (n + 1) =ᵐ[P] fun ω ↦ nextAction n (hist A R' n ω) := by have hA := h.measurable_A have hR' := h.measurable_R exact ae_eq_of_condDistrib_eq_deterministic (by fun_prop) (by fun_prop) (by fun_prop) (h.hasCondDistrib_action n).condDistrib_eq lemma action_detAlgorithm_ae_all_eq - (h : IsAlgEnvSeq A R' (detAlgorithm nextaction h_next action0) env P) : - ∀ᵐ ω ∂P, A 0 ω = action0 ∧ ∀ n, A (n + 1) ω = nextaction n (hist A R' n ω) := by + (h : IsAlgEnvSeq A R' (detAlgorithm nextAction h_next action0) env P) : + ∀ᵐ ω ∂P, A 0 ω = action0 ∧ ∀ n, A (n + 1) ω = nextAction n (hist A R' n ω) := by rw [eventually_and, ae_all_iff] exact ⟨action_zero_detAlgorithm h, action_detAlgorithm_ae_eq h⟩ @@ -69,7 +69,7 @@ end IsAlgEnvSeq namespace IT -local notation "𝔓" => trajMeasure (detAlgorithm nextaction h_next action0) env +local notation "𝔓" => trajMeasure (detAlgorithm nextAction h_next action0) env lemma HasLaw_action_zero_detAlgorithm : HasLaw (IT.action 0) (Measure.dirac action0) 𝔓 where map_eq := (IT.hasLaw_action_zero _ _).map_eq @@ -82,13 +82,13 @@ lemma action_zero_detAlgorithm [MeasurableSingletonClass α] : exact ae_of_ae_map (by fun_prop) h_eq lemma action_detAlgorithm_ae_eq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] - [Nonempty R] (n : ℕ) : IT.action (n + 1) =ᵐ[𝔓] fun h ↦ nextaction n (IT.hist n h) := + [Nonempty R] (n : ℕ) : IT.action (n + 1) =ᵐ[𝔓] fun h ↦ nextAction n (IT.hist n h) := ae_eq_of_condDistrib_eq_deterministic (by fun_prop) (by fun_prop) (by fun_prop) - (IT.condDistrib_action (detAlgorithm nextaction h_next action0) env n) + (IT.condDistrib_action (detAlgorithm nextAction h_next action0) env n) lemma action_detAlgorithm_ae_all_eq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] : - ∀ᵐ h ∂𝔓, IT.action 0 h = action0 ∧ ∀ n, IT.action (n + 1) h = nextaction n (IT.hist n h) := by + ∀ᵐ h ∂𝔓, IT.action 0 h = action0 ∧ ∀ n, IT.action (n + 1) h = nextAction n (IT.hist n h) := by rw [eventually_and, ae_all_iff] exact ⟨action_zero_detAlgorithm, action_detAlgorithm_ae_eq⟩ diff --git a/LeanBandits/SequentialLearning/StationaryEnv.lean b/LeanBandits/SequentialLearning/StationaryEnv.lean index 35d44157..a52b4a0c 100644 --- a/LeanBandits/SequentialLearning/StationaryEnv.lean +++ b/LeanBandits/SequentialLearning/StationaryEnv.lean @@ -87,46 +87,30 @@ local notation "𝔓" => trajMeasure alg (stationaryEnv ν) /-- The conditional distribution of the reward at time `n` given the action at time `n` is `ν`. -/ lemma condDistrib_reward_stationaryEnv (n : ℕ) : - condDistrib (IT.reward n) (IT.action n) 𝔓 =ᵐ[(𝔓).map (IT.action n)] ν := by - cases n with - | zero => - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] - change (𝔓).map (IT.step 0) = (𝔓).map (IT.action 0) ⊗ₘ ν - rw [(IT.hasLaw_action_zero alg (stationaryEnv ν)).map_eq, - (IT.hasLaw_step_zero alg (stationaryEnv ν)).map_eq, stationaryEnv_ν0] - | succ n => - have h_eq := IT.condDistrib_reward alg (stationaryEnv ν) n - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h_eq ⊢ - have : (𝔓).map (IT.action (n + 1)) = - ((𝔓).map (fun x ↦ (IT.hist n x, IT.action (n + 1) x))).snd := by - rw [Measure.snd_map_prodMk (by fun_prop)] - simp only [stationaryEnv_feedback] at h_eq - rw [this, ← Measure.snd_prodAssoc_compProd_prodMkLeft, ← h_eq, - Measure.snd_map_prodMk (by fun_prop), Measure.map_map (by fun_prop) (by fun_prop)] - congr + condDistrib (IT.reward n) (IT.action n) 𝔓 =ᵐ[(𝔓).map (IT.action n)] ν := + IsAlgEnvSeq.condDistrib_reward_stationaryEnv + (IT.isAlgEnvSeq_trajMeasure alg (stationaryEnv ν)) n /-- The reward at time `n + 1` is conditionally independent of the history up to time `n` given the action at time `n + 1`. -/ lemma condIndepFun_reward_hist_action (n : ℕ) : IT.reward (n + 1) ⟂ᵢ[IT.action (n + 1), IT.measurable_action _ ; 𝔓] IT.hist n := - condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft - (by fun_prop) (by fun_prop) (by fun_prop) (IT.condDistrib_reward alg (stationaryEnv ν) n) + IsAlgEnvSeq.condIndepFun_reward_hist_action + (IT.isAlgEnvSeq_trajMeasure alg (stationaryEnv ν)) n lemma condIndepFun_reward_hist_action_action {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) : reward (n + 1) ⟂ᵢ[action (n + 1), measurable_action (n + 1); trajMeasure alg (stationaryEnv ν)] - (fun ω ↦ (hist n ω, action (n + 1) ω)) := by - have h_indep : reward (n + 1) ⟂ᵢ[action (n + 1), measurable_action (n + 1); - trajMeasure alg (stationaryEnv ν)] hist n := by - convert condIndepFun_reward_hist_action (alg := alg) (ν := ν) n - exact h_indep.prod_right (by fun_prop) (by fun_prop) (by fun_prop) + (fun ω ↦ (hist n ω, action (n + 1) ω)) := + IsAlgEnvSeq.condIndepFun_reward_hist_action_action + (IT.isAlgEnvSeq_trajMeasure alg (stationaryEnv ν)) n lemma condIndepFun_reward_hist_action_action' {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) (hn : n ≠ 0) : reward n ⟂ᵢ[action n, measurable_action n; trajMeasure alg (stationaryEnv ν)] - (fun ω ↦ (hist (n - 1) ω, action n ω)) := by - have := condIndepFun_reward_hist_action_action (alg := alg) (ν := ν) (n - 1) - grind + (fun ω ↦ (hist (n - 1) ω, action n ω)) := + IsAlgEnvSeq.condIndepFun_reward_hist_action_action' + (IT.isAlgEnvSeq_trajMeasure alg (stationaryEnv ν)) n hn end IT From 63f5e4208fa0f592c5aae4789b730a939850f003 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 14 Jan 2026 18:05:51 +0100 Subject: [PATCH 22/30] add indepedence proof from Claude --- LeanBandits/Bandit/Bandit.lean | 53 ++++++++++++++++++++++--- LeanBandits/ForMathlib/CondDistrib.lean | 8 ++-- 2 files changed, 52 insertions(+), 9 deletions(-) diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 0e19def2..1b94cc0f 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -436,11 +436,54 @@ lemma hasCondDistrib_reward_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMa · simp · rwa [Measure.map_apply (by fun_prop) (by simp)] at ha -omit [DecidableEq α] in -lemma indepFun_fst_add_one_aux (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : +-- proved by Claude, then slightly golfed +omit [DecidableEq α] [Nonempty α] [StandardBorelSpace α] [Countable α] [StandardBorelSpace R] + [Nonempty R] in +lemma indepFun_fst_add_one_aux (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : (fun ω ↦ ω.1 (n + 1)) ⟂ᵢ[arrayMeasure ν] (fun ω ↦ (fun (i : Iic n) ↦ ω.1 i, ω.2)) := by - rw [indepFun_iff_map_prod_eq_prod_map_map (by fun_prop) (by fun_prop)] - sorry + let μ₁ : Measure (ℕ → I) := Measure.infinitePi fun _ ↦ volume + let μ₂ : Measure (ℕ → α → R) := Bandit.streamMeasure ν + -- Coordinates of μ₁ are independent + have h_indep : iIndepFun (fun i (ω : ℕ → I) ↦ ω i) μ₁ := + iIndepFun_infinitePi (fun _ ↦ measurable_id) + have h_indep_n : IndepFun (fun ω ↦ ω (n + 1)) (fun ω ↦ fun i : Iic n ↦ ω i) μ₁ := by + have h := h_indep.indepFun_finset₀ {n + 1} (Iic n) (by simp) + (fun i ↦ (measurable_pi_apply i).aemeasurable) + convert h.comp (measurable_pi_apply ⟨n + 1, by simp⟩) measurable_id using 1 + rw [indepFun_iff_measure_inter_preimage_eq_mul] + intro s t hs ht + let X : (ℕ → I) × (ℕ → α → R) → I := fun ω ↦ ω.1 (n + 1) + let Y : (ℕ → I) × (ℕ → α → R) → (Iic n → I) × (ℕ → α → R) := fun ω ↦ (fun i ↦ ω.1 i, ω.2) + change (μ₁.prod μ₂) (X ⁻¹' s ∩ Y ⁻¹' t) = (μ₁.prod μ₂) (X ⁻¹' s) * (μ₁.prod μ₂) (Y ⁻¹' t) + -- Rewrite using Fubini + rw [Measure.prod_apply (hs.preimage (by fun_prop : Measurable X)), + Measure.prod_apply (ht.preimage (by fun_prop : Measurable Y)), + Measure.prod_apply ((hs.preimage (by fun_prop : Measurable X)).inter + (ht.preimage (by fun_prop : Measurable Y)))] + -- Compute fibers + have hX_fst ω₁ : μ₂ (Prod.mk ω₁ ⁻¹' (X ⁻¹' s)) = s.indicator 1 (ω₁ (n + 1)) := by + simp only [X, Set.preimage_preimage] + by_cases h : ω₁ (n + 1) ∈ s <;> simp [h] + have hY_fst ω₁ : μ₂ (Prod.mk ω₁ ⁻¹' (Y ⁻¹' t)) = μ₂ {y | ((fun i : Iic n ↦ ω₁ i), y) ∈ t} := rfl + have hXY ω₁ : μ₂ (Prod.mk ω₁ ⁻¹' (X ⁻¹' s ∩ Y ⁻¹' t)) = + s.indicator 1 (ω₁ (n + 1)) * μ₂ {y | ((fun i : Iic n ↦ ω₁ i), y) ∈ t} := by + simp only [X, Y, Set.preimage_inter, Set.preimage_preimage] + by_cases h : ω₁ (n + 1) ∈ s + · simp [h] + grind + · simp [h] + simp_rw [hY_fst, hX_fst, hXY] + -- Factor the integral using independence + let g : (Iic n → I) → ENNReal := fun x ↦ μ₂ {y | (x, y) ∈ t} + have hg_meas : Measurable g := measurable_measure_prodMk_left ht + have hf_meas : Measurable (fun ω₁ : ℕ → I ↦ s.indicator (1 : I → ENNReal) (ω₁ (n + 1))) := + (measurable_one.indicator hs).comp (measurable_pi_apply _) + have hindep_fg : IndepFun (fun ω₁ ↦ s.indicator (1 : I → ENNReal) (ω₁ (n + 1))) + (fun ω₁ ↦ g (fun i ↦ ω₁ i)) μ₁ := + h_indep_n.comp (measurable_one.indicator hs) hg_meas + have h_eq (ω₁ : ℕ → I) : μ₂ {y | ((fun i : Iic n ↦ ω₁ i), y) ∈ t} = g (fun i ↦ ω₁ i) := rfl + simp_rw [h_eq] + exact lintegral_mul_eq_lintegral_mul_lintegral_of_indepFun hf_meas (by fun_prop) hindep_fg omit [StandardBorelSpace R] [Nonempty R] in lemma measurable_hist_todo (alg : Algorithm α R) (n : ℕ) : @@ -520,7 +563,7 @@ lemma measurable_hist_todo (alg : Algorithm α R) (n : ℕ) : lemma indepFun_fst_add_one_hist (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : IndepFun (fun ω ↦ ω.1 (n + 1)) (hist alg · n) (arrayMeasure ν) := - (indepFun_fst_add_one_aux alg ν n).of_measurable_right (measurable_hist_todo alg n) + (indepFun_fst_add_one_aux ν n).of_measurable_right (measurable_hist_todo alg n) lemma hasCondDistrib_action' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : HasCondDistrib (action alg (n + 1)) (hist alg · n) (alg.policy n) (arrayMeasure ν) := by diff --git a/LeanBandits/ForMathlib/CondDistrib.lean b/LeanBandits/ForMathlib/CondDistrib.lean index bf137f29..030e1bd6 100644 --- a/LeanBandits/ForMathlib/CondDistrib.lean +++ b/LeanBandits/ForMathlib/CondDistrib.lean @@ -97,14 +97,14 @@ lemma condDistrib_prod_self_left [StandardBorelSpace β] [Nonempty β] [Standard have h_fst' := (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_fst filter_upwards [h_prod, h_fst'] with z hz1 hz2 rw [hz1] + simp only [Kernel.deterministic_apply] at hz2 + change ∀ᵐ y ∂(condDistrib X T μ z), condDistrib T (fun ω ↦ (T ω, X ω)) μ (z, y) = Measure.dirac z + at hz2 ext t ht rw [Kernel.compProd_apply ht] - have hz2' : ∀ᵐ y ∂(condDistrib X T μ z), - condDistrib T (fun ω ↦ (T ω, X ω)) μ (z, y) = Measure.dirac z := by - filter_upwards [hz2] with y hy; convert hy using 2 calc ∫⁻ y, condDistrib T (fun ω ↦ (T ω, X ω)) μ (z, y) (Prod.mk y ⁻¹' t) ∂condDistrib X T μ z _ = ∫⁻ y, (Measure.dirac z) (Prod.mk y ⁻¹' t) ∂condDistrib X T μ z := - lintegral_congr_ae (hz2'.mono fun y hy ↦ by simp only [hy]) + lintegral_congr_ae (hz2.mono fun y hy ↦ by simp only [hy]) _ = ∫⁻ y, (Prod.mk y ⁻¹' t).indicator 1 z ∂condDistrib X T μ z := lintegral_congr fun y ↦ Measure.dirac_apply' _ (ht.preimage (by fun_prop)) _ = (condDistrib X T μ z) ((fun y ↦ (y, z)) ⁻¹' t) := by From b8bca0607da560045a469289a7ccda5742398cfb Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 14 Jan 2026 20:53:24 +0100 Subject: [PATCH 23/30] reorder Bandit file --- LeanBandits/Bandit/Bandit.lean | 288 +++++++++++++++------------------ 1 file changed, 127 insertions(+), 161 deletions(-) diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 1b94cc0f..cd6e2d7f 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -165,6 +165,8 @@ namespace ArrayModel open unitInterval +section ProbabilitySpace + variable (α R) in /-- Probability space for the array model of stochastic bandits. -/ def probSpace : Type _ := (ℕ → I) × (ℕ → α → R) @@ -213,6 +215,12 @@ lemma measurable_algFunction (alg : Algorithm α R) (n : ℕ) : Measurable (Function.uncurry (algFunction alg n)) := (Kernel.representation (alg.policy n)).choose_spec.1 +end ProbabilitySpace + +variable [Nonempty α] [StandardBorelSpace α] + +section HistoryActionReward + /-- History of actions and rewards up to time `n` in the array model. -/ noncomputable def hist [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) : (n : ℕ) → Iic n → α × R @@ -366,7 +374,95 @@ lemma hist_add_one_eq_IicSuccProd [DecidableEq α] (alg : Algorithm α R) (ω : (hist alg ω n, (action alg (n + 1) ω, reward alg (n + 1) ω)) := by rw [hist_add_one_eq_IicSuccProd', reward_add_one, action_add_one_eq] -variable [DecidableEq α] [Countable α] +end HistoryActionReward + +variable [DecidableEq α] + +section Congruence + +-- very useful to prove measurability +lemma hist_congr (alg : Algorithm α R) (n : ℕ) {ω ω' : probSpace α R} + (hω1 : ∀ i ≤ n, ω.1 i = ω'.1 i) + (hω2 : ∀ i a, i < pullCount (action alg) a (n + 1) ω → ω.2 i a = ω'.2 i a) : + hist alg ω n = hist alg ω' n := by + induction n with + | zero => + simp only [zero_add, pullCount_one] at hω2 + simp_rw [hist_zero] + ext i : 1 + simp only [le_refl, hω1, Prod.mk.injEq, true_and] + refine hω2 0 _ ?_ + simp [action, hω1] + | succ n hn => + simp_rw [hist_add_one_eq_IicSuccProd] + specialize hn fun i hin ↦ hω1 i (by grind) + have h_hist : hist alg ω n = hist alg ω' n := by + refine hn fun i a hi ↦ hω2 i a (hi.trans_le ?_) + exact pullCount_mono _ (by lia) _ + have h_action : action alg (n + 1) ω = action alg (n + 1) ω' := by + simp_rw [action_add_one_eq] + rw [h_hist, hω1 _ le_rfl] + congr 3 + simp only [reward_add_one, h_hist, h_action] + refine hω2 _ _ ?_ + rw [pullCount_add_one, h_action] + simp only [↓reduceIte] + rw [pullCount_eq_pullCount' (R' := reward alg) (by simp)] + simp only [Nat.add_one_sub_one] + rw [← h_hist, hist_eq] + change pullCount' n (fun i ↦ (action alg i ω, reward alg i ω)) (action alg (n + 1) ω') < + pullCount' n (fun i ↦ (action alg i ω, reward alg i ω)) (action alg (n + 1) ω') + 1 + grind + +lemma stepsUntil_congr_aux (alg : Algorithm α R) + (a : α) (m n : ℕ) {ω ω' : probSpace α R} + (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) + (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) + (h_eq : action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m) : + action alg (n + 1) ω' = a ∧ pullCount (action alg) a (n + 1) ω' = m := by + obtain ⟨h_action, h_pc⟩ := h_eq + have h_hist := hist_congr alg n (ω := ω) (ω' := ω') (by grind) fun i b hi ↦ ?_ + swap + · rcases eq_or_ne b a with (rfl | hba) + · refine hω2_eq i ?_ + rw [h_pc] at hi + grind + · grind + constructor + · rw [← h_action, action_add_one_eq] + simp [h_hist, hω1] + · simp_rw [← h_pc, pullCount_eq_sum] + refine Finset.sum_congr rfl fun i hi ↦ ?_ + congr 2 + rw [hist_eq _ _ n, hist_eq _ _ n, funext_iff] at h_hist + unfold action + specialize h_hist ⟨i, by grind⟩ + simp only at h_hist + rw [h_hist] + +lemma stepsUntil_congr (alg : Algorithm α R) (a : α) (m n : ℕ) {ω ω' : probSpace α R} + (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) + (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) : + (action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m) ↔ + (action alg (n + 1) ω' = a ∧ pullCount (action alg) a (n + 1) ω' = m) := + ⟨stepsUntil_congr_aux alg a m n hω1 hω2_ne hω2_eq, + stepsUntil_congr_aux alg a m n (by grind) (by grind) (by grind)⟩ + +lemma stepsUntil_indicator_congr (alg : Algorithm α R) (a : α) (m n : ℕ) {ω ω' : probSpace α R} + (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) + (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) : + {ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}.indicator (fun _ ↦ 1) + ω = + {ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}.indicator + (fun _ ↦ 1) ω' := by + simp only [Set.indicator_apply, Set.mem_setOf_eq] + simp_rw [stepsUntil_congr alg a m n hω1 hω2_ne hω2_eq] + +end Congruence + +section Laws + +variable [Countable α] lemma hasLaw_action_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : HasLaw (action alg 0) alg.p0 (arrayMeasure ν) where @@ -381,23 +477,18 @@ lemma hasLaw_action_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKern rw [(measurePreserving_eval_infinitePi (fun _ ↦ volume) 0).map_eq] _ = alg.p0 := initAlgFunction_map alg -variable [StandardBorelSpace R] [Nonempty R] - -omit [Nonempty α] [StandardBorelSpace α] [DecidableEq α] [Countable α] - [StandardBorelSpace R] [Nonempty R] in +omit [Nonempty α] [StandardBorelSpace α] [DecidableEq α] [Countable α] in lemma indepFun_fst_snd (ν : Kernel α R) [IsMarkovKernel ν] : IndepFun Prod.fst Prod.snd (arrayMeasure ν) := indepFun_prod measurable_id measurable_id -omit [Nonempty α] [StandardBorelSpace α] [DecidableEq α] [Countable α] - [StandardBorelSpace R] [Nonempty R] in +omit [Nonempty α] [StandardBorelSpace α] [DecidableEq α] [Countable α] in lemma indepFun_fst_zero_snd_zero_action (ν : Kernel α R) [IsMarkovKernel ν] (a : α) : IndepFun (fun ω ↦ ω.1 0) (fun ω ↦ ω.2 0 a) (arrayMeasure ν) := indepFun_prod (X := fun ω : ℕ → I ↦ ω 0) (Y := fun ω : ℕ → α → R ↦ ω 0 a) (by fun_prop) (by fun_prop) -omit [Nonempty α] [StandardBorelSpace α] [DecidableEq α] [Countable α] - [StandardBorelSpace R] [Nonempty R] in +omit [Nonempty α] [StandardBorelSpace α] [DecidableEq α] [Countable α] in lemma map_snd_apply_arrayMeasure {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) (a : α) : (arrayMeasure ν).map (fun ω ↦ ω.2 n a) = ν a := by calc (arrayMeasure ν).map (fun ω ↦ ω.2 n a) @@ -410,6 +501,8 @@ lemma map_snd_apply_arrayMeasure {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ rw [this, ← Measure.map_map (by fun_prop) (by fun_prop), Measure.infinitePi_map_eval, Measure.infinitePi_map_eval] +variable [StandardBorelSpace R] [Nonempty R] + lemma hasCondDistrib_reward_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : HasCondDistrib (reward alg 0) (action alg 0) ν (arrayMeasure ν) where condDistrib_eq := by @@ -489,77 +582,18 @@ omit [StandardBorelSpace R] [Nonempty R] in lemma measurable_hist_todo (alg : Algorithm α R) (n : ℕ) : Measurable[MeasurableSpace.comap (fun ω ↦ (fun (i : Iic n) ↦ ω.1 i, ω.2)) inferInstance] (hist alg · n) := by - induction n with - | zero => - simp only [hist_zero] - have : (fun (ω : probSpace α R) (i : Iic 0) ↦ - (initAlgFunction alg (ω.1 0), ω.2 0 (initAlgFunction alg (ω.1 0)))) = - (fun (p : (Iic 0 → I) × (ℕ → α → R)) (i : Iic 0) ↦ (initAlgFunction alg (p.1 ⟨0, by simp⟩), - p.2 0 (initAlgFunction alg (p.1 ⟨0, by simp⟩)))) ∘ - (fun (ω : probSpace α R) ↦ (fun (i : Iic 0) ↦ ω.1 i, ω.2)) := rfl - rw [this] - have h_meas : Measurable (fun (p : (Iic 0 → I) × (ℕ → α → R)) (i : Iic 0) ↦ - (initAlgFunction alg (p.1 ⟨0, by simp⟩), - p.2 0 (initAlgFunction alg (p.1 ⟨0, by simp⟩)))) := by - rw [measurable_pi_iff] - intro i - refine Measurable.prodMk (by fun_prop) ?_ - change Measurable ((fun x : (α × (ℕ → α → R)) ↦ x.2 0 x.1) ∘ - (fun x : (Iic 0 → I) × (ℕ → α → R) ↦ (initAlgFunction alg (x.1 ⟨0, by simp⟩), x.2))) - have : Measurable (fun x : (α × (ℕ → α → R)) ↦ x.2 0 x.1) := - measurable_from_prod_countable_right fun p ↦ by simp only; fun_prop - exact this.comp (by fun_prop) - refine Measurable.comp h_meas ?_ - exact Measurable.of_comap_le le_rfl - | succ n hn => - simp_rw [hist_add_one_eq_IicSuccProd] - have h_hist : Measurable[MeasurableSpace.comap (fun ω ↦ (fun (i : Iic (n + 1)) ↦ ω.1 i, ω.2)) - inferInstance] (hist alg · n) := by - rw [measurable_iff_comap_le] at hn ⊢ - refine hn.trans ?_ - rw [← measurable_iff_comap_le] - have : (fun (ω : probSpace α R) ↦ (fun (i : Iic n) ↦ ω.1 i, ω.2)) = - (fun (p : (Iic (n + 1) → I) × (ℕ → α → R)) ↦ (fun (i : Iic n) ↦ p.1 ⟨i, by grind⟩, p.2)) ∘ - (fun (ω : probSpace α R) ↦ (fun (i : Iic (n + 1)) ↦ ω.1 i, ω.2)) := rfl - rw [this] - exact Measurable.comp (by fun_prop) (Measurable.of_comap_le le_rfl) - refine (MeasurableEquiv.measurable _).comp (Measurable.prodMk h_hist ?_) - have h_action : Measurable[MeasurableSpace.comap (fun ω ↦ (fun (i : Iic (n + 1)) ↦ ω.1 i, ω.2)) - inferInstance] (action alg (n + 1)) := by - rw [action_add_one_eq] - have : (fun ω ↦ algFunction alg n (hist alg ω n) (ω.1 (n + 1))) = - (Function.uncurry (algFunction alg n)) ∘ (fun ω ↦ (hist alg ω n, ω.1 (n + 1))) := rfl - rw [this] - refine (measurable_algFunction alg n).comp (h_hist.prodMk ?_) - have : (fun ω : probSpace α R ↦ ω.1 (n + 1)) = - (fun (p : (Iic (n + 1) → I) × (ℕ → α → R)) ↦ p.1 ⟨n + 1, by simp⟩) ∘ - (fun ω ↦ (fun (i : Iic (n + 1)) ↦ ω.1 i, ω.2)) := rfl - rw [this] - exact Measurable.comp (by fun_prop) (Measurable.of_comap_le le_rfl) - refine h_action.prodMk ?_ - rw [reward_add_one] - have : (fun ω ↦ ω.2 (pullCount' n (hist alg ω n) (action alg (n + 1) ω)) - (action alg (n + 1) ω)) = - (fun p : ((Iic (n + 1) → I) × (ℕ → α → R)) × - (Iic n → α × R) × α ↦ p.1.2 (pullCount' n p.2.1 p.2.2) p.2.2) ∘ - (fun ω ↦ ((fun i : Iic (n + 1) ↦ ω.1 i, ω.2), hist alg ω n, action alg (n + 1) ω)) := rfl - rw [this] - have h_meas : Measurable - (fun p : ((Iic (n + 1) → I) × (ℕ → α → R)) × (Iic n → α × R) × α ↦ - p.1.2 (pullCount' n p.2.1 p.2.2) p.2.2) := by - have : (fun p : ((Iic (n + 1) → I) × (ℕ → α → R)) × (Iic n → α × R) × α ↦ - p.1.2 (pullCount' n p.2.1 p.2.2) p.2.2) = - (fun (x : (ℕ → α → R) × ℕ × α) ↦ x.1 x.2.1 x.2.2) ∘ - (fun p : ((Iic (n + 1) → I) × (ℕ → α → R)) × (Iic n → α × R) × α ↦ - (p.1.2, pullCount' n p.2.1 p.2.2, p.2.2)) := rfl - rw [this] - refine Measurable.comp (measurable_from_prod_countable_left (m := inferInstance) fun p ↦ ?_) - ?_ - · simp only; fun_prop - refine Measurable.prodMk (by fun_prop) (Measurable.prodMk ?_ (by fun_prop)) - exact (measurable_uncurry_pullCount' (α := α) (mR := mR) n).comp (by fun_prop) - refine h_meas.comp (Measurable.prodMk ?_ (Measurable.prodMk h_hist h_action)) - exact Measurable.of_comap_le le_rfl + have h_eq : (hist alg · n) = + ((hist alg · n) ∘ (fun p ↦ (fun i : ℕ ↦ p.1 ⟨min i n, by grind⟩, p.2))) ∘ + (fun ω ↦ (fun (i : Iic n) ↦ ω.1 i, ω.2)) := by + ext ω : 1 + exact hist_congr alg n (by grind) (by simp) + rw [h_eq] + refine measurable_comp_comap _ (Measurable.comp (by fun_prop) ?_) + refine Measurable.prodMk ?_ (by fun_prop) + rw [measurable_pi_iff] + intro i + change Measurable ((fun p ↦ p ⟨min i n, by simp⟩) ∘ (fun x : (Iic n → I) × (ℕ → α → R) ↦ x.1)) + exact Measurable.comp (by fun_prop) measurable_fst lemma indepFun_fst_add_one_hist (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : IndepFun (fun ω ↦ ω.1 (n + 1)) (hist alg · n) (arrayMeasure ν) := @@ -620,40 +654,6 @@ lemma hasCondDistrib_action' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkov · exact hs · exact hs.preimage (by fun_prop) -omit [Countable α] [StandardBorelSpace R] [Nonempty R] in -lemma hist_congr (alg : Algorithm α R) (n : ℕ) {ω ω' : probSpace α R} - (hω1 : ∀ i ≤ n, ω.1 i = ω'.1 i) - (hω2 : ∀ i a, i < pullCount (action alg) a (n + 1) ω → ω.2 i a = ω'.2 i a) : - hist alg ω n = hist alg ω' n := by - induction n with - | zero => - simp only [zero_add, pullCount_one] at hω2 - simp_rw [hist_zero] - ext i : 1 - simp only [le_refl, hω1, Prod.mk.injEq, true_and] - refine hω2 0 _ ?_ - simp [action, hω1] - | succ n hn => - simp_rw [hist_add_one_eq_IicSuccProd] - specialize hn fun i hin ↦ hω1 i (by grind) - have h_hist : hist alg ω n = hist alg ω' n := by - refine hn fun i a hi ↦ hω2 i a (hi.trans_le ?_) - exact pullCount_mono _ (by lia) _ - have h_action : action alg (n + 1) ω = action alg (n + 1) ω' := by - simp_rw [action_add_one_eq] - rw [h_hist, hω1 _ le_rfl] - congr 3 - simp only [reward_add_one, h_hist, h_action] - refine hω2 _ _ ?_ - rw [pullCount_add_one, h_action] - simp only [↓reduceIte] - rw [pullCount_eq_pullCount' (R' := reward alg) (by simp)] - simp only [Nat.add_one_sub_one] - rw [← h_hist, hist_eq] - change pullCount' n (fun i ↦ (action alg i ω, reward alg i ω)) (action alg (n + 1) ω') < - pullCount' n (fun i ↦ (action alg i ω, reward alg i ω)) (action alg (n + 1) ω') + 1 - grind - -- very bad name /-- All random variables in the space, except for the unseen rewards for action `a` after time `n`. -/ @@ -721,53 +721,6 @@ lemma indepFun_snd_apply_aux (alg : Algorithm α R) (fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b)) := by sorry -omit [Countable α] [StandardBorelSpace R] [Nonempty R] in -lemma stepsUntil_congr_aux (alg : Algorithm α R) - (a : α) (m n : ℕ) {ω ω' : probSpace α R} - (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) - (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) - (h_eq : action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m) : - action alg (n + 1) ω' = a ∧ pullCount (action alg) a (n + 1) ω' = m := by - obtain ⟨h_action, h_pc⟩ := h_eq - have h_hist := hist_congr alg n (ω := ω) (ω' := ω') (by grind) fun i b hi ↦ ?_ - swap - · rcases eq_or_ne b a with (rfl | hba) - · refine hω2_eq i ?_ - rw [h_pc] at hi - grind - · grind - constructor - · rw [← h_action, action_add_one_eq] - simp [h_hist, hω1] - · simp_rw [← h_pc, pullCount_eq_sum] - refine Finset.sum_congr rfl fun i hi ↦ ?_ - congr 2 - rw [hist_eq _ _ n, hist_eq _ _ n, funext_iff] at h_hist - unfold action - specialize h_hist ⟨i, by grind⟩ - simp only at h_hist - rw [h_hist] - -omit [Countable α] [StandardBorelSpace R] [Nonempty R] in -lemma stepsUntil_congr (alg : Algorithm α R) (a : α) (m n : ℕ) {ω ω' : probSpace α R} - (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) - (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) : - (action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m) ↔ - (action alg (n + 1) ω' = a ∧ pullCount (action alg) a (n + 1) ω' = m) := - ⟨stepsUntil_congr_aux alg a m n hω1 hω2_ne hω2_eq, - stepsUntil_congr_aux alg a m n (by grind) (by grind) (by grind)⟩ - -omit [Countable α] [StandardBorelSpace R] [Nonempty R] in -lemma stepsUntil_indicator_congr (alg : Algorithm α R) (a : α) (m n : ℕ) {ω ω' : probSpace α R} - (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) - (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) : - {ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}.indicator (fun _ ↦ 1) - ω = - {ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}.indicator - (fun _ ↦ 1) ω' := by - simp only [Set.indicator_apply, Set.mem_setOf_eq] - simp_rw [stepsUntil_congr alg a m n hω1 hω2_ne hω2_eq] - omit [StandardBorelSpace R] [Nonempty R] in lemma measurable_stepsUntil (alg : Algorithm α R) (a : α) (m n : ℕ) : Measurable[MeasurableSpace.comap @@ -803,6 +756,9 @@ lemma measurable_pullCount_action_add_one (alg : Algorithm α R) (n : ℕ) : (fun ω : probSpace α R ↦ (ω, action alg (n + 1) ω))) exact (measurable_uncurry_pullCount (by fun_prop) _).comp (by fun_prop) +/-- The conditional distribution of the reward at time `n + 1`, given the action at time `n + 1` +and the number of times that action has been pulled before time `n + 1`, is equal to +the kernel `ν`. -/ lemma hasCondDistrib_reward_pullCount_action (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : HasCondDistrib (reward alg (n + 1)) @@ -931,6 +887,9 @@ lemma indepFun_snd_hist_cond (alg : Algorithm α R) refine Measurable.ite ?_ (by fun_prop) (by fun_prop) exact MeasurableSet.const _ +/-- The conditional distribution of the reward at time `n + 1`, given the history up to time `n`, +the action at time `n + 1`, and the number of times that action has been pulled before time `n + 1`, +is equal to the kernel `ν`. -/ lemma hasCondDistrib_reward_hist_action_pullCount (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : HasCondDistrib (reward alg (n + 1)) @@ -973,6 +932,9 @@ lemma hasCondDistrib_reward_hist_action_pullCount intro ha simp [ha] +/-- The reward at time `n + 1` is conditionally independent of the history up to time `n`, +given the action at time `n + 1` and the number of times that action has been pulled before +time `n + 1`. -/ lemma condIndepFun_reward_hist (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : (reward alg (n + 1)) ⟂ᵢ[(fun ω ↦ (action alg (n + 1) ω, pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)), @@ -1000,6 +962,8 @@ lemma measurable_pullCount_action_add_one_hist (alg : Algorithm α R) (n : ℕ) rw [this] exact measurable_comp_comap _ (Measurable.prodMk (by fun_prop) (by fun_prop)) +/-- The conditional distribution of the reward at time `n + 1`, given the history up to time `n` +and the action at time `n + 1`, is equal to the kernel `ν`. -/ lemma hasCondDistrib_reward' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : HasCondDistrib (reward alg (n + 1)) (fun ω ↦ (hist alg ω n, action alg (n + 1) ω)) (ν.prodMkLeft _) (arrayMeasure ν) := by @@ -1074,6 +1038,8 @@ lemma isAlgEnvSeq_arrayMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMark hasCondDistrib_action := hasCondDistrib_action alg ν hasCondDistrib_reward := hasCondDistrib_reward alg ν +end Laws + end ArrayModel end MeasureSpace From 89ed9a47c2f9ca93ae4d1237192ac10511f1d759 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 14 Jan 2026 21:24:47 +0100 Subject: [PATCH 24/30] comment out sorrys in RewardByCountMeasure --- LeanBandits/Bandit/RewardByCountMeasure.lean | 132 +++++++++---------- 1 file changed, 66 insertions(+), 66 deletions(-) diff --git a/LeanBandits/Bandit/RewardByCountMeasure.lean b/LeanBandits/Bandit/RewardByCountMeasure.lean index 60a2b369..85b8aa46 100644 --- a/LeanBandits/Bandit/RewardByCountMeasure.lean +++ b/LeanBandits/Bandit/RewardByCountMeasure.lean @@ -233,43 +233,43 @@ lemma identDistrib_rewardByCount_eval [StandardBorelSpace Ω] [Countable α] (identDistrib_rewardByCount_id h a n hn).trans (identDistrib_eval_eval_id_streamMeasure ν m a).symm -lemma indepFun_rewardByCount_Iic [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] - (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) - (n : ℕ) : - (rewardByCount A R a (n + 1)) ⟂ᵢ[𝔓'] fun ω (i : Iic n) ↦ rewardByCount A R a i ω := by - sorry +-- lemma indepFun_rewardByCount_Iic [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) +-- (n : ℕ) : +-- (rewardByCount A R a (n + 1)) ⟂ᵢ[𝔓'] fun ω (i : Iic n) ↦ rewardByCount A R a i ω := by +-- sorry -lemma iIndepFun_rewardByCount' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] - (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) : - iIndepFun (rewardByCount A R a) 𝔓' := by - have hA := h.measurable_A - have hR := h.measurable_R - rw [iIndepFun_nat_iff_forall_indepFun (by fun_prop)] - exact indepFun_rewardByCount_Iic h a +-- lemma iIndepFun_rewardByCount' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) : +-- iIndepFun (rewardByCount A R a) 𝔓' := by +-- have hA := h.measurable_A +-- have hR := h.measurable_R +-- rw [iIndepFun_nat_iff_forall_indepFun (by fun_prop)] +-- exact indepFun_rewardByCount_Iic h a -lemma iIndepFun_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] - (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) : - iIndepFun (fun (p : α × ℕ) ↦ rewardByCount A R p.1 (p.2 + 1)) 𝔓' := by - sorry +-- lemma iIndepFun_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) : +-- iIndepFun (fun (p : α × ℕ) ↦ rewardByCount A R p.1 (p.2 + 1)) 𝔓' := by +-- sorry -lemma identDistrib_rewardByCount_stream_all [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] - (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) : - IdentDistrib (fun ω (p : α × ℕ) ↦ rewardByCount A R p.1 (p.2 + 1) ω) - (fun ω p ↦ ω p.2 p.1) 𝔓' (Bandit.streamMeasure ν) := by - refine IdentDistrib.pi (fun p ↦ ?_) ?_ ?_ - · refine identDistrib_rewardByCount_eval h p.1 (p.2 + 1) p.2 (by simp) (ν := ν) - · exact iIndepFun_rewardByCount h - · sorry +-- lemma identDistrib_rewardByCount_stream_all [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) : +-- IdentDistrib (fun ω (p : α × ℕ) ↦ rewardByCount A R p.1 (p.2 + 1) ω) +-- (fun ω p ↦ ω p.2 p.1) 𝔓' (Bandit.streamMeasure ν) := by +-- refine IdentDistrib.pi (fun p ↦ ?_) ?_ ?_ +-- · refine identDistrib_rewardByCount_eval h p.1 (p.2 + 1) p.2 (by simp) (ν := ν) +-- · exact iIndepFun_rewardByCount h +-- · sorry -lemma identDistrib_rewardByCount_stream' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] - (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) : - IdentDistrib (fun ω n ↦ rewardByCount A R a (n + 1) ω) (fun ω n ↦ ω n a) - 𝔓' (Bandit.streamMeasure ν) := by - refine IdentDistrib.pi (fun n ↦ ?_) ?_ ?_ - · refine identDistrib_rewardByCount_eval h a (n + 1) n (by simp) (ν := ν) - · have h_indep := iIndepFun_rewardByCount' h a - exact iIndepFun.precomp (g := fun n ↦ n + 1) (fun i j hij ↦ by grind) h_indep - · exact iIndepFun_eval_streamMeasure'' ν a +-- lemma identDistrib_rewardByCount_stream' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) : +-- IdentDistrib (fun ω n ↦ rewardByCount A R a (n + 1) ω) (fun ω n ↦ ω n a) +-- 𝔓' (Bandit.streamMeasure ν) := by +-- refine IdentDistrib.pi (fun n ↦ ?_) ?_ ?_ +-- · refine identDistrib_rewardByCount_eval h a (n + 1) n (by simp) (ν := ν) +-- · have h_indep := iIndepFun_rewardByCount' h a +-- exact iIndepFun.precomp (g := fun n ↦ n + 1) (fun i j hij ↦ by grind) h_indep +-- · exact iIndepFun_eval_streamMeasure'' ν a omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in lemma identDistrib_eval_streamMeasure_measure (a : α) : @@ -286,40 +286,40 @@ lemma identDistrib_eval_streamMeasure_measure (a : α) : rw [← Measure.snd, Bandit.snd_measure] exact iIndepFun_eval_streamMeasure'' ν a -lemma identDistrib_rewardByCount_stream [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] - (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) : - IdentDistrib (fun ω n ↦ rewardByCount A R a (n + 1) ω) (fun ω n ↦ ω.2 n a) 𝔓' 𝔓 := - (identDistrib_rewardByCount_stream' h a).trans (identDistrib_eval_streamMeasure_measure a) +-- lemma identDistrib_rewardByCount_stream [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) : +-- IdentDistrib (fun ω n ↦ rewardByCount A R a (n + 1) ω) (fun ω n ↦ ω.2 n a) 𝔓' 𝔓 := +-- (identDistrib_rewardByCount_stream' h a).trans (identDistrib_eval_streamMeasure_measure a) -lemma indepFun_rewardByCount_of_ne [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] - (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {a b : α} (hab : a ≠ b) : - IndepFun (fun ω s ↦ rewardByCount A R a s ω) (fun ω s ↦ rewardByCount A R b s ω) 𝔓' := by - sorry +-- lemma indepFun_rewardByCount_of_ne [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {a b : α} (hab : a ≠ b) : +-- IndepFun (fun ω s ↦ rewardByCount A R a s ω) (fun ω s ↦ rewardByCount A R b s ω) 𝔓' := by +-- sorry -lemma identDistrib_sum_Icc_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] - (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (m : ℕ) (a : α) : - IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount A R a s ω) - (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓' 𝔓 := by - have h1 (a : α) : - IdentDistrib (fun ω s ↦ rewardByCount A R a (s + 1) ω) (fun ω s ↦ ω.2 s a) 𝔓' 𝔓 := - identDistrib_rewardByCount_stream h a - have h_eq (ω : Ω × (ℕ → α → ℝ)) : ∑ s ∈ Icc 1 m, rewardByCount A R a s ω - = ∑ s ∈ range m, rewardByCount A R a (s + 1) ω := by - let e : Icc 1 m ≃ range m := - { toFun x := ⟨x - 1, by have h := x.2; simp only [mem_Icc] at h; simp; grind⟩ - invFun x := ⟨x + 1, by - have h := x.2 - simp only [mem_Icc, le_add_iff_nonneg_left, zero_le, true_and, ge_iff_le] - simp only [mem_range] at h - grind⟩ - left_inv x := by have h := x.2; simp only [mem_Icc] at h; grind - right_inv x := by have h := x.2; grind } - rw [← sum_coe_sort (Icc 1 m), ← sum_coe_sort (range m), sum_equiv e] - · simp - · simp only [univ_eq_attach, mem_attach, forall_const, Subtype.forall, mem_Icc, - forall_and_index] - grind - simp_rw [h_eq] - exact IdentDistrib.comp (h1 a) (u := fun p ↦ ∑ s ∈ range m, p s) (by fun_prop) +-- lemma identDistrib_sum_Icc_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (m : ℕ) (a : α) : +-- IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount A R a s ω) +-- (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓' 𝔓 := by +-- have h1 (a : α) : +-- IdentDistrib (fun ω s ↦ rewardByCount A R a (s + 1) ω) (fun ω s ↦ ω.2 s a) 𝔓' 𝔓 := +-- identDistrib_rewardByCount_stream h a +-- have h_eq (ω : Ω × (ℕ → α → ℝ)) : ∑ s ∈ Icc 1 m, rewardByCount A R a s ω +-- = ∑ s ∈ range m, rewardByCount A R a (s + 1) ω := by +-- let e : Icc 1 m ≃ range m := +-- { toFun x := ⟨x - 1, by have h := x.2; simp only [mem_Icc] at h; simp; grind⟩ +-- invFun x := ⟨x + 1, by +-- have h := x.2 +-- simp only [mem_Icc, le_add_iff_nonneg_left, zero_le, true_and, ge_iff_le] +-- simp only [mem_range] at h +-- grind⟩ +-- left_inv x := by have h := x.2; simp only [mem_Icc] at h; grind +-- right_inv x := by have h := x.2; grind } +-- rw [← sum_coe_sort (Icc 1 m), ← sum_coe_sort (range m), sum_equiv e] +-- · simp +-- · simp only [univ_eq_attach, mem_attach, forall_const, Subtype.forall, mem_Icc, +-- forall_and_index] +-- grind +-- simp_rw [h_eq] +-- exact IdentDistrib.comp (h1 a) (u := fun p ↦ ∑ s ∈ range m, p s) (by fun_prop) end Bandits From 1a991b717547d2169b602bca17dddfb665240356 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 14 Jan 2026 21:50:20 +0100 Subject: [PATCH 25/30] minor cleanup --- LeanBandits/BanditAlgorithms/ETC.lean | 26 -------------------------- LeanBandits/BanditAlgorithms/UCB.lean | 2 -- 2 files changed, 28 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/ETC.lean b/LeanBandits/BanditAlgorithms/ETC.lean index c173d916..14445cd2 100644 --- a/LeanBandits/BanditAlgorithms/ETC.lean +++ b/LeanBandits/BanditAlgorithms/ETC.lean @@ -68,8 +68,6 @@ variable {hK : 0 < K} {m : ℕ} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] {P : Measure Ω} [IsProbabilityMeasure P] {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} -local notation "𝔓" => P.prod (Bandit.streamMeasure ν) - lemma arm_zero [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) : A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ := by @@ -193,30 +191,6 @@ lemma sumRewards_bestArm_le_of_arm_mul_eq [Nonempty (Fin K)] · simp [ha, hm] · simp [h_best, hm] --- lemma identDistrib_aux [Nonempty (Fin K)] --- (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (a b : Fin K) : --- IdentDistrib --- (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount A R a s ω, ∑ s ∈ Icc 1 m, rewardByCount A R b s ω)) --- (fun ω ↦ (∑ s ∈ range m, ω.2 s a, ∑ s ∈ range m, ω.2 s b)) --- 𝔓 (Bandit.measure (etcAlgorithm hK m) ν) := by --- have h2 (a : Fin K) : IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount A R a s ω) --- (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓 (Bandit.measure (etcAlgorithm hK m) ν) := --- identDistrib_sum_Icc_rewardByCount h m a --- by_cases hab : a = b --- · simp only [hab] --- exact (h2 b).comp (u := fun p ↦ (p, p)) (by fun_prop) --- refine (h2 a).prodMk (h2 b) ?_ ?_ --- · suffices IndepFun (fun ω s ↦ rewardByCount A R a s ω) (fun ω s ↦ rewardByCount A R b s ω) --- 𝔓 by --- exact this.comp (φ := fun p ↦ ∑ i ∈ Icc 1 m, p i) (ψ := fun p ↦ ∑ j ∈ Icc 1 m, p j) --- (by fun_prop) (by fun_prop) --- exact indepFun_rewardByCount_of_ne h hab --- · suffices IndepFun (fun ω s ↦ ω.2 s a) (fun ω s ↦ ω.2 s b) --- (Bandit.measure (etcAlgorithm hK m) ν) by --- exact this.comp (φ := fun p ↦ ∑ i ∈ range m, p i) (ψ := fun p ↦ ∑ j ∈ range m, p j) --- (by fun_prop) (by fun_prop) --- exact indepFun_eval_snd_measure _ ν hab - lemma probReal_sumRewards_le_sumRewards_le [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) : diff --git a/LeanBandits/BanditAlgorithms/UCB.lean b/LeanBandits/BanditAlgorithms/UCB.lean index bd300d8a..121bb0e8 100644 --- a/LeanBandits/BanditAlgorithms/UCB.lean +++ b/LeanBandits/BanditAlgorithms/UCB.lean @@ -77,8 +77,6 @@ lemma ucbWidth_eq_ucbWidth' (c : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) (hn : n norm_cast grind -local notation "𝔓" => P.prod (Bandit.streamMeasure ν) - lemma arm_zero [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) : A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ := by From 634743ab5aed417a6e714674a3cfe6634a75d0ae Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 14 Jan 2026 23:09:58 +0100 Subject: [PATCH 26/30] sorry-free --- LeanBandits/Bandit/Bandit.lean | 272 ++++++++++++++++++++++++++--- LeanBandits/Bandit/SumRewards.lean | 8 +- 2 files changed, 249 insertions(+), 31 deletions(-) diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index cd6e2d7f..8c81862f 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -417,7 +417,7 @@ lemma hist_congr (alg : Algorithm α R) (n : ℕ) {ω ω' : probSpace α R} lemma stepsUntil_congr_aux (alg : Algorithm α R) (a : α) (m n : ℕ) {ω ω' : probSpace α R} (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) - (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) + (hω2_eq : ∀ i, i + 1 ≤ m → ω.2 i a = ω'.2 i a) (h_eq : action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m) : action alg (n + 1) ω' = a ∧ pullCount (action alg) a (n + 1) ω' = m := by obtain ⟨h_action, h_pc⟩ := h_eq @@ -442,7 +442,7 @@ lemma stepsUntil_congr_aux (alg : Algorithm α R) lemma stepsUntil_congr (alg : Algorithm α R) (a : α) (m n : ℕ) {ω ω' : probSpace α R} (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) - (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) : + (hω2_eq : ∀ i, i + 1 ≤ m → ω.2 i a = ω'.2 i a) : (action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m) ↔ (action alg (n + 1) ω' = a ∧ pullCount (action alg) a (n + 1) ω' = m) := ⟨stepsUntil_congr_aux alg a m n hω1 hω2_ne hω2_eq, @@ -450,7 +450,7 @@ lemma stepsUntil_congr (alg : Algorithm α R) (a : α) (m n : ℕ) {ω ω' : pro lemma stepsUntil_indicator_congr (alg : Algorithm α R) (a : α) (m n : ℕ) {ω ω' : probSpace α R} (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) - (hω2_eq : ∀ i ≤ m - 1, ω.2 i a = ω'.2 i a) : + (hω2_eq : ∀ i, i + 1 ≤ m → ω.2 i a = ω'.2 i a) : {ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}.indicator (fun _ ↦ 1) ω = {ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}.indicator @@ -660,17 +660,27 @@ time `n`. -/ noncomputable def truePast (alg : Algorithm α R) (a : α) (n : ℕ) (ω : probSpace α R) : probSpace α R := - (ω.1, fun i b ↦ if b = a then ω.2 (min i ((pullCount (action alg) a (n + 1) ω) - 1)) a + (ω.1, fun i b ↦ if b = a then if pullCount (action alg) a (n + 1) ω ≠ 0 then + ω.2 (min i ((pullCount (action alg) a (n + 1) ω) - 1)) a else Nonempty.some inferInstance else ω.2 i b) -omit [Countable α] [StandardBorelSpace R] [Nonempty R] in +omit [Countable α] [StandardBorelSpace R] in lemma truePast_eq_of_pullCount_eq (alg : Algorithm α R) (a : α) (n m : ℕ) (ω : probSpace α R) (h_pc : pullCount (action alg) a (n + 1) ω = m) : - truePast alg a n ω = (ω.1, fun i b ↦ if b = a then ω.2 (min i (m - 1)) a else ω.2 i b) := by + truePast alg a n ω = (ω.1, fun i b ↦ if b = a then if m ≠ 0 then + ω.2 (min i (m - 1)) a else Nonempty.some inferInstance else ω.2 i b) := by simp [truePast, h_pc] -omit [StandardBorelSpace R] [Nonempty R] in +omit [Countable α] [StandardBorelSpace R] in +lemma truePast_eq_of_pullCount_eq_of_ne_zero (alg : Algorithm α R) + (a : α) (n m : ℕ) (ω : probSpace α R) + (h_pc : pullCount (action alg) a (n + 1) ω = m) (hm : m ≠ 0) : + truePast alg a n ω = (ω.1, fun i b ↦ if b = a then + ω.2 (min i (m - 1)) a else ω.2 i b) := by + simp [truePast, h_pc, hm] + +omit [StandardBorelSpace R] in lemma measurable_hist_truePast (alg : Algorithm α R) (a : α) (n : ℕ) : Measurable[MeasurableSpace.comap (truePast alg a n) inferInstance] (hist alg · n) := by @@ -680,14 +690,14 @@ lemma measurable_hist_truePast (alg : Algorithm α R) by_cases hb : b = a · subst hb simp only [truePast, ↓reduceIte] - rw [min_eq_left] + rw [min_eq_left, if_pos (by grind)] grind · simp [truePast, hb] rw [h_eq] refine Measurable.comp ?_ (Measurable.of_comap_le le_rfl) fun_prop -omit [StandardBorelSpace R] [Nonempty R] in +omit [StandardBorelSpace R] in lemma measurable_action_add_one_truePast (alg : Algorithm α R) (a : α) (n : ℕ) : Measurable[MeasurableSpace.comap (truePast alg a n) inferInstance] @@ -702,7 +712,7 @@ lemma measurable_action_add_one_truePast (alg : Algorithm α R) rw [this] exact Measurable.comp (by fun_prop) (Measurable.of_comap_le le_rfl) -omit [StandardBorelSpace R] [Nonempty R] in +omit [StandardBorelSpace R] in lemma measurable_pullCount_add_one_truePast (alg : Algorithm α R) (a : α) (n : ℕ) : Measurable[MeasurableSpace.comap (truePast alg a n) inferInstance] (pullCount (action alg) a (n + 1)) := by @@ -715,38 +725,246 @@ lemma measurable_pullCount_add_one_truePast (alg : Algorithm α R) (a : α) (n : simp_rw [hist_eq _ _ n, @measurable_pi_iff] at h_meas exact (h_meas ⟨i, by grind⟩).fst -lemma indepFun_snd_apply_aux (alg : Algorithm α R) - (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m n : ℕ) : +omit [Nonempty α] [StandardBorelSpace α] [Countable α] [StandardBorelSpace R] in +lemma indepFun_snd_apply_aux (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m : ℕ) : (fun ω ↦ ω.2 m a) ⟂ᵢ[arrayMeasure ν] - (fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b)) := by - sorry - -omit [StandardBorelSpace R] [Nonempty R] in + (fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b)) := by + unfold arrayMeasure + let μ₁ : Measure (ℕ → I) := Measure.infinitePi fun _ ↦ volume + let μ₂ : Measure (ℕ → α → R) := Measure.infinitePi fun _ ↦ Measure.infinitePi ν + -- Independence within μ₂: coordinates ω i are independent + have h_indep₂ : iIndepFun (fun i (ω : ℕ → α → R) ↦ ω i) μ₂ := + iIndepFun_infinitePi (fun _ ↦ measurable_id) + -- Independence within each infinitePi ν: coordinates f b are independent + have h_indep_inner : iIndepFun (fun (b : α) (f : α → R) ↦ f b) (Measure.infinitePi ν) := + iIndepFun_infinitePi (fun _ ↦ measurable_id) + rw [indepFun_iff_measure_inter_preimage_eq_mul] + intro s t hs ht + let X : (ℕ → I) × (ℕ → α → R) → R := fun ω ↦ ω.2 m a + let Y : (ℕ → I) × (ℕ → α → R) → (ℕ → I) × (ℕ → α → R) := + fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b) + have hX_meas : Measurable X := + (measurable_pi_apply a).comp ((measurable_pi_apply m).comp measurable_snd) + have hY_meas : Measurable Y := by + change Measurable (fun ω : (ℕ → I) × (ℕ → α → R) ↦ + (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b)) + refine Measurable.prod measurable_fst ?_ + refine measurable_pi_lambda _ (fun k ↦ ?_) + refine measurable_pi_lambda _ (fun b ↦ ?_) + by_cases hb : b = a + · simp only [hb, ↓reduceIte] + by_cases hm : m ≠ 0 + · simp only [ne_eq, hm, not_false_eq_true, ↓reduceIte] + exact (measurable_pi_apply a).comp + ((measurable_pi_apply (min k (m - 1))).comp measurable_snd) + · simp only [hm, ↓reduceIte] + exact measurable_const + · simp only [hb, ↓reduceIte] + exact (measurable_pi_apply b).comp ((measurable_pi_apply k).comp measurable_snd) + change (μ₁.prod μ₂) (X ⁻¹' s ∩ Y ⁻¹' t) = (μ₁.prod μ₂) (X ⁻¹' s) * (μ₁.prod μ₂) (Y ⁻¹' t) + -- Use Fubini on μ₁.prod μ₂ + rw [Measure.prod_apply (hs.preimage hX_meas), + Measure.prod_apply (ht.preimage hY_meas), + Measure.prod_apply ((hs.preimage hX_meas).inter (ht.preimage hY_meas))] + -- X only depends on ω₂, so its fiber is constant in ω₁ + have hX_fst : ∀ ω₁, μ₂ (Prod.mk ω₁ ⁻¹' (X ⁻¹' s)) = μ₂ ((fun ω₂ ↦ ω₂ m a) ⁻¹' s) := fun _ ↦ rfl + simp_rw [hX_fst] + -- The LHS integral: fiber of X ∩ Y at ω₁ + -- Key: X depends only on ω₂ m a, while Y's dependence on ω₂ avoids (m, a) + -- Define the "truncation" map on ω₂ + let trunc : (ℕ → α → R) → (ℕ → α → R) := + fun ω₂ k b ↦ if b = a then if m ≠ 0 then ω₂ (min k (m - 1)) b + else Nonempty.some inferInstance else ω₂ k b + -- The fiber of Y at ω₁ only depends on trunc(ω₂) + have hY_fiber : ∀ ω₁, Prod.mk ω₁ ⁻¹' (Y ⁻¹' t) = (fun ω₂ ↦ (ω₁, trunc ω₂)) ⁻¹' t := fun _ ↦ rfl + -- The fiber of X ∩ Y factors + have hXY_fiber : ∀ ω₁, Prod.mk ω₁ ⁻¹' (X ⁻¹' s ∩ Y ⁻¹' t) = + ((fun ω₂ ↦ ω₂ m a) ⁻¹' s) ∩ ((fun ω₂ ↦ (ω₁, trunc ω₂)) ⁻¹' t) := fun _ ↦ rfl + simp_rw [hXY_fiber, hY_fiber] + -- Now we use independence in μ₂: (ω₂ m a) is independent of (trunc ω₂) + -- because trunc only uses indices (k, a) with k < m, and (k, b) with b ≠ a + have h_trunc_meas : Measurable trunc := by + refine measurable_pi_lambda _ (fun k ↦ ?_) + refine measurable_pi_lambda _ (fun b ↦ ?_) + simp only [trunc] + by_cases hb : b = a + · simp only [hb, ↓reduceIte] + by_cases hm : m = 0 + · simp only [hm] + exact measurable_const + · simp only [ne_eq, hm, not_false_eq_true, ↓reduceIte] + exact (measurable_pi_apply a).comp (measurable_pi_apply (min k (m - 1))) + · simp only [hb, ↓reduceIte] + exact (measurable_pi_apply b).comp (measurable_pi_apply k) + -- Key independence: (ω₂ m a) ⟂ trunc because trunc only uses coordinates ≠ (m, a) + have h_indep_trunc : IndepFun (fun ω₂ ↦ ω₂ m a) trunc μ₂ := by + -- Factor trunc through proj which extracts the relevant coordinates + let proj : (ℕ → α → R) → ((ℕ → R) × (ℕ → {b : α // b ≠ a} → R)) := fun ω₂ ↦ + (fun k ↦ if m ≠ 0 then ω₂ (min k (m - 1)) a else Nonempty.some inferInstance, + fun k ⟨b, _⟩ ↦ ω₂ k b) + have h_trunc_proj : ∀ ω₂, trunc ω₂ = (fun p k b ↦ + if h : b = a then if m ≠ 0 then p.1 k + else Nonempty.some inferInstance else p.2 k ⟨b, h⟩) (proj ω₂) := by + intro ω₂; ext k b; simp only [trunc, proj]; by_cases hb : b = a <;> simp [hb]; grind + have h_proj_meas : Measurable proj := by + refine Measurable.prod ?_ ?_ + · refine measurable_pi_lambda _ fun k ↦ ?_ + by_cases hm : m ≠ 0 + · simp only [proj, ne_eq, hm, not_false_eq_true, ↓reduceIte] + exact (measurable_pi_apply a).comp (measurable_pi_apply (min k (m - 1))) + · simp [proj, hm] + · exact measurable_pi_lambda _ (fun k ↦ measurable_pi_lambda _ (fun ⟨b, _⟩ ↦ + (measurable_pi_apply b).comp (measurable_pi_apply k))) + have h_g_meas : Measurable (fun p : (ℕ → R) × (ℕ → {b : α // b ≠ a} → R) ↦ + (fun k b ↦ if h : b = a then if m ≠ 0 then p.1 k else Nonempty.some inferInstance + else p.2 k ⟨b, h⟩)) := by + refine measurable_pi_lambda _ (fun k ↦ measurable_pi_lambda _ (fun b ↦ ?_)) + by_cases hb : b = a + · simp only [hb, ↓reduceDIte] + by_cases hm : m ≠ 0 + · simp only [ne_eq, hm, not_false_eq_true] + exact (measurable_pi_apply k).comp measurable_fst + · simp [hm] + · simp only [hb, ↓reduceDIte] + exact (measurable_pi_apply (⟨b, hb⟩ : {b : α // b ≠ a})).comp + ((measurable_pi_apply k).comp measurable_snd) + -- Show (ω₂ m a) ⟂ proj: proj uses coordinates disjoint from (m, a) + have h_indep_proj : IndepFun (fun ω₂ ↦ ω₂ m a) proj μ₂ := by + have h_row_bound (hm : m ≠ 0) : ∀ k, min k (m - 1) < m := by + intro k + calc min k (m - 1) ≤ m - 1 := Nat.min_le_right k (m - 1) + _ < m := Nat.sub_lt (by grind) Nat.one_pos + rw [indepFun_iff_measure_inter_preimage_eq_mul] + intro s t' hs ht' + -- rows_lt_m extracts column a at rows < m, other_cols extracts columns ≠ a + let rows_lt_m : (ℕ → α → R) → (Iio m → R) := fun ω₂ ⟨j, _⟩ ↦ ω₂ j a + let other_cols : (ℕ → α → R) → (ℕ → {b : α // b ≠ a} → R) := fun ω₂ k ⟨b, _⟩ ↦ ω₂ k b + have h_proj_factor : ∀ ω₂, proj ω₂ = + ((fun r k ↦ if hm : m ≠ 0 then r ⟨min k (m - 1), Finset.mem_Iio.mpr (h_row_bound hm k)⟩ + else Nonempty.some inferInstance) (rows_lt_m ω₂), + other_cols ω₂) := by + intro ω₂; ext1 + · ext k + by_cases hm : m ≠ 0 + · simp [proj, rows_lt_m, hm] + · simp [proj, hm] + · rfl + -- Use iIndepFun structure of the doubly-indexed infinite product + have h_iindep : iIndepFun (fun (p : ℕ × α) ω ↦ ω p.1 p.2) μ₂ := + iIndepFun_uncurry_infinitePi' (X := fun _ _ ↦ id) (fun _ ↦ ν) (by fun_prop) + have h_rows_meas : Measurable rows_lt_m := + measurable_pi_lambda _ (fun ⟨j, _⟩ ↦ (measurable_pi_apply a).comp (measurable_pi_apply j)) + have h_other_meas : Measurable other_cols := + measurable_pi_lambda _ (fun k ↦ measurable_pi_lambda _ (fun ⟨b, _⟩ ↦ + (measurable_pi_apply b).comp (measurable_pi_apply k))) + -- Show (ω₂ m a) ⟂ (rows_lt_m, other_cols) via indep_iSup_of_disjoint + have h_indep_combined : IndepFun (fun ω₂ ↦ ω₂ m a) + (fun ω₂ ↦ (rows_lt_m ω₂, other_cols ω₂)) μ₂ := by + rw [IndepFun_iff_Indep] + have h_comap_le : (MeasurableSpace.pi.prod MeasurableSpace.pi).comap + (fun ω₂ ↦ (rows_lt_m ω₂, other_cols ω₂)) ≤ + ⨆ (p : {p : ℕ × α // p ≠ (m, a)}), mR.comap (fun ω ↦ ω p.val.1 p.val.2) := by + rw [MeasurableSpace.comap_prodMk] + refine sup_le ?_ ?_ + · rw [MeasurableSpace.comap_pi] + refine iSup_le (fun ⟨j, hj⟩ ↦ ?_) + have h_ne : (j, a) ≠ (m, a) := fun h ↦ (Finset.mem_Iio.mp hj).ne (Prod.mk.inj h).1 + exact le_iSup_of_le ⟨(j, a), h_ne⟩ le_rfl + · rw [MeasurableSpace.comap_pi] + refine iSup_le (fun k ↦ ?_) + rw [MeasurableSpace.comap_pi] + refine iSup_le (fun ⟨b, hb⟩ ↦ ?_) + have h_ne : (k, b) ≠ (m, a) := fun h ↦ hb (Prod.mk.inj h).2 + exact le_iSup_of_le ⟨(k, b), h_ne⟩ le_rfl + refine indep_of_indep_of_le_right ?_ h_comap_le + have h_disjoint : Disjoint ({(m, a)} : Set (ℕ × α)) {p | p ≠ (m, a)} := by simp + have h_le : ∀ p : ℕ × α, mR.comap (fun ω : ℕ → α → R ↦ ω p.1 p.2) ≤ + MeasurableSpace.pi (m := fun _ ↦ MeasurableSpace.pi) := fun p ↦ + Measurable.comap_le ((measurable_pi_apply p.2).comp (measurable_pi_apply p.1)) + have h_iindep' : iIndep (fun p : ℕ × α ↦ mR.comap (fun ω : ℕ → α → R ↦ ω p.1 p.2)) μ₂ := + h_iindep.iIndep + have h_indep := indep_iSup_of_disjoint h_le h_iindep' h_disjoint + convert h_indep using 2 + · simp only [Set.mem_singleton_iff, iSup_iSup_eq_left] + · simp only [ne_eq, Set.mem_setOf_eq, iSup_subtype'] + have h_proj_preimage : proj ⁻¹' t' = (fun ω₂ ↦ (rows_lt_m ω₂, other_cols ω₂)) ⁻¹' + {p | ((fun r k ↦ if hm : m ≠ 0 then + r ⟨min k (m - 1), Finset.mem_Iio.mpr (h_row_bound hm k)⟩ + else Nonempty.some inferInstance) p.1, p.2) ∈ t'} + := by ext ω₂; simp only [Set.mem_preimage, Set.mem_setOf_eq, h_proj_factor] + rw [indepFun_iff_measure_inter_preimage_eq_mul] at h_indep_combined + rw [h_proj_preimage] + let T : Set ((Iio m → R) × (ℕ → {b : α // b ≠ a} → R)) := + {p | ((fun r k ↦ if hm : m ≠ 0 then + r ⟨min k (m - 1), Finset.mem_Iio.mpr (h_row_bound hm k)⟩ + else Nonempty.some inferInstance) p.1, p.2) ∈ t'} + have hT_meas : MeasurableSet T := by + refine ht'.preimage (Measurable.prod ?_ measurable_snd) + refine measurable_pi_lambda _ (fun k ↦ ?_) + by_cases hm : m ≠ 0 + · simp only [ne_eq, hm, not_false_eq_true, ↓reduceDIte] + exact (measurable_pi_apply (⟨min k (m - 1), Finset.mem_Iio.mpr (h_row_bound hm k)⟩ : + Iio m)).comp measurable_fst + · simp [hm] + change μ₂ ((fun ω₂ ↦ ω₂ m a) ⁻¹' s ∩ (fun ω₂ ↦ (rows_lt_m ω₂, other_cols ω₂)) ⁻¹' T) = + μ₂ ((fun ω₂ ↦ ω₂ m a) ⁻¹' s) * + μ₂ ((fun ω₂ ↦ (rows_lt_m ω₂, other_cols ω₂)) ⁻¹' T) + exact h_indep_combined s T hs hT_meas + have h_eq : trunc = (fun p k b ↦ if h : b = a then if m ≠ 0 then p.1 k + else Nonempty.some inferInstance else p.2 k ⟨b, h⟩) ∘ proj := by + funext ω₂; exact h_trunc_proj ω₂ + rw [h_eq] + exact h_indep_proj.comp measurable_id h_g_meas + rw [indepFun_iff_measure_inter_preimage_eq_mul] at h_indep_trunc + have h_const : ∀ ω₁, μ₂ (((fun ω₂ ↦ ω₂ m a) ⁻¹' s) ∩ ((fun ω₂ ↦ (ω₁, trunc ω₂)) ⁻¹' t)) = + μ₂ ((fun ω₂ ↦ ω₂ m a) ⁻¹' s) * μ₂ ((fun ω₂ ↦ (ω₁, trunc ω₂)) ⁻¹' t) := fun ω₁ ↦ + h_indep_trunc s _ hs (ht.preimage (by fun_prop)) + simp_rw [h_const] + let c := μ₂ ((fun ω₂ ↦ ω₂ m a) ⁻¹' s) + change ∫⁻ x, c * μ₂ ((fun ω₂ ↦ (x, trunc ω₂)) ⁻¹' t) ∂μ₁ = + (∫⁻ _, c ∂μ₁) * ∫⁻ x, μ₂ ((fun ω₂ ↦ (x, trunc ω₂)) ⁻¹' t) ∂μ₁ + have h_preimage : ∀ x, (fun ω₂ ↦ (x, trunc ω₂)) ⁻¹' t = trunc ⁻¹' (Prod.mk x ⁻¹' t) := fun _ ↦ rfl + simp_rw [h_preimage] + have h_map : ∀ x, μ₂ (trunc ⁻¹' (Prod.mk x ⁻¹' t)) = (μ₂.map trunc) (Prod.mk x ⁻¹' t) := by + intro x; rw [Measure.map_apply h_trunc_meas (ht.preimage (by fun_prop))] + simp_rw [h_map] + rw [lintegral_const_mul _ (measurable_measure_prodMk_left_finite ht), + lintegral_const, measure_univ, mul_one] + + +omit [StandardBorelSpace R] in lemma measurable_stepsUntil (alg : Algorithm α R) (a : α) (m n : ℕ) : Measurable[MeasurableSpace.comap - (fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b)) inferInstance] + (fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b)) inferInstance] (({ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}).indicator (fun _ ↦ 1)) := by let f := ({ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}).indicator (fun _ ↦ 1) have h_eq : f = f ∘ - fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b) := by + fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b) := by ext ω exact stepsUntil_indicator_congr alg a m n (by grind) (by grind) (by grind) change Measurable[MeasurableSpace.comap - (fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b)) inferInstance] f + (fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b)) inferInstance] f rw [h_eq] refine Measurable.comp ?_ (Measurable.of_comap_le le_rfl) refine Measurable.indicator (by fun_prop) ?_ exact MeasurableSet.inter ((measurableSet_singleton _).preimage (by fun_prop)) ((measurableSet_singleton _).preimage (by fun_prop)) +omit [StandardBorelSpace R] in lemma indepFun_snd_apply_pullCount_action (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m n : ℕ) : (fun ω ↦ ω.2 m a) ⟂ᵢ[arrayMeasure ν] ({ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}).indicator (fun _ ↦ 1) := - (indepFun_snd_apply_aux alg ν a m n).of_measurable_right (measurable_stepsUntil alg a m n) + (indepFun_snd_apply_aux ν a m).of_measurable_right (measurable_stepsUntil alg a m n) omit [StandardBorelSpace R] [Nonempty R] in @[fun_prop] @@ -848,7 +1066,8 @@ lemma indepFun_snd_hist_cond (alg : Algorithm α R) refine IndepFun.of_measurable_right ?_ h_meas have h_ae_eq : truePast alg a n =ᵐ[(arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}]] - (fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b)) := by + (fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b)) := by refine ae_cond_of_forall_mem ?_ fun x hx ↦ ?_ · refine (measurableSet_singleton _).preimage ?_ have h_meas_pc : Measurable fun ω ↦ @@ -868,7 +1087,8 @@ lemma indepFun_snd_hist_cond (alg : Algorithm α R) refine IndepFun.congr ?_ EventuallyEq.rfl h_ae_eq.symm suffices (fun ω ↦ ω.2 m a) ⟂ᵢ[(arrayMeasure ν)[|(({ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}).indicator (fun _ ↦ 1)) ⁻¹' {1}]] - fun ω ↦ (ω.1, fun k b ↦ if b = a then ω.2 (min k (m - 1)) b else ω.2 k b) by + fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b) by convert this ext ω simp only [Set.mem_preimage, Set.mem_singleton_iff, Prod.mk.injEq, Set.indicator_apply, @@ -880,12 +1100,12 @@ lemma indepFun_snd_hist_cond (alg : Algorithm α R) obtain ⟨f, hf, hf_eq⟩ := h_meas.exists_eq_measurable_comp simp_rw [hf_eq] refine indepFun_todo (Z := f) (z := 1) ?_ ?_ hf - · exact indepFun_snd_apply_aux alg ν a m n + · exact indepFun_snd_apply_aux ν a m · refine Measurable.prodMk (by fun_prop) ?_ simp_rw [measurable_pi_iff] - intro m a - refine Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact MeasurableSet.const _ + intro i b + refine Measurable.ite (MeasurableSet.const _) ?_ (by fun_prop) + refine Measurable.ite (MeasurableSet.const _) (by fun_prop) (by fun_prop) /-- The conditional distribution of the reward at time `n + 1`, given the history up to time `n`, the action at time `n + 1`, and the number of times that action has been pulled before time `n + 1`, diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 7ac161c4..65ecf7e8 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -306,22 +306,20 @@ variable {α Ω Ω' : Type*} [DecidableEq α] {mα : MeasurableSpace α} {mΩ : {A : ℕ → Ω → α} {R : ℕ → Ω → ℝ} {A₂ : ℕ → Ω' → α} {R₂ : ℕ → Ω' → ℝ} {ω : Ω} {m n t : ℕ} {a : α} -variable [StandardBorelSpace α] [Nonempty α] - -omit [Nonempty α] in lemma sumRewards_eq_comp : sumRewards A R a n = (fun p ↦ ∑ i ∈ range n, if (p i).1 = a then (p i).2 else 0) ∘ (fun ω n ↦ (A n ω, R n ω)) := by ext simp [sumRewards] -omit [Nonempty α] in lemma pullCount_eq_comp : pullCount A a n = (fun p ↦ ∑ i ∈ range n, if (p i).1 = a then 1 else 0) ∘ (fun ω n ↦ (A n ω, R n ω)) := by ext simp [pullCount] +variable [StandardBorelSpace α] [Nonempty α] + -- todo: write those lemmas with IdentDistrib instead of equality of maps lemma _root_.Learning.IsAlgEnvSeq.law_sumRewards_unique (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) @@ -467,7 +465,7 @@ lemma prob_pullCount_mem_and_sumRewards_mem_le [Countable α] exists_eq_right, mem_filter, mem_range] at hk simp [hk.2.1] -lemma todo [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) +lemma prob_sumRewards_mem_le [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {B : Set ℝ} (hB : MeasurableSet B) : P (sumRewards A R a n ⁻¹' B) ≤ ∑ k ∈ range (n + 1), Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := by From 15f361dd7588d57c7cc30c3a92d6b792d1f16dff Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 15 Jan 2026 07:55:11 +0100 Subject: [PATCH 27/30] fix blueprint refs --- blueprint/lean_decls | 3 --- blueprint/src/chapters/bandit.tex | 10 ++-------- blueprint/src/macros/print.tex | 2 +- 3 files changed, 3 insertions(+), 12 deletions(-) diff --git a/blueprint/lean_decls b/blueprint/lean_decls index 671d8fec..09abdc59 100644 --- a/blueprint/lean_decls +++ b/blueprint/lean_decls @@ -50,9 +50,6 @@ Bandits.reward_cond_stepsUntil ProbabilityTheory.condDistrib_ae_eq_cond Bandits.condDistrib_rewardByCount_stepsUntil Bandits.hasLaw_rewardByCount -Bandits.iIndepFun_rewardByCount' -Bandits.identDistrib_rewardByCount_stream -Bandits.identDistrib_sum_Icc_rewardByCount Bandits.regret Bandits.gap Learning.sum_pullCount_mul diff --git a/blueprint/src/chapters/bandit.tex b/blueprint/src/chapters/bandit.tex index e504f4ca..ce9f3c09 100644 --- a/blueprint/src/chapters/bandit.tex +++ b/blueprint/src/chapters/bandit.tex @@ -205,8 +205,6 @@ \section{Alternative models: rewards indexed by time or pull count}\label{sec:al \begin{lemma}\label{lem:iIndepFun_rewardByCount} \uses{def:rewardByCount} - \leanok - \lean{Bandits.iIndepFun_rewardByCount'} The rewards $(Y_{n,a})_{n \in \mathbb{N}}$ are independent. \end{lemma} @@ -235,12 +233,10 @@ \section{Alternative models: rewards indexed by time or pull count}\label{sec:al \begin{lemma}\label{lem:identDistrib_rewardByCount_stream} \uses{def:rewardByCount} - \leanok - \lean{Bandits.identDistrib_rewardByCount_stream} The random sequences $(Y_{n+1,a})_{n \in \mathbb{N}}$ and $(Z_{n,a})_{n \in \mathbb{N}}$ are identically distributed. \end{lemma} -\begin{proof}\leanok +\begin{proof} \uses{lem:hasLaw_rewardByCount, lem:iIndepFun_rewardByCount} \end{proof} @@ -248,12 +244,10 @@ \section{Alternative models: rewards indexed by time or pull count}\label{sec:al \begin{lemma}\label{lem:identDistrib_sum_Icc_rewardByCount} \uses{def:rewardByCount} - \leanok - \lean{Bandits.identDistrib_sum_Icc_rewardByCount} The random variables $\sum_{i=1}^n Y_{i,a}$ and $\sum_{i=0}^{n-1} Z_{i,a}$ are identically distributed. \end{lemma} -\begin{proof}\leanok +\begin{proof} \uses{lem:identDistrib_rewardByCount_stream} Immediate consequence of Lemma~\ref{lem:identDistrib_rewardByCount_stream}. \end{proof} diff --git a/blueprint/src/macros/print.tex b/blueprint/src/macros/print.tex index 668fec13..78708343 100644 --- a/blueprint/src/macros/print.tex +++ b/blueprint/src/macros/print.tex @@ -26,4 +26,4 @@ \NewDocumentCommand{\proves}{m} {\clist_map_inline:nn{#1}{\vphantom{\ref{##1}}}% \ignorespaces} -\ExplSyntaxOff \ No newline at end of file +\ExplSyntaxOff From 75a3b702d0cea4d0f20dc4b6008a2dd4eab33d80 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 15 Jan 2026 09:14:34 +0100 Subject: [PATCH 28/30] add and adapt Paulo's regret lemmas --- LeanBandits/Bandit/Regret.lean | 74 ++++++++++++++++++++++++++++------ 1 file changed, 62 insertions(+), 12 deletions(-) diff --git a/LeanBandits/Bandit/Regret.lean b/LeanBandits/Bandit/Regret.lean index 89adb008..fdbfb8da 100644 --- a/LeanBandits/Bandit/Regret.lean +++ b/LeanBandits/Bandit/Regret.lean @@ -6,7 +6,7 @@ Authors: Rémy Degenne, Paulo Rauber import LeanBandits.SequentialLearning.FiniteActions /-! -# Regret +# Regret, gap, best arm -/ @@ -21,13 +21,6 @@ variable {α Ω : Type*} [DecidableEq α] {mα : MeasurableSpace α} {mΩ : Meas {A : ℕ → Ω → α} {R : ℕ → Ω → ℝ} {ω : Ω} {m n t : ℕ} {a : α} -/-! ### Definitions of regret, gaps, pull counts -/ - -/-- Regret of a sequence of pulls `k : ℕ → α` at time `t` for the reward kernel `ν ; Kernel α ℝ`. -/ -noncomputable -def regret (ν : Kernel α ℝ) (A : ℕ → Ω → α) (t : ℕ) (ω : Ω) : ℝ := - t * (⨆ a, (ν a)[id]) - ∑ s ∈ range t, (ν (A s ω))[id] - /-- Gap of an action `a`: difference between the highest mean of the actions and the mean of `a`. -/ noncomputable def gap (ν : Kernel α ℝ) (a : α) : ℝ := (⨆ i, (ν i)[id]) - (ν a)[id] @@ -37,13 +30,29 @@ lemma gap_nonneg [Fintype α] : 0 ≤ gap ν a := by rw [gap, sub_nonneg] exact le_ciSup (f := fun i ↦ (ν i)[id]) (by simp) a -section RewardByCount +/-- Regret of a sequence of pulls `k : ℕ → α` at time `t` for the reward kernel `ν ; Kernel α ℝ`. -/ +noncomputable +def regret (ν : Kernel α ℝ) (A : ℕ → Ω → α) (t : ℕ) (ω : Ω) : ℝ := + t * (⨆ a, (ν a)[id]) - ∑ s ∈ range t, (ν (A s ω))[id] + +omit [DecidableEq α] in +lemma regret_eq_sum_gap : regret ν A t ω = ∑ s ∈ range t, gap ν (A s ω) := by + simp [regret, gap] + +omit [DecidableEq α] in +lemma regret_nonneg [Fintype α] : 0 ≤ regret ν A t ω := by + rw [regret_eq_sum_gap] + exact sum_nonneg (fun _ _ ↦ gap_nonneg) + +omit [DecidableEq α] in +lemma gap_eq_zero_of_regret_eq_zero [Fintype α] (hr : regret ν A t ω = 0) {s : ℕ} (hs : s < t) : + gap ν (A s ω) = 0 := by + rw [regret_eq_sum_gap] at hr + exact (sum_eq_zero_iff_of_nonneg fun _ _ ↦ gap_nonneg).1 hr s (mem_range.2 hs) lemma regret_eq_sum_pullCount_mul_gap [Fintype α] : regret ν A t ω = ∑ a, pullCount A a t ω * gap ν a := by - simp [sum_pullCount_mul, regret, gap, sum_sub_distrib] - -end RewardByCount + simp_rw [regret_eq_sum_gap, sum_pullCount_mul] section bestArm @@ -70,6 +79,47 @@ omit [DecidableEq α] in lemma gap_bestArm : gap ν (bestArm ν) = 0 := by rw [gap_eq_bestArm_sub, sub_self] +omit [DecidableEq α] in +lemma integral_eq_of_gap_eq_zero (hg : gap ν a = 0) : (ν (bestArm ν))[id] = (ν a)[id] := by + rwa [← sub_eq_zero, ← gap_eq_bestArm_sub] + end bestArm +section Asymptotics + +omit [DecidableEq α] in +/-- If the regret is sublinear, the average mean reward tends to the highest mean of the arms. -/ +lemma avg_mean_reward_tendsto_of_sublinear_regret + (hr : (regret ν A · ω) =o[atTop] fun t ↦ (t : ℝ)) : + Tendsto (fun t ↦ (∑ s ∈ range t, (ν (A s ω))[id]) / (t : ℝ)) + atTop (nhds (⨆ a, (ν a)[id])) := by + have ht : Tendsto (fun t ↦ (⨆ a, (ν a)[id]) - regret ν A t ω / t) + atTop (nhds (⨆ a, (ν a)[id])) := by + simpa using tendsto_const_nhds.sub hr.tendsto_div_nhds_zero + apply ht.congr' + filter_upwards [eventually_ne_atTop 0] with t ht + rw [regret] + field_simp + ring + +/-- If the regret is sublinear, the rate of suboptimal arm pulls tends to zero. -/ +lemma pullCount_rate_tendsto_of_sublinear_regret [Fintype α] + (hr : (regret ν A · ω) =o[atTop] fun t ↦ (t : ℝ)) (hg : 0 < gap ν a) : + Tendsto (fun t ↦ (pullCount A a t ω : ℝ) / t) atTop (nhds 0) := by + have hb (t : ℕ) : (pullCount A a t ω : ℝ) * gap ν a ≤ regret ν A t ω := by + rw [regret_eq_sum_pullCount_mul_gap] + exact single_le_sum (f := fun a ↦ pullCount A a t ω * gap ν a) + (fun _ _ ↦ mul_nonneg (Nat.cast_nonneg _) gap_nonneg) (mem_univ a) + have hb' (t : ℕ) : (pullCount A a t ω : ℝ) / t ≤ regret ν A t ω / t / gap ν a := by + obtain ht | ht := eq_or_ne t 0 + · simp [ht] + · calc (pullCount A a t ω : ℝ) / t + = pullCount A a t ω * gap ν a / gap ν a / t := by field_simp + _ ≤ regret ν A t ω / gap ν a / t := by gcongr; exact hb t + _ = regret ν A t ω / t / gap ν a := by ring + apply squeeze_zero' (Eventually.of_forall fun _ ↦ by positivity) (Eventually.of_forall hb') + simpa using hr.tendsto_div_nhds_zero.div_const (gap ν a) + +end Asymptotics + end Bandits From 7f3107863086cd4a63b83ffa8d3f3313ff454602 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 15 Jan 2026 11:17:57 +0100 Subject: [PATCH 29/30] partial blueprint update --- blueprint/lean_decls | 26 +++- blueprint/src/biblio.bib | 7 + blueprint/src/chapters/algorithm.tex | 223 +++++++++++++++++++-------- blueprint/src/chapters/bandit.tex | 48 ++++++ blueprint/src/print.tex | 2 +- blueprint/src/web.tex | 2 +- 6 files changed, 232 insertions(+), 76 deletions(-) diff --git a/blueprint/lean_decls b/blueprint/lean_decls index 09abdc59..54612c19 100644 --- a/blueprint/lean_decls +++ b/blueprint/lean_decls @@ -2,6 +2,16 @@ Learning.Algorithm Learning.Environment Learning.detAlgorithm Learning.stationaryEnv +Learning.IsAlgEnvSeq +Learning.IsAlgEnvSeq.hist +Learning.IsAlgEnvSeq.step +Learning.IsAlgEnvSeq.hasLaw_step_zero +Learning.IsAlgEnvSeq.hasCondDistrib_step +Learning.IsAlgEnvSeq.filtration +Learning.IsAlgEnvSeq.filtrationAction +Learning.isAlgEnvSeq_unique +Learning.IsAlgEnvSeq.condDistrib_reward_stationaryEnv +Learning.IsAlgEnvSeq.condIndepFun_reward_hist_action ProbabilityTheory.Kernel.traj ProbabilityTheory.Kernel.trajMeasure Learning.IT.step @@ -15,12 +25,11 @@ Learning.IT.action Learning.IT.reward Learning.IT.adapted_action Learning.IT.adapted_reward -Learning.IsAlgEnvSeq.hasCondDistrib_action -Learning.IsAlgEnvSeq.hasCondDistrib_reward -Learning.IsAlgEnvSeq.hasLaw_action_zero -Learning.IsAlgEnvSeq.hasCondDistrib_reward_zero -Learning.IsAlgEnvSeq.condDistrib_reward_stationaryEnv -Learning.IsAlgEnvSeq.condIndepFun_reward_hist_action +Learning.IT.condDistrib_action +Learning.IT.condDistrib_reward +Learning.IT.hasLaw_action_zero +Learning.IT.condDistrib_reward_zero +Learning.IT.isAlgEnvSeq_trajMeasure Learning.pullCount Learning.pullCount_zero Learning.pullCount_mono @@ -44,6 +53,11 @@ Learning.empMean Learning.sum_rewardByCount_eq_sumRewards Bandits.Bandit.trajMeasure Bandits.Bandit.measure +Bandits.ArrayModel.probSpace +Bandits.ArrayModel.arrayMeasure +Bandits.ArrayModel.algFunction +Bandits.ArrayModel.initAlgFunction +Bandits.ArrayModel.isAlgEnvSeq_arrayMeasure Learning.measurable_comap_indicator_stepsUntil_eq Bandits.condIndepFun_reward_stepsUntil_action Bandits.reward_cond_stepsUntil diff --git a/blueprint/src/biblio.bib b/blueprint/src/biblio.bib index 82ef848f..f6e1818e 100644 --- a/blueprint/src/biblio.bib +++ b/blueprint/src/biblio.bib @@ -269,3 +269,10 @@ @inproceedings{hirata2023semantic year={2023}, organization={Schloss Dagstuhl--Leibniz-Zentrum f{\"u}r Informatik} } + +@book{lattimore2020bandit, + title={Bandit algorithms}, + author={Lattimore, Tor and Szepesv{\'a}ri, Csaba}, + year={2020}, + publisher={Cambridge University Press} +} diff --git a/blueprint/src/chapters/algorithm.tex b/blueprint/src/chapters/algorithm.tex index dbbf6e74..044be591 100644 --- a/blueprint/src/chapters/algorithm.tex +++ b/blueprint/src/chapters/algorithm.tex @@ -1,12 +1,19 @@ \chapter{Iterative stochastic algorithms} +Warning: all times start at zero. + +TODO: notations + +All measurable spaces are assumed to be standard Borel. + + \begin{definition}[Algorithm]\label{def:algorithm} \leanok \lean{Learning.Algorithm} A sequential, stochastic algorithm with actions in a measurable space $\mathcal{A}$ and observations in a measurable space $\mathcal{R}$ is described by the following data: \begin{itemize} - \item for all $t \in \mathbb{N}$, a policy $\pi_t : (\mathcal{A} \times \mathcal{R})^{t+1} \rightsquigarrow \mathcal{A}$, a Markov kernel which gives the distribution of the action of the algorithm at time $t+1$ given the history of previous pulls and observations, + \item for all $t \in \mathbb{N}$, a policy $\pi_t : (\mathcal{A} \times \mathcal{R})^{t+1} \rightsquigarrow \mathcal{A}$, a Markov kernel which gives the distribution of the action of the algorithm at time $t+1$ given the history of previous actions and observations, \item $P_0 \in \mathcal{P}(\mathcal{A})$, a probability measure that gives the distribution of the first action. \end{itemize} \end{definition} @@ -40,10 +47,7 @@ \chapter{Iterative stochastic algorithms} An environment is stationary if there exists a Markov kernel $\nu : \mathcal{A} \rightsquigarrow \mathcal{R}$ such that $\nu'_0 = \nu$ and for all $t \in \mathbb{N}$, for all $h_t \in (\mathcal{A} \times \mathcal{R})^{t+1}$, for all $a \in \mathcal{A}$, $\nu_t(h_t, a) = \nu(a)$. \end{definition} - -\begin{remark}[Lean remark: properties vs constructors] -There are several ways to implement the last two definitions in Lean. We could write them as properties of algorithms and environments, or we can implement constructors that create algorithms and environments from the data in the definitions. We chose the latter option. Time will tell if it was a good choice. -\end{remark} +TODO: possibly change the ``stationary'' name. Let's detail four examples of interactions between an algorithm and an environment. @@ -64,18 +68,122 @@ \chapter{Iterative stochastic algorithms} \end{enumerate} +We will want to make global probabilistic statements about the whole sequence of actions and observations. +For example, we may want to prove that an optimization algorithm converges to the minimum of a function almost surely. +For such a statement to make sense, we need a probability space on which the whole sequence of actions and observations is defined as a random variable. + +We denote by $P[X \mid Y]$ the conditional distribution of a random variable $X$ given another random variable $Y$ under a probability measure $P$. + + +\begin{definition}[Algorithm-environment interaction]\label{def:IsAlgEnvSeq} + \uses{def:algorithm, def:environment} + \leanok + \lean{Learning.IsAlgEnvSeq} +Let $\mathfrak{A}$ be an algorithm as in Definition~\ref{def:algorithm} and $\mathfrak{E}$ be an environment as in Definition~\ref{def:environment}. +A probability space $(\Omega, P)$ and two sequences of random variables $A : \mathbb{N} \to \Omega \to \mathcal{A}$ and $R : \mathbb{N} \to \Omega \to \mathcal{R}$ form an algorithm-environment interaction for $\mathfrak{A}$ and $\mathfrak{E}$ if the following conditions hold: +\begin{enumerate} + \item The law of $A_0$ is $P_0$. + \item $P \left[ R_0 \mid A_0 \right] = \nu'_0$. + \item For all $t \in \mathbb{N}$, $P\left[A_{t+1} \mid A_0, R_0, \ldots, A_t, R_t \right] = \pi_t$. + \item For all $t \in \mathbb{N}$, $P\left[R_{t+1} \mid A_0, R_0, \ldots, A_t, R_t, A_{t+1}\right] = \nu_t$. +\end{enumerate} +\end{definition} + + +\begin{definition}[History]\label{def:history} + \leanok + \lean{Learning.IsAlgEnvSeq.hist, Learning.IsAlgEnvSeq.step} +For two sequences of random variables $A : \mathbb{N} \to \Omega \to \mathcal{A}$ and $R : \mathbb{N} \to \Omega \to \mathcal{R}$ (actions and observations), we call step of the interaction at time $t$ the random variable $X_t : \Omega \to \mathcal{A} \times \mathcal{R}$ defined by $X_t(\omega) = (A_t(\omega), R_t(\omega))$. +We call history up to time $t$ the random variable $H_t : \Omega \to (\mathcal{A} \times \mathcal{R})^{t+1}$ defined by $H_t(\omega) = (X_0(\omega), \ldots, X_t(\omega))$. +\end{definition} + + +\begin{lemma}\label{lem:law_step} + \uses{def:IsAlgEnvSeq, def:history} + \leanok + \lean{Learning.IsAlgEnvSeq.hasLaw_step_zero, Learning.IsAlgEnvSeq.hasCondDistrib_step} +In an algorithm-environment interaction $(A, R, P)$ as in Definition~\ref{def:IsAlgEnvSeq}, +\begin{itemize} + \item the law of the initial step $X_0$ is $P_0 \otimes \nu'_0$, + \item for all $t \in \mathbb{N}$, $P \left[ X_{t+1} \mid H_t \right] = \pi_t \otimes \nu_t$. +\end{itemize} +\end{lemma} + +\begin{proof}\leanok +Immediate from the properties of an algorithm-environment interaction. +\end{proof} + + +\begin{definition}\label{def:IsAlgEnvSeq.filtration} + \uses{def:IsAlgEnvSeq, def:history} + \leanok + \lean{Learning.IsAlgEnvSeq.filtration, Learning.IsAlgEnvSeq.filtrationAction} +For an algorithm-environment interaction $(A, R, P)$ as in Definition~\ref{def:IsAlgEnvSeq}, we denote by $\mathcal{F}_t$ the sigma-algebra generated by the history up to time $t$: $\mathcal{F}_t = \sigma(H_t)$. +We denote by $\mathcal{F}^A_t$ the sigma-algebra generated by the history up to time $t-1$ and the action at time $t$: $\mathcal{F}^A_t = \sigma(H_{t-1}, A_t)$. +\end{definition} + + +\begin{theorem}[\cite{lattimore2020bandit}, Proposition 4.8]\label{thm:isAlgEnvSeq_unique} + \uses{def:IsAlgEnvSeq} + \leanok + \lean{Learning.isAlgEnvSeq_unique} +If $(A, R, P)$ and $(A', R', P')$ are two algorithm-environment interactions for the same algorithm $\mathfrak{A}$ and environment $\mathfrak{E}$, then the joint distributions of the sequences of actions and observations are equal: the law of $(A_i, R_i)_{i \in \mathbb{N}}$ under $P$ is equal to the law of $(A'_i, R'_i)_{i \in \mathbb{N}}$ under $P'$. +\end{theorem} + +\begin{proof}\leanok + +\end{proof} + + + +\section{Stationary environment} + +Recall that in a stationary environment, there exists a Markov kernel $\nu : \mathcal{A} \rightsquigarrow \mathcal{R}$ such that $\nu'_0 = \nu$ and for all $t \in \mathbb{N}$, for all $h_t \in (\mathcal{A} \times \mathcal{R})^{t+1}$, for all $a \in \mathcal{A}$, $\nu_t(h_t, a) = \nu(a)$. + +Let $(A, R, P)$ be an algorithm-environment interaction in a stationary environment with kernel $\nu$. + +\begin{lemma}\label{lem:condDistrib_reward_stationaryEnv} + \uses{def:IsAlgEnvSeq, def:stationaryEnv} + \leanok + \lean{Learning.IsAlgEnvSeq.condDistrib_reward_stationaryEnv} +In a stationary environment, for any $t \in \mathbb{N}$, the conditional distribution $P\left[R_t \mid A_t\right]$ is $(A_{t*} P_{\mathcal{T}})$-almost surely equal to $\nu$. +\end{lemma} + +\begin{proof}\leanok + \uses{lem:law_step, def:stationaryEnv} + +\end{proof} + + +\begin{lemma}\label{lem:condIndepFun_reward_hist_action} + \uses{def:IsAlgEnvSeq, def:stationaryEnv} + \leanok + \lean{Learning.IsAlgEnvSeq.condIndepFun_reward_hist_action} +In a stationary environment, for any $t \in \mathbb{N}$, the reward $R_{t+1}$ is conditionally independent of the history $H_t$ given the action $A_{t+1}$ (more succinctly, $R_{t+1} \ind H_t \mid A_{t+1}$). +\end{lemma} + +\begin{proof}\leanok + +\end{proof} + \section{Probability space: Ionescu-Tulcea theorem} +In Theorem~\ref{thm:isAlgEnvSeq_unique}, we saw that the distribution of the sequence of actions and observations in a suitable probability space is uniquely determined by the algorithm and the environment. +We now show that such a probability space actually exists: for any algorithm and environment, we build an algorithm-environment interaction as in Definition~\ref{def:IsAlgEnvSeq}. + + + +\subsection{Ionescu-Tulcea theorem} + If we group together the policy of the algorithm and the kernel of the environment at each time step, we get a sequence of Markov kernels $(\kappa_t)_{t \in \mathbb{N}}$, with $\kappa_t : (\mathcal{A} \times \mathcal{R})^{t+1} \rightsquigarrow (\mathcal{A} \times \mathcal{R})$. -We will want to make global probabilistic statements about the whole sequence of actions and observations. -For example, we may want to prove that an optimization algorithm converges to the minimum of a function almost surely. -For such a statement to make sense, we need a probability space on which the whole sequence of actions and observations is defined as a random variable. + We now abstract that situation and consider a sequence of measurable spaces $(\Omega_t)_{t \in \mathbb{N}}$, a probability measure $\mu$ on $\Omega_0$ and a sequence of Markov kernels $\kappa_t : \prod_{s=0}^t \Omega_s \rightsquigarrow \Omega_{t+1}$. The Ionescu-Tulcea theorem builds a probability space from the sequence of kernels and the initial measure. + \begin{theorem}[Ionescu-Tulcea]\label{thm:ionescu-tulcea} \mathlibok \lean{ProbabilityTheory.Kernel.traj} @@ -100,7 +208,7 @@ \section{Probability space: Ionescu-Tulcea theorem} \end{definition} -\begin{definition}[Step and history]\label{def:history} +\begin{definition}[Step and history]\label{def:IT.history} \leanok \lean{Learning.IT.step, Learning.IT.hist} For $t \in \mathbb{N}$, we denote by $X_t \in \Omega_t$ the random variable describing the time step $t$, and by $H_t \in \prod_{s=0}^t \Omega_s$ the history up to time $t$. @@ -110,8 +218,8 @@ \section{Probability space: Ionescu-Tulcea theorem} Note: $(X_t)_{t \in \mathbb{N}}$ is the canonical process on $\Omega_{\mathcal{T}}$. $H_t$ is equal to $\pi_{[0,t]}$. -\begin{definition}[Filtration]\label{def:filtration} - \uses{def:history} +\begin{definition}[Filtration]\label{def:IT.filtration} + \uses{def:IT.history} \leanok \lean{Learning.IT.filtration} For $t \in \mathbb{N}$, we denote by $\mathcal{F}_t$ the sigma-algebra generated by the history up to time $t$: $\mathcal{F}_t = \sigma(H_t)$. @@ -122,7 +230,7 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:adapted_history} - \uses{def:history, def:filtration} + \uses{def:IT.history, def:IT.filtration} \leanok \lean{Learning.IT.adapted_step, Learning.IT.adapted_hist} The random variables $X_t$ and $H_t$ are $\mathcal{F}_t$-measurable. @@ -134,11 +242,8 @@ \section{Probability space: Ionescu-Tulcea theorem} \end{proof} -We now list properties of those random variables that follow from the construction of the trajectory measure. -We write $P[X \mid Y]$ for the conditional distribution of a random variable $X$ given another random variable $Y$ under a probability measure $P$. - \begin{lemma}\label{lem:condDistrib_X_add_one} - \uses{def:history, def:trajMeasure} + \uses{def:IT.history, def:trajMeasure} \leanok \lean{ProbabilityTheory.Kernel.condDistrib_trajMeasure} For any $t \in \mathbb{N}$, the conditional distribution $P_{\mathcal{T}}\left[X_{t+1} \mid H_t\right]$ is $((H_t)_* P_{\mathcal{T}})$-almost surely equal to $\kappa_t$. @@ -152,7 +257,7 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:law_X_zero} - \uses{def:history, def:trajMeasure} + \uses{def:IT.history, def:trajMeasure} \leanok \lean{Learning.IsAlgEnvSeq.hasLaw_step_zero} The law of $X_0$ under $P_{\mathcal{T}}$ is $\mu$. @@ -163,26 +268,26 @@ \section{Probability space: Ionescu-Tulcea theorem} \end{proof} -\paragraph{Case of an algorithm-environment interaction.} -We suppose now that, as in the algorithm-environment interaction, $\Omega_t = \mathcal{A}_t \times \mathcal{R}_t$ for some measurable spaces $\mathcal{A}_t$ and $\mathcal{R}_t$, and that for all $t \in \mathbb{N}$, $\kappa_t = \pi_t \otimes \nu_t$ for policy kernels $\pi_t : \prod_{s=0}^t(\mathcal{A}_s \times \mathcal{R}_s) \rightsquigarrow \mathcal{A}$ and feedback kernels $\nu_t : \prod_{s=0}^t(\mathcal{A}_s \times \mathcal{R}_s) \times \mathcal{A} \rightsquigarrow \mathcal{R}$. -Likewise, $\mu = \alpha_0 \otimes \nu'_0$ for a probability measure $\alpha_0$ on $\mathcal{A}_0$ and a Markov kernel $\nu'_0 : \mathcal{A}_0 \rightsquigarrow \mathcal{R}_0$. -The step random variable $X_t$ takes values in $\mathcal{A}_t \times \mathcal{R}_t$. -TODO: the code does not have $\mathcal{A}_t$ but a unique $\mathcal{A}$, same for $\mathcal{R}$. +\subsection{Case of an algorithm-environment interaction} -\begin{definition}\label{def:actionReward} - \uses{def:history} +We now go back to the setting of an algorithm interacting with an environment and suppose that $\Omega_t = \mathcal{A} \times \mathcal{R}$ for some measurable spaces $\mathcal{A}$ and $\mathcal{R}$, and that for all $t \in \mathbb{N}$, $\kappa_t = \pi_t \otimes \nu_t$ for policy kernels $\pi_t : (\mathcal{A} \times \mathcal{R})^{t+1} \rightsquigarrow \mathcal{A}$ and feedback kernels $\nu_t : (\mathcal{A} \times \mathcal{R})^{t+1} \times \mathcal{A} \rightsquigarrow \mathcal{R}$. +Likewise, $\mu = P_0 \otimes \nu'_0$ for a probability measure $P_0$ on $\mathcal{A}$ and a Markov kernel $\nu'_0 : \mathcal{A}_\rightsquigarrow \mathcal{R}$. +The step random variable $X_t$ takes values in $\mathcal{A} \times \mathcal{R}$. + +\begin{definition}\label{def:IT.actionReward} + \uses{def:IT.history} \leanok \lean{Learning.IT.action, Learning.IT.reward} -We write $A_t$ and $R_t$ for the projections of $X_t$ on $\mathcal{A}_t$ and $\mathcal{R}_t$ respectively. +We write $A_t$ and $R_t$ for the projections of $X_t$ on $\mathcal{A}$ and $\mathcal{R}$ respectively. $A_t$ is the action taken at time $t$ and $R_t$ is the reward received at time $t$. -Formally, $A_t(\omega) = \omega_{t,1}$ and $R_t(\omega) = \omega_{t,2}$ for $\omega = \prod_{t=0}^{+\infty}(\omega_{t,1}, \omega_{t,2}) \in \prod_{t=0}^{+\infty} \mathcal{A}_t \times \mathcal{R}_t$. +Formally, $A_t(\omega) = \omega_{t,1}$ and $R_t(\omega) = \omega_{t,2}$ for $\omega = \prod_{t=0}^{+\infty}(\omega_{t,1}, \omega_{t,2}) \in \Omega_{\mathcal{T}} = \prod_{t=0}^{+\infty} \mathcal{A} \times \mathcal{R}$. \end{definition} \begin{lemma}\label{lem:adapted_action_reward} - \uses{def:actionReward, def:filtration} + \uses{def:IT.actionReward, def:IT.filtration} \leanok \lean{Learning.IT.adapted_action, Learning.IT.adapted_reward} The random variables $A_t$ and $R_t$ are $\mathcal{F}_t$-measurable. @@ -198,9 +303,9 @@ \section{Probability space: Ionescu-Tulcea theorem} We need to check that the random variables $A_t$ and $R_t$ have the expected conditional distributions. \begin{lemma}\label{lem:condDistrib_A_add_one} - \uses{def:actionReward, def:trajMeasure, def:algorithm, def:environment} + \uses{def:IT.actionReward, def:trajMeasure, def:algorithm, def:environment} \leanok - \lean{Learning.IsAlgEnvSeq.hasCondDistrib_action} + \lean{Learning.IT.condDistrib_action} For any $t \in \mathbb{N}$, the conditional distribution $P_{\mathcal{T}}\left[A_{t+1} \mid H_t\right]$ is $((H_t)_* P_{\mathcal{T}})$-almost surely equal to $\pi_t$. \end{lemma} @@ -212,9 +317,9 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:condDistrib_R_add_one} - \uses{def:actionReward, def:trajMeasure, def:algorithm, def:environment} + \uses{def:IT.actionReward, def:trajMeasure, def:algorithm, def:environment} \leanok - \lean{Learning.IsAlgEnvSeq.hasCondDistrib_reward} + \lean{Learning.IT.condDistrib_reward} For any $t \in \mathbb{N}$, the conditional distribution $P_{\mathcal{T}}\left[R_{t+1} \mid H_t, A_{t+1}\right]$ is $((H_t, A_{t+1})_* P_{\mathcal{T}})$-almost surely equal to $\nu_t$. \end{lemma} @@ -231,9 +336,9 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:law_A_zero} - \uses{def:actionReward, def:trajMeasure, def:algorithm, def:environment} + \uses{def:IT.actionReward, def:trajMeasure, def:algorithm, def:environment} \leanok - \lean{Learning.IsAlgEnvSeq.hasLaw_action_zero} + \lean{Learning.IT.hasLaw_action_zero} The law of $A_0$ under $P_{\mathcal{T}}$ is $\alpha_0$. \end{lemma} @@ -244,9 +349,9 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:condDistrib_R_zero} - \uses{def:actionReward, def:trajMeasure, def:algorithm, def:environment} + \uses{def:IT.actionReward, def:trajMeasure, def:algorithm, def:environment} \leanok - \lean{Learning.IsAlgEnvSeq.hasCondDistrib_reward_zero} + \lean{Learning.IT.condDistrib_reward_zero} The conditional distribution $P_{\mathcal{T}}\left[R_0 \mid A_0\right]$ is $(A_{0*} P_{\mathcal{T}})$-almost surely equal to $\nu'_0$. \end{lemma} @@ -260,34 +365,16 @@ \section{Probability space: Ionescu-Tulcea theorem} \end{proof} - -\section{Stationary environment} - -Recall that in a stationary environment, there exists a Markov kernel $\nu : \mathcal{A} \rightsquigarrow \mathcal{R}$ such that $\nu'_0 = \nu$ and for all $t \in \mathbb{N}$, for all $h_t \in (\mathcal{A} \times \mathcal{R})^{t+1}$, for all $a \in \mathcal{A}$, $\nu_t(h_t, a) = \nu(a)$. - - -\begin{lemma}\label{lem:condDistrib_reward_stationaryEnv} - \uses{def:actionReward, def:trajMeasure, def:algorithm, def:stationaryEnv} +\begin{theorem}\label{thm:isAlgEnvSeq_trajMeasure} + \uses{def:IsAlgEnvSeq, def:trajMeasure} \leanok - \lean{Learning.IsAlgEnvSeq.condDistrib_reward_stationaryEnv} -In a stationary environment, for any $t \in \mathbb{N}$, the conditional distribution $P_{\mathcal{T}}\left[R_t \mid A_t\right]$ is $(A_{t*} P_{\mathcal{T}})$-almost surely equal to $\nu$. -\end{lemma} - -\begin{proof}\leanok - \uses{lem:condDistrib_R_add_one, def:stationaryEnv} - -\end{proof} - - -\begin{lemma}\label{lem:condIndepFun_reward_hist_action} - \uses{def:actionReward, def:trajMeasure, def:algorithm, def:stationaryEnv} - \leanok - \lean{Learning.IsAlgEnvSeq.condIndepFun_reward_hist_action} -In a stationary environment, for any $t \in \mathbb{N}$, the reward $R_{t+1}$ is conditionally independent of the history $H_t$ given the action $A_{t+1}$ (more succinctly, $R_{t+1} \ind H_t \mid A_{t+1}$). -\end{lemma} + \lean{Learning.IT.isAlgEnvSeq_trajMeasure} +In the probability space $(\Omega_{\mathcal{T}}, P_{\mathcal{T}})$ constructed from an algorithm $\mathfrak{A}$ and an environment $\mathfrak{E}$ as above, the sequences of random variables $A : \mathbb{N} \to \Omega_{\mathcal{T}} \to \mathcal{A}$ and $R : \mathbb{N} \to \Omega_{\mathcal{T}} \to \mathcal{R}$ form an algorithm-environment interaction for $\mathfrak{A}$ and $\mathfrak{E}$. +\end{theorem} \begin{proof}\leanok - + \uses{lem:law_A_zero, lem:condDistrib_R_zero, lem:condDistrib_A_add_one, lem:condDistrib_R_add_one} +The four conditions of Definition~\ref{def:IsAlgEnvSeq} are exactly the statements of Lemmas~\ref{lem:law_A_zero}, \ref{lem:condDistrib_R_zero}, \ref{lem:condDistrib_A_add_one} and \ref{lem:condDistrib_R_add_one}. \end{proof} @@ -298,7 +385,7 @@ \section{Finitely many actions} We can also define the time step at which an action was chosen a certain number of times, and the value of the reward obtained when pulling an action for the $m$-th time. \begin{definition}[Pull counts]\label{def:pullCount} - \uses{def:actionReward} + \uses{def:IT.actionReward} \leanok \lean{Learning.pullCount} For an action $a \in \mathcal{A}$ and a time $t \in \mathbb{N}$, we denote by $N_{t,a}$ the number of times that action $a$ has been chosen before time $t$, that is $N_{t,a} = \sum_{s=0}^{t-1} \mathbb{I}\{A_s = a\}$. @@ -312,15 +399,15 @@ \section{Finitely many actions} That means that any tool used to define a policy must be a function defined on $(\mathcal{A} \times \mathcal{R})^{t+1}$. For example a definition of the empirical mean of an action must be a function $t : \mathbb{N} \to (\mathcal{A} \times \mathcal{R})^{t+1} \to \mathbb{R}$. -When we analyze an algorithm, we work on the other hand on the bandit probability space $(\Omega, \mathbb{P})$, in which $\Omega = (\mathcal{A} \times \mathcal{R})^{\mathbb{N}}$ is the full history, which describes the whole sequence of actions and rewards. +When we analyze an algorithm, we work on the other hand on a probability space $(\Omega, P)$, in which $\Omega$ could be for example $(\mathcal{A} \times \mathcal{R})^{\mathbb{N}}$, the full history, which describes the whole sequence of actions and rewards. As a stochastic process, the empirical mean of an action is a function $\mathbb{N} \to (\mathcal{A} \times \mathcal{R})^{\mathbb{N}} \to \mathbb{R}$. -Thus there are two similar but still distinct types of objects: those defined on the partial history, which are used to build algorithms, and those defined on the full history, which are used to analyze algorithms. +Thus there are two similar but still distinct types of objects: those defined on the partial history, which are used to build algorithms, and those defined on a generic probability space (the full history in the Ionescu-Tulcea construction), which are used to analyze algorithms. \end{remark} \begin{lemma}\label{lem:pullCount_basic} - \uses{def:pullCount, def:actionReward} + \uses{def:pullCount, def:IT.actionReward} \leanok \lean{Learning.pullCount_zero, Learning.pullCount_mono, Learning.pullCount_add_one, Learning.pullCount_le, Learning.pullCount_congr} We note the following basic properties of $N_{t,a}$: @@ -339,10 +426,10 @@ \section{Finitely many actions} \begin{lemma}\label{lem:predictable_pullCount} - \uses{def:filtration, def:pullCount} + \uses{def:IsAlgEnvSeq.filtration, def:pullCount} \leanok \lean{Learning.isPredictable_pullCount} -Let $a \in \mathcal{A}$. The process $(N_{t,a})_{t \in \mathbb{N}}$ is predictable with respect to the filtration $\mathcal{F}$. +Let $a \in \mathcal{A}$. The process $(N_{t,a})_{t \in \mathbb{N}}$ is predictable with respect to the filtration $\mathcal{F}$ of the algorithm-environment interaction. \end{lemma} \begin{proof}\leanok @@ -384,7 +471,7 @@ \section{Finitely many actions} \begin{lemma}\label{lem:isStoppingTime_stepsUntil} - \uses{def:filtration, def:stepsUntil} + \uses{def:IsAlgEnvSeq.filtration, def:stepsUntil} \leanok \lean{Learning.isStoppingTime_stepsUntil} Let $a \in \mathcal{A}$. For any $n > 0$, the random variable $T_{n,a}$ is a stopping time with respect to the filtration $\mathcal{F}$. @@ -441,7 +528,7 @@ \section{Scalar rewards} \begin{definition}[Sum of rewards]\label{def:sumRewards} - \uses{def:actionReward} + \uses{def:IT.actionReward} \leanok \lean{Learning.sumRewards} Let $S_{t, a} = \sum_{s=0}^{t-1} R_s \mathbb{I}\{A_s = a\}$ be the sum of the rewards obtained by chosing action $a$ before time $t$. diff --git a/blueprint/src/chapters/bandit.tex b/blueprint/src/chapters/bandit.tex index ce9f3c09..2b2fe870 100644 --- a/blueprint/src/chapters/bandit.tex +++ b/blueprint/src/chapters/bandit.tex @@ -34,6 +34,54 @@ \section{Algorithm, bandit and probability space} \end{definition} +\section{The array model of rewards} + +We previously built a probability space on which we can define the sequence of arms and rewards generated by the interaction between the algorithm and the bandit, using the Ionescu-Tulcea theorem. +From Theorem~\ref{thm:isAlgEnvSeq_unique}, we know that the law of the sequence of arms and rewards is independent of the probability space used to define them. +Nonetheless, we now build an alternative model of the rewards, on which it will be easier to prove concentration inequalities. +By the uniqueness of the law, these statements will then transfer to any algorithm-environment interaction. + + +\begin{definition}\label{def:arrayMeasure} + \leanok + \lean{Bandits.ArrayModel.probSpace, Bandits.ArrayModel.arrayMeasure} +Let $I = [0,1]$ and let $P_U$ be the uniform distribution on $I$. We define the probability space $(\Omega_{\mathcal{A}}, P_{\mathcal{A}})$, where +\begin{align*} + \Omega_{\mathcal{A}} &:= I^{\mathbb{N}} \times \mathcal{R}^{\mathbb{N} \times \mathcal{A}} + \: , \\ + P_{\mathcal{A}} &:= \left( \bigotimes_{n \in \mathbb{N}} P_U \right) \otimes \left( \bigotimes_{n \in \mathbb{N}, a \in \mathcal{A}} \nu(a) \right) + \: . +\end{align*} +\end{definition} + + +\begin{definition}\label{def:algFunction} + \uses{def:algorithm} + \leanok + \lean{Bandits.ArrayModel.algFunction, Bandits.ArrayModel.initAlgFunction} +Since $\mathcal{A}$ and $\mathcal{R}$ are standard Borel spaces, there exists jointly measurable functions $f'_0 : I \to \mathcal{A}$ and $f_t : (\mathcal{A} \times \mathcal{R})^{t+1} \times I \to \mathcal{A}$ such that +\begin{itemize} + \item the law of $f'_0$ is $P_0$, + \item for all history $h_t \in (\mathcal{A} \times \mathcal{R})^{t+1}$, the law of $f_t(h_t, \cdot)$ is $\pi_t(h_t)$. +\end{itemize} +\end{definition} + + +TODO: lots of results + + +\begin{theorem}\label{thm:isAlgEnvSeq_arrayMeasure} + \uses{def:bandit, def:arrayMeasure} + \leanok + \lean{Bandits.ArrayModel.isAlgEnvSeq_arrayMeasure} +TODO +\end{theorem} + +\begin{proof} + +\end{proof} + + \section{Alternative models: rewards indexed by time or pull count}\label{sec:alt_model} The description of the bandit model above considers that at time $t$, a reward $R_t$ is generated, depending on the arm $A_t$ pulled at that time. diff --git a/blueprint/src/print.tex b/blueprint/src/print.tex index bf1c8685..4445c547 100644 --- a/blueprint/src/print.tex +++ b/blueprint/src/print.tex @@ -25,7 +25,7 @@ \input{macros/print} \title{LeanBandits\\ \Large{A Lean package for bandit algorithms}} -\author{Rémy Degenne} +\author{Rémy Degenne, Paulo Rauber} \begin{document} \maketitle diff --git a/blueprint/src/web.tex b/blueprint/src/web.tex index d7590660..8f3ac506 100644 --- a/blueprint/src/web.tex +++ b/blueprint/src/web.tex @@ -21,7 +21,7 @@ \dochome{https://RemyDegenne.github.io/lean-bandits/docs} \title{LeanBandits} -\author{Rémy Degenne} +\author{Rémy Degenne, Paulo Rauber} \begin{document} \maketitle From a260892c4e81d476f7c568009bdb0ca8630a319b Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 15 Jan 2026 11:28:15 +0100 Subject: [PATCH 30/30] fix --- blueprint/src/chapters/bandit.tex | 4 ++-- blueprint/src/chapters/ucb.tex | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/blueprint/src/chapters/bandit.tex b/blueprint/src/chapters/bandit.tex index 2b2fe870..e89628b6 100644 --- a/blueprint/src/chapters/bandit.tex +++ b/blueprint/src/chapters/bandit.tex @@ -109,7 +109,7 @@ \section{Alternative models: rewards indexed by time or pull count}\label{sec:al \begin{lemma}\label{lem:condIndepFun_reward_stepsUntil_arm} - \uses{def:stepsUntil, def:actionReward, def:Bandit.measure} + \uses{def:stepsUntil, def:IT.actionReward, def:Bandit.measure} \leanok \lean{Bandits.condIndepFun_reward_stepsUntil_action} For $t > 0$, $R_t \ind \mathbb{I}\{T_{n, a} = t\} \mid A_t$. @@ -312,7 +312,7 @@ \section{Regret and other bandit quantities} \begin{definition}[Regret]\label{def:regret} - \uses{def:armMean, def:actionReward} + \uses{def:armMean, def:IT.actionReward} \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: diff --git a/blueprint/src/chapters/ucb.tex b/blueprint/src/chapters/ucb.tex index 1cfb5177..54ffc48a 100644 --- a/blueprint/src/chapters/ucb.tex +++ b/blueprint/src/chapters/ucb.tex @@ -1,7 +1,7 @@ \section{UCB} \begin{definition}[UCB algorithm]\label{def:ucbAlgorithm} - \uses{def:actionReward, def:pullCount, def:empMean} + \uses{def:IT.actionReward, def:pullCount, def:empMean} \leanok \lean{Bandits.UCB.nextArm, Bandits.ucbAlgorithm} The UCB algorithm with parameter $c \in \mathbb{R}_+$ is defined as follows: