From eea8b6694d8b8d5b0b90e0426a322e19203a0a9e Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 19 Jan 2026 11:23:21 +0000 Subject: [PATCH 001/155] Draft TS --- LeanBandits/BanditAlgorithms/TS.lean | 50 ++++++++++++++++++++++++++++ 1 file changed, 50 insertions(+) create mode 100644 LeanBandits/BanditAlgorithms/TS.lean diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean new file mode 100644 index 00000000..a62a2476 --- /dev/null +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -0,0 +1,50 @@ +/- +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 + +open MeasureTheory ProbabilityTheory Finset +open Learning + +variable {α R : Type*} [mα : MeasurableSpace α] [mR : MeasurableSpace R] +variable {ℰ : Type*} [mℰ : MeasurableSpace ℰ] +variable {Ω : Type*} [mΩ : MeasurableSpace Ω] + +structure isStationaryBayesAlgEnvSeq + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + (alg : Algorithm α R) + (Q : Measure ℰ) [IsProbabilityMeasure Q] (κ : Kernel (ℰ × α) R) [IsMarkovKernel κ] + (E : Ω → ℰ) (A : ℕ → Ω → α) (R' : ℕ → Ω → R) + (P : Measure Ω) [IsFiniteMeasure P] : Prop where + measurable_E : Measurable E := by fun_prop + measurable_A n : Measurable (A n) := by fun_prop + measurable_R n : Measurable (R' n) := by fun_prop + hasLaw_env : HasLaw E Q P + hasLaw_action_zero : HasLaw (fun ω ↦ (A 0 ω)) alg.p0 P + hasCondDistrib_reward_zero : HasCondDistrib (R' 0) (fun ω ↦ (E ω, A 0 ω)) κ P + hasCondDistrib_action n : + HasCondDistrib (A (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, E ω)) + ((alg.policy n).prodMkRight _) P + hasCondDistrib_reward n : + HasCondDistrib (R' (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, E ω, A (n + 1) ω)) + (κ.prodMkLeft _) P + + + +-- structure Algorithm (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] where +-- /-- Policy or sampling rule: distribution of the next action. -/ +-- policy : (n : ℕ) → Kernel (Iic n → α × R) α +-- [h_policy : ∀ n, IsMarkovKernel (policy n)] +-- /-- Distribution of the first action. -/ +-- p0 : Measure α +-- [hp0 : IsProbabilityMeasure p0] + + +-- noncomputable +-- def tsAlgorithm : Algorithm (Fin K) ℝ where +-- policy := tsPolicy hK μ ℓ +-- h_policy := isMarkovKernel_tsPolicy hK μ ℓ +-- p0 := tsInitialPolicy hK μ ℓ +-- hp0 := isProbabilityMeasure_tsInitialPolicy hK μ ℓ From 4f305997823b13e607070d782146f5ae864a9fcf Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 19 Jan 2026 13:44:34 +0000 Subject: [PATCH 002/155] Draft TS --- LeanBandits/BanditAlgorithms/TS.lean | 79 +++++++++++++++++++++++----- 1 file changed, 66 insertions(+), 13 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index a62a2476..dfd1b685 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -3,11 +3,15 @@ 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 Mathlib.Probability.Distributions.Uniform +import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.SequentialLearning.Algorithm open MeasureTheory ProbabilityTheory Finset open Learning +section StationaryBayesAlgEnvSeq + variable {α R : Type*} [mα : MeasurableSpace α] [mR : MeasurableSpace R] variable {ℰ : Type*} [mℰ : MeasurableSpace ℰ] variable {Ω : Type*} [mΩ : MeasurableSpace Ω] @@ -31,20 +35,69 @@ structure isStationaryBayesAlgEnvSeq HasCondDistrib (R' (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, E ω, A (n + 1) ω)) (κ.prodMkLeft _) P +end StationaryBayesAlgEnvSeq + +section Uniform + +variable {K : ℕ} (hK : 0 < K) + +noncomputable +def uniformAlgorithm : Algorithm (Fin K) ℝ := + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + { policy _ := Kernel.const _ (PMF.uniformOfFintype (Fin K)).toMeasure + p0 := (PMF.uniformOfFintype (Fin K)).toMeasure } + +end Uniform + +section ThompsonSampling + +variable {K : ℕ} (hK : 0 < K) + +variable {ℰ : Type*} [mℰ : MeasurableSpace ℰ] [StandardBorelSpace ℰ] [Nonempty ℰ] +variable {Ω : Type*} [mΩ : MeasurableSpace Ω] + +variable (Q : Measure ℰ) [IsProbabilityMeasure Q] (κ : Kernel (ℰ × (Fin K)) ℝ) [IsMarkovKernel κ] +variable (E : Ω → ℰ) (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → ℝ) +variable (P : Measure Ω) [IsFiniteMeasure P] + +noncomputable +def tsPosterior (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) ℰ := + condDistrib E (IsAlgEnvSeq.hist A R' n) P + +noncomputable +def isMarkovKernel_tsPosterior (n : ℕ) : IsMarkovKernel (tsPosterior E A R' P n) := by + unfold tsPosterior + infer_instance + +noncomputable +def tsPolicy (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + (tsPosterior E A R' P n).map (measurableArgmax (fun e k ↦ (κ (e, k))[id])) + +def isMarkovKernel_tsPolicy (n : ℕ) : IsMarkovKernel (tsPolicy hK κ E A R' P n) := by + have : IsMarkovKernel (tsPosterior E A R' P n) := isMarkovKernel_tsPosterior E A R' P n + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + apply Kernel.IsMarkovKernel.map + exact measurable_measurableArgmax fun k => + (stronglyMeasurable_id.integral_kernel (κ := κ.comap (·, k) (by fun_prop))).measurable +noncomputable +def tsInitPolicy : Measure (Fin K) := + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + Q.map (measurableArgmax (fun e k ↦ (κ (e, k))[id])) --- structure Algorithm (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] where --- /-- Policy or sampling rule: distribution of the next action. -/ --- policy : (n : ℕ) → Kernel (Iic n → α × R) α --- [h_policy : ∀ n, IsMarkovKernel (policy n)] --- /-- Distribution of the first action. -/ --- p0 : Measure α --- [hp0 : IsProbabilityMeasure p0] +def isProbabilityMeasure_tsInitPolicy : IsProbabilityMeasure (tsInitPolicy hK Q κ) := by + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + apply Measure.isProbabilityMeasure_map + apply Measurable.aemeasurable + exact (measurable_measurableArgmax fun k => + (stronglyMeasurable_id.integral_kernel (κ := κ.comap (·, k) (by fun_prop))).measurable) +noncomputable +def tsAlgorithm : Algorithm (Fin K) ℝ where + policy := tsPolicy hK κ E A R' P + h_policy := isMarkovKernel_tsPolicy hK κ E A R' P + p0 := tsInitPolicy hK Q κ + hp0 := isProbabilityMeasure_tsInitPolicy hK Q κ --- noncomputable --- def tsAlgorithm : Algorithm (Fin K) ℝ where --- policy := tsPolicy hK μ ℓ --- h_policy := isMarkovKernel_tsPolicy hK μ ℓ --- p0 := tsInitialPolicy hK μ ℓ --- hp0 := isProbabilityMeasure_tsInitialPolicy hK μ ℓ +end ThompsonSampling From 747c29f59a94244379154336e7bc54a44745dd22 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 20 Jan 2026 11:16:43 +0000 Subject: [PATCH 003/155] Organize --- LeanBandits/BanditAlgorithms/TS.lean | 105 +++++++++++++++------------ 1 file changed, 59 insertions(+), 46 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index dfd1b685..e6953538 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -3,79 +3,80 @@ 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 Mathlib.Probability.Distributions.Uniform import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.SequentialLearning.Algorithm open MeasureTheory ProbabilityTheory Finset open Learning -section StationaryBayesAlgEnvSeq +section Algorithm -- SequentialLearning/Algorithm.lean variable {α R : Type*} [mα : MeasurableSpace α] [mR : MeasurableSpace R] -variable {ℰ : Type*} [mℰ : MeasurableSpace ℰ] -variable {Ω : Type*} [mΩ : MeasurableSpace Ω] -structure isStationaryBayesAlgEnvSeq +namespace Learning + +def Algorithm.prod_left (E : Type*) [MeasurableSpace E] (alg : Algorithm α R) : + Algorithm α (E × R) where + policy n := (alg.policy n).comap (fun h i ↦ ((h i).1, (h i).2.2)) (by fun_prop) + p0 := alg.p0 + +variable {Ω E : Type*} [mΩ : MeasurableSpace Ω] [mE : MeasurableSpace E] + +def IsPOAlgEnvSeq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) - (Q : Measure ℰ) [IsProbabilityMeasure Q] (κ : Kernel (ℰ × α) R) [IsMarkovKernel κ] - (E : Ω → ℰ) (A : ℕ → Ω → α) (R' : ℕ → Ω → R) - (P : Measure Ω) [IsFiniteMeasure P] : Prop where - measurable_E : Measurable E := by fun_prop - measurable_A n : Measurable (A n) := by fun_prop - measurable_R n : Measurable (R' n) := by fun_prop - hasLaw_env : HasLaw E Q P - hasLaw_action_zero : HasLaw (fun ω ↦ (A 0 ω)) alg.p0 P - hasCondDistrib_reward_zero : HasCondDistrib (R' 0) (fun ω ↦ (E ω, A 0 ω)) κ P - hasCondDistrib_action n : - HasCondDistrib (A (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, E ω)) - ((alg.policy n).prodMkRight _) P - hasCondDistrib_reward n : - HasCondDistrib (R' (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, E ω, A (n + 1) ω)) - (κ.prodMkLeft _) P - -end StationaryBayesAlgEnvSeq - -section Uniform - -variable {K : ℕ} (hK : 0 < K) + [StandardBorelSpace E] [Nonempty E] + (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (E' : ℕ → Ω → E) + (alg : Algorithm α R) (env : Environment α (E × R)) (P : Measure Ω) [IsFiniteMeasure P] + := IsAlgEnvSeq A (fun n ω ↦ (E' n ω, R' n ω)) (alg.prod_left E) env P -noncomputable -def uniformAlgorithm : Algorithm (Fin K) ℝ := - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - { policy _ := Kernel.const _ (PMF.uniformOfFintype (Fin K)).toMeasure - p0 := (PMF.uniformOfFintype (Fin K)).toMeasure } +end Learning + +end Algorithm -end Uniform +section StationaryEnv -- SequentialLearning/StationaryEnv.lean -section ThompsonSampling +variable {α R E : Type*} [mα : MeasurableSpace α] [mR : MeasurableSpace R] [mE : MeasurableSpace E] -variable {K : ℕ} (hK : 0 < K) +noncomputable +def BayesStationaryEnv (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) + [IsMarkovKernel κ] : Environment α (E × R) where + feedback n := + let g : (Iic n → α × (E × R)) × α → (α × E) := fun (h, a) => (a, (h ⟨0, by simp⟩).2.1) + (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) + h_feedback := inferInstance + ν0 := (Kernel.const α Q) ⊗ₖ κ + hp0 := Kernel.IsMarkovKernel.compProd _ _ + +end StationaryEnv -variable {ℰ : Type*} [mℰ : MeasurableSpace ℰ] [StandardBorelSpace ℰ] [Nonempty ℰ] +section ThompsonSampling -- BanditAlgorithms/TS.lean + +variable {K : ℕ} +variable {E : Type*} [mE : MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] variable {Ω : Type*} [mΩ : MeasurableSpace Ω] -variable (Q : Measure ℰ) [IsProbabilityMeasure Q] (κ : Kernel (ℰ × (Fin K)) ℝ) [IsMarkovKernel κ] -variable (E : Ω → ℰ) (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → ℝ) +variable (hK : 0 < K) +variable (Q : Measure E) [IsProbabilityMeasure Q] +variable (κ : Kernel (E × Fin K) ℝ) [IsMarkovKernel κ] +variable (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (E' : ℕ → Ω → E) variable (P : Measure Ω) [IsFiniteMeasure P] noncomputable -def tsPosterior (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) ℰ := - condDistrib E (IsAlgEnvSeq.hist A R' n) P +def tsPosterior (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) E := + condDistrib (E' 0) (IsAlgEnvSeq.hist A R' n) P noncomputable -def isMarkovKernel_tsPosterior (n : ℕ) : IsMarkovKernel (tsPosterior E A R' P n) := by +def isMarkovKernel_tsPosterior (n : ℕ) : IsMarkovKernel (tsPosterior A R' E' P n) := by unfold tsPosterior infer_instance noncomputable def tsPolicy (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - (tsPosterior E A R' P n).map (measurableArgmax (fun e k ↦ (κ (e, k))[id])) + (tsPosterior A R' E' P n).map (measurableArgmax (fun e k ↦ (κ (e, k))[id])) -def isMarkovKernel_tsPolicy (n : ℕ) : IsMarkovKernel (tsPolicy hK κ E A R' P n) := by - have : IsMarkovKernel (tsPosterior E A R' P n) := isMarkovKernel_tsPosterior E A R' P n +def isMarkovKernel_tsPolicy (n : ℕ) : IsMarkovKernel (tsPolicy hK κ A R' E' P n) := by + have : IsMarkovKernel (tsPosterior A R' E' P n) := isMarkovKernel_tsPosterior A R' E' P n have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK apply Kernel.IsMarkovKernel.map exact measurable_measurableArgmax fun k => @@ -95,9 +96,21 @@ def isProbabilityMeasure_tsInitPolicy : IsProbabilityMeasure (tsInitPolicy hK Q noncomputable def tsAlgorithm : Algorithm (Fin K) ℝ where - policy := tsPolicy hK κ E A R' P - h_policy := isMarkovKernel_tsPolicy hK κ E A R' P + policy := tsPolicy hK κ A R' E' P + h_policy := isMarkovKernel_tsPolicy hK κ A R' E' P p0 := tsInitPolicy hK Q κ hp0 := isProbabilityMeasure_tsInitPolicy hK Q κ end ThompsonSampling + +-- section Uniform + +-- variable {K : ℕ} (hK : 0 < K) + +-- noncomputable +-- def uniformAlgorithm : Algorithm (Fin K) ℝ := +-- have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK +-- { policy _ := Kernel.const _ (PMF.uniformOfFintype (Fin K)).toMeasure +-- p0 := (PMF.uniformOfFintype (Fin K)).toMeasure } + +-- end Uniform From 5b2b5e5e7da24d3c26847c15a72a3962a27cb90f Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 20 Jan 2026 13:28:47 +0000 Subject: [PATCH 004/155] Integrate --- LeanBandits/BanditAlgorithms/TS.lean | 64 +++++++++++++++++----------- 1 file changed, 39 insertions(+), 25 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index e6953538..852b71ca 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.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 Mathlib.Probability.Distributions.Uniform import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.SequentialLearning.Algorithm @@ -38,7 +39,7 @@ section StationaryEnv -- SequentialLearning/StationaryEnv.lean variable {α R E : Type*} [mα : MeasurableSpace α] [mR : MeasurableSpace R] [mE : MeasurableSpace E] noncomputable -def BayesStationaryEnv (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) +def bayesStationaryEnv (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] : Environment α (E × R) where feedback n := let g : (Iic n → α × (E × R)) × α → (α × E) := fun (h, a) => (a, (h ⟨0, by simp⟩).2.1) @@ -47,8 +48,31 @@ def BayesStationaryEnv (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α ν0 := (Kernel.const α Q) ⊗ₖ κ hp0 := Kernel.IsMarkovKernel.compProd _ _ +noncomputable +def bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) + [IsMarkovKernel κ] (alg : Algorithm α R) := + trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) + +instance isProbabilityMeasure_bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] + (κ : Kernel (α × E) R) [IsMarkovKernel κ] (alg : Algorithm α R) : + IsProbabilityMeasure (bayesTrajMeasure Q κ alg) := by + unfold bayesTrajMeasure + infer_instance + end StationaryEnv +section Uniform -- BanditAlgorithms/Uniform.lean + +variable {K : ℕ} (hK : 0 < K) + +noncomputable +def uniformAlgorithm : Algorithm (Fin K) ℝ := + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + { policy _ := Kernel.const _ (PMF.uniformOfFintype (Fin K)).toMeasure + p0 := (PMF.uniformOfFintype (Fin K)).toMeasure } + +end Uniform + section ThompsonSampling -- BanditAlgorithms/TS.lean variable {K : ℕ} @@ -57,60 +81,50 @@ variable {Ω : Type*} [mΩ : MeasurableSpace Ω] variable (hK : 0 < K) variable (Q : Measure E) [IsProbabilityMeasure Q] -variable (κ : Kernel (E × Fin K) ℝ) [IsMarkovKernel κ] -variable (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (E' : ℕ → Ω → E) -variable (P : Measure Ω) [IsFiniteMeasure P] +variable (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] noncomputable def tsPosterior (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) E := + let P := bayesTrajMeasure Q κ (uniformAlgorithm hK) + let E' : ℕ → (ℕ → (Fin K) × (E × ℝ)) → E := fun n ω ↦ (ω n).2.1 + let A : ℕ → (ℕ → (Fin K) × (E × ℝ)) → (Fin K) := fun n ω ↦ (ω n).1 + let R' : ℕ → (ℕ → (Fin K) × (E × ℝ)) → ℝ := fun n ω ↦ (ω n).2.2 condDistrib (E' 0) (IsAlgEnvSeq.hist A R' n) P noncomputable -def isMarkovKernel_tsPosterior (n : ℕ) : IsMarkovKernel (tsPosterior A R' E' P n) := by +def isMarkovKernel_tsPosterior (n : ℕ) : IsMarkovKernel (tsPosterior hK Q κ n) := by unfold tsPosterior infer_instance noncomputable def tsPolicy (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - (tsPosterior A R' E' P n).map (measurableArgmax (fun e k ↦ (κ (e, k))[id])) + (tsPosterior hK Q κ n).map (measurableArgmax (fun e k ↦ (κ (k, e))[id])) -def isMarkovKernel_tsPolicy (n : ℕ) : IsMarkovKernel (tsPolicy hK κ A R' E' P n) := by - have : IsMarkovKernel (tsPosterior A R' E' P n) := isMarkovKernel_tsPosterior A R' E' P n +def isMarkovKernel_tsPolicy (n : ℕ) : IsMarkovKernel (tsPolicy hK Q κ n) := by + have : IsMarkovKernel (tsPosterior hK Q κ n) := isMarkovKernel_tsPosterior hK Q κ n have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK apply Kernel.IsMarkovKernel.map exact measurable_measurableArgmax fun k => - (stronglyMeasurable_id.integral_kernel (κ := κ.comap (·, k) (by fun_prop))).measurable + (stronglyMeasurable_id.integral_kernel (κ := κ.comap (k, ·) (by fun_prop))).measurable noncomputable def tsInitPolicy : Measure (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - Q.map (measurableArgmax (fun e k ↦ (κ (e, k))[id])) + Q.map (measurableArgmax (fun e k ↦ (κ (k, e))[id])) def isProbabilityMeasure_tsInitPolicy : IsProbabilityMeasure (tsInitPolicy hK Q κ) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK apply Measure.isProbabilityMeasure_map apply Measurable.aemeasurable exact (measurable_measurableArgmax fun k => - (stronglyMeasurable_id.integral_kernel (κ := κ.comap (·, k) (by fun_prop))).measurable) + (stronglyMeasurable_id.integral_kernel (κ := κ.comap (k, ·) (by fun_prop))).measurable) noncomputable def tsAlgorithm : Algorithm (Fin K) ℝ where - policy := tsPolicy hK κ A R' E' P - h_policy := isMarkovKernel_tsPolicy hK κ A R' E' P + policy := tsPolicy hK Q κ + h_policy := isMarkovKernel_tsPolicy hK Q κ p0 := tsInitPolicy hK Q κ hp0 := isProbabilityMeasure_tsInitPolicy hK Q κ end ThompsonSampling - --- section Uniform - --- variable {K : ℕ} (hK : 0 < K) - --- noncomputable --- def uniformAlgorithm : Algorithm (Fin K) ℝ := --- have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK --- { policy _ := Kernel.const _ (PMF.uniformOfFintype (Fin K)).toMeasure --- p0 := (PMF.uniformOfFintype (Fin K)).toMeasure } - --- end Uniform From 52202283fc275931d6c165c37e78404668853043 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 20 Jan 2026 14:10:20 +0000 Subject: [PATCH 005/155] Organize TS --- LeanBandits/BanditAlgorithms/TS.lean | 39 +++++++++++++++++----------- 1 file changed, 24 insertions(+), 15 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 852b71ca..2420da49 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -6,6 +6,7 @@ Authors: Rémy Degenne, Paulo Rauber import Mathlib.Probability.Distributions.Uniform import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.SequentialLearning.Algorithm +import LeanBandits.SequentialLearning.IonescuTulceaSpace open MeasureTheory ProbabilityTheory Finset open Learning @@ -21,15 +22,6 @@ def Algorithm.prod_left (E : Type*) [MeasurableSpace E] (alg : Algorithm α R) : policy n := (alg.policy n).comap (fun h i ↦ ((h i).1, (h i).2.2)) (by fun_prop) p0 := alg.p0 -variable {Ω E : Type*} [mΩ : MeasurableSpace Ω] [mE : MeasurableSpace E] - -def IsPOAlgEnvSeq - [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - [StandardBorelSpace E] [Nonempty E] - (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (E' : ℕ → Ω → E) - (alg : Algorithm α R) (env : Environment α (E × R)) (P : Measure Ω) [IsFiniteMeasure P] - := IsAlgEnvSeq A (fun n ω ↦ (E' n ω, R' n ω)) (alg.prod_left E) env P - end Learning end Algorithm @@ -59,8 +51,30 @@ instance isProbabilityMeasure_bayesTrajMeasure (Q : Measure E) [IsProbabilityMea unfold bayesTrajMeasure infer_instance +lemma isAlgEnvSeq_bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) + [IsMarkovKernel κ] [StandardBorelSpace α] [Nonempty α] + [StandardBorelSpace R] [StandardBorelSpace E] [Nonempty E] [Nonempty R] (alg : Algorithm α R) : + IsAlgEnvSeq IT.action IT.reward (alg.prod_left E) (bayesStationaryEnv Q κ) + (bayesTrajMeasure Q κ alg) := + IT.isAlgEnvSeq_trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) + end StationaryEnv +namespace POTraj + +variable {α R E : Type*} + +def action (n : ℕ) (ω : ℕ → α × (E × R)) : α := (ω n).1 + +def reward (n : ℕ) (ω : ℕ → α × (E × R)) : R := (ω n).2.2 + +def hist (n : ℕ) (ω : ℕ → α × (E × R)) : Iic n → α × R := + fun i ↦ (action i ω, reward i ω) + +def latent (n : ℕ) (ω : ℕ → α × (E × R)) : E := (ω n).2.1 + +end POTraj + section Uniform -- BanditAlgorithms/Uniform.lean variable {K : ℕ} (hK : 0 < K) @@ -77,7 +91,6 @@ section ThompsonSampling -- BanditAlgorithms/TS.lean variable {K : ℕ} variable {E : Type*} [mE : MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] -variable {Ω : Type*} [mΩ : MeasurableSpace Ω] variable (hK : 0 < K) variable (Q : Measure E) [IsProbabilityMeasure Q] @@ -85,11 +98,7 @@ variable (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] noncomputable def tsPosterior (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) E := - let P := bayesTrajMeasure Q κ (uniformAlgorithm hK) - let E' : ℕ → (ℕ → (Fin K) × (E × ℝ)) → E := fun n ω ↦ (ω n).2.1 - let A : ℕ → (ℕ → (Fin K) × (E × ℝ)) → (Fin K) := fun n ω ↦ (ω n).1 - let R' : ℕ → (ℕ → (Fin K) × (E × ℝ)) → ℝ := fun n ω ↦ (ω n).2.2 - condDistrib (E' 0) (IsAlgEnvSeq.hist A R' n) P + condDistrib (POTraj.latent 0) (POTraj.hist n) (bayesTrajMeasure Q κ (uniformAlgorithm hK)) noncomputable def isMarkovKernel_tsPosterior (n : ℕ) : IsMarkovKernel (tsPosterior hK Q κ n) := by From 7188eb39f08e411e8de3dac9e88fce8d6b9c2052 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 20 Jan 2026 14:32:25 +0000 Subject: [PATCH 006/155] Organize TS --- LeanBandits.lean | 3 + LeanBandits/BanditAlgorithms/TS.lean | 85 +------------------ LeanBandits/BanditAlgorithms/Uniform.lean | 18 ++++ LeanBandits/SequentialLearning/Algorithm.lean | 5 ++ .../BayesStationaryEnv.lean | 54 ++++++++++++ 5 files changed, 82 insertions(+), 83 deletions(-) create mode 100644 LeanBandits/BanditAlgorithms/Uniform.lean create mode 100644 LeanBandits/SequentialLearning/BayesStationaryEnv.lean diff --git a/LeanBandits.lean b/LeanBandits.lean index c1a8f31c..fa049b2c 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -5,6 +5,8 @@ import LeanBandits.Bandit.SumRewards import LeanBandits.BanditAlgorithms.AuxSums import LeanBandits.BanditAlgorithms.ETC import LeanBandits.BanditAlgorithms.UCB +import LeanBandits.BanditAlgorithms.Uniform +import LeanBandits.BanditAlgorithms.TS import LeanBandits.ForMathlib.CondDistrib import LeanBandits.ForMathlib.CondIndepFun import LeanBandits.ForMathlib.HasCondDistrib @@ -18,6 +20,7 @@ import LeanBandits.ForMathlib.StandardBorel import LeanBandits.ForMathlib.SubGaussian import LeanBandits.ForMathlib.Traj import LeanBandits.SequentialLearning.Algorithm +import LeanBandits.SequentialLearning.BayesStationaryEnv import LeanBandits.SequentialLearning.Deterministic import LeanBandits.SequentialLearning.FiniteActions import LeanBandits.SequentialLearning.IonescuTulceaSpace diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 2420da49..1bb23da6 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -3,92 +3,13 @@ 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 Mathlib.Probability.Distributions.Uniform import LeanBandits.ForMathlib.MeasurableArgMax -import LeanBandits.SequentialLearning.Algorithm -import LeanBandits.SequentialLearning.IonescuTulceaSpace +import LeanBandits.BanditAlgorithms.Uniform +import LeanBandits.SequentialLearning.BayesStationaryEnv open MeasureTheory ProbabilityTheory Finset open Learning -section Algorithm -- SequentialLearning/Algorithm.lean - -variable {α R : Type*} [mα : MeasurableSpace α] [mR : MeasurableSpace R] - -namespace Learning - -def Algorithm.prod_left (E : Type*) [MeasurableSpace E] (alg : Algorithm α R) : - Algorithm α (E × R) where - policy n := (alg.policy n).comap (fun h i ↦ ((h i).1, (h i).2.2)) (by fun_prop) - p0 := alg.p0 - -end Learning - -end Algorithm - -section StationaryEnv -- SequentialLearning/StationaryEnv.lean - -variable {α R E : Type*} [mα : MeasurableSpace α] [mR : MeasurableSpace R] [mE : MeasurableSpace E] - -noncomputable -def bayesStationaryEnv (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) - [IsMarkovKernel κ] : Environment α (E × R) where - feedback n := - let g : (Iic n → α × (E × R)) × α → (α × E) := fun (h, a) => (a, (h ⟨0, by simp⟩).2.1) - (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) - h_feedback := inferInstance - ν0 := (Kernel.const α Q) ⊗ₖ κ - hp0 := Kernel.IsMarkovKernel.compProd _ _ - -noncomputable -def bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) - [IsMarkovKernel κ] (alg : Algorithm α R) := - trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) - -instance isProbabilityMeasure_bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] - (κ : Kernel (α × E) R) [IsMarkovKernel κ] (alg : Algorithm α R) : - IsProbabilityMeasure (bayesTrajMeasure Q κ alg) := by - unfold bayesTrajMeasure - infer_instance - -lemma isAlgEnvSeq_bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) - [IsMarkovKernel κ] [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [StandardBorelSpace E] [Nonempty E] [Nonempty R] (alg : Algorithm α R) : - IsAlgEnvSeq IT.action IT.reward (alg.prod_left E) (bayesStationaryEnv Q κ) - (bayesTrajMeasure Q κ alg) := - IT.isAlgEnvSeq_trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) - -end StationaryEnv - -namespace POTraj - -variable {α R E : Type*} - -def action (n : ℕ) (ω : ℕ → α × (E × R)) : α := (ω n).1 - -def reward (n : ℕ) (ω : ℕ → α × (E × R)) : R := (ω n).2.2 - -def hist (n : ℕ) (ω : ℕ → α × (E × R)) : Iic n → α × R := - fun i ↦ (action i ω, reward i ω) - -def latent (n : ℕ) (ω : ℕ → α × (E × R)) : E := (ω n).2.1 - -end POTraj - -section Uniform -- BanditAlgorithms/Uniform.lean - -variable {K : ℕ} (hK : 0 < K) - -noncomputable -def uniformAlgorithm : Algorithm (Fin K) ℝ := - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - { policy _ := Kernel.const _ (PMF.uniformOfFintype (Fin K)).toMeasure - p0 := (PMF.uniformOfFintype (Fin K)).toMeasure } - -end Uniform - -section ThompsonSampling -- BanditAlgorithms/TS.lean - variable {K : ℕ} variable {E : Type*} [mE : MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] @@ -135,5 +56,3 @@ def tsAlgorithm : Algorithm (Fin K) ℝ where h_policy := isMarkovKernel_tsPolicy hK Q κ p0 := tsInitPolicy hK Q κ hp0 := isProbabilityMeasure_tsInitPolicy hK Q κ - -end ThompsonSampling diff --git a/LeanBandits/BanditAlgorithms/Uniform.lean b/LeanBandits/BanditAlgorithms/Uniform.lean new file mode 100644 index 00000000..16cd3a58 --- /dev/null +++ b/LeanBandits/BanditAlgorithms/Uniform.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, Paulo Rauber +-/ +import Mathlib.Probability.Distributions.Uniform +import LeanBandits.SequentialLearning.Algorithm + +open ProbabilityTheory +open Learning + +variable {K : ℕ} (hK : 0 < K) + +noncomputable +def uniformAlgorithm : Algorithm (Fin K) ℝ := + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + { policy _ := Kernel.const _ (PMF.uniformOfFintype (Fin K)).toMeasure + p0 := (PMF.uniformOfFintype (Fin K)).toMeasure } diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index df3c528b..74fb6b17 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -30,6 +30,11 @@ structure Algorithm (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] wher instance (alg : Algorithm α R) (n : ℕ) : IsMarkovKernel (alg.policy n) := alg.h_policy n instance (alg : Algorithm α R) : IsProbabilityMeasure alg.p0 := alg.hp0 +def Algorithm.prod_left (E : Type*) [MeasurableSpace E] (alg : Algorithm α R) : + Algorithm α (E × R) where + policy n := (alg.policy n).comap (fun h i ↦ ((h i).1, (h i).2.2)) (by fun_prop) + p0 := alg.p0 + /-- A stochastic environment. -/ structure Environment (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] where /-- Distribution of the next observation as function of the past history. -/ diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean new file mode 100644 index 00000000..ea5c5cb1 --- /dev/null +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -0,0 +1,54 @@ +/- +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.IonescuTulceaSpace + +open MeasureTheory ProbabilityTheory Finset +open Learning + +variable {α R E : Type*} [mα : MeasurableSpace α] [mR : MeasurableSpace R] [mE : MeasurableSpace E] + +noncomputable +def bayesStationaryEnv (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) + [IsMarkovKernel κ] : Environment α (E × R) where + feedback n := + let g : (Iic n → α × (E × R)) × α → (α × E) := fun (h, a) => (a, (h ⟨0, by simp⟩).2.1) + (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) + h_feedback := inferInstance + ν0 := (Kernel.const α Q) ⊗ₖ κ + hp0 := Kernel.IsMarkovKernel.compProd _ _ + +noncomputable +def bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) + [IsMarkovKernel κ] (alg : Algorithm α R) := + trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) + +instance isProbabilityMeasure_bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] + (κ : Kernel (α × E) R) [IsMarkovKernel κ] (alg : Algorithm α R) : + IsProbabilityMeasure (bayesTrajMeasure Q κ alg) := by + unfold bayesTrajMeasure + infer_instance + +lemma isAlgEnvSeq_bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) + [IsMarkovKernel κ] [StandardBorelSpace α] [Nonempty α] + [StandardBorelSpace R] [StandardBorelSpace E] [Nonempty E] [Nonempty R] (alg : Algorithm α R) : + IsAlgEnvSeq IT.action IT.reward (alg.prod_left E) (bayesStationaryEnv Q κ) + (bayesTrajMeasure Q κ alg) := + IT.isAlgEnvSeq_trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) + +namespace POTraj + +variable {α R E : Type*} + +def action (n : ℕ) (ω : ℕ → α × (E × R)) : α := (ω n).1 + +def reward (n : ℕ) (ω : ℕ → α × (E × R)) : R := (ω n).2.2 + +def hist (n : ℕ) (ω : ℕ → α × (E × R)) : Iic n → α × R := + fun i ↦ (action i ω, reward i ω) + +def latent (n : ℕ) (ω : ℕ → α × (E × R)) : E := (ω n).2.1 + +end POTraj From 1d6292259bdcb84f2a4f2e7cea04bb384048b713 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 20 Jan 2026 14:43:43 +0000 Subject: [PATCH 007/155] Reorder imports --- LeanBandits.lean | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/LeanBandits.lean b/LeanBandits.lean index fa049b2c..f84e8da5 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -4,9 +4,9 @@ import LeanBandits.Bandit.RewardByCountMeasure import LeanBandits.Bandit.SumRewards import LeanBandits.BanditAlgorithms.AuxSums import LeanBandits.BanditAlgorithms.ETC +import LeanBandits.BanditAlgorithms.TS import LeanBandits.BanditAlgorithms.UCB import LeanBandits.BanditAlgorithms.Uniform -import LeanBandits.BanditAlgorithms.TS import LeanBandits.ForMathlib.CondDistrib import LeanBandits.ForMathlib.CondIndepFun import LeanBandits.ForMathlib.HasCondDistrib From 904adff5a941db1db504d0a315856871f0f09ef8 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 20 Jan 2026 15:12:25 +0000 Subject: [PATCH 008/155] Fix namespaces --- LeanBandits/BanditAlgorithms/TS.lean | 7 ++++-- LeanBandits/BanditAlgorithms/Uniform.lean | 7 ++++-- .../BayesStationaryEnv.lean | 23 ++++++++++--------- 3 files changed, 22 insertions(+), 15 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 1bb23da6..273bc6a7 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -7,8 +7,9 @@ import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.BanditAlgorithms.Uniform import LeanBandits.SequentialLearning.BayesStationaryEnv -open MeasureTheory ProbabilityTheory Finset -open Learning +open MeasureTheory ProbabilityTheory Finset Learning + +namespace Bandits variable {K : ℕ} variable {E : Type*} [mE : MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] @@ -56,3 +57,5 @@ def tsAlgorithm : Algorithm (Fin K) ℝ where h_policy := isMarkovKernel_tsPolicy hK Q κ p0 := tsInitPolicy hK Q κ hp0 := isProbabilityMeasure_tsInitPolicy hK Q κ + +end Bandits diff --git a/LeanBandits/BanditAlgorithms/Uniform.lean b/LeanBandits/BanditAlgorithms/Uniform.lean index 16cd3a58..9e67deee 100644 --- a/LeanBandits/BanditAlgorithms/Uniform.lean +++ b/LeanBandits/BanditAlgorithms/Uniform.lean @@ -6,8 +6,9 @@ Authors: Rémy Degenne, Paulo Rauber import Mathlib.Probability.Distributions.Uniform import LeanBandits.SequentialLearning.Algorithm -open ProbabilityTheory -open Learning +open ProbabilityTheory Learning + +namespace Bandits variable {K : ℕ} (hK : 0 < K) @@ -16,3 +17,5 @@ def uniformAlgorithm : Algorithm (Fin K) ℝ := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK { policy _ := Kernel.const _ (PMF.uniformOfFintype (Fin K)).toMeasure p0 := (PMF.uniformOfFintype (Fin K)).toMeasure } + +end Bandits diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index ea5c5cb1..9ba543e7 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -6,13 +6,14 @@ Authors: Rémy Degenne, Paulo Rauber import LeanBandits.SequentialLearning.IonescuTulceaSpace open MeasureTheory ProbabilityTheory Finset -open Learning + +namespace Learning variable {α R E : Type*} [mα : MeasurableSpace α] [mR : MeasurableSpace R] [mE : MeasurableSpace E] +variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] noncomputable -def bayesStationaryEnv (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) - [IsMarkovKernel κ] : Environment α (E × R) where +def bayesStationaryEnv : Environment α (E × R) where feedback n := let g : (Iic n → α × (E × R)) × α → (α × E) := fun (h, a) => (a, (h ⟨0, by simp⟩).2.1) (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) @@ -21,20 +22,18 @@ def bayesStationaryEnv (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α hp0 := Kernel.IsMarkovKernel.compProd _ _ noncomputable -def bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) - [IsMarkovKernel κ] (alg : Algorithm α R) := +def bayesTrajMeasure (alg : Algorithm α R) := trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) -instance isProbabilityMeasure_bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] - (κ : Kernel (α × E) R) [IsMarkovKernel κ] (alg : Algorithm α R) : +instance isProbabilityMeasure_bayesTrajMeasure (alg : Algorithm α R) : IsProbabilityMeasure (bayesTrajMeasure Q κ alg) := by unfold bayesTrajMeasure infer_instance -lemma isAlgEnvSeq_bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) - [IsMarkovKernel κ] [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [StandardBorelSpace E] [Nonempty E] [Nonempty R] (alg : Algorithm α R) : - IsAlgEnvSeq IT.action IT.reward (alg.prod_left E) (bayesStationaryEnv Q κ) +lemma isAlgEnvSeq_bayesTrajMeasure + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + [StandardBorelSpace E] [Nonempty E] (alg : Algorithm α R) : + IsAlgEnvSeq IT.action IT.reward (alg.prod_left E) (bayesStationaryEnv Q κ) (bayesTrajMeasure Q κ alg) := IT.isAlgEnvSeq_trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) @@ -52,3 +51,5 @@ def hist (n : ℕ) (ω : ℕ → α × (E × R)) : Iic n → α × R := def latent (n : ℕ) (ω : ℕ → α × (E × R)) : E := (ω n).2.1 end POTraj + +end Learning From b5d5d1201abab7a21eca257f5a54e948be907c93 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 20 Jan 2026 15:53:31 +0000 Subject: [PATCH 009/155] Define posterior --- LeanBandits/BanditAlgorithms/TS.lean | 4 +- .../BayesStationaryEnv.lean | 49 ++++++++++--------- 2 files changed, 27 insertions(+), 26 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 273bc6a7..9e5e8bd5 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -20,11 +20,11 @@ variable (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] noncomputable def tsPosterior (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) E := - condDistrib (POTraj.latent 0) (POTraj.hist n) (bayesTrajMeasure Q κ (uniformAlgorithm hK)) + Learning.Bayes.posterior Q κ n (uniformAlgorithm hK) noncomputable def isMarkovKernel_tsPosterior (n : ℕ) : IsMarkovKernel (tsPosterior hK Q κ n) := by - unfold tsPosterior + unfold tsPosterior Bayes.posterior infer_instance noncomputable diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 9ba543e7..6733a1c4 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -7,49 +7,50 @@ import LeanBandits.SequentialLearning.IonescuTulceaSpace open MeasureTheory ProbabilityTheory Finset -namespace Learning +namespace Learning.Bayes variable {α R E : Type*} [mα : MeasurableSpace α] [mR : MeasurableSpace R] [mE : MeasurableSpace E] variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] noncomputable -def bayesStationaryEnv : Environment α (E × R) where +def StationaryEnv : Environment α (E × R) where feedback n := - let g : (Iic n → α × (E × R)) × α → (α × E) := fun (h, a) => (a, (h ⟨0, by simp⟩).2.1) + let g : (Iic n → α × E × R) × α → α × E := fun (h, a) => (a, (h ⟨0, by simp⟩).2.1) (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) h_feedback := inferInstance ν0 := (Kernel.const α Q) ⊗ₖ κ hp0 := Kernel.IsMarkovKernel.compProd _ _ noncomputable -def bayesTrajMeasure (alg : Algorithm α R) := - trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) +def trajMeasure (alg : Algorithm α R) := + Learning.trajMeasure (alg.prod_left E) (StationaryEnv Q κ) -instance isProbabilityMeasure_bayesTrajMeasure (alg : Algorithm α R) : - IsProbabilityMeasure (bayesTrajMeasure Q κ alg) := by - unfold bayesTrajMeasure +instance (alg : Algorithm α R) : IsProbabilityMeasure (trajMeasure Q κ alg) := by + unfold trajMeasure infer_instance -lemma isAlgEnvSeq_bayesTrajMeasure - [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - [StandardBorelSpace E] [Nonempty E] (alg : Algorithm α R) : - IsAlgEnvSeq IT.action IT.reward (alg.prod_left E) (bayesStationaryEnv Q κ) - (bayesTrajMeasure Q κ alg) := - IT.isAlgEnvSeq_trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) +def action (n : ℕ) (ω : ℕ → α × E × R) : α := (ω n).1 -namespace POTraj +def reward (n : ℕ) (ω : ℕ → α × E × R) : R := (ω n).2.2 -variable {α R E : Type*} +def hist (n : ℕ) (ω : ℕ → α × E × R) : Iic n → α × R := fun i ↦ (action i ω, reward i ω) -def action (n : ℕ) (ω : ℕ → α × (E × R)) : α := (ω n).1 +def env (ω : ℕ → α × E × R) : E := (ω 0).2.1 -def reward (n : ℕ) (ω : ℕ → α × (E × R)) : R := (ω n).2.2 - -def hist (n : ℕ) (ω : ℕ → α × (E × R)) : Iic n → α × R := - fun i ↦ (action i ω, reward i ω) +noncomputable +def posterior [StandardBorelSpace E] [Nonempty E] (n : ℕ) (alg : Algorithm α R) : + Kernel (Iic n → α × R) E := + condDistrib env (hist n) (trajMeasure Q κ alg) -def latent (n : ℕ) (ω : ℕ → α × (E × R)) : E := (ω n).2.1 +instance [StandardBorelSpace E] [Nonempty E] (n : ℕ) (alg : Algorithm α R) : + IsMarkovKernel (posterior Q κ n alg) := by + unfold posterior + infer_instance -end POTraj +lemma isAlgEnvSeq_trajMeasure + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] [StandardBorelSpace E] + [Nonempty E] (alg : Algorithm α R) : + IsAlgEnvSeq action IT.reward (alg.prod_left E) (StationaryEnv Q κ) (trajMeasure Q κ alg) := + IT.isAlgEnvSeq_trajMeasure _ _ -end Learning +end Learning.Bayes From 13a5245ed159015be447f899756ecf0d231dceda Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 20 Jan 2026 16:03:39 +0000 Subject: [PATCH 010/155] Reorder arguments --- LeanBandits/BanditAlgorithms/TS.lean | 2 +- LeanBandits/SequentialLearning/BayesStationaryEnv.lean | 6 +++--- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 9e5e8bd5..0ca61a81 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -20,7 +20,7 @@ variable (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] noncomputable def tsPosterior (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) E := - Learning.Bayes.posterior Q κ n (uniformAlgorithm hK) + Learning.Bayes.posterior Q κ (uniformAlgorithm hK) n noncomputable def isMarkovKernel_tsPosterior (n : ℕ) : IsMarkovKernel (tsPosterior hK Q κ n) := by diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 6733a1c4..c73c9d95 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -38,12 +38,12 @@ def hist (n : ℕ) (ω : ℕ → α × E × R) : Iic n → α × R := fun i ↦ def env (ω : ℕ → α × E × R) : E := (ω 0).2.1 noncomputable -def posterior [StandardBorelSpace E] [Nonempty E] (n : ℕ) (alg : Algorithm α R) : +def posterior [StandardBorelSpace E] [Nonempty E] (alg : Algorithm α R) (n : ℕ) : Kernel (Iic n → α × R) E := condDistrib env (hist n) (trajMeasure Q κ alg) -instance [StandardBorelSpace E] [Nonempty E] (n : ℕ) (alg : Algorithm α R) : - IsMarkovKernel (posterior Q κ n alg) := by +instance [StandardBorelSpace E] [Nonempty E] (alg : Algorithm α R) (n : ℕ) : + IsMarkovKernel (posterior Q κ alg n) := by unfold posterior infer_instance From e1a0b277f5e36cbaec63a1c6df7602893157e242 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 21 Jan 2026 09:27:34 +0000 Subject: [PATCH 011/155] Adopt suggestions --- LeanBandits/BanditAlgorithms/TS.lean | 2 +- LeanBandits/BanditAlgorithms/Uniform.lean | 12 +++++++----- .../SequentialLearning/BayesStationaryEnv.lean | 9 ++------- 3 files changed, 10 insertions(+), 13 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 0ca61a81..0b8db08c 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -1,5 +1,5 @@ /- -Copyright (c) 2025 Rémy Degenne. All rights reserved. +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 -/ diff --git a/LeanBandits/BanditAlgorithms/Uniform.lean b/LeanBandits/BanditAlgorithms/Uniform.lean index 9e67deee..1827f695 100644 --- a/LeanBandits/BanditAlgorithms/Uniform.lean +++ b/LeanBandits/BanditAlgorithms/Uniform.lean @@ -1,12 +1,12 @@ /- -Copyright (c) 2025 Rémy Degenne. All rights reserved. +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 Mathlib.Probability.Distributions.Uniform +import Mathlib.Probability.UniformOn import LeanBandits.SequentialLearning.Algorithm -open ProbabilityTheory Learning +open MeasureTheory ProbabilityTheory Learning namespace Bandits @@ -15,7 +15,9 @@ variable {K : ℕ} (hK : 0 < K) noncomputable def uniformAlgorithm : Algorithm (Fin K) ℝ := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - { policy _ := Kernel.const _ (PMF.uniformOfFintype (Fin K)).toMeasure - p0 := (PMF.uniformOfFintype (Fin K)).toMeasure } + have : IsProbabilityMeasure (uniformOn (Set.univ : Set (Fin K))) := + uniformOn_isProbabilityMeasure Set.finite_univ Set.univ_nonempty + { policy _ := Kernel.const _ (uniformOn Set.univ) + p0 := (uniformOn Set.univ) } end Bandits diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index c73c9d95..cce5c782 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -1,5 +1,5 @@ /- -Copyright (c) 2025 Rémy Degenne. All rights reserved. +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 -/ @@ -17,17 +17,12 @@ def StationaryEnv : Environment α (E × R) where feedback n := let g : (Iic n → α × E × R) × α → α × E := fun (h, a) => (a, (h ⟨0, by simp⟩).2.1) (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) - h_feedback := inferInstance ν0 := (Kernel.const α Q) ⊗ₖ κ - hp0 := Kernel.IsMarkovKernel.compProd _ _ noncomputable def trajMeasure (alg : Algorithm α R) := Learning.trajMeasure (alg.prod_left E) (StationaryEnv Q κ) - -instance (alg : Algorithm α R) : IsProbabilityMeasure (trajMeasure Q κ alg) := by - unfold trajMeasure - infer_instance +deriving IsProbabilityMeasure def action (n : ℕ) (ω : ℕ → α × E × R) : α := (ω n).1 From 564e904f3fdc6e0439e41e64ff36f5808d8935b9 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 21 Jan 2026 10:22:46 +0000 Subject: [PATCH 012/155] Add documentation --- LeanBandits/BanditAlgorithms/TS.lean | 18 +++++++++++------- LeanBandits/BanditAlgorithms/Uniform.lean | 3 ++- LeanBandits/SequentialLearning/Algorithm.lean | 2 ++ .../SequentialLearning/BayesStationaryEnv.lean | 15 +++++++++++++++ 4 files changed, 30 insertions(+), 8 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 0b8db08c..5bc72135 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -18,44 +18,48 @@ variable (hK : 0 < K) variable (Q : Measure E) [IsProbabilityMeasure Q] variable (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] +/-- The posterior over "environments" for every given history for TS. Note that we pretend that the +data was generated by an algorithm that chooses actions uniformly at random to avoid circularity in +the definition of `tsAlgorithm`. -/ noncomputable def tsPosterior (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) E := Learning.Bayes.posterior Q κ (uniformAlgorithm hK) n -noncomputable -def isMarkovKernel_tsPosterior (n : ℕ) : IsMarkovKernel (tsPosterior hK Q κ n) := by +instance (n : ℕ) : IsMarkovKernel (tsPosterior hK Q κ n) := by unfold tsPosterior Bayes.posterior infer_instance +/-- The distribution over actions for every given history for TS. -/ noncomputable def tsPolicy (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK (tsPosterior hK Q κ n).map (measurableArgmax (fun e k ↦ (κ (k, e))[id])) -def isMarkovKernel_tsPolicy (n : ℕ) : IsMarkovKernel (tsPolicy hK Q κ n) := by - have : IsMarkovKernel (tsPosterior hK Q κ n) := isMarkovKernel_tsPosterior hK Q κ n +instance (n : ℕ) : IsMarkovKernel (tsPolicy hK Q κ n) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK apply Kernel.IsMarkovKernel.map exact measurable_measurableArgmax fun k => (stronglyMeasurable_id.integral_kernel (κ := κ.comap (k, ·) (by fun_prop))).measurable +/-- The initial distribution over actions for TS. -/ noncomputable def tsInitPolicy : Measure (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK Q.map (measurableArgmax (fun e k ↦ (κ (k, e))[id])) -def isProbabilityMeasure_tsInitPolicy : IsProbabilityMeasure (tsInitPolicy hK Q κ) := by +instance : IsProbabilityMeasure (tsInitPolicy hK Q κ) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK apply Measure.isProbabilityMeasure_map apply Measurable.aemeasurable exact (measurable_measurableArgmax fun k => (stronglyMeasurable_id.integral_kernel (κ := κ.comap (k, ·) (by fun_prop))).measurable) +/-- The Thompson Sampling (TS) algorithm: actions are chosen according to the probability that they +are optimal given prior knowledge represented by a prior distribution `Q` and a data generation +model represented by a kernel `κ`. -/ noncomputable def tsAlgorithm : Algorithm (Fin K) ℝ where policy := tsPolicy hK Q κ - h_policy := isMarkovKernel_tsPolicy hK Q κ p0 := tsInitPolicy hK Q κ - hp0 := isProbabilityMeasure_tsInitPolicy hK Q κ end Bandits diff --git a/LeanBandits/BanditAlgorithms/Uniform.lean b/LeanBandits/BanditAlgorithms/Uniform.lean index 1827f695..cd054d93 100644 --- a/LeanBandits/BanditAlgorithms/Uniform.lean +++ b/LeanBandits/BanditAlgorithms/Uniform.lean @@ -12,12 +12,13 @@ namespace Bandits variable {K : ℕ} (hK : 0 < K) +/-- The Uniform algorithm: actions are chosen uniformly at random. -/ noncomputable def uniformAlgorithm : Algorithm (Fin K) ℝ := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK have : IsProbabilityMeasure (uniformOn (Set.univ : Set (Fin K))) := uniformOn_isProbabilityMeasure Set.finite_univ Set.univ_nonempty { policy _ := Kernel.const _ (uniformOn Set.univ) - p0 := (uniformOn Set.univ) } + p0 := uniformOn Set.univ } end Bandits diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index 74fb6b17..17b60d56 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -30,6 +30,8 @@ structure Algorithm (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] wher instance (alg : Algorithm α R) (n : ℕ) : IsMarkovKernel (alg.policy n) := alg.h_policy n instance (alg : Algorithm α R) : IsProbabilityMeasure alg.p0 := alg.hp0 +/-- An algorithm that receives observations in `E × R` created form an algorithm that receives +observations in `R` by ignoring the additional information. -/ def Algorithm.prod_left (E : Type*) [MeasurableSpace E] (alg : Algorithm α R) : Algorithm α (E × R) where policy n := (alg.policy n).comap (fun h i ↦ ((h i).1, (h i).2.2)) (by fun_prop) diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index cce5c782..4dc91692 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -12,6 +12,13 @@ namespace Learning.Bayes variable {α R E : Type*} [mα : MeasurableSpace α] [mR : MeasurableSpace R] [mE : MeasurableSpace E] variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] +/-- Given a prior distribution `Q` over "environments" and a kernel `k` that defines a reward +distribution `κ (a, e)` for each action `a : α` and "environment" `e : E`, a StationaryEnv +represents an environment (with an observation space `E × R`) that draws an environment `e : E` at +the very first step which, together with `k`, defines how the bandit process behaves. Because the +"environment" `e` is repeated at every step and reveals the best arm, it only makes sense to study +algorithms that ignore it, which is why `trajMeasure` is created from an algorithm whose +observation space is just `R` (rather than `E × R`). -/ noncomputable def StationaryEnv : Environment α (E × R) where feedback n := @@ -19,19 +26,27 @@ def StationaryEnv : Environment α (E × R) where (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) ν0 := (Kernel.const α Q) ⊗ₖ κ +/-- Measure on the sequence of actions and observations (in `E × R`) generated by the +algorithm/environment. -/ noncomputable def trajMeasure (alg : Algorithm α R) := Learning.trajMeasure (alg.prod_left E) (StationaryEnv Q κ) deriving IsProbabilityMeasure +/-- `action n` is the action pulled at time `n`. -/ def action (n : ℕ) (ω : ℕ → α × E × R) : α := (ω n).1 +/-- `reward n` is the reward at time `n`. -/ def reward (n : ℕ) (ω : ℕ → α × E × R) : R := (ω n).2.2 +/-- `hist n` is the (observable) history up to time `n`. -/ def hist (n : ℕ) (ω : ℕ → α × E × R) : Iic n → α × R := fun i ↦ (action i ω, reward i ω) +/-- `env` is the "environment" distributed according to `Q` that, together with the kernel `k`, +defines how the rewards are generated given actions. -/ def env (ω : ℕ → α × E × R) : E := (ω 0).2.1 +/-- The posterior over "environments" for every given history (for a fixed algorithm). -/ noncomputable def posterior [StandardBorelSpace E] [Nonempty E] (alg : Algorithm α R) (n : ℕ) : Kernel (Iic n → α × R) E := From 36ff1520ef45e827a4d7c2894ad6b2e9f18b9653 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 21 Jan 2026 16:13:42 +0000 Subject: [PATCH 013/155] Add basic properties of BayesStationaryEnv --- LeanBandits/ForMathlib/HasCondDistrib.lean | 63 +++++++++ .../BayesStationaryEnv.lean | 124 +++++++++++++++++- 2 files changed, 184 insertions(+), 3 deletions(-) diff --git a/LeanBandits/ForMathlib/HasCondDistrib.lean b/LeanBandits/ForMathlib/HasCondDistrib.lean index a151124e..19dfec6d 100644 --- a/LeanBandits/ForMathlib/HasCondDistrib.lean +++ b/LeanBandits/ForMathlib/HasCondDistrib.lean @@ -182,6 +182,19 @@ lemma hasCondDistrib_prod_right_iff [IsFiniteMeasure μ] [IsFiniteKernel κ] (X rw [← Measure.map_prod_map _ _ (by fun_prop) (by fun_prop), Measure.map_id, Measure.map_dirac (by fun_prop)] +-- Revise (Claude) +lemma HasCondDistrib.hasLaw_of_const {Q : Measure Ω} + [IsProbabilityMeasure μ] [IsProbabilityMeasure Q] + (h : HasCondDistrib Y X (Kernel.const β Q) μ) : HasLaw Y Q μ where + aemeasurable := h.aemeasurable_fst + map_eq := by + have : IsProbabilityMeasure (μ.map X) := Measure.isProbabilityMeasure_map h.aemeasurable_snd + rw [← Measure.snd_prod (μ := μ.map X) (ν := Q), ← Measure.compProd_const, + ← (condDistrib_ae_eq_iff_measure_eq_compProd X h.aemeasurable_fst _).1 h.condDistrib_eq, + Measure.snd, AEMeasurable.map_map_of_aemeasurable (by fun_prop) + (h.aemeasurable_snd.prodMk h.aemeasurable_fst)] + rfl + lemma HasLaw.prod_of_hasCondDistrib {P : Measure β} [IsFiniteMeasure μ] [IsSFiniteKernel κ] (h1 : HasLaw X P μ) (h2 : HasCondDistrib Y X κ μ) : HasLaw (fun ω ↦ (X ω, Y ω)) (P ⊗ₘ κ) μ := by @@ -193,6 +206,32 @@ lemma HasLaw.prod_of_hasCondDistrib {P : Measure β} [IsFiniteMeasure μ] [IsSFi rw [← h1.map_eq] exact h2.condDistrib_eq +-- Revise (Claude) +lemma HasCondDistrib.of_compProd [IsFiniteMeasure μ] [IsFiniteKernel κ] + {Z : α → Ω'} {η : Kernel (β × Ω) Ω'} [IsMarkovKernel η] + (h : HasCondDistrib (fun ω ↦ (Y ω, Z ω)) X (κ ⊗ₖ η) μ) : + HasCondDistrib Z (fun ω ↦ (X ω, Y ω)) η μ := by + have hY : AEMeasurable Y μ := h.aemeasurable_fst.fst + have hZ : AEMeasurable Z μ := h.aemeasurable_fst.snd + have hX : AEMeasurable X μ := h.aemeasurable_snd + refine ⟨hZ, by fun_prop, ?_⟩ + have h_eq := h.condDistrib_eq + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop) _] at h_eq ⊢ + have h_assoc : (μ.map X ⊗ₘ κ ⊗ₘ η).map MeasurableEquiv.prodAssoc = μ.map X ⊗ₘ (κ ⊗ₖ η) := + Measure.compProd_assoc' + calc μ.map (fun ω ↦ ((X ω, Y ω), Z ω)) + _ = (μ.map (fun ω ↦ (X ω, Y ω, Z ω))).map MeasurableEquiv.prodAssoc.symm := by + rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)]; rfl + _ = (μ.map X ⊗ₘ (κ ⊗ₖ η)).map MeasurableEquiv.prodAssoc.symm := by rw [h_eq] + _ = μ.map X ⊗ₘ κ ⊗ₘ η := by + rw [← h_assoc, Measure.map_map (by fun_prop) (by fun_prop)] + simp only [MeasurableEquiv.symm_comp_self, Measure.map_id] + _ = μ.map (fun ω ↦ (X ω, Y ω)) ⊗ₘ η := by + have h_fst := h.fst + rw [Kernel.fst_compProd] at h_fst + have h_fst_eq := (condDistrib_ae_eq_iff_measure_eq_compProd X hY _).mp h_fst.condDistrib_eq + rw [h_fst_eq] + lemma HasCondDistrib.prod [IsFiniteMeasure μ] [IsFiniteKernel κ] {Z : α → Ω'} {η : Kernel (β × Ω) Ω'} [IsFiniteKernel η] (h1 : HasCondDistrib Y X κ μ) (h2 : HasCondDistrib Z (fun ω ↦ (X ω, Y ω)) η μ) : @@ -216,4 +255,28 @@ lemma HasCondDistrib.prod [IsFiniteMeasure μ] [IsFiniteKernel κ] AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] rfl +-- Revise (Claude) +lemma HasCondDistrib.comp_left [IsFiniteMeasure μ] [IsFiniteKernel κ] {f : γ → β} + (hf : Measurable f) {Z : α → γ} (h : HasCondDistrib Y Z (κ.comap f hf) μ) : + HasCondDistrib Y (f ∘ Z) κ μ where + aemeasurable_fst := h.aemeasurable_fst + aemeasurable_snd := hf.comp_aemeasurable h.aemeasurable_snd + condDistrib_eq := by + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ h.aemeasurable_fst] + calc μ.map (fun ω ↦ ((f ∘ Z) ω, Y ω)) + _ = (μ.map (fun ω ↦ (Z ω, Y ω))).map (Prod.map f id) := by + rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) + (h.aemeasurable_snd.prodMk h.aemeasurable_fst)]; rfl + _ = (μ.map Z ⊗ₘ κ.comap f hf).map (Prod.map f id) := by + rw [(condDistrib_ae_eq_iff_measure_eq_compProd Z h.aemeasurable_fst _).mp h.condDistrib_eq] + _ = μ.map (f ∘ Z) ⊗ₘ κ := by + rw [← AEMeasurable.map_map_of_aemeasurable hf.aemeasurable h.aemeasurable_snd] + ext s hs + rw [Measure.map_apply (by fun_prop) hs, Measure.compProd_apply hs, + Measure.compProd_apply (hs.preimage (by fun_prop))] + rw [lintegral_map (Kernel.measurable_kernel_prodMk_left hs) hf] + refine lintegral_congr fun x ↦ ?_ + rw [Kernel.comap_apply] + congr 1 + end ProbabilityTheory diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 4dc91692..5681e3a5 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.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.SequentialLearning.IonescuTulceaSpace +import LeanBandits.Bandit.Regret open MeasureTheory ProbabilityTheory Finset @@ -46,6 +47,22 @@ def hist (n : ℕ) (ω : ℕ → α × E × R) : Iic n → α × R := fun i ↦ defines how the rewards are generated given actions. -/ def env (ω : ℕ → α × E × R) : E := (ω 0).2.1 +@[fun_prop] +lemma measurable_action (n : ℕ) : Measurable (@action α E R n) := by + unfold action; fun_prop + +@[fun_prop] +lemma measurable_reward (n : ℕ) : Measurable (@reward α E R n) := by + unfold reward; fun_prop + +@[fun_prop] +lemma measurable_env : Measurable (@env α E R) := by + unfold env; fun_prop + +@[fun_prop] +lemma measurable_hist (n : ℕ) : Measurable (@hist α E R n) := by + unfold hist; fun_prop + /-- The posterior over "environments" for every given history (for a fixed algorithm). -/ noncomputable def posterior [StandardBorelSpace E] [Nonempty E] (alg : Algorithm α R) (n : ℕ) : @@ -57,10 +74,111 @@ instance [StandardBorelSpace E] [Nonempty E] (alg : Algorithm α R) (n : ℕ) : unfold posterior infer_instance -lemma isAlgEnvSeq_trajMeasure - [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] [StandardBorelSpace E] - [Nonempty E] (alg : Algorithm α R) : +section Laws + +variable (alg : Algorithm α R) +variable [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] +variable [StandardBorelSpace E] [Nonempty E] + +lemma isAlgEnvSeq_trajMeasure : IsAlgEnvSeq action IT.reward (alg.prod_left E) (StationaryEnv Q κ) (trajMeasure Q κ alg) := IT.isAlgEnvSeq_trajMeasure _ _ +-- Revise (Claude) +lemma hasLaw_env : HasLaw env Q (trajMeasure Q κ alg) := by + apply HasCondDistrib.hasLaw_of_const + simpa [StationaryEnv] using (isAlgEnvSeq_trajMeasure Q κ alg).hasCondDistrib_reward_zero.fst + +-- Revise (Claude) +lemma hasCondDistrib_action (n : ℕ) : + HasCondDistrib (action (n + 1)) (hist n) (alg.policy n) (trajMeasure Q κ alg) := + ((isAlgEnvSeq_trajMeasure Q κ alg).hasCondDistrib_action n).comp_left (by fun_prop) + +-- Revise (Claude) +lemma hasCondDistrib_reward_zero : + HasCondDistrib (reward 0) (fun ω ↦ (action 0 ω, env ω)) κ (trajMeasure Q κ alg) := by + simpa [StationaryEnv] using + (isAlgEnvSeq_trajMeasure Q κ alg).hasCondDistrib_reward_zero.of_compProd + +-- Revise (Claude) +lemma hasCondDistrib_reward (n : ℕ) : + HasCondDistrib (reward (n + 1)) (fun ω ↦ (action (n + 1) ω, env ω)) κ + (trajMeasure Q κ alg) := by + have h := ((isAlgEnvSeq_trajMeasure Q κ alg).hasCondDistrib_reward n).snd + simp_rw [StationaryEnv, Kernel.snd_prod] at h + exact h.comp_left (by fun_prop) + +-- Revise (Claude) +lemma hasCondDistrib_action_env_hist (n : ℕ) : + HasCondDistrib (action (n + 1)) (fun ω ↦ (env ω, hist n ω)) + ((alg.policy n).prodMkLeft E) (trajMeasure Q κ alg) := by + have h := (isAlgEnvSeq_trajMeasure Q κ alg).hasCondDistrib_action n + simp only [Algorithm.prod_left, Kernel.prodMkLeft] at h ⊢ + have hf : Measurable (fun h : Iic n → α × E × R ↦ + ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2))) := by fun_prop + have h' : HasCondDistrib (action (n + 1)) (IsAlgEnvSeq.hist action IT.reward n) + (((alg.policy n).comap Prod.snd (by fun_prop)).comap + (fun h : Iic n → α × E × R ↦ ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2))) hf) + (trajMeasure Q κ alg) := by convert h using 2 + convert h'.comp_left hf using 2 + +-- Revise (Claude) +lemma hasCondDistrib_reward_hist (n : ℕ) : + HasCondDistrib (reward (n + 1)) (fun ω ↦ (hist n ω, action (n + 1) ω, env ω)) + (κ.prodMkLeft _) (trajMeasure Q κ alg) := by + have h := ((isAlgEnvSeq_trajMeasure Q κ alg).hasCondDistrib_reward n).snd + simp_rw [StationaryEnv, Kernel.snd_prod, Kernel.prodMkLeft] at h ⊢ + have hf : Measurable (fun (p : (Iic n → α × (E × R)) × α) ↦ + ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1)) := by fun_prop + have h' : HasCondDistrib (reward (n + 1)) + (fun ω ↦ (IsAlgEnvSeq.hist action IT.reward n ω, action (n + 1) ω)) + ((κ.comap Prod.snd (by fun_prop)).comap (fun (p : (Iic n → α × (E × R)) × α) ↦ + ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1)) hf) + (trajMeasure Q κ alg) := by convert h using 2 + convert h'.comp_left hf using 2 + +-- Revise (Claude) +lemma condIndepFun_action_env_hist (n : ℕ) : + action (n + 1) ⟂ᵢ[hist n, (by fun_prop); trajMeasure Q κ alg] env := + condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (by fun_prop) (by fun_prop) (by fun_prop) + (hasCondDistrib_action_env_hist Q κ alg n).condDistrib_eq + +-- Revise (Claude) +lemma condIndepFun_reward_hist (n : ℕ) : + reward (n + 1) + ⟂ᵢ[(fun ω ↦ (action (n + 1) ω, env ω)), (by fun_prop); trajMeasure Q κ alg] hist n := + condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (by fun_prop) (by fun_prop) (by fun_prop) + (hasCondDistrib_reward_hist Q κ alg n).condDistrib_eq + +end Laws + +section Regret + +variable (κ : Kernel (α × E) ℝ) + +noncomputable +def regretAt (t : ℕ) (ω : ℕ → α × E × ℝ) : ℝ := + Bandits.regret (κ.comap (·, env ω) (by fun_prop)) action t ω + +-- Revise (Claude) +lemma measurable_regretAt [Fintype α] (t : ℕ) : Measurable (regretAt κ t) := by + unfold regretAt Bandits.regret + have hmean : Measurable fun (p : α × E) ↦ (κ p)[id] := + stronglyMeasurable_id.integral_kernel.measurable + apply Measurable.sub + · apply Measurable.const_mul + have h : ∀ a, Measurable fun (ω : ℕ → α × E × ℝ) ↦ (κ (a, env ω))[id] := fun a ↦ + hmean.comp (measurable_const.prodMk measurable_env) + exact Measurable.iSup h + · apply Finset.measurable_sum + intro s _ + exact hmean.comp ((measurable_action s).prodMk measurable_env) + +variable (alg : Algorithm α ℝ) + +noncomputable +def regret [IsMarkovKernel κ] (t : ℕ) : ℝ := (trajMeasure Q κ alg)[regretAt κ t] + +end Regret + end Learning.Bayes From f56c8b73cd909e2bf9594dbfc1ad8434a244f192 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 23 Jan 2026 11:34:16 +0000 Subject: [PATCH 014/155] Remove IT space from BayesStationaryEnv (WIP) --- LeanBandits/BanditAlgorithms/TS.lean | 7 +- .../BayesStationaryEnv.lean | 170 +++++++++--------- 2 files changed, 89 insertions(+), 88 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 5bc72135..04457e1c 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -6,6 +6,7 @@ Authors: Rémy Degenne, Paulo Rauber import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.BanditAlgorithms.Uniform import LeanBandits.SequentialLearning.BayesStationaryEnv +import LeanBandits.SequentialLearning.IonescuTulceaSpace open MeasureTheory ProbabilityTheory Finset Learning @@ -23,10 +24,12 @@ data was generated by an algorithm that chooses actions uniformly at random to a the definition of `tsAlgorithm`. -/ noncomputable def tsPosterior (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) E := - Learning.Bayes.posterior Q κ (uniformAlgorithm hK) n + Learning.Bayes.IsAlgEnvSeq.posterior + (trajMeasure ((uniformAlgorithm hK).prod_left E) (Bayes.StationaryEnv Q κ)) + IT.action IT.reward n instance (n : ℕ) : IsMarkovKernel (tsPosterior hK Q κ n) := by - unfold tsPosterior Bayes.posterior + unfold tsPosterior Bayes.IsAlgEnvSeq.posterior infer_instance /-- The distribution over actions for every given history for TS. -/ diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 5681e3a5..8f5a897a 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -3,7 +3,6 @@ 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.SequentialLearning.IonescuTulceaSpace import LeanBandits.Bandit.Regret open MeasureTheory ProbabilityTheory Finset @@ -15,11 +14,10 @@ variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsM /-- Given a prior distribution `Q` over "environments" and a kernel `k` that defines a reward distribution `κ (a, e)` for each action `a : α` and "environment" `e : E`, a StationaryEnv -represents an environment (with an observation space `E × R`) that draws an environment `e : E` at +represents an environment (with an observation space `E × R`) that draws an "environment" `e : E` at the very first step which, together with `k`, defines how the bandit process behaves. Because the "environment" `e` is repeated at every step and reveals the best arm, it only makes sense to study -algorithms that ignore it, which is why `trajMeasure` is created from an algorithm whose -observation space is just `R` (rather than `E × R`). -/ +algorithms that ignore the information in `E` and just receive the information in `R`. -/ noncomputable def StationaryEnv : Environment α (E × R) where feedback n := @@ -27,158 +25,158 @@ def StationaryEnv : Environment α (E × R) where (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) ν0 := (Kernel.const α Q) ⊗ₖ κ -/-- Measure on the sequence of actions and observations (in `E × R`) generated by the -algorithm/environment. -/ -noncomputable -def trajMeasure (alg : Algorithm α R) := - Learning.trajMeasure (alg.prod_left E) (StationaryEnv Q κ) -deriving IsProbabilityMeasure +variable {Ω : Type*} [MeasurableSpace Ω] +variable (P : Measure Ω) [IsProbabilityMeasure P] +variable (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) -/-- `action n` is the action pulled at time `n`. -/ -def action (n : ℕ) (ω : ℕ → α × E × R) : α := (ω n).1 +namespace IsAlgEnvSeq -/-- `reward n` is the reward at time `n`. -/ -def reward (n : ℕ) (ω : ℕ → α × E × R) : R := (ω n).2.2 +def env (ω : Ω) : E := (R' 0 ω).1 -/-- `hist n` is the (observable) history up to time `n`. -/ -def hist (n : ℕ) (ω : ℕ → α × E × R) : Iic n → α × R := fun i ↦ (action i ω, reward i ω) +@[fun_prop] +lemma measurable_env (hR' : ∀ n, Measurable (R' n)) : Measurable (env R'):= (hR' 0).fst -/-- `env` is the "environment" distributed according to `Q` that, together with the kernel `k`, -defines how the rewards are generated given actions. -/ -def env (ω : ℕ → α × E × R) : E := (ω 0).2.1 +def reward (n : ℕ) (ω : Ω) : R := (R' n ω).2 @[fun_prop] -lemma measurable_action (n : ℕ) : Measurable (@action α E R n) := by - unfold action; fun_prop +lemma measurable_reward (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : Measurable (reward R' n) := by + unfold reward + fun_prop -@[fun_prop] -lemma measurable_reward (n : ℕ) : Measurable (@reward α E R n) := by - unfold reward; fun_prop +def hist (n : ℕ) (ω : Ω) : Iic n → α × R := fun i ↦ (A i ω, (R' i ω).2) @[fun_prop] -lemma measurable_env : Measurable (@env α E R) := by - unfold env; fun_prop +lemma measurable_hist (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : + Measurable (hist A R' n) := by + unfold hist + fun_prop -@[fun_prop] -lemma measurable_hist (n : ℕ) : Measurable (@hist α E R n) := by - unfold hist; fun_prop +variable [StandardBorelSpace α] [Nonempty α] +variable [StandardBorelSpace E] [Nonempty E] +variable [StandardBorelSpace R] [Nonempty R] /-- The posterior over "environments" for every given history (for a fixed algorithm). -/ noncomputable -def posterior [StandardBorelSpace E] [Nonempty E] (alg : Algorithm α R) (n : ℕ) : - Kernel (Iic n → α × R) E := - condDistrib env (hist n) (trajMeasure Q κ alg) +def posterior (n : ℕ) : Kernel (Iic n → α × R) E := + condDistrib (env R') (hist A R' n) P -instance [StandardBorelSpace E] [Nonempty E] (alg : Algorithm α R) (n : ℕ) : - IsMarkovKernel (posterior Q κ alg n) := by +instance [StandardBorelSpace E] [Nonempty E] (n : ℕ) : IsMarkovKernel (posterior P A R' n) := by unfold posterior infer_instance section Laws variable (alg : Algorithm α R) -variable [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] -variable [StandardBorelSpace E] [Nonempty E] - -lemma isAlgEnvSeq_trajMeasure : - IsAlgEnvSeq action IT.reward (alg.prod_left E) (StationaryEnv Q κ) (trajMeasure Q κ alg) := - IT.isAlgEnvSeq_trajMeasure _ _ -- Revise (Claude) -lemma hasLaw_env : HasLaw env Q (trajMeasure Q κ alg) := by +lemma hasLaw_env (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) : + HasLaw (env R') Q P := by apply HasCondDistrib.hasLaw_of_const - simpa [StationaryEnv] using (isAlgEnvSeq_trajMeasure Q κ alg).hasCondDistrib_reward_zero.fst + simpa [StationaryEnv] using h.hasCondDistrib_reward_zero.fst -- Revise (Claude) -lemma hasCondDistrib_action (n : ℕ) : - HasCondDistrib (action (n + 1)) (hist n) (alg.policy n) (trajMeasure Q κ alg) := - ((isAlgEnvSeq_trajMeasure Q κ alg).hasCondDistrib_action n).comp_left (by fun_prop) +lemma hasCondDistrib_action (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) (n : ℕ) : + HasCondDistrib (A (n + 1)) (hist A R' n) (alg.policy n) P := + (h.hasCondDistrib_action n).comp_left (by fun_prop) -- Revise (Claude) -lemma hasCondDistrib_reward_zero : - HasCondDistrib (reward 0) (fun ω ↦ (action 0 ω, env ω)) κ (trajMeasure Q κ alg) := by - simpa [StationaryEnv] using - (isAlgEnvSeq_trajMeasure Q κ alg).hasCondDistrib_reward_zero.of_compProd +lemma hasCondDistrib_reward_zero (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) : + HasCondDistrib (reward R' 0) (fun ω ↦ (A 0 ω, env R' ω)) κ P := by + simpa [StationaryEnv] using h.hasCondDistrib_reward_zero.of_compProd -- Revise (Claude) -lemma hasCondDistrib_reward (n : ℕ) : - HasCondDistrib (reward (n + 1)) (fun ω ↦ (action (n + 1) ω, env ω)) κ - (trajMeasure Q κ alg) := by - have h := ((isAlgEnvSeq_trajMeasure Q κ alg).hasCondDistrib_reward n).snd +lemma hasCondDistrib_reward (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) (n : ℕ) : + HasCondDistrib (reward R' (n + 1)) (fun ω ↦ (A (n + 1) ω, env R' ω)) κ P := by + have h := (h.hasCondDistrib_reward n).snd simp_rw [StationaryEnv, Kernel.snd_prod] at h exact h.comp_left (by fun_prop) -- Revise (Claude) -lemma hasCondDistrib_action_env_hist (n : ℕ) : - HasCondDistrib (action (n + 1)) (fun ω ↦ (env ω, hist n ω)) - ((alg.policy n).prodMkLeft E) (trajMeasure Q κ alg) := by - have h := (isAlgEnvSeq_trajMeasure Q κ alg).hasCondDistrib_action n +lemma hasCondDistrib_action_env_hist (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) + (n : ℕ) : + HasCondDistrib (A (n + 1)) (fun ω ↦ (env R' ω, hist A R' n ω)) + ((alg.policy n).prodMkLeft E) P := by + have h := h.hasCondDistrib_action n simp only [Algorithm.prod_left, Kernel.prodMkLeft] at h ⊢ have hf : Measurable (fun h : Iic n → α × E × R ↦ ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2))) := by fun_prop - have h' : HasCondDistrib (action (n + 1)) (IsAlgEnvSeq.hist action IT.reward n) + have h' : HasCondDistrib (A (n + 1)) (Learning.IsAlgEnvSeq.hist A R' n) (((alg.policy n).comap Prod.snd (by fun_prop)).comap (fun h : Iic n → α × E × R ↦ ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2))) hf) - (trajMeasure Q κ alg) := by convert h using 2 + P := by convert h using 2 convert h'.comp_left hf using 2 --- Revise (Claude) -lemma hasCondDistrib_reward_hist (n : ℕ) : - HasCondDistrib (reward (n + 1)) (fun ω ↦ (hist n ω, action (n + 1) ω, env ω)) - (κ.prodMkLeft _) (trajMeasure Q κ alg) := by - have h := ((isAlgEnvSeq_trajMeasure Q κ alg).hasCondDistrib_reward n).snd +-- -- Revise (Claude) +lemma hasCondDistrib_reward_hist (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) + (n : ℕ) : + HasCondDistrib (reward R' (n + 1)) (fun ω ↦ (hist A R' n ω, A (n + 1) ω, env R' ω)) + (κ.prodMkLeft _) P := by + have h := (h.hasCondDistrib_reward n).snd simp_rw [StationaryEnv, Kernel.snd_prod, Kernel.prodMkLeft] at h ⊢ have hf : Measurable (fun (p : (Iic n → α × (E × R)) × α) ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1)) := by fun_prop - have h' : HasCondDistrib (reward (n + 1)) - (fun ω ↦ (IsAlgEnvSeq.hist action IT.reward n ω, action (n + 1) ω)) + have h' : HasCondDistrib (reward R' (n + 1)) + (fun ω ↦ (Learning.IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) ((κ.comap Prod.snd (by fun_prop)).comap (fun (p : (Iic n → α × (E × R)) × α) ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1)) hf) - (trajMeasure Q κ alg) := by convert h using 2 + P := by convert h using 2 convert h'.comp_left hf using 2 +variable [StandardBorelSpace Ω] + -- Revise (Claude) -lemma condIndepFun_action_env_hist (n : ℕ) : - action (n + 1) ⟂ᵢ[hist n, (by fun_prop); trajMeasure Q κ alg] env := - condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (by fun_prop) (by fun_prop) (by fun_prop) - (hasCondDistrib_action_env_hist Q κ alg n).condDistrib_eq +lemma condIndepFun_action_env_hist (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) + (n : ℕ) : + A (n + 1) ⟂ᵢ[hist A R' n, measurable_hist A R' h.measurable_A h.measurable_R n; P] (env R') := + condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (measurable_env R' h.measurable_R) + (h.measurable_A _) (measurable_hist A R' h.measurable_A h.measurable_R n) + (hasCondDistrib_action_env_hist Q κ P A R' alg h n).condDistrib_eq -- Revise (Claude) -lemma condIndepFun_reward_hist (n : ℕ) : - reward (n + 1) - ⟂ᵢ[(fun ω ↦ (action (n + 1) ω, env ω)), (by fun_prop); trajMeasure Q κ alg] hist n := - condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (by fun_prop) (by fun_prop) (by fun_prop) - (hasCondDistrib_reward_hist Q κ alg n).condDistrib_eq +lemma condIndepFun_reward_hist (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) + (n : ℕ) : + reward R' (n + 1) + ⟂ᵢ[(fun ω ↦ (A (n + 1) ω, env R' ω)), + (h.measurable_A _).prodMk (measurable_env R' h.measurable_R); P] + hist A R' n := + condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft + (measurable_hist A R' h.measurable_A h.measurable_R n) + (measurable_reward R' h.measurable_R _) + ((h.measurable_A _).prodMk (measurable_env R' h.measurable_R)) + (hasCondDistrib_reward_hist Q κ P A R' alg h n).condDistrib_eq end Laws section Regret variable (κ : Kernel (α × E) ℝ) +variable (alg : Algorithm α ℝ) noncomputable -def regretAt (t : ℕ) (ω : ℕ → α × E × ℝ) : ℝ := - Bandits.regret (κ.comap (·, env ω) (by fun_prop)) action t ω +def regretAt (t : ℕ) (ω : Ω) : ℝ := + Bandits.regret (κ.comap (·, env R' ω) (by fun_prop)) A t ω --- Revise (Claude) -lemma measurable_regretAt [Fintype α] (t : ℕ) : Measurable (regretAt κ t) := by +omit [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] + [StandardBorelSpace R] [Nonempty R] in +lemma measurable_regretAt [Fintype α] (hA : ∀ n, Measurable (A n)) + (hR' : ∀ n, Measurable (R' n)) (t : ℕ) : Measurable (regretAt A R' κ t) := by unfold regretAt Bandits.regret have hmean : Measurable fun (p : α × E) ↦ (κ p)[id] := stronglyMeasurable_id.integral_kernel.measurable apply Measurable.sub · apply Measurable.const_mul - have h : ∀ a, Measurable fun (ω : ℕ → α × E × ℝ) ↦ (κ (a, env ω))[id] := fun a ↦ - hmean.comp (measurable_const.prodMk measurable_env) + have h : ∀ a, Measurable fun (ω : Ω) ↦ (κ (a, env R' ω))[id] := fun a ↦ + hmean.comp (measurable_const.prodMk (measurable_env R' hR')) exact Measurable.iSup h · apply Finset.measurable_sum intro s _ - exact hmean.comp ((measurable_action s).prodMk measurable_env) - -variable (alg : Algorithm α ℝ) + exact hmean.comp ((hA s).prodMk (measurable_env R' hR')) noncomputable -def regret [IsMarkovKernel κ] (t : ℕ) : ℝ := (trajMeasure Q κ alg)[regretAt κ t] +def regret [IsMarkovKernel κ] (t : ℕ) : ℝ := P[regretAt A R' κ t] end Regret +end IsAlgEnvSeq + end Learning.Bayes From abd75f831132dea5e115b038813f2728a738c8bd Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 23 Jan 2026 13:13:20 +0000 Subject: [PATCH 015/155] Move posterior over best arm to BayesStationarEnv --- LeanBandits/BanditAlgorithms/TS.lean | 23 ++++--------------- .../BayesStationaryEnv.lean | 18 +++++++++++++++ 2 files changed, 23 insertions(+), 18 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 04457e1c..7a4b6570 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -19,30 +19,17 @@ variable (hK : 0 < K) variable (Q : Measure E) [IsProbabilityMeasure Q] variable (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] -/-- The posterior over "environments" for every given history for TS. Note that we pretend that the -data was generated by an algorithm that chooses actions uniformly at random to avoid circularity in -the definition of `tsAlgorithm`. -/ -noncomputable -def tsPosterior (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) E := - Learning.Bayes.IsAlgEnvSeq.posterior - (trajMeasure ((uniformAlgorithm hK).prod_left E) (Bayes.StationaryEnv Q κ)) - IT.action IT.reward n - -instance (n : ℕ) : IsMarkovKernel (tsPosterior hK Q κ n) := by - unfold tsPosterior Bayes.IsAlgEnvSeq.posterior - infer_instance - /-- The distribution over actions for every given history for TS. -/ noncomputable def tsPolicy (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - (tsPosterior hK Q κ n).map (measurableArgmax (fun e k ↦ (κ (k, e))[id])) + Learning.Bayes.IsAlgEnvSeq.condDistribBestArm + (trajMeasure ((uniformAlgorithm hK).prod_left E) (Bayes.StationaryEnv Q κ)) IT.action + IT.reward κ n instance (n : ℕ) : IsMarkovKernel (tsPolicy hK Q κ n) := by - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - apply Kernel.IsMarkovKernel.map - exact measurable_measurableArgmax fun k => - (stronglyMeasurable_id.integral_kernel (κ := κ.comap (k, ·) (by fun_prop))).measurable + unfold tsPolicy Bayes.IsAlgEnvSeq.condDistribBestArm + infer_instance /-- The initial distribution over actions for TS. -/ noncomputable diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 8f5a897a..dcc6ee57 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -3,6 +3,7 @@ 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.MeasurableArgMax import LeanBandits.Bandit.Regret open MeasureTheory ProbabilityTheory Finset @@ -152,6 +153,23 @@ section Regret variable (κ : Kernel (α × E) ℝ) variable (alg : Algorithm α ℝ) +noncomputable +def value (a : α) (ω : Ω) : ℝ := (κ (a, env R' ω))[id] + +noncomputable +def bestArm [Fintype α] [Encodable α] [MeasurableSingletonClass α] (ω : Ω) : α := + measurableArgmax (fun ω' a ↦ value R' κ a ω') ω + +noncomputable +def condDistribBestArm [Fintype α] [Encodable α] [MeasurableSingletonClass α] (n : ℕ) : + Kernel (Iic n → α × R) α := + condDistrib (bestArm R' κ) (hist A R' n) P + +instance (n : ℕ) [Fintype α] [Encodable α] [MeasurableSingletonClass α] : + IsMarkovKernel (condDistribBestArm P A R' κ n) := by + unfold condDistribBestArm + infer_instance + noncomputable def regretAt (t : ℕ) (ω : Ω) : ℝ := Bandits.regret (κ.comap (·, env R' ω) (by fun_prop)) A t ω From 3a9a5b477b9818dea2e3cafefdfee3edd44508f7 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 23 Jan 2026 15:40:07 +0000 Subject: [PATCH 016/155] Organize BayesStationaryEnv --- LeanBandits/BanditAlgorithms/TS.lean | 8 +- .../BayesStationaryEnv.lean | 176 ++++++++++-------- 2 files changed, 100 insertions(+), 84 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 7a4b6570..7f083ae7 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -6,7 +6,6 @@ Authors: Rémy Degenne, Paulo Rauber import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.BanditAlgorithms.Uniform import LeanBandits.SequentialLearning.BayesStationaryEnv -import LeanBandits.SequentialLearning.IonescuTulceaSpace open MeasureTheory ProbabilityTheory Finset Learning @@ -23,12 +22,11 @@ variable (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] noncomputable def tsPolicy (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - Learning.Bayes.IsAlgEnvSeq.condDistribBestArm - (trajMeasure ((uniformAlgorithm hK).prod_left E) (Bayes.StationaryEnv Q κ)) IT.action - IT.reward κ n + IsBayesianAlgEnvSeq.condDistribBestAction (IT.bayesianTrajMeasure Q κ (uniformAlgorithm hK)) κ + IT.action IT.reward n instance (n : ℕ) : IsMarkovKernel (tsPolicy hK Q κ n) := by - unfold tsPolicy Bayes.IsAlgEnvSeq.condDistribBestArm + unfold tsPolicy infer_instance /-- The initial distribution over actions for TS. -/ diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index dcc6ee57..15787218 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -5,97 +5,99 @@ Authors: Rémy Degenne, Paulo Rauber -/ import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.Bandit.Regret +import LeanBandits.SequentialLearning.IonescuTulceaSpace open MeasureTheory ProbabilityTheory Finset -namespace Learning.Bayes +namespace Learning -variable {α R E : Type*} [mα : MeasurableSpace α] [mR : MeasurableSpace R] [mE : MeasurableSpace E] +variable {α E R : Type*} [mα : MeasurableSpace α] [mE : MeasurableSpace E] [mR : MeasurableSpace R] variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] -/-- Given a prior distribution `Q` over "environments" and a kernel `k` that defines a reward -distribution `κ (a, e)` for each action `a : α` and "environment" `e : E`, a StationaryEnv -represents an environment (with an observation space `E × R`) that draws an "environment" `e : E` at -the very first step which, together with `k`, defines how the bandit process behaves. Because the -"environment" `e` is repeated at every step and reveals the best arm, it only makes sense to study -algorithms that ignore the information in `E` and just receive the information in `R`. -/ noncomputable -def StationaryEnv : Environment α (E × R) where +def BayesianStationaryEnv : Environment α (E × R) where feedback n := let g : (Iic n → α × E × R) × α → α × E := fun (h, a) => (a, (h ⟨0, by simp⟩).2.1) (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) ν0 := (Kernel.const α Q) ⊗ₖ κ -variable {Ω : Type*} [MeasurableSpace Ω] -variable (P : Measure Ω) [IsProbabilityMeasure P] -variable (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) +variable {Ω : Type*} [mΩ : MeasurableSpace Ω] + +structure IsBayesianAlgEnvSeq + [StandardBorelSpace α] [Nonempty α] + [StandardBorelSpace E] [Nonempty E] + [StandardBorelSpace R] [Nonempty R] + (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (alg : Algorithm α R) + (P : Measure Ω) [IsFiniteMeasure P] + extends IsAlgEnvSeq A R' (alg.prod_left E) (BayesianStationaryEnv Q κ) P -namespace IsAlgEnvSeq +variable (P : Measure Ω) [IsProbabilityMeasure P] +variable (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (alg : Algorithm α R) -def env (ω : Ω) : E := (R' 0 ω).1 +def IsBayesianAlgEnvSeq.env (ω : Ω) : E := (R' 0 ω).1 @[fun_prop] -lemma measurable_env (hR' : ∀ n, Measurable (R' n)) : Measurable (env R'):= (hR' 0).fst +lemma IsBayesianAlgEnvSeq.measurable_env (hR' : ∀ n, Measurable (R' n)) : Measurable (env R'):= + (hR' 0).fst -def reward (n : ℕ) (ω : Ω) : R := (R' n ω).2 +def IsBayesianAlgEnvSeq.reward (n : ℕ) (ω : Ω) : R := (R' n ω).2 @[fun_prop] -lemma measurable_reward (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : Measurable (reward R' n) := by +lemma IsBayesianAlgEnvSeq.measurable_reward (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : + Measurable (reward R' n) := by unfold reward fun_prop -def hist (n : ℕ) (ω : Ω) : Iic n → α × R := fun i ↦ (A i ω, (R' i ω).2) +def IsBayesianAlgEnvSeq.hist (n : ℕ) (ω : Ω) : Iic n → α × R := fun i ↦ (A i ω, (R' i ω).2) @[fun_prop] -lemma measurable_hist (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : - Measurable (hist A R' n) := by +lemma IsBayesianAlgEnvSeq.measurable_hist (hA : ∀ n, Measurable (A n)) + (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : Measurable (hist A R' n) := by unfold hist fun_prop -variable [StandardBorelSpace α] [Nonempty α] -variable [StandardBorelSpace E] [Nonempty E] -variable [StandardBorelSpace R] [Nonempty R] - -/-- The posterior over "environments" for every given history (for a fixed algorithm). -/ noncomputable -def posterior (n : ℕ) : Kernel (Iic n → α × R) E := +def IsBayesianAlgEnvSeq.posterior [StandardBorelSpace E] [Nonempty E] (n : ℕ) : + Kernel (Iic n → α × R) E := condDistrib (env R') (hist A R' n) P -instance [StandardBorelSpace E] [Nonempty E] (n : ℕ) : IsMarkovKernel (posterior P A R' n) := by - unfold posterior +instance (n : ℕ) [StandardBorelSpace E] [Nonempty E] : + IsMarkovKernel (IsBayesianAlgEnvSeq.posterior P A R' n) := by + unfold IsBayesianAlgEnvSeq.posterior infer_instance section Laws -variable (alg : Algorithm α R) +variable [StandardBorelSpace α] [Nonempty α] +variable [StandardBorelSpace E] [Nonempty E] +variable [StandardBorelSpace R] [Nonempty R] --- Revise (Claude) -lemma hasLaw_env (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) : +-- (Claude) +lemma IsBayesianAlgEnvSeq.hasLaw_env (h : IsBayesianAlgEnvSeq Q κ A R' alg P) : HasLaw (env R') Q P := by apply HasCondDistrib.hasLaw_of_const - simpa [StationaryEnv] using h.hasCondDistrib_reward_zero.fst + simpa [BayesianStationaryEnv] using h.hasCondDistrib_reward_zero.fst --- Revise (Claude) -lemma hasCondDistrib_action (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) (n : ℕ) : +-- (Claude) +lemma IsBayesianAlgEnvSeq.hasCondDistrib_action' (h : IsBayesianAlgEnvSeq Q κ A R' alg P) (n : ℕ) : HasCondDistrib (A (n + 1)) (hist A R' n) (alg.policy n) P := (h.hasCondDistrib_action n).comp_left (by fun_prop) --- Revise (Claude) -lemma hasCondDistrib_reward_zero (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) : +--(Claude) +lemma IsBayesianAlgEnvSeq.hasCondDistrib_reward_zero' (h : IsBayesianAlgEnvSeq Q κ A R' alg P) : HasCondDistrib (reward R' 0) (fun ω ↦ (A 0 ω, env R' ω)) κ P := by - simpa [StationaryEnv] using h.hasCondDistrib_reward_zero.of_compProd + simpa [BayesianStationaryEnv] using h.hasCondDistrib_reward_zero.of_compProd --- Revise (Claude) -lemma hasCondDistrib_reward (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) (n : ℕ) : +-- (Claude) +lemma IsBayesianAlgEnvSeq.hasCondDistrib_reward' (h : IsBayesianAlgEnvSeq Q κ A R' alg P) (n : ℕ) : HasCondDistrib (reward R' (n + 1)) (fun ω ↦ (A (n + 1) ω, env R' ω)) κ P := by have h := (h.hasCondDistrib_reward n).snd - simp_rw [StationaryEnv, Kernel.snd_prod] at h + simp_rw [BayesianStationaryEnv, Kernel.snd_prod] at h exact h.comp_left (by fun_prop) --- Revise (Claude) -lemma hasCondDistrib_action_env_hist (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) - (n : ℕ) : - HasCondDistrib (A (n + 1)) (fun ω ↦ (env R' ω, hist A R' n ω)) +-- (Claude) +lemma IsBayesianAlgEnvSeq.hasCondDistrib_action_env_hist (h : IsBayesianAlgEnvSeq Q κ A R' alg P) + (n : ℕ) : HasCondDistrib (A (n + 1)) (fun ω ↦ (env R' ω, hist A R' n ω)) ((alg.policy n).prodMkLeft E) P := by have h := h.hasCondDistrib_action n simp only [Algorithm.prod_left, Kernel.prodMkLeft] at h ⊢ @@ -107,13 +109,13 @@ lemma hasCondDistrib_action_env_hist (h : IsAlgEnvSeq A R' (alg.prod_left E) (St P := by convert h using 2 convert h'.comp_left hf using 2 --- -- Revise (Claude) -lemma hasCondDistrib_reward_hist (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) +-- -- (Claude) +lemma IsBayesianAlgEnvSeq.hasCondDistrib_reward_hist (h : IsBayesianAlgEnvSeq Q κ A R' alg P) (n : ℕ) : HasCondDistrib (reward R' (n + 1)) (fun ω ↦ (hist A R' n ω, A (n + 1) ω, env R' ω)) (κ.prodMkLeft _) P := by have h := (h.hasCondDistrib_reward n).snd - simp_rw [StationaryEnv, Kernel.snd_prod, Kernel.prodMkLeft] at h ⊢ + simp_rw [BayesianStationaryEnv, Kernel.snd_prod, Kernel.prodMkLeft] at h ⊢ have hf : Measurable (fun (p : (Iic n → α × (E × R)) × α) ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1)) := by fun_prop have h' : HasCondDistrib (reward R' (n + 1)) @@ -123,19 +125,17 @@ lemma hasCondDistrib_reward_hist (h : IsAlgEnvSeq A R' (alg.prod_left E) (Statio P := by convert h using 2 convert h'.comp_left hf using 2 -variable [StandardBorelSpace Ω] - --- Revise (Claude) -lemma condIndepFun_action_env_hist (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) - (n : ℕ) : +-- (Claude) +lemma IsBayesianAlgEnvSeq.condIndepFun_action_env_hist [StandardBorelSpace Ω] + (h : IsBayesianAlgEnvSeq Q κ A R' alg P) (n : ℕ) : A (n + 1) ⟂ᵢ[hist A R' n, measurable_hist A R' h.measurable_A h.measurable_R n; P] (env R') := condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (measurable_env R' h.measurable_R) (h.measurable_A _) (measurable_hist A R' h.measurable_A h.measurable_R n) (hasCondDistrib_action_env_hist Q κ P A R' alg h n).condDistrib_eq --- Revise (Claude) -lemma condIndepFun_reward_hist (h : IsAlgEnvSeq A R' (alg.prod_left E) (StationaryEnv Q κ) P) - (n : ℕ) : +-- Claude +lemma IsBayesianAlgEnvSeq.condIndepFun_reward_hist [StandardBorelSpace Ω] + (h : IsBayesianAlgEnvSeq Q κ A R' alg P) (n : ℕ) : reward R' (n + 1) ⟂ᵢ[(fun ω ↦ (A (n + 1) ω, env R' ω)), (h.measurable_A _).prodMk (measurable_env R' h.measurable_R); P] @@ -148,37 +148,38 @@ lemma condIndepFun_reward_hist (h : IsAlgEnvSeq A R' (alg.prod_left E) (Stationa end Laws -section Regret +section Real variable (κ : Kernel (α × E) ℝ) -variable (alg : Algorithm α ℝ) +variable (A : ℕ → Ω → α) (R' : ℕ → Ω → E × ℝ) (alg : Algorithm α ℝ) +-- κ? noncomputable -def value (a : α) (ω : Ω) : ℝ := (κ (a, env R' ω))[id] +def IsBayesianAlgEnvSeq.actionValue (ω : Ω) (a : α) : ℝ := (κ (a, env R' ω))[id] noncomputable -def bestArm [Fintype α] [Encodable α] [MeasurableSingletonClass α] (ω : Ω) : α := - measurableArgmax (fun ω' a ↦ value R' κ a ω') ω +def IsBayesianAlgEnvSeq.bestAction [Nonempty α] [Fintype α] [Encodable α] + [MeasurableSingletonClass α] (ω : Ω) : α := measurableArgmax (actionValue κ R') ω noncomputable -def condDistribBestArm [Fintype α] [Encodable α] [MeasurableSingletonClass α] (n : ℕ) : - Kernel (Iic n → α × R) α := - condDistrib (bestArm R' κ) (hist A R' n) P - -instance (n : ℕ) [Fintype α] [Encodable α] [MeasurableSingletonClass α] : - IsMarkovKernel (condDistribBestArm P A R' κ n) := by - unfold condDistribBestArm +def IsBayesianAlgEnvSeq.condDistribBestAction [StandardBorelSpace α] [Nonempty α] [Fintype α] + [Encodable α] (n : ℕ) : Kernel (Iic n → α × ℝ) α := + condDistrib (bestAction κ R') (hist A R' n) P + +instance [StandardBorelSpace α] [Nonempty α] [Fintype α] + [Encodable α] (n : ℕ) : + IsMarkovKernel (IsBayesianAlgEnvSeq.condDistribBestAction P κ A R' n) := by + unfold IsBayesianAlgEnvSeq.condDistribBestAction infer_instance noncomputable -def regretAt (t : ℕ) (ω : Ω) : ℝ := - Bandits.regret (κ.comap (·, env R' ω) (by fun_prop)) A t ω - -omit [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] - [StandardBorelSpace R] [Nonempty R] in -lemma measurable_regretAt [Fintype α] (hA : ∀ n, Measurable (A n)) - (hR' : ∀ n, Measurable (R' n)) (t : ℕ) : Measurable (regretAt A R' κ t) := by - unfold regretAt Bandits.regret +def IsBayesianAlgEnvSeq.regret (t : ℕ) (ω : Ω) : ℝ := + Bandits.regret (κ.comap (·, env R' ω) (by fun_prop)) A t ω + +-- Claude +lemma IsBayesianAlgEnvSeq.measurable_regret [Fintype α] (hA : ∀ n, Measurable (A n)) + (hR' : ∀ n, Measurable (R' n)) (t : ℕ) : Measurable (regret κ A R' t) := by + unfold regret Bandits.regret have hmean : Measurable fun (p : α × E) ↦ (κ p)[id] := stronglyMeasurable_id.integral_kernel.measurable apply Measurable.sub @@ -191,10 +192,27 @@ lemma measurable_regretAt [Fintype α] (hA : ∀ n, Measurable (A n)) exact hmean.comp ((hA s).prodMk (measurable_env R' hR')) noncomputable -def regret [IsMarkovKernel κ] (t : ℕ) : ℝ := P[regretAt A R' κ t] +def IsBayesianAlgEnvSeq.bayesianRegret + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] + [IsMarkovKernel κ] (t : ℕ) : ℝ := + P[IsBayesianAlgEnvSeq.regret κ A R' t] + +end Real + +namespace IT + +noncomputable +def bayesianTrajMeasure (alg : Algorithm α R) : Measure (ℕ → α × E × R) := + trajMeasure (alg.prod_left E) (BayesianStationaryEnv Q κ) +deriving IsProbabilityMeasure -end Regret +lemma isBayesianAlgEnvSeq_bayesianTrajMeasure + [StandardBorelSpace α] [Nonempty α] + [StandardBorelSpace E] [Nonempty E] + [StandardBorelSpace R] [Nonempty R] (alg : Algorithm α R) : + IsBayesianAlgEnvSeq Q κ action reward alg (bayesianTrajMeasure Q κ alg) := + ⟨isAlgEnvSeq_trajMeasure _ _⟩ -end IsAlgEnvSeq +end IT -end Learning.Bayes +end Learning From 5a0d994da4631597ac979a9dce777734430d62f5 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 26 Jan 2026 10:05:26 +0000 Subject: [PATCH 017/155] State Bayesian regret theorem --- LeanBandits/BanditAlgorithms/TS.lean | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 7f083ae7..625c2a78 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -18,6 +18,8 @@ variable (hK : 0 < K) variable (Q : Measure E) [IsProbabilityMeasure Q] variable (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] +namespace TS + /-- The distribution over actions for every given history for TS. -/ noncomputable def tsPolicy (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := @@ -50,4 +52,17 @@ def tsAlgorithm : Algorithm (Fin K) ℝ where policy := tsPolicy hK Q κ p0 := tsInitPolicy hK Q κ + +variable {Ω : Type*} [MeasurableSpace Ω] +variable {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → E × ℝ} +variable {P : Measure Ω} [IsFiniteMeasure P] + +def bayesian_regret_le [Nonempty (Fin K)] + (h : IsBayesianAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) : + ∃ C > 0, ∀ K > 0, ∀ n : ℕ, + (IsBayesianAlgEnvSeq.bayesianRegret P κ A R' n) ≤ C * √(K * n * Real.log n) := + sorry + +end TS + end Bandits From 73520cdc94e4c8a152814249be03da4135a5a810 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 26 Jan 2026 10:13:34 +0000 Subject: [PATCH 018/155] Add extra conditions for Bayesian regret statement --- LeanBandits/BanditAlgorithms/TS.lean | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 625c2a78..02428bf5 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -6,6 +6,7 @@ Authors: Rémy Degenne, Paulo Rauber import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.BanditAlgorithms.Uniform import LeanBandits.SequentialLearning.BayesStationaryEnv +import LeanBandits.ForMathlib.SubGaussian open MeasureTheory ProbabilityTheory Finset Learning @@ -58,7 +59,9 @@ variable {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → E × ℝ} variable {P : Measure Ω} [IsFiniteMeasure P] def bayesian_regret_le [Nonempty (Fin K)] - (h : IsBayesianAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) : + (h : IsBayesianAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) + (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) : ∃ C > 0, ∀ K > 0, ∀ n : ℕ, (IsBayesianAlgEnvSeq.bayesianRegret P κ A R' n) ≤ C * √(K * n * Real.log n) := sorry From 8f551571e21059b5fcb355912d42481490701bd4 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 26 Jan 2026 10:40:46 +0000 Subject: [PATCH 019/155] Fix statement --- LeanBandits/BanditAlgorithms/TS.lean | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 02428bf5..6a2108c8 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -62,7 +62,7 @@ def bayesian_regret_le [Nonempty (Fin K)] (h : IsBayesianAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) : - ∃ C > 0, ∀ K > 0, ∀ n : ℕ, + ∃ C > 0, ∀ n : ℕ, (IsBayesianAlgEnvSeq.bayesianRegret P κ A R' n) ≤ C * √(K * n * Real.log n) := sorry From 57033cf2d589cff73948c2a9ef78cd74f67326e1 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 27 Jan 2026 11:47:11 +0000 Subject: [PATCH 020/155] Refactor --- LeanBandits/BanditAlgorithms/TS.lean | 15 +- LeanBandits/BanditAlgorithms/Uniform.lean | 4 +- .../BayesStationaryEnv.lean | 248 ++++++++++-------- 3 files changed, 140 insertions(+), 127 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 6a2108c8..df4b5def 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -25,8 +25,7 @@ namespace TS noncomputable def tsPolicy (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - IsBayesianAlgEnvSeq.condDistribBestAction (IT.bayesianTrajMeasure Q κ (uniformAlgorithm hK)) κ - IT.action IT.reward n + IT.posteriorBestArm Q κ (uniformAlgorithm hK) n instance (n : ℕ) : IsMarkovKernel (tsPolicy hK Q κ n) := by unfold tsPolicy @@ -36,14 +35,12 @@ instance (n : ℕ) : IsMarkovKernel (tsPolicy hK Q κ n) := by noncomputable def tsInitPolicy : Measure (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - Q.map (measurableArgmax (fun e k ↦ (κ (k, e))[id])) + IT.priorBestArm Q κ (uniformAlgorithm hK) instance : IsProbabilityMeasure (tsInitPolicy hK Q κ) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - apply Measure.isProbabilityMeasure_map - apply Measurable.aemeasurable - exact (measurable_measurableArgmax fun k => - (stronglyMeasurable_id.integral_kernel (κ := κ.comap (k, ·) (by fun_prop))).measurable) + unfold tsInitPolicy + infer_instance /-- The Thompson Sampling (TS) algorithm: actions are chosen according to the probability that they are optimal given prior knowledge represented by a prior distribution `Q` and a data generation @@ -59,11 +56,11 @@ variable {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → E × ℝ} variable {P : Measure Ω} [IsFiniteMeasure P] def bayesian_regret_le [Nonempty (Fin K)] - (h : IsBayesianAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) : ∃ C > 0, ∀ n : ℕ, - (IsBayesianAlgEnvSeq.bayesianRegret P κ A R' n) ≤ C * √(K * n * Real.log n) := + (IsBayesAlgEnvSeq.bayesRegret κ A R' P n) ≤ C * √(K * n * Real.log n) := sorry end TS diff --git a/LeanBandits/BanditAlgorithms/Uniform.lean b/LeanBandits/BanditAlgorithms/Uniform.lean index cd054d93..7e8208b2 100644 --- a/LeanBandits/BanditAlgorithms/Uniform.lean +++ b/LeanBandits/BanditAlgorithms/Uniform.lean @@ -10,11 +10,11 @@ open MeasureTheory ProbabilityTheory Learning namespace Bandits -variable {K : ℕ} (hK : 0 < K) +variable {K : ℕ} /-- The Uniform algorithm: actions are chosen uniformly at random. -/ noncomputable -def uniformAlgorithm : Algorithm (Fin K) ℝ := +def uniformAlgorithm (hK : 0 < K) : Algorithm (Fin K) ℝ := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK have : IsProbabilityMeasure (uniformOn (Set.univ : Set (Fin K))) := uniformOn_isProbabilityMeasure Set.finite_univ Set.univ_nonempty diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 15787218..2e65b9cf 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -12,10 +12,11 @@ open MeasureTheory ProbabilityTheory Finset namespace Learning variable {α E R : Type*} [mα : MeasurableSpace α] [mE : MeasurableSpace E] [mR : MeasurableSpace R] -variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] noncomputable -def BayesianStationaryEnv : Environment α (E × R) where +def bayesStationaryEnv + (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] : + Environment α (E × R) where feedback n := let g : (Iic n → α × E × R) × α → α × E := fun (h, a) => (a, (h ⟨0, by simp⟩).2.1) (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) @@ -23,81 +24,126 @@ def BayesianStationaryEnv : Environment α (E × R) where variable {Ω : Type*} [mΩ : MeasurableSpace Ω] -structure IsBayesianAlgEnvSeq +def IsBayesAlgEnvSeq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] [StandardBorelSpace R] [Nonempty R] + (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (alg : Algorithm α R) (P : Measure Ω) [IsFiniteMeasure P] - extends IsAlgEnvSeq A R' (alg.prod_left E) (BayesianStationaryEnv Q κ) P + := IsAlgEnvSeq A R' (alg.prod_left E) (bayesStationaryEnv Q κ) P -variable (P : Measure Ω) [IsProbabilityMeasure P] -variable (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (alg : Algorithm α R) +namespace IsBayesAlgEnvSeq -def IsBayesianAlgEnvSeq.env (ω : Ω) : E := (R' 0 ω).1 +variable [StandardBorelSpace α] [Nonempty α] +variable [StandardBorelSpace E] [Nonempty E] +variable [StandardBorelSpace R] [Nonempty R] + +variable {Q : Measure E} [IsProbabilityMeasure Q] {κ : Kernel (α × E) R} [IsMarkovKernel κ] +variable {A : ℕ → Ω → α} {R' : ℕ → Ω → E × R} +variable {alg : Algorithm α R} +variable {P : Measure Ω} [IsProbabilityMeasure P] + +def env (R' : ℕ → Ω → E × R) (ω : Ω) : E := (R' 0 ω).1 + +@[fun_prop] +lemma measurable_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) : Measurable (env R') := + (h.measurable_R 0).fst + +def reward (R' : ℕ → Ω → E × R) (n : ℕ) (ω : Ω) : R := (R' n ω).2 @[fun_prop] -lemma IsBayesianAlgEnvSeq.measurable_env (hR' : ∀ n, Measurable (R' n)) : Measurable (env R'):= - (hR' 0).fst +lemma measurable_reward (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : + Measurable (reward R' n) := + (h.measurable_R n).snd -def IsBayesianAlgEnvSeq.reward (n : ℕ) (ω : Ω) : R := (R' n ω).2 +def hist (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (n : ℕ) (ω : Ω) : Iic n → α × R := + fun i ↦ (A i ω, (R' i ω).2) @[fun_prop] -lemma IsBayesianAlgEnvSeq.measurable_reward (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : - Measurable (reward R' n) := by - unfold reward - fun_prop +lemma measurable_hist (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : Measurable (hist A R' n) := + measurable_pi_lambda _ fun i => (h.measurable_A i).prodMk (h.measurable_R i).snd -def IsBayesianAlgEnvSeq.hist (n : ℕ) (ω : Ω) : Iic n → α × R := fun i ↦ (A i ω, (R' i ω).2) +def action_env (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (n : ℕ) (ω : Ω) : α × E := (A n ω, (R' 0 ω).1) @[fun_prop] -lemma IsBayesianAlgEnvSeq.measurable_hist (hA : ∀ n, Measurable (A n)) - (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : Measurable (hist A R' n) := by - unfold hist - fun_prop +lemma measurable_action_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : + Measurable (action_env A R' n) := + (h.measurable_A n).prodMk h.measurable_env + +section Real + +variable {R' : ℕ → Ω → E × ℝ} +variable {κ : Kernel (α × E) ℝ} [IsMarkovKernel κ] +variable {alg : Algorithm α ℝ} noncomputable -def IsBayesianAlgEnvSeq.posterior [StandardBorelSpace E] [Nonempty E] (n : ℕ) : - Kernel (Iic n → α × R) E := - condDistrib (env R') (hist A R' n) P +def armMean (κ : Kernel (α × E) ℝ) (R' : ℕ → Ω → E × ℝ) (a : α) (ω : Ω) : ℝ := + (κ (a, env R' ω))[id] -instance (n : ℕ) [StandardBorelSpace E] [Nonempty E] : - IsMarkovKernel (IsBayesianAlgEnvSeq.posterior P A R' n) := by - unfold IsBayesianAlgEnvSeq.posterior - infer_instance +lemma measurable_armMean (h : IsBayesAlgEnvSeq Q κ A R' alg P) (a : α) : + Measurable (armMean κ R' a) := + stronglyMeasurable_id.integral_kernel.measurable.comp + (measurable_const.prodMk h.measurable_env) -section Laws +noncomputable +def bestArm (κ : Kernel (α × E) ℝ) [Fintype α] [Encodable α] (R' : ℕ → Ω → E × ℝ) := + measurableArgmax (fun ω a ↦ armMean κ R' a ω) -variable [StandardBorelSpace α] [Nonempty α] -variable [StandardBorelSpace E] [Nonempty E] -variable [StandardBorelSpace R] [Nonempty R] +lemma measurable_bestArm [Fintype α] [Encodable α] [MeasurableSingletonClass α] + (h : IsBayesAlgEnvSeq Q κ A R' alg P) : + Measurable (bestArm κ R') := + measurable_measurableArgmax h.measurable_armMean --- (Claude) -lemma IsBayesianAlgEnvSeq.hasLaw_env (h : IsBayesianAlgEnvSeq Q κ A R' alg P) : - HasLaw (env R') Q P := by +noncomputable +def regret (κ : Kernel (α × E) ℝ) (A : ℕ → Ω → α) (R' : ℕ → Ω → E × ℝ) (t : ℕ) (ω : Ω) : ℝ := + Bandits.regret (κ.comap (·, env R' ω) (by fun_prop)) A t ω + +-- -- Claude +lemma measurable_regret [Fintype α] [Encodable α] (h : IsBayesAlgEnvSeq Q κ A R' alg P) (t : ℕ) : + Measurable (regret κ A R' t) := by + unfold regret Bandits.regret + apply Measurable.sub + · apply Measurable.const_mul + exact Measurable.iSup h.measurable_armMean + · apply Finset.measurable_sum + intro s _ + exact stronglyMeasurable_id.integral_kernel.measurable.comp + ((h.measurable_A s).prodMk h.measurable_env) + +noncomputable +def bayesRegret (κ : Kernel (α × E) ℝ) (A : ℕ → Ω → α) (R' : ℕ → Ω → E × ℝ) (P : Measure Ω) + (t : ℕ) : ℝ := + P[regret κ A R' t] + +end Real + +section Laws + +lemma hasLaw_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) : HasLaw (env R') Q P := by apply HasCondDistrib.hasLaw_of_const - simpa [BayesianStationaryEnv] using h.hasCondDistrib_reward_zero.fst + simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst -- (Claude) -lemma IsBayesianAlgEnvSeq.hasCondDistrib_action' (h : IsBayesianAlgEnvSeq Q κ A R' alg P) (n : ℕ) : +lemma hasCondDistrib_action' (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : HasCondDistrib (A (n + 1)) (hist A R' n) (alg.policy n) P := (h.hasCondDistrib_action n).comp_left (by fun_prop) --(Claude) -lemma IsBayesianAlgEnvSeq.hasCondDistrib_reward_zero' (h : IsBayesianAlgEnvSeq Q κ A R' alg P) : - HasCondDistrib (reward R' 0) (fun ω ↦ (A 0 ω, env R' ω)) κ P := by - simpa [BayesianStationaryEnv] using h.hasCondDistrib_reward_zero.of_compProd +lemma hasCondDistrib_reward_zero' (h : IsBayesAlgEnvSeq Q κ A R' alg P) : + HasCondDistrib (reward R' 0) (action_env A R' 0) κ P := by + simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.of_compProd -- (Claude) -lemma IsBayesianAlgEnvSeq.hasCondDistrib_reward' (h : IsBayesianAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - HasCondDistrib (reward R' (n + 1)) (fun ω ↦ (A (n + 1) ω, env R' ω)) κ P := by +lemma hasCondDistrib_reward' (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : + HasCondDistrib (reward R' (n + 1)) (action_env A R' (n + 1)) κ P := by have h := (h.hasCondDistrib_reward n).snd - simp_rw [BayesianStationaryEnv, Kernel.snd_prod] at h + simp_rw [bayesStationaryEnv, Kernel.snd_prod] at h exact h.comp_left (by fun_prop) -- (Claude) -lemma IsBayesianAlgEnvSeq.hasCondDistrib_action_env_hist (h : IsBayesianAlgEnvSeq Q κ A R' alg P) - (n : ℕ) : HasCondDistrib (A (n + 1)) (fun ω ↦ (env R' ω, hist A R' n ω)) +lemma hasCondDistrib_action_env_hist (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : + HasCondDistrib (A (n + 1)) (fun ω ↦ (env R' ω, hist A R' n ω)) ((alg.policy n).prodMkLeft E) P := by have h := h.hasCondDistrib_action n simp only [Algorithm.prod_left, Kernel.prodMkLeft] at h ⊢ @@ -110,12 +156,11 @@ lemma IsBayesianAlgEnvSeq.hasCondDistrib_action_env_hist (h : IsBayesianAlgEnvSe convert h'.comp_left hf using 2 -- -- (Claude) -lemma IsBayesianAlgEnvSeq.hasCondDistrib_reward_hist (h : IsBayesianAlgEnvSeq Q κ A R' alg P) - (n : ℕ) : +lemma hasCondDistrib_reward_hist_action_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : HasCondDistrib (reward R' (n + 1)) (fun ω ↦ (hist A R' n ω, A (n + 1) ω, env R' ω)) (κ.prodMkLeft _) P := by have h := (h.hasCondDistrib_reward n).snd - simp_rw [BayesianStationaryEnv, Kernel.snd_prod, Kernel.prodMkLeft] at h ⊢ + simp_rw [bayesStationaryEnv, Kernel.snd_prod, Kernel.prodMkLeft] at h ⊢ have hf : Measurable (fun (p : (Iic n → α × (E × R)) × α) ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1)) := by fun_prop have h' : HasCondDistrib (reward R' (n + 1)) @@ -126,92 +171,63 @@ lemma IsBayesianAlgEnvSeq.hasCondDistrib_reward_hist (h : IsBayesianAlgEnvSeq Q convert h'.comp_left hf using 2 -- (Claude) -lemma IsBayesianAlgEnvSeq.condIndepFun_action_env_hist [StandardBorelSpace Ω] - (h : IsBayesianAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - A (n + 1) ⟂ᵢ[hist A R' n, measurable_hist A R' h.measurable_A h.measurable_R n; P] (env R') := - condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (measurable_env R' h.measurable_R) - (h.measurable_A _) (measurable_hist A R' h.measurable_A h.measurable_R n) - (hasCondDistrib_action_env_hist Q κ P A R' alg h n).condDistrib_eq +lemma condIndepFun_action_env_hist [StandardBorelSpace Ω] (h : IsBayesAlgEnvSeq Q κ A R' alg P) + (n : ℕ) : A (n + 1) ⟂ᵢ[hist A R' n, h.measurable_hist n; P] (env R') := + condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft + (h.measurable_env) (h.measurable_A _) (h.measurable_hist n) + (hasCondDistrib_action_env_hist h n).condDistrib_eq -- Claude -lemma IsBayesianAlgEnvSeq.condIndepFun_reward_hist [StandardBorelSpace Ω] - (h : IsBayesianAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - reward R' (n + 1) - ⟂ᵢ[(fun ω ↦ (A (n + 1) ω, env R' ω)), - (h.measurable_A _).prodMk (measurable_env R' h.measurable_R); P] +lemma condIndepFun_reward_hist_action_env [StandardBorelSpace Ω] + (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : + reward R' (n + 1) ⟂ᵢ[action_env A R' (n + 1), h.measurable_action_env (n + 1); P] hist A R' n := condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft - (measurable_hist A R' h.measurable_A h.measurable_R n) - (measurable_reward R' h.measurable_R _) - ((h.measurable_A _).prodMk (measurable_env R' h.measurable_R)) - (hasCondDistrib_reward_hist Q κ P A R' alg h n).condDistrib_eq + (h.measurable_hist n) (h.measurable_reward _) (h.measurable_action_env _) + (hasCondDistrib_reward_hist_action_env h n).condDistrib_eq end Laws -section Real - -variable (κ : Kernel (α × E) ℝ) -variable (A : ℕ → Ω → α) (R' : ℕ → Ω → E × ℝ) (alg : Algorithm α ℝ) +end IsBayesAlgEnvSeq --- κ? -noncomputable -def IsBayesianAlgEnvSeq.actionValue (ω : Ω) (a : α) : ℝ := (κ (a, env R' ω))[id] +namespace IT -noncomputable -def IsBayesianAlgEnvSeq.bestAction [Nonempty α] [Fintype α] [Encodable α] - [MeasurableSingletonClass α] (ω : Ω) : α := measurableArgmax (actionValue κ R') ω +variable [StandardBorelSpace α] [Nonempty α] +variable [StandardBorelSpace E] [Nonempty E] +variable [StandardBorelSpace R] [Nonempty R] noncomputable -def IsBayesianAlgEnvSeq.condDistribBestAction [StandardBorelSpace α] [Nonempty α] [Fintype α] - [Encodable α] (n : ℕ) : Kernel (Iic n → α × ℝ) α := - condDistrib (bestAction κ R') (hist A R' n) P +def bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) + [IsMarkovKernel κ] (alg : Algorithm α R) : Measure (ℕ → α × E × R) := + trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) +deriving IsProbabilityMeasure -instance [StandardBorelSpace α] [Nonempty α] [Fintype α] - [Encodable α] (n : ℕ) : - IsMarkovKernel (IsBayesianAlgEnvSeq.condDistribBestAction P κ A R' n) := by - unfold IsBayesianAlgEnvSeq.condDistribBestAction - infer_instance +lemma isBayesAlgEnvSeq_bayesianTrajMeasure + (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] + (alg : Algorithm α R) : IsBayesAlgEnvSeq Q κ action reward alg (bayesTrajMeasure Q κ alg) := + isAlgEnvSeq_trajMeasure _ _ noncomputable -def IsBayesianAlgEnvSeq.regret (t : ℕ) (ω : Ω) : ℝ := - Bandits.regret (κ.comap (·, env R' ω) (by fun_prop)) A t ω - --- Claude -lemma IsBayesianAlgEnvSeq.measurable_regret [Fintype α] (hA : ∀ n, Measurable (A n)) - (hR' : ∀ n, Measurable (R' n)) (t : ℕ) : Measurable (regret κ A R' t) := by - unfold regret Bandits.regret - have hmean : Measurable fun (p : α × E) ↦ (κ p)[id] := - stronglyMeasurable_id.integral_kernel.measurable - apply Measurable.sub - · apply Measurable.const_mul - have h : ∀ a, Measurable fun (ω : Ω) ↦ (κ (a, env R' ω))[id] := fun a ↦ - hmean.comp (measurable_const.prodMk (measurable_env R' hR')) - exact Measurable.iSup h - · apply Finset.measurable_sum - intro s _ - exact hmean.comp ((hA s).prodMk (measurable_env R' hR')) - -noncomputable -def IsBayesianAlgEnvSeq.bayesianRegret - [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] - [IsMarkovKernel κ] (t : ℕ) : ℝ := - P[IsBayesianAlgEnvSeq.regret κ A R' t] - -end Real - -namespace IT +def posteriorBestArm [Fintype α] [Encodable α] (Q : Measure E) [IsProbabilityMeasure Q] + (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] (alg : Algorithm α ℝ) (n : ℕ) : + Kernel (Iic n → α × ℝ) α := + condDistrib (IsBayesAlgEnvSeq.bestArm κ reward) (IsBayesAlgEnvSeq.hist action reward n) + (IT.bayesTrajMeasure Q κ alg) +deriving IsMarkovKernel noncomputable -def bayesianTrajMeasure (alg : Algorithm α R) : Measure (ℕ → α × E × R) := - trajMeasure (alg.prod_left E) (BayesianStationaryEnv Q κ) -deriving IsProbabilityMeasure - -lemma isBayesianAlgEnvSeq_bayesianTrajMeasure - [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace E] [Nonempty E] - [StandardBorelSpace R] [Nonempty R] (alg : Algorithm α R) : - IsBayesianAlgEnvSeq Q κ action reward alg (bayesianTrajMeasure Q κ alg) := - ⟨isAlgEnvSeq_trajMeasure _ _⟩ +def priorBestArm [Fintype α] [Encodable α] (Q : Measure E) [IsProbabilityMeasure Q] + (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] (alg : Algorithm α ℝ) : Measure α := + (bayesTrajMeasure Q κ alg).map (IsBayesAlgEnvSeq.bestArm κ IT.reward) + +instance [Fintype α] [Encodable α] [MeasurableSingletonClass α] + (Q : Measure E) [IsProbabilityMeasure Q] + (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] (alg : Algorithm α ℝ) : + IsProbabilityMeasure (priorBestArm Q κ alg) := + Measure.isProbabilityMeasure_map + (measurable_measurableArgmax + (IsBayesAlgEnvSeq.measurable_armMean + (isBayesAlgEnvSeq_bayesianTrajMeasure Q κ alg))).aemeasurable end IT From d7f059426ea7ce0d70964c10a167a597a382d1e0 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 27 Jan 2026 11:48:58 +0000 Subject: [PATCH 021/155] Minor --- LeanBandits/BanditAlgorithms/TS.lean | 2 -- LeanBandits/SequentialLearning/BayesStationaryEnv.lean | 3 +-- 2 files changed, 1 insertion(+), 4 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index df4b5def..67f4c060 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -38,7 +38,6 @@ def tsInitPolicy : Measure (Fin K) := IT.priorBestArm Q κ (uniformAlgorithm hK) instance : IsProbabilityMeasure (tsInitPolicy hK Q κ) := by - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK unfold tsInitPolicy infer_instance @@ -50,7 +49,6 @@ def tsAlgorithm : Algorithm (Fin K) ℝ where policy := tsPolicy hK Q κ p0 := tsInitPolicy hK Q κ - variable {Ω : Type*} [MeasurableSpace Ω] variable {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → E × ℝ} variable {P : Measure Ω} [IsFiniteMeasure P] diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 2e65b9cf..4fce3a2e 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -220,8 +220,7 @@ def priorBestArm [Fintype α] [Encodable α] (Q : Measure E) [IsProbabilityMeasu (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] (alg : Algorithm α ℝ) : Measure α := (bayesTrajMeasure Q κ alg).map (IsBayesAlgEnvSeq.bestArm κ IT.reward) -instance [Fintype α] [Encodable α] [MeasurableSingletonClass α] - (Q : Measure E) [IsProbabilityMeasure Q] +instance [Fintype α] [Encodable α] (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] (alg : Algorithm α ℝ) : IsProbabilityMeasure (priorBestArm Q κ alg) := Measure.isProbabilityMeasure_map From 0211c1afcffe9ea4dc75b5b48b80bb1df9f465b4 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 27 Jan 2026 16:17:59 +0000 Subject: [PATCH 022/155] Finish refactor --- LeanBandits/BanditAlgorithms/TS.lean | 53 ++++--- LeanBandits/ForMathlib/HasCondDistrib.lean | 6 +- .../BayesStationaryEnv.lean | 150 +++++++++--------- 3 files changed, 110 insertions(+), 99 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 67f4c060..ce6c467f 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -4,63 +4,66 @@ Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne, Paulo Rauber -/ import LeanBandits.ForMathlib.MeasurableArgMax +import LeanBandits.ForMathlib.SubGaussian import LeanBandits.BanditAlgorithms.Uniform import LeanBandits.SequentialLearning.BayesStationaryEnv -import LeanBandits.ForMathlib.SubGaussian open MeasureTheory ProbabilityTheory Finset Learning namespace Bandits -variable {K : ℕ} -variable {E : Type*} [mE : MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] - -variable (hK : 0 < K) -variable (Q : Measure E) [IsProbabilityMeasure Q] -variable (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] - namespace TS +variable {K : ℕ} (hK : 0 < K) +variable {E : Type*} [mE : MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] +variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] + /-- The distribution over actions for every given history for TS. -/ noncomputable -def tsPolicy (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := +def policy (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK IT.posteriorBestArm Q κ (uniformAlgorithm hK) n - -instance (n : ℕ) : IsMarkovKernel (tsPolicy hK Q κ n) := by - unfold tsPolicy - infer_instance +deriving IsMarkovKernel /-- The initial distribution over actions for TS. -/ noncomputable -def tsInitPolicy : Measure (Fin K) := +def initialPolicy : Measure (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK IT.priorBestArm Q κ (uniformAlgorithm hK) -instance : IsProbabilityMeasure (tsInitPolicy hK Q κ) := by - unfold tsInitPolicy +instance : IsProbabilityMeasure (initialPolicy hK Q κ) := by + unfold initialPolicy infer_instance +end TS + +variable {K : ℕ} +variable {E : Type*} [mE : MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] + /-- The Thompson Sampling (TS) algorithm: actions are chosen according to the probability that they are optimal given prior knowledge represented by a prior distribution `Q` and a data generation model represented by a kernel `κ`. -/ noncomputable -def tsAlgorithm : Algorithm (Fin K) ℝ where - policy := tsPolicy hK Q κ - p0 := tsInitPolicy hK Q κ +def tsAlgorithm (hK : 0 < K) (Q : Measure E) [IsProbabilityMeasure Q] + (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] : Algorithm (Fin K) ℝ where + policy := TS.policy hK Q κ + p0 := TS.initialPolicy hK Q κ + +section Regret +variable (hK : 0 < K) variable {Ω : Type*} [MeasurableSpace Ω] -variable {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → E × ℝ} -variable {P : Measure Ω} [IsFiniteMeasure P] +variable (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → E × ℝ) +variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] +variable (P : Measure Ω) [IsFiniteMeasure P] -def bayesian_regret_le [Nonempty (Fin K)] +lemma TS.bayesRegret_le [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) : - ∃ C > 0, ∀ n : ℕ, - (IsBayesAlgEnvSeq.bayesRegret κ A R' P n) ≤ C * √(K * n * Real.log n) := + ∃ C, ∀ n, (IsBayesAlgEnvSeq.bayesRegret κ A R' P n) ≤ C * √(K * n * Real.log n) := sorry -end TS +end Regret end Bandits diff --git a/LeanBandits/ForMathlib/HasCondDistrib.lean b/LeanBandits/ForMathlib/HasCondDistrib.lean index 19dfec6d..18e78a84 100644 --- a/LeanBandits/ForMathlib/HasCondDistrib.lean +++ b/LeanBandits/ForMathlib/HasCondDistrib.lean @@ -182,7 +182,7 @@ lemma hasCondDistrib_prod_right_iff [IsFiniteMeasure μ] [IsFiniteKernel κ] (X rw [← Measure.map_prod_map _ _ (by fun_prop) (by fun_prop), Measure.map_id, Measure.map_dirac (by fun_prop)] --- Revise (Claude) +-- Claude lemma HasCondDistrib.hasLaw_of_const {Q : Measure Ω} [IsProbabilityMeasure μ] [IsProbabilityMeasure Q] (h : HasCondDistrib Y X (Kernel.const β Q) μ) : HasLaw Y Q μ where @@ -206,7 +206,7 @@ lemma HasLaw.prod_of_hasCondDistrib {P : Measure β} [IsFiniteMeasure μ] [IsSFi rw [← h1.map_eq] exact h2.condDistrib_eq --- Revise (Claude) +-- Claude lemma HasCondDistrib.of_compProd [IsFiniteMeasure μ] [IsFiniteKernel κ] {Z : α → Ω'} {η : Kernel (β × Ω) Ω'} [IsMarkovKernel η] (h : HasCondDistrib (fun ω ↦ (Y ω, Z ω)) X (κ ⊗ₖ η) μ) : @@ -255,7 +255,7 @@ lemma HasCondDistrib.prod [IsFiniteMeasure μ] [IsFiniteKernel κ] AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] rfl --- Revise (Claude) +-- Claude lemma HasCondDistrib.comp_left [IsFiniteMeasure μ] [IsFiniteKernel κ] {f : γ → β} (hf : Measurable f) {Z : α → γ} (h : HasCondDistrib Y Z (κ.comap f hf) μ) : HasCondDistrib Y (f ∘ Z) κ μ where diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 4fce3a2e..e0fac092 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -13,6 +13,10 @@ namespace Learning variable {α E R : Type*} [mα : MeasurableSpace α] [mE : MeasurableSpace E] [mR : MeasurableSpace R] +/-- Given a prior distribution `Q` over "environments" and a kernel `k` that defines a reward +distribution `κ (a, e)` for each action `a : α` and "environment" `e : E`, a `bayesStationaryEnv` +corresponds to an environment (with an observation space `E × R`) that draws an "environment" +`e : E` at the very first step and defines a stationary environment from `k (·, e)`. -/ noncomputable def bayesStationaryEnv (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] : @@ -24,6 +28,8 @@ def bayesStationaryEnv variable {Ω : Type*} [mΩ : MeasurableSpace Ω] +/-- A Bayesian algorithm-environment sequence: a sequence of actions and observations from an +algorithm that ignores the underlying "environment" while interacting with `bayesStationaryEnv`. -/ def IsBayesAlgEnvSeq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] @@ -44,26 +50,29 @@ variable {A : ℕ → Ω → α} {R' : ℕ → Ω → E × R} variable {alg : Algorithm α R} variable {P : Measure Ω} [IsProbabilityMeasure P] +/-- The underlying "environment". -/ def env (R' : ℕ → Ω → E × R) (ω : Ω) : E := (R' 0 ω).1 @[fun_prop] lemma measurable_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) : Measurable (env R') := (h.measurable_R 0).fst +/-- The reward at time `n`. -/ def reward (R' : ℕ → Ω → E × R) (n : ℕ) (ω : Ω) : R := (R' n ω).2 @[fun_prop] -lemma measurable_reward (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - Measurable (reward R' n) := +lemma measurable_reward (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : Measurable (reward R' n) := (h.measurable_R n).snd +/-- The history of actions and rewards up to time `n`. -/ def hist (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (n : ℕ) (ω : Ω) : Iic n → α × R := fun i ↦ (A i ω, (R' i ω).2) @[fun_prop] lemma measurable_hist (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : Measurable (hist A R' n) := - measurable_pi_lambda _ fun i => (h.measurable_A i).prodMk (h.measurable_R i).snd + measurable_pi_iff.2 fun i => ((h.measurable_A i).prodMk (h.measurable_R i).snd) +/-- The action at time `n` together with the underlying "environment". -/ def action_env (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (n : ℕ) (ω : Ω) : α × E := (A n ω, (R' 0 ω).1) @[fun_prop] @@ -77,40 +86,37 @@ variable {R' : ℕ → Ω → E × ℝ} variable {κ : Kernel (α × E) ℝ} [IsMarkovKernel κ] variable {alg : Algorithm α ℝ} +/-- The mean of action `a : α` in the underlying "environment". -/ noncomputable -def armMean (κ : Kernel (α × E) ℝ) (R' : ℕ → Ω → E × ℝ) (a : α) (ω : Ω) : ℝ := - (κ (a, env R' ω))[id] +def armMean (κ : Kernel (α × E) ℝ) (R' : ℕ → Ω → E × ℝ) (a : α) (ω : Ω) : ℝ := (κ (a, env R' ω))[id] lemma measurable_armMean (h : IsBayesAlgEnvSeq Q κ A R' alg P) (a : α) : Measurable (armMean κ R' a) := - stronglyMeasurable_id.integral_kernel.measurable.comp - (measurable_const.prodMk h.measurable_env) + stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk h.measurable_env) +/-- An action with the highest mean in the underlying "environment". -/ noncomputable -def bestArm (κ : Kernel (α × E) ℝ) [Fintype α] [Encodable α] (R' : ℕ → Ω → E × ℝ) := +def bestArm [Fintype α] [Encodable α] (κ : Kernel (α × E) ℝ) (R' : ℕ → Ω → E × ℝ) := measurableArgmax (fun ω a ↦ armMean κ R' a ω) -lemma measurable_bestArm [Fintype α] [Encodable α] [MeasurableSingletonClass α] - (h : IsBayesAlgEnvSeq Q κ A R' alg P) : +lemma measurable_bestArm [Fintype α] [Encodable α] (h : IsBayesAlgEnvSeq Q κ A R' alg P) : Measurable (bestArm κ R') := measurable_measurableArgmax h.measurable_armMean +/-- Regret of a sequence of pulls at time `t` considering the underlying "environment". -/ noncomputable def regret (κ : Kernel (α × E) ℝ) (A : ℕ → Ω → α) (R' : ℕ → Ω → E × ℝ) (t : ℕ) (ω : Ω) : ℝ := Bandits.regret (κ.comap (·, env R' ω) (by fun_prop)) A t ω --- -- Claude -lemma measurable_regret [Fintype α] [Encodable α] (h : IsBayesAlgEnvSeq Q κ A R' alg P) (t : ℕ) : +lemma measurable_regret [Encodable α] (h : IsBayesAlgEnvSeq Q κ A R' alg P) (t : ℕ) : Measurable (regret κ A R' t) := by - unfold regret Bandits.regret apply Measurable.sub - · apply Measurable.const_mul - exact Measurable.iSup h.measurable_armMean - · apply Finset.measurable_sum - intro s _ - exact stronglyMeasurable_id.integral_kernel.measurable.comp - ((h.measurable_A s).prodMk h.measurable_env) + · exact Measurable.const_mul (Measurable.iSup h.measurable_armMean) _ + · exact Finset.measurable_sum _ fun s _ ↦ + stronglyMeasurable_id.integral_kernel.measurable.comp (h.measurable_action_env s) +/-- The expected regret according to `P`, implicitly assuming that `IsBayesAlgEnvSeq Q κ A R' alg P` + for some prior distribution over "environments" `Q` and some algorithm `alg`. -/ noncomputable def bayesRegret (κ : Kernel (α × E) ℝ) (A : ℕ → Ω → α) (R' : ℕ → Ω → E × ℝ) (P : Measure Ω) (t : ℕ) : ℝ := @@ -124,105 +130,107 @@ lemma hasLaw_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) : HasLaw (env R') Q P := apply HasCondDistrib.hasLaw_of_const simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst --- (Claude) lemma hasCondDistrib_action' (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : HasCondDistrib (A (n + 1)) (hist A R' n) (alg.policy n) P := (h.hasCondDistrib_action n).comp_left (by fun_prop) ---(Claude) lemma hasCondDistrib_reward_zero' (h : IsBayesAlgEnvSeq Q κ A R' alg P) : HasCondDistrib (reward R' 0) (action_env A R' 0) κ P := by simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.of_compProd --- (Claude) lemma hasCondDistrib_reward' (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : HasCondDistrib (reward R' (n + 1)) (action_env A R' (n + 1)) κ P := by - have h := (h.hasCondDistrib_reward n).snd - simp_rw [bayesStationaryEnv, Kernel.snd_prod] at h - exact h.comp_left (by fun_prop) + have hr := (h.hasCondDistrib_reward n).snd + simp_rw [bayesStationaryEnv, Kernel.snd_prod] at hr + exact hr.comp_left (by fun_prop) --- (Claude) +-- Auxiliar lemma for `condIndepFun_action_env_hist` (Claude) lemma hasCondDistrib_action_env_hist (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : HasCondDistrib (A (n + 1)) (fun ω ↦ (env R' ω, hist A R' n ω)) ((alg.policy n).prodMkLeft E) P := by - have h := h.hasCondDistrib_action n - simp only [Algorithm.prod_left, Kernel.prodMkLeft] at h ⊢ - have hf : Measurable (fun h : Iic n → α × E × R ↦ - ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2))) := by fun_prop - have h' : HasCondDistrib (A (n + 1)) (Learning.IsAlgEnvSeq.hist A R' n) - (((alg.policy n).comap Prod.snd (by fun_prop)).comap - (fun h : Iic n → α × E × R ↦ ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2))) hf) - P := by convert h using 2 - convert h'.comp_left hf using 2 - --- -- (Claude) + let f : (Iic n → α × E × R) → E × (Iic n → α × R) := + fun h ↦ ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2)) + suffices h' : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) + (((alg.policy n).comap Prod.snd (by fun_prop)).comap f (by fun_prop)) P from h'.comp_left + exact h.hasCondDistrib_action n + +-- Auxiliar lemma for `condIndepFun_reward_hist_action_env` (Claude) lemma hasCondDistrib_reward_hist_action_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : HasCondDistrib (reward R' (n + 1)) (fun ω ↦ (hist A R' n ω, A (n + 1) ω, env R' ω)) (κ.prodMkLeft _) P := by - have h := (h.hasCondDistrib_reward n).snd - simp_rw [bayesStationaryEnv, Kernel.snd_prod, Kernel.prodMkLeft] at h ⊢ - have hf : Measurable (fun (p : (Iic n → α × (E × R)) × α) ↦ - ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1)) := by fun_prop - have h' : HasCondDistrib (reward R' (n + 1)) - (fun ω ↦ (Learning.IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) - ((κ.comap Prod.snd (by fun_prop)).comap (fun (p : (Iic n → α × (E × R)) × α) ↦ - ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1)) hf) - P := by convert h using 2 - convert h'.comp_left hf using 2 - --- (Claude) + let f : (Iic n → α × (E × R)) × α → (Iic n → α × R) × α × E := + fun p ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1) + have hf : Measurable f := by fun_prop + suffices h' : HasCondDistrib (reward R' (n + 1)) + (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) + ((κ.comap Prod.snd (by fun_prop)).comap f hf) P from by + convert h'.comp_left hf + simpa [bayesStationaryEnv, Kernel.snd_prod] using (h.hasCondDistrib_reward n).snd + +end Laws + +section Independence + +lemma indepFun_action_zero_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) : + IndepFun (A 0) (env R') P := by + rw [indepFun_iff_condDistrib_eq_const (h.measurable_A 0).aemeasurable + h.measurable_env.aemeasurable, h.hasLaw_env.map_eq] + simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst.condDistrib_eq + lemma condIndepFun_action_env_hist [StandardBorelSpace Ω] (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : A (n + 1) ⟂ᵢ[hist A R' n, h.measurable_hist n; P] (env R') := condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (h.measurable_env) (h.measurable_A _) (h.measurable_hist n) (hasCondDistrib_action_env_hist h n).condDistrib_eq --- Claude lemma condIndepFun_reward_hist_action_env [StandardBorelSpace Ω] (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - reward R' (n + 1) ⟂ᵢ[action_env A R' (n + 1), h.measurable_action_env (n + 1); P] - hist A R' n := + reward R' (n + 1) ⟂ᵢ[action_env A R' (n + 1), h.measurable_action_env (n + 1); P] hist A R' n := condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (h.measurable_hist n) (h.measurable_reward _) (h.measurable_action_env _) (hasCondDistrib_reward_hist_action_env h n).condDistrib_eq -end Laws +end Independence end IsBayesAlgEnvSeq namespace IT -variable [StandardBorelSpace α] [Nonempty α] -variable [StandardBorelSpace E] [Nonempty E] -variable [StandardBorelSpace R] [Nonempty R] - +/-- Measure on the sequence of actions and observations generated by an algorithm that ignores the +underlying "environment" while interacting with a `bayesStationaryEnv`. -/ noncomputable def bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] (alg : Algorithm α R) : Measure (ℕ → α × E × R) := - trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) + trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) deriving IsProbabilityMeasure lemma isBayesAlgEnvSeq_bayesianTrajMeasure - (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] - (alg : Algorithm α R) : IsBayesAlgEnvSeq Q κ action reward alg (bayesTrajMeasure Q κ alg) := + [StandardBorelSpace α] [Nonempty α] + [StandardBorelSpace E] [Nonempty E] + [StandardBorelSpace R] [Nonempty R] + (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] + (alg : Algorithm α R) : IsBayesAlgEnvSeq Q κ action reward alg (bayesTrajMeasure Q κ alg) := isAlgEnvSeq_trajMeasure _ _ +/-- The conditional distribution over the best arm given the observed history. -/ noncomputable -def posteriorBestArm [Fintype α] [Encodable α] (Q : Measure E) [IsProbabilityMeasure Q] - (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] (alg : Algorithm α ℝ) (n : ℕ) : - Kernel (Iic n → α × ℝ) α := +def posteriorBestArm [StandardBorelSpace α] [Nonempty α] [Fintype α] [Encodable α] + (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] + (alg : Algorithm α ℝ) (n : ℕ) : Kernel (Iic n → α × ℝ) α := condDistrib (IsBayesAlgEnvSeq.bestArm κ reward) (IsBayesAlgEnvSeq.hist action reward n) - (IT.bayesTrajMeasure Q κ alg) + (bayesTrajMeasure Q κ alg) deriving IsMarkovKernel +/-- The initial distribution over the best arm. -/ noncomputable -def priorBestArm [Fintype α] [Encodable α] (Q : Measure E) [IsProbabilityMeasure Q] - (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] (alg : Algorithm α ℝ) : Measure α := - (bayesTrajMeasure Q κ alg).map (IsBayesAlgEnvSeq.bestArm κ IT.reward) - -instance [Fintype α] [Encodable α] (Q : Measure E) [IsProbabilityMeasure Q] - (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] (alg : Algorithm α ℝ) : - IsProbabilityMeasure (priorBestArm Q κ alg) := +def priorBestArm [StandardBorelSpace α] [Nonempty α] [Fintype α] [Encodable α] + (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] + (alg : Algorithm α ℝ) : Measure α := + (bayesTrajMeasure Q κ alg).map (IsBayesAlgEnvSeq.bestArm κ reward) + +instance [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] [Fintype α] + [Encodable α] (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] + (alg : Algorithm α ℝ) : IsProbabilityMeasure (priorBestArm Q κ alg) := Measure.isProbabilityMeasure_map (measurable_measurableArgmax (IsBayesAlgEnvSeq.measurable_armMean From 3ba0128ecd329d0308bf879bf6d43bd1a4d63420 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 27 Jan 2026 16:42:15 +0000 Subject: [PATCH 023/155] Minor --- LeanBandits/BanditAlgorithms/TS.lean | 2 ++ LeanBandits/BanditAlgorithms/Uniform.lean | 2 ++ LeanBandits/SequentialLearning/BayesStationaryEnv.lean | 8 ++++++-- 3 files changed, 10 insertions(+), 2 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index ce6c467f..a9d08de7 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -8,6 +8,8 @@ import LeanBandits.ForMathlib.SubGaussian import LeanBandits.BanditAlgorithms.Uniform import LeanBandits.SequentialLearning.BayesStationaryEnv +/-! # The Thompson Sampling Algorithm -/ + open MeasureTheory ProbabilityTheory Finset Learning namespace Bandits diff --git a/LeanBandits/BanditAlgorithms/Uniform.lean b/LeanBandits/BanditAlgorithms/Uniform.lean index 7e8208b2..d422b689 100644 --- a/LeanBandits/BanditAlgorithms/Uniform.lean +++ b/LeanBandits/BanditAlgorithms/Uniform.lean @@ -6,6 +6,8 @@ Authors: Rémy Degenne, Paulo Rauber import Mathlib.Probability.UniformOn import LeanBandits.SequentialLearning.Algorithm +/-! # The Uniform Algorithm -/ + open MeasureTheory ProbabilityTheory Learning namespace Bandits diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index e0fac092..68592a80 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -7,6 +7,8 @@ import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.Bandit.Regret import LeanBandits.SequentialLearning.IonescuTulceaSpace +/-! # Bayesian stationary environments -/ + open MeasureTheory ProbabilityTheory Finset namespace Learning @@ -90,6 +92,7 @@ variable {alg : Algorithm α ℝ} noncomputable def armMean (κ : Kernel (α × E) ℝ) (R' : ℕ → Ω → E × ℝ) (a : α) (ω : Ω) : ℝ := (κ (a, env R' ω))[id] +@[fun_prop] lemma measurable_armMean (h : IsBayesAlgEnvSeq Q κ A R' alg P) (a : α) : Measurable (armMean κ R' a) := stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk h.measurable_env) @@ -99,6 +102,7 @@ noncomputable def bestArm [Fintype α] [Encodable α] (κ : Kernel (α × E) ℝ) (R' : ℕ → Ω → E × ℝ) := measurableArgmax (fun ω a ↦ armMean κ R' a ω) +@[fun_prop] lemma measurable_bestArm [Fintype α] [Encodable α] (h : IsBayesAlgEnvSeq Q κ A R' alg P) : Measurable (bestArm κ R') := measurable_measurableArgmax h.measurable_armMean @@ -144,7 +148,7 @@ lemma hasCondDistrib_reward' (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : simp_rw [bayesStationaryEnv, Kernel.snd_prod] at hr exact hr.comp_left (by fun_prop) --- Auxiliar lemma for `condIndepFun_action_env_hist` (Claude) +-- Auxiliary lemma for `condIndepFun_action_env_hist` (Claude) lemma hasCondDistrib_action_env_hist (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : HasCondDistrib (A (n + 1)) (fun ω ↦ (env R' ω, hist A R' n ω)) ((alg.policy n).prodMkLeft E) P := by @@ -154,7 +158,7 @@ lemma hasCondDistrib_action_env_hist (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : (((alg.policy n).comap Prod.snd (by fun_prop)).comap f (by fun_prop)) P from h'.comp_left exact h.hasCondDistrib_action n --- Auxiliar lemma for `condIndepFun_reward_hist_action_env` (Claude) +-- Auxiliary lemma for `condIndepFun_reward_hist_action_env` (Claude) lemma hasCondDistrib_reward_hist_action_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : HasCondDistrib (reward R' (n + 1)) (fun ω ↦ (hist A R' n ω, A (n + 1) ω, env R' ω)) (κ.prodMkLeft _) P := by From 8f0f5fee312a37033d163f936d6c2633e4c055a1 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 28 Jan 2026 13:09:14 +0000 Subject: [PATCH 024/155] Fix documentation --- LeanBandits/BanditAlgorithms/TS.lean | 1 - LeanBandits/SequentialLearning/BayesStationaryEnv.lean | 9 ++++----- 2 files changed, 4 insertions(+), 6 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index a9d08de7..0bc2bff3 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -3,7 +3,6 @@ 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.MeasurableArgMax import LeanBandits.ForMathlib.SubGaussian import LeanBandits.BanditAlgorithms.Uniform import LeanBandits.SequentialLearning.BayesStationaryEnv diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 68592a80..bd4436a4 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -119,8 +119,8 @@ lemma measurable_regret [Encodable α] (h : IsBayesAlgEnvSeq Q κ A R' alg P) (t · exact Finset.measurable_sum _ fun s _ ↦ stronglyMeasurable_id.integral_kernel.measurable.comp (h.measurable_action_env s) -/-- The expected regret according to `P`, implicitly assuming that `IsBayesAlgEnvSeq Q κ A R' alg P` - for some prior distribution over "environments" `Q` and some algorithm `alg`. -/ +/-- If `IsBayesAlgEnvSeq Q κ A R' alg P`, then `bayesRegret κ A R' P t` is the expected +regret at time `t` of the algorithm `alg` given a prior distribution over "environments" `Q`. -/ noncomputable def bayesRegret (κ : Kernel (α × E) ℝ) (A : ℕ → Ω → α) (R' : ℕ → Ω → E × ℝ) (P : Measure Ω) (t : ℕ) : ℝ := @@ -162,13 +162,12 @@ lemma hasCondDistrib_action_env_hist (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : lemma hasCondDistrib_reward_hist_action_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : HasCondDistrib (reward R' (n + 1)) (fun ω ↦ (hist A R' n ω, A (n + 1) ω, env R' ω)) (κ.prodMkLeft _) P := by - let f : (Iic n → α × (E × R)) × α → (Iic n → α × R) × α × E := + let f : (Iic n → α × E × R) × α → (Iic n → α × R) × α × E := fun p ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1) have hf : Measurable f := by fun_prop suffices h' : HasCondDistrib (reward R' (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) - ((κ.comap Prod.snd (by fun_prop)).comap f hf) P from by - convert h'.comp_left hf + ((κ.comap Prod.snd (by fun_prop)).comap f hf) P from h'.comp_left hf simpa [bayesStationaryEnv, Kernel.snd_prod] using (h.hasCondDistrib_reward n).snd end Laws From 56111b0b497ceb4de72fc3e7e44f7ecff37068d7 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 2 Feb 2026 10:03:39 +0000 Subject: [PATCH 025/155] Change statement --- LeanBandits/BanditAlgorithms/TS.lean | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 0bc2bff3..8db83a83 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -61,8 +61,8 @@ variable (P : Measure Ω) [IsFiniteMeasure P] lemma TS.bayesRegret_le [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) - (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) : - ∃ C, ∀ n, (IsBayesAlgEnvSeq.bayesRegret κ A R' P n) ≤ C * √(K * n * Real.log n) := + (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) (t : ℕ) : + IsBayesAlgEnvSeq.bayesRegret κ A R' P t ≤ 4 * K + 8 * √(K * t * Real.log t) := sorry end Regret From bba43969116895523692ebb617ab8961acfcdee8 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 10 Feb 2026 10:56:09 +0000 Subject: [PATCH 026/155] Complete Thompson sammpling Bayesian regret bound --- LeanBandits.lean | 1 + LeanBandits/Bandit/SumRewards.lean | 166 ++- LeanBandits/BanditAlgorithms/TS.lean | 1081 ++++++++++++++++- LeanBandits/BanditAlgorithms/UCB.lean | 4 +- LeanBandits/ForMathlib/Measurable.lean | 13 + LeanBandits/ForMathlib/MeasurableArgMax.lean | 50 +- .../BayesStationaryEnv.lean | 206 ++++ .../SequentialLearning/FiniteActions.lean | 44 +- .../SequentialLearning/HistoryDensity.lean | 568 +++++++++ 9 files changed, 2041 insertions(+), 92 deletions(-) create mode 100644 LeanBandits/SequentialLearning/HistoryDensity.lean diff --git a/LeanBandits.lean b/LeanBandits.lean index b566505d..ed67f3d0 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -25,5 +25,6 @@ import LeanBandits.SequentialLearning.Algorithm import LeanBandits.SequentialLearning.BayesStationaryEnv import LeanBandits.SequentialLearning.Deterministic import LeanBandits.SequentialLearning.FiniteActions +import LeanBandits.SequentialLearning.HistoryDensity import LeanBandits.SequentialLearning.IonescuTulceaSpace import LeanBandits.SequentialLearning.StationaryEnv diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 81ec80a3..318a577e 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -77,6 +77,15 @@ lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount' (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) +private lemma sum_Icc_one_eq_sum_range {m : ℕ} {f : ℕ → ℝ} : + ∑ i ∈ Icc 1 m, f (i - 1) = ∑ i ∈ range m, f i := by + have h : Icc 1 m = (range m).image (· + 1) := by + ext x; simp only [mem_Icc, mem_image, mem_range]; constructor + · intro ⟨h1, h2⟩; exact ⟨x - 1, by omega, by omega⟩ + · rintro ⟨a, ha, rfl⟩; omega + rw [h, Finset.sum_image (fun _ _ _ _ h => by omega)] + simp + 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 ω)) @@ -88,19 +97,7 @@ lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount (n : ℕ) : · infer_instance ext a : 1 congr 1 - 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] + exact sum_Icc_one_eq_sum_range.symm lemma identDistrib_pullCount_prod_sumRewards (n : ℕ) : IdentDistrib (fun ω a ↦ (pullCount A a n ω, sumRewards A R a n ω)) @@ -171,6 +168,32 @@ lemma identDistrib_sum_range_snd (a : α) (k : ℕ) : (ν := streamMeasure ν), Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)] rfl +omit [Countable α] in +lemma sumRewards_eq_sum_stream_of_pullCount_eq (a : α) (s m : ℕ) (ω : probSpace α ℝ) + (hpc : pullCount A a s ω = m) : + sumRewards A R a s ω = ∑ i ∈ range m, ω.2 i a := by + let ω' : probSpace α ℝ × (ℕ → α → ℝ) := (ω, ω.2) + have h_sum_rbc : sumRewards A R a s ω = ∑ i ∈ Icc 1 m, rewardByCount A R a i ω' := by + rw [← sum_rewardByCount_eq_sumRewards a s ω', hpc] + rw [h_sum_rbc] + have h_rbc_eq (i : ℕ) (hi : i ∈ Icc 1 m) : rewardByCount A R a i ω' = ω.2 (i - 1) a := by + have hi' := mem_Icc.mp hi + have hi_ne : i ≠ 0 := Nat.one_le_iff_ne_zero.mp hi'.1 + have h_i_le : i ≤ pullCount A a s ω := hpc ▸ hi'.2 + have hs_pos : 0 < s := + Nat.pos_of_ne_zero (by rintro rfl; simp [pullCount] at hpc; omega) + have h_exists : ∃ t, pullCount A a (t + 1) ω = i := + exists_pullCount_eq_of_le (n := s - 1) (Nat.sub_add_cancel hs_pos ▸ h_i_le) hi_ne + rw [rewardByCount_of_stepsUntil_ne_top (stepsUntil_ne_top h_exists)] + simp only [reward_eq] + have h_action : A (stepsUntil A a i ω).toNat ω = a := + action_stepsUntil («A» := A) hi_ne h_exists + congr! + rw [h_action, pullCount_stepsUntil hi_ne h_exists] + calc ∑ i ∈ Icc 1 m, rewardByCount A R a i ω' + _ = ∑ i ∈ Icc 1 m, ω.2 (i - 1) a := Finset.sum_congr rfl h_rbc_eq + _ = ∑ j ∈ range m, ω.2 j a := sum_Icc_one_eq_sum_range (f := fun i => ω.2 i a) + 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} ≤ @@ -222,6 +245,24 @@ lemma prob_pullCount_mem_and_sumRewards_mem_le (a : α) (n : ℕ) exists_eq_right, mem_filter, mem_range] at hk simp [hk.2.1] +lemma prob_exists_pullCount_eq_and_sumRewards_mem_le (a : α) (n m : ℕ) + {B : Set ℝ} (hB : MeasurableSet B) : + 𝔓 {ω | ∃ s, s ≤ n ∧ pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B} ≤ + streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by + calc 𝔓 {ω | ∃ s, s ≤ n ∧ pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B} + _ ≤ 𝔓 {ω | ∑ i ∈ range m, ω.2 i a ∈ B} := by + -- Show the containment: the existential set ⊆ {sum ∈ B} + apply measure_mono + intro ω ⟨s, _hs, hpc, hB'⟩ + -- When pullCount(s, ω) = m, sumRewards(s, ω) = ∑ i < m, ω.2 i a in the ArrayModel. + rw [sumRewards_eq_sum_stream_of_pullCount_eq a s m ω hpc] at hB' + exact hB' + _ = streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by + have := (identDistrib_sum_range_snd (ν := ν) a m).map_eq + rw [Measure.ext_iff] at this + specialize this B hB + rwa [Measure.map_apply (by fun_prop) hB, Measure.map_apply (by fun_prop) hB] at this + 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 ω} ≤ @@ -363,38 +404,9 @@ 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) + P'.map (fun ω ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) := + ((h1.law_pullCount_sumRewards_unique' h2 (n := n)).comp + (u := fun f ↦ f a) (by fun_prop)).map_eq -- this is what we will use for UCB lemma prob_pullCount_prod_sumRewards_mem_le [Countable α] @@ -454,6 +466,66 @@ lemma prob_pullCount_eq_and_sumRewards_mem_le [Countable α] have hm' : m < n + 1 := by lia simpa [hm'] using h_le +lemma prob_exists_pullCount_eq_and_sumRewards_mem_le [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + {m : ℕ} {B : Set ℝ} (hB : MeasurableSet B) : + P {ω | ∃ s, s ≤ n ∧ pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B} ≤ + streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by + have hA := h.measurable_A + have hR := h.measurable_R + have h_eq : {ω | ∃ s, s ≤ n ∧ pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B} = + ⋃ s ∈ range (n + 1), {ω | pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B} := by + ext ω; simp only [Set.mem_setOf_eq, Set.mem_iUnion, mem_range] + constructor <;> rintro ⟨s, hs, rest⟩ <;> exact ⟨s, by omega, rest⟩ + rw [h_eq] + have h_AM := ArrayModel.prob_exists_pullCount_eq_and_sumRewards_mem_le + (ν := ν) (alg := alg) a n m hB + let pc := fun (p : ℕ → α × ℝ) (s : ℕ) ↦ ∑ i ∈ range s, if (p i).1 = a then 1 else 0 + let sr := fun (p : ℕ → α × ℝ) (s : ℕ) ↦ ∑ i ∈ range s, if (p i).1 = a then (p i).2 else 0 + let S := ⋃ s ∈ range (n + 1), {p : ℕ → α × ℝ | pc p s = m ∧ sr p s ∈ B} + have hS : MeasurableSet S := by + simp only [S] + apply MeasurableSet.iUnion + intro s + apply MeasurableSet.iUnion + intro _ + apply MeasurableSet.inter + · exact (measurableSet_singleton _).preimage + (measurable_sum _ fun i _ ↦ Measurable.ite + ((measurableSet_singleton _).preimage (by fun_prop)) (by fun_prop) (by fun_prop)) + · exact hB.preimage + (measurable_sum _ fun i _ ↦ Measurable.ite + ((measurableSet_singleton _).preimage (by fun_prop)) (by fun_prop) (by fun_prop)) + have h_eq1 : (⋃ s ∈ range (n + 1), {ω | pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B}) = + (fun ω t ↦ (A t ω, R t ω)) ⁻¹' S := by + ext ω + simp only [Set.mem_iUnion, mem_range, Set.mem_setOf_eq, Set.mem_preimage, S, pc, sr, + pullCount, sumRewards, Finset.card_filter] + have h_eq2 : (⋃ s ∈ range (n + 1), {ω | pullCount (ArrayModel.action alg) a s ω = m ∧ + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a s ω ∈ B}) = + (fun ω t ↦ (ArrayModel.action alg t ω, ArrayModel.reward alg t ω)) ⁻¹' S := by + ext ω + simp only [Set.mem_iUnion, mem_range, Set.mem_setOf_eq, Set.mem_preimage, S, pc, sr, + pullCount, sumRewards, Finset.card_filter] + have h_unique := isAlgEnvSeq_unique h (ArrayModel.isAlgEnvSeq_arrayMeasure alg ν) + calc P (⋃ s ∈ range (n + 1), {ω | pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B}) + _ = (ArrayModel.arrayMeasure ν) + (⋃ s ∈ range (n + 1), {ω | pullCount (ArrayModel.action alg) a s ω = m ∧ + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a s ω ∈ B}) := by + rw [h_eq1, h_eq2, ← Measure.map_apply _ hS, ← Measure.map_apply _ hS, h_unique] + · rw [measurable_pi_iff]; intro t; exact (by fun_prop : Measurable fun ω ↦ + (ArrayModel.action alg t ω, ArrayModel.reward alg t ω)) + · rw [measurable_pi_iff]; intro t; exact (hA t).prodMk (hR t) + _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by + have h_set_eq : (⋃ s ∈ range (n + 1), {ω | pullCount (ArrayModel.action alg) a s ω = m ∧ + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a s ω ∈ B}) = + {ω | ∃ s, s ≤ n ∧ pullCount (ArrayModel.action alg) a s ω = m ∧ + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a s ω ∈ B} := by + ext ω; simp only [Set.mem_iUnion, mem_range, Set.mem_setOf_eq] + constructor <;> rintro ⟨s, hs, rest⟩ <;> exact ⟨s, by omega, rest⟩ + rw [h_set_eq] + exact h_AM + 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₂ ∧ @@ -578,7 +650,7 @@ lemma prob_sum_ge_sqrt_log {σ2 : ℝ≥0} open Real omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -lemma todo {σ2 : ℝ≥0} {c : ℝ} +lemma streamMeasure_sampleMean_add_sqrt_le {σ2 : ℝ≥0} {c : ℝ} (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (hσ2 : σ2 ≠ 0) (hc : 0 ≤ c) (a : α) (n k : ℕ) (hk : k ≠ 0) : streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(2 * c * σ2 * log (n + 1) / k) ≤ (ν a)[id]} ≤ @@ -604,7 +676,7 @@ lemma todo {σ2 : ℝ≥0} {c : ℝ} _ ≤ 1 / (n + 1) ^ c := prob_sum_le_sqrt_log hν hσ2 hc a k hk omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -lemma todo' {σ2 : ℝ≥0} {c : ℝ} +lemma streamMeasure_le_sampleMean_sub_sqrt {σ2 : ℝ≥0} {c : ℝ} (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (hσ2 : σ2 ≠ 0) (hc : 0 ≤ c) (a : α) (n k : ℕ) (hk : k ≠ 0) : streamMeasure ν diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 8db83a83..d08b49f2 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -5,12 +5,16 @@ Authors: Rémy Degenne, Paulo Rauber -/ import LeanBandits.ForMathlib.SubGaussian import LeanBandits.BanditAlgorithms.Uniform +import LeanBandits.BanditAlgorithms.UCB import LeanBandits.SequentialLearning.BayesStationaryEnv +import LeanBandits.SequentialLearning.HistoryDensity /-! # The Thompson Sampling Algorithm -/ open MeasureTheory ProbabilityTheory Finset Learning +open scoped ENNReal + namespace Bandits namespace TS @@ -56,14 +60,1083 @@ variable (hK : 0 < K) variable {Ω : Type*} [MeasurableSpace Ω] variable (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → E × ℝ) variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] -variable (P : Measure Ω) [IsFiniteMeasure P] +variable (P : Measure Ω) [IsProbabilityMeasure P] + +noncomputable +def ucbIndex (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → E × ℝ) (δ : ℝ) + (a : Fin K) (t : ℕ) (ω : Ω) : ℝ := + max 0 (min 1 + (empMean A (IsBayesAlgEnvSeq.reward R') a t ω + + √(2 * Real.log (1 / δ) / (max 1 (pullCount A a t ω) : ℝ)))) + +omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in +lemma ucbIndex_nonneg (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : 0 ≤ ucbIndex A R' δ a t ω := + le_max_left 0 _ + +omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in +lemma ucbIndex_le_one (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : ucbIndex A R' δ a t ω ≤ 1 := + max_le (by norm_num) (min_le_left 1 _) + +omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in +lemma ucbIndex_mem_Icc (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : + ucbIndex A R' δ a t ω ∈ Set.Icc 0 1 := + ⟨ucbIndex_nonneg A R' δ a t ω, ucbIndex_le_one A R' δ a t ω⟩ + +omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in +lemma norm_ucbIndex_le_one (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : + ‖ucbIndex A R' δ a t ω‖ ≤ 1 := by + rw [Real.norm_eq_abs, abs_of_nonneg (ucbIndex_nonneg A R' δ a t ω)] + exact ucbIndex_le_one A R' δ a t ω + +omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in +lemma abs_sub_le_one_of_mem_Icc {x y : ℝ} (hx : x ∈ Set.Icc 0 1) (hy : y ∈ Set.Icc 0 1) : + |x - y| ≤ 1 := by + rw [abs_le]; constructor <;> linarith [hx.1, hx.2, hy.1, hy.2] + +lemma sum_sqrt_le {ι : Type*} (s : Finset ι) (c : ι → ℝ) (hc : ∀ i, 0 ≤ c i) : + ∑ i ∈ s, √(c i) ≤ √(#s * ∑ i ∈ s, c i) := by + have h := Real.sum_sqrt_mul_sqrt_le s hc (fun _ => zero_le_one) + simp only [Real.sqrt_one, mul_one, sum_const, nsmul_eq_mul] at h + calc ∑ i ∈ s, √(c i) ≤ √(∑ i ∈ s, c i) * √↑(#s) := h + _ = _ := by rw [← Real.sqrt_mul (Finset.sum_nonneg (fun i _ => hc i)), mul_comm] + +omit [StandardBorelSpace E] [Nonempty E] in +lemma sum_inv_sqrt_max_one_le (N : ℕ) : + ∑ j ∈ range N, (1 / √(↑(max 1 j) : ℝ)) ≤ 2 * √↑N := by + suffices h : ∀ M : ℕ, 0 < M → + ∑ j ∈ range M, (1 / √(↑(max 1 j) : ℝ)) + 1 / √↑M ≤ 2 * √↑M by + cases N with + | zero => simp + | succ n => + have := h (n + 1) (Nat.succ_pos n) + linarith [div_nonneg zero_le_one (Real.sqrt_nonneg (↑(n + 1) : ℝ))] + intro M hM + induction M with + | zero => omega + | succ n ih => + rw [sum_range_succ] + by_cases hn : n = 0 + · subst hn; simp; norm_num + · have hn_pos : 0 < n := Nat.pos_of_ne_zero hn + have hmax : (↑(max 1 n) : ℝ) = ↑n := by + simp [Nat.max_eq_right (by omega : 1 ≤ n)] + rw [hmax] + have h_ih := ih hn_pos + suffices h_key : 1 / √(↑(n + 1) : ℝ) ≤ 2 * (√↑(n + 1) - √↑n) by linarith + have hns : (0 : ℝ) < ↑(n + 1) := by positivity + have hnn : (0 : ℝ) ≤ ↑n := by positivity + set a := √(↑(n + 1) : ℝ) + set b := √(↑n : ℝ) + have hsn : a * a = ↑(n + 1) := Real.mul_self_sqrt (le_of_lt hns) + have hs : b * b = ↑n := Real.mul_self_sqrt hnn + have hab : 2 * (a * b) ≤ ↑(n + 1) + ↑n := by + nlinarith [mul_self_nonneg (a - b)] + rw [div_le_iff₀ (by positivity : 0 < a)] + have h_expand : 2 * (a - b) * a = 2 * (a * a) - 2 * (a * b) := by ring + rw [h_expand, hsn] + have : (↑(n + 1) : ℝ) = ↑n + 1 := by push_cast; ring + linarith + +@[fun_prop] +lemma measurable_ucbIndex [Nonempty (Fin K)] + (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) + (δ : ℝ) (a : Fin K) (t : ℕ) : + Measurable (ucbIndex A R' δ a t) := by + unfold ucbIndex + apply Measurable.max measurable_const + apply Measurable.min measurable_const + apply Measurable.add + · exact measurable_empMean (fun n ↦ h.measurable_A n) + (fun n ↦ h.measurable_reward n) a t + · have hpc : Measurable (fun ω ↦ (pullCount A a t ω : ℝ)) := + measurable_from_top.comp (measurable_pullCount (fun n ↦ h.measurable_A n) a t) + exact (measurable_const.div (measurable_const.max hpc)).sqrt + +omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in +lemma armMean_le_ucbIndex (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) + (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) + (hconc : + |empMean A (IsBayesAlgEnvSeq.reward R') a t ω - IsBayesAlgEnvSeq.armMean κ R' a ω| + < √(2 * Real.log (1 / δ) / (max 1 (pullCount A a t ω) : ℝ))) : + IsBayesAlgEnvSeq.armMean κ R' a ω ≤ ucbIndex A R' δ a t ω := by + unfold ucbIndex + have hmean := hm a (IsBayesAlgEnvSeq.env R' ω) + simp only [IsBayesAlgEnvSeq.armMean] at hmean hconc ⊢ + have habs := abs_sub_lt_iff.mp hconc + refine le_max_of_le_right (le_min hmean.2 ?_) + linarith [habs.2] + +omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in +lemma ucbIndex_sub_armMean_le (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) + (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) + (hconc : + |empMean A (IsBayesAlgEnvSeq.reward R') a t ω - IsBayesAlgEnvSeq.armMean κ R' a ω| + < √(2 * Real.log (1 / δ) / (max 1 (pullCount A a t ω) : ℝ))) : + ucbIndex A R' δ a t ω - IsBayesAlgEnvSeq.armMean κ R' a ω + ≤ 2 * √(2 * Real.log (1 / δ) / (max 1 (pullCount A a t ω) : ℝ)) := by + unfold ucbIndex + simp only [IsBayesAlgEnvSeq.armMean] at hconc ⊢ + set w := √(2 * Real.log (1 / δ) / max 1 ↑(pullCount A a t ω)) + set emp := empMean A (IsBayesAlgEnvSeq.reward R') a t ω + have habs := abs_sub_lt_iff.mp hconc + have hmean := hm a (IsBayesAlgEnvSeq.env R' ω) + have h1 : max 0 (min 1 (emp + w)) ≤ emp + w := + max_le_iff.mpr ⟨by linarith [hmean.1, habs.2], min_le_right _ _⟩ + linarith [habs.2] + +lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] + (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) (t : ℕ) : + condDistrib (A (t + 1)) (IsBayesAlgEnvSeq.hist A R' t) P + =ᵐ[P.map (IsBayesAlgEnvSeq.hist A R' t)] + condDistrib (IsBayesAlgEnvSeq.bestArm κ R') (IsBayesAlgEnvSeq.hist A R' t) P := + (h.hasCondDistrib_action' t).condDistrib_eq.trans + (posteriorBestArm_eq_uniform Q κ h hK t).symm + +omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in +lemma le_armMean_bestArm [Nonempty (Fin K)] (ω : Ω) (i : Fin K) : + IsBayesAlgEnvSeq.armMean κ R' i ω ≤ + IsBayesAlgEnvSeq.armMean κ R' (IsBayesAlgEnvSeq.bestArm κ R' ω) ω := by + have := isMaxOn_measurableArgmax (fun ω a ↦ IsBayesAlgEnvSeq.armMean κ R' a ω) ω i + simp only [IsBayesAlgEnvSeq.bestArm]; convert this + +omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in +lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) + (ω : Ω) : ⨆ i, IsBayesAlgEnvSeq.armMean κ R' i ω = + IsBayesAlgEnvSeq.armMean κ R' (IsBayesAlgEnvSeq.bestArm κ R' ω) ω := + le_antisymm (ciSup_le (le_armMean_bestArm R' κ ω)) + (le_ciSup (f := fun i ↦ IsBayesAlgEnvSeq.armMean κ R' i ω) + ⟨1, by rintro _ ⟨i, rfl⟩; exact (hm i _).2⟩ _) + +omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in +lemma gap_eq_armMean_sub [Nonempty (Fin K)] (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) + (s : ℕ) (ω : Ω) : gap (κ.comap (·, IsBayesAlgEnvSeq.env R' ω) (by fun_prop)) (A s ω) = + IsBayesAlgEnvSeq.armMean κ R' (IsBayesAlgEnvSeq.bestArm κ R' ω) ω - + IsBayesAlgEnvSeq.armMean κ R' (A s ω) ω := by + simp only [gap, Kernel.comap_apply] + exact congr_arg (· - _) (iSup_armMean_eq_bestArm R' κ hm ω) + +lemma bayesRegret_eq_sum_integral_gap [Nonempty (Fin K)] + {alg : Algorithm (Fin K) ℝ} + (h : IsBayesAlgEnvSeq Q κ A R' alg P) + {C : ℝ} (hm : ∀ a e, |(κ (a, e))[id]| ≤ C) (t : ℕ) : + IsBayesAlgEnvSeq.bayesRegret κ A R' P t = + ∑ s ∈ range t, P[fun ω ↦ gap (κ.comap (·, IsBayesAlgEnvSeq.env R' ω) (by fun_prop)) + (A s ω)] := by + simp only [IsBayesAlgEnvSeq.bayesRegret, IsBayesAlgEnvSeq.regret, regret_eq_sum_gap] + refine integral_finset_sum _ (fun s _ => ?_) + have hmeas : Measurable (fun ω ↦ gap (κ.comap (·, IsBayesAlgEnvSeq.env R' ω) (by fun_prop)) + (A s ω)) := + (Measurable.iSup h.measurable_armMean).sub + (stronglyMeasurable_id.integral_kernel.measurable.comp (h.measurable_action_env s)) + refine ⟨hmeas.aestronglyMeasurable, HasFiniteIntegral.of_bounded (C := 2 * C) + (Filter.Eventually.of_forall fun ω => ?_)⟩ + simp only [Real.norm_eq_abs, gap, Kernel.comap_apply] + set e := IsBayesAlgEnvSeq.env R' ω + have hbdd : BddAbove (Set.range fun i => (κ (i, e))[id]) := + ⟨C, by rintro _ ⟨i, rfl⟩; exact le_of_abs_le (hm i e)⟩ + rw [abs_of_nonneg (sub_nonneg.mpr (le_ciSup hbdd _))] + have h1 : (⨆ i, (κ (i, e))[id]) ≤ C := ciSup_le fun i => le_of_abs_le (hm i e) + have h2 : -C ≤ (κ (A s ω, e))[id] := neg_le_of_abs_le (hm (A s ω) e) + linarith + +omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsProbabilityMeasure Q] + [IsMarkovKernel κ] [IsProbabilityMeasure P] in +lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : + ∑ s ∈ range n, f (pullCount A (A s ω) s ω) = + ∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω), f j := by + induction n with + | zero => simp + | succ n ih => + rw [sum_range_succ, ih] + suffices ∑ a, ∑ j ∈ range (pullCount A a (n + 1) ω), f j = + (∑ a, ∑ j ∈ range (pullCount A a n ω), f j) + + f (pullCount A (A n ω) n ω) by linarith + have h_eq : ∀ a, ∑ j ∈ range (pullCount A a (n + 1) ω), f j = + ∑ j ∈ range (pullCount A a n ω), f j + + if A n ω = a then f (pullCount A a n ω) else 0 := by + intro a + rw [pullCount_add_one] + split_ifs with h + · rw [sum_range_succ] + · simp + simp_rw [h_eq, sum_add_distrib] + congr 1 + simp + +omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsProbabilityMeasure Q] + [IsMarkovKernel κ] [IsProbabilityMeasure P] in +lemma sum_ucbIndex_sub_armMean_le (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) + (δ : ℝ) (n : ℕ) (ω : Ω) + (hconc : ∀ s < n, ∀ a, + |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - IsBayesAlgEnvSeq.armMean κ R' a ω| + < √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ))) : + ∑ s ∈ range n, (ucbIndex A R' δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ R' (A s ω) ω) + ≤ 2 * √(8 * Real.log (1 / δ)) * √(↑K * ↑n) := by + have hterm : ∀ s ∈ range n, + ucbIndex A R' δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ R' (A s ω) ω + ≤ 2 * √(2 * Real.log (1 / δ) / (max 1 (pullCount A (A s ω) s ω) : ℝ)) := + fun s hs => ucbIndex_sub_armMean_le A R' κ hm δ (A s ω) s ω (hconc s (mem_range.mp hs) _) + calc ∑ s ∈ range n, + (ucbIndex A R' δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ R' (A s ω) ω) + ≤ ∑ s ∈ range n, + 2 * √(2 * Real.log (1 / δ) / (max 1 (pullCount A (A s ω) s ω) : ℝ)) := + sum_le_sum hterm + _ ≤ 2 * √(8 * Real.log (1 / δ)) * √(↑K * ↑n) := by + set c := Real.log (1 / δ) + by_cases hc : 0 ≤ 2 * c + · open Real in + calc ∑ s ∈ range n, 2 * √(2 * c / max 1 ↑(pullCount A (A s ω) s ω)) + = ∑ s ∈ range n, √(8 * c) * + (1 / √(↑(max 1 (pullCount A (A s ω) s ω)) : ℝ)) := + sum_congr rfl fun s _ => by + rw [show (8 : ℝ) * c = (2 : ℝ) ^ 2 * (2 * c) from by ring] + rw [sqrt_mul (by positivity : (0:ℝ) ≤ 2 ^ 2), + sqrt_sq (by norm_num : (0:ℝ) ≤ 2)] + rw [sqrt_div (by linarith : 0 ≤ 2 * c)]; push_cast; ring + _ = √(8 * c) * ∑ s ∈ range n, + (1 / √(↑(max 1 (pullCount A (A s ω) s ω)) : ℝ)) := by + rw [mul_sum] + _ = √(8 * c) * ∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω), + (1 / √(↑(max 1 j) : ℝ)) := by + congr 1; exact sum_comp_pullCount A (fun j => 1 / √(↑(max 1 j) : ℝ)) n ω + _ ≤ √(8 * c) * ∑ a : Fin K, (2 * √↑(pullCount A a n ω)) := by + gcongr with a; exact sum_inv_sqrt_max_one_le _ + _ = √(8 * c) * (2 * ∑ a : Fin K, √↑(pullCount A a n ω)) := by + simp only [mul_sum] + _ ≤ √(8 * c) * (2 * √(↑K * ↑n)) := by + gcongr + calc ∑ a : Fin K, √↑(pullCount A a n ω) + ≤ √(↑(Finset.univ.card) * ∑ a, ↑(pullCount A a n ω)) := + sum_sqrt_le Finset.univ _ fun a => by positivity + _ = √(↑K * ↑n) := by + congr 1; rw [Finset.card_fin]; congr 1 + have h := sum_pullCount (A := A) (t := n) (ω := ω) + exact_mod_cast h + _ = 2 * √(8 * c) * √(↑K * ↑n) := by ring + · have h0 : ∀ s ∈ range n, + 2 * √(2 * c / max 1 ↑(pullCount A (A s ω) s ω)) = 0 := + fun s _ => by + open Real in + have : 2 * c / max 1 ↑(pullCount A (A s ω) s ω) ≤ 0 := + div_nonpos_of_nonpos_of_nonneg (by linarith) (by positivity) + simp [sqrt_eq_zero'.mpr this] + rw [sum_congr rfl h0]; simp only [sum_const_zero]; positivity + +lemma streamMeasure_concentration_le_delta {α : Type*} [MeasurableSpace α] + {ν : Kernel α ℝ} [IsMarkovKernel ν] + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + (a : α) (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : + streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(2 * Real.log (1 / δ) / k) ≤ (ν a)[id]} ≤ + ENNReal.ofReal δ := by + have hlog : 0 < Real.log (1 / δ) := + Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) + have hk_pos : (0 : ℝ) < k := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hk) + calc + streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(2 * Real.log (1 / δ) / k) ≤ (ν a)[id]} + _ = streamMeasure ν + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) / k ≤ -√(2 * Real.log (1 / δ) / k)} := by + congr with ω + field_simp + rw [Finset.sum_sub_distrib] + simp + grind + _ = streamMeasure ν + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ -√(2 * k * Real.log (1 / δ))} := by + congr with ω + field_simp + congr! 2 + rw [Real.sqrt_div (by positivity), ← mul_div_assoc, mul_comm, mul_div_assoc, Real.div_sqrt, + mul_assoc (k : ℝ), Real.sqrt_mul (x := (k : ℝ)) (by positivity), mul_comm] + _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * Real.log (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 _; exact (hν a).congr_identDistrib + ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) + _ = ENNReal.ofReal δ := by + rw [Real.sq_sqrt (by positivity)] + simp only [neg_div, Real.exp_neg, mul_one] + rw [mul_div_assoc, mul_div_cancel₀ _ (by positivity : (2 * k : ℝ) ≠ 0), + Real.exp_log (by positivity : (0 : ℝ) < 1 / δ), one_div, inv_inv] + +lemma streamMeasure_concentration_ge_delta {α : Type*} [MeasurableSpace α] + {ν : Kernel α ℝ} [IsMarkovKernel ν] + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + (a : α) (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : + streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - √(2 * Real.log (1 / δ) / k)} ≤ + ENNReal.ofReal δ := by + have hlog : 0 < Real.log (1 / δ) := + Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) + have hk_pos : (0 : ℝ) < k := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hk) + calc + streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - √(2 * Real.log (1 / δ) / k)} + _ = streamMeasure ν + {ω | √(2 * Real.log (1 / δ) / k) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id])) / k} := by + congr with ω + field_simp + rw [Finset.sum_sub_distrib] + simp + grind + _ = streamMeasure ν + {ω | √(2 * k * Real.log (1 / δ)) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} := by + congr with ω + field_simp + congr! 1 + rw [Real.sqrt_div (by positivity), ← mul_div_assoc, mul_comm, mul_div_assoc, Real.div_sqrt] + rw [← Real.sqrt_mul (by positivity : (0 : ℝ) ≤ 2 * Real.log (1 / δ)), mul_comm] + _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * Real.log (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 _; exact (hν a).congr_identDistrib + ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) + _ = ENNReal.ofReal δ := by + rw [Real.sq_sqrt (by positivity)] + simp only [neg_div, Real.exp_neg, mul_one] + rw [mul_div_assoc, mul_div_cancel₀ _ (by positivity : (2 * k : ℝ) ≠ 0)] + rw [Real.exp_log (by positivity), one_div, inv_inv] + +lemma prob_concentration_single_delta_cond [StandardBorelSpace Ω] [Nonempty Ω] [Nonempty (Fin K)] + (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) + (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) + (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) + (hδ_large : 1 < 2 * Real.log (1 / δ)) : + ∀ᵐ e ∂(P.map (IsBayesAlgEnvSeq.env R')), + (condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P e) + {ω | √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ + |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} ≤ + ENNReal.ofReal (2 * s * δ) := by + filter_upwards [IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h] with e h_isAlgEnvSeq + let ν := κ.comap (·, e) (by fun_prop) + have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) 1 (ν a') := fun a' ↦ by + simp only [ν, Kernel.comap_apply]; exact hs a' e + have h_mean : (ν a)[id] = (κ (a, e))[id] := by simp only [ν, Kernel.comap_apply] + rw [← h_mean] + let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P e + have h_law := h_isAlgEnvSeq.law_pullCount_sumRewards_unique' + (ArrayModel.isAlgEnvSeq_arrayMeasure (tsAlgorithm hK Q κ) ν) (n := s) + let B_low := fun m : ℕ ↦ {x : ℝ | x / m + √(2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} + let B_high := fun m : ℕ ↦ {x : ℝ | (ν a)[id] ≤ x / m - √(2 * Real.log (1 / δ) / m)} + have h_stream_bound : ∀ m : ℕ, m ≠ 0 → + streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ≤ + ENNReal.ofReal (2 * δ) := by + intro m hm0 + calc streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} + ≤ streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m} + + streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_high m} := by + have h_union : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ⊆ + {ω | ∑ i ∈ range m, ω i a ∈ B_low m} ∪ {ω | ∑ i ∈ range m, ω i a ∈ B_high m} := by + intro ω hω; simp only [Set.mem_setOf_eq, Set.mem_union] at hω ⊢; exact hω + exact (measure_mono h_union).trans (measure_union_le _ _) + _ ≤ ENNReal.ofReal δ + ENNReal.ofReal δ := by + gcongr + · have h_eq : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m} = + {ω | (∑ i ∈ range m, ω i a) / m + √(2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} := by + ext ω; simp only [Set.mem_setOf_eq, B_low] + rw [h_eq]; exact streamMeasure_concentration_le_delta h_subG a m hm0 δ hδ hδ1 + · have h_eq : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_high m} = + {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - √(2 * Real.log (1 / δ) / m)} := by + ext ω; simp only [Set.mem_setOf_eq, B_high] + rw [h_eq]; exact streamMeasure_concentration_ge_delta h_subG a m hm0 δ hδ hδ1 + _ = ENNReal.ofReal (2 * δ) := by + rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf + let badSet := {ω : ℕ → (Fin K) × ℝ | + √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ + |empMean IT.action IT.reward a s ω - (ν a)[id]|} + have h_bound_per_m : ∀ m : ℕ, m ≠ 0 → m ≤ s → + P' {ω | pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} ≤ + streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := by + intro m hm0 hms + have hB_meas : MeasurableSet (B_low m ∪ B_high m) := + MeasurableSet.union (measurableSet_le (by fun_prop) (by fun_prop)) + (measurableSet_le (by fun_prop) (by fun_prop)) + exact prob_pullCount_eq_and_sumRewards_mem_le h_isAlgEnvSeq hms hB_meas + have h_bad_subset : badSet ⊆ + ⋃ m ∈ (Finset.range (s + 1)).filter (· ≠ 0), + {ω | pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by + intro ω hω + simp only [Set.mem_setOf_eq, badSet] at hω + simp only [Set.mem_iUnion, Finset.mem_filter, Finset.mem_range, Set.mem_setOf_eq] + set m := pullCount IT.action a s ω with hm_def + have hms : m ≤ s := pullCount_le (A := IT.action) a s ω + by_cases hm0 : m = 0 + · have h_empMean_zero : empMean IT.action IT.reward a s ω = 0 := by + simp only [empMean, ← hm_def, hm0, Nat.cast_zero, div_zero] + simp only [hm0, Nat.cast_zero, h_empMean_zero] at hω + exfalso + have h_mu := hm a e + simp only [max_eq_left (zero_le_one' ℝ), div_one] at hω + rw [h_mean, zero_sub, abs_neg, abs_of_nonneg h_mu.1] at hω + have : 1 < √(2 * Real.log (1 / δ)) := by + rw [Real.lt_sqrt (by norm_num)]; simpa using hδ_large + linarith [h_mu.2] + · -- Case: m ≥ 1 + use m + refine ⟨⟨Nat.lt_succ_of_le hms, hm0⟩, rfl, ?_⟩ + simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] + have hm_pos : (0 : ℝ) < m := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hm0) + have hmax_eq : (max 1 (m : ℕ) : ℝ) = m := by + simp only [Nat.one_le_cast, Nat.one_le_iff_ne_zero.mpr hm0, max_eq_right] + rw [hmax_eq] at hω + have h_empMean : empMean IT.action IT.reward a s ω = + sumRewards IT.action IT.reward a s ω / m := by + simp only [empMean, hm_def] + rw [h_empMean] at hω + by_cases h_le : sumRewards IT.action IT.reward a s ω / m ≤ (ν a)[id] + · left + have habs : |sumRewards IT.action IT.reward a s ω / ↑m - (ν a)[id]| = + (ν a)[id] - sumRewards IT.action IT.reward a s ω / m := by + rw [abs_of_nonpos (sub_nonpos.mpr h_le), neg_sub] + rw [habs] at hω + linarith + · right + have h_gt : (ν a)[id] < sumRewards IT.action IT.reward a s ω / m := not_le.mp h_le + have habs : |sumRewards IT.action IT.reward a s ω / ↑m - (ν a)[id]| = + sumRewards IT.action IT.reward a s ω / m - (ν a)[id] := + abs_of_pos (sub_pos.mpr h_gt) + rw [habs] at hω + linarith + calc P' badSet + ≤ P' (⋃ m ∈ (Finset.range (s + 1)).filter (· ≠ 0), + {ω | pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m}) := + measure_mono h_bad_subset + _ ≤ ∑ m ∈ (Finset.range (s + 1)).filter (· ≠ 0), + P' {ω | pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := + measure_biUnion_finset_le _ _ + _ ≤ ∑ m ∈ (Finset.range (s + 1)).filter (· ≠ 0), + streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := by + apply Finset.sum_le_sum + intro m hm + have hm0 : m ≠ 0 := (Finset.mem_filter.mp hm).2 + have hms : m ≤ s := Nat.lt_succ_iff.mp (Finset.mem_range.mp (Finset.mem_filter.mp hm).1) + exact h_bound_per_m m hm0 hms + _ ≤ ∑ _m ∈ (Finset.range (s + 1)).filter (· ≠ 0), ENNReal.ofReal (2 * δ) := by + apply Finset.sum_le_sum + intro m hm + have hm0 : m ≠ 0 := (Finset.mem_filter.mp hm).2 + exact h_stream_bound m hm0 + _ = ((Finset.range (s + 1)).filter (· ≠ 0)).card • ENNReal.ofReal (2 * δ) := by + simp only [Finset.sum_const] + _ = s • ENNReal.ofReal (2 * δ) := by + congr 1 + have hS_eq : (Finset.range (s + 1)).filter (· ≠ 0) = Finset.Icc 1 s := by + ext m; simp only [Finset.mem_filter, Finset.mem_range, ne_eq, Finset.mem_Icc]; omega + rw [hS_eq, Nat.card_Icc, Nat.add_sub_cancel] + _ = ENNReal.ofReal (2 * s * δ) := by + rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast s, ← ENNReal.ofReal_mul (Nat.cast_nonneg s)] + congr 1; ring + +lemma prob_concentration_single_delta [StandardBorelSpace Ω] [Nonempty Ω] [Nonempty (Fin K)] + (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) + (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) + (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) + (hδ_large : 1 < 2 * Real.log (1 / δ)) : + P {ω | √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ + |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - IsBayesAlgEnvSeq.armMean κ R' a ω|} ≤ + ENNReal.ofReal (2 * s * δ) := by + let badSet : E → Set (ℕ → (Fin K) × ℝ) := fun e ↦ + {t | √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s t) : ℝ)) ≤ + |empMean IT.action IT.reward a s t - (κ (a, e))[id]|} + have h_set_eq : {ω | √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ + |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - IsBayesAlgEnvSeq.armMean κ R' a ω|} = + (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + {p | p.2 ∈ badSet p.1} := by + ext ω + simp only [Set.mem_setOf_eq, Set.mem_preimage, badSet, IsBayesAlgEnvSeq.armMean] + have h1 : pullCount A a s ω = pullCount IT.action a s (IsBayesAlgEnvSeq.traj A R' ω) := by + unfold pullCount IsBayesAlgEnvSeq.traj IT.action; rfl + have h2 : empMean A (IsBayesAlgEnvSeq.reward R') a s ω = + empMean IT.action IT.reward a s (IsBayesAlgEnvSeq.traj A R' ω) := by + unfold empMean IsBayesAlgEnvSeq.traj IsBayesAlgEnvSeq.reward IT.action IT.reward; rfl + rw [h1, h2] + have h_meas_pair : + Measurable (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) := + h.measurable_env.prodMk h.measurable_traj + have h_disint : P.map (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) = + P.map (IsBayesAlgEnvSeq.env R') ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P := + (compProd_map_condDistrib (h.measurable_traj.aemeasurable)).symm + have h_cond := prob_concentration_single_delta_cond hK A R' Q κ P h hs hm a s δ hδ hδ1 hδ_large + have h_kernel : Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (a, p.1))[id]) := + stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk measurable_fst) + have h_meas_set : MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | p.2 ∈ badSet p.1} := by + change MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | + √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s p.2) : ℝ)) ≤ + |empMean IT.action IT.reward a s p.2 - (κ (a, p.1))[id]|} + exact measurableSet_le (by fun_prop) + (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp measurable_snd).sub + h_kernel).abs + calc P _ = P ((fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + {p | p.2 ∈ badSet p.1}) := by rw [h_set_eq] + _ = (P.map (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω))) + {p | p.2 ∈ badSet p.1} := by + rw [Measure.map_apply h_meas_pair h_meas_set] + _ = (P.map (IsBayesAlgEnvSeq.env R') ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P) + {p | p.2 ∈ badSet p.1} := by rw [h_disint] + _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P e) + (badSet e) ∂(P.map (IsBayesAlgEnvSeq.env R')) := by + rw [Measure.compProd_apply h_meas_set]; rfl + _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * s * δ) ∂(P.map (IsBayesAlgEnvSeq.env R')) := by + apply lintegral_mono_ae + filter_upwards [h_cond] with e h_e; exact h_e + _ = ENNReal.ofReal (2 * s * δ) := by + rw [lintegral_const, Measure.map_apply h.measurable_env MeasurableSet.univ] + simp [measure_univ] + +lemma prob_concentration_fail_delta [StandardBorelSpace Ω] [Nonempty Ω] [Nonempty (Fin K)] + (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) + (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) + (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) + (hδ_large : 1 < 2 * Real.log (1 / δ)) : + P {ω | ∃ s < n, ∃ a, + √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ + |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - IsBayesAlgEnvSeq.armMean κ R' a ω|} + ≤ ENNReal.ofReal (2 * K * n * δ) := by + let badSet := fun (s : ℕ) (a : Fin K) ↦ {ω : Ω | + √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ + |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - IsBayesAlgEnvSeq.armMean κ R' a ω|} + have h_set_eq : {ω | ∃ s < n, ∃ a, √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ + |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - IsBayesAlgEnvSeq.armMean κ R' a ω|} = + ⋃ s ∈ Finset.range n, ⋃ a : Fin K, badSet s a := by + ext ω; simp only [Set.mem_setOf_eq, Finset.mem_range, Set.mem_iUnion, badSet, exists_prop] + rw [h_set_eq] + have h_reorg : ⋃ s ∈ Finset.range n, ⋃ a : Fin K, badSet s a = + ⋃ a : Fin K, ⋃ s ∈ Finset.range n, badSet s a := by + ext ω; simp only [Set.mem_iUnion, Finset.mem_range]; exact + ⟨fun ⟨s, hs, a, ha⟩ ↦ ⟨a, s, hs, ha⟩, fun ⟨a, s, hs, ha⟩ ↦ ⟨s, hs, a, ha⟩⟩ + rw [h_reorg] + have h_arm_bound : ∀ a : Fin K, P (⋃ s ∈ Finset.range n, badSet s a) ≤ + ENNReal.ofReal (2 * n * δ) := by + intro a + by_cases hn : n = 0 + · simp [hn] + have hn' : 0 < n := Nat.pos_of_ne_zero hn + let badSetIT := fun (s : ℕ) (e : E) ↦ {ω : ℕ → (Fin K) × ℝ | + √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ + |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} + have h_set_eq : ⋃ s ∈ Finset.range n, badSet s a = + (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by + ext ω + simp only [Set.mem_iUnion, Finset.mem_range, badSet, badSetIT, Set.mem_preimage, + Set.mem_setOf_eq, IsBayesAlgEnvSeq.armMean] + exact Iff.rfl + rw [h_set_eq] + have h_meas_pair : + Measurable (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) := + h.measurable_env.prodMk h.measurable_traj + have h_disint : P.map (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) = + P.map (IsBayesAlgEnvSeq.env R') ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P := + (compProd_map_condDistrib (h.measurable_traj.aemeasurable)).symm + have h_cond_bound : ∀ᵐ e ∂(P.map (IsBayesAlgEnvSeq.env R')), + (condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P e) + (⋃ s ∈ Finset.range n, badSetIT s e) ≤ ENNReal.ofReal (2 * n * δ) := by + filter_upwards [IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h] with e h_isAlgEnvSeq + let ν := κ.comap (·, e) (by fun_prop) + let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P e + have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) 1 (ν a') := fun a' ↦ by + simp only [ν, Kernel.comap_apply]; exact hs a' e + have h_mean' : ∀ a', (κ (a', e))[id] ∈ Set.Icc 0 1 := fun a' ↦ hm a' e + have h_mean : (ν a)[id] = (κ (a, e))[id] := by simp only [ν, Kernel.comap_apply] + let B_low := fun m : ℕ ↦ + {x : ℝ | x / m + √(2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} + let B_high := fun m : ℕ ↦ + {x : ℝ | (ν a)[id] ≤ x / m - √(2 * Real.log (1 / δ) / m)} + have h_stream_bound : ∀ m : ℕ, m ≠ 0 → + streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ≤ + ENNReal.ofReal (2 * δ) := by + intro m hm0 + calc streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} + ≤ streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m} + + streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_high m} := by + have h_union : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ⊆ + {ω | ∑ i ∈ range m, ω i a ∈ B_low m} ∪ + {ω | ∑ i ∈ range m, ω i a ∈ B_high m} := by + intro ω hω; simp only [Set.mem_setOf_eq, Set.mem_union] at hω ⊢; exact hω + exact (measure_mono h_union).trans (measure_union_le _ _) + _ ≤ ENNReal.ofReal δ + ENNReal.ofReal δ := by + gcongr + · have h_eq : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m} = + {ω | (∑ i ∈ range m, ω i a) / m + √(2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} := by + ext ω; simp only [Set.mem_setOf_eq, B_low] + rw [h_eq]; exact streamMeasure_concentration_le_delta h_subG a m hm0 δ hδ hδ1 + · have h_eq : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_high m} = + {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - √(2 * Real.log (1 / δ) / m)} := by + ext ω; simp only [Set.mem_setOf_eq, B_high] + rw [h_eq]; exact streamMeasure_concentration_ge_delta h_subG a m hm0 δ hδ hδ1 + _ = ENNReal.ofReal (2 * δ) := by + rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf + have hB_meas : ∀ m, MeasurableSet (B_low m ∪ B_high m) := fun m ↦ + MeasurableSet.union (measurableSet_le (by fun_prop) (by fun_prop)) + (measurableSet_le (by fun_prop) (by fun_prop)) + let S := Finset.Icc 1 (n - 1) + have hS_card : S.card = n - 1 := by simp only [Nat.card_Icc, S]; omega + have h_decomp : ⋃ s ∈ Finset.range n, badSetIT s e = + ⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by + ext ω + simp only [Set.mem_iUnion, Finset.mem_range, exists_prop, badSetIT, Set.mem_setOf_eq, + Finset.mem_Icc, S] + constructor + · rintro ⟨s, hs, hbad⟩ + let m := pullCount IT.action a s ω + have hm_pos : 0 < m := by + by_contra hm0; push_neg at hm0 + have hm0' : pullCount IT.action a s ω = 0 := by omega + simp only [hm0', Nat.cast_zero, max_eq_left (zero_le_one' ℝ), div_one, + empMean, div_zero] at hbad + rw [zero_sub, abs_neg, abs_of_nonneg (h_mean' a).1] at hbad + have : 1 < √(2 * Real.log (1 / δ)) := by + rw [Real.lt_sqrt (by norm_num)]; simpa using hδ_large + linarith [(h_mean' a).2] + have hm_le : m ≤ n - 1 := by + have h1 : m ≤ s := pullCount_le (A := IT.action) a s ω + omega + refine ⟨m, ⟨hm_pos, hm_le⟩, s, hs, rfl, ?_⟩ + simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] + have h_pc_pos : (0 : ℝ) < m := Nat.cast_pos.mpr hm_pos + simp only [empMean] at hbad + rw [show (max 1 (pullCount IT.action a s ω) : ℝ) = m by + simp only [m]; rw [max_eq_right]; exact Nat.one_le_cast.mpr hm_pos] at hbad + have h_abs := le_abs'.mp hbad + rcases h_abs with h_neg | h_pos + · left; linarith + · right; linarith + · rintro ⟨m, ⟨hm_pos, hm_le⟩, s, hs, hpc, hB⟩ + refine ⟨s, hs, ?_⟩ + simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] at hB + have h_pc_pos : (0 : ℝ) < m := Nat.cast_pos.mpr hm_pos + simp only [empMean, hpc] + rw [show (max 1 m : ℝ) = m by rw [max_eq_right]; exact Nat.one_le_cast.mpr hm_pos] + cases hB with + | inl h => + have h1 : sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id] ≤ + -√(2 * Real.log (1 / δ) / m) := by linarith + calc √(2 * Real.log (1 / δ) / m) + ≤ -(sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id]) := by linarith + _ ≤ |sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id]| := neg_le_abs _ + | inr h => + have h1 : √(2 * Real.log (1 / δ) / m) ≤ + sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id] := by linarith + exact h1.trans (le_abs_self _) + rw [h_decomp] + calc P' (⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m}) + ≤ ∑ m ∈ S, P' {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := + measure_biUnion_finset_le S _ + _ ≤ ∑ m ∈ S, streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := by + apply Finset.sum_le_sum + intro m hm + have hm_pos : m ≠ 0 := Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1 + have hm_le : m ≤ n - 1 := (Finset.mem_Icc.mp hm).2 + have h_contain : {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} ⊆ + {ω | pullCount IT.action a (n - 1) ω = m ∧ + sumRewards IT.action IT.reward a (n - 1) ω ∈ B_low m ∪ B_high m} ∪ + {ω | pullCount IT.action a (n - 1) ω > m ∧ + ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by + intro ω ⟨s, hs, hpc, hB⟩ + simp only [Set.mem_union, Set.mem_setOf_eq] + have hs' : s ≤ n - 1 := Nat.le_sub_one_of_lt hs + have h_pc_mono := pullCount_mono (A := IT.action) a hs' ω + by_cases h_eq : pullCount IT.action a (n - 1) ω = m + · left + refine ⟨h_eq, ?_⟩ + have h_pc_eq : pullCount IT.action a s ω = pullCount IT.action a (n - 1) ω := + hpc.symm ▸ h_eq.symm + rw [← sumRewards_eq_of_pullCount_eq hs' h_pc_eq] + exact hB + · right + exact ⟨by omega, s, hs, hpc, hB⟩ + calc P' {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} + ≤ P' {ω | ∃ s, s ≤ n - 1 ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by + apply measure_mono + intro ω ⟨s, hs, hpc, hB⟩ + exact ⟨s, Nat.le_sub_one_of_lt hs, hpc, hB⟩ + _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := + prob_exists_pullCount_eq_and_sumRewards_mem_le (n := n - 1) + h_isAlgEnvSeq (hB_meas m) + _ ≤ ∑ _m ∈ S, ENNReal.ofReal (2 * δ) := by + apply Finset.sum_le_sum + intro m hm + have hm_pos : m ≠ 0 := Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1 + exact h_stream_bound m hm_pos + _ = S.card • ENNReal.ofReal (2 * δ) := by simp only [Finset.sum_const] + _ = (n - 1) • ENNReal.ofReal (2 * δ) := by rw [hS_card] + _ ≤ ENNReal.ofReal (2 * n * δ) := by + rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast (n - 1), + ← ENNReal.ofReal_mul (Nat.cast_nonneg (n - 1))] + apply ENNReal.ofReal_le_ofReal + have h1 : (n - 1 : ℕ) ≤ n := Nat.sub_le n 1 + have h2 : (↑(n - 1) : ℝ) ≤ (↑n : ℝ) := Nat.cast_le.mpr h1 + nlinarith [h2, hδ.le] + have h_kernel : Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (a, p.1))[id]) := + stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk measurable_fst) + have h_meas_set : MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | + p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by + have h_eq : {p : E × (ℕ → (Fin K) × ℝ) | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} = + ⋃ s ∈ Finset.range n, {p | p.2 ∈ badSetIT s p.1} := by + ext p; simp only [Set.mem_setOf_eq, Set.mem_iUnion, Finset.mem_range] + rw [h_eq] + exact .biUnion (Finset.range n).countable_toSet fun s _ ↦ by + simp only [badSetIT] + change MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | + √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s p.2) : ℝ)) ≤ + |empMean IT.action IT.reward a s p.2 - (κ (a, p.1))[id]|} + exact measurableSet_le (by fun_prop) + (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp + measurable_snd).sub h_kernel).abs + calc P ((fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1}) + = (P.map (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω))) + {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by + rw [Measure.map_apply h_meas_pair h_meas_set] + _ = (P.map (IsBayesAlgEnvSeq.env R') ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P) + {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by + rw [h_disint] + _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P e) + (⋃ s ∈ Finset.range n, badSetIT s e) ∂(P.map (IsBayesAlgEnvSeq.env R')) := by + rw [Measure.compProd_apply h_meas_set]; rfl + _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * n * δ) ∂(P.map (IsBayesAlgEnvSeq.env R')) := by + apply lintegral_mono_ae h_cond_bound + _ = ENNReal.ofReal (2 * n * δ) := by + rw [lintegral_const, Measure.map_apply h.measurable_env MeasurableSet.univ] + simp [measure_univ] + calc P (⋃ a : Fin K, ⋃ s ∈ Finset.range n, badSet s a) + ≤ ∑ a : Fin K, P (⋃ s ∈ Finset.range n, badSet s a) := measure_iUnion_fintype_le _ _ + _ ≤ ∑ _a : Fin K, ENNReal.ofReal (2 * n * δ) := by + apply Finset.sum_le_sum; intro a _; exact h_arm_bound a + _ = K • ENNReal.ofReal (2 * n * δ) := by simp [Finset.sum_const] + _ = ENNReal.ofReal (2 * K * n * δ) := by + simp only [nsmul_eq_mul] + rw [← ENNReal.ofReal_natCast K, ← ENNReal.ofReal_mul (Nat.cast_nonneg K)] + congr 1; ring + +lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] + (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) + (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) + (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) + (hδ_large : 1 < 2 * Real.log (1 / δ)) : + IsBayesAlgEnvSeq.bayesRegret κ A R' P n + ≤ 4 * K * n ^ 2 * δ + 2 * √(8 * Real.log (1 / δ)) * √(↑K * ↑n) := by + let bestArm := IsBayesAlgEnvSeq.bestArm κ R' + let armMean := IsBayesAlgEnvSeq.armMean κ R' + let ucb := ucbIndex A R' δ + set Eδ : Set Ω := {ω | ∀ s < n, ∀ a, + |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - armMean a ω| + < √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ))} + have hm_ucb : ∀ a t, Measurable (ucbIndex A R' δ a t) := + fun a t ↦ measurable_ucbIndex hK A R' Q κ P h δ a t + have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.armMean κ R' a) := + fun a ↦ h.measurable_armMean a + have hm_best : Measurable (IsBayesAlgEnvSeq.bestArm κ R') := h.measurable_bestArm + have h_first_bound : ∀ ω, + |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| ≤ n := fun ω ↦ + calc |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| + ≤ ∑ s ∈ range n, |armMean (bestArm ω) ω - ucb (bestArm ω) s ω| := + Finset.abs_sum_le_sum_abs _ _ + _ ≤ ∑ s ∈ range n, (1 : ℝ) := Finset.sum_le_sum fun s _ ↦ + abs_sub_le_one_of_mem_Icc (hm _ _) (ucbIndex_mem_Icc A R' δ _ _ _) + _ = ↑n := by simp + have h_second_bound : ∀ ω, + |∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)| ≤ n := fun ω ↦ + calc |∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)| + ≤ ∑ s ∈ range n, |ucb (A s ω) s ω - armMean (A s ω) ω| := + Finset.abs_sum_le_sum_abs _ _ + _ ≤ ∑ s ∈ range n, (1 : ℝ) := Finset.sum_le_sum fun s _ ↦ + abs_sub_le_one_of_mem_Icc (ucbIndex_mem_Icc A R' δ _ _ _) (hm _ _) + _ = ↑n := by simp + have h_int_sum1 : Integrable (fun ω ↦ ∑ s ∈ range n, + (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)) P := by + apply Integrable.of_bound (C := (↑n)) + · exact (Finset.measurable_fun_sum _ fun s _ ↦ + (measurable_apply_fin hm_arm hm_best).sub + (measurable_apply_fin (fun a ↦ hm_ucb a s) hm_best)).aestronglyMeasurable + · filter_upwards with ω; rw [Real.norm_eq_abs]; exact h_first_bound ω + have h_int_sum2 : Integrable (fun ω ↦ ∑ s ∈ range n, + (ucb (A s ω) s ω - armMean (A s ω) ω)) P := by + apply Integrable.of_bound (C := (↑n)) + · exact (Finset.measurable_fun_sum _ fun s _ ↦ + (measurable_apply_fin (fun a ↦ hm_ucb a s) (h.measurable_A s)).sub + (measurable_apply_fin hm_arm (h.measurable_A s))).aestronglyMeasurable + · filter_upwards with ω; rw [Real.norm_eq_abs]; exact h_second_bound ω + have h_swap : + IsBayesAlgEnvSeq.bayesRegret κ A R' P n = + P[fun ω ↦ ∑ s ∈ range n, + (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)] + + P[fun ω ↦ ∑ s ∈ range n, + (ucb (A s ω) s ω - armMean (A s ω) ω)] := by + have hC : ∀ a e, |(κ (a, e))[id]| ≤ 1 := fun a e ↦ by + have := hm a e; rw [abs_le]; exact ⟨by linarith [this.1], this.2⟩ + have h_regret_gap := bayesRegret_eq_sum_integral_gap (h := h) (hm := hC) (t := n) + have h_regret_eq : IsBayesAlgEnvSeq.bayesRegret κ A R' P n = + ∑ s ∈ range n, P[fun ω ↦ armMean (bestArm ω) ω - armMean (A s ω) ω] := by + rw [h_regret_gap]; congr 1 with s + exact integral_congr_ae (ae_of_all _ fun ω ↦ gap_eq_armMean_sub A R' κ hm s ω) + have h_int_ucb_sub : ∀ s, Integrable (fun ω ↦ ucb (A s ω) s ω - ucb (bestArm ω) s ω) P := by + intro s + apply Integrable.sub + · exact ⟨(measurable_apply_fin (fun a ↦ hm_ucb a s) (h.measurable_A s)).aestronglyMeasurable, + HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ norm_ucbIndex_le_one A R' _ _ _ _)⟩ + · exact ⟨(measurable_apply_fin (fun a ↦ hm_ucb a s) hm_best).aestronglyMeasurable, + HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ norm_ucbIndex_le_one A R' _ _ _ _)⟩ + have h_ucb_zero : ∀ a (ω : Ω), ucbIndex A R' δ a 0 ω = + max 0 (min 1 (√(2 * Real.log (1 / δ)))) := by + intro a ω; unfold ucbIndex empMean sumRewards; simp [pullCount_zero] + have h_ucb_swap : ∀ s, ∫ ω, (ucb (A s ω) s ω - ucb (bestArm ω) s ω) ∂P = 0 := by + intro s + cases s with + | zero => + have : ∀ ω, ucb (A 0 ω) 0 ω - ucb (bestArm ω) 0 ω = 0 := fun ω ↦ by + change ucbIndex A R' δ _ 0 ω - ucbIndex A R' δ _ 0 ω = 0 + simp [h_ucb_zero] + exact (integral_congr_ae (ae_of_all _ this)).trans (integral_zero _ _) + | succ t => + have hts := ts_identity hK A R' Q κ P h t + have h_map_eq : P.map (fun ω ↦ (IsBayesAlgEnvSeq.hist A R' t ω, A (t + 1) ω)) = + P.map (fun ω ↦ (IsBayesAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestArm κ R' ω)) := by + rw [← compProd_map_condDistrib (hY := (h.measurable_A (t + 1)).aemeasurable), + ← compProd_map_condDistrib (hY := hm_best.aemeasurable)] + exact Measure.compProd_congr hts + have h_int_eq : ∀ (f : (Iic t → Fin K × ℝ) × Fin K → ℝ), Measurable f → + ∫ ω, f (IsBayesAlgEnvSeq.hist A R' t ω, A (t + 1) ω) ∂P = + ∫ ω, f (IsBayesAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestArm κ R' ω) ∂P := by + intro f hf + rw [← integral_map + ((h.measurable_hist t).prodMk (h.measurable_A (t + 1))).aemeasurable + hf.aestronglyMeasurable, + ← integral_map + ((h.measurable_hist t).prodMk hm_best).aemeasurable + hf.aestronglyMeasurable, + h_map_eq] + set g : (Iic t → Fin K × ℝ) × Fin K → ℝ := + fun p ↦ max 0 (min 1 (empMean' t p.1 p.2 + + √(2 * Real.log (1 / δ) / (max 1 (pullCount' t p.1 p.2) : ℝ)))) + have h_hist_eq : ∀ (ω : Ω), + (fun (i : Iic t) ↦ (A (↑i) ω, + IsBayesAlgEnvSeq.reward R' (↑i) ω)) = + IsBayesAlgEnvSeq.hist A R' t ω := by + intro ω; rfl + have hg_eq : ∀ a (ω : Ω), ucbIndex A R' δ a (t + 1) ω = + g (IsBayesAlgEnvSeq.hist A R' t ω, a) := by + intro a ω + simp only [g, ucbIndex] + rw [empMean_add_one_eq_empMean' (A := A) (R' := IsBayesAlgEnvSeq.reward R'), + pullCount_add_one_eq_pullCount' (A := A) (R' := IsBayesAlgEnvSeq.reward R'), + h_hist_eq] + have hg_meas : Measurable g := by + apply Measurable.max measurable_const + apply Measurable.min measurable_const + apply Measurable.add + · exact measurable_apply_fin (fun a ↦ (measurable_empMean' t a).comp measurable_fst) + measurable_snd + · apply Measurable.sqrt + apply Measurable.div measurable_const + apply Measurable.max measurable_const + exact measurable_apply_fin + (fun a ↦ measurable_from_top.comp ((measurable_pullCount' t a).comp measurable_fst)) + measurable_snd + have h_eq_g1 : (fun ω ↦ ucb (A (t + 1) ω) (t + 1) ω) = + fun ω ↦ g (IsBayesAlgEnvSeq.hist A R' t ω, A (t + 1) ω) := + funext fun ω ↦ hg_eq _ _ + have h_eq_g2 : (fun ω ↦ ucb (bestArm ω) (t + 1) ω) = + fun ω ↦ g (IsBayesAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestArm κ R' ω) := + funext fun ω ↦ hg_eq _ _ + have h_int_ucb : ∀ {f : Ω → Fin K}, Measurable f → + Integrable (fun ω ↦ ucb (f ω) (t + 1) ω) P := fun hf ↦ + ⟨(measurable_apply_fin (fun a ↦ hm_ucb a (t + 1)) hf).aestronglyMeasurable, + HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ norm_ucbIndex_le_one A R' _ _ _ _)⟩ + have h_int1 := h_int_ucb (h.measurable_A (t + 1)) + have h_int2 := h_int_ucb hm_best + rw [show (fun ω ↦ ucb (A (t + 1) ω) (t + 1) ω - + ucb (bestArm ω) (t + 1) ω) = + fun ω ↦ (fun ω ↦ ucb (A (t + 1) ω) (t + 1) ω) ω - + (fun ω ↦ ucb (bestArm ω) (t + 1) ω) ω from rfl, + integral_sub h_int1 h_int2, h_eq_g1, h_eq_g2, + h_int_eq g hg_meas, sub_self] + have h_ucb_sum_zero : ∫ ω, ∑ s ∈ range n, + (ucb (A s ω) s ω - ucb (bestArm ω) s ω) ∂P = 0 := by + rw [integral_finset_sum _ (fun s _ ↦ h_int_ucb_sub s)] + exact Finset.sum_eq_zero fun s _ ↦ h_ucb_swap s + rw [h_regret_eq] + have h_pw : ∀ ω, (∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)) + + (∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)) = + (∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω)) + + (∑ s ∈ range n, (ucb (A s ω) s ω - ucb (bestArm ω) s ω)) := by + intro ω + simp only [← Finset.sum_add_distrib] + apply Finset.sum_congr rfl; intros; ring + have h_int_ucb_swap : Integrable + (fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - ucb (bestArm ω) s ω)) P := + integrable_finset_sum _ fun s _ ↦ h_int_ucb_sub s + have h_int_gap_s : ∀ s ∈ range n, + Integrable (fun ω ↦ armMean (bestArm ω) ω - armMean (A s ω) ω) P := by + intro s _ + refine Integrable.of_bound + ((measurable_apply_fin hm_arm hm_best).sub + (measurable_apply_fin hm_arm (h.measurable_A s))).aestronglyMeasurable + 1 (ae_of_all _ fun ω ↦ ?_) + rw [Real.norm_eq_abs] + exact abs_sub_le_one_of_mem_Icc (hm _ _) (hm _ _) + have h_int_gap : Integrable + (fun ω ↦ ∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω)) P := + integrable_finset_sum _ h_int_gap_s + calc ∑ s ∈ range n, ∫ (x : Ω), (fun ω ↦ armMean (bestArm ω) ω - armMean (A s ω) ω) x ∂P + = ∫ ω, ∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω) ∂P := + (integral_finset_sum _ h_int_gap_s).symm + _ = ∫ ω, ((∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω)) + + (∑ s ∈ range n, (ucb (A s ω) s ω - ucb (bestArm ω) s ω))) ∂P := by + rw [integral_add h_int_gap h_int_ucb_swap, h_ucb_sum_zero, add_zero] + _ = ∫ ω, ((∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)) + + (∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω))) ∂P := by + congr 1; ext ω; linarith [h_pw ω] + _ = _ := integral_add h_int_sum1 h_int_sum2 + have h_first_Eδ : ∀ ω ∈ Eδ, + ∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω) + ≤ 0 := by + intro ω hω + apply Finset.sum_nonpos + intro s hs + linarith [armMean_le_ucbIndex A R' κ hm δ + (bestArm ω) s ω (hω s (mem_range.mp hs) _)] + have h_second_Eδ : ∀ ω ∈ Eδ, + ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω) + ≤ 2 * √(8 * Real.log (1 / δ)) * √(↑K * ↑n) := by + intro ω hω + exact sum_ucbIndex_sub_armMean_le A R' κ hm δ n ω hω + have h_prob : P Eδᶜ ≤ ENNReal.ofReal (2 * K * n * δ) := by + have : Eδᶜ = {ω | ∃ s < n, ∃ a, √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ + |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - armMean a ω|} := by + ext ω; simp only [Eδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push_neg; rfl + rw [this] + exact prob_concentration_fail_delta (hK := hK) (A := A) (R' := R') + (Q := Q) (κ := κ) (P := P) h hs hm n δ hδ hδ1 hδ_large + have hm_emp : ∀ a s, Measurable (fun ω ↦ empMean A (IsBayesAlgEnvSeq.reward R') a s ω) := + fun a s ↦ measurable_empMean (fun n ↦ h.measurable_A n) (fun n ↦ h.measurable_reward n) a s + have hm_pc : ∀ a s, Measurable (fun ω ↦ (pullCount A a s ω : ℝ)) := + fun a s ↦ measurable_from_top.comp (measurable_pullCount (fun n ↦ h.measurable_A n) a s) + have hEδ_meas : MeasurableSet Eδ := by + suffices ∀ s a, MeasurableSet {ω | + |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - armMean a ω| + < √(2 * Real.log (1 / δ) / max 1 ↑(pullCount A a s ω))} by + simp only [Eδ, Set.setOf_forall] + exact .iInter fun s ↦ .iInter fun _ ↦ .iInter fun a ↦ this s a + intro s a + exact measurableSet_lt + ((hm_emp a s).sub (h.measurable_armMean a)).abs + ((measurable_const.div (measurable_const.max (hm_pc a s))).sqrt) + rw [h_swap] + set f1 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, + (armMean (bestArm ω) ω - ucb (bestArm ω) s ω) + set f2 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, + (ucb (A s ω) s ω - armMean (A s ω) ω) + set B := 2 * √(8 * Real.log (1 / δ)) * √(↑K * ↑n) + have s1 := (integral_add_compl hEδ_meas h_int_sum1).symm + have s2 := (integral_add_compl hEδ_meas h_int_sum2).symm + have h1g : ∫ ω in Eδ, f1 ω ∂P ≤ 0 := + setIntegral_nonpos hEδ_meas fun ω hω ↦ h_first_Eδ ω hω + have h_compl_bound : ∀ {f}, Integrable f P → (∀ ω, f ω ≤ ↑n) → + ∫ ω in Eδᶜ, f ω ∂P ≤ ↑n * P.real Eδᶜ := fun hint hle ↦ by + have := setIntegral_mono_on (hf := hint.integrableOn) (hg := integrableOn_const) + hEδ_meas.compl fun ω _ ↦ hle ω + rwa [setIntegral_const, smul_eq_mul, mul_comm] at this + have h1b := h_compl_bound h_int_sum1 fun ω ↦ (abs_le.mp (h_first_bound ω)).2 + have h2g : ∫ ω in Eδ, f2 ω ∂P ≤ B := by + have hB : 0 ≤ B := by positivity + have := setIntegral_mono_on (hf := h_int_sum2.integrableOn) + (hg := integrableOn_const) hEδ_meas + fun ω hω ↦ h_second_Eδ ω hω + rw [setIntegral_const, smul_eq_mul, mul_comm] at this + exact le_trans this (mul_le_of_le_one_right hB measureReal_le_one) + have h2b := h_compl_bound h_int_sum2 fun ω ↦ (abs_le.mp (h_second_bound ω)).2 + have hP : P.real Eδᶜ ≤ 2 * ↑K * ↑n * δ := + ENNReal.toReal_le_of_le_ofReal (by positivity) h_prob + rw [s1, s2] + have hP0 : 0 ≤ P.real Eδᶜ := by positivity + nlinarith -lemma TS.bayesRegret_le [Nonempty (Fin K)] +lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) (t : ℕ) : - IsBayesAlgEnvSeq.bayesRegret κ A R' P t ≤ 4 * K + 8 * √(K * t * Real.log t) := - sorry + IsBayesAlgEnvSeq.bayesRegret κ A R' P t ≤ 4 * K + 8 * √(K * t * Real.log t) := by + by_cases ht : t = 0 + · simp [ht, IsBayesAlgEnvSeq.bayesRegret, IsBayesAlgEnvSeq.regret, regret] + by_cases ht1_eq : t = 1 + · subst ht1_eq + simp only [Nat.cast_one, Real.log_one, mul_zero, Real.sqrt_zero, mul_zero, add_zero] + calc IsBayesAlgEnvSeq.bayesRegret κ A R' P 1 + ≤ 1 := by + unfold IsBayesAlgEnvSeq.bayesRegret IsBayesAlgEnvSeq.regret Bandits.regret + simp only [Finset.range_one, Finset.sum_singleton, Nat.cast_one, one_mul, + Kernel.comap_apply] + refine (integral_mono_of_nonneg (ae_of_all _ fun ω ↦ sub_nonneg.mpr + (le_ciSup ⟨1, by rintro _ ⟨a, rfl⟩; exact (hm a _).2⟩ _)) + (integrable_const 1) (ae_of_all _ fun ω ↦ by + linarith [ciSup_le fun a ↦ (hm a (IsBayesAlgEnvSeq.env R' ω)).2, + (hm (A 0 ω) (IsBayesAlgEnvSeq.env R' ω)).1])).trans ?_ + simp + _ ≤ 4 * (K : ℝ) := by + nlinarith [show (1 : ℝ) ≤ K from Nat.one_le_cast.mpr (Nat.one_le_of_lt hK)] + -- For t ≥ 2, we have δ = 1/t² < 1 + · have ht2 : 2 ≤ t := by omega + have htpos : (0 : ℝ) < t := Nat.cast_pos.mpr (Nat.pos_of_ne_zero ht) + have _ht1 : (1 : ℝ) ≤ t := Nat.one_le_cast.mpr (Nat.pos_of_ne_zero ht) + have hδ : (0 : ℝ) < 1 / (t : ℝ) ^ 2 := div_pos one_pos (pow_pos htpos 2) + have hδ1 : 1 / (t : ℝ) ^ 2 < 1 := by + rw [div_lt_one (pow_pos htpos 2)] + have ht2_real : (2 : ℝ) ≤ t := Nat.ofNat_le_cast.mpr ht2 + calc (1 : ℝ) < 2 ^ 2 := by norm_num + _ ≤ (t : ℝ) ^ 2 := by gcongr + -- First term simplification: 4K · t² · (1/t²) = 4K + have h_first : 4 * (K : ℝ) * ↑t ^ 2 * (1 / (↑t) ^ 2) = 4 * ↑K := by + field_simp + -- Second term simplification: log(1/(1/t²)) = log(t²) = 2 log(t) + have h_log : Real.log (1 / (1 / (↑t : ℝ) ^ 2)) = 2 * Real.log ↑t := by + rw [one_div_one_div, Real.log_pow]; norm_cast + -- For t ≥ 2, we have 2 * log(t²) = 4 log(t) ≥ 4 log(2) > 1 (since log(2) > 0.69) + have hδ_large : 1 < 2 * Real.log (1 / (1 / (↑t : ℝ) ^ 2)) := by + rw [h_log] + have h_log2 : (1 : ℝ) / 2 < Real.log 2 := by + rw [Real.lt_log_iff_exp_lt (by norm_num : (0 : ℝ) < 2)] + calc Real.exp (1 / 2) = Real.sqrt (Real.exp 1) := by rw [← Real.exp_half] + _ < Real.sqrt 4 := by gcongr; linarith [Real.exp_one_lt_d9] + _ = 2 := by rw [show (4 : ℝ) = 2 ^ 2 by norm_num, Real.sqrt_sq (by norm_num)] + have ht2_real : (2 : ℝ) ≤ t := Nat.ofNat_le_cast.mpr ht2 + linarith [Real.log_le_log (by norm_num : (0 : ℝ) < 2) ht2_real] + have hK_real_pos : (0 : ℝ) < K := Nat.cast_pos.mpr hK + have hKt_nonneg : (0 : ℝ) ≤ ↑K * ↑t := by positivity + calc IsBayesAlgEnvSeq.bayesRegret κ A R' P t + ≤ 4 * ↑K * ↑t ^ 2 * (1 / (↑t) ^ 2) + + 2 * √(8 * Real.log (1 / (1 / (↑t) ^ 2))) * √(↑K * ↑t) := + bayesRegret_le_of_delta (hK := hK) (A := A) (R' := R') (Q := Q) + (κ := κ) (P := P) h hs hm t (1 / (↑t) ^ 2) hδ hδ1 hδ_large + _ = 4 * ↑K + 2 * √(16 * Real.log ↑t) * √(↑K * ↑t) := by rw [h_first, h_log]; ring_nf + _ = 4 * ↑K + 8 * (√(Real.log ↑t) * √(↑K * ↑t)) := by + rw [show (16 : ℝ) = 4 ^ 2 by norm_num, Real.sqrt_mul (by norm_num : (0 : ℝ) ≤ 4 ^ 2), + Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 4)]; ring + _ = 4 * ↑K + 8 * √(↑K * ↑t * Real.log ↑t) := by + rw [← Real.sqrt_mul (Real.log_nonneg (Nat.one_le_cast.mpr (Nat.pos_of_ne_zero ht)))] + ring_nf end Regret diff --git a/LeanBandits/BanditAlgorithms/UCB.lean b/LeanBandits/BanditAlgorithms/UCB.lean index 806a2a12..e8491763 100644 --- a/LeanBandits/BanditAlgorithms/UCB.lean +++ b/LeanBandits/BanditAlgorithms/UCB.lean @@ -241,7 +241,7 @@ lemma prob_ucbIndex_le [Nonempty (Fin K)] {alg : Algorithm (Fin K) ℝ} grind _ ≤ ∑ k ∈ Icc 1 n, (1 : ℝ≥0∞) / (n + 1) ^ c := by gcongr with k hk - exact todo hν hσ2 hc a n k (by grind) + exact streamMeasure_sampleMean_add_sqrt_le hν hσ2 hc a n k (by grind) _ ≤ (n + 1) * (1 : ℝ≥0∞) / (n + 1) ^ c := 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)] @@ -284,7 +284,7 @@ lemma prob_ucbIndex_ge [Nonempty (Fin K)] {alg : Algorithm (Fin K) ℝ} grind _ ≤ ∑ k ∈ Icc 1 n, (1 : ℝ≥0∞) / (n + 1) ^ c := by gcongr with k hk - exact todo' hν hσ2 hc a n k (by grind) + exact streamMeasure_le_sampleMean_sub_sqrt hν hσ2 hc a n k (by grind) _ ≤ (n + 1) * (1 : ℝ≥0∞) / (n + 1) ^ c := 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)] diff --git a/LeanBandits/ForMathlib/Measurable.lean b/LeanBandits/ForMathlib/Measurable.lean index 5a36cdb4..5b6b03f7 100644 --- a/LeanBandits/ForMathlib/Measurable.lean +++ b/LeanBandits/ForMathlib/Measurable.lean @@ -87,4 +87,17 @@ lemma measurable_sum_Icc_of_le {f : ℕ → α → ℝ} {g : α → ℕ} {n : refine Measurable.ite ?_ (by fun_prop) (by fun_prop) exact (measurableSet_singleton _).preimage (by fun_prop) +lemma measurable_apply_fin {α' : Type*} [MeasurableSpace α'] [Fintype α'] + [MeasurableSingletonClass α'] + {f : α' → α → ℝ} {g : α → α'} + (hf : ∀ a, Measurable (f a)) (hg : Measurable g) : + Measurable (fun ω ↦ f (g ω) ω) := by + classical + have : (fun ω ↦ f (g ω) ω) = fun ω ↦ ∑ a : α', if g ω = a then f a ω else 0 := by + ext ω; simp [Finset.sum_ite_eq] + rw [this] + apply Finset.measurable_fun_sum + intro a _ + exact Measurable.ite (hg (measurableSet_singleton a)) (hf a) measurable_const + end MeasureTheory diff --git a/LeanBandits/ForMathlib/MeasurableArgMax.lean b/LeanBandits/ForMathlib/MeasurableArgMax.lean index 0769a25d..02fe3cc2 100644 --- a/LeanBandits/ForMathlib/MeasurableArgMax.lean +++ b/LeanBandits/ForMathlib/MeasurableArgMax.lean @@ -18,9 +18,7 @@ lemma measurable_encode {α : Type*} {_ : MeasurableSpace α} [Encodable α] [MeasurableSingletonClass α] : Measurable (Encodable.encode (α := α)) := by refine measurable_to_nat fun a ↦ ?_ - have : Encodable.encode ⁻¹' {Encodable.encode a} = {a} := by ext; simp - rw [this] - exact measurableSet_singleton _ + rw [show Encodable.encode ⁻¹' {Encodable.encode a} = {a} from by ext; simp]; measurability lemma measurableEmbedding_encode (α : Type*) {_ : MeasurableSpace α} [Encodable α] [MeasurableSingletonClass α] : @@ -39,13 +37,12 @@ lemma measurableSet_isMax [Countable 𝓨] {f : 𝓧 → 𝓨 → α} (hf : ∀ y, Measurable (fun x ↦ f x y)) (y : 𝓨) : MeasurableSet {x | ∀ z, f x z ≤ f x y} := by rw [show {x | ∀ y', f x y' ≤ f x y} = ⋂ y', {x | f x y' ≤ f x y} by ext; simp] - exact MeasurableSet.iInter fun z ↦ measurableSet_le (by fun_prop) (by fun_prop) + exact .iInter fun z ↦ measurableSet_le (hf z) (hf y) lemma exists_isMaxOn' {α : Type*} [LinearOrder α] [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] (f : 𝓧 → 𝓨 → α) (x : 𝓧) : - ∃ n : ℕ, ∃ y, n = Encodable.encode y ∧ ∀ z, f x z ≤ f x y := by - obtain ⟨y, h⟩ := Finite.exists_max (f x) - exact ⟨Encodable.encode y, y, rfl, h⟩ + ∃ n : ℕ, ∃ y, n = Encodable.encode y ∧ ∀ z, f x z ≤ f x y := + let ⟨y, h⟩ := Finite.exists_max (f x); ⟨Encodable.encode y, y, rfl, h⟩ /-- A measurable argmax function. -/ noncomputable @@ -61,13 +58,11 @@ lemma measurable_measurableArgmax [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y] (hf : ∀ y, Measurable (fun x ↦ f x y)) : Measurable (measurableArgmax f) := by - refine (MeasurableEmbedding.measurable_invFun (measurableEmbedding_encode 𝓨)).comp ?_ - refine measurable_find _ fun n ↦ ?_ - have : {x | ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y} - = ⋃ y, ({x | n = Encodable.encode y} ∩ {x | ∀ z, f x z ≤ f x y}) := by ext; simp - rw [this] - refine MeasurableSet.iUnion fun y ↦ (MeasurableSet.inter (by simp) ?_) - exact measurableSet_isMax (by fun_prop) y + refine (MeasurableEmbedding.measurable_invFun (measurableEmbedding_encode 𝓨)).comp + (measurable_find _ fun n ↦ ?_) + rw [show {x | ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y} + = ⋃ y, ({x | n = Encodable.encode y} ∩ {x | ∀ z, f x z ≤ f x y}) from by ext; simp] + exact .iUnion fun y ↦ .inter (by simp) (measurableSet_isMax hf y) lemma isMaxOn_measurableArgmax {α : Type*} [LinearOrder α] [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨] @@ -76,9 +71,32 @@ lemma isMaxOn_measurableArgmax {α : Type*} [LinearOrder α] (x : 𝓧) (z : 𝓨) : f x z ≤ f x (measurableArgmax f x) := by obtain ⟨y, h_eq, h_le⟩ := Nat.find_spec (exists_isMaxOn' f x) - refine le_trans (h_le z) (le_of_eq ?_) - rw [measurableArgmax, h_eq, + exact (h_le z).trans_eq <| by rw [measurableArgmax, h_eq, MeasurableEmbedding.leftInverse_invFun (measurableEmbedding_encode 𝓨) y] +/-- Congruence lemma: measurableArgmax only depends on the function values at the point. -/ +lemma measurableArgmax_congr {𝓧₁ 𝓧₂ : Type*} {α : Type*} [LinearOrder α] + [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨] + (f₁ : 𝓧₁ → 𝓨 → α) (f₂ : 𝓧₂ → 𝓨 → α) + [∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f₁ x z ≤ f₁ x y] + [∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f₂ x z ≤ f₂ x y] + (x₁ : 𝓧₁) (x₂ : 𝓧₂) (h : f₁ x₁ = f₂ x₂) : + measurableArgmax f₁ x₁ = measurableArgmax f₂ x₂ := by + simp only [measurableArgmax]; congr 1 + exact Nat.find_congr' fun {_} => + ⟨fun ⟨y, hn, hy⟩ => ⟨y, hn, h ▸ hy⟩, fun ⟨y, hn, hy⟩ => ⟨y, hn, h.symm ▸ hy⟩⟩ + +/-- measurableArgmax is independent of the DecidablePred instance used. + This follows from Nat.find_congr' which handles different decidability instances. -/ +lemma measurableArgmax_eq_of_eq {α : Type*} [LinearOrder α] + [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨] + (f : 𝓧 → 𝓨 → α) + (d1 : ∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y) + (d2 : ∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y) + (x : 𝓧) : + @measurableArgmax 𝓧 𝓨 α _ _ _ _ _ _ f d1 x = @measurableArgmax 𝓧 𝓨 α _ _ _ _ _ _ f d2 x := by + simp only [measurableArgmax]; congr 1 + exact @Nat.find_congr' _ _ (d1 x) (d2 x) _ _ (fun {_} ↦ Iff.rfl) + end Finite end MeasurableArgmax diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index bd4436a4..b6e3f402 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -6,6 +6,8 @@ Authors: Rémy Degenne, Paulo Rauber import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.Bandit.Regret import LeanBandits.SequentialLearning.IonescuTulceaSpace +import LeanBandits.SequentialLearning.StationaryEnv +import Mathlib.Probability.Kernel.Posterior /-! # Bayesian stationary environments -/ @@ -195,6 +197,210 @@ lemma condIndepFun_reward_hist_action_env [StandardBorelSpace Ω] end Independence +section Posterior + +/-- The posterior on the environment given history equals Mathlib's `posterior` applied to the +likelihood kernel and prior. This is the measure-theoretic formulation of Bayes' rule. -/ +lemma condDistrib_env_hist_eq_posterior [StandardBorelSpace Ω] + (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : + condDistrib (env R') (hist A R' n) P + =ᵐ[P.map (hist A R' n)] posterior (condDistrib (hist A R' n) (env R') P) Q := by + -- The key is to show P.map (env, hist) = Q ⊗ₘ condDistrib hist env P + -- Then use compProd_posterior_eq_map_swap and uniqueness of conditional kernels + have h_env_meas : Measurable (env R') := h.measurable_env + have h_hist_meas : Measurable (hist A R' n) := h.measurable_hist n + set κ' := condDistrib (hist A R' n) (env R') P with hκ' + have h_disint : P.map (fun ω => (env R' ω, hist A R' n ω)) = Q ⊗ₘ κ' := by + rw [← h.hasLaw_env.map_eq, compProd_map_condDistrib (h_hist_meas.aemeasurable)] + have h_marg : P.map (hist A R' n) = κ' ∘ₘ Q := by + have : P.map (hist A R' n) = (P.map (fun ω => (env R' ω, hist A R' n ω))).snd := by + rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)]; rfl + rw [this, h_disint, Measure.snd_compProd] + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] + rw [show P.map (fun ω => (hist A R' n ω, env R' ω)) = (Q ⊗ₘ κ').map Prod.swap from by + rw [← h_disint, Measure.map_map (by fun_prop) (by fun_prop)]; rfl] + rw [← compProd_posterior_eq_map_swap (κ := κ') (μ := Q), h_marg] + +end Posterior + +section StationaryEnvConnection + +def traj (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (ω : Ω) : ℕ → α × R := + fun n => (A n ω, (R' n ω).2) + +@[fun_prop] +lemma measurable_traj (h : IsBayesAlgEnvSeq Q κ A R' alg P) : Measurable (traj A R') := + measurable_pi_lambda _ fun n => (h.measurable_A n).prodMk (h.measurable_R n).snd + +omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] + [StandardBorelSpace R] [Nonempty R] in +/-- The traj function commutes with IT projections: IT.action n ∘ traj = A n -/ +lemma IT_action_comp_traj (n : ℕ) : IT.action n ∘ traj A R' = A n := by + ext ω; simp [IT.action, traj] + +omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] + [StandardBorelSpace R] [Nonempty R] in +/-- The traj function commutes with IT projections: IT.reward n ∘ traj = reward n -/ +lemma IT_reward_comp_traj (n : ℕ) : IT.reward n ∘ traj A R' = reward R' n := by + ext ω; simp [IT.reward, traj, reward] + +omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] + [StandardBorelSpace R] [Nonempty R] in +/-- The traj function commutes with IT projections: IT.hist n ∘ traj = hist n -/ +lemma IT_hist_comp_traj (n : ℕ) : IT.hist n ∘ traj A R' = hist A R' n := by + ext ω i : 2 + simp only [Function.comp_apply, IT.hist, traj, hist] + +omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] + [StandardBorelSpace R] [Nonempty R] in +/-- The pair (IT.hist n, IT.action (n+1)) commutes with traj. -/ +lemma IT_hist_action_comp_traj (n : ℕ) : + (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) ∘ traj A R' = + fun ω ↦ (hist A R' n ω, A (n + 1) ω) := by + ext ω : 1 + simp only [Function.comp_apply, IT.action, traj, Prod.mk.injEq] + exact ⟨funext fun _ => rfl, trivial⟩ + +lemma condDistrib_traj_action_zero (h : IsBayesAlgEnvSeq Q κ A R' alg P) : + ∀ᵐ e ∂(P.map (env R')), + (condDistrib (traj A R') (env R') P e).map (IT.action 0) = alg.p0 := by + have h_comp : condDistrib (IT.action 0 ∘ traj A R') (env R') P + =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map (IT.action 0) := + condDistrib_comp (env R') (h.measurable_traj.aemeasurable) (IT.measurable_action 0) + rw [IT_action_comp_traj] at h_comp + filter_upwards [h_comp, condDistrib_of_indepFun h.indepFun_action_zero_env.symm + h.measurable_env.aemeasurable (h.measurable_A 0).aemeasurable] with e h_comp_e h_indep_e + rw [← Kernel.map_apply _ (IT.measurable_action 0), ← h_comp_e, h_indep_e, Kernel.const_apply] + simp only [h.hasLaw_action_zero.map_eq, Algorithm.prod_left] + +lemma hasCondDistrib_reward_zero_condDistrib (h : IsBayesAlgEnvSeq Q κ A R' alg P) : + ∀ᵐ e ∂(P.map (env R')), + HasCondDistrib (IT.reward 0) (IT.action 0) (κ.comap (·, e) (by fun_prop)) + (condDistrib (traj A R') (env R') P e) := by + have h_swap : HasCondDistrib (reward R' 0) (fun ω ↦ (env R' ω, A 0 ω)) + (κ.comap Prod.swap (by fun_prop)) P := by + convert h.hasCondDistrib_reward_zero'.comp_right + (MeasurableEquiv.prodComm : α × E ≃ᵐ E × α) using 2 + have h_prod := condDistrib_prod_left (h.measurable_A 0).aemeasurable + (h.measurable_reward 0).aemeasurable h.measurable_env.aemeasurable (μ := P) + have h_comp_pair : condDistrib ((fun ω ↦ (IT.action 0 ω, IT.reward 0 ω)) ∘ traj A R') (env R') P + =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map + (fun ω ↦ (IT.action 0 ω, IT.reward 0 ω)) := + condDistrib_comp (env R') h.measurable_traj.aemeasurable (by fun_prop) + have h_comp_action : condDistrib (IT.action 0 ∘ traj A R') (env R') P + =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map (IT.action 0) := + condDistrib_comp (env R') h.measurable_traj.aemeasurable (IT.measurable_action 0) + rw [show (fun ω ↦ (IT.action 0 ω, IT.reward 0 ω)) ∘ traj A R' = + fun ω ↦ (A 0 ω, reward R' 0 ω) from by + ext ω : 1; simp only [Function.comp_apply, IT.action, IT.reward, traj, reward]] at h_comp_pair + rw [IT_action_comp_traj] at h_comp_action + have h_swap_eq := h_swap.condDistrib_eq + rw [(compProd_map_condDistrib (h.measurable_A 0).aemeasurable).symm] at h_swap_eq + filter_upwards [h_prod, h_comp_pair, h_comp_action, + (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_swap_eq] + with e h_prod_e h_pair_e h_act_e h_nested_e + refine ⟨by fun_prop, by fun_prop, ?_⟩ + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] + rw [← Kernel.map_apply _ (by fun_prop), ← h_pair_e] + conv_rhs => rw [← Kernel.map_apply _ (IT.measurable_action 0), ← h_act_e] + rw [h_prod_e, Kernel.compProd_apply_eq_compProd_sectR] + refine Measure.compProd_congr ?_ + filter_upwards [h_nested_e] with a ha + ext s _ + rw [Kernel.sectR_apply, Kernel.comap_apply, ha, Kernel.comap_apply]; rfl + +lemma hasCondDistrib_action_condDistrib (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : + ∀ᵐ e ∂(P.map (env R')), + HasCondDistrib (IT.action (n + 1)) (IT.hist n) (alg.policy n) + (condDistrib (traj A R') (env R') P e) := by + have h_prod := condDistrib_prod_left (h.measurable_hist n).aemeasurable + (h.measurable_A (n + 1)).aemeasurable h.measurable_env.aemeasurable (μ := P) + have h_action_env := (hasCondDistrib_action_env_hist h n).condDistrib_eq + have h_comp_pair : condDistrib ((fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) ∘ traj A R') + (env R') P =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map + (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) := + condDistrib_comp (env R') h.measurable_traj.aemeasurable (by fun_prop) + have h_comp_hist : condDistrib (IT.hist n ∘ traj A R') (env R') P + =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map (IT.hist n) := + condDistrib_comp (env R') h.measurable_traj.aemeasurable (IT.measurable_hist n) + rw [IT_hist_action_comp_traj] at h_comp_pair + rw [IT_hist_comp_traj] at h_comp_hist + rw [(compProd_map_condDistrib (h.measurable_hist n).aemeasurable).symm] at h_action_env + filter_upwards [h_prod, h_comp_pair, h_comp_hist, + (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_action_env] + with e h_prod_e h_pair_e h_hist_e h_nested_e + refine ⟨by fun_prop, by fun_prop, ?_⟩ + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] + rw [← Kernel.map_apply _ (by fun_prop), ← h_pair_e] + conv_rhs => rw [← Kernel.map_apply _ (IT.measurable_hist n), ← h_hist_e] + rw [h_prod_e, Kernel.compProd_apply_eq_compProd_sectR] + refine Measure.compProd_congr ?_ + filter_upwards [h_nested_e] with _ ha + ext s _ + rw [Kernel.sectR_apply, ha, Kernel.prodMkLeft_apply] + +lemma hasCondDistrib_reward_condDistrib (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : + ∀ᵐ e ∂(P.map (env R')), + HasCondDistrib (IT.reward (n + 1)) (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) + ((κ.comap (·, e) (by fun_prop)).prodMkLeft _) + (condDistrib (traj A R') (env R') P e) := by + have h_prod := condDistrib_prod_left + (Measurable.prodMk (h.measurable_hist n) (h.measurable_A (n + 1))).aemeasurable + (h.measurable_reward (n + 1)).aemeasurable h.measurable_env.aemeasurable (μ := P) + have h_swap : HasCondDistrib (reward R' (n + 1)) (fun ω ↦ (env R' ω, hist A R' n ω, A (n + 1) ω)) + (κ.comap (fun p ↦ (p.2.2, p.1)) (by fun_prop)) P := + (hasCondDistrib_reward_hist_action_env h n).comp_right + (MeasurableEquiv.prodAssoc.symm.trans MeasurableEquiv.prodComm) + have h_swap_eq := h_swap.condDistrib_eq + have h_comp_triple : condDistrib + ((fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), IT.reward (n + 1) ω)) ∘ traj A R') (env R') P + =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map + (fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), IT.reward (n + 1) ω)) := + condDistrib_comp (env R') h.measurable_traj.aemeasurable (by fun_prop) + have h_comp_pair : condDistrib ((fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) ∘ traj A R') + (env R') P =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map + (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) := + condDistrib_comp (env R') h.measurable_traj.aemeasurable (by fun_prop) + rw [show (fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), IT.reward (n + 1) ω)) ∘ + traj A R' = fun ω ↦ ((hist A R' n ω, A (n + 1) ω), reward R' (n + 1) ω) from by + ext ω : 1 + simp only [Function.comp_apply, IT.action, IT.reward, traj, reward, Prod.mk.injEq] + exact ⟨⟨funext fun i => rfl, trivial⟩, trivial⟩] at h_comp_triple + rw [IT_hist_action_comp_traj] at h_comp_pair + rw [(compProd_map_condDistrib (Measurable.prodMk (h.measurable_hist n) + (h.measurable_A (n + 1))).aemeasurable).symm] at h_swap_eq + filter_upwards [h_prod, h_comp_triple, h_comp_pair, + (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_swap_eq] + with e h_prod_e h_triple_e h_pair_e h_nested_e + refine ⟨by fun_prop, by fun_prop, ?_⟩ + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] + rw [← Kernel.map_apply _ (by fun_prop), ← h_triple_e] + conv_rhs => rw [← Kernel.map_apply _ (by fun_prop), ← h_pair_e] + rw [h_prod_e, Kernel.compProd_apply_eq_compProd_sectR] + refine Measure.compProd_congr ?_ + filter_upwards [h_nested_e] with _ ha + ext s _ + rw [Kernel.sectR_apply, ha, Kernel.comap_apply, Kernel.prodMkLeft_apply, Kernel.comap_apply] + +lemma condDistrib_traj_isAlgEnvSeq [StandardBorelSpace Ω] [Nonempty Ω] + (h : IsBayesAlgEnvSeq Q κ A R' alg P) : + ∀ᵐ e ∂(P.map (env R')), + IsAlgEnvSeq IT.action IT.reward alg + (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (condDistrib (traj A R') (env R') P e) := by + filter_upwards [condDistrib_traj_action_zero h, + hasCondDistrib_reward_zero_condDistrib h, + ae_all_iff.2 (hasCondDistrib_action_condDistrib h), + ae_all_iff.2 (hasCondDistrib_reward_condDistrib h)] with e h_law h_r0 h_a h_r + exact { + hasLaw_action_zero := ⟨by fun_prop, h_law⟩ + hasCondDistrib_reward_zero := h_r0 + hasCondDistrib_action := h_a + hasCondDistrib_reward := h_r + } + +end StationaryEnvConnection + end IsBayesAlgEnvSeq namespace IT diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index 527db1f6..9b1f53ab 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -95,10 +95,7 @@ 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 + | succ n => simp [pullCount_add_one_eq_pullCount' (R' := R')] 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 @@ -274,11 +271,8 @@ lemma stepsUntil_zero_of_ne (hka : A 0 ω ≠ a) : stepsUntil A a 0 ω = 0 := by 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 + intro n; exact (this.trans_le (monotone_pullCount _ _ (by omega))).ne' + 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)] : @@ -366,11 +360,7 @@ lemma stepsUntil_eq_zero_iff : 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] + by_cases hka : A 0 ω = a <;> simp_all · cases h' with | inl h => rw [h.1, stepsUntil_zero_of_ne h.2] @@ -644,11 +634,6 @@ theorem isStoppingTime_stepsUntil_filtrationAction [MeasurableSingletonClass α] · 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 : ℕ) : --- MeasurableSpace (ℕ → α × R) := --- (isStoppingTime_stepsUntil_filtrationAction a m (mR := mR)).measurableSpace - end Measurability end StepsUntil @@ -767,6 +752,22 @@ noncomputable def empMean' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := (sumRewards' n h a) / (pullCount' n h a) +lemma sumRewards_eq_of_pullCount_eq {R' : ℕ → Ω → ℝ} {s t : ℕ} (hst : s ≤ t) + (h_eq : pullCount A a s ω = pullCount A a t ω) : + sumRewards A R' a s ω = sumRewards A R' a t ω := by + induction t, hst using Nat.le_induction with + | base => rfl + | succ t hst' ih => + have h_mono : pullCount A a s ω ≤ pullCount A a t ω := pullCount_mono a hst' ω + have h_mono' : pullCount A a t ω ≤ pullCount A a (t + 1) ω := pullCount_mono a (Nat.le_succ t) ω + have h_eq_t : pullCount A a s ω = pullCount A a t ω := le_antisymm h_mono (h_eq ▸ h_mono') + have hne : A t ω ≠ a := by + intro ha + have h1 := ha ▸ pullCount_action_eq_pullCount_add_one (A := A) t ω + omega + simp only [sumRewards, sum_range_succ, if_neg hne, add_zero] + exact ih h_eq_t + 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 @@ -796,10 +797,7 @@ lemma sumRewards_eq_sumRewards' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (h 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 + | succ n => simp [sumRewards_add_one_eq_sumRewards'] lemma empMean_add_one_eq_empMean' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} : empMean A R' a (n + 1) ω = empMean' n (fun i ↦ (A i ω, R' i ω)) a := by diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean new file mode 100644 index 00000000..5c3c4df5 --- /dev/null +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -0,0 +1,568 @@ +/- +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.SequentialLearning.StationaryEnv +import LeanBandits.SequentialLearning.BayesStationaryEnv +import LeanBandits.BanditAlgorithms.Uniform +import Mathlib.Probability.Kernel.RadonNikodym + +/-! +# Algorithm-Independence of Bayesian Posteriors + +The key result: the posterior distribution on the environment (and therefore on the best arm) +given the observed history is independent of the algorithm used to generate the data. + +The proof routes through a uniform algorithm as reference measure. The history distribution +under any algorithm is absolutely continuous w.r.t. the uniform algorithm's, with a density +that depends only on action probabilities (not the environment). This density factorization +implies the posteriors agree. +-/ + +open MeasureTheory ProbabilityTheory Finset Preorder + +open scoped ENNReal NNReal + +namespace Learning + +variable {K : ℕ} + +section UniformFullSupport + +variable {K : ℕ} (hK : 0 < K) + +/-- The uniform algorithm gives positive probability to every action. -/ +lemma uniformAlgorithm_p0_pos (a : Fin K) : + (Bandits.uniformAlgorithm hK).p0 {a} > 0 := by + simp only [Bandits.uniformAlgorithm, uniformOn] + refine cond_pos_of_inter_ne_zero MeasurableSet.univ ?_ + simp only [Set.univ_inter, Measure.count_singleton, ne_eq, one_ne_zero, not_false_eq_true] + +/-- The uniform algorithm's policy gives positive probability to every action. -/ +lemma uniformAlgorithm_policy_pos (n : ℕ) (h : Iic n → Fin K × ℝ) (a : Fin K) : + (Bandits.uniformAlgorithm hK).policy n h {a} > 0 := by + simp only [Bandits.uniformAlgorithm, Kernel.const_apply, uniformOn] + refine cond_pos_of_inter_ne_zero MeasurableSet.univ ?_ + simp only [Set.univ_inter, Measure.count_singleton, ne_eq, one_ne_zero, not_false_eq_true] + +/-- Any measure on a finite type is absolutely continuous wrt any measure giving positive mass + to all singletons. -/ +lemma absolutelyContinuous_of_forall_singleton_pos {α : Type*} [MeasurableSpace α] + [MeasurableSingletonClass α] [Fintype α] + {μ ν : Measure α} [IsFiniteMeasure μ] + (hν : ∀ a : α, ν {a} > 0) : μ ≪ ν := by + intro s hs + have h_empty : s = ∅ := by + by_contra h + obtain ⟨a, ha⟩ := Set.nonempty_iff_ne_empty.mpr h + have h1 : ν {a} ≤ ν s := measure_mono (Set.singleton_subset_iff.mpr ha) + exact absurd (le_antisymm (hs ▸ h1) (zero_le _)) (ne_of_gt (hν a)) + rw [h_empty, measure_empty] + +/-- `rnDeriv` is pointwise finite when the reference measure has full support on singletons. -/ +lemma rnDeriv_ne_top_of_forall_singleton_pos {α : Type*} [MeasurableSpace α] + [MeasurableSingletonClass α] [Fintype α] + {μ ν : Measure α} [IsFiniteMeasure μ] [IsFiniteMeasure ν] + (hν : ∀ a, ν {a} > 0) (a : α) : μ.rnDeriv ν a ≠ ⊤ := by + intro h_eq + have h_mem : a ∈ {x | ¬ (μ.rnDeriv ν x < ⊤)} := by simp [h_eq] + have h_null : ν {x | ¬ (μ.rnDeriv ν x < ⊤)} = 0 := + ae_iff.mp (Measure.rnDeriv_lt_top μ ν) + exact absurd (le_antisymm ((measure_mono (Set.singleton_subset_iff.mpr h_mem)).trans + (le_of_eq h_null)) (zero_le _)) (ne_of_gt (hν a)) + +/-- Kernel `rnDeriv` is pointwise finite when the reference kernel has full support + on singletons. -/ +lemma kernel_rnDeriv_ne_top_of_forall_singleton_pos + {α' β' : Type*} [MeasurableSpace α'] [MeasurableSpace β'] + [MeasurableSingletonClass β'] [Fintype β'] + [MeasurableSpace.CountableOrCountablyGenerated α' β'] + {κ η : Kernel α' β'} [IsFiniteKernel κ] [IsFiniteKernel η] + (hη : ∀ a b, η a {b} > 0) (a : α') (b : β') : + Kernel.rnDeriv κ η a b ≠ ⊤ := by + intro h_eq + have h_mem : b ∈ {x | ¬ (Kernel.rnDeriv κ η a x < ⊤)} := by simp [h_eq] + have h_null : η a {x | ¬ (Kernel.rnDeriv κ η a x < ⊤)} = 0 := + ae_iff.mp (Kernel.rnDeriv_lt_top κ η) + exact absurd (le_antisymm ((measure_mono (Set.singleton_subset_iff.mpr h_mem)).trans + (le_of_eq h_null)) (zero_le _)) (ne_of_gt (hη a b)) + +end UniformFullSupport + +section WithDensityHelpers + +variable {α' β' : Type*} {mα' : MeasurableSpace α'} {mβ' : MeasurableSpace β'} + +/-- Composing `withDensity` on the measure side of a `compProd`: +`(μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst)`. -/ +private lemma withDensity_compProd_left + {μ : Measure α'} [SFinite μ] {κ : Kernel α' β'} [IsSFiniteKernel κ] + {f : α' → ENNReal} (hf : Measurable f) [SFinite (μ.withDensity f)] : + (μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst) := by + ext s hs + rw [Measure.compProd_apply hs, withDensity_apply _ hs, + lintegral_withDensity_eq_lintegral_mul₀ hf.aemeasurable + (Kernel.measurable_kernel_prodMk_left hs).aemeasurable, + ← lintegral_indicator hs, + Measure.lintegral_compProd ((hf.comp measurable_fst).indicator hs)] + congr 1; ext a; simp_rw [Pi.mul_apply] + have : (fun b => s.indicator (f ∘ Prod.fst) (a, b)) = + fun b => (Prod.mk a ⁻¹' s).indicator (fun _ => f a) b := by + ext b; simp only [Function.comp, Set.indicator, Set.mem_preimage]; rfl + rw [this, lintegral_indicator_const (hs.preimage (by fun_prop))] + +/-- Mapping a `withDensity` through `MeasurableEquiv.symm`: +`(μ.withDensity f).map e.symm = (μ.map e.symm).withDensity (f ∘ e)`. -/ +private lemma withDensity_map_equiv_symm + {μ : Measure β'} {e : α' ≃ᵐ β'} {f : β' → ENNReal} (hf : Measurable f) : + (μ.withDensity f).map e.symm = (μ.map e.symm).withDensity (f ∘ e) := by + ext s hs + rw [Measure.map_apply e.symm.measurable hs, + withDensity_apply _ (e.symm.measurable hs), + withDensity_apply _ hs, Measure.restrict_map e.symm.measurable hs, + lintegral_map (hf.comp e.measurable) e.symm.measurable] + simp_rw [Function.comp_apply, e.apply_symm_apply] + +/-- Mapping a `withDensity` through a `MeasurableEquiv` from the snd component. -/ +private lemma map_swap_withDensity_fst + {μ : Measure (α' × β')} [SFinite μ] + {f : β' → ENNReal} (hf : Measurable f) : + (μ.withDensity (f ∘ Prod.snd)).map Prod.swap + = (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := by + ext s hs + rw [Measure.map_apply measurable_swap hs, withDensity_apply _ (measurable_swap hs), + withDensity_apply _ hs, Measure.restrict_map measurable_swap hs] + exact (lintegral_map (hf.comp measurable_fst) measurable_swap).symm + +/-- `(μ.withDensity (h ∘ g)).map g = (μ.map g).withDensity h`. -/ +private lemma withDensity_map_eq' + {γ' : Type*} {mγ' : MeasurableSpace γ'} + {μ : Measure α'} {g : α' → γ'} {h : γ' → ENNReal} + (hg : Measurable g) (hh : Measurable h) : + (μ.withDensity (h ∘ g)).map g = (μ.map g).withDensity h := by + ext s hs + rw [Measure.map_apply hg hs, withDensity_apply _ (hg hs), withDensity_apply _ hs] + conv_rhs => rw [Measure.restrict_map hg hs] + rw [lintegral_map hh hg]; rfl + +/-- `(κ.withDensity (fun _ => ρ)) ∘ₘ Q = (κ ∘ₘ Q).withDensity ρ`. -/ +private lemma comp_withDensity_const + {γ' : Type*} {mγ' : MeasurableSpace γ'} + {Q : Measure α'} [SFinite Q] + {κ : Kernel α' γ'} [IsSFiniteKernel κ] + {ρ : γ' → ENNReal} (hρ : Measurable ρ) + [IsSFiniteKernel (κ.withDensity (fun _ => ρ))] : + (κ.withDensity (fun _ => ρ)) ∘ₘ Q = (κ ∘ₘ Q).withDensity ρ := by + rw [← Measure.snd_compProd Q (κ.withDensity (fun _ => ρ)), + Measure.compProd_withDensity (show Measurable (Function.uncurry (fun (_ : α') => ρ)) from + hρ.comp measurable_snd), + ← Measure.snd_compProd Q κ, Measure.snd, Measure.snd] + exact withDensity_map_eq' measurable_snd hρ + +/-- `(μ.withDensity f) ⊗ₘ (κ.withDensity g) = (μ ⊗ₘ κ).withDensity (f ∘ fst * uncurry g)`. -/ +private lemma withDensity_compProd_withDensity + {γ' : Type*} {mγ' : MeasurableSpace γ'} + {μ : Measure α'} [SFinite μ] + {κ : Kernel α' γ'} [IsSFiniteKernel κ] + {f : α' → ENNReal} {g : α' → γ' → ENNReal} + (hf : Measurable f) (hg : Measurable (Function.uncurry g)) + [SFinite (μ.withDensity f)] [IsSFiniteKernel (κ.withDensity g)] : + (μ.withDensity f) ⊗ₘ (κ.withDensity g) + = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst * Function.uncurry g) := by + rw [Measure.compProd_withDensity hg, withDensity_compProd_left hf] + exact (withDensity_mul _ (hf.comp measurable_fst) hg).symm + +end WithDensityHelpers + +section AbsolutelyContinuousHist + +variable {K : ℕ} [Nonempty (Fin K)] + +omit [Nonempty (Fin K)] in +/-- The step kernel for a stationary environment decomposes as a product of the policy + measure and the reward kernel. -/ +private lemma absolutelyContinuous_stepKernel_stationary + (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) (ν : Kernel (Fin K) ℝ) + [IsMarkovKernel ν] (n : ℕ) (h : Iic n → Fin K × ℝ) : + stepKernel alg (stationaryEnv ν) n h ≪ + stepKernel (Bandits.uniformAlgorithm hK) (stationaryEnv ν) n h := by + have h1 : stepKernel alg (stationaryEnv ν) n h = (alg.policy n h) ⊗ₘ ν := by + simp only [stepKernel, stationaryEnv]; ext s hs + simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] + have h2 : stepKernel (Bandits.uniformAlgorithm hK) (stationaryEnv ν) n h = + ((Bandits.uniformAlgorithm hK).policy n h) ⊗ₘ ν := by + simp only [stepKernel, stationaryEnv]; ext s hs + simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] + rw [h1, h2] + exact Measure.AbsolutelyContinuous.compProd_left + (absolutelyContinuous_of_forall_singleton_pos (uniformAlgorithm_policy_pos hK n h)) _ + +-- `compProd` unfolding requires extra heartbeats +/-- The history distribution at time `n + 1` decomposes as a compProd of the history at time `n` + and the step kernel, composed with `IicSuccProd.symm`. -/ +private lemma map_hist_succ_eq_compProd_map + (alg : Algorithm (Fin K) ℝ) (env : Environment (Fin K) ℝ) (n : ℕ) : + (trajMeasure alg env).map (IT.hist (n + 1)) = + ((trajMeasure alg env).map (IT.hist n) ⊗ₘ stepKernel alg env n).map + (MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n).symm := by + set P := trajMeasure alg env + set e := MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n + have h_func : IT.hist (α := Fin K) (R := ℝ) (n + 1) = e.symm ∘ + (fun (ω : ℕ → Fin K × ℝ) => (IT.hist n ω, IT.step (n + 1) ω)) := by + funext ω + change frestrictLe (n + 1) ω = e.symm (frestrictLe n ω, IT.step (n + 1) ω) + rw [show IT.step (n + 1) ω = ω (n + 1) from Prod.mk.eta] + change frestrictLe (n + 1) ω = e.symm (e (frestrictLe (n + 1) ω)) + rw [e.symm_apply_apply] + rw [h_func, (Measure.map_map e.symm.measurable (by fun_prop : + Measurable (fun (ω : ℕ → Fin K × ℝ) => + (IT.hist n ω, IT.step (n + 1) ω)))).symm] + congr 1 + have h_cd := (IT.isAlgEnvSeq_trajMeasure alg env).hasCondDistrib_step n + exact ((condDistrib_ae_eq_iff_measure_eq_compProd _ + (by fun_prop : AEMeasurable (IsAlgEnvSeq.step IT.action IT.reward (n + 1)) P) + (stepKernel alg env n)).mp h_cd.condDistrib_eq) + +/-- The history distribution under any algorithm is absolutely continuous w.r.t. the + history distribution under the uniform algorithm, for a stationary environment. -/ +private lemma absolutelyContinuous_map_hist_stationary + (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) + (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] (t : ℕ) : + (trajMeasure alg (stationaryEnv ν)).map (IT.hist t) ≪ + (trajMeasure (Bandits.uniformAlgorithm hK) (stationaryEnv ν)).map + (IT.hist t) := by + induction t with + | zero => + simp only [IT.hist_eq_frestrictLe, trajMeasure, + Kernel.trajMeasure_map_frestrictLe, Kernel.partialTraj_self, + Measure.id_comp, stationaryEnv_ν0] + exact (Measure.AbsolutelyContinuous.compProd_left + (absolutelyContinuous_of_forall_singleton_pos + (uniformAlgorithm_p0_pos hK)) _).map + (MeasurableEquiv.piUnique _).symm.measurable + | succ n ih => + rw [map_hist_succ_eq_compProd_map, map_hist_succ_eq_compProd_map] + exact (Measure.AbsolutelyContinuous.compProd ih + (Filter.Eventually.of_forall fun h => + absolutelyContinuous_stepKernel_stationary hK alg ν n h)).map + (MeasurableEquiv.IicSuccProd _ n).symm.measurable + +-- `condDistrib_comp` + `eq_trajMeasure` chain requires extra heartbeats +/-- The conditional distribution of the history given the environment equals the trajectory + measure's history marginal (for Bayesian stationary environments). -/ +private lemma condDistrib_hist_env_eq_traj + {E : Type*} [MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] + (Q : Measure E) [IsProbabilityMeasure Q] + (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] + {Ω : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] + {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → E × ℝ} + {alg : Algorithm (Fin K) ℝ} {P : Measure Ω} [IsProbabilityMeasure P] + (h : IsBayesAlgEnvSeq Q κ A R' alg P) (t : ℕ) : + ∀ᵐ e ∂Q, + condDistrib (IsBayesAlgEnvSeq.hist A R' t) + (IsBayesAlgEnvSeq.env R') P e = + (trajMeasure alg + (stationaryEnv (κ.comap (·, e) (by fun_prop)))).map (IT.hist t) := by + rw [← h.hasLaw_env.map_eq] + have h_comp : condDistrib (IT.hist t ∘ IsBayesAlgEnvSeq.traj A R') + (IsBayesAlgEnvSeq.env R') P + =ᵐ[P.map (IsBayesAlgEnvSeq.env R')] + (condDistrib (IsBayesAlgEnvSeq.traj A R') + (IsBayesAlgEnvSeq.env R') P).map (IT.hist t) := + condDistrib_comp (IsBayesAlgEnvSeq.env R') + h.measurable_traj.aemeasurable (IT.measurable_hist t) + rw [IsBayesAlgEnvSeq.IT_hist_comp_traj (E := E) t] at h_comp + filter_upwards [h_comp, h.condDistrib_traj_isAlgEnvSeq] with e hc he + rw [hc, Kernel.map_apply _ (IT.measurable_hist t)] + congr 1 + have h' := eq_trajMeasure_of_isAlgEnvSeq he + have hid : + (fun (ω : ℕ → Fin K × ℝ) n => (IT.action n ω, IT.reward n ω)) = id := by + funext ω n; exact Prod.mk.eta + rw [hid, Measure.map_id] at h'; exact h' + +end AbsolutelyContinuousHist + +section PosteriorEquality + +variable {E' X' : Type*} {mE' : MeasurableSpace E'} {mX' : MeasurableSpace X'} + +-- `compProd` unfolding requires extra heartbeats +/-- If `κ₁ =ᵐ[Q] κ₂.withDensity (fun _ => ρ)`, then the posteriors agree: +`κ₂†Q =ᵐ[κ₁ ∘ₘ Q] κ₁†Q`. -/ +private theorem posterior_eq_of_withDensity_ae_eq + [StandardBorelSpace E'] [Nonempty E'] + {Q : Measure E'} [IsFiniteMeasure Q] + {κ₁ κ₂ : Kernel E' X'} [IsFiniteKernel κ₁] [IsFiniteKernel κ₂] + {ρ : X' → ENNReal} (hρ : Measurable ρ) + [IsSFiniteKernel (κ₂.withDensity (fun _ => ρ))] + [SFinite ((κ₂ ∘ₘ Q).withDensity ρ)] + (h_ae : κ₁ =ᵐ[Q] κ₂.withDensity (fun _ => ρ)) : + κ₂†Q =ᵐ[κ₁ ∘ₘ Q] κ₁†Q := by + apply ae_eq_posterior_of_compProd_eq + have h2 : Q ⊗ₘ (κ₂.withDensity (fun _ => ρ)) + = (Q ⊗ₘ κ₂).withDensity (ρ ∘ Prod.snd) := by + have := Measure.compProd_withDensity (κ := κ₂) (μ := Q) + (show Measurable (Function.uncurry (fun _ => ρ)) from hρ.comp measurable_snd) + convert this using 1 + calc (κ₁ ∘ₘ Q) ⊗ₘ (κ₂†Q) + = ((κ₂ ∘ₘ Q).withDensity ρ) ⊗ₘ (κ₂†Q) := by + congr 1; rw [← comp_withDensity_const hρ]; exact Measure.bind_congr_right h_ae + _ = ((κ₂ ∘ₘ Q) ⊗ₘ (κ₂†Q)).withDensity (ρ ∘ Prod.fst) := + withDensity_compProd_left hρ + _ = ((Q ⊗ₘ κ₂).map Prod.swap).withDensity (ρ ∘ Prod.fst) := by + rw [compProd_posterior_eq_map_swap] + _ = ((Q ⊗ₘ κ₂).withDensity (ρ ∘ Prod.snd)).map Prod.swap := by + rw [map_swap_withDensity_fst hρ] + _ = (Q ⊗ₘ (κ₂.withDensity (fun _ => ρ))).map Prod.swap := by rw [h2] + _ = (Q ⊗ₘ κ₁).map Prod.swap := by rw [Measure.compProd_congr h_ae] + +end PosteriorEquality + +section DensityIndependence + +variable {K : ℕ} [Nonempty (Fin K)] + +/-- The history distribution under any algorithm is a `withDensity` of the history distribution +under the uniform algorithm, with a density that does not depend on the reward kernel `ν`. +This is the key factorization property: the density ratio only involves action probabilities. -/ +private lemma exists_density_independent_of_env + (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) (t : ℕ) : + ∃ ρ : (Iic t → Fin K × ℝ) → ENNReal, Measurable ρ ∧ (∀ h, ρ h ≠ ⊤) ∧ + ∀ (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν], + (trajMeasure alg (stationaryEnv ν)).map (IT.hist t) = + ((trajMeasure (Bandits.uniformAlgorithm hK) (stationaryEnv ν)).map + (IT.hist t)).withDensity ρ := by + set unif := Bandits.uniformAlgorithm hK + induction t with + | zero => + set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) + refine ⟨(alg.p0.rnDeriv unif.p0 ∘ Prod.fst) ∘ e, + (Measure.measurable_rnDeriv _ _).comp (measurable_fst.comp e.measurable), + fun h => rnDeriv_ne_top_of_forall_singleton_pos (uniformAlgorithm_p0_pos hK) _, ?_⟩ + intro ν _ + have h_ac : alg.p0 ≪ unif.p0 := + absolutelyContinuous_of_forall_singleton_pos (uniformAlgorithm_p0_pos hK) + simp only [IT.hist_eq_frestrictLe, trajMeasure, + Kernel.trajMeasure_map_frestrictLe, Kernel.partialTraj_self, + Measure.id_comp, stationaryEnv_ν0] + conv_lhs => rw [← Measure.withDensity_rnDeriv_eq _ _ h_ac] + rw [withDensity_compProd_left (Measure.measurable_rnDeriv _ _)] + exact withDensity_map_equiv_symm + ((Measure.measurable_rnDeriv _ _).comp measurable_fst) + | succ n ih => + obtain ⟨ρ_n, hρ_n_meas, hρ_n_ne_top, hρ_n⟩ := ih + let σ : (Iic n → Fin K × ℝ) → (Fin K × ℝ) → ENNReal := + fun h ar => Kernel.rnDeriv (alg.policy n) (unif.policy n) h ar.1 + let e := MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n + have hσ_meas : Measurable (Function.uncurry σ) := + (Kernel.measurable_rnDeriv _ _).comp + (measurable_fst.prodMk (measurable_fst.comp measurable_snd)) + refine ⟨(ρ_n ∘ Prod.fst * Function.uncurry σ) ∘ e, ?_, ?_, ?_⟩ + · exact ((hρ_n_meas.comp measurable_fst).mul hσ_meas).comp e.measurable + · intro h + exact ENNReal.mul_ne_top (hρ_n_ne_top _) + (kernel_rnDeriv_ne_top_of_forall_singleton_pos + (fun h' a => uniformAlgorithm_policy_pos hK n h' a) _ _) + · intro ν _inst + have h_step : stepKernel alg (stationaryEnv ν) n = + (stepKernel unif (stationaryEnv ν) n).withDensity σ := by + ext h : 1 + rw [Kernel.withDensity_apply _ hσ_meas] + have h_alg : stepKernel alg (stationaryEnv ν) n h = (alg.policy n h) ⊗ₘ ν := by + ext s hs + simp only [stepKernel, stationaryEnv, Kernel.compProd_apply hs, + Measure.compProd_apply hs, Kernel.prodMkLeft_apply] + have h_unif : stepKernel unif (stationaryEnv ν) n h = (unif.policy n h) ⊗ₘ ν := by + ext s hs + simp only [stepKernel, stationaryEnv, Kernel.compProd_apply hs, + Measure.compProd_apply hs, Kernel.prodMkLeft_apply] + have h_wd : ((unif.policy n) h).withDensity + (Kernel.rnDeriv (alg.policy n) (unif.policy n) h) = alg.policy n h := by + rw [← Kernel.withDensity_apply _ (Kernel.measurable_rnDeriv _ _)] + exact Kernel.withDensity_rnDeriv_eq (κ := alg.policy n) (η := unif.policy n) + (absolutelyContinuous_of_forall_singleton_pos (uniformAlgorithm_policy_pos hK n h)) + rw [h_alg, h_unif, ← h_wd] + haveI : SFinite ((unif.policy n h).withDensity + (Kernel.rnDeriv (alg.policy n) (unif.policy n) h)) := by + rw [h_wd]; infer_instance + exact withDensity_compProd_left + (Kernel.measurable_rnDeriv (alg.policy n) (unif.policy n)).of_uncurry_left + haveI : IsSFiniteKernel ((stepKernel unif (stationaryEnv ν) n).withDensity σ) := by + rw [← h_step]; infer_instance + rw [map_hist_succ_eq_compProd_map alg (stationaryEnv ν) n, + map_hist_succ_eq_compProd_map unif (stationaryEnv ν) n, + hρ_n ν, h_step, + withDensity_compProd_withDensity hρ_n_meas hσ_meas] + exact withDensity_map_equiv_symm + ((hρ_n_meas.comp measurable_fst).mul hσ_meas) + +end DensityIndependence + +section PosteriorIndependence + +/-! ### Algorithm-independence of the posterior + +The key theorem: the posterior distribution on the best arm given the observed history +is independent of the algorithm used to generate the data. The proof routes through +the uniform algorithm as a reference measure. The posterior on the environment given history +is algorithm-independent (ae wrt the algorithm's own history distribution), and this +transfers to the posterior on the best arm via `condDistrib_comp`. +-/ + +variable {K : ℕ} [Nonempty (Fin K)] +variable {E : Type*} [MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] +variable (Q : Measure E) [IsProbabilityMeasure Q] +variable (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] +variable {Ω : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] +variable {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → E × ℝ} +variable {alg : Algorithm (Fin K) ℝ} +variable {P : Measure Ω} [IsProbabilityMeasure P] + +/-- Maps an environment to the best arm (the arm with highest mean reward). -/ +noncomputable def envToBestArm (κ : Kernel (Fin K × E) ℝ) : E → Fin K := + measurableArgmax fun e a ↦ (κ (a, e))[id] + +omit [StandardBorelSpace E] [Nonempty E] [IsMarkovKernel κ] in +lemma measurable_envToBestArm : Measurable (envToBestArm κ) := + measurable_measurableArgmax fun _ ↦ + stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk measurable_id) + +omit [StandardBorelSpace E] [Nonempty E] [IsMarkovKernel κ] [MeasurableSpace Ω] + [IsProbabilityMeasure P] [Nonempty Ω] in +lemma bestArm_eq_envToBestArm_comp_env : + IsBayesAlgEnvSeq.bestArm κ R' = envToBestArm κ ∘ IsBayesAlgEnvSeq.env R' := by + funext ω + simp only [Function.comp_apply, IsBayesAlgEnvSeq.bestArm, envToBestArm, + IsBayesAlgEnvSeq.env] + exact (measurableArgmax_eq_of_eq _ _ _ ω).trans (measurableArgmax_congr _ _ ω _ rfl) + +/-- The marginal on the history equals `condDistrib (hist) (env) P ∘ₘ Q`. -/ +private lemma map_hist_eq_condDistrib_comp + {Ω' : Type*} [MeasurableSpace Ω'] [StandardBorelSpace Ω'] [Nonempty Ω'] + {A' : ℕ → Ω' → Fin K} {R'' : ℕ → Ω' → E × ℝ} + {alg' : Algorithm (Fin K) ℝ} {P' : Measure Ω'} [IsProbabilityMeasure P'] + (h' : IsBayesAlgEnvSeq Q κ A' R'' alg' P') (t : ℕ) : + P'.map (IsBayesAlgEnvSeq.hist A' R'' t) = + condDistrib (IsBayesAlgEnvSeq.hist A' R'' t) (IsBayesAlgEnvSeq.env R'') P' ∘ₘ Q := by + calc P'.map (IsBayesAlgEnvSeq.hist A' R'' t) + _ = (P'.map (fun ω => (IsBayesAlgEnvSeq.env R'' ω, + IsBayesAlgEnvSeq.hist A' R'' t ω))).snd := + (Measure.snd_map_prodMk h'.measurable_env).symm + _ = (P'.map (IsBayesAlgEnvSeq.env R'') ⊗ₘ condDistrib + (IsBayesAlgEnvSeq.hist A' R'' t) (IsBayesAlgEnvSeq.env R'') P').snd := by + rw [compProd_map_condDistrib (h'.measurable_hist t).aemeasurable] + _ = (Q ⊗ₘ condDistrib (IsBayesAlgEnvSeq.hist A' R'' t) + (IsBayesAlgEnvSeq.env R'') P').snd := by rw [h'.hasLaw_env.map_eq] + _ = _ := Measure.snd_compProd Q _ + +/-- The history distribution under any algorithm is absolutely continuous w.r.t. the + history distribution under the uniform algorithm (since uniform gives positive + probability to every action). -/ +lemma absolutelyContinuous_map_hist_uniform + (h : IsBayesAlgEnvSeq Q κ A R' alg P) (hK : 0 < K) + {Ωu : Type*} [MeasurableSpace Ωu] [StandardBorelSpace Ωu] [Nonempty Ωu] + {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → E × ℝ} + {Pu : Measure Ωu} [IsProbabilityMeasure Pu] + (hu : IsBayesAlgEnvSeq Q κ Au Ru (Bandits.uniformAlgorithm hK) Pu) + (t : ℕ) : + P.map (IsBayesAlgEnvSeq.hist A R' t) ≪ + Pu.map (IsBayesAlgEnvSeq.hist Au Ru t) := by + set κ_alg := condDistrib (IsBayesAlgEnvSeq.hist A R' t) + (IsBayesAlgEnvSeq.env R') P + set κ_unif := condDistrib (IsBayesAlgEnvSeq.hist Au Ru t) + (IsBayesAlgEnvSeq.env Ru) Pu + rw [map_hist_eq_condDistrib_comp Q κ h t, map_hist_eq_condDistrib_comp Q κ hu t, + ← Measure.snd_compProd, ← Measure.snd_compProd] + exact (Measure.AbsolutelyContinuous.compProd_right + (show ∀ᵐ e ∂Q, κ_alg e ≪ κ_unif e from by + filter_upwards [condDistrib_hist_env_eq_traj Q κ h t, + condDistrib_hist_env_eq_traj Q κ hu t] with e he_alg he_unif + rw [he_alg, he_unif] + exact absolutelyContinuous_map_hist_stationary hK alg _ t)).map + measurable_snd + +/-- The posterior on the environment given history is algorithm-independent. -/ +lemma condDistrib_env_hist_alg_indep + (h : IsBayesAlgEnvSeq Q κ A R' alg P) (hK : 0 < K) + {Ωu : Type*} [MeasurableSpace Ωu] [StandardBorelSpace Ωu] [Nonempty Ωu] + {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → E × ℝ} + {Pu : Measure Ωu} [IsProbabilityMeasure Pu] + (hu : IsBayesAlgEnvSeq Q κ Au Ru (Bandits.uniformAlgorithm hK) Pu) + (t : ℕ) : + condDistrib (IsBayesAlgEnvSeq.env R') (IsBayesAlgEnvSeq.hist A R' t) P + =ᵐ[P.map (IsBayesAlgEnvSeq.hist A R' t)] + condDistrib (IsBayesAlgEnvSeq.env Ru) (IsBayesAlgEnvSeq.hist Au Ru t) Pu := by + set κ_alg := condDistrib (IsBayesAlgEnvSeq.hist A R' t) (IsBayesAlgEnvSeq.env R') P + set κ_unif := condDistrib (IsBayesAlgEnvSeq.hist Au Ru t) (IsBayesAlgEnvSeq.env Ru) Pu + obtain ⟨ρ, hρ_meas, hρ_ne_top, hρ⟩ := exists_density_independent_of_env hK alg t + -- Key factorization: κ_alg =ᵐ[Q] κ_unif.withDensity (fun _ => ρ) + have h_wd_ae : κ_alg =ᵐ[Q] κ_unif.withDensity (fun _ => ρ) := by + filter_upwards [condDistrib_hist_env_eq_traj Q κ h t, + condDistrib_hist_env_eq_traj Q κ hu t] with e he_alg he_unif + rw [Kernel.withDensity_apply _ + (show Measurable (Function.uncurry (fun (_ : E) => ρ)) from hρ_meas.comp measurable_snd), + he_alg, he_unif] + exact hρ _ + haveI : IsSFiniteKernel (κ_unif.withDensity (fun _ => ρ)) := + Kernel.IsSFiniteKernel.withDensity _ (fun _ b => hρ_ne_top b) + -- Posterior equality via density factorization + have h_post : posterior κ_unif Q + =ᵐ[P.map (IsBayesAlgEnvSeq.hist A R' t)] posterior κ_alg Q := by + rw [map_hist_eq_condDistrib_comp Q κ h t] + exact posterior_eq_of_withDensity_ae_eq hρ_meas h_wd_ae + -- Bayes' rule for both algorithms + have h1 := h.condDistrib_env_hist_eq_posterior t + have h2' : condDistrib (IsBayesAlgEnvSeq.env Ru) (IsBayesAlgEnvSeq.hist Au Ru t) Pu + =ᵐ[P.map (IsBayesAlgEnvSeq.hist A R' t)] posterior κ_unif Q := + (absolutelyContinuous_map_hist_uniform Q κ h hK hu t).ae_le + (hu.condDistrib_env_hist_eq_posterior t) + exact h1.trans (h_post.symm.trans h2'.symm) + +/-- The posterior on the best arm equals the uniform algorithm's posterior. -/ +lemma posteriorBestArm_eq_uniform + (h : IsBayesAlgEnvSeq Q κ A R' alg P) (hK : 0 < K) (t : ℕ) : + condDistrib (IsBayesAlgEnvSeq.bestArm κ R') (IsBayesAlgEnvSeq.hist A R' t) P + =ᵐ[P.map (IsBayesAlgEnvSeq.hist A R' t)] + IT.posteriorBestArm Q κ (Bandits.uniformAlgorithm hK) t := by + unfold IT.posteriorBestArm + set Pu := IT.bayesTrajMeasure Q κ (Bandits.uniformAlgorithm hK) + set histf := IsBayesAlgEnvSeq.hist A R' t + set histfu := IsBayesAlgEnvSeq.hist IT.action IT.reward t + set envf := IsBayesAlgEnvSeq.env R' + set envfu := IsBayesAlgEnvSeq.env (E := E) IT.reward + set bau := IsBayesAlgEnvSeq.bestArm (Ω := ℕ → Fin K × E × ℝ) κ IT.reward + have h_ITu := IT.isBayesAlgEnvSeq_bayesianTrajMeasure Q κ (Bandits.uniformAlgorithm hK) + -- LHS: condDistrib (bestArm κ R') histf P + -- =ᵐ (condDistrib envf histf P).map (envToBestArm κ) + have h_comp_alg : condDistrib (IsBayesAlgEnvSeq.bestArm κ R') histf P + =ᵐ[P.map histf] (condDistrib envf histf P).map (envToBestArm κ) := by + rw [bestArm_eq_envToBestArm_comp_env κ] + exact condDistrib_comp (mβ := MeasurableSpace.pi) histf + h.measurable_env.aemeasurable (measurable_envToBestArm κ) + -- RHS: condDistrib bau histfu Pu + -- =ᵐ (condDistrib envfu histfu Pu).map (envToBestArm κ) + have h_comp_unif : condDistrib bau histfu Pu + =ᵐ[Pu.map histfu] (condDistrib envfu histfu Pu).map (envToBestArm κ) := by + change condDistrib (IsBayesAlgEnvSeq.bestArm κ IT.reward) histfu Pu + =ᵐ[Pu.map histfu] (condDistrib envfu histfu Pu).map (envToBestArm κ) + rw [bestArm_eq_envToBestArm_comp_env κ] + exact condDistrib_comp (mβ := MeasurableSpace.pi) histfu + h_ITu.measurable_env.aemeasurable (measurable_envToBestArm κ) + -- Environment posterior independence + have h_env_indep := condDistrib_env_hist_alg_indep Q κ h hK h_ITu t + -- Map both sides by envToBestArm + have h_map_indep : (condDistrib envf histf P).map (envToBestArm κ) + =ᵐ[P.map histf] (condDistrib envfu histfu Pu).map (envToBestArm κ) := by + filter_upwards [h_env_indep] with x hx + simp only [Kernel.map_apply _ (measurable_envToBestArm κ)] + rw [hx] + -- Transfer h_comp_unif from ae[Pu.map histfu] to ae[P.map histf] + exact h_comp_alg.trans (h_map_indep.trans + (h_comp_unif.filter_mono + (absolutelyContinuous_map_hist_uniform Q κ h hK h_ITu t).ae_le).symm) + +end PosteriorIndependence + +end Learning From 347f2c8e76fc79964d672cd003438ad1dc0a05ac Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Tue, 10 Feb 2026 14:23:47 +0100 Subject: [PATCH 027/155] fix: add import --- LeanBandits/BanditAlgorithms/TS.lean | 1 + 1 file changed, 1 insertion(+) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index d08b49f2..a52df0f2 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -8,6 +8,7 @@ import LeanBandits.BanditAlgorithms.Uniform import LeanBandits.BanditAlgorithms.UCB import LeanBandits.SequentialLearning.BayesStationaryEnv import LeanBandits.SequentialLearning.HistoryDensity +import Mathlib.Analysis.Complex.ExponentialBounds /-! # The Thompson Sampling Algorithm -/ From e6144378ad8abc2d3c36d609e30ba310c8309ffc Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Tue, 10 Feb 2026 14:26:57 +0100 Subject: [PATCH 028/155] remove private; move lemma --- LeanBandits/Bandit/SumRewards.lean | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 318a577e..5e7729d2 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -17,6 +17,15 @@ namespace Bandits namespace ArrayModel +lemma sum_Icc_one_eq_sum_range {m : ℕ} {f : ℕ → ℝ} : + ∑ i ∈ Icc 1 m, f (i - 1) = ∑ i ∈ range m, f i := by + have h : Icc 1 m = (range m).image (· + 1) := by + ext x; simp only [mem_Icc, mem_image, mem_range]; constructor + · intro ⟨h1, h2⟩; exact ⟨x - 1, by omega, by omega⟩ + · rintro ⟨a, ha, rfl⟩; omega + rw [h, Finset.sum_image (fun _ _ _ _ h => by omega)] + simp + variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [Countable α] [StandardBorelSpace α] [Nonempty α] {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] @@ -77,15 +86,6 @@ lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount' (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) -private lemma sum_Icc_one_eq_sum_range {m : ℕ} {f : ℕ → ℝ} : - ∑ i ∈ Icc 1 m, f (i - 1) = ∑ i ∈ range m, f i := by - have h : Icc 1 m = (range m).image (· + 1) := by - ext x; simp only [mem_Icc, mem_image, mem_range]; constructor - · intro ⟨h1, h2⟩; exact ⟨x - 1, by omega, by omega⟩ - · rintro ⟨a, ha, rfl⟩; omega - rw [h, Finset.sum_image (fun _ _ _ _ h => by omega)] - simp - 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 ω)) From 09ddc7094cbdd93e45105867ac4c095277216502 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Tue, 10 Feb 2026 14:28:44 +0100 Subject: [PATCH 029/155] undo undesirable minor changes --- LeanBandits/ForMathlib/MeasurableArgMax.lean | 26 ++++++++++++-------- 1 file changed, 16 insertions(+), 10 deletions(-) diff --git a/LeanBandits/ForMathlib/MeasurableArgMax.lean b/LeanBandits/ForMathlib/MeasurableArgMax.lean index 02fe3cc2..ca6f64bd 100644 --- a/LeanBandits/ForMathlib/MeasurableArgMax.lean +++ b/LeanBandits/ForMathlib/MeasurableArgMax.lean @@ -18,7 +18,9 @@ lemma measurable_encode {α : Type*} {_ : MeasurableSpace α} [Encodable α] [MeasurableSingletonClass α] : Measurable (Encodable.encode (α := α)) := by refine measurable_to_nat fun a ↦ ?_ - rw [show Encodable.encode ⁻¹' {Encodable.encode a} = {a} from by ext; simp]; measurability + have : Encodable.encode ⁻¹' {Encodable.encode a} = {a} := by ext; simp + rw [this] + exact measurableSet_singleton _ lemma measurableEmbedding_encode (α : Type*) {_ : MeasurableSpace α} [Encodable α] [MeasurableSingletonClass α] : @@ -37,12 +39,13 @@ lemma measurableSet_isMax [Countable 𝓨] {f : 𝓧 → 𝓨 → α} (hf : ∀ y, Measurable (fun x ↦ f x y)) (y : 𝓨) : MeasurableSet {x | ∀ z, f x z ≤ f x y} := by rw [show {x | ∀ y', f x y' ≤ f x y} = ⋂ y', {x | f x y' ≤ f x y} by ext; simp] - exact .iInter fun z ↦ measurableSet_le (hf z) (hf y) + exact MeasurableSet.iInter fun z ↦ measurableSet_le (by fun_prop) (by fun_prop) lemma exists_isMaxOn' {α : Type*} [LinearOrder α] [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] (f : 𝓧 → 𝓨 → α) (x : 𝓧) : - ∃ n : ℕ, ∃ y, n = Encodable.encode y ∧ ∀ z, f x z ≤ f x y := - let ⟨y, h⟩ := Finite.exists_max (f x); ⟨Encodable.encode y, y, rfl, h⟩ + ∃ n : ℕ, ∃ y, n = Encodable.encode y ∧ ∀ z, f x z ≤ f x y := by + obtain ⟨y, h⟩ := Finite.exists_max (f x) + exact ⟨Encodable.encode y, y, rfl, h⟩ /-- A measurable argmax function. -/ noncomputable @@ -58,11 +61,13 @@ lemma measurable_measurableArgmax [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y] (hf : ∀ y, Measurable (fun x ↦ f x y)) : Measurable (measurableArgmax f) := by - refine (MeasurableEmbedding.measurable_invFun (measurableEmbedding_encode 𝓨)).comp - (measurable_find _ fun n ↦ ?_) - rw [show {x | ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y} - = ⋃ y, ({x | n = Encodable.encode y} ∩ {x | ∀ z, f x z ≤ f x y}) from by ext; simp] - exact .iUnion fun y ↦ .inter (by simp) (measurableSet_isMax hf y) + refine (MeasurableEmbedding.measurable_invFun (measurableEmbedding_encode 𝓨)).comp ?_ + refine measurable_find _ fun n ↦ ?_ + have : {x | ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y} + = ⋃ y, ({x | n = Encodable.encode y} ∩ {x | ∀ z, f x z ≤ f x y}) := by ext; simp + rw [this] + refine MeasurableSet.iUnion fun y ↦ (MeasurableSet.inter (by simp) ?_) + exact measurableSet_isMax (by fun_prop) y lemma isMaxOn_measurableArgmax {α : Type*} [LinearOrder α] [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨] @@ -71,7 +76,8 @@ lemma isMaxOn_measurableArgmax {α : Type*} [LinearOrder α] (x : 𝓧) (z : 𝓨) : f x z ≤ f x (measurableArgmax f x) := by obtain ⟨y, h_eq, h_le⟩ := Nat.find_spec (exists_isMaxOn' f x) - exact (h_le z).trans_eq <| by rw [measurableArgmax, h_eq, + refine le_trans (h_le z) (le_of_eq ?_) + rw [measurableArgmax, h_eq, MeasurableEmbedding.leftInverse_invFun (measurableEmbedding_encode 𝓨) y] /-- Congruence lemma: measurableArgmax only depends on the function values at the point. -/ From 8c7175808ae811fe85a99b246da1a39f1940280b Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Tue, 10 Feb 2026 14:30:15 +0100 Subject: [PATCH 030/155] minor fixes --- LeanBandits/SequentialLearning/FiniteActions.lean | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index 3469b61f..29416ba9 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -266,9 +266,10 @@ lemma stepsUntil_zero_of_ne (hka : A 0 ω ≠ a) : stepsUntil A a 0 ω = 0 := by 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; exact (this.trans_le (monotone_pullCount _ _ (by omega))).ne' - rw [← hka, ← zero_add 1, pullCount_action_eq_pullCount_add_one]; simp + suffices 0 < pullCount A a 1 ω from + fun _ ↦ (this.trans_le (monotone_pullCount _ _ (by omega))).ne' + 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)] : @@ -749,6 +750,7 @@ noncomputable def empMean' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := (sumRewards' n h a) / (pullCount' n h a) +@[simp] lemma sumRewards_zero {R' : ℕ → Ω → ℝ} : sumRewards A R' a 0 = 0 := by ext; simp [sumRewards] lemma sumRewards_add_one {R' : ℕ → Ω → ℝ} : @@ -772,7 +774,6 @@ lemma sumRewards_eq_of_pullCount_eq {R' : ℕ → Ω → ℝ} {s t : ℕ} (hst : simp only [sumRewards, sum_range_succ, if_neg hne, add_zero] exact ih h_eq_t -@[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 From e15f795f537d18876b29d0b0b70ce065e0b49b43 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Tue, 10 Feb 2026 14:33:27 +0100 Subject: [PATCH 031/155] lint --- LeanBandits/BanditAlgorithms/Uniform.lean | 2 +- LeanBandits/ForMathlib/Measurable.lean | 3 ++- LeanBandits/SequentialLearning/HistoryDensity.lean | 6 +++--- 3 files changed, 6 insertions(+), 5 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/Uniform.lean b/LeanBandits/BanditAlgorithms/Uniform.lean index d422b689..7ced2d55 100644 --- a/LeanBandits/BanditAlgorithms/Uniform.lean +++ b/LeanBandits/BanditAlgorithms/Uniform.lean @@ -19,7 +19,7 @@ noncomputable def uniformAlgorithm (hK : 0 < K) : Algorithm (Fin K) ℝ := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK have : IsProbabilityMeasure (uniformOn (Set.univ : Set (Fin K))) := - uniformOn_isProbabilityMeasure Set.finite_univ Set.univ_nonempty + isProbabilityMeasure_uniformOn Set.finite_univ Set.univ_nonempty { policy _ := Kernel.const _ (uniformOn Set.univ) p0 := uniformOn Set.univ } diff --git a/LeanBandits/ForMathlib/Measurable.lean b/LeanBandits/ForMathlib/Measurable.lean index 5b6b03f7..438e8317 100644 --- a/LeanBandits/ForMathlib/Measurable.lean +++ b/LeanBandits/ForMathlib/Measurable.lean @@ -87,12 +87,13 @@ lemma measurable_sum_Icc_of_le {f : ℕ → α → ℝ} {g : α → ℕ} {n : refine Measurable.ite ?_ (by fun_prop) (by fun_prop) exact (measurableSet_singleton _).preimage (by fun_prop) -lemma measurable_apply_fin {α' : Type*} [MeasurableSpace α'] [Fintype α'] +lemma measurable_apply_fin {α' : Type*} [MeasurableSpace α'] [Finite α'] [MeasurableSingletonClass α'] {f : α' → α → ℝ} {g : α → α'} (hf : ∀ a, Measurable (f a)) (hg : Measurable g) : Measurable (fun ω ↦ f (g ω) ω) := by classical + have := Fintype.ofFinite α' have : (fun ω ↦ f (g ω) ω) = fun ω ↦ ∑ a : α', if g ω = a then f a ω else 0 := by ext ω; simp [Finset.sum_ite_eq] rw [this] diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 5c3c4df5..7929b2f1 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -49,7 +49,7 @@ lemma uniformAlgorithm_policy_pos (n : ℕ) (h : Iic n → Fin K × ℝ) (a : Fi /-- Any measure on a finite type is absolutely continuous wrt any measure giving positive mass to all singletons. -/ lemma absolutelyContinuous_of_forall_singleton_pos {α : Type*} [MeasurableSpace α] - [MeasurableSingletonClass α] [Fintype α] + [MeasurableSingletonClass α] [Finite α] {μ ν : Measure α} [IsFiniteMeasure μ] (hν : ∀ a : α, ν {a} > 0) : μ ≪ ν := by intro s hs @@ -62,7 +62,7 @@ lemma absolutelyContinuous_of_forall_singleton_pos {α : Type*} [MeasurableSpace /-- `rnDeriv` is pointwise finite when the reference measure has full support on singletons. -/ lemma rnDeriv_ne_top_of_forall_singleton_pos {α : Type*} [MeasurableSpace α] - [MeasurableSingletonClass α] [Fintype α] + [MeasurableSingletonClass α] [Finite α] {μ ν : Measure α} [IsFiniteMeasure μ] [IsFiniteMeasure ν] (hν : ∀ a, ν {a} > 0) (a : α) : μ.rnDeriv ν a ≠ ⊤ := by intro h_eq @@ -76,7 +76,7 @@ lemma rnDeriv_ne_top_of_forall_singleton_pos {α : Type*} [MeasurableSpace α] on singletons. -/ lemma kernel_rnDeriv_ne_top_of_forall_singleton_pos {α' β' : Type*} [MeasurableSpace α'] [MeasurableSpace β'] - [MeasurableSingletonClass β'] [Fintype β'] + [MeasurableSingletonClass β'] [Finite β'] [MeasurableSpace.CountableOrCountablyGenerated α' β'] {κ η : Kernel α' β'} [IsFiniteKernel κ] [IsFiniteKernel η] (hη : ∀ a b, η a {b} > 0) (a : α') (b : β') : From b8e5a04501669427d6414910c3e6605cfe5e1193 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Tue, 10 Feb 2026 14:38:18 +0100 Subject: [PATCH 032/155] lint --- LeanBandits/BanditAlgorithms/TS.lean | 6 +++--- LeanBandits/SequentialLearning/BayesStationaryEnv.lean | 6 ++---- LeanBandits/SequentialLearning/HistoryDensity.lean | 9 +++------ 3 files changed, 8 insertions(+), 13 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index a52df0f2..69c8a85e 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -401,7 +401,7 @@ lemma streamMeasure_concentration_ge_delta {α : Type*} [MeasurableSpace α] rw [mul_div_assoc, mul_div_cancel₀ _ (by positivity : (2 * k : ℝ) ≠ 0)] rw [Real.exp_log (by positivity), one_div, inv_inv] -lemma prob_concentration_single_delta_cond [StandardBorelSpace Ω] [Nonempty Ω] [Nonempty (Fin K)] +lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) @@ -536,7 +536,7 @@ lemma prob_concentration_single_delta_cond [StandardBorelSpace Ω] [Nonempty Ω] rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast s, ← ENNReal.ofReal_mul (Nat.cast_nonneg s)] congr 1; ring -lemma prob_concentration_single_delta [StandardBorelSpace Ω] [Nonempty Ω] [Nonempty (Fin K)] +lemma prob_concentration_single_delta [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) @@ -595,7 +595,7 @@ lemma prob_concentration_single_delta [StandardBorelSpace Ω] [Nonempty Ω] [Non rw [lintegral_const, Measure.map_apply h.measurable_env MeasurableSet.univ] simp [measure_univ] -lemma prob_concentration_fail_delta [StandardBorelSpace Ω] [Nonempty Ω] [Nonempty (Fin K)] +lemma prob_concentration_fail_delta [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index b6e3f402..7fd3da21 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -201,8 +201,7 @@ section Posterior /-- The posterior on the environment given history equals Mathlib's `posterior` applied to the likelihood kernel and prior. This is the measure-theoretic formulation of Bayes' rule. -/ -lemma condDistrib_env_hist_eq_posterior [StandardBorelSpace Ω] - (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : +lemma condDistrib_env_hist_eq_posterior (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : condDistrib (env R') (hist A R' n) P =ᵐ[P.map (hist A R' n)] posterior (condDistrib (hist A R' n) (env R') P) Q := by -- The key is to show P.map (env, hist) = Q ⊗ₘ condDistrib hist env P @@ -382,8 +381,7 @@ lemma hasCondDistrib_reward_condDistrib (h : IsBayesAlgEnvSeq Q κ A R' alg P) ( ext s _ rw [Kernel.sectR_apply, ha, Kernel.comap_apply, Kernel.prodMkLeft_apply, Kernel.comap_apply] -lemma condDistrib_traj_isAlgEnvSeq [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ A R' alg P) : +lemma condDistrib_traj_isAlgEnvSeq (h : IsBayesAlgEnvSeq Q κ A R' alg P) : ∀ᵐ e ∂(P.map (env R')), IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv (κ.comap (·, e) (by fun_prop))) diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 7929b2f1..4b3b9b26 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -48,9 +48,8 @@ lemma uniformAlgorithm_policy_pos (n : ℕ) (h : Iic n → Fin K × ℝ) (a : Fi /-- Any measure on a finite type is absolutely continuous wrt any measure giving positive mass to all singletons. -/ -lemma absolutelyContinuous_of_forall_singleton_pos {α : Type*} [MeasurableSpace α] - [MeasurableSingletonClass α] [Finite α] - {μ ν : Measure α} [IsFiniteMeasure μ] +lemma absolutelyContinuous_of_forall_singleton_pos {α : Type*} {mα : MeasurableSpace α} + {μ ν : Measure α} (hν : ∀ a : α, ν {a} > 0) : μ ≪ ν := by intro s hs have h_empty : s = ∅ := by @@ -62,8 +61,7 @@ lemma absolutelyContinuous_of_forall_singleton_pos {α : Type*} [MeasurableSpace /-- `rnDeriv` is pointwise finite when the reference measure has full support on singletons. -/ lemma rnDeriv_ne_top_of_forall_singleton_pos {α : Type*} [MeasurableSpace α] - [MeasurableSingletonClass α] [Finite α] - {μ ν : Measure α} [IsFiniteMeasure μ] [IsFiniteMeasure ν] + {μ ν : Measure α} [SigmaFinite μ] (hν : ∀ a, ν {a} > 0) (a : α) : μ.rnDeriv ν a ≠ ⊤ := by intro h_eq have h_mem : a ∈ {x | ¬ (μ.rnDeriv ν x < ⊤)} := by simp [h_eq] @@ -76,7 +74,6 @@ lemma rnDeriv_ne_top_of_forall_singleton_pos {α : Type*} [MeasurableSpace α] on singletons. -/ lemma kernel_rnDeriv_ne_top_of_forall_singleton_pos {α' β' : Type*} [MeasurableSpace α'] [MeasurableSpace β'] - [MeasurableSingletonClass β'] [Finite β'] [MeasurableSpace.CountableOrCountablyGenerated α' β'] {κ η : Kernel α' β'} [IsFiniteKernel κ] [IsFiniteKernel η] (hη : ∀ a b, η a {b} > 0) (a : α') (b : β') : From 165024e7b7b517923452eda5a2b7e363357f5992 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Tue, 10 Feb 2026 14:46:27 +0100 Subject: [PATCH 033/155] move uniform lemmas --- LeanBandits/BanditAlgorithms/TS.lean | 4 +-- LeanBandits/BanditAlgorithms/Uniform.lean | 15 ++++++++++- .../SequentialLearning/HistoryDensity.lean | 27 +++++-------------- 3 files changed, 22 insertions(+), 24 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 69c8a85e..664ab1e5 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -1100,9 +1100,9 @@ lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] nlinarith [show (1 : ℝ) ≤ K from Nat.one_le_cast.mpr (Nat.one_le_of_lt hK)] -- For t ≥ 2, we have δ = 1/t² < 1 · have ht2 : 2 ≤ t := by omega - have htpos : (0 : ℝ) < t := Nat.cast_pos.mpr (Nat.pos_of_ne_zero ht) + have htpos : (0 : ℝ) < t := by positivity have _ht1 : (1 : ℝ) ≤ t := Nat.one_le_cast.mpr (Nat.pos_of_ne_zero ht) - have hδ : (0 : ℝ) < 1 / (t : ℝ) ^ 2 := div_pos one_pos (pow_pos htpos 2) + have hδ : (0 : ℝ) < 1 / (t : ℝ) ^ 2 := by positivity have hδ1 : 1 / (t : ℝ) ^ 2 < 1 := by rw [div_lt_one (pow_pos htpos 2)] have ht2_real : (2 : ℝ) ≤ t := Nat.ofNat_le_cast.mpr ht2 diff --git a/LeanBandits/BanditAlgorithms/Uniform.lean b/LeanBandits/BanditAlgorithms/Uniform.lean index 7ced2d55..de148f93 100644 --- a/LeanBandits/BanditAlgorithms/Uniform.lean +++ b/LeanBandits/BanditAlgorithms/Uniform.lean @@ -12,7 +12,7 @@ open MeasureTheory ProbabilityTheory Learning namespace Bandits -variable {K : ℕ} +variable {K : ℕ} {hK : 0 < K} /-- The Uniform algorithm: actions are chosen uniformly at random. -/ noncomputable @@ -23,4 +23,17 @@ def uniformAlgorithm (hK : 0 < K) : Algorithm (Fin K) ℝ := { policy _ := Kernel.const _ (uniformOn Set.univ) p0 := uniformOn Set.univ } +/-- The uniform algorithm gives positive probability to every action. -/ +lemma uniformAlgorithm_p0_pos (a : Fin K) : (uniformAlgorithm hK).p0 {a} > 0 := by + simp only [uniformAlgorithm, uniformOn] + refine cond_pos_of_inter_ne_zero MeasurableSet.univ ?_ + simp only [Set.univ_inter, Measure.count_singleton, ne_eq, one_ne_zero, not_false_eq_true] + +/-- The uniform algorithm's policy gives positive probability to every action. -/ +lemma uniformAlgorithm_policy_pos {n : ℕ} (h : Finset.Iic n → Fin K × ℝ) (a : Fin K) : + (uniformAlgorithm hK).policy n h {a} > 0 := by + simp only [uniformAlgorithm, Kernel.const_apply, uniformOn] + refine cond_pos_of_inter_ne_zero MeasurableSet.univ ?_ + simp only [Set.univ_inter, Measure.count_singleton, ne_eq, one_ne_zero, not_false_eq_true] + end Bandits diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 4b3b9b26..79d56d6d 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -32,20 +32,6 @@ section UniformFullSupport variable {K : ℕ} (hK : 0 < K) -/-- The uniform algorithm gives positive probability to every action. -/ -lemma uniformAlgorithm_p0_pos (a : Fin K) : - (Bandits.uniformAlgorithm hK).p0 {a} > 0 := by - simp only [Bandits.uniformAlgorithm, uniformOn] - refine cond_pos_of_inter_ne_zero MeasurableSet.univ ?_ - simp only [Set.univ_inter, Measure.count_singleton, ne_eq, one_ne_zero, not_false_eq_true] - -/-- The uniform algorithm's policy gives positive probability to every action. -/ -lemma uniformAlgorithm_policy_pos (n : ℕ) (h : Iic n → Fin K × ℝ) (a : Fin K) : - (Bandits.uniformAlgorithm hK).policy n h {a} > 0 := by - simp only [Bandits.uniformAlgorithm, Kernel.const_apply, uniformOn] - refine cond_pos_of_inter_ne_zero MeasurableSet.univ ?_ - simp only [Set.univ_inter, Measure.count_singleton, ne_eq, one_ne_zero, not_false_eq_true] - /-- Any measure on a finite type is absolutely continuous wrt any measure giving positive mass to all singletons. -/ lemma absolutelyContinuous_of_forall_singleton_pos {α : Type*} {mα : MeasurableSpace α} @@ -193,7 +179,7 @@ private lemma absolutelyContinuous_stepKernel_stationary simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] rw [h1, h2] exact Measure.AbsolutelyContinuous.compProd_left - (absolutelyContinuous_of_forall_singleton_pos (uniformAlgorithm_policy_pos hK n h)) _ + (absolutelyContinuous_of_forall_singleton_pos (Bandits.uniformAlgorithm_policy_pos h)) _ -- `compProd` unfolding requires extra heartbeats /-- The history distribution at time `n + 1` decomposes as a compProd of the history at time `n` @@ -235,8 +221,7 @@ private lemma absolutelyContinuous_map_hist_stationary Kernel.trajMeasure_map_frestrictLe, Kernel.partialTraj_self, Measure.id_comp, stationaryEnv_ν0] exact (Measure.AbsolutelyContinuous.compProd_left - (absolutelyContinuous_of_forall_singleton_pos - (uniformAlgorithm_p0_pos hK)) _).map + (absolutelyContinuous_of_forall_singleton_pos Bandits.uniformAlgorithm_p0_pos) _).map (MeasurableEquiv.piUnique _).symm.measurable | succ n ih => rw [map_hist_succ_eq_compProd_map, map_hist_succ_eq_compProd_map] @@ -337,10 +322,10 @@ private lemma exists_density_independent_of_env set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) refine ⟨(alg.p0.rnDeriv unif.p0 ∘ Prod.fst) ∘ e, (Measure.measurable_rnDeriv _ _).comp (measurable_fst.comp e.measurable), - fun h => rnDeriv_ne_top_of_forall_singleton_pos (uniformAlgorithm_p0_pos hK) _, ?_⟩ + fun h => rnDeriv_ne_top_of_forall_singleton_pos Bandits.uniformAlgorithm_p0_pos _, ?_⟩ intro ν _ have h_ac : alg.p0 ≪ unif.p0 := - absolutelyContinuous_of_forall_singleton_pos (uniformAlgorithm_p0_pos hK) + absolutelyContinuous_of_forall_singleton_pos Bandits.uniformAlgorithm_p0_pos simp only [IT.hist_eq_frestrictLe, trajMeasure, Kernel.trajMeasure_map_frestrictLe, Kernel.partialTraj_self, Measure.id_comp, stationaryEnv_ν0] @@ -361,7 +346,7 @@ private lemma exists_density_independent_of_env · intro h exact ENNReal.mul_ne_top (hρ_n_ne_top _) (kernel_rnDeriv_ne_top_of_forall_singleton_pos - (fun h' a => uniformAlgorithm_policy_pos hK n h' a) _ _) + (fun h' a => Bandits.uniformAlgorithm_policy_pos h' a) _ _) · intro ν _inst have h_step : stepKernel alg (stationaryEnv ν) n = (stepKernel unif (stationaryEnv ν) n).withDensity σ := by @@ -379,7 +364,7 @@ private lemma exists_density_independent_of_env (Kernel.rnDeriv (alg.policy n) (unif.policy n) h) = alg.policy n h := by rw [← Kernel.withDensity_apply _ (Kernel.measurable_rnDeriv _ _)] exact Kernel.withDensity_rnDeriv_eq (κ := alg.policy n) (η := unif.policy n) - (absolutelyContinuous_of_forall_singleton_pos (uniformAlgorithm_policy_pos hK n h)) + (absolutelyContinuous_of_forall_singleton_pos (Bandits.uniformAlgorithm_policy_pos h)) rw [h_alg, h_unif, ← h_wd] haveI : SFinite ((unif.policy n h).withDensity (Kernel.rnDeriv (alg.policy n) (unif.policy n) h)) := by From 9b687fe4343b5a1389d3726a5edea60bff55b485 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Tue, 10 Feb 2026 14:54:44 +0100 Subject: [PATCH 034/155] minor --- .../SequentialLearning/HistoryDensity.lean | 57 ++++++++----------- 1 file changed, 25 insertions(+), 32 deletions(-) diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 79d56d6d..dadaec41 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -26,17 +26,14 @@ open scoped ENNReal NNReal namespace Learning -variable {K : ℕ} +variable {α β : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {μ ν : Measure α} {K : ℕ} section UniformFullSupport -variable {K : ℕ} (hK : 0 < K) +variable (hK : 0 < K) -/-- Any measure on a finite type is absolutely continuous wrt any measure giving positive mass - to all singletons. -/ -lemma absolutelyContinuous_of_forall_singleton_pos {α : Type*} {mα : MeasurableSpace α} - {μ ν : Measure α} - (hν : ∀ a : α, ν {a} > 0) : μ ≪ ν := by +/-- Any measure is absolutely continuous wrt any measure giving positive mass to all singletons. -/ +lemma absolutelyContinuous_of_forall_singleton_pos (hν : ∀ a : α, ν {a} > 0) : μ ≪ ν := by intro s hs have h_empty : s = ∅ := by by_contra h @@ -46,8 +43,7 @@ lemma absolutelyContinuous_of_forall_singleton_pos {α : Type*} {mα : Measurabl rw [h_empty, measure_empty] /-- `rnDeriv` is pointwise finite when the reference measure has full support on singletons. -/ -lemma rnDeriv_ne_top_of_forall_singleton_pos {α : Type*} [MeasurableSpace α] - {μ ν : Measure α} [SigmaFinite μ] +lemma rnDeriv_ne_top_of_forall_singleton_pos [SigmaFinite μ] (hν : ∀ a, ν {a} > 0) (a : α) : μ.rnDeriv ν a ≠ ⊤ := by intro h_eq have h_mem : a ∈ {x | ¬ (μ.rnDeriv ν x < ⊤)} := by simp [h_eq] @@ -59,10 +55,9 @@ lemma rnDeriv_ne_top_of_forall_singleton_pos {α : Type*} [MeasurableSpace α] /-- Kernel `rnDeriv` is pointwise finite when the reference kernel has full support on singletons. -/ lemma kernel_rnDeriv_ne_top_of_forall_singleton_pos - {α' β' : Type*} [MeasurableSpace α'] [MeasurableSpace β'] - [MeasurableSpace.CountableOrCountablyGenerated α' β'] - {κ η : Kernel α' β'} [IsFiniteKernel κ] [IsFiniteKernel η] - (hη : ∀ a b, η a {b} > 0) (a : α') (b : β') : + [MeasurableSpace.CountableOrCountablyGenerated α β] + {κ η : Kernel α β} [IsFiniteKernel κ] [IsFiniteKernel η] + (hη : ∀ a b, η a {b} > 0) (a : α) (b : β) : Kernel.rnDeriv κ η a b ≠ ⊤ := by intro h_eq have h_mem : b ∈ {x | ¬ (Kernel.rnDeriv κ η a x < ⊤)} := by simp [h_eq] @@ -75,13 +70,11 @@ end UniformFullSupport section WithDensityHelpers -variable {α' β' : Type*} {mα' : MeasurableSpace α'} {mβ' : MeasurableSpace β'} - /-- Composing `withDensity` on the measure side of a `compProd`: `(μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst)`. -/ private lemma withDensity_compProd_left - {μ : Measure α'} [SFinite μ] {κ : Kernel α' β'} [IsSFiniteKernel κ] - {f : α' → ENNReal} (hf : Measurable f) [SFinite (μ.withDensity f)] : + {μ : Measure α} [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] + {f : α → ℝ≥0∞} (hf : Measurable f) [SFinite (μ.withDensity f)] : (μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst) := by ext s hs rw [Measure.compProd_apply hs, withDensity_apply _ hs, @@ -98,7 +91,7 @@ private lemma withDensity_compProd_left /-- Mapping a `withDensity` through `MeasurableEquiv.symm`: `(μ.withDensity f).map e.symm = (μ.map e.symm).withDensity (f ∘ e)`. -/ private lemma withDensity_map_equiv_symm - {μ : Measure β'} {e : α' ≃ᵐ β'} {f : β' → ENNReal} (hf : Measurable f) : + {μ : Measure β} {e : α ≃ᵐ β} {f : β → ℝ≥0∞} (hf : Measurable f) : (μ.withDensity f).map e.symm = (μ.map e.symm).withDensity (f ∘ e) := by ext s hs rw [Measure.map_apply e.symm.measurable hs, @@ -109,8 +102,8 @@ private lemma withDensity_map_equiv_symm /-- Mapping a `withDensity` through a `MeasurableEquiv` from the snd component. -/ private lemma map_swap_withDensity_fst - {μ : Measure (α' × β')} [SFinite μ] - {f : β' → ENNReal} (hf : Measurable f) : + {μ : Measure (α × β)} [SFinite μ] + {f : β → ℝ≥0∞} (hf : Measurable f) : (μ.withDensity (f ∘ Prod.snd)).map Prod.swap = (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := by ext s hs @@ -121,7 +114,7 @@ private lemma map_swap_withDensity_fst /-- `(μ.withDensity (h ∘ g)).map g = (μ.map g).withDensity h`. -/ private lemma withDensity_map_eq' {γ' : Type*} {mγ' : MeasurableSpace γ'} - {μ : Measure α'} {g : α' → γ'} {h : γ' → ENNReal} + {μ : Measure α} {g : α → γ'} {h : γ' → ℝ≥0∞} (hg : Measurable g) (hh : Measurable h) : (μ.withDensity (h ∘ g)).map g = (μ.map g).withDensity h := by ext s hs @@ -132,13 +125,13 @@ private lemma withDensity_map_eq' /-- `(κ.withDensity (fun _ => ρ)) ∘ₘ Q = (κ ∘ₘ Q).withDensity ρ`. -/ private lemma comp_withDensity_const {γ' : Type*} {mγ' : MeasurableSpace γ'} - {Q : Measure α'} [SFinite Q] - {κ : Kernel α' γ'} [IsSFiniteKernel κ] - {ρ : γ' → ENNReal} (hρ : Measurable ρ) + {Q : Measure α} [SFinite Q] + {κ : Kernel α γ'} [IsSFiniteKernel κ] + {ρ : γ' → ℝ≥0∞} (hρ : Measurable ρ) [IsSFiniteKernel (κ.withDensity (fun _ => ρ))] : (κ.withDensity (fun _ => ρ)) ∘ₘ Q = (κ ∘ₘ Q).withDensity ρ := by rw [← Measure.snd_compProd Q (κ.withDensity (fun _ => ρ)), - Measure.compProd_withDensity (show Measurable (Function.uncurry (fun (_ : α') => ρ)) from + Measure.compProd_withDensity (show Measurable (Function.uncurry (fun (_ : α) => ρ)) from hρ.comp measurable_snd), ← Measure.snd_compProd Q κ, Measure.snd, Measure.snd] exact withDensity_map_eq' measurable_snd hρ @@ -146,9 +139,9 @@ private lemma comp_withDensity_const /-- `(μ.withDensity f) ⊗ₘ (κ.withDensity g) = (μ ⊗ₘ κ).withDensity (f ∘ fst * uncurry g)`. -/ private lemma withDensity_compProd_withDensity {γ' : Type*} {mγ' : MeasurableSpace γ'} - {μ : Measure α'} [SFinite μ] - {κ : Kernel α' γ'} [IsSFiniteKernel κ] - {f : α' → ENNReal} {g : α' → γ' → ENNReal} + {μ : Measure α} [SFinite μ] + {κ : Kernel α γ'} [IsSFiniteKernel κ] + {f : α → ℝ≥0∞} {g : α → γ' → ℝ≥0∞} (hf : Measurable f) (hg : Measurable (Function.uncurry g)) [SFinite (μ.withDensity f)] [IsSFiniteKernel (κ.withDensity g)] : (μ.withDensity f) ⊗ₘ (κ.withDensity g) @@ -160,7 +153,7 @@ end WithDensityHelpers section AbsolutelyContinuousHist -variable {K : ℕ} [Nonempty (Fin K)] +variable [Nonempty (Fin K)] omit [Nonempty (Fin K)] in /-- The step kernel for a stationary environment decomposes as a product of the policy @@ -277,7 +270,7 @@ private theorem posterior_eq_of_withDensity_ae_eq [StandardBorelSpace E'] [Nonempty E'] {Q : Measure E'} [IsFiniteMeasure Q] {κ₁ κ₂ : Kernel E' X'} [IsFiniteKernel κ₁] [IsFiniteKernel κ₂] - {ρ : X' → ENNReal} (hρ : Measurable ρ) + {ρ : X' → ℝ≥0∞} (hρ : Measurable ρ) [IsSFiniteKernel (κ₂.withDensity (fun _ => ρ))] [SFinite ((κ₂ ∘ₘ Q).withDensity ρ)] (h_ae : κ₁ =ᵐ[Q] κ₂.withDensity (fun _ => ρ)) : @@ -311,7 +304,7 @@ under the uniform algorithm, with a density that does not depend on the reward k This is the key factorization property: the density ratio only involves action probabilities. -/ private lemma exists_density_independent_of_env (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) (t : ℕ) : - ∃ ρ : (Iic t → Fin K × ℝ) → ENNReal, Measurable ρ ∧ (∀ h, ρ h ≠ ⊤) ∧ + ∃ ρ : (Iic t → Fin K × ℝ) → ℝ≥0∞, Measurable ρ ∧ (∀ h, ρ h ≠ ⊤) ∧ ∀ (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν], (trajMeasure alg (stationaryEnv ν)).map (IT.hist t) = ((trajMeasure (Bandits.uniformAlgorithm hK) (stationaryEnv ν)).map @@ -335,7 +328,7 @@ private lemma exists_density_independent_of_env ((Measure.measurable_rnDeriv _ _).comp measurable_fst) | succ n ih => obtain ⟨ρ_n, hρ_n_meas, hρ_n_ne_top, hρ_n⟩ := ih - let σ : (Iic n → Fin K × ℝ) → (Fin K × ℝ) → ENNReal := + let σ : (Iic n → Fin K × ℝ) → (Fin K × ℝ) → ℝ≥0∞ := fun h ar => Kernel.rnDeriv (alg.policy n) (unif.policy n) h ar.1 let e := MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n have hσ_meas : Measurable (Function.uncurry σ) := From 8f60d122dd53fdc08f9b7f5af55303d4a03e861b Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sun, 15 Feb 2026 20:46:00 +0100 Subject: [PATCH 035/155] minor --- LeanBandits/SequentialLearning/BayesStationaryEnv.lean | 5 ++--- LeanBandits/SequentialLearning/FiniteActions.lean | 7 +++---- 2 files changed, 5 insertions(+), 7 deletions(-) diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 7fd3da21..8f0fa092 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -5,7 +5,6 @@ Authors: Rémy Degenne, Paulo Rauber -/ import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.Bandit.Regret -import LeanBandits.SequentialLearning.IonescuTulceaSpace import LeanBandits.SequentialLearning.StationaryEnv import Mathlib.Probability.Kernel.Posterior @@ -40,8 +39,8 @@ def IsBayesAlgEnvSeq [StandardBorelSpace R] [Nonempty R] (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (alg : Algorithm α R) - (P : Measure Ω) [IsFiniteMeasure P] - := IsAlgEnvSeq A R' (alg.prod_left E) (bayesStationaryEnv Q κ) P + (P : Measure Ω) [IsFiniteMeasure P] := + IsAlgEnvSeq A R' (alg.prod_left E) (bayesStationaryEnv Q κ) P namespace IsBayesAlgEnvSeq diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index 29416ba9..acea00ba 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -764,15 +764,14 @@ lemma sumRewards_eq_of_pullCount_eq {R' : ℕ → Ω → ℝ} {s t : ℕ} (hst : induction t, hst using Nat.le_induction with | base => rfl | succ t hst' ih => - have h_mono : pullCount A a s ω ≤ pullCount A a t ω := pullCount_mono a hst' ω have h_mono' : pullCount A a t ω ≤ pullCount A a (t + 1) ω := pullCount_mono a (Nat.le_succ t) ω - have h_eq_t : pullCount A a s ω = pullCount A a t ω := le_antisymm h_mono (h_eq ▸ h_mono') + have h_eq_t : pullCount A a s ω = pullCount A a t ω := + le_antisymm (pullCount_mono a hst' ω) (h_eq ▸ h_mono') have hne : A t ω ≠ a := by intro ha have h1 := ha ▸ pullCount_action_eq_pullCount_add_one (A := A) t ω omega - simp only [sumRewards, sum_range_succ, if_neg hne, add_zero] - exact ih h_eq_t + rw [sumRewards_add_one, if_neg hne, add_zero, ih h_eq_t] lemma sumRewards_eq_pullCount_mul_empMean {R' : ℕ → Ω → ℝ} {ω : Ω} (h_pull : pullCount A a t ω ≠ 0) : From 87d12e3c7157628ab6a1ff7a2a477a922e9bb1e4 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sun, 15 Feb 2026 21:29:36 +0100 Subject: [PATCH 036/155] fix --- LeanBandits/BanditAlgorithms/TS.lean | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 664ab1e5..45722a13 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -760,7 +760,7 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] refine ⟨h_eq, ?_⟩ have h_pc_eq : pullCount IT.action a s ω = pullCount IT.action a (n - 1) ω := hpc.symm ▸ h_eq.symm - rw [← sumRewards_eq_of_pullCount_eq hs' h_pc_eq] + rw [← sumRewards_eq_of_pullCount_eq h_pc_eq] exact hB · right exact ⟨by omega, s, hs, hpc, hB⟩ From 59f01d9f9b397d641f421c96f2e4c9e04c85a0ee Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sun, 15 Feb 2026 21:59:43 +0100 Subject: [PATCH 037/155] minor --- .../SequentialLearning/HistoryDensity.lean | 37 +++++++++---------- 1 file changed, 18 insertions(+), 19 deletions(-) diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index dadaec41..dc611b6a 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -70,11 +70,12 @@ end UniformFullSupport section WithDensityHelpers +variable {γ : Type*} {mγ : MeasurableSpace γ} + /-- Composing `withDensity` on the measure side of a `compProd`: `(μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst)`. -/ -private lemma withDensity_compProd_left - {μ : Measure α} [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] - {f : α → ℝ≥0∞} (hf : Measurable f) [SFinite (μ.withDensity f)] : +private lemma withDensity_compProd_left [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] + {f : α → ℝ≥0∞} (hf : Measurable f) : (μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst) := by ext s hs rw [Measure.compProd_apply hs, withDensity_apply _ hs, @@ -82,9 +83,11 @@ private lemma withDensity_compProd_left (Kernel.measurable_kernel_prodMk_left hs).aemeasurable, ← lintegral_indicator hs, Measure.lintegral_compProd ((hf.comp measurable_fst).indicator hs)] - congr 1; ext a; simp_rw [Pi.mul_apply] - have : (fun b => s.indicator (f ∘ Prod.fst) (a, b)) = - fun b => (Prod.mk a ⁻¹' s).indicator (fun _ => f a) b := by + congr 1 + ext a + simp_rw [Pi.mul_apply] + have : (fun b ↦ s.indicator (f ∘ Prod.fst) (a, b)) = + fun b ↦ (Prod.mk a ⁻¹' s).indicator (fun _ ↦ f a) b := by ext b; simp only [Function.comp, Set.indicator, Set.mem_preimage]; rfl rw [this, lintegral_indicator_const (hs.preimage (by fun_prop))] @@ -113,8 +116,7 @@ private lemma map_swap_withDensity_fst /-- `(μ.withDensity (h ∘ g)).map g = (μ.map g).withDensity h`. -/ private lemma withDensity_map_eq' - {γ' : Type*} {mγ' : MeasurableSpace γ'} - {μ : Measure α} {g : α → γ'} {h : γ' → ℝ≥0∞} + {μ : Measure α} {g : α → γ} {h : γ → ℝ≥0∞} (hg : Measurable g) (hh : Measurable h) : (μ.withDensity (h ∘ g)).map g = (μ.map g).withDensity h := by ext s hs @@ -124,10 +126,9 @@ private lemma withDensity_map_eq' /-- `(κ.withDensity (fun _ => ρ)) ∘ₘ Q = (κ ∘ₘ Q).withDensity ρ`. -/ private lemma comp_withDensity_const - {γ' : Type*} {mγ' : MeasurableSpace γ'} {Q : Measure α} [SFinite Q] - {κ : Kernel α γ'} [IsSFiniteKernel κ] - {ρ : γ' → ℝ≥0∞} (hρ : Measurable ρ) + {κ : Kernel α γ} [IsSFiniteKernel κ] + {ρ : γ → ℝ≥0∞} (hρ : Measurable ρ) [IsSFiniteKernel (κ.withDensity (fun _ => ρ))] : (κ.withDensity (fun _ => ρ)) ∘ₘ Q = (κ ∘ₘ Q).withDensity ρ := by rw [← Measure.snd_compProd Q (κ.withDensity (fun _ => ρ)), @@ -137,15 +138,13 @@ private lemma comp_withDensity_const exact withDensity_map_eq' measurable_snd hρ /-- `(μ.withDensity f) ⊗ₘ (κ.withDensity g) = (μ ⊗ₘ κ).withDensity (f ∘ fst * uncurry g)`. -/ -private lemma withDensity_compProd_withDensity - {γ' : Type*} {mγ' : MeasurableSpace γ'} - {μ : Measure α} [SFinite μ] - {κ : Kernel α γ'} [IsSFiniteKernel κ] - {f : α → ℝ≥0∞} {g : α → γ' → ℝ≥0∞} +private lemma withDensity_compProd_withDensity [SFinite μ] + {κ : Kernel α γ} [IsSFiniteKernel κ] + {f : α → ℝ≥0∞} {g : α → γ → ℝ≥0∞} (hf : Measurable f) (hg : Measurable (Function.uncurry g)) - [SFinite (μ.withDensity f)] [IsSFiniteKernel (κ.withDensity g)] : - (μ.withDensity f) ⊗ₘ (κ.withDensity g) - = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst * Function.uncurry g) := by + [IsSFiniteKernel (κ.withDensity g)] : + (μ.withDensity f) ⊗ₘ (κ.withDensity g) = + (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst * Function.uncurry g) := by rw [Measure.compProd_withDensity hg, withDensity_compProd_left hf] exact (withDensity_mul _ (hf.comp measurable_fst) hg).symm From dbec0b7728e41c60a7d6d8ebc022d8abf316b5d5 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 18 Feb 2026 11:57:21 +0000 Subject: [PATCH 038/155] Redefine IsBayesAlgEnvSeq --- LeanBandits/BanditAlgorithms/TS.lean | 247 +++++------ .../BayesStationaryEnv.lean | 403 +++++++++--------- .../SequentialLearning/HistoryDensity.lean | 118 +++-- 3 files changed, 396 insertions(+), 372 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 45722a13..993fc636 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -59,15 +59,15 @@ section Regret variable (hK : 0 < K) variable {Ω : Type*} [MeasurableSpace Ω] -variable (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → E × ℝ) +variable (E' : Ω → E) (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → ℝ) variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] variable (P : Measure Ω) [IsProbabilityMeasure P] noncomputable -def ucbIndex (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → E × ℝ) (δ : ℝ) +def ucbIndex (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : ℝ := max 0 (min 1 - (empMean A (IsBayesAlgEnvSeq.reward R') a t ω + (empMean A R' a t ω + √(2 * Real.log (1 / δ) / (max 1 (pullCount A a t ω) : ℝ)))) omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in @@ -140,7 +140,7 @@ lemma sum_inv_sqrt_max_one_le (N : ℕ) : @[fun_prop] lemma measurable_ucbIndex [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) (δ : ℝ) (a : Fin K) (t : ℕ) : Measurable (ucbIndex A R' δ a t) := by unfold ucbIndex @@ -148,7 +148,7 @@ lemma measurable_ucbIndex [Nonempty (Fin K)] apply Measurable.min measurable_const apply Measurable.add · exact measurable_empMean (fun n ↦ h.measurable_A n) - (fun n ↦ h.measurable_reward n) a t + (fun n ↦ h.measurable_R n) a t · have hpc : Measurable (fun ω ↦ (pullCount A a t ω : ℝ)) := measurable_from_top.comp (measurable_pullCount (fun n ↦ h.measurable_A n) a t) exact (measurable_const.div (measurable_const.max hpc)).sqrt @@ -157,11 +157,11 @@ omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ lemma armMean_le_ucbIndex (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hconc : - |empMean A (IsBayesAlgEnvSeq.reward R') a t ω - IsBayesAlgEnvSeq.armMean κ R' a ω| + |empMean A R' a t ω - IsBayesAlgEnvSeq.armMean κ E' a ω| < √(2 * Real.log (1 / δ) / (max 1 (pullCount A a t ω) : ℝ))) : - IsBayesAlgEnvSeq.armMean κ R' a ω ≤ ucbIndex A R' δ a t ω := by + IsBayesAlgEnvSeq.armMean κ E' a ω ≤ ucbIndex A R' δ a t ω := by unfold ucbIndex - have hmean := hm a (IsBayesAlgEnvSeq.env R' ω) + have hmean := hm a (E' ω) simp only [IsBayesAlgEnvSeq.armMean] at hmean hconc ⊢ have habs := abs_sub_lt_iff.mp hconc refine le_max_of_le_right (le_min hmean.2 ?_) @@ -171,68 +171,70 @@ omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ lemma ucbIndex_sub_armMean_le (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hconc : - |empMean A (IsBayesAlgEnvSeq.reward R') a t ω - IsBayesAlgEnvSeq.armMean κ R' a ω| + |empMean A R' a t ω - IsBayesAlgEnvSeq.armMean κ E' a ω| < √(2 * Real.log (1 / δ) / (max 1 (pullCount A a t ω) : ℝ))) : - ucbIndex A R' δ a t ω - IsBayesAlgEnvSeq.armMean κ R' a ω + ucbIndex A R' δ a t ω - IsBayesAlgEnvSeq.armMean κ E' a ω ≤ 2 * √(2 * Real.log (1 / δ) / (max 1 (pullCount A a t ω) : ℝ)) := by unfold ucbIndex simp only [IsBayesAlgEnvSeq.armMean] at hconc ⊢ set w := √(2 * Real.log (1 / δ) / max 1 ↑(pullCount A a t ω)) - set emp := empMean A (IsBayesAlgEnvSeq.reward R') a t ω + set emp := empMean A R' a t ω have habs := abs_sub_lt_iff.mp hconc - have hmean := hm a (IsBayesAlgEnvSeq.env R' ω) + have hmean := hm a (E' ω) have h1 : max 0 (min 1 (emp + w)) ≤ emp + w := max_le_iff.mpr ⟨by linarith [hmean.1, habs.2], min_le_right _ _⟩ linarith [habs.2] lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) (t : ℕ) : - condDistrib (A (t + 1)) (IsBayesAlgEnvSeq.hist A R' t) P - =ᵐ[P.map (IsBayesAlgEnvSeq.hist A R' t)] - condDistrib (IsBayesAlgEnvSeq.bestArm κ R') (IsBayesAlgEnvSeq.hist A R' t) P := + (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) (t : ℕ) : + condDistrib (A (t + 1)) (IsAlgEnvSeq.hist A R' t) P + =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] + condDistrib (IsBayesAlgEnvSeq.bestArm κ E') (IsAlgEnvSeq.hist A R' t) P := (h.hasCondDistrib_action' t).condDistrib_eq.trans (posteriorBestArm_eq_uniform Q κ h hK t).symm omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma le_armMean_bestArm [Nonempty (Fin K)] (ω : Ω) (i : Fin K) : - IsBayesAlgEnvSeq.armMean κ R' i ω ≤ - IsBayesAlgEnvSeq.armMean κ R' (IsBayesAlgEnvSeq.bestArm κ R' ω) ω := by - have := isMaxOn_measurableArgmax (fun ω a ↦ IsBayesAlgEnvSeq.armMean κ R' a ω) ω i + IsBayesAlgEnvSeq.armMean κ E' i ω ≤ + IsBayesAlgEnvSeq.armMean κ E' (IsBayesAlgEnvSeq.bestArm κ E' ω) ω := by + have := isMaxOn_measurableArgmax (fun ω a ↦ IsBayesAlgEnvSeq.armMean κ E' a ω) ω i simp only [IsBayesAlgEnvSeq.bestArm]; convert this omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) - (ω : Ω) : ⨆ i, IsBayesAlgEnvSeq.armMean κ R' i ω = - IsBayesAlgEnvSeq.armMean κ R' (IsBayesAlgEnvSeq.bestArm κ R' ω) ω := - le_antisymm (ciSup_le (le_armMean_bestArm R' κ ω)) - (le_ciSup (f := fun i ↦ IsBayesAlgEnvSeq.armMean κ R' i ω) + (ω : Ω) : ⨆ i, IsBayesAlgEnvSeq.armMean κ E' i ω = + IsBayesAlgEnvSeq.armMean κ E' (IsBayesAlgEnvSeq.bestArm κ E' ω) ω := + le_antisymm (ciSup_le (le_armMean_bestArm E' κ ω)) + (le_ciSup (f := fun i ↦ IsBayesAlgEnvSeq.armMean κ E' i ω) ⟨1, by rintro _ ⟨i, rfl⟩; exact (hm i _).2⟩ _) omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma gap_eq_armMean_sub [Nonempty (Fin K)] (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) - (s : ℕ) (ω : Ω) : gap (κ.comap (·, IsBayesAlgEnvSeq.env R' ω) (by fun_prop)) (A s ω) = - IsBayesAlgEnvSeq.armMean κ R' (IsBayesAlgEnvSeq.bestArm κ R' ω) ω - - IsBayesAlgEnvSeq.armMean κ R' (A s ω) ω := by + (s : ℕ) (ω : Ω) : gap (κ.comap (·, E' ω) (by fun_prop)) (A s ω) = + IsBayesAlgEnvSeq.armMean κ E' (IsBayesAlgEnvSeq.bestArm κ E' ω) ω - + IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω := by simp only [gap, Kernel.comap_apply] - exact congr_arg (· - _) (iSup_armMean_eq_bestArm R' κ hm ω) + exact congr_arg (· - _) (iSup_armMean_eq_bestArm E' κ hm ω) +omit [StandardBorelSpace E] [Nonempty E] in lemma bayesRegret_eq_sum_integral_gap [Nonempty (Fin K)] {alg : Algorithm (Fin K) ℝ} - (h : IsBayesAlgEnvSeq Q κ A R' alg P) + (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) {C : ℝ} (hm : ∀ a e, |(κ (a, e))[id]| ≤ C) (t : ℕ) : - IsBayesAlgEnvSeq.bayesRegret κ A R' P t = - ∑ s ∈ range t, P[fun ω ↦ gap (κ.comap (·, IsBayesAlgEnvSeq.env R' ω) (by fun_prop)) + IsBayesAlgEnvSeq.bayesRegret κ A E' P t = + ∑ s ∈ range t, P[fun ω ↦ gap (κ.comap (·, E' ω) (by fun_prop)) (A s ω)] := by simp only [IsBayesAlgEnvSeq.bayesRegret, IsBayesAlgEnvSeq.regret, regret_eq_sum_gap] refine integral_finset_sum _ (fun s _ => ?_) - have hmeas : Measurable (fun ω ↦ gap (κ.comap (·, IsBayesAlgEnvSeq.env R' ω) (by fun_prop)) + have hmeas : Measurable (fun ω ↦ gap (κ.comap (·, E' ω) (by fun_prop)) (A s ω)) := (Measurable.iSup h.measurable_armMean).sub - (stronglyMeasurable_id.integral_kernel.measurable.comp (h.measurable_action_env s)) + (stronglyMeasurable_id.integral_kernel.measurable.comp + ((h.measurable_A s).prodMk h.measurable_E)) refine ⟨hmeas.aestronglyMeasurable, HasFiniteIntegral.of_bounded (C := 2 * C) (Filter.Eventually.of_forall fun ω => ?_)⟩ simp only [Real.norm_eq_abs, gap, Kernel.comap_apply] - set e := IsBayesAlgEnvSeq.env R' ω + set e := E' ω have hbdd : BddAbove (Set.range fun i => (κ (i, e))[id]) := ⟨C, by rintro _ ⟨i, rfl⟩; exact le_of_abs_le (hm i e)⟩ rw [abs_of_nonneg (sub_nonneg.mpr (le_ciSup hbdd _))] @@ -269,16 +271,16 @@ omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsProbabilityMeas lemma sum_ucbIndex_sub_armMean_le (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) (δ : ℝ) (n : ℕ) (ω : Ω) (hconc : ∀ s < n, ∀ a, - |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - IsBayesAlgEnvSeq.armMean κ R' a ω| + |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω| < √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ))) : - ∑ s ∈ range n, (ucbIndex A R' δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ R' (A s ω) ω) + ∑ s ∈ range n, (ucbIndex A R' δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) ≤ 2 * √(8 * Real.log (1 / δ)) * √(↑K * ↑n) := by have hterm : ∀ s ∈ range n, - ucbIndex A R' δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ R' (A s ω) ω + ucbIndex A R' δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω ≤ 2 * √(2 * Real.log (1 / δ) / (max 1 (pullCount A (A s ω) s ω) : ℝ)) := - fun s hs => ucbIndex_sub_armMean_le A R' κ hm δ (A s ω) s ω (hconc s (mem_range.mp hs) _) + fun s hs => ucbIndex_sub_armMean_le E' A R' κ hm δ (A s ω) s ω (hconc s (mem_range.mp hs) _) calc ∑ s ∈ range n, - (ucbIndex A R' δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ R' (A s ω) ω) + (ucbIndex A R' δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) ≤ ∑ s ∈ range n, 2 * √(2 * Real.log (1 / δ) / (max 1 (pullCount A (A s ω) s ω) : ℝ)) := sum_le_sum hterm @@ -402,13 +404,13 @@ lemma streamMeasure_concentration_ge_delta {α : Type*} [MeasurableSpace α] rw [Real.exp_log (by positivity), one_div, inv_inv] lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) (hδ_large : 1 < 2 * Real.log (1 / δ)) : - ∀ᵐ e ∂(P.map (IsBayesAlgEnvSeq.env R')), - (condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P e) + ∀ᵐ e ∂(P.map (E')), + (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) {ω | √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} ≤ ENNReal.ofReal (2 * s * δ) := by @@ -418,7 +420,7 @@ lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] simp only [ν, Kernel.comap_apply]; exact hs a' e have h_mean : (ν a)[id] = (κ (a, e))[id] := by simp only [ν, Kernel.comap_apply] rw [← h_mean] - let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P e + let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e have h_law := h_isAlgEnvSeq.law_pullCount_sumRewards_unique' (ArrayModel.isAlgEnvSeq_arrayMeasure (tsAlgorithm hK Q κ) ν) (n := s) let B_low := fun m : ℕ ↦ {x : ℝ | x / m + √(2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} @@ -537,37 +539,37 @@ lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] congr 1; ring lemma prob_concentration_single_delta [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) (hδ_large : 1 < 2 * Real.log (1 / δ)) : P {ω | √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ - |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - IsBayesAlgEnvSeq.armMean κ R' a ω|} ≤ + |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} ≤ ENNReal.ofReal (2 * s * δ) := by let badSet : E → Set (ℕ → (Fin K) × ℝ) := fun e ↦ {t | √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s t) : ℝ)) ≤ |empMean IT.action IT.reward a s t - (κ (a, e))[id]|} have h_set_eq : {ω | √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ - |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - IsBayesAlgEnvSeq.armMean κ R' a ω|} = - (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} = + (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' {p | p.2 ∈ badSet p.1} := by ext ω simp only [Set.mem_setOf_eq, Set.mem_preimage, badSet, IsBayesAlgEnvSeq.armMean] have h1 : pullCount A a s ω = pullCount IT.action a s (IsBayesAlgEnvSeq.traj A R' ω) := by unfold pullCount IsBayesAlgEnvSeq.traj IT.action; rfl - have h2 : empMean A (IsBayesAlgEnvSeq.reward R') a s ω = + have h2 : empMean A R' a s ω = empMean IT.action IT.reward a s (IsBayesAlgEnvSeq.traj A R' ω) := by - unfold empMean IsBayesAlgEnvSeq.traj IsBayesAlgEnvSeq.reward IT.action IT.reward; rfl + unfold empMean IsBayesAlgEnvSeq.traj IT.action IT.reward; rfl rw [h1, h2] have h_meas_pair : - Measurable (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) := - h.measurable_env.prodMk h.measurable_traj - have h_disint : P.map (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) = - P.map (IsBayesAlgEnvSeq.env R') ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P := + Measurable (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) := + h.measurable_E.prodMk h.measurable_traj + have h_disint : P.map (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) = + P.map (E') ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.traj A R') E' P := (compProd_map_condDistrib (h.measurable_traj.aemeasurable)).symm - have h_cond := prob_concentration_single_delta_cond hK A R' Q κ P h hs hm a s δ hδ hδ1 hδ_large + have h_cond := prob_concentration_single_delta_cond hK E' A R' Q κ P h hs hm a s δ hδ hδ1 hδ_large have h_kernel : Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (a, p.1))[id]) := stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk measurable_fst) have h_meas_set : MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | p.2 ∈ badSet p.1} := by @@ -577,39 +579,39 @@ lemma prob_concentration_single_delta [Nonempty (Fin K)] exact measurableSet_le (by fun_prop) (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp measurable_snd).sub h_kernel).abs - calc P _ = P ((fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + calc P _ = P ((fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' {p | p.2 ∈ badSet p.1}) := by rw [h_set_eq] - _ = (P.map (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω))) + _ = (P.map (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω))) {p | p.2 ∈ badSet p.1} := by rw [Measure.map_apply h_meas_pair h_meas_set] - _ = (P.map (IsBayesAlgEnvSeq.env R') ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P) + _ = (P.map (E') ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.traj A R') E' P) {p | p.2 ∈ badSet p.1} := by rw [h_disint] - _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P e) - (badSet e) ∂(P.map (IsBayesAlgEnvSeq.env R')) := by + _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + (badSet e) ∂(P.map (E')) := by rw [Measure.compProd_apply h_meas_set]; rfl - _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * s * δ) ∂(P.map (IsBayesAlgEnvSeq.env R')) := by + _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * s * δ) ∂(P.map (E')) := by apply lintegral_mono_ae filter_upwards [h_cond] with e h_e; exact h_e _ = ENNReal.ofReal (2 * s * δ) := by - rw [lintegral_const, Measure.map_apply h.measurable_env MeasurableSet.univ] + rw [lintegral_const, Measure.map_apply h.measurable_E MeasurableSet.univ] simp [measure_univ] lemma prob_concentration_fail_delta [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) (hδ_large : 1 < 2 * Real.log (1 / δ)) : P {ω | ∃ s < n, ∃ a, √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ - |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - IsBayesAlgEnvSeq.armMean κ R' a ω|} + |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} ≤ ENNReal.ofReal (2 * K * n * δ) := by let badSet := fun (s : ℕ) (a : Fin K) ↦ {ω : Ω | √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ - |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - IsBayesAlgEnvSeq.armMean κ R' a ω|} + |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} have h_set_eq : {ω | ∃ s < n, ∃ a, √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ - |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - IsBayesAlgEnvSeq.armMean κ R' a ω|} = + |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} = ⋃ s ∈ Finset.range n, ⋃ a : Fin K, badSet s a := by ext ω; simp only [Set.mem_setOf_eq, Finset.mem_range, Set.mem_iUnion, badSet, exists_prop] rw [h_set_eq] @@ -628,7 +630,7 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} have h_set_eq : ⋃ s ∈ Finset.range n, badSet s a = - (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by ext ω simp only [Set.mem_iUnion, Finset.mem_range, badSet, badSetIT, Set.mem_preimage, @@ -636,18 +638,18 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] exact Iff.rfl rw [h_set_eq] have h_meas_pair : - Measurable (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) := - h.measurable_env.prodMk h.measurable_traj - have h_disint : P.map (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) = - P.map (IsBayesAlgEnvSeq.env R') ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P := + Measurable (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) := + h.measurable_E.prodMk h.measurable_traj + have h_disint : P.map (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) = + P.map (E') ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.traj A R') E' P := (compProd_map_condDistrib (h.measurable_traj.aemeasurable)).symm - have h_cond_bound : ∀ᵐ e ∂(P.map (IsBayesAlgEnvSeq.env R')), - (condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P e) + have h_cond_bound : ∀ᵐ e ∂(P.map (E')), + (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) (⋃ s ∈ Finset.range n, badSetIT s e) ≤ ENNReal.ofReal (2 * n * δ) := by filter_upwards [IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h] with e h_isAlgEnvSeq let ν := κ.comap (·, e) (by fun_prop) - let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P e + let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) 1 (ν a') := fun a' ↦ by simp only [ν, Kernel.comap_apply]; exact hs a' e have h_mean' : ∀ a', (κ (a', e))[id] ∈ Set.Icc 0 1 := fun a' ↦ hm a' e @@ -804,22 +806,22 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] exact measurableSet_le (by fun_prop) (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp measurable_snd).sub h_kernel).abs - calc P ((fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + calc P ((fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1}) - = (P.map (fun ω ↦ (IsBayesAlgEnvSeq.env R' ω, IsBayesAlgEnvSeq.traj A R' ω))) + = (P.map (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω))) {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by rw [Measure.map_apply h_meas_pair h_meas_set] - _ = (P.map (IsBayesAlgEnvSeq.env R') ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P) + _ = (P.map (E') ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.traj A R') E' P) {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by rw [h_disint] - _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.traj A R') (IsBayesAlgEnvSeq.env R') P e) - (⋃ s ∈ Finset.range n, badSetIT s e) ∂(P.map (IsBayesAlgEnvSeq.env R')) := by + _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + (⋃ s ∈ Finset.range n, badSetIT s e) ∂(P.map (E')) := by rw [Measure.compProd_apply h_meas_set]; rfl - _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * n * δ) ∂(P.map (IsBayesAlgEnvSeq.env R')) := by + _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * n * δ) ∂(P.map (E')) := by apply lintegral_mono_ae h_cond_bound _ = ENNReal.ofReal (2 * n * δ) := by - rw [lintegral_const, Measure.map_apply h.measurable_env MeasurableSet.univ] + rw [lintegral_const, Measure.map_apply h.measurable_E MeasurableSet.univ] simp [measure_univ] calc P (⋃ a : Fin K, ⋃ s ∈ Finset.range n, badSet s a) ≤ ∑ a : Fin K, P (⋃ s ∈ Finset.range n, badSet s a) := measure_iUnion_fintype_le _ _ @@ -832,24 +834,24 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] congr 1; ring lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) (hδ_large : 1 < 2 * Real.log (1 / δ)) : - IsBayesAlgEnvSeq.bayesRegret κ A R' P n + IsBayesAlgEnvSeq.bayesRegret κ A E' P n ≤ 4 * K * n ^ 2 * δ + 2 * √(8 * Real.log (1 / δ)) * √(↑K * ↑n) := by - let bestArm := IsBayesAlgEnvSeq.bestArm κ R' - let armMean := IsBayesAlgEnvSeq.armMean κ R' + let bestArm := IsBayesAlgEnvSeq.bestArm κ E' + let armMean := IsBayesAlgEnvSeq.armMean κ E' let ucb := ucbIndex A R' δ set Eδ : Set Ω := {ω | ∀ s < n, ∀ a, - |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - armMean a ω| + |empMean A R' a s ω - armMean a ω| < √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ))} have hm_ucb : ∀ a t, Measurable (ucbIndex A R' δ a t) := - fun a t ↦ measurable_ucbIndex hK A R' Q κ P h δ a t - have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.armMean κ R' a) := + fun a t ↦ measurable_ucbIndex hK E' A R' Q κ P h δ a t + have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.armMean κ E' a) := fun a ↦ h.measurable_armMean a - have hm_best : Measurable (IsBayesAlgEnvSeq.bestArm κ R') := h.measurable_bestArm + have hm_best : Measurable (IsBayesAlgEnvSeq.bestArm κ E') := h.measurable_bestArm have h_first_bound : ∀ ω, |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| ≤ n := fun ω ↦ calc |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| @@ -881,7 +883,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (measurable_apply_fin hm_arm (h.measurable_A s))).aestronglyMeasurable · filter_upwards with ω; rw [Real.norm_eq_abs]; exact h_second_bound ω have h_swap : - IsBayesAlgEnvSeq.bayesRegret κ A R' P n = + IsBayesAlgEnvSeq.bayesRegret κ A E' P n = P[fun ω ↦ ∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)] + P[fun ω ↦ ∑ s ∈ range n, @@ -889,10 +891,10 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp have hC : ∀ a e, |(κ (a, e))[id]| ≤ 1 := fun a e ↦ by have := hm a e; rw [abs_le]; exact ⟨by linarith [this.1], this.2⟩ have h_regret_gap := bayesRegret_eq_sum_integral_gap (h := h) (hm := hC) (t := n) - have h_regret_eq : IsBayesAlgEnvSeq.bayesRegret κ A R' P n = + have h_regret_eq : IsBayesAlgEnvSeq.bayesRegret κ A E' P n = ∑ s ∈ range n, P[fun ω ↦ armMean (bestArm ω) ω - armMean (A s ω) ω] := by rw [h_regret_gap]; congr 1 with s - exact integral_congr_ae (ae_of_all _ fun ω ↦ gap_eq_armMean_sub A R' κ hm s ω) + exact integral_congr_ae (ae_of_all _ fun ω ↦ gap_eq_armMean_sub E' A κ hm s ω) have h_int_ucb_sub : ∀ s, Integrable (fun ω ↦ ucb (A s ω) s ω - ucb (bestArm ω) s ω) P := by intro s apply Integrable.sub @@ -912,21 +914,22 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp simp [h_ucb_zero] exact (integral_congr_ae (ae_of_all _ this)).trans (integral_zero _ _) | succ t => - have hts := ts_identity hK A R' Q κ P h t - have h_map_eq : P.map (fun ω ↦ (IsBayesAlgEnvSeq.hist A R' t ω, A (t + 1) ω)) = - P.map (fun ω ↦ (IsBayesAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestArm κ R' ω)) := by + have hts := ts_identity hK E' A R' Q κ P h t + have h_map_eq : P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω)) = + P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestArm κ E' ω)) := by rw [← compProd_map_condDistrib (hY := (h.measurable_A (t + 1)).aemeasurable), ← compProd_map_condDistrib (hY := hm_best.aemeasurable)] exact Measure.compProd_congr hts have h_int_eq : ∀ (f : (Iic t → Fin K × ℝ) × Fin K → ℝ), Measurable f → - ∫ ω, f (IsBayesAlgEnvSeq.hist A R' t ω, A (t + 1) ω) ∂P = - ∫ ω, f (IsBayesAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestArm κ R' ω) ∂P := by + ∫ ω, f (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω) ∂P = + ∫ ω, f (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestArm κ E' ω) ∂P := by intro f hf + have hm_hist := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t rw [← integral_map - ((h.measurable_hist t).prodMk (h.measurable_A (t + 1))).aemeasurable + (hm_hist.prodMk (h.measurable_A (t + 1))).aemeasurable hf.aestronglyMeasurable, ← integral_map - ((h.measurable_hist t).prodMk hm_best).aemeasurable + (hm_hist.prodMk hm_best).aemeasurable hf.aestronglyMeasurable, h_map_eq] set g : (Iic t → Fin K × ℝ) × Fin K → ℝ := @@ -934,15 +937,15 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp √(2 * Real.log (1 / δ) / (max 1 (pullCount' t p.1 p.2) : ℝ)))) have h_hist_eq : ∀ (ω : Ω), (fun (i : Iic t) ↦ (A (↑i) ω, - IsBayesAlgEnvSeq.reward R' (↑i) ω)) = - IsBayesAlgEnvSeq.hist A R' t ω := by + R' (↑i) ω)) = + IsAlgEnvSeq.hist A R' t ω := by intro ω; rfl have hg_eq : ∀ a (ω : Ω), ucbIndex A R' δ a (t + 1) ω = - g (IsBayesAlgEnvSeq.hist A R' t ω, a) := by + g (IsAlgEnvSeq.hist A R' t ω, a) := by intro a ω simp only [g, ucbIndex] - rw [empMean_add_one_eq_empMean' (A := A) (R' := IsBayesAlgEnvSeq.reward R'), - pullCount_add_one_eq_pullCount' (A := A) (R' := IsBayesAlgEnvSeq.reward R'), + rw [empMean_add_one_eq_empMean' (A := A) (R' := R'), + pullCount_add_one_eq_pullCount' (A := A) (R' := R'), h_hist_eq] have hg_meas : Measurable g := by apply Measurable.max measurable_const @@ -957,10 +960,10 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (fun a ↦ measurable_from_top.comp ((measurable_pullCount' t a).comp measurable_fst)) measurable_snd have h_eq_g1 : (fun ω ↦ ucb (A (t + 1) ω) (t + 1) ω) = - fun ω ↦ g (IsBayesAlgEnvSeq.hist A R' t ω, A (t + 1) ω) := + fun ω ↦ g (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω) := funext fun ω ↦ hg_eq _ _ have h_eq_g2 : (fun ω ↦ ucb (bestArm ω) (t + 1) ω) = - fun ω ↦ g (IsBayesAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestArm κ R' ω) := + fun ω ↦ g (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestArm κ E' ω) := funext fun ω ↦ hg_eq _ _ have h_int_ucb : ∀ {f : Ω → Fin K}, Measurable f → Integrable (fun ω ↦ ucb (f ω) (t + 1) ω) P := fun hf ↦ @@ -1017,27 +1020,27 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp intro ω hω apply Finset.sum_nonpos intro s hs - linarith [armMean_le_ucbIndex A R' κ hm δ + linarith [armMean_le_ucbIndex E' A R' κ hm δ (bestArm ω) s ω (hω s (mem_range.mp hs) _)] have h_second_Eδ : ∀ ω ∈ Eδ, ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω) ≤ 2 * √(8 * Real.log (1 / δ)) * √(↑K * ↑n) := by intro ω hω - exact sum_ucbIndex_sub_armMean_le A R' κ hm δ n ω hω + exact sum_ucbIndex_sub_armMean_le E' A R' κ hm δ n ω hω have h_prob : P Eδᶜ ≤ ENNReal.ofReal (2 * K * n * δ) := by have : Eδᶜ = {ω | ∃ s < n, ∃ a, √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ - |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - armMean a ω|} := by + |empMean A R' a s ω - armMean a ω|} := by ext ω; simp only [Eδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push_neg; rfl rw [this] - exact prob_concentration_fail_delta (hK := hK) (A := A) (R' := R') + exact prob_concentration_fail_delta (hK := hK) (E' := E') (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) h hs hm n δ hδ hδ1 hδ_large - have hm_emp : ∀ a s, Measurable (fun ω ↦ empMean A (IsBayesAlgEnvSeq.reward R') a s ω) := - fun a s ↦ measurable_empMean (fun n ↦ h.measurable_A n) (fun n ↦ h.measurable_reward n) a s + have hm_emp : ∀ a s, Measurable (fun ω ↦ empMean A R' a s ω) := + fun a s ↦ measurable_empMean (fun n ↦ h.measurable_A n) (fun n ↦ h.measurable_R n) a s have hm_pc : ∀ a s, Measurable (fun ω ↦ (pullCount A a s ω : ℝ)) := fun a s ↦ measurable_from_top.comp (measurable_pullCount (fun n ↦ h.measurable_A n) a s) have hEδ_meas : MeasurableSet Eδ := by suffices ∀ s a, MeasurableSet {ω | - |empMean A (IsBayesAlgEnvSeq.reward R') a s ω - armMean a ω| + |empMean A R' a s ω - armMean a ω| < √(2 * Real.log (1 / δ) / max 1 ↑(pullCount A a s ω))} by simp only [Eδ, Set.setOf_forall] exact .iInter fun s ↦ .iInter fun _ ↦ .iInter fun a ↦ this s a @@ -1076,16 +1079,16 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp nlinarith lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) (t : ℕ) : - IsBayesAlgEnvSeq.bayesRegret κ A R' P t ≤ 4 * K + 8 * √(K * t * Real.log t) := by + IsBayesAlgEnvSeq.bayesRegret κ A E' P t ≤ 4 * K + 8 * √(K * t * Real.log t) := by by_cases ht : t = 0 · simp [ht, IsBayesAlgEnvSeq.bayesRegret, IsBayesAlgEnvSeq.regret, regret] by_cases ht1_eq : t = 1 · subst ht1_eq simp only [Nat.cast_one, Real.log_one, mul_zero, Real.sqrt_zero, mul_zero, add_zero] - calc IsBayesAlgEnvSeq.bayesRegret κ A R' P 1 + calc IsBayesAlgEnvSeq.bayesRegret κ A E' P 1 ≤ 1 := by unfold IsBayesAlgEnvSeq.bayesRegret IsBayesAlgEnvSeq.regret Bandits.regret simp only [Finset.range_one, Finset.sum_singleton, Nat.cast_one, one_mul, @@ -1093,8 +1096,8 @@ lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] refine (integral_mono_of_nonneg (ae_of_all _ fun ω ↦ sub_nonneg.mpr (le_ciSup ⟨1, by rintro _ ⟨a, rfl⟩; exact (hm a _).2⟩ _)) (integrable_const 1) (ae_of_all _ fun ω ↦ by - linarith [ciSup_le fun a ↦ (hm a (IsBayesAlgEnvSeq.env R' ω)).2, - (hm (A 0 ω) (IsBayesAlgEnvSeq.env R' ω)).1])).trans ?_ + linarith [ciSup_le fun a ↦ (hm a (E' ω)).2, + (hm (A 0 ω) (E' ω)).1])).trans ?_ simp _ ≤ 4 * (K : ℝ) := by nlinarith [show (1 : ℝ) ≤ K from Nat.one_le_cast.mpr (Nat.one_le_of_lt hK)] @@ -1126,10 +1129,10 @@ lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] linarith [Real.log_le_log (by norm_num : (0 : ℝ) < 2) ht2_real] have hK_real_pos : (0 : ℝ) < K := Nat.cast_pos.mpr hK have hKt_nonneg : (0 : ℝ) ≤ ↑K * ↑t := by positivity - calc IsBayesAlgEnvSeq.bayesRegret κ A R' P t + calc IsBayesAlgEnvSeq.bayesRegret κ A E' P t ≤ 4 * ↑K * ↑t ^ 2 * (1 / (↑t) ^ 2) + 2 * √(8 * Real.log (1 / (1 / (↑t) ^ 2))) * √(↑K * ↑t) := - bayesRegret_le_of_delta (hK := hK) (A := A) (R' := R') (Q := Q) + bayesRegret_le_of_delta (hK := hK) (E' := E') (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) h hs hm t (1 / (↑t) ^ 2) hδ hδ1 hδ_large _ = 4 * ↑K + 2 * √(16 * Real.log ↑t) * √(↑K * ↑t) := by rw [h_first, h_log]; ring_nf _ = 4 * ↑K + 8 * (√(Real.log ↑t) * √(↑K * ↑t)) := by diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 8f0fa092..ec9b017a 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -32,190 +32,160 @@ def bayesStationaryEnv variable {Ω : Type*} [mΩ : MeasurableSpace Ω] /-- A Bayesian algorithm-environment sequence: a sequence of actions and observations from an -algorithm that ignores the underlying "environment" while interacting with `bayesStationaryEnv`. -/ -def IsBayesAlgEnvSeq - [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace E] [Nonempty E] - [StandardBorelSpace R] [Nonempty R] +algorithm that ignores the underlying "environment" while interacting with a Bayesian stationary +environment. The environment `E'` is drawn from a prior `Q`, and rewards follow a kernel `κ` +conditioned on the action and environment. -/ +structure IsBayesAlgEnvSeq + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] - (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (alg : Algorithm α R) - (P : Measure Ω) [IsFiniteMeasure P] := - IsAlgEnvSeq A R' (alg.prod_left E) (bayesStationaryEnv Q κ) P + (E' : Ω → E) (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (alg : Algorithm α R) + (P : Measure Ω) [IsFiniteMeasure P] : Prop where + measurable_E : Measurable E' := by fun_prop + measurable_A n : Measurable (A n) := by fun_prop + measurable_R n : Measurable (R' n) := by fun_prop + hasLaw_env : HasLaw E' Q P + hasCondDistrib_action_zero : HasCondDistrib (A 0) E' (Kernel.const _ alg.p0) P + hasCondDistrib_reward_zero : HasCondDistrib (R' 0) (fun ω ↦ (A 0 ω, E' ω)) κ P + hasCondDistrib_action n : + HasCondDistrib (A (n + 1)) (fun ω ↦ (E' ω, IsAlgEnvSeq.hist A R' n ω)) + ((alg.policy n).prodMkLeft _) P + hasCondDistrib_reward n : + HasCondDistrib (R' (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω, E' ω)) + (κ.prodMkLeft _) P namespace IsBayesAlgEnvSeq variable [StandardBorelSpace α] [Nonempty α] -variable [StandardBorelSpace E] [Nonempty E] variable [StandardBorelSpace R] [Nonempty R] variable {Q : Measure E} [IsProbabilityMeasure Q] {κ : Kernel (α × E) R} [IsMarkovKernel κ] -variable {A : ℕ → Ω → α} {R' : ℕ → Ω → E × R} +variable {E' : Ω → E} {A : ℕ → Ω → α} {R' : ℕ → Ω → R} variable {alg : Algorithm α R} variable {P : Measure Ω} [IsProbabilityMeasure P] -/-- The underlying "environment". -/ -def env (R' : ℕ → Ω → E × R) (ω : Ω) : E := (R' 0 ω).1 - -@[fun_prop] -lemma measurable_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) : Measurable (env R') := - (h.measurable_R 0).fst - -/-- The reward at time `n`. -/ -def reward (R' : ℕ → Ω → E × R) (n : ℕ) (ω : Ω) : R := (R' n ω).2 - -@[fun_prop] -lemma measurable_reward (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : Measurable (reward R' n) := - (h.measurable_R n).snd - -/-- The history of actions and rewards up to time `n`. -/ -def hist (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (n : ℕ) (ω : Ω) : Iic n → α × R := - fun i ↦ (A i ω, (R' i ω).2) - -@[fun_prop] -lemma measurable_hist (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : Measurable (hist A R' n) := - measurable_pi_iff.2 fun i => ((h.measurable_A i).prodMk (h.measurable_R i).snd) - -/-- The action at time `n` together with the underlying "environment". -/ -def action_env (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (n : ℕ) (ω : Ω) : α × E := (A n ω, (R' 0 ω).1) +/-- The trajectory of actions and rewards as a function into the IT space. -/ +def traj (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (ω : Ω) : ℕ → α × R := + fun n => (A n ω, R' n ω) @[fun_prop] -lemma measurable_action_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - Measurable (action_env A R' n) := - (h.measurable_A n).prodMk h.measurable_env +lemma measurable_traj (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : Measurable (traj A R') := + measurable_pi_lambda _ fun n => (h.measurable_A n).prodMk (h.measurable_R n) section Real -variable {R' : ℕ → Ω → E × ℝ} +variable {E' : Ω → E} {R' : ℕ → Ω → ℝ} variable {κ : Kernel (α × E) ℝ} [IsMarkovKernel κ] variable {alg : Algorithm α ℝ} /-- The mean of action `a : α` in the underlying "environment". -/ noncomputable -def armMean (κ : Kernel (α × E) ℝ) (R' : ℕ → Ω → E × ℝ) (a : α) (ω : Ω) : ℝ := (κ (a, env R' ω))[id] +def armMean (κ : Kernel (α × E) ℝ) (E' : Ω → E) (a : α) (ω : Ω) : ℝ := (κ (a, E' ω))[id] @[fun_prop] -lemma measurable_armMean (h : IsBayesAlgEnvSeq Q κ A R' alg P) (a : α) : - Measurable (armMean κ R' a) := - stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk h.measurable_env) +lemma measurable_armMean (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (a : α) : + Measurable (armMean κ E' a) := + stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk h.measurable_E) /-- An action with the highest mean in the underlying "environment". -/ noncomputable -def bestArm [Fintype α] [Encodable α] (κ : Kernel (α × E) ℝ) (R' : ℕ → Ω → E × ℝ) := - measurableArgmax (fun ω a ↦ armMean κ R' a ω) +def bestArm [Fintype α] [Encodable α] (κ : Kernel (α × E) ℝ) (E' : Ω → E) := + measurableArgmax (fun ω a ↦ armMean κ E' a ω) @[fun_prop] -lemma measurable_bestArm [Fintype α] [Encodable α] (h : IsBayesAlgEnvSeq Q κ A R' alg P) : - Measurable (bestArm κ R') := +lemma measurable_bestArm [Fintype α] [Encodable α] (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : + Measurable (bestArm κ E') := measurable_measurableArgmax h.measurable_armMean /-- Regret of a sequence of pulls at time `t` considering the underlying "environment". -/ noncomputable -def regret (κ : Kernel (α × E) ℝ) (A : ℕ → Ω → α) (R' : ℕ → Ω → E × ℝ) (t : ℕ) (ω : Ω) : ℝ := - Bandits.regret (κ.comap (·, env R' ω) (by fun_prop)) A t ω +def regret (κ : Kernel (α × E) ℝ) (A : ℕ → Ω → α) (E' : Ω → E) (t : ℕ) (ω : Ω) : ℝ := + Bandits.regret (κ.comap (·, E' ω) (by fun_prop)) A t ω -lemma measurable_regret [Encodable α] (h : IsBayesAlgEnvSeq Q κ A R' alg P) (t : ℕ) : - Measurable (regret κ A R' t) := by +lemma measurable_regret [Encodable α] (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (t : ℕ) : + Measurable (regret κ A E' t) := by apply Measurable.sub · exact Measurable.const_mul (Measurable.iSup h.measurable_armMean) _ · exact Finset.measurable_sum _ fun s _ ↦ - stronglyMeasurable_id.integral_kernel.measurable.comp (h.measurable_action_env s) + stronglyMeasurable_id.integral_kernel.measurable.comp + ((h.measurable_A s).prodMk h.measurable_E) -/-- If `IsBayesAlgEnvSeq Q κ A R' alg P`, then `bayesRegret κ A R' P t` is the expected +/-- If `IsBayesAlgEnvSeq Q κ E' A R' alg P`, then `bayesRegret κ A E' P t` is the expected regret at time `t` of the algorithm `alg` given a prior distribution over "environments" `Q`. -/ noncomputable -def bayesRegret (κ : Kernel (α × E) ℝ) (A : ℕ → Ω → α) (R' : ℕ → Ω → E × ℝ) (P : Measure Ω) +def bayesRegret (κ : Kernel (α × E) ℝ) (A : ℕ → Ω → α) (E' : Ω → E) (P : Measure Ω) (t : ℕ) : ℝ := - P[regret κ A R' t] + P[regret κ A E' t] end Real section Laws -lemma hasLaw_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) : HasLaw (env R') Q P := by - apply HasCondDistrib.hasLaw_of_const - simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst +lemma hasLaw_action_zero (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : + HasLaw (A 0) alg.p0 P := + h.hasCondDistrib_action_zero.hasLaw_of_const + +lemma indepFun_action_zero_env [StandardBorelSpace E] [Nonempty E] + (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : + IndepFun (A 0) E' P := + ((indepFun_iff_condDistrib_eq_const h.measurable_E.aemeasurable + (h.measurable_A 0).aemeasurable).2 (by + rw [h.hasLaw_action_zero.map_eq]; exact h.hasCondDistrib_action_zero.condDistrib_eq)).symm -lemma hasCondDistrib_action' (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - HasCondDistrib (A (n + 1)) (hist A R' n) (alg.policy n) P := +lemma hasCondDistrib_action' (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : + HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) (alg.policy n) P := (h.hasCondDistrib_action n).comp_left (by fun_prop) -lemma hasCondDistrib_reward_zero' (h : IsBayesAlgEnvSeq Q κ A R' alg P) : - HasCondDistrib (reward R' 0) (action_env A R' 0) κ P := by - simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.of_compProd - -lemma hasCondDistrib_reward' (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - HasCondDistrib (reward R' (n + 1)) (action_env A R' (n + 1)) κ P := by - have hr := (h.hasCondDistrib_reward n).snd - simp_rw [bayesStationaryEnv, Kernel.snd_prod] at hr - exact hr.comp_left (by fun_prop) - --- Auxiliary lemma for `condIndepFun_action_env_hist` (Claude) -lemma hasCondDistrib_action_env_hist (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - HasCondDistrib (A (n + 1)) (fun ω ↦ (env R' ω, hist A R' n ω)) - ((alg.policy n).prodMkLeft E) P := by - let f : (Iic n → α × E × R) → E × (Iic n → α × R) := - fun h ↦ ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2)) - suffices h' : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) - (((alg.policy n).comap Prod.snd (by fun_prop)).comap f (by fun_prop)) P from h'.comp_left - exact h.hasCondDistrib_action n - --- Auxiliary lemma for `condIndepFun_reward_hist_action_env` (Claude) -lemma hasCondDistrib_reward_hist_action_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - HasCondDistrib (reward R' (n + 1)) (fun ω ↦ (hist A R' n ω, A (n + 1) ω, env R' ω)) - (κ.prodMkLeft _) P := by - let f : (Iic n → α × E × R) × α → (Iic n → α × R) × α × E := - fun p ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1) - have hf : Measurable f := by fun_prop - suffices h' : HasCondDistrib (reward R' (n + 1)) - (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) - ((κ.comap Prod.snd (by fun_prop)).comap f hf) P from h'.comp_left hf - simpa [bayesStationaryEnv, Kernel.snd_prod] using (h.hasCondDistrib_reward n).snd +lemma hasCondDistrib_reward' (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : + HasCondDistrib (R' (n + 1)) (fun ω ↦ (A (n + 1) ω, E' ω)) κ P := + (h.hasCondDistrib_reward n).comp_left (by fun_prop) end Laws section Independence -lemma indepFun_action_zero_env (h : IsBayesAlgEnvSeq Q κ A R' alg P) : - IndepFun (A 0) (env R') P := by - rw [indepFun_iff_condDistrib_eq_const (h.measurable_A 0).aemeasurable - h.measurable_env.aemeasurable, h.hasLaw_env.map_eq] - simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst.condDistrib_eq - -lemma condIndepFun_action_env_hist [StandardBorelSpace Ω] (h : IsBayesAlgEnvSeq Q κ A R' alg P) - (n : ℕ) : A (n + 1) ⟂ᵢ[hist A R' n, h.measurable_hist n; P] (env R') := +lemma condIndepFun_action_env_hist [StandardBorelSpace E] [Nonempty E] + [StandardBorelSpace Ω] (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) + (n : ℕ) : + A (n + 1) ⟂ᵢ[IsAlgEnvSeq.hist A R' n, + IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n; P] E' := condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft - (h.measurable_env) (h.measurable_A _) (h.measurable_hist n) - (hasCondDistrib_action_env_hist h n).condDistrib_eq + h.measurable_E (h.measurable_A _) (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n) + (h.hasCondDistrib_action n).condDistrib_eq -lemma condIndepFun_reward_hist_action_env [StandardBorelSpace Ω] - (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - reward R' (n + 1) ⟂ᵢ[action_env A R' (n + 1), h.measurable_action_env (n + 1); P] hist A R' n := +lemma condIndepFun_reward_hist [StandardBorelSpace Ω] + (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : + R' (n + 1) ⟂ᵢ[fun ω ↦ (A (n + 1) ω, E' ω), + (h.measurable_A (n + 1)).prodMk h.measurable_E; P] IsAlgEnvSeq.hist A R' n := condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft - (h.measurable_hist n) (h.measurable_reward _) (h.measurable_action_env _) - (hasCondDistrib_reward_hist_action_env h n).condDistrib_eq + (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n) (h.measurable_R _) + ((h.measurable_A _).prodMk h.measurable_E) + (h.hasCondDistrib_reward n).condDistrib_eq end Independence section Posterior +variable [StandardBorelSpace E] [Nonempty E] + /-- The posterior on the environment given history equals Mathlib's `posterior` applied to the likelihood kernel and prior. This is the measure-theoretic formulation of Bayes' rule. -/ -lemma condDistrib_env_hist_eq_posterior (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - condDistrib (env R') (hist A R' n) P - =ᵐ[P.map (hist A R' n)] posterior (condDistrib (hist A R' n) (env R') P) Q := by - -- The key is to show P.map (env, hist) = Q ⊗ₘ condDistrib hist env P - -- Then use compProd_posterior_eq_map_swap and uniqueness of conditional kernels - have h_env_meas : Measurable (env R') := h.measurable_env - have h_hist_meas : Measurable (hist A R' n) := h.measurable_hist n - set κ' := condDistrib (hist A R' n) (env R') P with hκ' - have h_disint : P.map (fun ω => (env R' ω, hist A R' n ω)) = Q ⊗ₘ κ' := by +lemma condDistrib_env_hist_eq_posterior (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : + condDistrib E' (IsAlgEnvSeq.hist A R' n) P + =ᵐ[P.map (IsAlgEnvSeq.hist A R' n)] + posterior (condDistrib (IsAlgEnvSeq.hist A R' n) E' P) Q := by + have h_env_meas : Measurable E' := h.measurable_E + have h_hist_meas : Measurable (IsAlgEnvSeq.hist A R' n) := + IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n + set κ' := condDistrib (IsAlgEnvSeq.hist A R' n) E' P with hκ' + have h_disint : P.map (fun ω => (E' ω, IsAlgEnvSeq.hist A R' n ω)) = Q ⊗ₘ κ' := by rw [← h.hasLaw_env.map_eq, compProd_map_condDistrib (h_hist_meas.aemeasurable)] - have h_marg : P.map (hist A R' n) = κ' ∘ₘ Q := by - have : P.map (hist A R' n) = (P.map (fun ω => (env R' ω, hist A R' n ω))).snd := by + have h_marg : P.map (IsAlgEnvSeq.hist A R' n) = κ' ∘ₘ Q := by + have : P.map (IsAlgEnvSeq.hist A R' n) = + (P.map (fun ω => (E' ω, IsAlgEnvSeq.hist A R' n ω))).snd := by rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)]; rfl rw [this, h_disint, Measure.snd_compProd] rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] - rw [show P.map (fun ω => (hist A R' n ω, env R' ω)) = (Q ⊗ₘ κ').map Prod.swap from by + rw [show P.map (fun ω => (IsAlgEnvSeq.hist A R' n ω, E' ω)) = (Q ⊗ₘ κ').map Prod.swap from by rw [← h_disint, Measure.map_map (by fun_prop) (by fun_prop)]; rfl] rw [← compProd_posterior_eq_map_swap (κ := κ') (μ := Q), h_marg] @@ -223,74 +193,71 @@ end Posterior section StationaryEnvConnection -def traj (A : ℕ → Ω → α) (R' : ℕ → Ω → E × R) (ω : Ω) : ℕ → α × R := - fun n => (A n ω, (R' n ω).2) - -@[fun_prop] -lemma measurable_traj (h : IsBayesAlgEnvSeq Q κ A R' alg P) : Measurable (traj A R') := - measurable_pi_lambda _ fun n => (h.measurable_A n).prodMk (h.measurable_R n).snd +variable [StandardBorelSpace E] [Nonempty E] -omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] +omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] in /-- The traj function commutes with IT projections: IT.action n ∘ traj = A n -/ lemma IT_action_comp_traj (n : ℕ) : IT.action n ∘ traj A R' = A n := by ext ω; simp [IT.action, traj] -omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] +omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] in -/-- The traj function commutes with IT projections: IT.reward n ∘ traj = reward n -/ -lemma IT_reward_comp_traj (n : ℕ) : IT.reward n ∘ traj A R' = reward R' n := by - ext ω; simp [IT.reward, traj, reward] +/-- The traj function commutes with IT projections: IT.reward n ∘ traj = R' n -/ +lemma IT_reward_comp_traj (n : ℕ) : IT.reward n ∘ traj A R' = R' n := by + ext ω; simp [IT.reward, traj] -omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] +omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] in -/-- The traj function commutes with IT projections: IT.hist n ∘ traj = hist n -/ -lemma IT_hist_comp_traj (n : ℕ) : IT.hist n ∘ traj A R' = hist A R' n := by +/-- The traj function commutes with IT projections: IT.hist n ∘ traj = IsAlgEnvSeq.hist n -/ +lemma IT_hist_comp_traj (n : ℕ) : + IT.hist n ∘ traj A R' = IsAlgEnvSeq.hist A R' n := by ext ω i : 2 - simp only [Function.comp_apply, IT.hist, traj, hist] + simp only [Function.comp_apply, IT.hist, traj, IsAlgEnvSeq.hist] -omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] +omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] in /-- The pair (IT.hist n, IT.action (n+1)) commutes with traj. -/ lemma IT_hist_action_comp_traj (n : ℕ) : (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) ∘ traj A R' = - fun ω ↦ (hist A R' n ω, A (n + 1) ω) := by + fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω) := by ext ω : 1 simp only [Function.comp_apply, IT.action, traj, Prod.mk.injEq] exact ⟨funext fun _ => rfl, trivial⟩ -lemma condDistrib_traj_action_zero (h : IsBayesAlgEnvSeq Q κ A R' alg P) : - ∀ᵐ e ∂(P.map (env R')), - (condDistrib (traj A R') (env R') P e).map (IT.action 0) = alg.p0 := by - have h_comp : condDistrib (IT.action 0 ∘ traj A R') (env R') P - =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map (IT.action 0) := - condDistrib_comp (env R') (h.measurable_traj.aemeasurable) (IT.measurable_action 0) +lemma condDistrib_traj_action_zero (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : + ∀ᵐ e ∂(P.map E'), + (condDistrib (traj A R') E' P e).map (IT.action 0) = alg.p0 := by + have h_comp : condDistrib (IT.action 0 ∘ traj A R') E' P + =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map (IT.action 0) := + condDistrib_comp E' (h.measurable_traj.aemeasurable) (IT.measurable_action 0) rw [IT_action_comp_traj] at h_comp filter_upwards [h_comp, condDistrib_of_indepFun h.indepFun_action_zero_env.symm - h.measurable_env.aemeasurable (h.measurable_A 0).aemeasurable] with e h_comp_e h_indep_e + h.measurable_E.aemeasurable (h.measurable_A 0).aemeasurable] with e h_comp_e h_indep_e rw [← Kernel.map_apply _ (IT.measurable_action 0), ← h_comp_e, h_indep_e, Kernel.const_apply] - simp only [h.hasLaw_action_zero.map_eq, Algorithm.prod_left] + exact h.hasLaw_action_zero.map_eq -lemma hasCondDistrib_reward_zero_condDistrib (h : IsBayesAlgEnvSeq Q κ A R' alg P) : - ∀ᵐ e ∂(P.map (env R')), +omit [StandardBorelSpace E] [Nonempty E] in +lemma hasCondDistrib_reward_zero_condDistrib (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : + ∀ᵐ e ∂(P.map E'), HasCondDistrib (IT.reward 0) (IT.action 0) (κ.comap (·, e) (by fun_prop)) - (condDistrib (traj A R') (env R') P e) := by - have h_swap : HasCondDistrib (reward R' 0) (fun ω ↦ (env R' ω, A 0 ω)) + (condDistrib (traj A R') E' P e) := by + have h_swap : HasCondDistrib (R' 0) (fun ω ↦ (E' ω, A 0 ω)) (κ.comap Prod.swap (by fun_prop)) P := by - convert h.hasCondDistrib_reward_zero'.comp_right + convert h.hasCondDistrib_reward_zero.comp_right (MeasurableEquiv.prodComm : α × E ≃ᵐ E × α) using 2 have h_prod := condDistrib_prod_left (h.measurable_A 0).aemeasurable - (h.measurable_reward 0).aemeasurable h.measurable_env.aemeasurable (μ := P) - have h_comp_pair : condDistrib ((fun ω ↦ (IT.action 0 ω, IT.reward 0 ω)) ∘ traj A R') (env R') P - =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map + (h.measurable_R 0).aemeasurable h.measurable_E.aemeasurable (μ := P) + have h_comp_pair : condDistrib ((fun ω ↦ (IT.action 0 ω, IT.reward 0 ω)) ∘ traj A R') E' P + =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map (fun ω ↦ (IT.action 0 ω, IT.reward 0 ω)) := - condDistrib_comp (env R') h.measurable_traj.aemeasurable (by fun_prop) - have h_comp_action : condDistrib (IT.action 0 ∘ traj A R') (env R') P - =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map (IT.action 0) := - condDistrib_comp (env R') h.measurable_traj.aemeasurable (IT.measurable_action 0) + condDistrib_comp E' h.measurable_traj.aemeasurable (by fun_prop) + have h_comp_action : condDistrib (IT.action 0 ∘ traj A R') E' P + =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map (IT.action 0) := + condDistrib_comp E' h.measurable_traj.aemeasurable (IT.measurable_action 0) rw [show (fun ω ↦ (IT.action 0 ω, IT.reward 0 ω)) ∘ traj A R' = - fun ω ↦ (A 0 ω, reward R' 0 ω) from by - ext ω : 1; simp only [Function.comp_apply, IT.action, IT.reward, traj, reward]] at h_comp_pair + fun ω ↦ (A 0 ω, R' 0 ω) from by + ext ω : 1; simp only [Function.comp_apply, IT.action, IT.reward, traj]] at h_comp_pair rw [IT_action_comp_traj] at h_comp_action have h_swap_eq := h_swap.condDistrib_eq rw [(compProd_map_condDistrib (h.measurable_A 0).aemeasurable).symm] at h_swap_eq @@ -307,23 +274,25 @@ lemma hasCondDistrib_reward_zero_condDistrib (h : IsBayesAlgEnvSeq Q κ A R' alg ext s _ rw [Kernel.sectR_apply, Kernel.comap_apply, ha, Kernel.comap_apply]; rfl -lemma hasCondDistrib_action_condDistrib (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - ∀ᵐ e ∂(P.map (env R')), +omit [StandardBorelSpace E] [Nonempty E] in +lemma hasCondDistrib_action_condDistrib (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : + ∀ᵐ e ∂(P.map E'), HasCondDistrib (IT.action (n + 1)) (IT.hist n) (alg.policy n) - (condDistrib (traj A R') (env R') P e) := by - have h_prod := condDistrib_prod_left (h.measurable_hist n).aemeasurable - (h.measurable_A (n + 1)).aemeasurable h.measurable_env.aemeasurable (μ := P) - have h_action_env := (hasCondDistrib_action_env_hist h n).condDistrib_eq + (condDistrib (traj A R') E' P e) := by + have h_hist_meas := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n + have h_prod := condDistrib_prod_left h_hist_meas.aemeasurable + (h.measurable_A (n + 1)).aemeasurable h.measurable_E.aemeasurable (μ := P) + have h_action_env := (h.hasCondDistrib_action n).condDistrib_eq have h_comp_pair : condDistrib ((fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) ∘ traj A R') - (env R') P =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map + E' P =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) := - condDistrib_comp (env R') h.measurable_traj.aemeasurable (by fun_prop) - have h_comp_hist : condDistrib (IT.hist n ∘ traj A R') (env R') P - =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map (IT.hist n) := - condDistrib_comp (env R') h.measurable_traj.aemeasurable (IT.measurable_hist n) + condDistrib_comp E' h.measurable_traj.aemeasurable (by fun_prop) + have h_comp_hist : condDistrib (IT.hist n ∘ traj A R') E' P + =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map (IT.hist n) := + condDistrib_comp E' h.measurable_traj.aemeasurable (IT.measurable_hist n) rw [IT_hist_action_comp_traj] at h_comp_pair rw [IT_hist_comp_traj] at h_comp_hist - rw [(compProd_map_condDistrib (h.measurable_hist n).aemeasurable).symm] at h_action_env + rw [(compProd_map_condDistrib h_hist_meas.aemeasurable).symm] at h_action_env filter_upwards [h_prod, h_comp_pair, h_comp_hist, (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_action_env] with e h_prod_e h_pair_e h_hist_e h_nested_e @@ -337,35 +306,38 @@ lemma hasCondDistrib_action_condDistrib (h : IsBayesAlgEnvSeq Q κ A R' alg P) ( ext s _ rw [Kernel.sectR_apply, ha, Kernel.prodMkLeft_apply] -lemma hasCondDistrib_reward_condDistrib (h : IsBayesAlgEnvSeq Q κ A R' alg P) (n : ℕ) : - ∀ᵐ e ∂(P.map (env R')), +omit [StandardBorelSpace E] [Nonempty E] in +lemma hasCondDistrib_reward_condDistrib (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : + ∀ᵐ e ∂(P.map E'), HasCondDistrib (IT.reward (n + 1)) (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) ((κ.comap (·, e) (by fun_prop)).prodMkLeft _) - (condDistrib (traj A R') (env R') P e) := by + (condDistrib (traj A R') E' P e) := by + have h_hist_meas := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n have h_prod := condDistrib_prod_left - (Measurable.prodMk (h.measurable_hist n) (h.measurable_A (n + 1))).aemeasurable - (h.measurable_reward (n + 1)).aemeasurable h.measurable_env.aemeasurable (μ := P) - have h_swap : HasCondDistrib (reward R' (n + 1)) (fun ω ↦ (env R' ω, hist A R' n ω, A (n + 1) ω)) + (Measurable.prodMk h_hist_meas (h.measurable_A (n + 1))).aemeasurable + (h.measurable_R (n + 1)).aemeasurable h.measurable_E.aemeasurable (μ := P) + have h_swap : HasCondDistrib (R' (n + 1)) + (fun ω ↦ (E' ω, IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) (κ.comap (fun p ↦ (p.2.2, p.1)) (by fun_prop)) P := - (hasCondDistrib_reward_hist_action_env h n).comp_right + (h.hasCondDistrib_reward n).comp_right (MeasurableEquiv.prodAssoc.symm.trans MeasurableEquiv.prodComm) have h_swap_eq := h_swap.condDistrib_eq have h_comp_triple : condDistrib - ((fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), IT.reward (n + 1) ω)) ∘ traj A R') (env R') P - =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map + ((fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), IT.reward (n + 1) ω)) ∘ traj A R') E' P + =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map (fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), IT.reward (n + 1) ω)) := - condDistrib_comp (env R') h.measurable_traj.aemeasurable (by fun_prop) + condDistrib_comp E' h.measurable_traj.aemeasurable (by fun_prop) have h_comp_pair : condDistrib ((fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) ∘ traj A R') - (env R') P =ᵐ[P.map (env R')] (condDistrib (traj A R') (env R') P).map + E' P =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) := - condDistrib_comp (env R') h.measurable_traj.aemeasurable (by fun_prop) + condDistrib_comp E' h.measurable_traj.aemeasurable (by fun_prop) rw [show (fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), IT.reward (n + 1) ω)) ∘ - traj A R' = fun ω ↦ ((hist A R' n ω, A (n + 1) ω), reward R' (n + 1) ω) from by + traj A R' = fun ω ↦ ((IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω), R' (n + 1) ω) from by ext ω : 1 - simp only [Function.comp_apply, IT.action, IT.reward, traj, reward, Prod.mk.injEq] + simp only [Function.comp_apply, IT.action, IT.reward, traj, Prod.mk.injEq] exact ⟨⟨funext fun i => rfl, trivial⟩, trivial⟩] at h_comp_triple rw [IT_hist_action_comp_traj] at h_comp_pair - rw [(compProd_map_condDistrib (Measurable.prodMk (h.measurable_hist n) + rw [(compProd_map_condDistrib (Measurable.prodMk h_hist_meas (h.measurable_A (n + 1))).aemeasurable).symm] at h_swap_eq filter_upwards [h_prod, h_comp_triple, h_comp_pair, (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_swap_eq] @@ -380,11 +352,11 @@ lemma hasCondDistrib_reward_condDistrib (h : IsBayesAlgEnvSeq Q κ A R' alg P) ( ext s _ rw [Kernel.sectR_apply, ha, Kernel.comap_apply, Kernel.prodMkLeft_apply, Kernel.comap_apply] -lemma condDistrib_traj_isAlgEnvSeq (h : IsBayesAlgEnvSeq Q κ A R' alg P) : - ∀ᵐ e ∂(P.map (env R')), +lemma condDistrib_traj_isAlgEnvSeq (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : + ∀ᵐ e ∂(P.map E'), IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv (κ.comap (·, e) (by fun_prop))) - (condDistrib (traj A R') (env R') P e) := by + (condDistrib (traj A R') E' P e) := by filter_upwards [condDistrib_traj_action_zero h, hasCondDistrib_reward_zero_condDistrib h, ae_all_iff.2 (hasCondDistrib_action_condDistrib h), @@ -400,6 +372,54 @@ end StationaryEnvConnection end IsBayesAlgEnvSeq +/-- Bridge theorem: an `IsAlgEnvSeq` for `(alg.prod_left E)` and `(bayesStationaryEnv Q κ)` +gives rise to an `IsBayesAlgEnvSeq`. -/ +theorem IsAlgEnvSeq.toIsBayesAlgEnvSeq + [StandardBorelSpace α] [Nonempty α] + [StandardBorelSpace E] [Nonempty E] + [StandardBorelSpace R] [Nonempty R] + {Q : Measure E} [IsProbabilityMeasure Q] {κ : Kernel (α × E) R} [IsMarkovKernel κ] + {A : ℕ → Ω → α} {R'' : ℕ → Ω → E × R} {alg : Algorithm α R} + {P : Measure Ω} [IsProbabilityMeasure P] + (h : IsAlgEnvSeq A R'' (alg.prod_left E) (bayesStationaryEnv Q κ) P) : + IsBayesAlgEnvSeq Q κ (fun ω ↦ (R'' 0 ω).1) A (fun n ω ↦ (R'' n ω).2) alg P where + measurable_E := (h.measurable_R 0).fst + measurable_A := h.measurable_A + measurable_R n := (h.measurable_R n).snd + hasLaw_env := by + apply HasCondDistrib.hasLaw_of_const + simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst + hasCondDistrib_action_zero := by + have hfst : HasCondDistrib (fun ω ↦ (R'' 0 ω).1) (A 0) (Kernel.const α Q) P := by + simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst + -- E' | A 0 is constant Q = P.map E', so A 0 and E' are independent + have h_indep : IndepFun (A 0) (fun ω ↦ (R'' 0 ω).1) P := by + rw [indepFun_iff_condDistrib_eq_const (h.measurable_A 0).aemeasurable + (h.measurable_R 0).fst.aemeasurable, hfst.hasLaw_of_const.map_eq] + exact hfst.condDistrib_eq + -- From independence: condDistrib (A 0) E' P = const (P.map (A 0)) = const alg.p0 + have hcd := condDistrib_of_indepFun h_indep.symm (h.measurable_R 0).fst.aemeasurable + (h.measurable_A 0).aemeasurable + simp only [h.hasLaw_action_zero.map_eq, Algorithm.prod_left] at hcd + exact ⟨(h.measurable_A 0).aemeasurable, (h.measurable_R 0).fst.aemeasurable, hcd⟩ + hasCondDistrib_reward_zero := by + simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.of_compProd + hasCondDistrib_action n := by + let f : (Iic n → α × E × R) → E × (Iic n → α × R) := + fun h ↦ ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2)) + suffices h' : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R'' n) + (((alg.policy n).comap Prod.snd (by fun_prop)).comap f (by fun_prop)) P from + h'.comp_left (f := f) + exact h.hasCondDistrib_action n + hasCondDistrib_reward n := by + let f : (Iic n → α × E × R) × α → (Iic n → α × R) × α × E := + fun p ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1) + have hf : Measurable f := by fun_prop + suffices h' : HasCondDistrib (fun ω ↦ (R'' (n + 1) ω).2) + (fun ω ↦ (IsAlgEnvSeq.hist A R'' n ω, A (n + 1) ω)) + ((κ.comap Prod.snd (by fun_prop)).comap f hf) P from h'.comp_left hf + simpa [bayesStationaryEnv, Kernel.snd_prod] using (h.hasCondDistrib_reward n).snd + namespace IT /-- Measure on the sequence of actions and observations generated by an algorithm that ignores the @@ -415,15 +435,18 @@ lemma isBayesAlgEnvSeq_bayesianTrajMeasure [StandardBorelSpace E] [Nonempty E] [StandardBorelSpace R] [Nonempty R] (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] - (alg : Algorithm α R) : IsBayesAlgEnvSeq Q κ action reward alg (bayesTrajMeasure Q κ alg) := - isAlgEnvSeq_trajMeasure _ _ + (alg : Algorithm α R) : + IsBayesAlgEnvSeq Q κ (fun ω ↦ (ω 0).2.1) action (fun n ω ↦ (ω n).2.2) + alg (bayesTrajMeasure Q κ alg) := + (isAlgEnvSeq_trajMeasure _ _).toIsBayesAlgEnvSeq /-- The conditional distribution over the best arm given the observed history. -/ noncomputable def posteriorBestArm [StandardBorelSpace α] [Nonempty α] [Fintype α] [Encodable α] (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] (alg : Algorithm α ℝ) (n : ℕ) : Kernel (Iic n → α × ℝ) α := - condDistrib (IsBayesAlgEnvSeq.bestArm κ reward) (IsBayesAlgEnvSeq.hist action reward n) + condDistrib (IsBayesAlgEnvSeq.bestArm κ (fun ω ↦ (ω 0).2.1)) + (IsAlgEnvSeq.hist action (fun n ω ↦ (ω n).2.2) n) (bayesTrajMeasure Q κ alg) deriving IsMarkovKernel @@ -432,7 +455,7 @@ noncomputable def priorBestArm [StandardBorelSpace α] [Nonempty α] [Fintype α] [Encodable α] (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] (alg : Algorithm α ℝ) : Measure α := - (bayesTrajMeasure Q κ alg).map (IsBayesAlgEnvSeq.bestArm κ reward) + (bayesTrajMeasure Q κ alg).map (IsBayesAlgEnvSeq.bestArm κ (fun ω ↦ (ω 0).2.1)) instance [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] [Fintype α] [Encodable α] (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index dc611b6a..25782b27 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -230,23 +230,22 @@ private lemma condDistrib_hist_env_eq_traj (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] {Ω : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] - {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → E × ℝ} + {E' : Ω → E} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} {alg : Algorithm (Fin K) ℝ} {P : Measure Ω} [IsProbabilityMeasure P] - (h : IsBayesAlgEnvSeq Q κ A R' alg P) (t : ℕ) : + (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (t : ℕ) : ∀ᵐ e ∂Q, - condDistrib (IsBayesAlgEnvSeq.hist A R' t) - (IsBayesAlgEnvSeq.env R') P e = + condDistrib (IsAlgEnvSeq.hist A R' t) E' P e = (trajMeasure alg (stationaryEnv (κ.comap (·, e) (by fun_prop)))).map (IT.hist t) := by rw [← h.hasLaw_env.map_eq] have h_comp : condDistrib (IT.hist t ∘ IsBayesAlgEnvSeq.traj A R') - (IsBayesAlgEnvSeq.env R') P - =ᵐ[P.map (IsBayesAlgEnvSeq.env R')] + E' P + =ᵐ[P.map E'] (condDistrib (IsBayesAlgEnvSeq.traj A R') - (IsBayesAlgEnvSeq.env R') P).map (IT.hist t) := - condDistrib_comp (IsBayesAlgEnvSeq.env R') + E' P).map (IT.hist t) := + condDistrib_comp E' h.measurable_traj.aemeasurable (IT.measurable_hist t) - rw [IsBayesAlgEnvSeq.IT_hist_comp_traj (E := E) t] at h_comp + rw [IsBayesAlgEnvSeq.IT_hist_comp_traj t] at h_comp filter_upwards [h_comp, h.condDistrib_traj_isAlgEnvSeq] with e hc he rw [hc, Kernel.map_apply _ (IT.measurable_hist t)] congr 1 @@ -390,7 +389,7 @@ variable {E : Type*} [MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] variable (Q : Measure E) [IsProbabilityMeasure Q] variable (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] variable {Ω : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] -variable {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → E × ℝ} +variable {E' : Ω → E} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} variable {alg : Algorithm (Fin K) ℝ} variable {P : Measure Ω} [IsProbabilityMeasure P] @@ -406,47 +405,46 @@ lemma measurable_envToBestArm : Measurable (envToBestArm κ) := omit [StandardBorelSpace E] [Nonempty E] [IsMarkovKernel κ] [MeasurableSpace Ω] [IsProbabilityMeasure P] [Nonempty Ω] in lemma bestArm_eq_envToBestArm_comp_env : - IsBayesAlgEnvSeq.bestArm κ R' = envToBestArm κ ∘ IsBayesAlgEnvSeq.env R' := by - funext ω - simp only [Function.comp_apply, IsBayesAlgEnvSeq.bestArm, envToBestArm, - IsBayesAlgEnvSeq.env] + IsBayesAlgEnvSeq.bestArm κ E' = envToBestArm κ ∘ E' := by + funext ω; simp only [Function.comp_apply] + unfold IsBayesAlgEnvSeq.bestArm IsBayesAlgEnvSeq.armMean envToBestArm exact (measurableArgmax_eq_of_eq _ _ _ ω).trans (measurableArgmax_congr _ _ ω _ rfl) +omit [StandardBorelSpace E] [Nonempty E] in /-- The marginal on the history equals `condDistrib (hist) (env) P ∘ₘ Q`. -/ private lemma map_hist_eq_condDistrib_comp {Ω' : Type*} [MeasurableSpace Ω'] [StandardBorelSpace Ω'] [Nonempty Ω'] - {A' : ℕ → Ω' → Fin K} {R'' : ℕ → Ω' → E × ℝ} + {E'' : Ω' → E} {A' : ℕ → Ω' → Fin K} {R'' : ℕ → Ω' → ℝ} {alg' : Algorithm (Fin K) ℝ} {P' : Measure Ω'} [IsProbabilityMeasure P'] - (h' : IsBayesAlgEnvSeq Q κ A' R'' alg' P') (t : ℕ) : - P'.map (IsBayesAlgEnvSeq.hist A' R'' t) = - condDistrib (IsBayesAlgEnvSeq.hist A' R'' t) (IsBayesAlgEnvSeq.env R'') P' ∘ₘ Q := by - calc P'.map (IsBayesAlgEnvSeq.hist A' R'' t) - _ = (P'.map (fun ω => (IsBayesAlgEnvSeq.env R'' ω, - IsBayesAlgEnvSeq.hist A' R'' t ω))).snd := - (Measure.snd_map_prodMk h'.measurable_env).symm - _ = (P'.map (IsBayesAlgEnvSeq.env R'') ⊗ₘ condDistrib - (IsBayesAlgEnvSeq.hist A' R'' t) (IsBayesAlgEnvSeq.env R'') P').snd := by - rw [compProd_map_condDistrib (h'.measurable_hist t).aemeasurable] - _ = (Q ⊗ₘ condDistrib (IsBayesAlgEnvSeq.hist A' R'' t) - (IsBayesAlgEnvSeq.env R'') P').snd := by rw [h'.hasLaw_env.map_eq] + (h' : IsBayesAlgEnvSeq Q κ E'' A' R'' alg' P') (t : ℕ) : + P'.map (IsAlgEnvSeq.hist A' R'' t) = + condDistrib (IsAlgEnvSeq.hist A' R'' t) E'' P' ∘ₘ Q := by + calc P'.map (IsAlgEnvSeq.hist A' R'' t) + _ = (P'.map (fun ω => (E'' ω, + IsAlgEnvSeq.hist A' R'' t ω))).snd := + (Measure.snd_map_prodMk h'.measurable_E).symm + _ = (P'.map E'' ⊗ₘ condDistrib + (IsAlgEnvSeq.hist A' R'' t) E'' P').snd := by + rw [compProd_map_condDistrib + (IsAlgEnvSeq.measurable_hist h'.measurable_A h'.measurable_R t).aemeasurable] + _ = (Q ⊗ₘ condDistrib (IsAlgEnvSeq.hist A' R'' t) + E'' P').snd := by rw [h'.hasLaw_env.map_eq] _ = _ := Measure.snd_compProd Q _ /-- The history distribution under any algorithm is absolutely continuous w.r.t. the history distribution under the uniform algorithm (since uniform gives positive probability to every action). -/ lemma absolutelyContinuous_map_hist_uniform - (h : IsBayesAlgEnvSeq Q κ A R' alg P) (hK : 0 < K) + (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (hK : 0 < K) {Ωu : Type*} [MeasurableSpace Ωu] [StandardBorelSpace Ωu] [Nonempty Ωu] - {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → E × ℝ} + {Eu : Ωu → E} {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → ℝ} {Pu : Measure Ωu} [IsProbabilityMeasure Pu] - (hu : IsBayesAlgEnvSeq Q κ Au Ru (Bandits.uniformAlgorithm hK) Pu) + (hu : IsBayesAlgEnvSeq Q κ Eu Au Ru (Bandits.uniformAlgorithm hK) Pu) (t : ℕ) : - P.map (IsBayesAlgEnvSeq.hist A R' t) ≪ - Pu.map (IsBayesAlgEnvSeq.hist Au Ru t) := by - set κ_alg := condDistrib (IsBayesAlgEnvSeq.hist A R' t) - (IsBayesAlgEnvSeq.env R') P - set κ_unif := condDistrib (IsBayesAlgEnvSeq.hist Au Ru t) - (IsBayesAlgEnvSeq.env Ru) Pu + P.map (IsAlgEnvSeq.hist A R' t) ≪ + Pu.map (IsAlgEnvSeq.hist Au Ru t) := by + set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' t) E' P + set κ_unif := condDistrib (IsAlgEnvSeq.hist Au Ru t) Eu Pu rw [map_hist_eq_condDistrib_comp Q κ h t, map_hist_eq_condDistrib_comp Q κ hu t, ← Measure.snd_compProd, ← Measure.snd_compProd] exact (Measure.AbsolutelyContinuous.compProd_right @@ -459,17 +457,17 @@ lemma absolutelyContinuous_map_hist_uniform /-- The posterior on the environment given history is algorithm-independent. -/ lemma condDistrib_env_hist_alg_indep - (h : IsBayesAlgEnvSeq Q κ A R' alg P) (hK : 0 < K) + (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (hK : 0 < K) {Ωu : Type*} [MeasurableSpace Ωu] [StandardBorelSpace Ωu] [Nonempty Ωu] - {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → E × ℝ} + {Eu : Ωu → E} {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → ℝ} {Pu : Measure Ωu} [IsProbabilityMeasure Pu] - (hu : IsBayesAlgEnvSeq Q κ Au Ru (Bandits.uniformAlgorithm hK) Pu) + (hu : IsBayesAlgEnvSeq Q κ Eu Au Ru (Bandits.uniformAlgorithm hK) Pu) (t : ℕ) : - condDistrib (IsBayesAlgEnvSeq.env R') (IsBayesAlgEnvSeq.hist A R' t) P - =ᵐ[P.map (IsBayesAlgEnvSeq.hist A R' t)] - condDistrib (IsBayesAlgEnvSeq.env Ru) (IsBayesAlgEnvSeq.hist Au Ru t) Pu := by - set κ_alg := condDistrib (IsBayesAlgEnvSeq.hist A R' t) (IsBayesAlgEnvSeq.env R') P - set κ_unif := condDistrib (IsBayesAlgEnvSeq.hist Au Ru t) (IsBayesAlgEnvSeq.env Ru) Pu + condDistrib E' (IsAlgEnvSeq.hist A R' t) P + =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] + condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu := by + set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' t) E' P + set κ_unif := condDistrib (IsAlgEnvSeq.hist Au Ru t) Eu Pu obtain ⟨ρ, hρ_meas, hρ_ne_top, hρ⟩ := exists_density_independent_of_env hK alg t -- Key factorization: κ_alg =ᵐ[Q] κ_unif.withDensity (fun _ => ρ) have h_wd_ae : κ_alg =ᵐ[Q] κ_unif.withDensity (fun _ => ρ) := by @@ -483,47 +481,47 @@ lemma condDistrib_env_hist_alg_indep Kernel.IsSFiniteKernel.withDensity _ (fun _ b => hρ_ne_top b) -- Posterior equality via density factorization have h_post : posterior κ_unif Q - =ᵐ[P.map (IsBayesAlgEnvSeq.hist A R' t)] posterior κ_alg Q := by + =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] posterior κ_alg Q := by rw [map_hist_eq_condDistrib_comp Q κ h t] exact posterior_eq_of_withDensity_ae_eq hρ_meas h_wd_ae -- Bayes' rule for both algorithms have h1 := h.condDistrib_env_hist_eq_posterior t - have h2' : condDistrib (IsBayesAlgEnvSeq.env Ru) (IsBayesAlgEnvSeq.hist Au Ru t) Pu - =ᵐ[P.map (IsBayesAlgEnvSeq.hist A R' t)] posterior κ_unif Q := + have h2' : condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu + =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] posterior κ_unif Q := (absolutelyContinuous_map_hist_uniform Q κ h hK hu t).ae_le (hu.condDistrib_env_hist_eq_posterior t) exact h1.trans (h_post.symm.trans h2'.symm) /-- The posterior on the best arm equals the uniform algorithm's posterior. -/ lemma posteriorBestArm_eq_uniform - (h : IsBayesAlgEnvSeq Q κ A R' alg P) (hK : 0 < K) (t : ℕ) : - condDistrib (IsBayesAlgEnvSeq.bestArm κ R') (IsBayesAlgEnvSeq.hist A R' t) P - =ᵐ[P.map (IsBayesAlgEnvSeq.hist A R' t)] + (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (hK : 0 < K) (t : ℕ) : + condDistrib (IsBayesAlgEnvSeq.bestArm κ E') (IsAlgEnvSeq.hist A R' t) P + =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] IT.posteriorBestArm Q κ (Bandits.uniformAlgorithm hK) t := by unfold IT.posteriorBestArm set Pu := IT.bayesTrajMeasure Q κ (Bandits.uniformAlgorithm hK) - set histf := IsBayesAlgEnvSeq.hist A R' t - set histfu := IsBayesAlgEnvSeq.hist IT.action IT.reward t - set envf := IsBayesAlgEnvSeq.env R' - set envfu := IsBayesAlgEnvSeq.env (E := E) IT.reward - set bau := IsBayesAlgEnvSeq.bestArm (Ω := ℕ → Fin K × E × ℝ) κ IT.reward + set histf := IsAlgEnvSeq.hist A R' t + set histfu := IsAlgEnvSeq.hist IT.action (fun n (ω : ℕ → Fin K × E × ℝ) ↦ (ω n).2.2) t + set envf := E' + set envfu : (ℕ → Fin K × E × ℝ) → E := fun ω ↦ (ω 0).2.1 + set bau := IsBayesAlgEnvSeq.bestArm (Ω := ℕ → Fin K × E × ℝ) κ envfu have h_ITu := IT.isBayesAlgEnvSeq_bayesianTrajMeasure Q κ (Bandits.uniformAlgorithm hK) - -- LHS: condDistrib (bestArm κ R') histf P + -- LHS: condDistrib (bestArm κ E') histf P -- =ᵐ (condDistrib envf histf P).map (envToBestArm κ) - have h_comp_alg : condDistrib (IsBayesAlgEnvSeq.bestArm κ R') histf P + have h_comp_alg : condDistrib (IsBayesAlgEnvSeq.bestArm κ E') histf P =ᵐ[P.map histf] (condDistrib envf histf P).map (envToBestArm κ) := by rw [bestArm_eq_envToBestArm_comp_env κ] exact condDistrib_comp (mβ := MeasurableSpace.pi) histf - h.measurable_env.aemeasurable (measurable_envToBestArm κ) + h.measurable_E.aemeasurable (measurable_envToBestArm κ) -- RHS: condDistrib bau histfu Pu -- =ᵐ (condDistrib envfu histfu Pu).map (envToBestArm κ) have h_comp_unif : condDistrib bau histfu Pu =ᵐ[Pu.map histfu] (condDistrib envfu histfu Pu).map (envToBestArm κ) := by - change condDistrib (IsBayesAlgEnvSeq.bestArm κ IT.reward) histfu Pu + change condDistrib (IsBayesAlgEnvSeq.bestArm κ envfu) histfu Pu =ᵐ[Pu.map histfu] (condDistrib envfu histfu Pu).map (envToBestArm κ) rw [bestArm_eq_envToBestArm_comp_env κ] exact condDistrib_comp (mβ := MeasurableSpace.pi) histfu - h_ITu.measurable_env.aemeasurable (measurable_envToBestArm κ) + h_ITu.measurable_E.aemeasurable (measurable_envToBestArm κ) -- Environment posterior independence have h_env_indep := condDistrib_env_hist_alg_indep Q κ h hK h_ITu t -- Map both sides by envToBestArm From 6f0de3229cecc731918738167b4a86df4c06ff90 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 18 Feb 2026 14:07:02 +0000 Subject: [PATCH 039/155] Refactor HistoryDensity --- .../BayesStationaryEnv.lean | 42 ++- .../SequentialLearning/HistoryDensity.lean | 337 +++++++++++------- 2 files changed, 234 insertions(+), 145 deletions(-) diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index ec9b017a..ca0d947f 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -169,25 +169,29 @@ variable [StandardBorelSpace E] [Nonempty E] /-- The posterior on the environment given history equals Mathlib's `posterior` applied to the likelihood kernel and prior. This is the measure-theoretic formulation of Bayes' rule. -/ -lemma condDistrib_env_hist_eq_posterior (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : - condDistrib E' (IsAlgEnvSeq.hist A R' n) P - =ᵐ[P.map (IsAlgEnvSeq.hist A R' n)] - posterior (condDistrib (IsAlgEnvSeq.hist A R' n) E' P) Q := by - have h_env_meas : Measurable E' := h.measurable_E - have h_hist_meas : Measurable (IsAlgEnvSeq.hist A R' n) := - IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n - set κ' := condDistrib (IsAlgEnvSeq.hist A R' n) E' P with hκ' - have h_disint : P.map (fun ω => (E' ω, IsAlgEnvSeq.hist A R' n ω)) = Q ⊗ₘ κ' := by - rw [← h.hasLaw_env.map_eq, compProd_map_condDistrib (h_hist_meas.aemeasurable)] - have h_marg : P.map (IsAlgEnvSeq.hist A R' n) = κ' ∘ₘ Q := by - have : P.map (IsAlgEnvSeq.hist A R' n) = - (P.map (fun ω => (E' ω, IsAlgEnvSeq.hist A R' n ω))).snd := by - rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)]; rfl - rw [this, h_disint, Measure.snd_compProd] - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] - rw [show P.map (fun ω => (IsAlgEnvSeq.hist A R' n ω, E' ω)) = (Q ⊗ₘ κ').map Prod.swap from by - rw [← h_disint, Measure.map_map (by fun_prop) (by fun_prop)]; rfl] - rw [← compProd_posterior_eq_map_swap (κ := κ') (μ := Q), h_marg] +lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : + HasCondDistrib E' (IsAlgEnvSeq.hist A R' n) + (posterior (condDistrib (IsAlgEnvSeq.hist A R' n) E' P) Q) P where + aemeasurable_fst := h.measurable_E.aemeasurable + aemeasurable_snd := + (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable + condDistrib_eq := by + have h_env_meas : Measurable E' := h.measurable_E + have h_hist_meas : Measurable (IsAlgEnvSeq.hist A R' n) := + IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n + set κ' := condDistrib (IsAlgEnvSeq.hist A R' n) E' P with hκ' + have h_disint : P.map (fun ω => (E' ω, IsAlgEnvSeq.hist A R' n ω)) = Q ⊗ₘ κ' := by + rw [← h.hasLaw_env.map_eq, compProd_map_condDistrib (h_hist_meas.aemeasurable)] + have h_marg : P.map (IsAlgEnvSeq.hist A R' n) = κ' ∘ₘ Q := by + have : P.map (IsAlgEnvSeq.hist A R' n) = + (P.map (fun ω => (E' ω, IsAlgEnvSeq.hist A R' n ω))).snd := by + rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)]; rfl + rw [this, h_disint, Measure.snd_compProd] + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ h_env_meas.aemeasurable] + rw [show P.map (fun ω => (IsAlgEnvSeq.hist A R' n ω, E' ω)) = + (Q ⊗ₘ κ').map Prod.swap from by + rw [← h_disint, Measure.map_map (by fun_prop) (by fun_prop)]; rfl] + rw [← compProd_posterior_eq_map_swap (κ := κ') (μ := Q), h_marg] end Posterior diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 25782b27..c3afa83e 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -173,88 +173,76 @@ private lemma absolutelyContinuous_stepKernel_stationary exact Measure.AbsolutelyContinuous.compProd_left (absolutelyContinuous_of_forall_singleton_pos (Bandits.uniformAlgorithm_policy_pos h)) _ --- `compProd` unfolding requires extra heartbeats /-- The history distribution at time `n + 1` decomposes as a compProd of the history at time `n` and the step kernel, composed with `IicSuccProd.symm`. -/ private lemma map_hist_succ_eq_compProd_map - (alg : Algorithm (Fin K) ℝ) (env : Environment (Fin K) ℝ) (n : ℕ) : - (trajMeasure alg env).map (IT.hist (n + 1)) = - ((trajMeasure alg env).map (IT.hist n) ⊗ₘ stepKernel alg env n).map + {Ω : Type*} [MeasurableSpace Ω] + {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} + {alg : Algorithm (Fin K) ℝ} {env : Environment (Fin K) ℝ} + {P : Measure Ω} [IsFiniteMeasure P] + (h : IsAlgEnvSeq A R' alg env P) (n : ℕ) : + P.map (IsAlgEnvSeq.hist A R' (n + 1)) = + (P.map (IsAlgEnvSeq.hist A R' n) ⊗ₘ stepKernel alg env n).map (MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n).symm := by - set P := trajMeasure alg env set e := MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n - have h_func : IT.hist (α := Fin K) (R := ℝ) (n + 1) = e.symm ∘ - (fun (ω : ℕ → Fin K × ℝ) => (IT.hist n ω, IT.step (n + 1) ω)) := by - funext ω - change frestrictLe (n + 1) ω = e.symm (frestrictLe n ω, IT.step (n + 1) ω) - rw [show IT.step (n + 1) ω = ω (n + 1) from Prod.mk.eta] - change frestrictLe (n + 1) ω = e.symm (e (frestrictLe (n + 1) ω)) + have hA := h.measurable_A; have hR := h.measurable_R + have h_func : IsAlgEnvSeq.hist A R' (n + 1) = e.symm ∘ + (fun ω => (IsAlgEnvSeq.hist A R' n ω, IsAlgEnvSeq.step A R' (n + 1) ω)) := by + funext ω; simp only [Function.comp_apply] + change frestrictLe (n + 1) (fun k => IsAlgEnvSeq.step A R' k ω) = + e.symm (frestrictLe n (fun k => IsAlgEnvSeq.step A R' k ω), + IsAlgEnvSeq.step A R' (n + 1) ω) + change frestrictLe (n + 1) (fun k => IsAlgEnvSeq.step A R' k ω) = + e.symm (e (frestrictLe (n + 1) (fun k => IsAlgEnvSeq.step A R' k ω))) rw [e.symm_apply_apply] - rw [h_func, (Measure.map_map e.symm.measurable (by fun_prop : - Measurable (fun (ω : ℕ → Fin K × ℝ) => - (IT.hist n ω, IT.step (n + 1) ω)))).symm] + rw [h_func, (Measure.map_map e.symm.measurable + ((IsAlgEnvSeq.measurable_hist hA hR n).prodMk + (IsAlgEnvSeq.measurable_step (n + 1) (hA _) (hR _)))).symm] congr 1 - have h_cd := (IT.isAlgEnvSeq_trajMeasure alg env).hasCondDistrib_step n + have h_cd := h.hasCondDistrib_step n exact ((condDistrib_ae_eq_iff_measure_eq_compProd _ - (by fun_prop : AEMeasurable (IsAlgEnvSeq.step IT.action IT.reward (n + 1)) P) + (IsAlgEnvSeq.measurable_step (n + 1) (hA _) (hR _)).aemeasurable (stepKernel alg env n)).mp h_cd.condDistrib_eq) /-- The history distribution under any algorithm is absolutely continuous w.r.t. the history distribution under the uniform algorithm, for a stationary environment. -/ private lemma absolutelyContinuous_map_hist_stationary (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) - (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] (t : ℕ) : - (trajMeasure alg (stationaryEnv ν)).map (IT.hist t) ≪ - (trajMeasure (Bandits.uniformAlgorithm hK) (stationaryEnv ν)).map - (IT.hist t) := by + (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] + {Ω₁ : Type*} [MeasurableSpace Ω₁] + {A₁ : ℕ → Ω₁ → Fin K} {R₁ : ℕ → Ω₁ → ℝ} + {P₁ : Measure Ω₁} [IsProbabilityMeasure P₁] + (h₁ : IsAlgEnvSeq A₁ R₁ alg (stationaryEnv ν) P₁) + {Ω₂ : Type*} [MeasurableSpace Ω₂] + {A₂ : ℕ → Ω₂ → Fin K} {R₂ : ℕ → Ω₂ → ℝ} + {P₂ : Measure Ω₂} [IsProbabilityMeasure P₂] + (h₂ : IsAlgEnvSeq A₂ R₂ (Bandits.uniformAlgorithm hK) (stationaryEnv ν) P₂) + (t : ℕ) : + P₁.map (IsAlgEnvSeq.hist A₁ R₁ t) ≪ P₂.map (IsAlgEnvSeq.hist A₂ R₂ t) := by induction t with | zero => - simp only [IT.hist_eq_frestrictLe, trajMeasure, - Kernel.trajMeasure_map_frestrictLe, Kernel.partialTraj_self, - Measure.id_comp, stationaryEnv_ν0] + set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) + have h_hist₁ : IsAlgEnvSeq.hist A₁ R₁ 0 = e.symm ∘ IsAlgEnvSeq.step A₁ R₁ 0 := by + funext ω ⟨i, hi⟩; have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl + have h_hist₂ : IsAlgEnvSeq.hist A₂ R₂ 0 = e.symm ∘ IsAlgEnvSeq.step A₂ R₂ 0 := by + funext ω ⟨i, hi⟩; have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl + rw [h_hist₁, h_hist₂, + ← Measure.map_map e.symm.measurable + (IsAlgEnvSeq.measurable_step 0 (h₁.measurable_A _) (h₁.measurable_R _)), + ← Measure.map_map e.symm.measurable + (IsAlgEnvSeq.measurable_step 0 (h₂.measurable_A _) (h₂.measurable_R _)), + h₁.hasLaw_step_zero.map_eq, h₂.hasLaw_step_zero.map_eq] + simp only [stationaryEnv_ν0] exact (Measure.AbsolutelyContinuous.compProd_left (absolutelyContinuous_of_forall_singleton_pos Bandits.uniformAlgorithm_p0_pos) _).map - (MeasurableEquiv.piUnique _).symm.measurable + e.symm.measurable | succ n ih => - rw [map_hist_succ_eq_compProd_map, map_hist_succ_eq_compProd_map] + rw [map_hist_succ_eq_compProd_map h₁, map_hist_succ_eq_compProd_map h₂] exact (Measure.AbsolutelyContinuous.compProd ih (Filter.Eventually.of_forall fun h => absolutelyContinuous_stepKernel_stationary hK alg ν n h)).map (MeasurableEquiv.IicSuccProd _ n).symm.measurable --- `condDistrib_comp` + `eq_trajMeasure` chain requires extra heartbeats -/-- The conditional distribution of the history given the environment equals the trajectory - measure's history marginal (for Bayesian stationary environments). -/ -private lemma condDistrib_hist_env_eq_traj - {E : Type*} [MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] - (Q : Measure E) [IsProbabilityMeasure Q] - (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] - {Ω : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] - {E' : Ω → E} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} - {alg : Algorithm (Fin K) ℝ} {P : Measure Ω} [IsProbabilityMeasure P] - (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (t : ℕ) : - ∀ᵐ e ∂Q, - condDistrib (IsAlgEnvSeq.hist A R' t) E' P e = - (trajMeasure alg - (stationaryEnv (κ.comap (·, e) (by fun_prop)))).map (IT.hist t) := by - rw [← h.hasLaw_env.map_eq] - have h_comp : condDistrib (IT.hist t ∘ IsBayesAlgEnvSeq.traj A R') - E' P - =ᵐ[P.map E'] - (condDistrib (IsBayesAlgEnvSeq.traj A R') - E' P).map (IT.hist t) := - condDistrib_comp E' - h.measurable_traj.aemeasurable (IT.measurable_hist t) - rw [IsBayesAlgEnvSeq.IT_hist_comp_traj t] at h_comp - filter_upwards [h_comp, h.condDistrib_traj_isAlgEnvSeq] with e hc he - rw [hc, Kernel.map_apply _ (IT.measurable_hist t)] - congr 1 - have h' := eq_trajMeasure_of_isAlgEnvSeq he - have hid : - (fun (ω : ℕ → Fin K × ℝ) n => (IT.action n ω, IT.reward n ω)) = id := by - funext ω n; exact Prod.mk.eta - rw [hid, Measure.map_id] at h'; exact h' - end AbsolutelyContinuousHist section PosteriorEquality @@ -297,79 +285,118 @@ section DensityIndependence variable {K : ℕ} [Nonempty (Fin K)] -/-- The history distribution under any algorithm is a `withDensity` of the history distribution -under the uniform algorithm, with a density that does not depend on the reward kernel `ν`. -This is the key factorization property: the density ratio only involves action probabilities. -/ -private lemma exists_density_independent_of_env - (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) (t : ℕ) : - ∃ ρ : (Iic t → Fin K × ℝ) → ℝ≥0∞, Measurable ρ ∧ (∀ h, ρ h ≠ ⊤) ∧ - ∀ (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν], - (trajMeasure alg (stationaryEnv ν)).map (IT.hist t) = - ((trajMeasure (Bandits.uniformAlgorithm hK) (stationaryEnv ν)).map - (IT.hist t)).withDensity ρ := by +/-- The density of the history distribution under `alg` w.r.t. the uniform algorithm. +This density depends only on the algorithm's action probabilities, not on the reward kernel. -/ +private noncomputable def historyDensity + (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) : + (t : ℕ) → (Iic t → Fin K × ℝ) → ℝ≥0∞ + | 0 => (alg.p0.rnDeriv (Bandits.uniformAlgorithm hK).p0 ∘ Prod.fst) ∘ + MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) + | n + 1 => + let σ : (Iic n → Fin K × ℝ) → (Fin K × ℝ) → ℝ≥0∞ := + fun h ar => Kernel.rnDeriv (alg.policy n) + ((Bandits.uniformAlgorithm hK).policy n) h ar.1 + (historyDensity hK alg n ∘ Prod.fst * Function.uncurry σ) ∘ + MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n + +omit [Nonempty (Fin K)] in +private lemma measurable_historyDensity (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) (t : ℕ) : + Measurable (historyDensity hK alg t) := by + induction t with + | zero => + exact (Measure.measurable_rnDeriv _ _).comp + (measurable_fst.comp (MeasurableEquiv.piUnique _).measurable) + | succ n ih => + exact ((ih.comp measurable_fst).mul + ((Kernel.measurable_rnDeriv _ _).comp + (measurable_fst.prodMk (measurable_fst.comp measurable_snd)))).comp + (MeasurableEquiv.IicSuccProd _ n).measurable + +omit [Nonempty (Fin K)] in +private lemma historyDensity_ne_top (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) (t : ℕ) + (h : Iic t → Fin K × ℝ) : historyDensity hK alg t h ≠ ⊤ := by + induction t with + | zero => exact rnDeriv_ne_top_of_forall_singleton_pos Bandits.uniformAlgorithm_p0_pos _ + | succ n ih => + exact ENNReal.mul_ne_top (ih _) + (kernel_rnDeriv_ne_top_of_forall_singleton_pos + (fun h' a => Bandits.uniformAlgorithm_policy_pos h' a) _ _) + +/-- The history distribution under any algorithm equals the uniform algorithm's history +distribution weighted by `historyDensity`, for any stationary environment. -/ +private lemma map_hist_eq_withDensity_historyDensity + (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) (t : ℕ) + (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] + {Ω₁ : Type*} [MeasurableSpace Ω₁] + {A₁ : ℕ → Ω₁ → Fin K} {R₁ : ℕ → Ω₁ → ℝ} + {P₁ : Measure Ω₁} [IsProbabilityMeasure P₁] + (h₁ : IsAlgEnvSeq A₁ R₁ alg (stationaryEnv ν) P₁) + {Ω₂ : Type*} [MeasurableSpace Ω₂] + {A₂ : ℕ → Ω₂ → Fin K} {R₂ : ℕ → Ω₂ → ℝ} + {P₂ : Measure Ω₂} [IsProbabilityMeasure P₂] + (h₂ : IsAlgEnvSeq A₂ R₂ (Bandits.uniformAlgorithm hK) (stationaryEnv ν) P₂) : + P₁.map (IsAlgEnvSeq.hist A₁ R₁ t) = + (P₂.map (IsAlgEnvSeq.hist A₂ R₂ t)).withDensity (historyDensity hK alg t) := by set unif := Bandits.uniformAlgorithm hK induction t with | zero => set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) - refine ⟨(alg.p0.rnDeriv unif.p0 ∘ Prod.fst) ∘ e, - (Measure.measurable_rnDeriv _ _).comp (measurable_fst.comp e.measurable), - fun h => rnDeriv_ne_top_of_forall_singleton_pos Bandits.uniformAlgorithm_p0_pos _, ?_⟩ - intro ν _ have h_ac : alg.p0 ≪ unif.p0 := absolutelyContinuous_of_forall_singleton_pos Bandits.uniformAlgorithm_p0_pos - simp only [IT.hist_eq_frestrictLe, trajMeasure, - Kernel.trajMeasure_map_frestrictLe, Kernel.partialTraj_self, - Measure.id_comp, stationaryEnv_ν0] + have h_hist₁ : IsAlgEnvSeq.hist A₁ R₁ 0 = e.symm ∘ IsAlgEnvSeq.step A₁ R₁ 0 := by + funext ω ⟨i, hi⟩ + have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl + have h_hist₂ : IsAlgEnvSeq.hist A₂ R₂ 0 = e.symm ∘ IsAlgEnvSeq.step A₂ R₂ 0 := by + funext ω ⟨i, hi⟩ + have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl + rw [h_hist₁, h_hist₂, + ← Measure.map_map e.symm.measurable + (IsAlgEnvSeq.measurable_step 0 (h₁.measurable_A _) (h₁.measurable_R _)), + ← Measure.map_map e.symm.measurable + (IsAlgEnvSeq.measurable_step 0 (h₂.measurable_A _) (h₂.measurable_R _)), + h₁.hasLaw_step_zero.map_eq, h₂.hasLaw_step_zero.map_eq] + simp only [stationaryEnv_ν0] conv_lhs => rw [← Measure.withDensity_rnDeriv_eq _ _ h_ac] rw [withDensity_compProd_left (Measure.measurable_rnDeriv _ _)] exact withDensity_map_equiv_symm ((Measure.measurable_rnDeriv _ _).comp measurable_fst) | succ n ih => - obtain ⟨ρ_n, hρ_n_meas, hρ_n_ne_top, hρ_n⟩ := ih let σ : (Iic n → Fin K × ℝ) → (Fin K × ℝ) → ℝ≥0∞ := fun h ar => Kernel.rnDeriv (alg.policy n) (unif.policy n) h ar.1 - let e := MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n have hσ_meas : Measurable (Function.uncurry σ) := (Kernel.measurable_rnDeriv _ _).comp (measurable_fst.prodMk (measurable_fst.comp measurable_snd)) - refine ⟨(ρ_n ∘ Prod.fst * Function.uncurry σ) ∘ e, ?_, ?_, ?_⟩ - · exact ((hρ_n_meas.comp measurable_fst).mul hσ_meas).comp e.measurable - · intro h - exact ENNReal.mul_ne_top (hρ_n_ne_top _) - (kernel_rnDeriv_ne_top_of_forall_singleton_pos - (fun h' a => Bandits.uniformAlgorithm_policy_pos h' a) _ _) - · intro ν _inst - have h_step : stepKernel alg (stationaryEnv ν) n = - (stepKernel unif (stationaryEnv ν) n).withDensity σ := by - ext h : 1 - rw [Kernel.withDensity_apply _ hσ_meas] - have h_alg : stepKernel alg (stationaryEnv ν) n h = (alg.policy n h) ⊗ₘ ν := by - ext s hs - simp only [stepKernel, stationaryEnv, Kernel.compProd_apply hs, - Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - have h_unif : stepKernel unif (stationaryEnv ν) n h = (unif.policy n h) ⊗ₘ ν := by - ext s hs - simp only [stepKernel, stationaryEnv, Kernel.compProd_apply hs, - Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - have h_wd : ((unif.policy n) h).withDensity - (Kernel.rnDeriv (alg.policy n) (unif.policy n) h) = alg.policy n h := by - rw [← Kernel.withDensity_apply _ (Kernel.measurable_rnDeriv _ _)] - exact Kernel.withDensity_rnDeriv_eq (κ := alg.policy n) (η := unif.policy n) - (absolutelyContinuous_of_forall_singleton_pos (Bandits.uniformAlgorithm_policy_pos h)) - rw [h_alg, h_unif, ← h_wd] - haveI : SFinite ((unif.policy n h).withDensity - (Kernel.rnDeriv (alg.policy n) (unif.policy n) h)) := by - rw [h_wd]; infer_instance - exact withDensity_compProd_left - (Kernel.measurable_rnDeriv (alg.policy n) (unif.policy n)).of_uncurry_left - haveI : IsSFiniteKernel ((stepKernel unif (stationaryEnv ν) n).withDensity σ) := by - rw [← h_step]; infer_instance - rw [map_hist_succ_eq_compProd_map alg (stationaryEnv ν) n, - map_hist_succ_eq_compProd_map unif (stationaryEnv ν) n, - hρ_n ν, h_step, - withDensity_compProd_withDensity hρ_n_meas hσ_meas] - exact withDensity_map_equiv_symm - ((hρ_n_meas.comp measurable_fst).mul hσ_meas) + have h_step : stepKernel alg (stationaryEnv ν) n = + (stepKernel unif (stationaryEnv ν) n).withDensity σ := by + ext h : 1 + rw [Kernel.withDensity_apply _ hσ_meas] + have h_alg : stepKernel alg (stationaryEnv ν) n h = (alg.policy n h) ⊗ₘ ν := by + ext s hs + simp only [stepKernel, stationaryEnv, Kernel.compProd_apply hs, + Measure.compProd_apply hs, Kernel.prodMkLeft_apply] + have h_unif : stepKernel unif (stationaryEnv ν) n h = (unif.policy n h) ⊗ₘ ν := by + ext s hs + simp only [stepKernel, stationaryEnv, Kernel.compProd_apply hs, + Measure.compProd_apply hs, Kernel.prodMkLeft_apply] + have h_wd : ((unif.policy n) h).withDensity + (Kernel.rnDeriv (alg.policy n) (unif.policy n) h) = alg.policy n h := by + rw [← Kernel.withDensity_apply _ (Kernel.measurable_rnDeriv _ _)] + exact Kernel.withDensity_rnDeriv_eq (κ := alg.policy n) (η := unif.policy n) + (absolutelyContinuous_of_forall_singleton_pos (Bandits.uniformAlgorithm_policy_pos h)) + rw [h_alg, h_unif, ← h_wd] + haveI : SFinite ((unif.policy n h).withDensity + (Kernel.rnDeriv (alg.policy n) (unif.policy n) h)) := by + rw [h_wd]; infer_instance + exact withDensity_compProd_left + (Kernel.measurable_rnDeriv (alg.policy n) (unif.policy n)).of_uncurry_left + haveI : IsSFiniteKernel ((stepKernel unif (stationaryEnv ν) n).withDensity σ) := by + rw [← h_step]; infer_instance + rw [map_hist_succ_eq_compProd_map h₁ n, + map_hist_succ_eq_compProd_map h₂ n, + ih, h_step, + withDensity_compProd_withDensity (measurable_historyDensity hK alg n) hσ_meas] + exact withDensity_map_equiv_symm + (((measurable_historyDensity hK alg n).comp measurable_fst).mul hσ_meas) end DensityIndependence @@ -449,10 +476,38 @@ lemma absolutelyContinuous_map_hist_uniform ← Measure.snd_compProd, ← Measure.snd_compProd] exact (Measure.AbsolutelyContinuous.compProd_right (show ∀ᵐ e ∂Q, κ_alg e ≪ κ_unif e from by - filter_upwards [condDistrib_hist_env_eq_traj Q κ h t, - condDistrib_hist_env_eq_traj Q κ hu t] with e he_alg he_unif - rw [he_alg, he_unif] - exact absolutelyContinuous_map_hist_stationary hK alg _ t)).map + have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : + (ℕ → Fin K × ℝ) → (Iic t → Fin K × ℝ)) = IT.hist t := + funext fun ω => funext fun i => Prod.mk.eta + have h_cd₁ : ∀ᵐ e ∂Q, + κ_alg e = (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e).map (IT.hist t) := by + rw [← h.hasLaw_env.map_eq] + have h_comp : condDistrib (IT.hist t ∘ IsBayesAlgEnvSeq.traj A R') E' P + =ᵐ[P.map E'] (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P).map (IT.hist t) := + condDistrib_comp E' h.measurable_traj.aemeasurable (IT.measurable_hist t) + rw [IsBayesAlgEnvSeq.IT_hist_comp_traj] at h_comp + filter_upwards [h_comp] with e he + rw [he, Kernel.map_apply _ (IT.measurable_hist t)] + have h_cd₂ : ∀ᵐ e ∂Q, + κ_unif e = (condDistrib (IsBayesAlgEnvSeq.traj Au Ru) Eu Pu e).map (IT.hist t) := by + rw [← hu.hasLaw_env.map_eq] + have h_comp : condDistrib (IT.hist t ∘ IsBayesAlgEnvSeq.traj Au Ru) Eu Pu + =ᵐ[Pu.map Eu] (condDistrib (IsBayesAlgEnvSeq.traj Au Ru) Eu Pu).map (IT.hist t) := + condDistrib_comp Eu hu.measurable_traj.aemeasurable (IT.measurable_hist t) + rw [IsBayesAlgEnvSeq.IT_hist_comp_traj] at h_comp + filter_upwards [h_comp] with e he + rw [he, Kernel.map_apply _ (IT.measurable_hist t)] + have hae₁ : ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward alg + (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) := by + rw [← h.hasLaw_env.map_eq]; exact h.condDistrib_traj_isAlgEnvSeq + have hae₂ : ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward (Bandits.uniformAlgorithm hK) + (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (condDistrib (IsBayesAlgEnvSeq.traj Au Ru) Eu Pu e) := by + rw [← hu.hasLaw_env.map_eq]; exact hu.condDistrib_traj_isAlgEnvSeq + filter_upwards [h_cd₁, h_cd₂, hae₁, hae₂] with e he₁ he₂ hae₁ hae₂ + rw [he₁, he₂, ← h_IT_hist] + exact absolutelyContinuous_map_hist_stationary hK alg _ hae₁ hae₂ t)).map measurable_snd /-- The posterior on the environment given history is algorithm-independent. -/ @@ -468,15 +523,45 @@ lemma condDistrib_env_hist_alg_indep condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu := by set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' t) E' P set κ_unif := condDistrib (IsAlgEnvSeq.hist Au Ru t) Eu Pu - obtain ⟨ρ, hρ_meas, hρ_ne_top, hρ⟩ := exists_density_independent_of_env hK alg t + set ρ := historyDensity hK alg t + have hρ_meas := measurable_historyDensity hK alg t + have hρ_ne_top := historyDensity_ne_top hK alg t -- Key factorization: κ_alg =ᵐ[Q] κ_unif.withDensity (fun _ => ρ) have h_wd_ae : κ_alg =ᵐ[Q] κ_unif.withDensity (fun _ => ρ) := by - filter_upwards [condDistrib_hist_env_eq_traj Q κ h t, - condDistrib_hist_env_eq_traj Q κ hu t] with e he_alg he_unif + have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : + (ℕ → Fin K × ℝ) → (Iic t → Fin K × ℝ)) = IT.hist t := + funext fun ω => funext fun i => Prod.mk.eta + have h_cd₁ : ∀ᵐ e ∂Q, + κ_alg e = (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e).map (IT.hist t) := by + rw [← h.hasLaw_env.map_eq] + have h_comp : condDistrib (IT.hist t ∘ IsBayesAlgEnvSeq.traj A R') E' P + =ᵐ[P.map E'] (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P).map (IT.hist t) := + condDistrib_comp E' h.measurable_traj.aemeasurable (IT.measurable_hist t) + rw [IsBayesAlgEnvSeq.IT_hist_comp_traj] at h_comp + filter_upwards [h_comp] with e he + rw [he, Kernel.map_apply _ (IT.measurable_hist t)] + have h_cd₂ : ∀ᵐ e ∂Q, + κ_unif e = (condDistrib (IsBayesAlgEnvSeq.traj Au Ru) Eu Pu e).map (IT.hist t) := by + rw [← hu.hasLaw_env.map_eq] + have h_comp : condDistrib (IT.hist t ∘ IsBayesAlgEnvSeq.traj Au Ru) Eu Pu + =ᵐ[Pu.map Eu] (condDistrib (IsBayesAlgEnvSeq.traj Au Ru) Eu Pu).map (IT.hist t) := + condDistrib_comp Eu hu.measurable_traj.aemeasurable (IT.measurable_hist t) + rw [IsBayesAlgEnvSeq.IT_hist_comp_traj] at h_comp + filter_upwards [h_comp] with e he + rw [he, Kernel.map_apply _ (IT.measurable_hist t)] + have hae₁ : ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward alg + (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) := by + rw [← h.hasLaw_env.map_eq]; exact h.condDistrib_traj_isAlgEnvSeq + have hae₂ : ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward (Bandits.uniformAlgorithm hK) + (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (condDistrib (IsBayesAlgEnvSeq.traj Au Ru) Eu Pu e) := by + rw [← hu.hasLaw_env.map_eq]; exact hu.condDistrib_traj_isAlgEnvSeq + filter_upwards [h_cd₁, h_cd₂, hae₁, hae₂] with e he₁ he₂ hae₁ hae₂ rw [Kernel.withDensity_apply _ (show Measurable (Function.uncurry (fun (_ : E) => ρ)) from hρ_meas.comp measurable_snd), - he_alg, he_unif] - exact hρ _ + he₁, he₂, ← h_IT_hist] + exact map_hist_eq_withDensity_historyDensity hK alg t _ hae₁ hae₂ haveI : IsSFiniteKernel (κ_unif.withDensity (fun _ => ρ)) := Kernel.IsSFiniteKernel.withDensity _ (fun _ b => hρ_ne_top b) -- Posterior equality via density factorization @@ -485,11 +570,11 @@ lemma condDistrib_env_hist_alg_indep rw [map_hist_eq_condDistrib_comp Q κ h t] exact posterior_eq_of_withDensity_ae_eq hρ_meas h_wd_ae -- Bayes' rule for both algorithms - have h1 := h.condDistrib_env_hist_eq_posterior t + have h1 := (h.hasCondDistrib_env_hist t).condDistrib_eq have h2' : condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] posterior κ_unif Q := (absolutelyContinuous_map_hist_uniform Q κ h hK hu t).ae_le - (hu.condDistrib_env_hist_eq_posterior t) + (hu.hasCondDistrib_env_hist t).condDistrib_eq exact h1.trans (h_post.symm.trans h2'.symm) /-- The posterior on the best arm equals the uniform algorithm's posterior. -/ From 34a11b344a770ec6a007e087c6006cece5c24c12 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 18 Feb 2026 19:38:36 +0000 Subject: [PATCH 040/155] Generalize TS bound (loose) --- LeanBandits/BanditAlgorithms/TS.lean | 707 ++++++++++++++++----------- 1 file changed, 423 insertions(+), 284 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 993fc636..355d77b3 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -14,7 +14,7 @@ import Mathlib.Analysis.Complex.ExponentialBounds open MeasureTheory ProbabilityTheory Finset Learning -open scoped ENNReal +open scoped ENNReal NNReal namespace Bandits @@ -64,34 +64,43 @@ variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (Fin K × E) ℝ) variable (P : Measure Ω) [IsProbabilityMeasure P] noncomputable -def ucbIndex (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (δ : ℝ) +def ucbIndex (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (σ2 lo hi δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : ℝ := - max 0 (min 1 + if pullCount A a t ω = 0 then hi + else max lo (min hi (empMean A R' a t ω - + √(2 * Real.log (1 / δ) / (max 1 (pullCount A a t ω) : ℝ)))) + + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ)))) omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in -lemma ucbIndex_nonneg (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : 0 ≤ ucbIndex A R' δ a t ω := - le_max_left 0 _ +lemma lo_le_ucbIndex (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : ℕ) (ω : Ω) : + lo ≤ ucbIndex A R' σ2 lo hi δ a t ω := by + unfold ucbIndex; split_ifs <;> [exact hlo; exact le_max_left lo _] omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in -lemma ucbIndex_le_one (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : ucbIndex A R' δ a t ω ≤ 1 := - max_le (by norm_num) (min_le_left 1 _) +lemma ucbIndex_le_hi (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : ℕ) (ω : Ω) : + ucbIndex A R' σ2 lo hi δ a t ω ≤ hi := by + unfold ucbIndex; split_ifs <;> [exact le_refl _; exact max_le hlo (min_le_left hi _)] omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in -lemma ucbIndex_mem_Icc (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : - ucbIndex A R' δ a t ω ∈ Set.Icc 0 1 := - ⟨ucbIndex_nonneg A R' δ a t ω, ucbIndex_le_one A R' δ a t ω⟩ +lemma ucbIndex_mem_Icc (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : ℕ) (ω : Ω) : + ucbIndex A R' σ2 lo hi δ a t ω ∈ Set.Icc lo hi := + ⟨lo_le_ucbIndex A R' σ2 lo hi δ hlo a t ω, ucbIndex_le_hi A R' σ2 lo hi δ hlo a t ω⟩ omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in -lemma norm_ucbIndex_le_one (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : - ‖ucbIndex A R' δ a t ω‖ ≤ 1 := by - rw [Real.norm_eq_abs, abs_of_nonneg (ucbIndex_nonneg A R' δ a t ω)] - exact ucbIndex_le_one A R' δ a t ω +lemma abs_ucbIndex_le (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : ℕ) (ω : Ω) : + |ucbIndex A R' σ2 lo hi δ a t ω| ≤ max |lo| |hi| := by + have hmem := ucbIndex_mem_Icc A R' σ2 lo hi δ hlo a t ω + exact abs_le_max_abs_abs hmem.1 hmem.2 omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in -lemma abs_sub_le_one_of_mem_Icc {x y : ℝ} (hx : x ∈ Set.Icc 0 1) (hy : y ∈ Set.Icc 0 1) : - |x - y| ≤ 1 := by +lemma norm_ucbIndex_le (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : ℕ) (ω : Ω) : + ‖ucbIndex A R' σ2 lo hi δ a t ω‖ ≤ max |lo| |hi| := by + rw [Real.norm_eq_abs]; exact abs_ucbIndex_le A R' σ2 lo hi δ hlo a t ω + +omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in +lemma abs_sub_le_of_mem_Icc {lo hi x y : ℝ} (hx : x ∈ Set.Icc lo hi) + (hy : y ∈ Set.Icc lo hi) : + |x - y| ≤ hi - lo := by rw [abs_le]; constructor <;> linarith [hx.1, hx.2, hy.1, hy.2] lemma sum_sqrt_le {ι : Type*} (s : Finset ι) (c : ι → ℝ) (hc : ∀ i, 0 ≤ c i) : @@ -141,47 +150,50 @@ lemma sum_inv_sqrt_max_one_le (N : ℕ) : @[fun_prop] lemma measurable_ucbIndex [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) - (δ : ℝ) (a : Fin K) (t : ℕ) : - Measurable (ucbIndex A R' δ a t) := by + (σ2 lo hi δ : ℝ) (a : Fin K) (t : ℕ) : + Measurable (ucbIndex A R' σ2 lo hi δ a t) := by unfold ucbIndex - apply Measurable.max measurable_const - apply Measurable.min measurable_const - apply Measurable.add - · exact measurable_empMean (fun n ↦ h.measurable_A n) - (fun n ↦ h.measurable_R n) a t - · have hpc : Measurable (fun ω ↦ (pullCount A a t ω : ℝ)) := - measurable_from_top.comp (measurable_pullCount (fun n ↦ h.measurable_A n) a t) - exact (measurable_const.div (measurable_const.max hpc)).sqrt + have hpc : Measurable (fun ω ↦ (pullCount A a t ω : ℝ)) := + measurable_from_top.comp (measurable_pullCount (fun n ↦ h.measurable_A n) a t) + refine Measurable.ite ?_ measurable_const ?_ + · exact (measurable_pullCount (fun n ↦ h.measurable_A n) a t) (measurableSet_singleton 0) + · exact (Measurable.max measurable_const (Measurable.min measurable_const + (Measurable.add (measurable_empMean (fun n ↦ h.measurable_A n) + (fun n ↦ h.measurable_R n) a t) + (measurable_const.div hpc).sqrt))) omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in -lemma armMean_le_ucbIndex (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) - (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) - (hconc : +lemma armMean_le_ucbIndex {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) + (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) + (hconc : pullCount A a t ω ≠ 0 → |empMean A R' a t ω - IsBayesAlgEnvSeq.armMean κ E' a ω| - < √(2 * Real.log (1 / δ) / (max 1 (pullCount A a t ω) : ℝ))) : - IsBayesAlgEnvSeq.armMean κ E' a ω ≤ ucbIndex A R' δ a t ω := by + < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ))) : + IsBayesAlgEnvSeq.armMean κ E' a ω ≤ ucbIndex A R' σ2 lo hi δ a t ω := by unfold ucbIndex have hmean := hm a (E' ω) simp only [IsBayesAlgEnvSeq.armMean] at hmean hconc ⊢ - have habs := abs_sub_lt_iff.mp hconc - refine le_max_of_le_right (le_min hmean.2 ?_) - linarith [habs.2] + split_ifs with h0 + · exact hmean.2 + · have habs := abs_sub_lt_iff.mp (hconc h0) + refine le_max_of_le_right (le_min hmean.2 ?_) + linarith [habs.2] omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in -lemma ucbIndex_sub_armMean_le (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) - (δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) +lemma ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) + (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hpc : pullCount A a t ω ≠ 0) (hconc : |empMean A R' a t ω - IsBayesAlgEnvSeq.armMean κ E' a ω| - < √(2 * Real.log (1 / δ) / (max 1 (pullCount A a t ω) : ℝ))) : - ucbIndex A R' δ a t ω - IsBayesAlgEnvSeq.armMean κ E' a ω - ≤ 2 * √(2 * Real.log (1 / δ) / (max 1 (pullCount A a t ω) : ℝ)) := by + < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ))) : + ucbIndex A R' σ2 lo hi δ a t ω - IsBayesAlgEnvSeq.armMean κ E' a ω + ≤ 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ)) := by unfold ucbIndex simp only [IsBayesAlgEnvSeq.armMean] at hconc ⊢ - set w := √(2 * Real.log (1 / δ) / max 1 ↑(pullCount A a t ω)) + rw [if_neg hpc] + set w := √(2 * σ2 * Real.log (1 / δ) / ↑(pullCount A a t ω)) set emp := empMean A R' a t ω have habs := abs_sub_lt_iff.mp hconc have hmean := hm a (E' ω) - have h1 : max 0 (min 1 (emp + w)) ≤ emp + w := + have h1 : max lo (min hi (emp + w)) ≤ emp + w := max_le_iff.mpr ⟨by linarith [hmean.1, habs.2], min_le_right _ _⟩ linarith [habs.2] @@ -201,15 +213,17 @@ lemma le_armMean_bestArm [Nonempty (Fin K)] (ω : Ω) (i : Fin K) : simp only [IsBayesAlgEnvSeq.bestArm]; convert this omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in -lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) +lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] {lo hi : ℝ} + (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc lo hi) (ω : Ω) : ⨆ i, IsBayesAlgEnvSeq.armMean κ E' i ω = IsBayesAlgEnvSeq.armMean κ E' (IsBayesAlgEnvSeq.bestArm κ E' ω) ω := le_antisymm (ciSup_le (le_armMean_bestArm E' κ ω)) (le_ciSup (f := fun i ↦ IsBayesAlgEnvSeq.armMean κ E' i ω) - ⟨1, by rintro _ ⟨i, rfl⟩; exact (hm i _).2⟩ _) + ⟨hi, by rintro _ ⟨i, rfl⟩; exact (hm i _).2⟩ _) omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in -lemma gap_eq_armMean_sub [Nonempty (Fin K)] (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) +lemma gap_eq_armMean_sub [Nonempty (Fin K)] {lo hi : ℝ} + (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc lo hi) (s : ℕ) (ω : Ω) : gap (κ.comap (·, E' ω) (by fun_prop)) (A s ω) = IsBayesAlgEnvSeq.armMean κ E' (IsBayesAlgEnvSeq.bestArm κ E' ω) ω - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω := by @@ -268,163 +282,243 @@ lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsProbabilityMeasure Q] [IsMarkovKernel κ] [IsProbabilityMeasure P] in -lemma sum_ucbIndex_sub_armMean_le (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) - (δ : ℝ) (n : ℕ) (ω : Ω) - (hconc : ∀ s < n, ∀ a, +lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) + (hlo : lo ≤ hi) (σ2 δ : ℝ) (n : ℕ) (ω : Ω) + (hconc : ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω| - < √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ))) : - ∑ s ∈ range n, (ucbIndex A R' δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) - ≤ 2 * √(8 * Real.log (1 / δ)) * √(↑K * ↑n) := by - have hterm : ∀ s ∈ range n, - ucbIndex A R' δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω - ≤ 2 * √(2 * Real.log (1 / δ) / (max 1 (pullCount A (A s ω) s ω) : ℝ)) := - fun s hs => ucbIndex_sub_armMean_le E' A R' κ hm δ (A s ω) s ω (hconc s (mem_range.mp hs) _) - calc ∑ s ∈ range n, - (ucbIndex A R' δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) - ≤ ∑ s ∈ range n, - 2 * √(2 * Real.log (1 / δ) / (max 1 (pullCount A (A s ω) s ω) : ℝ)) := - sum_le_sum hterm - _ ≤ 2 * √(8 * Real.log (1 / δ)) * √(↑K * ↑n) := by - set c := Real.log (1 / δ) - by_cases hc : 0 ≤ 2 * c - · open Real in - calc ∑ s ∈ range n, 2 * √(2 * c / max 1 ↑(pullCount A (A s ω) s ω)) - = ∑ s ∈ range n, √(8 * c) * - (1 / √(↑(max 1 (pullCount A (A s ω) s ω)) : ℝ)) := - sum_congr rfl fun s _ => by - rw [show (8 : ℝ) * c = (2 : ℝ) ^ 2 * (2 * c) from by ring] - rw [sqrt_mul (by positivity : (0:ℝ) ≤ 2 ^ 2), - sqrt_sq (by norm_num : (0:ℝ) ≤ 2)] - rw [sqrt_div (by linarith : 0 ≤ 2 * c)]; push_cast; ring - _ = √(8 * c) * ∑ s ∈ range n, - (1 / √(↑(max 1 (pullCount A (A s ω) s ω)) : ℝ)) := by - rw [mul_sum] - _ = √(8 * c) * ∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω), - (1 / √(↑(max 1 j) : ℝ)) := by - congr 1; exact sum_comp_pullCount A (fun j => 1 / √(↑(max 1 j) : ℝ)) n ω - _ ≤ √(8 * c) * ∑ a : Fin K, (2 * √↑(pullCount A a n ω)) := by - gcongr with a; exact sum_inv_sqrt_max_one_le _ - _ = √(8 * c) * (2 * ∑ a : Fin K, √↑(pullCount A a n ω)) := by - simp only [mul_sum] - _ ≤ √(8 * c) * (2 * √(↑K * ↑n)) := by - gcongr - calc ∑ a : Fin K, √↑(pullCount A a n ω) - ≤ √(↑(Finset.univ.card) * ∑ a, ↑(pullCount A a n ω)) := - sum_sqrt_le Finset.univ _ fun a => by positivity - _ = √(↑K * ↑n) := by - congr 1; rw [Finset.card_fin]; congr 1 - have h := sum_pullCount (A := A) (t := n) (ω := ω) - exact_mod_cast h - _ = 2 * √(8 * c) * √(↑K * ↑n) := by ring - · have h0 : ∀ s ∈ range n, - 2 * √(2 * c / max 1 ↑(pullCount A (A s ω) s ω)) = 0 := - fun s _ => by - open Real in - have : 2 * c / max 1 ↑(pullCount A (A s ω) s ω) ≤ 0 := - div_nonpos_of_nonpos_of_nonneg (by linarith) (by positivity) - simp [sqrt_eq_zero'.mpr this] - rw [sum_congr rfl h0]; simp only [sum_const_zero]; positivity + < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ))) : + ∑ s ∈ range n, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - + IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) + ≤ (hi - lo) * ↑K + 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by + -- Split range n into first-pull (pc=0) and non-first-pull (pc≠0) sets + set S0 := (range n).filter (fun s => pullCount A (A s ω) s ω = 0) + set S1 := (range n).filter (fun s => pullCount A (A s ω) s ω ≠ 0) + have hpart : range n = S0 ∪ S1 := (Finset.filter_union_filter_not_eq _ _).symm + have hdisj : Disjoint S0 S1 := Finset.disjoint_filter_filter_not _ _ _ + conv_lhs => rw [hpart] + rw [Finset.sum_union hdisj] + -- We bound ∑_{S0} and ∑_{S1} separately, then combine + suffices h_S0 : ∑ s ∈ S0, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - + IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) ≤ (hi - lo) * ↑K by + suffices h_S1 : ∑ s ∈ S1, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - + IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) + ≤ 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) by + have := Finset.sum_union hdisj (f := fun s => + ucbIndex A R' σ2 lo hi δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) + rw [← hpart] at this; linarith + -- Bound ∑_{S1}: each term ≤ 2√(2σ2c/pc) = 2√(2σ2c/max(1,pc)), so ≤ full sum + calc ∑ s ∈ S1, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - + IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) + ≤ ∑ s ∈ S1, + 2 * √(2 * σ2 * Real.log (1 / δ) / (max 1 (pullCount A (A s ω) s ω) : ℝ)) := + sum_le_sum fun s hs => by + have hpc : pullCount A (A s ω) s ω ≠ 0 := (Finset.mem_filter.mp hs).2 + have hpc_eq : (max 1 (pullCount A (A s ω) s ω) : ℝ) = + (pullCount A (A s ω) s ω : ℝ) := by + simp [Nat.one_le_iff_ne_zero.mpr hpc] + rw [hpc_eq] + exact ucbIndex_sub_armMean_le E' A R' κ hm σ2 δ (A s ω) s ω hpc + (hconc s (mem_range.mp (Finset.mem_filter.mp hs).1) _ hpc) + _ ≤ ∑ s ∈ range n, + 2 * √(2 * σ2 * Real.log (1 / δ) / (max 1 (pullCount A (A s ω) s ω) : ℝ)) := + Finset.sum_le_sum_of_subset_of_nonneg + (Finset.filter_subset _ _) fun s _ _ => by positivity + _ ≤ 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by + set c := Real.log (1 / δ) + by_cases hc : 0 ≤ 2 * σ2 * c + · open Real in + calc ∑ s ∈ range n, 2 * √(2 * σ2 * c / max 1 ↑(pullCount A (A s ω) s ω)) + = ∑ s ∈ range n, √(8 * σ2 * c) * + (1 / √(↑(max 1 (pullCount A (A s ω) s ω)) : ℝ)) := + sum_congr rfl fun s _ => by + rw [show (8 : ℝ) * σ2 * c = (2 : ℝ) ^ 2 * (2 * σ2 * c) from by ring] + rw [sqrt_mul (by positivity : (0:ℝ) ≤ 2 ^ 2), + sqrt_sq (by norm_num : (0:ℝ) ≤ 2)] + rw [sqrt_div (by linarith : 0 ≤ 2 * σ2 * c)]; push_cast; ring + _ = √(8 * σ2 * c) * ∑ s ∈ range n, + (1 / √(↑(max 1 (pullCount A (A s ω) s ω)) : ℝ)) := by + rw [mul_sum] + _ = √(8 * σ2 * c) * ∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω), + (1 / √(↑(max 1 j) : ℝ)) := by + congr 1; exact sum_comp_pullCount A (fun j => 1 / √(↑(max 1 j) : ℝ)) n ω + _ ≤ √(8 * σ2 * c) * ∑ a : Fin K, (2 * √↑(pullCount A a n ω)) := by + gcongr with a; exact sum_inv_sqrt_max_one_le _ + _ = √(8 * σ2 * c) * (2 * ∑ a : Fin K, √↑(pullCount A a n ω)) := by + simp only [mul_sum] + _ ≤ √(8 * σ2 * c) * (2 * √(↑K * ↑n)) := by + gcongr + calc ∑ a : Fin K, √↑(pullCount A a n ω) + ≤ √(↑(Finset.univ.card) * ∑ a, ↑(pullCount A a n ω)) := + sum_sqrt_le Finset.univ _ fun a => by positivity + _ = √(↑K * ↑n) := by + congr 1; rw [Finset.card_fin]; congr 1 + have h := sum_pullCount (A := A) (t := n) (ω := ω) + exact_mod_cast h + _ = 2 * √(8 * σ2 * c) * √(↑K * ↑n) := by ring + · have h0 : ∀ s ∈ range n, + 2 * √(2 * σ2 * c / max 1 ↑(pullCount A (A s ω) s ω)) = 0 := + fun s _ => by + open Real in + have : 2 * σ2 * c / max 1 ↑(pullCount A (A s ω) s ω) ≤ 0 := + div_nonpos_of_nonpos_of_nonneg (by linarith) (by positivity) + simp [sqrt_eq_zero'.mpr this] + rw [sum_congr rfl h0]; simp only [sum_const_zero]; positivity + -- Bound ∑_{S0}: each term = hi - armMean ≤ hi - lo, and #S0 ≤ K + have hterm_S0 : ∀ s ∈ S0, ucbIndex A R' σ2 lo hi δ (A s ω) s ω - + IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω ≤ hi - lo := fun s hs => by + have hpc : pullCount A (A s ω) s ω = 0 := (Finset.mem_filter.mp hs).2 + simp only [ucbIndex, hpc, ↓reduceIte, IsBayesAlgEnvSeq.armMean] + linarith [(hm (A s ω) (E' ω)).1] + have h_card_S0 : #S0 ≤ K := by + calc #S0 ≤ #(Finset.univ : Finset (Fin K)) := + Finset.card_le_card_of_injOn (fun s => A s ω) + (fun _ _ => Finset.mem_coe.mpr (Finset.mem_univ _)) (by + intro s₁ hs₁ s₂ hs₂ heq + have hpc₁ := (Finset.mem_filter.mp (Finset.mem_coe.mp hs₁)).2 + have hpc₂ := (Finset.mem_filter.mp (Finset.mem_coe.mp hs₂)).2 + by_contra h_ne + rcases lt_or_gt_of_ne h_ne with h_lt | h_lt + · have : s₁ ∈ (range s₂).filter (fun i => A i ω = A s₂ ω) := by + simp [mem_range.mpr h_lt, heq] + exact absurd hpc₂ (show pullCount A (A s₂ ω) s₂ ω ≠ 0 from + Finset.card_ne_zero_of_mem this) + · have : s₂ ∈ (range s₁).filter (fun i => A i ω = A s₁ ω) := by + simp [mem_range.mpr h_lt, ← heq] + exact absurd hpc₁ (show pullCount A (A s₁ ω) s₁ ω ≠ 0 from + Finset.card_ne_zero_of_mem this)) + _ = K := Finset.card_fin K + calc ∑ s ∈ S0, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - + IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) + ≤ ∑ _s ∈ S0, (hi - lo) := sum_le_sum hterm_S0 + _ = #S0 * (hi - lo) := by rw [sum_const, nsmul_eq_mul] + _ ≤ ↑K * (hi - lo) := by + apply mul_le_mul_of_nonneg_right _ (by linarith) + exact_mod_cast h_card_S0 + _ = (hi - lo) * ↑K := by ring lemma streamMeasure_concentration_le_delta {α : Type*} [MeasurableSpace α] - {ν : Kernel α ℝ} [IsMarkovKernel ν] - (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + {ν : Kernel α ℝ} [IsMarkovKernel ν] {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (a : α) (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(2 * Real.log (1 / δ) / k) ≤ (ν a)[id]} ≤ + streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + + √(2 * ↑σ2 * Real.log (1 / δ) / k) ≤ (ν a)[id]} ≤ ENNReal.ofReal δ := by have hlog : 0 < Real.log (1 / δ) := Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) have hk_pos : (0 : ℝ) < k := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hk) + have hσ2_pos : (0 : ℝ) < ↑σ2 := NNReal.coe_pos.mpr (pos_iff_ne_zero.mpr hσ2) calc - streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(2 * Real.log (1 / δ) / k) ≤ (ν a)[id]} + streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + + √(2 * ↑σ2 * Real.log (1 / δ) / k) ≤ (ν a)[id]} _ = streamMeasure ν - {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) / k ≤ -√(2 * Real.log (1 / δ) / k)} := by + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) / k ≤ + -√(2 * ↑σ2 * Real.log (1 / δ) / k)} := by congr with ω field_simp rw [Finset.sum_sub_distrib] simp grind _ = streamMeasure ν - {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ -√(2 * k * Real.log (1 / δ))} := by + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ + -√(2 * k * ↑σ2 * Real.log (1 / δ))} := by congr with ω field_simp congr! 2 - rw [Real.sqrt_div (by positivity), ← mul_div_assoc, mul_comm, mul_div_assoc, Real.div_sqrt, - mul_assoc (k : ℝ), Real.sqrt_mul (x := (k : ℝ)) (by positivity), mul_comm] - _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * Real.log (1 / δ)))^2 / (2 * k * 1))) := by + rw [Real.sqrt_div (by positivity : 0 ≤ 2 * ↑σ2 * Real.log (1 / δ)), + show ↑k * 2 * ↑σ2 * Real.log (1 / δ) = ↑k * (2 * ↑σ2 * Real.log (1 / δ)) from by ring, + Real.sqrt_mul (by positivity : (0 : ℝ) ≤ ↑k), ← mul_div_assoc, + mul_div_right_comm, Real.div_sqrt] + _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / + (2 * k * ↑σ2))) := by rw [← ofReal_measureReal] gcongr - refine HasSubgaussianMGF.measure_sum_range_le_le_of_iIndepFun (c := 1) ?_ ?_ (by positivity) + refine HasSubgaussianMGF.measure_sum_range_le_le_of_iIndepFun (c := σ2) ?_ ?_ + (by positivity) · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) (fun _ ↦ by fun_prop) · intro i _; exact (hν a).congr_identDistrib ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) _ = ENNReal.ofReal δ := by rw [Real.sq_sqrt (by positivity)] - simp only [neg_div, Real.exp_neg, mul_one] - rw [mul_div_assoc, mul_div_cancel₀ _ (by positivity : (2 * k : ℝ) ≠ 0), - Real.exp_log (by positivity : (0 : ℝ) < 1 / δ), one_div, inv_inv] + simp only [neg_div, Real.exp_neg] + rw [show 2 * (k : ℝ) * ↑σ2 * Real.log (1 / δ) / (2 * k * ↑σ2) = + Real.log (1 / δ) from by field_simp [ne_of_gt hσ2_pos, ne_of_gt hk_pos]] + rw [Real.exp_log (by positivity : (0 : ℝ) < 1 / δ), one_div, inv_inv] lemma streamMeasure_concentration_ge_delta {α : Type*} [MeasurableSpace α] - {ν : Kernel α ℝ} [IsMarkovKernel ν] - (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + {ν : Kernel α ℝ} [IsMarkovKernel ν] {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (a : α) (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - √(2 * Real.log (1 / δ) / k)} ≤ + streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - + √(2 * ↑σ2 * Real.log (1 / δ) / k)} ≤ ENNReal.ofReal δ := by have hlog : 0 < Real.log (1 / δ) := Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) have hk_pos : (0 : ℝ) < k := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hk) + have hσ2_pos : (0 : ℝ) < ↑σ2 := NNReal.coe_pos.mpr (pos_iff_ne_zero.mpr hσ2) calc - streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - √(2 * Real.log (1 / δ) / k)} + streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - + √(2 * ↑σ2 * Real.log (1 / δ) / k)} _ = streamMeasure ν - {ω | √(2 * Real.log (1 / δ) / k) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id])) / k} := by + {ω | √(2 * ↑σ2 * Real.log (1 / δ) / k) ≤ + (∑ s ∈ range k, (ω s a - (ν a)[id])) / k} := by congr with ω field_simp rw [Finset.sum_sub_distrib] simp grind _ = streamMeasure ν - {ω | √(2 * k * Real.log (1 / δ)) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} := by + {ω | √(2 * k * ↑σ2 * Real.log (1 / δ)) ≤ + (∑ s ∈ range k, (ω s a - (ν a)[id]))} := by congr with ω field_simp congr! 1 - rw [Real.sqrt_div (by positivity), ← mul_div_assoc, mul_comm, mul_div_assoc, Real.div_sqrt] - rw [← Real.sqrt_mul (by positivity : (0 : ℝ) ≤ 2 * Real.log (1 / δ)), mul_comm] - _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * Real.log (1 / δ)))^2 / (2 * k * 1))) := by + rw [Real.sqrt_div (by positivity : 0 ≤ 2 * ↑σ2 * Real.log (1 / δ)), + show 2 * ↑σ2 * Real.log (1 / δ) * ↑k = ↑k * (2 * ↑σ2 * Real.log (1 / δ)) from by ring, + Real.sqrt_mul (by positivity : (0 : ℝ) ≤ ↑k), ← mul_div_assoc, + mul_div_right_comm, Real.div_sqrt] + _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / + (2 * k * ↑σ2))) := by rw [← ofReal_measureReal] gcongr - refine HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := 1) ?_ ?_ (by positivity) + refine HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := σ2) ?_ ?_ + (by positivity) · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) (fun _ ↦ by fun_prop) · intro i _; exact (hν a).congr_identDistrib ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) _ = ENNReal.ofReal δ := by rw [Real.sq_sqrt (by positivity)] - simp only [neg_div, Real.exp_neg, mul_one] - rw [mul_div_assoc, mul_div_cancel₀ _ (by positivity : (2 * k : ℝ) ≠ 0)] - rw [Real.exp_log (by positivity), one_div, inv_inv] + simp only [neg_div, Real.exp_neg] + rw [show 2 * (k : ℝ) * ↑σ2 * Real.log (1 / δ) / (2 * k * ↑σ2) = + Real.log (1 / δ) from by field_simp [ne_of_gt hσ2_pos, ne_of_gt hk_pos]] + rw [Real.exp_log (by positivity : (0 : ℝ) < 1 / δ), one_div, inv_inv] lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) - (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) + {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) + {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc lo hi) (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) - (hδ_large : 1 < 2 * Real.log (1 / δ)) : + (hδ_large : max |lo| |hi| < √(2 * ↑σ2 * Real.log (1 / δ))) : ∀ᵐ e ∂(P.map (E')), (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) - {ω | √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ + {ω | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} ≤ ENNReal.ofReal (2 * s * δ) := by filter_upwards [IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h] with e h_isAlgEnvSeq let ν := κ.comap (·, e) (by fun_prop) - have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) 1 (ν a') := fun a' ↦ by + have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by simp only [ν, Kernel.comap_apply]; exact hs a' e have h_mean : (ν a)[id] = (κ (a, e))[id] := by simp only [ν, Kernel.comap_apply] rw [← h_mean] let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e have h_law := h_isAlgEnvSeq.law_pullCount_sumRewards_unique' (ArrayModel.isAlgEnvSeq_arrayMeasure (tsAlgorithm hK Q κ) ν) (n := s) - let B_low := fun m : ℕ ↦ {x : ℝ | x / m + √(2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} - let B_high := fun m : ℕ ↦ {x : ℝ | (ν a)[id] ≤ x / m - √(2 * Real.log (1 / δ) / m)} + let B_low := fun m : ℕ ↦ + {x : ℝ | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} + let B_high := fun m : ℕ ↦ + {x : ℝ | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)} have h_stream_bound : ∀ m : ℕ, m ≠ 0 → streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ≤ ENNReal.ofReal (2 * δ) := by @@ -439,17 +533,19 @@ lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] _ ≤ ENNReal.ofReal δ + ENNReal.ofReal δ := by gcongr · have h_eq : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m} = - {ω | (∑ i ∈ range m, ω i a) / m + √(2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} := by + {ω | (∑ i ∈ range m, ω i a) / m + + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} := by ext ω; simp only [Set.mem_setOf_eq, B_low] - rw [h_eq]; exact streamMeasure_concentration_le_delta h_subG a m hm0 δ hδ hδ1 + rw [h_eq]; exact streamMeasure_concentration_le_delta hσ2 h_subG a m hm0 δ hδ hδ1 · have h_eq : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_high m} = - {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - √(2 * Real.log (1 / δ) / m)} := by + {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - + √(2 * ↑σ2 * Real.log (1 / δ) / m)} := by ext ω; simp only [Set.mem_setOf_eq, B_high] - rw [h_eq]; exact streamMeasure_concentration_ge_delta h_subG a m hm0 δ hδ hδ1 + rw [h_eq]; exact streamMeasure_concentration_ge_delta hσ2 h_subG a m hm0 δ hδ hδ1 _ = ENNReal.ofReal (2 * δ) := by rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf let badSet := {ω : ℕ → (Fin K) × ℝ | - √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ + √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ |empMean IT.action IT.reward a s ω - (ν a)[id]|} have h_bound_per_m : ∀ m : ℕ, m ≠ 0 → m ≤ s → P' {ω | pullCount IT.action a s ω = m ∧ @@ -476,10 +572,17 @@ lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] exfalso have h_mu := hm a e simp only [max_eq_left (zero_le_one' ℝ), div_one] at hω - rw [h_mean, zero_sub, abs_neg, abs_of_nonneg h_mu.1] at hω - have : 1 < √(2 * Real.log (1 / δ)) := by - rw [Real.lt_sqrt (by norm_num)]; simpa using hδ_large - linarith [h_mu.2] + rw [h_mean, zero_sub, abs_neg] at hω + have h_abs : |(κ (a, e))[id]| ≤ max |lo| |hi| := by + apply abs_le.mpr + constructor + · calc -max |lo| |hi| ≤ -|lo| := neg_le_neg (le_max_left _ _) + _ ≤ lo := neg_abs_le lo + _ ≤ (κ (a, e))[id] := h_mu.1 + · calc (κ (a, e))[id] ≤ hi := h_mu.2 + _ ≤ |hi| := le_abs_self _ + _ ≤ max |lo| |hi| := le_max_right _ _ + linarith · -- Case: m ≥ 1 use m refine ⟨⟨Nat.lt_succ_of_le hms, hm0⟩, rfl, ?_⟩ @@ -540,17 +643,19 @@ lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] lemma prob_concentration_single_delta [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) - (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) + {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) + {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc lo hi) (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) - (hδ_large : 1 < 2 * Real.log (1 / δ)) : - P {ω | √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ + (hδ_large : max |lo| |hi| < √(2 * ↑σ2 * Real.log (1 / δ))) : + P {ω | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} ≤ ENNReal.ofReal (2 * s * δ) := by let badSet : E → Set (ℕ → (Fin K) × ℝ) := fun e ↦ - {t | √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s t) : ℝ)) ≤ + {t | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s t) : ℝ)) ≤ |empMean IT.action IT.reward a s t - (κ (a, e))[id]|} - have h_set_eq : {ω | √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ + have h_set_eq : {ω | √(2 * ↑σ2 * Real.log (1 / δ) / + (max 1 (pullCount A a s ω) : ℝ)) ≤ |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} = (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' {p | p.2 ∈ badSet p.1} := by @@ -569,12 +674,13 @@ lemma prob_concentration_single_delta [Nonempty (Fin K)] P.map (E') ⊗ₘ condDistrib (IsBayesAlgEnvSeq.traj A R') E' P := (compProd_map_condDistrib (h.measurable_traj.aemeasurable)).symm - have h_cond := prob_concentration_single_delta_cond hK E' A R' Q κ P h hs hm a s δ hδ hδ1 hδ_large + have h_cond := prob_concentration_single_delta_cond hK E' A R' Q κ P h hσ2 hs hm a s δ hδ hδ1 + hδ_large have h_kernel : Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (a, p.1))[id]) := stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk measurable_fst) have h_meas_set : MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | p.2 ∈ badSet p.1} := by change MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | - √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s p.2) : ℝ)) ≤ + √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s p.2) : ℝ)) ≤ |empMean IT.action IT.reward a s p.2 - (κ (a, p.1))[id]|} exact measurableSet_le (by fun_prop) (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp measurable_snd).sub @@ -599,18 +705,18 @@ lemma prob_concentration_single_delta [Nonempty (Fin K)] lemma prob_concentration_fail_delta [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) - (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc 0 1) - (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) - (hδ_large : 1 < 2 * Real.log (1 / δ)) : - P {ω | ∃ s < n, ∃ a, - √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ + {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) + (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : + P {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} ≤ ENNReal.ofReal (2 * K * n * δ) := by - let badSet := fun (s : ℕ) (a : Fin K) ↦ {ω : Ω | - √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ + let badSet := fun (s : ℕ) (a : Fin K) ↦ {ω : Ω | pullCount A a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} - have h_set_eq : {ω | ∃ s < n, ∃ a, √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ + have h_set_eq : {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} = ⋃ s ∈ Finset.range n, ⋃ a : Fin K, badSet s a := by ext ω; simp only [Set.mem_setOf_eq, Finset.mem_range, Set.mem_iUnion, badSet, exists_prop] @@ -627,8 +733,9 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] · simp [hn] have hn' : 0 < n := Nat.pos_of_ne_zero hn let badSetIT := fun (s : ℕ) (e : E) ↦ {ω : ℕ → (Fin K) × ℝ | - √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} + pullCount IT.action a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ + |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} have h_set_eq : ⋃ s ∈ Finset.range n, badSet s a = (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by @@ -650,14 +757,13 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] filter_upwards [IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h] with e h_isAlgEnvSeq let ν := κ.comap (·, e) (by fun_prop) let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e - have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) 1 (ν a') := fun a' ↦ by + have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by simp only [ν, Kernel.comap_apply]; exact hs a' e - have h_mean' : ∀ a', (κ (a', e))[id] ∈ Set.Icc 0 1 := fun a' ↦ hm a' e have h_mean : (ν a)[id] = (κ (a, e))[id] := by simp only [ν, Kernel.comap_apply] let B_low := fun m : ℕ ↦ - {x : ℝ | x / m + √(2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} + {x : ℝ | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} let B_high := fun m : ℕ ↦ - {x : ℝ | (ν a)[id] ≤ x / m - √(2 * Real.log (1 / δ) / m)} + {x : ℝ | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)} have h_stream_bound : ∀ m : ℕ, m ≠ 0 → streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ≤ ENNReal.ofReal (2 * δ) := by @@ -672,14 +778,20 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] exact (measure_mono h_union).trans (measure_union_le _ _) _ ≤ ENNReal.ofReal δ + ENNReal.ofReal δ := by gcongr - · have h_eq : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m} = - {ω | (∑ i ∈ range m, ω i a) / m + √(2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} := by + · have h_eq : {ω : ℕ → Fin K → ℝ | + ∑ i ∈ range m, ω i a ∈ B_low m} = + {ω | (∑ i ∈ range m, ω i a) / m + + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} := by ext ω; simp only [Set.mem_setOf_eq, B_low] - rw [h_eq]; exact streamMeasure_concentration_le_delta h_subG a m hm0 δ hδ hδ1 - · have h_eq : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_high m} = - {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - √(2 * Real.log (1 / δ) / m)} := by + rw [h_eq] + exact streamMeasure_concentration_le_delta hσ2 h_subG a m hm0 δ hδ hδ1 + · have h_eq : {ω : ℕ → Fin K → ℝ | + ∑ i ∈ range m, ω i a ∈ B_high m} = + {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - + √(2 * ↑σ2 * Real.log (1 / δ) / m)} := by ext ω; simp only [Set.mem_setOf_eq, B_high] - rw [h_eq]; exact streamMeasure_concentration_ge_delta h_subG a m hm0 δ hδ hδ1 + rw [h_eq] + exact streamMeasure_concentration_ge_delta hσ2 h_subG a m hm0 δ hδ hδ1 _ = ENNReal.ofReal (2 * δ) := by rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf have hB_meas : ∀ m, MeasurableSet (B_low m ∪ B_high m) := fun m ↦ @@ -696,15 +808,7 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] constructor · rintro ⟨s, hs, hbad⟩ let m := pullCount IT.action a s ω - have hm_pos : 0 < m := by - by_contra hm0; push_neg at hm0 - have hm0' : pullCount IT.action a s ω = 0 := by omega - simp only [hm0', Nat.cast_zero, max_eq_left (zero_le_one' ℝ), div_one, - empMean, div_zero] at hbad - rw [zero_sub, abs_neg, abs_of_nonneg (h_mean' a).1] at hbad - have : 1 < √(2 * Real.log (1 / δ)) := by - rw [Real.lt_sqrt (by norm_num)]; simpa using hδ_large - linarith [(h_mean' a).2] + have hm_pos : 0 < m := Nat.pos_of_ne_zero hbad.1 have hm_le : m ≤ n - 1 := by have h1 : m ≤ s := pullCount_le (A := IT.action) a s ω omega @@ -712,9 +816,7 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] have h_pc_pos : (0 : ℝ) < m := Nat.cast_pos.mpr hm_pos simp only [empMean] at hbad - rw [show (max 1 (pullCount IT.action a s ω) : ℝ) = m by - simp only [m]; rw [max_eq_right]; exact Nat.one_le_cast.mpr hm_pos] at hbad - have h_abs := le_abs'.mp hbad + have h_abs := le_abs'.mp hbad.2 rcases h_abs with h_neg | h_pos · left; linarith · right; linarith @@ -723,16 +825,16 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] at hB have h_pc_pos : (0 : ℝ) < m := Nat.cast_pos.mpr hm_pos simp only [empMean, hpc] - rw [show (max 1 m : ℝ) = m by rw [max_eq_right]; exact Nat.one_le_cast.mpr hm_pos] + refine ⟨Nat.one_le_iff_ne_zero.mp hm_pos, ?_⟩ cases hB with | inl h => have h1 : sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id] ≤ - -√(2 * Real.log (1 / δ) / m) := by linarith - calc √(2 * Real.log (1 / δ) / m) + -√(2 * ↑σ2 * Real.log (1 / δ) / m) := by linarith + calc √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ -(sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id]) := by linarith _ ≤ |sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id]| := neg_le_abs _ | inr h => - have h1 : √(2 * Real.log (1 / δ) / m) ≤ + have h1 : √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id] := by linarith exact h1.trans (le_abs_self _) rw [h_decomp] @@ -801,11 +903,15 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] exact .biUnion (Finset.range n).countable_toSet fun s _ ↦ by simp only [badSetIT] change MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | - √(2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s p.2) : ℝ)) ≤ - |empMean IT.action IT.reward a s p.2 - (κ (a, p.1))[id]|} - exact measurableSet_le (by fun_prop) - (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp - measurable_snd).sub h_kernel).abs + pullCount IT.action a s p.2 ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s p.2 : ℝ)) ≤ + |empMean IT.action IT.reward a s p.2 - (κ (a, p.1))[id]|} + exact MeasurableSet.inter + (((measurable_pullCount IT.measurable_action a s).comp measurable_snd) + (measurableSet_singleton (0 : ℕ)).compl) + (measurableSet_le (by fun_prop) + (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp + measurable_snd).sub h_kernel).abs) calc P ((fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1}) = (P.map (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω))) @@ -835,49 +941,54 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) - (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) - (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) - (hδ_large : 1 < 2 * Real.log (1 / δ)) : + {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) + {lo hi : ℝ} (hlo : lo ≤ hi) (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) + (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : IsBayesAlgEnvSeq.bayesRegret κ A E' P n - ≤ 4 * K * n ^ 2 * δ + 2 * √(8 * Real.log (1 / δ)) * √(↑K * ↑n) := by + ≤ (hi - lo) * ↑K + 4 * (hi - lo) * K * n ^ 2 * δ + + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by let bestArm := IsBayesAlgEnvSeq.bestArm κ E' let armMean := IsBayesAlgEnvSeq.armMean κ E' - let ucb := ucbIndex A R' δ - set Eδ : Set Ω := {ω | ∀ s < n, ∀ a, + let ucb := ucbIndex A R' (↑σ2) lo hi δ + set Eδ : Set Ω := {ω | ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → |empMean A R' a s ω - armMean a ω| - < √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ))} - have hm_ucb : ∀ a t, Measurable (ucbIndex A R' δ a t) := - fun a t ↦ measurable_ucbIndex hK E' A R' Q κ P h δ a t + < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ))} + have hm_ucb : ∀ a t, Measurable (ucbIndex A R' (↑σ2) lo hi δ a t) := + fun a t ↦ measurable_ucbIndex hK E' A R' Q κ P h (↑σ2) lo hi δ a t have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.armMean κ E' a) := fun a ↦ h.measurable_armMean a have hm_best : Measurable (IsBayesAlgEnvSeq.bestArm κ E') := h.measurable_bestArm have h_first_bound : ∀ ω, - |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| ≤ n := fun ω ↦ + |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| + ≤ n * (hi - lo) := fun ω ↦ calc |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| ≤ ∑ s ∈ range n, |armMean (bestArm ω) ω - ucb (bestArm ω) s ω| := Finset.abs_sum_le_sum_abs _ _ - _ ≤ ∑ s ∈ range n, (1 : ℝ) := Finset.sum_le_sum fun s _ ↦ - abs_sub_le_one_of_mem_Icc (hm _ _) (ucbIndex_mem_Icc A R' δ _ _ _) - _ = ↑n := by simp + _ ≤ ∑ s ∈ range n, (hi - lo) := Finset.sum_le_sum fun s _ ↦ + abs_sub_le_of_mem_Icc (hm _ _) (ucbIndex_mem_Icc A R' (↑σ2) lo hi δ hlo _ _ _) + _ = ↑n * (hi - lo) := by + rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] have h_second_bound : ∀ ω, - |∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)| ≤ n := fun ω ↦ + |∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)| + ≤ n * (hi - lo) := fun ω ↦ calc |∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)| ≤ ∑ s ∈ range n, |ucb (A s ω) s ω - armMean (A s ω) ω| := Finset.abs_sum_le_sum_abs _ _ - _ ≤ ∑ s ∈ range n, (1 : ℝ) := Finset.sum_le_sum fun s _ ↦ - abs_sub_le_one_of_mem_Icc (ucbIndex_mem_Icc A R' δ _ _ _) (hm _ _) - _ = ↑n := by simp + _ ≤ ∑ s ∈ range n, (hi - lo) := Finset.sum_le_sum fun s _ ↦ + abs_sub_le_of_mem_Icc (ucbIndex_mem_Icc A R' (↑σ2) lo hi δ hlo _ _ _) (hm _ _) + _ = ↑n * (hi - lo) := by + rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] have h_int_sum1 : Integrable (fun ω ↦ ∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)) P := by - apply Integrable.of_bound (C := (↑n)) + apply Integrable.of_bound (C := ↑n * (hi - lo)) · exact (Finset.measurable_fun_sum _ fun s _ ↦ (measurable_apply_fin hm_arm hm_best).sub (measurable_apply_fin (fun a ↦ hm_ucb a s) hm_best)).aestronglyMeasurable · filter_upwards with ω; rw [Real.norm_eq_abs]; exact h_first_bound ω have h_int_sum2 : Integrable (fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)) P := by - apply Integrable.of_bound (C := (↑n)) + apply Integrable.of_bound (C := ↑n * (hi - lo)) · exact (Finset.measurable_fun_sum _ fun s _ ↦ (measurable_apply_fin (fun a ↦ hm_ucb a s) (h.measurable_A s)).sub (measurable_apply_fin hm_arm (h.measurable_A s))).aestronglyMeasurable @@ -888,8 +999,8 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)] + P[fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)] := by - have hC : ∀ a e, |(κ (a, e))[id]| ≤ 1 := fun a e ↦ by - have := hm a e; rw [abs_le]; exact ⟨by linarith [this.1], this.2⟩ + have hC : ∀ a e, |(κ (a, e))[id]| ≤ max |lo| |hi| := fun a e ↦ + abs_le_max_abs_abs (hm a e).1 (hm a e).2 have h_regret_gap := bayesRegret_eq_sum_integral_gap (h := h) (hm := hC) (t := n) have h_regret_eq : IsBayesAlgEnvSeq.bayesRegret κ A E' P n = ∑ s ∈ range n, P[fun ω ↦ armMean (bestArm ω) ω - armMean (A s ω) ω] := by @@ -899,18 +1010,19 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp intro s apply Integrable.sub · exact ⟨(measurable_apply_fin (fun a ↦ hm_ucb a s) (h.measurable_A s)).aestronglyMeasurable, - HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ norm_ucbIndex_le_one A R' _ _ _ _)⟩ + HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ + norm_ucbIndex_le A R' (↑σ2) lo hi δ hlo _ _ _)⟩ · exact ⟨(measurable_apply_fin (fun a ↦ hm_ucb a s) hm_best).aestronglyMeasurable, - HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ norm_ucbIndex_le_one A R' _ _ _ _)⟩ - have h_ucb_zero : ∀ a (ω : Ω), ucbIndex A R' δ a 0 ω = - max 0 (min 1 (√(2 * Real.log (1 / δ)))) := by - intro a ω; unfold ucbIndex empMean sumRewards; simp [pullCount_zero] + HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ + norm_ucbIndex_le A R' (↑σ2) lo hi δ hlo _ _ _)⟩ + have h_ucb_zero : ∀ a (ω : Ω), ucbIndex A R' (↑σ2) lo hi δ a 0 ω = hi := by + intro a ω; unfold ucbIndex; simp [pullCount_zero] have h_ucb_swap : ∀ s, ∫ ω, (ucb (A s ω) s ω - ucb (bestArm ω) s ω) ∂P = 0 := by intro s cases s with | zero => have : ∀ ω, ucb (A 0 ω) 0 ω - ucb (bestArm ω) 0 ω = 0 := fun ω ↦ by - change ucbIndex A R' δ _ 0 ω - ucbIndex A R' δ _ 0 ω = 0 + change ucbIndex A R' (↑σ2) lo hi δ _ 0 ω - ucbIndex A R' (↑σ2) lo hi δ _ 0 ω = 0 simp [h_ucb_zero] exact (integral_congr_ae (ae_of_all _ this)).trans (integral_zero _ _) | succ t => @@ -933,14 +1045,15 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp hf.aestronglyMeasurable, h_map_eq] set g : (Iic t → Fin K × ℝ) × Fin K → ℝ := - fun p ↦ max 0 (min 1 (empMean' t p.1 p.2 + - √(2 * Real.log (1 / δ) / (max 1 (pullCount' t p.1 p.2) : ℝ)))) + fun p ↦ if pullCount' t p.1 p.2 = 0 then hi + else max lo (min hi (empMean' t p.1 p.2 + + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount' t p.1 p.2 : ℝ)))) have h_hist_eq : ∀ (ω : Ω), (fun (i : Iic t) ↦ (A (↑i) ω, R' (↑i) ω)) = IsAlgEnvSeq.hist A R' t ω := by intro ω; rfl - have hg_eq : ∀ a (ω : Ω), ucbIndex A R' δ a (t + 1) ω = + have hg_eq : ∀ a (ω : Ω), ucbIndex A R' (↑σ2) lo hi δ a (t + 1) ω = g (IsAlgEnvSeq.hist A R' t ω, a) := by intro a ω simp only [g, ucbIndex] @@ -948,17 +1061,27 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp pullCount_add_one_eq_pullCount' (A := A) (R' := R'), h_hist_eq] have hg_meas : Measurable g := by - apply Measurable.max measurable_const - apply Measurable.min measurable_const - apply Measurable.add - · exact measurable_apply_fin (fun a ↦ (measurable_empMean' t a).comp measurable_fst) - measurable_snd - · apply Measurable.sqrt - apply Measurable.div measurable_const - apply Measurable.max measurable_const - exact measurable_apply_fin - (fun a ↦ measurable_from_top.comp ((measurable_pullCount' t a).comp measurable_fst)) - measurable_snd + apply Measurable.ite + · have : MeasurableSet {p : (Iic t → Fin K × ℝ) × Fin K | + (pullCount' t p.1 p.2 : ℝ) = (0 : ℝ)} := + measurableSet_eq_fun + (measurable_apply_fin + (fun a ↦ measurable_from_top.comp + ((measurable_pullCount' t a).comp measurable_fst)) + measurable_snd) + measurable_const + simp only [Nat.cast_eq_zero] at this; exact this + · exact measurable_const + · apply Measurable.max measurable_const + apply Measurable.min measurable_const + apply Measurable.add + · exact measurable_apply_fin (fun a ↦ (measurable_empMean' t a).comp measurable_fst) + measurable_snd + · apply Measurable.sqrt + apply Measurable.div measurable_const + exact measurable_apply_fin + (fun a ↦ measurable_from_top.comp ((measurable_pullCount' t a).comp measurable_fst)) + measurable_snd have h_eq_g1 : (fun ω ↦ ucb (A (t + 1) ω) (t + 1) ω) = fun ω ↦ g (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω) := funext fun ω ↦ hg_eq _ _ @@ -968,7 +1091,8 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp have h_int_ucb : ∀ {f : Ω → Fin K}, Measurable f → Integrable (fun ω ↦ ucb (f ω) (t + 1) ω) P := fun hf ↦ ⟨(measurable_apply_fin (fun a ↦ hm_ucb a (t + 1)) hf).aestronglyMeasurable, - HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ norm_ucbIndex_le_one A R' _ _ _ _)⟩ + HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ + norm_ucbIndex_le A R' (↑σ2) lo hi δ hlo _ _ _)⟩ have h_int1 := h_int_ucb (h.measurable_A (t + 1)) have h_int2 := h_int_ucb hm_best rw [show (fun ω ↦ ucb (A (t + 1) ω) (t + 1) ω - @@ -998,9 +1122,9 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp refine Integrable.of_bound ((measurable_apply_fin hm_arm hm_best).sub (measurable_apply_fin hm_arm (h.measurable_A s))).aestronglyMeasurable - 1 (ae_of_all _ fun ω ↦ ?_) + (hi - lo) (ae_of_all _ fun ω ↦ ?_) rw [Real.norm_eq_abs] - exact abs_sub_le_one_of_mem_Icc (hm _ _) (hm _ _) + exact abs_sub_le_of_mem_Icc (hm _ _) (hm _ _) have h_int_gap : Integrable (fun ω ↦ ∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω)) P := integrable_finset_sum _ h_int_gap_s @@ -1020,52 +1144,66 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp intro ω hω apply Finset.sum_nonpos intro s hs - linarith [armMean_le_ucbIndex E' A R' κ hm δ + linarith [armMean_le_ucbIndex E' A R' κ hm (↑σ2) δ (bestArm ω) s ω (hω s (mem_range.mp hs) _)] have h_second_Eδ : ∀ ω ∈ Eδ, ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω) - ≤ 2 * √(8 * Real.log (1 / δ)) * √(↑K * ↑n) := by + ≤ (hi - lo) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by intro ω hω - exact sum_ucbIndex_sub_armMean_le E' A R' κ hm δ n ω hω + exact sum_ucbIndex_sub_armMean_le E' A R' κ hm hlo (↑σ2) δ n ω hω have h_prob : P Eδᶜ ≤ ENNReal.ofReal (2 * K * n * δ) := by - have : Eδᶜ = {ω | ∃ s < n, ∃ a, √(2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ + have : Eδᶜ = {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ |empMean A R' a s ω - armMean a ω|} := by ext ω; simp only [Eδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push_neg; rfl rw [this] exact prob_concentration_fail_delta (hK := hK) (E' := E') (A := A) (R' := R') - (Q := Q) (κ := κ) (P := P) h hs hm n δ hδ hδ1 hδ_large + (Q := Q) (κ := κ) (P := P) h hσ2 hs n δ hδ hδ1 have hm_emp : ∀ a s, Measurable (fun ω ↦ empMean A R' a s ω) := fun a s ↦ measurable_empMean (fun n ↦ h.measurable_A n) (fun n ↦ h.measurable_R n) a s have hm_pc : ∀ a s, Measurable (fun ω ↦ (pullCount A a s ω : ℝ)) := fun a s ↦ measurable_from_top.comp (measurable_pullCount (fun n ↦ h.measurable_A n) a s) have hEδ_meas : MeasurableSet Eδ := by - suffices ∀ s a, MeasurableSet {ω | + suffices ∀ s a, MeasurableSet {ω : Ω | pullCount A a s ω ≠ 0 → |empMean A R' a s ω - armMean a ω| - < √(2 * Real.log (1 / δ) / max 1 ↑(pullCount A a s ω))} by + < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} by simp only [Eδ, Set.setOf_forall] exact .iInter fun s ↦ .iInter fun _ ↦ .iInter fun a ↦ this s a intro s a - exact measurableSet_lt - ((hm_emp a s).sub (h.measurable_armMean a)).abs - ((measurable_const.div (measurable_const.max (hm_pc a s))).sqrt) + have h_eq : {ω : Ω | pullCount A a s ω ≠ 0 → + |empMean A R' a s ω - armMean a ω| + < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} = + {ω | pullCount A a s ω = 0} ∪ {ω | + |empMean A R' a s ω - armMean a ω| + < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} := by + ext ω; simp only [Set.mem_setOf_eq, Set.mem_union]; tauto + rw [h_eq] + have h_eq0 : {ω : Ω | pullCount A a s ω = 0} = + {ω : Ω | (pullCount A a s ω : ℝ) = 0} := by + ext ω; simp [Nat.cast_eq_zero] + exact MeasurableSet.union (h_eq0 ▸ hm_pc a s (measurableSet_singleton (0 : ℝ))) + (measurableSet_lt + ((hm_emp a s).sub (h.measurable_armMean a)).abs + ((measurable_const.div (hm_pc a s)).sqrt)) rw [h_swap] set f1 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω) set f2 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω) - set B := 2 * √(8 * Real.log (1 / δ)) * √(↑K * ↑n) + set B := (hi - lo) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) have s1 := (integral_add_compl hEδ_meas h_int_sum1).symm have s2 := (integral_add_compl hEδ_meas h_int_sum2).symm have h1g : ∫ ω in Eδ, f1 ω ∂P ≤ 0 := setIntegral_nonpos hEδ_meas fun ω hω ↦ h_first_Eδ ω hω - have h_compl_bound : ∀ {f}, Integrable f P → (∀ ω, f ω ≤ ↑n) → - ∫ ω in Eδᶜ, f ω ∂P ≤ ↑n * P.real Eδᶜ := fun hint hle ↦ by + have h_compl_bound : ∀ {f}, Integrable f P → (∀ ω, f ω ≤ ↑n * (hi - lo)) → + ∫ ω in Eδᶜ, f ω ∂P ≤ ↑n * (hi - lo) * P.real Eδᶜ := fun hint hle ↦ by have := setIntegral_mono_on (hf := hint.integrableOn) (hg := integrableOn_const) hEδ_meas.compl fun ω _ ↦ hle ω rwa [setIntegral_const, smul_eq_mul, mul_comm] at this have h1b := h_compl_bound h_int_sum1 fun ω ↦ (abs_le.mp (h_first_bound ω)).2 have h2g : ∫ ω in Eδ, f2 ω ∂P ≤ B := by - have hB : 0 ≤ B := by positivity + have hB : 0 ≤ B := add_nonneg (mul_nonneg (sub_nonneg.mpr hlo) (Nat.cast_nonneg K)) + (by positivity) have := setIntegral_mono_on (hf := h_int_sum2.integrableOn) (hg := integrableOn_const) hEδ_meas fun ω hω ↦ h_second_Eδ ω hω @@ -1076,31 +1214,40 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp ENNReal.toReal_le_of_le_ofReal (by positivity) h_prob rw [s1, s2] have hP0 : 0 ≤ P.real Eδᶜ := by positivity + have hhi_lo : (0 : ℝ) ≤ hi - lo := sub_nonneg.mpr hlo + have h_key := mul_le_mul_of_nonneg_left hP + (mul_nonneg (Nat.cast_nonneg n) hhi_lo) nlinarith lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) 1 (κ (a, e))) - (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc 0 1)) (t : ℕ) : - IsBayesAlgEnvSeq.bayesRegret κ A E' P t ≤ 4 * K + 8 * √(K * t * Real.log t) := by + {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) + {lo hi : ℝ} (hlo : lo ≤ hi) + (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) (t : ℕ) : + IsBayesAlgEnvSeq.bayesRegret κ A E' P t + ≤ 5 * (hi - lo) * K + 8 * √(σ2 * K * t * Real.log t) := by by_cases ht : t = 0 · simp [ht, IsBayesAlgEnvSeq.bayesRegret, IsBayesAlgEnvSeq.regret, regret] + nlinarith [sub_nonneg.mpr hlo, Nat.cast_nonneg (α := ℝ) K, + Real.sqrt_nonneg (↑σ2 * ↑K * (0 : ℝ) * Real.log (0 : ℝ))] by_cases ht1_eq : t = 1 · subst ht1_eq simp only [Nat.cast_one, Real.log_one, mul_zero, Real.sqrt_zero, mul_zero, add_zero] calc IsBayesAlgEnvSeq.bayesRegret κ A E' P 1 - ≤ 1 := by + ≤ hi - lo := by unfold IsBayesAlgEnvSeq.bayesRegret IsBayesAlgEnvSeq.regret Bandits.regret simp only [Finset.range_one, Finset.sum_singleton, Nat.cast_one, one_mul, Kernel.comap_apply] refine (integral_mono_of_nonneg (ae_of_all _ fun ω ↦ sub_nonneg.mpr - (le_ciSup ⟨1, by rintro _ ⟨a, rfl⟩; exact (hm a _).2⟩ _)) - (integrable_const 1) (ae_of_all _ fun ω ↦ by + (le_ciSup ⟨hi, by rintro _ ⟨a, rfl⟩; exact (hm a _).2⟩ _)) + (integrable_const (hi - lo)) (ae_of_all _ fun ω ↦ by linarith [ciSup_le fun a ↦ (hm a (E' ω)).2, (hm (A 0 ω) (E' ω)).1])).trans ?_ simp - _ ≤ 4 * (K : ℝ) := by - nlinarith [show (1 : ℝ) ≤ K from Nat.one_le_cast.mpr (Nat.one_le_of_lt hK)] + _ ≤ 5 * (hi - lo) * (K : ℝ) := by + nlinarith [show (1 : ℝ) ≤ K from Nat.one_le_cast.mpr (Nat.one_le_of_lt hK), + sub_nonneg.mpr hlo] -- For t ≥ 2, we have δ = 1/t² < 1 · have ht2 : 2 ≤ t := by omega have htpos : (0 : ℝ) < t := by positivity @@ -1111,36 +1258,28 @@ lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] have ht2_real : (2 : ℝ) ≤ t := Nat.ofNat_le_cast.mpr ht2 calc (1 : ℝ) < 2 ^ 2 := by norm_num _ ≤ (t : ℝ) ^ 2 := by gcongr - -- First term simplification: 4K · t² · (1/t²) = 4K - have h_first : 4 * (K : ℝ) * ↑t ^ 2 * (1 / (↑t) ^ 2) = 4 * ↑K := by - field_simp + -- First term simplification: (hi-lo)*K + 4*(hi-lo)*K*t²*(1/t²) = 5*(hi-lo)*K + have h_first : (hi - lo) * ↑K + 4 * (hi - lo) * ↑K * ↑t ^ 2 * (1 / (↑t) ^ 2) + = 5 * (hi - lo) * ↑K := by + field_simp; ring -- Second term simplification: log(1/(1/t²)) = log(t²) = 2 log(t) have h_log : Real.log (1 / (1 / (↑t : ℝ) ^ 2)) = 2 * Real.log ↑t := by rw [one_div_one_div, Real.log_pow]; norm_cast - -- For t ≥ 2, we have 2 * log(t²) = 4 log(t) ≥ 4 log(2) > 1 (since log(2) > 0.69) - have hδ_large : 1 < 2 * Real.log (1 / (1 / (↑t : ℝ) ^ 2)) := by - rw [h_log] - have h_log2 : (1 : ℝ) / 2 < Real.log 2 := by - rw [Real.lt_log_iff_exp_lt (by norm_num : (0 : ℝ) < 2)] - calc Real.exp (1 / 2) = Real.sqrt (Real.exp 1) := by rw [← Real.exp_half] - _ < Real.sqrt 4 := by gcongr; linarith [Real.exp_one_lt_d9] - _ = 2 := by rw [show (4 : ℝ) = 2 ^ 2 by norm_num, Real.sqrt_sq (by norm_num)] - have ht2_real : (2 : ℝ) ≤ t := Nat.ofNat_le_cast.mpr ht2 - linarith [Real.log_le_log (by norm_num : (0 : ℝ) < 2) ht2_real] - have hK_real_pos : (0 : ℝ) < K := Nat.cast_pos.mpr hK - have hKt_nonneg : (0 : ℝ) ≤ ↑K * ↑t := by positivity calc IsBayesAlgEnvSeq.bayesRegret κ A E' P t - ≤ 4 * ↑K * ↑t ^ 2 * (1 / (↑t) ^ 2) - + 2 * √(8 * Real.log (1 / (1 / (↑t) ^ 2))) * √(↑K * ↑t) := + ≤ (hi - lo) * ↑K + 4 * (hi - lo) * ↑K * ↑t ^ 2 * (1 / (↑t) ^ 2) + + 2 * √(8 * ↑σ2 * Real.log (1 / (1 / (↑t) ^ 2))) * √(↑K * ↑t) := bayesRegret_le_of_delta (hK := hK) (E' := E') (A := A) (R' := R') (Q := Q) - (κ := κ) (P := P) h hs hm t (1 / (↑t) ^ 2) hδ hδ1 hδ_large - _ = 4 * ↑K + 2 * √(16 * Real.log ↑t) * √(↑K * ↑t) := by rw [h_first, h_log]; ring_nf - _ = 4 * ↑K + 8 * (√(Real.log ↑t) * √(↑K * ↑t)) := by - rw [show (16 : ℝ) = 4 ^ 2 by norm_num, Real.sqrt_mul (by norm_num : (0 : ℝ) ≤ 4 ^ 2), - Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 4)]; ring - _ = 4 * ↑K + 8 * √(↑K * ↑t * Real.log ↑t) := by - rw [← Real.sqrt_mul (Real.log_nonneg (Nat.one_le_cast.mpr (Nat.pos_of_ne_zero ht)))] - ring_nf + (κ := κ) (P := P) h hσ2 hs hlo hm t (1 / (↑t) ^ 2) hδ hδ1 + _ = 5 * (hi - lo) * ↑K + 8 * (√(↑σ2 * Real.log ↑t) * √(↑K * ↑t)) := by + rw [h_first, h_log, + show (8 : ℝ) * ↑σ2 * (2 * Real.log ↑t) = 4 ^ 2 * (↑σ2 * Real.log ↑t) by ring, + Real.sqrt_mul (by positivity : (0 : ℝ) ≤ 4 ^ 2), + Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 4)] + ring + _ = 5 * (hi - lo) * ↑K + 8 * √(↑σ2 * ↑K * ↑t * Real.log ↑t) := by + rw [← Real.sqrt_mul (mul_nonneg (NNReal.coe_nonneg σ2) + (Real.log_nonneg (Nat.one_le_cast.mpr (Nat.pos_of_ne_zero ht))))] + congr 1; congr 1; congr 1; ring end Regret From 0a62626de54605ae0736605fdc057850c9431052 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 18 Feb 2026 20:29:56 +0000 Subject: [PATCH 041/155] Generalize TS Bayesian regret bound (tighter) --- LeanBandits/BanditAlgorithms/TS.lean | 342 ++++++++++++++++++++++++--- 1 file changed, 308 insertions(+), 34 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 355d77b3..6517652e 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -939,21 +939,262 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] rw [← ENNReal.ofReal_natCast K, ← ENNReal.ofReal_mul (Nat.cast_nonneg K)] congr 1; ring +lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] + (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) + {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) + (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : + P {ω | ∃ s < n, pullCount A (IsBayesAlgEnvSeq.bestArm κ E' ω) s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / + (pullCount A (IsBayesAlgEnvSeq.bestArm κ E' ω) s ω : ℝ)) ≤ + |empMean A R' (IsBayesAlgEnvSeq.bestArm κ E' ω) s ω - + IsBayesAlgEnvSeq.armMean κ E' (IsBayesAlgEnvSeq.bestArm κ E' ω) ω|} + ≤ ENNReal.ofReal (2 * n * δ) := by + by_cases hn : n = 0 + · simp [hn] + have hn' : 0 < n := Nat.pos_of_ne_zero hn + rw [show IsBayesAlgEnvSeq.bestArm κ E' = envToBestArm κ ∘ E' from + bestArm_eq_envToBestArm_comp_env κ] + let badSetIT := fun (a : Fin K) (s : ℕ) (e : E) ↦ {ω : ℕ → (Fin K) × ℝ | + pullCount IT.action a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ + |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} + have h_set_eq : {ω | ∃ s < n, pullCount A ((envToBestArm κ ∘ E') ω) s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / + (pullCount A ((envToBestArm κ ∘ E') ω) s ω : ℝ)) ≤ + |empMean A R' ((envToBestArm κ ∘ E') ω) s ω - + IsBayesAlgEnvSeq.armMean κ E' ((envToBestArm κ ∘ E') ω) ω|} = + (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1} := by + ext ω + simp only [Set.mem_setOf_eq, Finset.mem_range, Set.mem_preimage, Set.mem_iUnion, + badSetIT, IsBayesAlgEnvSeq.armMean, Function.comp_apply, exists_prop] + rfl + rw [h_set_eq] + have h_meas_pair : + Measurable (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) := + h.measurable_E.prodMk h.measurable_traj + have h_disint : P.map (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) = + P.map E' ⊗ₘ condDistrib (IsBayesAlgEnvSeq.traj A R') E' P := + (compProd_map_condDistrib (h.measurable_traj.aemeasurable)).symm + have h_cond_bound : ∀ᵐ e ∂(P.map E'), ∀ a : Fin K, + (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + (⋃ s ∈ Finset.range n, badSetIT a s e) ≤ ENNReal.ofReal (2 * n * δ) := by + filter_upwards [IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h] with e h_isAlgEnvSeq + intro a + let ν := κ.comap (·, e) (by fun_prop) + let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e + have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by + simp only [ν, Kernel.comap_apply]; exact hs a' e + have h_mean : (ν a)[id] = (κ (a, e))[id] := by simp only [ν, Kernel.comap_apply] + let B_low := fun m : ℕ ↦ + {x : ℝ | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} + let B_high := fun m : ℕ ↦ + {x : ℝ | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)} + have h_stream_bound : ∀ m : ℕ, m ≠ 0 → + streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ≤ + ENNReal.ofReal (2 * δ) := by + intro m hm0 + calc streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} + ≤ streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m} + + streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_high m} := by + have h_union : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ⊆ + {ω | ∑ i ∈ range m, ω i a ∈ B_low m} ∪ + {ω | ∑ i ∈ range m, ω i a ∈ B_high m} := by + intro ω hω; simp only [Set.mem_setOf_eq, Set.mem_union] at hω ⊢; exact hω + exact (measure_mono h_union).trans (measure_union_le _ _) + _ ≤ ENNReal.ofReal δ + ENNReal.ofReal δ := by + gcongr + · have h_eq : {ω : ℕ → Fin K → ℝ | + ∑ i ∈ range m, ω i a ∈ B_low m} = + {ω | (∑ i ∈ range m, ω i a) / m + + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} := by + ext ω; simp only [Set.mem_setOf_eq, B_low] + rw [h_eq] + exact streamMeasure_concentration_le_delta hσ2 h_subG a m hm0 δ hδ hδ1 + · have h_eq : {ω : ℕ → Fin K → ℝ | + ∑ i ∈ range m, ω i a ∈ B_high m} = + {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - + √(2 * ↑σ2 * Real.log (1 / δ) / m)} := by + ext ω; simp only [Set.mem_setOf_eq, B_high] + rw [h_eq] + exact streamMeasure_concentration_ge_delta hσ2 h_subG a m hm0 δ hδ hδ1 + _ = ENNReal.ofReal (2 * δ) := by + rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf + have hB_meas : ∀ m, MeasurableSet (B_low m ∪ B_high m) := fun m ↦ + MeasurableSet.union (measurableSet_le (by fun_prop) (by fun_prop)) + (measurableSet_le (by fun_prop) (by fun_prop)) + let S := Finset.Icc 1 (n - 1) + have hS_card : S.card = n - 1 := by simp only [Nat.card_Icc, S]; omega + have h_decomp : ⋃ s ∈ Finset.range n, badSetIT a s e = + ⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by + ext ω + simp only [Set.mem_iUnion, Finset.mem_range, exists_prop, badSetIT, Set.mem_setOf_eq, + Finset.mem_Icc, S] + constructor + · rintro ⟨s, hs, hbad⟩ + let m := pullCount IT.action a s ω + have hm_pos : 0 < m := Nat.pos_of_ne_zero hbad.1 + have hm_le : m ≤ n - 1 := by + have h1 : m ≤ s := pullCount_le (A := IT.action) a s ω + omega + refine ⟨m, ⟨hm_pos, hm_le⟩, s, hs, rfl, ?_⟩ + simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] + have h_pc_pos : (0 : ℝ) < m := Nat.cast_pos.mpr hm_pos + simp only [empMean] at hbad + have h_abs := le_abs'.mp hbad.2 + rcases h_abs with h_neg | h_pos + · left; linarith + · right; linarith + · rintro ⟨m, ⟨hm_pos, hm_le⟩, s, hs, hpc, hB⟩ + refine ⟨s, hs, ?_⟩ + simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] at hB + have h_pc_pos : (0 : ℝ) < m := Nat.cast_pos.mpr hm_pos + simp only [empMean, hpc] + refine ⟨Nat.one_le_iff_ne_zero.mp hm_pos, ?_⟩ + cases hB with + | inl h => + have h1 : sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id] ≤ + -√(2 * ↑σ2 * Real.log (1 / δ) / m) := by linarith + calc √(2 * ↑σ2 * Real.log (1 / δ) / m) + ≤ -(sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id]) := by linarith + _ ≤ |sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id]| := neg_le_abs _ + | inr h => + have h1 : √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ + sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id] := by linarith + exact h1.trans (le_abs_self _) + rw [h_decomp] + calc P' (⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m}) + ≤ ∑ m ∈ S, P' {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := + measure_biUnion_finset_le S _ + _ ≤ ∑ m ∈ S, streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := by + apply Finset.sum_le_sum + intro m hm + have hm_pos : m ≠ 0 := Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1 + have hm_le : m ≤ n - 1 := (Finset.mem_Icc.mp hm).2 + have h_contain : {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} ⊆ + {ω | pullCount IT.action a (n - 1) ω = m ∧ + sumRewards IT.action IT.reward a (n - 1) ω ∈ B_low m ∪ B_high m} ∪ + {ω | pullCount IT.action a (n - 1) ω > m ∧ + ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by + intro ω ⟨s, hs, hpc, hB⟩ + simp only [Set.mem_union, Set.mem_setOf_eq] + have hs' : s ≤ n - 1 := Nat.le_sub_one_of_lt hs + have h_pc_mono := pullCount_mono (A := IT.action) a hs' ω + by_cases h_eq : pullCount IT.action a (n - 1) ω = m + · left + refine ⟨h_eq, ?_⟩ + have h_pc_eq : pullCount IT.action a s ω = pullCount IT.action a (n - 1) ω := + hpc.symm ▸ h_eq.symm + rw [← sumRewards_eq_of_pullCount_eq h_pc_eq] + exact hB + · right + exact ⟨by omega, s, hs, hpc, hB⟩ + calc P' {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} + ≤ P' {ω | ∃ s, s ≤ n - 1 ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by + apply measure_mono + intro ω ⟨s, hs, hpc, hB⟩ + exact ⟨s, Nat.le_sub_one_of_lt hs, hpc, hB⟩ + _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := + prob_exists_pullCount_eq_and_sumRewards_mem_le (n := n - 1) + h_isAlgEnvSeq (hB_meas m) + _ ≤ ∑ _m ∈ S, ENNReal.ofReal (2 * δ) := by + apply Finset.sum_le_sum + intro m hm + have hm_pos : m ≠ 0 := Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1 + exact h_stream_bound m hm_pos + _ = S.card • ENNReal.ofReal (2 * δ) := by simp only [Finset.sum_const] + _ = (n - 1) • ENNReal.ofReal (2 * δ) := by rw [hS_card] + _ ≤ ENNReal.ofReal (2 * n * δ) := by + rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast (n - 1), + ← ENNReal.ofReal_mul (Nat.cast_nonneg (n - 1))] + apply ENNReal.ofReal_le_ofReal + have h1 : (n - 1 : ℕ) ≤ n := Nat.sub_le n 1 + have h2 : (↑(n - 1) : ℝ) ≤ (↑n : ℝ) := Nat.cast_le.mpr h1 + nlinarith [h2, hδ.le] + have h_cond_best : ∀ᵐ e ∂(P.map E'), + (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + (⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ e) s e) ≤ + ENNReal.ofReal (2 * n * δ) := by + filter_upwards [h_cond_bound] with e he + exact he (envToBestArm κ e) + have h_kernel : ∀ a, Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (a, p.1))[id]) := + fun a ↦ stronglyMeasurable_id.integral_kernel.measurable.comp + (measurable_const.prodMk measurable_fst) + have h_meas_badSetIT : ∀ a s, MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | + p.2 ∈ badSetIT a s p.1} := by + intro a s + simp only [badSetIT] + change MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | + pullCount IT.action a s p.2 ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s p.2 : ℝ)) ≤ + |empMean IT.action IT.reward a s p.2 - (κ (a, p.1))[id]|} + exact MeasurableSet.inter + (((measurable_pullCount IT.measurable_action a s).comp measurable_snd) + (measurableSet_singleton (0 : ℕ)).compl) + (measurableSet_le (by fun_prop) + (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp + measurable_snd).sub (h_kernel a)).abs) + have h_meas_set : MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | + p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1} := by + have h_eq : {p : E × (ℕ → (Fin K) × ℝ) | + p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1} = + ⋃ a : Fin K, ((envToBestArm κ ∘ Prod.fst) ⁻¹' {a} ∩ + ⋃ s ∈ Finset.range n, {p | p.2 ∈ badSetIT a s p.1}) := by + ext p; simp only [Set.mem_setOf_eq, Set.mem_iUnion, Set.mem_inter_iff, + Set.mem_preimage, Function.comp_apply, Set.mem_singleton_iff, Finset.mem_range] + constructor + · intro ⟨s, hs, hm⟩; exact ⟨envToBestArm κ p.1, rfl, s, hs, hm⟩ + · rintro ⟨a, ha, s, hs, hm⟩; exact ⟨s, hs, ha ▸ hm⟩ + rw [h_eq] + exact .iUnion fun a ↦ .inter + ((measurable_envToBestArm (κ := κ) |>.comp measurable_fst) (measurableSet_singleton a)) + (.biUnion (Finset.range n).countable_toSet fun s _ ↦ h_meas_badSetIT a s) + calc P ((fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1}) + = (P.map (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω))) + {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1} := by + rw [Measure.map_apply h_meas_pair h_meas_set] + _ = (P.map E' ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.traj A R') E' P) + {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1} := by + rw [h_disint] + _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + (⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ e) s e) ∂(P.map E') := by + rw [Measure.compProd_apply h_meas_set]; rfl + _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * n * δ) ∂(P.map E') := by + apply lintegral_mono_ae h_cond_best + _ = ENNReal.ofReal (2 * n * δ) := by + rw [lintegral_const, Measure.map_apply h.measurable_E MeasurableSet.univ] + simp [measure_univ] + lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) - {lo hi : ℝ} (hlo : lo ≤ hi) (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) + {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : IsBayesAlgEnvSeq.bayesRegret κ A E' P n - ≤ (hi - lo) * ↑K + 4 * (hi - lo) * K * n ^ 2 * δ + + ≤ (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * n ^ 2 * δ + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by + have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _) + have hlo : lo ≤ hi := h1.trans h2 let bestArm := IsBayesAlgEnvSeq.bestArm κ E' let armMean := IsBayesAlgEnvSeq.armMean κ E' let ucb := ucbIndex A R' (↑σ2) lo hi δ set Eδ : Set Ω := {ω | ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → |empMean A R' a s ω - armMean a ω| < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ))} + set Fδ : Set Ω := {ω | ∀ s < n, pullCount A (bestArm ω) s ω ≠ 0 → + |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω| + < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A (bestArm ω) s ω : ℝ))} have hm_ucb : ∀ a t, Measurable (ucbIndex A R' (↑σ2) lo hi δ a t) := fun a t ↦ measurable_ucbIndex hK E' A R' Q κ P h (↑σ2) lo hi δ a t have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.armMean κ E' a) := @@ -1138,14 +1379,14 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω))) ∂P := by congr 1; ext ω; linarith [h_pw ω] _ = _ := integral_add h_int_sum1 h_int_sum2 - have h_first_Eδ : ∀ ω ∈ Eδ, + have h_first_Fδ : ∀ ω ∈ Fδ, ∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω) ≤ 0 := by intro ω hω apply Finset.sum_nonpos intro s hs linarith [armMean_le_ucbIndex E' A R' κ hm (↑σ2) δ - (bestArm ω) s ω (hω s (mem_range.mp hs) _)] + (bestArm ω) s ω (hω s (mem_range.mp hs))] have h_second_Eδ : ∀ ω ∈ Eδ, ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω) ≤ (hi - lo) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by @@ -1163,12 +1404,9 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp fun a s ↦ measurable_empMean (fun n ↦ h.measurable_A n) (fun n ↦ h.measurable_R n) a s have hm_pc : ∀ a s, Measurable (fun ω ↦ (pullCount A a s ω : ℝ)) := fun a s ↦ measurable_from_top.comp (measurable_pullCount (fun n ↦ h.measurable_A n) a s) - have hEδ_meas : MeasurableSet Eδ := by - suffices ∀ s a, MeasurableSet {ω : Ω | pullCount A a s ω ≠ 0 → - |empMean A R' a s ω - armMean a ω| - < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} by - simp only [Eδ, Set.setOf_forall] - exact .iInter fun s ↦ .iInter fun _ ↦ .iInter fun a ↦ this s a + have h_arm_meas : ∀ s a, MeasurableSet {ω : Ω | pullCount A a s ω ≠ 0 → + |empMean A R' a s ω - armMean a ω| + < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} := by intro s a have h_eq : {ω : Ω | pullCount A a s ω ≠ 0 → |empMean A R' a s ω - armMean a ω| @@ -1185,22 +1423,49 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (measurableSet_lt ((hm_emp a s).sub (h.measurable_armMean a)).abs ((measurable_const.div (hm_pc a s)).sqrt)) + have hEδ_meas : MeasurableSet Eδ := by + simp only [Eδ, Set.setOf_forall] + exact .iInter fun s ↦ .iInter fun _ ↦ .iInter fun a ↦ h_arm_meas s a + have hFδ_meas : MeasurableSet Fδ := by + simp only [Fδ, Set.setOf_forall] + apply MeasurableSet.iInter; intro s + apply MeasurableSet.iInter; intro _ + have : {ω : Ω | pullCount A (bestArm ω) s ω ≠ 0 → + |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω| + < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A (bestArm ω) s ω))} = + ⋃ a : Fin K, (bestArm ⁻¹' {a}) ∩ {ω | pullCount A a s ω ≠ 0 → + |empMean A R' a s ω - armMean a ω| + < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} := by + ext ω + simp only [Set.mem_iUnion, Set.mem_inter_iff, Set.mem_preimage, + Set.mem_singleton_iff, Set.mem_setOf_eq] + constructor + · exact fun h => ⟨_, rfl, h⟩ + · rintro ⟨_, rfl, h⟩; exact h + rw [this] + exact .iUnion fun a => .inter (hm_best (measurableSet_singleton a)) (h_arm_meas s a) + have h_prob_F : P Fδᶜ ≤ ENNReal.ofReal (2 * ↑n * δ) := by + have : Fδᶜ = {ω | ∃ s < n, pullCount A (bestArm ω) s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A (bestArm ω) s ω : ℝ)) ≤ + |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω|} := by + ext ω; simp only [Fδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push_neg; rfl + rw [this] + exact prob_concentration_bestArm_fail_delta (hK := hK) (E' := E') (A := A) (R' := R') + (Q := Q) (κ := κ) (P := P) h hσ2 hs n δ hδ hδ1 rw [h_swap] set f1 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω) set f2 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω) set B := (hi - lo) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) - have s1 := (integral_add_compl hEδ_meas h_int_sum1).symm + have s1 := (integral_add_compl hFδ_meas h_int_sum1).symm have s2 := (integral_add_compl hEδ_meas h_int_sum2).symm - have h1g : ∫ ω in Eδ, f1 ω ∂P ≤ 0 := - setIntegral_nonpos hEδ_meas fun ω hω ↦ h_first_Eδ ω hω - have h_compl_bound : ∀ {f}, Integrable f P → (∀ ω, f ω ≤ ↑n * (hi - lo)) → - ∫ ω in Eδᶜ, f ω ∂P ≤ ↑n * (hi - lo) * P.real Eδᶜ := fun hint hle ↦ by - have := setIntegral_mono_on (hf := hint.integrableOn) (hg := integrableOn_const) - hEδ_meas.compl fun ω _ ↦ hle ω + have h1g : ∫ ω in Fδ, f1 ω ∂P ≤ 0 := + setIntegral_nonpos hFδ_meas fun ω hω ↦ h_first_Fδ ω hω + have h1b : ∫ ω in Fδᶜ, f1 ω ∂P ≤ ↑n * (hi - lo) * P.real Fδᶜ := by + have := setIntegral_mono_on (hf := h_int_sum1.integrableOn) (hg := integrableOn_const) + hFδ_meas.compl fun ω _ ↦ (abs_le.mp (h_first_bound ω)).2 rwa [setIntegral_const, smul_eq_mul, mul_comm] at this - have h1b := h_compl_bound h_int_sum1 fun ω ↦ (abs_le.mp (h_first_bound ω)).2 have h2g : ∫ ω in Eδ, f2 ω ∂P ≤ B := by have hB : 0 ≤ B := add_nonneg (mul_nonneg (sub_nonneg.mpr hlo) (Nat.cast_nonneg K)) (by positivity) @@ -1209,13 +1474,21 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp fun ω hω ↦ h_second_Eδ ω hω rw [setIntegral_const, smul_eq_mul, mul_comm] at this exact le_trans this (mul_le_of_le_one_right hB measureReal_le_one) - have h2b := h_compl_bound h_int_sum2 fun ω ↦ (abs_le.mp (h_second_bound ω)).2 - have hP : P.real Eδᶜ ≤ 2 * ↑K * ↑n * δ := + have h2b : ∫ ω in Eδᶜ, f2 ω ∂P ≤ ↑n * (hi - lo) * P.real Eδᶜ := by + have := setIntegral_mono_on (hf := h_int_sum2.integrableOn) (hg := integrableOn_const) + hEδ_meas.compl fun ω _ ↦ (abs_le.mp (h_second_bound ω)).2 + rwa [setIntegral_const, smul_eq_mul, mul_comm] at this + have hPF : P.real Fδᶜ ≤ 2 * ↑n * δ := + ENNReal.toReal_le_of_le_ofReal (by positivity) h_prob_F + have hPE : P.real Eδᶜ ≤ 2 * ↑K * ↑n * δ := ENNReal.toReal_le_of_le_ofReal (by positivity) h_prob rw [s1, s2] - have hP0 : 0 ≤ P.real Eδᶜ := by positivity + have hPF0 : 0 ≤ P.real Fδᶜ := by positivity + have hPE0 : 0 ≤ P.real Eδᶜ := by positivity have hhi_lo : (0 : ℝ) ≤ hi - lo := sub_nonneg.mpr hlo - have h_key := mul_le_mul_of_nonneg_left hP + have h_key_F := mul_le_mul_of_nonneg_left hPF + (mul_nonneg (Nat.cast_nonneg n) hhi_lo) + have h_key_E := mul_le_mul_of_nonneg_left hPE (mul_nonneg (Nat.cast_nonneg n) hhi_lo) nlinarith @@ -1223,13 +1496,14 @@ lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) - {lo hi : ℝ} (hlo : lo ≤ hi) - (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) (t : ℕ) : + {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) (t : ℕ) : IsBayesAlgEnvSeq.bayesRegret κ A E' P t - ≤ 5 * (hi - lo) * K + 8 * √(σ2 * K * t * Real.log t) := by + ≤ (3 * K + 2) * (hi - lo) + 8 * √(σ2 * K * t * Real.log t) := by + have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _) + have hlo : lo ≤ hi := h1.trans h2 by_cases ht : t = 0 · simp [ht, IsBayesAlgEnvSeq.bayesRegret, IsBayesAlgEnvSeq.regret, regret] - nlinarith [sub_nonneg.mpr hlo, Nat.cast_nonneg (α := ℝ) K, + nlinarith [sub_nonneg.mpr hlo, show (0 : ℝ) < K from Nat.cast_pos.mpr hK, Real.sqrt_nonneg (↑σ2 * ↑K * (0 : ℝ) * Real.log (0 : ℝ))] by_cases ht1_eq : t = 1 · subst ht1_eq @@ -1245,7 +1519,7 @@ lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] linarith [ciSup_le fun a ↦ (hm a (E' ω)).2, (hm (A 0 ω) (E' ω)).1])).trans ?_ simp - _ ≤ 5 * (hi - lo) * (K : ℝ) := by + _ ≤ (3 * ↑K + 2) * (hi - lo) := by nlinarith [show (1 : ℝ) ≤ K from Nat.one_le_cast.mpr (Nat.one_le_of_lt hK), sub_nonneg.mpr hlo] -- For t ≥ 2, we have δ = 1/t² < 1 @@ -1258,25 +1532,25 @@ lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] have ht2_real : (2 : ℝ) ≤ t := Nat.ofNat_le_cast.mpr ht2 calc (1 : ℝ) < 2 ^ 2 := by norm_num _ ≤ (t : ℝ) ^ 2 := by gcongr - -- First term simplification: (hi-lo)*K + 4*(hi-lo)*K*t²*(1/t²) = 5*(hi-lo)*K - have h_first : (hi - lo) * ↑K + 4 * (hi - lo) * ↑K * ↑t ^ 2 * (1 / (↑t) ^ 2) - = 5 * (hi - lo) * ↑K := by + -- First term: (hi-lo)*K + 2*(K+1)*(hi-lo)*t²*(1/t²) = (3K+2)*(hi-lo) + have h_first : (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * ↑t ^ 2 * (1 / (↑t) ^ 2) + = (3 * ↑K + 2) * (hi - lo) := by field_simp; ring -- Second term simplification: log(1/(1/t²)) = log(t²) = 2 log(t) have h_log : Real.log (1 / (1 / (↑t : ℝ) ^ 2)) = 2 * Real.log ↑t := by rw [one_div_one_div, Real.log_pow]; norm_cast calc IsBayesAlgEnvSeq.bayesRegret κ A E' P t - ≤ (hi - lo) * ↑K + 4 * (hi - lo) * ↑K * ↑t ^ 2 * (1 / (↑t) ^ 2) + ≤ (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * ↑t ^ 2 * (1 / (↑t) ^ 2) + 2 * √(8 * ↑σ2 * Real.log (1 / (1 / (↑t) ^ 2))) * √(↑K * ↑t) := bayesRegret_le_of_delta (hK := hK) (E' := E') (A := A) (R' := R') (Q := Q) - (κ := κ) (P := P) h hσ2 hs hlo hm t (1 / (↑t) ^ 2) hδ hδ1 - _ = 5 * (hi - lo) * ↑K + 8 * (√(↑σ2 * Real.log ↑t) * √(↑K * ↑t)) := by + (κ := κ) (P := P) h hσ2 hs hm t (1 / (↑t) ^ 2) hδ hδ1 + _ = (3 * ↑K + 2) * (hi - lo) + 8 * (√(↑σ2 * Real.log ↑t) * √(↑K * ↑t)) := by rw [h_first, h_log, show (8 : ℝ) * ↑σ2 * (2 * Real.log ↑t) = 4 ^ 2 * (↑σ2 * Real.log ↑t) by ring, Real.sqrt_mul (by positivity : (0 : ℝ) ≤ 4 ^ 2), Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 4)] ring - _ = 5 * (hi - lo) * ↑K + 8 * √(↑σ2 * ↑K * ↑t * Real.log ↑t) := by + _ = (3 * ↑K + 2) * (hi - lo) + 8 * √(↑σ2 * ↑K * ↑t * Real.log ↑t) := by rw [← Real.sqrt_mul (mul_nonneg (NNReal.coe_nonneg σ2) (Real.log_nonneg (Nat.one_le_cast.mpr (Nat.pos_of_ne_zero ht))))] congr 1; congr 1; congr 1; ring From e29b6a6787222d0e6428a3ce3247f9ca1374ee4a Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 18 Feb 2026 21:57:33 +0000 Subject: [PATCH 042/155] Minor --- LeanBandits/BanditAlgorithms/TS.lean | 597 +++++++++------------------ 1 file changed, 196 insertions(+), 401 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 6517652e..b9e655ed 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -44,7 +44,10 @@ instance : IsProbabilityMeasure (initialPolicy hK Q κ) := by end TS variable {K : ℕ} -variable {E : Type*} [mE : MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] + +section Algorithm + +variable {E : Type*} [MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] /-- The Thompson Sampling (TS) algorithm: actions are chosen according to the probability that they are optimal given prior knowledge represented by a prior distribution `Q` and a data generation @@ -55,8 +58,11 @@ def tsAlgorithm (hK : 0 < K) (Q : Measure E) [IsProbabilityMeasure Q] policy := TS.policy hK Q κ p0 := TS.initialPolicy hK Q κ +end Algorithm + section Regret +variable {E : Type*} [mE : MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] variable (hK : 0 < K) variable {Ω : Type*} [MeasurableSpace Ω] variable (E' : Ω → E) (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → ℝ) @@ -248,13 +254,11 @@ lemma bayesRegret_eq_sum_integral_gap [Nonempty (Fin K)] refine ⟨hmeas.aestronglyMeasurable, HasFiniteIntegral.of_bounded (C := 2 * C) (Filter.Eventually.of_forall fun ω => ?_)⟩ simp only [Real.norm_eq_abs, gap, Kernel.comap_apply] - set e := E' ω - have hbdd : BddAbove (Set.range fun i => (κ (i, e))[id]) := - ⟨C, by rintro _ ⟨i, rfl⟩; exact le_of_abs_le (hm i e)⟩ + have hbdd : BddAbove (Set.range fun i => (κ (i, E' ω))[id]) := + ⟨C, by rintro _ ⟨i, rfl⟩; exact le_of_abs_le (hm i _)⟩ rw [abs_of_nonneg (sub_nonneg.mpr (le_ciSup hbdd _))] - have h1 : (⨆ i, (κ (i, e))[id]) ≤ C := ciSup_le fun i => le_of_abs_le (hm i e) - have h2 : -C ≤ (κ (A s ω, e))[id] := neg_le_of_abs_le (hm (A s ω) e) - linarith + linarith [ciSup_le fun i => le_of_abs_le (hm i (E' ω)), + neg_le_of_abs_le (hm (A s ω) (E' ω))] omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsProbabilityMeasure Q] [IsMarkovKernel κ] [IsProbabilityMeasure P] in @@ -494,6 +498,30 @@ lemma streamMeasure_concentration_ge_delta {α : Type*} [MeasurableSpace α] Real.log (1 / δ) from by field_simp [ne_of_gt hσ2_pos, ne_of_gt hk_pos]] rw [Real.exp_log (by positivity : (0 : ℝ) < 1 / δ), one_div, inv_inv] +private lemma streamMeasure_concentration_bound {α : Type*} [MeasurableSpace α] + {ν : Kernel α ℝ} [IsMarkovKernel ν] {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) + (a : α) {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) (m : ℕ) (hm : m ≠ 0) : + streamMeasure ν {ω : ℕ → α → ℝ | ∑ i ∈ range m, ω i a ∈ + {x | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} ∪ + {x | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)}} ≤ + ENNReal.ofReal (2 * δ) := + calc streamMeasure ν {ω : ℕ → α → ℝ | ∑ i ∈ range m, ω i a ∈ + {x | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} ∪ + {x | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)}} + ≤ streamMeasure ν {ω | (∑ i ∈ range m, ω i a) / m + + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} + + streamMeasure ν {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - + √(2 * ↑σ2 * Real.log (1 / δ) / m)} := by + apply (measure_mono (fun ω hω ↦ ?_)).trans (measure_union_le _ _) + simp only [Set.mem_setOf_eq, Set.mem_union] at hω ⊢; exact hω + _ ≤ ENNReal.ofReal δ + ENNReal.ofReal δ := by + gcongr + · exact streamMeasure_concentration_le_delta hσ2 hν a m hm δ hδ hδ1 + · exact streamMeasure_concentration_ge_delta hσ2 hν a m hm δ hδ hδ1 + _ = ENNReal.ofReal (2 * δ) := by + rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf + lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) @@ -521,29 +549,8 @@ lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] {x : ℝ | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)} have h_stream_bound : ∀ m : ℕ, m ≠ 0 → streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ≤ - ENNReal.ofReal (2 * δ) := by - intro m hm0 - calc streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} - ≤ streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m} + - streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_high m} := by - have h_union : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ⊆ - {ω | ∑ i ∈ range m, ω i a ∈ B_low m} ∪ {ω | ∑ i ∈ range m, ω i a ∈ B_high m} := by - intro ω hω; simp only [Set.mem_setOf_eq, Set.mem_union] at hω ⊢; exact hω - exact (measure_mono h_union).trans (measure_union_le _ _) - _ ≤ ENNReal.ofReal δ + ENNReal.ofReal δ := by - gcongr - · have h_eq : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m} = - {ω | (∑ i ∈ range m, ω i a) / m + - √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} := by - ext ω; simp only [Set.mem_setOf_eq, B_low] - rw [h_eq]; exact streamMeasure_concentration_le_delta hσ2 h_subG a m hm0 δ hδ hδ1 - · have h_eq : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_high m} = - {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - - √(2 * ↑σ2 * Real.log (1 / δ) / m)} := by - ext ω; simp only [Set.mem_setOf_eq, B_high] - rw [h_eq]; exact streamMeasure_concentration_ge_delta hσ2 h_subG a m hm0 δ hδ hδ1 - _ = ENNReal.ofReal (2 * δ) := by - rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf + ENNReal.ofReal (2 * δ) := + fun m hm0 ↦ streamMeasure_concentration_bound hσ2 h_subG a hδ hδ1 m hm0 let badSet := {ω : ℕ → (Fin K) × ℝ | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ |empMean IT.action IT.reward a s ω - (ν a)[id]|} @@ -566,49 +573,24 @@ lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] set m := pullCount IT.action a s ω with hm_def have hms : m ≤ s := pullCount_le (A := IT.action) a s ω by_cases hm0 : m = 0 - · have h_empMean_zero : empMean IT.action IT.reward a s ω = 0 := by + · exfalso + have h_empMean_zero : empMean IT.action IT.reward a s ω = 0 := by simp only [empMean, ← hm_def, hm0, Nat.cast_zero, div_zero] - simp only [hm0, Nat.cast_zero, h_empMean_zero] at hω - exfalso - have h_mu := hm a e - simp only [max_eq_left (zero_le_one' ℝ), div_one] at hω + simp only [hm0, Nat.cast_zero, h_empMean_zero, max_eq_left (zero_le_one' ℝ), div_one] at hω rw [h_mean, zero_sub, abs_neg] at hω - have h_abs : |(κ (a, e))[id]| ≤ max |lo| |hi| := by - apply abs_le.mpr - constructor - · calc -max |lo| |hi| ≤ -|lo| := neg_le_neg (le_max_left _ _) - _ ≤ lo := neg_abs_le lo - _ ≤ (κ (a, e))[id] := h_mu.1 - · calc (κ (a, e))[id] ≤ hi := h_mu.2 - _ ≤ |hi| := le_abs_self _ - _ ≤ max |lo| |hi| := le_max_right _ _ - linarith + linarith [abs_le_max_abs_abs (hm a e).1 (hm a e).2] · -- Case: m ≥ 1 use m refine ⟨⟨Nat.lt_succ_of_le hms, hm0⟩, rfl, ?_⟩ simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] - have hm_pos : (0 : ℝ) < m := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hm0) have hmax_eq : (max 1 (m : ℕ) : ℝ) = m := by simp only [Nat.one_le_cast, Nat.one_le_iff_ne_zero.mpr hm0, max_eq_right] rw [hmax_eq] at hω - have h_empMean : empMean IT.action IT.reward a s ω = - sumRewards IT.action IT.reward a s ω / m := by - simp only [empMean, hm_def] - rw [h_empMean] at hω + rw [show empMean IT.action IT.reward a s ω = + sumRewards IT.action IT.reward a s ω / m from by simp only [empMean, hm_def]] at hω by_cases h_le : sumRewards IT.action IT.reward a s ω / m ≤ (ν a)[id] - · left - have habs : |sumRewards IT.action IT.reward a s ω / ↑m - (ν a)[id]| = - (ν a)[id] - sumRewards IT.action IT.reward a s ω / m := by - rw [abs_of_nonpos (sub_nonpos.mpr h_le), neg_sub] - rw [habs] at hω - linarith - · right - have h_gt : (ν a)[id] < sumRewards IT.action IT.reward a s ω / m := not_le.mp h_le - have habs : |sumRewards IT.action IT.reward a s ω / ↑m - (ν a)[id]| = - sumRewards IT.action IT.reward a s ω / m - (ν a)[id] := - abs_of_pos (sub_pos.mpr h_gt) - rw [habs] at hω - linarith + · left; rw [abs_of_nonpos (sub_nonpos.mpr h_le), neg_sub] at hω; linarith + · right; rw [abs_of_pos (sub_pos.mpr (not_le.mp h_le))] at hω; linarith calc P' badSet ≤ P' (⋃ m ∈ (Finset.range (s + 1)).filter (· ≠ 0), {ω | pullCount IT.action a s ω = m ∧ @@ -619,17 +601,11 @@ lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := measure_biUnion_finset_le _ _ _ ≤ ∑ m ∈ (Finset.range (s + 1)).filter (· ≠ 0), - streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := by - apply Finset.sum_le_sum - intro m hm - have hm0 : m ≠ 0 := (Finset.mem_filter.mp hm).2 - have hms : m ≤ s := Nat.lt_succ_iff.mp (Finset.mem_range.mp (Finset.mem_filter.mp hm).1) - exact h_bound_per_m m hm0 hms - _ ≤ ∑ _m ∈ (Finset.range (s + 1)).filter (· ≠ 0), ENNReal.ofReal (2 * δ) := by - apply Finset.sum_le_sum - intro m hm - have hm0 : m ≠ 0 := (Finset.mem_filter.mp hm).2 - exact h_stream_bound m hm0 + streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := + Finset.sum_le_sum fun m hm ↦ h_bound_per_m m (Finset.mem_filter.mp hm).2 + (Nat.lt_succ_iff.mp (Finset.mem_range.mp (Finset.mem_filter.mp hm).1)) + _ ≤ ∑ _m ∈ (Finset.range (s + 1)).filter (· ≠ 0), ENNReal.ofReal (2 * δ) := + Finset.sum_le_sum fun m hm ↦ h_stream_bound m (Finset.mem_filter.mp hm).2 _ = ((Finset.range (s + 1)).filter (· ≠ 0)).card • ENNReal.ofReal (2 * δ) := by simp only [Finset.sum_const] _ = s • ENNReal.ofReal (2 * δ) := by @@ -703,6 +679,118 @@ lemma prob_concentration_single_delta [Nonempty (Fin K)] rw [lintegral_const, Measure.map_apply h.measurable_E MeasurableSet.univ] simp [measure_univ] +private lemma concentration_cond_bound [Nonempty (Fin K)] + {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) + {n : ℕ} (hn : 0 < n) {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) + (e : E) (h_isAlgEnvSeq : IsAlgEnvSeq IT.action IT.reward (tsAlgorithm hK Q κ) + (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e)) + (a : Fin K) : + (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + (⋃ s ∈ Finset.range n, {ω | pullCount IT.action a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ + |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|}) ≤ + ENNReal.ofReal (2 * n * δ) := by + let ν := κ.comap (·, e) (by fun_prop) + let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e + have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by + simp only [ν, Kernel.comap_apply]; exact hs a' e + have h_mean : (ν a)[id] = (κ (a, e))[id] := by simp only [ν, Kernel.comap_apply] + let B_low := fun m : ℕ ↦ + {x : ℝ | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} + let B_high := fun m : ℕ ↦ + {x : ℝ | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)} + have h_stream_bound : ∀ m : ℕ, m ≠ 0 → + streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ≤ + ENNReal.ofReal (2 * δ) := + fun m hm0 ↦ streamMeasure_concentration_bound hσ2 h_subG a hδ hδ1 m hm0 + have hB_meas : ∀ m, MeasurableSet (B_low m ∪ B_high m) := fun m ↦ + MeasurableSet.union (measurableSet_le (by fun_prop) (by fun_prop)) + (measurableSet_le (by fun_prop) (by fun_prop)) + let badSetIT := fun (s : ℕ) ↦ {ω : ℕ → (Fin K) × ℝ | + pullCount IT.action a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ + |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} + let S := Finset.Icc 1 (n - 1) + have hS_card : S.card = n - 1 := by simp only [Nat.card_Icc, S]; omega + have h_decomp : ⋃ s ∈ Finset.range n, badSetIT s = + ⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by + ext ω + simp only [Set.mem_iUnion, Finset.mem_range, exists_prop, badSetIT, Set.mem_setOf_eq, + Finset.mem_Icc, S] + constructor + · rintro ⟨s, hs, hbad⟩ + let m := pullCount IT.action a s ω + have hm_pos : 0 < m := Nat.pos_of_ne_zero hbad.1 + have hm_le : m ≤ n - 1 := by + have h1 : m ≤ s := pullCount_le (A := IT.action) a s ω + omega + refine ⟨m, ⟨hm_pos, hm_le⟩, s, hs, rfl, ?_⟩ + simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] + simp only [empMean] at hbad + rcases le_abs'.mp hbad.2 with h | h <;> [left; right] <;> linarith + · rintro ⟨m, ⟨hm_pos, hm_le⟩, s, hs, hpc, hB⟩ + refine ⟨s, hs, ?_⟩ + simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] at hB + simp only [empMean, hpc] + refine ⟨Nat.one_le_iff_ne_zero.mp hm_pos, ?_⟩ + rcases hB with h | h + · exact le_abs.mpr (.inr (by linarith)) + · exact le_abs.mpr (.inl (by linarith)) + rw [h_decomp] + calc P' (⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m}) + ≤ ∑ m ∈ S, P' {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := + measure_biUnion_finset_le S _ + _ ≤ ∑ m ∈ S, streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := by + apply Finset.sum_le_sum + intro m hm + have hm_pos : m ≠ 0 := Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1 + have hm_le : m ≤ n - 1 := (Finset.mem_Icc.mp hm).2 + have h_contain : {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} ⊆ + {ω | pullCount IT.action a (n - 1) ω = m ∧ + sumRewards IT.action IT.reward a (n - 1) ω ∈ B_low m ∪ B_high m} ∪ + {ω | pullCount IT.action a (n - 1) ω > m ∧ + ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by + intro ω ⟨s, hs, hpc, hB⟩ + simp only [Set.mem_union, Set.mem_setOf_eq] + have hs' : s ≤ n - 1 := Nat.le_sub_one_of_lt hs + have h_pc_mono := pullCount_mono (A := IT.action) a hs' ω + by_cases h_eq : pullCount IT.action a (n - 1) ω = m + · left + refine ⟨h_eq, ?_⟩ + have h_pc_eq : pullCount IT.action a s ω = pullCount IT.action a (n - 1) ω := + hpc.symm ▸ h_eq.symm + rw [← sumRewards_eq_of_pullCount_eq h_pc_eq] + exact hB + · right + exact ⟨by omega, s, hs, hpc, hB⟩ + calc P' {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} + ≤ P' {ω | ∃ s, s ≤ n - 1 ∧ pullCount IT.action a s ω = m ∧ + sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by + apply measure_mono + intro ω ⟨s, hs, hpc, hB⟩ + exact ⟨s, Nat.le_sub_one_of_lt hs, hpc, hB⟩ + _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := + prob_exists_pullCount_eq_and_sumRewards_mem_le (n := n - 1) + h_isAlgEnvSeq (hB_meas m) + _ ≤ ∑ _m ∈ S, ENNReal.ofReal (2 * δ) := + Finset.sum_le_sum fun m hm ↦ + h_stream_bound m (Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1) + _ = (n - 1) • ENNReal.ofReal (2 * δ) := by + simp only [Finset.sum_const, hS_card] + _ ≤ ENNReal.ofReal (2 * n * δ) := by + rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast (n - 1), + ← ENNReal.ofReal_mul (Nat.cast_nonneg (n - 1))] + exact ENNReal.ofReal_le_ofReal (by + nlinarith [(Nat.cast_le (α := ℝ)).mpr (Nat.sub_le n 1), hδ.le]) + lemma prob_concentration_fail_delta [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) @@ -755,143 +843,8 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) (⋃ s ∈ Finset.range n, badSetIT s e) ≤ ENNReal.ofReal (2 * n * δ) := by filter_upwards [IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h] with e h_isAlgEnvSeq - let ν := κ.comap (·, e) (by fun_prop) - let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e - have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by - simp only [ν, Kernel.comap_apply]; exact hs a' e - have h_mean : (ν a)[id] = (κ (a, e))[id] := by simp only [ν, Kernel.comap_apply] - let B_low := fun m : ℕ ↦ - {x : ℝ | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} - let B_high := fun m : ℕ ↦ - {x : ℝ | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)} - have h_stream_bound : ∀ m : ℕ, m ≠ 0 → - streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ≤ - ENNReal.ofReal (2 * δ) := by - intro m hm0 - calc streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} - ≤ streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m} + - streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_high m} := by - have h_union : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ⊆ - {ω | ∑ i ∈ range m, ω i a ∈ B_low m} ∪ - {ω | ∑ i ∈ range m, ω i a ∈ B_high m} := by - intro ω hω; simp only [Set.mem_setOf_eq, Set.mem_union] at hω ⊢; exact hω - exact (measure_mono h_union).trans (measure_union_le _ _) - _ ≤ ENNReal.ofReal δ + ENNReal.ofReal δ := by - gcongr - · have h_eq : {ω : ℕ → Fin K → ℝ | - ∑ i ∈ range m, ω i a ∈ B_low m} = - {ω | (∑ i ∈ range m, ω i a) / m + - √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} := by - ext ω; simp only [Set.mem_setOf_eq, B_low] - rw [h_eq] - exact streamMeasure_concentration_le_delta hσ2 h_subG a m hm0 δ hδ hδ1 - · have h_eq : {ω : ℕ → Fin K → ℝ | - ∑ i ∈ range m, ω i a ∈ B_high m} = - {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - - √(2 * ↑σ2 * Real.log (1 / δ) / m)} := by - ext ω; simp only [Set.mem_setOf_eq, B_high] - rw [h_eq] - exact streamMeasure_concentration_ge_delta hσ2 h_subG a m hm0 δ hδ hδ1 - _ = ENNReal.ofReal (2 * δ) := by - rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf - have hB_meas : ∀ m, MeasurableSet (B_low m ∪ B_high m) := fun m ↦ - MeasurableSet.union (measurableSet_le (by fun_prop) (by fun_prop)) - (measurableSet_le (by fun_prop) (by fun_prop)) - let S := Finset.Icc 1 (n - 1) - have hS_card : S.card = n - 1 := by simp only [Nat.card_Icc, S]; omega - have h_decomp : ⋃ s ∈ Finset.range n, badSetIT s e = - ⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by - ext ω - simp only [Set.mem_iUnion, Finset.mem_range, exists_prop, badSetIT, Set.mem_setOf_eq, - Finset.mem_Icc, S] - constructor - · rintro ⟨s, hs, hbad⟩ - let m := pullCount IT.action a s ω - have hm_pos : 0 < m := Nat.pos_of_ne_zero hbad.1 - have hm_le : m ≤ n - 1 := by - have h1 : m ≤ s := pullCount_le (A := IT.action) a s ω - omega - refine ⟨m, ⟨hm_pos, hm_le⟩, s, hs, rfl, ?_⟩ - simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] - have h_pc_pos : (0 : ℝ) < m := Nat.cast_pos.mpr hm_pos - simp only [empMean] at hbad - have h_abs := le_abs'.mp hbad.2 - rcases h_abs with h_neg | h_pos - · left; linarith - · right; linarith - · rintro ⟨m, ⟨hm_pos, hm_le⟩, s, hs, hpc, hB⟩ - refine ⟨s, hs, ?_⟩ - simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] at hB - have h_pc_pos : (0 : ℝ) < m := Nat.cast_pos.mpr hm_pos - simp only [empMean, hpc] - refine ⟨Nat.one_le_iff_ne_zero.mp hm_pos, ?_⟩ - cases hB with - | inl h => - have h1 : sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id] ≤ - -√(2 * ↑σ2 * Real.log (1 / δ) / m) := by linarith - calc √(2 * ↑σ2 * Real.log (1 / δ) / m) - ≤ -(sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id]) := by linarith - _ ≤ |sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id]| := neg_le_abs _ - | inr h => - have h1 : √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ - sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id] := by linarith - exact h1.trans (le_abs_self _) - rw [h_decomp] - calc P' (⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m}) - ≤ ∑ m ∈ S, P' {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := - measure_biUnion_finset_le S _ - _ ≤ ∑ m ∈ S, streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := by - apply Finset.sum_le_sum - intro m hm - have hm_pos : m ≠ 0 := Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1 - have hm_le : m ≤ n - 1 := (Finset.mem_Icc.mp hm).2 - have h_contain : {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} ⊆ - {ω | pullCount IT.action a (n - 1) ω = m ∧ - sumRewards IT.action IT.reward a (n - 1) ω ∈ B_low m ∪ B_high m} ∪ - {ω | pullCount IT.action a (n - 1) ω > m ∧ - ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by - intro ω ⟨s, hs, hpc, hB⟩ - simp only [Set.mem_union, Set.mem_setOf_eq] - have hs' : s ≤ n - 1 := Nat.le_sub_one_of_lt hs - have h_pc_mono := pullCount_mono (A := IT.action) a hs' ω - by_cases h_eq : pullCount IT.action a (n - 1) ω = m - · left - refine ⟨h_eq, ?_⟩ - have h_pc_eq : pullCount IT.action a s ω = pullCount IT.action a (n - 1) ω := - hpc.symm ▸ h_eq.symm - rw [← sumRewards_eq_of_pullCount_eq h_pc_eq] - exact hB - · right - exact ⟨by omega, s, hs, hpc, hB⟩ - calc P' {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} - ≤ P' {ω | ∃ s, s ≤ n - 1 ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by - apply measure_mono - intro ω ⟨s, hs, hpc, hB⟩ - exact ⟨s, Nat.le_sub_one_of_lt hs, hpc, hB⟩ - _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := - prob_exists_pullCount_eq_and_sumRewards_mem_le (n := n - 1) - h_isAlgEnvSeq (hB_meas m) - _ ≤ ∑ _m ∈ S, ENNReal.ofReal (2 * δ) := by - apply Finset.sum_le_sum - intro m hm - have hm_pos : m ≠ 0 := Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1 - exact h_stream_bound m hm_pos - _ = S.card • ENNReal.ofReal (2 * δ) := by simp only [Finset.sum_const] - _ = (n - 1) • ENNReal.ofReal (2 * δ) := by rw [hS_card] - _ ≤ ENNReal.ofReal (2 * n * δ) := by - rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast (n - 1), - ← ENNReal.ofReal_mul (Nat.cast_nonneg (n - 1))] - apply ENNReal.ofReal_le_ofReal - have h1 : (n - 1 : ℕ) ≤ n := Nat.sub_le n 1 - have h2 : (↑(n - 1) : ℝ) ≤ (↑n : ℝ) := Nat.cast_le.mpr h1 - nlinarith [h2, hδ.le] + exact concentration_cond_bound (hK := hK) (E' := E') (A := A) (R' := R') + (Q := Q) (κ := κ) (P := P) hσ2 hs hn' hδ hδ1 e h_isAlgEnvSeq a have h_kernel : Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (a, p.1))[id]) := stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk measurable_fst) have h_meas_set : MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | @@ -931,8 +884,8 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] simp [measure_univ] calc P (⋃ a : Fin K, ⋃ s ∈ Finset.range n, badSet s a) ≤ ∑ a : Fin K, P (⋃ s ∈ Finset.range n, badSet s a) := measure_iUnion_fintype_le _ _ - _ ≤ ∑ _a : Fin K, ENNReal.ofReal (2 * n * δ) := by - apply Finset.sum_le_sum; intro a _; exact h_arm_bound a + _ ≤ ∑ _a : Fin K, ENNReal.ofReal (2 * n * δ) := + Finset.sum_le_sum fun a _ ↦ h_arm_bound a _ = K • ENNReal.ofReal (2 * n * δ) := by simp [Finset.sum_const] _ = ENNReal.ofReal (2 * K * n * δ) := by simp only [nsmul_eq_mul] @@ -982,143 +935,8 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] (⋃ s ∈ Finset.range n, badSetIT a s e) ≤ ENNReal.ofReal (2 * n * δ) := by filter_upwards [IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h] with e h_isAlgEnvSeq intro a - let ν := κ.comap (·, e) (by fun_prop) - let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e - have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by - simp only [ν, Kernel.comap_apply]; exact hs a' e - have h_mean : (ν a)[id] = (κ (a, e))[id] := by simp only [ν, Kernel.comap_apply] - let B_low := fun m : ℕ ↦ - {x : ℝ | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} - let B_high := fun m : ℕ ↦ - {x : ℝ | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)} - have h_stream_bound : ∀ m : ℕ, m ≠ 0 → - streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ≤ - ENNReal.ofReal (2 * δ) := by - intro m hm0 - calc streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} - ≤ streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m} + - streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_high m} := by - have h_union : {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ⊆ - {ω | ∑ i ∈ range m, ω i a ∈ B_low m} ∪ - {ω | ∑ i ∈ range m, ω i a ∈ B_high m} := by - intro ω hω; simp only [Set.mem_setOf_eq, Set.mem_union] at hω ⊢; exact hω - exact (measure_mono h_union).trans (measure_union_le _ _) - _ ≤ ENNReal.ofReal δ + ENNReal.ofReal δ := by - gcongr - · have h_eq : {ω : ℕ → Fin K → ℝ | - ∑ i ∈ range m, ω i a ∈ B_low m} = - {ω | (∑ i ∈ range m, ω i a) / m + - √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} := by - ext ω; simp only [Set.mem_setOf_eq, B_low] - rw [h_eq] - exact streamMeasure_concentration_le_delta hσ2 h_subG a m hm0 δ hδ hδ1 - · have h_eq : {ω : ℕ → Fin K → ℝ | - ∑ i ∈ range m, ω i a ∈ B_high m} = - {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - - √(2 * ↑σ2 * Real.log (1 / δ) / m)} := by - ext ω; simp only [Set.mem_setOf_eq, B_high] - rw [h_eq] - exact streamMeasure_concentration_ge_delta hσ2 h_subG a m hm0 δ hδ hδ1 - _ = ENNReal.ofReal (2 * δ) := by - rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf - have hB_meas : ∀ m, MeasurableSet (B_low m ∪ B_high m) := fun m ↦ - MeasurableSet.union (measurableSet_le (by fun_prop) (by fun_prop)) - (measurableSet_le (by fun_prop) (by fun_prop)) - let S := Finset.Icc 1 (n - 1) - have hS_card : S.card = n - 1 := by simp only [Nat.card_Icc, S]; omega - have h_decomp : ⋃ s ∈ Finset.range n, badSetIT a s e = - ⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by - ext ω - simp only [Set.mem_iUnion, Finset.mem_range, exists_prop, badSetIT, Set.mem_setOf_eq, - Finset.mem_Icc, S] - constructor - · rintro ⟨s, hs, hbad⟩ - let m := pullCount IT.action a s ω - have hm_pos : 0 < m := Nat.pos_of_ne_zero hbad.1 - have hm_le : m ≤ n - 1 := by - have h1 : m ≤ s := pullCount_le (A := IT.action) a s ω - omega - refine ⟨m, ⟨hm_pos, hm_le⟩, s, hs, rfl, ?_⟩ - simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] - have h_pc_pos : (0 : ℝ) < m := Nat.cast_pos.mpr hm_pos - simp only [empMean] at hbad - have h_abs := le_abs'.mp hbad.2 - rcases h_abs with h_neg | h_pos - · left; linarith - · right; linarith - · rintro ⟨m, ⟨hm_pos, hm_le⟩, s, hs, hpc, hB⟩ - refine ⟨s, hs, ?_⟩ - simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] at hB - have h_pc_pos : (0 : ℝ) < m := Nat.cast_pos.mpr hm_pos - simp only [empMean, hpc] - refine ⟨Nat.one_le_iff_ne_zero.mp hm_pos, ?_⟩ - cases hB with - | inl h => - have h1 : sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id] ≤ - -√(2 * ↑σ2 * Real.log (1 / δ) / m) := by linarith - calc √(2 * ↑σ2 * Real.log (1 / δ) / m) - ≤ -(sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id]) := by linarith - _ ≤ |sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id]| := neg_le_abs _ - | inr h => - have h1 : √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ - sumRewards IT.action IT.reward a s ω / m - (κ (a, e))[id] := by linarith - exact h1.trans (le_abs_self _) - rw [h_decomp] - calc P' (⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m}) - ≤ ∑ m ∈ S, P' {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := - measure_biUnion_finset_le S _ - _ ≤ ∑ m ∈ S, streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := by - apply Finset.sum_le_sum - intro m hm - have hm_pos : m ≠ 0 := Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1 - have hm_le : m ≤ n - 1 := (Finset.mem_Icc.mp hm).2 - have h_contain : {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} ⊆ - {ω | pullCount IT.action a (n - 1) ω = m ∧ - sumRewards IT.action IT.reward a (n - 1) ω ∈ B_low m ∪ B_high m} ∪ - {ω | pullCount IT.action a (n - 1) ω > m ∧ - ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by - intro ω ⟨s, hs, hpc, hB⟩ - simp only [Set.mem_union, Set.mem_setOf_eq] - have hs' : s ≤ n - 1 := Nat.le_sub_one_of_lt hs - have h_pc_mono := pullCount_mono (A := IT.action) a hs' ω - by_cases h_eq : pullCount IT.action a (n - 1) ω = m - · left - refine ⟨h_eq, ?_⟩ - have h_pc_eq : pullCount IT.action a s ω = pullCount IT.action a (n - 1) ω := - hpc.symm ▸ h_eq.symm - rw [← sumRewards_eq_of_pullCount_eq h_pc_eq] - exact hB - · right - exact ⟨by omega, s, hs, hpc, hB⟩ - calc P' {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} - ≤ P' {ω | ∃ s, s ≤ n - 1 ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by - apply measure_mono - intro ω ⟨s, hs, hpc, hB⟩ - exact ⟨s, Nat.le_sub_one_of_lt hs, hpc, hB⟩ - _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := - prob_exists_pullCount_eq_and_sumRewards_mem_le (n := n - 1) - h_isAlgEnvSeq (hB_meas m) - _ ≤ ∑ _m ∈ S, ENNReal.ofReal (2 * δ) := by - apply Finset.sum_le_sum - intro m hm - have hm_pos : m ≠ 0 := Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1 - exact h_stream_bound m hm_pos - _ = S.card • ENNReal.ofReal (2 * δ) := by simp only [Finset.sum_const] - _ = (n - 1) • ENNReal.ofReal (2 * δ) := by rw [hS_card] - _ ≤ ENNReal.ofReal (2 * n * δ) := by - rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast (n - 1), - ← ENNReal.ofReal_mul (Nat.cast_nonneg (n - 1))] - apply ENNReal.ofReal_le_ofReal - have h1 : (n - 1 : ℕ) ≤ n := Nat.sub_le n 1 - have h2 : (↑(n - 1) : ℝ) ≤ (↑n : ℝ) := Nat.cast_le.mpr h1 - nlinarith [h2, hδ.le] + exact concentration_cond_bound (hK := hK) (E' := E') (A := A) (R' := R') + (Q := Q) (κ := κ) (P := P) hσ2 hs hn' hδ hδ1 e h_isAlgEnvSeq a have h_cond_best : ∀ᵐ e ∂(P.map E'), (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) (⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ e) s e) ≤ @@ -1240,22 +1058,19 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)] + P[fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)] := by - have hC : ∀ a e, |(κ (a, e))[id]| ≤ max |lo| |hi| := fun a e ↦ - abs_le_max_abs_abs (hm a e).1 (hm a e).2 - have h_regret_gap := bayesRegret_eq_sum_integral_gap (h := h) (hm := hC) (t := n) have h_regret_eq : IsBayesAlgEnvSeq.bayesRegret κ A E' P n = ∑ s ∈ range n, P[fun ω ↦ armMean (bestArm ω) ω - armMean (A s ω) ω] := by - rw [h_regret_gap]; congr 1 with s + rw [bayesRegret_eq_sum_integral_gap (h := h) + (hm := fun a e ↦ abs_le_max_abs_abs (hm a e).1 (hm a e).2) (t := n)] + congr 1 with s exact integral_congr_ae (ae_of_all _ fun ω ↦ gap_eq_armMean_sub E' A κ hm s ω) - have h_int_ucb_sub : ∀ s, Integrable (fun ω ↦ ucb (A s ω) s ω - ucb (bestArm ω) s ω) P := by - intro s - apply Integrable.sub - · exact ⟨(measurable_apply_fin (fun a ↦ hm_ucb a s) (h.measurable_A s)).aestronglyMeasurable, - HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ - norm_ucbIndex_le A R' (↑σ2) lo hi δ hlo _ _ _)⟩ - · exact ⟨(measurable_apply_fin (fun a ↦ hm_ucb a s) hm_best).aestronglyMeasurable, - HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ - norm_ucbIndex_le A R' (↑σ2) lo hi δ hlo _ _ _)⟩ + have h_int_ucb : ∀ s {f : Ω → Fin K}, Measurable f → + Integrable (fun ω ↦ ucb (f ω) s ω) P := fun s {_} hf ↦ + ⟨(measurable_apply_fin (fun a ↦ hm_ucb a s) hf).aestronglyMeasurable, + HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ + norm_ucbIndex_le A R' (↑σ2) lo hi δ hlo _ _ _)⟩ + have h_int_ucb_sub : ∀ s, Integrable (fun ω ↦ ucb (A s ω) s ω - ucb (bestArm ω) s ω) P := + fun s ↦ (h_int_ucb s (h.measurable_A s)).sub (h_int_ucb s hm_best) have h_ucb_zero : ∀ a (ω : Ω), ucbIndex A R' (↑σ2) lo hi δ a 0 ω = hi := by intro a ω; unfold ucbIndex; simp [pullCount_zero] have h_ucb_swap : ∀ s, ∫ ω, (ucb (A s ω) s ω - ucb (bestArm ω) s ω) ∂P = 0 := by @@ -1289,11 +1104,8 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp fun p ↦ if pullCount' t p.1 p.2 = 0 then hi else max lo (min hi (empMean' t p.1 p.2 + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount' t p.1 p.2 : ℝ)))) - have h_hist_eq : ∀ (ω : Ω), - (fun (i : Iic t) ↦ (A (↑i) ω, - R' (↑i) ω)) = - IsAlgEnvSeq.hist A R' t ω := by - intro ω; rfl + have h_hist_eq : ∀ ω : Ω, (fun (i : Iic t) ↦ (A (↑i) ω, R' (↑i) ω)) = + IsAlgEnvSeq.hist A R' t ω := fun ω ↦ rfl have hg_eq : ∀ a (ω : Ω), ucbIndex A R' (↑σ2) lo hi δ a (t + 1) ω = g (IsAlgEnvSeq.hist A R' t ω, a) := by intro a ω @@ -1323,24 +1135,13 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp exact measurable_apply_fin (fun a ↦ measurable_from_top.comp ((measurable_pullCount' t a).comp measurable_fst)) measurable_snd - have h_eq_g1 : (fun ω ↦ ucb (A (t + 1) ω) (t + 1) ω) = - fun ω ↦ g (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω) := - funext fun ω ↦ hg_eq _ _ - have h_eq_g2 : (fun ω ↦ ucb (bestArm ω) (t + 1) ω) = - fun ω ↦ g (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestArm κ E' ω) := - funext fun ω ↦ hg_eq _ _ - have h_int_ucb : ∀ {f : Ω → Fin K}, Measurable f → - Integrable (fun ω ↦ ucb (f ω) (t + 1) ω) P := fun hf ↦ - ⟨(measurable_apply_fin (fun a ↦ hm_ucb a (t + 1)) hf).aestronglyMeasurable, - HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ - norm_ucbIndex_le A R' (↑σ2) lo hi δ hlo _ _ _)⟩ - have h_int1 := h_int_ucb (h.measurable_A (t + 1)) - have h_int2 := h_int_ucb hm_best rw [show (fun ω ↦ ucb (A (t + 1) ω) (t + 1) ω - ucb (bestArm ω) (t + 1) ω) = fun ω ↦ (fun ω ↦ ucb (A (t + 1) ω) (t + 1) ω) ω - (fun ω ↦ ucb (bestArm ω) (t + 1) ω) ω from rfl, - integral_sub h_int1 h_int2, h_eq_g1, h_eq_g2, + integral_sub (h_int_ucb (t + 1) (h.measurable_A (t + 1))) + (h_int_ucb (t + 1) hm_best), + funext fun ω ↦ hg_eq _ _, funext fun ω ↦ hg_eq _ _, h_int_eq g hg_meas, sub_self] have h_ucb_sum_zero : ∫ ω, ∑ s ∈ range n, (ucb (A s ω) s ω - ucb (bestArm ω) s ω) ∂P = 0 := by @@ -1408,18 +1209,15 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp |empMean A R' a s ω - armMean a ω| < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} := by intro s a - have h_eq : {ω : Ω | pullCount A a s ω ≠ 0 → + have : {ω : Ω | pullCount A a s ω ≠ 0 → |empMean A R' a s ω - armMean a ω| < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} = - {ω | pullCount A a s ω = 0} ∪ {ω | + {ω | (pullCount A a s ω : ℝ) = 0} ∪ {ω | |empMean A R' a s ω - armMean a ω| < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} := by - ext ω; simp only [Set.mem_setOf_eq, Set.mem_union]; tauto - rw [h_eq] - have h_eq0 : {ω : Ω | pullCount A a s ω = 0} = - {ω : Ω | (pullCount A a s ω : ℝ) = 0} := by - ext ω; simp [Nat.cast_eq_zero] - exact MeasurableSet.union (h_eq0 ▸ hm_pc a s (measurableSet_singleton (0 : ℝ))) + ext ω; simp only [Set.mem_setOf_eq, Set.mem_union, Nat.cast_eq_zero]; tauto + rw [this] + exact MeasurableSet.union (hm_pc a s (measurableSet_singleton (0 : ℝ))) (measurableSet_lt ((hm_emp a s).sub (h.measurable_armMean a)).abs ((measurable_const.div (hm_pc a s)).sqrt)) @@ -1458,8 +1256,6 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp set f2 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω) set B := (hi - lo) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) - have s1 := (integral_add_compl hFδ_meas h_int_sum1).symm - have s2 := (integral_add_compl hEδ_meas h_int_sum2).symm have h1g : ∫ ω in Fδ, f1 ω ∂P ≤ 0 := setIntegral_nonpos hFδ_meas fun ω hω ↦ h_first_Fδ ω hω have h1b : ∫ ω in Fδᶜ, f1 ω ∂P ≤ ↑n * (hi - lo) * P.real Fδᶜ := by @@ -1482,15 +1278,14 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp ENNReal.toReal_le_of_le_ofReal (by positivity) h_prob_F have hPE : P.real Eδᶜ ≤ 2 * ↑K * ↑n * δ := ENNReal.toReal_le_of_le_ofReal (by positivity) h_prob - rw [s1, s2] - have hPF0 : 0 ≤ P.real Fδᶜ := by positivity - have hPE0 : 0 ≤ P.real Eδᶜ := by positivity - have hhi_lo : (0 : ℝ) ≤ hi - lo := sub_nonneg.mpr hlo - have h_key_F := mul_le_mul_of_nonneg_left hPF - (mul_nonneg (Nat.cast_nonneg n) hhi_lo) - have h_key_E := mul_le_mul_of_nonneg_left hPE - (mul_nonneg (Nat.cast_nonneg n) hhi_lo) - nlinarith + rw [(integral_add_compl hFδ_meas h_int_sum1).symm, + (integral_add_compl hEδ_meas h_int_sum2).symm] + nlinarith [mul_le_mul_of_nonneg_left hPF + (mul_nonneg (Nat.cast_nonneg n) (sub_nonneg.mpr hlo)), + mul_le_mul_of_nonneg_left hPE + (mul_nonneg (Nat.cast_nonneg n) (sub_nonneg.mpr hlo)), + measureReal_nonneg (μ := P) (s := Fδᶜ), + measureReal_nonneg (μ := P) (s := Eδᶜ)] lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) From fa035fe9b08b678df62bc3d4032bad6271997b1c Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 20 Feb 2026 20:30:09 +0000 Subject: [PATCH 043/155] Refactor BayesStationaryEnv (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 249 ++++----- LeanBandits/ForMathlib/HasCondDistrib.lean | 2 +- .../BayesStationaryEnv.lean | 472 +++++++----------- .../SequentialLearning/HistoryDensity.lean | 214 ++++---- 4 files changed, 416 insertions(+), 521 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index b9e655ed..281140c5 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -155,7 +155,7 @@ lemma sum_inv_sqrt_max_one_le (N : ℕ) : @[fun_prop] lemma measurable_ucbIndex [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) (σ2 lo hi δ : ℝ) (a : Fin K) (t : ℕ) : Measurable (ucbIndex A R' σ2 lo hi δ a t) := by unfold ucbIndex @@ -172,12 +172,12 @@ omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ lemma armMean_le_ucbIndex {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hconc : pullCount A a t ω ≠ 0 → - |empMean A R' a t ω - IsBayesAlgEnvSeq.armMean κ E' a ω| + |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E' a ω| < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ))) : - IsBayesAlgEnvSeq.armMean κ E' a ω ≤ ucbIndex A R' σ2 lo hi δ a t ω := by + IsBayesAlgEnvSeq.actionMean κ E' a ω ≤ ucbIndex A R' σ2 lo hi δ a t ω := by unfold ucbIndex have hmean := hm a (E' ω) - simp only [IsBayesAlgEnvSeq.armMean] at hmean hconc ⊢ + simp only [IsBayesAlgEnvSeq.actionMean] at hmean hconc ⊢ split_ifs with h0 · exact hmean.2 · have habs := abs_sub_lt_iff.mp (hconc h0) @@ -188,12 +188,12 @@ omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ lemma ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hpc : pullCount A a t ω ≠ 0) (hconc : - |empMean A R' a t ω - IsBayesAlgEnvSeq.armMean κ E' a ω| + |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E' a ω| < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ))) : - ucbIndex A R' σ2 lo hi δ a t ω - IsBayesAlgEnvSeq.armMean κ E' a ω + ucbIndex A R' σ2 lo hi δ a t ω - IsBayesAlgEnvSeq.actionMean κ E' a ω ≤ 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ)) := by unfold ucbIndex - simp only [IsBayesAlgEnvSeq.armMean] at hconc ⊢ + simp only [IsBayesAlgEnvSeq.actionMean] at hconc ⊢ rw [if_neg hpc] set w := √(2 * σ2 * Real.log (1 / δ) / ↑(pullCount A a t ω)) set emp := empMean A R' a t ω @@ -204,51 +204,52 @@ lemma ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ ( linarith [habs.2] lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) (t : ℕ) : + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) (t : ℕ) : condDistrib (A (t + 1)) (IsAlgEnvSeq.hist A R' t) P =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - condDistrib (IsBayesAlgEnvSeq.bestArm κ E') (IsAlgEnvSeq.hist A R' t) P := + condDistrib (IsBayesAlgEnvSeq.bestAction κ E') (IsAlgEnvSeq.hist A R' t) P := (h.hasCondDistrib_action' t).condDistrib_eq.trans (posteriorBestArm_eq_uniform Q κ h hK t).symm omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma le_armMean_bestArm [Nonempty (Fin K)] (ω : Ω) (i : Fin K) : - IsBayesAlgEnvSeq.armMean κ E' i ω ≤ - IsBayesAlgEnvSeq.armMean κ E' (IsBayesAlgEnvSeq.bestArm κ E' ω) ω := by - have := isMaxOn_measurableArgmax (fun ω a ↦ IsBayesAlgEnvSeq.armMean κ E' a ω) ω i - simp only [IsBayesAlgEnvSeq.bestArm]; convert this + IsBayesAlgEnvSeq.actionMean κ E' i ω ≤ + IsBayesAlgEnvSeq.actionMean κ E' (IsBayesAlgEnvSeq.bestAction κ E' ω) ω := by + have := isMaxOn_measurableArgmax (fun ω a ↦ IsBayesAlgEnvSeq.actionMean κ E' a ω) ω i + simp only [IsBayesAlgEnvSeq.bestAction]; convert this omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc lo hi) - (ω : Ω) : ⨆ i, IsBayesAlgEnvSeq.armMean κ E' i ω = - IsBayesAlgEnvSeq.armMean κ E' (IsBayesAlgEnvSeq.bestArm κ E' ω) ω := + (ω : Ω) : ⨆ i, IsBayesAlgEnvSeq.actionMean κ E' i ω = + IsBayesAlgEnvSeq.actionMean κ E' (IsBayesAlgEnvSeq.bestAction κ E' ω) ω := le_antisymm (ciSup_le (le_armMean_bestArm E' κ ω)) - (le_ciSup (f := fun i ↦ IsBayesAlgEnvSeq.armMean κ E' i ω) + (le_ciSup (f := fun i ↦ IsBayesAlgEnvSeq.actionMean κ E' i ω) ⟨hi, by rintro _ ⟨i, rfl⟩; exact (hm i _).2⟩ _) omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma gap_eq_armMean_sub [Nonempty (Fin K)] {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc lo hi) (s : ℕ) (ω : Ω) : gap (κ.comap (·, E' ω) (by fun_prop)) (A s ω) = - IsBayesAlgEnvSeq.armMean κ E' (IsBayesAlgEnvSeq.bestArm κ E' ω) ω - - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω := by + IsBayesAlgEnvSeq.actionMean κ E' (IsBayesAlgEnvSeq.bestAction κ E' ω) ω - + IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω := by simp only [gap, Kernel.comap_apply] exact congr_arg (· - _) (iSup_armMean_eq_bestArm E' κ hm ω) -omit [StandardBorelSpace E] [Nonempty E] in +omit [StandardBorelSpace E] [Nonempty E] [IsProbabilityMeasure Q] [IsMarkovKernel κ] in lemma bayesRegret_eq_sum_integral_gap [Nonempty (Fin K)] {alg : Algorithm (Fin K) ℝ} - (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) + (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) {C : ℝ} (hm : ∀ a e, |(κ (a, e))[id]| ≤ C) (t : ℕ) : - IsBayesAlgEnvSeq.bayesRegret κ A E' P t = + P[IsBayesAlgEnvSeq.regret κ E' A t] = ∑ s ∈ range t, P[fun ω ↦ gap (κ.comap (·, E' ω) (by fun_prop)) (A s ω)] := by - simp only [IsBayesAlgEnvSeq.bayesRegret, IsBayesAlgEnvSeq.regret, regret_eq_sum_gap] + simp only [IsBayesAlgEnvSeq.regret, regret_eq_sum_gap] refine integral_finset_sum _ (fun s _ => ?_) have hmeas : Measurable (fun ω ↦ gap (κ.comap (·, E' ω) (by fun_prop)) (A s ω)) := - (Measurable.iSup h.measurable_armMean).sub + (Measurable.iSup (fun a ↦ IsBayesAlgEnvSeq.measurable_actionMean + (a := a) h.measurable_E)).sub (stronglyMeasurable_id.integral_kernel.measurable.comp ((h.measurable_A s).prodMk h.measurable_E)) refine ⟨hmeas.aestronglyMeasurable, HasFiniteIntegral.of_bounded (C := 2 * C) @@ -289,10 +290,10 @@ omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsProbabilityMeas lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) (hlo : lo ≤ hi) (σ2 δ : ℝ) (n : ℕ) (ω : Ω) (hconc : ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → - |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω| + |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω| < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ))) : ∑ s ∈ range n, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) + IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω) ≤ (hi - lo) * ↑K + 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by -- Split range n into first-pull (pc=0) and non-first-pull (pc≠0) sets set S0 := (range n).filter (fun s => pullCount A (A s ω) s ω = 0) @@ -303,16 +304,16 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] rw [Finset.sum_union hdisj] -- We bound ∑_{S0} and ∑_{S1} separately, then combine suffices h_S0 : ∑ s ∈ S0, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) ≤ (hi - lo) * ↑K by + IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω) ≤ (hi - lo) * ↑K by suffices h_S1 : ∑ s ∈ S1, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) + IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω) ≤ 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) by have := Finset.sum_union hdisj (f := fun s => - ucbIndex A R' σ2 lo hi δ (A s ω) s ω - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) + ucbIndex A R' σ2 lo hi δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω) rw [← hpart] at this; linarith -- Bound ∑_{S1}: each term ≤ 2√(2σ2c/pc) = 2√(2σ2c/max(1,pc)), so ≤ full sum calc ∑ s ∈ S1, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) + IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω) ≤ ∑ s ∈ S1, 2 * √(2 * σ2 * Real.log (1 / δ) / (max 1 (pullCount A (A s ω) s ω) : ℝ)) := sum_le_sum fun s hs => by @@ -369,9 +370,9 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] rw [sum_congr rfl h0]; simp only [sum_const_zero]; positivity -- Bound ∑_{S0}: each term = hi - armMean ≤ hi - lo, and #S0 ≤ K have hterm_S0 : ∀ s ∈ S0, ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω ≤ hi - lo := fun s hs => by + IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω ≤ hi - lo := fun s hs => by have hpc : pullCount A (A s ω) s ω = 0 := (Finset.mem_filter.mp hs).2 - simp only [ucbIndex, hpc, ↓reduceIte, IsBayesAlgEnvSeq.armMean] + simp only [ucbIndex, hpc, ↓reduceIte, IsBayesAlgEnvSeq.actionMean] linarith [(hm (A s ω) (E' ω)).1] have h_card_S0 : #S0 ≤ K := by calc #S0 ≤ #(Finset.univ : Finset (Fin K)) := @@ -392,7 +393,7 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] Finset.card_ne_zero_of_mem this)) _ = K := Finset.card_fin K calc ∑ s ∈ S0, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.armMean κ E' (A s ω) ω) + IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω) ≤ ∑ _s ∈ S0, (hi - lo) := sum_le_sum hterm_S0 _ = #S0 * (hi - lo) := by rw [sum_const, nsmul_eq_mul] _ ≤ ↑K * (hi - lo) := by @@ -523,24 +524,28 @@ private lemma streamMeasure_concentration_bound {α : Type*} [MeasurableSpace α rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc lo hi) (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) (hδ_large : max |lo| |hi| < √(2 * ↑σ2 * Real.log (1 / δ))) : ∀ᵐ e ∂(P.map (E')), - (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) {ω | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} ≤ ENNReal.ofReal (2 * s * δ) := by - filter_upwards [IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h] with e h_isAlgEnvSeq + have h_cond_ae : ∀ᵐ e ∂(P.map E'), IsAlgEnvSeq IT.action IT.reward + (tsAlgorithm hK Q κ) (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E' P e) := by + rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h + filter_upwards [h_cond_ae] with e h_isAlgEnvSeq let ν := κ.comap (·, e) (by fun_prop) have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by simp only [ν, Kernel.comap_apply]; exact hs a' e have h_mean : (ν a)[id] = (κ (a, e))[id] := by simp only [ν, Kernel.comap_apply] rw [← h_mean] - let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e + let P' := condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e have h_law := h_isAlgEnvSeq.law_pullCount_sumRewards_unique' (ArrayModel.isAlgEnvSeq_arrayMeasure (tsAlgorithm hK Q κ) ν) (n := s) let B_low := fun m : ℕ ↦ @@ -618,38 +623,40 @@ lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] congr 1; ring lemma prob_concentration_single_delta [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc lo hi) (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) (hδ_large : max |lo| |hi| < √(2 * ↑σ2 * Real.log (1 / δ))) : P {ω | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} ≤ + |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω|} ≤ ENNReal.ofReal (2 * s * δ) := by let badSet : E → Set (ℕ → (Fin K) × ℝ) := fun e ↦ {t | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s t) : ℝ)) ≤ |empMean IT.action IT.reward a s t - (κ (a, e))[id]|} have h_set_eq : {ω | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} = - (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω|} = + (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' {p | p.2 ∈ badSet p.1} := by ext ω - simp only [Set.mem_setOf_eq, Set.mem_preimage, badSet, IsBayesAlgEnvSeq.armMean] - have h1 : pullCount A a s ω = pullCount IT.action a s (IsBayesAlgEnvSeq.traj A R' ω) := by - unfold pullCount IsBayesAlgEnvSeq.traj IT.action; rfl + simp only [Set.mem_setOf_eq, Set.mem_preimage, badSet, IsBayesAlgEnvSeq.actionMean] + have h1 : pullCount A a s ω = pullCount IT.action a s ((fun ω n => (A n ω, R' n ω)) ω) := by + unfold pullCount IT.action; rfl have h2 : empMean A R' a s ω = - empMean IT.action IT.reward a s (IsBayesAlgEnvSeq.traj A R' ω) := by - unfold empMean IsBayesAlgEnvSeq.traj IT.action IT.reward; rfl + empMean IT.action IT.reward a s ((fun ω n => (A n ω, R' n ω)) ω) := by + unfold empMean IT.action IT.reward; rfl rw [h1, h2] have h_meas_pair : - Measurable (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) := - h.measurable_E.prodMk h.measurable_traj - have h_disint : P.map (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) = + Measurable (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) := + h.measurable_E.prodMk (measurable_pi_lambda _ fun n => + (h.measurable_A n).prodMk (h.measurable_R n)) + have h_disint : P.map (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) = P.map (E') ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.traj A R') E' P := - (compProd_map_condDistrib (h.measurable_traj.aemeasurable)).symm + condDistrib ((fun ω n => (A n ω, R' n ω))) E' P := + (compProd_map_condDistrib ((measurable_pi_lambda _ fun n => + (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable)).symm have h_cond := prob_concentration_single_delta_cond hK E' A R' Q κ P h hσ2 hs hm a s δ hδ hδ1 hδ_large have h_kernel : Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (a, p.1))[id]) := @@ -661,15 +668,15 @@ lemma prob_concentration_single_delta [Nonempty (Fin K)] exact measurableSet_le (by fun_prop) (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp measurable_snd).sub h_kernel).abs - calc P _ = P ((fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + calc P _ = P ((fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' {p | p.2 ∈ badSet p.1}) := by rw [h_set_eq] - _ = (P.map (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω))) + _ = (P.map (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω))) {p | p.2 ∈ badSet p.1} := by rw [Measure.map_apply h_meas_pair h_meas_set] _ = (P.map (E') ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.traj A R') E' P) + condDistrib ((fun ω n => (A n ω, R' n ω))) E' P) {p | p.2 ∈ badSet p.1} := by rw [h_disint] - _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + _ = ∫⁻ e, (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) (badSet e) ∂(P.map (E')) := by rw [Measure.compProd_apply h_meas_set]; rfl _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * s * δ) ∂(P.map (E')) := by @@ -685,15 +692,15 @@ private lemma concentration_cond_bound [Nonempty (Fin K)] {n : ℕ} (hn : 0 < n) {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) (e : E) (h_isAlgEnvSeq : IsAlgEnvSeq IT.action IT.reward (tsAlgorithm hK Q κ) (stationaryEnv (κ.comap (·, e) (by fun_prop))) - (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e)) + (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e)) (a : Fin K) : - (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) (⋃ s ∈ Finset.range n, {ω | pullCount IT.action a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|}) ≤ ENNReal.ofReal (2 * n * δ) := by let ν := κ.comap (·, e) (by fun_prop) - let P' := condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e + let P' := condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by simp only [ν, Kernel.comap_apply]; exact hs a' e have h_mean : (ν a)[id] = (κ (a, e))[id] := by simp only [ν, Kernel.comap_apply] @@ -792,20 +799,20 @@ private lemma concentration_cond_bound [Nonempty (Fin K)] nlinarith [(Nat.cast_le (α := ℝ)).mpr (Nat.sub_le n 1), hδ.le]) lemma prob_concentration_fail_delta [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : P {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} + |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω|} ≤ ENNReal.ofReal (2 * K * n * δ) := by let badSet := fun (s : ℕ) (a : Fin K) ↦ {ω : Ω | pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} + |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω|} have h_set_eq : {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.armMean κ E' a ω|} = + |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω|} = ⋃ s ∈ Finset.range n, ⋃ a : Fin K, badSet s a := by ext ω; simp only [Set.mem_setOf_eq, Finset.mem_range, Set.mem_iUnion, badSet, exists_prop] rw [h_set_eq] @@ -825,24 +832,30 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} have h_set_eq : ⋃ s ∈ Finset.range n, badSet s a = - (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by ext ω simp only [Set.mem_iUnion, Finset.mem_range, badSet, badSetIT, Set.mem_preimage, - Set.mem_setOf_eq, IsBayesAlgEnvSeq.armMean] + Set.mem_setOf_eq, IsBayesAlgEnvSeq.actionMean] exact Iff.rfl rw [h_set_eq] have h_meas_pair : - Measurable (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) := - h.measurable_E.prodMk h.measurable_traj - have h_disint : P.map (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) = + Measurable (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) := + h.measurable_E.prodMk (measurable_pi_lambda _ fun n => + (h.measurable_A n).prodMk (h.measurable_R n)) + have h_disint : P.map (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) = P.map (E') ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.traj A R') E' P := - (compProd_map_condDistrib (h.measurable_traj.aemeasurable)).symm + condDistrib ((fun ω n => (A n ω, R' n ω))) E' P := + (compProd_map_condDistrib ((measurable_pi_lambda _ fun n => + (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable)).symm have h_cond_bound : ∀ᵐ e ∂(P.map (E')), - (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) (⋃ s ∈ Finset.range n, badSetIT s e) ≤ ENNReal.ofReal (2 * n * δ) := by - filter_upwards [IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h] with e h_isAlgEnvSeq + have h_cond_ae : ∀ᵐ e ∂(P.map E'), IsAlgEnvSeq IT.action IT.reward + (tsAlgorithm hK Q κ) (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E' P e) := by + rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h + filter_upwards [h_cond_ae] with e h_isAlgEnvSeq exact concentration_cond_bound (hK := hK) (E' := E') (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) hσ2 hs hn' hδ hδ1 e h_isAlgEnvSeq a have h_kernel : Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (a, p.1))[id]) := @@ -865,16 +878,16 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] (measurableSet_le (by fun_prop) (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp measurable_snd).sub h_kernel).abs) - calc P ((fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + calc P ((fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1}) - = (P.map (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω))) + = (P.map (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω))) {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by rw [Measure.map_apply h_meas_pair h_meas_set] _ = (P.map (E') ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.traj A R') E' P) + condDistrib ((fun ω n => (A n ω, R' n ω))) E' P) {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by rw [h_disint] - _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + _ = ∫⁻ e, (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) (⋃ s ∈ Finset.range n, badSetIT s e) ∂(P.map (E')) := by rw [Measure.compProd_apply h_meas_set]; rfl _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * n * δ) ∂(P.map (E')) := by @@ -893,21 +906,21 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] congr 1; ring lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - P {ω | ∃ s < n, pullCount A (IsBayesAlgEnvSeq.bestArm κ E' ω) s ω ≠ 0 ∧ + P {ω | ∃ s < n, pullCount A (IsBayesAlgEnvSeq.bestAction κ E' ω) s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / - (pullCount A (IsBayesAlgEnvSeq.bestArm κ E' ω) s ω : ℝ)) ≤ - |empMean A R' (IsBayesAlgEnvSeq.bestArm κ E' ω) s ω - - IsBayesAlgEnvSeq.armMean κ E' (IsBayesAlgEnvSeq.bestArm κ E' ω) ω|} + (pullCount A (IsBayesAlgEnvSeq.bestAction κ E' ω) s ω : ℝ)) ≤ + |empMean A R' (IsBayesAlgEnvSeq.bestAction κ E' ω) s ω - + IsBayesAlgEnvSeq.actionMean κ E' (IsBayesAlgEnvSeq.bestAction κ E' ω) ω|} ≤ ENNReal.ofReal (2 * n * δ) := by by_cases hn : n = 0 · simp [hn] have hn' : 0 < n := Nat.pos_of_ne_zero hn - rw [show IsBayesAlgEnvSeq.bestArm κ E' = envToBestArm κ ∘ E' from - bestArm_eq_envToBestArm_comp_env κ] + rw [show IsBayesAlgEnvSeq.bestAction κ E' = envToBestArm κ ∘ E' from + bestAction_eq_envToBestArm_comp_env κ] let badSetIT := fun (a : Fin K) (s : ℕ) (e : E) ↦ {ω : ℕ → (Fin K) × ℝ | pullCount IT.action a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ @@ -916,29 +929,35 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A ((envToBestArm κ ∘ E') ω) s ω : ℝ)) ≤ |empMean A R' ((envToBestArm κ ∘ E') ω) s ω - - IsBayesAlgEnvSeq.armMean κ E' ((envToBestArm κ ∘ E') ω) ω|} = - (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + IsBayesAlgEnvSeq.actionMean κ E' ((envToBestArm κ ∘ E') ω) ω|} = + (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1} := by ext ω simp only [Set.mem_setOf_eq, Finset.mem_range, Set.mem_preimage, Set.mem_iUnion, - badSetIT, IsBayesAlgEnvSeq.armMean, Function.comp_apply, exists_prop] + badSetIT, IsBayesAlgEnvSeq.actionMean, Function.comp_apply, exists_prop] rfl rw [h_set_eq] have h_meas_pair : - Measurable (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) := - h.measurable_E.prodMk h.measurable_traj - have h_disint : P.map (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) = - P.map E' ⊗ₘ condDistrib (IsBayesAlgEnvSeq.traj A R') E' P := - (compProd_map_condDistrib (h.measurable_traj.aemeasurable)).symm + Measurable (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) := + h.measurable_E.prodMk (measurable_pi_lambda _ fun n => + (h.measurable_A n).prodMk (h.measurable_R n)) + have h_disint : P.map (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) = + P.map E' ⊗ₘ condDistrib ((fun ω n => (A n ω, R' n ω))) E' P := + (compProd_map_condDistrib ((measurable_pi_lambda _ fun n => + (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable)).symm have h_cond_bound : ∀ᵐ e ∂(P.map E'), ∀ a : Fin K, - (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) (⋃ s ∈ Finset.range n, badSetIT a s e) ≤ ENNReal.ofReal (2 * n * δ) := by - filter_upwards [IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h] with e h_isAlgEnvSeq + have h_cond_ae : ∀ᵐ e ∂(P.map E'), IsAlgEnvSeq IT.action IT.reward + (tsAlgorithm hK Q κ) (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E' P e) := by + rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h + filter_upwards [h_cond_ae] with e h_isAlgEnvSeq intro a exact concentration_cond_bound (hK := hK) (E' := E') (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) hσ2 hs hn' hδ hδ1 e h_isAlgEnvSeq a have h_cond_best : ∀ᵐ e ∂(P.map E'), - (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) (⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ e) s e) ≤ ENNReal.ofReal (2 * n * δ) := by filter_upwards [h_cond_bound] with e he @@ -975,16 +994,16 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] exact .iUnion fun a ↦ .inter ((measurable_envToBestArm (κ := κ) |>.comp measurable_fst) (measurableSet_singleton a)) (.biUnion (Finset.range n).countable_toSet fun s _ ↦ h_meas_badSetIT a s) - calc P ((fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω)) ⁻¹' + calc P ((fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1}) - = (P.map (fun ω ↦ (E' ω, IsBayesAlgEnvSeq.traj A R' ω))) + = (P.map (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω))) {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1} := by rw [Measure.map_apply h_meas_pair h_meas_set] _ = (P.map E' ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.traj A R') E' P) + condDistrib ((fun ω n => (A n ω, R' n ω))) E' P) {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1} := by rw [h_disint] - _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) + _ = ∫⁻ e, (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) (⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ e) s e) ∂(P.map E') := by rw [Measure.compProd_apply h_meas_set]; rfl _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * n * δ) ∂(P.map E') := by @@ -994,18 +1013,18 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] simp [measure_univ] lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - IsBayesAlgEnvSeq.bayesRegret κ A E' P n + P[IsBayesAlgEnvSeq.regret κ E' A n] ≤ (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * n ^ 2 * δ + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _) have hlo : lo ≤ hi := h1.trans h2 - let bestArm := IsBayesAlgEnvSeq.bestArm κ E' - let armMean := IsBayesAlgEnvSeq.armMean κ E' + let bestArm := IsBayesAlgEnvSeq.bestAction κ E' + let armMean := IsBayesAlgEnvSeq.actionMean κ E' let ucb := ucbIndex A R' (↑σ2) lo hi δ set Eδ : Set Ω := {ω | ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → |empMean A R' a s ω - armMean a ω| @@ -1015,9 +1034,10 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A (bestArm ω) s ω : ℝ))} have hm_ucb : ∀ a t, Measurable (ucbIndex A R' (↑σ2) lo hi δ a t) := fun a t ↦ measurable_ucbIndex hK E' A R' Q κ P h (↑σ2) lo hi δ a t - have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.armMean κ E' a) := - fun a ↦ h.measurable_armMean a - have hm_best : Measurable (IsBayesAlgEnvSeq.bestArm κ E') := h.measurable_bestArm + have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.actionMean κ E' a) := + fun a ↦ IsBayesAlgEnvSeq.measurable_actionMean (a := a) h.measurable_E + have hm_best : Measurable (IsBayesAlgEnvSeq.bestAction κ E') := + IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E have h_first_bound : ∀ ω, |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| ≤ n * (hi - lo) := fun ω ↦ @@ -1053,12 +1073,12 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (measurable_apply_fin hm_arm (h.measurable_A s))).aestronglyMeasurable · filter_upwards with ω; rw [Real.norm_eq_abs]; exact h_second_bound ω have h_swap : - IsBayesAlgEnvSeq.bayesRegret κ A E' P n = + P[IsBayesAlgEnvSeq.regret κ E' A n] = P[fun ω ↦ ∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)] + P[fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)] := by - have h_regret_eq : IsBayesAlgEnvSeq.bayesRegret κ A E' P n = + have h_regret_eq : P[IsBayesAlgEnvSeq.regret κ E' A n] = ∑ s ∈ range n, P[fun ω ↦ armMean (bestArm ω) ω - armMean (A s ω) ω] := by rw [bayesRegret_eq_sum_integral_gap (h := h) (hm := fun a e ↦ abs_le_max_abs_abs (hm a e).1 (hm a e).2) (t := n)] @@ -1084,13 +1104,13 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp | succ t => have hts := ts_identity hK E' A R' Q κ P h t have h_map_eq : P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω)) = - P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestArm κ E' ω)) := by + P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E' ω)) := by rw [← compProd_map_condDistrib (hY := (h.measurable_A (t + 1)).aemeasurable), ← compProd_map_condDistrib (hY := hm_best.aemeasurable)] exact Measure.compProd_congr hts have h_int_eq : ∀ (f : (Iic t → Fin K × ℝ) × Fin K → ℝ), Measurable f → ∫ ω, f (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω) ∂P = - ∫ ω, f (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestArm κ E' ω) ∂P := by + ∫ ω, f (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E' ω) ∂P := by intro f hf have hm_hist := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t rw [← integral_map @@ -1219,7 +1239,8 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp rw [this] exact MeasurableSet.union (hm_pc a s (measurableSet_singleton (0 : ℝ))) (measurableSet_lt - ((hm_emp a s).sub (h.measurable_armMean a)).abs + ((hm_emp a s).sub + (IsBayesAlgEnvSeq.measurable_actionMean (a := a) h.measurable_E)).abs ((measurable_const.div (hm_pc a s)).sqrt)) have hEδ_meas : MeasurableSet Eδ := by simp only [Eδ, Set.setOf_forall] @@ -1288,24 +1309,24 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp measureReal_nonneg (μ := P) (s := Eδᶜ)] lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ E' A R' (tsAlgorithm hK Q κ) P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) (t : ℕ) : - IsBayesAlgEnvSeq.bayesRegret κ A E' P t + P[IsBayesAlgEnvSeq.regret κ E' A t] ≤ (3 * K + 2) * (hi - lo) + 8 * √(σ2 * K * t * Real.log t) := by have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _) have hlo : lo ≤ hi := h1.trans h2 by_cases ht : t = 0 - · simp [ht, IsBayesAlgEnvSeq.bayesRegret, IsBayesAlgEnvSeq.regret, regret] + · simp [ht, IsBayesAlgEnvSeq.regret, Bandits.regret] nlinarith [sub_nonneg.mpr hlo, show (0 : ℝ) < K from Nat.cast_pos.mpr hK, Real.sqrt_nonneg (↑σ2 * ↑K * (0 : ℝ) * Real.log (0 : ℝ))] by_cases ht1_eq : t = 1 · subst ht1_eq simp only [Nat.cast_one, Real.log_one, mul_zero, Real.sqrt_zero, mul_zero, add_zero] - calc IsBayesAlgEnvSeq.bayesRegret κ A E' P 1 + calc P[IsBayesAlgEnvSeq.regret κ E' A 1] ≤ hi - lo := by - unfold IsBayesAlgEnvSeq.bayesRegret IsBayesAlgEnvSeq.regret Bandits.regret + unfold IsBayesAlgEnvSeq.regret Bandits.regret simp only [Finset.range_one, Finset.sum_singleton, Nat.cast_one, one_mul, Kernel.comap_apply] refine (integral_mono_of_nonneg (ae_of_all _ fun ω ↦ sub_nonneg.mpr @@ -1334,7 +1355,7 @@ lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] -- Second term simplification: log(1/(1/t²)) = log(t²) = 2 log(t) have h_log : Real.log (1 / (1 / (↑t : ℝ) ^ 2)) = 2 * Real.log ↑t := by rw [one_div_one_div, Real.log_pow]; norm_cast - calc IsBayesAlgEnvSeq.bayesRegret κ A E' P t + calc P[IsBayesAlgEnvSeq.regret κ E' A t] ≤ (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * ↑t ^ 2 * (1 / (↑t) ^ 2) + 2 * √(8 * ↑σ2 * Real.log (1 / (1 / (↑t) ^ 2))) * √(↑K * ↑t) := bayesRegret_le_of_delta (hK := hK) (E' := E') (A := A) (R' := R') (Q := Q) diff --git a/LeanBandits/ForMathlib/HasCondDistrib.lean b/LeanBandits/ForMathlib/HasCondDistrib.lean index 18e78a84..5d427c54 100644 --- a/LeanBandits/ForMathlib/HasCondDistrib.lean +++ b/LeanBandits/ForMathlib/HasCondDistrib.lean @@ -184,7 +184,7 @@ lemma hasCondDistrib_prod_right_iff [IsFiniteMeasure μ] [IsFiniteKernel κ] (X -- Claude lemma HasCondDistrib.hasLaw_of_const {Q : Measure Ω} - [IsProbabilityMeasure μ] [IsProbabilityMeasure Q] + [IsProbabilityMeasure μ] [IsFiniteMeasure Q] (h : HasCondDistrib Y X (Kernel.const β Q) μ) : HasLaw Y Q μ where aemeasurable := h.aemeasurable_fst map_eq := by diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index ca0d947f..f2a5a5b8 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -3,10 +3,9 @@ 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.MeasurableArgMax import LeanBandits.Bandit.Regret +import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.SequentialLearning.StationaryEnv -import Mathlib.Probability.Kernel.Posterior /-! # Bayesian stationary environments -/ @@ -14,255 +13,84 @@ open MeasureTheory ProbabilityTheory Finset namespace Learning -variable {α E R : Type*} [mα : MeasurableSpace α] [mE : MeasurableSpace E] [mR : MeasurableSpace R] - -/-- Given a prior distribution `Q` over "environments" and a kernel `k` that defines a reward -distribution `κ (a, e)` for each action `a : α` and "environment" `e : E`, a `bayesStationaryEnv` -corresponds to an environment (with an observation space `E × R`) that draws an "environment" -`e : E` at the very first step and defines a stationary environment from `k (·, e)`. -/ -noncomputable -def bayesStationaryEnv - (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] : - Environment α (E × R) where - feedback n := - let g : (Iic n → α × E × R) × α → α × E := fun (h, a) => (a, (h ⟨0, by simp⟩).2.1) - (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) - ν0 := (Kernel.const α Q) ⊗ₖ κ - -variable {Ω : Type*} [mΩ : MeasurableSpace Ω] +variable {α R 𝓔 : Type*} [MeasurableSpace α] [MeasurableSpace R] [MeasurableSpace 𝓔] +variable {Ω : Type*} [MeasurableSpace Ω] -/-- A Bayesian algorithm-environment sequence: a sequence of actions and observations from an -algorithm that ignores the underlying "environment" while interacting with a Bayesian stationary -environment. The environment `E'` is drawn from a prior `Q`, and rewards follow a kernel `κ` -conditioned on the action and environment. -/ structure IsBayesAlgEnvSeq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] - (E' : Ω → E) (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (alg : Algorithm α R) + (Q : Measure 𝓔) (κ : Kernel (α × 𝓔) R) (alg : Algorithm α R) + (E : Ω → 𝓔) (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (P : Measure Ω) [IsFiniteMeasure P] : Prop where - measurable_E : Measurable E' := by fun_prop + measurable_E : Measurable E := by fun_prop measurable_A n : Measurable (A n) := by fun_prop measurable_R n : Measurable (R' n) := by fun_prop - hasLaw_env : HasLaw E' Q P - hasCondDistrib_action_zero : HasCondDistrib (A 0) E' (Kernel.const _ alg.p0) P - hasCondDistrib_reward_zero : HasCondDistrib (R' 0) (fun ω ↦ (A 0 ω, E' ω)) κ P + hasLaw_env : HasLaw E Q P + hasCondDistrib_action_zero : HasCondDistrib (A 0) E (Kernel.const _ alg.p0) P + hasCondDistrib_reward_zero : HasCondDistrib (R' 0) (fun ω ↦ (A 0 ω, E ω)) κ P hasCondDistrib_action n : - HasCondDistrib (A (n + 1)) (fun ω ↦ (E' ω, IsAlgEnvSeq.hist A R' n ω)) + HasCondDistrib (A (n + 1)) (fun ω ↦ (E ω, IsAlgEnvSeq.hist A R' n ω)) ((alg.policy n).prodMkLeft _) P hasCondDistrib_reward n : - HasCondDistrib (R' (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω, E' ω)) + HasCondDistrib (R' (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω, E ω)) (κ.prodMkLeft _) P namespace IsBayesAlgEnvSeq -variable [StandardBorelSpace α] [Nonempty α] -variable [StandardBorelSpace R] [Nonempty R] - -variable {Q : Measure E} [IsProbabilityMeasure Q] {κ : Kernel (α × E) R} [IsMarkovKernel κ] -variable {E' : Ω → E} {A : ℕ → Ω → α} {R' : ℕ → Ω → R} -variable {alg : Algorithm α R} -variable {P : Measure Ω} [IsProbabilityMeasure P] - -/-- The trajectory of actions and rewards as a function into the IT space. -/ -def traj (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (ω : Ω) : ℕ → α × R := - fun n => (A n ω, R' n ω) - -@[fun_prop] -lemma measurable_traj (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : Measurable (traj A R') := - measurable_pi_lambda _ fun n => (h.measurable_A n).prodMk (h.measurable_R n) - -section Real - -variable {E' : Ω → E} {R' : ℕ → Ω → ℝ} -variable {κ : Kernel (α × E) ℝ} [IsMarkovKernel κ] -variable {alg : Algorithm α ℝ} - -/-- The mean of action `a : α` in the underlying "environment". -/ -noncomputable -def armMean (κ : Kernel (α × E) ℝ) (E' : Ω → E) (a : α) (ω : Ω) : ℝ := (κ (a, E' ω))[id] - -@[fun_prop] -lemma measurable_armMean (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (a : α) : - Measurable (armMean κ E' a) := - stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk h.measurable_E) - -/-- An action with the highest mean in the underlying "environment". -/ -noncomputable -def bestArm [Fintype α] [Encodable α] (κ : Kernel (α × E) ℝ) (E' : Ω → E) := - measurableArgmax (fun ω a ↦ armMean κ E' a ω) - -@[fun_prop] -lemma measurable_bestArm [Fintype α] [Encodable α] (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : - Measurable (bestArm κ E') := - measurable_measurableArgmax h.measurable_armMean - -/-- Regret of a sequence of pulls at time `t` considering the underlying "environment". -/ -noncomputable -def regret (κ : Kernel (α × E) ℝ) (A : ℕ → Ω → α) (E' : Ω → E) (t : ℕ) (ω : Ω) : ℝ := - Bandits.regret (κ.comap (·, E' ω) (by fun_prop)) A t ω - -lemma measurable_regret [Encodable α] (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (t : ℕ) : - Measurable (regret κ A E' t) := by - apply Measurable.sub - · exact Measurable.const_mul (Measurable.iSup h.measurable_armMean) _ - · exact Finset.measurable_sum _ fun s _ ↦ - stronglyMeasurable_id.integral_kernel.measurable.comp - ((h.measurable_A s).prodMk h.measurable_E) - -/-- If `IsBayesAlgEnvSeq Q κ E' A R' alg P`, then `bayesRegret κ A E' P t` is the expected -regret at time `t` of the algorithm `alg` given a prior distribution over "environments" `Q`. -/ -noncomputable -def bayesRegret (κ : Kernel (α × E) ℝ) (A : ℕ → Ω → α) (E' : Ω → E) (P : Measure Ω) - (t : ℕ) : ℝ := - P[regret κ A E' t] - -end Real - section Laws -lemma hasLaw_action_zero (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : +variable [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] +variable {Q : Measure 𝓔} {κ : Kernel (α × 𝓔) R} {alg : Algorithm α R} +variable {E : Ω → 𝓔} {A : ℕ → Ω → α} {R' : ℕ → Ω → R} +variable {P : Measure Ω} [IsFiniteMeasure P] + +lemma hasLaw_action_zero [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : HasLaw (A 0) alg.p0 P := h.hasCondDistrib_action_zero.hasLaw_of_const -lemma indepFun_action_zero_env [StandardBorelSpace E] [Nonempty E] - (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : - IndepFun (A 0) E' P := - ((indepFun_iff_condDistrib_eq_const h.measurable_E.aemeasurable - (h.measurable_A 0).aemeasurable).2 (by - rw [h.hasLaw_action_zero.map_eq]; exact h.hasCondDistrib_action_zero.condDistrib_eq)).symm - -lemma hasCondDistrib_action' (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : +lemma hasCondDistrib_action' (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) (alg.policy n) P := (h.hasCondDistrib_action n).comp_left (by fun_prop) -lemma hasCondDistrib_reward' (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : - HasCondDistrib (R' (n + 1)) (fun ω ↦ (A (n + 1) ω, E' ω)) κ P := +lemma hasCondDistrib_reward' [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : + HasCondDistrib (R' (n + 1)) (fun ω ↦ (A (n + 1) ω, E ω)) κ P := (h.hasCondDistrib_reward n).comp_left (by fun_prop) -end Laws - -section Independence - -lemma condIndepFun_action_env_hist [StandardBorelSpace E] [Nonempty E] - [StandardBorelSpace Ω] (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) - (n : ℕ) : - A (n + 1) ⟂ᵢ[IsAlgEnvSeq.hist A R' n, - IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n; P] E' := - condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft - h.measurable_E (h.measurable_A _) (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n) - (h.hasCondDistrib_action n).condDistrib_eq - -lemma condIndepFun_reward_hist [StandardBorelSpace Ω] - (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : - R' (n + 1) ⟂ᵢ[fun ω ↦ (A (n + 1) ω, E' ω), - (h.measurable_A (n + 1)).prodMk h.measurable_E; P] IsAlgEnvSeq.hist A R' n := - condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft - (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n) (h.measurable_R _) - ((h.measurable_A _).prodMk h.measurable_E) - (h.hasCondDistrib_reward n).condDistrib_eq - -end Independence - -section Posterior - -variable [StandardBorelSpace E] [Nonempty E] - -/-- The posterior on the environment given history equals Mathlib's `posterior` applied to the -likelihood kernel and prior. This is the measure-theoretic formulation of Bayes' rule. -/ -lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : - HasCondDistrib E' (IsAlgEnvSeq.hist A R' n) - (posterior (condDistrib (IsAlgEnvSeq.hist A R' n) E' P) Q) P where - aemeasurable_fst := h.measurable_E.aemeasurable - aemeasurable_snd := - (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable - condDistrib_eq := by - have h_env_meas : Measurable E' := h.measurable_E - have h_hist_meas : Measurable (IsAlgEnvSeq.hist A R' n) := - IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n - set κ' := condDistrib (IsAlgEnvSeq.hist A R' n) E' P with hκ' - have h_disint : P.map (fun ω => (E' ω, IsAlgEnvSeq.hist A R' n ω)) = Q ⊗ₘ κ' := by - rw [← h.hasLaw_env.map_eq, compProd_map_condDistrib (h_hist_meas.aemeasurable)] - have h_marg : P.map (IsAlgEnvSeq.hist A R' n) = κ' ∘ₘ Q := by - have : P.map (IsAlgEnvSeq.hist A R' n) = - (P.map (fun ω => (E' ω, IsAlgEnvSeq.hist A R' n ω))).snd := by - rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)]; rfl - rw [this, h_disint, Measure.snd_compProd] - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ h_env_meas.aemeasurable] - rw [show P.map (fun ω => (IsAlgEnvSeq.hist A R' n ω, E' ω)) = - (Q ⊗ₘ κ').map Prod.swap from by - rw [← h_disint, Measure.map_map (by fun_prop) (by fun_prop)]; rfl] - rw [← compProd_posterior_eq_map_swap (κ := κ') (μ := Q), h_marg] - -end Posterior - -section StationaryEnvConnection - -variable [StandardBorelSpace E] [Nonempty E] - -omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] in -/-- The traj function commutes with IT projections: IT.action n ∘ traj = A n -/ -lemma IT_action_comp_traj (n : ℕ) : IT.action n ∘ traj A R' = A n := by - ext ω; simp [IT.action, traj] - -omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] in -/-- The traj function commutes with IT projections: IT.reward n ∘ traj = R' n -/ -lemma IT_reward_comp_traj (n : ℕ) : IT.reward n ∘ traj A R' = R' n := by - ext ω; simp [IT.reward, traj] - -omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] in -/-- The traj function commutes with IT projections: IT.hist n ∘ traj = IsAlgEnvSeq.hist n -/ -lemma IT_hist_comp_traj (n : ℕ) : - IT.hist n ∘ traj A R' = IsAlgEnvSeq.hist A R' n := by - ext ω i : 2 - simp only [Function.comp_apply, IT.hist, traj, IsAlgEnvSeq.hist] - -omit mα mE mR mΩ [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] in -/-- The pair (IT.hist n, IT.action (n+1)) commutes with traj. -/ -lemma IT_hist_action_comp_traj (n : ℕ) : - (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) ∘ traj A R' = - fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω) := by - ext ω : 1 - simp only [Function.comp_apply, IT.action, traj, Prod.mk.injEq] - exact ⟨funext fun _ => rfl, trivial⟩ - -lemma condDistrib_traj_action_zero (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : - ∀ᵐ e ∂(P.map E'), - (condDistrib (traj A R') E' P e).map (IT.action 0) = alg.p0 := by - have h_comp : condDistrib (IT.action 0 ∘ traj A R') E' P - =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map (IT.action 0) := - condDistrib_comp E' (h.measurable_traj.aemeasurable) (IT.measurable_action 0) - rw [IT_action_comp_traj] at h_comp - filter_upwards [h_comp, condDistrib_of_indepFun h.indepFun_action_zero_env.symm - h.measurable_E.aemeasurable (h.measurable_A 0).aemeasurable] with e h_comp_e h_indep_e - rw [← Kernel.map_apply _ (IT.measurable_action 0), ← h_comp_e, h_indep_e, Kernel.const_apply] - exact h.hasLaw_action_zero.map_eq - -omit [StandardBorelSpace E] [Nonempty E] in -lemma hasCondDistrib_reward_zero_condDistrib (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : - ∀ᵐ e ∂(P.map E'), - HasCondDistrib (IT.reward 0) (IT.action 0) (κ.comap (·, e) (by fun_prop)) - (condDistrib (traj A R') E' P e) := by - have h_swap : HasCondDistrib (R' 0) (fun ω ↦ (E' ω, A 0 ω)) +--- + +lemma hasLaw_action_zero_fiber (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : + ∀ᵐ e ∂Q, HasLaw (IT.action 0) alg.p0 + (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E P e) := by + rw [← h.hasLaw_env.map_eq] + have hW : AEMeasurable (fun ω n ↦ (A n ω, R' n ω)) P := + (measurable_pi_lambda _ fun n ↦ (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable + have h_comp : ⇑(condDistrib (A 0) E P) =ᶠ[ae (P.map E)] + ⇑((condDistrib (fun ω n ↦ (A n ω, R' n ω)) E P).map (IT.action 0)) := + condDistrib_comp E hW (IT.measurable_action 0) + filter_upwards [h_comp, h.hasCondDistrib_action_zero.condDistrib_eq] with e he hcd + exact ⟨(IT.measurable_action 0).aemeasurable, by + rw [← Kernel.map_apply _ (IT.measurable_action 0), ← he, hcd, Kernel.const_apply]⟩ + +lemma hasCondDistrib_reward_zero_fiber [IsFiniteKernel κ] + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : + ∀ᵐ e ∂Q, HasCondDistrib (IT.reward 0) (IT.action 0) + (κ.comap (·, e) (by fun_prop)) + (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E P e) := by + rw [← h.hasLaw_env.map_eq] + set W := fun ω n ↦ (A n ω, R' n ω) + have hW : AEMeasurable W P := + (measurable_pi_lambda _ fun n ↦ (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable + have h_swap : HasCondDistrib (R' 0) (fun ω ↦ (E ω, A 0 ω)) (κ.comap Prod.swap (by fun_prop)) P := by convert h.hasCondDistrib_reward_zero.comp_right - (MeasurableEquiv.prodComm : α × E ≃ᵐ E × α) using 2 + (MeasurableEquiv.prodComm : α × 𝓔 ≃ᵐ 𝓔 × α) using 2 have h_prod := condDistrib_prod_left (h.measurable_A 0).aemeasurable (h.measurable_R 0).aemeasurable h.measurable_E.aemeasurable (μ := P) - have h_comp_pair : condDistrib ((fun ω ↦ (IT.action 0 ω, IT.reward 0 ω)) ∘ traj A R') E' P - =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map - (fun ω ↦ (IT.action 0 ω, IT.reward 0 ω)) := - condDistrib_comp E' h.measurable_traj.aemeasurable (by fun_prop) - have h_comp_action : condDistrib (IT.action 0 ∘ traj A R') E' P - =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map (IT.action 0) := - condDistrib_comp E' h.measurable_traj.aemeasurable (IT.measurable_action 0) - rw [show (fun ω ↦ (IT.action 0 ω, IT.reward 0 ω)) ∘ traj A R' = - fun ω ↦ (A 0 ω, R' 0 ω) from by - ext ω : 1; simp only [Function.comp_apply, IT.action, IT.reward, traj]] at h_comp_pair - rw [IT_action_comp_traj] at h_comp_action + have h_comp_pair : ⇑(condDistrib (fun ω ↦ (A 0 ω, R' 0 ω)) E P) =ᶠ[ae (P.map E)] + ⇑((condDistrib W E P).map (fun ω ↦ (IT.action 0 ω, IT.reward 0 ω))) := + condDistrib_comp E hW ((IT.measurable_action 0).prodMk (IT.measurable_reward 0)) + have h_comp_action : ⇑(condDistrib (A 0) E P) =ᶠ[ae (P.map E)] + ⇑((condDistrib W E P).map (IT.action 0)) := + condDistrib_comp E hW (IT.measurable_action 0) have h_swap_eq := h_swap.condDistrib_eq rw [(compProd_map_condDistrib (h.measurable_A 0).aemeasurable).symm] at h_swap_eq filter_upwards [h_prod, h_comp_pair, h_comp_action, @@ -278,69 +106,77 @@ lemma hasCondDistrib_reward_zero_condDistrib (h : IsBayesAlgEnvSeq Q κ E' A R' ext s _ rw [Kernel.sectR_apply, Kernel.comap_apply, ha, Kernel.comap_apply]; rfl -omit [StandardBorelSpace E] [Nonempty E] in -lemma hasCondDistrib_action_condDistrib (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : - ∀ᵐ e ∂(P.map E'), - HasCondDistrib (IT.action (n + 1)) (IT.hist n) (alg.policy n) - (condDistrib (traj A R') E' P e) := by +lemma hasCondDistrib_action_fiber (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : + ∀ᵐ e ∂Q, HasCondDistrib (IT.action (n + 1)) + (IsAlgEnvSeq.hist IT.action IT.reward n) (alg.policy n) + (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E P e) := by + rw [← h.hasLaw_env.map_eq] + set W := fun ω n ↦ (A n ω, R' n ω) + have hW : AEMeasurable W P := + (measurable_pi_lambda _ fun n ↦ (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable have h_hist_meas := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n have h_prod := condDistrib_prod_left h_hist_meas.aemeasurable (h.measurable_A (n + 1)).aemeasurable h.measurable_E.aemeasurable (μ := P) have h_action_env := (h.hasCondDistrib_action n).condDistrib_eq - have h_comp_pair : condDistrib ((fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) ∘ traj A R') - E' P =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map - (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) := - condDistrib_comp E' h.measurable_traj.aemeasurable (by fun_prop) - have h_comp_hist : condDistrib (IT.hist n ∘ traj A R') E' P - =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map (IT.hist n) := - condDistrib_comp E' h.measurable_traj.aemeasurable (IT.measurable_hist n) - rw [IT_hist_action_comp_traj] at h_comp_pair - rw [IT_hist_comp_traj] at h_comp_hist + have h_hist_IT_meas : Measurable + (IsAlgEnvSeq.hist (IT.action (R := R)) (IT.reward (α := α)) n) := + IsAlgEnvSeq.measurable_hist (fun n ↦ IT.measurable_action n) (fun n ↦ IT.measurable_reward n) n + have h_comp_pair : ⇑(condDistrib (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) E P) + =ᶠ[ae (P.map E)] ⇑((condDistrib W E P).map + (fun ω ↦ (IsAlgEnvSeq.hist IT.action IT.reward n ω, IT.action (n + 1) ω))) := + condDistrib_comp E hW (h_hist_IT_meas.prodMk (IT.measurable_action (n + 1))) + have h_comp_hist : ⇑(condDistrib (IsAlgEnvSeq.hist A R' n) E P) =ᶠ[ae (P.map E)] + ⇑((condDistrib W E P).map (IsAlgEnvSeq.hist IT.action IT.reward n)) := + condDistrib_comp E hW h_hist_IT_meas rw [(compProd_map_condDistrib h_hist_meas.aemeasurable).symm] at h_action_env filter_upwards [h_prod, h_comp_pair, h_comp_hist, (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_action_env] with e h_prod_e h_pair_e h_hist_e h_nested_e refine ⟨by fun_prop, by fun_prop, ?_⟩ rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] - rw [← Kernel.map_apply _ (by fun_prop), ← h_pair_e] - conv_rhs => rw [← Kernel.map_apply _ (IT.measurable_hist n), ← h_hist_e] + rw [← Kernel.map_apply _ (h_hist_IT_meas.prodMk (IT.measurable_action (n + 1))), + ← h_pair_e] + conv_rhs => rw [← Kernel.map_apply _ h_hist_IT_meas, ← h_hist_e] rw [h_prod_e, Kernel.compProd_apply_eq_compProd_sectR] refine Measure.compProd_congr ?_ filter_upwards [h_nested_e] with _ ha ext s _ rw [Kernel.sectR_apply, ha, Kernel.prodMkLeft_apply] -omit [StandardBorelSpace E] [Nonempty E] in -lemma hasCondDistrib_reward_condDistrib (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (n : ℕ) : - ∀ᵐ e ∂(P.map E'), - HasCondDistrib (IT.reward (n + 1)) (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) - ((κ.comap (·, e) (by fun_prop)).prodMkLeft _) - (condDistrib (traj A R') E' P e) := by +lemma hasCondDistrib_reward_fiber [IsFiniteKernel κ] + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : + ∀ᵐ e ∂Q, HasCondDistrib (IT.reward (n + 1)) + (fun f ↦ (IsAlgEnvSeq.hist IT.action IT.reward n f, IT.action (n + 1) f)) + ((κ.comap (·, e) (by fun_prop)).prodMkLeft _) + (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E P e) := by + rw [← h.hasLaw_env.map_eq] + set W := fun ω n ↦ (A n ω, R' n ω) + have hW : AEMeasurable W P := + (measurable_pi_lambda _ fun n ↦ (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable have h_hist_meas := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n have h_prod := condDistrib_prod_left (Measurable.prodMk h_hist_meas (h.measurable_A (n + 1))).aemeasurable (h.measurable_R (n + 1)).aemeasurable h.measurable_E.aemeasurable (μ := P) have h_swap : HasCondDistrib (R' (n + 1)) - (fun ω ↦ (E' ω, IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) + (fun ω ↦ (E ω, IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) (κ.comap (fun p ↦ (p.2.2, p.1)) (by fun_prop)) P := (h.hasCondDistrib_reward n).comp_right (MeasurableEquiv.prodAssoc.symm.trans MeasurableEquiv.prodComm) have h_swap_eq := h_swap.condDistrib_eq - have h_comp_triple : condDistrib - ((fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), IT.reward (n + 1) ω)) ∘ traj A R') E' P - =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map - (fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), IT.reward (n + 1) ω)) := - condDistrib_comp E' h.measurable_traj.aemeasurable (by fun_prop) - have h_comp_pair : condDistrib ((fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) ∘ traj A R') - E' P =ᵐ[P.map E'] (condDistrib (traj A R') E' P).map - (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω)) := - condDistrib_comp E' h.measurable_traj.aemeasurable (by fun_prop) - rw [show (fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), IT.reward (n + 1) ω)) ∘ - traj A R' = fun ω ↦ ((IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω), R' (n + 1) ω) from by - ext ω : 1 - simp only [Function.comp_apply, IT.action, IT.reward, traj, Prod.mk.injEq] - exact ⟨⟨funext fun i => rfl, trivial⟩, trivial⟩] at h_comp_triple - rw [IT_hist_action_comp_traj] at h_comp_pair + have h_hist_IT_meas : Measurable + (IsAlgEnvSeq.hist (IT.action (R := R)) (IT.reward (α := α)) n) := + IsAlgEnvSeq.measurable_hist (fun n ↦ IT.measurable_action n) (fun n ↦ IT.measurable_reward n) n + have h_pair_meas := h_hist_IT_meas.prodMk (IT.measurable_action (n + 1)) + have h_comp_triple : ⇑(condDistrib + (fun ω ↦ ((IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω), R' (n + 1) ω)) E P) + =ᶠ[ae (P.map E)] ⇑((condDistrib W E P).map + (fun ω ↦ ((IsAlgEnvSeq.hist IT.action IT.reward n ω, IT.action (n + 1) ω), + IT.reward (n + 1) ω))) := + condDistrib_comp E hW (h_pair_meas.prodMk (IT.measurable_reward (n + 1))) + have h_comp_pair : ⇑(condDistrib (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) E P) + =ᶠ[ae (P.map E)] ⇑((condDistrib W E P).map + (fun ω ↦ (IsAlgEnvSeq.hist IT.action IT.reward n ω, IT.action (n + 1) ω))) := + condDistrib_comp E hW h_pair_meas rw [(compProd_map_condDistrib (Measurable.prodMk h_hist_meas (h.measurable_A (n + 1))).aemeasurable).symm] at h_swap_eq filter_upwards [h_prod, h_comp_triple, h_comp_pair, @@ -348,45 +184,80 @@ lemma hasCondDistrib_reward_condDistrib (h : IsBayesAlgEnvSeq Q κ E' A R' alg P with e h_prod_e h_triple_e h_pair_e h_nested_e refine ⟨by fun_prop, by fun_prop, ?_⟩ rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] - rw [← Kernel.map_apply _ (by fun_prop), ← h_triple_e] - conv_rhs => rw [← Kernel.map_apply _ (by fun_prop), ← h_pair_e] + rw [← Kernel.map_apply _ (h_pair_meas.prodMk (IT.measurable_reward (n + 1))), ← h_triple_e] + conv_rhs => rw [← Kernel.map_apply _ h_pair_meas, ← h_pair_e] rw [h_prod_e, Kernel.compProd_apply_eq_compProd_sectR] refine Measure.compProd_congr ?_ filter_upwards [h_nested_e] with _ ha ext s _ rw [Kernel.sectR_apply, ha, Kernel.comap_apply, Kernel.prodMkLeft_apply, Kernel.comap_apply] -lemma condDistrib_traj_isAlgEnvSeq (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) : - ∀ᵐ e ∂(P.map E'), - IsAlgEnvSeq IT.action IT.reward alg - (stationaryEnv (κ.comap (·, e) (by fun_prop))) - (condDistrib (traj A R') E' P e) := by - filter_upwards [condDistrib_traj_action_zero h, - hasCondDistrib_reward_zero_condDistrib h, - ae_all_iff.2 (hasCondDistrib_action_condDistrib h), - ae_all_iff.2 (hasCondDistrib_reward_condDistrib h)] with e h_law h_r0 h_a h_r +lemma condDistrib_traj_isAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : + ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E P e) := by + filter_upwards [hasLaw_action_zero_fiber h, + hasCondDistrib_reward_zero_fiber h, + ae_all_iff.2 (hasCondDistrib_action_fiber h), + ae_all_iff.2 (hasCondDistrib_reward_fiber h)] + with _ h_law h_r0 h_a h_r exact { - hasLaw_action_zero := ⟨by fun_prop, h_law⟩ + hasLaw_action_zero := h_law hasCondDistrib_reward_zero := h_r0 hasCondDistrib_action := h_a hasCondDistrib_reward := h_r } -end StationaryEnvConnection +end Laws + +section Real + +noncomputable +def actionMean (κ : Kernel (α × 𝓔) ℝ) (E : Ω → 𝓔) (a : α) (ω : Ω) : ℝ := (κ (a, E ω))[id] + +@[fun_prop] +lemma measurable_actionMean {κ : Kernel (α × 𝓔) ℝ} {E : Ω → 𝓔} {a : α} (hE : Measurable E) : + Measurable (actionMean κ E a) := + stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop) + +noncomputable +def bestAction [Fintype α] [Encodable α] [Nonempty α] [MeasurableSingletonClass α] + (κ : Kernel (α × 𝓔) ℝ) (E : Ω → 𝓔) (ω : Ω) : α := + measurableArgmax (fun ω' a ↦ actionMean κ E a ω') ω + +@[fun_prop] +lemma measurable_bestAction [Fintype α] [Encodable α] [Nonempty α] [MeasurableSingletonClass α] + {κ : Kernel (α × 𝓔) ℝ} {E : Ω → 𝓔} (hE : Measurable E) : Measurable (bestAction κ E) := + measurable_measurableArgmax (by fun_prop) + +noncomputable +def regret (κ : Kernel (α × 𝓔) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → α) (t : ℕ) (ω : Ω) : ℝ := + Bandits.regret (κ.comap (·, E ω) (by fun_prop)) A t ω + +end Real end IsBayesAlgEnvSeq +section StationaryEquivalence + +noncomputable +def bayesStationaryEnv (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) R) + [IsMarkovKernel κ] : Environment α (𝓔 × R) where + feedback n := + let g : (Iic n → α × 𝓔 × R) × α → α × 𝓔 := fun (h, a) => (a, (h ⟨0, by simp⟩).2.1) + (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) + ν0 := (Kernel.const _ Q) ⊗ₖ κ + /-- Bridge theorem: an `IsAlgEnvSeq` for `(alg.prod_left E)` and `(bayesStationaryEnv Q κ)` gives rise to an `IsBayesAlgEnvSeq`. -/ theorem IsAlgEnvSeq.toIsBayesAlgEnvSeq [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace E] [Nonempty E] + [StandardBorelSpace 𝓔] [Nonempty 𝓔] [StandardBorelSpace R] [Nonempty R] - {Q : Measure E} [IsProbabilityMeasure Q] {κ : Kernel (α × E) R} [IsMarkovKernel κ] - {A : ℕ → Ω → α} {R'' : ℕ → Ω → E × R} {alg : Algorithm α R} + {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (α × 𝓔) R} [IsMarkovKernel κ] + {A : ℕ → Ω → α} {R'' : ℕ → Ω → 𝓔 × R} {alg : Algorithm α R} {P : Measure Ω} [IsProbabilityMeasure P] - (h : IsAlgEnvSeq A R'' (alg.prod_left E) (bayesStationaryEnv Q κ) P) : - IsBayesAlgEnvSeq Q κ (fun ω ↦ (R'' 0 ω).1) A (fun n ω ↦ (R'' n ω).2) alg P where + (h : IsAlgEnvSeq A R'' (alg.prod_left 𝓔) (bayesStationaryEnv Q κ) P) : + IsBayesAlgEnvSeq Q κ alg (fun ω ↦ (R'' 0 ω).1) A (fun n ω ↦ (R'' n ω).2) P where measurable_E := (h.measurable_R 0).fst measurable_A := h.measurable_A measurable_R n := (h.measurable_R n).snd @@ -409,14 +280,14 @@ theorem IsAlgEnvSeq.toIsBayesAlgEnvSeq hasCondDistrib_reward_zero := by simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.of_compProd hasCondDistrib_action n := by - let f : (Iic n → α × E × R) → E × (Iic n → α × R) := + let f : (Iic n → α × 𝓔 × R) → 𝓔 × (Iic n → α × R) := fun h ↦ ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2)) suffices h' : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R'' n) (((alg.policy n).comap Prod.snd (by fun_prop)).comap f (by fun_prop)) P from h'.comp_left (f := f) exact h.hasCondDistrib_action n hasCondDistrib_reward n := by - let f : (Iic n → α × E × R) × α → (Iic n → α × R) × α × E := + let f : (Iic n → α × 𝓔 × R) × α → (Iic n → α × R) × α × 𝓔 := fun p ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1) have hf : Measurable f := by fun_prop suffices h' : HasCondDistrib (fun ω ↦ (R'' (n + 1) ω).2) @@ -426,30 +297,28 @@ theorem IsAlgEnvSeq.toIsBayesAlgEnvSeq namespace IT -/-- Measure on the sequence of actions and observations generated by an algorithm that ignores the -underlying "environment" while interacting with a `bayesStationaryEnv`. -/ noncomputable -def bayesTrajMeasure (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) - [IsMarkovKernel κ] (alg : Algorithm α R) : Measure (ℕ → α × E × R) := - trajMeasure (alg.prod_left E) (bayesStationaryEnv Q κ) +def bayesTrajMeasure (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) R) + [IsMarkovKernel κ] (alg : Algorithm α R) : Measure (ℕ → α × 𝓔 × R) := + trajMeasure (alg.prod_left 𝓔) (bayesStationaryEnv Q κ) deriving IsProbabilityMeasure lemma isBayesAlgEnvSeq_bayesianTrajMeasure [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace E] [Nonempty E] + [StandardBorelSpace 𝓔] [Nonempty 𝓔] [StandardBorelSpace R] [Nonempty R] - (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) R) [IsMarkovKernel κ] + (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) R) [IsMarkovKernel κ] (alg : Algorithm α R) : - IsBayesAlgEnvSeq Q κ (fun ω ↦ (ω 0).2.1) action (fun n ω ↦ (ω n).2.2) - alg (bayesTrajMeasure Q κ alg) := + IsBayesAlgEnvSeq Q κ alg (fun ω ↦ (ω 0).2.1) action (fun n ω ↦ (ω n).2.2) + (bayesTrajMeasure Q κ alg) := (isAlgEnvSeq_trajMeasure _ _).toIsBayesAlgEnvSeq /-- The conditional distribution over the best arm given the observed history. -/ noncomputable def posteriorBestArm [StandardBorelSpace α] [Nonempty α] [Fintype α] [Encodable α] - (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] + (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) ℝ) [IsMarkovKernel κ] (alg : Algorithm α ℝ) (n : ℕ) : Kernel (Iic n → α × ℝ) α := - condDistrib (IsBayesAlgEnvSeq.bestArm κ (fun ω ↦ (ω 0).2.1)) + condDistrib (IsBayesAlgEnvSeq.bestAction κ (fun ω ↦ (ω 0).2.1)) (IsAlgEnvSeq.hist action (fun n ω ↦ (ω n).2.2) n) (bayesTrajMeasure Q κ alg) deriving IsMarkovKernel @@ -457,18 +326,17 @@ deriving IsMarkovKernel /-- The initial distribution over the best arm. -/ noncomputable def priorBestArm [StandardBorelSpace α] [Nonempty α] [Fintype α] [Encodable α] - (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] + (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) ℝ) [IsMarkovKernel κ] (alg : Algorithm α ℝ) : Measure α := - (bayesTrajMeasure Q κ alg).map (IsBayesAlgEnvSeq.bestArm κ (fun ω ↦ (ω 0).2.1)) + (bayesTrajMeasure Q κ alg).map (IsBayesAlgEnvSeq.bestAction κ (fun ω ↦ (ω 0).2.1)) -instance [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace E] [Nonempty E] [Fintype α] - [Encodable α] (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (α × E) ℝ) [IsMarkovKernel κ] - (alg : Algorithm α ℝ) : IsProbabilityMeasure (priorBestArm Q κ alg) := - Measure.isProbabilityMeasure_map - (measurable_measurableArgmax - (IsBayesAlgEnvSeq.measurable_armMean - (isBayesAlgEnvSeq_bayesianTrajMeasure Q κ alg))).aemeasurable +instance [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace 𝓔] [Nonempty 𝓔] [Fintype α] + [Encodable α] (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) ℝ) + [IsMarkovKernel κ] (alg : Algorithm α ℝ) : IsProbabilityMeasure (priorBestArm Q κ alg) := + Measure.isProbabilityMeasure_map (by fun_prop) end IT +end StationaryEquivalence + end Learning diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index c3afa83e..00f03cec 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -245,42 +245,6 @@ private lemma absolutelyContinuous_map_hist_stationary end AbsolutelyContinuousHist -section PosteriorEquality - -variable {E' X' : Type*} {mE' : MeasurableSpace E'} {mX' : MeasurableSpace X'} - --- `compProd` unfolding requires extra heartbeats -/-- If `κ₁ =ᵐ[Q] κ₂.withDensity (fun _ => ρ)`, then the posteriors agree: -`κ₂†Q =ᵐ[κ₁ ∘ₘ Q] κ₁†Q`. -/ -private theorem posterior_eq_of_withDensity_ae_eq - [StandardBorelSpace E'] [Nonempty E'] - {Q : Measure E'} [IsFiniteMeasure Q] - {κ₁ κ₂ : Kernel E' X'} [IsFiniteKernel κ₁] [IsFiniteKernel κ₂] - {ρ : X' → ℝ≥0∞} (hρ : Measurable ρ) - [IsSFiniteKernel (κ₂.withDensity (fun _ => ρ))] - [SFinite ((κ₂ ∘ₘ Q).withDensity ρ)] - (h_ae : κ₁ =ᵐ[Q] κ₂.withDensity (fun _ => ρ)) : - κ₂†Q =ᵐ[κ₁ ∘ₘ Q] κ₁†Q := by - apply ae_eq_posterior_of_compProd_eq - have h2 : Q ⊗ₘ (κ₂.withDensity (fun _ => ρ)) - = (Q ⊗ₘ κ₂).withDensity (ρ ∘ Prod.snd) := by - have := Measure.compProd_withDensity (κ := κ₂) (μ := Q) - (show Measurable (Function.uncurry (fun _ => ρ)) from hρ.comp measurable_snd) - convert this using 1 - calc (κ₁ ∘ₘ Q) ⊗ₘ (κ₂†Q) - = ((κ₂ ∘ₘ Q).withDensity ρ) ⊗ₘ (κ₂†Q) := by - congr 1; rw [← comp_withDensity_const hρ]; exact Measure.bind_congr_right h_ae - _ = ((κ₂ ∘ₘ Q) ⊗ₘ (κ₂†Q)).withDensity (ρ ∘ Prod.fst) := - withDensity_compProd_left hρ - _ = ((Q ⊗ₘ κ₂).map Prod.swap).withDensity (ρ ∘ Prod.fst) := by - rw [compProd_posterior_eq_map_swap] - _ = ((Q ⊗ₘ κ₂).withDensity (ρ ∘ Prod.snd)).map Prod.swap := by - rw [map_swap_withDensity_fst hρ] - _ = (Q ⊗ₘ (κ₂.withDensity (fun _ => ρ))).map Prod.swap := by rw [h2] - _ = (Q ⊗ₘ κ₁).map Prod.swap := by rw [Measure.compProd_congr h_ae] - -end PosteriorEquality - section DensityIndependence variable {K : ℕ} [Nonempty (Fin K)] @@ -431,19 +395,19 @@ lemma measurable_envToBestArm : Measurable (envToBestArm κ) := omit [StandardBorelSpace E] [Nonempty E] [IsMarkovKernel κ] [MeasurableSpace Ω] [IsProbabilityMeasure P] [Nonempty Ω] in -lemma bestArm_eq_envToBestArm_comp_env : - IsBayesAlgEnvSeq.bestArm κ E' = envToBestArm κ ∘ E' := by +lemma bestAction_eq_envToBestArm_comp_env : + IsBayesAlgEnvSeq.bestAction κ E' = envToBestArm κ ∘ E' := by funext ω; simp only [Function.comp_apply] - unfold IsBayesAlgEnvSeq.bestArm IsBayesAlgEnvSeq.armMean envToBestArm + unfold IsBayesAlgEnvSeq.bestAction IsBayesAlgEnvSeq.actionMean envToBestArm exact (measurableArgmax_eq_of_eq _ _ _ ω).trans (measurableArgmax_congr _ _ ω _ rfl) -omit [StandardBorelSpace E] [Nonempty E] in +omit [StandardBorelSpace E] [Nonempty E] [IsMarkovKernel κ] in /-- The marginal on the history equals `condDistrib (hist) (env) P ∘ₘ Q`. -/ private lemma map_hist_eq_condDistrib_comp {Ω' : Type*} [MeasurableSpace Ω'] [StandardBorelSpace Ω'] [Nonempty Ω'] {E'' : Ω' → E} {A' : ℕ → Ω' → Fin K} {R'' : ℕ → Ω' → ℝ} {alg' : Algorithm (Fin K) ℝ} {P' : Measure Ω'} [IsProbabilityMeasure P'] - (h' : IsBayesAlgEnvSeq Q κ E'' A' R'' alg' P') (t : ℕ) : + (h' : IsBayesAlgEnvSeq Q κ alg' E'' A' R'' P') (t : ℕ) : P'.map (IsAlgEnvSeq.hist A' R'' t) = condDistrib (IsAlgEnvSeq.hist A' R'' t) E'' P' ∘ₘ Q := by calc P'.map (IsAlgEnvSeq.hist A' R'' t) @@ -458,15 +422,16 @@ private lemma map_hist_eq_condDistrib_comp E'' P').snd := by rw [h'.hasLaw_env.map_eq] _ = _ := Measure.snd_compProd Q _ +omit [StandardBorelSpace E] [Nonempty E] in /-- The history distribution under any algorithm is absolutely continuous w.r.t. the history distribution under the uniform algorithm (since uniform gives positive probability to every action). -/ lemma absolutelyContinuous_map_hist_uniform - (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (hK : 0 < K) + (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) (hK : 0 < K) {Ωu : Type*} [MeasurableSpace Ωu] [StandardBorelSpace Ωu] [Nonempty Ωu] {Eu : Ωu → E} {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → ℝ} {Pu : Measure Ωu} [IsProbabilityMeasure Pu] - (hu : IsBayesAlgEnvSeq Q κ Eu Au Ru (Bandits.uniformAlgorithm hK) Pu) + (hu : IsBayesAlgEnvSeq Q κ (Bandits.uniformAlgorithm hK) Eu Au Ru Pu) (t : ℕ) : P.map (IsAlgEnvSeq.hist A R' t) ≪ Pu.map (IsAlgEnvSeq.hist Au Ru t) := by @@ -474,49 +439,46 @@ lemma absolutelyContinuous_map_hist_uniform set κ_unif := condDistrib (IsAlgEnvSeq.hist Au Ru t) Eu Pu rw [map_hist_eq_condDistrib_comp Q κ h t, map_hist_eq_condDistrib_comp Q κ hu t, ← Measure.snd_compProd, ← Measure.snd_compProd] + have hW_meas : Measurable (fun (ω : Ω) (n : ℕ) => (A n ω, R' n ω)) := + measurable_pi_lambda _ fun n => (h.measurable_A n).prodMk (h.measurable_R n) + have hWu_meas : Measurable (fun (ω : Ωu) (n : ℕ) => (Au n ω, Ru n ω)) := + measurable_pi_lambda _ fun n => (hu.measurable_A n).prodMk (hu.measurable_R n) exact (Measure.AbsolutelyContinuous.compProd_right (show ∀ᵐ e ∂Q, κ_alg e ≪ κ_unif e from by have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : (ℕ → Fin K × ℝ) → (Iic t → Fin K × ℝ)) = IT.hist t := funext fun ω => funext fun i => Prod.mk.eta - have h_cd₁ : ∀ᵐ e ∂Q, - κ_alg e = (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e).map (IT.hist t) := by + have h_cd₁ : ∀ᵐ e ∂Q, κ_alg e = + (condDistrib (fun ω n => (A n ω, R' n ω)) E' P e).map (IT.hist t) := by rw [← h.hasLaw_env.map_eq] - have h_comp : condDistrib (IT.hist t ∘ IsBayesAlgEnvSeq.traj A R') E' P - =ᵐ[P.map E'] (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P).map (IT.hist t) := - condDistrib_comp E' h.measurable_traj.aemeasurable (IT.measurable_hist t) - rw [IsBayesAlgEnvSeq.IT_hist_comp_traj] at h_comp + have h_comp : κ_alg + =ᵐ[P.map E'] (condDistrib (fun ω n => (A n ω, R' n ω)) E' P).map (IT.hist t) := + condDistrib_comp E' hW_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have h_cd₂ : ∀ᵐ e ∂Q, - κ_unif e = (condDistrib (IsBayesAlgEnvSeq.traj Au Ru) Eu Pu e).map (IT.hist t) := by + have h_cd₂ : ∀ᵐ e ∂Q, κ_unif e = + (condDistrib (fun ω n => (Au n ω, Ru n ω)) Eu Pu e).map (IT.hist t) := by rw [← hu.hasLaw_env.map_eq] - have h_comp : condDistrib (IT.hist t ∘ IsBayesAlgEnvSeq.traj Au Ru) Eu Pu - =ᵐ[Pu.map Eu] (condDistrib (IsBayesAlgEnvSeq.traj Au Ru) Eu Pu).map (IT.hist t) := - condDistrib_comp Eu hu.measurable_traj.aemeasurable (IT.measurable_hist t) - rw [IsBayesAlgEnvSeq.IT_hist_comp_traj] at h_comp + have h_comp : κ_unif + =ᵐ[Pu.map Eu] (condDistrib (fun ω n => (Au n ω, Ru n ω)) Eu Pu).map (IT.hist t) := + condDistrib_comp Eu hWu_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have hae₁ : ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward alg - (stationaryEnv (κ.comap (·, e) (by fun_prop))) - (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) := by - rw [← h.hasLaw_env.map_eq]; exact h.condDistrib_traj_isAlgEnvSeq - have hae₂ : ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward (Bandits.uniformAlgorithm hK) - (stationaryEnv (κ.comap (·, e) (by fun_prop))) - (condDistrib (IsBayesAlgEnvSeq.traj Au Ru) Eu Pu e) := by - rw [← hu.hasLaw_env.map_eq]; exact hu.condDistrib_traj_isAlgEnvSeq + have hae₁ := h.condDistrib_traj_isAlgEnvSeq + have hae₂ := hu.condDistrib_traj_isAlgEnvSeq filter_upwards [h_cd₁, h_cd₂, hae₁, hae₂] with e he₁ he₂ hae₁ hae₂ rw [he₁, he₂, ← h_IT_hist] exact absolutelyContinuous_map_hist_stationary hK alg _ hae₁ hae₂ t)).map measurable_snd +omit [StandardBorelSpace Ω] [Nonempty Ω] in /-- The posterior on the environment given history is algorithm-independent. -/ lemma condDistrib_env_hist_alg_indep - (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (hK : 0 < K) + (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) (hK : 0 < K) {Ωu : Type*} [MeasurableSpace Ωu] [StandardBorelSpace Ωu] [Nonempty Ωu] {Eu : Ωu → E} {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → ℝ} {Pu : Measure Ωu} [IsProbabilityMeasure Pu] - (hu : IsBayesAlgEnvSeq Q κ Eu Au Ru (Bandits.uniformAlgorithm hK) Pu) + (hu : IsBayesAlgEnvSeq Q κ (Bandits.uniformAlgorithm hK) Eu Au Ru Pu) (t : ℕ) : condDistrib E' (IsAlgEnvSeq.hist A R' t) P =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] @@ -526,37 +488,33 @@ lemma condDistrib_env_hist_alg_indep set ρ := historyDensity hK alg t have hρ_meas := measurable_historyDensity hK alg t have hρ_ne_top := historyDensity_ne_top hK alg t + have hW_meas : Measurable (fun (ω : Ω) (n : ℕ) => (A n ω, R' n ω)) := + measurable_pi_lambda _ fun n => (h.measurable_A n).prodMk (h.measurable_R n) + have hWu_meas : Measurable (fun (ω : Ωu) (n : ℕ) => (Au n ω, Ru n ω)) := + measurable_pi_lambda _ fun n => (hu.measurable_A n).prodMk (hu.measurable_R n) -- Key factorization: κ_alg =ᵐ[Q] κ_unif.withDensity (fun _ => ρ) have h_wd_ae : κ_alg =ᵐ[Q] κ_unif.withDensity (fun _ => ρ) := by have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : (ℕ → Fin K × ℝ) → (Iic t → Fin K × ℝ)) = IT.hist t := funext fun ω => funext fun i => Prod.mk.eta - have h_cd₁ : ∀ᵐ e ∂Q, - κ_alg e = (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e).map (IT.hist t) := by + have h_cd₁ : ∀ᵐ e ∂Q, κ_alg e = + (condDistrib (fun ω n => (A n ω, R' n ω)) E' P e).map (IT.hist t) := by rw [← h.hasLaw_env.map_eq] - have h_comp : condDistrib (IT.hist t ∘ IsBayesAlgEnvSeq.traj A R') E' P - =ᵐ[P.map E'] (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P).map (IT.hist t) := - condDistrib_comp E' h.measurable_traj.aemeasurable (IT.measurable_hist t) - rw [IsBayesAlgEnvSeq.IT_hist_comp_traj] at h_comp + have h_comp : κ_alg + =ᵐ[P.map E'] (condDistrib (fun ω n => (A n ω, R' n ω)) E' P).map (IT.hist t) := + condDistrib_comp E' hW_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have h_cd₂ : ∀ᵐ e ∂Q, - κ_unif e = (condDistrib (IsBayesAlgEnvSeq.traj Au Ru) Eu Pu e).map (IT.hist t) := by + have h_cd₂ : ∀ᵐ e ∂Q, κ_unif e = + (condDistrib (fun ω n => (Au n ω, Ru n ω)) Eu Pu e).map (IT.hist t) := by rw [← hu.hasLaw_env.map_eq] - have h_comp : condDistrib (IT.hist t ∘ IsBayesAlgEnvSeq.traj Au Ru) Eu Pu - =ᵐ[Pu.map Eu] (condDistrib (IsBayesAlgEnvSeq.traj Au Ru) Eu Pu).map (IT.hist t) := - condDistrib_comp Eu hu.measurable_traj.aemeasurable (IT.measurable_hist t) - rw [IsBayesAlgEnvSeq.IT_hist_comp_traj] at h_comp + have h_comp : κ_unif + =ᵐ[Pu.map Eu] (condDistrib (fun ω n => (Au n ω, Ru n ω)) Eu Pu).map (IT.hist t) := + condDistrib_comp Eu hWu_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have hae₁ : ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward alg - (stationaryEnv (κ.comap (·, e) (by fun_prop))) - (condDistrib (IsBayesAlgEnvSeq.traj A R') E' P e) := by - rw [← h.hasLaw_env.map_eq]; exact h.condDistrib_traj_isAlgEnvSeq - have hae₂ : ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward (Bandits.uniformAlgorithm hK) - (stationaryEnv (κ.comap (·, e) (by fun_prop))) - (condDistrib (IsBayesAlgEnvSeq.traj Au Ru) Eu Pu e) := by - rw [← hu.hasLaw_env.map_eq]; exact hu.condDistrib_traj_isAlgEnvSeq + have hae₁ := h.condDistrib_traj_isAlgEnvSeq + have hae₂ := hu.condDistrib_traj_isAlgEnvSeq filter_upwards [h_cd₁, h_cd₂, hae₁, hae₂] with e he₁ he₂ hae₁ hae₂ rw [Kernel.withDensity_apply _ (show Measurable (Function.uncurry (fun (_ : E) => ρ)) from hρ_meas.comp measurable_snd), @@ -564,23 +522,71 @@ lemma condDistrib_env_hist_alg_indep exact map_hist_eq_withDensity_historyDensity hK alg t _ hae₁ hae₂ haveI : IsSFiniteKernel (κ_unif.withDensity (fun _ => ρ)) := Kernel.IsSFiniteKernel.withDensity _ (fun _ b => hρ_ne_top b) - -- Posterior equality via density factorization - have h_post : posterior κ_unif Q - =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] posterior κ_alg Q := by - rw [map_hist_eq_condDistrib_comp Q κ h t] - exact posterior_eq_of_withDensity_ae_eq hρ_meas h_wd_ae - -- Bayes' rule for both algorithms - have h1 := (h.hasCondDistrib_env_hist t).condDistrib_eq - have h2' : condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu - =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] posterior κ_unif Q := - (absolutelyContinuous_map_hist_uniform Q κ h hK hu t).ae_le - (hu.hasCondDistrib_env_hist t).condDistrib_eq - exact h1.trans (h_post.symm.trans h2'.symm) + -- Direct condDistrib equality via joint measure argument + -- Show: P.map (hist, E') = P.map hist ⊗ₘ condDistrib Eu hist_u Pu + -- using the density factorization and disintegration + have h_joint₁ : P.map (fun ω => (E' ω, IsAlgEnvSeq.hist A R' t ω)) = Q ⊗ₘ κ_alg := by + rw [← h.hasLaw_env.map_eq] + exact (compProd_map_condDistrib + (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t).aemeasurable).symm + have h_joint₂ : Pu.map (fun ω => (Eu ω, IsAlgEnvSeq.hist Au Ru t ω)) = Q ⊗ₘ κ_unif := by + rw [← hu.hasLaw_env.map_eq] + exact (compProd_map_condDistrib + (IsAlgEnvSeq.measurable_hist hu.measurable_A hu.measurable_R t).aemeasurable).symm + -- The swapped joint of P equals P.map hist ⊗ₘ condDistrib Eu hist_u Pu + have h_meas_hist := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t + have h_meas_hist_u := IsAlgEnvSeq.measurable_hist hu.measurable_A hu.measurable_R t + -- P.map hist = (Pu.map hist_u).withDensity ρ + have h_hist : P.map (IsAlgEnvSeq.hist A R' t) + = (Pu.map (IsAlgEnvSeq.hist Au Ru t)).withDensity ρ := by + have h_marg₁ : P.map (IsAlgEnvSeq.hist A R' t) = (Q ⊗ₘ κ_alg).map Prod.snd := by + rw [← h_joint₁] + exact (Measure.map_map measurable_snd (h.measurable_E.prodMk h_meas_hist)).symm + have h_marg₂ : Pu.map (IsAlgEnvSeq.hist Au Ru t) = (Q ⊗ₘ κ_unif).map Prod.snd := by + rw [← h_joint₂] + exact (Measure.map_map measurable_snd (hu.measurable_E.prodMk h_meas_hist_u)).symm + rw [h_marg₁, h_marg₂, Measure.compProd_congr h_wd_ae, + Measure.compProd_withDensity + (show Measurable (Function.uncurry (fun (_ : E) => ρ)) from hρ_meas.comp measurable_snd)] + exact withDensity_map_eq' measurable_snd hρ_meas + have h_swap : P.map (fun ω => (IsAlgEnvSeq.hist A R' t ω, E' ω)) + = P.map (IsAlgEnvSeq.hist A R' t) ⊗ₘ condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu := by + have h_uncurry_meas : Measurable (Function.uncurry (fun (_ : E) => ρ)) := + hρ_meas.comp measurable_snd + calc P.map (fun ω => (IsAlgEnvSeq.hist A R' t ω, E' ω)) + _ = (Q ⊗ₘ κ_alg).map Prod.swap := by + rw [← h_joint₁] + exact (Measure.map_map measurable_swap + (h.measurable_E.prodMk h_meas_hist)).symm + _ = (Q ⊗ₘ (κ_unif.withDensity (fun _ => ρ))).map Prod.swap := by + rw [Measure.compProd_congr h_wd_ae] + _ = ((Q ⊗ₘ κ_unif).withDensity (ρ ∘ Prod.snd)).map Prod.swap := by + congr 1; exact Measure.compProd_withDensity h_uncurry_meas + _ = ((Q ⊗ₘ κ_unif).map Prod.swap).withDensity (ρ ∘ Prod.fst) := + map_swap_withDensity_fst hρ_meas + _ = (Pu.map (fun ω => (IsAlgEnvSeq.hist Au Ru t ω, Eu ω))).withDensity + (ρ ∘ Prod.fst) := by + congr 1; rw [← h_joint₂] + exact Measure.map_map measurable_swap + (hu.measurable_E.prodMk h_meas_hist_u) + _ = (Pu.map (IsAlgEnvSeq.hist Au Ru t) ⊗ₘ + condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu).withDensity + (ρ ∘ Prod.fst) := by + rw [← compProd_map_condDistrib hu.measurable_E.aemeasurable] + _ = (Pu.map (IsAlgEnvSeq.hist Au Ru t)).withDensity ρ ⊗ₘ + condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu := + (withDensity_compProd_left hρ_meas).symm + _ = P.map (IsAlgEnvSeq.hist A R' t) ⊗ₘ + condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu := by + rw [h_hist] + -- By uniqueness of disintegration + exact (condDistrib_ae_eq_iff_measure_eq_compProd _ + h.measurable_E.aemeasurable (condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu)).mpr h_swap /-- The posterior on the best arm equals the uniform algorithm's posterior. -/ lemma posteriorBestArm_eq_uniform - (h : IsBayesAlgEnvSeq Q κ E' A R' alg P) (hK : 0 < K) (t : ℕ) : - condDistrib (IsBayesAlgEnvSeq.bestArm κ E') (IsAlgEnvSeq.hist A R' t) P + (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) (hK : 0 < K) (t : ℕ) : + condDistrib (IsBayesAlgEnvSeq.bestAction κ E') (IsAlgEnvSeq.hist A R' t) P =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] IT.posteriorBestArm Q κ (Bandits.uniformAlgorithm hK) t := by unfold IT.posteriorBestArm @@ -589,22 +595,22 @@ lemma posteriorBestArm_eq_uniform set histfu := IsAlgEnvSeq.hist IT.action (fun n (ω : ℕ → Fin K × E × ℝ) ↦ (ω n).2.2) t set envf := E' set envfu : (ℕ → Fin K × E × ℝ) → E := fun ω ↦ (ω 0).2.1 - set bau := IsBayesAlgEnvSeq.bestArm (Ω := ℕ → Fin K × E × ℝ) κ envfu + set bau := IsBayesAlgEnvSeq.bestAction (Ω := ℕ → Fin K × E × ℝ) κ envfu have h_ITu := IT.isBayesAlgEnvSeq_bayesianTrajMeasure Q κ (Bandits.uniformAlgorithm hK) -- LHS: condDistrib (bestArm κ E') histf P -- =ᵐ (condDistrib envf histf P).map (envToBestArm κ) - have h_comp_alg : condDistrib (IsBayesAlgEnvSeq.bestArm κ E') histf P + have h_comp_alg : condDistrib (IsBayesAlgEnvSeq.bestAction κ E') histf P =ᵐ[P.map histf] (condDistrib envf histf P).map (envToBestArm κ) := by - rw [bestArm_eq_envToBestArm_comp_env κ] + rw [bestAction_eq_envToBestArm_comp_env κ] exact condDistrib_comp (mβ := MeasurableSpace.pi) histf h.measurable_E.aemeasurable (measurable_envToBestArm κ) -- RHS: condDistrib bau histfu Pu -- =ᵐ (condDistrib envfu histfu Pu).map (envToBestArm κ) have h_comp_unif : condDistrib bau histfu Pu =ᵐ[Pu.map histfu] (condDistrib envfu histfu Pu).map (envToBestArm κ) := by - change condDistrib (IsBayesAlgEnvSeq.bestArm κ envfu) histfu Pu + change condDistrib (IsBayesAlgEnvSeq.bestAction κ envfu) histfu Pu =ᵐ[Pu.map histfu] (condDistrib envfu histfu Pu).map (envToBestArm κ) - rw [bestArm_eq_envToBestArm_comp_env κ] + rw [bestAction_eq_envToBestArm_comp_env κ] exact condDistrib_comp (mβ := MeasurableSpace.pi) histfu h_ITu.measurable_E.aemeasurable (measurable_envToBestArm κ) -- Environment posterior independence From 725941b19a2bb60c57f7e06f51366bcbec60e21d Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Sat, 21 Feb 2026 04:31:22 +0000 Subject: [PATCH 044/155] Refactor BayesStationaryEnv (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 51 ++-- .../BayesStationaryEnv.lean | 242 ++++++++---------- .../SequentialLearning/HistoryDensity.lean | 56 +--- 3 files changed, 159 insertions(+), 190 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 281140c5..c3658199 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -21,25 +21,32 @@ namespace Bandits namespace TS variable {K : ℕ} (hK : 0 < K) -variable {E : Type*} [mE : MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] -variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] +variable {𝓔 : Type*} [mE : MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] +variable (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (Fin K × 𝓔) ℝ) [IsMarkovKernel κ] /-- The distribution over actions for every given history for TS. -/ noncomputable def policy (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - IT.posteriorBestArm Q κ (uniformAlgorithm hK) n -deriving IsMarkovKernel + (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) n).map + (IsBayesAlgEnvSeq.bestAction κ id) + +instance (n : ℕ) : IsMarkovKernel (policy hK Q κ n) := by + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + unfold policy + exact Kernel.IsMarkovKernel.map _ + (IsBayesAlgEnvSeq.measurable_bestAction measurable_id) /-- The initial distribution over actions for TS. -/ noncomputable def initialPolicy : Measure (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - IT.priorBestArm Q κ (uniformAlgorithm hK) + Q.map (IsBayesAlgEnvSeq.bestAction κ id) instance : IsProbabilityMeasure (initialPolicy hK Q κ) := by - unfold initialPolicy - infer_instance + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + exact Measure.isProbabilityMeasure_map + (IsBayesAlgEnvSeq.measurable_bestAction (by fun_prop)).aemeasurable end TS @@ -47,14 +54,14 @@ variable {K : ℕ} section Algorithm -variable {E : Type*} [MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] +variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] /-- The Thompson Sampling (TS) algorithm: actions are chosen according to the probability that they are optimal given prior knowledge represented by a prior distribution `Q` and a data generation model represented by a kernel `κ`. -/ noncomputable -def tsAlgorithm (hK : 0 < K) (Q : Measure E) [IsProbabilityMeasure Q] - (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] : Algorithm (Fin K) ℝ where +def tsAlgorithm (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] + (κ : Kernel (Fin K × 𝓔) ℝ) [IsMarkovKernel κ] : Algorithm (Fin K) ℝ where policy := TS.policy hK Q κ p0 := TS.initialPolicy hK Q κ @@ -208,8 +215,22 @@ lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] condDistrib (A (t + 1)) (IsAlgEnvSeq.hist A R' t) P =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] condDistrib (IsBayesAlgEnvSeq.bestAction κ E') (IsAlgEnvSeq.hist A R' t) P := - (h.hasCondDistrib_action' t).condDistrib_eq.trans - (posteriorBestArm_eq_uniform Q κ h hK t).symm + by + have h_ba_comp : IsBayesAlgEnvSeq.bestAction κ E' + = IsBayesAlgEnvSeq.bestAction κ id ∘ E' := by + rw [bestAction_eq_envToBestArm_comp_env κ (E' := E'), + bestAction_eq_envToBestArm_comp_env κ (E' := id), Function.comp_id] + rw [h_ba_comp] + have hm := IsBayesAlgEnvSeq.measurable_bestAction (κ := κ) measurable_id + have h_comp := condDistrib_comp (mβ := MeasurableSpace.pi) (μ := P) + (IsAlgEnvSeq.hist A R' t) h.measurable_E.aemeasurable hm + have h_map : (condDistrib E' (IsAlgEnvSeq.hist A R' t) P).map + (IsBayesAlgEnvSeq.bestAction κ id) =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] + (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) t).map + (IsBayesAlgEnvSeq.bestAction κ id) := by + filter_upwards [posterior_eq_uniform Q κ h hK t] with x hx + simp only [Kernel.map_apply _ hm, hx] + exact (h.hasCondDistrib_action' t).condDistrib_eq.trans (h_comp.trans h_map).symm omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma le_armMean_bestArm [Nonempty (Fin K)] (ω : Ω) (i : Fin K) : @@ -538,7 +559,7 @@ lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] have h_cond_ae : ∀ᵐ e ∂(P.map E'), IsAlgEnvSeq IT.action IT.reward (tsAlgorithm hK Q κ) (stationaryEnv (κ.comap (·, e) (by fun_prop))) (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E' P e) := by - rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h + rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h filter_upwards [h_cond_ae] with e h_isAlgEnvSeq let ν := κ.comap (·, e) (by fun_prop) have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by @@ -854,7 +875,7 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] have h_cond_ae : ∀ᵐ e ∂(P.map E'), IsAlgEnvSeq IT.action IT.reward (tsAlgorithm hK Q κ) (stationaryEnv (κ.comap (·, e) (by fun_prop))) (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E' P e) := by - rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h + rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h filter_upwards [h_cond_ae] with e h_isAlgEnvSeq exact concentration_cond_bound (hK := hK) (E' := E') (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) hσ2 hs hn' hδ hδ1 e h_isAlgEnvSeq a @@ -951,7 +972,7 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] have h_cond_ae : ∀ᵐ e ∂(P.map E'), IsAlgEnvSeq IT.action IT.reward (tsAlgorithm hK Q κ) (stationaryEnv (κ.comap (·, e) (by fun_prop))) (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E' P e) := by - rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.condDistrib_traj_isAlgEnvSeq h + rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h filter_upwards [h_cond_ae] with e h_isAlgEnvSeq intro a exact concentration_cond_bound (hK := hK) (E' := E') (A := A) (R' := R') diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index f2a5a5b8..c560fca1 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -13,8 +13,8 @@ open MeasureTheory ProbabilityTheory Finset namespace Learning -variable {α R 𝓔 : Type*} [MeasurableSpace α] [MeasurableSpace R] [MeasurableSpace 𝓔] -variable {Ω : Type*} [MeasurableSpace Ω] +variable {𝓔 α R Ω : Type*} +variable [MeasurableSpace 𝓔] [MeasurableSpace α] [MeasurableSpace R] [MeasurableSpace Ω] structure IsBayesAlgEnvSeq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] @@ -36,13 +36,55 @@ structure IsBayesAlgEnvSeq namespace IsBayesAlgEnvSeq -section Laws +def trajectory (A : ℕ → Ω → α) (R' : ℕ → Ω → R) : Ω → ℕ → α × R := fun ω n ↦ (A n ω, R' n ω) + +@[fun_prop] +lemma measurable_trajectory {A : ℕ → Ω → α} {R' : ℕ → Ω → R} (hA : ∀ n, Measurable (A n)) + (hR : ∀ n, Measurable (R' n)) : Measurable (trajectory A R') := by + unfold trajectory + fun_prop + +section Real + +noncomputable +def actionMean (κ : Kernel (α × 𝓔) ℝ) (E : Ω → 𝓔) (a : α) (ω : Ω) : ℝ := (κ (a, E ω))[id] + +@[fun_prop] +lemma measurable_actionMean {κ : Kernel (α × 𝓔) ℝ} {E : Ω → 𝓔} {a : α} (hE : Measurable E) : + Measurable (actionMean κ E a) := + stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop) + +noncomputable +def bestAction [Fintype α] [Encodable α] [Nonempty α] [MeasurableSingletonClass α] + (κ : Kernel (α × 𝓔) ℝ) (E : Ω → 𝓔) (ω : Ω) : α := + measurableArgmax (fun ω' a ↦ actionMean κ E a ω') ω + +@[fun_prop] +lemma measurable_bestAction [Nonempty α] [Fintype α] [Encodable α] [MeasurableSingletonClass α] + {κ : Kernel (α × 𝓔) ℝ} {E : Ω → 𝓔} (hE : Measurable E) : Measurable (bestAction κ E) := + measurable_measurableArgmax (by fun_prop) + +noncomputable +def regret (κ : Kernel (α × 𝓔) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → α) (t : ℕ) (ω : Ω) : ℝ := + Bandits.regret (κ.comap (·, E ω) (by fun_prop)) A t ω + +@[fun_prop] +lemma measurable_regret [Countable α] {κ : Kernel (α × 𝓔) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → α} {t : ℕ} + (hE : Measurable E) (hA : ∀ n, Measurable (A n)) : + Measurable (regret κ E A t) := by + have hm := (stronglyMeasurable_id.integral_kernel (κ := κ)).measurable + exact (Measurable.const_mul (Measurable.iSup fun _ ↦ hm.comp (by fun_prop)) _).sub + (Finset.measurable_sum _ fun _ _ ↦ hm.comp (by fun_prop)) + +end Real variable [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] variable {Q : Measure 𝓔} {κ : Kernel (α × 𝓔) R} {alg : Algorithm α R} variable {E : Ω → 𝓔} {A : ℕ → Ω → α} {R' : ℕ → Ω → R} variable {P : Measure Ω} [IsFiniteMeasure P] +section Laws + lemma hasLaw_action_zero [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : HasLaw (A 0) alg.p0 P := h.hasCondDistrib_action_zero.hasLaw_of_const @@ -55,30 +97,25 @@ lemma hasCondDistrib_reward' [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg HasCondDistrib (R' (n + 1)) (fun ω ↦ (A (n + 1) ω, E ω)) κ P := (h.hasCondDistrib_reward n).comp_left (by fun_prop) ---- +end Laws + +section CondDistribIsAlgEnvSeq -lemma hasLaw_action_zero_fiber (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : - ∀ᵐ e ∂Q, HasLaw (IT.action 0) alg.p0 - (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E P e) := by +lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : + ∀ᵐ e ∂Q, HasLaw (IT.action 0) alg.p0 (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] - have hW : AEMeasurable (fun ω n ↦ (A n ω, R' n ω)) P := - (measurable_pi_lambda _ fun n ↦ (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable - have h_comp : ⇑(condDistrib (A 0) E P) =ᶠ[ae (P.map E)] - ⇑((condDistrib (fun ω n ↦ (A n ω, R' n ω)) E P).map (IT.action 0)) := - condDistrib_comp E hW (IT.measurable_action 0) - filter_upwards [h_comp, h.hasCondDistrib_action_zero.condDistrib_eq] with e he hcd + filter_upwards [condDistrib_comp E + (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable (IT.measurable_action 0), + h.hasCondDistrib_action_zero.condDistrib_eq] with e he hcd exact ⟨(IT.measurable_action 0).aemeasurable, by - rw [← Kernel.map_apply _ (IT.measurable_action 0), ← he, hcd, Kernel.const_apply]⟩ + rw [← Kernel.map_apply _ (IT.measurable_action 0), ← he, + show IT.action 0 ∘ trajectory A R' = A 0 from rfl, hcd, Kernel.const_apply]⟩ -lemma hasCondDistrib_reward_zero_fiber [IsFiniteKernel κ] - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : - ∀ᵐ e ∂Q, HasCondDistrib (IT.reward 0) (IT.action 0) - (κ.comap (·, e) (by fun_prop)) - (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E P e) := by +lemma hasCondDistrib_IT_reward_zero [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : + ∀ᵐ e ∂Q, HasCondDistrib (IT.reward 0) (IT.action 0) (κ.comap (·, e) (by fun_prop)) + (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] - set W := fun ω n ↦ (A n ω, R' n ω) - have hW : AEMeasurable W P := - (measurable_pi_lambda _ fun n ↦ (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable + have hW := (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable (μ := P) have h_swap : HasCondDistrib (R' 0) (fun ω ↦ (E ω, A 0 ω)) (κ.comap Prod.swap (by fun_prop)) P := by convert h.hasCondDistrib_reward_zero.comp_right @@ -86,10 +123,10 @@ lemma hasCondDistrib_reward_zero_fiber [IsFiniteKernel κ] have h_prod := condDistrib_prod_left (h.measurable_A 0).aemeasurable (h.measurable_R 0).aemeasurable h.measurable_E.aemeasurable (μ := P) have h_comp_pair : ⇑(condDistrib (fun ω ↦ (A 0 ω, R' 0 ω)) E P) =ᶠ[ae (P.map E)] - ⇑((condDistrib W E P).map (fun ω ↦ (IT.action 0 ω, IT.reward 0 ω))) := + ⇑((condDistrib (trajectory A R') E P).map (fun ω ↦ (IT.action 0 ω, IT.reward 0 ω))) := condDistrib_comp E hW ((IT.measurable_action 0).prodMk (IT.measurable_reward 0)) have h_comp_action : ⇑(condDistrib (A 0) E P) =ᶠ[ae (P.map E)] - ⇑((condDistrib W E P).map (IT.action 0)) := + ⇑((condDistrib (trajectory A R') E P).map (IT.action 0)) := condDistrib_comp E hW (IT.measurable_action 0) have h_swap_eq := h_swap.condDistrib_eq rw [(compProd_map_condDistrib (h.measurable_A 0).aemeasurable).symm] at h_swap_eq @@ -106,53 +143,42 @@ lemma hasCondDistrib_reward_zero_fiber [IsFiniteKernel κ] ext s _ rw [Kernel.sectR_apply, Kernel.comap_apply, ha, Kernel.comap_apply]; rfl -lemma hasCondDistrib_action_fiber (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : - ∀ᵐ e ∂Q, HasCondDistrib (IT.action (n + 1)) - (IsAlgEnvSeq.hist IT.action IT.reward n) (alg.policy n) - (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E P e) := by +lemma hasCondDistrib_IT_action (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : + ∀ᵐ e ∂Q, HasCondDistrib (IT.action (n + 1)) (IT.hist n) (alg.policy n) + (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] - set W := fun ω n ↦ (A n ω, R' n ω) - have hW : AEMeasurable W P := - (measurable_pi_lambda _ fun n ↦ (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable + have hW := (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable (μ := P) have h_hist_meas := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n have h_prod := condDistrib_prod_left h_hist_meas.aemeasurable (h.measurable_A (n + 1)).aemeasurable h.measurable_E.aemeasurable (μ := P) have h_action_env := (h.hasCondDistrib_action n).condDistrib_eq - have h_hist_IT_meas : Measurable - (IsAlgEnvSeq.hist (IT.action (R := R)) (IT.reward (α := α)) n) := - IsAlgEnvSeq.measurable_hist (fun n ↦ IT.measurable_action n) (fun n ↦ IT.measurable_reward n) n have h_comp_pair : ⇑(condDistrib (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) E P) - =ᶠ[ae (P.map E)] ⇑((condDistrib W E P).map - (fun ω ↦ (IsAlgEnvSeq.hist IT.action IT.reward n ω, IT.action (n + 1) ω))) := - condDistrib_comp E hW (h_hist_IT_meas.prodMk (IT.measurable_action (n + 1))) + =ᶠ[ae (P.map E)] ⇑((condDistrib (trajectory A R') E P).map + (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω))) := + condDistrib_comp E hW ((IT.measurable_hist n).prodMk (IT.measurable_action (n + 1))) have h_comp_hist : ⇑(condDistrib (IsAlgEnvSeq.hist A R' n) E P) =ᶠ[ae (P.map E)] - ⇑((condDistrib W E P).map (IsAlgEnvSeq.hist IT.action IT.reward n)) := - condDistrib_comp E hW h_hist_IT_meas + ⇑((condDistrib (trajectory A R') E P).map (IT.hist n)) := + condDistrib_comp E hW (IT.measurable_hist n) rw [(compProd_map_condDistrib h_hist_meas.aemeasurable).symm] at h_action_env filter_upwards [h_prod, h_comp_pair, h_comp_hist, (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_action_env] with e h_prod_e h_pair_e h_hist_e h_nested_e refine ⟨by fun_prop, by fun_prop, ?_⟩ rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] - rw [← Kernel.map_apply _ (h_hist_IT_meas.prodMk (IT.measurable_action (n + 1))), + rw [← Kernel.map_apply _ ((IT.measurable_hist n).prodMk (IT.measurable_action (n + 1))), ← h_pair_e] - conv_rhs => rw [← Kernel.map_apply _ h_hist_IT_meas, ← h_hist_e] + conv_rhs => rw [← Kernel.map_apply _ (IT.measurable_hist n), ← h_hist_e] rw [h_prod_e, Kernel.compProd_apply_eq_compProd_sectR] refine Measure.compProd_congr ?_ filter_upwards [h_nested_e] with _ ha ext s _ rw [Kernel.sectR_apply, ha, Kernel.prodMkLeft_apply] -lemma hasCondDistrib_reward_fiber [IsFiniteKernel κ] - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : - ∀ᵐ e ∂Q, HasCondDistrib (IT.reward (n + 1)) - (fun f ↦ (IsAlgEnvSeq.hist IT.action IT.reward n f, IT.action (n + 1) f)) - ((κ.comap (·, e) (by fun_prop)).prodMkLeft _) - (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E P e) := by +lemma hasCondDistrib_IT_reward [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : + ∀ᵐ e ∂Q, HasCondDistrib (IT.reward (n + 1)) (fun x ↦ (IT.hist n x, IT.action (n + 1) x)) + ((κ.comap (·, e) (by fun_prop)).prodMkLeft _) (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] - set W := fun ω n ↦ (A n ω, R' n ω) - have hW : AEMeasurable W P := - (measurable_pi_lambda _ fun n ↦ (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable + have hW := (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable (μ := P) have h_hist_meas := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n have h_prod := condDistrib_prod_left (Measurable.prodMk h_hist_meas (h.measurable_A (n + 1))).aemeasurable @@ -163,19 +189,17 @@ lemma hasCondDistrib_reward_fiber [IsFiniteKernel κ] (h.hasCondDistrib_reward n).comp_right (MeasurableEquiv.prodAssoc.symm.trans MeasurableEquiv.prodComm) have h_swap_eq := h_swap.condDistrib_eq - have h_hist_IT_meas : Measurable - (IsAlgEnvSeq.hist (IT.action (R := R)) (IT.reward (α := α)) n) := - IsAlgEnvSeq.measurable_hist (fun n ↦ IT.measurable_action n) (fun n ↦ IT.measurable_reward n) n - have h_pair_meas := h_hist_IT_meas.prodMk (IT.measurable_action (n + 1)) + have h_pair_meas : Measurable + (fun f : ℕ → α × R ↦ (IT.hist n f, IT.action (n + 1) f)) := by fun_prop have h_comp_triple : ⇑(condDistrib (fun ω ↦ ((IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω), R' (n + 1) ω)) E P) - =ᶠ[ae (P.map E)] ⇑((condDistrib W E P).map - (fun ω ↦ ((IsAlgEnvSeq.hist IT.action IT.reward n ω, IT.action (n + 1) ω), + =ᶠ[ae (P.map E)] ⇑((condDistrib (trajectory A R') E P).map + (fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), IT.reward (n + 1) ω))) := condDistrib_comp E hW (h_pair_meas.prodMk (IT.measurable_reward (n + 1))) have h_comp_pair : ⇑(condDistrib (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) E P) - =ᶠ[ae (P.map E)] ⇑((condDistrib W E P).map - (fun ω ↦ (IsAlgEnvSeq.hist IT.action IT.reward n ω, IT.action (n + 1) ω))) := + =ᶠ[ae (P.map E)] ⇑((condDistrib (trajectory A R') E P).map + (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω))) := condDistrib_comp E hW h_pair_meas rw [(compProd_map_condDistrib (Measurable.prodMk h_hist_meas (h.measurable_A (n + 1))).aemeasurable).symm] at h_swap_eq @@ -192,52 +216,24 @@ lemma hasCondDistrib_reward_fiber [IsFiniteKernel κ] ext s _ rw [Kernel.sectR_apply, ha, Kernel.comap_apply, Kernel.prodMkLeft_apply, Kernel.comap_apply] -lemma condDistrib_traj_isAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : +lemma ae_IsAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv (κ.comap (·, e) (by fun_prop))) - (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E P e) := by - filter_upwards [hasLaw_action_zero_fiber h, - hasCondDistrib_reward_zero_fiber h, - ae_all_iff.2 (hasCondDistrib_action_fiber h), - ae_all_iff.2 (hasCondDistrib_reward_fiber h)] - with _ h_law h_r0 h_a h_r + (condDistrib (trajectory A R') E P e) := by + filter_upwards [hasLaw_IT_action_zero h, hasCondDistrib_IT_reward_zero h, + ae_all_iff.2 (hasCondDistrib_IT_action h), ae_all_iff.2 (hasCondDistrib_IT_reward h)] + with _ h_a0 h_r0 h_a h_r exact { - hasLaw_action_zero := h_law + hasLaw_action_zero := h_a0 hasCondDistrib_reward_zero := h_r0 hasCondDistrib_action := h_a hasCondDistrib_reward := h_r } -end Laws - -section Real - -noncomputable -def actionMean (κ : Kernel (α × 𝓔) ℝ) (E : Ω → 𝓔) (a : α) (ω : Ω) : ℝ := (κ (a, E ω))[id] - -@[fun_prop] -lemma measurable_actionMean {κ : Kernel (α × 𝓔) ℝ} {E : Ω → 𝓔} {a : α} (hE : Measurable E) : - Measurable (actionMean κ E a) := - stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop) - -noncomputable -def bestAction [Fintype α] [Encodable α] [Nonempty α] [MeasurableSingletonClass α] - (κ : Kernel (α × 𝓔) ℝ) (E : Ω → 𝓔) (ω : Ω) : α := - measurableArgmax (fun ω' a ↦ actionMean κ E a ω') ω - -@[fun_prop] -lemma measurable_bestAction [Fintype α] [Encodable α] [Nonempty α] [MeasurableSingletonClass α] - {κ : Kernel (α × 𝓔) ℝ} {E : Ω → 𝓔} (hE : Measurable E) : Measurable (bestAction κ E) := - measurable_measurableArgmax (by fun_prop) - -noncomputable -def regret (κ : Kernel (α × 𝓔) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → α) (t : ℕ) (ω : Ω) : ℝ := - Bandits.regret (κ.comap (·, E ω) (by fun_prop)) A t ω - -end Real +end CondDistribIsAlgEnvSeq end IsBayesAlgEnvSeq -section StationaryEquivalence +section IsAlgEnvSeq noncomputable def bayesStationaryEnv (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) R) @@ -247,17 +243,15 @@ def bayesStationaryEnv (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) ν0 := (Kernel.const _ Q) ⊗ₖ κ -/-- Bridge theorem: an `IsAlgEnvSeq` for `(alg.prod_left E)` and `(bayesStationaryEnv Q κ)` -gives rise to an `IsBayesAlgEnvSeq`. -/ -theorem IsAlgEnvSeq.toIsBayesAlgEnvSeq - [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace 𝓔] [Nonempty 𝓔] - [StandardBorelSpace R] [Nonempty R] - {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (α × 𝓔) R} [IsMarkovKernel κ] - {A : ℕ → Ω → α} {R'' : ℕ → Ω → 𝓔 × R} {alg : Algorithm α R} - {P : Measure Ω} [IsProbabilityMeasure P] - (h : IsAlgEnvSeq A R'' (alg.prod_left 𝓔) (bayesStationaryEnv Q κ) P) : - IsBayesAlgEnvSeq Q κ alg (fun ω ↦ (R'' 0 ω).1) A (fun n ω ↦ (R'' n ω).2) P where +variable [Nonempty α] [Nonempty 𝓔] [Nonempty R] +variable [StandardBorelSpace α] [StandardBorelSpace 𝓔] [StandardBorelSpace R] +variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (α × 𝓔) R} [IsMarkovKernel κ] +variable {alg : Algorithm α R} {A : ℕ → Ω → α} {R' : ℕ → Ω → 𝓔 × R} +variable {P : Measure Ω} [IsProbabilityMeasure P] + +lemma IsAlgEnvSeq.isBayesAlgEnvSeq + (h : IsAlgEnvSeq A R' (alg.prod_left 𝓔) (bayesStationaryEnv Q κ) P) : + IsBayesAlgEnvSeq Q κ alg (fun ω ↦ (R' 0 ω).1) A (fun n ω ↦ (R' n ω).2) P where measurable_E := (h.measurable_R 0).fst measurable_A := h.measurable_A measurable_R n := (h.measurable_R n).snd @@ -265,10 +259,10 @@ theorem IsAlgEnvSeq.toIsBayesAlgEnvSeq apply HasCondDistrib.hasLaw_of_const simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst hasCondDistrib_action_zero := by - have hfst : HasCondDistrib (fun ω ↦ (R'' 0 ω).1) (A 0) (Kernel.const α Q) P := by + have hfst : HasCondDistrib (fun ω ↦ (R' 0 ω).1) (A 0) (Kernel.const α Q) P := by simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst -- E' | A 0 is constant Q = P.map E', so A 0 and E' are independent - have h_indep : IndepFun (A 0) (fun ω ↦ (R'' 0 ω).1) P := by + have h_indep : IndepFun (A 0) (fun ω ↦ (R' 0 ω).1) P := by rw [indepFun_iff_condDistrib_eq_const (h.measurable_A 0).aemeasurable (h.measurable_R 0).fst.aemeasurable, hfst.hasLaw_of_const.map_eq] exact hfst.condDistrib_eq @@ -282,7 +276,7 @@ theorem IsAlgEnvSeq.toIsBayesAlgEnvSeq hasCondDistrib_action n := by let f : (Iic n → α × 𝓔 × R) → 𝓔 × (Iic n → α × R) := fun h ↦ ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2)) - suffices h' : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R'' n) + suffices h' : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) (((alg.policy n).comap Prod.snd (by fun_prop)).comap f (by fun_prop)) P from h'.comp_left (f := f) exact h.hasCondDistrib_action n @@ -290,11 +284,13 @@ theorem IsAlgEnvSeq.toIsBayesAlgEnvSeq let f : (Iic n → α × 𝓔 × R) × α → (Iic n → α × R) × α × 𝓔 := fun p ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1) have hf : Measurable f := by fun_prop - suffices h' : HasCondDistrib (fun ω ↦ (R'' (n + 1) ω).2) - (fun ω ↦ (IsAlgEnvSeq.hist A R'' n ω, A (n + 1) ω)) + suffices h' : HasCondDistrib (fun ω ↦ (R' (n + 1) ω).2) + (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) ((κ.comap Prod.snd (by fun_prop)).comap f hf) P from h'.comp_left hf simpa [bayesStationaryEnv, Kernel.snd_prod] using (h.hasCondDistrib_reward n).snd +end IsAlgEnvSeq + namespace IT noncomputable @@ -303,7 +299,7 @@ def bayesTrajMeasure (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel ( trajMeasure (alg.prod_left 𝓔) (bayesStationaryEnv Q κ) deriving IsProbabilityMeasure -lemma isBayesAlgEnvSeq_bayesianTrajMeasure +lemma isBayesAlgEnvSeq_bayesTrajMeasure [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace 𝓔] [Nonempty 𝓔] [StandardBorelSpace R] [Nonempty R] @@ -311,32 +307,16 @@ lemma isBayesAlgEnvSeq_bayesianTrajMeasure (alg : Algorithm α R) : IsBayesAlgEnvSeq Q κ alg (fun ω ↦ (ω 0).2.1) action (fun n ω ↦ (ω n).2.2) (bayesTrajMeasure Q κ alg) := - (isAlgEnvSeq_trajMeasure _ _).toIsBayesAlgEnvSeq + (isAlgEnvSeq_trajMeasure _ _).isBayesAlgEnvSeq -/-- The conditional distribution over the best arm given the observed history. -/ noncomputable -def posteriorBestArm [StandardBorelSpace α] [Nonempty α] [Fintype α] [Encodable α] +def bayesTrajMeasurePosterior [StandardBorelSpace 𝓔] [Nonempty 𝓔] (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) ℝ) [IsMarkovKernel κ] - (alg : Algorithm α ℝ) (n : ℕ) : Kernel (Iic n → α × ℝ) α := - condDistrib (IsBayesAlgEnvSeq.bestAction κ (fun ω ↦ (ω 0).2.1)) - (IsAlgEnvSeq.hist action (fun n ω ↦ (ω n).2.2) n) + (alg : Algorithm α ℝ) (n : ℕ) : Kernel (Iic n → α × ℝ) 𝓔 := + condDistrib (fun ω ↦ (ω 0).2.1) (IsAlgEnvSeq.hist action (fun n ω ↦ (ω n).2.2) n) (bayesTrajMeasure Q κ alg) deriving IsMarkovKernel -/-- The initial distribution over the best arm. -/ -noncomputable -def priorBestArm [StandardBorelSpace α] [Nonempty α] [Fintype α] [Encodable α] - (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) ℝ) [IsMarkovKernel κ] - (alg : Algorithm α ℝ) : Measure α := - (bayesTrajMeasure Q κ alg).map (IsBayesAlgEnvSeq.bestAction κ (fun ω ↦ (ω 0).2.1)) - -instance [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace 𝓔] [Nonempty 𝓔] [Fintype α] - [Encodable α] (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) ℝ) - [IsMarkovKernel κ] (alg : Algorithm α ℝ) : IsProbabilityMeasure (priorBestArm Q κ alg) := - Measure.isProbabilityMeasure_map (by fun_prop) - end IT -end StationaryEquivalence - end Learning diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 00f03cec..6719e9ef 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -464,8 +464,8 @@ lemma absolutelyContinuous_map_hist_uniform condDistrib_comp Eu hWu_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have hae₁ := h.condDistrib_traj_isAlgEnvSeq - have hae₂ := hu.condDistrib_traj_isAlgEnvSeq + have hae₁ := h.ae_IsAlgEnvSeq + have hae₂ := hu.ae_IsAlgEnvSeq filter_upwards [h_cd₁, h_cd₂, hae₁, hae₂] with e he₁ he₂ hae₁ hae₂ rw [he₁, he₂, ← h_IT_hist] exact absolutelyContinuous_map_hist_stationary hK alg _ hae₁ hae₂ t)).map @@ -513,8 +513,8 @@ lemma condDistrib_env_hist_alg_indep condDistrib_comp Eu hWu_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have hae₁ := h.condDistrib_traj_isAlgEnvSeq - have hae₂ := hu.condDistrib_traj_isAlgEnvSeq + have hae₁ := h.ae_IsAlgEnvSeq + have hae₂ := hu.ae_IsAlgEnvSeq filter_upwards [h_cd₁, h_cd₂, hae₁, hae₂] with e he₁ he₂ hae₁ hae₂ rw [Kernel.withDensity_apply _ (show Measurable (Function.uncurry (fun (_ : E) => ρ)) from hρ_meas.comp measurable_snd), @@ -583,48 +583,16 @@ lemma condDistrib_env_hist_alg_indep exact (condDistrib_ae_eq_iff_measure_eq_compProd _ h.measurable_E.aemeasurable (condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu)).mpr h_swap -/-- The posterior on the best arm equals the uniform algorithm's posterior. -/ -lemma posteriorBestArm_eq_uniform +omit [StandardBorelSpace Ω] [Nonempty Ω] in +/-- The environment posterior is algorithm-independent: it equals the posterior under the +uniform algorithm, which is `IsBayesAlgEnvSeq.posterior`. -/ +lemma posterior_eq_uniform (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) (hK : 0 < K) (t : ℕ) : - condDistrib (IsBayesAlgEnvSeq.bestAction κ E') (IsAlgEnvSeq.hist A R' t) P + condDistrib E' (IsAlgEnvSeq.hist A R' t) P =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - IT.posteriorBestArm Q κ (Bandits.uniformAlgorithm hK) t := by - unfold IT.posteriorBestArm - set Pu := IT.bayesTrajMeasure Q κ (Bandits.uniformAlgorithm hK) - set histf := IsAlgEnvSeq.hist A R' t - set histfu := IsAlgEnvSeq.hist IT.action (fun n (ω : ℕ → Fin K × E × ℝ) ↦ (ω n).2.2) t - set envf := E' - set envfu : (ℕ → Fin K × E × ℝ) → E := fun ω ↦ (ω 0).2.1 - set bau := IsBayesAlgEnvSeq.bestAction (Ω := ℕ → Fin K × E × ℝ) κ envfu - have h_ITu := IT.isBayesAlgEnvSeq_bayesianTrajMeasure Q κ (Bandits.uniformAlgorithm hK) - -- LHS: condDistrib (bestArm κ E') histf P - -- =ᵐ (condDistrib envf histf P).map (envToBestArm κ) - have h_comp_alg : condDistrib (IsBayesAlgEnvSeq.bestAction κ E') histf P - =ᵐ[P.map histf] (condDistrib envf histf P).map (envToBestArm κ) := by - rw [bestAction_eq_envToBestArm_comp_env κ] - exact condDistrib_comp (mβ := MeasurableSpace.pi) histf - h.measurable_E.aemeasurable (measurable_envToBestArm κ) - -- RHS: condDistrib bau histfu Pu - -- =ᵐ (condDistrib envfu histfu Pu).map (envToBestArm κ) - have h_comp_unif : condDistrib bau histfu Pu - =ᵐ[Pu.map histfu] (condDistrib envfu histfu Pu).map (envToBestArm κ) := by - change condDistrib (IsBayesAlgEnvSeq.bestAction κ envfu) histfu Pu - =ᵐ[Pu.map histfu] (condDistrib envfu histfu Pu).map (envToBestArm κ) - rw [bestAction_eq_envToBestArm_comp_env κ] - exact condDistrib_comp (mβ := MeasurableSpace.pi) histfu - h_ITu.measurable_E.aemeasurable (measurable_envToBestArm κ) - -- Environment posterior independence - have h_env_indep := condDistrib_env_hist_alg_indep Q κ h hK h_ITu t - -- Map both sides by envToBestArm - have h_map_indep : (condDistrib envf histf P).map (envToBestArm κ) - =ᵐ[P.map histf] (condDistrib envfu histfu Pu).map (envToBestArm κ) := by - filter_upwards [h_env_indep] with x hx - simp only [Kernel.map_apply _ (measurable_envToBestArm κ)] - rw [hx] - -- Transfer h_comp_unif from ae[Pu.map histfu] to ae[P.map histf] - exact h_comp_alg.trans (h_map_indep.trans - (h_comp_unif.filter_mono - (absolutelyContinuous_map_hist_uniform Q κ h hK h_ITu t).ae_le).symm) + IT.bayesTrajMeasurePosterior Q κ (Bandits.uniformAlgorithm hK) t := + condDistrib_env_hist_alg_indep Q κ h hK + (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (Bandits.uniformAlgorithm hK)) t end PosteriorIndependence From ca1c1ec2fb51437fc8b20e40e99d96cf72bfb20d Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 23 Feb 2026 12:56:48 +0000 Subject: [PATCH 045/155] Minor --- .../BayesStationaryEnv.lean | 27 +++++++++---------- 1 file changed, 13 insertions(+), 14 deletions(-) diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index c560fca1..5b5affa5 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -36,7 +36,7 @@ structure IsBayesAlgEnvSeq namespace IsBayesAlgEnvSeq -def trajectory (A : ℕ → Ω → α) (R' : ℕ → Ω → R) : Ω → ℕ → α × R := fun ω n ↦ (A n ω, R' n ω) +def trajectory (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (ω : Ω) : ℕ → α × R := fun n ↦ (A n ω, R' n ω) @[fun_prop] lemma measurable_trajectory {A : ℕ → Ω → α} {R' : ℕ → Ω → R} (hA : ∀ n, Measurable (A n)) @@ -55,7 +55,7 @@ lemma measurable_actionMean {κ : Kernel (α × 𝓔) ℝ} {E : Ω → 𝓔} {a stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop) noncomputable -def bestAction [Fintype α] [Encodable α] [Nonempty α] [MeasurableSingletonClass α] +def bestAction [Nonempty α] [Fintype α] [Encodable α] [MeasurableSingletonClass α] (κ : Kernel (α × 𝓔) ℝ) (E : Ω → 𝓔) (ω : Ω) : α := measurableArgmax (fun ω' a ↦ actionMean κ E a ω') ω @@ -73,7 +73,7 @@ lemma measurable_regret [Countable α] {κ : Kernel (α × 𝓔) ℝ} {E : Ω (hE : Measurable E) (hA : ∀ n, Measurable (A n)) : Measurable (regret κ E A t) := by have hm := (stronglyMeasurable_id.integral_kernel (κ := κ)).measurable - exact (Measurable.const_mul (Measurable.iSup fun _ ↦ hm.comp (by fun_prop)) _).sub + exact (Measurable.const_mul (Measurable.iSup fun _ ↦ (hm.comp (by fun_prop))) _).sub (Finset.measurable_sum _ fun _ _ ↦ hm.comp (by fun_prop)) end Real @@ -86,8 +86,7 @@ variable {P : Measure Ω} [IsFiniteMeasure P] section Laws lemma hasLaw_action_zero [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : - HasLaw (A 0) alg.p0 P := - h.hasCondDistrib_action_zero.hasLaw_of_const + HasLaw (A 0) alg.p0 P := h.hasCondDistrib_action_zero.hasLaw_of_const lemma hasCondDistrib_action' (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) (alg.policy n) P := @@ -103,19 +102,19 @@ section CondDistribIsAlgEnvSeq lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : ∀ᵐ e ∂Q, HasLaw (IT.action 0) alg.p0 (condDistrib (trajectory A R') E P e) := by + have hmt := (measurable_trajectory h.measurable_A h.measurable_R) + have hma : Measurable (IT.action 0) := IT.measurable_action (α := α) (R := R) 0 rw [← h.hasLaw_env.map_eq] - filter_upwards [condDistrib_comp E - (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable (IT.measurable_action 0), - h.hasCondDistrib_action_zero.condDistrib_eq] with e he hcd - exact ⟨(IT.measurable_action 0).aemeasurable, by - rw [← Kernel.map_apply _ (IT.measurable_action 0), ← he, - show IT.action 0 ∘ trajectory A R' = A 0 from rfl, hcd, Kernel.const_apply]⟩ + filter_upwards [condDistrib_comp E (hmt.aemeasurable) hma, + h.hasCondDistrib_action_zero.condDistrib_eq] with e hc hcd + have hat : IT.action 0 ∘ trajectory A R' = A 0 := rfl + exact ⟨hma.aemeasurable, by rw [← Kernel.map_apply _ hma, ← hc, hat, hcd, Kernel.const_apply]⟩ lemma hasCondDistrib_IT_reward_zero [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : ∀ᵐ e ∂Q, HasCondDistrib (IT.reward 0) (IT.action 0) (κ.comap (·, e) (by fun_prop)) (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] - have hW := (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable (μ := P) + have hmt := (measurable_trajectory h.measurable_A h.measurable_R) have h_swap : HasCondDistrib (R' 0) (fun ω ↦ (E ω, A 0 ω)) (κ.comap Prod.swap (by fun_prop)) P := by convert h.hasCondDistrib_reward_zero.comp_right @@ -124,10 +123,10 @@ lemma hasCondDistrib_IT_reward_zero [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q (h.measurable_R 0).aemeasurable h.measurable_E.aemeasurable (μ := P) have h_comp_pair : ⇑(condDistrib (fun ω ↦ (A 0 ω, R' 0 ω)) E P) =ᶠ[ae (P.map E)] ⇑((condDistrib (trajectory A R') E P).map (fun ω ↦ (IT.action 0 ω, IT.reward 0 ω))) := - condDistrib_comp E hW ((IT.measurable_action 0).prodMk (IT.measurable_reward 0)) + condDistrib_comp E hmt.aemeasurable ((IT.measurable_action 0).prodMk (IT.measurable_reward 0)) have h_comp_action : ⇑(condDistrib (A 0) E P) =ᶠ[ae (P.map E)] ⇑((condDistrib (trajectory A R') E P).map (IT.action 0)) := - condDistrib_comp E hW (IT.measurable_action 0) + condDistrib_comp E hmt.aemeasurable (IT.measurable_action 0) have h_swap_eq := h_swap.condDistrib_eq rw [(compProd_map_condDistrib (h.measurable_A 0).aemeasurable).symm] at h_swap_eq filter_upwards [h_prod, h_comp_pair, h_comp_action, From 8f7dfa5a8c660b2a4bfb26497e08d9522080242b Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 23 Feb 2026 14:04:56 +0000 Subject: [PATCH 046/155] Swap kernel --- LeanBandits/BanditAlgorithms/TS.lean | 106 +++++++-------- .../BayesStationaryEnv.lean | 125 +++++++++--------- .../SequentialLearning/HistoryDensity.lean | 8 +- 3 files changed, 118 insertions(+), 121 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index c3658199..daceceef 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -22,7 +22,7 @@ namespace TS variable {K : ℕ} (hK : 0 < K) variable {𝓔 : Type*} [mE : MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] -variable (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (Fin K × 𝓔) ℝ) [IsMarkovKernel κ] +variable (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] /-- The distribution over actions for every given history for TS. -/ noncomputable @@ -61,7 +61,7 @@ are optimal given prior knowledge represented by a prior distribution `Q` and a model represented by a kernel `κ`. -/ noncomputable def tsAlgorithm (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] - (κ : Kernel (Fin K × 𝓔) ℝ) [IsMarkovKernel κ] : Algorithm (Fin K) ℝ where + (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] : Algorithm (Fin K) ℝ where policy := TS.policy hK Q κ p0 := TS.initialPolicy hK Q κ @@ -73,7 +73,7 @@ variable {E : Type*} [mE : MeasurableSpace E] [StandardBorelSpace E] [Nonempty E variable (hK : 0 < K) variable {Ω : Type*} [MeasurableSpace Ω] variable (E' : Ω → E) (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → ℝ) -variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] +variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (E × Fin K) ℝ) [IsMarkovKernel κ] variable (P : Measure Ω) [IsProbabilityMeasure P] noncomputable @@ -176,7 +176,7 @@ lemma measurable_ucbIndex [Nonempty (Fin K)] (measurable_const.div hpc).sqrt))) omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in -lemma armMean_le_ucbIndex {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) +lemma armMean_le_ucbIndex {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hconc : pullCount A a t ω ≠ 0 → |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E' a ω| @@ -192,7 +192,7 @@ lemma armMean_le_ucbIndex {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set. linarith [habs.2] omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in -lemma ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) +lemma ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hpc : pullCount A a t ω ≠ 0) (hconc : |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E' a ω| @@ -241,7 +241,7 @@ lemma le_armMean_bestArm [Nonempty (Fin K)] (ω : Ω) (i : Fin K) : omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] {lo hi : ℝ} - (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc lo hi) + (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) (ω : Ω) : ⨆ i, IsBayesAlgEnvSeq.actionMean κ E' i ω = IsBayesAlgEnvSeq.actionMean κ E' (IsBayesAlgEnvSeq.bestAction κ E' ω) ω := le_antisymm (ciSup_le (le_armMean_bestArm E' κ ω)) @@ -250,33 +250,33 @@ lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] {lo hi : ℝ} omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma gap_eq_armMean_sub [Nonempty (Fin K)] {lo hi : ℝ} - (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc lo hi) - (s : ℕ) (ω : Ω) : gap (κ.comap (·, E' ω) (by fun_prop)) (A s ω) = + (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) + (s : ℕ) (ω : Ω) : gap (κ.sectR (E' ω)) (A s ω) = IsBayesAlgEnvSeq.actionMean κ E' (IsBayesAlgEnvSeq.bestAction κ E' ω) ω - IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω := by - simp only [gap, Kernel.comap_apply] + simp only [gap, Kernel.sectR_apply] exact congr_arg (· - _) (iSup_armMean_eq_bestArm E' κ hm ω) omit [StandardBorelSpace E] [Nonempty E] [IsProbabilityMeasure Q] [IsMarkovKernel κ] in lemma bayesRegret_eq_sum_integral_gap [Nonempty (Fin K)] {alg : Algorithm (Fin K) ℝ} (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) - {C : ℝ} (hm : ∀ a e, |(κ (a, e))[id]| ≤ C) (t : ℕ) : + {C : ℝ} (hm : ∀ a e, |(κ (e, a))[id]| ≤ C) (t : ℕ) : P[IsBayesAlgEnvSeq.regret κ E' A t] = - ∑ s ∈ range t, P[fun ω ↦ gap (κ.comap (·, E' ω) (by fun_prop)) + ∑ s ∈ range t, P[fun ω ↦ gap (κ.sectR (E' ω)) (A s ω)] := by simp only [IsBayesAlgEnvSeq.regret, regret_eq_sum_gap] refine integral_finset_sum _ (fun s _ => ?_) - have hmeas : Measurable (fun ω ↦ gap (κ.comap (·, E' ω) (by fun_prop)) + have hmeas : Measurable (fun ω ↦ gap (κ.sectR (E' ω)) (A s ω)) := (Measurable.iSup (fun a ↦ IsBayesAlgEnvSeq.measurable_actionMean (a := a) h.measurable_E)).sub (stronglyMeasurable_id.integral_kernel.measurable.comp - ((h.measurable_A s).prodMk h.measurable_E)) + (h.measurable_E.prodMk (h.measurable_A s))) refine ⟨hmeas.aestronglyMeasurable, HasFiniteIntegral.of_bounded (C := 2 * C) (Filter.Eventually.of_forall fun ω => ?_)⟩ - simp only [Real.norm_eq_abs, gap, Kernel.comap_apply] - have hbdd : BddAbove (Set.range fun i => (κ (i, E' ω))[id]) := + simp only [Real.norm_eq_abs, gap, Kernel.sectR_apply] + have hbdd : BddAbove (Set.range fun i => (κ (E' ω, i))[id]) := ⟨C, by rintro _ ⟨i, rfl⟩; exact le_of_abs_le (hm i _)⟩ rw [abs_of_nonneg (sub_nonneg.mpr (le_ciSup hbdd _))] linarith [ciSup_le fun i => le_of_abs_le (hm i (E' ω)), @@ -308,7 +308,7 @@ lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsProbabilityMeasure Q] [IsMarkovKernel κ] [IsProbabilityMeasure P] in -lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) +lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (hlo : lo ≤ hi) (σ2 δ : ℝ) (n : ℕ) (ω : Ω) (hconc : ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω| @@ -547,24 +547,24 @@ private lemma streamMeasure_concentration_bound {α : Type*} [MeasurableSpace α lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) - {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc lo hi) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) (hδ_large : max |lo| |hi| < √(2 * ↑σ2 * Real.log (1 / δ))) : ∀ᵐ e ∂(P.map (E')), (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) {ω | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} ≤ + |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} ≤ ENNReal.ofReal (2 * s * δ) := by have h_cond_ae : ∀ᵐ e ∂(P.map E'), IsAlgEnvSeq IT.action IT.reward - (tsAlgorithm hK Q κ) (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (tsAlgorithm hK Q κ) (stationaryEnv (κ.sectR e)) (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E' P e) := by rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h filter_upwards [h_cond_ae] with e h_isAlgEnvSeq - let ν := κ.comap (·, e) (by fun_prop) + let ν := κ.sectR e have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by - simp only [ν, Kernel.comap_apply]; exact hs a' e - have h_mean : (ν a)[id] = (κ (a, e))[id] := by simp only [ν, Kernel.comap_apply] + simp only [ν, Kernel.sectR_apply]; exact hs a' e + have h_mean : (ν a)[id] = (κ (e, a))[id] := by simp only [ν, Kernel.sectR_apply] rw [← h_mean] let P' := condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e have h_law := h_isAlgEnvSeq.law_pullCount_sumRewards_unique' @@ -646,8 +646,8 @@ lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] lemma prob_concentration_single_delta [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) - {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ Set.Icc lo hi) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) (hδ_large : max |lo| |hi| < √(2 * ↑σ2 * Real.log (1 / δ))) : P {ω | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ @@ -655,7 +655,7 @@ lemma prob_concentration_single_delta [Nonempty (Fin K)] ENNReal.ofReal (2 * s * δ) := by let badSet : E → Set (ℕ → (Fin K) × ℝ) := fun e ↦ {t | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s t) : ℝ)) ≤ - |empMean IT.action IT.reward a s t - (κ (a, e))[id]|} + |empMean IT.action IT.reward a s t - (κ (e, a))[id]|} have h_set_eq : {ω | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω|} = @@ -680,12 +680,12 @@ lemma prob_concentration_single_delta [Nonempty (Fin K)] (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable)).symm have h_cond := prob_concentration_single_delta_cond hK E' A R' Q κ P h hσ2 hs hm a s δ hδ hδ1 hδ_large - have h_kernel : Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (a, p.1))[id]) := - stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk measurable_fst) + have h_kernel : Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := + stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_fst.prodMk measurable_const) have h_meas_set : MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | p.2 ∈ badSet p.1} := by change MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s p.2) : ℝ)) ≤ - |empMean IT.action IT.reward a s p.2 - (κ (a, p.1))[id]|} + |empMean IT.action IT.reward a s p.2 - (κ (p.1, a))[id]|} exact measurableSet_le (by fun_prop) (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp measurable_snd).sub h_kernel).abs @@ -709,22 +709,22 @@ lemma prob_concentration_single_delta [Nonempty (Fin K)] private lemma concentration_cond_bound [Nonempty (Fin K)] {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {n : ℕ} (hn : 0 < n) {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) (e : E) (h_isAlgEnvSeq : IsAlgEnvSeq IT.action IT.reward (tsAlgorithm hK Q κ) - (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (stationaryEnv (κ.sectR e)) (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e)) (a : Fin K) : (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) (⋃ s ∈ Finset.range n, {ω | pullCount IT.action a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|}) ≤ + |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|}) ≤ ENNReal.ofReal (2 * n * δ) := by - let ν := κ.comap (·, e) (by fun_prop) + let ν := κ.sectR e let P' := condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by - simp only [ν, Kernel.comap_apply]; exact hs a' e - have h_mean : (ν a)[id] = (κ (a, e))[id] := by simp only [ν, Kernel.comap_apply] + simp only [ν, Kernel.sectR_apply]; exact hs a' e + have h_mean : (ν a)[id] = (κ (e, a))[id] := by simp only [ν, Kernel.sectR_apply] let B_low := fun m : ℕ ↦ {x : ℝ | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} let B_high := fun m : ℕ ↦ @@ -739,7 +739,7 @@ private lemma concentration_cond_bound [Nonempty (Fin K)] let badSetIT := fun (s : ℕ) ↦ {ω : ℕ → (Fin K) × ℝ | pullCount IT.action a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} + |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} let S := Finset.Icc 1 (n - 1) have hS_card : S.card = n - 1 := by simp only [Nat.card_Icc, S]; omega have h_decomp : ⋃ s ∈ Finset.range n, badSetIT s = @@ -822,7 +822,7 @@ private lemma concentration_cond_bound [Nonempty (Fin K)] lemma prob_concentration_fail_delta [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : P {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ @@ -851,7 +851,7 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] let badSetIT := fun (s : ℕ) (e : E) ↦ {ω : ℕ → (Fin K) × ℝ | pullCount IT.action a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} + |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} have h_set_eq : ⋃ s ∈ Finset.range n, badSet s a = (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by @@ -873,14 +873,14 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) (⋃ s ∈ Finset.range n, badSetIT s e) ≤ ENNReal.ofReal (2 * n * δ) := by have h_cond_ae : ∀ᵐ e ∂(P.map E'), IsAlgEnvSeq IT.action IT.reward - (tsAlgorithm hK Q κ) (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (tsAlgorithm hK Q κ) (stationaryEnv (κ.sectR e)) (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E' P e) := by rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h filter_upwards [h_cond_ae] with e h_isAlgEnvSeq exact concentration_cond_bound (hK := hK) (E' := E') (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) hσ2 hs hn' hδ hδ1 e h_isAlgEnvSeq a - have h_kernel : Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (a, p.1))[id]) := - stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk measurable_fst) + have h_kernel : Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := + stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_fst.prodMk measurable_const) have h_meas_set : MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by have h_eq : {p : E × (ℕ → (Fin K) × ℝ) | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} = @@ -892,7 +892,7 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] change MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | pullCount IT.action a s p.2 ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s p.2 : ℝ)) ≤ - |empMean IT.action IT.reward a s p.2 - (κ (a, p.1))[id]|} + |empMean IT.action IT.reward a s p.2 - (κ (p.1, a))[id]|} exact MeasurableSet.inter (((measurable_pullCount IT.measurable_action a s).comp measurable_snd) (measurableSet_singleton (0 : ℕ)).compl) @@ -929,7 +929,7 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : P {ω | ∃ s < n, pullCount A (IsBayesAlgEnvSeq.bestAction κ E' ω) s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / @@ -945,7 +945,7 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] let badSetIT := fun (a : Fin K) (s : ℕ) (e : E) ↦ {ω : ℕ → (Fin K) × ℝ | pullCount IT.action a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (κ (a, e))[id]|} + |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} have h_set_eq : {ω | ∃ s < n, pullCount A ((envToBestArm κ ∘ E') ω) s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A ((envToBestArm κ ∘ E') ω) s ω : ℝ)) ≤ @@ -970,7 +970,7 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) (⋃ s ∈ Finset.range n, badSetIT a s e) ≤ ENNReal.ofReal (2 * n * δ) := by have h_cond_ae : ∀ᵐ e ∂(P.map E'), IsAlgEnvSeq IT.action IT.reward - (tsAlgorithm hK Q κ) (stationaryEnv (κ.comap (·, e) (by fun_prop))) + (tsAlgorithm hK Q κ) (stationaryEnv (κ.sectR e)) (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E' P e) := by rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h filter_upwards [h_cond_ae] with e h_isAlgEnvSeq @@ -983,9 +983,9 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] ENNReal.ofReal (2 * n * δ) := by filter_upwards [h_cond_bound] with e he exact he (envToBestArm κ e) - have h_kernel : ∀ a, Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (a, p.1))[id]) := + have h_kernel : ∀ a, Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := fun a ↦ stronglyMeasurable_id.integral_kernel.measurable.comp - (measurable_const.prodMk measurable_fst) + (measurable_fst.prodMk measurable_const) have h_meas_badSetIT : ∀ a s, MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | p.2 ∈ badSetIT a s p.1} := by intro a s @@ -993,7 +993,7 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] change MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | pullCount IT.action a s p.2 ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s p.2 : ℝ)) ≤ - |empMean IT.action IT.reward a s p.2 - (κ (a, p.1))[id]|} + |empMean IT.action IT.reward a s p.2 - (κ (p.1, a))[id]|} exact MeasurableSet.inter (((measurable_pullCount IT.measurable_action a s).comp measurable_snd) (measurableSet_singleton (0 : ℕ)).compl) @@ -1036,8 +1036,8 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) - {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : P[IsBayesAlgEnvSeq.regret κ E' A n] ≤ (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * n ^ 2 * δ + @@ -1332,8 +1332,8 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (a, e))[id]) σ2 (κ (a, e))) - {lo hi : ℝ} (hm : ∀ a e, (κ (a, e))[id] ∈ (Set.Icc lo hi)) (t : ℕ) : + (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (t : ℕ) : P[IsBayesAlgEnvSeq.regret κ E' A t] ≤ (3 * K + 2) * (hi - lo) + 8 * √(σ2 * K * t * Real.log t) := by have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _) @@ -1349,7 +1349,7 @@ lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] ≤ hi - lo := by unfold IsBayesAlgEnvSeq.regret Bandits.regret simp only [Finset.range_one, Finset.sum_singleton, Nat.cast_one, one_mul, - Kernel.comap_apply] + Kernel.sectR_apply] refine (integral_mono_of_nonneg (ae_of_all _ fun ω ↦ sub_nonneg.mpr (le_ciSup ⟨hi, by rintro _ ⟨a, rfl⟩; exact (hm a _).2⟩ _)) (integrable_const (hi - lo)) (ae_of_all _ fun ω ↦ by diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 5b5affa5..8f9a8a6c 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -18,7 +18,7 @@ variable [MeasurableSpace 𝓔] [MeasurableSpace α] [MeasurableSpace R] [Measur structure IsBayesAlgEnvSeq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (Q : Measure 𝓔) (κ : Kernel (α × 𝓔) R) (alg : Algorithm α R) + (Q : Measure 𝓔) (κ : Kernel (𝓔 × α) R) (alg : Algorithm α R) (E : Ω → 𝓔) (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (P : Measure Ω) [IsFiniteMeasure P] : Prop where measurable_E : Measurable E := by fun_prop @@ -26,12 +26,12 @@ structure IsBayesAlgEnvSeq measurable_R n : Measurable (R' n) := by fun_prop hasLaw_env : HasLaw E Q P hasCondDistrib_action_zero : HasCondDistrib (A 0) E (Kernel.const _ alg.p0) P - hasCondDistrib_reward_zero : HasCondDistrib (R' 0) (fun ω ↦ (A 0 ω, E ω)) κ P + hasCondDistrib_reward_zero : HasCondDistrib (R' 0) (fun ω ↦ (E ω, A 0 ω)) κ P hasCondDistrib_action n : HasCondDistrib (A (n + 1)) (fun ω ↦ (E ω, IsAlgEnvSeq.hist A R' n ω)) ((alg.policy n).prodMkLeft _) P hasCondDistrib_reward n : - HasCondDistrib (R' (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω, E ω)) + HasCondDistrib (R' (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, E ω, A (n + 1) ω)) (κ.prodMkLeft _) P namespace IsBayesAlgEnvSeq @@ -47,29 +47,29 @@ lemma measurable_trajectory {A : ℕ → Ω → α} {R' : ℕ → Ω → R} (hA section Real noncomputable -def actionMean (κ : Kernel (α × 𝓔) ℝ) (E : Ω → 𝓔) (a : α) (ω : Ω) : ℝ := (κ (a, E ω))[id] +def actionMean (κ : Kernel (𝓔 × α) ℝ) (E : Ω → 𝓔) (a : α) (ω : Ω) : ℝ := (κ (E ω, a))[id] @[fun_prop] -lemma measurable_actionMean {κ : Kernel (α × 𝓔) ℝ} {E : Ω → 𝓔} {a : α} (hE : Measurable E) : +lemma measurable_actionMean {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} {a : α} (hE : Measurable E) : Measurable (actionMean κ E a) := stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop) noncomputable def bestAction [Nonempty α] [Fintype α] [Encodable α] [MeasurableSingletonClass α] - (κ : Kernel (α × 𝓔) ℝ) (E : Ω → 𝓔) (ω : Ω) : α := + (κ : Kernel (𝓔 × α) ℝ) (E : Ω → 𝓔) (ω : Ω) : α := measurableArgmax (fun ω' a ↦ actionMean κ E a ω') ω @[fun_prop] lemma measurable_bestAction [Nonempty α] [Fintype α] [Encodable α] [MeasurableSingletonClass α] - {κ : Kernel (α × 𝓔) ℝ} {E : Ω → 𝓔} (hE : Measurable E) : Measurable (bestAction κ E) := + {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} (hE : Measurable E) : Measurable (bestAction κ E) := measurable_measurableArgmax (by fun_prop) noncomputable -def regret (κ : Kernel (α × 𝓔) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → α) (t : ℕ) (ω : Ω) : ℝ := - Bandits.regret (κ.comap (·, E ω) (by fun_prop)) A t ω +def regret (κ : Kernel (𝓔 × α) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → α) (t : ℕ) (ω : Ω) : ℝ := + Bandits.regret (κ.sectR (E ω)) A t ω @[fun_prop] -lemma measurable_regret [Countable α] {κ : Kernel (α × 𝓔) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → α} {t : ℕ} +lemma measurable_regret [Countable α] {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → α} {t : ℕ} (hE : Measurable E) (hA : ∀ n, Measurable (A n)) : Measurable (regret κ E A t) := by have hm := (stronglyMeasurable_id.integral_kernel (κ := κ)).measurable @@ -79,7 +79,7 @@ lemma measurable_regret [Countable α] {κ : Kernel (α × 𝓔) ℝ} {E : Ω end Real variable [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] -variable {Q : Measure 𝓔} {κ : Kernel (α × 𝓔) R} {alg : Algorithm α R} +variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × α) R} {alg : Algorithm α R} variable {E : Ω → 𝓔} {A : ℕ → Ω → α} {R' : ℕ → Ω → R} variable {P : Measure Ω} [IsFiniteMeasure P] @@ -93,7 +93,7 @@ lemma hasCondDistrib_action' (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) (h.hasCondDistrib_action n).comp_left (by fun_prop) lemma hasCondDistrib_reward' [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : - HasCondDistrib (R' (n + 1)) (fun ω ↦ (A (n + 1) ω, E ω)) κ P := + HasCondDistrib (R' (n + 1)) (fun ω ↦ (E ω, A (n + 1) ω)) κ P := (h.hasCondDistrib_reward n).comp_left (by fun_prop) end Laws @@ -111,26 +111,23 @@ lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : exact ⟨hma.aemeasurable, by rw [← Kernel.map_apply _ hma, ← hc, hat, hcd, Kernel.const_apply]⟩ lemma hasCondDistrib_IT_reward_zero [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : - ∀ᵐ e ∂Q, HasCondDistrib (IT.reward 0) (IT.action 0) (κ.comap (·, e) (by fun_prop)) + ∀ᵐ e ∂Q, HasCondDistrib (IT.reward 0) (IT.action 0) (κ.sectR e) (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] have hmt := (measurable_trajectory h.measurable_A h.measurable_R) - have h_swap : HasCondDistrib (R' 0) (fun ω ↦ (E ω, A 0 ω)) - (κ.comap Prod.swap (by fun_prop)) P := by - convert h.hasCondDistrib_reward_zero.comp_right - (MeasurableEquiv.prodComm : α × 𝓔 ≃ᵐ 𝓔 × α) using 2 have h_prod := condDistrib_prod_left (h.measurable_A 0).aemeasurable (h.measurable_R 0).aemeasurable h.measurable_E.aemeasurable (μ := P) have h_comp_pair : ⇑(condDistrib (fun ω ↦ (A 0 ω, R' 0 ω)) E P) =ᶠ[ae (P.map E)] ⇑((condDistrib (trajectory A R') E P).map (fun ω ↦ (IT.action 0 ω, IT.reward 0 ω))) := - condDistrib_comp E hmt.aemeasurable ((IT.measurable_action 0).prodMk (IT.measurable_reward 0)) + condDistrib_comp E hmt.aemeasurable + ((IT.measurable_action 0).prodMk (IT.measurable_reward 0)) have h_comp_action : ⇑(condDistrib (A 0) E P) =ᶠ[ae (P.map E)] ⇑((condDistrib (trajectory A R') E P).map (IT.action 0)) := condDistrib_comp E hmt.aemeasurable (IT.measurable_action 0) - have h_swap_eq := h_swap.condDistrib_eq - rw [(compProd_map_condDistrib (h.measurable_A 0).aemeasurable).symm] at h_swap_eq + have h_reward_eq := h.hasCondDistrib_reward_zero.condDistrib_eq + rw [(compProd_map_condDistrib (h.measurable_A 0).aemeasurable).symm] at h_reward_eq filter_upwards [h_prod, h_comp_pair, h_comp_action, - (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_swap_eq] + (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_reward_eq] with e h_prod_e h_pair_e h_act_e h_nested_e refine ⟨by fun_prop, by fun_prop, ?_⟩ rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] @@ -140,7 +137,7 @@ lemma hasCondDistrib_IT_reward_zero [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q refine Measure.compProd_congr ?_ filter_upwards [h_nested_e] with a ha ext s _ - rw [Kernel.sectR_apply, Kernel.comap_apply, ha, Kernel.comap_apply]; rfl + simp_rw [Kernel.sectR_apply, ha] lemma hasCondDistrib_IT_action (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : ∀ᵐ e ∂Q, HasCondDistrib (IT.action (n + 1)) (IT.hist n) (alg.policy n) @@ -175,48 +172,46 @@ lemma hasCondDistrib_IT_action (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ lemma hasCondDistrib_IT_reward [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : ∀ᵐ e ∂Q, HasCondDistrib (IT.reward (n + 1)) (fun x ↦ (IT.hist n x, IT.action (n + 1) x)) - ((κ.comap (·, e) (by fun_prop)).prodMkLeft _) (condDistrib (trajectory A R') E P e) := by + ((κ.sectR e).prodMkLeft _) (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] - have hW := (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable (μ := P) - have h_hist_meas := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n - have h_prod := condDistrib_prod_left - (Measurable.prodMk h_hist_meas (h.measurable_A (n + 1))).aemeasurable - (h.measurable_R (n + 1)).aemeasurable h.measurable_E.aemeasurable (μ := P) - have h_swap : HasCondDistrib (R' (n + 1)) - (fun ω ↦ (E ω, IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) - (κ.comap (fun p ↦ (p.2.2, p.1)) (by fun_prop)) P := - (h.hasCondDistrib_reward n).comp_right - (MeasurableEquiv.prodAssoc.symm.trans MeasurableEquiv.prodComm) - have h_swap_eq := h_swap.condDistrib_eq - have h_pair_meas : Measurable + have hmt := measurable_trajectory h.measurable_A h.measurable_R + have hm := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n + have hm_pair : Measurable (fun f : ℕ → α × R ↦ (IT.hist n f, IT.action (n + 1) f)) := by fun_prop + have h_reorder : HasCondDistrib (R' (n + 1)) + (fun ω ↦ (E ω, IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) + (κ.comap (fun (e, _, a) ↦ (e, a)) (by fun_prop)) P := by + convert (h.hasCondDistrib_reward n).comp_right (MeasurableEquiv.prodAssoc.symm.trans + ((MeasurableEquiv.prodComm.prodCongr (MeasurableEquiv.refl α)).trans + MeasurableEquiv.prodAssoc)) using 2 + have h_eq := h_reorder.condDistrib_eq + rw [(compProd_map_condDistrib (hm.prodMk (h.measurable_A (n + 1))).aemeasurable).symm] at h_eq have h_comp_triple : ⇑(condDistrib (fun ω ↦ ((IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω), R' (n + 1) ω)) E P) =ᶠ[ae (P.map E)] ⇑((condDistrib (trajectory A R') E P).map - (fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), - IT.reward (n + 1) ω))) := - condDistrib_comp E hW (h_pair_meas.prodMk (IT.measurable_reward (n + 1))) + (fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), IT.reward (n + 1) ω))) := + condDistrib_comp E hmt.aemeasurable (hm_pair.prodMk (IT.measurable_reward (n + 1))) have h_comp_pair : ⇑(condDistrib (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) E P) =ᶠ[ae (P.map E)] ⇑((condDistrib (trajectory A R') E P).map (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω))) := - condDistrib_comp E hW h_pair_meas - rw [(compProd_map_condDistrib (Measurable.prodMk h_hist_meas - (h.measurable_A (n + 1))).aemeasurable).symm] at h_swap_eq - filter_upwards [h_prod, h_comp_triple, h_comp_pair, - (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_swap_eq] - with e h_prod_e h_triple_e h_pair_e h_nested_e + condDistrib_comp E hmt.aemeasurable hm_pair + filter_upwards [ + condDistrib_prod_left (hm.prodMk (h.measurable_A (n + 1))).aemeasurable + (h.measurable_R (n + 1)).aemeasurable h.measurable_E.aemeasurable (μ := P), + h_comp_triple, h_comp_pair, + (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_eq] + with e h_prod h_triple h_pair h_inner refine ⟨by fun_prop, by fun_prop, ?_⟩ rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] - rw [← Kernel.map_apply _ (h_pair_meas.prodMk (IT.measurable_reward (n + 1))), ← h_triple_e] - conv_rhs => rw [← Kernel.map_apply _ h_pair_meas, ← h_pair_e] - rw [h_prod_e, Kernel.compProd_apply_eq_compProd_sectR] - refine Measure.compProd_congr ?_ - filter_upwards [h_nested_e] with _ ha - ext s _ - rw [Kernel.sectR_apply, ha, Kernel.comap_apply, Kernel.prodMkLeft_apply, Kernel.comap_apply] + rw [← Kernel.map_apply _ (hm_pair.prodMk (IT.measurable_reward (n + 1))), ← h_triple] + conv_rhs => rw [← Kernel.map_apply _ hm_pair, ← h_pair] + rw [h_prod, Kernel.compProd_apply_eq_compProd_sectR] + exact Measure.compProd_congr (by + filter_upwards [h_inner] with a ha; ext s _ + simp only [Kernel.sectR_apply, ha, Kernel.comap_apply, Kernel.prodMkLeft_apply]) lemma ae_IsAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : - ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv (κ.comap (·, e) (by fun_prop))) + ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv (κ.sectR e)) (condDistrib (trajectory A R') E P e) := by filter_upwards [hasLaw_IT_action_zero h, hasCondDistrib_IT_reward_zero h, ae_all_iff.2 (hasCondDistrib_IT_action h), ae_all_iff.2 (hasCondDistrib_IT_reward h)] @@ -235,16 +230,16 @@ end IsBayesAlgEnvSeq section IsAlgEnvSeq noncomputable -def bayesStationaryEnv (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) R) +def bayesStationaryEnv (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × α) R) [IsMarkovKernel κ] : Environment α (𝓔 × R) where feedback n := - let g : (Iic n → α × 𝓔 × R) × α → α × 𝓔 := fun (h, a) => (a, (h ⟨0, by simp⟩).2.1) - (Kernel.deterministic (Prod.snd ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) - ν0 := (Kernel.const _ Q) ⊗ₖ κ + let g : (Iic n → α × 𝓔 × R) × α → 𝓔 × α := fun (h, a) => ((h ⟨0, by simp⟩).2.1, a) + (Kernel.deterministic (Prod.fst ∘ g) (by fun_prop)) ×ₖ (κ.comap g (by fun_prop)) + ν0 := (Kernel.const _ Q) ⊗ₖ κ.swapLeft variable [Nonempty α] [Nonempty 𝓔] [Nonempty R] variable [StandardBorelSpace α] [StandardBorelSpace 𝓔] [StandardBorelSpace R] -variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (α × 𝓔) R} [IsMarkovKernel κ] +variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × α) R} [IsMarkovKernel κ] variable {alg : Algorithm α R} {A : ℕ → Ω → α} {R' : ℕ → Ω → 𝓔 × R} variable {P : Measure Ω} [IsProbabilityMeasure P] @@ -271,7 +266,9 @@ lemma IsAlgEnvSeq.isBayesAlgEnvSeq simp only [h.hasLaw_action_zero.map_eq, Algorithm.prod_left] at hcd exact ⟨(h.measurable_A 0).aemeasurable, (h.measurable_R 0).fst.aemeasurable, hcd⟩ hasCondDistrib_reward_zero := by - simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.of_compProd + have h0 := h.hasCondDistrib_reward_zero + simp only [bayesStationaryEnv] at h0 + convert h0.of_compProd.comp_right (MeasurableEquiv.prodComm : α × 𝓔 ≃ᵐ 𝓔 × α) using 2 hasCondDistrib_action n := by let f : (Iic n → α × 𝓔 × R) → 𝓔 × (Iic n → α × R) := fun h ↦ ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2)) @@ -280,12 +277,12 @@ lemma IsAlgEnvSeq.isBayesAlgEnvSeq h'.comp_left (f := f) exact h.hasCondDistrib_action n hasCondDistrib_reward n := by - let f : (Iic n → α × 𝓔 × R) × α → (Iic n → α × R) × α × 𝓔 := - fun p ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), p.2, (p.1 ⟨0, by simp⟩).2.1) + let f : (Iic n → α × 𝓔 × R) × α → (Iic n → α × R) × 𝓔 × α := + fun p ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), (p.1 ⟨0, by simp⟩).2.1, p.2) have hf : Measurable f := by fun_prop suffices h' : HasCondDistrib (fun ω ↦ (R' (n + 1) ω).2) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) - ((κ.comap Prod.snd (by fun_prop)).comap f hf) P from h'.comp_left hf + ((Kernel.prodMkLeft (↥(Iic n) → α × R) κ).comap f hf) P from h'.comp_left hf simpa [bayesStationaryEnv, Kernel.snd_prod] using (h.hasCondDistrib_reward n).snd end IsAlgEnvSeq @@ -293,7 +290,7 @@ end IsAlgEnvSeq namespace IT noncomputable -def bayesTrajMeasure (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) R) +def bayesTrajMeasure (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × α) R) [IsMarkovKernel κ] (alg : Algorithm α R) : Measure (ℕ → α × 𝓔 × R) := trajMeasure (alg.prod_left 𝓔) (bayesStationaryEnv Q κ) deriving IsProbabilityMeasure @@ -302,7 +299,7 @@ lemma isBayesAlgEnvSeq_bayesTrajMeasure [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace 𝓔] [Nonempty 𝓔] [StandardBorelSpace R] [Nonempty R] - (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) R) [IsMarkovKernel κ] + (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × α) R) [IsMarkovKernel κ] (alg : Algorithm α R) : IsBayesAlgEnvSeq Q κ alg (fun ω ↦ (ω 0).2.1) action (fun n ω ↦ (ω n).2.2) (bayesTrajMeasure Q κ alg) := @@ -310,7 +307,7 @@ lemma isBayesAlgEnvSeq_bayesTrajMeasure noncomputable def bayesTrajMeasurePosterior [StandardBorelSpace 𝓔] [Nonempty 𝓔] - (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (α × 𝓔) ℝ) [IsMarkovKernel κ] + (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × α) ℝ) [IsMarkovKernel κ] (alg : Algorithm α ℝ) (n : ℕ) : Kernel (Iic n → α × ℝ) 𝓔 := condDistrib (fun ω ↦ (ω 0).2.1) (IsAlgEnvSeq.hist action (fun n ω ↦ (ω n).2.2) n) (bayesTrajMeasure Q κ alg) diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 6719e9ef..9dfdc41d 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -378,20 +378,20 @@ transfers to the posterior on the best arm via `condDistrib_comp`. variable {K : ℕ} [Nonempty (Fin K)] variable {E : Type*} [MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] variable (Q : Measure E) [IsProbabilityMeasure Q] -variable (κ : Kernel (Fin K × E) ℝ) [IsMarkovKernel κ] +variable (κ : Kernel (E × Fin K) ℝ) [IsMarkovKernel κ] variable {Ω : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] variable {E' : Ω → E} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} variable {alg : Algorithm (Fin K) ℝ} variable {P : Measure Ω} [IsProbabilityMeasure P] /-- Maps an environment to the best arm (the arm with highest mean reward). -/ -noncomputable def envToBestArm (κ : Kernel (Fin K × E) ℝ) : E → Fin K := - measurableArgmax fun e a ↦ (κ (a, e))[id] +noncomputable def envToBestArm (κ : Kernel (E × Fin K) ℝ) : E → Fin K := + measurableArgmax fun e a ↦ (κ (e, a))[id] omit [StandardBorelSpace E] [Nonempty E] [IsMarkovKernel κ] in lemma measurable_envToBestArm : Measurable (envToBestArm κ) := measurable_measurableArgmax fun _ ↦ - stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_const.prodMk measurable_id) + stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_id.prodMk measurable_const) omit [StandardBorelSpace E] [Nonempty E] [IsMarkovKernel κ] [MeasurableSpace Ω] [IsProbabilityMeasure P] [Nonempty Ω] in From d3e9b9956d25814535916f11fd2e55bb52b827fc Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 23 Feb 2026 16:36:41 +0000 Subject: [PATCH 047/155] Refactor BayesStationaryEnv (in progress) --- LeanBandits/ForMathlib/CondDistrib.lean | 13 +- LeanBandits/ForMathlib/HasCondDistrib.lean | 59 ++++++++ .../BayesStationaryEnv.lean | 126 ++++-------------- 3 files changed, 99 insertions(+), 99 deletions(-) diff --git a/LeanBandits/ForMathlib/CondDistrib.lean b/LeanBandits/ForMathlib/CondDistrib.lean index 9d57f6de..0f76b355 100644 --- a/LeanBandits/ForMathlib/CondDistrib.lean +++ b/LeanBandits/ForMathlib/CondDistrib.lean @@ -24,6 +24,16 @@ section CondDistrib variable [IsFiniteMeasure μ] +/-- An ae-equality of kernels wrt the joint law `μ.map (X, Y)` is equivalent to an ae-equality +fiberwise via the conditional distribution of `Y` given `X`. -/ +lemma Kernel.ae_eq_map_prod_iff_ae_condDistrib + [MeasurableSpace.CountableOrCountablyGenerated (β × Ω) δ] + (hY : AEMeasurable Y μ) {f g : Kernel (β × Ω) δ} [IsFiniteKernel f] [IsFiniteKernel g] : + f =ᵐ[μ.map (fun ω ↦ (X ω, Y ω))] g ↔ + ∀ᵐ x ∂(μ.map X), ∀ᵐ y ∂(condDistrib Y X μ x), f (x, y) = g (x, y) := by + rw [← compProd_map_condDistrib hY] + exact Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _) + lemma condDistrib_prod_left [StandardBorelSpace β] [Nonempty β] (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) (hT : AEMeasurable T μ) : condDistrib (fun ω ↦ (X ω, Y ω)) T μ @@ -39,8 +49,7 @@ lemma condDistrib_prod_self_left [StandardBorelSpace β] [Nonempty β] [Standard 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 + have h_fst' := (Kernel.ae_eq_map_prod_iff_ae_condDistrib hX).mp h_fst filter_upwards [h_prod, h_fst'] with z hz1 hz2 rw [hz1] simp only [Kernel.deterministic_apply] at hz2 diff --git a/LeanBandits/ForMathlib/HasCondDistrib.lean b/LeanBandits/ForMathlib/HasCondDistrib.lean index 5d427c54..42fb4b6d 100644 --- a/LeanBandits/ForMathlib/HasCondDistrib.lean +++ b/LeanBandits/ForMathlib/HasCondDistrib.lean @@ -195,6 +195,18 @@ lemma HasCondDistrib.hasLaw_of_const {Q : Measure Ω} (h.aemeasurable_snd.prodMk h.aemeasurable_fst)] rfl +lemma HasCondDistrib.swap_const {Q : Measure Ω} + [StandardBorelSpace β] [Nonempty β] + [IsProbabilityMeasure μ] [IsFiniteMeasure Q] + (h : HasCondDistrib Y X (Kernel.const β Q) μ) : + HasCondDistrib X Y (Kernel.const Ω (μ.map X)) μ := by + have h_indep : IndepFun X Y μ := by + rw [indepFun_iff_condDistrib_eq_const h.aemeasurable_snd h.aemeasurable_fst, + h.hasLaw_of_const.map_eq] + exact h.condDistrib_eq + exact ⟨h.aemeasurable_snd, h.aemeasurable_fst, + condDistrib_of_indepFun h_indep.symm h.aemeasurable_fst h.aemeasurable_snd⟩ + lemma HasLaw.prod_of_hasCondDistrib {P : Measure β} [IsFiniteMeasure μ] [IsSFiniteKernel κ] (h1 : HasLaw X P μ) (h2 : HasCondDistrib Y X κ μ) : HasLaw (fun ω ↦ (X ω, Y ω)) (P ⊗ₘ κ) μ := by @@ -279,4 +291,51 @@ lemma HasCondDistrib.comp_left [IsFiniteMeasure μ] [IsFiniteKernel κ] {f : γ rw [Kernel.comap_apply] congr 1 +/-- Transfer a `HasCondDistrib` from the outer probability space to the conditional distribution +of `W` given `Z`. If `g ∘ W` is conditionally distributed as `η` given `(Z, f ∘ W)`, then in the +conditional space given `Z = z`, `g` is conditionally distributed as `η.sectR z` given `f`. -/ +lemma HasCondDistrib.ae_hasCondDistrib_sectR [IsFiniteMeasure μ] + [StandardBorelSpace β] [Nonempty β] + {δ : Type*} [MeasurableSpace δ] [StandardBorelSpace δ] [Nonempty δ] + {W : α → δ} {Z : α → γ} + {f : δ → β} {g : δ → Ω} + {η : Kernel (γ × β) Ω} [IsFiniteKernel η] + (hf : Measurable f) (hg : Measurable g) + (hW : AEMeasurable W μ) (hZ : AEMeasurable Z μ) + (hcd : HasCondDistrib (g ∘ W) (fun ω ↦ (Z ω, (f ∘ W) ω)) η μ) : + ∀ᵐ z ∂(μ.map Z), HasCondDistrib g f (η.sectR z) (condDistrib W Z μ z) := by + have hfW := hf.comp_aemeasurable hW + have h_prod := condDistrib_prod_left hfW (hg.comp_aemeasurable hW) hZ (μ := μ) + have h_comp_pair : (condDistrib (fun ω ↦ ((f ∘ W) ω, (g ∘ W) ω)) Z μ) + =ᵐ[μ.map Z] (condDistrib W Z μ).map (fun w ↦ (f w, g w)) := + condDistrib_comp Z hW (hf.prodMk hg) + have h_comp_fst : (condDistrib (f ∘ W) Z μ) + =ᵐ[μ.map Z] (condDistrib W Z μ).map f := + condDistrib_comp Z hW hf + have h_nested := (Kernel.ae_eq_map_prod_iff_ae_condDistrib hfW).mp hcd.condDistrib_eq + filter_upwards [h_prod, h_comp_pair, h_comp_fst, h_nested] + with z h_prod_z h_pair_z h_fst_z h_nested_z + refine ⟨hg.aemeasurable, hf.aemeasurable, ?_⟩ + rw [condDistrib_ae_eq_iff_measure_eq_compProd f hg.aemeasurable, + ← Kernel.map_apply _ (hf.prodMk hg), ← h_pair_z, + ← Kernel.map_apply _ hf, ← h_fst_z, + h_prod_z, Kernel.compProd_apply_eq_compProd_sectR] + exact Measure.compProd_congr (h_nested_z.mono fun a ha ↦ by + simp only [Kernel.sectR_apply]; exact ha) + +/-- Variant of `ae_hasCondDistrib_sectR` where `Z` appears second in the conditioning pair. +If `g ∘ W` is conditionally distributed as `η` given `(f ∘ W, Z)`, then in the conditional space +given `Z = z`, `g` is conditionally distributed as `η.sectL z` given `f`. -/ +lemma HasCondDistrib.ae_hasCondDistrib_sectL [IsFiniteMeasure μ] + [StandardBorelSpace β] [Nonempty β] + {δ : Type*} [MeasurableSpace δ] [StandardBorelSpace δ] [Nonempty δ] + {W : α → δ} {Z : α → γ} + {f : δ → β} {g : δ → Ω} + {η : Kernel (β × γ) Ω} [IsFiniteKernel η] + (hf : Measurable f) (hg : Measurable g) + (hW : AEMeasurable W μ) (hZ : AEMeasurable Z μ) + (hcd : HasCondDistrib (g ∘ W) (fun ω ↦ ((f ∘ W) ω, Z ω)) η μ) : + ∀ᵐ z ∂(μ.map Z), HasCondDistrib g f (η.sectL z) (condDistrib W Z μ z) := + (hcd.comp_right .prodComm).ae_hasCondDistrib_sectR hf hg hW hZ + end ProbabilityTheory diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 8f9a8a6c..3228c840 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -102,73 +102,33 @@ section CondDistribIsAlgEnvSeq lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : ∀ᵐ e ∂Q, HasLaw (IT.action 0) alg.p0 (condDistrib (trajectory A R') E P e) := by - have hmt := (measurable_trajectory h.measurable_A h.measurable_R) - have hma : Measurable (IT.action 0) := IT.measurable_action (α := α) (R := R) 0 rw [← h.hasLaw_env.map_eq] - filter_upwards [condDistrib_comp E (hmt.aemeasurable) hma, + filter_upwards [condDistrib_comp E + (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable + (IT.measurable_action (α := α) (R := R) 0), h.hasCondDistrib_action_zero.condDistrib_eq] with e hc hcd - have hat : IT.action 0 ∘ trajectory A R' = A 0 := rfl - exact ⟨hma.aemeasurable, by rw [← Kernel.map_apply _ hma, ← hc, hat, hcd, Kernel.const_apply]⟩ + exact ⟨(IT.measurable_action 0).aemeasurable, by + rw [← Kernel.map_apply _ (IT.measurable_action 0), ← hc, + show IT.action 0 ∘ trajectory A R' = A 0 from rfl, hcd, Kernel.const_apply]⟩ lemma hasCondDistrib_IT_reward_zero [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : ∀ᵐ e ∂Q, HasCondDistrib (IT.reward 0) (IT.action 0) (κ.sectR e) (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] - have hmt := (measurable_trajectory h.measurable_A h.measurable_R) - have h_prod := condDistrib_prod_left (h.measurable_A 0).aemeasurable - (h.measurable_R 0).aemeasurable h.measurable_E.aemeasurable (μ := P) - have h_comp_pair : ⇑(condDistrib (fun ω ↦ (A 0 ω, R' 0 ω)) E P) =ᶠ[ae (P.map E)] - ⇑((condDistrib (trajectory A R') E P).map (fun ω ↦ (IT.action 0 ω, IT.reward 0 ω))) := - condDistrib_comp E hmt.aemeasurable - ((IT.measurable_action 0).prodMk (IT.measurable_reward 0)) - have h_comp_action : ⇑(condDistrib (A 0) E P) =ᶠ[ae (P.map E)] - ⇑((condDistrib (trajectory A R') E P).map (IT.action 0)) := - condDistrib_comp E hmt.aemeasurable (IT.measurable_action 0) - have h_reward_eq := h.hasCondDistrib_reward_zero.condDistrib_eq - rw [(compProd_map_condDistrib (h.measurable_A 0).aemeasurable).symm] at h_reward_eq - filter_upwards [h_prod, h_comp_pair, h_comp_action, - (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_reward_eq] - with e h_prod_e h_pair_e h_act_e h_nested_e - refine ⟨by fun_prop, by fun_prop, ?_⟩ - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] - rw [← Kernel.map_apply _ (by fun_prop), ← h_pair_e] - conv_rhs => rw [← Kernel.map_apply _ (IT.measurable_action 0), ← h_act_e] - rw [h_prod_e, Kernel.compProd_apply_eq_compProd_sectR] - refine Measure.compProd_congr ?_ - filter_upwards [h_nested_e] with a ha - ext s _ - simp_rw [Kernel.sectR_apply, ha] + exact h.hasCondDistrib_reward_zero.ae_hasCondDistrib_sectR + (IT.measurable_action 0) (IT.measurable_reward 0) + (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable + h.measurable_E.aemeasurable lemma hasCondDistrib_IT_action (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : ∀ᵐ e ∂Q, HasCondDistrib (IT.action (n + 1)) (IT.hist n) (alg.policy n) (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] - have hW := (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable (μ := P) - have h_hist_meas := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n - have h_prod := condDistrib_prod_left h_hist_meas.aemeasurable - (h.measurable_A (n + 1)).aemeasurable h.measurable_E.aemeasurable (μ := P) - have h_action_env := (h.hasCondDistrib_action n).condDistrib_eq - have h_comp_pair : ⇑(condDistrib (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) E P) - =ᶠ[ae (P.map E)] ⇑((condDistrib (trajectory A R') E P).map - (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω))) := - condDistrib_comp E hW ((IT.measurable_hist n).prodMk (IT.measurable_action (n + 1))) - have h_comp_hist : ⇑(condDistrib (IsAlgEnvSeq.hist A R' n) E P) =ᶠ[ae (P.map E)] - ⇑((condDistrib (trajectory A R') E P).map (IT.hist n)) := - condDistrib_comp E hW (IT.measurable_hist n) - rw [(compProd_map_condDistrib h_hist_meas.aemeasurable).symm] at h_action_env - filter_upwards [h_prod, h_comp_pair, h_comp_hist, - (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_action_env] - with e h_prod_e h_pair_e h_hist_e h_nested_e - refine ⟨by fun_prop, by fun_prop, ?_⟩ - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] - rw [← Kernel.map_apply _ ((IT.measurable_hist n).prodMk (IT.measurable_action (n + 1))), - ← h_pair_e] - conv_rhs => rw [← Kernel.map_apply _ (IT.measurable_hist n), ← h_hist_e] - rw [h_prod_e, Kernel.compProd_apply_eq_compProd_sectR] - refine Measure.compProd_congr ?_ - filter_upwards [h_nested_e] with _ ha - ext s _ - rw [Kernel.sectR_apply, ha, Kernel.prodMkLeft_apply] + filter_upwards [(h.hasCondDistrib_action n).ae_hasCondDistrib_sectR + (IT.measurable_hist n) (IT.measurable_action (n + 1)) + (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable + h.measurable_E.aemeasurable] with e he + rwa [Kernel.sectR_prodMkLeft] at he lemma hasCondDistrib_IT_reward [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : ∀ᵐ e ∂Q, HasCondDistrib (IT.reward (n + 1)) (fun x ↦ (IT.hist n x, IT.action (n + 1) x)) @@ -176,39 +136,20 @@ lemma hasCondDistrib_IT_reward [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ al rw [← h.hasLaw_env.map_eq] have hmt := measurable_trajectory h.measurable_A h.measurable_R have hm := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n - have hm_pair : Measurable - (fun f : ℕ → α × R ↦ (IT.hist n f, IT.action (n + 1) f)) := by fun_prop have h_reorder : HasCondDistrib (R' (n + 1)) - (fun ω ↦ (E ω, IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) - (κ.comap (fun (e, _, a) ↦ (e, a)) (by fun_prop)) P := by - convert (h.hasCondDistrib_reward n).comp_right (MeasurableEquiv.prodAssoc.symm.trans - ((MeasurableEquiv.prodComm.prodCongr (MeasurableEquiv.refl α)).trans - MeasurableEquiv.prodAssoc)) using 2 - have h_eq := h_reorder.condDistrib_eq - rw [(compProd_map_condDistrib (hm.prodMk (h.measurable_A (n + 1))).aemeasurable).symm] at h_eq - have h_comp_triple : ⇑(condDistrib - (fun ω ↦ ((IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω), R' (n + 1) ω)) E P) - =ᶠ[ae (P.map E)] ⇑((condDistrib (trajectory A R') E P).map - (fun ω ↦ ((IT.hist n ω, IT.action (n + 1) ω), IT.reward (n + 1) ω))) := - condDistrib_comp E hmt.aemeasurable (hm_pair.prodMk (IT.measurable_reward (n + 1))) - have h_comp_pair : ⇑(condDistrib (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) E P) - =ᶠ[ae (P.map E)] ⇑((condDistrib (trajectory A R') E P).map - (fun ω ↦ (IT.hist n ω, IT.action (n + 1) ω))) := - condDistrib_comp E hmt.aemeasurable hm_pair - filter_upwards [ - condDistrib_prod_left (hm.prodMk (h.measurable_A (n + 1))).aemeasurable - (h.measurable_R (n + 1)).aemeasurable h.measurable_E.aemeasurable (μ := P), - h_comp_triple, h_comp_pair, - (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_eq] - with e h_prod h_triple h_pair h_inner - refine ⟨by fun_prop, by fun_prop, ?_⟩ - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] - rw [← Kernel.map_apply _ (hm_pair.prodMk (IT.measurable_reward (n + 1))), ← h_triple] - conv_rhs => rw [← Kernel.map_apply _ hm_pair, ← h_pair] - rw [h_prod, Kernel.compProd_apply_eq_compProd_sectR] - exact Measure.compProd_congr (by - filter_upwards [h_inner] with a ha; ext s _ - simp only [Kernel.sectR_apply, ha, Kernel.comap_apply, Kernel.prodMkLeft_apply]) + (fun ω ↦ ((IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω), E ω)) + ((κ.prodMkLeft _).comap (fun ((h, a), e) ↦ (h, (e, a))) (by fun_prop)) P := by + convert (h.hasCondDistrib_reward n).comp_right + ((MeasurableEquiv.prodCongr (.refl _) .prodComm).trans MeasurableEquiv.prodAssoc.symm) using 2 + filter_upwards [h_reorder.ae_hasCondDistrib_sectL + ((IT.measurable_hist n).prodMk (IT.measurable_action (n + 1))) + (IT.measurable_reward (n + 1)) + hmt.aemeasurable h.measurable_E.aemeasurable] with e he + have hk : ((κ.prodMkLeft _).comap (fun ((h, a), e) ↦ (h, (e, a))) (by fun_prop)).sectL e = + (κ.sectR e).prodMkLeft (↥(Iic n) → α × R) := + Kernel.ext fun ⟨_, a⟩ ↦ by + simp [Kernel.sectL_apply, Kernel.comap_apply, Kernel.prodMkLeft_apply] + rw [hk] at he; exact he lemma ae_IsAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv (κ.sectR e)) @@ -255,16 +196,7 @@ lemma IsAlgEnvSeq.isBayesAlgEnvSeq hasCondDistrib_action_zero := by have hfst : HasCondDistrib (fun ω ↦ (R' 0 ω).1) (A 0) (Kernel.const α Q) P := by simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst - -- E' | A 0 is constant Q = P.map E', so A 0 and E' are independent - have h_indep : IndepFun (A 0) (fun ω ↦ (R' 0 ω).1) P := by - rw [indepFun_iff_condDistrib_eq_const (h.measurable_A 0).aemeasurable - (h.measurable_R 0).fst.aemeasurable, hfst.hasLaw_of_const.map_eq] - exact hfst.condDistrib_eq - -- From independence: condDistrib (A 0) E' P = const (P.map (A 0)) = const alg.p0 - have hcd := condDistrib_of_indepFun h_indep.symm (h.measurable_R 0).fst.aemeasurable - (h.measurable_A 0).aemeasurable - simp only [h.hasLaw_action_zero.map_eq, Algorithm.prod_left] at hcd - exact ⟨(h.measurable_A 0).aemeasurable, (h.measurable_R 0).fst.aemeasurable, hcd⟩ + simpa [h.hasLaw_action_zero.map_eq, Algorithm.prod_left] using hfst.swap_const hasCondDistrib_reward_zero := by have h0 := h.hasCondDistrib_reward_zero simp only [bayesStationaryEnv] at h0 From a3747a9e79c2512427a4e7d353ea7a024db3fbb6 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 24 Feb 2026 15:46:23 +0000 Subject: [PATCH 048/155] Refactor BayesStationaryEnv --- .../BayesStationaryEnv.lean | 66 +++++++------------ 1 file changed, 25 insertions(+), 41 deletions(-) diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 3228c840..b134020e 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -104,7 +104,7 @@ lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : ∀ᵐ e ∂Q, HasLaw (IT.action 0) alg.p0 (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] filter_upwards [condDistrib_comp E - (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable + ((measurable_trajectory h.measurable_A h.measurable_R).aemeasurable) (IT.measurable_action (α := α) (R := R) 0), h.hasCondDistrib_action_zero.condDistrib_eq] with e hc hcd exact ⟨(IT.measurable_action 0).aemeasurable, by @@ -131,38 +131,25 @@ lemma hasCondDistrib_IT_action (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ rwa [Kernel.sectR_prodMkLeft] at he lemma hasCondDistrib_IT_reward [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : - ∀ᵐ e ∂Q, HasCondDistrib (IT.reward (n + 1)) (fun x ↦ (IT.hist n x, IT.action (n + 1) x)) + ∀ᵐ e ∂Q, HasCondDistrib (IT.reward (n + 1)) (fun τ ↦ (IT.hist n τ, IT.action (n + 1) τ)) ((κ.sectR e).prodMkLeft _) (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] - have hmt := measurable_trajectory h.measurable_A h.measurable_R - have hm := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n have h_reorder : HasCondDistrib (R' (n + 1)) - (fun ω ↦ ((IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω), E ω)) - ((κ.prodMkLeft _).comap (fun ((h, a), e) ↦ (h, (e, a))) (by fun_prop)) P := by - convert (h.hasCondDistrib_reward n).comp_right - ((MeasurableEquiv.prodCongr (.refl _) .prodComm).trans MeasurableEquiv.prodAssoc.symm) using 2 - filter_upwards [h_reorder.ae_hasCondDistrib_sectL - ((IT.measurable_hist n).prodMk (IT.measurable_action (n + 1))) - (IT.measurable_reward (n + 1)) - hmt.aemeasurable h.measurable_E.aemeasurable] with e he - have hk : ((κ.prodMkLeft _).comap (fun ((h, a), e) ↦ (h, (e, a))) (by fun_prop)).sectL e = - (κ.sectR e).prodMkLeft (↥(Iic n) → α × R) := - Kernel.ext fun ⟨_, a⟩ ↦ by - simp [Kernel.sectL_apply, Kernel.comap_apply, Kernel.prodMkLeft_apply] - rw [hk] at he; exact he + (fun ω ↦ (E ω, IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) + (κ.comap (fun (e, _, a) ↦ (e, a)) (by fun_prop)) P := + (h.hasCondDistrib_reward n).comp_right (MeasurableEquiv.prodAssoc.symm.trans + ((MeasurableEquiv.prodCongr .prodComm (.refl _)).trans .prodAssoc)) + exact h_reorder.ae_hasCondDistrib_sectR ((IT.measurable_hist n).prodMk + (IT.measurable_action (n + 1))) (IT.measurable_reward (n + 1)) + (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable h.measurable_E.aemeasurable lemma ae_IsAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv (κ.sectR e)) (condDistrib (trajectory A R') E P e) := by filter_upwards [hasLaw_IT_action_zero h, hasCondDistrib_IT_reward_zero h, ae_all_iff.2 (hasCondDistrib_IT_action h), ae_all_iff.2 (hasCondDistrib_IT_reward h)] - with _ h_a0 h_r0 h_a h_r - exact { - hasLaw_action_zero := h_a0 - hasCondDistrib_reward_zero := h_r0 - hasCondDistrib_action := h_a - hasCondDistrib_reward := h_r - } + with _ ha0 hr0 hA hR + exact ⟨IT.measurable_action, IT.measurable_reward, ha0, hr0, hA, hR⟩ end CondDistribIsAlgEnvSeq @@ -194,28 +181,26 @@ lemma IsAlgEnvSeq.isBayesAlgEnvSeq apply HasCondDistrib.hasLaw_of_const simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst hasCondDistrib_action_zero := by - have hfst : HasCondDistrib (fun ω ↦ (R' 0 ω).1) (A 0) (Kernel.const α Q) P := by + have hfst : HasCondDistrib (fun ω ↦ (R' 0 ω).1) (A 0) (Kernel.const _ Q) P := by simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst simpa [h.hasLaw_action_zero.map_eq, Algorithm.prod_left] using hfst.swap_const - hasCondDistrib_reward_zero := by - have h0 := h.hasCondDistrib_reward_zero - simp only [bayesStationaryEnv] at h0 - convert h0.of_compProd.comp_right (MeasurableEquiv.prodComm : α × 𝓔 ≃ᵐ 𝓔 × α) using 2 + hasCondDistrib_reward_zero := + h.hasCondDistrib_reward_zero.of_compProd.comp_right MeasurableEquiv.prodComm hasCondDistrib_action n := by let f : (Iic n → α × 𝓔 × R) → 𝓔 × (Iic n → α × R) := fun h ↦ ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2)) - suffices h' : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) - (((alg.policy n).comap Prod.snd (by fun_prop)).comap f (by fun_prop)) P from - h'.comp_left (f := f) - exact h.hasCondDistrib_action n + have hc : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) + (((alg.policy n).comap Prod.snd (by fun_prop)).comap f (by fun_prop)) P := + h.hasCondDistrib_action n + exact hc.comp_left (f := f) hasCondDistrib_reward n := by let f : (Iic n → α × 𝓔 × R) × α → (Iic n → α × R) × 𝓔 × α := fun p ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), (p.1 ⟨0, by simp⟩).2.1, p.2) - have hf : Measurable f := by fun_prop - suffices h' : HasCondDistrib (fun ω ↦ (R' (n + 1) ω).2) + have hc : HasCondDistrib (fun ω ↦ (R' (n + 1) ω).2) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) - ((Kernel.prodMkLeft (↥(Iic n) → α × R) κ).comap f hf) P from h'.comp_left hf - simpa [bayesStationaryEnv, Kernel.snd_prod] using (h.hasCondDistrib_reward n).snd + ((Kernel.prodMkLeft ((Iic n) → α × R) κ).comap f (by fun_prop)) P := by + simpa [bayesStationaryEnv, Kernel.snd_prod] using (h.hasCondDistrib_reward n).snd + exact hc.comp_left (by fun_prop) end IsAlgEnvSeq @@ -234,13 +219,12 @@ lemma isBayesAlgEnvSeq_bayesTrajMeasure (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × α) R) [IsMarkovKernel κ] (alg : Algorithm α R) : IsBayesAlgEnvSeq Q κ alg (fun ω ↦ (ω 0).2.1) action (fun n ω ↦ (ω n).2.2) - (bayesTrajMeasure Q κ alg) := - (isAlgEnvSeq_trajMeasure _ _).isBayesAlgEnvSeq + (bayesTrajMeasure Q κ alg) := (isAlgEnvSeq_trajMeasure _ _).isBayesAlgEnvSeq noncomputable def bayesTrajMeasurePosterior [StandardBorelSpace 𝓔] [Nonempty 𝓔] - (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × α) ℝ) [IsMarkovKernel κ] - (alg : Algorithm α ℝ) (n : ℕ) : Kernel (Iic n → α × ℝ) 𝓔 := + (Q : Measure 𝓔) (κ : Kernel (𝓔 × α) R) [IsMarkovKernel κ] + (alg : Algorithm α R) (n : ℕ) : Kernel (Iic n → α × R) 𝓔 := condDistrib (fun ω ↦ (ω 0).2.1) (IsAlgEnvSeq.hist action (fun n ω ↦ (ω n).2.2) n) (bayesTrajMeasure Q κ alg) deriving IsMarkovKernel From afc7cde44795ce95841def3eb66276bd593c340d Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 24 Feb 2026 15:51:40 +0000 Subject: [PATCH 049/155] Minor --- LeanBandits/SequentialLearning/BayesStationaryEnv.lean | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index b134020e..0c9cddb0 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -134,12 +134,12 @@ lemma hasCondDistrib_IT_reward [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ al ∀ᵐ e ∂Q, HasCondDistrib (IT.reward (n + 1)) (fun τ ↦ (IT.hist n τ, IT.action (n + 1) τ)) ((κ.sectR e).prodMkLeft _) (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] - have h_reorder : HasCondDistrib (R' (n + 1)) + have hc : HasCondDistrib (R' (n + 1)) (fun ω ↦ (E ω, IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) (κ.comap (fun (e, _, a) ↦ (e, a)) (by fun_prop)) P := (h.hasCondDistrib_reward n).comp_right (MeasurableEquiv.prodAssoc.symm.trans ((MeasurableEquiv.prodCongr .prodComm (.refl _)).trans .prodAssoc)) - exact h_reorder.ae_hasCondDistrib_sectR ((IT.measurable_hist n).prodMk + exact hc.ae_hasCondDistrib_sectR ((IT.measurable_hist n).prodMk (IT.measurable_action (n + 1))) (IT.measurable_reward (n + 1)) (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable h.measurable_E.aemeasurable @@ -223,7 +223,7 @@ lemma isBayesAlgEnvSeq_bayesTrajMeasure noncomputable def bayesTrajMeasurePosterior [StandardBorelSpace 𝓔] [Nonempty 𝓔] - (Q : Measure 𝓔) (κ : Kernel (𝓔 × α) R) [IsMarkovKernel κ] + (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × α) R) [IsMarkovKernel κ] (alg : Algorithm α R) (n : ℕ) : Kernel (Iic n → α × R) 𝓔 := condDistrib (fun ω ↦ (ω 0).2.1) (IsAlgEnvSeq.hist action (fun n ω ↦ (ω n).2.2) n) (bayesTrajMeasure Q κ alg) From f7a3e81fb6ff9ee45168e8c6e97e61a78b376122 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 24 Feb 2026 15:52:36 +0000 Subject: [PATCH 050/155] Minor --- LeanBandits/SequentialLearning/BayesStationaryEnv.lean | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 0c9cddb0..60fccbfd 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -181,9 +181,9 @@ lemma IsAlgEnvSeq.isBayesAlgEnvSeq apply HasCondDistrib.hasLaw_of_const simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst hasCondDistrib_action_zero := by - have hfst : HasCondDistrib (fun ω ↦ (R' 0 ω).1) (A 0) (Kernel.const _ Q) P := by + have hc : HasCondDistrib (fun ω ↦ (R' 0 ω).1) (A 0) (Kernel.const _ Q) P := by simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst - simpa [h.hasLaw_action_zero.map_eq, Algorithm.prod_left] using hfst.swap_const + simpa [h.hasLaw_action_zero.map_eq, Algorithm.prod_left] using hc.swap_const hasCondDistrib_reward_zero := h.hasCondDistrib_reward_zero.of_compProd.comp_right MeasurableEquiv.prodComm hasCondDistrib_action n := by From 3bf70f57c885a49a5f94c1912f11b0ae6c75d1dc Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 24 Feb 2026 17:48:45 +0000 Subject: [PATCH 051/155] Refactor HistoryDensity (in progress) --- LeanBandits.lean | 2 + LeanBandits/BanditAlgorithms/TS.lean | 48 +-- LeanBandits/ForMathlib/FullSupport.lean | 45 +++ LeanBandits/ForMathlib/WithDensity.lean | 95 ++++++ .../SequentialLearning/HistoryDensity.lean | 323 ++++++------------ 5 files changed, 267 insertions(+), 246 deletions(-) create mode 100644 LeanBandits/ForMathlib/FullSupport.lean create mode 100644 LeanBandits/ForMathlib/WithDensity.lean diff --git a/LeanBandits.lean b/LeanBandits.lean index ed67f3d0..0fbf92ee 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -10,6 +10,7 @@ import LeanBandits.BanditAlgorithms.UCB import LeanBandits.BanditAlgorithms.Uniform import LeanBandits.ForMathlib.CondDistrib import LeanBandits.ForMathlib.CondIndepFun +import LeanBandits.ForMathlib.FullSupport import LeanBandits.ForMathlib.HasCondDistrib import LeanBandits.ForMathlib.IndepFun import LeanBandits.ForMathlib.IndepInfinitePi @@ -21,6 +22,7 @@ import LeanBandits.ForMathlib.MeasurableArgMax import LeanBandits.ForMathlib.StandardBorel import LeanBandits.ForMathlib.SubGaussian import LeanBandits.ForMathlib.Traj +import LeanBandits.ForMathlib.WithDensity import LeanBandits.SequentialLearning.Algorithm import LeanBandits.SequentialLearning.BayesStationaryEnv import LeanBandits.SequentialLearning.Deterministic diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index daceceef..f0b9e193 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -217,9 +217,7 @@ lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] condDistrib (IsBayesAlgEnvSeq.bestAction κ E') (IsAlgEnvSeq.hist A R' t) P := by have h_ba_comp : IsBayesAlgEnvSeq.bestAction κ E' - = IsBayesAlgEnvSeq.bestAction κ id ∘ E' := by - rw [bestAction_eq_envToBestArm_comp_env κ (E' := E'), - bestAction_eq_envToBestArm_comp_env κ (E' := id), Function.comp_id] + = IsBayesAlgEnvSeq.bestAction κ id ∘ E' := rfl rw [h_ba_comp] have hm := IsBayesAlgEnvSeq.measurable_bestAction (κ := κ) measurable_id have h_comp := condDistrib_comp (mβ := MeasurableSpace.pi) (μ := P) @@ -228,7 +226,7 @@ lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (IsBayesAlgEnvSeq.bestAction κ id) =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) t).map (IsBayesAlgEnvSeq.bestAction κ id) := by - filter_upwards [posterior_eq_uniform Q κ h hK t] with x hx + filter_upwards [posterior_eq_ref Q κ h (Bandits.uniformAlgorithm_hasFullSupport hK) t] with x hx simp only [Kernel.map_apply _ hm, hx] exact (h.hasCondDistrib_action' t).condDistrib_eq.trans (h_comp.trans h_map).symm @@ -940,19 +938,20 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] by_cases hn : n = 0 · simp [hn] have hn' : 0 < n := Nat.pos_of_ne_zero hn - rw [show IsBayesAlgEnvSeq.bestAction κ E' = envToBestArm κ ∘ E' from - bestAction_eq_envToBestArm_comp_env κ] + rw [show IsBayesAlgEnvSeq.bestAction κ E' = IsBayesAlgEnvSeq.bestAction κ id ∘ E' from + rfl] let badSetIT := fun (a : Fin K) (s : ℕ) (e : E) ↦ {ω : ℕ → (Fin K) × ℝ | pullCount IT.action a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} - have h_set_eq : {ω | ∃ s < n, pullCount A ((envToBestArm κ ∘ E') ω) s ω ≠ 0 ∧ + have h_set_eq : {ω | ∃ s < n, pullCount A ((IsBayesAlgEnvSeq.bestAction κ id ∘ E') ω) s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / - (pullCount A ((envToBestArm κ ∘ E') ω) s ω : ℝ)) ≤ - |empMean A R' ((envToBestArm κ ∘ E') ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E' ((envToBestArm κ ∘ E') ω) ω|} = + (pullCount A ((IsBayesAlgEnvSeq.bestAction κ id ∘ E') ω) s ω : ℝ)) ≤ + |empMean A R' ((IsBayesAlgEnvSeq.bestAction κ id ∘ E') ω) s ω - + IsBayesAlgEnvSeq.actionMean κ E' ((IsBayesAlgEnvSeq.bestAction κ id ∘ E') ω) ω|} = (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' - {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1} := by + {p | p.2 ∈ ⋃ s ∈ Finset.range n, + badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} := by ext ω simp only [Set.mem_setOf_eq, Finset.mem_range, Set.mem_preimage, Set.mem_iUnion, badSetIT, IsBayesAlgEnvSeq.actionMean, Function.comp_apply, exists_prop] @@ -979,10 +978,10 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] (Q := Q) (κ := κ) (P := P) hσ2 hs hn' hδ hδ1 e h_isAlgEnvSeq a have h_cond_best : ∀ᵐ e ∂(P.map E'), (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) - (⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ e) s e) ≤ + (⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id e) s e) ≤ ENNReal.ofReal (2 * n * δ) := by filter_upwards [h_cond_bound] with e he - exact he (envToBestArm κ e) + exact he (IsBayesAlgEnvSeq.bestAction κ id e) have h_kernel : ∀ a, Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := fun a ↦ stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_fst.prodMk measurable_const) @@ -1001,31 +1000,36 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp measurable_snd).sub (h_kernel a)).abs) have h_meas_set : MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | - p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1} := by + p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} := by have h_eq : {p : E × (ℕ → (Fin K) × ℝ) | - p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1} = - ⋃ a : Fin K, ((envToBestArm κ ∘ Prod.fst) ⁻¹' {a} ∩ + p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} = + ⋃ a : Fin K, ((IsBayesAlgEnvSeq.bestAction κ id ∘ Prod.fst) ⁻¹' {a} ∩ ⋃ s ∈ Finset.range n, {p | p.2 ∈ badSetIT a s p.1}) := by ext p; simp only [Set.mem_setOf_eq, Set.mem_iUnion, Set.mem_inter_iff, Set.mem_preimage, Function.comp_apply, Set.mem_singleton_iff, Finset.mem_range] constructor - · intro ⟨s, hs, hm⟩; exact ⟨envToBestArm κ p.1, rfl, s, hs, hm⟩ + · intro ⟨s, hs, hm⟩; exact ⟨IsBayesAlgEnvSeq.bestAction κ id p.1, rfl, s, hs, hm⟩ · rintro ⟨a, ha, s, hs, hm⟩; exact ⟨s, hs, ha ▸ hm⟩ rw [h_eq] exact .iUnion fun a ↦ .inter - ((measurable_envToBestArm (κ := κ) |>.comp measurable_fst) (measurableSet_singleton a)) + ((IsBayesAlgEnvSeq.measurable_bestAction (κ := κ) measurable_id |>.comp + measurable_fst) (measurableSet_singleton a)) (.biUnion (Finset.range n).countable_toSet fun s _ ↦ h_meas_badSetIT a s) calc P ((fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' - {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1}) + {p | p.2 ∈ ⋃ s ∈ Finset.range n, + badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1}) = (P.map (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω))) - {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1} := by + {p | p.2 ∈ ⋃ s ∈ Finset.range n, + badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} := by rw [Measure.map_apply h_meas_pair h_meas_set] _ = (P.map E' ⊗ₘ condDistrib ((fun ω n => (A n ω, R' n ω))) E' P) - {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ p.1) s p.1} := by + {p | p.2 ∈ ⋃ s ∈ Finset.range n, + badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} := by rw [h_disint] _ = ∫⁻ e, (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) - (⋃ s ∈ Finset.range n, badSetIT (envToBestArm κ e) s e) ∂(P.map E') := by + (⋃ s ∈ Finset.range n, + badSetIT (IsBayesAlgEnvSeq.bestAction κ id e) s e) ∂(P.map E') := by rw [Measure.compProd_apply h_meas_set]; rfl _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * n * δ) ∂(P.map E') := by apply lintegral_mono_ae h_cond_best diff --git a/LeanBandits/ForMathlib/FullSupport.lean b/LeanBandits/ForMathlib/FullSupport.lean new file mode 100644 index 00000000..bbcaae01 --- /dev/null +++ b/LeanBandits/ForMathlib/FullSupport.lean @@ -0,0 +1,45 @@ +/- +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 Mathlib.Probability.Kernel.RadonNikodym + +/-! +# Absolute continuity and rnDeriv finiteness from full support + +When a reference measure gives positive mass to every singleton, any measure is absolutely +continuous with respect to it, and the Radon-Nikodym derivative is pointwise finite. +-/ + +open MeasureTheory ProbabilityTheory + +variable {α β : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {μ ν : Measure α} + +/-- Any measure is absolutely continuous wrt any measure giving positive mass to all singletons. -/ +lemma absolutelyContinuous_of_forall_singleton_pos (hν : ∀ a : α, ν {a} > 0) : μ ≪ ν := by + intro s hs + rcases s.eq_empty_or_nonempty with rfl | ⟨a, ha⟩ + · exact measure_empty + · exact absurd (measure_mono_null (Set.singleton_subset_iff.mpr ha) hs) (hν a).ne' + +/-- An ae property holds everywhere when the reference measure gives positive mass + to every singleton. -/ +lemma forall_of_ae_of_forall_singleton_pos (hν : ∀ a, ν {a} > 0) {p : α → Prop} + (hp : ∀ᵐ a ∂ν, p a) (a : α) : p a := by + by_contra h + exact absurd (measure_mono_null (Set.singleton_subset_iff.mpr h) (ae_iff.mp hp)) (hν a).ne' + +/-- `rnDeriv` is pointwise finite when the reference measure has full support on singletons. -/ +lemma rnDeriv_ne_top_of_forall_singleton_pos [SigmaFinite μ] + (hν : ∀ a, ν {a} > 0) (a : α) : μ.rnDeriv ν a ≠ ⊤ := + (forall_of_ae_of_forall_singleton_pos hν (Measure.rnDeriv_lt_top μ ν) a).ne + +/-- Kernel `rnDeriv` is pointwise finite when the reference kernel has full support + on singletons. -/ +lemma kernel_rnDeriv_ne_top_of_forall_singleton_pos + [MeasurableSpace.CountableOrCountablyGenerated α β] + {κ η : Kernel α β} [IsFiniteKernel κ] [IsFiniteKernel η] + (hη : ∀ a b, η a {b} > 0) (a : α) (b : β) : + Kernel.rnDeriv κ η a b ≠ ⊤ := + (forall_of_ae_of_forall_singleton_pos (hη a) (Kernel.rnDeriv_lt_top κ η) b).ne diff --git a/LeanBandits/ForMathlib/WithDensity.lean b/LeanBandits/ForMathlib/WithDensity.lean new file mode 100644 index 00000000..6c209038 --- /dev/null +++ b/LeanBandits/ForMathlib/WithDensity.lean @@ -0,0 +1,95 @@ +/- +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 Mathlib.Probability.Kernel.CompProdEqIff +import Mathlib.Probability.Kernel.Composition.MeasureComp +/-! +# Interactions of `withDensity` with `compProd`, `map`, and `swap` + +Lemmas for pushing `Measure.withDensity` and `Kernel.withDensity` through +`compProd`, `MeasurableEquiv.map`, `Prod.swap`, and composition. +-/ + +open MeasureTheory ProbabilityTheory + +open scoped ENNReal + +variable {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {μ : Measure α} + +/-- Composing `withDensity` on the measure side of a `compProd`: +`(μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst)`. -/ +lemma withDensity_compProd_left [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] + {f : α → ℝ≥0∞} (hf : Measurable f) : + (μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst) := by + ext s hs + rw [Measure.compProd_apply hs, withDensity_apply _ hs, + lintegral_withDensity_eq_lintegral_mul₀ hf.aemeasurable + (Kernel.measurable_kernel_prodMk_left hs).aemeasurable, + ← lintegral_indicator hs, + Measure.lintegral_compProd ((hf.comp measurable_fst).indicator hs)] + congr 1 + ext a + simp_rw [Pi.mul_apply] + have : (fun b ↦ s.indicator (f ∘ Prod.fst) (a, b)) = + fun b ↦ (Prod.mk a ⁻¹' s).indicator (fun _ ↦ f a) b := by + ext b; simp only [Function.comp, Set.indicator, Set.mem_preimage]; rfl + rw [this, lintegral_indicator_const (hs.preimage (by fun_prop))] + +/-- Mapping a `withDensity` through `MeasurableEquiv.symm`: +`(μ.withDensity f).map e.symm = (μ.map e.symm).withDensity (f ∘ e)`. -/ +lemma withDensity_map_equiv_symm + {μ : Measure β} {e : α ≃ᵐ β} {f : β → ℝ≥0∞} (hf : Measurable f) : + (μ.withDensity f).map e.symm = (μ.map e.symm).withDensity (f ∘ e) := by + ext s hs + rw [Measure.map_apply e.symm.measurable hs, + withDensity_apply _ (e.symm.measurable hs), + withDensity_apply _ hs, Measure.restrict_map e.symm.measurable hs, + lintegral_map (hf.comp e.measurable) e.symm.measurable] + simp_rw [Function.comp_apply, e.apply_symm_apply] + +/-- Mapping a `withDensity` through a `MeasurableEquiv` from the snd component. -/ +lemma map_swap_withDensity_fst + {μ : Measure (α × β)} [SFinite μ] + {f : β → ℝ≥0∞} (hf : Measurable f) : + (μ.withDensity (f ∘ Prod.snd)).map Prod.swap + = (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := by + ext s hs + rw [Measure.map_apply measurable_swap hs, withDensity_apply _ (measurable_swap hs), + withDensity_apply _ hs, Measure.restrict_map measurable_swap hs] + exact (lintegral_map (hf.comp measurable_fst) measurable_swap).symm + +/-- `(μ.withDensity (f ∘ g)).map g = (μ.map g).withDensity f`. -/ +lemma map_withDensity_comp + {g : α → γ} {f : γ → ℝ≥0∞} + (hg : Measurable g) (hf : Measurable f) : + (μ.withDensity (f ∘ g)).map g = (μ.map g).withDensity f := by + ext s hs + simp only [Measure.map_apply hg hs, withDensity_apply _ (hg hs), withDensity_apply _ hs, + setLIntegral_map hs hf hg, Function.comp] + +/-- `(κ.withDensity (fun _ => f)) ∘ₘ μ = (κ ∘ₘ μ).withDensity f`. -/ +lemma comp_withDensity_const + [SFinite μ] + {κ : Kernel α γ} [IsSFiniteKernel κ] + {f : γ → ℝ≥0∞} (hf : Measurable f) + [IsSFiniteKernel (κ.withDensity (fun _ => f))] : + (κ.withDensity (fun _ => f)) ∘ₘ μ = (κ ∘ₘ μ).withDensity f := by + rw [← Measure.snd_compProd μ (κ.withDensity (fun _ => f)), + Measure.compProd_withDensity (show Measurable (Function.uncurry (fun (_ : α) => f)) from + hf.comp measurable_snd), + ← Measure.snd_compProd μ κ, Measure.snd, Measure.snd] + exact map_withDensity_comp measurable_snd hf + +/-- `(μ.withDensity f) ⊗ₘ (κ.withDensity g) = (μ ⊗ₘ κ).withDensity (f ∘ fst * uncurry g)`. -/ +lemma withDensity_compProd_withDensity [SFinite μ] + {κ : Kernel α γ} [IsSFiniteKernel κ] + {f : α → ℝ≥0∞} {g : α → γ → ℝ≥0∞} + (hf : Measurable f) (hg : Measurable (Function.uncurry g)) + [IsSFiniteKernel (κ.withDensity g)] : + (μ.withDensity f) ⊗ₘ (κ.withDensity g) = + (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst * Function.uncurry g) := by + rw [Measure.compProd_withDensity hg, withDensity_compProd_left hf] + exact (withDensity_mul _ (hf.comp measurable_fst) hg).symm diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 9dfdc41d..0bf2065c 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -6,7 +6,8 @@ Authors: Rémy Degenne, Paulo Rauber import LeanBandits.SequentialLearning.StationaryEnv import LeanBandits.SequentialLearning.BayesStationaryEnv import LeanBandits.BanditAlgorithms.Uniform -import Mathlib.Probability.Kernel.RadonNikodym +import LeanBandits.ForMathlib.FullSupport +import LeanBandits.ForMathlib.WithDensity /-! # Algorithm-Independence of Bayesian Posteriors @@ -14,8 +15,8 @@ import Mathlib.Probability.Kernel.RadonNikodym The key result: the posterior distribution on the environment (and therefore on the best arm) given the observed history is independent of the algorithm used to generate the data. -The proof routes through a uniform algorithm as reference measure. The history distribution -under any algorithm is absolutely continuous w.r.t. the uniform algorithm's, with a density +The proof routes through a reference algorithm with full support. The history distribution +under any algorithm is absolutely continuous w.r.t. the reference algorithm's, with a density that depends only on action probabilities (not the environment). This density factorization implies the posteriors agree. -/ @@ -26,129 +27,17 @@ open scoped ENNReal NNReal namespace Learning -variable {α β : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {μ ν : Measure α} {K : ℕ} - -section UniformFullSupport - -variable (hK : 0 < K) - -/-- Any measure is absolutely continuous wrt any measure giving positive mass to all singletons. -/ -lemma absolutelyContinuous_of_forall_singleton_pos (hν : ∀ a : α, ν {a} > 0) : μ ≪ ν := by - intro s hs - have h_empty : s = ∅ := by - by_contra h - obtain ⟨a, ha⟩ := Set.nonempty_iff_ne_empty.mpr h - have h1 : ν {a} ≤ ν s := measure_mono (Set.singleton_subset_iff.mpr ha) - exact absurd (le_antisymm (hs ▸ h1) (zero_le _)) (ne_of_gt (hν a)) - rw [h_empty, measure_empty] - -/-- `rnDeriv` is pointwise finite when the reference measure has full support on singletons. -/ -lemma rnDeriv_ne_top_of_forall_singleton_pos [SigmaFinite μ] - (hν : ∀ a, ν {a} > 0) (a : α) : μ.rnDeriv ν a ≠ ⊤ := by - intro h_eq - have h_mem : a ∈ {x | ¬ (μ.rnDeriv ν x < ⊤)} := by simp [h_eq] - have h_null : ν {x | ¬ (μ.rnDeriv ν x < ⊤)} = 0 := - ae_iff.mp (Measure.rnDeriv_lt_top μ ν) - exact absurd (le_antisymm ((measure_mono (Set.singleton_subset_iff.mpr h_mem)).trans - (le_of_eq h_null)) (zero_le _)) (ne_of_gt (hν a)) - -/-- Kernel `rnDeriv` is pointwise finite when the reference kernel has full support - on singletons. -/ -lemma kernel_rnDeriv_ne_top_of_forall_singleton_pos - [MeasurableSpace.CountableOrCountablyGenerated α β] - {κ η : Kernel α β} [IsFiniteKernel κ] [IsFiniteKernel η] - (hη : ∀ a b, η a {b} > 0) (a : α) (b : β) : - Kernel.rnDeriv κ η a b ≠ ⊤ := by - intro h_eq - have h_mem : b ∈ {x | ¬ (Kernel.rnDeriv κ η a x < ⊤)} := by simp [h_eq] - have h_null : η a {x | ¬ (Kernel.rnDeriv κ η a x < ⊤)} = 0 := - ae_iff.mp (Kernel.rnDeriv_lt_top κ η) - exact absurd (le_antisymm ((measure_mono (Set.singleton_subset_iff.mpr h_mem)).trans - (le_of_eq h_null)) (zero_le _)) (ne_of_gt (hη a b)) - -end UniformFullSupport - -section WithDensityHelpers - -variable {γ : Type*} {mγ : MeasurableSpace γ} - -/-- Composing `withDensity` on the measure side of a `compProd`: -`(μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst)`. -/ -private lemma withDensity_compProd_left [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] - {f : α → ℝ≥0∞} (hf : Measurable f) : - (μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst) := by - ext s hs - rw [Measure.compProd_apply hs, withDensity_apply _ hs, - lintegral_withDensity_eq_lintegral_mul₀ hf.aemeasurable - (Kernel.measurable_kernel_prodMk_left hs).aemeasurable, - ← lintegral_indicator hs, - Measure.lintegral_compProd ((hf.comp measurable_fst).indicator hs)] - congr 1 - ext a - simp_rw [Pi.mul_apply] - have : (fun b ↦ s.indicator (f ∘ Prod.fst) (a, b)) = - fun b ↦ (Prod.mk a ⁻¹' s).indicator (fun _ ↦ f a) b := by - ext b; simp only [Function.comp, Set.indicator, Set.mem_preimage]; rfl - rw [this, lintegral_indicator_const (hs.preimage (by fun_prop))] - -/-- Mapping a `withDensity` through `MeasurableEquiv.symm`: -`(μ.withDensity f).map e.symm = (μ.map e.symm).withDensity (f ∘ e)`. -/ -private lemma withDensity_map_equiv_symm - {μ : Measure β} {e : α ≃ᵐ β} {f : β → ℝ≥0∞} (hf : Measurable f) : - (μ.withDensity f).map e.symm = (μ.map e.symm).withDensity (f ∘ e) := by - ext s hs - rw [Measure.map_apply e.symm.measurable hs, - withDensity_apply _ (e.symm.measurable hs), - withDensity_apply _ hs, Measure.restrict_map e.symm.measurable hs, - lintegral_map (hf.comp e.measurable) e.symm.measurable] - simp_rw [Function.comp_apply, e.apply_symm_apply] - -/-- Mapping a `withDensity` through a `MeasurableEquiv` from the snd component. -/ -private lemma map_swap_withDensity_fst - {μ : Measure (α × β)} [SFinite μ] - {f : β → ℝ≥0∞} (hf : Measurable f) : - (μ.withDensity (f ∘ Prod.snd)).map Prod.swap - = (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := by - ext s hs - rw [Measure.map_apply measurable_swap hs, withDensity_apply _ (measurable_swap hs), - withDensity_apply _ hs, Measure.restrict_map measurable_swap hs] - exact (lintegral_map (hf.comp measurable_fst) measurable_swap).symm - -/-- `(μ.withDensity (h ∘ g)).map g = (μ.map g).withDensity h`. -/ -private lemma withDensity_map_eq' - {μ : Measure α} {g : α → γ} {h : γ → ℝ≥0∞} - (hg : Measurable g) (hh : Measurable h) : - (μ.withDensity (h ∘ g)).map g = (μ.map g).withDensity h := by - ext s hs - rw [Measure.map_apply hg hs, withDensity_apply _ (hg hs), withDensity_apply _ hs] - conv_rhs => rw [Measure.restrict_map hg hs] - rw [lintegral_map hh hg]; rfl - -/-- `(κ.withDensity (fun _ => ρ)) ∘ₘ Q = (κ ∘ₘ Q).withDensity ρ`. -/ -private lemma comp_withDensity_const - {Q : Measure α} [SFinite Q] - {κ : Kernel α γ} [IsSFiniteKernel κ] - {ρ : γ → ℝ≥0∞} (hρ : Measurable ρ) - [IsSFiniteKernel (κ.withDensity (fun _ => ρ))] : - (κ.withDensity (fun _ => ρ)) ∘ₘ Q = (κ ∘ₘ Q).withDensity ρ := by - rw [← Measure.snd_compProd Q (κ.withDensity (fun _ => ρ)), - Measure.compProd_withDensity (show Measurable (Function.uncurry (fun (_ : α) => ρ)) from - hρ.comp measurable_snd), - ← Measure.snd_compProd Q κ, Measure.snd, Measure.snd] - exact withDensity_map_eq' measurable_snd hρ - -/-- `(μ.withDensity f) ⊗ₘ (κ.withDensity g) = (μ ⊗ₘ κ).withDensity (f ∘ fst * uncurry g)`. -/ -private lemma withDensity_compProd_withDensity [SFinite μ] - {κ : Kernel α γ} [IsSFiniteKernel κ] - {f : α → ℝ≥0∞} {g : α → γ → ℝ≥0∞} - (hf : Measurable f) (hg : Measurable (Function.uncurry g)) - [IsSFiniteKernel (κ.withDensity g)] : - (μ.withDensity f) ⊗ₘ (κ.withDensity g) = - (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst * Function.uncurry g) := by - rw [Measure.compProd_withDensity hg, withDensity_compProd_left hf] - exact (withDensity_mul _ (hf.comp measurable_fst) hg).symm - -end WithDensityHelpers +variable {K : ℕ} + +/-- An algorithm has full support if every action has positive probability under both +the initial measure and every policy. -/ +def Algorithm.HasFullSupport (ref : Algorithm (Fin K) ℝ) : Prop := + (∀ a, ref.p0 {a} > 0) ∧ + (∀ n (h : Iic n → Fin K × ℝ) (a : Fin K), ref.policy n h {a} > 0) + +lemma Bandits.uniformAlgorithm_hasFullSupport (hK : 0 < K) : + (Bandits.uniformAlgorithm hK).HasFullSupport := + ⟨Bandits.uniformAlgorithm_p0_pos, fun _ h a => Bandits.uniformAlgorithm_policy_pos h a⟩ section AbsolutelyContinuousHist @@ -158,20 +47,20 @@ omit [Nonempty (Fin K)] in /-- The step kernel for a stationary environment decomposes as a product of the policy measure and the reward kernel. -/ private lemma absolutelyContinuous_stepKernel_stationary - (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) (ν : Kernel (Fin K) ℝ) - [IsMarkovKernel ν] (n : ℕ) (h : Iic n → Fin K × ℝ) : + (alg ref : Algorithm (Fin K) ℝ) (href : ref.HasFullSupport) + (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] (n : ℕ) (h : Iic n → Fin K × ℝ) : stepKernel alg (stationaryEnv ν) n h ≪ - stepKernel (Bandits.uniformAlgorithm hK) (stationaryEnv ν) n h := by + stepKernel ref (stationaryEnv ν) n h := by have h1 : stepKernel alg (stationaryEnv ν) n h = (alg.policy n h) ⊗ₘ ν := by simp only [stepKernel, stationaryEnv]; ext s hs simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - have h2 : stepKernel (Bandits.uniformAlgorithm hK) (stationaryEnv ν) n h = - ((Bandits.uniformAlgorithm hK).policy n h) ⊗ₘ ν := by + have h2 : stepKernel ref (stationaryEnv ν) n h = + (ref.policy n h) ⊗ₘ ν := by simp only [stepKernel, stationaryEnv]; ext s hs simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] rw [h1, h2] exact Measure.AbsolutelyContinuous.compProd_left - (absolutelyContinuous_of_forall_singleton_pos (Bandits.uniformAlgorithm_policy_pos h)) _ + (absolutelyContinuous_of_forall_singleton_pos (href.2 n h)) _ /-- The history distribution at time `n + 1` decomposes as a compProd of the history at time `n` and the step kernel, composed with `IicSuccProd.symm`. -/ @@ -205,9 +94,10 @@ private lemma map_hist_succ_eq_compProd_map (stepKernel alg env n)).mp h_cd.condDistrib_eq) /-- The history distribution under any algorithm is absolutely continuous w.r.t. the - history distribution under the uniform algorithm, for a stationary environment. -/ + history distribution under a reference algorithm with full support, + for a stationary environment. -/ private lemma absolutelyContinuous_map_hist_stationary - (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) + (alg ref : Algorithm (Fin K) ℝ) (href : ref.HasFullSupport) (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] {Ω₁ : Type*} [MeasurableSpace Ω₁] {A₁ : ℕ → Ω₁ → Fin K} {R₁ : ℕ → Ω₁ → ℝ} @@ -216,7 +106,7 @@ private lemma absolutelyContinuous_map_hist_stationary {Ω₂ : Type*} [MeasurableSpace Ω₂] {A₂ : ℕ → Ω₂ → Fin K} {R₂ : ℕ → Ω₂ → ℝ} {P₂ : Measure Ω₂} [IsProbabilityMeasure P₂] - (h₂ : IsAlgEnvSeq A₂ R₂ (Bandits.uniformAlgorithm hK) (stationaryEnv ν) P₂) + (h₂ : IsAlgEnvSeq A₂ R₂ ref (stationaryEnv ν) P₂) (t : ℕ) : P₁.map (IsAlgEnvSeq.hist A₁ R₁ t) ≪ P₂.map (IsAlgEnvSeq.hist A₂ R₂ t) := by induction t with @@ -234,13 +124,13 @@ private lemma absolutelyContinuous_map_hist_stationary h₁.hasLaw_step_zero.map_eq, h₂.hasLaw_step_zero.map_eq] simp only [stationaryEnv_ν0] exact (Measure.AbsolutelyContinuous.compProd_left - (absolutelyContinuous_of_forall_singleton_pos Bandits.uniformAlgorithm_p0_pos) _).map + (absolutelyContinuous_of_forall_singleton_pos href.1) _).map e.symm.measurable | succ n ih => rw [map_hist_succ_eq_compProd_map h₁, map_hist_succ_eq_compProd_map h₂] exact (Measure.AbsolutelyContinuous.compProd ih (Filter.Eventually.of_forall fun h => - absolutelyContinuous_stepKernel_stationary hK alg ν n h)).map + absolutelyContinuous_stepKernel_stationary alg ref href ν n h)).map (MeasurableEquiv.IicSuccProd _ n).symm.measurable end AbsolutelyContinuousHist @@ -249,23 +139,23 @@ section DensityIndependence variable {K : ℕ} [Nonempty (Fin K)] -/-- The density of the history distribution under `alg` w.r.t. the uniform algorithm. +/-- The density of the history distribution under `alg` w.r.t. a reference algorithm. This density depends only on the algorithm's action probabilities, not on the reward kernel. -/ private noncomputable def historyDensity - (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) : + (alg ref : Algorithm (Fin K) ℝ) : (t : ℕ) → (Iic t → Fin K × ℝ) → ℝ≥0∞ - | 0 => (alg.p0.rnDeriv (Bandits.uniformAlgorithm hK).p0 ∘ Prod.fst) ∘ + | 0 => (alg.p0.rnDeriv ref.p0 ∘ Prod.fst) ∘ MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) | n + 1 => let σ : (Iic n → Fin K × ℝ) → (Fin K × ℝ) → ℝ≥0∞ := fun h ar => Kernel.rnDeriv (alg.policy n) - ((Bandits.uniformAlgorithm hK).policy n) h ar.1 - (historyDensity hK alg n ∘ Prod.fst * Function.uncurry σ) ∘ + (ref.policy n) h ar.1 + (historyDensity alg ref n ∘ Prod.fst * Function.uncurry σ) ∘ MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n omit [Nonempty (Fin K)] in -private lemma measurable_historyDensity (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) (t : ℕ) : - Measurable (historyDensity hK alg t) := by +private lemma measurable_historyDensity (alg ref : Algorithm (Fin K) ℝ) (t : ℕ) : + Measurable (historyDensity alg ref t) := by induction t with | zero => exact (Measure.measurable_rnDeriv _ _).comp @@ -277,19 +167,20 @@ private lemma measurable_historyDensity (hK : 0 < K) (alg : Algorithm (Fin K) (MeasurableEquiv.IicSuccProd _ n).measurable omit [Nonempty (Fin K)] in -private lemma historyDensity_ne_top (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) (t : ℕ) - (h : Iic t → Fin K × ℝ) : historyDensity hK alg t h ≠ ⊤ := by +private lemma historyDensity_ne_top (alg ref : Algorithm (Fin K) ℝ) + (href : ref.HasFullSupport) (t : ℕ) + (h : Iic t → Fin K × ℝ) : historyDensity alg ref t h ≠ ⊤ := by induction t with - | zero => exact rnDeriv_ne_top_of_forall_singleton_pos Bandits.uniformAlgorithm_p0_pos _ + | zero => exact rnDeriv_ne_top_of_forall_singleton_pos href.1 _ | succ n ih => exact ENNReal.mul_ne_top (ih _) (kernel_rnDeriv_ne_top_of_forall_singleton_pos - (fun h' a => Bandits.uniformAlgorithm_policy_pos h' a) _ _) + (fun h' a => href.2 n h' a) _ _) -/-- The history distribution under any algorithm equals the uniform algorithm's history +/-- The history distribution under any algorithm equals the reference algorithm's history distribution weighted by `historyDensity`, for any stationary environment. -/ private lemma map_hist_eq_withDensity_historyDensity - (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) (t : ℕ) + (alg ref : Algorithm (Fin K) ℝ) (href : ref.HasFullSupport) (t : ℕ) (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] {Ω₁ : Type*} [MeasurableSpace Ω₁] {A₁ : ℕ → Ω₁ → Fin K} {R₁ : ℕ → Ω₁ → ℝ} @@ -298,15 +189,14 @@ private lemma map_hist_eq_withDensity_historyDensity {Ω₂ : Type*} [MeasurableSpace Ω₂] {A₂ : ℕ → Ω₂ → Fin K} {R₂ : ℕ → Ω₂ → ℝ} {P₂ : Measure Ω₂} [IsProbabilityMeasure P₂] - (h₂ : IsAlgEnvSeq A₂ R₂ (Bandits.uniformAlgorithm hK) (stationaryEnv ν) P₂) : + (h₂ : IsAlgEnvSeq A₂ R₂ ref (stationaryEnv ν) P₂) : P₁.map (IsAlgEnvSeq.hist A₁ R₁ t) = - (P₂.map (IsAlgEnvSeq.hist A₂ R₂ t)).withDensity (historyDensity hK alg t) := by - set unif := Bandits.uniformAlgorithm hK + (P₂.map (IsAlgEnvSeq.hist A₂ R₂ t)).withDensity (historyDensity alg ref t) := by induction t with | zero => set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) - have h_ac : alg.p0 ≪ unif.p0 := - absolutelyContinuous_of_forall_singleton_pos Bandits.uniformAlgorithm_p0_pos + have h_ac : alg.p0 ≪ ref.p0 := + absolutelyContinuous_of_forall_singleton_pos href.1 have h_hist₁ : IsAlgEnvSeq.hist A₁ R₁ 0 = e.symm ∘ IsAlgEnvSeq.step A₁ R₁ 0 := by funext ω ⟨i, hi⟩ have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl @@ -326,41 +216,41 @@ private lemma map_hist_eq_withDensity_historyDensity ((Measure.measurable_rnDeriv _ _).comp measurable_fst) | succ n ih => let σ : (Iic n → Fin K × ℝ) → (Fin K × ℝ) → ℝ≥0∞ := - fun h ar => Kernel.rnDeriv (alg.policy n) (unif.policy n) h ar.1 + fun h ar => Kernel.rnDeriv (alg.policy n) (ref.policy n) h ar.1 have hσ_meas : Measurable (Function.uncurry σ) := (Kernel.measurable_rnDeriv _ _).comp (measurable_fst.prodMk (measurable_fst.comp measurable_snd)) have h_step : stepKernel alg (stationaryEnv ν) n = - (stepKernel unif (stationaryEnv ν) n).withDensity σ := by + (stepKernel ref (stationaryEnv ν) n).withDensity σ := by ext h : 1 rw [Kernel.withDensity_apply _ hσ_meas] have h_alg : stepKernel alg (stationaryEnv ν) n h = (alg.policy n h) ⊗ₘ ν := by ext s hs simp only [stepKernel, stationaryEnv, Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - have h_unif : stepKernel unif (stationaryEnv ν) n h = (unif.policy n h) ⊗ₘ ν := by + have h_ref : stepKernel ref (stationaryEnv ν) n h = (ref.policy n h) ⊗ₘ ν := by ext s hs simp only [stepKernel, stationaryEnv, Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - have h_wd : ((unif.policy n) h).withDensity - (Kernel.rnDeriv (alg.policy n) (unif.policy n) h) = alg.policy n h := by + have h_wd : ((ref.policy n) h).withDensity + (Kernel.rnDeriv (alg.policy n) (ref.policy n) h) = alg.policy n h := by rw [← Kernel.withDensity_apply _ (Kernel.measurable_rnDeriv _ _)] - exact Kernel.withDensity_rnDeriv_eq (κ := alg.policy n) (η := unif.policy n) - (absolutelyContinuous_of_forall_singleton_pos (Bandits.uniformAlgorithm_policy_pos h)) - rw [h_alg, h_unif, ← h_wd] - haveI : SFinite ((unif.policy n h).withDensity - (Kernel.rnDeriv (alg.policy n) (unif.policy n) h)) := by + exact Kernel.withDensity_rnDeriv_eq (κ := alg.policy n) (η := ref.policy n) + (absolutelyContinuous_of_forall_singleton_pos (href.2 n h)) + rw [h_alg, h_ref, ← h_wd] + haveI : SFinite ((ref.policy n h).withDensity + (Kernel.rnDeriv (alg.policy n) (ref.policy n) h)) := by rw [h_wd]; infer_instance exact withDensity_compProd_left - (Kernel.measurable_rnDeriv (alg.policy n) (unif.policy n)).of_uncurry_left - haveI : IsSFiniteKernel ((stepKernel unif (stationaryEnv ν) n).withDensity σ) := by + (Kernel.measurable_rnDeriv (alg.policy n) (ref.policy n)).of_uncurry_left + haveI : IsSFiniteKernel ((stepKernel ref (stationaryEnv ν) n).withDensity σ) := by rw [← h_step]; infer_instance rw [map_hist_succ_eq_compProd_map h₁ n, map_hist_succ_eq_compProd_map h₂ n, ih, h_step, - withDensity_compProd_withDensity (measurable_historyDensity hK alg n) hσ_meas] + withDensity_compProd_withDensity (measurable_historyDensity alg ref n) hσ_meas] exact withDensity_map_equiv_symm - (((measurable_historyDensity hK alg n).comp measurable_fst).mul hσ_meas) + (((measurable_historyDensity alg ref n).comp measurable_fst).mul hσ_meas) end DensityIndependence @@ -370,7 +260,7 @@ section PosteriorIndependence The key theorem: the posterior distribution on the best arm given the observed history is independent of the algorithm used to generate the data. The proof routes through -the uniform algorithm as a reference measure. The posterior on the environment given history +a reference algorithm with full support. The posterior on the environment given history is algorithm-independent (ae wrt the algorithm's own history distribution), and this transfers to the posterior on the best arm via `condDistrib_comp`. -/ @@ -384,23 +274,6 @@ variable {E' : Ω → E} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} variable {alg : Algorithm (Fin K) ℝ} variable {P : Measure Ω} [IsProbabilityMeasure P] -/-- Maps an environment to the best arm (the arm with highest mean reward). -/ -noncomputable def envToBestArm (κ : Kernel (E × Fin K) ℝ) : E → Fin K := - measurableArgmax fun e a ↦ (κ (e, a))[id] - -omit [StandardBorelSpace E] [Nonempty E] [IsMarkovKernel κ] in -lemma measurable_envToBestArm : Measurable (envToBestArm κ) := - measurable_measurableArgmax fun _ ↦ - stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_id.prodMk measurable_const) - -omit [StandardBorelSpace E] [Nonempty E] [IsMarkovKernel κ] [MeasurableSpace Ω] - [IsProbabilityMeasure P] [Nonempty Ω] in -lemma bestAction_eq_envToBestArm_comp_env : - IsBayesAlgEnvSeq.bestAction κ E' = envToBestArm κ ∘ E' := by - funext ω; simp only [Function.comp_apply] - unfold IsBayesAlgEnvSeq.bestAction IsBayesAlgEnvSeq.actionMean envToBestArm - exact (measurableArgmax_eq_of_eq _ _ _ ω).trans (measurableArgmax_congr _ _ ω _ rfl) - omit [StandardBorelSpace E] [Nonempty E] [IsMarkovKernel κ] in /-- The marginal on the history equals `condDistrib (hist) (env) P ∘ₘ Q`. -/ private lemma map_hist_eq_condDistrib_comp @@ -424,19 +297,19 @@ private lemma map_hist_eq_condDistrib_comp omit [StandardBorelSpace E] [Nonempty E] in /-- The history distribution under any algorithm is absolutely continuous w.r.t. the - history distribution under the uniform algorithm (since uniform gives positive - probability to every action). -/ -lemma absolutelyContinuous_map_hist_uniform - (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) (hK : 0 < K) + history distribution under a reference algorithm with full support. -/ +lemma absolutelyContinuous_map_hist + (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) + {ref : Algorithm (Fin K) ℝ} (href : ref.HasFullSupport) {Ωu : Type*} [MeasurableSpace Ωu] [StandardBorelSpace Ωu] [Nonempty Ωu] {Eu : Ωu → E} {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → ℝ} {Pu : Measure Ωu} [IsProbabilityMeasure Pu] - (hu : IsBayesAlgEnvSeq Q κ (Bandits.uniformAlgorithm hK) Eu Au Ru Pu) + (hu : IsBayesAlgEnvSeq Q κ ref Eu Au Ru Pu) (t : ℕ) : P.map (IsAlgEnvSeq.hist A R' t) ≪ Pu.map (IsAlgEnvSeq.hist Au Ru t) := by set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' t) E' P - set κ_unif := condDistrib (IsAlgEnvSeq.hist Au Ru t) Eu Pu + set κ_ref := condDistrib (IsAlgEnvSeq.hist Au Ru t) Eu Pu rw [map_hist_eq_condDistrib_comp Q κ h t, map_hist_eq_condDistrib_comp Q κ hu t, ← Measure.snd_compProd, ← Measure.snd_compProd] have hW_meas : Measurable (fun (ω : Ω) (n : ℕ) => (A n ω, R' n ω)) := @@ -444,7 +317,7 @@ lemma absolutelyContinuous_map_hist_uniform have hWu_meas : Measurable (fun (ω : Ωu) (n : ℕ) => (Au n ω, Ru n ω)) := measurable_pi_lambda _ fun n => (hu.measurable_A n).prodMk (hu.measurable_R n) exact (Measure.AbsolutelyContinuous.compProd_right - (show ∀ᵐ e ∂Q, κ_alg e ≪ κ_unif e from by + (show ∀ᵐ e ∂Q, κ_alg e ≪ κ_ref e from by have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : (ℕ → Fin K × ℝ) → (Iic t → Fin K × ℝ)) = IT.hist t := funext fun ω => funext fun i => Prod.mk.eta @@ -456,10 +329,10 @@ lemma absolutelyContinuous_map_hist_uniform condDistrib_comp E' hW_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have h_cd₂ : ∀ᵐ e ∂Q, κ_unif e = + have h_cd₂ : ∀ᵐ e ∂Q, κ_ref e = (condDistrib (fun ω n => (Au n ω, Ru n ω)) Eu Pu e).map (IT.hist t) := by rw [← hu.hasLaw_env.map_eq] - have h_comp : κ_unif + have h_comp : κ_ref =ᵐ[Pu.map Eu] (condDistrib (fun ω n => (Au n ω, Ru n ω)) Eu Pu).map (IT.hist t) := condDistrib_comp Eu hWu_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he @@ -468,32 +341,33 @@ lemma absolutelyContinuous_map_hist_uniform have hae₂ := hu.ae_IsAlgEnvSeq filter_upwards [h_cd₁, h_cd₂, hae₁, hae₂] with e he₁ he₂ hae₁ hae₂ rw [he₁, he₂, ← h_IT_hist] - exact absolutelyContinuous_map_hist_stationary hK alg _ hae₁ hae₂ t)).map + exact absolutelyContinuous_map_hist_stationary alg ref href _ hae₁ hae₂ t)).map measurable_snd omit [StandardBorelSpace Ω] [Nonempty Ω] in /-- The posterior on the environment given history is algorithm-independent. -/ lemma condDistrib_env_hist_alg_indep - (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) (hK : 0 < K) + (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) + {ref : Algorithm (Fin K) ℝ} (href : ref.HasFullSupport) {Ωu : Type*} [MeasurableSpace Ωu] [StandardBorelSpace Ωu] [Nonempty Ωu] {Eu : Ωu → E} {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → ℝ} {Pu : Measure Ωu} [IsProbabilityMeasure Pu] - (hu : IsBayesAlgEnvSeq Q κ (Bandits.uniformAlgorithm hK) Eu Au Ru Pu) + (hu : IsBayesAlgEnvSeq Q κ ref Eu Au Ru Pu) (t : ℕ) : condDistrib E' (IsAlgEnvSeq.hist A R' t) P =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu := by set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' t) E' P - set κ_unif := condDistrib (IsAlgEnvSeq.hist Au Ru t) Eu Pu - set ρ := historyDensity hK alg t - have hρ_meas := measurable_historyDensity hK alg t - have hρ_ne_top := historyDensity_ne_top hK alg t + set κ_ref := condDistrib (IsAlgEnvSeq.hist Au Ru t) Eu Pu + set ρ := historyDensity alg ref t + have hρ_meas := measurable_historyDensity alg ref t + have hρ_ne_top := historyDensity_ne_top alg ref href t have hW_meas : Measurable (fun (ω : Ω) (n : ℕ) => (A n ω, R' n ω)) := measurable_pi_lambda _ fun n => (h.measurable_A n).prodMk (h.measurable_R n) have hWu_meas : Measurable (fun (ω : Ωu) (n : ℕ) => (Au n ω, Ru n ω)) := measurable_pi_lambda _ fun n => (hu.measurable_A n).prodMk (hu.measurable_R n) - -- Key factorization: κ_alg =ᵐ[Q] κ_unif.withDensity (fun _ => ρ) - have h_wd_ae : κ_alg =ᵐ[Q] κ_unif.withDensity (fun _ => ρ) := by + -- Key factorization: κ_alg =ᵐ[Q] κ_ref.withDensity (fun _ => ρ) + have h_wd_ae : κ_alg =ᵐ[Q] κ_ref.withDensity (fun _ => ρ) := by have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : (ℕ → Fin K × ℝ) → (Iic t → Fin K × ℝ)) = IT.hist t := funext fun ω => funext fun i => Prod.mk.eta @@ -505,10 +379,10 @@ lemma condDistrib_env_hist_alg_indep condDistrib_comp E' hW_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have h_cd₂ : ∀ᵐ e ∂Q, κ_unif e = + have h_cd₂ : ∀ᵐ e ∂Q, κ_ref e = (condDistrib (fun ω n => (Au n ω, Ru n ω)) Eu Pu e).map (IT.hist t) := by rw [← hu.hasLaw_env.map_eq] - have h_comp : κ_unif + have h_comp : κ_ref =ᵐ[Pu.map Eu] (condDistrib (fun ω n => (Au n ω, Ru n ω)) Eu Pu).map (IT.hist t) := condDistrib_comp Eu hWu_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he @@ -519,36 +393,36 @@ lemma condDistrib_env_hist_alg_indep rw [Kernel.withDensity_apply _ (show Measurable (Function.uncurry (fun (_ : E) => ρ)) from hρ_meas.comp measurable_snd), he₁, he₂, ← h_IT_hist] - exact map_hist_eq_withDensity_historyDensity hK alg t _ hae₁ hae₂ - haveI : IsSFiniteKernel (κ_unif.withDensity (fun _ => ρ)) := + exact map_hist_eq_withDensity_historyDensity alg ref href t _ hae₁ hae₂ + haveI : IsSFiniteKernel (κ_ref.withDensity (fun _ => ρ)) := Kernel.IsSFiniteKernel.withDensity _ (fun _ b => hρ_ne_top b) -- Direct condDistrib equality via joint measure argument - -- Show: P.map (hist, E') = P.map hist ⊗ₘ condDistrib Eu hist_u Pu + -- Show: P.map (hist, E') = P.map hist ⊗ₘ condDistrib Eu hist_ref Pu -- using the density factorization and disintegration have h_joint₁ : P.map (fun ω => (E' ω, IsAlgEnvSeq.hist A R' t ω)) = Q ⊗ₘ κ_alg := by rw [← h.hasLaw_env.map_eq] exact (compProd_map_condDistrib (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t).aemeasurable).symm - have h_joint₂ : Pu.map (fun ω => (Eu ω, IsAlgEnvSeq.hist Au Ru t ω)) = Q ⊗ₘ κ_unif := by + have h_joint₂ : Pu.map (fun ω => (Eu ω, IsAlgEnvSeq.hist Au Ru t ω)) = Q ⊗ₘ κ_ref := by rw [← hu.hasLaw_env.map_eq] exact (compProd_map_condDistrib (IsAlgEnvSeq.measurable_hist hu.measurable_A hu.measurable_R t).aemeasurable).symm - -- The swapped joint of P equals P.map hist ⊗ₘ condDistrib Eu hist_u Pu + -- The swapped joint of P equals P.map hist ⊗ₘ condDistrib Eu hist_ref Pu have h_meas_hist := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t have h_meas_hist_u := IsAlgEnvSeq.measurable_hist hu.measurable_A hu.measurable_R t - -- P.map hist = (Pu.map hist_u).withDensity ρ + -- P.map hist = (Pu.map hist_ref).withDensity ρ have h_hist : P.map (IsAlgEnvSeq.hist A R' t) = (Pu.map (IsAlgEnvSeq.hist Au Ru t)).withDensity ρ := by have h_marg₁ : P.map (IsAlgEnvSeq.hist A R' t) = (Q ⊗ₘ κ_alg).map Prod.snd := by rw [← h_joint₁] exact (Measure.map_map measurable_snd (h.measurable_E.prodMk h_meas_hist)).symm - have h_marg₂ : Pu.map (IsAlgEnvSeq.hist Au Ru t) = (Q ⊗ₘ κ_unif).map Prod.snd := by + have h_marg₂ : Pu.map (IsAlgEnvSeq.hist Au Ru t) = (Q ⊗ₘ κ_ref).map Prod.snd := by rw [← h_joint₂] exact (Measure.map_map measurable_snd (hu.measurable_E.prodMk h_meas_hist_u)).symm rw [h_marg₁, h_marg₂, Measure.compProd_congr h_wd_ae, Measure.compProd_withDensity (show Measurable (Function.uncurry (fun (_ : E) => ρ)) from hρ_meas.comp measurable_snd)] - exact withDensity_map_eq' measurable_snd hρ_meas + exact map_withDensity_comp measurable_snd hρ_meas have h_swap : P.map (fun ω => (IsAlgEnvSeq.hist A R' t ω, E' ω)) = P.map (IsAlgEnvSeq.hist A R' t) ⊗ₘ condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu := by have h_uncurry_meas : Measurable (Function.uncurry (fun (_ : E) => ρ)) := @@ -558,11 +432,11 @@ lemma condDistrib_env_hist_alg_indep rw [← h_joint₁] exact (Measure.map_map measurable_swap (h.measurable_E.prodMk h_meas_hist)).symm - _ = (Q ⊗ₘ (κ_unif.withDensity (fun _ => ρ))).map Prod.swap := by + _ = (Q ⊗ₘ (κ_ref.withDensity (fun _ => ρ))).map Prod.swap := by rw [Measure.compProd_congr h_wd_ae] - _ = ((Q ⊗ₘ κ_unif).withDensity (ρ ∘ Prod.snd)).map Prod.swap := by + _ = ((Q ⊗ₘ κ_ref).withDensity (ρ ∘ Prod.snd)).map Prod.swap := by congr 1; exact Measure.compProd_withDensity h_uncurry_meas - _ = ((Q ⊗ₘ κ_unif).map Prod.swap).withDensity (ρ ∘ Prod.fst) := + _ = ((Q ⊗ₘ κ_ref).map Prod.swap).withDensity (ρ ∘ Prod.fst) := map_swap_withDensity_fst hρ_meas _ = (Pu.map (fun ω => (IsAlgEnvSeq.hist Au Ru t ω, Eu ω))).withDensity (ρ ∘ Prod.fst) := by @@ -585,14 +459,15 @@ lemma condDistrib_env_hist_alg_indep omit [StandardBorelSpace Ω] [Nonempty Ω] in /-- The environment posterior is algorithm-independent: it equals the posterior under the -uniform algorithm, which is `IsBayesAlgEnvSeq.posterior`. -/ -lemma posterior_eq_uniform - (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) (hK : 0 < K) (t : ℕ) : +reference algorithm, which is `IsBayesAlgEnvSeq.posterior`. -/ +lemma posterior_eq_ref + (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) + {ref : Algorithm (Fin K) ℝ} (href : ref.HasFullSupport) (t : ℕ) : condDistrib E' (IsAlgEnvSeq.hist A R' t) P =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - IT.bayesTrajMeasurePosterior Q κ (Bandits.uniformAlgorithm hK) t := - condDistrib_env_hist_alg_indep Q κ h hK - (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (Bandits.uniformAlgorithm hK)) t + IT.bayesTrajMeasurePosterior Q κ ref t := + condDistrib_env_hist_alg_indep Q κ h href + (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ ref) t end PosteriorIndependence From 3abd375167cca5fcc2005cef2c53698e1439159b Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 25 Feb 2026 09:51:43 +0000 Subject: [PATCH 052/155] Refactor HistoryDensity (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 2 +- LeanBandits/BanditAlgorithms/Uniform.lean | 15 +++--------- LeanBandits/SequentialLearning/Algorithm.lean | 3 +++ .../SequentialLearning/HistoryDensity.lean | 24 ++++++------------- 4 files changed, 14 insertions(+), 30 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index f0b9e193..d4e9f6f8 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -226,7 +226,7 @@ lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (IsBayesAlgEnvSeq.bestAction κ id) =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) t).map (IsBayesAlgEnvSeq.bestAction κ id) := by - filter_upwards [posterior_eq_ref Q κ h (Bandits.uniformAlgorithm_hasFullSupport hK) t] with x hx + filter_upwards [posterior_eq_ref Q κ h (uniformAlgorithm_IsPositive hK) t] with x hx simp only [Kernel.map_apply _ hm, hx] exact (h.hasCondDistrib_action' t).condDistrib_eq.trans (h_comp.trans h_map).symm diff --git a/LeanBandits/BanditAlgorithms/Uniform.lean b/LeanBandits/BanditAlgorithms/Uniform.lean index de148f93..553003d7 100644 --- a/LeanBandits/BanditAlgorithms/Uniform.lean +++ b/LeanBandits/BanditAlgorithms/Uniform.lean @@ -23,17 +23,8 @@ def uniformAlgorithm (hK : 0 < K) : Algorithm (Fin K) ℝ := { policy _ := Kernel.const _ (uniformOn Set.univ) p0 := uniformOn Set.univ } -/-- The uniform algorithm gives positive probability to every action. -/ -lemma uniformAlgorithm_p0_pos (a : Fin K) : (uniformAlgorithm hK).p0 {a} > 0 := by - simp only [uniformAlgorithm, uniformOn] - refine cond_pos_of_inter_ne_zero MeasurableSet.univ ?_ - simp only [Set.univ_inter, Measure.count_singleton, ne_eq, one_ne_zero, not_false_eq_true] - -/-- The uniform algorithm's policy gives positive probability to every action. -/ -lemma uniformAlgorithm_policy_pos {n : ℕ} (h : Finset.Iic n → Fin K × ℝ) (a : Fin K) : - (uniformAlgorithm hK).policy n h {a} > 0 := by - simp only [uniformAlgorithm, Kernel.const_apply, uniformOn] - refine cond_pos_of_inter_ne_zero MeasurableSet.univ ?_ - simp only [Set.univ_inter, Measure.count_singleton, ne_eq, one_ne_zero, not_false_eq_true] +lemma uniformAlgorithm_IsPositive (hK : 0 < K) : (uniformAlgorithm hK).IsPositive := by + constructor + all_goals simp [uniformAlgorithm, uniformOn, cond_pos_of_inter_ne_zero] end Bandits diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index 0510656f..88b6f514 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -30,6 +30,9 @@ structure Algorithm (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] wher instance (alg : Algorithm α R) (n : ℕ) : IsMarkovKernel (alg.policy n) := alg.h_policy n instance (alg : Algorithm α R) : IsProbabilityMeasure alg.p0 := alg.hp0 +def Algorithm.IsPositive (alg : Algorithm α R) : Prop := + (∀ a, alg.p0 {a} > 0) ∧ (∀ n h a, alg.policy n h {a} > 0) + /-- An algorithm that receives observations in `E × R` created form an algorithm that receives observations in `R` by ignoring the additional information. -/ def Algorithm.prod_left (E : Type*) [MeasurableSpace E] (alg : Algorithm α R) : diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 0bf2065c..6f9a1861 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -29,16 +29,6 @@ namespace Learning variable {K : ℕ} -/-- An algorithm has full support if every action has positive probability under both -the initial measure and every policy. -/ -def Algorithm.HasFullSupport (ref : Algorithm (Fin K) ℝ) : Prop := - (∀ a, ref.p0 {a} > 0) ∧ - (∀ n (h : Iic n → Fin K × ℝ) (a : Fin K), ref.policy n h {a} > 0) - -lemma Bandits.uniformAlgorithm_hasFullSupport (hK : 0 < K) : - (Bandits.uniformAlgorithm hK).HasFullSupport := - ⟨Bandits.uniformAlgorithm_p0_pos, fun _ h a => Bandits.uniformAlgorithm_policy_pos h a⟩ - section AbsolutelyContinuousHist variable [Nonempty (Fin K)] @@ -47,7 +37,7 @@ omit [Nonempty (Fin K)] in /-- The step kernel for a stationary environment decomposes as a product of the policy measure and the reward kernel. -/ private lemma absolutelyContinuous_stepKernel_stationary - (alg ref : Algorithm (Fin K) ℝ) (href : ref.HasFullSupport) + (alg ref : Algorithm (Fin K) ℝ) (href : ref.IsPositive) (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] (n : ℕ) (h : Iic n → Fin K × ℝ) : stepKernel alg (stationaryEnv ν) n h ≪ stepKernel ref (stationaryEnv ν) n h := by @@ -97,7 +87,7 @@ private lemma map_hist_succ_eq_compProd_map history distribution under a reference algorithm with full support, for a stationary environment. -/ private lemma absolutelyContinuous_map_hist_stationary - (alg ref : Algorithm (Fin K) ℝ) (href : ref.HasFullSupport) + (alg ref : Algorithm (Fin K) ℝ) (href : ref.IsPositive) (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] {Ω₁ : Type*} [MeasurableSpace Ω₁] {A₁ : ℕ → Ω₁ → Fin K} {R₁ : ℕ → Ω₁ → ℝ} @@ -168,7 +158,7 @@ private lemma measurable_historyDensity (alg ref : Algorithm (Fin K) ℝ) (t : omit [Nonempty (Fin K)] in private lemma historyDensity_ne_top (alg ref : Algorithm (Fin K) ℝ) - (href : ref.HasFullSupport) (t : ℕ) + (href : ref.IsPositive) (t : ℕ) (h : Iic t → Fin K × ℝ) : historyDensity alg ref t h ≠ ⊤ := by induction t with | zero => exact rnDeriv_ne_top_of_forall_singleton_pos href.1 _ @@ -180,7 +170,7 @@ private lemma historyDensity_ne_top (alg ref : Algorithm (Fin K) ℝ) /-- The history distribution under any algorithm equals the reference algorithm's history distribution weighted by `historyDensity`, for any stationary environment. -/ private lemma map_hist_eq_withDensity_historyDensity - (alg ref : Algorithm (Fin K) ℝ) (href : ref.HasFullSupport) (t : ℕ) + (alg ref : Algorithm (Fin K) ℝ) (href : ref.IsPositive) (t : ℕ) (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] {Ω₁ : Type*} [MeasurableSpace Ω₁] {A₁ : ℕ → Ω₁ → Fin K} {R₁ : ℕ → Ω₁ → ℝ} @@ -300,7 +290,7 @@ omit [StandardBorelSpace E] [Nonempty E] in history distribution under a reference algorithm with full support. -/ lemma absolutelyContinuous_map_hist (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) - {ref : Algorithm (Fin K) ℝ} (href : ref.HasFullSupport) + {ref : Algorithm (Fin K) ℝ} (href : ref.IsPositive) {Ωu : Type*} [MeasurableSpace Ωu] [StandardBorelSpace Ωu] [Nonempty Ωu] {Eu : Ωu → E} {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → ℝ} {Pu : Measure Ωu} [IsProbabilityMeasure Pu] @@ -348,7 +338,7 @@ omit [StandardBorelSpace Ω] [Nonempty Ω] in /-- The posterior on the environment given history is algorithm-independent. -/ lemma condDistrib_env_hist_alg_indep (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) - {ref : Algorithm (Fin K) ℝ} (href : ref.HasFullSupport) + {ref : Algorithm (Fin K) ℝ} (href : ref.IsPositive) {Ωu : Type*} [MeasurableSpace Ωu] [StandardBorelSpace Ωu] [Nonempty Ωu] {Eu : Ωu → E} {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → ℝ} {Pu : Measure Ωu} [IsProbabilityMeasure Pu] @@ -462,7 +452,7 @@ omit [StandardBorelSpace Ω] [Nonempty Ω] in reference algorithm, which is `IsBayesAlgEnvSeq.posterior`. -/ lemma posterior_eq_ref (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) - {ref : Algorithm (Fin K) ℝ} (href : ref.HasFullSupport) (t : ℕ) : + {ref : Algorithm (Fin K) ℝ} (href : ref.IsPositive) (t : ℕ) : condDistrib E' (IsAlgEnvSeq.hist A R' t) P =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] IT.bayesTrajMeasurePosterior Q κ ref t := From fc462fa5a6ecbbaff75ab47137539762025e2af8 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 25 Feb 2026 10:58:56 +0000 Subject: [PATCH 053/155] Refactor HistoryDensity (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 5 +- .../SequentialLearning/HistoryDensity.lean | 519 ++++++++---------- 2 files changed, 239 insertions(+), 285 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index d4e9f6f8..16a98ac3 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -226,8 +226,9 @@ lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (IsBayesAlgEnvSeq.bestAction κ id) =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) t).map (IsBayesAlgEnvSeq.bestAction κ id) := by - filter_upwards [posterior_eq_ref Q κ h (uniformAlgorithm_IsPositive hK) t] with x hx - simp only [Kernel.map_apply _ hm, hx] + filter_upwards [h.condDistrib_env_hist_alg_indep (uniformAlgorithm_IsPositive hK) + (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) t] with x hx + simp only [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hx] exact (h.hasCondDistrib_action' t).condDistrib_eq.trans (h_comp.trans h_map).symm omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 6f9a1861..8b386655 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -9,18 +9,6 @@ import LeanBandits.BanditAlgorithms.Uniform import LeanBandits.ForMathlib.FullSupport import LeanBandits.ForMathlib.WithDensity -/-! -# Algorithm-Independence of Bayesian Posteriors - -The key result: the posterior distribution on the environment (and therefore on the best arm) -given the observed history is independent of the algorithm used to generate the data. - -The proof routes through a reference algorithm with full support. The history distribution -under any algorithm is absolutely continuous w.r.t. the reference algorithm's, with a density -that depends only on action probabilities (not the environment). This density factorization -implies the posteriors agree. --/ - open MeasureTheory ProbabilityTheory Finset Preorder open scoped ENNReal NNReal @@ -29,32 +17,69 @@ namespace Learning variable {K : ℕ} -section AbsolutelyContinuousHist - -variable [Nonempty (Fin K)] - -omit [Nonempty (Fin K)] in -/-- The step kernel for a stationary environment decomposes as a product of the policy - measure and the reward kernel. -/ -private lemma absolutelyContinuous_stepKernel_stationary - (alg ref : Algorithm (Fin K) ℝ) (href : ref.IsPositive) - (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] (n : ℕ) (h : Iic n → Fin K × ℝ) : +/-- The step kernel for a stationary environment under a positive algorithm absolutely + continuously dominates any other algorithm's step kernel. -/ +lemma Algorithm.IsPositive.absolutelyContinuous_stepKernel_stationary + {alg₀ : Algorithm (Fin K) ℝ} (hpos : alg₀.IsPositive) + (alg : Algorithm (Fin K) ℝ) (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] + (n : ℕ) (h : Iic n → Fin K × ℝ) : stepKernel alg (stationaryEnv ν) n h ≪ - stepKernel ref (stationaryEnv ν) n h := by + stepKernel alg₀ (stationaryEnv ν) n h := by have h1 : stepKernel alg (stationaryEnv ν) n h = (alg.policy n h) ⊗ₘ ν := by simp only [stepKernel, stationaryEnv]; ext s hs simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - have h2 : stepKernel ref (stationaryEnv ν) n h = - (ref.policy n h) ⊗ₘ ν := by + have h2 : stepKernel alg₀ (stationaryEnv ν) n h = + (alg₀.policy n h) ⊗ₘ ν := by simp only [stepKernel, stationaryEnv]; ext s hs simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] rw [h1, h2] exact Measure.AbsolutelyContinuous.compProd_left - (absolutelyContinuous_of_forall_singleton_pos (href.2 n h)) _ + (absolutelyContinuous_of_forall_singleton_pos (hpos.2 n h)) _ + +/-- The density of the history distribution under `alg` w.r.t. a positive reference algorithm. +This density depends only on the algorithm's action probabilities, not on the reward kernel. -/ +noncomputable def historyDensity + (alg alg₀ : Algorithm (Fin K) ℝ) : + (t : ℕ) → (Iic t → Fin K × ℝ) → ℝ≥0∞ + | 0 => (alg.p0.rnDeriv alg₀.p0 ∘ Prod.fst) ∘ + MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) + | n + 1 => + let σ : (Iic n → Fin K × ℝ) → (Fin K × ℝ) → ℝ≥0∞ := + fun h ar => Kernel.rnDeriv (alg.policy n) + (alg₀.policy n) h ar.1 + (historyDensity alg alg₀ n ∘ Prod.fst * Function.uncurry σ) ∘ + MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n + +@[fun_prop] +lemma measurable_historyDensity (alg alg₀ : Algorithm (Fin K) ℝ) (t : ℕ) : + Measurable (historyDensity alg alg₀ t) := by + induction t with + | zero => + exact (Measure.measurable_rnDeriv _ _).comp + (measurable_fst.comp (MeasurableEquiv.piUnique _).measurable) + | succ n ih => + exact ((ih.comp measurable_fst).mul + ((Kernel.measurable_rnDeriv _ _).comp + (measurable_fst.prodMk (measurable_fst.comp measurable_snd)))).comp + (MeasurableEquiv.IicSuccProd _ n).measurable + +lemma historyDensity_ne_top (alg alg₀ : Algorithm (Fin K) ℝ) + (hpos : alg₀.IsPositive) (t : ℕ) + (h : Iic t → Fin K × ℝ) : historyDensity alg alg₀ t h ≠ ⊤ := by + induction t with + | zero => exact rnDeriv_ne_top_of_forall_singleton_pos hpos.1 _ + | succ n ih => + exact ENNReal.mul_ne_top (ih _) + (kernel_rnDeriv_ne_top_of_forall_singleton_pos + (fun h' a => hpos.2 n h' a) _ _) + +namespace IsAlgEnvSeq + +variable [Nonempty (Fin K)] /-- The history distribution at time `n + 1` decomposes as a compProd of the history at time `n` and the step kernel, composed with `IicSuccProd.symm`. -/ -private lemma map_hist_succ_eq_compProd_map +lemma map_hist_succ_eq_compProd_map {Ω : Type*} [MeasurableSpace Ω] {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} {alg : Algorithm (Fin K) ℝ} {env : Environment (Fin K) ℝ} @@ -83,122 +108,74 @@ private lemma map_hist_succ_eq_compProd_map (IsAlgEnvSeq.measurable_step (n + 1) (hA _) (hR _)).aemeasurable (stepKernel alg env n)).mp h_cd.condDistrib_eq) +variable {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] +variable {Ω : Type*} [MeasurableSpace Ω] +variable {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} +variable {alg : Algorithm (Fin K) ℝ} +variable {P : Measure Ω} [IsProbabilityMeasure P] +variable {alg₀ : Algorithm (Fin K) ℝ} +variable {Ω₀ : Type*} [MeasurableSpace Ω₀] +variable {A₀ : ℕ → Ω₀ → Fin K} {R₀ : ℕ → Ω₀ → ℝ} +variable {P₀ : Measure Ω₀} [IsProbabilityMeasure P₀] + /-- The history distribution under any algorithm is absolutely continuous w.r.t. the - history distribution under a reference algorithm with full support, + history distribution under a positive reference algorithm, for a stationary environment. -/ -private lemma absolutelyContinuous_map_hist_stationary - (alg ref : Algorithm (Fin K) ℝ) (href : ref.IsPositive) - (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] - {Ω₁ : Type*} [MeasurableSpace Ω₁] - {A₁ : ℕ → Ω₁ → Fin K} {R₁ : ℕ → Ω₁ → ℝ} - {P₁ : Measure Ω₁} [IsProbabilityMeasure P₁] - (h₁ : IsAlgEnvSeq A₁ R₁ alg (stationaryEnv ν) P₁) - {Ω₂ : Type*} [MeasurableSpace Ω₂] - {A₂ : ℕ → Ω₂ → Fin K} {R₂ : ℕ → Ω₂ → ℝ} - {P₂ : Measure Ω₂} [IsProbabilityMeasure P₂] - (h₂ : IsAlgEnvSeq A₂ R₂ ref (stationaryEnv ν) P₂) +lemma absolutelyContinuous_map_hist_stationary + (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) + (hpos : alg₀.IsPositive) + (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ (stationaryEnv ν) P₀) (t : ℕ) : - P₁.map (IsAlgEnvSeq.hist A₁ R₁ t) ≪ P₂.map (IsAlgEnvSeq.hist A₂ R₂ t) := by + P.map (IsAlgEnvSeq.hist A R' t) ≪ P₀.map (IsAlgEnvSeq.hist A₀ R₀ t) := by induction t with | zero => set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) - have h_hist₁ : IsAlgEnvSeq.hist A₁ R₁ 0 = e.symm ∘ IsAlgEnvSeq.step A₁ R₁ 0 := by + have h_hist : IsAlgEnvSeq.hist A R' 0 = e.symm ∘ IsAlgEnvSeq.step A R' 0 := by funext ω ⟨i, hi⟩; have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl - have h_hist₂ : IsAlgEnvSeq.hist A₂ R₂ 0 = e.symm ∘ IsAlgEnvSeq.step A₂ R₂ 0 := by + have h_hist₀ : IsAlgEnvSeq.hist A₀ R₀ 0 = e.symm ∘ IsAlgEnvSeq.step A₀ R₀ 0 := by funext ω ⟨i, hi⟩; have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl - rw [h_hist₁, h_hist₂, + rw [h_hist, h_hist₀, ← Measure.map_map e.symm.measurable - (IsAlgEnvSeq.measurable_step 0 (h₁.measurable_A _) (h₁.measurable_R _)), + (IsAlgEnvSeq.measurable_step 0 (h.measurable_A _) (h.measurable_R _)), ← Measure.map_map e.symm.measurable - (IsAlgEnvSeq.measurable_step 0 (h₂.measurable_A _) (h₂.measurable_R _)), - h₁.hasLaw_step_zero.map_eq, h₂.hasLaw_step_zero.map_eq] + (IsAlgEnvSeq.measurable_step 0 (h₀.measurable_A _) (h₀.measurable_R _)), + h.hasLaw_step_zero.map_eq, h₀.hasLaw_step_zero.map_eq] simp only [stationaryEnv_ν0] exact (Measure.AbsolutelyContinuous.compProd_left - (absolutelyContinuous_of_forall_singleton_pos href.1) _).map + (absolutelyContinuous_of_forall_singleton_pos hpos.1) _).map e.symm.measurable | succ n ih => - rw [map_hist_succ_eq_compProd_map h₁, map_hist_succ_eq_compProd_map h₂] + rw [h.map_hist_succ_eq_compProd_map, h₀.map_hist_succ_eq_compProd_map] exact (Measure.AbsolutelyContinuous.compProd ih - (Filter.Eventually.of_forall fun h => - absolutelyContinuous_stepKernel_stationary alg ref href ν n h)).map + (Filter.Eventually.of_forall fun x => + hpos.absolutelyContinuous_stepKernel_stationary alg ν n x)).map (MeasurableEquiv.IicSuccProd _ n).symm.measurable -end AbsolutelyContinuousHist - -section DensityIndependence - -variable {K : ℕ} [Nonempty (Fin K)] - -/-- The density of the history distribution under `alg` w.r.t. a reference algorithm. -This density depends only on the algorithm's action probabilities, not on the reward kernel. -/ -private noncomputable def historyDensity - (alg ref : Algorithm (Fin K) ℝ) : - (t : ℕ) → (Iic t → Fin K × ℝ) → ℝ≥0∞ - | 0 => (alg.p0.rnDeriv ref.p0 ∘ Prod.fst) ∘ - MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) - | n + 1 => - let σ : (Iic n → Fin K × ℝ) → (Fin K × ℝ) → ℝ≥0∞ := - fun h ar => Kernel.rnDeriv (alg.policy n) - (ref.policy n) h ar.1 - (historyDensity alg ref n ∘ Prod.fst * Function.uncurry σ) ∘ - MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n - -omit [Nonempty (Fin K)] in -private lemma measurable_historyDensity (alg ref : Algorithm (Fin K) ℝ) (t : ℕ) : - Measurable (historyDensity alg ref t) := by - induction t with - | zero => - exact (Measure.measurable_rnDeriv _ _).comp - (measurable_fst.comp (MeasurableEquiv.piUnique _).measurable) - | succ n ih => - exact ((ih.comp measurable_fst).mul - ((Kernel.measurable_rnDeriv _ _).comp - (measurable_fst.prodMk (measurable_fst.comp measurable_snd)))).comp - (MeasurableEquiv.IicSuccProd _ n).measurable - -omit [Nonempty (Fin K)] in -private lemma historyDensity_ne_top (alg ref : Algorithm (Fin K) ℝ) - (href : ref.IsPositive) (t : ℕ) - (h : Iic t → Fin K × ℝ) : historyDensity alg ref t h ≠ ⊤ := by - induction t with - | zero => exact rnDeriv_ne_top_of_forall_singleton_pos href.1 _ - | succ n ih => - exact ENNReal.mul_ne_top (ih _) - (kernel_rnDeriv_ne_top_of_forall_singleton_pos - (fun h' a => href.2 n h' a) _ _) - -/-- The history distribution under any algorithm equals the reference algorithm's history +/-- The history distribution under any algorithm equals the positive reference algorithm's history distribution weighted by `historyDensity`, for any stationary environment. -/ -private lemma map_hist_eq_withDensity_historyDensity - (alg ref : Algorithm (Fin K) ℝ) (href : ref.IsPositive) (t : ℕ) - (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] - {Ω₁ : Type*} [MeasurableSpace Ω₁] - {A₁ : ℕ → Ω₁ → Fin K} {R₁ : ℕ → Ω₁ → ℝ} - {P₁ : Measure Ω₁} [IsProbabilityMeasure P₁] - (h₁ : IsAlgEnvSeq A₁ R₁ alg (stationaryEnv ν) P₁) - {Ω₂ : Type*} [MeasurableSpace Ω₂] - {A₂ : ℕ → Ω₂ → Fin K} {R₂ : ℕ → Ω₂ → ℝ} - {P₂ : Measure Ω₂} [IsProbabilityMeasure P₂] - (h₂ : IsAlgEnvSeq A₂ R₂ ref (stationaryEnv ν) P₂) : - P₁.map (IsAlgEnvSeq.hist A₁ R₁ t) = - (P₂.map (IsAlgEnvSeq.hist A₂ R₂ t)).withDensity (historyDensity alg ref t) := by +lemma map_hist_eq_withDensity_historyDensity + (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) + (hpos : alg₀.IsPositive) (t : ℕ) + (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ (stationaryEnv ν) P₀) : + P.map (IsAlgEnvSeq.hist A R' t) = + (P₀.map (IsAlgEnvSeq.hist A₀ R₀ t)).withDensity (historyDensity alg alg₀ t) := by induction t with | zero => set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) - have h_ac : alg.p0 ≪ ref.p0 := - absolutelyContinuous_of_forall_singleton_pos href.1 - have h_hist₁ : IsAlgEnvSeq.hist A₁ R₁ 0 = e.symm ∘ IsAlgEnvSeq.step A₁ R₁ 0 := by + have h_ac : alg.p0 ≪ alg₀.p0 := + absolutelyContinuous_of_forall_singleton_pos hpos.1 + have h_hist : IsAlgEnvSeq.hist A R' 0 = e.symm ∘ IsAlgEnvSeq.step A R' 0 := by funext ω ⟨i, hi⟩ have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl - have h_hist₂ : IsAlgEnvSeq.hist A₂ R₂ 0 = e.symm ∘ IsAlgEnvSeq.step A₂ R₂ 0 := by + have h_hist₀ : IsAlgEnvSeq.hist A₀ R₀ 0 = e.symm ∘ IsAlgEnvSeq.step A₀ R₀ 0 := by funext ω ⟨i, hi⟩ have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl - rw [h_hist₁, h_hist₂, + rw [h_hist, h_hist₀, ← Measure.map_map e.symm.measurable - (IsAlgEnvSeq.measurable_step 0 (h₁.measurable_A _) (h₁.measurable_R _)), + (IsAlgEnvSeq.measurable_step 0 (h.measurable_A _) (h.measurable_R _)), ← Measure.map_map e.symm.measurable - (IsAlgEnvSeq.measurable_step 0 (h₂.measurable_A _) (h₂.measurable_R _)), - h₁.hasLaw_step_zero.map_eq, h₂.hasLaw_step_zero.map_eq] + (IsAlgEnvSeq.measurable_step 0 (h₀.measurable_A _) (h₀.measurable_R _)), + h.hasLaw_step_zero.map_eq, h₀.hasLaw_step_zero.map_eq] simp only [stationaryEnv_ν0] conv_lhs => rw [← Measure.withDensity_rnDeriv_eq _ _ h_ac] rw [withDensity_compProd_left (Measure.measurable_rnDeriv _ _)] @@ -206,259 +183,235 @@ private lemma map_hist_eq_withDensity_historyDensity ((Measure.measurable_rnDeriv _ _).comp measurable_fst) | succ n ih => let σ : (Iic n → Fin K × ℝ) → (Fin K × ℝ) → ℝ≥0∞ := - fun h ar => Kernel.rnDeriv (alg.policy n) (ref.policy n) h ar.1 + fun x ar => Kernel.rnDeriv (alg.policy n) (alg₀.policy n) x ar.1 have hσ_meas : Measurable (Function.uncurry σ) := (Kernel.measurable_rnDeriv _ _).comp (measurable_fst.prodMk (measurable_fst.comp measurable_snd)) have h_step : stepKernel alg (stationaryEnv ν) n = - (stepKernel ref (stationaryEnv ν) n).withDensity σ := by - ext h : 1 + (stepKernel alg₀ (stationaryEnv ν) n).withDensity σ := by + ext x : 1 rw [Kernel.withDensity_apply _ hσ_meas] - have h_alg : stepKernel alg (stationaryEnv ν) n h = (alg.policy n h) ⊗ₘ ν := by + have h_alg : stepKernel alg (stationaryEnv ν) n x = (alg.policy n x) ⊗ₘ ν := by ext s hs simp only [stepKernel, stationaryEnv, Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - have h_ref : stepKernel ref (stationaryEnv ν) n h = (ref.policy n h) ⊗ₘ ν := by + have h_alg₀ : stepKernel alg₀ (stationaryEnv ν) n x = (alg₀.policy n x) ⊗ₘ ν := by ext s hs simp only [stepKernel, stationaryEnv, Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - have h_wd : ((ref.policy n) h).withDensity - (Kernel.rnDeriv (alg.policy n) (ref.policy n) h) = alg.policy n h := by + have h_wd : ((alg₀.policy n) x).withDensity + (Kernel.rnDeriv (alg.policy n) (alg₀.policy n) x) = alg.policy n x := by rw [← Kernel.withDensity_apply _ (Kernel.measurable_rnDeriv _ _)] - exact Kernel.withDensity_rnDeriv_eq (κ := alg.policy n) (η := ref.policy n) - (absolutelyContinuous_of_forall_singleton_pos (href.2 n h)) - rw [h_alg, h_ref, ← h_wd] - haveI : SFinite ((ref.policy n h).withDensity - (Kernel.rnDeriv (alg.policy n) (ref.policy n) h)) := by + exact Kernel.withDensity_rnDeriv_eq (κ := alg.policy n) (η := alg₀.policy n) + (absolutelyContinuous_of_forall_singleton_pos (hpos.2 n x)) + rw [h_alg, h_alg₀, ← h_wd] + haveI : SFinite ((alg₀.policy n x).withDensity + (Kernel.rnDeriv (alg.policy n) (alg₀.policy n) x)) := by rw [h_wd]; infer_instance exact withDensity_compProd_left - (Kernel.measurable_rnDeriv (alg.policy n) (ref.policy n)).of_uncurry_left - haveI : IsSFiniteKernel ((stepKernel ref (stationaryEnv ν) n).withDensity σ) := by + (Kernel.measurable_rnDeriv (alg.policy n) (alg₀.policy n)).of_uncurry_left + haveI : IsSFiniteKernel ((stepKernel alg₀ (stationaryEnv ν) n).withDensity σ) := by rw [← h_step]; infer_instance - rw [map_hist_succ_eq_compProd_map h₁ n, - map_hist_succ_eq_compProd_map h₂ n, + rw [h.map_hist_succ_eq_compProd_map n, + h₀.map_hist_succ_eq_compProd_map n, ih, h_step, - withDensity_compProd_withDensity (measurable_historyDensity alg ref n) hσ_meas] + withDensity_compProd_withDensity (measurable_historyDensity alg alg₀ n) hσ_meas] exact withDensity_map_equiv_symm - (((measurable_historyDensity alg ref n).comp measurable_fst).mul hσ_meas) - -end DensityIndependence + (((measurable_historyDensity alg alg₀ n).comp measurable_fst).mul hσ_meas) -section PosteriorIndependence +end IsAlgEnvSeq -/-! ### Algorithm-independence of the posterior - -The key theorem: the posterior distribution on the best arm given the observed history -is independent of the algorithm used to generate the data. The proof routes through -a reference algorithm with full support. The posterior on the environment given history -is algorithm-independent (ae wrt the algorithm's own history distribution), and this -transfers to the posterior on the best arm via `condDistrib_comp`. --/ +namespace IsBayesAlgEnvSeq -variable {K : ℕ} [Nonempty (Fin K)] -variable {E : Type*} [MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] -variable (Q : Measure E) [IsProbabilityMeasure Q] -variable (κ : Kernel (E × Fin K) ℝ) [IsMarkovKernel κ] -variable {Ω : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] -variable {E' : Ω → E} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} -variable {alg : Algorithm (Fin K) ℝ} -variable {P : Measure Ω} [IsProbabilityMeasure P] +variable [Nonempty (Fin K)] +variable {𝓔 : Type*} [MeasurableSpace 𝓔] +variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] +variable {κ : Kernel (𝓔 × Fin K) ℝ} -omit [StandardBorelSpace E] [Nonempty E] [IsMarkovKernel κ] in /-- The marginal on the history equals `condDistrib (hist) (env) P ∘ₘ Q`. -/ -private lemma map_hist_eq_condDistrib_comp - {Ω' : Type*} [MeasurableSpace Ω'] [StandardBorelSpace Ω'] [Nonempty Ω'] - {E'' : Ω' → E} {A' : ℕ → Ω' → Fin K} {R'' : ℕ → Ω' → ℝ} - {alg' : Algorithm (Fin K) ℝ} {P' : Measure Ω'} [IsProbabilityMeasure P'] - (h' : IsBayesAlgEnvSeq Q κ alg' E'' A' R'' P') (t : ℕ) : - P'.map (IsAlgEnvSeq.hist A' R'' t) = - condDistrib (IsAlgEnvSeq.hist A' R'' t) E'' P' ∘ₘ Q := by - calc P'.map (IsAlgEnvSeq.hist A' R'' t) - _ = (P'.map (fun ω => (E'' ω, - IsAlgEnvSeq.hist A' R'' t ω))).snd := - (Measure.snd_map_prodMk h'.measurable_E).symm - _ = (P'.map E'' ⊗ₘ condDistrib - (IsAlgEnvSeq.hist A' R'' t) E'' P').snd := by +lemma map_hist_eq_condDistrib_comp + {Ω : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] + {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} + {alg : Algorithm (Fin K) ℝ} {P : Measure Ω} [IsProbabilityMeasure P] + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (t : ℕ) : + P.map (IsAlgEnvSeq.hist A R' t) = + condDistrib (IsAlgEnvSeq.hist A R' t) E P ∘ₘ Q := by + calc P.map (IsAlgEnvSeq.hist A R' t) + _ = (P.map (fun ω => (E ω, + IsAlgEnvSeq.hist A R' t ω))).snd := + (Measure.snd_map_prodMk h.measurable_E).symm + _ = (P.map E ⊗ₘ condDistrib + (IsAlgEnvSeq.hist A R' t) E P).snd := by rw [compProd_map_condDistrib - (IsAlgEnvSeq.measurable_hist h'.measurable_A h'.measurable_R t).aemeasurable] - _ = (Q ⊗ₘ condDistrib (IsAlgEnvSeq.hist A' R'' t) - E'' P').snd := by rw [h'.hasLaw_env.map_eq] + (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t).aemeasurable] + _ = (Q ⊗ₘ condDistrib (IsAlgEnvSeq.hist A R' t) + E P).snd := by rw [h.hasLaw_env.map_eq] _ = _ := Measure.snd_compProd Q _ -omit [StandardBorelSpace E] [Nonempty E] in +variable {Ω : Type*} [MeasurableSpace Ω] +variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} +variable {alg : Algorithm (Fin K) ℝ} +variable {P : Measure Ω} [IsProbabilityMeasure P] +variable {alg₀ : Algorithm (Fin K) ℝ} +variable {Ω₀ : Type*} [MeasurableSpace Ω₀] +variable {E₀ : Ω₀ → 𝓔} {A₀ : ℕ → Ω₀ → Fin K} {R₀ : ℕ → Ω₀ → ℝ} +variable {P₀ : Measure Ω₀} [IsProbabilityMeasure P₀] + /-- The history distribution under any algorithm is absolutely continuous w.r.t. the - history distribution under a reference algorithm with full support. -/ + history distribution under a positive reference algorithm. -/ lemma absolutelyContinuous_map_hist - (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) - {ref : Algorithm (Fin K) ℝ} (href : ref.IsPositive) - {Ωu : Type*} [MeasurableSpace Ωu] [StandardBorelSpace Ωu] [Nonempty Ωu] - {Eu : Ωu → E} {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → ℝ} - {Pu : Measure Ωu} [IsProbabilityMeasure Pu] - (hu : IsBayesAlgEnvSeq Q κ ref Eu Au Ru Pu) + [IsMarkovKernel κ] [StandardBorelSpace Ω] [Nonempty Ω] + [StandardBorelSpace Ω₀] [Nonempty Ω₀] + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) + (hpos : alg₀.IsPositive) + (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (t : ℕ) : P.map (IsAlgEnvSeq.hist A R' t) ≪ - Pu.map (IsAlgEnvSeq.hist Au Ru t) := by - set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' t) E' P - set κ_ref := condDistrib (IsAlgEnvSeq.hist Au Ru t) Eu Pu - rw [map_hist_eq_condDistrib_comp Q κ h t, map_hist_eq_condDistrib_comp Q κ hu t, + P₀.map (IsAlgEnvSeq.hist A₀ R₀ t) := by + set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' t) E P + set κ₀ := condDistrib (IsAlgEnvSeq.hist A₀ R₀ t) E₀ P₀ + rw [h.map_hist_eq_condDistrib_comp t, h₀.map_hist_eq_condDistrib_comp t, ← Measure.snd_compProd, ← Measure.snd_compProd] have hW_meas : Measurable (fun (ω : Ω) (n : ℕ) => (A n ω, R' n ω)) := measurable_pi_lambda _ fun n => (h.measurable_A n).prodMk (h.measurable_R n) - have hWu_meas : Measurable (fun (ω : Ωu) (n : ℕ) => (Au n ω, Ru n ω)) := - measurable_pi_lambda _ fun n => (hu.measurable_A n).prodMk (hu.measurable_R n) + have hW₀_meas : Measurable (fun (ω : Ω₀) (n : ℕ) => (A₀ n ω, R₀ n ω)) := + measurable_pi_lambda _ fun n => (h₀.measurable_A n).prodMk (h₀.measurable_R n) exact (Measure.AbsolutelyContinuous.compProd_right - (show ∀ᵐ e ∂Q, κ_alg e ≪ κ_ref e from by + (show ∀ᵐ e ∂Q, κ_alg e ≪ κ₀ e from by have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : (ℕ → Fin K × ℝ) → (Iic t → Fin K × ℝ)) = IT.hist t := funext fun ω => funext fun i => Prod.mk.eta - have h_cd₁ : ∀ᵐ e ∂Q, κ_alg e = - (condDistrib (fun ω n => (A n ω, R' n ω)) E' P e).map (IT.hist t) := by + have h_cd : ∀ᵐ e ∂Q, κ_alg e = + (condDistrib (fun ω n => (A n ω, R' n ω)) E P e).map (IT.hist t) := by rw [← h.hasLaw_env.map_eq] have h_comp : κ_alg - =ᵐ[P.map E'] (condDistrib (fun ω n => (A n ω, R' n ω)) E' P).map (IT.hist t) := - condDistrib_comp E' hW_meas.aemeasurable (IT.measurable_hist t) + =ᵐ[P.map E] (condDistrib (fun ω n => (A n ω, R' n ω)) E P).map (IT.hist t) := + condDistrib_comp E hW_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have h_cd₂ : ∀ᵐ e ∂Q, κ_ref e = - (condDistrib (fun ω n => (Au n ω, Ru n ω)) Eu Pu e).map (IT.hist t) := by - rw [← hu.hasLaw_env.map_eq] - have h_comp : κ_ref - =ᵐ[Pu.map Eu] (condDistrib (fun ω n => (Au n ω, Ru n ω)) Eu Pu).map (IT.hist t) := - condDistrib_comp Eu hWu_meas.aemeasurable (IT.measurable_hist t) + have h_cd₀ : ∀ᵐ e ∂Q, κ₀ e = + (condDistrib (fun ω n => (A₀ n ω, R₀ n ω)) E₀ P₀ e).map (IT.hist t) := by + rw [← h₀.hasLaw_env.map_eq] + have h_comp : κ₀ + =ᵐ[P₀.map E₀] (condDistrib (fun ω n => (A₀ n ω, R₀ n ω)) E₀ P₀).map (IT.hist t) := + condDistrib_comp E₀ hW₀_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have hae₁ := h.ae_IsAlgEnvSeq - have hae₂ := hu.ae_IsAlgEnvSeq - filter_upwards [h_cd₁, h_cd₂, hae₁, hae₂] with e he₁ he₂ hae₁ hae₂ - rw [he₁, he₂, ← h_IT_hist] - exact absolutelyContinuous_map_hist_stationary alg ref href _ hae₁ hae₂ t)).map + have hae := h.ae_IsAlgEnvSeq + have hae₀ := h₀.ae_IsAlgEnvSeq + filter_upwards [h_cd, h_cd₀, hae, hae₀] with e he he₀ hae hae₀ + rw [he, he₀, ← h_IT_hist] + exact hae.absolutelyContinuous_map_hist_stationary hpos hae₀ t)).map measurable_snd -omit [StandardBorelSpace Ω] [Nonempty Ω] in +variable [StandardBorelSpace 𝓔] [Nonempty 𝓔] [IsMarkovKernel κ] + /-- The posterior on the environment given history is algorithm-independent. -/ lemma condDistrib_env_hist_alg_indep - (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) - {ref : Algorithm (Fin K) ℝ} (href : ref.IsPositive) - {Ωu : Type*} [MeasurableSpace Ωu] [StandardBorelSpace Ωu] [Nonempty Ωu] - {Eu : Ωu → E} {Au : ℕ → Ωu → Fin K} {Ru : ℕ → Ωu → ℝ} - {Pu : Measure Ωu} [IsProbabilityMeasure Pu] - (hu : IsBayesAlgEnvSeq Q κ ref Eu Au Ru Pu) + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) + (hpos : alg₀.IsPositive) + (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (t : ℕ) : - condDistrib E' (IsAlgEnvSeq.hist A R' t) P + condDistrib E (IsAlgEnvSeq.hist A R' t) P =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu := by - set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' t) E' P - set κ_ref := condDistrib (IsAlgEnvSeq.hist Au Ru t) Eu Pu - set ρ := historyDensity alg ref t - have hρ_meas := measurable_historyDensity alg ref t - have hρ_ne_top := historyDensity_ne_top alg ref href t + condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀ := by + set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' t) E P + set κ₀ := condDistrib (IsAlgEnvSeq.hist A₀ R₀ t) E₀ P₀ + set ρ := historyDensity alg alg₀ t + have hρ_meas := measurable_historyDensity alg alg₀ t + have hρ_ne_top := historyDensity_ne_top alg alg₀ hpos t have hW_meas : Measurable (fun (ω : Ω) (n : ℕ) => (A n ω, R' n ω)) := measurable_pi_lambda _ fun n => (h.measurable_A n).prodMk (h.measurable_R n) - have hWu_meas : Measurable (fun (ω : Ωu) (n : ℕ) => (Au n ω, Ru n ω)) := - measurable_pi_lambda _ fun n => (hu.measurable_A n).prodMk (hu.measurable_R n) - -- Key factorization: κ_alg =ᵐ[Q] κ_ref.withDensity (fun _ => ρ) - have h_wd_ae : κ_alg =ᵐ[Q] κ_ref.withDensity (fun _ => ρ) := by + have hW₀_meas : Measurable (fun (ω : Ω₀) (n : ℕ) => (A₀ n ω, R₀ n ω)) := + measurable_pi_lambda _ fun n => (h₀.measurable_A n).prodMk (h₀.measurable_R n) + -- Key factorization: κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) + have h_wd_ae : κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) := by have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : (ℕ → Fin K × ℝ) → (Iic t → Fin K × ℝ)) = IT.hist t := funext fun ω => funext fun i => Prod.mk.eta - have h_cd₁ : ∀ᵐ e ∂Q, κ_alg e = - (condDistrib (fun ω n => (A n ω, R' n ω)) E' P e).map (IT.hist t) := by + have h_cd : ∀ᵐ e ∂Q, κ_alg e = + (condDistrib (fun ω n => (A n ω, R' n ω)) E P e).map (IT.hist t) := by rw [← h.hasLaw_env.map_eq] have h_comp : κ_alg - =ᵐ[P.map E'] (condDistrib (fun ω n => (A n ω, R' n ω)) E' P).map (IT.hist t) := - condDistrib_comp E' hW_meas.aemeasurable (IT.measurable_hist t) + =ᵐ[P.map E] (condDistrib (fun ω n => (A n ω, R' n ω)) E P).map (IT.hist t) := + condDistrib_comp E hW_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have h_cd₂ : ∀ᵐ e ∂Q, κ_ref e = - (condDistrib (fun ω n => (Au n ω, Ru n ω)) Eu Pu e).map (IT.hist t) := by - rw [← hu.hasLaw_env.map_eq] - have h_comp : κ_ref - =ᵐ[Pu.map Eu] (condDistrib (fun ω n => (Au n ω, Ru n ω)) Eu Pu).map (IT.hist t) := - condDistrib_comp Eu hWu_meas.aemeasurable (IT.measurable_hist t) + have h_cd₀ : ∀ᵐ e ∂Q, κ₀ e = + (condDistrib (fun ω n => (A₀ n ω, R₀ n ω)) E₀ P₀ e).map (IT.hist t) := by + rw [← h₀.hasLaw_env.map_eq] + have h_comp : κ₀ + =ᵐ[P₀.map E₀] (condDistrib (fun ω n => (A₀ n ω, R₀ n ω)) E₀ P₀).map (IT.hist t) := + condDistrib_comp E₀ hW₀_meas.aemeasurable (IT.measurable_hist t) filter_upwards [h_comp] with e he rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have hae₁ := h.ae_IsAlgEnvSeq - have hae₂ := hu.ae_IsAlgEnvSeq - filter_upwards [h_cd₁, h_cd₂, hae₁, hae₂] with e he₁ he₂ hae₁ hae₂ + have hae := h.ae_IsAlgEnvSeq + have hae₀ := h₀.ae_IsAlgEnvSeq + filter_upwards [h_cd, h_cd₀, hae, hae₀] with e he he₀ hae hae₀ rw [Kernel.withDensity_apply _ - (show Measurable (Function.uncurry (fun (_ : E) => ρ)) from hρ_meas.comp measurable_snd), - he₁, he₂, ← h_IT_hist] - exact map_hist_eq_withDensity_historyDensity alg ref href t _ hae₁ hae₂ - haveI : IsSFiniteKernel (κ_ref.withDensity (fun _ => ρ)) := + (show Measurable (Function.uncurry (fun (_ : 𝓔) => ρ)) from hρ_meas.comp measurable_snd), + he, he₀, ← h_IT_hist] + exact hae.map_hist_eq_withDensity_historyDensity hpos t hae₀ + haveI : IsSFiniteKernel (κ₀.withDensity (fun _ => ρ)) := Kernel.IsSFiniteKernel.withDensity _ (fun _ b => hρ_ne_top b) -- Direct condDistrib equality via joint measure argument - -- Show: P.map (hist, E') = P.map hist ⊗ₘ condDistrib Eu hist_ref Pu - -- using the density factorization and disintegration - have h_joint₁ : P.map (fun ω => (E' ω, IsAlgEnvSeq.hist A R' t ω)) = Q ⊗ₘ κ_alg := by + have h_joint : P.map (fun ω => (E ω, IsAlgEnvSeq.hist A R' t ω)) = Q ⊗ₘ κ_alg := by rw [← h.hasLaw_env.map_eq] exact (compProd_map_condDistrib (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t).aemeasurable).symm - have h_joint₂ : Pu.map (fun ω => (Eu ω, IsAlgEnvSeq.hist Au Ru t ω)) = Q ⊗ₘ κ_ref := by - rw [← hu.hasLaw_env.map_eq] + have h_joint₀ : P₀.map (fun ω => (E₀ ω, IsAlgEnvSeq.hist A₀ R₀ t ω)) = Q ⊗ₘ κ₀ := by + rw [← h₀.hasLaw_env.map_eq] exact (compProd_map_condDistrib - (IsAlgEnvSeq.measurable_hist hu.measurable_A hu.measurable_R t).aemeasurable).symm - -- The swapped joint of P equals P.map hist ⊗ₘ condDistrib Eu hist_ref Pu + (IsAlgEnvSeq.measurable_hist h₀.measurable_A h₀.measurable_R t).aemeasurable).symm have h_meas_hist := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t - have h_meas_hist_u := IsAlgEnvSeq.measurable_hist hu.measurable_A hu.measurable_R t - -- P.map hist = (Pu.map hist_ref).withDensity ρ + have h_meas_hist₀ := IsAlgEnvSeq.measurable_hist h₀.measurable_A h₀.measurable_R t + -- P.map hist = (P₀.map hist₀).withDensity ρ have h_hist : P.map (IsAlgEnvSeq.hist A R' t) - = (Pu.map (IsAlgEnvSeq.hist Au Ru t)).withDensity ρ := by - have h_marg₁ : P.map (IsAlgEnvSeq.hist A R' t) = (Q ⊗ₘ κ_alg).map Prod.snd := by - rw [← h_joint₁] + = (P₀.map (IsAlgEnvSeq.hist A₀ R₀ t)).withDensity ρ := by + have h_marg : P.map (IsAlgEnvSeq.hist A R' t) = (Q ⊗ₘ κ_alg).map Prod.snd := by + rw [← h_joint] exact (Measure.map_map measurable_snd (h.measurable_E.prodMk h_meas_hist)).symm - have h_marg₂ : Pu.map (IsAlgEnvSeq.hist Au Ru t) = (Q ⊗ₘ κ_ref).map Prod.snd := by - rw [← h_joint₂] - exact (Measure.map_map measurable_snd (hu.measurable_E.prodMk h_meas_hist_u)).symm - rw [h_marg₁, h_marg₂, Measure.compProd_congr h_wd_ae, + have h_marg₀ : P₀.map (IsAlgEnvSeq.hist A₀ R₀ t) = (Q ⊗ₘ κ₀).map Prod.snd := by + rw [← h_joint₀] + exact (Measure.map_map measurable_snd (h₀.measurable_E.prodMk h_meas_hist₀)).symm + rw [h_marg, h_marg₀, Measure.compProd_congr h_wd_ae, Measure.compProd_withDensity - (show Measurable (Function.uncurry (fun (_ : E) => ρ)) from hρ_meas.comp measurable_snd)] + (show Measurable (Function.uncurry (fun (_ : 𝓔) => ρ)) from hρ_meas.comp measurable_snd)] exact map_withDensity_comp measurable_snd hρ_meas - have h_swap : P.map (fun ω => (IsAlgEnvSeq.hist A R' t ω, E' ω)) - = P.map (IsAlgEnvSeq.hist A R' t) ⊗ₘ condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu := by - have h_uncurry_meas : Measurable (Function.uncurry (fun (_ : E) => ρ)) := + have h_swap : P.map (fun ω => (IsAlgEnvSeq.hist A R' t ω, E ω)) + = P.map (IsAlgEnvSeq.hist A R' t) ⊗ₘ condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀ := by + have h_uncurry_meas : Measurable (Function.uncurry (fun (_ : 𝓔) => ρ)) := hρ_meas.comp measurable_snd - calc P.map (fun ω => (IsAlgEnvSeq.hist A R' t ω, E' ω)) + calc P.map (fun ω => (IsAlgEnvSeq.hist A R' t ω, E ω)) _ = (Q ⊗ₘ κ_alg).map Prod.swap := by - rw [← h_joint₁] + rw [← h_joint] exact (Measure.map_map measurable_swap (h.measurable_E.prodMk h_meas_hist)).symm - _ = (Q ⊗ₘ (κ_ref.withDensity (fun _ => ρ))).map Prod.swap := by + _ = (Q ⊗ₘ (κ₀.withDensity (fun _ => ρ))).map Prod.swap := by rw [Measure.compProd_congr h_wd_ae] - _ = ((Q ⊗ₘ κ_ref).withDensity (ρ ∘ Prod.snd)).map Prod.swap := by + _ = ((Q ⊗ₘ κ₀).withDensity (ρ ∘ Prod.snd)).map Prod.swap := by congr 1; exact Measure.compProd_withDensity h_uncurry_meas - _ = ((Q ⊗ₘ κ_ref).map Prod.swap).withDensity (ρ ∘ Prod.fst) := + _ = ((Q ⊗ₘ κ₀).map Prod.swap).withDensity (ρ ∘ Prod.fst) := map_swap_withDensity_fst hρ_meas - _ = (Pu.map (fun ω => (IsAlgEnvSeq.hist Au Ru t ω, Eu ω))).withDensity + _ = (P₀.map (fun ω => (IsAlgEnvSeq.hist A₀ R₀ t ω, E₀ ω))).withDensity (ρ ∘ Prod.fst) := by - congr 1; rw [← h_joint₂] + congr 1; rw [← h_joint₀] exact Measure.map_map measurable_swap - (hu.measurable_E.prodMk h_meas_hist_u) - _ = (Pu.map (IsAlgEnvSeq.hist Au Ru t) ⊗ₘ - condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu).withDensity + (h₀.measurable_E.prodMk h_meas_hist₀) + _ = (P₀.map (IsAlgEnvSeq.hist A₀ R₀ t) ⊗ₘ + condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀).withDensity (ρ ∘ Prod.fst) := by - rw [← compProd_map_condDistrib hu.measurable_E.aemeasurable] - _ = (Pu.map (IsAlgEnvSeq.hist Au Ru t)).withDensity ρ ⊗ₘ - condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu := + rw [← compProd_map_condDistrib h₀.measurable_E.aemeasurable] + _ = (P₀.map (IsAlgEnvSeq.hist A₀ R₀ t)).withDensity ρ ⊗ₘ + condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀ := (withDensity_compProd_left hρ_meas).symm _ = P.map (IsAlgEnvSeq.hist A R' t) ⊗ₘ - condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu := by + condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀ := by rw [h_hist] -- By uniqueness of disintegration exact (condDistrib_ae_eq_iff_measure_eq_compProd _ - h.measurable_E.aemeasurable (condDistrib Eu (IsAlgEnvSeq.hist Au Ru t) Pu)).mpr h_swap - -omit [StandardBorelSpace Ω] [Nonempty Ω] in -/-- The environment posterior is algorithm-independent: it equals the posterior under the -reference algorithm, which is `IsBayesAlgEnvSeq.posterior`. -/ -lemma posterior_eq_ref - (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) - {ref : Algorithm (Fin K) ℝ} (href : ref.IsPositive) (t : ℕ) : - condDistrib E' (IsAlgEnvSeq.hist A R' t) P - =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - IT.bayesTrajMeasurePosterior Q κ ref t := - condDistrib_env_hist_alg_indep Q κ h href - (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ ref) t + h.measurable_E.aemeasurable (condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀)).mpr h_swap -end PosteriorIndependence +end IsBayesAlgEnvSeq end Learning From f8293fd9ee7beda1783e4c397baf1e32b951e2c5 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 25 Feb 2026 12:58:42 +0000 Subject: [PATCH 054/155] Refactor HistoryDensity (in progress) --- .../BayesStationaryEnv.lean | 10 ++ .../SequentialLearning/HistoryDensity.lean | 126 ++++++++---------- 2 files changed, 65 insertions(+), 71 deletions(-) diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 60fccbfd..eb301e75 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -98,6 +98,16 @@ lemma hasCondDistrib_reward' [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg end Laws +section Maps + +lemma map_hist_eq_condDistrib_comp [SFinite Q] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (t : ℕ) : + P.map (IsAlgEnvSeq.hist A R' t) = condDistrib (IsAlgEnvSeq.hist A R' t) E P ∘ₘ Q := by + rw [← Measure.snd_map_prodMk h.measurable_E, ← compProd_map_condDistrib + (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t).aemeasurable, + h.hasLaw_env.map_eq, Measure.snd_compProd] + +end Maps + section CondDistribIsAlgEnvSeq lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 8b386655..7c9be229 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -3,11 +3,9 @@ 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.SequentialLearning.StationaryEnv -import LeanBandits.SequentialLearning.BayesStationaryEnv -import LeanBandits.BanditAlgorithms.Uniform import LeanBandits.ForMathlib.FullSupport import LeanBandits.ForMathlib.WithDensity +import LeanBandits.SequentialLearning.BayesStationaryEnv open MeasureTheory ProbabilityTheory Finset Preorder @@ -15,43 +13,28 @@ open scoped ENNReal NNReal namespace Learning -variable {K : ℕ} +variable {α : Type*} {R : Type*} [MeasurableSpace α] [MeasurableSpace R] -/-- The step kernel for a stationary environment under a positive algorithm absolutely - continuously dominates any other algorithm's step kernel. -/ -lemma Algorithm.IsPositive.absolutelyContinuous_stepKernel_stationary - {alg₀ : Algorithm (Fin K) ℝ} (hpos : alg₀.IsPositive) - (alg : Algorithm (Fin K) ℝ) (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] - (n : ℕ) (h : Iic n → Fin K × ℝ) : - stepKernel alg (stationaryEnv ν) n h ≪ - stepKernel alg₀ (stationaryEnv ν) n h := by - have h1 : stepKernel alg (stationaryEnv ν) n h = (alg.policy n h) ⊗ₘ ν := by - simp only [stepKernel, stationaryEnv]; ext s hs - simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - have h2 : stepKernel alg₀ (stationaryEnv ν) n h = - (alg₀.policy n h) ⊗ₘ ν := by - simp only [stepKernel, stationaryEnv]; ext s hs - simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - rw [h1, h2] - exact Measure.AbsolutelyContinuous.compProd_left - (absolutelyContinuous_of_forall_singleton_pos (hpos.2 n h)) _ +section HistoryDensity + +variable [MeasurableSpace.CountablyGenerated α] /-- The density of the history distribution under `alg` w.r.t. a positive reference algorithm. This density depends only on the algorithm's action probabilities, not on the reward kernel. -/ noncomputable def historyDensity - (alg alg₀ : Algorithm (Fin K) ℝ) : - (t : ℕ) → (Iic t → Fin K × ℝ) → ℝ≥0∞ + (alg alg₀ : Algorithm α R) : + (t : ℕ) → (Iic t → α × R) → ℝ≥0∞ | 0 => (alg.p0.rnDeriv alg₀.p0 ∘ Prod.fst) ∘ - MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) + MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => α × R) | n + 1 => - let σ : (Iic n → Fin K × ℝ) → (Fin K × ℝ) → ℝ≥0∞ := + let σ : (Iic n → α × R) → (α × R) → ℝ≥0∞ := fun h ar => Kernel.rnDeriv (alg.policy n) (alg₀.policy n) h ar.1 (historyDensity alg alg₀ n ∘ Prod.fst * Function.uncurry σ) ∘ - MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n + MeasurableEquiv.IicSuccProd (fun _ : ℕ => α × R) n @[fun_prop] -lemma measurable_historyDensity (alg alg₀ : Algorithm (Fin K) ℝ) (t : ℕ) : +lemma measurable_historyDensity (alg alg₀ : Algorithm α R) (t : ℕ) : Measurable (historyDensity alg alg₀ t) := by induction t with | zero => @@ -63,9 +46,9 @@ lemma measurable_historyDensity (alg alg₀ : Algorithm (Fin K) ℝ) (t : ℕ) : (measurable_fst.prodMk (measurable_fst.comp measurable_snd)))).comp (MeasurableEquiv.IicSuccProd _ n).measurable -lemma historyDensity_ne_top (alg alg₀ : Algorithm (Fin K) ℝ) +lemma historyDensity_ne_top (alg alg₀ : Algorithm α R) (hpos : alg₀.IsPositive) (t : ℕ) - (h : Iic t → Fin K × ℝ) : historyDensity alg alg₀ t h ≠ ⊤ := by + (h : Iic t → α × R) : historyDensity alg alg₀ t h ≠ ⊤ := by induction t with | zero => exact rnDeriv_ne_top_of_forall_singleton_pos hpos.1 _ | succ n ih => @@ -73,22 +56,43 @@ lemma historyDensity_ne_top (alg alg₀ : Algorithm (Fin K) ℝ) (kernel_rnDeriv_ne_top_of_forall_singleton_pos (fun h' a => hpos.2 n h' a) _ _) +end HistoryDensity + +/-- The step kernel for a stationary environment under a positive algorithm absolutely + continuously dominates any other algorithm's step kernel. -/ +lemma Algorithm.IsPositive.absolutelyContinuous_stepKernel_stationary + {alg₀ : Algorithm α R} (hpos : alg₀.IsPositive) + (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] + (n : ℕ) (h : Iic n → α × R) : + stepKernel alg (stationaryEnv ν) n h ≪ + stepKernel alg₀ (stationaryEnv ν) n h := by + have h1 : stepKernel alg (stationaryEnv ν) n h = (alg.policy n h) ⊗ₘ ν := by + simp only [stepKernel, stationaryEnv]; ext s hs + simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] + have h2 : stepKernel alg₀ (stationaryEnv ν) n h = + (alg₀.policy n h) ⊗ₘ ν := by + simp only [stepKernel, stationaryEnv]; ext s hs + simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] + rw [h1, h2] + exact Measure.AbsolutelyContinuous.compProd_left + (absolutelyContinuous_of_forall_singleton_pos (hpos.2 n h)) _ + namespace IsAlgEnvSeq -variable [Nonempty (Fin K)] +variable [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] /-- The history distribution at time `n + 1` decomposes as a compProd of the history at time `n` and the step kernel, composed with `IicSuccProd.symm`. -/ lemma map_hist_succ_eq_compProd_map {Ω : Type*} [MeasurableSpace Ω] - {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} - {alg : Algorithm (Fin K) ℝ} {env : Environment (Fin K) ℝ} + {A : ℕ → Ω → α} {R' : ℕ → Ω → R} + {alg : Algorithm α R} {env : Environment α R} {P : Measure Ω} [IsFiniteMeasure P] (h : IsAlgEnvSeq A R' alg env P) (n : ℕ) : P.map (IsAlgEnvSeq.hist A R' (n + 1)) = (P.map (IsAlgEnvSeq.hist A R' n) ⊗ₘ stepKernel alg env n).map - (MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n).symm := by - set e := MeasurableEquiv.IicSuccProd (fun _ : ℕ => Fin K × ℝ) n + (MeasurableEquiv.IicSuccProd (fun _ : ℕ => α × R) n).symm := by + set e := MeasurableEquiv.IicSuccProd (fun _ : ℕ => α × R) n have hA := h.measurable_A; have hR := h.measurable_R have h_func : IsAlgEnvSeq.hist A R' (n + 1) = e.symm ∘ (fun ω => (IsAlgEnvSeq.hist A R' n ω, IsAlgEnvSeq.step A R' (n + 1) ω)) := by @@ -108,14 +112,14 @@ lemma map_hist_succ_eq_compProd_map (IsAlgEnvSeq.measurable_step (n + 1) (hA _) (hR _)).aemeasurable (stepKernel alg env n)).mp h_cd.condDistrib_eq) -variable {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] +variable {ν : Kernel α R} [IsMarkovKernel ν] variable {Ω : Type*} [MeasurableSpace Ω] -variable {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} -variable {alg : Algorithm (Fin K) ℝ} +variable {A : ℕ → Ω → α} {R' : ℕ → Ω → R} +variable {alg : Algorithm α R} variable {P : Measure Ω} [IsProbabilityMeasure P] -variable {alg₀ : Algorithm (Fin K) ℝ} +variable {alg₀ : Algorithm α R} variable {Ω₀ : Type*} [MeasurableSpace Ω₀] -variable {A₀ : ℕ → Ω₀ → Fin K} {R₀ : ℕ → Ω₀ → ℝ} +variable {A₀ : ℕ → Ω₀ → α} {R₀ : ℕ → Ω₀ → R} variable {P₀ : Measure Ω₀} [IsProbabilityMeasure P₀] /-- The history distribution under any algorithm is absolutely continuous w.r.t. the @@ -129,7 +133,7 @@ lemma absolutelyContinuous_map_hist_stationary P.map (IsAlgEnvSeq.hist A R' t) ≪ P₀.map (IsAlgEnvSeq.hist A₀ R₀ t) := by induction t with | zero => - set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) + set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => α × R) have h_hist : IsAlgEnvSeq.hist A R' 0 = e.symm ∘ IsAlgEnvSeq.step A R' 0 := by funext ω ⟨i, hi⟩; have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl have h_hist₀ : IsAlgEnvSeq.hist A₀ R₀ 0 = e.symm ∘ IsAlgEnvSeq.step A₀ R₀ 0 := by @@ -161,7 +165,7 @@ lemma map_hist_eq_withDensity_historyDensity (P₀.map (IsAlgEnvSeq.hist A₀ R₀ t)).withDensity (historyDensity alg alg₀ t) := by induction t with | zero => - set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => Fin K × ℝ) + set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => α × R) have h_ac : alg.p0 ≪ alg₀.p0 := absolutelyContinuous_of_forall_singleton_pos hpos.1 have h_hist : IsAlgEnvSeq.hist A R' 0 = e.symm ∘ IsAlgEnvSeq.step A R' 0 := by @@ -182,7 +186,7 @@ lemma map_hist_eq_withDensity_historyDensity exact withDensity_map_equiv_symm ((Measure.measurable_rnDeriv _ _).comp measurable_fst) | succ n ih => - let σ : (Iic n → Fin K × ℝ) → (Fin K × ℝ) → ℝ≥0∞ := + let σ : (Iic n → α × R) → (α × R) → ℝ≥0∞ := fun x ar => Kernel.rnDeriv (alg.policy n) (alg₀.policy n) x ar.1 have hσ_meas : Measurable (Function.uncurry σ) := (Kernel.measurable_rnDeriv _ _).comp @@ -223,38 +227,18 @@ end IsAlgEnvSeq namespace IsBayesAlgEnvSeq -variable [Nonempty (Fin K)] variable {𝓔 : Type*} [MeasurableSpace 𝓔] +variable [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] -variable {κ : Kernel (𝓔 × Fin K) ℝ} - -/-- The marginal on the history equals `condDistrib (hist) (env) P ∘ₘ Q`. -/ -lemma map_hist_eq_condDistrib_comp - {Ω : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω] - {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} - {alg : Algorithm (Fin K) ℝ} {P : Measure Ω} [IsProbabilityMeasure P] - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (t : ℕ) : - P.map (IsAlgEnvSeq.hist A R' t) = - condDistrib (IsAlgEnvSeq.hist A R' t) E P ∘ₘ Q := by - calc P.map (IsAlgEnvSeq.hist A R' t) - _ = (P.map (fun ω => (E ω, - IsAlgEnvSeq.hist A R' t ω))).snd := - (Measure.snd_map_prodMk h.measurable_E).symm - _ = (P.map E ⊗ₘ condDistrib - (IsAlgEnvSeq.hist A R' t) E P).snd := by - rw [compProd_map_condDistrib - (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t).aemeasurable] - _ = (Q ⊗ₘ condDistrib (IsAlgEnvSeq.hist A R' t) - E P).snd := by rw [h.hasLaw_env.map_eq] - _ = _ := Measure.snd_compProd Q _ +variable {κ : Kernel (𝓔 × α) R} variable {Ω : Type*} [MeasurableSpace Ω] -variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} -variable {alg : Algorithm (Fin K) ℝ} +variable {E : Ω → 𝓔} {A : ℕ → Ω → α} {R' : ℕ → Ω → R} +variable {alg : Algorithm α R} variable {P : Measure Ω} [IsProbabilityMeasure P] -variable {alg₀ : Algorithm (Fin K) ℝ} +variable {alg₀ : Algorithm α R} variable {Ω₀ : Type*} [MeasurableSpace Ω₀] -variable {E₀ : Ω₀ → 𝓔} {A₀ : ℕ → Ω₀ → Fin K} {R₀ : ℕ → Ω₀ → ℝ} +variable {E₀ : Ω₀ → 𝓔} {A₀ : ℕ → Ω₀ → α} {R₀ : ℕ → Ω₀ → R} variable {P₀ : Measure Ω₀} [IsProbabilityMeasure P₀] /-- The history distribution under any algorithm is absolutely continuous w.r.t. the @@ -279,7 +263,7 @@ lemma absolutelyContinuous_map_hist exact (Measure.AbsolutelyContinuous.compProd_right (show ∀ᵐ e ∂Q, κ_alg e ≪ κ₀ e from by have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : - (ℕ → Fin K × ℝ) → (Iic t → Fin K × ℝ)) = IT.hist t := + (ℕ → α × R) → (Iic t → α × R)) = IT.hist t := funext fun ω => funext fun i => Prod.mk.eta have h_cd : ∀ᵐ e ∂Q, κ_alg e = (condDistrib (fun ω n => (A n ω, R' n ω)) E P e).map (IT.hist t) := by @@ -327,7 +311,7 @@ lemma condDistrib_env_hist_alg_indep -- Key factorization: κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) have h_wd_ae : κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) := by have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : - (ℕ → Fin K × ℝ) → (Iic t → Fin K × ℝ)) = IT.hist t := + (ℕ → α × R) → (Iic t → α × R)) = IT.hist t := funext fun ω => funext fun i => Prod.mk.eta have h_cd : ∀ᵐ e ∂Q, κ_alg e = (condDistrib (fun ω n => (A n ω, R' n ω)) E P e).map (IT.hist t) := by From 129ed0b09689e7da4b82f6450ecd7915edcc894d Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 25 Feb 2026 15:29:14 +0000 Subject: [PATCH 055/155] Refactor HistoryDensity (in progress) --- LeanBandits/ForMathlib/FullSupport.lean | 12 ++- LeanBandits/ForMathlib/WithDensity.lean | 4 + .../BayesStationaryEnv.lean | 10 +++ .../SequentialLearning/HistoryDensity.lean | 82 +++++-------------- 4 files changed, 44 insertions(+), 64 deletions(-) diff --git a/LeanBandits/ForMathlib/FullSupport.lean b/LeanBandits/ForMathlib/FullSupport.lean index bbcaae01..84ae492d 100644 --- a/LeanBandits/ForMathlib/FullSupport.lean +++ b/LeanBandits/ForMathlib/FullSupport.lean @@ -16,6 +16,8 @@ open MeasureTheory ProbabilityTheory variable {α β : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {μ ν : Measure α} +namespace Measure + /-- Any measure is absolutely continuous wrt any measure giving positive mass to all singletons. -/ lemma absolutelyContinuous_of_forall_singleton_pos (hν : ∀ a : α, ν {a} > 0) : μ ≪ ν := by intro s hs @@ -35,11 +37,17 @@ lemma rnDeriv_ne_top_of_forall_singleton_pos [SigmaFinite μ] (hν : ∀ a, ν {a} > 0) (a : α) : μ.rnDeriv ν a ≠ ⊤ := (forall_of_ae_of_forall_singleton_pos hν (Measure.rnDeriv_lt_top μ ν) a).ne +end Measure + +namespace Kernel + /-- Kernel `rnDeriv` is pointwise finite when the reference kernel has full support on singletons. -/ -lemma kernel_rnDeriv_ne_top_of_forall_singleton_pos +lemma rnDeriv_ne_top_of_forall_singleton_pos [MeasurableSpace.CountableOrCountablyGenerated α β] {κ η : Kernel α β} [IsFiniteKernel κ] [IsFiniteKernel η] (hη : ∀ a b, η a {b} > 0) (a : α) (b : β) : Kernel.rnDeriv κ η a b ≠ ⊤ := - (forall_of_ae_of_forall_singleton_pos (hη a) (Kernel.rnDeriv_lt_top κ η) b).ne + (Measure.forall_of_ae_of_forall_singleton_pos (hη a) (Kernel.rnDeriv_lt_top κ η) b).ne + +end Kernel diff --git a/LeanBandits/ForMathlib/WithDensity.lean b/LeanBandits/ForMathlib/WithDensity.lean index 6c209038..d1917d65 100644 --- a/LeanBandits/ForMathlib/WithDensity.lean +++ b/LeanBandits/ForMathlib/WithDensity.lean @@ -19,6 +19,8 @@ open scoped ENNReal variable {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} {μ : Measure α} +namespace Measure + /-- Composing `withDensity` on the measure side of a `compProd`: `(μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst)`. -/ lemma withDensity_compProd_left [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] @@ -93,3 +95,5 @@ lemma withDensity_compProd_withDensity [SFinite μ] (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst * Function.uncurry g) := by rw [Measure.compProd_withDensity hg, withDensity_compProd_left hf] exact (withDensity_mul _ (hf.comp measurable_fst) hg).symm + +end Measure diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index eb301e75..840052de 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -121,6 +121,16 @@ lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : rw [← Kernel.map_apply _ (IT.measurable_action 0), ← hc, show IT.action 0 ∘ trajectory A R' = A 0 from rfl, hcd, Kernel.const_apply]⟩ +lemma hasLaw_IT_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : + ∀ᵐ e ∂Q, HasLaw (IT.hist n) (condDistrib (IsAlgEnvSeq.hist A R' n) E P e) + (condDistrib (trajectory A R') E P e) := by + rw [← h.hasLaw_env.map_eq, show IsAlgEnvSeq.hist A R' n = IT.hist n ∘ trajectory A R' from rfl] + filter_upwards [condDistrib_comp E + (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable + (IT.measurable_hist n)] with e he + exact ⟨(IT.measurable_hist n).aemeasurable, by + rw [← Kernel.map_apply _ (IT.measurable_hist n), he]⟩ + lemma hasCondDistrib_IT_reward_zero [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : ∀ᵐ e ∂Q, HasCondDistrib (IT.reward 0) (IT.action 0) (κ.sectR e) (condDistrib (trajectory A R') E P e) := by diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 7c9be229..300144e0 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -50,10 +50,10 @@ lemma historyDensity_ne_top (alg alg₀ : Algorithm α R) (hpos : alg₀.IsPositive) (t : ℕ) (h : Iic t → α × R) : historyDensity alg alg₀ t h ≠ ⊤ := by induction t with - | zero => exact rnDeriv_ne_top_of_forall_singleton_pos hpos.1 _ + | zero => exact Measure.rnDeriv_ne_top_of_forall_singleton_pos hpos.1 _ | succ n ih => exact ENNReal.mul_ne_top (ih _) - (kernel_rnDeriv_ne_top_of_forall_singleton_pos + (Kernel.rnDeriv_ne_top_of_forall_singleton_pos (fun h' a => hpos.2 n h' a) _ _) end HistoryDensity @@ -75,7 +75,7 @@ lemma Algorithm.IsPositive.absolutelyContinuous_stepKernel_stationary simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] rw [h1, h2] exact Measure.AbsolutelyContinuous.compProd_left - (absolutelyContinuous_of_forall_singleton_pos (hpos.2 n h)) _ + (Measure.absolutelyContinuous_of_forall_singleton_pos (hpos.2 n h)) _ namespace IsAlgEnvSeq @@ -146,7 +146,7 @@ lemma absolutelyContinuous_map_hist_stationary h.hasLaw_step_zero.map_eq, h₀.hasLaw_step_zero.map_eq] simp only [stationaryEnv_ν0] exact (Measure.AbsolutelyContinuous.compProd_left - (absolutelyContinuous_of_forall_singleton_pos hpos.1) _).map + (Measure.absolutelyContinuous_of_forall_singleton_pos hpos.1) _).map e.symm.measurable | succ n ih => rw [h.map_hist_succ_eq_compProd_map, h₀.map_hist_succ_eq_compProd_map] @@ -167,7 +167,7 @@ lemma map_hist_eq_withDensity_historyDensity | zero => set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => α × R) have h_ac : alg.p0 ≪ alg₀.p0 := - absolutelyContinuous_of_forall_singleton_pos hpos.1 + Measure.absolutelyContinuous_of_forall_singleton_pos hpos.1 have h_hist : IsAlgEnvSeq.hist A R' 0 = e.symm ∘ IsAlgEnvSeq.step A R' 0 := by funext ω ⟨i, hi⟩ have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl @@ -182,8 +182,8 @@ lemma map_hist_eq_withDensity_historyDensity h.hasLaw_step_zero.map_eq, h₀.hasLaw_step_zero.map_eq] simp only [stationaryEnv_ν0] conv_lhs => rw [← Measure.withDensity_rnDeriv_eq _ _ h_ac] - rw [withDensity_compProd_left (Measure.measurable_rnDeriv _ _)] - exact withDensity_map_equiv_symm + rw [Measure.withDensity_compProd_left (Measure.measurable_rnDeriv _ _)] + exact Measure.withDensity_map_equiv_symm ((Measure.measurable_rnDeriv _ _).comp measurable_fst) | succ n ih => let σ : (Iic n → α × R) → (α × R) → ℝ≥0∞ := @@ -207,20 +207,20 @@ lemma map_hist_eq_withDensity_historyDensity (Kernel.rnDeriv (alg.policy n) (alg₀.policy n) x) = alg.policy n x := by rw [← Kernel.withDensity_apply _ (Kernel.measurable_rnDeriv _ _)] exact Kernel.withDensity_rnDeriv_eq (κ := alg.policy n) (η := alg₀.policy n) - (absolutelyContinuous_of_forall_singleton_pos (hpos.2 n x)) + (Measure.absolutelyContinuous_of_forall_singleton_pos (hpos.2 n x)) rw [h_alg, h_alg₀, ← h_wd] haveI : SFinite ((alg₀.policy n x).withDensity (Kernel.rnDeriv (alg.policy n) (alg₀.policy n) x)) := by rw [h_wd]; infer_instance - exact withDensity_compProd_left + exact Measure.withDensity_compProd_left (Kernel.measurable_rnDeriv (alg.policy n) (alg₀.policy n)).of_uncurry_left haveI : IsSFiniteKernel ((stepKernel alg₀ (stationaryEnv ν) n).withDensity σ) := by rw [← h_step]; infer_instance rw [h.map_hist_succ_eq_compProd_map n, h₀.map_hist_succ_eq_compProd_map n, ih, h_step, - withDensity_compProd_withDensity (measurable_historyDensity alg alg₀ n) hσ_meas] - exact withDensity_map_equiv_symm + Measure.withDensity_compProd_withDensity (measurable_historyDensity alg alg₀ n) hσ_meas] + exact Measure.withDensity_map_equiv_symm (((measurable_historyDensity alg alg₀ n).comp measurable_fst).mul hσ_meas) end IsAlgEnvSeq @@ -256,35 +256,14 @@ lemma absolutelyContinuous_map_hist set κ₀ := condDistrib (IsAlgEnvSeq.hist A₀ R₀ t) E₀ P₀ rw [h.map_hist_eq_condDistrib_comp t, h₀.map_hist_eq_condDistrib_comp t, ← Measure.snd_compProd, ← Measure.snd_compProd] - have hW_meas : Measurable (fun (ω : Ω) (n : ℕ) => (A n ω, R' n ω)) := - measurable_pi_lambda _ fun n => (h.measurable_A n).prodMk (h.measurable_R n) - have hW₀_meas : Measurable (fun (ω : Ω₀) (n : ℕ) => (A₀ n ω, R₀ n ω)) := - measurable_pi_lambda _ fun n => (h₀.measurable_A n).prodMk (h₀.measurable_R n) exact (Measure.AbsolutelyContinuous.compProd_right (show ∀ᵐ e ∂Q, κ_alg e ≪ κ₀ e from by have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : (ℕ → α × R) → (Iic t → α × R)) = IT.hist t := funext fun ω => funext fun i => Prod.mk.eta - have h_cd : ∀ᵐ e ∂Q, κ_alg e = - (condDistrib (fun ω n => (A n ω, R' n ω)) E P e).map (IT.hist t) := by - rw [← h.hasLaw_env.map_eq] - have h_comp : κ_alg - =ᵐ[P.map E] (condDistrib (fun ω n => (A n ω, R' n ω)) E P).map (IT.hist t) := - condDistrib_comp E hW_meas.aemeasurable (IT.measurable_hist t) - filter_upwards [h_comp] with e he - rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have h_cd₀ : ∀ᵐ e ∂Q, κ₀ e = - (condDistrib (fun ω n => (A₀ n ω, R₀ n ω)) E₀ P₀ e).map (IT.hist t) := by - rw [← h₀.hasLaw_env.map_eq] - have h_comp : κ₀ - =ᵐ[P₀.map E₀] (condDistrib (fun ω n => (A₀ n ω, R₀ n ω)) E₀ P₀).map (IT.hist t) := - condDistrib_comp E₀ hW₀_meas.aemeasurable (IT.measurable_hist t) - filter_upwards [h_comp] with e he - rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have hae := h.ae_IsAlgEnvSeq - have hae₀ := h₀.ae_IsAlgEnvSeq - filter_upwards [h_cd, h_cd₀, hae, hae₀] with e he he₀ hae hae₀ - rw [he, he₀, ← h_IT_hist] + filter_upwards [h.hasLaw_IT_hist t, h₀.hasLaw_IT_hist t, + h.ae_IsAlgEnvSeq, h₀.ae_IsAlgEnvSeq] with e he he₀ hae hae₀ + rw [← he.map_eq, ← he₀.map_eq, ← h_IT_hist] exact hae.absolutelyContinuous_map_hist_stationary hpos hae₀ t)).map measurable_snd @@ -304,37 +283,16 @@ lemma condDistrib_env_hist_alg_indep set ρ := historyDensity alg alg₀ t have hρ_meas := measurable_historyDensity alg alg₀ t have hρ_ne_top := historyDensity_ne_top alg alg₀ hpos t - have hW_meas : Measurable (fun (ω : Ω) (n : ℕ) => (A n ω, R' n ω)) := - measurable_pi_lambda _ fun n => (h.measurable_A n).prodMk (h.measurable_R n) - have hW₀_meas : Measurable (fun (ω : Ω₀) (n : ℕ) => (A₀ n ω, R₀ n ω)) := - measurable_pi_lambda _ fun n => (h₀.measurable_A n).prodMk (h₀.measurable_R n) -- Key factorization: κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) have h_wd_ae : κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) := by have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : (ℕ → α × R) → (Iic t → α × R)) = IT.hist t := funext fun ω => funext fun i => Prod.mk.eta - have h_cd : ∀ᵐ e ∂Q, κ_alg e = - (condDistrib (fun ω n => (A n ω, R' n ω)) E P e).map (IT.hist t) := by - rw [← h.hasLaw_env.map_eq] - have h_comp : κ_alg - =ᵐ[P.map E] (condDistrib (fun ω n => (A n ω, R' n ω)) E P).map (IT.hist t) := - condDistrib_comp E hW_meas.aemeasurable (IT.measurable_hist t) - filter_upwards [h_comp] with e he - rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have h_cd₀ : ∀ᵐ e ∂Q, κ₀ e = - (condDistrib (fun ω n => (A₀ n ω, R₀ n ω)) E₀ P₀ e).map (IT.hist t) := by - rw [← h₀.hasLaw_env.map_eq] - have h_comp : κ₀ - =ᵐ[P₀.map E₀] (condDistrib (fun ω n => (A₀ n ω, R₀ n ω)) E₀ P₀).map (IT.hist t) := - condDistrib_comp E₀ hW₀_meas.aemeasurable (IT.measurable_hist t) - filter_upwards [h_comp] with e he - rw [he, Kernel.map_apply _ (IT.measurable_hist t)] - have hae := h.ae_IsAlgEnvSeq - have hae₀ := h₀.ae_IsAlgEnvSeq - filter_upwards [h_cd, h_cd₀, hae, hae₀] with e he he₀ hae hae₀ + filter_upwards [h.hasLaw_IT_hist t, h₀.hasLaw_IT_hist t, + h.ae_IsAlgEnvSeq, h₀.ae_IsAlgEnvSeq] with e he he₀ hae hae₀ rw [Kernel.withDensity_apply _ (show Measurable (Function.uncurry (fun (_ : 𝓔) => ρ)) from hρ_meas.comp measurable_snd), - he, he₀, ← h_IT_hist] + ← he.map_eq, ← he₀.map_eq, ← h_IT_hist] exact hae.map_hist_eq_withDensity_historyDensity hpos t hae₀ haveI : IsSFiniteKernel (κ₀.withDensity (fun _ => ρ)) := Kernel.IsSFiniteKernel.withDensity _ (fun _ b => hρ_ne_top b) @@ -361,7 +319,7 @@ lemma condDistrib_env_hist_alg_indep rw [h_marg, h_marg₀, Measure.compProd_congr h_wd_ae, Measure.compProd_withDensity (show Measurable (Function.uncurry (fun (_ : 𝓔) => ρ)) from hρ_meas.comp measurable_snd)] - exact map_withDensity_comp measurable_snd hρ_meas + exact Measure.map_withDensity_comp measurable_snd hρ_meas have h_swap : P.map (fun ω => (IsAlgEnvSeq.hist A R' t ω, E ω)) = P.map (IsAlgEnvSeq.hist A R' t) ⊗ₘ condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀ := by have h_uncurry_meas : Measurable (Function.uncurry (fun (_ : 𝓔) => ρ)) := @@ -376,7 +334,7 @@ lemma condDistrib_env_hist_alg_indep _ = ((Q ⊗ₘ κ₀).withDensity (ρ ∘ Prod.snd)).map Prod.swap := by congr 1; exact Measure.compProd_withDensity h_uncurry_meas _ = ((Q ⊗ₘ κ₀).map Prod.swap).withDensity (ρ ∘ Prod.fst) := - map_swap_withDensity_fst hρ_meas + Measure.map_swap_withDensity_fst hρ_meas _ = (P₀.map (fun ω => (IsAlgEnvSeq.hist A₀ R₀ t ω, E₀ ω))).withDensity (ρ ∘ Prod.fst) := by congr 1; rw [← h_joint₀] @@ -388,7 +346,7 @@ lemma condDistrib_env_hist_alg_indep rw [← compProd_map_condDistrib h₀.measurable_E.aemeasurable] _ = (P₀.map (IsAlgEnvSeq.hist A₀ R₀ t)).withDensity ρ ⊗ₘ condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀ := - (withDensity_compProd_left hρ_meas).symm + (Measure.withDensity_compProd_left hρ_meas).symm _ = P.map (IsAlgEnvSeq.hist A R' t) ⊗ₘ condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀ := by rw [h_hist] From 4a60bf1bab750af0c5f3ab8088e5706e4a334722 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 27 Feb 2026 13:50:59 +0000 Subject: [PATCH 056/155] Refactor HistoryDensity (in progress) --- LeanBandits/ForMathlib/FullSupport.lean | 13 ++ .../SequentialLearning/HistoryDensity.lean | 131 +++++------------- 2 files changed, 51 insertions(+), 93 deletions(-) diff --git a/LeanBandits/ForMathlib/FullSupport.lean b/LeanBandits/ForMathlib/FullSupport.lean index 84ae492d..48fed438 100644 --- a/LeanBandits/ForMathlib/FullSupport.lean +++ b/LeanBandits/ForMathlib/FullSupport.lean @@ -4,6 +4,7 @@ Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne, Paulo Rauber -/ import Mathlib.Probability.Kernel.RadonNikodym +import Mathlib.Probability.Kernel.Composition.MeasureCompProd /-! # Absolute continuity and rnDeriv finiteness from full support @@ -51,3 +52,15 @@ lemma rnDeriv_ne_top_of_forall_singleton_pos (Measure.forall_of_ae_of_forall_singleton_pos (hη a) (Kernel.rnDeriv_lt_top κ η) b).ne end Kernel + +variable {γ : Type*} {mγ : MeasurableSpace γ} + +namespace Measure.AbsolutelyContinuous + +/-- If `κ a` is absolutely continuous wrt `η a`, then so is the kernel compProd at `a`. -/ +lemma kernel_compProd_left {κ η : Kernel α β} [IsSFiniteKernel κ] [IsSFiniteKernel η] + {ξ : Kernel (α × β) γ} [IsSFiniteKernel ξ] {a : α} (hac : κ a ≪ η a) : + (κ ⊗ₖ ξ) a ≪ (η ⊗ₖ ξ) a := by + simp_rw [Kernel.compProd_apply_eq_compProd_sectR, hac.compProd_left _] + +end Measure.AbsolutelyContinuous diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 300144e0..31785d3c 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -7,118 +7,62 @@ import LeanBandits.ForMathlib.FullSupport import LeanBandits.ForMathlib.WithDensity import LeanBandits.SequentialLearning.BayesStationaryEnv -open MeasureTheory ProbabilityTheory Finset Preorder +open MeasureTheory ProbabilityTheory Finset -open scoped ENNReal NNReal +open scoped ENNReal namespace Learning -variable {α : Type*} {R : Type*} [MeasurableSpace α] [MeasurableSpace R] +variable {α R Ω : Type*} [MeasurableSpace α] [MeasurableSpace R] [MeasurableSpace Ω] -section HistoryDensity - -variable [MeasurableSpace.CountablyGenerated α] - -/-- The density of the history distribution under `alg` w.r.t. a positive reference algorithm. -This density depends only on the algorithm's action probabilities, not on the reward kernel. -/ -noncomputable def historyDensity - (alg alg₀ : Algorithm α R) : - (t : ℕ) → (Iic t → α × R) → ℝ≥0∞ - | 0 => (alg.p0.rnDeriv alg₀.p0 ∘ Prod.fst) ∘ - MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => α × R) - | n + 1 => - let σ : (Iic n → α × R) → (α × R) → ℝ≥0∞ := - fun h ar => Kernel.rnDeriv (alg.policy n) - (alg₀.policy n) h ar.1 - (historyDensity alg alg₀ n ∘ Prod.fst * Function.uncurry σ) ∘ - MeasurableEquiv.IicSuccProd (fun _ : ℕ => α × R) n +noncomputable +def historyDensity [MeasurableSpace.CountablyGenerated α] (alg alg₀ : Algorithm α R) : + (n : ℕ) → (Iic n → α × R) → ℝ≥0∞ + | 0, h => (alg.p0.rnDeriv alg₀.p0 (h ⟨0, by simp⟩).1) + | n + 1, h => let p := MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n h + historyDensity alg alg₀ n p.1 * (alg.policy n).rnDeriv (alg₀.policy n) p.1 p.2.1 @[fun_prop] -lemma measurable_historyDensity (alg alg₀ : Algorithm α R) (t : ℕ) : - Measurable (historyDensity alg alg₀ t) := by +lemma measurable_historyDensity [MeasurableSpace.CountablyGenerated α] (alg alg₀ : Algorithm α R) + (t : ℕ) : Measurable (historyDensity alg alg₀ t) := by induction t with - | zero => - exact (Measure.measurable_rnDeriv _ _).comp - (measurable_fst.comp (MeasurableEquiv.piUnique _).measurable) - | succ n ih => - exact ((ih.comp measurable_fst).mul - ((Kernel.measurable_rnDeriv _ _).comp - (measurable_fst.prodMk (measurable_fst.comp measurable_snd)))).comp - (MeasurableEquiv.IicSuccProd _ n).measurable + | zero => simp_rw [historyDensity]; fun_prop + | succ n ih => simp_rw [historyDensity]; fun_prop -lemma historyDensity_ne_top (alg alg₀ : Algorithm α R) - (hpos : alg₀.IsPositive) (t : ℕ) - (h : Iic t → α × R) : historyDensity alg alg₀ t h ≠ ⊤ := by - induction t with - | zero => exact Measure.rnDeriv_ne_top_of_forall_singleton_pos hpos.1 _ +lemma Algorithm.IsPositive.historyDensity_ne_top [MeasurableSpace.CountablyGenerated α] + {alg₀ : Algorithm α R} (hp : alg₀.IsPositive) (alg : Algorithm α R) (n : ℕ) + (h : Iic n → α × R) : historyDensity alg alg₀ n h ≠ ⊤ := by + induction n with + | zero => exact Measure.rnDeriv_ne_top_of_forall_singleton_pos hp.1 _ | succ n ih => - exact ENNReal.mul_ne_top (ih _) - (Kernel.rnDeriv_ne_top_of_forall_singleton_pos - (fun h' a => hpos.2 n h' a) _ _) - -end HistoryDensity - -/-- The step kernel for a stationary environment under a positive algorithm absolutely - continuously dominates any other algorithm's step kernel. -/ -lemma Algorithm.IsPositive.absolutelyContinuous_stepKernel_stationary - {alg₀ : Algorithm α R} (hpos : alg₀.IsPositive) - (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] - (n : ℕ) (h : Iic n → α × R) : - stepKernel alg (stationaryEnv ν) n h ≪ - stepKernel alg₀ (stationaryEnv ν) n h := by - have h1 : stepKernel alg (stationaryEnv ν) n h = (alg.policy n h) ⊗ₘ ν := by - simp only [stepKernel, stationaryEnv]; ext s hs - simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - have h2 : stepKernel alg₀ (stationaryEnv ν) n h = - (alg₀.policy n h) ⊗ₘ ν := by - simp only [stepKernel, stationaryEnv]; ext s hs - simp only [Kernel.compProd_apply hs, Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - rw [h1, h2] - exact Measure.AbsolutelyContinuous.compProd_left - (Measure.absolutelyContinuous_of_forall_singleton_pos (hpos.2 n h)) _ + exact ENNReal.mul_ne_top (ih _) (Kernel.rnDeriv_ne_top_of_forall_singleton_pos (hp.2 n) _ _) namespace IsAlgEnvSeq variable [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] +variable {alg : Algorithm α R} {env : Environment α R} +variable {A : ℕ → Ω → α} {R' : ℕ → Ω → R} +variable {P : Measure Ω} [IsFiniteMeasure P] -/-- The history distribution at time `n + 1` decomposes as a compProd of the history at time `n` - and the step kernel, composed with `IicSuccProd.symm`. -/ -lemma map_hist_succ_eq_compProd_map - {Ω : Type*} [MeasurableSpace Ω] - {A : ℕ → Ω → α} {R' : ℕ → Ω → R} - {alg : Algorithm α R} {env : Environment α R} - {P : Measure Ω} [IsFiniteMeasure P] - (h : IsAlgEnvSeq A R' alg env P) (n : ℕ) : +-- Algorithm.lean? +lemma map_hist_succ_eq_compProd_map (h : IsAlgEnvSeq A R' alg env P) (n : ℕ) : P.map (IsAlgEnvSeq.hist A R' (n + 1)) = (P.map (IsAlgEnvSeq.hist A R' n) ⊗ₘ stepKernel alg env n).map - (MeasurableEquiv.IicSuccProd (fun _ : ℕ => α × R) n).symm := by - set e := MeasurableEquiv.IicSuccProd (fun _ : ℕ => α × R) n - have hA := h.measurable_A; have hR := h.measurable_R - have h_func : IsAlgEnvSeq.hist A R' (n + 1) = e.symm ∘ - (fun ω => (IsAlgEnvSeq.hist A R' n ω, IsAlgEnvSeq.step A R' (n + 1) ω)) := by - funext ω; simp only [Function.comp_apply] - change frestrictLe (n + 1) (fun k => IsAlgEnvSeq.step A R' k ω) = - e.symm (frestrictLe n (fun k => IsAlgEnvSeq.step A R' k ω), - IsAlgEnvSeq.step A R' (n + 1) ω) - change frestrictLe (n + 1) (fun k => IsAlgEnvSeq.step A R' k ω) = - e.symm (e (frestrictLe (n + 1) (fun k => IsAlgEnvSeq.step A R' k ω))) - rw [e.symm_apply_apply] - rw [h_func, (Measure.map_map e.symm.measurable - ((IsAlgEnvSeq.measurable_hist hA hR n).prodMk - (IsAlgEnvSeq.measurable_step (n + 1) (hA _) (hR _)))).symm] + (MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n).symm := by + set e := (MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n).symm + have hf : IsAlgEnvSeq.hist A R' (n + 1) = e ∘ + (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, IsAlgEnvSeq.step A R' (n + 1) ω)) := + funext fun _ ↦ (e.apply_symm_apply _).symm + have hA := h.measurable_A + have hR := h.measurable_R + rw [hf, ← Measure.map_map e.measurable (by fun_prop)] congr 1 - have h_cd := h.hasCondDistrib_step n - exact ((condDistrib_ae_eq_iff_measure_eq_compProd _ - (IsAlgEnvSeq.measurable_step (n + 1) (hA _) (hR _)).aemeasurable - (stepKernel alg env n)).mp h_cd.condDistrib_eq) + apply (condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop) (stepKernel alg env n)).1 + exact (h.hasCondDistrib_step n).condDistrib_eq variable {ν : Kernel α R} [IsMarkovKernel ν] -variable {Ω : Type*} [MeasurableSpace Ω] -variable {A : ℕ → Ω → α} {R' : ℕ → Ω → R} -variable {alg : Algorithm α R} -variable {P : Measure Ω} [IsProbabilityMeasure P] -variable {alg₀ : Algorithm α R} variable {Ω₀ : Type*} [MeasurableSpace Ω₀] +variable {alg₀ : Algorithm α R} variable {A₀ : ℕ → Ω₀ → α} {R₀ : ℕ → Ω₀ → R} variable {P₀ : Measure Ω₀} [IsProbabilityMeasure P₀] @@ -152,7 +96,8 @@ lemma absolutelyContinuous_map_hist_stationary rw [h.map_hist_succ_eq_compProd_map, h₀.map_hist_succ_eq_compProd_map] exact (Measure.AbsolutelyContinuous.compProd ih (Filter.Eventually.of_forall fun x => - hpos.absolutelyContinuous_stepKernel_stationary alg ν n x)).map + Measure.AbsolutelyContinuous.kernel_compProd_left + (Measure.absolutelyContinuous_of_forall_singleton_pos (hpos.2 n x)))).map (MeasurableEquiv.IicSuccProd _ n).symm.measurable /-- The history distribution under any algorithm equals the positive reference algorithm's history @@ -282,7 +227,7 @@ lemma condDistrib_env_hist_alg_indep set κ₀ := condDistrib (IsAlgEnvSeq.hist A₀ R₀ t) E₀ P₀ set ρ := historyDensity alg alg₀ t have hρ_meas := measurable_historyDensity alg alg₀ t - have hρ_ne_top := historyDensity_ne_top alg alg₀ hpos t + have hρ_ne_top := hpos.historyDensity_ne_top alg t -- Key factorization: κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) have h_wd_ae : κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) := by have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : From 3d014179aaa532f12b05f39ca2f7cf3901296b20 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 3 Mar 2026 14:49:17 +0000 Subject: [PATCH 057/155] Refactor HistoryDensity (in progress) --- LeanBandits/ForMathlib/WithDensity.lean | 74 +++++-- .../SequentialLearning/HistoryDensity.lean | 203 +++++++----------- 2 files changed, 135 insertions(+), 142 deletions(-) diff --git a/LeanBandits/ForMathlib/WithDensity.lean b/LeanBandits/ForMathlib/WithDensity.lean index d1917d65..174297dc 100644 --- a/LeanBandits/ForMathlib/WithDensity.lean +++ b/LeanBandits/ForMathlib/WithDensity.lean @@ -40,17 +40,17 @@ lemma withDensity_compProd_left [SFinite μ] {κ : Kernel α β} [IsSFiniteKerne ext b; simp only [Function.comp, Set.indicator, Set.mem_preimage]; rfl rw [this, lintegral_indicator_const (hs.preimage (by fun_prop))] -/-- Mapping a `withDensity` through `MeasurableEquiv.symm`: -`(μ.withDensity f).map e.symm = (μ.map e.symm).withDensity (f ∘ e)`. -/ -lemma withDensity_map_equiv_symm - {μ : Measure β} {e : α ≃ᵐ β} {f : β → ℝ≥0∞} (hf : Measurable f) : - (μ.withDensity f).map e.symm = (μ.map e.symm).withDensity (f ∘ e) := by +/-- Pushing a `withDensity` through a `MeasurableEquiv`: +`(μ.withDensity f).map e = (μ.map e).withDensity (f ∘ e.symm)`. -/ +lemma withDensity_map_equiv + {e : α ≃ᵐ β} {f : α → ℝ≥0∞} (hf : Measurable f) : + (μ.withDensity f).map e = (μ.map e).withDensity (f ∘ e.symm) := by ext s hs - rw [Measure.map_apply e.symm.measurable hs, - withDensity_apply _ (e.symm.measurable hs), - withDensity_apply _ hs, Measure.restrict_map e.symm.measurable hs, - lintegral_map (hf.comp e.measurable) e.symm.measurable] - simp_rw [Function.comp_apply, e.apply_symm_apply] + rw [Measure.map_apply e.measurable hs, + withDensity_apply _ (e.measurable hs), + withDensity_apply _ hs, Measure.restrict_map e.measurable hs, + lintegral_map (hf.comp e.symm.measurable) e.measurable] + simp_rw [Function.comp_apply, e.symm_apply_apply] /-- Mapping a `withDensity` through a `MeasurableEquiv` from the snd component. -/ lemma map_swap_withDensity_fst @@ -72,6 +72,21 @@ lemma map_withDensity_comp simp only [Measure.map_apply hg hs, withDensity_apply _ (hg hs), withDensity_apply _ hs, setLIntegral_map hs hf hg, Function.comp] +/-- `(μ.withDensity f) ⊗ₘ (κ.withDensity g) = (μ ⊗ₘ κ).withDensity (f ∘ fst * uncurry g)`. -/ +lemma withDensity_compProd_withDensity [SFinite μ] + {κ : Kernel α γ} [IsSFiniteKernel κ] + {f : α → ℝ≥0∞} {g : α → γ → ℝ≥0∞} + (hf : Measurable f) (hg : Measurable (Function.uncurry g)) + [IsSFiniteKernel (κ.withDensity g)] : + (μ.withDensity f) ⊗ₘ (κ.withDensity g) = + (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst * Function.uncurry g) := by + rw [Measure.compProd_withDensity hg, withDensity_compProd_left hf] + exact (withDensity_mul _ (hf.comp measurable_fst) hg).symm + +end Measure + +namespace ProbabilityTheory.Kernel + /-- `(κ.withDensity (fun _ => f)) ∘ₘ μ = (κ ∘ₘ μ).withDensity f`. -/ lemma comp_withDensity_const [SFinite μ] @@ -83,17 +98,32 @@ lemma comp_withDensity_const Measure.compProd_withDensity (show Measurable (Function.uncurry (fun (_ : α) => f)) from hf.comp measurable_snd), ← Measure.snd_compProd μ κ, Measure.snd, Measure.snd] - exact map_withDensity_comp measurable_snd hf + exact Measure.map_withDensity_comp measurable_snd hf -/-- `(μ.withDensity f) ⊗ₘ (κ.withDensity g) = (μ ⊗ₘ κ).withDensity (f ∘ fst * uncurry g)`. -/ -lemma withDensity_compProd_withDensity [SFinite μ] - {κ : Kernel α γ} [IsSFiniteKernel κ] - {f : α → ℝ≥0∞} {g : α → γ → ℝ≥0∞} - (hf : Measurable f) (hg : Measurable (Function.uncurry g)) - [IsSFiniteKernel (κ.withDensity g)] : - (μ.withDensity f) ⊗ₘ (κ.withDensity g) = - (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst * Function.uncurry g) := by - rw [Measure.compProd_withDensity hg, withDensity_compProd_left hf] - exact (withDensity_mul _ (hf.comp measurable_fst) hg).symm +/-- Composing `Kernel.withDensity` on the left kernel of `Kernel.compProd`: +`(κ.withDensity f) ⊗ₖ η = (κ ⊗ₖ η).withDensity (fun a (b, _) => f a b)`. -/ +lemma withDensity_compProd_left + {κ : Kernel α β} {η : Kernel (α × β) γ} {f : α → β → ℝ≥0∞} + [IsSFiniteKernel κ] [IsSFiniteKernel η] [IsSFiniteKernel (κ.withDensity f)] + (hf : Measurable (Function.uncurry f)) : + (κ.withDensity f) ⊗ₖ η = + (κ ⊗ₖ η).withDensity (fun a (b, _) ↦ f a b) := by + have hg : Measurable (Function.uncurry (fun a (bc : β × γ) => f a bc.1)) := + hf.comp (measurable_fst.prodMk (measurable_fst.comp measurable_snd)) + ext x : 1 + haveI : SFinite ((κ x).withDensity (f x)) := by + rw [← Kernel.withDensity_apply _ hf]; infer_instance + simp only [Kernel.compProd_apply_eq_compProd_sectR, Kernel.withDensity_apply _ hf, + Kernel.withDensity_apply _ hg] + exact Measure.withDensity_compProd_left hf.of_uncurry_left -end Measure +/-- If `κ a ≪ η a` for all `a`, then `η.withDensity (κ.rnDeriv η) = κ`. -/ +lemma withDensity_rnDeriv_eq' {κ η : Kernel α β} + [MeasurableSpace.CountableOrCountablyGenerated α β] + [IsFiniteKernel κ] [IsFiniteKernel η] + (h : ∀ a, κ a ≪ η a) : + η.withDensity (κ.rnDeriv η) = κ := by + ext a : 1 + exact Kernel.withDensity_rnDeriv_eq (h a) + +end ProbabilityTheory.Kernel diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 31785d3c..e12740bd 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -13,21 +13,26 @@ open scoped ENNReal namespace Learning -variable {α R Ω : Type*} [MeasurableSpace α] [MeasurableSpace R] [MeasurableSpace Ω] +variable {α R : Type*} [MeasurableSpace α] [MeasurableSpace R] noncomputable def historyDensity [MeasurableSpace.CountablyGenerated α] (alg alg₀ : Algorithm α R) : (n : ℕ) → (Iic n → α × R) → ℝ≥0∞ | 0, h => (alg.p0.rnDeriv alg₀.p0 (h ⟨0, by simp⟩).1) - | n + 1, h => let p := MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n h + | n + 1, h => + let p := MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n h historyDensity alg alg₀ n p.1 * (alg.policy n).rnDeriv (alg₀.policy n) p.1 p.2.1 @[fun_prop] lemma measurable_historyDensity [MeasurableSpace.CountablyGenerated α] (alg alg₀ : Algorithm α R) - (t : ℕ) : Measurable (historyDensity alg alg₀ t) := by - induction t with - | zero => simp_rw [historyDensity]; fun_prop - | succ n ih => simp_rw [historyDensity]; fun_prop + (n : ℕ) : Measurable (historyDensity alg alg₀ n) := by + induction n with + | zero => + simp_rw [historyDensity] + fun_prop + | succ n ih => + simp_rw [historyDensity] + fun_prop lemma Algorithm.IsPositive.historyDensity_ne_top [MeasurableSpace.CountablyGenerated α] {alg₀ : Algorithm α R} (hp : alg₀.IsPositive) (alg : Algorithm α R) (n : ℕ) @@ -39,134 +44,92 @@ lemma Algorithm.IsPositive.historyDensity_ne_top [MeasurableSpace.CountablyGener namespace IsAlgEnvSeq +variable {Ω : Type*} [MeasurableSpace Ω] variable [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] variable {alg : Algorithm α R} {env : Environment α R} variable {A : ℕ → Ω → α} {R' : ℕ → Ω → R} variable {P : Measure Ω} [IsFiniteMeasure P] --- Algorithm.lean? -lemma map_hist_succ_eq_compProd_map (h : IsAlgEnvSeq A R' alg env P) (n : ℕ) : - P.map (IsAlgEnvSeq.hist A R' (n + 1)) = - (P.map (IsAlgEnvSeq.hist A R' n) ⊗ₘ stepKernel alg env n).map - (MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n).symm := by - set e := (MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n).symm - have hf : IsAlgEnvSeq.hist A R' (n + 1) = e ∘ - (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, IsAlgEnvSeq.step A R' (n + 1) ω)) := - funext fun _ ↦ (e.apply_symm_apply _).symm - have hA := h.measurable_A - have hR := h.measurable_R - rw [hf, ← Measure.map_map e.measurable (by fun_prop)] - congr 1 - apply (condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop) (stepKernel alg env n)).1 - exact (h.hasCondDistrib_step n).condDistrib_eq +lemma hasLaw_hist_zero (h : IsAlgEnvSeq A R' alg env P) : HasLaw (hist A R' 0) + ((P.map (step A R' 0)).map + (MeasurableEquiv.piUnique (fun _ : Iic 0 ↦ α × R)).symm) P where + aemeasurable := (measurable_hist h.measurable_A h.measurable_R 0).aemeasurable + map_eq := by + have he : (MeasurableEquiv.piUnique (fun _ : Iic 0 ↦ α × R)).symm ∘ step A R' 0 = + hist A R' 0 := by + funext _ ⟨0, _⟩ + rfl + rw [← he] + have hA := h.measurable_A + have hR := h.measurable_R + exact (Measure.map_map (by fun_prop) (by fun_prop)).symm + +lemma hasLaw_hist_succ (h : IsAlgEnvSeq A R' alg env P) (n : ℕ) : HasLaw (hist A R' (n + 1)) + ((P.map (hist A R' n) ⊗ₘ condDistrib (step A R' (n + 1)) (hist A R' n) P).map + (MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n).symm) P where + aemeasurable := (measurable_hist h.measurable_A h.measurable_R (n + 1)).aemeasurable + map_eq := by + have he : (MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n).symm ∘ + (fun ω ↦ (hist A R' n ω, step A R' (n + 1) ω)) = hist A R' (n + 1) := by + funext ω + exact (MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n).symm_apply_apply (hist A R' (n + 1) ω) + have hA := h.measurable_A + have hR := h.measurable_R + rw [← he, ← Measure.map_map (by fun_prop) (by fun_prop)] + congr + exact (compProd_map_condDistrib (by fun_prop)).symm -variable {ν : Kernel α R} [IsMarkovKernel ν] variable {Ω₀ : Type*} [MeasurableSpace Ω₀] variable {alg₀ : Algorithm α R} variable {A₀ : ℕ → Ω₀ → α} {R₀ : ℕ → Ω₀ → R} variable {P₀ : Measure Ω₀} [IsProbabilityMeasure P₀] -/-- The history distribution under any algorithm is absolutely continuous w.r.t. the - history distribution under a positive reference algorithm, - for a stationary environment. -/ -lemma absolutelyContinuous_map_hist_stationary - (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) - (hpos : alg₀.IsPositive) - (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ (stationaryEnv ν) P₀) - (t : ℕ) : - P.map (IsAlgEnvSeq.hist A R' t) ≪ P₀.map (IsAlgEnvSeq.hist A₀ R₀ t) := by - induction t with +lemma absolutelyContinuous_map_hist (h : IsAlgEnvSeq A R' alg env P) + (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ env P₀) (hp : alg₀.IsPositive) (n : ℕ) : + P.map (IsAlgEnvSeq.hist A R' n) ≪ P₀.map (IsAlgEnvSeq.hist A₀ R₀ n) := by + induction n with | zero => - set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => α × R) - have h_hist : IsAlgEnvSeq.hist A R' 0 = e.symm ∘ IsAlgEnvSeq.step A R' 0 := by - funext ω ⟨i, hi⟩; have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl - have h_hist₀ : IsAlgEnvSeq.hist A₀ R₀ 0 = e.symm ∘ IsAlgEnvSeq.step A₀ R₀ 0 := by - funext ω ⟨i, hi⟩; have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl - rw [h_hist, h_hist₀, - ← Measure.map_map e.symm.measurable - (IsAlgEnvSeq.measurable_step 0 (h.measurable_A _) (h.measurable_R _)), - ← Measure.map_map e.symm.measurable - (IsAlgEnvSeq.measurable_step 0 (h₀.measurable_A _) (h₀.measurable_R _)), - h.hasLaw_step_zero.map_eq, h₀.hasLaw_step_zero.map_eq] - simp only [stationaryEnv_ν0] - exact (Measure.AbsolutelyContinuous.compProd_left - (Measure.absolutelyContinuous_of_forall_singleton_pos hpos.1) _).map - e.symm.measurable + rw [h.hasLaw_hist_zero.map_eq, h₀.hasLaw_hist_zero.map_eq] + apply Measure.AbsolutelyContinuous.map _ (by fun_prop) + rw [h.hasLaw_step_zero.map_eq, h₀.hasLaw_step_zero.map_eq] + apply Measure.AbsolutelyContinuous.compProd_left + exact Measure.absolutelyContinuous_of_forall_singleton_pos hp.1 | succ n ih => - rw [h.map_hist_succ_eq_compProd_map, h₀.map_hist_succ_eq_compProd_map] - exact (Measure.AbsolutelyContinuous.compProd ih - (Filter.Eventually.of_forall fun x => - Measure.AbsolutelyContinuous.kernel_compProd_left - (Measure.absolutelyContinuous_of_forall_singleton_pos (hpos.2 n x)))).map - (MeasurableEquiv.IicSuccProd _ n).symm.measurable + rw [(h.hasLaw_hist_succ n).map_eq, (h₀.hasLaw_hist_succ n).map_eq] + apply Measure.AbsolutelyContinuous.map _ (by fun_prop) + rw [Measure.compProd_congr (h.hasCondDistrib_step n).condDistrib_eq, + Measure.compProd_congr (h₀.hasCondDistrib_step n).condDistrib_eq] + apply Measure.AbsolutelyContinuous.compProd ih + filter_upwards with h' + apply Measure.AbsolutelyContinuous.kernel_compProd_left + exact Measure.absolutelyContinuous_of_forall_singleton_pos (hp.2 n h') -/-- The history distribution under any algorithm equals the positive reference algorithm's history -distribution weighted by `historyDensity`, for any stationary environment. -/ -lemma map_hist_eq_withDensity_historyDensity - (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) - (hpos : alg₀.IsPositive) (t : ℕ) - (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ (stationaryEnv ν) P₀) : - P.map (IsAlgEnvSeq.hist A R' t) = - (P₀.map (IsAlgEnvSeq.hist A₀ R₀ t)).withDensity (historyDensity alg alg₀ t) := by - induction t with +lemma map_hist_eq_withDensity_historyDensity (h : IsAlgEnvSeq A R' alg env P) + (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ env P₀) (hp : alg₀.IsPositive) (n : ℕ) : + P.map (IsAlgEnvSeq.hist A R' n) = + (P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity (historyDensity alg alg₀ n) := by + induction n with | zero => - set e := MeasurableEquiv.piUnique (fun _ : Iic (0 : ℕ) => α × R) - have h_ac : alg.p0 ≪ alg₀.p0 := - Measure.absolutelyContinuous_of_forall_singleton_pos hpos.1 - have h_hist : IsAlgEnvSeq.hist A R' 0 = e.symm ∘ IsAlgEnvSeq.step A R' 0 := by - funext ω ⟨i, hi⟩ - have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl - have h_hist₀ : IsAlgEnvSeq.hist A₀ R₀ 0 = e.symm ∘ IsAlgEnvSeq.step A₀ R₀ 0 := by - funext ω ⟨i, hi⟩ - have : i = 0 := Nat.le_zero.mp (Finset.mem_Iic.mp hi); subst this; rfl - rw [h_hist, h_hist₀, - ← Measure.map_map e.symm.measurable - (IsAlgEnvSeq.measurable_step 0 (h.measurable_A _) (h.measurable_R _)), - ← Measure.map_map e.symm.measurable - (IsAlgEnvSeq.measurable_step 0 (h₀.measurable_A _) (h₀.measurable_R _)), - h.hasLaw_step_zero.map_eq, h₀.hasLaw_step_zero.map_eq] - simp only [stationaryEnv_ν0] - conv_lhs => rw [← Measure.withDensity_rnDeriv_eq _ _ h_ac] - rw [Measure.withDensity_compProd_left (Measure.measurable_rnDeriv _ _)] - exact Measure.withDensity_map_equiv_symm - ((Measure.measurable_rnDeriv _ _).comp measurable_fst) + rw [h.hasLaw_hist_zero.map_eq, h₀.hasLaw_hist_zero.map_eq, h.hasLaw_step_zero.map_eq, + h₀.hasLaw_step_zero.map_eq] + have ha : alg.p0 ≪ alg₀.p0 := Measure.absolutelyContinuous_of_forall_singleton_pos hp.1 + rw [← Measure.withDensity_rnDeriv_eq _ _ ha, Measure.withDensity_compProd_left (by fun_prop)] + exact Measure.withDensity_map_equiv (by fun_prop) | succ n ih => - let σ : (Iic n → α × R) → (α × R) → ℝ≥0∞ := - fun x ar => Kernel.rnDeriv (alg.policy n) (alg₀.policy n) x ar.1 - have hσ_meas : Measurable (Function.uncurry σ) := - (Kernel.measurable_rnDeriv _ _).comp - (measurable_fst.prodMk (measurable_fst.comp measurable_snd)) - have h_step : stepKernel alg (stationaryEnv ν) n = - (stepKernel alg₀ (stationaryEnv ν) n).withDensity σ := by - ext x : 1 - rw [Kernel.withDensity_apply _ hσ_meas] - have h_alg : stepKernel alg (stationaryEnv ν) n x = (alg.policy n x) ⊗ₘ ν := by - ext s hs - simp only [stepKernel, stationaryEnv, Kernel.compProd_apply hs, - Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - have h_alg₀ : stepKernel alg₀ (stationaryEnv ν) n x = (alg₀.policy n x) ⊗ₘ ν := by - ext s hs - simp only [stepKernel, stationaryEnv, Kernel.compProd_apply hs, - Measure.compProd_apply hs, Kernel.prodMkLeft_apply] - have h_wd : ((alg₀.policy n) x).withDensity - (Kernel.rnDeriv (alg.policy n) (alg₀.policy n) x) = alg.policy n x := by - rw [← Kernel.withDensity_apply _ (Kernel.measurable_rnDeriv _ _)] - exact Kernel.withDensity_rnDeriv_eq (κ := alg.policy n) (η := alg₀.policy n) - (Measure.absolutelyContinuous_of_forall_singleton_pos (hpos.2 n x)) - rw [h_alg, h_alg₀, ← h_wd] - haveI : SFinite ((alg₀.policy n x).withDensity - (Kernel.rnDeriv (alg.policy n) (alg₀.policy n) x)) := by - rw [h_wd]; infer_instance - exact Measure.withDensity_compProd_left - (Kernel.measurable_rnDeriv (alg.policy n) (alg₀.policy n)).of_uncurry_left - haveI : IsSFiniteKernel ((stepKernel alg₀ (stationaryEnv ν) n).withDensity σ) := by - rw [← h_step]; infer_instance - rw [h.map_hist_succ_eq_compProd_map n, - h₀.map_hist_succ_eq_compProd_map n, - ih, h_step, - Measure.withDensity_compProd_withDensity (measurable_historyDensity alg alg₀ n) hσ_meas] - exact Measure.withDensity_map_equiv_symm - (((measurable_historyDensity alg alg₀ n).comp measurable_fst).mul hσ_meas) + let ρ h' (ar : α × R) := Kernel.rnDeriv (alg.policy n) (alg₀.policy n) h' ar.1 + have hpo h' : alg.policy n h' ≪ alg₀.policy n h' := + Measure.absolutelyContinuous_of_forall_singleton_pos (hp.2 n h') + have hs : stepKernel alg env n = (stepKernel alg₀ env n).withDensity ρ := by + rw [stepKernel, ← Kernel.withDensity_rnDeriv_eq' hpo] + exact Kernel.withDensity_compProd_left (Kernel.measurable_rnDeriv _ _) + have : IsMarkovKernel ((stepKernel alg₀ env n).withDensity ρ) := by + rw [← hs] + infer_instance + rw [(h.hasLaw_hist_succ n).map_eq, (h₀.hasLaw_hist_succ n).map_eq, + Measure.compProd_congr (h.hasCondDistrib_step n).condDistrib_eq, + Measure.compProd_congr (h₀.hasCondDistrib_step n).condDistrib_eq, ih, hs, + Measure.withDensity_compProd_withDensity (by fun_prop) (by fun_prop)] + exact Measure.withDensity_map_equiv (by fun_prop) end IsAlgEnvSeq @@ -209,7 +172,7 @@ lemma absolutelyContinuous_map_hist filter_upwards [h.hasLaw_IT_hist t, h₀.hasLaw_IT_hist t, h.ae_IsAlgEnvSeq, h₀.ae_IsAlgEnvSeq] with e he he₀ hae hae₀ rw [← he.map_eq, ← he₀.map_eq, ← h_IT_hist] - exact hae.absolutelyContinuous_map_hist_stationary hpos hae₀ t)).map + exact hae.absolutelyContinuous_map_hist hae₀ hpos t)).map measurable_snd variable [StandardBorelSpace 𝓔] [Nonempty 𝓔] [IsMarkovKernel κ] @@ -238,7 +201,7 @@ lemma condDistrib_env_hist_alg_indep rw [Kernel.withDensity_apply _ (show Measurable (Function.uncurry (fun (_ : 𝓔) => ρ)) from hρ_meas.comp measurable_snd), ← he.map_eq, ← he₀.map_eq, ← h_IT_hist] - exact hae.map_hist_eq_withDensity_historyDensity hpos t hae₀ + exact hae.map_hist_eq_withDensity_historyDensity hae₀ hpos t haveI : IsSFiniteKernel (κ₀.withDensity (fun _ => ρ)) := Kernel.IsSFiniteKernel.withDensity _ (fun _ b => hρ_ne_top b) -- Direct condDistrib equality via joint measure argument From 8afa7c41976af8632f13f8a0ef102926158084ea Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 3 Mar 2026 14:51:39 +0000 Subject: [PATCH 058/155] Minor --- LeanBandits/ForMathlib/WithDensity.lean | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/LeanBandits/ForMathlib/WithDensity.lean b/LeanBandits/ForMathlib/WithDensity.lean index 174297dc..90487287 100644 --- a/LeanBandits/ForMathlib/WithDensity.lean +++ b/LeanBandits/ForMathlib/WithDensity.lean @@ -72,14 +72,14 @@ lemma map_withDensity_comp simp only [Measure.map_apply hg hs, withDensity_apply _ (hg hs), withDensity_apply _ hs, setLIntegral_map hs hf hg, Function.comp] -/-- `(μ.withDensity f) ⊗ₘ (κ.withDensity g) = (μ ⊗ₘ κ).withDensity (f ∘ fst * uncurry g)`. -/ +/-- `(f · μ) ⊗ₘ (g · κ) = ((a, c) ↦ f a * g a c) · (μ ⊗ₘ κ)`. -/ lemma withDensity_compProd_withDensity [SFinite μ] {κ : Kernel α γ} [IsSFiniteKernel κ] {f : α → ℝ≥0∞} {g : α → γ → ℝ≥0∞} (hf : Measurable f) (hg : Measurable (Function.uncurry g)) [IsSFiniteKernel (κ.withDensity g)] : (μ.withDensity f) ⊗ₘ (κ.withDensity g) = - (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst * Function.uncurry g) := by + (μ ⊗ₘ κ).withDensity (fun (a, c) => f a * g a c) := by rw [Measure.compProd_withDensity hg, withDensity_compProd_left hf] exact (withDensity_mul _ (hf.comp measurable_fst) hg).symm From a4e7449f6357b13ce9e547c7a5ed987757be9941 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 3 Mar 2026 16:13:37 +0000 Subject: [PATCH 059/155] Refactor HistoryDensity (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 5 +- .../BayesStationaryEnv.lean | 18 +- .../SequentialLearning/HistoryDensity.lean | 252 ++++++++++-------- 3 files changed, 146 insertions(+), 129 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 16a98ac3..2371a195 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -226,8 +226,9 @@ lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (IsBayesAlgEnvSeq.bestAction κ id) =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) t).map (IsBayesAlgEnvSeq.bestAction κ id) := by - filter_upwards [h.condDistrib_env_hist_alg_indep (uniformAlgorithm_IsPositive hK) - (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) t] with x hx + filter_upwards [(h.hasCondDistrib_env_hist (uniformAlgorithm_IsPositive hK) + (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) t).condDistrib_eq] + with x hx simp only [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hx] exact (h.hasCondDistrib_action' t).condDistrib_eq.trans (h_comp.trans h_map).symm diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 840052de..242e1db1 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -96,17 +96,15 @@ lemma hasCondDistrib_reward' [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg HasCondDistrib (R' (n + 1)) (fun ω ↦ (E ω, A (n + 1) ω)) κ P := (h.hasCondDistrib_reward n).comp_left (by fun_prop) -end Laws - -section Maps +lemma hasLaw_hist [SFinite Q] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : + HasLaw (IsAlgEnvSeq.hist A R' n) (condDistrib (IsAlgEnvSeq.hist A R' n) E P ∘ₘ Q) P where + aemeasurable := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable + map_eq := by + rw [← Measure.snd_map_prodMk h.measurable_E, ← compProd_map_condDistrib + (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable, + h.hasLaw_env.map_eq, Measure.snd_compProd] -lemma map_hist_eq_condDistrib_comp [SFinite Q] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (t : ℕ) : - P.map (IsAlgEnvSeq.hist A R' t) = condDistrib (IsAlgEnvSeq.hist A R' t) E P ∘ₘ Q := by - rw [← Measure.snd_map_prodMk h.measurable_E, ← compProd_map_condDistrib - (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t).aemeasurable, - h.hasLaw_env.map_eq, Measure.snd_compProd] - -end Maps +end Laws section CondDistribIsAlgEnvSeq diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index e12740bd..bece15a9 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -104,32 +104,34 @@ lemma absolutelyContinuous_map_hist (h : IsAlgEnvSeq A R' alg env P) apply Measure.AbsolutelyContinuous.kernel_compProd_left exact Measure.absolutelyContinuous_of_forall_singleton_pos (hp.2 n h') -lemma map_hist_eq_withDensity_historyDensity (h : IsAlgEnvSeq A R' alg env P) +lemma hasLaw_hist_withDensity (h : IsAlgEnvSeq A R' alg env P) (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ env P₀) (hp : alg₀.IsPositive) (n : ℕ) : - P.map (IsAlgEnvSeq.hist A R' n) = - (P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity (historyDensity alg alg₀ n) := by - induction n with - | zero => - rw [h.hasLaw_hist_zero.map_eq, h₀.hasLaw_hist_zero.map_eq, h.hasLaw_step_zero.map_eq, - h₀.hasLaw_step_zero.map_eq] - have ha : alg.p0 ≪ alg₀.p0 := Measure.absolutelyContinuous_of_forall_singleton_pos hp.1 - rw [← Measure.withDensity_rnDeriv_eq _ _ ha, Measure.withDensity_compProd_left (by fun_prop)] - exact Measure.withDensity_map_equiv (by fun_prop) - | succ n ih => - let ρ h' (ar : α × R) := Kernel.rnDeriv (alg.policy n) (alg₀.policy n) h' ar.1 - have hpo h' : alg.policy n h' ≪ alg₀.policy n h' := - Measure.absolutelyContinuous_of_forall_singleton_pos (hp.2 n h') - have hs : stepKernel alg env n = (stepKernel alg₀ env n).withDensity ρ := by - rw [stepKernel, ← Kernel.withDensity_rnDeriv_eq' hpo] - exact Kernel.withDensity_compProd_left (Kernel.measurable_rnDeriv _ _) - have : IsMarkovKernel ((stepKernel alg₀ env n).withDensity ρ) := by - rw [← hs] - infer_instance - rw [(h.hasLaw_hist_succ n).map_eq, (h₀.hasLaw_hist_succ n).map_eq, - Measure.compProd_congr (h.hasCondDistrib_step n).condDistrib_eq, - Measure.compProd_congr (h₀.hasCondDistrib_step n).condDistrib_eq, ih, hs, - Measure.withDensity_compProd_withDensity (by fun_prop) (by fun_prop)] - exact Measure.withDensity_map_equiv (by fun_prop) + HasLaw (IsAlgEnvSeq.hist A R' n) + ((P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity (historyDensity alg alg₀ n)) P where + aemeasurable := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable + map_eq := by + induction n with + | zero => + rw [h.hasLaw_hist_zero.map_eq, h₀.hasLaw_hist_zero.map_eq, h.hasLaw_step_zero.map_eq, + h₀.hasLaw_step_zero.map_eq] + have ha : alg.p0 ≪ alg₀.p0 := Measure.absolutelyContinuous_of_forall_singleton_pos hp.1 + rw [← Measure.withDensity_rnDeriv_eq _ _ ha, Measure.withDensity_compProd_left (by fun_prop)] + exact Measure.withDensity_map_equiv (by fun_prop) + | succ n ih => + let ρ h' (ar : α × R) := Kernel.rnDeriv (alg.policy n) (alg₀.policy n) h' ar.1 + have hpo h' : alg.policy n h' ≪ alg₀.policy n h' := + Measure.absolutelyContinuous_of_forall_singleton_pos (hp.2 n h') + have hs : stepKernel alg env n = (stepKernel alg₀ env n).withDensity ρ := by + rw [stepKernel, ← Kernel.withDensity_rnDeriv_eq' hpo] + exact Kernel.withDensity_compProd_left (Kernel.measurable_rnDeriv _ _) + have : IsMarkovKernel ((stepKernel alg₀ env n).withDensity ρ) := by + rw [← hs] + infer_instance + rw [(h.hasLaw_hist_succ n).map_eq, (h₀.hasLaw_hist_succ n).map_eq, + Measure.compProd_congr (h.hasCondDistrib_step n).condDistrib_eq, + Measure.compProd_congr (h₀.hasCondDistrib_step n).condDistrib_eq, ih, hs, + Measure.withDensity_compProd_withDensity (by fun_prop) (by fun_prop)] + exact Measure.withDensity_map_equiv (by fun_prop) end IsAlgEnvSeq @@ -137,102 +139,104 @@ namespace IsBayesAlgEnvSeq variable {𝓔 : Type*} [MeasurableSpace 𝓔] variable [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] -variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] -variable {κ : Kernel (𝓔 × α) R} +variable {Q : Measure 𝓔} +variable {κ : Kernel (𝓔 × α) R} [IsMarkovKernel κ] variable {Ω : Type*} [MeasurableSpace Ω] variable {E : Ω → 𝓔} {A : ℕ → Ω → α} {R' : ℕ → Ω → R} variable {alg : Algorithm α R} variable {P : Measure Ω} [IsProbabilityMeasure P] -variable {alg₀ : Algorithm α R} + variable {Ω₀ : Type*} [MeasurableSpace Ω₀] variable {E₀ : Ω₀ → 𝓔} {A₀ : ℕ → Ω₀ → α} {R₀ : ℕ → Ω₀ → R} +variable {alg₀ : Algorithm α R} variable {P₀ : Measure Ω₀} [IsProbabilityMeasure P₀] -/-- The history distribution under any algorithm is absolutely continuous w.r.t. the - history distribution under a positive reference algorithm. -/ -lemma absolutelyContinuous_map_hist - [IsMarkovKernel κ] [StandardBorelSpace Ω] [Nonempty Ω] - [StandardBorelSpace Ω₀] [Nonempty Ω₀] +lemma hasCondDistrib_hist_withDensity_historyDensity (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - (hpos : alg₀.IsPositive) + (hp : alg₀.IsPositive) (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) - (t : ℕ) : - P.map (IsAlgEnvSeq.hist A R' t) ≪ - P₀.map (IsAlgEnvSeq.hist A₀ R₀ t) := by - set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' t) E P - set κ₀ := condDistrib (IsAlgEnvSeq.hist A₀ R₀ t) E₀ P₀ - rw [h.map_hist_eq_condDistrib_comp t, h₀.map_hist_eq_condDistrib_comp t, - ← Measure.snd_compProd, ← Measure.snd_compProd] - exact (Measure.AbsolutelyContinuous.compProd_right - (show ∀ᵐ e ∂Q, κ_alg e ≪ κ₀ e from by - have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : - (ℕ → α × R) → (Iic t → α × R)) = IT.hist t := - funext fun ω => funext fun i => Prod.mk.eta - filter_upwards [h.hasLaw_IT_hist t, h₀.hasLaw_IT_hist t, - h.ae_IsAlgEnvSeq, h₀.ae_IsAlgEnvSeq] with e he he₀ hae hae₀ - rw [← he.map_eq, ← he₀.map_eq, ← h_IT_hist] - exact hae.absolutelyContinuous_map_hist hae₀ hpos t)).map - measurable_snd + (n : ℕ) : + HasCondDistrib (IsAlgEnvSeq.hist A R' n) E + ((condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀).withDensity + (fun _ => historyDensity alg alg₀ n)) P where + aemeasurable_fst := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable + aemeasurable_snd := h.measurable_E.aemeasurable + condDistrib_eq := by + rw [h.hasLaw_env.map_eq] + have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward n : + (ℕ → α × R) → (Iic n → α × R)) = IT.hist n := + funext fun ω => funext fun i => Prod.mk.eta + filter_upwards [h.hasLaw_IT_hist n, h₀.hasLaw_IT_hist n, + h.ae_IsAlgEnvSeq, h₀.ae_IsAlgEnvSeq] with e he he₀ hae hae₀ + rw [Kernel.withDensity_apply _ (by fun_prop), + ← he.map_eq, ← he₀.map_eq, ← h_IT_hist] + exact (hae.hasLaw_hist_withDensity hae₀ hp n).map_eq -variable [StandardBorelSpace 𝓔] [Nonempty 𝓔] [IsMarkovKernel κ] +variable [IsProbabilityMeasure Q] -/-- The posterior on the environment given history is algorithm-independent. -/ -lemma condDistrib_env_hist_alg_indep +/-- The history distribution under any algorithm equals the reference algorithm's + history distribution, weighted by the history density ratio. -/ +lemma hasLaw_hist_withDensity (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - (hpos : alg₀.IsPositive) + (hp : alg₀.IsPositive) (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) - (t : ℕ) : - condDistrib E (IsAlgEnvSeq.hist A R' t) P - =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀ := by - set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' t) E P - set κ₀ := condDistrib (IsAlgEnvSeq.hist A₀ R₀ t) E₀ P₀ - set ρ := historyDensity alg alg₀ t - have hρ_meas := measurable_historyDensity alg alg₀ t - have hρ_ne_top := hpos.historyDensity_ne_top alg t - -- Key factorization: κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) - have h_wd_ae : κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) := by - have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward t : - (ℕ → α × R) → (Iic t → α × R)) = IT.hist t := - funext fun ω => funext fun i => Prod.mk.eta - filter_upwards [h.hasLaw_IT_hist t, h₀.hasLaw_IT_hist t, - h.ae_IsAlgEnvSeq, h₀.ae_IsAlgEnvSeq] with e he he₀ hae hae₀ - rw [Kernel.withDensity_apply _ - (show Measurable (Function.uncurry (fun (_ : 𝓔) => ρ)) from hρ_meas.comp measurable_snd), - ← he.map_eq, ← he₀.map_eq, ← h_IT_hist] - exact hae.map_hist_eq_withDensity_historyDensity hae₀ hpos t - haveI : IsSFiniteKernel (κ₀.withDensity (fun _ => ρ)) := - Kernel.IsSFiniteKernel.withDensity _ (fun _ b => hρ_ne_top b) - -- Direct condDistrib equality via joint measure argument - have h_joint : P.map (fun ω => (E ω, IsAlgEnvSeq.hist A R' t ω)) = Q ⊗ₘ κ_alg := by - rw [← h.hasLaw_env.map_eq] - exact (compProd_map_condDistrib - (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t).aemeasurable).symm - have h_joint₀ : P₀.map (fun ω => (E₀ ω, IsAlgEnvSeq.hist A₀ R₀ t ω)) = Q ⊗ₘ κ₀ := by - rw [← h₀.hasLaw_env.map_eq] - exact (compProd_map_condDistrib - (IsAlgEnvSeq.measurable_hist h₀.measurable_A h₀.measurable_R t).aemeasurable).symm - have h_meas_hist := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t - have h_meas_hist₀ := IsAlgEnvSeq.measurable_hist h₀.measurable_A h₀.measurable_R t - -- P.map hist = (P₀.map hist₀).withDensity ρ - have h_hist : P.map (IsAlgEnvSeq.hist A R' t) - = (P₀.map (IsAlgEnvSeq.hist A₀ R₀ t)).withDensity ρ := by - have h_marg : P.map (IsAlgEnvSeq.hist A R' t) = (Q ⊗ₘ κ_alg).map Prod.snd := by - rw [← h_joint] - exact (Measure.map_map measurable_snd (h.measurable_E.prodMk h_meas_hist)).symm - have h_marg₀ : P₀.map (IsAlgEnvSeq.hist A₀ R₀ t) = (Q ⊗ₘ κ₀).map Prod.snd := by - rw [← h_joint₀] - exact (Measure.map_map measurable_snd (h₀.measurable_E.prodMk h_meas_hist₀)).symm - rw [h_marg, h_marg₀, Measure.compProd_congr h_wd_ae, - Measure.compProd_withDensity - (show Measurable (Function.uncurry (fun (_ : 𝓔) => ρ)) from hρ_meas.comp measurable_snd)] - exact Measure.map_withDensity_comp measurable_snd hρ_meas - have h_swap : P.map (fun ω => (IsAlgEnvSeq.hist A R' t ω, E ω)) - = P.map (IsAlgEnvSeq.hist A R' t) ⊗ₘ condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀ := by - have h_uncurry_meas : Measurable (Function.uncurry (fun (_ : 𝓔) => ρ)) := - hρ_meas.comp measurable_snd - calc P.map (fun ω => (IsAlgEnvSeq.hist A R' t ω, E ω)) + (n : ℕ) : + HasLaw (IsAlgEnvSeq.hist A R' n) + ((P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity (historyDensity alg alg₀ n)) P where + aemeasurable := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable + map_eq := by + set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' n) E P + set κ₀ := condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀ + set ρ := historyDensity alg alg₀ n + have hρ_meas := measurable_historyDensity alg alg₀ n + have hρ_ne_top := hp.historyDensity_ne_top alg n + have h_wd_ae : κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) := by + rw [← h.hasLaw_env.map_eq] + exact (h.hasCondDistrib_hist_withDensity_historyDensity hp h₀ n).condDistrib_eq + haveI : IsSFiniteKernel (κ₀.withDensity (fun _ => ρ)) := + Kernel.IsSFiniteKernel.withDensity _ (fun _ b => hρ_ne_top b) + rw [(h.hasLaw_hist n).map_eq, (h₀.hasLaw_hist n).map_eq, + ← Measure.snd_compProd Q κ_alg, Measure.compProd_congr h_wd_ae, + Measure.snd_compProd Q (κ₀.withDensity (fun _ => ρ))] + exact Kernel.comp_withDensity_const hρ_meas + +variable [StandardBorelSpace 𝓔] [Nonempty 𝓔] + +/-- The joint distribution of (history, environment) under any algorithm equals + the history marginal compProd with the reference algorithm's posterior. -/ +lemma hasLaw_hist_env + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) + (hp : alg₀.IsPositive) + (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) + (n : ℕ) : + HasLaw (fun ω => (IsAlgEnvSeq.hist A R' n ω, E ω)) + (P.map (IsAlgEnvSeq.hist A R' n) ⊗ₘ + condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀) P where + aemeasurable := + ((IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).prodMk + h.measurable_E).aemeasurable + map_eq := by + set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' n) E P + set κ₀ := condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀ + set ρ := historyDensity alg alg₀ n + have hρ_meas := measurable_historyDensity alg alg₀ n + have hρ_ne_top := hp.historyDensity_ne_top alg n + have h_wd_ae : κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) := by + rw [← h.hasLaw_env.map_eq] + exact (h.hasCondDistrib_hist_withDensity_historyDensity hp h₀ n).condDistrib_eq + haveI : IsSFiniteKernel (κ₀.withDensity (fun _ => ρ)) := + Kernel.IsSFiniteKernel.withDensity _ (fun _ b => hρ_ne_top b) + have h_hist := (h.hasLaw_hist_withDensity hp h₀ n).map_eq + have h_meas_hist := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n + have h_meas_hist₀ := IsAlgEnvSeq.measurable_hist h₀.measurable_A h₀.measurable_R n + have h_joint : P.map (fun ω => (E ω, IsAlgEnvSeq.hist A R' n ω)) = Q ⊗ₘ κ_alg := by + rw [← h.hasLaw_env.map_eq] + exact (compProd_map_condDistrib h_meas_hist.aemeasurable).symm + have h_joint₀ : P₀.map (fun ω => (E₀ ω, IsAlgEnvSeq.hist A₀ R₀ n ω)) = Q ⊗ₘ κ₀ := by + rw [← h₀.hasLaw_env.map_eq] + exact (compProd_map_condDistrib h_meas_hist₀.aemeasurable).symm + calc P.map (fun ω => (IsAlgEnvSeq.hist A R' n ω, E ω)) _ = (Q ⊗ₘ κ_alg).map Prod.swap := by rw [← h_joint] exact (Measure.map_map measurable_swap @@ -240,27 +244,41 @@ lemma condDistrib_env_hist_alg_indep _ = (Q ⊗ₘ (κ₀.withDensity (fun _ => ρ))).map Prod.swap := by rw [Measure.compProd_congr h_wd_ae] _ = ((Q ⊗ₘ κ₀).withDensity (ρ ∘ Prod.snd)).map Prod.swap := by - congr 1; exact Measure.compProd_withDensity h_uncurry_meas + congr 1; exact Measure.compProd_withDensity (by fun_prop) _ = ((Q ⊗ₘ κ₀).map Prod.swap).withDensity (ρ ∘ Prod.fst) := Measure.map_swap_withDensity_fst hρ_meas - _ = (P₀.map (fun ω => (IsAlgEnvSeq.hist A₀ R₀ t ω, E₀ ω))).withDensity + _ = (P₀.map (fun ω => (IsAlgEnvSeq.hist A₀ R₀ n ω, E₀ ω))).withDensity (ρ ∘ Prod.fst) := by congr 1; rw [← h_joint₀] exact Measure.map_map measurable_swap (h₀.measurable_E.prodMk h_meas_hist₀) - _ = (P₀.map (IsAlgEnvSeq.hist A₀ R₀ t) ⊗ₘ - condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀).withDensity + _ = (P₀.map (IsAlgEnvSeq.hist A₀ R₀ n) ⊗ₘ + condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀).withDensity (ρ ∘ Prod.fst) := by rw [← compProd_map_condDistrib h₀.measurable_E.aemeasurable] - _ = (P₀.map (IsAlgEnvSeq.hist A₀ R₀ t)).withDensity ρ ⊗ₘ - condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀ := + _ = (P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity ρ ⊗ₘ + condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀ := (Measure.withDensity_compProd_left hρ_meas).symm - _ = P.map (IsAlgEnvSeq.hist A R' t) ⊗ₘ - condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀ := by + _ = P.map (IsAlgEnvSeq.hist A R' n) ⊗ₘ + condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀ := by rw [h_hist] - -- By uniqueness of disintegration - exact (condDistrib_ae_eq_iff_measure_eq_compProd _ - h.measurable_E.aemeasurable (condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ t) P₀)).mpr h_swap + +/-- The posterior on the environment given history is algorithm-independent. -/ +lemma hasCondDistrib_env_hist + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) + (hp : alg₀.IsPositive) + (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) + (n : ℕ) : + HasCondDistrib E (IsAlgEnvSeq.hist A R' n) + (condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀) P where + aemeasurable_fst := h.measurable_E.aemeasurable + aemeasurable_snd := + (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable + condDistrib_eq := + (condDistrib_ae_eq_iff_measure_eq_compProd _ + h.measurable_E.aemeasurable + (condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀)).mpr + (h.hasLaw_hist_env hp h₀ n).map_eq end IsBayesAlgEnvSeq From 37dbef211034a3443c0f4f4e2abd235e658c933b Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 4 Mar 2026 10:25:28 +0000 Subject: [PATCH 060/155] Refactor HistoryDensity (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 2 +- LeanBandits/BanditAlgorithms/Uniform.lean | 11 ++-- LeanBandits/ForMathlib/WithDensity.lean | 49 +++++++++++--- LeanBandits/SequentialLearning/Algorithm.lean | 7 +- .../SequentialLearning/HistoryDensity.lean | 66 ++++++------------- 5 files changed, 74 insertions(+), 61 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 2371a195..06895931 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -226,7 +226,7 @@ lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (IsBayesAlgEnvSeq.bestAction κ id) =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) t).map (IsBayesAlgEnvSeq.bestAction κ id) := by - filter_upwards [(h.hasCondDistrib_env_hist (uniformAlgorithm_IsPositive hK) + filter_upwards [(h.hasCondDistrib_env_hist (absolutelyContinuous_uniformAlgorithm hK _) (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) t).condDistrib_eq] with x hx simp only [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hx] diff --git a/LeanBandits/BanditAlgorithms/Uniform.lean b/LeanBandits/BanditAlgorithms/Uniform.lean index 553003d7..89705a3d 100644 --- a/LeanBandits/BanditAlgorithms/Uniform.lean +++ b/LeanBandits/BanditAlgorithms/Uniform.lean @@ -3,7 +3,7 @@ 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 Mathlib.Probability.UniformOn +import LeanBandits.ForMathlib.FullSupport import LeanBandits.SequentialLearning.Algorithm /-! # The Uniform Algorithm -/ @@ -23,8 +23,11 @@ def uniformAlgorithm (hK : 0 < K) : Algorithm (Fin K) ℝ := { policy _ := Kernel.const _ (uniformOn Set.univ) p0 := uniformOn Set.univ } -lemma uniformAlgorithm_IsPositive (hK : 0 < K) : (uniformAlgorithm hK).IsPositive := by - constructor - all_goals simp [uniformAlgorithm, uniformOn, cond_pos_of_inter_ne_zero] +lemma absolutelyContinuous_uniformAlgorithm (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) : + alg ≪ₐ uniformAlgorithm hK where + p0 := Measure.absolutelyContinuous_of_forall_singleton_pos + (by simp [uniformAlgorithm, uniformOn, cond_pos_of_inter_ne_zero]) + policy n h := Measure.absolutelyContinuous_of_forall_singleton_pos + (by simp [uniformAlgorithm, uniformOn, cond_pos_of_inter_ne_zero]) end Bandits diff --git a/LeanBandits/ForMathlib/WithDensity.lean b/LeanBandits/ForMathlib/WithDensity.lean index 90487287..b806f5aa 100644 --- a/LeanBandits/ForMathlib/WithDensity.lean +++ b/LeanBandits/ForMathlib/WithDensity.lean @@ -83,22 +83,53 @@ lemma withDensity_compProd_withDensity [SFinite μ] rw [Measure.compProd_withDensity hg, withDensity_compProd_left hf] exact (withDensity_mul _ (hf.comp measurable_fst) hg).symm +/-- If `κ =ᵐ[μ] η.withDensity (fun _ => f)`, then `μ ⊗ₘ κ = (μ ⊗ₘ η).withDensity (f ∘ Prod.snd)`. + Unlike `compProd_congr` + `compProd_withDensity`, does not need + `IsSFiniteKernel (η.withDensity (fun _ => f))`. -/ +lemma compProd_eq_compProd_withDensity [SFinite μ] + {κ η : Kernel α β} [IsSFiniteKernel κ] [IsSFiniteKernel η] + {f : β → ℝ≥0∞} (hf : Measurable f) + (h : κ =ᵐ[μ] η.withDensity (fun _ b ↦ f b)) : + μ ⊗ₘ κ = (μ ⊗ₘ η).withDensity (f ∘ Prod.snd) := by + have hf_uncurry : Measurable (Function.uncurry (fun (_ : α) => f)) := + hf.comp measurable_snd + ext s hs + have lhs : (μ ⊗ₘ κ) s = ∫⁻ a, (κ a) (Prod.mk a ⁻¹' s) ∂μ := + Measure.compProd_apply hs + have rhs : ((μ ⊗ₘ η).withDensity (f ∘ Prod.snd)) s = + ∫⁻ a, ∫⁻ b in Prod.mk a ⁻¹' s, f b ∂(η a) ∂μ := by + rw [withDensity_apply _ hs, ← lintegral_indicator hs, + Measure.lintegral_compProd ((hf.comp measurable_snd).indicator hs)] + congr 1; ext a + have : (fun b => s.indicator (f ∘ Prod.snd) (a, b)) = (Prod.mk a ⁻¹' s).indicator f := by + ext b; simp only [Set.indicator, Set.mem_preimage]; rfl + rw [this, lintegral_indicator (hs.preimage measurable_prodMk_left)] + rw [lhs, rhs] + apply lintegral_congr_ae + filter_upwards [h] with a ha + rw [ha, Kernel.withDensity_apply _ hf_uncurry, withDensity_apply _ (hs.preimage (by fun_prop))] + end Measure namespace ProbabilityTheory.Kernel /-- `(κ.withDensity (fun _ => f)) ∘ₘ μ = (κ ∘ₘ μ).withDensity f`. -/ lemma comp_withDensity_const - [SFinite μ] {κ : Kernel α γ} [IsSFiniteKernel κ] - {f : γ → ℝ≥0∞} (hf : Measurable f) - [IsSFiniteKernel (κ.withDensity (fun _ => f))] : - (κ.withDensity (fun _ => f)) ∘ₘ μ = (κ ∘ₘ μ).withDensity f := by - rw [← Measure.snd_compProd μ (κ.withDensity (fun _ => f)), - Measure.compProd_withDensity (show Measurable (Function.uncurry (fun (_ : α) => f)) from - hf.comp measurable_snd), - ← Measure.snd_compProd μ κ, Measure.snd, Measure.snd] - exact Measure.map_withDensity_comp measurable_snd hf + {f : γ → ℝ≥0∞} (hf : Measurable f) : + (κ.withDensity (fun _ c ↦ f c)) ∘ₘ μ = (κ ∘ₘ μ).withDensity f := by + have hf_uncurry : Measurable (Function.uncurry (fun (_ : α) => f)) := + hf.comp measurable_snd + ext s hs + have lhs : ((κ.withDensity (fun _ => f)) ∘ₘ μ) s = ∫⁻ a, ∫⁻ x in s, f x ∂(κ a) ∂μ := by + rw [Measure.bind_apply hs (Kernel.measurable _).aemeasurable] + congr 1; ext a + rw [Kernel.withDensity_apply _ hf_uncurry, withDensity_apply _ hs] + have rhs : ((κ ∘ₘ μ).withDensity f) s = ∫⁻ a, ∫⁻ x in s, f x ∂(κ a) ∂μ := by + rw [withDensity_apply _ hs, ← lintegral_indicator hs f, + Measure.lintegral_bind (Kernel.measurable _).aemeasurable ((hf.indicator hs).aemeasurable)] + congr 1; ext a; rw [lintegral_indicator hs f] + rw [lhs, rhs] /-- Composing `Kernel.withDensity` on the left kernel of `Kernel.compProd`: `(κ.withDensity f) ⊗ₖ η = (κ ⊗ₖ η).withDensity (fun a (b, _) => f a b)`. -/ diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index 88b6f514..926ecb9c 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -30,8 +30,11 @@ structure Algorithm (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] wher instance (alg : Algorithm α R) (n : ℕ) : IsMarkovKernel (alg.policy n) := alg.h_policy n instance (alg : Algorithm α R) : IsProbabilityMeasure alg.p0 := alg.hp0 -def Algorithm.IsPositive (alg : Algorithm α R) : Prop := - (∀ a, alg.p0 {a} > 0) ∧ (∀ n h a, alg.policy n h {a} > 0) +structure Algorithm.AbsolutelyContinuous (alg alg₀ : Algorithm α R) : Prop where + p0 : alg.p0 ≪ alg₀.p0 + policy n h : alg.policy n h ≪ alg₀.policy n h + +scoped notation:50 alg " ≪ₐ " alg₀ => Algorithm.AbsolutelyContinuous alg alg₀ /-- An algorithm that receives observations in `E × R` created form an algorithm that receives observations in `R` by ignoring the additional information. -/ diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index bece15a9..223986f6 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -34,14 +34,6 @@ lemma measurable_historyDensity [MeasurableSpace.CountablyGenerated α] (alg alg simp_rw [historyDensity] fun_prop -lemma Algorithm.IsPositive.historyDensity_ne_top [MeasurableSpace.CountablyGenerated α] - {alg₀ : Algorithm α R} (hp : alg₀.IsPositive) (alg : Algorithm α R) (n : ℕ) - (h : Iic n → α × R) : historyDensity alg alg₀ n h ≠ ⊤ := by - induction n with - | zero => exact Measure.rnDeriv_ne_top_of_forall_singleton_pos hp.1 _ - | succ n ih => - exact ENNReal.mul_ne_top (ih _) (Kernel.rnDeriv_ne_top_of_forall_singleton_pos (hp.2 n) _ _) - namespace IsAlgEnvSeq variable {Ω : Type*} [MeasurableSpace Ω] @@ -51,8 +43,7 @@ variable {A : ℕ → Ω → α} {R' : ℕ → Ω → R} variable {P : Measure Ω} [IsFiniteMeasure P] lemma hasLaw_hist_zero (h : IsAlgEnvSeq A R' alg env P) : HasLaw (hist A R' 0) - ((P.map (step A R' 0)).map - (MeasurableEquiv.piUnique (fun _ : Iic 0 ↦ α × R)).symm) P where + ((P.map (step A R' 0)).map (MeasurableEquiv.piUnique (fun _ : Iic 0 ↦ α × R)).symm) P where aemeasurable := (measurable_hist h.measurable_A h.measurable_R 0).aemeasurable map_eq := by have he : (MeasurableEquiv.piUnique (fun _ : Iic 0 ↦ α × R)).symm ∘ step A R' 0 = @@ -85,15 +76,14 @@ variable {A₀ : ℕ → Ω₀ → α} {R₀ : ℕ → Ω₀ → R} variable {P₀ : Measure Ω₀} [IsProbabilityMeasure P₀] lemma absolutelyContinuous_map_hist (h : IsAlgEnvSeq A R' alg env P) - (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ env P₀) (hp : alg₀.IsPositive) (n : ℕ) : + (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ env P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : P.map (IsAlgEnvSeq.hist A R' n) ≪ P₀.map (IsAlgEnvSeq.hist A₀ R₀ n) := by induction n with | zero => rw [h.hasLaw_hist_zero.map_eq, h₀.hasLaw_hist_zero.map_eq] apply Measure.AbsolutelyContinuous.map _ (by fun_prop) rw [h.hasLaw_step_zero.map_eq, h₀.hasLaw_step_zero.map_eq] - apply Measure.AbsolutelyContinuous.compProd_left - exact Measure.absolutelyContinuous_of_forall_singleton_pos hp.1 + exact Measure.AbsolutelyContinuous.compProd_left hc.p0 _ | succ n ih => rw [(h.hasLaw_hist_succ n).map_eq, (h₀.hasLaw_hist_succ n).map_eq] apply Measure.AbsolutelyContinuous.map _ (by fun_prop) @@ -101,12 +91,10 @@ lemma absolutelyContinuous_map_hist (h : IsAlgEnvSeq A R' alg env P) Measure.compProd_congr (h₀.hasCondDistrib_step n).condDistrib_eq] apply Measure.AbsolutelyContinuous.compProd ih filter_upwards with h' - apply Measure.AbsolutelyContinuous.kernel_compProd_left - exact Measure.absolutelyContinuous_of_forall_singleton_pos (hp.2 n h') + exact Measure.AbsolutelyContinuous.kernel_compProd_left (hc.policy n h') -lemma hasLaw_hist_withDensity (h : IsAlgEnvSeq A R' alg env P) - (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ env P₀) (hp : alg₀.IsPositive) (n : ℕ) : - HasLaw (IsAlgEnvSeq.hist A R' n) +lemma hasLaw_hist_withDensity (h : IsAlgEnvSeq A R' alg env P) (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ env P₀) + (hc : alg ≪ₐ alg₀) (n : ℕ) : HasLaw (IsAlgEnvSeq.hist A R' n) ((P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity (historyDensity alg alg₀ n)) P where aemeasurable := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable map_eq := by @@ -114,15 +102,13 @@ lemma hasLaw_hist_withDensity (h : IsAlgEnvSeq A R' alg env P) | zero => rw [h.hasLaw_hist_zero.map_eq, h₀.hasLaw_hist_zero.map_eq, h.hasLaw_step_zero.map_eq, h₀.hasLaw_step_zero.map_eq] - have ha : alg.p0 ≪ alg₀.p0 := Measure.absolutelyContinuous_of_forall_singleton_pos hp.1 - rw [← Measure.withDensity_rnDeriv_eq _ _ ha, Measure.withDensity_compProd_left (by fun_prop)] + rw [← Measure.withDensity_rnDeriv_eq _ _ hc.p0, + Measure.withDensity_compProd_left (by fun_prop)] exact Measure.withDensity_map_equiv (by fun_prop) | succ n ih => let ρ h' (ar : α × R) := Kernel.rnDeriv (alg.policy n) (alg₀.policy n) h' ar.1 - have hpo h' : alg.policy n h' ≪ alg₀.policy n h' := - Measure.absolutelyContinuous_of_forall_singleton_pos (hp.2 n h') have hs : stepKernel alg env n = (stepKernel alg₀ env n).withDensity ρ := by - rw [stepKernel, ← Kernel.withDensity_rnDeriv_eq' hpo] + rw [stepKernel, ← Kernel.withDensity_rnDeriv_eq' (hc.policy n)] exact Kernel.withDensity_compProd_left (Kernel.measurable_rnDeriv _ _) have : IsMarkovKernel ((stepKernel alg₀ env n).withDensity ρ) := by rw [← hs] @@ -154,7 +140,7 @@ variable {P₀ : Measure Ω₀} [IsProbabilityMeasure P₀] lemma hasCondDistrib_hist_withDensity_historyDensity (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - (hp : alg₀.IsPositive) + (hc : alg ≪ₐ alg₀) (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (n : ℕ) : HasCondDistrib (IsAlgEnvSeq.hist A R' n) E @@ -171,7 +157,7 @@ lemma hasCondDistrib_hist_withDensity_historyDensity h.ae_IsAlgEnvSeq, h₀.ae_IsAlgEnvSeq] with e he he₀ hae hae₀ rw [Kernel.withDensity_apply _ (by fun_prop), ← he.map_eq, ← he₀.map_eq, ← h_IT_hist] - exact (hae.hasLaw_hist_withDensity hae₀ hp n).map_eq + exact (hae.hasLaw_hist_withDensity hae₀ hc n).map_eq variable [IsProbabilityMeasure Q] @@ -179,7 +165,7 @@ variable [IsProbabilityMeasure Q] history distribution, weighted by the history density ratio. -/ lemma hasLaw_hist_withDensity (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - (hp : alg₀.IsPositive) + (hc : alg ≪ₐ alg₀) (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (n : ℕ) : HasLaw (IsAlgEnvSeq.hist A R' n) @@ -190,16 +176,11 @@ lemma hasLaw_hist_withDensity set κ₀ := condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀ set ρ := historyDensity alg alg₀ n have hρ_meas := measurable_historyDensity alg alg₀ n - have hρ_ne_top := hp.historyDensity_ne_top alg n have h_wd_ae : κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) := by rw [← h.hasLaw_env.map_eq] - exact (h.hasCondDistrib_hist_withDensity_historyDensity hp h₀ n).condDistrib_eq - haveI : IsSFiniteKernel (κ₀.withDensity (fun _ => ρ)) := - Kernel.IsSFiniteKernel.withDensity _ (fun _ b => hρ_ne_top b) - rw [(h.hasLaw_hist n).map_eq, (h₀.hasLaw_hist n).map_eq, - ← Measure.snd_compProd Q κ_alg, Measure.compProd_congr h_wd_ae, - Measure.snd_compProd Q (κ₀.withDensity (fun _ => ρ))] - exact Kernel.comp_withDensity_const hρ_meas + exact (h.hasCondDistrib_hist_withDensity_historyDensity hc h₀ n).condDistrib_eq + rw [(h.hasLaw_hist n).map_eq, Measure.bind_congr_right h_wd_ae, + Kernel.comp_withDensity_const hρ_meas, (h₀.hasLaw_hist n).map_eq] variable [StandardBorelSpace 𝓔] [Nonempty 𝓔] @@ -207,7 +188,7 @@ variable [StandardBorelSpace 𝓔] [Nonempty 𝓔] the history marginal compProd with the reference algorithm's posterior. -/ lemma hasLaw_hist_env (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - (hp : alg₀.IsPositive) + (hc : alg ≪ₐ alg₀) (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (n : ℕ) : HasLaw (fun ω => (IsAlgEnvSeq.hist A R' n ω, E ω)) @@ -221,13 +202,10 @@ lemma hasLaw_hist_env set κ₀ := condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀ set ρ := historyDensity alg alg₀ n have hρ_meas := measurable_historyDensity alg alg₀ n - have hρ_ne_top := hp.historyDensity_ne_top alg n have h_wd_ae : κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) := by rw [← h.hasLaw_env.map_eq] - exact (h.hasCondDistrib_hist_withDensity_historyDensity hp h₀ n).condDistrib_eq - haveI : IsSFiniteKernel (κ₀.withDensity (fun _ => ρ)) := - Kernel.IsSFiniteKernel.withDensity _ (fun _ b => hρ_ne_top b) - have h_hist := (h.hasLaw_hist_withDensity hp h₀ n).map_eq + exact (h.hasCondDistrib_hist_withDensity_historyDensity hc h₀ n).condDistrib_eq + have h_hist := (h.hasLaw_hist_withDensity hc h₀ n).map_eq have h_meas_hist := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n have h_meas_hist₀ := IsAlgEnvSeq.measurable_hist h₀.measurable_A h₀.measurable_R n have h_joint : P.map (fun ω => (E ω, IsAlgEnvSeq.hist A R' n ω)) = Q ⊗ₘ κ_alg := by @@ -241,10 +219,8 @@ lemma hasLaw_hist_env rw [← h_joint] exact (Measure.map_map measurable_swap (h.measurable_E.prodMk h_meas_hist)).symm - _ = (Q ⊗ₘ (κ₀.withDensity (fun _ => ρ))).map Prod.swap := by - rw [Measure.compProd_congr h_wd_ae] _ = ((Q ⊗ₘ κ₀).withDensity (ρ ∘ Prod.snd)).map Prod.swap := by - congr 1; exact Measure.compProd_withDensity (by fun_prop) + congr 1; exact Measure.compProd_eq_compProd_withDensity hρ_meas h_wd_ae _ = ((Q ⊗ₘ κ₀).map Prod.swap).withDensity (ρ ∘ Prod.fst) := Measure.map_swap_withDensity_fst hρ_meas _ = (P₀.map (fun ω => (IsAlgEnvSeq.hist A₀ R₀ n ω, E₀ ω))).withDensity @@ -266,7 +242,7 @@ lemma hasLaw_hist_env /-- The posterior on the environment given history is algorithm-independent. -/ lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - (hp : alg₀.IsPositive) + (hc : alg ≪ₐ alg₀) (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (n : ℕ) : HasCondDistrib E (IsAlgEnvSeq.hist A R' n) @@ -278,7 +254,7 @@ lemma hasCondDistrib_env_hist (condDistrib_ae_eq_iff_measure_eq_compProd _ h.measurable_E.aemeasurable (condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀)).mpr - (h.hasLaw_hist_env hp h₀ n).map_eq + (h.hasLaw_hist_env hc h₀ n).map_eq end IsBayesAlgEnvSeq From e9370437c26aee291a1e430c37608dfd7901310a Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 4 Mar 2026 14:14:08 +0000 Subject: [PATCH 061/155] Refactor HistoryDensity (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 5 +- LeanBandits/ForMathlib/CondDistrib.lean | 5 + .../BayesStationaryEnv.lean | 8 - .../SequentialLearning/HistoryDensity.lean | 147 ++++++------------ 4 files changed, 53 insertions(+), 112 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 06895931..294ca436 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -226,8 +226,9 @@ lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (IsBayesAlgEnvSeq.bestAction κ id) =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) t).map (IsBayesAlgEnvSeq.bestAction κ id) := by - filter_upwards [(h.hasCondDistrib_env_hist (absolutelyContinuous_uniformAlgorithm hK _) - (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) t).condDistrib_eq] + filter_upwards [(h.hasCondDistrib_env_hist + (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) + (absolutelyContinuous_uniformAlgorithm hK _) t).condDistrib_eq] with x hx simp only [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hx] exact (h.hasCondDistrib_action' t).condDistrib_eq.trans (h_comp.trans h_map).symm diff --git a/LeanBandits/ForMathlib/CondDistrib.lean b/LeanBandits/ForMathlib/CondDistrib.lean index 0f76b355..231a4787 100644 --- a/LeanBandits/ForMathlib/CondDistrib.lean +++ b/LeanBandits/ForMathlib/CondDistrib.lean @@ -392,6 +392,11 @@ lemma ae_eq_of_condDistrib_eq_deterministic {f : β → Ω} (hf : Measurable f) rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h hfX exact ae_eq_of_map_prodMk_eq hf hX hY (hfX ▸ h) +/-- The marginal law of `Y` is obtained by integrating `condDistrib Y X μ` against `μ.map X`. -/ +lemma map_bind_condDistrib (hX : Measurable X) (hY : AEMeasurable Y μ) : + (μ.map X).bind (condDistrib Y X μ) = μ.map Y := by + rw [← Measure.snd_compProd, compProd_map_condDistrib hY, Measure.snd_map_prodMk hX] + end CondDistrib section Cond diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 242e1db1..e5890852 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -96,14 +96,6 @@ lemma hasCondDistrib_reward' [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg HasCondDistrib (R' (n + 1)) (fun ω ↦ (E ω, A (n + 1) ω)) κ P := (h.hasCondDistrib_reward n).comp_left (by fun_prop) -lemma hasLaw_hist [SFinite Q] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : - HasLaw (IsAlgEnvSeq.hist A R' n) (condDistrib (IsAlgEnvSeq.hist A R' n) E P ∘ₘ Q) P where - aemeasurable := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable - map_eq := by - rw [← Measure.snd_map_prodMk h.measurable_E, ← compProd_map_condDistrib - (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable, - h.hasLaw_env.map_eq, Measure.snd_compProd] - end Laws section CondDistribIsAlgEnvSeq diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index 223986f6..f0be051d 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -90,8 +90,7 @@ lemma absolutelyContinuous_map_hist (h : IsAlgEnvSeq A R' alg env P) rw [Measure.compProd_congr (h.hasCondDistrib_step n).condDistrib_eq, Measure.compProd_congr (h₀.hasCondDistrib_step n).condDistrib_eq] apply Measure.AbsolutelyContinuous.compProd ih - filter_upwards with h' - exact Measure.AbsolutelyContinuous.kernel_compProd_left (hc.policy n h') + filter_upwards with h' using Measure.AbsolutelyContinuous.kernel_compProd_left (hc.policy n h') lemma hasLaw_hist_withDensity (h : IsAlgEnvSeq A R' alg env P) (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ env P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : HasLaw (IsAlgEnvSeq.hist A R' n) @@ -138,123 +137,67 @@ variable {E₀ : Ω₀ → 𝓔} {A₀ : ℕ → Ω₀ → α} {R₀ : ℕ → variable {alg₀ : Algorithm α R} variable {P₀ : Measure Ω₀} [IsProbabilityMeasure P₀] -lemma hasCondDistrib_hist_withDensity_historyDensity - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - (hc : alg ≪ₐ alg₀) - (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) - (n : ℕ) : +lemma hasCondDistrib_hist_condDistrib_withDensity (h : IsBayesAlgEnvSeq Q κ alg E A R' P) + (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : HasCondDistrib (IsAlgEnvSeq.hist A R' n) E ((condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀).withDensity - (fun _ => historyDensity alg alg₀ n)) P where + (fun _ ↦ historyDensity alg alg₀ n)) P where aemeasurable_fst := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable aemeasurable_snd := h.measurable_E.aemeasurable condDistrib_eq := by rw [h.hasLaw_env.map_eq] - have h_IT_hist : (IsAlgEnvSeq.hist IT.action IT.reward n : - (ℕ → α × R) → (Iic n → α × R)) = IT.hist n := - funext fun ω => funext fun i => Prod.mk.eta - filter_upwards [h.hasLaw_IT_hist n, h₀.hasLaw_IT_hist n, - h.ae_IsAlgEnvSeq, h₀.ae_IsAlgEnvSeq] with e he he₀ hae hae₀ - rw [Kernel.withDensity_apply _ (by fun_prop), - ← he.map_eq, ← he₀.map_eq, ← h_IT_hist] + filter_upwards [h.ae_IsAlgEnvSeq, h₀.ae_IsAlgEnvSeq, h.hasLaw_IT_hist n, h₀.hasLaw_IT_hist n] + with _ hae hae₀ he he₀ + rw [Kernel.withDensity_apply _ (by fun_prop), ← he.map_eq, ← he₀.map_eq] exact (hae.hasLaw_hist_withDensity hae₀ hc n).map_eq -variable [IsProbabilityMeasure Q] - -/-- The history distribution under any algorithm equals the reference algorithm's - history distribution, weighted by the history density ratio. -/ -lemma hasLaw_hist_withDensity - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - (hc : alg ≪ₐ alg₀) - (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) - (n : ℕ) : - HasLaw (IsAlgEnvSeq.hist A R' n) - ((P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity (historyDensity alg alg₀ n)) P where - aemeasurable := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable - map_eq := by - set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' n) E P - set κ₀ := condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀ - set ρ := historyDensity alg alg₀ n - have hρ_meas := measurable_historyDensity alg alg₀ n - have h_wd_ae : κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) := by - rw [← h.hasLaw_env.map_eq] - exact (h.hasCondDistrib_hist_withDensity_historyDensity hc h₀ n).condDistrib_eq - rw [(h.hasLaw_hist n).map_eq, Measure.bind_congr_right h_wd_ae, - Kernel.comp_withDensity_const hρ_meas, (h₀.hasLaw_hist n).map_eq] - variable [StandardBorelSpace 𝓔] [Nonempty 𝓔] +variable [IsProbabilityMeasure Q] -/-- The joint distribution of (history, environment) under any algorithm equals - the history marginal compProd with the reference algorithm's posterior. -/ -lemma hasLaw_hist_env - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - (hc : alg ≪ₐ alg₀) - (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) - (n : ℕ) : +lemma hasLaw_hist_env (h : IsBayesAlgEnvSeq Q κ alg E A R' P) + (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : HasLaw (fun ω => (IsAlgEnvSeq.hist A R' n ω, E ω)) - (P.map (IsAlgEnvSeq.hist A R' n) ⊗ₘ - condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀) P where - aemeasurable := - ((IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).prodMk + (P.map (IsAlgEnvSeq.hist A R' n) ⊗ₘ condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀) P where + aemeasurable := ((IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).prodMk h.measurable_E).aemeasurable map_eq := by - set κ_alg := condDistrib (IsAlgEnvSeq.hist A R' n) E P - set κ₀ := condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀ + have hA := h.measurable_A + have hR := h.measurable_R + have hA₀ := h₀.measurable_A + have hR₀ := h₀.measurable_R + have hE := h.measurable_E + have hE₀ := h₀.measurable_E set ρ := historyDensity alg alg₀ n - have hρ_meas := measurable_historyDensity alg alg₀ n - have h_wd_ae : κ_alg =ᵐ[Q] κ₀.withDensity (fun _ => ρ) := by - rw [← h.hasLaw_env.map_eq] - exact (h.hasCondDistrib_hist_withDensity_historyDensity hc h₀ n).condDistrib_eq - have h_hist := (h.hasLaw_hist_withDensity hc h₀ n).map_eq - have h_meas_hist := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n - have h_meas_hist₀ := IsAlgEnvSeq.measurable_hist h₀.measurable_A h₀.measurable_R n - have h_joint : P.map (fun ω => (E ω, IsAlgEnvSeq.hist A R' n ω)) = Q ⊗ₘ κ_alg := by - rw [← h.hasLaw_env.map_eq] - exact (compProd_map_condDistrib h_meas_hist.aemeasurable).symm - have h_joint₀ : P₀.map (fun ω => (E₀ ω, IsAlgEnvSeq.hist A₀ R₀ n ω)) = Q ⊗ₘ κ₀ := by - rw [← h₀.hasLaw_env.map_eq] - exact (compProd_map_condDistrib h_meas_hist₀.aemeasurable).symm - calc P.map (fun ω => (IsAlgEnvSeq.hist A R' n ω, E ω)) - _ = (Q ⊗ₘ κ_alg).map Prod.swap := by - rw [← h_joint] - exact (Measure.map_map measurable_swap - (h.measurable_E.prodMk h_meas_hist)).symm - _ = ((Q ⊗ₘ κ₀).withDensity (ρ ∘ Prod.snd)).map Prod.swap := by - congr 1; exact Measure.compProd_eq_compProd_withDensity hρ_meas h_wd_ae - _ = ((Q ⊗ₘ κ₀).map Prod.swap).withDensity (ρ ∘ Prod.fst) := - Measure.map_swap_withDensity_fst hρ_meas - _ = (P₀.map (fun ω => (IsAlgEnvSeq.hist A₀ R₀ n ω, E₀ ω))).withDensity - (ρ ∘ Prod.fst) := by - congr 1; rw [← h_joint₀] - exact Measure.map_map measurable_swap - (h₀.measurable_E.prodMk h_meas_hist₀) - _ = (P₀.map (IsAlgEnvSeq.hist A₀ R₀ n) ⊗ₘ - condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀).withDensity - (ρ ∘ Prod.fst) := by - rw [← compProd_map_condDistrib h₀.measurable_E.aemeasurable] - _ = (P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity ρ ⊗ₘ - condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀ := - (Measure.withDensity_compProd_left hρ_meas).symm - _ = P.map (IsAlgEnvSeq.hist A R' n) ⊗ₘ - condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀ := by - rw [h_hist] - -/-- The posterior on the environment given history is algorithm-independent. -/ -lemma hasCondDistrib_env_hist - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - (hc : alg ≪ₐ alg₀) - (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) - (n : ℕ) : + have h_wd_ae : condDistrib (IsAlgEnvSeq.hist A R' n) E P =ᵐ[Q] + (condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀).withDensity (fun _ => ρ) := + h.hasLaw_env.map_eq ▸ (h.hasCondDistrib_hist_condDistrib_withDensity h₀ hc n).condDistrib_eq + have h_hist : P.map (IsAlgEnvSeq.hist A R' n) = + (P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity ρ := by + rw [← map_bind_condDistrib hE (by fun_prop), h.hasLaw_env.map_eq, + Measure.bind_congr_right h_wd_ae, Kernel.comp_withDensity_const (by fun_prop), + ← h₀.hasLaw_env.map_eq, map_bind_condDistrib hE₀ (by fun_prop)] + have h_swap : P.map (fun ω => (IsAlgEnvSeq.hist A R' n ω, E ω)) = + (Q ⊗ₘ condDistrib (IsAlgEnvSeq.hist A R' n) E P).map Prod.swap := by + rw [← h.hasLaw_env.map_eq, compProd_map_condDistrib (by fun_prop)] + symm; exact Measure.map_map measurable_swap (by fun_prop) + have h_swap₀ : (Q ⊗ₘ condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀).map Prod.swap = + P₀.map (fun ω => (IsAlgEnvSeq.hist A₀ R₀ n ω, E₀ ω)) := by + rw [← h₀.hasLaw_env.map_eq, compProd_map_condDistrib (by fun_prop)] + exact Measure.map_map measurable_swap (by fun_prop) + rw [h_swap, Measure.compProd_eq_compProd_withDensity (by fun_prop) h_wd_ae, + Measure.map_swap_withDensity_fst (by fun_prop), h_swap₀, + ← compProd_map_condDistrib (by fun_prop), + ← Measure.withDensity_compProd_left (by fun_prop), ← h_hist] + +lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) + (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : HasCondDistrib E (IsAlgEnvSeq.hist A R' n) (condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀) P where aemeasurable_fst := h.measurable_E.aemeasurable - aemeasurable_snd := - (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable - condDistrib_eq := - (condDistrib_ae_eq_iff_measure_eq_compProd _ - h.measurable_E.aemeasurable - (condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀)).mpr - (h.hasLaw_hist_env hc h₀ n).map_eq + aemeasurable_snd := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable + condDistrib_eq := by + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ h.measurable_E.aemeasurable] + exact (h.hasLaw_hist_env h₀ hc n).map_eq end IsBayesAlgEnvSeq From d536774cf7584cbd666624b397a0108c8578e37b Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 4 Mar 2026 16:25:14 +0000 Subject: [PATCH 062/155] Refactor HistoryDensity (in progress) --- LeanBandits/ForMathlib/CondDistrib.lean | 6 ++++ .../SequentialLearning/HistoryDensity.lean | 29 ++++++++----------- 2 files changed, 18 insertions(+), 17 deletions(-) diff --git a/LeanBandits/ForMathlib/CondDistrib.lean b/LeanBandits/ForMathlib/CondDistrib.lean index 231a4787..c79d7266 100644 --- a/LeanBandits/ForMathlib/CondDistrib.lean +++ b/LeanBandits/ForMathlib/CondDistrib.lean @@ -397,6 +397,12 @@ lemma map_bind_condDistrib (hX : Measurable X) (hY : AEMeasurable Y μ) : (μ.map X).bind (condDistrib Y X μ) = μ.map Y := by rw [← Measure.snd_compProd, compProd_map_condDistrib hY, Measure.snd_map_prodMk hX] +/-- The joint law of `(Y, X)` equals the compProd of `μ.map X` and `condDistrib Y X μ`, swapped. -/ +lemma compProd_map_condDistrib_swap (hX : Measurable X) (hY : Measurable Y) : + (μ.map X ⊗ₘ condDistrib Y X μ).map Prod.swap = μ.map (fun ω ↦ (Y ω, X ω)) := by + rw [compProd_map_condDistrib hY.aemeasurable, Measure.map_map measurable_swap (hX.prodMk hY)] + rfl + end CondDistrib section Cond diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/HistoryDensity.lean index f0be051d..265ddaf9 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/HistoryDensity.lean @@ -156,7 +156,7 @@ variable [IsProbabilityMeasure Q] lemma hasLaw_hist_env (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : - HasLaw (fun ω => (IsAlgEnvSeq.hist A R' n ω, E ω)) + HasLaw (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, E ω)) (P.map (IsAlgEnvSeq.hist A R' n) ⊗ₘ condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀) P where aemeasurable := ((IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).prodMk h.measurable_E).aemeasurable @@ -168,26 +168,21 @@ lemma hasLaw_hist_env (h : IsBayesAlgEnvSeq Q κ alg E A R' P) have hE := h.measurable_E have hE₀ := h₀.measurable_E set ρ := historyDensity alg alg₀ n - have h_wd_ae : condDistrib (IsAlgEnvSeq.hist A R' n) E P =ᵐ[Q] - (condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀).withDensity (fun _ => ρ) := - h.hasLaw_env.map_eq ▸ (h.hasCondDistrib_hist_condDistrib_withDensity h₀ hc n).condDistrib_eq - have h_hist : P.map (IsAlgEnvSeq.hist A R' n) = + have hcd : condDistrib (IsAlgEnvSeq.hist A R' n) E P =ᵐ[Q] + (condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀).withDensity (fun _ ↦ ρ) := by + rw [← h.hasLaw_env.map_eq] + exact (h.hasCondDistrib_hist_condDistrib_withDensity h₀ hc n).condDistrib_eq + have hm : P.map (IsAlgEnvSeq.hist A R' n) = (P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity ρ := by rw [← map_bind_condDistrib hE (by fun_prop), h.hasLaw_env.map_eq, - Measure.bind_congr_right h_wd_ae, Kernel.comp_withDensity_const (by fun_prop), + Measure.bind_congr_right hcd, Kernel.comp_withDensity_const (by fun_prop), ← h₀.hasLaw_env.map_eq, map_bind_condDistrib hE₀ (by fun_prop)] - have h_swap : P.map (fun ω => (IsAlgEnvSeq.hist A R' n ω, E ω)) = - (Q ⊗ₘ condDistrib (IsAlgEnvSeq.hist A R' n) E P).map Prod.swap := by - rw [← h.hasLaw_env.map_eq, compProd_map_condDistrib (by fun_prop)] - symm; exact Measure.map_map measurable_swap (by fun_prop) - have h_swap₀ : (Q ⊗ₘ condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀).map Prod.swap = - P₀.map (fun ω => (IsAlgEnvSeq.hist A₀ R₀ n ω, E₀ ω)) := by - rw [← h₀.hasLaw_env.map_eq, compProd_map_condDistrib (by fun_prop)] - exact Measure.map_map measurable_swap (by fun_prop) - rw [h_swap, Measure.compProd_eq_compProd_withDensity (by fun_prop) h_wd_ae, - Measure.map_swap_withDensity_fst (by fun_prop), h_swap₀, + rw [← compProd_map_condDistrib_swap hE (by fun_prop), h.hasLaw_env.map_eq, + Measure.compProd_eq_compProd_withDensity (by fun_prop) hcd, + Measure.map_swap_withDensity_fst (by fun_prop), + ← h₀.hasLaw_env.map_eq, compProd_map_condDistrib_swap hE₀ (by fun_prop), ← compProd_map_condDistrib (by fun_prop), - ← Measure.withDensity_compProd_left (by fun_prop), ← h_hist] + ← Measure.withDensity_compProd_left (by fun_prop), ← hm] lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : From 10b4efd11a8b83ae0b3400b532861af7b476693f Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 5 Mar 2026 10:45:01 +0000 Subject: [PATCH 063/155] Refactor AlgorithmDensity (HistoryDensity, in progress) --- LeanBandits.lean | 2 +- LeanBandits/BanditAlgorithms/TS.lean | 2 +- ...toryDensity.lean => AlgorithmDensity.lean} | 22 +++++++++++-------- .../BayesStationaryEnv.lean | 20 ++++++++--------- 4 files changed, 25 insertions(+), 21 deletions(-) rename LeanBandits/SequentialLearning/{HistoryDensity.lean => AlgorithmDensity.lean} (93%) diff --git a/LeanBandits.lean b/LeanBandits.lean index 0fbf92ee..2e47600a 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -24,9 +24,9 @@ import LeanBandits.ForMathlib.SubGaussian import LeanBandits.ForMathlib.Traj import LeanBandits.ForMathlib.WithDensity import LeanBandits.SequentialLearning.Algorithm +import LeanBandits.SequentialLearning.AlgorithmDensity import LeanBandits.SequentialLearning.BayesStationaryEnv import LeanBandits.SequentialLearning.Deterministic import LeanBandits.SequentialLearning.FiniteActions -import LeanBandits.SequentialLearning.HistoryDensity import LeanBandits.SequentialLearning.IonescuTulceaSpace import LeanBandits.SequentialLearning.StationaryEnv diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 294ca436..e9eb1473 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -7,7 +7,7 @@ import LeanBandits.ForMathlib.SubGaussian import LeanBandits.BanditAlgorithms.Uniform import LeanBandits.BanditAlgorithms.UCB import LeanBandits.SequentialLearning.BayesStationaryEnv -import LeanBandits.SequentialLearning.HistoryDensity +import LeanBandits.SequentialLearning.AlgorithmDensity import Mathlib.Analysis.Complex.ExponentialBounds /-! # The Thompson Sampling Algorithm -/ diff --git a/LeanBandits/SequentialLearning/HistoryDensity.lean b/LeanBandits/SequentialLearning/AlgorithmDensity.lean similarity index 93% rename from LeanBandits/SequentialLearning/HistoryDensity.lean rename to LeanBandits/SequentialLearning/AlgorithmDensity.lean index 265ddaf9..9127c0c9 100644 --- a/LeanBandits/SequentialLearning/HistoryDensity.lean +++ b/LeanBandits/SequentialLearning/AlgorithmDensity.lean @@ -15,25 +15,29 @@ namespace Learning variable {α R : Type*} [MeasurableSpace α] [MeasurableSpace R] +namespace Algorithm + noncomputable -def historyDensity [MeasurableSpace.CountablyGenerated α] (alg alg₀ : Algorithm α R) : +def density [MeasurableSpace.CountablyGenerated α] (alg alg₀ : Algorithm α R) : (n : ℕ) → (Iic n → α × R) → ℝ≥0∞ | 0, h => (alg.p0.rnDeriv alg₀.p0 (h ⟨0, by simp⟩).1) | n + 1, h => let p := MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n h - historyDensity alg alg₀ n p.1 * (alg.policy n).rnDeriv (alg₀.policy n) p.1 p.2.1 + alg.density alg₀ n p.1 * (alg.policy n).rnDeriv (alg₀.policy n) p.1 p.2.1 @[fun_prop] -lemma measurable_historyDensity [MeasurableSpace.CountablyGenerated α] (alg alg₀ : Algorithm α R) - (n : ℕ) : Measurable (historyDensity alg alg₀ n) := by +lemma measurable_density [MeasurableSpace.CountablyGenerated α] (alg alg₀ : Algorithm α R) (n : ℕ) : + Measurable (alg.density alg₀ n) := by induction n with | zero => - simp_rw [historyDensity] + simp_rw [density] fun_prop | succ n ih => - simp_rw [historyDensity] + simp_rw [density] fun_prop +end Algorithm + namespace IsAlgEnvSeq variable {Ω : Type*} [MeasurableSpace Ω] @@ -94,7 +98,7 @@ lemma absolutelyContinuous_map_hist (h : IsAlgEnvSeq A R' alg env P) lemma hasLaw_hist_withDensity (h : IsAlgEnvSeq A R' alg env P) (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ env P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : HasLaw (IsAlgEnvSeq.hist A R' n) - ((P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity (historyDensity alg alg₀ n)) P where + ((P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity (alg.density alg₀ n)) P where aemeasurable := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable map_eq := by induction n with @@ -141,7 +145,7 @@ lemma hasCondDistrib_hist_condDistrib_withDensity (h : IsBayesAlgEnvSeq Q κ alg (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : HasCondDistrib (IsAlgEnvSeq.hist A R' n) E ((condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀).withDensity - (fun _ ↦ historyDensity alg alg₀ n)) P where + (fun _ ↦ alg.density alg₀ n)) P where aemeasurable_fst := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable aemeasurable_snd := h.measurable_E.aemeasurable condDistrib_eq := by @@ -167,7 +171,7 @@ lemma hasLaw_hist_env (h : IsBayesAlgEnvSeq Q κ alg E A R' P) have hR₀ := h₀.measurable_R have hE := h.measurable_E have hE₀ := h₀.measurable_E - set ρ := historyDensity alg alg₀ n + set ρ := alg.density alg₀ n have hcd : condDistrib (IsAlgEnvSeq.hist A R' n) E P =ᵐ[Q] (condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀).withDensity (fun _ ↦ ρ) := by rw [← h.hasLaw_env.map_eq] diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index e5890852..ac548e70 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -111,16 +111,6 @@ lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : rw [← Kernel.map_apply _ (IT.measurable_action 0), ← hc, show IT.action 0 ∘ trajectory A R' = A 0 from rfl, hcd, Kernel.const_apply]⟩ -lemma hasLaw_IT_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : - ∀ᵐ e ∂Q, HasLaw (IT.hist n) (condDistrib (IsAlgEnvSeq.hist A R' n) E P e) - (condDistrib (trajectory A R') E P e) := by - rw [← h.hasLaw_env.map_eq, show IsAlgEnvSeq.hist A R' n = IT.hist n ∘ trajectory A R' from rfl] - filter_upwards [condDistrib_comp E - (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable - (IT.measurable_hist n)] with e he - exact ⟨(IT.measurable_hist n).aemeasurable, by - rw [← Kernel.map_apply _ (IT.measurable_hist n), he]⟩ - lemma hasCondDistrib_IT_reward_zero [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : ∀ᵐ e ∂Q, HasCondDistrib (IT.reward 0) (IT.action 0) (κ.sectR e) (condDistrib (trajectory A R') E P e) := by @@ -153,6 +143,16 @@ lemma hasCondDistrib_IT_reward [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ al (IT.measurable_action (n + 1))) (IT.measurable_reward (n + 1)) (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable h.measurable_E.aemeasurable +lemma hasLaw_IT_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : + ∀ᵐ e ∂Q, HasLaw (IT.hist n) (condDistrib (IsAlgEnvSeq.hist A R' n) E P e) + (condDistrib (trajectory A R') E P e) := by + rw [← h.hasLaw_env.map_eq, show IsAlgEnvSeq.hist A R' n = IT.hist n ∘ trajectory A R' from rfl] + filter_upwards [condDistrib_comp E + (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable + (IT.measurable_hist n)] with e he + exact ⟨(IT.measurable_hist n).aemeasurable, by + rw [← Kernel.map_apply _ (IT.measurable_hist n), he]⟩ + lemma ae_IsAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv (κ.sectR e)) (condDistrib (trajectory A R') E P e) := by From fbbc1720374fee8fff584a64b41c29df94329a04 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 5 Mar 2026 11:06:40 +0000 Subject: [PATCH 064/155] Minor --- .../SequentialLearning/AlgorithmDensity.lean | 18 +++++------------- 1 file changed, 5 insertions(+), 13 deletions(-) diff --git a/LeanBandits/SequentialLearning/AlgorithmDensity.lean b/LeanBandits/SequentialLearning/AlgorithmDensity.lean index 9127c0c9..016cfe78 100644 --- a/LeanBandits/SequentialLearning/AlgorithmDensity.lean +++ b/LeanBandits/SequentialLearning/AlgorithmDensity.lean @@ -141,15 +141,11 @@ variable {E₀ : Ω₀ → 𝓔} {A₀ : ℕ → Ω₀ → α} {R₀ : ℕ → variable {alg₀ : Algorithm α R} variable {P₀ : Measure Ω₀} [IsProbabilityMeasure P₀] -lemma hasCondDistrib_hist_condDistrib_withDensity (h : IsBayesAlgEnvSeq Q κ alg E A R' P) +lemma condDistrib_hist_eq_condDistrib_hist_withDensity (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : - HasCondDistrib (IsAlgEnvSeq.hist A R' n) E + condDistrib (IsAlgEnvSeq.hist A R' n) E P =ᵐ[Q] ((condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀).withDensity - (fun _ ↦ alg.density alg₀ n)) P where - aemeasurable_fst := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable - aemeasurable_snd := h.measurable_E.aemeasurable - condDistrib_eq := by - rw [h.hasLaw_env.map_eq] + (fun _ ↦ alg.density alg₀ n)) := by filter_upwards [h.ae_IsAlgEnvSeq, h₀.ae_IsAlgEnvSeq, h.hasLaw_IT_hist n, h₀.hasLaw_IT_hist n] with _ hae hae₀ he he₀ rw [Kernel.withDensity_apply _ (by fun_prop), ← he.map_eq, ← he₀.map_eq] @@ -171,13 +167,9 @@ lemma hasLaw_hist_env (h : IsBayesAlgEnvSeq Q κ alg E A R' P) have hR₀ := h₀.measurable_R have hE := h.measurable_E have hE₀ := h₀.measurable_E - set ρ := alg.density alg₀ n - have hcd : condDistrib (IsAlgEnvSeq.hist A R' n) E P =ᵐ[Q] - (condDistrib (IsAlgEnvSeq.hist A₀ R₀ n) E₀ P₀).withDensity (fun _ ↦ ρ) := by - rw [← h.hasLaw_env.map_eq] - exact (h.hasCondDistrib_hist_condDistrib_withDensity h₀ hc n).condDistrib_eq + have hcd := h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n have hm : P.map (IsAlgEnvSeq.hist A R' n) = - (P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity ρ := by + (P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity (alg.density alg₀ n) := by rw [← map_bind_condDistrib hE (by fun_prop), h.hasLaw_env.map_eq, Measure.bind_congr_right hcd, Kernel.comp_withDensity_const (by fun_prop), ← h₀.hasLaw_env.map_eq, map_bind_condDistrib hE₀ (by fun_prop)] From e51a8043c36e6b6e4cf4e0dd551c2e4bd3bf1610 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 6 Mar 2026 08:53:38 +0000 Subject: [PATCH 065/155] Refactor AlgorithmDensity (in progress) --- .../SequentialLearning/AlgorithmDensity.lean | 48 ++++++++++--------- 1 file changed, 26 insertions(+), 22 deletions(-) diff --git a/LeanBandits/SequentialLearning/AlgorithmDensity.lean b/LeanBandits/SequentialLearning/AlgorithmDensity.lean index 016cfe78..54cb6049 100644 --- a/LeanBandits/SequentialLearning/AlgorithmDensity.lean +++ b/LeanBandits/SequentialLearning/AlgorithmDensity.lean @@ -151,15 +151,11 @@ lemma condDistrib_hist_eq_condDistrib_hist_withDensity (h : IsBayesAlgEnvSeq Q rw [Kernel.withDensity_apply _ (by fun_prop), ← he.map_eq, ← he₀.map_eq] exact (hae.hasLaw_hist_withDensity hae₀ hc n).map_eq -variable [StandardBorelSpace 𝓔] [Nonempty 𝓔] -variable [IsProbabilityMeasure Q] - -lemma hasLaw_hist_env (h : IsBayesAlgEnvSeq Q κ alg E A R' P) +lemma hasLaw_hist_withDensity (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : - HasLaw (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, E ω)) - (P.map (IsAlgEnvSeq.hist A R' n) ⊗ₘ condDistrib E₀ (IsAlgEnvSeq.hist A₀ R₀ n) P₀) P where - aemeasurable := ((IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).prodMk - h.measurable_E).aemeasurable + HasLaw (IsAlgEnvSeq.hist A R' n) + ((P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity (alg.density alg₀ n)) P where + aemeasurable := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable map_eq := by have hA := h.measurable_A have hR := h.measurable_R @@ -167,18 +163,13 @@ lemma hasLaw_hist_env (h : IsBayesAlgEnvSeq Q κ alg E A R' P) have hR₀ := h₀.measurable_R have hE := h.measurable_E have hE₀ := h₀.measurable_E - have hcd := h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n - have hm : P.map (IsAlgEnvSeq.hist A R' n) = - (P₀.map (IsAlgEnvSeq.hist A₀ R₀ n)).withDensity (alg.density alg₀ n) := by - rw [← map_bind_condDistrib hE (by fun_prop), h.hasLaw_env.map_eq, - Measure.bind_congr_right hcd, Kernel.comp_withDensity_const (by fun_prop), - ← h₀.hasLaw_env.map_eq, map_bind_condDistrib hE₀ (by fun_prop)] - rw [← compProd_map_condDistrib_swap hE (by fun_prop), h.hasLaw_env.map_eq, - Measure.compProd_eq_compProd_withDensity (by fun_prop) hcd, - Measure.map_swap_withDensity_fst (by fun_prop), - ← h₀.hasLaw_env.map_eq, compProd_map_condDistrib_swap hE₀ (by fun_prop), - ← compProd_map_condDistrib (by fun_prop), - ← Measure.withDensity_compProd_left (by fun_prop), ← hm] + rw [← map_bind_condDistrib hE (by fun_prop), h.hasLaw_env.map_eq, + Measure.bind_congr_right (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), + Kernel.comp_withDensity_const (by fun_prop), + ← h₀.hasLaw_env.map_eq, map_bind_condDistrib hE₀ (by fun_prop)] + +variable [StandardBorelSpace 𝓔] [Nonempty 𝓔] +variable [IsProbabilityMeasure Q] lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ R₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : @@ -187,8 +178,21 @@ lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) aemeasurable_fst := h.measurable_E.aemeasurable aemeasurable_snd := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable condDistrib_eq := by - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ h.measurable_E.aemeasurable] - exact (h.hasLaw_hist_env h₀ hc n).map_eq + have hA := h.measurable_A + have hR := h.measurable_R + have hA₀ := h₀.measurable_A + have hR₀ := h₀.measurable_R + have hE := h.measurable_E + have hE₀ := h₀.measurable_E + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ h.measurable_E.aemeasurable, + ← compProd_map_condDistrib_swap hE (by fun_prop), h.hasLaw_env.map_eq, + Measure.compProd_eq_compProd_withDensity (by fun_prop) + (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), + Measure.map_swap_withDensity_fst (by fun_prop), + ← h₀.hasLaw_env.map_eq, compProd_map_condDistrib_swap hE₀ (by fun_prop), + ← compProd_map_condDistrib (by fun_prop), + ← Measure.withDensity_compProd_left (by fun_prop), + ← (hasLaw_hist_withDensity h h₀ hc n).map_eq] end IsBayesAlgEnvSeq From 190131c6d8f07d23f10532df65fe3b7ff44d9e28 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 6 Mar 2026 11:32:19 +0000 Subject: [PATCH 066/155] Minor --- LeanBandits/SequentialLearning/BayesStationaryEnv.lean | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index ac548e70..2ebbfa5c 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -106,7 +106,7 @@ lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : filter_upwards [condDistrib_comp E ((measurable_trajectory h.measurable_A h.measurable_R).aemeasurable) (IT.measurable_action (α := α) (R := R) 0), - h.hasCondDistrib_action_zero.condDistrib_eq] with e hc hcd + h.hasCondDistrib_action_zero.condDistrib_eq] with _ hc hcd exact ⟨(IT.measurable_action 0).aemeasurable, by rw [← Kernel.map_apply _ (IT.measurable_action 0), ← hc, show IT.action 0 ∘ trajectory A R' = A 0 from rfl, hcd, Kernel.const_apply]⟩ @@ -127,7 +127,7 @@ lemma hasCondDistrib_IT_action (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ filter_upwards [(h.hasCondDistrib_action n).ae_hasCondDistrib_sectR (IT.measurable_hist n) (IT.measurable_action (n + 1)) (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable - h.measurable_E.aemeasurable] with e he + h.measurable_E.aemeasurable] with _ he rwa [Kernel.sectR_prodMkLeft] at he lemma hasCondDistrib_IT_reward [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : @@ -149,7 +149,7 @@ lemma hasLaw_IT_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : rw [← h.hasLaw_env.map_eq, show IsAlgEnvSeq.hist A R' n = IT.hist n ∘ trajectory A R' from rfl] filter_upwards [condDistrib_comp E (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable - (IT.measurable_hist n)] with e he + (IT.measurable_hist n)] with _ he exact ⟨(IT.measurable_hist n).aemeasurable, by rw [← Kernel.map_apply _ (IT.measurable_hist n), he]⟩ From f9bd60006672071df6e20db17d53327b5634b6c2 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 6 Mar 2026 13:07:05 +0000 Subject: [PATCH 067/155] Remove unused code --- LeanBandits/BanditAlgorithms/Uniform.lean | 2 +- LeanBandits/ForMathlib/FullSupport.lean | 33 -------------------- LeanBandits/ForMathlib/MeasurableArgMax.lean | 24 -------------- 3 files changed, 1 insertion(+), 58 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/Uniform.lean b/LeanBandits/BanditAlgorithms/Uniform.lean index 89705a3d..f385a454 100644 --- a/LeanBandits/BanditAlgorithms/Uniform.lean +++ b/LeanBandits/BanditAlgorithms/Uniform.lean @@ -12,7 +12,7 @@ open MeasureTheory ProbabilityTheory Learning namespace Bandits -variable {K : ℕ} {hK : 0 < K} +variable {K : ℕ} /-- The Uniform algorithm: actions are chosen uniformly at random. -/ noncomputable diff --git a/LeanBandits/ForMathlib/FullSupport.lean b/LeanBandits/ForMathlib/FullSupport.lean index 48fed438..51bd3bc7 100644 --- a/LeanBandits/ForMathlib/FullSupport.lean +++ b/LeanBandits/ForMathlib/FullSupport.lean @@ -3,16 +3,8 @@ 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 Mathlib.Probability.Kernel.RadonNikodym import Mathlib.Probability.Kernel.Composition.MeasureCompProd -/-! -# Absolute continuity and rnDeriv finiteness from full support - -When a reference measure gives positive mass to every singleton, any measure is absolutely -continuous with respect to it, and the Radon-Nikodym derivative is pointwise finite. --/ - open MeasureTheory ProbabilityTheory variable {α β : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {μ ν : Measure α} @@ -26,33 +18,8 @@ lemma absolutelyContinuous_of_forall_singleton_pos (hν : ∀ a : α, ν {a} > 0 · exact measure_empty · exact absurd (measure_mono_null (Set.singleton_subset_iff.mpr ha) hs) (hν a).ne' -/-- An ae property holds everywhere when the reference measure gives positive mass - to every singleton. -/ -lemma forall_of_ae_of_forall_singleton_pos (hν : ∀ a, ν {a} > 0) {p : α → Prop} - (hp : ∀ᵐ a ∂ν, p a) (a : α) : p a := by - by_contra h - exact absurd (measure_mono_null (Set.singleton_subset_iff.mpr h) (ae_iff.mp hp)) (hν a).ne' - -/-- `rnDeriv` is pointwise finite when the reference measure has full support on singletons. -/ -lemma rnDeriv_ne_top_of_forall_singleton_pos [SigmaFinite μ] - (hν : ∀ a, ν {a} > 0) (a : α) : μ.rnDeriv ν a ≠ ⊤ := - (forall_of_ae_of_forall_singleton_pos hν (Measure.rnDeriv_lt_top μ ν) a).ne - end Measure -namespace Kernel - -/-- Kernel `rnDeriv` is pointwise finite when the reference kernel has full support - on singletons. -/ -lemma rnDeriv_ne_top_of_forall_singleton_pos - [MeasurableSpace.CountableOrCountablyGenerated α β] - {κ η : Kernel α β} [IsFiniteKernel κ] [IsFiniteKernel η] - (hη : ∀ a b, η a {b} > 0) (a : α) (b : β) : - Kernel.rnDeriv κ η a b ≠ ⊤ := - (Measure.forall_of_ae_of_forall_singleton_pos (hη a) (Kernel.rnDeriv_lt_top κ η) b).ne - -end Kernel - variable {γ : Type*} {mγ : MeasurableSpace γ} namespace Measure.AbsolutelyContinuous diff --git a/LeanBandits/ForMathlib/MeasurableArgMax.lean b/LeanBandits/ForMathlib/MeasurableArgMax.lean index ca6f64bd..0769a25d 100644 --- a/LeanBandits/ForMathlib/MeasurableArgMax.lean +++ b/LeanBandits/ForMathlib/MeasurableArgMax.lean @@ -80,29 +80,5 @@ lemma isMaxOn_measurableArgmax {α : Type*} [LinearOrder α] rw [measurableArgmax, h_eq, MeasurableEmbedding.leftInverse_invFun (measurableEmbedding_encode 𝓨) y] -/-- Congruence lemma: measurableArgmax only depends on the function values at the point. -/ -lemma measurableArgmax_congr {𝓧₁ 𝓧₂ : Type*} {α : Type*} [LinearOrder α] - [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨] - (f₁ : 𝓧₁ → 𝓨 → α) (f₂ : 𝓧₂ → 𝓨 → α) - [∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f₁ x z ≤ f₁ x y] - [∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f₂ x z ≤ f₂ x y] - (x₁ : 𝓧₁) (x₂ : 𝓧₂) (h : f₁ x₁ = f₂ x₂) : - measurableArgmax f₁ x₁ = measurableArgmax f₂ x₂ := by - simp only [measurableArgmax]; congr 1 - exact Nat.find_congr' fun {_} => - ⟨fun ⟨y, hn, hy⟩ => ⟨y, hn, h ▸ hy⟩, fun ⟨y, hn, hy⟩ => ⟨y, hn, h.symm ▸ hy⟩⟩ - -/-- measurableArgmax is independent of the DecidablePred instance used. - This follows from Nat.find_congr' which handles different decidability instances. -/ -lemma measurableArgmax_eq_of_eq {α : Type*} [LinearOrder α] - [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨] - (f : 𝓧 → 𝓨 → α) - (d1 : ∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y) - (d2 : ∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y) - (x : 𝓧) : - @measurableArgmax 𝓧 𝓨 α _ _ _ _ _ _ f d1 x = @measurableArgmax 𝓧 𝓨 α _ _ _ _ _ _ f d2 x := by - simp only [measurableArgmax]; congr 1 - exact @Nat.find_congr' _ _ (d1 x) (d2 x) _ _ (fun {_} ↦ Iff.rfl) - end Finite end MeasurableArgmax From 8d0d429abc8ca3cfd347c6f0e6a65378e50a7865 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 6 Mar 2026 13:09:52 +0000 Subject: [PATCH 068/155] Revert rename --- LeanBandits/Bandit/SumRewards.lean | 4 ++-- LeanBandits/BanditAlgorithms/UCB.lean | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 9ed155cc..a80151ce 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -649,7 +649,7 @@ lemma prob_sum_ge_sqrt_log {σ2 : ℝ≥0} open Real omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -lemma streamMeasure_sampleMean_add_sqrt_le {σ2 : ℝ≥0} {c : ℝ} +lemma todo {σ2 : ℝ≥0} {c : ℝ} (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (hσ2 : σ2 ≠ 0) (hc : 0 ≤ c) (a : α) (n k : ℕ) (hk : k ≠ 0) : streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(2 * c * σ2 * log (n + 1) / k) ≤ (ν a)[id]} ≤ @@ -675,7 +675,7 @@ lemma streamMeasure_sampleMean_add_sqrt_le {σ2 : ℝ≥0} {c : ℝ} _ ≤ 1 / (n + 1) ^ c := prob_sum_le_sqrt_log hν hσ2 hc a k hk omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -lemma streamMeasure_le_sampleMean_sub_sqrt {σ2 : ℝ≥0} {c : ℝ} +lemma todo' {σ2 : ℝ≥0} {c : ℝ} (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (hσ2 : σ2 ≠ 0) (hc : 0 ≤ c) (a : α) (n k : ℕ) (hk : k ≠ 0) : streamMeasure ν diff --git a/LeanBandits/BanditAlgorithms/UCB.lean b/LeanBandits/BanditAlgorithms/UCB.lean index e8491763..806a2a12 100644 --- a/LeanBandits/BanditAlgorithms/UCB.lean +++ b/LeanBandits/BanditAlgorithms/UCB.lean @@ -241,7 +241,7 @@ lemma prob_ucbIndex_le [Nonempty (Fin K)] {alg : Algorithm (Fin K) ℝ} grind _ ≤ ∑ k ∈ Icc 1 n, (1 : ℝ≥0∞) / (n + 1) ^ c := by gcongr with k hk - exact streamMeasure_sampleMean_add_sqrt_le hν hσ2 hc a n k (by grind) + exact todo hν hσ2 hc a n k (by grind) _ ≤ (n + 1) * (1 : ℝ≥0∞) / (n + 1) ^ c := 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)] @@ -284,7 +284,7 @@ lemma prob_ucbIndex_ge [Nonempty (Fin K)] {alg : Algorithm (Fin K) ℝ} grind _ ≤ ∑ k ∈ Icc 1 n, (1 : ℝ≥0∞) / (n + 1) ^ c := by gcongr with k hk - exact streamMeasure_le_sampleMean_sub_sqrt hν hσ2 hc a n k (by grind) + exact todo' hν hσ2 hc a n k (by grind) _ ≤ (n + 1) * (1 : ℝ≥0∞) / (n + 1) ^ c := 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)] From cbfd74b3e279c859188286992cc49fc6c7484a6c Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 9 Mar 2026 14:09:59 +0000 Subject: [PATCH 069/155] Refactor TS.lean (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 586 ++++++++++++--------------- 1 file changed, 270 insertions(+), 316 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index e9eb1473..9c568263 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -69,48 +69,18 @@ end Algorithm section Regret -variable {E : Type*} [mE : MeasurableSpace E] [StandardBorelSpace E] [Nonempty E] +variable {𝓔 : Type*} [m𝓔 : MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] variable (hK : 0 < K) variable {Ω : Type*} [MeasurableSpace Ω] -variable (E' : Ω → E) (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → ℝ) -variable (Q : Measure E) [IsProbabilityMeasure Q] (κ : Kernel (E × Fin K) ℝ) [IsMarkovKernel κ] +variable (E : Ω → 𝓔) (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → ℝ) +variable (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] variable (P : Measure Ω) [IsProbabilityMeasure P] -noncomputable -def ucbIndex (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (σ2 lo hi δ : ℝ) - (a : Fin K) (t : ℕ) (ω : Ω) : ℝ := - if pullCount A a t ω = 0 then hi - else max lo (min hi - (empMean A R' a t ω - + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ)))) - -omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in -lemma lo_le_ucbIndex (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : ℕ) (ω : Ω) : - lo ≤ ucbIndex A R' σ2 lo hi δ a t ω := by - unfold ucbIndex; split_ifs <;> [exact hlo; exact le_max_left lo _] - -omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in -lemma ucbIndex_le_hi (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : ℕ) (ω : Ω) : - ucbIndex A R' σ2 lo hi δ a t ω ≤ hi := by - unfold ucbIndex; split_ifs <;> [exact le_refl _; exact max_le hlo (min_le_left hi _)] - -omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in -lemma ucbIndex_mem_Icc (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : ℕ) (ω : Ω) : - ucbIndex A R' σ2 lo hi δ a t ω ∈ Set.Icc lo hi := - ⟨lo_le_ucbIndex A R' σ2 lo hi δ hlo a t ω, ucbIndex_le_hi A R' σ2 lo hi δ hlo a t ω⟩ - -omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in -lemma abs_ucbIndex_le (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : ℕ) (ω : Ω) : - |ucbIndex A R' σ2 lo hi δ a t ω| ≤ max |lo| |hi| := by - have hmem := ucbIndex_mem_Icc A R' σ2 lo hi δ hlo a t ω - exact abs_le_max_abs_abs hmem.1 hmem.2 +namespace TS -omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in -lemma norm_ucbIndex_le (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : ℕ) (ω : Ω) : - ‖ucbIndex A R' σ2 lo hi δ a t ω‖ ≤ max |lo| |hi| := by - rw [Real.norm_eq_abs]; exact abs_ucbIndex_le A R' σ2 lo hi δ hlo a t ω +/-! ### Auxiliary real-analysis lemmas (candidates for migration to a utility file) -/ -omit mE [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] in +omit m𝓔 [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] in lemma abs_sub_le_of_mem_Icc {lo hi x y : ℝ} (hx : x ∈ Set.Icc lo hi) (hy : y ∈ Set.Icc lo hi) : |x - y| ≤ hi - lo := by @@ -123,67 +93,66 @@ lemma sum_sqrt_le {ι : Type*} (s : Finset ι) (c : ι → ℝ) (hc : ∀ i, 0 calc ∑ i ∈ s, √(c i) ≤ √(∑ i ∈ s, c i) * √↑(#s) := h _ = _ := by rw [← Real.sqrt_mul (Finset.sum_nonneg (fun i _ => hc i)), mul_comm] -omit [StandardBorelSpace E] [Nonempty E] in -lemma sum_inv_sqrt_max_one_le (N : ℕ) : - ∑ j ∈ range N, (1 / √(↑(max 1 j) : ℝ)) ≤ 2 * √↑N := by - suffices h : ∀ M : ℕ, 0 < M → - ∑ j ∈ range M, (1 / √(↑(max 1 j) : ℝ)) + 1 / √↑M ≤ 2 * √↑M by - cases N with - | zero => simp - | succ n => - have := h (n + 1) (Nat.succ_pos n) - linarith [div_nonneg zero_le_one (Real.sqrt_nonneg (↑(n + 1) : ℝ))] - intro M hM +omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] in +lemma sum_inv_sqrt_le (M : ℕ) (hM : 0 < M) : + ∑ j ∈ range M, (1 / √(↑j : ℝ)) + 1 / √↑M ≤ 2 * √↑M := by induction M with | zero => omega | succ n ih => rw [sum_range_succ] by_cases hn : n = 0 - · subst hn; simp; norm_num + · subst hn; simp · have hn_pos : 0 < n := Nat.pos_of_ne_zero hn - have hmax : (↑(max 1 n) : ℝ) = ↑n := by - simp [Nat.max_eq_right (by omega : 1 ≤ n)] - rw [hmax] have h_ih := ih hn_pos suffices h_key : 1 / √(↑(n + 1) : ℝ) ≤ 2 * (√↑(n + 1) - √↑n) by linarith - have hns : (0 : ℝ) < ↑(n + 1) := by positivity - have hnn : (0 : ℝ) ≤ ↑n := by positivity - set a := √(↑(n + 1) : ℝ) - set b := √(↑n : ℝ) - have hsn : a * a = ↑(n + 1) := Real.mul_self_sqrt (le_of_lt hns) - have hs : b * b = ↑n := Real.mul_self_sqrt hnn - have hab : 2 * (a * b) ≤ ↑(n + 1) + ↑n := by - nlinarith [mul_self_nonneg (a - b)] - rw [div_le_iff₀ (by positivity : 0 < a)] - have h_expand : 2 * (a - b) * a = 2 * (a * a) - 2 * (a * b) := by ring - rw [h_expand, hsn] - have : (↑(n + 1) : ℝ) = ↑n + 1 := by push_cast; ring - linarith + rw [div_le_iff₀ (Real.sqrt_pos.mpr (by positivity : (0 : ℝ) < ↑(n + 1)))] + nlinarith [Real.mul_self_sqrt (show (0 : ℝ) ≤ ↑(n + 1) by positivity), + Real.mul_self_sqrt (show (0 : ℝ) ≤ ↑n by positivity), + mul_self_nonneg (√(↑(n + 1) : ℝ) - √(↑n : ℝ)), + show (↑(n + 1) : ℝ) = ↑n + 1 from by push_cast; ring] + +/-! ### UCB index definition and properties -/ + +noncomputable +def ucbIndex (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (σ2 lo hi δ : ℝ) + (a : Fin K) (t : ℕ) (ω : Ω) : ℝ := + if pullCount A a t ω = 0 then hi + else max lo (min hi + (empMean A R' a t ω + + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ)))) + +omit m𝓔 [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] in +lemma ucbIndex_mem_Icc (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : ℕ) (ω : Ω) : + ucbIndex A R' σ2 lo hi δ a t ω ∈ Set.Icc lo hi := by + unfold ucbIndex + split_ifs <;> constructor + · exact hlo + · exact le_refl _ + · exact le_max_left lo _ + · exact max_le hlo (min_le_left hi _) @[fun_prop] lemma measurable_ucbIndex [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) (σ2 lo hi δ : ℝ) (a : Fin K) (t : ℕ) : Measurable (ucbIndex A R' σ2 lo hi δ a t) := by unfold ucbIndex - have hpc : Measurable (fun ω ↦ (pullCount A a t ω : ℝ)) := + have : Measurable (fun ω ↦ (pullCount A a t ω : ℝ)) := measurable_from_top.comp (measurable_pullCount (fun n ↦ h.measurable_A n) a t) - refine Measurable.ite ?_ measurable_const ?_ - · exact (measurable_pullCount (fun n ↦ h.measurable_A n) a t) (measurableSet_singleton 0) - · exact (Measurable.max measurable_const (Measurable.min measurable_const - (Measurable.add (measurable_empMean (fun n ↦ h.measurable_A n) - (fun n ↦ h.measurable_R n) a t) - (measurable_const.div hpc).sqrt))) + have := measurable_empMean (fun n ↦ h.measurable_A n) (fun n ↦ h.measurable_R n) a t + have := measurable_pullCount (fun n ↦ h.measurable_A n) a t + exact .ite ((measurable_pullCount (fun n ↦ h.measurable_A n) a t) + (measurableSet_singleton 0)) measurable_const (by fun_prop) -omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in +omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma armMean_le_ucbIndex {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hconc : pullCount A a t ω ≠ 0 → - |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E' a ω| + |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E a ω| < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ))) : - IsBayesAlgEnvSeq.actionMean κ E' a ω ≤ ucbIndex A R' σ2 lo hi δ a t ω := by + IsBayesAlgEnvSeq.actionMean κ E a ω ≤ ucbIndex A R' σ2 lo hi δ a t ω := by unfold ucbIndex - have hmean := hm a (E' ω) + have hmean := hm a (E ω) simp only [IsBayesAlgEnvSeq.actionMean] at hmean hconc ⊢ split_ifs with h0 · exact hmean.2 @@ -191,13 +160,13 @@ lemma armMean_le_ucbIndex {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set. refine le_max_of_le_right (le_min hmean.2 ?_) linarith [habs.2] -omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in +omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hpc : pullCount A a t ω ≠ 0) (hconc : - |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E' a ω| + |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E a ω| < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ))) : - ucbIndex A R' σ2 lo hi δ a t ω - IsBayesAlgEnvSeq.actionMean κ E' a ω + ucbIndex A R' σ2 lo hi δ a t ω - IsBayesAlgEnvSeq.actionMean κ E a ω ≤ 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ)) := by unfold ucbIndex simp only [IsBayesAlgEnvSeq.actionMean] at hconc ⊢ @@ -205,24 +174,24 @@ lemma ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ ( set w := √(2 * σ2 * Real.log (1 / δ) / ↑(pullCount A a t ω)) set emp := empMean A R' a t ω have habs := abs_sub_lt_iff.mp hconc - have hmean := hm a (E' ω) + have hmean := hm a (E ω) have h1 : max lo (min hi (emp + w)) ≤ emp + w := max_le_iff.mpr ⟨by linarith [hmean.1, habs.2], min_le_right _ _⟩ linarith [habs.2] lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) (t : ℕ) : + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) (t : ℕ) : condDistrib (A (t + 1)) (IsAlgEnvSeq.hist A R' t) P =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - condDistrib (IsBayesAlgEnvSeq.bestAction κ E') (IsAlgEnvSeq.hist A R' t) P := + condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' t) P := by - have h_ba_comp : IsBayesAlgEnvSeq.bestAction κ E' - = IsBayesAlgEnvSeq.bestAction κ id ∘ E' := rfl + have h_ba_comp : IsBayesAlgEnvSeq.bestAction κ E + = IsBayesAlgEnvSeq.bestAction κ id ∘ E := rfl rw [h_ba_comp] have hm := IsBayesAlgEnvSeq.measurable_bestAction (κ := κ) measurable_id have h_comp := condDistrib_comp (mβ := MeasurableSpace.pi) (μ := P) (IsAlgEnvSeq.hist A R' t) h.measurable_E.aemeasurable hm - have h_map : (condDistrib E' (IsAlgEnvSeq.hist A R' t) P).map + have h_map : (condDistrib E (IsAlgEnvSeq.hist A R' t) P).map (IsBayesAlgEnvSeq.bestAction κ id) =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) t).map (IsBayesAlgEnvSeq.bestAction κ id) := by @@ -233,42 +202,37 @@ lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] simp only [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hx] exact (h.hasCondDistrib_action' t).condDistrib_eq.trans (h_comp.trans h_map).symm -omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in -lemma le_armMean_bestArm [Nonempty (Fin K)] (ω : Ω) (i : Fin K) : - IsBayesAlgEnvSeq.actionMean κ E' i ω ≤ - IsBayesAlgEnvSeq.actionMean κ E' (IsBayesAlgEnvSeq.bestAction κ E' ω) ω := by - have := isMaxOn_measurableArgmax (fun ω a ↦ IsBayesAlgEnvSeq.actionMean κ E' a ω) ω i - simp only [IsBayesAlgEnvSeq.bestAction]; convert this - -omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in +omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) - (ω : Ω) : ⨆ i, IsBayesAlgEnvSeq.actionMean κ E' i ω = - IsBayesAlgEnvSeq.actionMean κ E' (IsBayesAlgEnvSeq.bestAction κ E' ω) ω := - le_antisymm (ciSup_le (le_armMean_bestArm E' κ ω)) - (le_ciSup (f := fun i ↦ IsBayesAlgEnvSeq.actionMean κ E' i ω) + (ω : Ω) : ⨆ i, IsBayesAlgEnvSeq.actionMean κ E i ω = + IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω := + le_antisymm (ciSup_le fun i ↦ by + have := isMaxOn_measurableArgmax (fun ω a ↦ IsBayesAlgEnvSeq.actionMean κ E a ω) ω i + simp only [IsBayesAlgEnvSeq.bestAction]; convert this) + (le_ciSup (f := fun i ↦ IsBayesAlgEnvSeq.actionMean κ E i ω) ⟨hi, by rintro _ ⟨i, rfl⟩; exact (hm i _).2⟩ _) -omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsMarkovKernel κ] in +omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma gap_eq_armMean_sub [Nonempty (Fin K)] {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) - (s : ℕ) (ω : Ω) : gap (κ.sectR (E' ω)) (A s ω) = - IsBayesAlgEnvSeq.actionMean κ E' (IsBayesAlgEnvSeq.bestAction κ E' ω) ω - - IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω := by + (s : ℕ) (ω : Ω) : gap (κ.sectR (E ω)) (A s ω) = + IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω := by simp only [gap, Kernel.sectR_apply] - exact congr_arg (· - _) (iSup_armMean_eq_bestArm E' κ hm ω) + exact congr_arg (· - _) (iSup_armMean_eq_bestArm E κ hm ω) -omit [StandardBorelSpace E] [Nonempty E] [IsProbabilityMeasure Q] [IsMarkovKernel κ] in +omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [IsProbabilityMeasure Q] [IsMarkovKernel κ] in lemma bayesRegret_eq_sum_integral_gap [Nonempty (Fin K)] {alg : Algorithm (Fin K) ℝ} - (h : IsBayesAlgEnvSeq Q κ alg E' A R' P) + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {C : ℝ} (hm : ∀ a e, |(κ (e, a))[id]| ≤ C) (t : ℕ) : - P[IsBayesAlgEnvSeq.regret κ E' A t] = - ∑ s ∈ range t, P[fun ω ↦ gap (κ.sectR (E' ω)) + P[IsBayesAlgEnvSeq.regret κ E A t] = + ∑ s ∈ range t, P[fun ω ↦ gap (κ.sectR (E ω)) (A s ω)] := by simp only [IsBayesAlgEnvSeq.regret, regret_eq_sum_gap] refine integral_finset_sum _ (fun s _ => ?_) - have hmeas : Measurable (fun ω ↦ gap (κ.sectR (E' ω)) + have hmeas : Measurable (fun ω ↦ gap (κ.sectR (E ω)) (A s ω)) := (Measurable.iSup (fun a ↦ IsBayesAlgEnvSeq.measurable_actionMean (a := a) h.measurable_E)).sub @@ -277,13 +241,13 @@ lemma bayesRegret_eq_sum_integral_gap [Nonempty (Fin K)] refine ⟨hmeas.aestronglyMeasurable, HasFiniteIntegral.of_bounded (C := 2 * C) (Filter.Eventually.of_forall fun ω => ?_)⟩ simp only [Real.norm_eq_abs, gap, Kernel.sectR_apply] - have hbdd : BddAbove (Set.range fun i => (κ (E' ω, i))[id]) := + have hbdd : BddAbove (Set.range fun i => (κ (E ω, i))[id]) := ⟨C, by rintro _ ⟨i, rfl⟩; exact le_of_abs_le (hm i _)⟩ rw [abs_of_nonneg (sub_nonneg.mpr (le_ciSup hbdd _))] - linarith [ciSup_le fun i => le_of_abs_le (hm i (E' ω)), - neg_le_of_abs_le (hm (A s ω) (E' ω))] + linarith [ciSup_le fun i => le_of_abs_le (hm i (E ω)), + neg_le_of_abs_le (hm (A s ω) (E ω))] -omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsProbabilityMeasure Q] +omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] [IsProbabilityMeasure Q] [IsMarkovKernel κ] [IsProbabilityMeasure P] in lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : ∑ s ∈ range n, f (pullCount A (A s ω) s ω) = @@ -307,15 +271,15 @@ lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : congr 1 simp -omit [StandardBorelSpace E] [Nonempty E] [MeasurableSpace Ω] [IsProbabilityMeasure Q] +omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] [IsProbabilityMeasure Q] [IsMarkovKernel κ] [IsProbabilityMeasure P] in lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (hlo : lo ≤ hi) (σ2 δ : ℝ) (n : ℕ) (ω : Ω) (hconc : ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → - |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω| + |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω| < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ))) : ∑ s ∈ range n, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω) + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ≤ (hi - lo) * ↑K + 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by -- Split range n into first-pull (pc=0) and non-first-pull (pc≠0) sets set S0 := (range n).filter (fun s => pullCount A (A s ω) s ω = 0) @@ -326,50 +290,51 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] rw [Finset.sum_union hdisj] -- We bound ∑_{S0} and ∑_{S1} separately, then combine suffices h_S0 : ∑ s ∈ S0, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω) ≤ (hi - lo) * ↑K by + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ≤ (hi - lo) * ↑K by suffices h_S1 : ∑ s ∈ S1, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω) + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ≤ 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) by have := Finset.sum_union hdisj (f := fun s => - ucbIndex A R' σ2 lo hi δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω) + ucbIndex A R' σ2 lo hi δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) rw [← hpart] at this; linarith -- Bound ∑_{S1}: each term ≤ 2√(2σ2c/pc) = 2√(2σ2c/max(1,pc)), so ≤ full sum calc ∑ s ∈ S1, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω) + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ≤ ∑ s ∈ S1, - 2 * √(2 * σ2 * Real.log (1 / δ) / (max 1 (pullCount A (A s ω) s ω) : ℝ)) := + 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) := sum_le_sum fun s hs => by have hpc : pullCount A (A s ω) s ω ≠ 0 := (Finset.mem_filter.mp hs).2 - have hpc_eq : (max 1 (pullCount A (A s ω) s ω) : ℝ) = - (pullCount A (A s ω) s ω : ℝ) := by - simp [Nat.one_le_iff_ne_zero.mpr hpc] - rw [hpc_eq] - exact ucbIndex_sub_armMean_le E' A R' κ hm σ2 δ (A s ω) s ω hpc + exact ucbIndex_sub_armMean_le E A R' κ hm σ2 δ (A s ω) s ω hpc (hconc s (mem_range.mp (Finset.mem_filter.mp hs).1) _ hpc) _ ≤ ∑ s ∈ range n, - 2 * √(2 * σ2 * Real.log (1 / δ) / (max 1 (pullCount A (A s ω) s ω) : ℝ)) := + 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) := Finset.sum_le_sum_of_subset_of_nonneg (Finset.filter_subset _ _) fun s _ _ => by positivity _ ≤ 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by set c := Real.log (1 / δ) by_cases hc : 0 ≤ 2 * σ2 * c · open Real in - calc ∑ s ∈ range n, 2 * √(2 * σ2 * c / max 1 ↑(pullCount A (A s ω) s ω)) + calc ∑ s ∈ range n, 2 * √(2 * σ2 * c / ↑(pullCount A (A s ω) s ω)) = ∑ s ∈ range n, √(8 * σ2 * c) * - (1 / √(↑(max 1 (pullCount A (A s ω) s ω)) : ℝ)) := + (1 / √(↑(pullCount A (A s ω) s ω) : ℝ)) := sum_congr rfl fun s _ => by rw [show (8 : ℝ) * σ2 * c = (2 : ℝ) ^ 2 * (2 * σ2 * c) from by ring] rw [sqrt_mul (by positivity : (0:ℝ) ≤ 2 ^ 2), sqrt_sq (by norm_num : (0:ℝ) ≤ 2)] - rw [sqrt_div (by linarith : 0 ≤ 2 * σ2 * c)]; push_cast; ring + rw [sqrt_div (by linarith : 0 ≤ 2 * σ2 * c)]; ring _ = √(8 * σ2 * c) * ∑ s ∈ range n, - (1 / √(↑(max 1 (pullCount A (A s ω) s ω)) : ℝ)) := by + (1 / √(↑(pullCount A (A s ω) s ω) : ℝ)) := by rw [mul_sum] _ = √(8 * σ2 * c) * ∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω), - (1 / √(↑(max 1 j) : ℝ)) := by - congr 1; exact sum_comp_pullCount A (fun j => 1 / √(↑(max 1 j) : ℝ)) n ω + (1 / √(↑j : ℝ)) := by + congr 1; exact sum_comp_pullCount A (fun j => 1 / √(↑j : ℝ)) n ω _ ≤ √(8 * σ2 * c) * ∑ a : Fin K, (2 * √↑(pullCount A a n ω)) := by - gcongr with a; exact sum_inv_sqrt_max_one_le _ + gcongr with a + by_cases ha : pullCount A a n ω = 0 + · simp [ha] + · have := sum_inv_sqrt_le _ (Nat.pos_of_ne_zero ha) + linarith [div_nonneg zero_le_one + (Real.sqrt_nonneg (↑(pullCount A a n ω) : ℝ))] _ = √(8 * σ2 * c) * (2 * ∑ a : Fin K, √↑(pullCount A a n ω)) := by simp only [mul_sum] _ ≤ √(8 * σ2 * c) * (2 * √(↑K * ↑n)) := by @@ -383,19 +348,19 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] exact_mod_cast h _ = 2 * √(8 * σ2 * c) * √(↑K * ↑n) := by ring · have h0 : ∀ s ∈ range n, - 2 * √(2 * σ2 * c / max 1 ↑(pullCount A (A s ω) s ω)) = 0 := + 2 * √(2 * σ2 * c / ↑(pullCount A (A s ω) s ω)) = 0 := fun s _ => by open Real in - have : 2 * σ2 * c / max 1 ↑(pullCount A (A s ω) s ω) ≤ 0 := - div_nonpos_of_nonpos_of_nonneg (by linarith) (by positivity) + have : 2 * σ2 * c / ↑(pullCount A (A s ω) s ω) ≤ 0 := + div_nonpos_of_nonpos_of_nonneg (by linarith) (Nat.cast_nonneg _) simp [sqrt_eq_zero'.mpr this] rw [sum_congr rfl h0]; simp only [sum_const_zero]; positivity -- Bound ∑_{S0}: each term = hi - armMean ≤ hi - lo, and #S0 ≤ K have hterm_S0 : ∀ s ∈ S0, ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω ≤ hi - lo := fun s hs => by + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω ≤ hi - lo := fun s hs => by have hpc : pullCount A (A s ω) s ω = 0 := (Finset.mem_filter.mp hs).2 simp only [ucbIndex, hpc, ↓reduceIte, IsBayesAlgEnvSeq.actionMean] - linarith [(hm (A s ω) (E' ω)).1] + linarith [(hm (A s ω) (E ω)).1] have h_card_S0 : #S0 ≤ K := by calc #S0 ≤ #(Finset.univ : Finset (Fin K)) := Finset.card_le_card_of_injOn (fun s => A s ω) @@ -415,7 +380,7 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] Finset.card_ne_zero_of_mem this)) _ = K := Finset.card_fin K calc ∑ s ∈ S0, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E' (A s ω) ω) + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ≤ ∑ _s ∈ S0, (hi - lo) := sum_le_sum hterm_S0 _ = #S0 * (hi - lo) := by rw [sum_const, nsmul_eq_mul] _ ≤ ↑K * (hi - lo) := by @@ -423,6 +388,20 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] exact_mod_cast h_card_S0 _ = (hi - lo) * ↑K := by ring +private lemma exp_neg_sq_div_eq_delta {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : + ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / + (2 * k * ↑σ2))) = ENNReal.ofReal δ := by + have hk_pos : (0 : ℝ) < k := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hk) + have hσ2_pos : (0 : ℝ) < ↑σ2 := NNReal.coe_pos.mpr (pos_iff_ne_zero.mpr hσ2) + have hlog : 0 < Real.log (1 / δ) := + Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) + rw [Real.sq_sqrt (by positivity)] + simp only [neg_div, Real.exp_neg] + rw [show 2 * (k : ℝ) * ↑σ2 * Real.log (1 / δ) / (2 * k * ↑σ2) = + Real.log (1 / δ) from by field_simp [ne_of_gt hσ2_pos, ne_of_gt hk_pos]] + rw [Real.exp_log (by positivity : (0 : ℝ) < 1 / δ), one_div, inv_inv] + lemma streamMeasure_concentration_le_delta {α : Type*} [MeasurableSpace α] {ν : Kernel α ℝ} [IsMarkovKernel ν] {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) @@ -432,8 +411,6 @@ lemma streamMeasure_concentration_le_delta {α : Type*} [MeasurableSpace α] ENNReal.ofReal δ := by have hlog : 0 < Real.log (1 / δ) := Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) - have hk_pos : (0 : ℝ) < k := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hk) - have hσ2_pos : (0 : ℝ) < ↑σ2 := NNReal.coe_pos.mpr (pos_iff_ne_zero.mpr hσ2) calc streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(2 * ↑σ2 * Real.log (1 / δ) / k) ≤ (ν a)[id]} @@ -465,12 +442,7 @@ lemma streamMeasure_concentration_le_delta {α : Type*} [MeasurableSpace α] (fun i ω ↦ ω - (ν a)[id]) (fun _ ↦ by fun_prop) · intro i _; exact (hν a).congr_identDistrib ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) - _ = ENNReal.ofReal δ := by - rw [Real.sq_sqrt (by positivity)] - simp only [neg_div, Real.exp_neg] - rw [show 2 * (k : ℝ) * ↑σ2 * Real.log (1 / δ) / (2 * k * ↑σ2) = - Real.log (1 / δ) from by field_simp [ne_of_gt hσ2_pos, ne_of_gt hk_pos]] - rw [Real.exp_log (by positivity : (0 : ℝ) < 1 / δ), one_div, inv_inv] + _ = ENNReal.ofReal δ := exp_neg_sq_div_eq_delta hσ2 k hk δ hδ hδ1 lemma streamMeasure_concentration_ge_delta {α : Type*} [MeasurableSpace α] {ν : Kernel α ℝ} [IsMarkovKernel ν] {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) @@ -481,8 +453,6 @@ lemma streamMeasure_concentration_ge_delta {α : Type*} [MeasurableSpace α] ENNReal.ofReal δ := by have hlog : 0 < Real.log (1 / δ) := Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) - have hk_pos : (0 : ℝ) < k := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hk) - have hσ2_pos : (0 : ℝ) < ↑σ2 := NNReal.coe_pos.mpr (pos_iff_ne_zero.mpr hσ2) calc streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - √(2 * ↑σ2 * Real.log (1 / δ) / k)} @@ -514,12 +484,7 @@ lemma streamMeasure_concentration_ge_delta {α : Type*} [MeasurableSpace α] (fun _ ↦ by fun_prop) · intro i _; exact (hν a).congr_identDistrib ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) - _ = ENNReal.ofReal δ := by - rw [Real.sq_sqrt (by positivity)] - simp only [neg_div, Real.exp_neg] - rw [show 2 * (k : ℝ) * ↑σ2 * Real.log (1 / δ) / (2 * k * ↑σ2) = - Real.log (1 / δ) from by field_simp [ne_of_gt hσ2_pos, ne_of_gt hk_pos]] - rw [Real.exp_log (by positivity : (0 : ℝ) < 1 / δ), one_div, inv_inv] + _ = ENNReal.ofReal δ := exp_neg_sq_div_eq_delta hσ2 k hk δ hδ hδ1 private lemma streamMeasure_concentration_bound {α : Type*} [MeasurableSpace α] {ν : Kernel α ℝ} [IsMarkovKernel ν] {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) @@ -546,20 +511,20 @@ private lemma streamMeasure_concentration_bound {α : Type*} [MeasurableSpace α rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) (hδ_large : max |lo| |hi| < √(2 * ↑σ2 * Real.log (1 / δ))) : - ∀ᵐ e ∂(P.map (E')), - (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) + ∀ᵐ e ∂(P.map (E)), + (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) {ω | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} ≤ ENNReal.ofReal (2 * s * δ) := by - have h_cond_ae : ∀ᵐ e ∂(P.map E'), IsAlgEnvSeq IT.action IT.reward + have h_cond_ae : ∀ᵐ e ∂(P.map E), IsAlgEnvSeq IT.action IT.reward (tsAlgorithm hK Q κ) (stationaryEnv (κ.sectR e)) - (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E' P e) := by + (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) := by rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h filter_upwards [h_cond_ae] with e h_isAlgEnvSeq let ν := κ.sectR e @@ -567,7 +532,7 @@ lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] simp only [ν, Kernel.sectR_apply]; exact hs a' e have h_mean : (ν a)[id] = (κ (e, a))[id] := by simp only [ν, Kernel.sectR_apply] rw [← h_mean] - let P' := condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e + let P' := condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e have h_law := h_isAlgEnvSeq.law_pullCount_sumRewards_unique' (ArrayModel.isAlgEnvSeq_arrayMeasure (tsAlgorithm hK Q κ) ν) (n := s) let B_low := fun m : ℕ ↦ @@ -645,63 +610,63 @@ lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] congr 1; ring lemma prob_concentration_single_delta [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) (hδ_large : max |lo| |hi| < √(2 * ↑σ2 * Real.log (1 / δ))) : P {ω | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω|} ≤ + |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} ≤ ENNReal.ofReal (2 * s * δ) := by - let badSet : E → Set (ℕ → (Fin K) × ℝ) := fun e ↦ + let badSet : 𝓔 → Set (ℕ → (Fin K) × ℝ) := fun e ↦ {t | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s t) : ℝ)) ≤ |empMean IT.action IT.reward a s t - (κ (e, a))[id]|} have h_set_eq : {ω | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω|} = - (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' + |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} = + (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' {p | p.2 ∈ badSet p.1} := by ext ω simp only [Set.mem_setOf_eq, Set.mem_preimage, badSet, IsBayesAlgEnvSeq.actionMean] - have h1 : pullCount A a s ω = pullCount IT.action a s ((fun ω n => (A n ω, R' n ω)) ω) := by + have h1 : pullCount A a s ω = pullCount IT.action a s (IsBayesAlgEnvSeq.trajectory A R' ω) := by unfold pullCount IT.action; rfl have h2 : empMean A R' a s ω = - empMean IT.action IT.reward a s ((fun ω n => (A n ω, R' n ω)) ω) := by + empMean IT.action IT.reward a s (IsBayesAlgEnvSeq.trajectory A R' ω) := by unfold empMean IT.action IT.reward; rfl rw [h1, h2] have h_meas_pair : - Measurable (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) := - h.measurable_E.prodMk (measurable_pi_lambda _ fun n => - (h.measurable_A n).prodMk (h.measurable_R n)) - have h_disint : P.map (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) = - P.map (E') ⊗ₘ - condDistrib ((fun ω n => (A n ω, R' n ω))) E' P := - (compProd_map_condDistrib ((measurable_pi_lambda _ fun n => - (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable)).symm - have h_cond := prob_concentration_single_delta_cond hK E' A R' Q κ P h hσ2 hs hm a s δ hδ hδ1 + Measurable (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) := + h.measurable_E.prodMk (IsBayesAlgEnvSeq.measurable_trajectory h.measurable_A h.measurable_R) + have h_disint : P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) = + P.map (E) ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P := + (compProd_map_condDistrib + (IsBayesAlgEnvSeq.measurable_trajectory + h.measurable_A h.measurable_R).aemeasurable).symm + have h_cond := prob_concentration_single_delta_cond hK E A R' Q κ P h hσ2 hs hm a s δ hδ hδ1 hδ_large - have h_kernel : Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := + have h_kernel : Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_fst.prodMk measurable_const) - have h_meas_set : MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | p.2 ∈ badSet p.1} := by - change MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | + have h_meas_set : MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ badSet p.1} := by + change MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s p.2) : ℝ)) ≤ |empMean IT.action IT.reward a s p.2 - (κ (p.1, a))[id]|} exact measurableSet_le (by fun_prop) (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp measurable_snd).sub h_kernel).abs - calc P _ = P ((fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' + calc P _ = P ((fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' {p | p.2 ∈ badSet p.1}) := by rw [h_set_eq] - _ = (P.map (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω))) + _ = (P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω))) {p | p.2 ∈ badSet p.1} := by rw [Measure.map_apply h_meas_pair h_meas_set] - _ = (P.map (E') ⊗ₘ - condDistrib ((fun ω n => (A n ω, R' n ω))) E' P) + _ = (P.map (E) ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P) {p | p.2 ∈ badSet p.1} := by rw [h_disint] - _ = ∫⁻ e, (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) - (badSet e) ∂(P.map (E')) := by + _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) + (badSet e) ∂(P.map (E)) := by rw [Measure.compProd_apply h_meas_set]; rfl - _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * s * δ) ∂(P.map (E')) := by + _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * s * δ) ∂(P.map (E)) := by apply lintegral_mono_ae filter_upwards [h_cond] with e h_e; exact h_e _ = ENNReal.ofReal (2 * s * δ) := by @@ -712,17 +677,17 @@ private lemma concentration_cond_bound [Nonempty (Fin K)] {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {n : ℕ} (hn : 0 < n) {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) - (e : E) (h_isAlgEnvSeq : IsAlgEnvSeq IT.action IT.reward (tsAlgorithm hK Q κ) + (e : 𝓔) (h_isAlgEnvSeq : IsAlgEnvSeq IT.action IT.reward (tsAlgorithm hK Q κ) (stationaryEnv (κ.sectR e)) - (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e)) + (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e)) (a : Fin K) : - (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) + (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) (⋃ s ∈ Finset.range n, {ω | pullCount IT.action a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|}) ≤ ENNReal.ofReal (2 * n * δ) := by let ν := κ.sectR e - let P' := condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e + let P' := condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by simp only [ν, Kernel.sectR_apply]; exact hs a' e have h_mean : (ν a)[id] = (κ (e, a))[id] := by simp only [ν, Kernel.sectR_apply] @@ -821,20 +786,20 @@ private lemma concentration_cond_bound [Nonempty (Fin K)] nlinarith [(Nat.cast_le (α := ℝ)).mpr (Nat.sub_le n 1), hδ.le]) lemma prob_concentration_fail_delta [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : P {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω|} + |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} ≤ ENNReal.ofReal (2 * K * n * δ) := by let badSet := fun (s : ℕ) (a : Fin K) ↦ {ω : Ω | pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω|} + |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} have h_set_eq : {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E' a ω|} = + |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} = ⋃ s ∈ Finset.range n, ⋃ a : Fin K, badSet s a := by ext ω; simp only [Set.mem_setOf_eq, Finset.mem_range, Set.mem_iUnion, badSet, exists_prop] rw [h_set_eq] @@ -849,12 +814,12 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] by_cases hn : n = 0 · simp [hn] have hn' : 0 < n := Nat.pos_of_ne_zero hn - let badSetIT := fun (s : ℕ) (e : E) ↦ {ω : ℕ → (Fin K) × ℝ | + let badSetIT := fun (s : ℕ) (e : 𝓔) ↦ {ω : ℕ → (Fin K) × ℝ | pullCount IT.action a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} have h_set_eq : ⋃ s ∈ Finset.range n, badSet s a = - (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' + (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by ext ω simp only [Set.mem_iUnion, Finset.mem_range, badSet, badSetIT, Set.mem_preimage, @@ -862,35 +827,35 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] exact Iff.rfl rw [h_set_eq] have h_meas_pair : - Measurable (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) := - h.measurable_E.prodMk (measurable_pi_lambda _ fun n => - (h.measurable_A n).prodMk (h.measurable_R n)) - have h_disint : P.map (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) = - P.map (E') ⊗ₘ - condDistrib ((fun ω n => (A n ω, R' n ω))) E' P := - (compProd_map_condDistrib ((measurable_pi_lambda _ fun n => - (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable)).symm - have h_cond_bound : ∀ᵐ e ∂(P.map (E')), - (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) + Measurable (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) := + h.measurable_E.prodMk (IsBayesAlgEnvSeq.measurable_trajectory h.measurable_A h.measurable_R) + have h_disint : P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) = + P.map (E) ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P := + (compProd_map_condDistrib + (IsBayesAlgEnvSeq.measurable_trajectory + h.measurable_A h.measurable_R).aemeasurable).symm + have h_cond_bound : ∀ᵐ e ∂(P.map (E)), + (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) (⋃ s ∈ Finset.range n, badSetIT s e) ≤ ENNReal.ofReal (2 * n * δ) := by - have h_cond_ae : ∀ᵐ e ∂(P.map E'), IsAlgEnvSeq IT.action IT.reward + have h_cond_ae : ∀ᵐ e ∂(P.map E), IsAlgEnvSeq IT.action IT.reward (tsAlgorithm hK Q κ) (stationaryEnv (κ.sectR e)) - (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E' P e) := by + (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) := by rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h filter_upwards [h_cond_ae] with e h_isAlgEnvSeq - exact concentration_cond_bound (hK := hK) (E' := E') (A := A) (R' := R') + exact concentration_cond_bound (hK := hK) (E := E) (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) hσ2 hs hn' hδ hδ1 e h_isAlgEnvSeq a - have h_kernel : Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := + have h_kernel : Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_fst.prodMk measurable_const) - have h_meas_set : MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | + have h_meas_set : MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by - have h_eq : {p : E × (ℕ → (Fin K) × ℝ) | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} = + have h_eq : {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} = ⋃ s ∈ Finset.range n, {p | p.2 ∈ badSetIT s p.1} := by ext p; simp only [Set.mem_setOf_eq, Set.mem_iUnion, Finset.mem_range] rw [h_eq] exact .biUnion (Finset.range n).countable_toSet fun s _ ↦ by simp only [badSetIT] - change MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | + change MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | pullCount IT.action a s p.2 ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s p.2 : ℝ)) ≤ |empMean IT.action IT.reward a s p.2 - (κ (p.1, a))[id]|} @@ -900,19 +865,19 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] (measurableSet_le (by fun_prop) (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp measurable_snd).sub h_kernel).abs) - calc P ((fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' + calc P ((fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1}) - = (P.map (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω))) + = (P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω))) {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by rw [Measure.map_apply h_meas_pair h_meas_set] - _ = (P.map (E') ⊗ₘ - condDistrib ((fun ω n => (A n ω, R' n ω))) E' P) + _ = (P.map (E) ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P) {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by rw [h_disint] - _ = ∫⁻ e, (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) - (⋃ s ∈ Finset.range n, badSetIT s e) ∂(P.map (E')) := by + _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) + (⋃ s ∈ Finset.range n, badSetIT s e) ∂(P.map (E)) := by rw [Measure.compProd_apply h_meas_set]; rfl - _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * n * δ) ∂(P.map (E')) := by + _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * n * δ) ∂(P.map (E)) := by apply lintegral_mono_ae h_cond_bound _ = ENNReal.ofReal (2 * n * δ) := by rw [lintegral_const, Measure.map_apply h.measurable_E MeasurableSet.univ] @@ -928,31 +893,31 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] congr 1; ring lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - P {ω | ∃ s < n, pullCount A (IsBayesAlgEnvSeq.bestAction κ E' ω) s ω ≠ 0 ∧ + P {ω | ∃ s < n, pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / - (pullCount A (IsBayesAlgEnvSeq.bestAction κ E' ω) s ω : ℝ)) ≤ - |empMean A R' (IsBayesAlgEnvSeq.bestAction κ E' ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E' (IsBayesAlgEnvSeq.bestAction κ E' ω) ω|} + (pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) s ω : ℝ)) ≤ + |empMean A R' (IsBayesAlgEnvSeq.bestAction κ E ω) s ω - + IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω|} ≤ ENNReal.ofReal (2 * n * δ) := by by_cases hn : n = 0 · simp [hn] have hn' : 0 < n := Nat.pos_of_ne_zero hn - rw [show IsBayesAlgEnvSeq.bestAction κ E' = IsBayesAlgEnvSeq.bestAction κ id ∘ E' from + rw [show IsBayesAlgEnvSeq.bestAction κ E = IsBayesAlgEnvSeq.bestAction κ id ∘ E from rfl] - let badSetIT := fun (a : Fin K) (s : ℕ) (e : E) ↦ {ω : ℕ → (Fin K) × ℝ | + let badSetIT := fun (a : Fin K) (s : ℕ) (e : 𝓔) ↦ {ω : ℕ → (Fin K) × ℝ | pullCount IT.action a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} - have h_set_eq : {ω | ∃ s < n, pullCount A ((IsBayesAlgEnvSeq.bestAction κ id ∘ E') ω) s ω ≠ 0 ∧ + have h_set_eq : {ω | ∃ s < n, pullCount A ((IsBayesAlgEnvSeq.bestAction κ id ∘ E) ω) s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / - (pullCount A ((IsBayesAlgEnvSeq.bestAction κ id ∘ E') ω) s ω : ℝ)) ≤ - |empMean A R' ((IsBayesAlgEnvSeq.bestAction κ id ∘ E') ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E' ((IsBayesAlgEnvSeq.bestAction κ id ∘ E') ω) ω|} = - (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' + (pullCount A ((IsBayesAlgEnvSeq.bestAction κ id ∘ E) ω) s ω : ℝ)) ≤ + |empMean A R' ((IsBayesAlgEnvSeq.bestAction κ id ∘ E) ω) s ω - + IsBayesAlgEnvSeq.actionMean κ E ((IsBayesAlgEnvSeq.bestAction κ id ∘ E) ω) ω|} = + (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} := by ext ω @@ -961,38 +926,38 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] rfl rw [h_set_eq] have h_meas_pair : - Measurable (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) := - h.measurable_E.prodMk (measurable_pi_lambda _ fun n => - (h.measurable_A n).prodMk (h.measurable_R n)) - have h_disint : P.map (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) = - P.map E' ⊗ₘ condDistrib ((fun ω n => (A n ω, R' n ω))) E' P := - (compProd_map_condDistrib ((measurable_pi_lambda _ fun n => - (h.measurable_A n).prodMk (h.measurable_R n)).aemeasurable)).symm - have h_cond_bound : ∀ᵐ e ∂(P.map E'), ∀ a : Fin K, - (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) + Measurable (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) := + h.measurable_E.prodMk (IsBayesAlgEnvSeq.measurable_trajectory h.measurable_A h.measurable_R) + have h_disint : P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) = + P.map E ⊗ₘ condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P := + (compProd_map_condDistrib + (IsBayesAlgEnvSeq.measurable_trajectory + h.measurable_A h.measurable_R).aemeasurable).symm + have h_cond_bound : ∀ᵐ e ∂(P.map E), ∀ a : Fin K, + (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) (⋃ s ∈ Finset.range n, badSetIT a s e) ≤ ENNReal.ofReal (2 * n * δ) := by - have h_cond_ae : ∀ᵐ e ∂(P.map E'), IsAlgEnvSeq IT.action IT.reward + have h_cond_ae : ∀ᵐ e ∂(P.map E), IsAlgEnvSeq IT.action IT.reward (tsAlgorithm hK Q κ) (stationaryEnv (κ.sectR e)) - (condDistrib (fun ω n ↦ (A n ω, R' n ω)) E' P e) := by + (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) := by rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h filter_upwards [h_cond_ae] with e h_isAlgEnvSeq intro a - exact concentration_cond_bound (hK := hK) (E' := E') (A := A) (R' := R') + exact concentration_cond_bound (hK := hK) (E := E) (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) hσ2 hs hn' hδ hδ1 e h_isAlgEnvSeq a - have h_cond_best : ∀ᵐ e ∂(P.map E'), - (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) + have h_cond_best : ∀ᵐ e ∂(P.map E), + (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) (⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id e) s e) ≤ ENNReal.ofReal (2 * n * δ) := by filter_upwards [h_cond_bound] with e he exact he (IsBayesAlgEnvSeq.bestAction κ id e) - have h_kernel : ∀ a, Measurable (fun p : E × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := + have h_kernel : ∀ a, Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := fun a ↦ stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_fst.prodMk measurable_const) - have h_meas_badSetIT : ∀ a s, MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | + have h_meas_badSetIT : ∀ a s, MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ badSetIT a s p.1} := by intro a s simp only [badSetIT] - change MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | + change MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | pullCount IT.action a s p.2 ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s p.2 : ℝ)) ≤ |empMean IT.action IT.reward a s p.2 - (κ (p.1, a))[id]|} @@ -1002,9 +967,9 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] (measurableSet_le (by fun_prop) (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp measurable_snd).sub (h_kernel a)).abs) - have h_meas_set : MeasurableSet {p : E × (ℕ → (Fin K) × ℝ) | + have h_meas_set : MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} := by - have h_eq : {p : E × (ℕ → (Fin K) × ℝ) | + have h_eq : {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} = ⋃ a : Fin K, ((IsBayesAlgEnvSeq.bestAction κ id ∘ Prod.fst) ⁻¹' {a} ∩ ⋃ s ∈ Finset.range n, {p | p.2 ∈ badSetIT a s p.1}) := by @@ -1018,41 +983,41 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] ((IsBayesAlgEnvSeq.measurable_bestAction (κ := κ) measurable_id |>.comp measurable_fst) (measurableSet_singleton a)) (.biUnion (Finset.range n).countable_toSet fun s _ ↦ h_meas_badSetIT a s) - calc P ((fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω)) ⁻¹' + calc P ((fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1}) - = (P.map (fun ω ↦ (E' ω, (fun ω n => (A n ω, R' n ω)) ω))) + = (P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω))) {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} := by rw [Measure.map_apply h_meas_pair h_meas_set] - _ = (P.map E' ⊗ₘ - condDistrib ((fun ω n => (A n ω, R' n ω))) E' P) + _ = (P.map E ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P) {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} := by rw [h_disint] - _ = ∫⁻ e, (condDistrib ((fun ω n => (A n ω, R' n ω))) E' P e) + _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) (⋃ s ∈ Finset.range n, - badSetIT (IsBayesAlgEnvSeq.bestAction κ id e) s e) ∂(P.map E') := by + badSetIT (IsBayesAlgEnvSeq.bestAction κ id e) s e) ∂(P.map E) := by rw [Measure.compProd_apply h_meas_set]; rfl - _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * n * δ) ∂(P.map E') := by + _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * n * δ) ∂(P.map E) := by apply lintegral_mono_ae h_cond_best _ = ENNReal.ofReal (2 * n * δ) := by rw [lintegral_const, Measure.map_apply h.measurable_E MeasurableSet.univ] simp [measure_univ] lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - P[IsBayesAlgEnvSeq.regret κ E' A n] + P[IsBayesAlgEnvSeq.regret κ E A n] ≤ (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * n ^ 2 * δ + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _) have hlo : lo ≤ hi := h1.trans h2 - let bestArm := IsBayesAlgEnvSeq.bestAction κ E' - let armMean := IsBayesAlgEnvSeq.actionMean κ E' + let bestArm := IsBayesAlgEnvSeq.bestAction κ E + let armMean := IsBayesAlgEnvSeq.actionMean κ E let ucb := ucbIndex A R' (↑σ2) lo hi δ set Eδ : Set Ω := {ω | ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → |empMean A R' a s ω - armMean a ω| @@ -1061,10 +1026,10 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω| < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A (bestArm ω) s ω : ℝ))} have hm_ucb : ∀ a t, Measurable (ucbIndex A R' (↑σ2) lo hi δ a t) := - fun a t ↦ measurable_ucbIndex hK E' A R' Q κ P h (↑σ2) lo hi δ a t - have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.actionMean κ E' a) := + fun a t ↦ measurable_ucbIndex hK E A R' Q κ P h (↑σ2) lo hi δ a t + have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.actionMean κ E a) := fun a ↦ IsBayesAlgEnvSeq.measurable_actionMean (a := a) h.measurable_E - have hm_best : Measurable (IsBayesAlgEnvSeq.bestAction κ E') := + have hm_best : Measurable (IsBayesAlgEnvSeq.bestAction κ E) := IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E have h_first_bound : ∀ ω, |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| @@ -1101,22 +1066,24 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (measurable_apply_fin hm_arm (h.measurable_A s))).aestronglyMeasurable · filter_upwards with ω; rw [Real.norm_eq_abs]; exact h_second_bound ω have h_swap : - P[IsBayesAlgEnvSeq.regret κ E' A n] = + P[IsBayesAlgEnvSeq.regret κ E A n] = P[fun ω ↦ ∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)] + P[fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)] := by - have h_regret_eq : P[IsBayesAlgEnvSeq.regret κ E' A n] = + have h_regret_eq : P[IsBayesAlgEnvSeq.regret κ E A n] = ∑ s ∈ range n, P[fun ω ↦ armMean (bestArm ω) ω - armMean (A s ω) ω] := by rw [bayesRegret_eq_sum_integral_gap (h := h) (hm := fun a e ↦ abs_le_max_abs_abs (hm a e).1 (hm a e).2) (t := n)] congr 1 with s - exact integral_congr_ae (ae_of_all _ fun ω ↦ gap_eq_armMean_sub E' A κ hm s ω) + exact integral_congr_ae (ae_of_all _ fun ω ↦ gap_eq_armMean_sub E A κ hm s ω) have h_int_ucb : ∀ s {f : Ω → Fin K}, Measurable f → Integrable (fun ω ↦ ucb (f ω) s ω) P := fun s {_} hf ↦ ⟨(measurable_apply_fin (fun a ↦ hm_ucb a s) hf).aestronglyMeasurable, - HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ - norm_ucbIndex_le A R' (↑σ2) lo hi δ hlo _ _ _)⟩ + HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ by + rw [Real.norm_eq_abs] + exact abs_le_max_abs_abs (ucbIndex_mem_Icc A R' (↑σ2) lo hi δ hlo _ _ _).1 + (ucbIndex_mem_Icc A R' (↑σ2) lo hi δ hlo _ _ _).2)⟩ have h_int_ucb_sub : ∀ s, Integrable (fun ω ↦ ucb (A s ω) s ω - ucb (bestArm ω) s ω) P := fun s ↦ (h_int_ucb s (h.measurable_A s)).sub (h_int_ucb s hm_best) have h_ucb_zero : ∀ a (ω : Ω), ucbIndex A R' (↑σ2) lo hi δ a 0 ω = hi := by @@ -1130,15 +1097,15 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp simp [h_ucb_zero] exact (integral_congr_ae (ae_of_all _ this)).trans (integral_zero _ _) | succ t => - have hts := ts_identity hK E' A R' Q κ P h t + have hts := ts_identity hK E A R' Q κ P h t have h_map_eq : P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω)) = - P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E' ω)) := by + P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E ω)) := by rw [← compProd_map_condDistrib (hY := (h.measurable_A (t + 1)).aemeasurable), ← compProd_map_condDistrib (hY := hm_best.aemeasurable)] exact Measure.compProd_congr hts have h_int_eq : ∀ (f : (Iic t → Fin K × ℝ) × Fin K → ℝ), Measurable f → ∫ ω, f (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω) ∂P = - ∫ ω, f (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E' ω) ∂P := by + ∫ ω, f (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E ω) ∂P := by intro f hf have hm_hist := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t rw [← integral_map @@ -1173,16 +1140,14 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp measurable_const simp only [Nat.cast_eq_zero] at this; exact this · exact measurable_const - · apply Measurable.max measurable_const - apply Measurable.min measurable_const - apply Measurable.add - · exact measurable_apply_fin (fun a ↦ (measurable_empMean' t a).comp measurable_fst) - measurable_snd - · apply Measurable.sqrt - apply Measurable.div measurable_const - exact measurable_apply_fin - (fun a ↦ measurable_from_top.comp ((measurable_pullCount' t a).comp measurable_fst)) - measurable_snd + · exact .max measurable_const (.min measurable_const + (.add (measurable_apply_fin + (fun a ↦ (measurable_empMean' t a).comp measurable_fst) + measurable_snd) + (measurable_const.div (measurable_apply_fin + (fun a ↦ measurable_from_top.comp + ((measurable_pullCount' t a).comp measurable_fst)) + measurable_snd)).sqrt)) rw [show (fun ω ↦ ucb (A (t + 1) ω) (t + 1) ω - ucb (bestArm ω) (t + 1) ω) = fun ω ↦ (fun ω ↦ ucb (A (t + 1) ω) (t + 1) ω) ω - @@ -1234,20 +1199,20 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp intro ω hω apply Finset.sum_nonpos intro s hs - linarith [armMean_le_ucbIndex E' A R' κ hm (↑σ2) δ + linarith [armMean_le_ucbIndex E A R' κ hm (↑σ2) δ (bestArm ω) s ω (hω s (mem_range.mp hs))] have h_second_Eδ : ∀ ω ∈ Eδ, ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω) ≤ (hi - lo) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by intro ω hω - exact sum_ucbIndex_sub_armMean_le E' A R' κ hm hlo (↑σ2) δ n ω hω + exact sum_ucbIndex_sub_armMean_le E A R' κ hm hlo (↑σ2) δ n ω hω have h_prob : P Eδᶜ ≤ ENNReal.ofReal (2 * K * n * δ) := by have : Eδᶜ = {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ |empMean A R' a s ω - armMean a ω|} := by ext ω; simp only [Eδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push_neg; rfl rw [this] - exact prob_concentration_fail_delta (hK := hK) (E' := E') (A := A) (R' := R') + exact prob_concentration_fail_delta (hK := hK) (E := E) (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) h hσ2 hs n δ hδ hδ1 have hm_emp : ∀ a s, Measurable (fun ω ↦ empMean A R' a s ω) := fun a s ↦ measurable_empMean (fun n ↦ h.measurable_A n) (fun n ↦ h.measurable_R n) a s @@ -1265,39 +1230,26 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} := by ext ω; simp only [Set.mem_setOf_eq, Set.mem_union, Nat.cast_eq_zero]; tauto rw [this] - exact MeasurableSet.union (hm_pc a s (measurableSet_singleton (0 : ℝ))) - (measurableSet_lt - ((hm_emp a s).sub - (IsBayesAlgEnvSeq.measurable_actionMean (a := a) h.measurable_E)).abs - ((measurable_const.div (hm_pc a s)).sqrt)) + exact .union (hm_pc a s (measurableSet_singleton _)) + (measurableSet_lt (by fun_prop) (by fun_prop)) have hEδ_meas : MeasurableSet Eδ := by simp only [Eδ, Set.setOf_forall] exact .iInter fun s ↦ .iInter fun _ ↦ .iInter fun a ↦ h_arm_meas s a have hFδ_meas : MeasurableSet Fδ := by simp only [Fδ, Set.setOf_forall] - apply MeasurableSet.iInter; intro s - apply MeasurableSet.iInter; intro _ - have : {ω : Ω | pullCount A (bestArm ω) s ω ≠ 0 → - |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω| - < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A (bestArm ω) s ω))} = - ⋃ a : Fin K, (bestArm ⁻¹' {a}) ∩ {ω | pullCount A a s ω ≠ 0 → - |empMean A R' a s ω - armMean a ω| - < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} := by - ext ω - simp only [Set.mem_iUnion, Set.mem_inter_iff, Set.mem_preimage, - Set.mem_singleton_iff, Set.mem_setOf_eq] - constructor - · exact fun h => ⟨_, rfl, h⟩ - · rintro ⟨_, rfl, h⟩; exact h - rw [this] - exact .iUnion fun a => .inter (hm_best (measurableSet_singleton a)) (h_arm_meas s a) + refine .iInter fun s ↦ .iInter fun _ ↦ ?_ + convert MeasurableSet.iUnion fun a ↦ + (hm_best (measurableSet_singleton a)).inter (h_arm_meas s a) using 1 + ext ω; simp only [Set.mem_iUnion, Set.mem_inter_iff, Set.mem_preimage, + Set.mem_singleton_iff, Set.mem_setOf_eq] + exact ⟨fun h => ⟨_, rfl, h⟩, fun ⟨_, rfl, h⟩ => h⟩ have h_prob_F : P Fδᶜ ≤ ENNReal.ofReal (2 * ↑n * δ) := by have : Fδᶜ = {ω | ∃ s < n, pullCount A (bestArm ω) s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A (bestArm ω) s ω : ℝ)) ≤ |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω|} := by ext ω; simp only [Fδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push_neg; rfl rw [this] - exact prob_concentration_bestArm_fail_delta (hK := hK) (E' := E') (A := A) (R' := R') + exact prob_concentration_bestArm_fail_delta (hK := hK) (E := E) (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) h hσ2 hs n δ hδ hδ1 rw [h_swap] set f1 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, @@ -1336,12 +1288,12 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp measureReal_nonneg (μ := P) (s := Fδᶜ), measureReal_nonneg (μ := P) (s := Eδᶜ)] -lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E' A R' P) +lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (t : ℕ) : - P[IsBayesAlgEnvSeq.regret κ E' A t] + P[IsBayesAlgEnvSeq.regret κ E A t] ≤ (3 * K + 2) * (hi - lo) + 8 * √(σ2 * K * t * Real.log t) := by have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _) have hlo : lo ≤ hi := h1.trans h2 @@ -1352,7 +1304,7 @@ lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] by_cases ht1_eq : t = 1 · subst ht1_eq simp only [Nat.cast_one, Real.log_one, mul_zero, Real.sqrt_zero, mul_zero, add_zero] - calc P[IsBayesAlgEnvSeq.regret κ E' A 1] + calc P[IsBayesAlgEnvSeq.regret κ E A 1] ≤ hi - lo := by unfold IsBayesAlgEnvSeq.regret Bandits.regret simp only [Finset.range_one, Finset.sum_singleton, Nat.cast_one, one_mul, @@ -1360,8 +1312,8 @@ lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] refine (integral_mono_of_nonneg (ae_of_all _ fun ω ↦ sub_nonneg.mpr (le_ciSup ⟨hi, by rintro _ ⟨a, rfl⟩; exact (hm a _).2⟩ _)) (integrable_const (hi - lo)) (ae_of_all _ fun ω ↦ by - linarith [ciSup_le fun a ↦ (hm a (E' ω)).2, - (hm (A 0 ω) (E' ω)).1])).trans ?_ + linarith [ciSup_le fun a ↦ (hm a (E ω)).2, + (hm (A 0 ω) (E ω)).1])).trans ?_ simp _ ≤ (3 * ↑K + 2) * (hi - lo) := by nlinarith [show (1 : ℝ) ≤ K from Nat.one_le_cast.mpr (Nat.one_le_of_lt hK), @@ -1383,10 +1335,10 @@ lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] -- Second term simplification: log(1/(1/t²)) = log(t²) = 2 log(t) have h_log : Real.log (1 / (1 / (↑t : ℝ) ^ 2)) = 2 * Real.log ↑t := by rw [one_div_one_div, Real.log_pow]; norm_cast - calc P[IsBayesAlgEnvSeq.regret κ E' A t] + calc P[IsBayesAlgEnvSeq.regret κ E A t] ≤ (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * ↑t ^ 2 * (1 / (↑t) ^ 2) + 2 * √(8 * ↑σ2 * Real.log (1 / (1 / (↑t) ^ 2))) * √(↑K * ↑t) := - bayesRegret_le_of_delta (hK := hK) (E' := E') (A := A) (R' := R') (Q := Q) + bayesRegret_le_of_delta (hK := hK) (E := E) (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) h hσ2 hs hm t (1 / (↑t) ^ 2) hδ hδ1 _ = (3 * ↑K + 2) * (hi - lo) + 8 * (√(↑σ2 * Real.log ↑t) * √(↑K * ↑t)) := by rw [h_first, h_log, @@ -1399,6 +1351,8 @@ lemma TS.bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (Real.log_nonneg (Nat.one_le_cast.mpr (Nat.pos_of_ne_zero ht))))] congr 1; congr 1; congr 1; ring +end TS + end Regret end Bandits From 31dc8606c829092ae141dc73dabc6e106bbb72ef Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 10 Mar 2026 10:51:48 +0000 Subject: [PATCH 070/155] Refactor TS.lean (in progress) --- LeanBandits/Bandit/SumRewards.lean | 126 ++++ LeanBandits/BanditAlgorithms/TS.lean | 836 ++++++++++----------------- 2 files changed, 442 insertions(+), 520 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index a80151ce..1b98cff0 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -559,6 +559,132 @@ lemma probReal_sumRewards_le_sumRewards_le [Fintype α] (h : IsAlgEnvSeq A R alg section Subgaussian +/-! ### Sub-Gaussian concentration (δ-parameterized) -/ + +private lemma exp_neg_sq_div_eq_delta {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : + ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / + (2 * k * ↑σ2))) = ENNReal.ofReal δ := by + have hk_pos : (0 : ℝ) < k := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hk) + have hσ2_pos : (0 : ℝ) < ↑σ2 := NNReal.coe_pos.mpr (pos_iff_ne_zero.mpr hσ2) + have hlog : 0 < Real.log (1 / δ) := + Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) + rw [Real.sq_sqrt (by positivity)] + simp only [neg_div, Real.exp_neg] + rw [show 2 * (k : ℝ) * ↑σ2 * Real.log (1 / δ) / (2 * k * ↑σ2) = + Real.log (1 / δ) from by field_simp [ne_of_gt hσ2_pos, ne_of_gt hk_pos]] + rw [Real.exp_log (by positivity : (0 : ℝ) < 1 / δ), one_div, inv_inv] + +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in +/-- Claude: δ-parameterized one-sided concentration for the stream measure. Setting `δ = 1/(n+1)^c` +recovers `todo` and `todo'` (case-split on `c = 0`) -/ +lemma streamMeasure_concentration_le_delta {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) + (a : α) (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : + streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + + √(2 * ↑σ2 * Real.log (1 / δ) / k) ≤ (ν a)[id]} ≤ + ENNReal.ofReal δ := by + have hlog : 0 < Real.log (1 / δ) := + Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) + calc + streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + + √(2 * ↑σ2 * Real.log (1 / δ) / k) ≤ (ν a)[id]} + _ = streamMeasure ν + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) / k ≤ + -√(2 * ↑σ2 * Real.log (1 / δ) / k)} := by + congr with ω + field_simp + rw [Finset.sum_sub_distrib] + simp + grind + _ = streamMeasure ν + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ + -√(2 * k * ↑σ2 * Real.log (1 / δ))} := by + congr with ω + field_simp + congr! 2 + rw [Real.sqrt_div (by positivity : 0 ≤ 2 * ↑σ2 * Real.log (1 / δ)), + show ↑k * 2 * ↑σ2 * Real.log (1 / δ) = ↑k * (2 * ↑σ2 * Real.log (1 / δ)) from by ring, + Real.sqrt_mul (by positivity : (0 : ℝ) ≤ ↑k), ← mul_div_assoc, + mul_div_right_comm, Real.div_sqrt] + _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / + (2 * k * ↑σ2))) := by + rw [← ofReal_measureReal] + gcongr + refine HasSubgaussianMGF.measure_sum_range_le_le_of_iIndepFun (c := σ2) ?_ ?_ + (by positivity) + · exact (iIndepFun_eval_streamMeasure'' ν a).comp + (fun i ω ↦ ω - (ν a)[id]) (fun _ ↦ by fun_prop) + · intro i _; exact (hν a).congr_identDistrib + ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) + _ = ENNReal.ofReal δ := exp_neg_sq_div_eq_delta hσ2 k hk δ hδ hδ1 + +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in +lemma streamMeasure_concentration_ge_delta {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) + (a : α) (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : + streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - + √(2 * ↑σ2 * Real.log (1 / δ) / k)} ≤ + ENNReal.ofReal δ := by + have hlog : 0 < Real.log (1 / δ) := + Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) + calc + streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - + √(2 * ↑σ2 * Real.log (1 / δ) / k)} + _ = streamMeasure ν + {ω | √(2 * ↑σ2 * Real.log (1 / δ) / k) ≤ + (∑ s ∈ range k, (ω s a - (ν a)[id])) / k} := by + congr with ω + field_simp + rw [Finset.sum_sub_distrib] + simp + grind + _ = streamMeasure ν + {ω | √(2 * k * ↑σ2 * Real.log (1 / δ)) ≤ + (∑ s ∈ range k, (ω s a - (ν a)[id]))} := by + congr with ω + field_simp + congr! 1 + rw [Real.sqrt_div (by positivity : 0 ≤ 2 * ↑σ2 * Real.log (1 / δ)), + show 2 * ↑σ2 * Real.log (1 / δ) * ↑k = ↑k * (2 * ↑σ2 * Real.log (1 / δ)) from by ring, + Real.sqrt_mul (by positivity : (0 : ℝ) ≤ ↑k), ← mul_div_assoc, + mul_div_right_comm, Real.div_sqrt] + _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / + (2 * k * ↑σ2))) := by + rw [← ofReal_measureReal] + gcongr + refine HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := σ2) ?_ ?_ + (by positivity) + · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) + (fun _ ↦ by fun_prop) + · intro i _; exact (hν a).congr_identDistrib + ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) + _ = ENNReal.ofReal δ := exp_neg_sq_div_eq_delta hσ2 k hk δ hδ hδ1 + +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in +lemma streamMeasure_concentration_bound {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) + (a : α) {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) (m : ℕ) (hm : m ≠ 0) : + streamMeasure ν {ω : ℕ → α → ℝ | ∑ i ∈ range m, ω i a ∈ + {x | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} ∪ + {x | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)}} ≤ + ENNReal.ofReal (2 * δ) := + calc streamMeasure ν {ω : ℕ → α → ℝ | ∑ i ∈ range m, ω i a ∈ + {x | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} ∪ + {x | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)}} + ≤ streamMeasure ν {ω | (∑ i ∈ range m, ω i a) / m + + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} + + streamMeasure ν {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - + √(2 * ↑σ2 * Real.log (1 / δ) / m)} := by + apply (measure_mono (fun ω hω ↦ ?_)).trans (measure_union_le _ _) + simp only [Set.mem_setOf_eq, Set.mem_union] at hω ⊢; exact hω + _ ≤ ENNReal.ofReal δ + ENNReal.ofReal δ := by + gcongr + · exact streamMeasure_concentration_le_delta hσ2 hν a m hm δ hδ hδ1 + · exact streamMeasure_concentration_ge_delta hσ2 hν a m hm δ hδ hδ1 + _ = ENNReal.ofReal (2 * δ) := by + rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf + omit [DecidableEq α] [StandardBorelSpace α] in lemma probReal_sum_le_sum_streamMeasure [Fintype α] {c : ℝ≥0} (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) c (ν a)) (a : α) (m : ℕ) : diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 9c568263..807e0889 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -3,12 +3,9 @@ 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.SubGaussian +import LeanBandits.Bandit.SumRewards import LeanBandits.BanditAlgorithms.Uniform -import LeanBandits.BanditAlgorithms.UCB -import LeanBandits.SequentialLearning.BayesStationaryEnv import LeanBandits.SequentialLearning.AlgorithmDensity -import Mathlib.Analysis.Complex.ExponentialBounds /-! # The Thompson Sampling Algorithm -/ @@ -21,7 +18,7 @@ namespace Bandits namespace TS variable {K : ℕ} (hK : 0 < K) -variable {𝓔 : Type*} [mE : MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] +variable {𝓔 : Type*} [m𝓔 : MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] variable (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] /-- The distribution over actions for every given history for TS. -/ @@ -67,20 +64,10 @@ def tsAlgorithm (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] end Algorithm -section Regret - -variable {𝓔 : Type*} [m𝓔 : MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] -variable (hK : 0 < K) -variable {Ω : Type*} [MeasurableSpace Ω] -variable (E : Ω → 𝓔) (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → ℝ) -variable (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] -variable (P : Measure Ω) [IsProbabilityMeasure P] - namespace TS -/-! ### Auxiliary real-analysis lemmas (candidates for migration to a utility file) -/ +/-! ### Auxiliary real-analysis lemmas -/ -omit m𝓔 [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] in lemma abs_sub_le_of_mem_Icc {lo hi x y : ℝ} (hx : x ∈ Set.Icc lo hi) (hy : y ∈ Set.Icc lo hi) : |x - y| ≤ hi - lo := by @@ -93,7 +80,6 @@ lemma sum_sqrt_le {ι : Type*} (s : Finset ι) (c : ι → ℝ) (hc : ∀ i, 0 calc ∑ i ∈ s, √(c i) ≤ √(∑ i ∈ s, c i) * √↑(#s) := h _ = _ := by rw [← Real.sqrt_mul (Finset.sum_nonneg (fun i _ => hc i)), mul_comm] -omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] in lemma sum_inv_sqrt_le (M : ℕ) (hM : 0 < M) : ∑ j ∈ range M, (1 / √(↑j : ℝ)) + 1 / √↑M ≤ 2 * √↑M := by induction M with @@ -111,7 +97,14 @@ lemma sum_inv_sqrt_le (M : ℕ) (hM : 0 < M) : mul_self_nonneg (√(↑(n + 1) : ℝ) - √(↑n : ℝ)), show (↑(n + 1) : ℝ) = ↑n + 1 from by push_cast; ring] -/-! ### UCB index definition and properties -/ +/-! ### UCB index: definition and deterministic bounds -/ + +section Deterministic + +variable {Ω : Type*} +variable {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} +variable {𝓔 : Type*} [MeasurableSpace 𝓔] +variable (E : Ω → 𝓔) (κ : Kernel (𝓔 × Fin K) ℝ) noncomputable def ucbIndex (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (σ2 lo hi δ : ℝ) @@ -121,7 +114,6 @@ def ucbIndex (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (σ2 lo hi δ (empMean A R' a t ω + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ)))) -omit m𝓔 [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] in lemma ucbIndex_mem_Icc (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : ℕ) (ω : Ω) : ucbIndex A R' σ2 lo hi δ a t ω ∈ Set.Icc lo hi := by unfold ucbIndex @@ -132,19 +124,52 @@ lemma ucbIndex_mem_Icc (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : · exact max_le hlo (min_le_left hi _) @[fun_prop] -lemma measurable_ucbIndex [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) - (σ2 lo hi δ : ℝ) (a : Fin K) (t : ℕ) : +lemma measurable_ucbIndex [MeasurableSpace Ω] + {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} + {σ2 lo hi δ : ℝ} {a : Fin K} {t : ℕ} + (hA : ∀ n, Measurable (A n)) (hR : ∀ n, Measurable (R' n)) : Measurable (ucbIndex A R' σ2 lo hi δ a t) := by unfold ucbIndex have : Measurable (fun ω ↦ (pullCount A a t ω : ℝ)) := - measurable_from_top.comp (measurable_pullCount (fun n ↦ h.measurable_A n) a t) - have := measurable_empMean (fun n ↦ h.measurable_A n) (fun n ↦ h.measurable_R n) a t - have := measurable_pullCount (fun n ↦ h.measurable_A n) a t - exact .ite ((measurable_pullCount (fun n ↦ h.measurable_A n) a t) + measurable_from_top.comp (measurable_pullCount hA a t) + have := measurable_empMean hA hR a t + have := measurable_pullCount hA a t + exact .ite ((measurable_pullCount hA a t) (measurableSet_singleton 0)) measurable_const (by fun_prop) -omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] [IsMarkovKernel κ] in +/-- History-based UCB index: like `ucbIndex` but takes a history `h : Iic t → Fin K × ℝ` +directly instead of the random variables `A` and `R'`. -/ +noncomputable +def ucbIndex' (σ2 lo hi δ : ℝ) (t : ℕ) (h : Iic t → Fin K × ℝ) (a : Fin K) : ℝ := + if pullCount' t h a = 0 then hi + else max lo (min hi (empMean' t h a + + √(2 * σ2 * Real.log (1 / δ) / (pullCount' t h a : ℝ)))) + +@[fun_prop] +lemma measurable_ucbIndex' {σ2 lo hi δ : ℝ} {a : Fin K} {t : ℕ} : + Measurable (fun h : Iic t → Fin K × ℝ ↦ ucbIndex' σ2 lo hi δ t h a) := by + unfold ucbIndex' + apply Measurable.ite + · have : MeasurableSet {h : Iic t → Fin K × ℝ | (pullCount' t h a : ℝ) = (0 : ℝ)} := + measurableSet_eq_fun + (measurable_from_top.comp (measurable_pullCount' t a)) + measurable_const + simp only [Nat.cast_eq_zero] at this; exact this + · exact measurable_const + · exact .max measurable_const (.min measurable_const + (.add (measurable_empMean' t a) + (measurable_const.div (measurable_from_top.comp (measurable_pullCount' t a))).sqrt)) + +lemma ucbIndex_succ_eq_ucbIndex'_hist (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) + (σ2 lo hi δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : + ucbIndex A R' σ2 lo hi δ a (t + 1) ω = + ucbIndex' σ2 lo hi δ t (IsAlgEnvSeq.hist A R' t ω) a := by + have hpc : pullCount A a (t + 1) ω = pullCount' t (IsAlgEnvSeq.hist A R' t ω) a := + pullCount_add_one_eq_pullCount' + have hem : empMean A R' a (t + 1) ω = empMean' t (IsAlgEnvSeq.hist A R' t ω) a := + empMean_add_one_eq_empMean' + simp only [ucbIndex, ucbIndex', hpc, hem] + lemma armMean_le_ucbIndex {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hconc : pullCount A a t ω ≠ 0 → @@ -160,7 +185,6 @@ lemma armMean_le_ucbIndex {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set. refine le_max_of_le_right (le_min hmean.2 ?_) linarith [habs.2] -omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hpc : pullCount A a t ω ≠ 0) (hconc : @@ -179,30 +203,6 @@ lemma ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ ( max_le_iff.mpr ⟨by linarith [hmean.1, habs.2], min_le_right _ _⟩ linarith [habs.2] -lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) (t : ℕ) : - condDistrib (A (t + 1)) (IsAlgEnvSeq.hist A R' t) P - =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' t) P := - by - have h_ba_comp : IsBayesAlgEnvSeq.bestAction κ E - = IsBayesAlgEnvSeq.bestAction κ id ∘ E := rfl - rw [h_ba_comp] - have hm := IsBayesAlgEnvSeq.measurable_bestAction (κ := κ) measurable_id - have h_comp := condDistrib_comp (mβ := MeasurableSpace.pi) (μ := P) - (IsAlgEnvSeq.hist A R' t) h.measurable_E.aemeasurable hm - have h_map : (condDistrib E (IsAlgEnvSeq.hist A R' t) P).map - (IsBayesAlgEnvSeq.bestAction κ id) =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) t).map - (IsBayesAlgEnvSeq.bestAction κ id) := by - filter_upwards [(h.hasCondDistrib_env_hist - (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) - (absolutelyContinuous_uniformAlgorithm hK _) t).condDistrib_eq] - with x hx - simp only [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hx] - exact (h.hasCondDistrib_action' t).condDistrib_eq.trans (h_comp.trans h_map).symm - -omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) (ω : Ω) : ⨆ i, IsBayesAlgEnvSeq.actionMean κ E i ω = @@ -213,7 +213,6 @@ lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] {lo hi : ℝ} (le_ciSup (f := fun i ↦ IsBayesAlgEnvSeq.actionMean κ E i ω) ⟨hi, by rintro _ ⟨i, rfl⟩; exact (hm i _).2⟩ _) -omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] [IsMarkovKernel κ] in lemma gap_eq_armMean_sub [Nonempty (Fin K)] {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) (s : ℕ) (ω : Ω) : gap (κ.sectR (E ω)) (A s ω) = @@ -222,33 +221,6 @@ lemma gap_eq_armMean_sub [Nonempty (Fin K)] {lo hi : ℝ} simp only [gap, Kernel.sectR_apply] exact congr_arg (· - _) (iSup_armMean_eq_bestArm E κ hm ω) -omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [IsProbabilityMeasure Q] [IsMarkovKernel κ] in -lemma bayesRegret_eq_sum_integral_gap [Nonempty (Fin K)] - {alg : Algorithm (Fin K) ℝ} - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - {C : ℝ} (hm : ∀ a e, |(κ (e, a))[id]| ≤ C) (t : ℕ) : - P[IsBayesAlgEnvSeq.regret κ E A t] = - ∑ s ∈ range t, P[fun ω ↦ gap (κ.sectR (E ω)) - (A s ω)] := by - simp only [IsBayesAlgEnvSeq.regret, regret_eq_sum_gap] - refine integral_finset_sum _ (fun s _ => ?_) - have hmeas : Measurable (fun ω ↦ gap (κ.sectR (E ω)) - (A s ω)) := - (Measurable.iSup (fun a ↦ IsBayesAlgEnvSeq.measurable_actionMean - (a := a) h.measurable_E)).sub - (stronglyMeasurable_id.integral_kernel.measurable.comp - (h.measurable_E.prodMk (h.measurable_A s))) - refine ⟨hmeas.aestronglyMeasurable, HasFiniteIntegral.of_bounded (C := 2 * C) - (Filter.Eventually.of_forall fun ω => ?_)⟩ - simp only [Real.norm_eq_abs, gap, Kernel.sectR_apply] - have hbdd : BddAbove (Set.range fun i => (κ (E ω, i))[id]) := - ⟨C, by rintro _ ⟨i, rfl⟩; exact le_of_abs_le (hm i _)⟩ - rw [abs_of_nonneg (sub_nonneg.mpr (le_ciSup hbdd _))] - linarith [ciSup_le fun i => le_of_abs_le (hm i (E ω)), - neg_le_of_abs_le (hm (A s ω) (E ω))] - -omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] [IsProbabilityMeasure Q] - [IsMarkovKernel κ] [IsProbabilityMeasure P] in lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : ∑ s ∈ range n, f (pullCount A (A s ω) s ω) = ∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω), f j := by @@ -271,8 +243,6 @@ lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : congr 1 simp -omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] [MeasurableSpace Ω] [IsProbabilityMeasure Q] - [IsMarkovKernel κ] [IsProbabilityMeasure P] in lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (hlo : lo ≤ hi) (σ2 δ : ℝ) (n : ℕ) (ω : Ω) (hconc : ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → @@ -304,7 +274,7 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) := sum_le_sum fun s hs => by have hpc : pullCount A (A s ω) s ω ≠ 0 := (Finset.mem_filter.mp hs).2 - exact ucbIndex_sub_armMean_le E A R' κ hm σ2 δ (A s ω) s ω hpc + exact ucbIndex_sub_armMean_le E κ hm σ2 δ (A s ω) s ω hpc (hconc s (mem_range.mp (Finset.mem_filter.mp hs).1) _ hpc) _ ≤ ∑ s ∈ range n, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) := @@ -327,7 +297,7 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] rw [mul_sum] _ = √(8 * σ2 * c) * ∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω), (1 / √(↑j : ℝ)) := by - congr 1; exact sum_comp_pullCount A (fun j => 1 / √(↑j : ℝ)) n ω + congr 1; exact sum_comp_pullCount (fun j => 1 / √(↑j : ℝ)) n ω _ ≤ √(8 * σ2 * c) * ∑ a : Fin K, (2 * √↑(pullCount A a n ω)) := by gcongr with a by_cases ha : pullCount A a n ω = 0 @@ -383,314 +353,39 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ≤ ∑ _s ∈ S0, (hi - lo) := sum_le_sum hterm_S0 _ = #S0 * (hi - lo) := by rw [sum_const, nsmul_eq_mul] - _ ≤ ↑K * (hi - lo) := by - apply mul_le_mul_of_nonneg_right _ (by linarith) - exact_mod_cast h_card_S0 + _ ≤ ↑K * (hi - lo) := by gcongr; linarith _ = (hi - lo) * ↑K := by ring -private lemma exp_neg_sq_div_eq_delta {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / - (2 * k * ↑σ2))) = ENNReal.ofReal δ := by - have hk_pos : (0 : ℝ) < k := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hk) - have hσ2_pos : (0 : ℝ) < ↑σ2 := NNReal.coe_pos.mpr (pos_iff_ne_zero.mpr hσ2) - have hlog : 0 < Real.log (1 / δ) := - Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) - rw [Real.sq_sqrt (by positivity)] - simp only [neg_div, Real.exp_neg] - rw [show 2 * (k : ℝ) * ↑σ2 * Real.log (1 / δ) / (2 * k * ↑σ2) = - Real.log (1 / δ) from by field_simp [ne_of_gt hσ2_pos, ne_of_gt hk_pos]] - rw [Real.exp_log (by positivity : (0 : ℝ) < 1 / δ), one_div, inv_inv] - -lemma streamMeasure_concentration_le_delta {α : Type*} [MeasurableSpace α] - {ν : Kernel α ℝ} [IsMarkovKernel ν] {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - (a : α) (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + - √(2 * ↑σ2 * Real.log (1 / δ) / k) ≤ (ν a)[id]} ≤ - ENNReal.ofReal δ := by - have hlog : 0 < Real.log (1 / δ) := - Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) - calc - streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + - √(2 * ↑σ2 * Real.log (1 / δ) / k) ≤ (ν a)[id]} - _ = streamMeasure ν - {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) / k ≤ - -√(2 * ↑σ2 * Real.log (1 / δ) / k)} := by - congr with ω - field_simp - rw [Finset.sum_sub_distrib] - simp - grind - _ = streamMeasure ν - {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ - -√(2 * k * ↑σ2 * Real.log (1 / δ))} := by - congr with ω - field_simp - congr! 2 - rw [Real.sqrt_div (by positivity : 0 ≤ 2 * ↑σ2 * Real.log (1 / δ)), - show ↑k * 2 * ↑σ2 * Real.log (1 / δ) = ↑k * (2 * ↑σ2 * Real.log (1 / δ)) from by ring, - Real.sqrt_mul (by positivity : (0 : ℝ) ≤ ↑k), ← mul_div_assoc, - mul_div_right_comm, Real.div_sqrt] - _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / - (2 * k * ↑σ2))) := by - rw [← ofReal_measureReal] - gcongr - refine HasSubgaussianMGF.measure_sum_range_le_le_of_iIndepFun (c := σ2) ?_ ?_ - (by positivity) - · exact (iIndepFun_eval_streamMeasure'' ν a).comp - (fun i ω ↦ ω - (ν a)[id]) (fun _ ↦ by fun_prop) - · intro i _; exact (hν a).congr_identDistrib - ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) - _ = ENNReal.ofReal δ := exp_neg_sq_div_eq_delta hσ2 k hk δ hδ hδ1 - -lemma streamMeasure_concentration_ge_delta {α : Type*} [MeasurableSpace α] - {ν : Kernel α ℝ} [IsMarkovKernel ν] {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - (a : α) (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - - √(2 * ↑σ2 * Real.log (1 / δ) / k)} ≤ - ENNReal.ofReal δ := by - have hlog : 0 < Real.log (1 / δ) := - Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) - calc - streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - - √(2 * ↑σ2 * Real.log (1 / δ) / k)} - _ = streamMeasure ν - {ω | √(2 * ↑σ2 * Real.log (1 / δ) / k) ≤ - (∑ s ∈ range k, (ω s a - (ν a)[id])) / k} := by - congr with ω - field_simp - rw [Finset.sum_sub_distrib] - simp - grind - _ = streamMeasure ν - {ω | √(2 * k * ↑σ2 * Real.log (1 / δ)) ≤ - (∑ s ∈ range k, (ω s a - (ν a)[id]))} := by - congr with ω - field_simp - congr! 1 - rw [Real.sqrt_div (by positivity : 0 ≤ 2 * ↑σ2 * Real.log (1 / δ)), - show 2 * ↑σ2 * Real.log (1 / δ) * ↑k = ↑k * (2 * ↑σ2 * Real.log (1 / δ)) from by ring, - Real.sqrt_mul (by positivity : (0 : ℝ) ≤ ↑k), ← mul_div_assoc, - mul_div_right_comm, Real.div_sqrt] - _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / - (2 * k * ↑σ2))) := by - rw [← ofReal_measureReal] - gcongr - refine HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := σ2) ?_ ?_ - (by positivity) - · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) - (fun _ ↦ by fun_prop) - · intro i _; exact (hν a).congr_identDistrib - ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) - _ = ENNReal.ofReal δ := exp_neg_sq_div_eq_delta hσ2 k hk δ hδ hδ1 - -private lemma streamMeasure_concentration_bound {α : Type*} [MeasurableSpace α] - {ν : Kernel α ℝ} [IsMarkovKernel ν] {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - (a : α) {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) (m : ℕ) (hm : m ≠ 0) : - streamMeasure ν {ω : ℕ → α → ℝ | ∑ i ∈ range m, ω i a ∈ - {x | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} ∪ - {x | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)}} ≤ - ENNReal.ofReal (2 * δ) := - calc streamMeasure ν {ω : ℕ → α → ℝ | ∑ i ∈ range m, ω i a ∈ - {x | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} ∪ - {x | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)}} - ≤ streamMeasure ν {ω | (∑ i ∈ range m, ω i a) / m + - √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} + - streamMeasure ν {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - - √(2 * ↑σ2 * Real.log (1 / δ) / m)} := by - apply (measure_mono (fun ω hω ↦ ?_)).trans (measure_union_le _ _) - simp only [Set.mem_setOf_eq, Set.mem_union] at hω ⊢; exact hω - _ ≤ ENNReal.ofReal δ + ENNReal.ofReal δ := by - gcongr - · exact streamMeasure_concentration_le_delta hσ2 hν a m hm δ hδ hδ1 - · exact streamMeasure_concentration_ge_delta hσ2 hν a m hm δ hδ hδ1 - _ = ENNReal.ofReal (2 * δ) := by - rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf - -lemma prob_concentration_single_delta_cond [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) - {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) - (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) - (hδ_large : max |lo| |hi| < √(2 * ↑σ2 * Real.log (1 / δ))) : - ∀ᵐ e ∂(P.map (E)), - (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) - {ω | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} ≤ - ENNReal.ofReal (2 * s * δ) := by - have h_cond_ae : ∀ᵐ e ∂(P.map E), IsAlgEnvSeq IT.action IT.reward - (tsAlgorithm hK Q κ) (stationaryEnv (κ.sectR e)) - (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) := by - rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h - filter_upwards [h_cond_ae] with e h_isAlgEnvSeq - let ν := κ.sectR e - have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by - simp only [ν, Kernel.sectR_apply]; exact hs a' e - have h_mean : (ν a)[id] = (κ (e, a))[id] := by simp only [ν, Kernel.sectR_apply] - rw [← h_mean] - let P' := condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e - have h_law := h_isAlgEnvSeq.law_pullCount_sumRewards_unique' - (ArrayModel.isAlgEnvSeq_arrayMeasure (tsAlgorithm hK Q κ) ν) (n := s) - let B_low := fun m : ℕ ↦ - {x : ℝ | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} - let B_high := fun m : ℕ ↦ - {x : ℝ | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)} - have h_stream_bound : ∀ m : ℕ, m ≠ 0 → - streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ≤ - ENNReal.ofReal (2 * δ) := - fun m hm0 ↦ streamMeasure_concentration_bound hσ2 h_subG a hδ hδ1 m hm0 - let badSet := {ω : ℕ → (Fin K) × ℝ | - √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s ω) : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (ν a)[id]|} - have h_bound_per_m : ∀ m : ℕ, m ≠ 0 → m ≤ s → - P' {ω | pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} ≤ - streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := by - intro m hm0 hms - have hB_meas : MeasurableSet (B_low m ∪ B_high m) := - MeasurableSet.union (measurableSet_le (by fun_prop) (by fun_prop)) - (measurableSet_le (by fun_prop) (by fun_prop)) - exact prob_pullCount_eq_and_sumRewards_mem_le h_isAlgEnvSeq hms hB_meas - have h_bad_subset : badSet ⊆ - ⋃ m ∈ (Finset.range (s + 1)).filter (· ≠ 0), - {ω | pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by - intro ω hω - simp only [Set.mem_setOf_eq, badSet] at hω - simp only [Set.mem_iUnion, Finset.mem_filter, Finset.mem_range, Set.mem_setOf_eq] - set m := pullCount IT.action a s ω with hm_def - have hms : m ≤ s := pullCount_le (A := IT.action) a s ω - by_cases hm0 : m = 0 - · exfalso - have h_empMean_zero : empMean IT.action IT.reward a s ω = 0 := by - simp only [empMean, ← hm_def, hm0, Nat.cast_zero, div_zero] - simp only [hm0, Nat.cast_zero, h_empMean_zero, max_eq_left (zero_le_one' ℝ), div_one] at hω - rw [h_mean, zero_sub, abs_neg] at hω - linarith [abs_le_max_abs_abs (hm a e).1 (hm a e).2] - · -- Case: m ≥ 1 - use m - refine ⟨⟨Nat.lt_succ_of_le hms, hm0⟩, rfl, ?_⟩ - simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] - have hmax_eq : (max 1 (m : ℕ) : ℝ) = m := by - simp only [Nat.one_le_cast, Nat.one_le_iff_ne_zero.mpr hm0, max_eq_right] - rw [hmax_eq] at hω - rw [show empMean IT.action IT.reward a s ω = - sumRewards IT.action IT.reward a s ω / m from by simp only [empMean, hm_def]] at hω - by_cases h_le : sumRewards IT.action IT.reward a s ω / m ≤ (ν a)[id] - · left; rw [abs_of_nonpos (sub_nonpos.mpr h_le), neg_sub] at hω; linarith - · right; rw [abs_of_pos (sub_pos.mpr (not_le.mp h_le))] at hω; linarith - calc P' badSet - ≤ P' (⋃ m ∈ (Finset.range (s + 1)).filter (· ≠ 0), - {ω | pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m}) := - measure_mono h_bad_subset - _ ≤ ∑ m ∈ (Finset.range (s + 1)).filter (· ≠ 0), - P' {ω | pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := - measure_biUnion_finset_le _ _ - _ ≤ ∑ m ∈ (Finset.range (s + 1)).filter (· ≠ 0), - streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := - Finset.sum_le_sum fun m hm ↦ h_bound_per_m m (Finset.mem_filter.mp hm).2 - (Nat.lt_succ_iff.mp (Finset.mem_range.mp (Finset.mem_filter.mp hm).1)) - _ ≤ ∑ _m ∈ (Finset.range (s + 1)).filter (· ≠ 0), ENNReal.ofReal (2 * δ) := - Finset.sum_le_sum fun m hm ↦ h_stream_bound m (Finset.mem_filter.mp hm).2 - _ = ((Finset.range (s + 1)).filter (· ≠ 0)).card • ENNReal.ofReal (2 * δ) := by - simp only [Finset.sum_const] - _ = s • ENNReal.ofReal (2 * δ) := by - congr 1 - have hS_eq : (Finset.range (s + 1)).filter (· ≠ 0) = Finset.Icc 1 s := by - ext m; simp only [Finset.mem_filter, Finset.mem_range, ne_eq, Finset.mem_Icc]; omega - rw [hS_eq, Nat.card_Icc, Nat.add_sub_cancel] - _ = ENNReal.ofReal (2 * s * δ) := by - rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast s, ← ENNReal.ofReal_mul (Nat.cast_nonneg s)] - congr 1; ring +end Deterministic -lemma prob_concentration_single_delta [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) - {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) - (a : Fin K) (s : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) - (hδ_large : max |lo| |hi| < √(2 * ↑σ2 * Real.log (1 / δ))) : - P {ω | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount A a s ω) : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} ≤ - ENNReal.ofReal (2 * s * δ) := by - let badSet : 𝓔 → Set (ℕ → (Fin K) × ℝ) := fun e ↦ - {t | √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s t) : ℝ)) ≤ - |empMean IT.action IT.reward a s t - (κ (e, a))[id]|} - have h_set_eq : {ω | √(2 * ↑σ2 * Real.log (1 / δ) / - (max 1 (pullCount A a s ω) : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} = - (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' - {p | p.2 ∈ badSet p.1} := by - ext ω - simp only [Set.mem_setOf_eq, Set.mem_preimage, badSet, IsBayesAlgEnvSeq.actionMean] - have h1 : pullCount A a s ω = pullCount IT.action a s (IsBayesAlgEnvSeq.trajectory A R' ω) := by - unfold pullCount IT.action; rfl - have h2 : empMean A R' a s ω = - empMean IT.action IT.reward a s (IsBayesAlgEnvSeq.trajectory A R' ω) := by - unfold empMean IT.action IT.reward; rfl - rw [h1, h2] - have h_meas_pair : - Measurable (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) := - h.measurable_E.prodMk (IsBayesAlgEnvSeq.measurable_trajectory h.measurable_A h.measurable_R) - have h_disint : P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) = - P.map (E) ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P := - (compProd_map_condDistrib - (IsBayesAlgEnvSeq.measurable_trajectory - h.measurable_A h.measurable_R).aemeasurable).symm - have h_cond := prob_concentration_single_delta_cond hK E A R' Q κ P h hσ2 hs hm a s δ hδ hδ1 - hδ_large - have h_kernel : Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := - stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_fst.prodMk measurable_const) - have h_meas_set : MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ badSet p.1} := by - change MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | - √(2 * ↑σ2 * Real.log (1 / δ) / (max 1 (pullCount IT.action a s p.2) : ℝ)) ≤ - |empMean IT.action IT.reward a s p.2 - (κ (p.1, a))[id]|} - exact measurableSet_le (by fun_prop) - (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp measurable_snd).sub - h_kernel).abs - calc P _ = P ((fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' - {p | p.2 ∈ badSet p.1}) := by rw [h_set_eq] - _ = (P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω))) - {p | p.2 ∈ badSet p.1} := by - rw [Measure.map_apply h_meas_pair h_meas_set] - _ = (P.map (E) ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P) - {p | p.2 ∈ badSet p.1} := by rw [h_disint] - _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) - (badSet e) ∂(P.map (E)) := by - rw [Measure.compProd_apply h_meas_set]; rfl - _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * s * δ) ∂(P.map (E)) := by - apply lintegral_mono_ae - filter_upwards [h_cond] with e h_e; exact h_e - _ = ENNReal.ofReal (2 * s * δ) := by - rw [lintegral_const, Measure.map_apply h.measurable_E MeasurableSet.univ] - simp [measure_univ] +/-! ### Concentration bounds (algorithm-generic) -private lemma concentration_cond_bound [Nonempty (Fin K)] - {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) +These lemmas take `{alg : Algorithm (Fin K) ℝ}` and +`(h : IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv ν) P')` — +they hold for any algorithm, not just TS. Placed in the `TS` namespace +following the convention of UCB.lean and ETC.lean (cf. `UCB.prob_ucbIndex_le`, +`ETC.probReal_sumRewards_le_sumRewards_le`). -/ + +section Concentration + +variable {K : ℕ} [Nonempty (Fin K)] +variable {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] +variable {P' : Measure (ℕ → (Fin K) × ℝ)} [IsProbabilityMeasure P'] +variable {σ2 : ℝ≥0} {alg : Algorithm (Fin K) ℝ} + +/-- Single-arm concentration bound. For any algorithm, the probability that the +empirical mean of arm `a` deviates from the true mean by more than +`√(2σ²·log(1/δ)/pullCount)` at some step before `n` is at most `2nδ`. -/ +lemma concentration_cond_bound + (hσ2 : σ2 ≠ 0) + (hs : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) {n : ℕ} (hn : 0 < n) {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) - (e : 𝓔) (h_isAlgEnvSeq : IsAlgEnvSeq IT.action IT.reward (tsAlgorithm hK Q κ) - (stationaryEnv (κ.sectR e)) - (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e)) + (h_isAlgEnvSeq : IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv ν) P') (a : Fin K) : - (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) - (⋃ s ∈ Finset.range n, {ω | pullCount IT.action a s ω ≠ 0 ∧ + P' (⋃ s ∈ Finset.range n, {ω | pullCount IT.action a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|}) ≤ + |empMean IT.action IT.reward a s ω - (ν a)[id]|}) ≤ ENNReal.ofReal (2 * n * δ) := by - let ν := κ.sectR e - let P' := condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e - have h_subG : ∀ a', HasSubgaussianMGF (fun x ↦ x - (ν a')[id]) σ2 (ν a') := fun a' ↦ by - simp only [ν, Kernel.sectR_apply]; exact hs a' e - have h_mean : (ν a)[id] = (κ (e, a))[id] := by simp only [ν, Kernel.sectR_apply] let B_low := fun m : ℕ ↦ {x : ℝ | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} let B_high := fun m : ℕ ↦ @@ -698,14 +393,14 @@ private lemma concentration_cond_bound [Nonempty (Fin K)] have h_stream_bound : ∀ m : ℕ, m ≠ 0 → streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ≤ ENNReal.ofReal (2 * δ) := - fun m hm0 ↦ streamMeasure_concentration_bound hσ2 h_subG a hδ hδ1 m hm0 + fun m hm0 ↦ streamMeasure_concentration_bound hσ2 hs a hδ hδ1 m hm0 have hB_meas : ∀ m, MeasurableSet (B_low m ∪ B_high m) := fun m ↦ MeasurableSet.union (measurableSet_le (by fun_prop) (by fun_prop)) (measurableSet_le (by fun_prop) (by fun_prop)) let badSetIT := fun (s : ℕ) ↦ {ω : ℕ → (Fin K) × ℝ | pullCount IT.action a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} + |empMean IT.action IT.reward a s ω - (ν a)[id]|} let S := Finset.Icc 1 (n - 1) have hS_card : S.card = n - 1 := by simp only [Nat.card_Icc, S]; omega have h_decomp : ⋃ s ∈ Finset.range n, badSetIT s = @@ -785,8 +480,94 @@ private lemma concentration_cond_bound [Nonempty (Fin K)] exact ENNReal.ofReal_le_ofReal (by nlinarith [(Nat.cast_le (α := ℝ)).mpr (Nat.sub_le n 1), hδ.le]) +/-- All-arms concentration bound. For any algorithm, the probability that +*some* arm's empirical mean deviates by more than the confidence width at some +step before `n` is at most `2Knδ`. -/ +lemma concentration_fail + (hσ2 : σ2 ≠ 0) + (hs : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) + (h : IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv ν) P') + (n : ℕ) {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) : + P' {ω | ∃ s < n, ∃ a, pullCount IT.action a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ + |empMean IT.action IT.reward a s ω - (ν a)[id]|} ≤ + ENNReal.ofReal (2 * K * n * δ) := by + let badSet := fun (a : Fin K) (s : ℕ) ↦ {ω : ℕ → (Fin K) × ℝ | + pullCount IT.action a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ + |empMean IT.action IT.reward a s ω - (ν a)[id]|} + have h_set_eq : {ω | ∃ s < n, ∃ a, pullCount IT.action a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ + |empMean IT.action IT.reward a s ω - (ν a)[id]|} = + ⋃ a : Fin K, ⋃ s ∈ Finset.range n, badSet a s := by + ext ω; simp only [Set.mem_setOf_eq, Set.mem_iUnion, Finset.mem_range, badSet, exists_prop] + exact ⟨fun ⟨s, hs, a, ha⟩ ↦ ⟨a, s, hs, ha⟩, fun ⟨a, s, hs, ha⟩ ↦ ⟨s, hs, a, ha⟩⟩ + rw [h_set_eq] + have h_arm_bound : ∀ a : Fin K, + P' (⋃ s ∈ Finset.range n, badSet a s) ≤ ENNReal.ofReal (2 * n * δ) := by + intro a + by_cases hn : n = 0 + · simp [hn] + exact concentration_cond_bound hσ2 hs (Nat.pos_of_ne_zero hn) hδ hδ1 h a + calc P' (⋃ a : Fin K, ⋃ s ∈ Finset.range n, badSet a s) + ≤ ∑ a : Fin K, P' (⋃ s ∈ Finset.range n, badSet a s) := + measure_iUnion_fintype_le _ _ + _ ≤ ∑ _a : Fin K, ENNReal.ofReal (2 * n * δ) := + Finset.sum_le_sum fun a _ ↦ h_arm_bound a + _ = K • ENNReal.ofReal (2 * n * δ) := by simp [Finset.sum_const] + _ = ENNReal.ofReal (2 * K * n * δ) := by + simp only [nsmul_eq_mul] + rw [← ENNReal.ofReal_natCast K, ← ENNReal.ofReal_mul (Nat.cast_nonneg K)] + congr 1; ring + +end Concentration + +end TS + +end Bandits + +open Bandits Bandits.TS + +/-! ### Algorithm-generic Bayesian lemmas -/ + +section BayesianConcentration + +variable {K : ℕ} {𝓔 : Type*} [MeasurableSpace 𝓔] {Ω : Type*} [MeasurableSpace Ω] +variable (E : Ω → 𝓔) (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → ℝ) +variable (Q : Measure 𝓔) (κ : Kernel (𝓔 × Fin K) ℝ) +variable (P : Measure Ω) [IsProbabilityMeasure P] + +namespace Learning.IsBayesAlgEnvSeq + +lemma bayesRegret_eq_sum_integral_gap [Nonempty (Fin K)] + {alg : Algorithm (Fin K) ℝ} + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) + {C : ℝ} (hm : ∀ a e, |(κ (e, a))[id]| ≤ C) (t : ℕ) : + P[IsBayesAlgEnvSeq.regret κ E A t] = + ∑ s ∈ range t, P[fun ω ↦ gap (κ.sectR (E ω)) + (A s ω)] := by + simp only [IsBayesAlgEnvSeq.regret, regret_eq_sum_gap] + refine integral_finset_sum _ (fun s _ => ?_) + have hmeas : Measurable (fun ω ↦ gap (κ.sectR (E ω)) + (A s ω)) := + (Measurable.iSup (fun a ↦ IsBayesAlgEnvSeq.measurable_actionMean + (a := a) h.measurable_E)).sub + (stronglyMeasurable_id.integral_kernel.measurable.comp + (h.measurable_E.prodMk (h.measurable_A s))) + refine ⟨hmeas.aestronglyMeasurable, HasFiniteIntegral.of_bounded (C := 2 * C) + (Filter.Eventually.of_forall fun ω => ?_)⟩ + simp only [Real.norm_eq_abs, gap, Kernel.sectR_apply] + have hbdd : BddAbove (Set.range fun i => (κ (E ω, i))[id]) := + ⟨C, by rintro _ ⟨i, rfl⟩; exact le_of_abs_le (hm i _)⟩ + rw [abs_of_nonneg (sub_nonneg.mpr (le_ciSup hbdd _))] + linarith [ciSup_le fun i => le_of_abs_le (hm i (E ω)), + neg_le_of_abs_le (hm (A s ω) (E ω))] + +variable [IsMarkovKernel κ] + lemma prob_concentration_fail_delta [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) + {alg : Algorithm (Fin K) ℝ} + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : @@ -794,106 +575,84 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} ≤ ENNReal.ofReal (2 * K * n * δ) := by - let badSet := fun (s : ℕ) (a : Fin K) ↦ {ω : Ω | pullCount A a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} + let badSetIT := fun (e : 𝓔) ↦ {ω : ℕ → (Fin K) × ℝ | + ∃ s < n, ∃ a, pullCount IT.action a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ + |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} have h_set_eq : {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} = - ⋃ s ∈ Finset.range n, ⋃ a : Fin K, badSet s a := by - ext ω; simp only [Set.mem_setOf_eq, Finset.mem_range, Set.mem_iUnion, badSet, exists_prop] + (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' + {p | p.2 ∈ badSetIT p.1} := by + ext ω + simp only [Set.mem_setOf_eq, Set.mem_preimage, badSetIT, IsBayesAlgEnvSeq.actionMean] + rfl rw [h_set_eq] - have h_reorg : ⋃ s ∈ Finset.range n, ⋃ a : Fin K, badSet s a = - ⋃ a : Fin K, ⋃ s ∈ Finset.range n, badSet s a := by - ext ω; simp only [Set.mem_iUnion, Finset.mem_range]; exact - ⟨fun ⟨s, hs, a, ha⟩ ↦ ⟨a, s, hs, ha⟩, fun ⟨a, s, hs, ha⟩ ↦ ⟨s, hs, a, ha⟩⟩ - rw [h_reorg] - have h_arm_bound : ∀ a : Fin K, P (⋃ s ∈ Finset.range n, badSet s a) ≤ - ENNReal.ofReal (2 * n * δ) := by - intro a - by_cases hn : n = 0 - · simp [hn] - have hn' : 0 < n := Nat.pos_of_ne_zero hn - let badSetIT := fun (s : ℕ) (e : 𝓔) ↦ {ω : ℕ → (Fin K) × ℝ | - pullCount IT.action a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} - have h_set_eq : ⋃ s ∈ Finset.range n, badSet s a = - (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' - {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by - ext ω - simp only [Set.mem_iUnion, Finset.mem_range, badSet, badSetIT, Set.mem_preimage, - Set.mem_setOf_eq, IsBayesAlgEnvSeq.actionMean] - exact Iff.rfl - rw [h_set_eq] - have h_meas_pair : - Measurable (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) := - h.measurable_E.prodMk (IsBayesAlgEnvSeq.measurable_trajectory h.measurable_A h.measurable_R) - have h_disint : P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) = - P.map (E) ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P := - (compProd_map_condDistrib + have h_meas_pair : + Measurable (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) := + h.measurable_E.prodMk (IsBayesAlgEnvSeq.measurable_trajectory h.measurable_A h.measurable_R) + have h_disint : P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) = + P.map E ⊗ₘ condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P := + (compProd_map_condDistrib (IsBayesAlgEnvSeq.measurable_trajectory h.measurable_A h.measurable_R).aemeasurable).symm - have h_cond_bound : ∀ᵐ e ∂(P.map (E)), - (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) - (⋃ s ∈ Finset.range n, badSetIT s e) ≤ ENNReal.ofReal (2 * n * δ) := by - have h_cond_ae : ∀ᵐ e ∂(P.map E), IsAlgEnvSeq IT.action IT.reward - (tsAlgorithm hK Q κ) (stationaryEnv (κ.sectR e)) - (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) := by - rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h - filter_upwards [h_cond_ae] with e h_isAlgEnvSeq - exact concentration_cond_bound (hK := hK) (E := E) (A := A) (R' := R') - (Q := Q) (κ := κ) (P := P) hσ2 hs hn' hδ hδ1 e h_isAlgEnvSeq a - have h_kernel : Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := - stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_fst.prodMk measurable_const) - have h_meas_set : MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | - p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by - have h_eq : {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} = - ⋃ s ∈ Finset.range n, {p | p.2 ∈ badSetIT s p.1} := by - ext p; simp only [Set.mem_setOf_eq, Set.mem_iUnion, Finset.mem_range] - rw [h_eq] - exact .biUnion (Finset.range n).countable_toSet fun s _ ↦ by - simp only [badSetIT] - change MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | - pullCount IT.action a s p.2 ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s p.2 : ℝ)) ≤ - |empMean IT.action IT.reward a s p.2 - (κ (p.1, a))[id]|} - exact MeasurableSet.inter + have h_kernel : ∀ a, Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := + fun a ↦ stronglyMeasurable_id.integral_kernel.measurable.comp + (measurable_fst.prodMk measurable_const) + have h_meas_set : MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ badSetIT p.1} := by + have h_eq : {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ badSetIT p.1} = + ⋃ s ∈ Finset.range n, ⋃ a : Fin K, {p | + pullCount IT.action a s p.2 ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s p.2 : ℝ)) ≤ + |empMean IT.action IT.reward a s p.2 - (κ (p.1, a))[id]|} := by + ext p; simp only [badSetIT, Set.mem_setOf_eq, Set.mem_iUnion, Finset.mem_range] + exact ⟨fun ⟨s, hs, a, ha⟩ ↦ ⟨s, hs, a, ha⟩, fun ⟨s, hs, a, ha⟩ ↦ ⟨s, hs, a, ha⟩⟩ + rw [h_eq] + exact .biUnion (Finset.range n).countable_toSet fun s _ ↦ + .iUnion fun a ↦ + MeasurableSet.inter (((measurable_pullCount IT.measurable_action a s).comp measurable_snd) (measurableSet_singleton (0 : ℕ)).compl) (measurableSet_le (by fun_prop) (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp - measurable_snd).sub h_kernel).abs) - calc P ((fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' - {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1}) - = (P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω))) - {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by - rw [Measure.map_apply h_meas_pair h_meas_set] - _ = (P.map (E) ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P) - {p | p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT s p.1} := by - rw [h_disint] - _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) - (⋃ s ∈ Finset.range n, badSetIT s e) ∂(P.map (E)) := by - rw [Measure.compProd_apply h_meas_set]; rfl - _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * n * δ) ∂(P.map (E)) := by - apply lintegral_mono_ae h_cond_bound - _ = ENNReal.ofReal (2 * n * δ) := by - rw [lintegral_const, Measure.map_apply h.measurable_E MeasurableSet.univ] - simp [measure_univ] - calc P (⋃ a : Fin K, ⋃ s ∈ Finset.range n, badSet s a) - ≤ ∑ a : Fin K, P (⋃ s ∈ Finset.range n, badSet s a) := measure_iUnion_fintype_le _ _ - _ ≤ ∑ _a : Fin K, ENNReal.ofReal (2 * n * δ) := - Finset.sum_le_sum fun a _ ↦ h_arm_bound a - _ = K • ENNReal.ofReal (2 * n * δ) := by simp [Finset.sum_const] + measurable_snd).sub (h_kernel a)).abs) + have h_cond_bound : ∀ᵐ e ∂(P.map E), + (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) (badSetIT e) ≤ + ENNReal.ofReal (2 * K * n * δ) := by + have h_cond_ae : ∀ᵐ e ∂(P.map E), IsAlgEnvSeq IT.action IT.reward + alg (stationaryEnv (κ.sectR e)) + (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) := by + rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h + filter_upwards [h_cond_ae] with e h_isAlgEnvSeq + have : badSetIT e = {ω | ∃ s < n, ∃ a, pullCount IT.action a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ + |empMean IT.action IT.reward a s ω - ((κ.sectR e) a)[id]|} := by + simp only [badSetIT, Kernel.sectR_apply] + rw [this] + exact TS.concentration_fail hσ2 + (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) + h_isAlgEnvSeq n hδ hδ1 + calc P ((fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' + {p | p.2 ∈ badSetIT p.1}) + = (P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω))) + {p | p.2 ∈ badSetIT p.1} := by + rw [Measure.map_apply h_meas_pair h_meas_set] + _ = (P.map E ⊗ₘ + condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P) + {p | p.2 ∈ badSetIT p.1} := by + rw [h_disint] + _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) + (badSetIT e) ∂(P.map E) := by + rw [Measure.compProd_apply h_meas_set]; rfl + _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * K * n * δ) ∂(P.map E) := by + apply lintegral_mono_ae h_cond_bound _ = ENNReal.ofReal (2 * K * n * δ) := by - simp only [nsmul_eq_mul] - rw [← ENNReal.ofReal_natCast K, ← ENNReal.ofReal_mul (Nat.cast_nonneg K)] - congr 1; ring + rw [lintegral_const, Measure.map_apply h.measurable_E MeasurableSet.univ] + simp [measure_univ] lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) + {alg : Algorithm (Fin K) ℝ} + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : @@ -933,23 +692,24 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] (compProd_map_condDistrib (IsBayesAlgEnvSeq.measurable_trajectory h.measurable_A h.measurable_R).aemeasurable).symm - have h_cond_bound : ∀ᵐ e ∂(P.map E), ∀ a : Fin K, + have h_cond_best : ∀ᵐ e ∂(P.map E), (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) - (⋃ s ∈ Finset.range n, badSetIT a s e) ≤ ENNReal.ofReal (2 * n * δ) := by + (⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id e) s e) ≤ + ENNReal.ofReal (2 * n * δ) := by have h_cond_ae : ∀ᵐ e ∂(P.map E), IsAlgEnvSeq IT.action IT.reward - (tsAlgorithm hK Q κ) (stationaryEnv (κ.sectR e)) + alg (stationaryEnv (κ.sectR e)) (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) := by rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h filter_upwards [h_cond_ae] with e h_isAlgEnvSeq - intro a - exact concentration_cond_bound (hK := hK) (E := E) (A := A) (R' := R') - (Q := Q) (κ := κ) (P := P) hσ2 hs hn' hδ hδ1 e h_isAlgEnvSeq a - have h_cond_best : ∀ᵐ e ∂(P.map E), - (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) - (⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id e) s e) ≤ - ENNReal.ofReal (2 * n * δ) := by - filter_upwards [h_cond_bound] with e he - exact he (IsBayesAlgEnvSeq.bestAction κ id e) + have h_eq : ∀ a, ⋃ s ∈ Finset.range n, badSetIT a s e = + ⋃ s ∈ Finset.range n, {ω | pullCount IT.action a s ω ≠ 0 ∧ + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ + |empMean IT.action IT.reward a s ω - ((κ.sectR e) a)[id]|} := by + intro a; simp only [badSetIT, Kernel.sectR_apply] + rw [h_eq] + exact TS.concentration_cond_bound hσ2 + (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) + hn' hδ hδ1 h_isAlgEnvSeq _ have h_kernel : ∀ a, Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := fun a ↦ stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_fst.prodMk measurable_const) @@ -1005,6 +765,48 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] rw [lintegral_const, Measure.map_apply h.measurable_E MeasurableSet.univ] simp [measure_univ] +end Learning.IsBayesAlgEnvSeq + +end BayesianConcentration + +/-! ### TS-specific regret bounds -/ + +namespace Bandits + +section TSRegret + +variable {K : ℕ} +variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] +variable (hK : 0 < K) +variable {Ω : Type*} [MeasurableSpace Ω] +variable (E : Ω → 𝓔) (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → ℝ) +variable (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] +variable (P : Measure Ω) [IsProbabilityMeasure P] + +namespace TS + +lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) (t : ℕ) : + condDistrib (A (t + 1)) (IsAlgEnvSeq.hist A R' t) P + =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] + condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' t) P := by + have h_ba_comp : IsBayesAlgEnvSeq.bestAction κ E + = IsBayesAlgEnvSeq.bestAction κ id ∘ E := rfl + rw [h_ba_comp] + have hm := IsBayesAlgEnvSeq.measurable_bestAction (κ := κ) measurable_id + have h_comp := condDistrib_comp (mβ := MeasurableSpace.pi) (μ := P) + (IsAlgEnvSeq.hist A R' t) h.measurable_E.aemeasurable hm + have h_map : (condDistrib E (IsAlgEnvSeq.hist A R' t) P).map + (IsBayesAlgEnvSeq.bestAction κ id) =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] + (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) t).map + (IsBayesAlgEnvSeq.bestAction κ id) := by + filter_upwards [(h.hasCondDistrib_env_hist + (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) + (absolutelyContinuous_uniformAlgorithm hK _) t).condDistrib_eq] + with x hx + simp only [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hx] + exact (h.hasCondDistrib_action' t).condDistrib_eq.trans (h_comp.trans h_map).symm + lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) @@ -1026,7 +828,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω| < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A (bestArm ω) s ω : ℝ))} have hm_ucb : ∀ a t, Measurable (ucbIndex A R' (↑σ2) lo hi δ a t) := - fun a t ↦ measurable_ucbIndex hK E A R' Q κ P h (↑σ2) lo hi δ a t + fun _ _ ↦ measurable_ucbIndex h.measurable_A h.measurable_R have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.actionMean κ E a) := fun a ↦ IsBayesAlgEnvSeq.measurable_actionMean (a := a) h.measurable_E have hm_best : Measurable (IsBayesAlgEnvSeq.bestAction κ E) := @@ -1037,8 +839,9 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp calc |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| ≤ ∑ s ∈ range n, |armMean (bestArm ω) ω - ucb (bestArm ω) s ω| := Finset.abs_sum_le_sum_abs _ _ - _ ≤ ∑ s ∈ range n, (hi - lo) := Finset.sum_le_sum fun s _ ↦ - abs_sub_le_of_mem_Icc (hm _ _) (ucbIndex_mem_Icc A R' (↑σ2) lo hi δ hlo _ _ _) + _ ≤ ∑ s ∈ range n, (hi - lo) := by + gcongr with s _ + exact abs_sub_le_of_mem_Icc (hm _ _) (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _) _ = ↑n * (hi - lo) := by rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] have h_second_bound : ∀ ω, @@ -1047,8 +850,9 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp calc |∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)| ≤ ∑ s ∈ range n, |ucb (A s ω) s ω - armMean (A s ω) ω| := Finset.abs_sum_le_sum_abs _ _ - _ ≤ ∑ s ∈ range n, (hi - lo) := Finset.sum_le_sum fun s _ ↦ - abs_sub_le_of_mem_Icc (ucbIndex_mem_Icc A R' (↑σ2) lo hi δ hlo _ _ _) (hm _ _) + _ ≤ ∑ s ∈ range n, (hi - lo) := by + gcongr with s _ + exact abs_sub_le_of_mem_Icc (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _) (hm _ _) _ = ↑n * (hi - lo) := by rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] have h_int_sum1 : Integrable (fun ω ↦ ∑ s ∈ range n, @@ -1073,17 +877,17 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (ucb (A s ω) s ω - armMean (A s ω) ω)] := by have h_regret_eq : P[IsBayesAlgEnvSeq.regret κ E A n] = ∑ s ∈ range n, P[fun ω ↦ armMean (bestArm ω) ω - armMean (A s ω) ω] := by - rw [bayesRegret_eq_sum_integral_gap (h := h) + rw [IsBayesAlgEnvSeq.bayesRegret_eq_sum_integral_gap (h := h) (hm := fun a e ↦ abs_le_max_abs_abs (hm a e).1 (hm a e).2) (t := n)] congr 1 with s - exact integral_congr_ae (ae_of_all _ fun ω ↦ gap_eq_armMean_sub E A κ hm s ω) + exact integral_congr_ae (ae_of_all _ fun ω ↦ gap_eq_armMean_sub E κ hm s ω) have h_int_ucb : ∀ s {f : Ω → Fin K}, Measurable f → Integrable (fun ω ↦ ucb (f ω) s ω) P := fun s {_} hf ↦ ⟨(measurable_apply_fin (fun a ↦ hm_ucb a s) hf).aestronglyMeasurable, HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ by rw [Real.norm_eq_abs] - exact abs_le_max_abs_abs (ucbIndex_mem_Icc A R' (↑σ2) lo hi δ hlo _ _ _).1 - (ucbIndex_mem_Icc A R' (↑σ2) lo hi δ hlo _ _ _).2)⟩ + exact abs_le_max_abs_abs (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _).1 + (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _).2)⟩ have h_int_ucb_sub : ∀ s, Integrable (fun ω ↦ ucb (A s ω) s ω - ucb (bestArm ω) s ω) P := fun s ↦ (h_int_ucb s (h.measurable_A s)).sub (h_int_ucb s hm_best) have h_ucb_zero : ∀ a (ω : Ω), ucbIndex A R' (↑σ2) lo hi δ a 0 ω = hi := by @@ -1119,15 +923,9 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp fun p ↦ if pullCount' t p.1 p.2 = 0 then hi else max lo (min hi (empMean' t p.1 p.2 + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount' t p.1 p.2 : ℝ)))) - have h_hist_eq : ∀ ω : Ω, (fun (i : Iic t) ↦ (A (↑i) ω, R' (↑i) ω)) = - IsAlgEnvSeq.hist A R' t ω := fun ω ↦ rfl have hg_eq : ∀ a (ω : Ω), ucbIndex A R' (↑σ2) lo hi δ a (t + 1) ω = - g (IsAlgEnvSeq.hist A R' t ω, a) := by - intro a ω - simp only [g, ucbIndex] - rw [empMean_add_one_eq_empMean' (A := A) (R' := R'), - pullCount_add_one_eq_pullCount' (A := A) (R' := R'), - h_hist_eq] + g (IsAlgEnvSeq.hist A R' t ω, a) := + fun a ω ↦ ucbIndex_succ_eq_ucbIndex'_hist A R' (↑σ2) lo hi δ a t ω have hg_meas : Measurable g := by apply Measurable.ite · have : MeasurableSet {p : (Iic t → Fin K × ℝ) × Fin K | @@ -1199,20 +997,20 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp intro ω hω apply Finset.sum_nonpos intro s hs - linarith [armMean_le_ucbIndex E A R' κ hm (↑σ2) δ + linarith [armMean_le_ucbIndex E κ hm (↑σ2) δ (bestArm ω) s ω (hω s (mem_range.mp hs))] have h_second_Eδ : ∀ ω ∈ Eδ, ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω) ≤ (hi - lo) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by intro ω hω - exact sum_ucbIndex_sub_armMean_le E A R' κ hm hlo (↑σ2) δ n ω hω + exact sum_ucbIndex_sub_armMean_le E κ hm hlo (↑σ2) δ n ω hω have h_prob : P Eδᶜ ≤ ENNReal.ofReal (2 * K * n * δ) := by have : Eδᶜ = {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ |empMean A R' a s ω - armMean a ω|} := by ext ω; simp only [Eδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push_neg; rfl rw [this] - exact prob_concentration_fail_delta (hK := hK) (E := E) (A := A) (R' := R') + exact IsBayesAlgEnvSeq.prob_concentration_fail_delta (E := E) (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) h hσ2 hs n δ hδ hδ1 have hm_emp : ∀ a s, Measurable (fun ω ↦ empMean A R' a s ω) := fun a s ↦ measurable_empMean (fun n ↦ h.measurable_A n) (fun n ↦ h.measurable_R n) a s @@ -1249,7 +1047,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω|} := by ext ω; simp only [Fδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push_neg; rfl rw [this] - exact prob_concentration_bestArm_fail_delta (hK := hK) (E := E) (A := A) (R' := R') + exact IsBayesAlgEnvSeq.prob_concentration_bestArm_fail_delta (E := E) (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) h hσ2 hs n δ hδ hδ1 rw [h_swap] set f1 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, @@ -1264,8 +1062,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp hFδ_meas.compl fun ω _ ↦ (abs_le.mp (h_first_bound ω)).2 rwa [setIntegral_const, smul_eq_mul, mul_comm] at this have h2g : ∫ ω in Eδ, f2 ω ∂P ≤ B := by - have hB : 0 ≤ B := add_nonneg (mul_nonneg (sub_nonneg.mpr hlo) (Nat.cast_nonneg K)) - (by positivity) + have hB : 0 ≤ B := by have : 0 ≤ hi - lo := sub_nonneg.mpr hlo; positivity have := setIntegral_mono_on (hf := h_int_sum2.integrableOn) (hg := integrableOn_const) hEδ_meas fun ω hω ↦ h_second_Eδ ω hω @@ -1281,10 +1078,9 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp ENNReal.toReal_le_of_le_ofReal (by positivity) h_prob rw [(integral_add_compl hFδ_meas h_int_sum1).symm, (integral_add_compl hEδ_meas h_int_sum2).symm] - nlinarith [mul_le_mul_of_nonneg_left hPF - (mul_nonneg (Nat.cast_nonneg n) (sub_nonneg.mpr hlo)), - mul_le_mul_of_nonneg_left hPE - (mul_nonneg (Nat.cast_nonneg n) (sub_nonneg.mpr hlo)), + have : 0 ≤ ↑n * (hi - lo) := by nlinarith + nlinarith [mul_le_mul_of_nonneg_left hPF this, + mul_le_mul_of_nonneg_left hPE this, measureReal_nonneg (μ := P) (s := Fδᶜ), measureReal_nonneg (μ := P) (s := Eδᶜ)] @@ -1299,7 +1095,7 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] have hlo : lo ≤ hi := h1.trans h2 by_cases ht : t = 0 · simp [ht, IsBayesAlgEnvSeq.regret, Bandits.regret] - nlinarith [sub_nonneg.mpr hlo, show (0 : ℝ) < K from Nat.cast_pos.mpr hK, + nlinarith [sub_nonneg.mpr hlo, Nat.cast_pos (α := ℝ).mpr hK, Real.sqrt_nonneg (↑σ2 * ↑K * (0 : ℝ) * Real.log (0 : ℝ))] by_cases ht1_eq : t = 1 · subst ht1_eq @@ -1316,12 +1112,12 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (hm (A 0 ω) (E ω)).1])).trans ?_ simp _ ≤ (3 * ↑K + 2) * (hi - lo) := by - nlinarith [show (1 : ℝ) ≤ K from Nat.one_le_cast.mpr (Nat.one_le_of_lt hK), + nlinarith [Nat.one_le_cast (α := ℝ).mpr (Nat.one_le_of_lt hK), sub_nonneg.mpr hlo] -- For t ≥ 2, we have δ = 1/t² < 1 · have ht2 : 2 ≤ t := by omega have htpos : (0 : ℝ) < t := by positivity - have _ht1 : (1 : ℝ) ≤ t := Nat.one_le_cast.mpr (Nat.pos_of_ne_zero ht) + have _ht1 : (1 : ℝ) ≤ t := by exact_mod_cast Nat.pos_of_ne_zero ht have hδ : (0 : ℝ) < 1 / (t : ℝ) ^ 2 := by positivity have hδ1 : 1 / (t : ℝ) ^ 2 < 1 := by rw [div_lt_one (pow_pos htpos 2)] @@ -1347,12 +1143,12 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 4)] ring _ = (3 * ↑K + 2) * (hi - lo) + 8 * √(↑σ2 * ↑K * ↑t * Real.log ↑t) := by - rw [← Real.sqrt_mul (mul_nonneg (NNReal.coe_nonneg σ2) - (Real.log_nonneg (Nat.one_le_cast.mpr (Nat.pos_of_ne_zero ht))))] - congr 1; congr 1; congr 1; ring + rw [← Real.sqrt_mul (by positivity : + 0 ≤ ↑σ2 * Real.log ↑t)] + congr 1; ring_nf end TS -end Regret +end TSRegret end Bandits From cb2bd8ae0c31a7cf8719bbf99fae234e92aec03d Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 11 Mar 2026 09:25:01 +0000 Subject: [PATCH 071/155] Refactor TS.lean (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 64 +++++++++++----------------- 1 file changed, 25 insertions(+), 39 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 807e0889..5da54ecc 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -11,62 +11,46 @@ import LeanBandits.SequentialLearning.AlgorithmDensity open MeasureTheory ProbabilityTheory Finset Learning -open scoped ENNReal NNReal +open scoped NNReal namespace Bandits -namespace TS +variable {K : ℕ} +variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] -variable {K : ℕ} (hK : 0 < K) -variable {𝓔 : Type*} [m𝓔 : MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] -variable (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] +namespace TS -/-- The distribution over actions for every given history for TS. -/ noncomputable -def policy (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := +def policy (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] + (hK : 0 < K) (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) n).map - (IsBayesAlgEnvSeq.bestAction κ id) + (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) n).map (IsBayesAlgEnvSeq.bestAction κ id) -instance (n : ℕ) : IsMarkovKernel (policy hK Q κ n) := by - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - unfold policy - exact Kernel.IsMarkovKernel.map _ - (IsBayesAlgEnvSeq.measurable_bestAction measurable_id) +instance {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] + {hK : 0 < K} (n : ℕ) : IsMarkovKernel (policy Q κ hK n) := + Kernel.IsMarkovKernel.map _ (by fun_prop) -/-- The initial distribution over actions for TS. -/ noncomputable -def initialPolicy : Measure (Fin K) := +def initialPolicy (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) + [IsMarkovKernel κ] (hK : 0 < K) : Measure (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK Q.map (IsBayesAlgEnvSeq.bestAction κ id) -instance : IsProbabilityMeasure (initialPolicy hK Q κ) := by - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - exact Measure.isProbabilityMeasure_map - (IsBayesAlgEnvSeq.measurable_bestAction (by fun_prop)).aemeasurable +instance {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] + {hK : 0 < K} : IsProbabilityMeasure (initialPolicy Q κ hK) := + Measure.isProbabilityMeasure_map (by fun_prop) end TS -variable {K : ℕ} - -section Algorithm - -variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] - -/-- The Thompson Sampling (TS) algorithm: actions are chosen according to the probability that they -are optimal given prior knowledge represented by a prior distribution `Q` and a data generation -model represented by a kernel `κ`. -/ noncomputable -def tsAlgorithm (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] - (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] : Algorithm (Fin K) ℝ where - policy := TS.policy hK Q κ - p0 := TS.initialPolicy hK Q κ - -end Algorithm +def tsAlgorithm (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) + [IsMarkovKernel κ] (hK : 0 < K) : Algorithm (Fin K) ℝ where + policy := TS.policy Q κ hK + p0 := TS.initialPolicy Q κ hK namespace TS -/-! ### Auxiliary real-analysis lemmas -/ +section Auxiliary lemma abs_sub_le_of_mem_Icc {lo hi x y : ℝ} (hx : x ∈ Set.Icc lo hi) (hy : y ∈ Set.Icc lo hi) : @@ -97,6 +81,8 @@ lemma sum_inv_sqrt_le (M : ℕ) (hM : 0 < M) : mul_self_nonneg (√(↑(n + 1) : ℝ) - √(↑n : ℝ)), show (↑(n + 1) : ℝ) = ↑n + 1 from by push_cast; ring] +end Auxiliary + /-! ### UCB index: definition and deterministic bounds -/ section Deterministic @@ -786,7 +772,7 @@ variable (P : Measure Ω) [IsProbabilityMeasure P] namespace TS lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) (t : ℕ) : + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (t : ℕ) : condDistrib (A (t + 1)) (IsAlgEnvSeq.hist A R' t) P =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' t) P := by @@ -808,7 +794,7 @@ lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] exact (h.hasCondDistrib_action' t).condDistrib_eq.trans (h_comp.trans h_map).symm lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) @@ -1085,7 +1071,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp measureReal_nonneg (μ := P) (s := Eδᶜ)] lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (t : ℕ) : From ef8da0806c4c5c460b1c9df6e44c683658101f78 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 13 Mar 2026 12:11:30 +0000 Subject: [PATCH 072/155] Refactor TS.lean (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 141 +++++++++++++-------------- 1 file changed, 67 insertions(+), 74 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 5da54ecc..9912637f 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -50,42 +50,7 @@ def tsAlgorithm (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 namespace TS -section Auxiliary - -lemma abs_sub_le_of_mem_Icc {lo hi x y : ℝ} (hx : x ∈ Set.Icc lo hi) - (hy : y ∈ Set.Icc lo hi) : - |x - y| ≤ hi - lo := by - rw [abs_le]; constructor <;> linarith [hx.1, hx.2, hy.1, hy.2] - -lemma sum_sqrt_le {ι : Type*} (s : Finset ι) (c : ι → ℝ) (hc : ∀ i, 0 ≤ c i) : - ∑ i ∈ s, √(c i) ≤ √(#s * ∑ i ∈ s, c i) := by - have h := Real.sum_sqrt_mul_sqrt_le s hc (fun _ => zero_le_one) - simp only [Real.sqrt_one, mul_one, sum_const, nsmul_eq_mul] at h - calc ∑ i ∈ s, √(c i) ≤ √(∑ i ∈ s, c i) * √↑(#s) := h - _ = _ := by rw [← Real.sqrt_mul (Finset.sum_nonneg (fun i _ => hc i)), mul_comm] - -lemma sum_inv_sqrt_le (M : ℕ) (hM : 0 < M) : - ∑ j ∈ range M, (1 / √(↑j : ℝ)) + 1 / √↑M ≤ 2 * √↑M := by - induction M with - | zero => omega - | succ n ih => - rw [sum_range_succ] - by_cases hn : n = 0 - · subst hn; simp - · have hn_pos : 0 < n := Nat.pos_of_ne_zero hn - have h_ih := ih hn_pos - suffices h_key : 1 / √(↑(n + 1) : ℝ) ≤ 2 * (√↑(n + 1) - √↑n) by linarith - rw [div_le_iff₀ (Real.sqrt_pos.mpr (by positivity : (0 : ℝ) < ↑(n + 1)))] - nlinarith [Real.mul_self_sqrt (show (0 : ℝ) ≤ ↑(n + 1) by positivity), - Real.mul_self_sqrt (show (0 : ℝ) ≤ ↑n by positivity), - mul_self_nonneg (√(↑(n + 1) : ℝ) - √(↑n : ℝ)), - show (↑(n + 1) : ℝ) = ↑n + 1 from by push_cast; ring] - -end Auxiliary - -/-! ### UCB index: definition and deterministic bounds -/ - -section Deterministic +section UCB variable {Ω : Type*} variable {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} @@ -123,38 +88,6 @@ lemma measurable_ucbIndex [MeasurableSpace Ω] exact .ite ((measurable_pullCount hA a t) (measurableSet_singleton 0)) measurable_const (by fun_prop) -/-- History-based UCB index: like `ucbIndex` but takes a history `h : Iic t → Fin K × ℝ` -directly instead of the random variables `A` and `R'`. -/ -noncomputable -def ucbIndex' (σ2 lo hi δ : ℝ) (t : ℕ) (h : Iic t → Fin K × ℝ) (a : Fin K) : ℝ := - if pullCount' t h a = 0 then hi - else max lo (min hi (empMean' t h a + - √(2 * σ2 * Real.log (1 / δ) / (pullCount' t h a : ℝ)))) - -@[fun_prop] -lemma measurable_ucbIndex' {σ2 lo hi δ : ℝ} {a : Fin K} {t : ℕ} : - Measurable (fun h : Iic t → Fin K × ℝ ↦ ucbIndex' σ2 lo hi δ t h a) := by - unfold ucbIndex' - apply Measurable.ite - · have : MeasurableSet {h : Iic t → Fin K × ℝ | (pullCount' t h a : ℝ) = (0 : ℝ)} := - measurableSet_eq_fun - (measurable_from_top.comp (measurable_pullCount' t a)) - measurable_const - simp only [Nat.cast_eq_zero] at this; exact this - · exact measurable_const - · exact .max measurable_const (.min measurable_const - (.add (measurable_empMean' t a) - (measurable_const.div (measurable_from_top.comp (measurable_pullCount' t a))).sqrt)) - -lemma ucbIndex_succ_eq_ucbIndex'_hist (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) - (σ2 lo hi δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : - ucbIndex A R' σ2 lo hi δ a (t + 1) ω = - ucbIndex' σ2 lo hi δ t (IsAlgEnvSeq.hist A R' t ω) a := by - have hpc : pullCount A a (t + 1) ω = pullCount' t (IsAlgEnvSeq.hist A R' t ω) a := - pullCount_add_one_eq_pullCount' - have hem : empMean A R' a (t + 1) ω = empMean' t (IsAlgEnvSeq.hist A R' t ω) a := - empMean_add_one_eq_empMean' - simp only [ucbIndex, ucbIndex', hpc, hem] lemma armMean_le_ucbIndex {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) @@ -229,6 +162,29 @@ lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : congr 1 simp +private lemma sum_sqrt_le {ι : Type*} {c : ι → ℝ} (s : Finset ι) (hc : ∀ i, 0 ≤ c i) : + ∑ i ∈ s, √(c i) ≤ √(#s * ∑ i ∈ s, c i) := by + have h := Real.sum_sqrt_mul_sqrt_le s hc (fun _ => zero_le_one) + simp only [Real.sqrt_one, mul_one, sum_const, nsmul_eq_mul] at h + rwa [Real.sqrt_mul (by positivity), mul_comm] + +private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 / √k ≤ 2 * √n - 1 := by + induction n with + | zero => simp at h + | succ n ih => + rw [sum_range_succ] + by_cases hn : n = 0 + · rw [hn] + simp + norm_num + · have hi := ih (Nat.pos_of_ne_zero hn) + suffices 1 / √↑(n + 1) ≤ 2 * (√↑(n + 1) - √n) by linarith + push_cast + field_simp + have : √(n + 1) * √(n + 1) = (n + 1) := Real.mul_self_sqrt (by positivity) + have : √n * √n = n := Real.mul_self_sqrt (by positivity) + nlinarith + lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (hlo : lo ≤ hi) (σ2 δ : ℝ) (n : ℕ) (ω : Ω) (hconc : ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → @@ -288,7 +244,8 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] gcongr with a by_cases ha : pullCount A a n ω = 0 · simp [ha] - · have := sum_inv_sqrt_le _ (Nat.pos_of_ne_zero ha) + · have := sum_inv_sqrt_le (Nat.pos_of_ne_zero ha) + rw [sum_range_succ] at this linarith [div_nonneg zero_le_one (Real.sqrt_nonneg (↑(pullCount A a n ω) : ℝ))] _ = √(8 * σ2 * c) * (2 * ∑ a : Fin K, √↑(pullCount A a n ω)) := by @@ -297,7 +254,7 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] gcongr calc ∑ a : Fin K, √↑(pullCount A a n ω) ≤ √(↑(Finset.univ.card) * ∑ a, ↑(pullCount A a n ω)) := - sum_sqrt_le Finset.univ _ fun a => by positivity + sum_sqrt_le Finset.univ fun a => by positivity _ = √(↑K * ↑n) := by congr 1; rw [Finset.card_fin]; congr 1 have h := sum_pullCount (A := A) (t := n) (ω := ω) @@ -342,7 +299,40 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] _ ≤ ↑K * (hi - lo) := by gcongr; linarith _ = (hi - lo) * ↑K := by ring -end Deterministic +/-- History-based UCB index: like `ucbIndex` but takes a history `h : Iic t → Fin K × ℝ` +directly instead of the random variables `A` and `R'`. -/ +noncomputable +def ucbIndex' (σ2 lo hi δ : ℝ) (t : ℕ) (h : Iic t → Fin K × ℝ) (a : Fin K) : ℝ := + if pullCount' t h a = 0 then hi + else max lo (min hi (empMean' t h a + + √(2 * σ2 * Real.log (1 / δ) / (pullCount' t h a : ℝ)))) + +@[fun_prop] +lemma measurable_ucbIndex' {σ2 lo hi δ : ℝ} {a : Fin K} {t : ℕ} : + Measurable (fun h : Iic t → Fin K × ℝ ↦ ucbIndex' σ2 lo hi δ t h a) := by + unfold ucbIndex' + apply Measurable.ite + · have : MeasurableSet {h : Iic t → Fin K × ℝ | (pullCount' t h a : ℝ) = (0 : ℝ)} := + measurableSet_eq_fun + (measurable_from_top.comp (measurable_pullCount' t a)) + measurable_const + simp only [Nat.cast_eq_zero] at this; exact this + · exact measurable_const + · exact .max measurable_const (.min measurable_const + (.add (measurable_empMean' t a) + (measurable_const.div (measurable_from_top.comp (measurable_pullCount' t a))).sqrt)) + +lemma ucbIndex_succ_eq_ucbIndex'_hist (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) + (σ2 lo hi δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : + ucbIndex A R' σ2 lo hi δ a (t + 1) ω = + ucbIndex' σ2 lo hi δ t (IsAlgEnvSeq.hist A R' t ω) a := by + have hpc : pullCount A a (t + 1) ω = pullCount' t (IsAlgEnvSeq.hist A R' t ω) a := + pullCount_add_one_eq_pullCount' + have hem : empMean A R' a (t + 1) ω = empMean' t (IsAlgEnvSeq.hist A R' t ω) a := + empMean_add_one_eq_empMean' + simp only [ucbIndex, ucbIndex', hpc, hem] + +end UCB /-! ### Concentration bounds (algorithm-generic) @@ -827,7 +817,9 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp Finset.abs_sum_le_sum_abs _ _ _ ≤ ∑ s ∈ range n, (hi - lo) := by gcongr with s _ - exact abs_sub_le_of_mem_Icc (hm _ _) (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _) + exact abs_sub_le_of_le_of_le (hm _ _).1 (hm _ _).2 + (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _).1 + (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _).2 _ = ↑n * (hi - lo) := by rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] have h_second_bound : ∀ ω, @@ -838,7 +830,8 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp Finset.abs_sum_le_sum_abs _ _ _ ≤ ∑ s ∈ range n, (hi - lo) := by gcongr with s _ - exact abs_sub_le_of_mem_Icc (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _) (hm _ _) + exact abs_sub_le_of_le_of_le (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _).1 + (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _).2 (hm _ _).1 (hm _ _).2 _ = ↑n * (hi - lo) := by rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] have h_int_sum1 : Integrable (fun ω ↦ ∑ s ∈ range n, @@ -963,7 +956,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (measurable_apply_fin hm_arm (h.measurable_A s))).aestronglyMeasurable (hi - lo) (ae_of_all _ fun ω ↦ ?_) rw [Real.norm_eq_abs] - exact abs_sub_le_of_mem_Icc (hm _ _) (hm _ _) + exact abs_sub_le_of_le_of_le (hm _ _).1 (hm _ _).2 (hm _ _).1 (hm _ _).2 have h_int_gap : Integrable (fun ω ↦ ∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω)) P := integrable_finset_sum _ h_int_gap_s From 196c9ca66d35e2a968a7a932b86db2a7c3998247 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 13 Mar 2026 12:25:59 +0000 Subject: [PATCH 073/155] Minor --- LeanBandits/BanditAlgorithms/TS.lean | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 9912637f..53d238d0 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -88,7 +88,6 @@ lemma measurable_ucbIndex [MeasurableSpace Ω] exact .ite ((measurable_pullCount hA a t) (measurableSet_singleton 0)) measurable_const (by fun_prop) - lemma armMean_le_ucbIndex {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hconc : pullCount A a t ω ≠ 0 → @@ -162,12 +161,14 @@ lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : congr 1 simp +/-- Helper for `sum_ucbIndex_sub_armMean_le`. -/ private lemma sum_sqrt_le {ι : Type*} {c : ι → ℝ} (s : Finset ι) (hc : ∀ i, 0 ≤ c i) : ∑ i ∈ s, √(c i) ≤ √(#s * ∑ i ∈ s, c i) := by have h := Real.sum_sqrt_mul_sqrt_le s hc (fun _ => zero_le_one) simp only [Real.sqrt_one, mul_one, sum_const, nsmul_eq_mul] at h rwa [Real.sqrt_mul (by positivity), mul_comm] +/-- Helper for `sum_ucbIndex_sub_armMean_le`. -/ private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 / √k ≤ 2 * √n - 1 := by induction n with | zero => simp at h From cd9515ec7b95e36a45086c977604513ff54e4378 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 13 Mar 2026 14:36:02 +0000 Subject: [PATCH 074/155] Refactor TS.lean (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 162 +++++++++++---------------- 1 file changed, 65 insertions(+), 97 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 53d238d0..67f88818 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -53,75 +53,38 @@ namespace TS section UCB variable {Ω : Type*} -variable {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} -variable {𝓔 : Type*} [MeasurableSpace 𝓔] -variable (E : Ω → 𝓔) (κ : Kernel (𝓔 × Fin K) ℝ) noncomputable -def ucbIndex (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (σ2 lo hi δ : ℝ) - (a : Fin K) (t : ℕ) (ω : Ω) : ℝ := - if pullCount A a t ω = 0 then hi - else max lo (min hi - (empMean A R' a t ω - + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ)))) +def ucbIndex (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ := + if pullCount A a n ω = 0 then u + else max l (min u (empMean A R' a n ω + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a n ω)))) -lemma ucbIndex_mem_Icc (σ2 lo hi δ : ℝ) (hlo : lo ≤ hi) (a : Fin K) (t : ℕ) (ω : Ω) : - ucbIndex A R' σ2 lo hi δ a t ω ∈ Set.Icc lo hi := by - unfold ucbIndex - split_ifs <;> constructor - · exact hlo - · exact le_refl _ - · exact le_max_left lo _ - · exact max_le hlo (min_le_left hi _) +variable {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} +variable {l u σ2 δ : ℝ} @[fun_prop] -lemma measurable_ucbIndex [MeasurableSpace Ω] - {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} - {σ2 lo hi δ : ℝ} {a : Fin K} {t : ℕ} - (hA : ∀ n, Measurable (A n)) (hR : ∀ n, Measurable (R' n)) : - Measurable (ucbIndex A R' σ2 lo hi δ a t) := by - unfold ucbIndex - have : Measurable (fun ω ↦ (pullCount A a t ω : ℝ)) := - measurable_from_top.comp (measurable_pullCount hA a t) - have := measurable_empMean hA hR a t - have := measurable_pullCount hA a t - exact .ite ((measurable_pullCount hA a t) - (measurableSet_singleton 0)) measurable_const (by fun_prop) +lemma measurable_ucbIndex [MeasurableSpace Ω] {a : Fin K} {n : ℕ} (hA : ∀ t, Measurable (A t)) + (hR : ∀ t, Measurable (R' t)) : Measurable (ucbIndex A R' l u σ2 δ a n) := + Measurable.ite (by measurability) (by fun_prop) (by fun_prop) -lemma armMean_le_ucbIndex {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) - (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) - (hconc : pullCount A a t ω ≠ 0 → - |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E a ω| - < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ))) : - IsBayesAlgEnvSeq.actionMean κ E a ω ≤ ucbIndex A R' σ2 lo hi δ a t ω := by +lemma ucbIndex_mem_Icc (h : l ≤ u) {a : Fin K} {n : ℕ} {ω : Ω} : + ucbIndex A R' l u σ2 δ a n ω ∈ Set.Icc l u := by unfold ucbIndex - have hmean := hm a (E ω) - simp only [IsBayesAlgEnvSeq.actionMean] at hmean hconc ⊢ - split_ifs with h0 - · exact hmean.2 - · have habs := abs_sub_lt_iff.mp (hconc h0) - refine le_max_of_le_right (le_min hmean.2 ?_) - linarith [habs.2] + grind -lemma ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) +lemma ucbIndex_sub_mean_le {lo hi μ : ℝ} (hμ : μ ∈ Set.Icc lo hi) (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hpc : pullCount A a t ω ≠ 0) - (hconc : - |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E a ω| + (h : + |empMean A R' a t ω - μ| < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ))) : - ucbIndex A R' σ2 lo hi δ a t ω - IsBayesAlgEnvSeq.actionMean κ E a ω + ucbIndex A R' lo hi σ2 δ a t ω - μ ≤ 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ)) := by unfold ucbIndex - simp only [IsBayesAlgEnvSeq.actionMean] at hconc ⊢ - rw [if_neg hpc] - set w := √(2 * σ2 * Real.log (1 / δ) / ↑(pullCount A a t ω)) - set emp := empMean A R' a t ω - have habs := abs_sub_lt_iff.mp hconc - have hmean := hm a (E ω) - have h1 : max lo (min hi (emp + w)) ≤ emp + w := - max_le_iff.mpr ⟨by linarith [hmean.1, habs.2], min_le_right _ _⟩ - linarith [habs.2] + grind -lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] {lo hi : ℝ} +omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] +lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] (E : Ω → 𝓔) + (κ : Kernel (𝓔 × Fin K) ℝ) {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) (ω : Ω) : ⨆ i, IsBayesAlgEnvSeq.actionMean κ E i ω = IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω := @@ -131,7 +94,8 @@ lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] {lo hi : ℝ} (le_ciSup (f := fun i ↦ IsBayesAlgEnvSeq.actionMean κ E i ω) ⟨hi, by rintro _ ⟨i, rfl⟩; exact (hm i _).2⟩ _) -lemma gap_eq_armMean_sub [Nonempty (Fin K)] {lo hi : ℝ} +lemma gap_eq_armMean_sub [Nonempty (Fin K)] (E : Ω → 𝓔) + (κ : Kernel (𝓔 × Fin K) ℝ) {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) (s : ℕ) (ω : Ω) : gap (κ.sectR (E ω)) (A s ω) = IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - @@ -186,13 +150,13 @@ private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 have : √n * √n = n := Real.mul_self_sqrt (by positivity) nlinarith -lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) +lemma sum_ucbIndex_sub_mean_le {lo hi : ℝ} {μ : Fin K → ℝ} + (hm : ∀ a, μ a ∈ Set.Icc lo hi) (hlo : lo ≤ hi) (σ2 δ : ℝ) (n : ℕ) (ω : Ω) (hconc : ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → - |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω| + |empMean A R' a s ω - μ a| < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ))) : - ∑ s ∈ range n, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) + ∑ s ∈ range n, (ucbIndex A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) ≤ (hi - lo) * ↑K + 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by -- Split range n into first-pull (pc=0) and non-first-pull (pc≠0) sets set S0 := (range n).filter (fun s => pullCount A (A s ω) s ω = 0) @@ -202,22 +166,21 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] conv_lhs => rw [hpart] rw [Finset.sum_union hdisj] -- We bound ∑_{S0} and ∑_{S1} separately, then combine - suffices h_S0 : ∑ s ∈ S0, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ≤ (hi - lo) * ↑K by - suffices h_S1 : ∑ s ∈ S1, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) + suffices h_S0 : ∑ s ∈ S0, (ucbIndex A R' lo hi σ2 δ (A s ω) s ω - + μ (A s ω)) ≤ (hi - lo) * ↑K by + suffices h_S1 : ∑ s ∈ S1, (ucbIndex A R' lo hi σ2 δ (A s ω) s ω - + μ (A s ω)) ≤ 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) by have := Finset.sum_union hdisj (f := fun s => - ucbIndex A R' σ2 lo hi δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) + ucbIndex A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) rw [← hpart] at this; linarith -- Bound ∑_{S1}: each term ≤ 2√(2σ2c/pc) = 2√(2σ2c/max(1,pc)), so ≤ full sum - calc ∑ s ∈ S1, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) + calc ∑ s ∈ S1, (ucbIndex A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) ≤ ∑ s ∈ S1, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) := sum_le_sum fun s hs => by have hpc : pullCount A (A s ω) s ω ≠ 0 := (Finset.mem_filter.mp hs).2 - exact ucbIndex_sub_armMean_le E κ hm σ2 δ (A s ω) s ω hpc + exact ucbIndex_sub_mean_le (hm _) σ2 δ (A s ω) s ω hpc (hconc s (mem_range.mp (Finset.mem_filter.mp hs).1) _ hpc) _ ≤ ∑ s ∈ range n, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) := @@ -269,12 +232,12 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] div_nonpos_of_nonpos_of_nonneg (by linarith) (Nat.cast_nonneg _) simp [sqrt_eq_zero'.mpr this] rw [sum_congr rfl h0]; simp only [sum_const_zero]; positivity - -- Bound ∑_{S0}: each term = hi - armMean ≤ hi - lo, and #S0 ≤ K - have hterm_S0 : ∀ s ∈ S0, ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω ≤ hi - lo := fun s hs => by + -- Bound ∑_{S0}: each term = hi - μ ≤ hi - lo, and #S0 ≤ K + have hterm_S0 : ∀ s ∈ S0, ucbIndex A R' lo hi σ2 δ (A s ω) s ω - + μ (A s ω) ≤ hi - lo := fun s hs => by have hpc : pullCount A (A s ω) s ω = 0 := (Finset.mem_filter.mp hs).2 - simp only [ucbIndex, hpc, ↓reduceIte, IsBayesAlgEnvSeq.actionMean] - linarith [(hm (A s ω) (E ω)).1] + simp only [ucbIndex, hpc, ↓reduceIte] + linarith [(hm (A s ω)).1] have h_card_S0 : #S0 ≤ K := by calc #S0 ≤ #(Finset.univ : Finset (Fin K)) := Finset.card_le_card_of_injOn (fun s => A s ω) @@ -293,8 +256,7 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] exact absurd hpc₁ (show pullCount A (A s₁ ω) s₁ ω ≠ 0 from Finset.card_ne_zero_of_mem this)) _ = K := Finset.card_fin K - calc ∑ s ∈ S0, (ucbIndex A R' σ2 lo hi δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) + calc ∑ s ∈ S0, (ucbIndex A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) ≤ ∑ _s ∈ S0, (hi - lo) := sum_le_sum hterm_S0 _ = #S0 * (hi - lo) := by rw [sum_const, nsmul_eq_mul] _ ≤ ↑K * (hi - lo) := by gcongr; linarith @@ -303,14 +265,14 @@ lemma sum_ucbIndex_sub_armMean_le {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] /-- History-based UCB index: like `ucbIndex` but takes a history `h : Iic t → Fin K × ℝ` directly instead of the random variables `A` and `R'`. -/ noncomputable -def ucbIndex' (σ2 lo hi δ : ℝ) (t : ℕ) (h : Iic t → Fin K × ℝ) (a : Fin K) : ℝ := +def ucbIndex' (lo hi σ2 δ : ℝ) (t : ℕ) (h : Iic t → Fin K × ℝ) (a : Fin K) : ℝ := if pullCount' t h a = 0 then hi else max lo (min hi (empMean' t h a + √(2 * σ2 * Real.log (1 / δ) / (pullCount' t h a : ℝ)))) @[fun_prop] -lemma measurable_ucbIndex' {σ2 lo hi δ : ℝ} {a : Fin K} {t : ℕ} : - Measurable (fun h : Iic t → Fin K × ℝ ↦ ucbIndex' σ2 lo hi δ t h a) := by +lemma measurable_ucbIndex' {lo hi σ2 δ : ℝ} {a : Fin K} {t : ℕ} : + Measurable (fun h : Iic t → Fin K × ℝ ↦ ucbIndex' lo hi σ2 δ t h a) := by unfold ucbIndex' apply Measurable.ite · have : MeasurableSet {h : Iic t → Fin K × ℝ | (pullCount' t h a : ℝ) = (0 : ℝ)} := @@ -324,9 +286,9 @@ lemma measurable_ucbIndex' {σ2 lo hi δ : ℝ} {a : Fin K} {t : ℕ} : (measurable_const.div (measurable_from_top.comp (measurable_pullCount' t a))).sqrt)) lemma ucbIndex_succ_eq_ucbIndex'_hist (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) - (σ2 lo hi δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : - ucbIndex A R' σ2 lo hi δ a (t + 1) ω = - ucbIndex' σ2 lo hi δ t (IsAlgEnvSeq.hist A R' t ω) a := by + (lo hi σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : + ucbIndex A R' lo hi σ2 δ a (t + 1) ω = + ucbIndex' lo hi σ2 δ t (IsAlgEnvSeq.hist A R' t ω) a := by have hpc : pullCount A a (t + 1) ω = pullCount' t (IsAlgEnvSeq.hist A R' t ω) a := pullCount_add_one_eq_pullCount' have hem : empMean A R' a (t + 1) ω = empMean' t (IsAlgEnvSeq.hist A R' t ω) a := @@ -797,14 +759,14 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp have hlo : lo ≤ hi := h1.trans h2 let bestArm := IsBayesAlgEnvSeq.bestAction κ E let armMean := IsBayesAlgEnvSeq.actionMean κ E - let ucb := ucbIndex A R' (↑σ2) lo hi δ + let ucb := ucbIndex A R' lo hi (↑σ2) δ set Eδ : Set Ω := {ω | ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → |empMean A R' a s ω - armMean a ω| < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ))} set Fδ : Set Ω := {ω | ∀ s < n, pullCount A (bestArm ω) s ω ≠ 0 → |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω| < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A (bestArm ω) s ω : ℝ))} - have hm_ucb : ∀ a t, Measurable (ucbIndex A R' (↑σ2) lo hi δ a t) := + have hm_ucb : ∀ a t, Measurable (ucbIndex A R' lo hi (↑σ2) δ a t) := fun _ _ ↦ measurable_ucbIndex h.measurable_A h.measurable_R have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.actionMean κ E a) := fun a ↦ IsBayesAlgEnvSeq.measurable_actionMean (a := a) h.measurable_E @@ -819,8 +781,8 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp _ ≤ ∑ s ∈ range n, (hi - lo) := by gcongr with s _ exact abs_sub_le_of_le_of_le (hm _ _).1 (hm _ _).2 - (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _).1 - (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _).2 + ((ucbIndex_mem_Icc hlo).1) + (ucbIndex_mem_Icc hlo).2 _ = ↑n * (hi - lo) := by rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] have h_second_bound : ∀ ω, @@ -831,8 +793,8 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp Finset.abs_sum_le_sum_abs _ _ _ ≤ ∑ s ∈ range n, (hi - lo) := by gcongr with s _ - exact abs_sub_le_of_le_of_le (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _).1 - (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _).2 (hm _ _).1 (hm _ _).2 + exact abs_sub_le_of_le_of_le (ucbIndex_mem_Icc hlo).1 + (ucbIndex_mem_Icc hlo).2 (hm _ _).1 (hm _ _).2 _ = ↑n * (hi - lo) := by rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] have h_int_sum1 : Integrable (fun ω ↦ ∑ s ∈ range n, @@ -866,18 +828,18 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp ⟨(measurable_apply_fin (fun a ↦ hm_ucb a s) hf).aestronglyMeasurable, HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ by rw [Real.norm_eq_abs] - exact abs_le_max_abs_abs (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _).1 - (ucbIndex_mem_Icc (↑σ2) lo hi δ hlo _ _ _).2)⟩ + exact abs_le_max_abs_abs (ucbIndex_mem_Icc hlo).1 + (ucbIndex_mem_Icc hlo).2)⟩ have h_int_ucb_sub : ∀ s, Integrable (fun ω ↦ ucb (A s ω) s ω - ucb (bestArm ω) s ω) P := fun s ↦ (h_int_ucb s (h.measurable_A s)).sub (h_int_ucb s hm_best) - have h_ucb_zero : ∀ a (ω : Ω), ucbIndex A R' (↑σ2) lo hi δ a 0 ω = hi := by + have h_ucb_zero : ∀ a (ω : Ω), ucbIndex A R' lo hi (↑σ2) δ a 0 ω = hi := by intro a ω; unfold ucbIndex; simp [pullCount_zero] have h_ucb_swap : ∀ s, ∫ ω, (ucb (A s ω) s ω - ucb (bestArm ω) s ω) ∂P = 0 := by intro s cases s with | zero => have : ∀ ω, ucb (A 0 ω) 0 ω - ucb (bestArm ω) 0 ω = 0 := fun ω ↦ by - change ucbIndex A R' (↑σ2) lo hi δ _ 0 ω - ucbIndex A R' (↑σ2) lo hi δ _ 0 ω = 0 + change ucbIndex A R' lo hi (↑σ2) δ _ 0 ω - ucbIndex A R' lo hi (↑σ2) δ _ 0 ω = 0 simp [h_ucb_zero] exact (integral_congr_ae (ae_of_all _ this)).trans (integral_zero _ _) | succ t => @@ -903,9 +865,9 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp fun p ↦ if pullCount' t p.1 p.2 = 0 then hi else max lo (min hi (empMean' t p.1 p.2 + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount' t p.1 p.2 : ℝ)))) - have hg_eq : ∀ a (ω : Ω), ucbIndex A R' (↑σ2) lo hi δ a (t + 1) ω = + have hg_eq : ∀ a (ω : Ω), ucbIndex A R' lo hi (↑σ2) δ a (t + 1) ω = g (IsAlgEnvSeq.hist A R' t ω, a) := - fun a ω ↦ ucbIndex_succ_eq_ucbIndex'_hist A R' (↑σ2) lo hi δ a t ω + fun a ω ↦ ucbIndex_succ_eq_ucbIndex'_hist A R' lo hi (↑σ2) δ a t ω have hg_meas : Measurable g := by apply Measurable.ite · have : MeasurableSet {p : (Iic t → Fin K × ℝ) × Fin K | @@ -977,13 +939,19 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp intro ω hω apply Finset.sum_nonpos intro s hs - linarith [armMean_le_ucbIndex E κ hm (↑σ2) δ - (bestArm ω) s ω (hω s (mem_range.mp hs))] + have : armMean (bestArm ω) ω ≤ ucb (bestArm ω) s ω := by + simp only [armMean, ucb]; unfold ucbIndex + split_ifs with h0 + · exact (hm (bestArm ω) (E ω)).2 + · have := abs_lt.mp ((hω s (mem_range.mp hs)) h0) + exact le_max_of_le_right (le_min (hm (bestArm ω) (E ω)).2 (by linarith)) + linarith have h_second_Eδ : ∀ ω ∈ Eδ, ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω) ≤ (hi - lo) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by intro ω hω - exact sum_ucbIndex_sub_armMean_le E κ hm hlo (↑σ2) δ n ω hω + exact sum_ucbIndex_sub_mean_le (μ := fun a => armMean a ω) + (fun a => hm a (E ω)) hlo (↑σ2) δ n ω hω have h_prob : P Eδᶜ ≤ ENNReal.ofReal (2 * K * n * δ) := by have : Eδᶜ = {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ From 6a7a8e57d4421aea66159813438310e7f31307d7 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 13 Mar 2026 14:40:22 +0000 Subject: [PATCH 075/155] Minor --- LeanBandits/BanditAlgorithms/TS.lean | 14 ++------------ 1 file changed, 2 insertions(+), 12 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 67f88818..0cd3d62b 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -72,16 +72,6 @@ lemma ucbIndex_mem_Icc (h : l ≤ u) {a : Fin K} {n : ℕ} {ω : Ω} : unfold ucbIndex grind -lemma ucbIndex_sub_mean_le {lo hi μ : ℝ} (hμ : μ ∈ Set.Icc lo hi) - (σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) (hpc : pullCount A a t ω ≠ 0) - (h : - |empMean A R' a t ω - μ| - < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ))) : - ucbIndex A R' lo hi σ2 δ a t ω - μ - ≤ 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω : ℝ)) := by - unfold ucbIndex - grind - omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] (E : Ω → 𝓔) (κ : Kernel (𝓔 × Fin K) ℝ) {lo hi : ℝ} @@ -180,8 +170,8 @@ lemma sum_ucbIndex_sub_mean_le {lo hi : ℝ} {μ : Fin K → ℝ} 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) := sum_le_sum fun s hs => by have hpc : pullCount A (A s ω) s ω ≠ 0 := (Finset.mem_filter.mp hs).2 - exact ucbIndex_sub_mean_le (hm _) σ2 δ (A s ω) s ω hpc - (hconc s (mem_range.mp (Finset.mem_filter.mp hs).1) _ hpc) + unfold ucbIndex + grind _ ≤ ∑ s ∈ range n, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) := Finset.sum_le_sum_of_subset_of_nonneg From 61d7f38062dd037fa39f695c81071eb093fc9d3d Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 17 Mar 2026 10:56:32 +0000 Subject: [PATCH 076/155] Refactor TS.lean (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 242 ++++++------------ .../BayesStationaryEnv.lean | 73 +++++- 2 files changed, 148 insertions(+), 167 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 0cd3d62b..67de435a 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -53,67 +53,41 @@ namespace TS section UCB variable {Ω : Type*} +variable {l u σ2 δ : ℝ} noncomputable def ucbIndex (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ := if pullCount A a n ω = 0 then u else max l (min u (empMean A R' a n ω + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a n ω)))) -variable {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} -variable {l u σ2 δ : ℝ} - @[fun_prop] -lemma measurable_ucbIndex [MeasurableSpace Ω] {a : Fin K} {n : ℕ} (hA : ∀ t, Measurable (A t)) - (hR : ∀ t, Measurable (R' t)) : Measurable (ucbIndex A R' l u σ2 δ a n) := +lemma measurable_ucbIndex [MeasurableSpace Ω] {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} {a : Fin K} + {n : ℕ} (hA : ∀ t, Measurable (A t)) (hR : ∀ t, Measurable (R' t)) : + Measurable (ucbIndex A R' l u σ2 δ a n) := Measurable.ite (by measurability) (by fun_prop) (by fun_prop) -lemma ucbIndex_mem_Icc (h : l ≤ u) {a : Fin K} {n : ℕ} {ω : Ω} : - ucbIndex A R' l u σ2 δ a n ω ∈ Set.Icc l u := by +lemma ucbIndex_mem_Icc {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} (h : l ≤ u) {a : Fin K} {n : ℕ} + {ω : Ω} : ucbIndex A R' l u σ2 δ a n ω ∈ Set.Icc l u := by unfold ucbIndex grind -omit [StandardBorelSpace 𝓔] [Nonempty 𝓔] -lemma iSup_armMean_eq_bestArm [Nonempty (Fin K)] (E : Ω → 𝓔) - (κ : Kernel (𝓔 × Fin K) ℝ) {lo hi : ℝ} - (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) - (ω : Ω) : ⨆ i, IsBayesAlgEnvSeq.actionMean κ E i ω = - IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω := - le_antisymm (ciSup_le fun i ↦ by - have := isMaxOn_measurableArgmax (fun ω a ↦ IsBayesAlgEnvSeq.actionMean κ E a ω) ω i - simp only [IsBayesAlgEnvSeq.bestAction]; convert this) - (le_ciSup (f := fun i ↦ IsBayesAlgEnvSeq.actionMean κ E i ω) - ⟨hi, by rintro _ ⟨i, rfl⟩; exact (hm i _).2⟩ _) +noncomputable +def ucbIndex' (n : ℕ) (h : Iic n → Fin K × ℝ) (l u σ2 δ : ℝ) (a : Fin K) : ℝ := + if pullCount' n h a = 0 then u + else max l (min u (empMean' n h a + √(2 * σ2 * Real.log (1 / δ) / (pullCount' n h a)))) -lemma gap_eq_armMean_sub [Nonempty (Fin K)] (E : Ω → 𝓔) - (κ : Kernel (𝓔 × Fin K) ℝ) {lo hi : ℝ} - (hm : ∀ a e, (κ (e, a))[id] ∈ Set.Icc lo hi) - (s : ℕ) (ω : Ω) : gap (κ.sectR (E ω)) (A s ω) = - IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω := by - simp only [gap, Kernel.sectR_apply] - exact congr_arg (· - _) (iSup_armMean_eq_bestArm E κ hm ω) +@[fun_prop] +lemma measurable_ucbIndex' {a : Fin K} {n : ℕ} : Measurable (fun h ↦ ucbIndex' n h l u σ2 δ a) := + Measurable.ite (by measurability) (by fun_prop) (by fun_prop) -lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : - ∑ s ∈ range n, f (pullCount A (A s ω) s ω) = - ∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω), f j := by - induction n with - | zero => simp - | succ n ih => - rw [sum_range_succ, ih] - suffices ∑ a, ∑ j ∈ range (pullCount A a (n + 1) ω), f j = - (∑ a, ∑ j ∈ range (pullCount A a n ω), f j) + - f (pullCount A (A n ω) n ω) by linarith - have h_eq : ∀ a, ∑ j ∈ range (pullCount A a (n + 1) ω), f j = - ∑ j ∈ range (pullCount A a n ω), f j + - if A n ω = a then f (pullCount A a n ω) else 0 := by - intro a - rw [pullCount_add_one] - split_ifs with h - · rw [sum_range_succ] - · simp - simp_rw [h_eq, sum_add_distrib] - congr 1 - simp +lemma ucbIndex_succ_eq_ucbIndex'_hist (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (a : Fin K) + (n : ℕ) (ω : Ω) : + ucbIndex A R' l u σ2 δ a (n + 1) ω = ucbIndex' n (IsAlgEnvSeq.hist A R' n ω) l u σ2 δ a := by + have hpc : pullCount A a (n + 1) ω = pullCount' n (IsAlgEnvSeq.hist A R' n ω) a := + pullCount_add_one_eq_pullCount' + have hem : empMean A R' a (n + 1) ω = empMean' n (IsAlgEnvSeq.hist A R' n ω) a := + empMean_add_one_eq_empMean' + simp_rw [ucbIndex, ucbIndex', hpc, hem] /-- Helper for `sum_ucbIndex_sub_armMean_le`. -/ private lemma sum_sqrt_le {ι : Type*} {c : ι → ℝ} (s : Finset ι) (hc : ∀ i, 0 ≤ c i) : @@ -140,6 +114,31 @@ private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 have : √n * √n = n := Real.mul_self_sqrt (by positivity) nlinarith +variable {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} + +/-- Helper for `sum_ucbIndex_sub_armMean_le`. -/ +lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : + ∑ s ∈ range n, f (pullCount A (A s ω) s ω) = + ∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω), f j := by + induction n with + | zero => simp + | succ n ih => + rw [sum_range_succ, ih] + suffices ∑ a, ∑ j ∈ range (pullCount A a (n + 1) ω), f j = + (∑ a, ∑ j ∈ range (pullCount A a n ω), f j) + + f (pullCount A (A n ω) n ω) by linarith + have h_eq : ∀ a, ∑ j ∈ range (pullCount A a (n + 1) ω), f j = + ∑ j ∈ range (pullCount A a n ω), f j + + if A n ω = a then f (pullCount A a n ω) else 0 := by + intro a + rw [pullCount_add_one] + split_ifs with h + · rw [sum_range_succ] + · simp + simp_rw [h_eq, sum_add_distrib] + congr 1 + simp + lemma sum_ucbIndex_sub_mean_le {lo hi : ℝ} {μ : Fin K → ℝ} (hm : ∀ a, μ a ∈ Set.Icc lo hi) (hlo : lo ≤ hi) (σ2 δ : ℝ) (n : ℕ) (ω : Ω) @@ -252,39 +251,6 @@ lemma sum_ucbIndex_sub_mean_le {lo hi : ℝ} {μ : Fin K → ℝ} _ ≤ ↑K * (hi - lo) := by gcongr; linarith _ = (hi - lo) * ↑K := by ring -/-- History-based UCB index: like `ucbIndex` but takes a history `h : Iic t → Fin K × ℝ` -directly instead of the random variables `A` and `R'`. -/ -noncomputable -def ucbIndex' (lo hi σ2 δ : ℝ) (t : ℕ) (h : Iic t → Fin K × ℝ) (a : Fin K) : ℝ := - if pullCount' t h a = 0 then hi - else max lo (min hi (empMean' t h a + - √(2 * σ2 * Real.log (1 / δ) / (pullCount' t h a : ℝ)))) - -@[fun_prop] -lemma measurable_ucbIndex' {lo hi σ2 δ : ℝ} {a : Fin K} {t : ℕ} : - Measurable (fun h : Iic t → Fin K × ℝ ↦ ucbIndex' lo hi σ2 δ t h a) := by - unfold ucbIndex' - apply Measurable.ite - · have : MeasurableSet {h : Iic t → Fin K × ℝ | (pullCount' t h a : ℝ) = (0 : ℝ)} := - measurableSet_eq_fun - (measurable_from_top.comp (measurable_pullCount' t a)) - measurable_const - simp only [Nat.cast_eq_zero] at this; exact this - · exact measurable_const - · exact .max measurable_const (.min measurable_const - (.add (measurable_empMean' t a) - (measurable_const.div (measurable_from_top.comp (measurable_pullCount' t a))).sqrt)) - -lemma ucbIndex_succ_eq_ucbIndex'_hist (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) - (lo hi σ2 δ : ℝ) (a : Fin K) (t : ℕ) (ω : Ω) : - ucbIndex A R' lo hi σ2 δ a (t + 1) ω = - ucbIndex' lo hi σ2 δ t (IsAlgEnvSeq.hist A R' t ω) a := by - have hpc : pullCount A a (t + 1) ω = pullCount' t (IsAlgEnvSeq.hist A R' t ω) a := - pullCount_add_one_eq_pullCount' - have hem : empMean A R' a (t + 1) ω = empMean' t (IsAlgEnvSeq.hist A R' t ω) a := - empMean_add_one_eq_empMean' - simp only [ucbIndex, ucbIndex', hpc, hem] - end UCB /-! ### Concentration bounds (algorithm-generic) @@ -468,30 +434,6 @@ variable (P : Measure Ω) [IsProbabilityMeasure P] namespace Learning.IsBayesAlgEnvSeq -lemma bayesRegret_eq_sum_integral_gap [Nonempty (Fin K)] - {alg : Algorithm (Fin K) ℝ} - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - {C : ℝ} (hm : ∀ a e, |(κ (e, a))[id]| ≤ C) (t : ℕ) : - P[IsBayesAlgEnvSeq.regret κ E A t] = - ∑ s ∈ range t, P[fun ω ↦ gap (κ.sectR (E ω)) - (A s ω)] := by - simp only [IsBayesAlgEnvSeq.regret, regret_eq_sum_gap] - refine integral_finset_sum _ (fun s _ => ?_) - have hmeas : Measurable (fun ω ↦ gap (κ.sectR (E ω)) - (A s ω)) := - (Measurable.iSup (fun a ↦ IsBayesAlgEnvSeq.measurable_actionMean - (a := a) h.measurable_E)).sub - (stronglyMeasurable_id.integral_kernel.measurable.comp - (h.measurable_E.prodMk (h.measurable_A s))) - refine ⟨hmeas.aestronglyMeasurable, HasFiniteIntegral.of_bounded (C := 2 * C) - (Filter.Eventually.of_forall fun ω => ?_)⟩ - simp only [Real.norm_eq_abs, gap, Kernel.sectR_apply] - have hbdd : BddAbove (Set.range fun i => (κ (E ω, i))[id]) := - ⟨C, by rintro _ ⟨i, rfl⟩; exact le_of_abs_le (hm i _)⟩ - rw [abs_of_nonneg (sub_nonneg.mpr (le_ciSup hbdd _))] - linarith [ciSup_le fun i => le_of_abs_le (hm i (E ω)), - neg_le_of_abs_le (hm (A s ω) (E ω))] - variable [IsMarkovKernel κ] lemma prob_concentration_fail_delta [Nonempty (Fin K)] @@ -740,23 +682,23 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) + {l u : ℝ} (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : P[IsBayesAlgEnvSeq.regret κ E A n] - ≤ (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * n ^ 2 * δ + + ≤ (u - l) * ↑K + 2 * (↑K + 1) * (u - l) * n ^ 2 * δ + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _) - have hlo : lo ≤ hi := h1.trans h2 + have hlo : l ≤ u := h1.trans h2 let bestArm := IsBayesAlgEnvSeq.bestAction κ E let armMean := IsBayesAlgEnvSeq.actionMean κ E - let ucb := ucbIndex A R' lo hi (↑σ2) δ + let ucb := ucbIndex A R' l u (↑σ2) δ set Eδ : Set Ω := {ω | ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → |empMean A R' a s ω - armMean a ω| < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ))} set Fδ : Set Ω := {ω | ∀ s < n, pullCount A (bestArm ω) s ω ≠ 0 → |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω| < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A (bestArm ω) s ω : ℝ))} - have hm_ucb : ∀ a t, Measurable (ucbIndex A R' lo hi (↑σ2) δ a t) := + have hm_ucb : ∀ a t, Measurable (ucbIndex A R' l u (↑σ2) δ a t) := fun _ _ ↦ measurable_ucbIndex h.measurable_A h.measurable_R have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.actionMean κ E a) := fun a ↦ IsBayesAlgEnvSeq.measurable_actionMean (a := a) h.measurable_E @@ -764,39 +706,39 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E have h_first_bound : ∀ ω, |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| - ≤ n * (hi - lo) := fun ω ↦ + ≤ n * (u - l) := fun ω ↦ calc |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| ≤ ∑ s ∈ range n, |armMean (bestArm ω) ω - ucb (bestArm ω) s ω| := Finset.abs_sum_le_sum_abs _ _ - _ ≤ ∑ s ∈ range n, (hi - lo) := by + _ ≤ ∑ s ∈ range n, (u - l) := by gcongr with s _ exact abs_sub_le_of_le_of_le (hm _ _).1 (hm _ _).2 ((ucbIndex_mem_Icc hlo).1) (ucbIndex_mem_Icc hlo).2 - _ = ↑n * (hi - lo) := by + _ = ↑n * (u - l) := by rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] have h_second_bound : ∀ ω, |∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)| - ≤ n * (hi - lo) := fun ω ↦ + ≤ n * (u - l) := fun ω ↦ calc |∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)| ≤ ∑ s ∈ range n, |ucb (A s ω) s ω - armMean (A s ω) ω| := Finset.abs_sum_le_sum_abs _ _ - _ ≤ ∑ s ∈ range n, (hi - lo) := by + _ ≤ ∑ s ∈ range n, (u - l) := by gcongr with s _ exact abs_sub_le_of_le_of_le (ucbIndex_mem_Icc hlo).1 (ucbIndex_mem_Icc hlo).2 (hm _ _).1 (hm _ _).2 - _ = ↑n * (hi - lo) := by + _ = ↑n * (u - l) := by rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] have h_int_sum1 : Integrable (fun ω ↦ ∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)) P := by - apply Integrable.of_bound (C := ↑n * (hi - lo)) + apply Integrable.of_bound (C := ↑n * (u - l)) · exact (Finset.measurable_fun_sum _ fun s _ ↦ (measurable_apply_fin hm_arm hm_best).sub (measurable_apply_fin (fun a ↦ hm_ucb a s) hm_best)).aestronglyMeasurable · filter_upwards with ω; rw [Real.norm_eq_abs]; exact h_first_bound ω have h_int_sum2 : Integrable (fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)) P := by - apply Integrable.of_bound (C := ↑n * (hi - lo)) + apply Integrable.of_bound (C := ↑n * (u - l)) · exact (Finset.measurable_fun_sum _ fun s _ ↦ (measurable_apply_fin (fun a ↦ hm_ucb a s) (h.measurable_A s)).sub (measurable_apply_fin hm_arm (h.measurable_A s))).aestronglyMeasurable @@ -807,12 +749,6 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)] + P[fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)] := by - have h_regret_eq : P[IsBayesAlgEnvSeq.regret κ E A n] = - ∑ s ∈ range n, P[fun ω ↦ armMean (bestArm ω) ω - armMean (A s ω) ω] := by - rw [IsBayesAlgEnvSeq.bayesRegret_eq_sum_integral_gap (h := h) - (hm := fun a e ↦ abs_le_max_abs_abs (hm a e).1 (hm a e).2) (t := n)] - congr 1 with s - exact integral_congr_ae (ae_of_all _ fun ω ↦ gap_eq_armMean_sub E κ hm s ω) have h_int_ucb : ∀ s {f : Ω → Fin K}, Measurable f → Integrable (fun ω ↦ ucb (f ω) s ω) P := fun s {_} hf ↦ ⟨(measurable_apply_fin (fun a ↦ hm_ucb a s) hf).aestronglyMeasurable, @@ -822,14 +758,14 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (ucbIndex_mem_Icc hlo).2)⟩ have h_int_ucb_sub : ∀ s, Integrable (fun ω ↦ ucb (A s ω) s ω - ucb (bestArm ω) s ω) P := fun s ↦ (h_int_ucb s (h.measurable_A s)).sub (h_int_ucb s hm_best) - have h_ucb_zero : ∀ a (ω : Ω), ucbIndex A R' lo hi (↑σ2) δ a 0 ω = hi := by + have h_ucb_zero : ∀ a (ω : Ω), ucbIndex A R' l u (↑σ2) δ a 0 ω = u := by intro a ω; unfold ucbIndex; simp [pullCount_zero] have h_ucb_swap : ∀ s, ∫ ω, (ucb (A s ω) s ω - ucb (bestArm ω) s ω) ∂P = 0 := by intro s cases s with | zero => have : ∀ ω, ucb (A 0 ω) 0 ω - ucb (bestArm ω) 0 ω = 0 := fun ω ↦ by - change ucbIndex A R' lo hi (↑σ2) δ _ 0 ω - ucbIndex A R' lo hi (↑σ2) δ _ 0 ω = 0 + change ucbIndex A R' l u (↑σ2) δ _ 0 ω - ucbIndex A R' l u (↑σ2) δ _ 0 ω = 0 simp [h_ucb_zero] exact (integral_congr_ae (ae_of_all _ this)).trans (integral_zero _ _) | succ t => @@ -852,12 +788,12 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp hf.aestronglyMeasurable, h_map_eq] set g : (Iic t → Fin K × ℝ) × Fin K → ℝ := - fun p ↦ if pullCount' t p.1 p.2 = 0 then hi - else max lo (min hi (empMean' t p.1 p.2 + + fun p ↦ if pullCount' t p.1 p.2 = 0 then u + else max l (min u (empMean' t p.1 p.2 + √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount' t p.1 p.2 : ℝ)))) - have hg_eq : ∀ a (ω : Ω), ucbIndex A R' lo hi (↑σ2) δ a (t + 1) ω = + have hg_eq : ∀ a (ω : Ω), ucbIndex A R' l u (↑σ2) δ a (t + 1) ω = g (IsAlgEnvSeq.hist A R' t ω, a) := - fun a ω ↦ ucbIndex_succ_eq_ucbIndex'_hist A R' lo hi (↑σ2) δ a t ω + fun a ω ↦ ucbIndex_succ_eq_ucbIndex'_hist A R' a t ω have hg_meas : Measurable g := by apply Measurable.ite · have : MeasurableSet {p : (Iic t → Fin K × ℝ) × Fin K | @@ -890,7 +826,9 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (ucb (A s ω) s ω - ucb (bestArm ω) s ω) ∂P = 0 := by rw [integral_finset_sum _ (fun s _ ↦ h_int_ucb_sub s)] exact Finset.sum_eq_zero fun s _ ↦ h_ucb_swap s - rw [h_regret_eq] + have h_int_gap : Integrable (fun ω ↦ IsBayesAlgEnvSeq.regret κ E A n ω) P := + IsBayesAlgEnvSeq.integrable_regret h.measurable_E (h.measurable_A) hm + simp_rw [IsBayesAlgEnvSeq.regret_eq_sum_gap, IsBayesAlgEnvSeq.gap_eq_sub] at h_int_gap ⊢ have h_pw : ∀ ω, (∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)) + (∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)) = (∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω)) + @@ -901,22 +839,8 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp have h_int_ucb_swap : Integrable (fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - ucb (bestArm ω) s ω)) P := integrable_finset_sum _ fun s _ ↦ h_int_ucb_sub s - have h_int_gap_s : ∀ s ∈ range n, - Integrable (fun ω ↦ armMean (bestArm ω) ω - armMean (A s ω) ω) P := by - intro s _ - refine Integrable.of_bound - ((measurable_apply_fin hm_arm hm_best).sub - (measurable_apply_fin hm_arm (h.measurable_A s))).aestronglyMeasurable - (hi - lo) (ae_of_all _ fun ω ↦ ?_) - rw [Real.norm_eq_abs] - exact abs_sub_le_of_le_of_le (hm _ _).1 (hm _ _).2 (hm _ _).1 (hm _ _).2 - have h_int_gap : Integrable - (fun ω ↦ ∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω)) P := - integrable_finset_sum _ h_int_gap_s - calc ∑ s ∈ range n, ∫ (x : Ω), (fun ω ↦ armMean (bestArm ω) ω - armMean (A s ω) ω) x ∂P - = ∫ ω, ∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω) ∂P := - (integral_finset_sum _ h_int_gap_s).symm - _ = ∫ ω, ((∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω)) + + calc ∫ ω, ∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω) ∂P + = ∫ ω, ((∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω)) + (∑ s ∈ range n, (ucb (A s ω) s ω - ucb (bestArm ω) s ω))) ∂P := by rw [integral_add h_int_gap h_int_ucb_swap, h_ucb_sum_zero, add_zero] _ = ∫ ω, ((∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)) + @@ -932,16 +856,16 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp have : armMean (bestArm ω) ω ≤ ucb (bestArm ω) s ω := by simp only [armMean, ucb]; unfold ucbIndex split_ifs with h0 - · exact (hm (bestArm ω) (E ω)).2 + · exact (hm (E ω) (bestArm ω)).2 · have := abs_lt.mp ((hω s (mem_range.mp hs)) h0) - exact le_max_of_le_right (le_min (hm (bestArm ω) (E ω)).2 (by linarith)) + exact le_max_of_le_right (le_min (hm (E ω) (bestArm ω)).2 (by linarith)) linarith have h_second_Eδ : ∀ ω ∈ Eδ, ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω) - ≤ (hi - lo) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by + ≤ (u - l) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by intro ω hω exact sum_ucbIndex_sub_mean_le (μ := fun a => armMean a ω) - (fun a => hm a (E ω)) hlo (↑σ2) δ n ω hω + (hm (E ω)) hlo (↑σ2) δ n ω hω have h_prob : P Eδᶜ ≤ ENNReal.ofReal (2 * K * n * δ) := by have : Eδᶜ = {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ @@ -992,21 +916,21 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (armMean (bestArm ω) ω - ucb (bestArm ω) s ω) set f2 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω) - set B := (hi - lo) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) + set B := (u - l) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) have h1g : ∫ ω in Fδ, f1 ω ∂P ≤ 0 := setIntegral_nonpos hFδ_meas fun ω hω ↦ h_first_Fδ ω hω - have h1b : ∫ ω in Fδᶜ, f1 ω ∂P ≤ ↑n * (hi - lo) * P.real Fδᶜ := by + have h1b : ∫ ω in Fδᶜ, f1 ω ∂P ≤ ↑n * (u - l) * P.real Fδᶜ := by have := setIntegral_mono_on (hf := h_int_sum1.integrableOn) (hg := integrableOn_const) hFδ_meas.compl fun ω _ ↦ (abs_le.mp (h_first_bound ω)).2 rwa [setIntegral_const, smul_eq_mul, mul_comm] at this have h2g : ∫ ω in Eδ, f2 ω ∂P ≤ B := by - have hB : 0 ≤ B := by have : 0 ≤ hi - lo := sub_nonneg.mpr hlo; positivity + have hB : 0 ≤ B := by have : 0 ≤ u - l := sub_nonneg.mpr hlo; positivity have := setIntegral_mono_on (hf := h_int_sum2.integrableOn) (hg := integrableOn_const) hEδ_meas fun ω hω ↦ h_second_Eδ ω hω rw [setIntegral_const, smul_eq_mul, mul_comm] at this exact le_trans this (mul_le_of_le_one_right hB measureReal_le_one) - have h2b : ∫ ω in Eδᶜ, f2 ω ∂P ≤ ↑n * (hi - lo) * P.real Eδᶜ := by + have h2b : ∫ ω in Eδᶜ, f2 ω ∂P ≤ ↑n * (u - l) * P.real Eδᶜ := by have := setIntegral_mono_on (hf := h_int_sum2.integrableOn) (hg := integrableOn_const) hEδ_meas.compl fun ω _ ↦ (abs_le.mp (h_second_bound ω)).2 rwa [setIntegral_const, smul_eq_mul, mul_comm] at this @@ -1016,7 +940,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp ENNReal.toReal_le_of_le_ofReal (by positivity) h_prob rw [(integral_add_compl hFδ_meas h_int_sum1).symm, (integral_add_compl hEδ_meas h_int_sum2).symm] - have : 0 ≤ ↑n * (hi - lo) := by nlinarith + have : 0 ≤ ↑n * (u - l) := by nlinarith nlinarith [mul_le_mul_of_nonneg_left hPF this, mul_le_mul_of_nonneg_left hPE this, measureReal_nonneg (μ := P) (s := Fδᶜ), @@ -1026,7 +950,7 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - {lo hi : ℝ} (hm : ∀ a e, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (t : ℕ) : + {lo hi : ℝ} (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (t : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A t] ≤ (3 * K + 2) * (hi - lo) + 8 * √(σ2 * K * t * Real.log t) := by have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _) @@ -1044,10 +968,10 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] simp only [Finset.range_one, Finset.sum_singleton, Nat.cast_one, one_mul, Kernel.sectR_apply] refine (integral_mono_of_nonneg (ae_of_all _ fun ω ↦ sub_nonneg.mpr - (le_ciSup ⟨hi, by rintro _ ⟨a, rfl⟩; exact (hm a _).2⟩ _)) + (le_ciSup ⟨hi, by rintro _ ⟨a, rfl⟩; exact (hm _ a).2⟩ _)) (integrable_const (hi - lo)) (ae_of_all _ fun ω ↦ by - linarith [ciSup_le fun a ↦ (hm a (E ω)).2, - (hm (A 0 ω) (E ω)).1])).trans ?_ + linarith [ciSup_le fun a ↦ (hm (E ω) a).2, + (hm (E ω) (A 0 ω)).1])).trans ?_ simp _ ≤ (3 * ↑K + 2) * (hi - lo) := by nlinarith [Nat.one_le_cast (α := ℝ).mpr (Nat.one_le_of_lt hK), diff --git a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean index 2ebbfa5c..16dff851 100644 --- a/LeanBandits/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanBandits/SequentialLearning/BayesStationaryEnv.lean @@ -64,17 +64,74 @@ lemma measurable_bestAction [Nonempty α] [Fintype α] [Encodable α] [Measurabl {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} (hE : Measurable E) : Measurable (bestAction κ E) := measurable_measurableArgmax (by fun_prop) +/-- The gap at time `n`. -/ noncomputable -def regret (κ : Kernel (𝓔 × α) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → α) (t : ℕ) (ω : Ω) : ℝ := - Bandits.regret (κ.sectR (E ω)) A t ω +def gap (κ : Kernel (𝓔 × α) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → α) (n : ℕ) (ω : Ω) : ℝ := + Bandits.gap (κ.sectR (E ω)) (A n ω) + +omit [MeasurableSpace Ω] in +/-- The gap is non-negative if the means are bounded by `u : ℝ` (even if `α` is not `Finite`). -/ +lemma gap_nonneg_of_le [Nonempty α] {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → α} {n : ℕ} + {ω : Ω} {u : ℝ} (h : ∀ e a, (κ (e, a))[id] ≤ u) : 0 ≤ gap κ E A n ω := by + simp_rw [gap, Bandits.gap, Kernel.sectR_apply] + linarith [le_ciSup ⟨u, Set.forall_mem_range.2 fun a ↦ (h (E ω) a)⟩ (A n ω)] + +omit [MeasurableSpace Ω] in +lemma gap_le_of_mem_Icc [Nonempty α] {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → α} {n : ℕ} + {ω : Ω} {l u : ℝ} (h : ∀ e a, (κ (e, a))[id] ∈ Set.Icc l u) : gap κ E A n ω ≤ u - l := by + simp_rw [gap, Bandits.gap, Kernel.sectR_apply] + grind [ciSup_le (fun a ↦ (h (E ω) a).2)] + +omit [MeasurableSpace Ω] in +lemma gap_eq_sub [Nonempty α] [Fintype α] [Encodable α] [MeasurableSingletonClass α] + {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → α} {n : ℕ} {ω : Ω} : + gap κ E A n ω = actionMean κ E (bestAction κ E ω) ω - actionMean κ E (A n ω) ω := by + rw [gap, Bandits.gap] + congr + apply le_antisymm + · exact ciSup_le (isMaxOn_measurableArgmax (fun ω' a ↦ actionMean κ E a ω') ω) + · exact Finite.le_ciSup (fun a ↦ actionMean κ E a ω) _ @[fun_prop] -lemma measurable_regret [Countable α] {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → α} {t : ℕ} - (hE : Measurable E) (hA : ∀ n, Measurable (A n)) : - Measurable (regret κ E A t) := by - have hm := (stronglyMeasurable_id.integral_kernel (κ := κ)).measurable - exact (Measurable.const_mul (Measurable.iSup fun _ ↦ (hm.comp (by fun_prop))) _).sub - (Finset.measurable_sum _ fun _ _ ↦ hm.comp (by fun_prop)) +lemma measurable_gap [Countable α] {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → α} {n : ℕ} + (hE : Measurable E) (hA : ∀ t, Measurable (A t)) : Measurable (gap κ E A n) := + (Measurable.iSup fun _ ↦ stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop)).sub + (stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop)) + +lemma integrable_gap [Countable α] [Nonempty α] {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} + {A : ℕ → Ω → α} {n : ℕ} {P : Measure Ω} [IsFiniteMeasure P] (hE : Measurable E) + (hA : ∀ t, Measurable (A t)) {l u : ℝ} (h : ∀ e a, (κ (e, a))[id] ∈ Set.Icc l u) : + Integrable (gap κ E A n) P := by + apply Integrable.of_bound (by fun_prop) (u - l) + filter_upwards with ω + rw [Real.norm_eq_abs, abs_of_nonneg (gap_nonneg_of_le (fun e a ↦ (h e a).2))] + exact gap_le_of_mem_Icc h + +noncomputable +def regret (κ : Kernel (𝓔 × α) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → α) (n : ℕ) (ω : Ω) : ℝ := + Bandits.regret (κ.sectR (E ω)) A n ω + +omit [MeasurableSpace Ω] in +lemma regret_eq_sum_gap {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → α} {n : ℕ} {ω : Ω} : + regret κ E A n ω = ∑ s ∈ range n, gap κ E A s ω := by + simp [regret, Bandits.regret, gap, Bandits.gap] + +omit [MeasurableSpace Ω] in +lemma regret_eq_sum_gap' {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → α} {n : ℕ} : + regret κ E A n = fun ω ↦ ∑ s ∈ range n, gap κ E A s ω := funext fun _ ↦ regret_eq_sum_gap + +@[fun_prop] +lemma measurable_regret [Countable α] {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → α} {n : ℕ} + (hE : Measurable E) (hA : ∀ t, Measurable (A t)) : Measurable (regret κ E A n) := by + rw [regret_eq_sum_gap'] + fun_prop + +lemma integrable_regret [Countable α] [Nonempty α] {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} + {A : ℕ → Ω → α} {n : ℕ} {P : Measure Ω} [IsFiniteMeasure P] (hE : Measurable E) + (hA : ∀ t, Measurable (A t)) {l u : ℝ} (h : ∀ e a, (κ (e, a))[id] ∈ Set.Icc l u) : + Integrable (regret κ E A n) P := by + rw [regret_eq_sum_gap'] + exact integrable_finset_sum _ (fun _ _ ↦ integrable_gap hE hA h) end Real From 4f36e384b133f84fcca4b013040cea256120c81f Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 17 Mar 2026 12:30:12 +0000 Subject: [PATCH 077/155] Refactor TS.lean (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 171 ++++++++---------- .../SequentialLearning/FiniteActions.lean | 17 ++ 2 files changed, 91 insertions(+), 97 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 67de435a..c142dc2c 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -56,38 +56,39 @@ variable {Ω : Type*} variable {l u σ2 δ : ℝ} noncomputable -def ucbIndex (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ := +def ucb (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ := if pullCount A a n ω = 0 then u else max l (min u (empMean A R' a n ω + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a n ω)))) @[fun_prop] -lemma measurable_ucbIndex [MeasurableSpace Ω] {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} {a : Fin K} +lemma measurable_ucb [MeasurableSpace Ω] {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} {a : Fin K} {n : ℕ} (hA : ∀ t, Measurable (A t)) (hR : ∀ t, Measurable (R' t)) : - Measurable (ucbIndex A R' l u σ2 δ a n) := + Measurable (ucb A R' l u σ2 δ a n) := Measurable.ite (by measurability) (by fun_prop) (by fun_prop) -lemma ucbIndex_mem_Icc {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} (h : l ≤ u) {a : Fin K} {n : ℕ} - {ω : Ω} : ucbIndex A R' l u σ2 δ a n ω ∈ Set.Icc l u := by - unfold ucbIndex +lemma ucb_mem_Icc {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} (h : l ≤ u) {a : Fin K} {n : ℕ} {ω : Ω} : + ucb A R' l u σ2 δ a n ω ∈ Set.Icc l u := by + unfold ucb grind noncomputable -def ucbIndex' (n : ℕ) (h : Iic n → Fin K × ℝ) (l u σ2 δ : ℝ) (a : Fin K) : ℝ := +def ucb' (n : ℕ) (h : Iic n → Fin K × ℝ) (l u σ2 δ : ℝ) (a : Fin K) : ℝ := if pullCount' n h a = 0 then u else max l (min u (empMean' n h a + √(2 * σ2 * Real.log (1 / δ) / (pullCount' n h a)))) @[fun_prop] -lemma measurable_ucbIndex' {a : Fin K} {n : ℕ} : Measurable (fun h ↦ ucbIndex' n h l u σ2 δ a) := +lemma measurable_uncurry_ucb' {n : ℕ} : + Measurable (fun p : (Iic n → Fin K × ℝ) × Fin K ↦ ucb' n p.1 l u σ2 δ p.2) := Measurable.ite (by measurability) (by fun_prop) (by fun_prop) lemma ucbIndex_succ_eq_ucbIndex'_hist (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : - ucbIndex A R' l u σ2 δ a (n + 1) ω = ucbIndex' n (IsAlgEnvSeq.hist A R' n ω) l u σ2 δ a := by + ucb A R' l u σ2 δ a (n + 1) ω = ucb' n (IsAlgEnvSeq.hist A R' n ω) l u σ2 δ a := by have hpc : pullCount A a (n + 1) ω = pullCount' n (IsAlgEnvSeq.hist A R' n ω) a := pullCount_add_one_eq_pullCount' have hem : empMean A R' a (n + 1) ω = empMean' n (IsAlgEnvSeq.hist A R' n ω) a := empMean_add_one_eq_empMean' - simp_rw [ucbIndex, ucbIndex', hpc, hem] + simp_rw [ucb, ucb', hpc, hem] /-- Helper for `sum_ucbIndex_sub_armMean_le`. -/ private lemma sum_sqrt_le {ι : Type*} {c : ι → ℝ} (s : Finset ι) (hc : ∀ i, 0 ≤ c i) : @@ -145,7 +146,7 @@ lemma sum_ucbIndex_sub_mean_le {lo hi : ℝ} {μ : Fin K → ℝ} (hconc : ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → |empMean A R' a s ω - μ a| < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ))) : - ∑ s ∈ range n, (ucbIndex A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) + ∑ s ∈ range n, (ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) ≤ (hi - lo) * ↑K + 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by -- Split range n into first-pull (pc=0) and non-first-pull (pc≠0) sets set S0 := (range n).filter (fun s => pullCount A (A s ω) s ω = 0) @@ -155,21 +156,21 @@ lemma sum_ucbIndex_sub_mean_le {lo hi : ℝ} {μ : Fin K → ℝ} conv_lhs => rw [hpart] rw [Finset.sum_union hdisj] -- We bound ∑_{S0} and ∑_{S1} separately, then combine - suffices h_S0 : ∑ s ∈ S0, (ucbIndex A R' lo hi σ2 δ (A s ω) s ω - + suffices h_S0 : ∑ s ∈ S0, (ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) ≤ (hi - lo) * ↑K by - suffices h_S1 : ∑ s ∈ S1, (ucbIndex A R' lo hi σ2 δ (A s ω) s ω - + suffices h_S1 : ∑ s ∈ S1, (ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) ≤ 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) by have := Finset.sum_union hdisj (f := fun s => - ucbIndex A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) + ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) rw [← hpart] at this; linarith -- Bound ∑_{S1}: each term ≤ 2√(2σ2c/pc) = 2√(2σ2c/max(1,pc)), so ≤ full sum - calc ∑ s ∈ S1, (ucbIndex A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) + calc ∑ s ∈ S1, (ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) ≤ ∑ s ∈ S1, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) := sum_le_sum fun s hs => by have hpc : pullCount A (A s ω) s ω ≠ 0 := (Finset.mem_filter.mp hs).2 - unfold ucbIndex + unfold ucb grind _ ≤ ∑ s ∈ range n, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) := @@ -222,10 +223,10 @@ lemma sum_ucbIndex_sub_mean_le {lo hi : ℝ} {μ : Fin K → ℝ} simp [sqrt_eq_zero'.mpr this] rw [sum_congr rfl h0]; simp only [sum_const_zero]; positivity -- Bound ∑_{S0}: each term = hi - μ ≤ hi - lo, and #S0 ≤ K - have hterm_S0 : ∀ s ∈ S0, ucbIndex A R' lo hi σ2 δ (A s ω) s ω - + have hterm_S0 : ∀ s ∈ S0, ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω) ≤ hi - lo := fun s hs => by have hpc : pullCount A (A s ω) s ω = 0 := (Finset.mem_filter.mp hs).2 - simp only [ucbIndex, hpc, ↓reduceIte] + simp only [ucb, hpc, ↓reduceIte] linarith [(hm (A s ω)).1] have h_card_S0 : #S0 ≤ K := by calc #S0 ≤ #(Finset.univ : Finset (Fin K)) := @@ -245,7 +246,7 @@ lemma sum_ucbIndex_sub_mean_le {lo hi : ℝ} {μ : Fin K → ℝ} exact absurd hpc₁ (show pullCount A (A s₁ ω) s₁ ω ≠ 0 from Finset.card_ne_zero_of_mem this)) _ = K := Finset.card_fin K - calc ∑ s ∈ S0, (ucbIndex A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) + calc ∑ s ∈ S0, (ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) ≤ ∑ _s ∈ S0, (hi - lo) := sum_le_sum hterm_S0 _ = #S0 * (hi - lo) := by rw [sum_const, nsmul_eq_mul] _ ≤ ↑K * (hi - lo) := by gcongr; linarith @@ -691,53 +692,53 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp have hlo : l ≤ u := h1.trans h2 let bestArm := IsBayesAlgEnvSeq.bestAction κ E let armMean := IsBayesAlgEnvSeq.actionMean κ E - let ucb := ucbIndex A R' l u (↑σ2) δ + let uc := ucb A R' l u (↑σ2) δ set Eδ : Set Ω := {ω | ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → |empMean A R' a s ω - armMean a ω| < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ))} set Fδ : Set Ω := {ω | ∀ s < n, pullCount A (bestArm ω) s ω ≠ 0 → |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω| < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A (bestArm ω) s ω : ℝ))} - have hm_ucb : ∀ a t, Measurable (ucbIndex A R' l u (↑σ2) δ a t) := - fun _ _ ↦ measurable_ucbIndex h.measurable_A h.measurable_R + have hm_ucb : ∀ a t, Measurable (ucb A R' l u (↑σ2) δ a t) := + fun _ _ ↦ measurable_ucb h.measurable_A h.measurable_R have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.actionMean κ E a) := fun a ↦ IsBayesAlgEnvSeq.measurable_actionMean (a := a) h.measurable_E have hm_best : Measurable (IsBayesAlgEnvSeq.bestAction κ E) := IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E have h_first_bound : ∀ ω, - |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| + |∑ s ∈ range n, (armMean (bestArm ω) ω - uc (bestArm ω) s ω)| ≤ n * (u - l) := fun ω ↦ - calc |∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)| - ≤ ∑ s ∈ range n, |armMean (bestArm ω) ω - ucb (bestArm ω) s ω| := + calc |∑ s ∈ range n, (armMean (bestArm ω) ω - uc (bestArm ω) s ω)| + ≤ ∑ s ∈ range n, |armMean (bestArm ω) ω - uc (bestArm ω) s ω| := Finset.abs_sum_le_sum_abs _ _ _ ≤ ∑ s ∈ range n, (u - l) := by gcongr with s _ exact abs_sub_le_of_le_of_le (hm _ _).1 (hm _ _).2 - ((ucbIndex_mem_Icc hlo).1) - (ucbIndex_mem_Icc hlo).2 + ((ucb_mem_Icc hlo).1) + (ucb_mem_Icc hlo).2 _ = ↑n * (u - l) := by rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] have h_second_bound : ∀ ω, - |∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)| + |∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω)| ≤ n * (u - l) := fun ω ↦ - calc |∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)| - ≤ ∑ s ∈ range n, |ucb (A s ω) s ω - armMean (A s ω) ω| := + calc |∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω)| + ≤ ∑ s ∈ range n, |uc (A s ω) s ω - armMean (A s ω) ω| := Finset.abs_sum_le_sum_abs _ _ _ ≤ ∑ s ∈ range n, (u - l) := by gcongr with s _ - exact abs_sub_le_of_le_of_le (ucbIndex_mem_Icc hlo).1 - (ucbIndex_mem_Icc hlo).2 (hm _ _).1 (hm _ _).2 + exact abs_sub_le_of_le_of_le (ucb_mem_Icc hlo).1 + (ucb_mem_Icc hlo).2 (hm _ _).1 (hm _ _).2 _ = ↑n * (u - l) := by rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] have h_int_sum1 : Integrable (fun ω ↦ ∑ s ∈ range n, - (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)) P := by + (armMean (bestArm ω) ω - uc (bestArm ω) s ω)) P := by apply Integrable.of_bound (C := ↑n * (u - l)) · exact (Finset.measurable_fun_sum _ fun s _ ↦ (measurable_apply_fin hm_arm hm_best).sub (measurable_apply_fin (fun a ↦ hm_ucb a s) hm_best)).aestronglyMeasurable · filter_upwards with ω; rw [Real.norm_eq_abs]; exact h_first_bound ω have h_int_sum2 : Integrable (fun ω ↦ ∑ s ∈ range n, - (ucb (A s ω) s ω - armMean (A s ω) ω)) P := by + (uc (A s ω) s ω - armMean (A s ω) ω)) P := by apply Integrable.of_bound (C := ↑n * (u - l)) · exact (Finset.measurable_fun_sum _ fun s _ ↦ (measurable_apply_fin (fun a ↦ hm_ucb a s) (h.measurable_A s)).sub @@ -746,26 +747,26 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp have h_swap : P[IsBayesAlgEnvSeq.regret κ E A n] = P[fun ω ↦ ∑ s ∈ range n, - (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)] + + (armMean (bestArm ω) ω - uc (bestArm ω) s ω)] + P[fun ω ↦ ∑ s ∈ range n, - (ucb (A s ω) s ω - armMean (A s ω) ω)] := by + (uc (A s ω) s ω - armMean (A s ω) ω)] := by have h_int_ucb : ∀ s {f : Ω → Fin K}, Measurable f → - Integrable (fun ω ↦ ucb (f ω) s ω) P := fun s {_} hf ↦ + Integrable (fun ω ↦ uc (f ω) s ω) P := fun s {_} hf ↦ ⟨(measurable_apply_fin (fun a ↦ hm_ucb a s) hf).aestronglyMeasurable, HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ by rw [Real.norm_eq_abs] - exact abs_le_max_abs_abs (ucbIndex_mem_Icc hlo).1 - (ucbIndex_mem_Icc hlo).2)⟩ - have h_int_ucb_sub : ∀ s, Integrable (fun ω ↦ ucb (A s ω) s ω - ucb (bestArm ω) s ω) P := + exact abs_le_max_abs_abs (ucb_mem_Icc hlo).1 + (ucb_mem_Icc hlo).2)⟩ + have h_int_ucb_sub : ∀ s, Integrable (fun ω ↦ uc (A s ω) s ω - uc (bestArm ω) s ω) P := fun s ↦ (h_int_ucb s (h.measurable_A s)).sub (h_int_ucb s hm_best) - have h_ucb_zero : ∀ a (ω : Ω), ucbIndex A R' l u (↑σ2) δ a 0 ω = u := by - intro a ω; unfold ucbIndex; simp [pullCount_zero] - have h_ucb_swap : ∀ s, ∫ ω, (ucb (A s ω) s ω - ucb (bestArm ω) s ω) ∂P = 0 := by + have h_ucb_zero : ∀ a (ω : Ω), ucb A R' l u (↑σ2) δ a 0 ω = u := by + intro a ω; unfold ucb; simp [pullCount_zero] + have h_ucb_swap : ∀ s, ∫ ω, (uc (A s ω) s ω - uc (bestArm ω) s ω) ∂P = 0 := by intro s cases s with | zero => - have : ∀ ω, ucb (A 0 ω) 0 ω - ucb (bestArm ω) 0 ω = 0 := fun ω ↦ by - change ucbIndex A R' l u (↑σ2) δ _ 0 ω - ucbIndex A R' l u (↑σ2) δ _ 0 ω = 0 + have : ∀ ω, uc (A 0 ω) 0 ω - uc (bestArm ω) 0 ω = 0 := fun ω ↦ by + change ucb A R' l u (↑σ2) δ _ 0 ω - ucb A R' l u (↑σ2) δ _ 0 ω = 0 simp [h_ucb_zero] exact (integral_congr_ae (ae_of_all _ this)).trans (integral_zero _ _) | succ t => @@ -787,81 +788,60 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (hm_hist.prodMk hm_best).aemeasurable hf.aestronglyMeasurable, h_map_eq] - set g : (Iic t → Fin K × ℝ) × Fin K → ℝ := - fun p ↦ if pullCount' t p.1 p.2 = 0 then u - else max l (min u (empMean' t p.1 p.2 + - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount' t p.1 p.2 : ℝ)))) - have hg_eq : ∀ a (ω : Ω), ucbIndex A R' l u (↑σ2) δ a (t + 1) ω = + let g : (Iic t → Fin K × ℝ) × Fin K → ℝ := + fun p ↦ ucb' t p.1 l u (↑σ2) δ p.2 + have hg_eq : ∀ a (ω : Ω), ucb A R' l u (↑σ2) δ a (t + 1) ω = g (IsAlgEnvSeq.hist A R' t ω, a) := fun a ω ↦ ucbIndex_succ_eq_ucbIndex'_hist A R' a t ω - have hg_meas : Measurable g := by - apply Measurable.ite - · have : MeasurableSet {p : (Iic t → Fin K × ℝ) × Fin K | - (pullCount' t p.1 p.2 : ℝ) = (0 : ℝ)} := - measurableSet_eq_fun - (measurable_apply_fin - (fun a ↦ measurable_from_top.comp - ((measurable_pullCount' t a).comp measurable_fst)) - measurable_snd) - measurable_const - simp only [Nat.cast_eq_zero] at this; exact this - · exact measurable_const - · exact .max measurable_const (.min measurable_const - (.add (measurable_apply_fin - (fun a ↦ (measurable_empMean' t a).comp measurable_fst) - measurable_snd) - (measurable_const.div (measurable_apply_fin - (fun a ↦ measurable_from_top.comp - ((measurable_pullCount' t a).comp measurable_fst)) - measurable_snd)).sqrt)) - rw [show (fun ω ↦ ucb (A (t + 1) ω) (t + 1) ω - - ucb (bestArm ω) (t + 1) ω) = - fun ω ↦ (fun ω ↦ ucb (A (t + 1) ω) (t + 1) ω) ω - - (fun ω ↦ ucb (bestArm ω) (t + 1) ω) ω from rfl, + have hg_meas : Measurable g := measurable_uncurry_ucb' + rw [show (fun ω ↦ uc (A (t + 1) ω) (t + 1) ω - + uc (bestArm ω) (t + 1) ω) = + fun ω ↦ (fun ω ↦ uc (A (t + 1) ω) (t + 1) ω) ω - + (fun ω ↦ uc (bestArm ω) (t + 1) ω) ω from rfl, integral_sub (h_int_ucb (t + 1) (h.measurable_A (t + 1))) (h_int_ucb (t + 1) hm_best), funext fun ω ↦ hg_eq _ _, funext fun ω ↦ hg_eq _ _, h_int_eq g hg_meas, sub_self] have h_ucb_sum_zero : ∫ ω, ∑ s ∈ range n, - (ucb (A s ω) s ω - ucb (bestArm ω) s ω) ∂P = 0 := by + (uc (A s ω) s ω - uc (bestArm ω) s ω) ∂P = 0 := by rw [integral_finset_sum _ (fun s _ ↦ h_int_ucb_sub s)] exact Finset.sum_eq_zero fun s _ ↦ h_ucb_swap s have h_int_gap : Integrable (fun ω ↦ IsBayesAlgEnvSeq.regret κ E A n ω) P := IsBayesAlgEnvSeq.integrable_regret h.measurable_E (h.measurable_A) hm simp_rw [IsBayesAlgEnvSeq.regret_eq_sum_gap, IsBayesAlgEnvSeq.gap_eq_sub] at h_int_gap ⊢ - have h_pw : ∀ ω, (∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)) + - (∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω)) = + have h_pw : ∀ ω, (∑ s ∈ range n, (armMean (bestArm ω) ω - uc (bestArm ω) s ω)) + + (∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω)) = (∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω)) + - (∑ s ∈ range n, (ucb (A s ω) s ω - ucb (bestArm ω) s ω)) := by + (∑ s ∈ range n, (uc (A s ω) s ω - uc (bestArm ω) s ω)) := by intro ω simp only [← Finset.sum_add_distrib] apply Finset.sum_congr rfl; intros; ring have h_int_ucb_swap : Integrable - (fun ω ↦ ∑ s ∈ range n, (ucb (A s ω) s ω - ucb (bestArm ω) s ω)) P := + (fun ω ↦ ∑ s ∈ range n, (uc (A s ω) s ω - uc (bestArm ω) s ω)) P := integrable_finset_sum _ fun s _ ↦ h_int_ucb_sub s calc ∫ ω, ∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω) ∂P = ∫ ω, ((∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω)) + - (∑ s ∈ range n, (ucb (A s ω) s ω - ucb (bestArm ω) s ω))) ∂P := by + (∑ s ∈ range n, (uc (A s ω) s ω - uc (bestArm ω) s ω))) ∂P := by rw [integral_add h_int_gap h_int_ucb_swap, h_ucb_sum_zero, add_zero] - _ = ∫ ω, ((∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω)) + - (∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω))) ∂P := by + _ = ∫ ω, ((∑ s ∈ range n, (armMean (bestArm ω) ω - uc (bestArm ω) s ω)) + + (∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω))) ∂P := by congr 1; ext ω; linarith [h_pw ω] _ = _ := integral_add h_int_sum1 h_int_sum2 have h_first_Fδ : ∀ ω ∈ Fδ, - ∑ s ∈ range n, (armMean (bestArm ω) ω - ucb (bestArm ω) s ω) + ∑ s ∈ range n, (armMean (bestArm ω) ω - uc (bestArm ω) s ω) ≤ 0 := by intro ω hω apply Finset.sum_nonpos intro s hs - have : armMean (bestArm ω) ω ≤ ucb (bestArm ω) s ω := by - simp only [armMean, ucb]; unfold ucbIndex + have : armMean (bestArm ω) ω ≤ uc (bestArm ω) s ω := by + simp only [armMean, uc]; unfold ucb split_ifs with h0 · exact (hm (E ω) (bestArm ω)).2 · have := abs_lt.mp ((hω s (mem_range.mp hs)) h0) exact le_max_of_le_right (le_min (hm (E ω) (bestArm ω)).2 (by linarith)) linarith have h_second_Eδ : ∀ ω ∈ Eδ, - ∑ s ∈ range n, (ucb (A s ω) s ω - armMean (A s ω) ω) + ∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω) ≤ (u - l) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by intro ω hω exact sum_ucbIndex_sub_mean_le (μ := fun a => armMean a ω) @@ -913,9 +893,9 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (Q := Q) (κ := κ) (P := P) h hσ2 hs n δ hδ hδ1 rw [h_swap] set f1 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, - (armMean (bestArm ω) ω - ucb (bestArm ω) s ω) + (armMean (bestArm ω) ω - uc (bestArm ω) s ω) set f2 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, - (ucb (A s ω) s ω - armMean (A s ω) ω) + (uc (A s ω) s ω - armMean (A s ω) ω) set B := (u - l) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) have h1g : ∫ ω in Fδ, f1 ω ∂P ≤ 0 := setIntegral_nonpos hFδ_meas fun ω hω ↦ h_first_Fδ ω hω @@ -964,15 +944,12 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] simp only [Nat.cast_one, Real.log_one, mul_zero, Real.sqrt_zero, mul_zero, add_zero] calc P[IsBayesAlgEnvSeq.regret κ E A 1] ≤ hi - lo := by - unfold IsBayesAlgEnvSeq.regret Bandits.regret - simp only [Finset.range_one, Finset.sum_singleton, Nat.cast_one, one_mul, - Kernel.sectR_apply] - refine (integral_mono_of_nonneg (ae_of_all _ fun ω ↦ sub_nonneg.mpr - (le_ciSup ⟨hi, by rintro _ ⟨a, rfl⟩; exact (hm _ a).2⟩ _)) - (integrable_const (hi - lo)) (ae_of_all _ fun ω ↦ by - linarith [ciSup_le fun a ↦ (hm (E ω) a).2, - (hm (E ω) (A 0 ω)).1])).trans ?_ - simp + rw [IsBayesAlgEnvSeq.regret_eq_sum_gap'] + simp only [Finset.range_one, Finset.sum_singleton] + exact (integral_mono_of_nonneg + (ae_of_all _ fun ω ↦ IsBayesAlgEnvSeq.gap_nonneg_of_le (fun e a ↦ (hm e a).2)) + (integrable_const _) + (ae_of_all _ fun ω ↦ IsBayesAlgEnvSeq.gap_le_of_mem_Icc hm)).trans (by simp) _ ≤ (3 * ↑K + 2) * (hi - lo) := by nlinarith [Nat.one_le_cast (α := ℝ).mpr (Nat.one_le_of_lt hK), sub_nonneg.mpr hlo] diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index b46415be..a2c2cb67 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -196,6 +196,7 @@ lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop +@[fun_prop] lemma measurable_uncurry_pullCount' [MeasurableEq α] (n : ℕ) : Measurable (fun p : (Iic n → α × R) × α ↦ pullCount' n p.1 p.2) := by simp_rw [pullCount'_eq_sum] @@ -861,12 +862,28 @@ lemma measurable_sumRewards' [MeasurableSingletonClass α] (n : ℕ) (a : α) : exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop +@[fun_prop] +lemma measurable_uncurry_sumRewards' [MeasurableEq α] (n : ℕ) : + Measurable (fun p : (Iic n → α × ℝ) × α ↦ sumRewards' n p.1 p.2) := by + simp_rw [sumRewards'] + have h_meas s : Measurable (fun p : (Iic n → α × ℝ) × α ↦ + if (p.1 s).1 = p.2 then (p.1 s).2 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_empMean' [MeasurableSingletonClass α] (n : ℕ) (a : α) : Measurable (fun h ↦ empMean' n h a) := by unfold empMean' fun_prop +@[fun_prop] +lemma measurable_uncurry_empMean' [MeasurableEq α] (n : ℕ) : + Measurable (fun p : (Iic n → α × ℝ) × α ↦ empMean' n p.1 p.2) := by + unfold empMean' + fun_prop + lemma IsAlgEnvSeq.isPredictable_sumRewards [StandardBorelSpace α] [Nonempty α] {R' : ℕ → Ω → ℝ} {alg : Algorithm α ℝ} {env : Environment α ℝ} (h : IsAlgEnvSeq A R' alg env P) (a : α) : From 35c01e6df32cefe080f1debb71b14bc8f5388a1d Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 17 Mar 2026 14:36:13 +0000 Subject: [PATCH 078/155] Refactor TS.lean (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 66 ++++++++++------------------ 1 file changed, 24 insertions(+), 42 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index c142dc2c..c331986b 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -53,6 +53,7 @@ namespace TS section UCB variable {Ω : Type*} +variable {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} variable {l u σ2 δ : ℝ} noncomputable @@ -61,12 +62,11 @@ def ucb (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) else max l (min u (empMean A R' a n ω + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a n ω)))) @[fun_prop] -lemma measurable_ucb [MeasurableSpace Ω] {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} {a : Fin K} - {n : ℕ} (hA : ∀ t, Measurable (A t)) (hR : ∀ t, Measurable (R' t)) : - Measurable (ucb A R' l u σ2 δ a n) := +lemma measurable_ucb [MeasurableSpace Ω] {a : Fin K} {n : ℕ} (hA : ∀ t, Measurable (A t)) + (hR : ∀ t, Measurable (R' t)) : Measurable (ucb A R' l u σ2 δ a n) := Measurable.ite (by measurability) (by fun_prop) (by fun_prop) -lemma ucb_mem_Icc {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} (h : l ≤ u) {a : Fin K} {n : ℕ} {ω : Ω} : +lemma ucb_mem_Icc (h : l ≤ u) {a : Fin K} {n : ℕ} {ω : Ω} : ucb A R' l u σ2 δ a n ω ∈ Set.Icc l u := by unfold ucb grind @@ -81,23 +81,22 @@ lemma measurable_uncurry_ucb' {n : ℕ} : Measurable (fun p : (Iic n → Fin K × ℝ) × Fin K ↦ ucb' n p.1 l u σ2 δ p.2) := Measurable.ite (by measurability) (by fun_prop) (by fun_prop) -lemma ucbIndex_succ_eq_ucbIndex'_hist (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (a : Fin K) - (n : ℕ) (ω : Ω) : +lemma ucb_succ_eq_ucb' {a : Fin K} {n : ℕ} {ω : Ω} : ucb A R' l u σ2 δ a (n + 1) ω = ucb' n (IsAlgEnvSeq.hist A R' n ω) l u σ2 δ a := by - have hpc : pullCount A a (n + 1) ω = pullCount' n (IsAlgEnvSeq.hist A R' n ω) a := + have hp : pullCount A a (n + 1) ω = pullCount' n (IsAlgEnvSeq.hist A R' n ω) a := pullCount_add_one_eq_pullCount' - have hem : empMean A R' a (n + 1) ω = empMean' n (IsAlgEnvSeq.hist A R' n ω) a := + have he : empMean A R' a (n + 1) ω = empMean' n (IsAlgEnvSeq.hist A R' n ω) a := empMean_add_one_eq_empMean' - simp_rw [ucb, ucb', hpc, hem] + rw [ucb, ucb', hp, he] -/-- Helper for `sum_ucbIndex_sub_armMean_le`. -/ +/-- Helper for `sum_ucb_sub_mean_le`. -/ private lemma sum_sqrt_le {ι : Type*} {c : ι → ℝ} (s : Finset ι) (hc : ∀ i, 0 ≤ c i) : ∑ i ∈ s, √(c i) ≤ √(#s * ∑ i ∈ s, c i) := by have h := Real.sum_sqrt_mul_sqrt_le s hc (fun _ => zero_le_one) simp only [Real.sqrt_one, mul_one, sum_const, nsmul_eq_mul] at h rwa [Real.sqrt_mul (by positivity), mul_comm] -/-- Helper for `sum_ucbIndex_sub_armMean_le`. -/ +/-- Helper for `sum_ucb_sub_mean_le`. -/ private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 / √k ≤ 2 * √n - 1 := by induction n with | zero => simp at h @@ -115,32 +114,22 @@ private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 have : √n * √n = n := Real.mul_self_sqrt (by positivity) nlinarith -variable {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} - -/-- Helper for `sum_ucbIndex_sub_armMean_le`. -/ -lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : +/-- Helper for `sum_ucb_sub_mean_le`. -/ +private lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : ∑ s ∈ range n, f (pullCount A (A s ω) s ω) = ∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω), f j := by induction n with | zero => simp | succ n ih => - rw [sum_range_succ, ih] - suffices ∑ a, ∑ j ∈ range (pullCount A a (n + 1) ω), f j = - (∑ a, ∑ j ∈ range (pullCount A a n ω), f j) + - f (pullCount A (A n ω) n ω) by linarith - have h_eq : ∀ a, ∑ j ∈ range (pullCount A a (n + 1) ω), f j = - ∑ j ∈ range (pullCount A a n ω), f j + - if A n ω = a then f (pullCount A a n ω) else 0 := by - intro a - rw [pullCount_add_one] - split_ifs with h - · rw [sum_range_succ] - · simp - simp_rw [h_eq, sum_add_distrib] - congr 1 - simp - -lemma sum_ucbIndex_sub_mean_le {lo hi : ℝ} {μ : Fin K → ℝ} + have hf : f (pullCount A (A n ω) n ω) = + ∑ a, if A n ω = a then f (pullCount A a n ω) else 0 := by simp + simp_rw [sum_range_succ, ih, hf, ← sum_add_distrib, pullCount_add_one] + congr 1 with a + split_ifs + · simp [sum_range_succ] + · simp + +lemma sum_ucb_sub_mean_le {lo hi : ℝ} {μ : Fin K → ℝ} (hm : ∀ a, μ a ∈ Set.Icc lo hi) (hlo : lo ≤ hi) (σ2 δ : ℝ) (n : ℕ) (ω : Ω) (hconc : ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → @@ -254,13 +243,7 @@ lemma sum_ucbIndex_sub_mean_le {lo hi : ℝ} {μ : Fin K → ℝ} end UCB -/-! ### Concentration bounds (algorithm-generic) - -These lemmas take `{alg : Algorithm (Fin K) ℝ}` and -`(h : IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv ν) P')` — -they hold for any algorithm, not just TS. Placed in the `TS` namespace -following the convention of UCB.lean and ETC.lean (cf. `UCB.prob_ucbIndex_le`, -`ETC.probReal_sumRewards_le_sumRewards_le`). -/ +/-! ### Concentration bounds -/ section Concentration @@ -791,8 +774,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp let g : (Iic t → Fin K × ℝ) × Fin K → ℝ := fun p ↦ ucb' t p.1 l u (↑σ2) δ p.2 have hg_eq : ∀ a (ω : Ω), ucb A R' l u (↑σ2) δ a (t + 1) ω = - g (IsAlgEnvSeq.hist A R' t ω, a) := - fun a ω ↦ ucbIndex_succ_eq_ucbIndex'_hist A R' a t ω + g (IsAlgEnvSeq.hist A R' t ω, a) := fun _ _ ↦ ucb_succ_eq_ucb' have hg_meas : Measurable g := measurable_uncurry_ucb' rw [show (fun ω ↦ uc (A (t + 1) ω) (t + 1) ω - uc (bestArm ω) (t + 1) ω) = @@ -844,7 +826,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp ∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω) ≤ (u - l) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by intro ω hω - exact sum_ucbIndex_sub_mean_le (μ := fun a => armMean a ω) + exact sum_ucb_sub_mean_le (μ := fun a => armMean a ω) (hm (E ω)) hlo (↑σ2) δ n ω hω have h_prob : P Eδᶜ ≤ ENNReal.ofReal (2 * K * n * δ) := by have : Eδᶜ = {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ From b702a65c2012edbcbf94436ff598072fceb1d60e Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 19 Mar 2026 10:27:54 +0000 Subject: [PATCH 079/155] Refactor TS.lean (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 211 ++++++------------ .../SequentialLearning/FiniteActions.lean | 13 ++ 2 files changed, 85 insertions(+), 139 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index c331986b..da7df53f 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -114,132 +114,68 @@ private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 have : √n * √n = n := Real.mul_self_sqrt (by positivity) nlinarith -/-- Helper for `sum_ucb_sub_mean_le`. -/ -private lemma sum_comp_pullCount (f : ℕ → ℝ) (n : ℕ) (ω : Ω) : - ∑ s ∈ range n, f (pullCount A (A s ω) s ω) = - ∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω), f j := by - induction n with - | zero => simp - | succ n ih => - have hf : f (pullCount A (A n ω) n ω) = - ∑ a, if A n ω = a then f (pullCount A a n ω) else 0 := by simp - simp_rw [sum_range_succ, ih, hf, ← sum_add_distrib, pullCount_add_one] - congr 1 with a - split_ifs - · simp [sum_range_succ] - · simp - -lemma sum_ucb_sub_mean_le {lo hi : ℝ} {μ : Fin K → ℝ} - (hm : ∀ a, μ a ∈ Set.Icc lo hi) - (hlo : lo ≤ hi) (σ2 δ : ℝ) (n : ℕ) (ω : Ω) - (hconc : ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → - |empMean A R' a s ω - μ a| - < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ))) : - ∑ s ∈ range n, (ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) - ≤ (hi - lo) * ↑K + 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by - -- Split range n into first-pull (pc=0) and non-first-pull (pc≠0) sets - set S0 := (range n).filter (fun s => pullCount A (A s ω) s ω = 0) - set S1 := (range n).filter (fun s => pullCount A (A s ω) s ω ≠ 0) - have hpart : range n = S0 ∪ S1 := (Finset.filter_union_filter_not_eq _ _).symm - have hdisj : Disjoint S0 S1 := Finset.disjoint_filter_filter_not _ _ _ - conv_lhs => rw [hpart] - rw [Finset.sum_union hdisj] - -- We bound ∑_{S0} and ∑_{S1} separately, then combine - suffices h_S0 : ∑ s ∈ S0, (ucb A R' lo hi σ2 δ (A s ω) s ω - - μ (A s ω)) ≤ (hi - lo) * ↑K by - suffices h_S1 : ∑ s ∈ S1, (ucb A R' lo hi σ2 δ (A s ω) s ω - - μ (A s ω)) - ≤ 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) by - have := Finset.sum_union hdisj (f := fun s => - ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) - rw [← hpart] at this; linarith - -- Bound ∑_{S1}: each term ≤ 2√(2σ2c/pc) = 2√(2σ2c/max(1,pc)), so ≤ full sum - calc ∑ s ∈ S1, (ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) - ≤ ∑ s ∈ S1, - 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) := - sum_le_sum fun s hs => by - have hpc : pullCount A (A s ω) s ω ≠ 0 := (Finset.mem_filter.mp hs).2 +lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, μ a ∈ Set.Icc l u) (hi : l ≤ u) + (hc : ∀ s < n, pullCount A (A s ω) s ω ≠ 0 → |empMean A R' (A s ω) s ω - μ (A s ω)| + < √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω))) : + ∑ s ∈ range n, (ucb A R' l u σ2 δ (A s ω) s ω - μ (A s ω)) + ≤ (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by + let S₀ := {s ∈ range n | pullCount A (A s ω) s ω = 0} + let S₁ := {s ∈ range n | pullCount A (A s ω) s ω ≠ 0} + have hu : S₀ ∪ S₁ = range n := filter_union_filter_not_eq _ _ + have hd : Disjoint S₀ S₁ := disjoint_filter_filter_not _ _ _ + rw [← hu, sum_union hd] + gcongr + · calc ∑ s ∈ S₀, (ucb A R' l u σ2 δ (A s ω) s ω - μ (A s ω)) + ≤ ∑ s ∈ S₀, (u - l) := + have (s : ℕ) : ucb A R' l u σ2 δ (A s ω) s ω ∈ Set.Icc l u := ucb_mem_Icc hi + sum_le_sum (by grind) + _ = ∑ s ∈ range n, if pullCount A (A s ω) s ω = 0 then (u - l) else 0 := by + rw [sum_filter] + _ = ∑ a, ∑ j ∈ range (pullCount A a n ω), if j = 0 then (u - l) else 0 := + sum_comp_pullCount (fun j => if j = 0 then (u - l) else 0) n ω + _ ≤ ∑ a, (u - l) := by + gcongr + rw [sum_ite_eq'] + grind + _ = (u - l) * K := by + rw [Fin.sum_const, nsmul_eq_mul, mul_comm] + · calc ∑ s ∈ S₁, (ucb A R' l u σ2 δ (A s ω) s ω - μ (A s ω)) + ≤ ∑ s ∈ S₁, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) := by + gcongr with s hs unfold ucb grind - _ ≤ ∑ s ∈ range n, - 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω : ℝ)) := - Finset.sum_le_sum_of_subset_of_nonneg - (Finset.filter_subset _ _) fun s _ _ => by positivity - _ ≤ 2 * √(8 * σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by - set c := Real.log (1 / δ) - by_cases hc : 0 ≤ 2 * σ2 * c - · open Real in - calc ∑ s ∈ range n, 2 * √(2 * σ2 * c / ↑(pullCount A (A s ω) s ω)) - = ∑ s ∈ range n, √(8 * σ2 * c) * - (1 / √(↑(pullCount A (A s ω) s ω) : ℝ)) := - sum_congr rfl fun s _ => by - rw [show (8 : ℝ) * σ2 * c = (2 : ℝ) ^ 2 * (2 * σ2 * c) from by ring] - rw [sqrt_mul (by positivity : (0:ℝ) ≤ 2 ^ 2), - sqrt_sq (by norm_num : (0:ℝ) ≤ 2)] - rw [sqrt_div (by linarith : 0 ≤ 2 * σ2 * c)]; ring - _ = √(8 * σ2 * c) * ∑ s ∈ range n, - (1 / √(↑(pullCount A (A s ω) s ω) : ℝ)) := by - rw [mul_sum] - _ = √(8 * σ2 * c) * ∑ a : Fin K, ∑ j ∈ range (pullCount A a n ω), - (1 / √(↑j : ℝ)) := by - congr 1; exact sum_comp_pullCount (fun j => 1 / √(↑j : ℝ)) n ω - _ ≤ √(8 * σ2 * c) * ∑ a : Fin K, (2 * √↑(pullCount A a n ω)) := by - gcongr with a - by_cases ha : pullCount A a n ω = 0 - · simp [ha] - · have := sum_inv_sqrt_le (Nat.pos_of_ne_zero ha) - rw [sum_range_succ] at this - linarith [div_nonneg zero_le_one - (Real.sqrt_nonneg (↑(pullCount A a n ω) : ℝ))] - _ = √(8 * σ2 * c) * (2 * ∑ a : Fin K, √↑(pullCount A a n ω)) := by - simp only [mul_sum] - _ ≤ √(8 * σ2 * c) * (2 * √(↑K * ↑n)) := by - gcongr - calc ∑ a : Fin K, √↑(pullCount A a n ω) - ≤ √(↑(Finset.univ.card) * ∑ a, ↑(pullCount A a n ω)) := - sum_sqrt_le Finset.univ fun a => by positivity - _ = √(↑K * ↑n) := by - congr 1; rw [Finset.card_fin]; congr 1 - have h := sum_pullCount (A := A) (t := n) (ω := ω) - exact_mod_cast h - _ = 2 * √(8 * σ2 * c) * √(↑K * ↑n) := by ring - · have h0 : ∀ s ∈ range n, - 2 * √(2 * σ2 * c / ↑(pullCount A (A s ω) s ω)) = 0 := - fun s _ => by - open Real in - have : 2 * σ2 * c / ↑(pullCount A (A s ω) s ω) ≤ 0 := - div_nonpos_of_nonpos_of_nonneg (by linarith) (Nat.cast_nonneg _) - simp [sqrt_eq_zero'.mpr this] - rw [sum_congr rfl h0]; simp only [sum_const_zero]; positivity - -- Bound ∑_{S0}: each term = hi - μ ≤ hi - lo, and #S0 ≤ K - have hterm_S0 : ∀ s ∈ S0, ucb A R' lo hi σ2 δ (A s ω) s ω - - μ (A s ω) ≤ hi - lo := fun s hs => by - have hpc : pullCount A (A s ω) s ω = 0 := (Finset.mem_filter.mp hs).2 - simp only [ucb, hpc, ↓reduceIte] - linarith [(hm (A s ω)).1] - have h_card_S0 : #S0 ≤ K := by - calc #S0 ≤ #(Finset.univ : Finset (Fin K)) := - Finset.card_le_card_of_injOn (fun s => A s ω) - (fun _ _ => Finset.mem_coe.mpr (Finset.mem_univ _)) (by - intro s₁ hs₁ s₂ hs₂ heq - have hpc₁ := (Finset.mem_filter.mp (Finset.mem_coe.mp hs₁)).2 - have hpc₂ := (Finset.mem_filter.mp (Finset.mem_coe.mp hs₂)).2 - by_contra h_ne - rcases lt_or_gt_of_ne h_ne with h_lt | h_lt - · have : s₁ ∈ (range s₂).filter (fun i => A i ω = A s₂ ω) := by - simp [mem_range.mpr h_lt, heq] - exact absurd hpc₂ (show pullCount A (A s₂ ω) s₂ ω ≠ 0 from - Finset.card_ne_zero_of_mem this) - · have : s₂ ∈ (range s₁).filter (fun i => A i ω = A s₁ ω) := by - simp [mem_range.mpr h_lt, ← heq] - exact absurd hpc₁ (show pullCount A (A s₁ ω) s₁ ω ≠ 0 from - Finset.card_ne_zero_of_mem this)) - _ = K := Finset.card_fin K - calc ∑ s ∈ S0, (ucb A R' lo hi σ2 δ (A s ω) s ω - μ (A s ω)) - ≤ ∑ _s ∈ S0, (hi - lo) := sum_le_sum hterm_S0 - _ = #S0 * (hi - lo) := by rw [sum_const, nsmul_eq_mul] - _ ≤ ↑K * (hi - lo) := by gcongr; linarith - _ = (hi - lo) * ↑K := by ring + _ ≤ ∑ s ∈ range n, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) := + sum_le_sum_of_subset_of_nonneg (filter_subset _ _) (fun _ _ _ => by positivity) + _ = 2 * √(2 * σ2 * Real.log (1 / δ)) * ∑ s ∈ range n, (1 / √(pullCount A (A s ω) s ω)) := by + rw [mul_sum] + congr with s + rw [Real.sqrt_div' _ (by positivity)] + ring + _ = 2 * √(2 * σ2 * Real.log (1 / δ)) * + ∑ a, ∑ j ∈ range (pullCount A a n ω), (1 / √j) := by + rw [sum_comp_pullCount (fun j => 1 / √j)] + _ ≤ 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * ∑ a, √(pullCount A a n ω)) := by + rw [mul_sum _ _ 2] + gcongr with a + by_cases ha : pullCount A a n ω = 0 + · simp [ha] + · have hi := sum_inv_sqrt_le (Nat.pos_of_ne_zero ha) + rw [sum_range_succ] at hi + have : 0 ≤ 1 / √(pullCount A a n ω) := by positivity + linarith + _ ≤ 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * √(K * ∑ a, (pullCount A a n ω))) := by + gcongr + have h := sum_sqrt_le Finset.univ (fun a => Nat.cast_nonneg (pullCount A a n ω)) + rw [Finset.card_fin] at h + exact_mod_cast h + _ = 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * √(K * n)) := by + congr + exact sum_pullCount (ω := ω) + _ = 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by + ring_nf + rw [← Real.sqrt_mul' _ (by positivity)] + ring_nf + end UCB @@ -670,7 +606,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : P[IsBayesAlgEnvSeq.regret κ E A n] ≤ (u - l) * ↑K + 2 * (↑K + 1) * (u - l) * n ^ 2 * δ + - 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by + 4 * √(2 * ↑σ2 * Real.log (1 / δ) * ↑K * ↑n) := by have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _) have hlo : l ≤ u := h1.trans h2 let bestArm := IsBayesAlgEnvSeq.bestAction κ E @@ -824,10 +760,10 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp linarith have h_second_Eδ : ∀ ω ∈ Eδ, ∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω) - ≤ (u - l) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) := by + ≤ (u - l) * ↑K + 4 * √(2 * ↑σ2 * Real.log (1 / δ) * ↑K * ↑n) := by intro ω hω - exact sum_ucb_sub_mean_le (μ := fun a => armMean a ω) - (hm (E ω)) hlo (↑σ2) δ n ω hω + exact sum_ucb_sub_mean_le (fun a ↦ armMean a ω) (hm (E ω)) hlo + (fun s hs hpc => hω s hs (A s ω) hpc) have h_prob : P Eδᶜ ≤ ENNReal.ofReal (2 * K * n * δ) := by have : Eδᶜ = {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ @@ -878,7 +814,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (armMean (bestArm ω) ω - uc (bestArm ω) s ω) set f2 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω) - set B := (u - l) * ↑K + 2 * √(8 * ↑σ2 * Real.log (1 / δ)) * √(↑K * ↑n) + set B := (u - l) * ↑K + 4 * √(2 * ↑σ2 * Real.log (1 / δ) * ↑K * ↑n) have h1g : ∫ ω in Fδ, f1 ω ∂P ≤ 0 := setIntegral_nonpos hFδ_meas fun ω hω ↦ h_first_Fδ ω hω have h1b : ∫ ω in Fδᶜ, f1 ω ∂P ≤ ↑n * (u - l) * P.real Fδᶜ := by @@ -954,19 +890,16 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] rw [one_div_one_div, Real.log_pow]; norm_cast calc P[IsBayesAlgEnvSeq.regret κ E A t] ≤ (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * ↑t ^ 2 * (1 / (↑t) ^ 2) - + 2 * √(8 * ↑σ2 * Real.log (1 / (1 / (↑t) ^ 2))) * √(↑K * ↑t) := + + 4 * √(2 * ↑σ2 * Real.log (1 / (1 / (↑t) ^ 2)) * ↑K * ↑t) := bayesRegret_le_of_delta (hK := hK) (E := E) (A := A) (R' := R') (Q := Q) (κ := κ) (P := P) h hσ2 hs hm t (1 / (↑t) ^ 2) hδ hδ1 - _ = (3 * ↑K + 2) * (hi - lo) + 8 * (√(↑σ2 * Real.log ↑t) * √(↑K * ↑t)) := by - rw [h_first, h_log, - show (8 : ℝ) * ↑σ2 * (2 * Real.log ↑t) = 4 ^ 2 * (↑σ2 * Real.log ↑t) by ring, - Real.sqrt_mul (by positivity : (0 : ℝ) ≤ 4 ^ 2), - Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 4)] - ring _ = (3 * ↑K + 2) * (hi - lo) + 8 * √(↑σ2 * ↑K * ↑t * Real.log ↑t) := by - rw [← Real.sqrt_mul (by positivity : - 0 ≤ ↑σ2 * Real.log ↑t)] - congr 1; ring_nf + rw [h_first, h_log]; congr 1 + rw [show (2 : ℝ) * ↑σ2 * (2 * Real.log ↑t) * ↑K * ↑t = + (2 : ℝ) ^ 2 * (↑σ2 * ↑K * ↑t * Real.log ↑t) from by ring, + Real.sqrt_mul (by positivity : (0 : ℝ) ≤ 2 ^ 2), + Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 2)] + ring end TS diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index a2c2cb67..1ef837d2 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -729,6 +729,19 @@ lemma sum_pullCount [Fintype α] {ω : Ω} : ∑ a, pullCount A a t ω = t := by rw [sum_pullCount_mul] simp +lemma sum_comp_pullCount [Fintype α] [AddCommMonoid R] (f : ℕ → R) (t : ℕ) (ω : Ω) : + ∑ s ∈ range t, f (pullCount A (A s ω) s ω) = ∑ a, ∑ j ∈ range (pullCount A a t ω), f j := by + induction t with + | zero => simp + | succ n ih => + have hf : f (pullCount A (A n ω) n ω) = + ∑ a, if A n ω = a then f (pullCount A a n ω) else 0 := by simp + simp_rw [sum_range_succ, ih, hf, ← sum_add_distrib, pullCount_add_one] + congr 1 with a + split_ifs + · simp [sum_range_succ] + · simp + section SumRewards /-- Sum of rewards obtained when pulling action `a` up to time `t` (exclusive). -/ From 93a31d0218abc696c1a4556c972bdba339ad03df Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 20 Mar 2026 14:28:31 +0000 Subject: [PATCH 080/155] Refactoring SumRewards.lean (in progress) --- LeanBandits/Bandit/SumRewards.lean | 198 +++++++++++++++++------- LeanBandits/BanditAlgorithms/TS.lean | 215 +++++++-------------------- 2 files changed, 192 insertions(+), 221 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 1b98cff0..0241476b 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -559,7 +559,7 @@ lemma probReal_sumRewards_le_sumRewards_le [Fintype α] (h : IsAlgEnvSeq A R alg section Subgaussian -/-! ### Sub-Gaussian concentration (δ-parameterized) -/ +/-! ### Sub-Gaussian tail bounds (δ-parameterized) -/ private lemma exp_neg_sq_div_eq_delta {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : @@ -576,37 +576,21 @@ private lemma exp_neg_sq_div_eq_delta {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) rw [Real.exp_log (by positivity : (0 : ℝ) < 1 / δ), one_div, inv_inv] omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -/-- Claude: δ-parameterized one-sided concentration for the stream measure. Setting `δ = 1/(n+1)^c` -recovers `todo` and `todo'` (case-split on `c = 0`) -/ -lemma streamMeasure_concentration_le_delta {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) +lemma streamMeasure_sum_sub_mean_le_le {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (a : α) (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + - √(2 * ↑σ2 * Real.log (1 / δ) / k) ≤ (ν a)[id]} ≤ + streamMeasure ν {ω | (∑ m ∈ range k, ω m a) + + √(2 * k * ↑σ2 * Real.log (1 / δ)) ≤ k * (ν a)[id]} ≤ ENNReal.ofReal δ := by - have hlog : 0 < Real.log (1 / δ) := - Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) calc - streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + - √(2 * ↑σ2 * Real.log (1 / δ) / k) ≤ (ν a)[id]} - _ = streamMeasure ν - {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) / k ≤ - -√(2 * ↑σ2 * Real.log (1 / δ) / k)} := by - congr with ω - field_simp - rw [Finset.sum_sub_distrib] - simp - grind + streamMeasure ν {ω | (∑ m ∈ range k, ω m a) + + √(2 * k * ↑σ2 * Real.log (1 / δ)) ≤ k * (ν a)[id]} _ = streamMeasure ν {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ -√(2 * k * ↑σ2 * Real.log (1 / δ))} := by congr with ω - field_simp - congr! 2 - rw [Real.sqrt_div (by positivity : 0 ≤ 2 * ↑σ2 * Real.log (1 / δ)), - show ↑k * 2 * ↑σ2 * Real.log (1 / δ) = ↑k * (2 * ↑σ2 * Real.log (1 / δ)) from by ring, - Real.sqrt_mul (by positivity : (0 : ℝ) ≤ ↑k), ← mul_div_assoc, - mul_div_right_comm, Real.div_sqrt] + rw [Finset.sum_sub_distrib, Finset.sum_const, Finset.card_range, nsmul_eq_mul] + constructor <;> intro h <;> linarith _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / (2 * k * ↑σ2))) := by rw [← ofReal_measureReal] @@ -620,35 +604,21 @@ lemma streamMeasure_concentration_le_delta {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) _ = ENNReal.ofReal δ := exp_neg_sq_div_eq_delta hσ2 k hk δ hδ hδ1 omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -lemma streamMeasure_concentration_ge_delta {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) +lemma streamMeasure_sum_sub_mean_ge_le {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (a : α) (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - - √(2 * ↑σ2 * Real.log (1 / δ) / k)} ≤ + streamMeasure ν {ω | k * (ν a)[id] + + √(2 * k * ↑σ2 * Real.log (1 / δ)) ≤ (∑ m ∈ range k, ω m a)} ≤ ENNReal.ofReal δ := by - have hlog : 0 < Real.log (1 / δ) := - Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) calc - streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - - √(2 * ↑σ2 * Real.log (1 / δ) / k)} - _ = streamMeasure ν - {ω | √(2 * ↑σ2 * Real.log (1 / δ) / k) ≤ - (∑ s ∈ range k, (ω s a - (ν a)[id])) / k} := by - congr with ω - field_simp - rw [Finset.sum_sub_distrib] - simp - grind + streamMeasure ν {ω | k * (ν a)[id] + + √(2 * k * ↑σ2 * Real.log (1 / δ)) ≤ (∑ m ∈ range k, ω m a)} _ = streamMeasure ν {ω | √(2 * k * ↑σ2 * Real.log (1 / δ)) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} := by congr with ω - field_simp - congr! 1 - rw [Real.sqrt_div (by positivity : 0 ≤ 2 * ↑σ2 * Real.log (1 / δ)), - show 2 * ↑σ2 * Real.log (1 / δ) * ↑k = ↑k * (2 * ↑σ2 * Real.log (1 / δ)) from by ring, - Real.sqrt_mul (by positivity : (0 : ℝ) ≤ ↑k), ← mul_div_assoc, - mul_div_right_comm, Real.div_sqrt] + rw [Finset.sum_sub_distrib, Finset.sum_const, Finset.card_range, nsmul_eq_mul] + constructor <;> intro h <;> linarith _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / (2 * k * ↑σ2))) := by rw [← ofReal_measureReal] @@ -662,29 +632,145 @@ lemma streamMeasure_concentration_ge_delta {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) _ = ENNReal.ofReal δ := exp_neg_sq_div_eq_delta hσ2 k hk δ hδ hδ1 omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -lemma streamMeasure_concentration_bound {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) +lemma streamMeasure_sum_sub_mean_mem_le {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (a : α) {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) (m : ℕ) (hm : m ≠ 0) : streamMeasure ν {ω : ℕ → α → ℝ | ∑ i ∈ range m, ω i a ∈ - {x | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} ∪ - {x | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)}} ≤ + {x | x + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ m * (ν a)[id]} ∪ + {x | m * (ν a)[id] + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ x}} ≤ ENNReal.ofReal (2 * δ) := calc streamMeasure ν {ω : ℕ → α → ℝ | ∑ i ∈ range m, ω i a ∈ - {x | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} ∪ - {x | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)}} - ≤ streamMeasure ν {ω | (∑ i ∈ range m, ω i a) / m + - √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} + - streamMeasure ν {ω | (ν a)[id] ≤ (∑ i ∈ range m, ω i a) / m - - √(2 * ↑σ2 * Real.log (1 / δ) / m)} := by + {x | x + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ m * (ν a)[id]} ∪ + {x | m * (ν a)[id] + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ x}} + ≤ streamMeasure ν {ω | (∑ i ∈ range m, ω i a) + + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ m * (ν a)[id]} + + streamMeasure ν {ω | m * (ν a)[id] + + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ (∑ i ∈ range m, ω i a)} := by apply (measure_mono (fun ω hω ↦ ?_)).trans (measure_union_le _ _) simp only [Set.mem_setOf_eq, Set.mem_union] at hω ⊢; exact hω _ ≤ ENNReal.ofReal δ + ENNReal.ofReal δ := by gcongr - · exact streamMeasure_concentration_le_delta hσ2 hν a m hm δ hδ hδ1 - · exact streamMeasure_concentration_ge_delta hσ2 hν a m hm δ hδ hδ1 + · exact streamMeasure_sum_sub_mean_le_le hσ2 hν a m hm δ hδ hδ1 + · exact streamMeasure_sum_sub_mean_ge_le hσ2 hν a m hm δ hδ hδ1 _ = ENNReal.ofReal (2 * δ) := by rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf +lemma prob_abs_sumRewards_sub_mean_ge_le [Countable α] + {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) : + P (⋃ s ∈ Finset.range n, {ω | pullCount A a s ω ≠ 0 ∧ + √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|}) ≤ + ENNReal.ofReal (2 * n * δ) := by + by_cases hn : n = 0 + · simp [hn] + have hn : 0 < n := Nat.pos_of_ne_zero hn + let B := fun m : ℕ ↦ + {x : ℝ | x + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ m * (ν a)[id]} ∪ + {x : ℝ | m * (ν a)[id] + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ x} + have h_stream_bound : ∀ m : ℕ, m ≠ 0 → + streamMeasure ν {ω : ℕ → α → ℝ | ∑ i ∈ range m, ω i a ∈ B m} ≤ + ENNReal.ofReal (2 * δ) := + fun m hm0 ↦ streamMeasure_sum_sub_mean_mem_le hσ2 hν a hδ hδ1 m hm0 + have hB_meas : ∀ m, MeasurableSet (B m) := fun m ↦ + MeasurableSet.union (measurableSet_le (by fun_prop) (by fun_prop)) + (measurableSet_le (by fun_prop) (by fun_prop)) + let S := Finset.Icc 1 (n - 1) + have hS_card : S.card = n - 1 := by simp only [Nat.card_Icc, S]; omega + have h_decomp : ⋃ s ∈ Finset.range n, {ω | pullCount A a s ω ≠ 0 ∧ + √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} = + ⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ + sumRewards A R a s ω ∈ B m} := by + ext ω + simp only [Set.mem_iUnion, Finset.mem_range, exists_prop, Set.mem_setOf_eq, + Finset.mem_Icc, S] + constructor + · rintro ⟨s, hs, hbad⟩ + let m := pullCount A a s ω + have hm_pos : 0 < m := Nat.pos_of_ne_zero hbad.1 + have hm_le : m ≤ n - 1 := by + have h1 : m ≤ s := pullCount_le (A := A) a s ω + omega + refine ⟨m, ⟨hm_pos, hm_le⟩, s, hs, rfl, ?_⟩ + simp only [Set.mem_union, B, Set.mem_setOf_eq] + rcases le_abs'.mp hbad.2 with h | h <;> [left; right] <;> linarith + · rintro ⟨m, ⟨hm_pos, hm_le⟩, s, hs, hpc, hB⟩ + refine ⟨s, hs, ?_, ?_⟩ + · rw [hpc]; omega + · simp only [hpc, Set.mem_union, B, Set.mem_setOf_eq] at hB ⊢ + rcases hB with h | h + · exact le_abs.mpr (.inr (by linarith)) + · exact le_abs.mpr (.inl (by linarith)) + rw [h_decomp] + calc P (⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ + sumRewards A R a s ω ∈ B m}) + ≤ ∑ m ∈ S, P {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ + sumRewards A R a s ω ∈ B m} := + measure_biUnion_finset_le S _ + _ ≤ ∑ m ∈ S, streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := by + apply Finset.sum_le_sum + intro m hm + calc P {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ + sumRewards A R a s ω ∈ B m} + ≤ P {ω | ∃ s, s ≤ n - 1 ∧ pullCount A a s ω = m ∧ + sumRewards A R a s ω ∈ B m} := by + apply measure_mono + intro ω ⟨s, hs, hpc, hB⟩ + exact ⟨s, Nat.le_sub_one_of_lt hs, hpc, hB⟩ + _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := + prob_exists_pullCount_eq_and_sumRewards_mem_le (n := n - 1) + h (hB_meas m) + _ ≤ ∑ _m ∈ S, ENNReal.ofReal (2 * δ) := + Finset.sum_le_sum fun m hm ↦ + h_stream_bound m (Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1) + _ = (n - 1) • ENNReal.ofReal (2 * δ) := by + simp only [Finset.sum_const, hS_card] + _ ≤ ENNReal.ofReal (2 * n * δ) := by + rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast (n - 1), + ← ENNReal.ofReal_mul (Nat.cast_nonneg (n - 1))] + exact ENNReal.ofReal_le_ofReal (by + nlinarith [(Nat.cast_le (α := ℝ)).mpr (Nat.sub_le n 1), hδ.le]) + +lemma prob_abs_sumRewards_sub_mean_ge_fintype_le [Fintype α] + {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) : + P {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ + √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} ≤ + ENNReal.ofReal (2 * Fintype.card α * n * δ) := by + let badSet := fun (a : α) (s : ℕ) ↦ {ω : Ω | + pullCount A a s ω ≠ 0 ∧ + √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} + have h_set_eq : {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ + √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} = + ⋃ a : α, ⋃ s ∈ Finset.range n, badSet a s := by + ext ω; simp only [Set.mem_setOf_eq, Set.mem_iUnion, Finset.mem_range, badSet, exists_prop] + exact ⟨fun ⟨s, hs, a, ha⟩ ↦ ⟨a, s, hs, ha⟩, fun ⟨a, s, hs, ha⟩ ↦ ⟨s, hs, a, ha⟩⟩ + rw [h_set_eq] + have h_arm_bound : ∀ a : α, + P (⋃ s ∈ Finset.range n, badSet a s) ≤ ENNReal.ofReal (2 * n * δ) := by + intro a + exact prob_abs_sumRewards_sub_mean_ge_le hσ2 hν h hδ hδ1 + calc P (⋃ a : α, ⋃ s ∈ Finset.range n, badSet a s) + ≤ ∑ a : α, P (⋃ s ∈ Finset.range n, badSet a s) := + measure_iUnion_fintype_le _ _ + _ ≤ ∑ _a : α, ENNReal.ofReal (2 * n * δ) := + Finset.sum_le_sum fun a _ ↦ h_arm_bound a + _ = Fintype.card α • ENNReal.ofReal (2 * n * δ) := by + simp [Finset.sum_const] + _ = ENNReal.ofReal (2 * Fintype.card α * n * δ) := by + simp only [nsmul_eq_mul] + rw [← ENNReal.ofReal_natCast (Fintype.card α), + ← ENNReal.ofReal_mul (Nat.cast_nonneg (Fintype.card α))] + congr 1; ring + omit [DecidableEq α] [StandardBorelSpace α] in lemma probReal_sum_le_sum_streamMeasure [Fintype α] {c : ℝ≥0} (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) c (ν a)) (a : α) (m : ℕ) : diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index da7df53f..3223da5c 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -176,172 +176,13 @@ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, rw [← Real.sqrt_mul' _ (by positivity)] ring_nf - end UCB -/-! ### Concentration bounds -/ - -section Concentration - -variable {K : ℕ} [Nonempty (Fin K)] -variable {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] -variable {P' : Measure (ℕ → (Fin K) × ℝ)} [IsProbabilityMeasure P'] -variable {σ2 : ℝ≥0} {alg : Algorithm (Fin K) ℝ} - -/-- Single-arm concentration bound. For any algorithm, the probability that the -empirical mean of arm `a` deviates from the true mean by more than -`√(2σ²·log(1/δ)/pullCount)` at some step before `n` is at most `2nδ`. -/ -lemma concentration_cond_bound - (hσ2 : σ2 ≠ 0) - (hs : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - {n : ℕ} (hn : 0 < n) {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) - (h_isAlgEnvSeq : IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv ν) P') - (a : Fin K) : - P' (⋃ s ∈ Finset.range n, {ω | pullCount IT.action a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (ν a)[id]|}) ≤ - ENNReal.ofReal (2 * n * δ) := by - let B_low := fun m : ℕ ↦ - {x : ℝ | x / m + √(2 * ↑σ2 * Real.log (1 / δ) / m) ≤ (ν a)[id]} - let B_high := fun m : ℕ ↦ - {x : ℝ | (ν a)[id] ≤ x / m - √(2 * ↑σ2 * Real.log (1 / δ) / m)} - have h_stream_bound : ∀ m : ℕ, m ≠ 0 → - streamMeasure ν {ω : ℕ → Fin K → ℝ | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} ≤ - ENNReal.ofReal (2 * δ) := - fun m hm0 ↦ streamMeasure_concentration_bound hσ2 hs a hδ hδ1 m hm0 - have hB_meas : ∀ m, MeasurableSet (B_low m ∪ B_high m) := fun m ↦ - MeasurableSet.union (measurableSet_le (by fun_prop) (by fun_prop)) - (measurableSet_le (by fun_prop) (by fun_prop)) - let badSetIT := fun (s : ℕ) ↦ {ω : ℕ → (Fin K) × ℝ | - pullCount IT.action a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (ν a)[id]|} - let S := Finset.Icc 1 (n - 1) - have hS_card : S.card = n - 1 := by simp only [Nat.card_Icc, S]; omega - have h_decomp : ⋃ s ∈ Finset.range n, badSetIT s = - ⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by - ext ω - simp only [Set.mem_iUnion, Finset.mem_range, exists_prop, badSetIT, Set.mem_setOf_eq, - Finset.mem_Icc, S] - constructor - · rintro ⟨s, hs, hbad⟩ - let m := pullCount IT.action a s ω - have hm_pos : 0 < m := Nat.pos_of_ne_zero hbad.1 - have hm_le : m ≤ n - 1 := by - have h1 : m ≤ s := pullCount_le (A := IT.action) a s ω - omega - refine ⟨m, ⟨hm_pos, hm_le⟩, s, hs, rfl, ?_⟩ - simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] - simp only [empMean] at hbad - rcases le_abs'.mp hbad.2 with h | h <;> [left; right] <;> linarith - · rintro ⟨m, ⟨hm_pos, hm_le⟩, s, hs, hpc, hB⟩ - refine ⟨s, hs, ?_⟩ - simp only [Set.mem_union, B_low, B_high, Set.mem_setOf_eq] at hB - simp only [empMean, hpc] - refine ⟨Nat.one_le_iff_ne_zero.mp hm_pos, ?_⟩ - rcases hB with h | h - · exact le_abs.mpr (.inr (by linarith)) - · exact le_abs.mpr (.inl (by linarith)) - rw [h_decomp] - calc P' (⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m}) - ≤ ∑ m ∈ S, P' {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := - measure_biUnion_finset_le S _ - _ ≤ ∑ m ∈ S, streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := by - apply Finset.sum_le_sum - intro m hm - have hm_pos : m ≠ 0 := Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1 - have hm_le : m ≤ n - 1 := (Finset.mem_Icc.mp hm).2 - have h_contain : {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} ⊆ - {ω | pullCount IT.action a (n - 1) ω = m ∧ - sumRewards IT.action IT.reward a (n - 1) ω ∈ B_low m ∪ B_high m} ∪ - {ω | pullCount IT.action a (n - 1) ω > m ∧ - ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by - intro ω ⟨s, hs, hpc, hB⟩ - simp only [Set.mem_union, Set.mem_setOf_eq] - have hs' : s ≤ n - 1 := Nat.le_sub_one_of_lt hs - have h_pc_mono := pullCount_mono (A := IT.action) a hs' ω - by_cases h_eq : pullCount IT.action a (n - 1) ω = m - · left - refine ⟨h_eq, ?_⟩ - have h_pc_eq : pullCount IT.action a s ω = pullCount IT.action a (n - 1) ω := - hpc.symm ▸ h_eq.symm - rw [← sumRewards_eq_of_pullCount_eq h_pc_eq] - exact hB - · right - exact ⟨by omega, s, hs, hpc, hB⟩ - calc P' {ω | ∃ s, s < n ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} - ≤ P' {ω | ∃ s, s ≤ n - 1 ∧ pullCount IT.action a s ω = m ∧ - sumRewards IT.action IT.reward a s ω ∈ B_low m ∪ B_high m} := by - apply measure_mono - intro ω ⟨s, hs, hpc, hB⟩ - exact ⟨s, Nat.le_sub_one_of_lt hs, hpc, hB⟩ - _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B_low m ∪ B_high m} := - prob_exists_pullCount_eq_and_sumRewards_mem_le (n := n - 1) - h_isAlgEnvSeq (hB_meas m) - _ ≤ ∑ _m ∈ S, ENNReal.ofReal (2 * δ) := - Finset.sum_le_sum fun m hm ↦ - h_stream_bound m (Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1) - _ = (n - 1) • ENNReal.ofReal (2 * δ) := by - simp only [Finset.sum_const, hS_card] - _ ≤ ENNReal.ofReal (2 * n * δ) := by - rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast (n - 1), - ← ENNReal.ofReal_mul (Nat.cast_nonneg (n - 1))] - exact ENNReal.ofReal_le_ofReal (by - nlinarith [(Nat.cast_le (α := ℝ)).mpr (Nat.sub_le n 1), hδ.le]) - -/-- All-arms concentration bound. For any algorithm, the probability that -*some* arm's empirical mean deviates by more than the confidence width at some -step before `n` is at most `2Knδ`. -/ -lemma concentration_fail - (hσ2 : σ2 ≠ 0) - (hs : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - (h : IsAlgEnvSeq IT.action IT.reward alg (stationaryEnv ν) P') - (n : ℕ) {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) : - P' {ω | ∃ s < n, ∃ a, pullCount IT.action a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (ν a)[id]|} ≤ - ENNReal.ofReal (2 * K * n * δ) := by - let badSet := fun (a : Fin K) (s : ℕ) ↦ {ω : ℕ → (Fin K) × ℝ | - pullCount IT.action a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (ν a)[id]|} - have h_set_eq : {ω | ∃ s < n, ∃ a, pullCount IT.action a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (ν a)[id]|} = - ⋃ a : Fin K, ⋃ s ∈ Finset.range n, badSet a s := by - ext ω; simp only [Set.mem_setOf_eq, Set.mem_iUnion, Finset.mem_range, badSet, exists_prop] - exact ⟨fun ⟨s, hs, a, ha⟩ ↦ ⟨a, s, hs, ha⟩, fun ⟨a, s, hs, ha⟩ ↦ ⟨s, hs, a, ha⟩⟩ - rw [h_set_eq] - have h_arm_bound : ∀ a : Fin K, - P' (⋃ s ∈ Finset.range n, badSet a s) ≤ ENNReal.ofReal (2 * n * δ) := by - intro a - by_cases hn : n = 0 - · simp [hn] - exact concentration_cond_bound hσ2 hs (Nat.pos_of_ne_zero hn) hδ hδ1 h a - calc P' (⋃ a : Fin K, ⋃ s ∈ Finset.range n, badSet a s) - ≤ ∑ a : Fin K, P' (⋃ s ∈ Finset.range n, badSet a s) := - measure_iUnion_fintype_le _ _ - _ ≤ ∑ _a : Fin K, ENNReal.ofReal (2 * n * δ) := - Finset.sum_le_sum fun a _ ↦ h_arm_bound a - _ = K • ENNReal.ofReal (2 * n * δ) := by simp [Finset.sum_const] - _ = ENNReal.ofReal (2 * K * n * δ) := by - simp only [nsmul_eq_mul] - rw [← ENNReal.ofReal_natCast K, ← ENNReal.ofReal_mul (Nat.cast_nonneg K)] - congr 1; ring - -end Concentration - end TS end Bandits -open Bandits Bandits.TS +open Bandits /-! ### Algorithm-generic Bayesian lemmas -/ @@ -420,9 +261,31 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] |empMean IT.action IT.reward a s ω - ((κ.sectR e) a)[id]|} := by simp only [badSetIT, Kernel.sectR_apply] rw [this] - exact TS.concentration_fail hσ2 - (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) - h_isAlgEnvSeq n hδ hδ1 + have h_cf := prob_abs_sumRewards_sub_mean_ge_fintype_le (n := n) hσ2 + (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) h_isAlgEnvSeq hδ hδ1 + simp only [Fintype.card_fin] at h_cf + refine le_trans (measure_mono fun ω hω ↦ ?_) h_cf + simp only [Set.mem_setOf_eq, empMean] at hω + obtain ⟨s, hs, a, hpc, hle⟩ := hω + simp only [Set.mem_setOf_eq] + have hk : (0 : ℝ) < pullCount IT.action a s ω := + Nat.cast_pos.mpr (Nat.pos_of_ne_zero hpc) + refine ⟨s, hs, a, hpc, ?_⟩ + rw [show sumRewards IT.action IT.reward a s ω / ↑(pullCount IT.action a s ω) - + ((κ.sectR e) a)[id] = (sumRewards IT.action IT.reward a s ω - + ↑(pullCount IT.action a s ω) * ((κ.sectR e) a)[id]) / + ↑(pullCount IT.action a s ω) from by field_simp, + abs_div, abs_of_pos hk, le_div_iff₀ hk] at hle + have hlog : (0 : ℝ) < Real.log (1 / δ) := Real.log_pos (by rw [lt_div_iff₀ hδ]; linarith) + rwa [show √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount IT.action a s ω)) * + ↑(pullCount IT.action a s ω) = + √(2 * ↑(pullCount IT.action a s ω) * ↑σ2 * Real.log (1 / δ)) from by + rw [show √_ * ↑(pullCount IT.action a s ω) = + √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount IT.action a s ω) * + ↑(pullCount IT.action a s ω) ^ 2) from by + rw [Real.sqrt_mul (by positivity), Real.sqrt_sq hk.le] + ] + congr 1; field_simp] at hle calc P ((fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' {p | p.2 ∈ badSetIT p.1}) = (P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω))) @@ -498,9 +361,31 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] |empMean IT.action IT.reward a s ω - ((κ.sectR e) a)[id]|} := by intro a; simp only [badSetIT, Kernel.sectR_apply] rw [h_eq] - exact TS.concentration_cond_bound hσ2 + set ba := IsBayesAlgEnvSeq.bestAction κ id e + have h_ccb := prob_abs_sumRewards_sub_mean_ge_le (a := ba) (n := n) hσ2 (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) - hn' hδ hδ1 h_isAlgEnvSeq _ + h_isAlgEnvSeq hδ hδ1 + refine le_trans (measure_mono fun ω hω ↦ ?_) h_ccb + simp only [Set.mem_iUnion, Finset.mem_range, Set.mem_setOf_eq] at hω ⊢ + obtain ⟨s, hs, hpc, hle⟩ := hω + simp only [empMean] at hle + have hk : (0 : ℝ) < pullCount IT.action ba s ω := + Nat.cast_pos.mpr (Nat.pos_of_ne_zero hpc) + refine ⟨s, hs, hpc, ?_⟩ + rw [show sumRewards IT.action IT.reward ba s ω / ↑(pullCount IT.action ba s ω) - + ((κ.sectR e) ba)[id] = (sumRewards IT.action IT.reward ba s ω - + ↑(pullCount IT.action ba s ω) * ((κ.sectR e) ba)[id]) / + ↑(pullCount IT.action ba s ω) from by field_simp, + abs_div, abs_of_pos hk, le_div_iff₀ hk] at hle + have hlog : (0 : ℝ) < Real.log (1 / δ) := Real.log_pos (by rw [lt_div_iff₀ hδ]; linarith) + rwa [show √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount IT.action ba s ω)) * + ↑(pullCount IT.action ba s ω) = + √(2 * ↑(pullCount IT.action ba s ω) * ↑σ2 * Real.log (1 / δ)) from by + rw [show √_ * ↑(pullCount IT.action ba s ω) = + √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount IT.action ba s ω) * + ↑(pullCount IT.action ba s ω) ^ 2) from by + rw [Real.sqrt_mul (by positivity), Real.sqrt_sq hk.le]] + congr 1; field_simp] at hle have h_kernel : ∀ a, Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := fun a ↦ stronglyMeasurable_id.integral_kernel.measurable.comp (measurable_fst.prodMk measurable_const) From a2014e162a6934a7929b5d5bda5e6ed0ab59f2a1 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 23 Mar 2026 13:58:08 +0000 Subject: [PATCH 081/155] Refactor SumRewards.lean (in progress) --- LeanBandits/Bandit/Bandit.lean | 10 ++++++++++ LeanBandits/Bandit/SumRewards.lean | 28 +--------------------------- 2 files changed, 11 insertions(+), 27 deletions(-) diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 67b8f986..c7bdb308 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -255,6 +255,16 @@ lemma reward_eq [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : rw [hist_eq] rfl +lemma sumRewards_eq [DecidableEq α] (alg : Algorithm α ℝ) (a : α) (n : ℕ) (ω : probSpace α ℝ) : + sumRewards (action alg) (reward alg) a n ω = + ∑ i ∈ range (pullCount (action alg) a n ω), ω.2 i a := by + induction n with + | zero => simp + | succ n ih => + by_cases ha : action alg n ω = a + · simp [ha, sumRewards_add_one, pullCount_add_one, sum_range_succ, ih, reward_eq] + · simp [ha, sumRewards_add_one, pullCount_eq_pullCount_of_action_ne, ih] + section Measurability lemma measurable_action_add_one' [DecidableEq α] {alg : Algorithm α R} diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 0241476b..0deade00 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -168,32 +168,6 @@ lemma identDistrib_sum_range_snd (a : α) (k : ℕ) : (ν := streamMeasure ν), Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)] rfl -omit [Countable α] in -lemma sumRewards_eq_sum_stream_of_pullCount_eq (a : α) (s m : ℕ) (ω : probSpace α ℝ) - (hpc : pullCount A a s ω = m) : - sumRewards A R a s ω = ∑ i ∈ range m, ω.2 i a := by - let ω' : probSpace α ℝ × (ℕ → α → ℝ) := (ω, ω.2) - have h_sum_rbc : sumRewards A R a s ω = ∑ i ∈ Icc 1 m, rewardByCount A R a i ω' := by - rw [← sum_rewardByCount_eq_sumRewards a s ω', hpc] - rw [h_sum_rbc] - have h_rbc_eq (i : ℕ) (hi : i ∈ Icc 1 m) : rewardByCount A R a i ω' = ω.2 (i - 1) a := by - have hi' := mem_Icc.mp hi - have hi_ne : i ≠ 0 := Nat.one_le_iff_ne_zero.mp hi'.1 - have h_i_le : i ≤ pullCount A a s ω := hpc ▸ hi'.2 - have hs_pos : 0 < s := - Nat.pos_of_ne_zero (by rintro rfl; simp [pullCount] at hpc; omega) - have h_exists : ∃ t, pullCount A a (t + 1) ω = i := - exists_pullCount_eq_of_le (n := s - 1) (Nat.sub_add_cancel hs_pos ▸ h_i_le) hi_ne - rw [rewardByCount_of_stepsUntil_ne_top (stepsUntil_ne_top h_exists)] - simp only [reward_eq] - have h_action : A (stepsUntil A a i ω).toNat ω = a := - action_stepsUntil («A» := A) hi_ne h_exists - congr! - rw [h_action, pullCount_stepsUntil hi_ne h_exists] - calc ∑ i ∈ Icc 1 m, rewardByCount A R a i ω' - _ = ∑ i ∈ Icc 1 m, ω.2 (i - 1) a := Finset.sum_congr rfl h_rbc_eq - _ = ∑ j ∈ range m, ω.2 j a := sum_Icc_one_eq_sum_range (f := fun i => ω.2 i a) - 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} ≤ @@ -255,7 +229,7 @@ lemma prob_exists_pullCount_eq_and_sumRewards_mem_le (a : α) (n m : ℕ) apply measure_mono intro ω ⟨s, _hs, hpc, hB'⟩ -- When pullCount(s, ω) = m, sumRewards(s, ω) = ∑ i < m, ω.2 i a in the ArrayModel. - rw [sumRewards_eq_sum_stream_of_pullCount_eq a s m ω hpc] at hB' + rw [sumRewards_eq alg a s ω, hpc] at hB' exact hB' _ = streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by have := (identDistrib_sum_range_snd (ν := ν) a m).map_eq From 319379e83f69c60978e85a7118c66a35447ed546 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 24 Mar 2026 09:23:27 +0000 Subject: [PATCH 082/155] Refactor SumRewards.lean (in progress) --- LeanBandits/Bandit/SumRewards.lean | 163 +---------------------- blueprint/lean_decls | 2 +- blueprint/src/chapters/concentration.tex | 3 +- 3 files changed, 6 insertions(+), 162 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 0deade00..e0375b9e 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -17,15 +17,6 @@ namespace Bandits namespace ArrayModel -lemma sum_Icc_one_eq_sum_range {m : ℕ} {f : ℕ → ℝ} : - ∑ i ∈ Icc 1 m, f (i - 1) = ∑ i ∈ range m, f i := by - have h : Icc 1 m = (range m).image (· + 1) := by - ext x; simp only [mem_Icc, mem_image, mem_range]; constructor - · intro ⟨h1, h2⟩; exact ⟨x - 1, by omega, by omega⟩ - · rintro ⟨a, ha, rfl⟩; omega - rw [h, Finset.sum_image (fun _ _ _ _ h => by omega)] - simp - variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [Countable α] [StandardBorelSpace α] [Nonempty α] {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] @@ -34,129 +25,6 @@ local notation "A" => action alg local notation "R" => reward alg local notation "𝔓" => arrayMeasure ν -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)) - ((𝔓).prod (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 (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] - 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 (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 a i ω hi - simp_rw [h_sum_eq] - conv_rhs => rw [← Measure.fst_prod (μ := 𝔓) (ν := streamMeasure ν), - Measure.fst] - rw [AEMeasurable.map_map_of_aemeasurable _ (by fun_prop)] - · rfl - simp only [Measure.map_fst_prod, measure_univ, one_smul] - 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_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)) - ((𝔓).prod (streamMeasure ν)) 𝔓 := by - convert identDistrib_pullCount_prod_sum_Icc_rewardByCount' n using 2 with ω - rotate_left - · infer_instance - · infer_instance - ext a : 1 - congr 1 - exact sum_Icc_one_eq_sum_range.symm - -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 (streamMeasure ν)) 𝔓 by - -- 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 (μ := 𝔓) (ν := 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 (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_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 (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 - 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) @@ -173,15 +41,7 @@ lemma prob_pullCount_prod_sumRewards_mem_le (a : α) (n : ℕ) 𝔓 {ω | (pullCount A a n ω, sumRewards A R a n ω) ∈ s} ≤ ∑ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), 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 ω, 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] - swap - · refine Measurable.prod (by fun_prop) ?_ - exact measurable_sum_range_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) + simp_rw [sumRewards_eq] 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ω ↦ ?_ @@ -242,25 +102,10 @@ lemma prob_sumRewards_le_sumRewards_le [Fintype α] (a : α) (n m₁ m₂ : ℕ) sumRewards A R (bestArm ν) n ω ≤ sumRewards A R a n ω} ≤ 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 + simp_rw [sumRewards_eq] 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) + ∑ i ∈ range (pullCount A (bestArm ν) n ω), ω.2 i (bestArm ν) ≤ + ∑ i ∈ range (pullCount A a n ω), ω.2 i a} _ ≤ 𝔓 ((fun ω ↦ (∑ i ∈ range m₁, ω.2 i (bestArm ν), ∑ i ∈ range m₂, ω.2 i a)) ⁻¹' {p | p.1 ≤ p.2}) := by refine measure_mono fun ω hω ↦ ?_ diff --git a/blueprint/lean_decls b/blueprint/lean_decls index 266a4a7d..eafbf8b0 100644 --- a/blueprint/lean_decls +++ b/blueprint/lean_decls @@ -105,7 +105,7 @@ ProbabilityTheory.HasSubgaussianMGF.add_of_indepFun ProbabilityTheory.HasSubgaussianMGF.measure_ge_le ProbabilityTheory.HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun ProbabilityTheory.HasSubgaussianMGF.measure_sum_le_sum_le' -Bandits.ArrayModel.identDistrib_pullCount_prod_sumRewards +# Bandits.ArrayModel.identDistrib_pullCount_prod_sumRewards Bandits.ArrayModel.identDistrib_sum_range_snd Bandits.ArrayModel.prob_pullCount_prod_sumRewards_mem_le Bandits.prob_pullCount_prod_sumRewards_mem_le diff --git a/blueprint/src/chapters/concentration.tex b/blueprint/src/chapters/concentration.tex index 5a888471..940ae194 100644 --- a/blueprint/src/chapters/concentration.tex +++ b/blueprint/src/chapters/concentration.tex @@ -89,10 +89,9 @@ \section{Sub-Gaussian random variables} \section{Concentration of the sums of rewards in bandit models} +%TODO: This lemma was removed. \begin{lemma}\label{lem:AM.identDistrib_pullCount_prod_sumRewards} \uses{def:arrayMeasure,def:AM.history,def:algorithm,def:sumRewards,def:pullCount} - \leanok - \lean{Bandits.ArrayModel.identDistrib_pullCount_prod_sumRewards} In the array model, for $t \in \mathbb{N}$, the random variable $(N_{t,a}, S_{t, a})_{a \in \mathcal{A}}$ has the same distribution as $(N_{t,a}, \sum_{s=0}^{N_{t,a}-1} \omega_{2, s, a})_{a \in \mathcal{A}}$. \end{lemma} From d37c489d0a8dddb52ede4913a0d8c5fe230ed78c Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 24 Mar 2026 10:56:28 +0000 Subject: [PATCH 083/155] Refactor SumRewards.lean (in progress) --- LeanBandits/Bandit/SumRewards.lean | 37 +++++++++++------------------- 1 file changed, 13 insertions(+), 24 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index e0375b9e..75fc25b3 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -79,23 +79,16 @@ lemma prob_pullCount_mem_and_sumRewards_mem_le (a : α) (n : ℕ) exists_eq_right, mem_filter, mem_range] at hk simp [hk.2.1] -lemma prob_exists_pullCount_eq_and_sumRewards_mem_le (a : α) (n m : ℕ) - {B : Set ℝ} (hB : MeasurableSet B) : - 𝔓 {ω | ∃ s, s ≤ n ∧ pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B} ≤ - streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by - calc 𝔓 {ω | ∃ s, s ≤ n ∧ pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B} +lemma prob_exists_pullCount_eq_and_sumRewards_mem_le (a : α) (m : ℕ) {B : Set ℝ} + (hB : MeasurableSet B) : 𝔓 {ω | ∃ n, pullCount A a n ω = m ∧ sumRewards A R a n ω ∈ B} ≤ + streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := + calc _ ≤ 𝔓 {ω | ∑ i ∈ range m, ω.2 i a ∈ B} := by - -- Show the containment: the existential set ⊆ {sum ∈ B} apply measure_mono - intro ω ⟨s, _hs, hpc, hB'⟩ - -- When pullCount(s, ω) = m, sumRewards(s, ω) = ∑ i < m, ω.2 i a in the ArrayModel. - rw [sumRewards_eq alg a s ω, hpc] at hB' - exact hB' - _ = streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by - have := (identDistrib_sum_range_snd (ν := ν) a m).map_eq - rw [Measure.ext_iff] at this - specialize this B hB - rwa [Measure.map_apply (by fun_prop) hB, Measure.map_apply (by fun_prop) hB] at this + intro ω ⟨s, hp, hs⟩ + rwa [sumRewards_eq alg a s ω, hp] at hs + _ = streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := + (identDistrib_sum_range_snd a m).measure_mem_eq hB lemma prob_sumRewards_le_sumRewards_le [Fintype α] (a : α) (n m₁ m₂ : ℕ) : (𝔓) {ω | pullCount A (bestArm ν) n ω = m₁ ∧ pullCount A a n ω = m₂ ∧ @@ -297,7 +290,7 @@ lemma prob_exists_pullCount_eq_and_sumRewards_mem_le [Countable α] constructor <;> rintro ⟨s, hs, rest⟩ <;> exact ⟨s, by omega, rest⟩ rw [h_eq] have h_AM := ArrayModel.prob_exists_pullCount_eq_and_sumRewards_mem_le - (ν := ν) (alg := alg) a n m hB + (ν := ν) (alg := alg) a m hB let pc := fun (p : ℕ → α × ℝ) (s : ℕ) ↦ ∑ i ∈ range s, if (p i).1 = a then 1 else 0 let sr := fun (p : ℕ → α × ℝ) (s : ℕ) ↦ ∑ i ∈ range s, if (p i).1 = a then (p i).2 else 0 let S := ⋃ s ∈ range (n + 1), {p : ℕ → α × ℝ | pc p s = m ∧ sr p s ∈ B} @@ -335,14 +328,10 @@ lemma prob_exists_pullCount_eq_and_sumRewards_mem_le [Countable α] (ArrayModel.action alg t ω, ArrayModel.reward alg t ω)) · rw [measurable_pi_iff]; intro t; exact (hA t).prodMk (hR t) _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by - have h_set_eq : (⋃ s ∈ range (n + 1), {ω | pullCount (ArrayModel.action alg) a s ω = m ∧ - sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a s ω ∈ B}) = - {ω | ∃ s, s ≤ n ∧ pullCount (ArrayModel.action alg) a s ω = m ∧ - sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a s ω ∈ B} := by - ext ω; simp only [Set.mem_iUnion, mem_range, Set.mem_setOf_eq] - constructor <;> rintro ⟨s, hs, rest⟩ <;> exact ⟨s, by omega, rest⟩ - rw [h_set_eq] - exact h_AM + apply le_trans (measure_mono _) h_AM + intro ω hω + simp only [Set.mem_iUnion, mem_range, Set.mem_setOf_eq] at hω + exact hω.elim fun s ⟨_, rest⟩ ↦ ⟨s, rest⟩ lemma probReal_sumRewards_le_sumRewards_le [Fintype α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n m₁ m₂ : ℕ) : From 68ead705494ee74fde357ea7de2cdbb7c473704e Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 24 Mar 2026 12:38:04 +0000 Subject: [PATCH 084/155] Refactor SumRewards.lean (in progress) --- LeanBandits/Bandit/SumRewards.lean | 121 +++++++++++++++-------------- 1 file changed, 61 insertions(+), 60 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 75fc25b3..463134f8 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -85,8 +85,8 @@ lemma prob_exists_pullCount_eq_and_sumRewards_mem_le (a : α) (m : ℕ) {B : Set calc _ ≤ 𝔓 {ω | ∑ i ∈ range m, ω.2 i a ∈ B} := by apply measure_mono - intro ω ⟨s, hp, hs⟩ - rwa [sumRewards_eq alg a s ω, hp] at hs + intro ω ⟨n, hp, hn⟩ + rwa [sumRewards_eq alg a n ω, hp] at hn _ = streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := (identDistrib_sum_range_snd a m).measure_mem_eq hB @@ -219,6 +219,48 @@ lemma _root_.Learning.IsAlgEnvSeq.law_pullCount_sumRewards_unique P'.map (fun ω ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) := ((h1.law_pullCount_sumRewards_unique' h2 (n := n)).comp (u := fun f ↦ f a) (by fun_prop)).map_eq +omit [DecidableEq α] in +lemma _root_.Learning.IsAlgEnvSeq.identDistrib_trajectory + (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + (h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') : + IdentDistrib (fun ω n ↦ (A n ω, R n ω)) (fun ω n ↦ (A₂ n ω, R₂ n ω)) P P' := + ⟨(measurable_pi_iff.mpr fun n ↦ (h1.measurable_A n).prodMk + (h1.measurable_R n)).aemeasurable, + (measurable_pi_iff.mpr fun n ↦ (h2.measurable_A n).prodMk + (h2.measurable_R n)).aemeasurable, + isAlgEnvSeq_unique h1 h2⟩ + +lemma _root_.Learning.IsAlgEnvSeq.identDistrib_pullCount_sumRewards + (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + (h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') : + IdentDistrib (fun ω n a ↦ (pullCount A a n ω, sumRewards A R a n ω)) + (fun ω n a ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) P P' := by + let F : (ℕ → α × ℝ) → ℕ → α → ℕ × ℝ := fun p n 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 n + rw [measurable_pi_iff] + intro a + simp only [F] + apply Measurable.prod + all_goals + dsimp only + refine measurable_sum _ fun i _ ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + have hid_traj := h1.identDistrib_trajectory h2 + have h_eq1 : (fun ω n a ↦ (pullCount A a n ω, sumRewards A R a n ω)) = + F ∘ (fun ω n ↦ (A n ω, R n ω)) := by + ext ω n a : 3 + simp only [Function.comp, F, pullCount, sumRewards, Finset.card_filter] + have h_eq2 : (fun ω n a ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) = + F ∘ (fun ω n ↦ (A₂ n ω, R₂ n ω)) := by + ext ω n a : 3 + simp only [Function.comp, F, pullCount, sumRewards, Finset.card_filter] + rw [h_eq1, h_eq2] + exact hid_traj.comp hF + -- this is what we will use for UCB lemma prob_pullCount_prod_sumRewards_mem_le [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) @@ -278,60 +320,22 @@ lemma prob_pullCount_eq_and_sumRewards_mem_le [Countable α] simpa [hm'] using h_le lemma prob_exists_pullCount_eq_and_sumRewards_mem_le [Countable α] - (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) - {m : ℕ} {B : Set ℝ} (hB : MeasurableSet B) : - P {ω | ∃ s, s ≤ n ∧ pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B} ≤ + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {m : ℕ} {B : Set ℝ} (hB : MeasurableSet B) : + P {ω | ∃ n, pullCount A a n ω = m ∧ sumRewards A R a n ω ∈ B} ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by - have hA := h.measurable_A - have hR := h.measurable_R - have h_eq : {ω | ∃ s, s ≤ n ∧ pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B} = - ⋃ s ∈ range (n + 1), {ω | pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B} := by - ext ω; simp only [Set.mem_setOf_eq, Set.mem_iUnion, mem_range] - constructor <;> rintro ⟨s, hs, rest⟩ <;> exact ⟨s, by omega, rest⟩ - rw [h_eq] - have h_AM := ArrayModel.prob_exists_pullCount_eq_and_sumRewards_mem_le - (ν := ν) (alg := alg) a m hB - let pc := fun (p : ℕ → α × ℝ) (s : ℕ) ↦ ∑ i ∈ range s, if (p i).1 = a then 1 else 0 - let sr := fun (p : ℕ → α × ℝ) (s : ℕ) ↦ ∑ i ∈ range s, if (p i).1 = a then (p i).2 else 0 - let S := ⋃ s ∈ range (n + 1), {p : ℕ → α × ℝ | pc p s = m ∧ sr p s ∈ B} + let S := {f : ℕ → α → ℕ × ℝ | ∃ s, (f s a).1 = m ∧ (f s a).2 ∈ B} have hS : MeasurableSet S := by - simp only [S] - apply MeasurableSet.iUnion - intro s - apply MeasurableSet.iUnion - intro _ - apply MeasurableSet.inter - · exact (measurableSet_singleton _).preimage - (measurable_sum _ fun i _ ↦ Measurable.ite - ((measurableSet_singleton _).preimage (by fun_prop)) (by fun_prop) (by fun_prop)) - · exact hB.preimage - (measurable_sum _ fun i _ ↦ Measurable.ite - ((measurableSet_singleton _).preimage (by fun_prop)) (by fun_prop) (by fun_prop)) - have h_eq1 : (⋃ s ∈ range (n + 1), {ω | pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B}) = - (fun ω t ↦ (A t ω, R t ω)) ⁻¹' S := by - ext ω - simp only [Set.mem_iUnion, mem_range, Set.mem_setOf_eq, Set.mem_preimage, S, pc, sr, - pullCount, sumRewards, Finset.card_filter] - have h_eq2 : (⋃ s ∈ range (n + 1), {ω | pullCount (ArrayModel.action alg) a s ω = m ∧ - sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a s ω ∈ B}) = - (fun ω t ↦ (ArrayModel.action alg t ω, ArrayModel.reward alg t ω)) ⁻¹' S := by - ext ω - simp only [Set.mem_iUnion, mem_range, Set.mem_setOf_eq, Set.mem_preimage, S, pc, sr, - pullCount, sumRewards, Finset.card_filter] - have h_unique := isAlgEnvSeq_unique h (ArrayModel.isAlgEnvSeq_arrayMeasure alg ν) - calc P (⋃ s ∈ range (n + 1), {ω | pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B}) + have : S = ⋃ s, (fun f ↦ f s a) ⁻¹' ({m} ×ˢ B) := by + ext f; simp [S, Set.mem_iUnion, Set.mem_prod] + rw [this] + exact .iUnion fun s ↦ ((measurableSet_singleton m).prod hB).preimage (by fun_prop) + calc _ _ = (ArrayModel.arrayMeasure ν) - (⋃ s ∈ range (n + 1), {ω | pullCount (ArrayModel.action alg) a s ω = m ∧ - sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a s ω ∈ B}) := by - rw [h_eq1, h_eq2, ← Measure.map_apply _ hS, ← Measure.map_apply _ hS, h_unique] - · rw [measurable_pi_iff]; intro t; exact (by fun_prop : Measurable fun ω ↦ - (ArrayModel.action alg t ω, ArrayModel.reward alg t ω)) - · rw [measurable_pi_iff]; intro t; exact (hA t).prodMk (hR t) - _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by - apply le_trans (measure_mono _) h_AM - intro ω hω - simp only [Set.mem_iUnion, mem_range, Set.mem_setOf_eq] at hω - exact hω.elim fun s ⟨_, rest⟩ ↦ ⟨s, rest⟩ + {ω | ∃ n, pullCount (ArrayModel.action alg) a n ω = m ∧ + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a n ω ∈ B} := + (h.identDistrib_pullCount_sumRewards + (ArrayModel.isAlgEnvSeq_arrayMeasure alg ν)).measure_mem_eq hS + _ ≤ _ := ArrayModel.prob_exists_pullCount_eq_and_sumRewards_mem_le a m hB lemma probReal_sumRewards_le_sumRewards_le [Fintype α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n m₁ m₂ : ℕ) : @@ -523,14 +527,11 @@ lemma prob_abs_sumRewards_sub_mean_ge_le [Countable α] intro m hm calc P {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B m} - ≤ P {ω | ∃ s, s ≤ n - 1 ∧ pullCount A a s ω = m ∧ - sumRewards A R a s ω ∈ B m} := by - apply measure_mono - intro ω ⟨s, hs, hpc, hB⟩ - exact ⟨s, Nat.le_sub_one_of_lt hs, hpc, hB⟩ + ≤ P {ω | ∃ s, pullCount A a s ω = m ∧ + sumRewards A R a s ω ∈ B m} := + measure_mono fun ω ⟨s, _, hpc, hB⟩ ↦ ⟨s, hpc, hB⟩ _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := - prob_exists_pullCount_eq_and_sumRewards_mem_le (n := n - 1) - h (hB_meas m) + prob_exists_pullCount_eq_and_sumRewards_mem_le h (hB_meas m) _ ≤ ∑ _m ∈ S, ENNReal.ofReal (2 * δ) := Finset.sum_le_sum fun m hm ↦ h_stream_bound m (Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1) From 9341a1f1a9600e28c134fcb1fb2bf8a310f05d59 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 24 Mar 2026 14:42:31 +0000 Subject: [PATCH 085/155] Refactor SumRewards.lean (in progress) --- LeanBandits/Bandit/SumRewards.lean | 59 ++++++++----------- LeanBandits/SequentialLearning/Algorithm.lean | 9 +++ 2 files changed, 32 insertions(+), 36 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 463134f8..2763b8b5 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -219,47 +219,34 @@ lemma _root_.Learning.IsAlgEnvSeq.law_pullCount_sumRewards_unique P'.map (fun ω ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) := ((h1.law_pullCount_sumRewards_unique' h2 (n := n)).comp (u := fun f ↦ f a) (by fun_prop)).map_eq -omit [DecidableEq α] in -lemma _root_.Learning.IsAlgEnvSeq.identDistrib_trajectory - (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) - (h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') : - IdentDistrib (fun ω n ↦ (A n ω, R n ω)) (fun ω n ↦ (A₂ n ω, R₂ n ω)) P P' := - ⟨(measurable_pi_iff.mpr fun n ↦ (h1.measurable_A n).prodMk - (h1.measurable_R n)).aemeasurable, - (measurable_pi_iff.mpr fun n ↦ (h2.measurable_A n).prodMk - (h2.measurable_R n)).aemeasurable, - isAlgEnvSeq_unique h1 h2⟩ - lemma _root_.Learning.IsAlgEnvSeq.identDistrib_pullCount_sumRewards (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') : IdentDistrib (fun ω n a ↦ (pullCount A a n ω, sumRewards A R a n ω)) - (fun ω n a ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) P P' := by - let F : (ℕ → α × ℝ) → ℕ → α → ℕ × ℝ := fun p n 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 n - rw [measurable_pi_iff] - intro a - simp only [F] - apply Measurable.prod - all_goals - dsimp only - refine measurable_sum _ fun i _ ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact (measurableSet_singleton _).preimage (by fun_prop) - have hid_traj := h1.identDistrib_trajectory h2 - have h_eq1 : (fun ω n a ↦ (pullCount A a n ω, sumRewards A R a n ω)) = - F ∘ (fun ω n ↦ (A n ω, R n ω)) := by - ext ω n a : 3 - simp only [Function.comp, F, pullCount, sumRewards, Finset.card_filter] - have h_eq2 : (fun ω n a ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) = - F ∘ (fun ω n ↦ (A₂ n ω, R₂ n ω)) := by + (fun ω' n a ↦ (pullCount A₂ a n ω', sumRewards A₂ R₂ a n ω')) P P' := by + let f (τ : ℕ → α × ℝ) (n : ℕ) (a : α) : ℕ × ℝ := + (∑ i ∈ range n, if (τ i).1 = a then 1 else 0, + ∑ i ∈ range n, if (τ i).1 = a then (τ i).2 else 0) + have hc1 : (fun ω n a ↦ (pullCount A a n ω, sumRewards A R a n ω)) = + f ∘ (fun ω n ↦ (A n ω, R n ω)) := by ext ω n a : 3 - simp only [Function.comp, F, pullCount, sumRewards, Finset.card_filter] - rw [h_eq1, h_eq2] - exact hid_traj.comp hF + simp_rw [Function.comp, f, pullCount, Finset.card_filter, sumRewards] + have hc2 : (fun ω' n a ↦ (pullCount A₂ a n ω', sumRewards A₂ R₂ a n ω')) = + f ∘ (fun ω' n ↦ (A₂ n ω', R₂ n ω')) := by + ext ω' n a : 3 + simp_rw [Function.comp, f, pullCount, Finset.card_filter, sumRewards] + have hf : Measurable f := by + simp_rw [f, measurable_pi_iff] + intro n a + apply Measurable.prod + · dsimp only + exact measurable_sum _ + (fun _ _ ↦ Measurable.ite (by measurability) (by fun_prop) (by fun_prop)) + · dsimp only + exact measurable_sum _ fun i _ ↦ .ite + ((measurableSet_singleton _).preimage (by fun_prop)) (by fun_prop) (by fun_prop) + rw [hc1, hc2] + exact (h1.identDistrib_trajectory h2).comp hf -- this is what we will use for UCB lemma prob_pullCount_prod_sumRewards_mem_le [Countable α] diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index 926ecb9c..383b2032 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -341,6 +341,15 @@ theorem isAlgEnvSeq_unique (h1 : 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] +lemma IsAlgEnvSeq.identDistrib_trajectory (h1 : IsAlgEnvSeq A₁ R₁ alg env P) + (h2 : IsAlgEnvSeq A₂ R₂ alg env P') : + IdentDistrib (fun ω n ↦ (A₁ n ω, R₁ n ω)) (fun ω' n ↦ (A₂ n ω', R₂ n ω')) P P' where + aemeasurable_fst := + (measurable_pi_iff.2 fun n ↦ (h1.measurable_A n).prodMk (h1.measurable_R n)).aemeasurable + aemeasurable_snd := + (measurable_pi_iff.2 fun n ↦ (h2.measurable_A n).prodMk (h2.measurable_R n)).aemeasurable + map_eq := isAlgEnvSeq_unique h1 h2 + theorem isAlgEnvSeqUntil_unique (h1 : IsAlgEnvSeqUntil A₁ R₁ alg env P N) (h2 : IsAlgEnvSeqUntil A₂ R₂ alg env P' N) : P.map (fun ω (n : Iic N) ↦ (A₁ n ω, R₁ n ω)) = From c6243d50fbc61faf586767f7d4fb2b5996807590 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 25 Mar 2026 10:51:46 +0000 Subject: [PATCH 086/155] Refactor SumRewards.lean (in progress) --- LeanBandits/Bandit/SumRewards.lean | 29 ++++++++++++++--------------- 1 file changed, 14 insertions(+), 15 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 2763b8b5..5ecee569 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -243,8 +243,8 @@ lemma _root_.Learning.IsAlgEnvSeq.identDistrib_pullCount_sumRewards exact measurable_sum _ (fun _ _ ↦ Measurable.ite (by measurability) (by fun_prop) (by fun_prop)) · dsimp only - exact measurable_sum _ fun i _ ↦ .ite - ((measurableSet_singleton _).preimage (by fun_prop)) (by fun_prop) (by fun_prop) + exact measurable_sum _ + (fun _ _ ↦ Measurable.ite (by measurability) (by fun_prop) (by fun_prop)) rw [hc1, hc2] exact (h1.identDistrib_trajectory h2).comp hf @@ -307,21 +307,20 @@ lemma prob_pullCount_eq_and_sumRewards_mem_le [Countable α] simpa [hm'] using h_le lemma prob_exists_pullCount_eq_and_sumRewards_mem_le [Countable α] - (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {m : ℕ} {B : Set ℝ} (hB : MeasurableSet B) : + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m : ℕ) {B : Set ℝ} + (hB : MeasurableSet B) : P {ω | ∃ n, pullCount A a n ω = m ∧ sumRewards A R a n ω ∈ B} ≤ - streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by - let S := {f : ℕ → α → ℕ × ℝ | ∃ s, (f s a).1 = m ∧ (f s a).2 ∈ B} - have hS : MeasurableSet S := by - have : S = ⋃ s, (fun f ↦ f s a) ⁻¹' ({m} ×ˢ B) := by - ext f; simp [S, Set.mem_iUnion, Set.mem_prod] - rw [this] - exact .iUnion fun s ↦ ((measurableSet_singleton m).prod hB).preimage (by fun_prop) - calc _ - _ = (ArrayModel.arrayMeasure ν) - {ω | ∃ n, pullCount (ArrayModel.action alg) a n ω = m ∧ + streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := + let s := {p : ℕ → α → ℕ × ℝ | ∃ n, (p n a).1 = m ∧ (p n a).2 ∈ B} + have : s = ⋃ n, (fun p ↦ p n a) ⁻¹' ({m} ×ˢ B) := by + ext p + simp [s] + have hs : MeasurableSet s := by measurability + calc P {ω | ∃ n, pullCount A a n ω = m ∧ sumRewards A R a n ω ∈ B} + _ = (ArrayModel.arrayMeasure ν) {ω | ∃ n, pullCount (ArrayModel.action alg) a n ω = m ∧ sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a n ω ∈ B} := (h.identDistrib_pullCount_sumRewards - (ArrayModel.isAlgEnvSeq_arrayMeasure alg ν)).measure_mem_eq hS + (ArrayModel.isAlgEnvSeq_arrayMeasure alg ν)).measure_mem_eq hs _ ≤ _ := ArrayModel.prob_exists_pullCount_eq_and_sumRewards_mem_le a m hB lemma probReal_sumRewards_le_sumRewards_le [Fintype α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) @@ -518,7 +517,7 @@ lemma prob_abs_sumRewards_sub_mean_ge_le [Countable α] sumRewards A R a s ω ∈ B m} := measure_mono fun ω ⟨s, _, hpc, hB⟩ ↦ ⟨s, hpc, hB⟩ _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := - prob_exists_pullCount_eq_and_sumRewards_mem_le h (hB_meas m) + prob_exists_pullCount_eq_and_sumRewards_mem_le h a m (hB_meas m) _ ≤ ∑ _m ∈ S, ENNReal.ofReal (2 * δ) := Finset.sum_le_sum fun m hm ↦ h_stream_bound m (Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1) From 96008d13efe17c0b7429d06c1ea46f4496a692f2 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 25 Mar 2026 15:58:01 +0000 Subject: [PATCH 087/155] Refactor SumRewards.lean (in progress) --- LeanBandits/Bandit/SumRewards.lean | 191 ++++++++++++--------------- LeanBandits/BanditAlgorithms/TS.lean | 17 +-- 2 files changed, 91 insertions(+), 117 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 5ecee569..97a1a9c6 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -357,107 +357,92 @@ lemma probReal_sumRewards_le_sumRewards_le [Fintype α] (h : IsAlgEnvSeq A R alg section Subgaussian -/-! ### Sub-Gaussian tail bounds (δ-parameterized) -/ - -private lemma exp_neg_sq_div_eq_delta {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) - (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / - (2 * k * ↑σ2))) = ENNReal.ofReal δ := by - have hk_pos : (0 : ℝ) < k := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hk) - have hσ2_pos : (0 : ℝ) < ↑σ2 := NNReal.coe_pos.mpr (pos_iff_ne_zero.mpr hσ2) - have hlog : 0 < Real.log (1 / δ) := - Real.log_pos (by rw [one_div]; exact one_lt_inv₀ hδ |>.mpr hδ1) - rw [Real.sq_sqrt (by positivity)] - simp only [neg_div, Real.exp_neg] - rw [show 2 * (k : ℝ) * ↑σ2 * Real.log (1 / δ) / (2 * k * ↑σ2) = - Real.log (1 / δ) from by field_simp [ne_of_gt hσ2_pos, ne_of_gt hk_pos]] - rw [Real.exp_log (by positivity : (0 : ℝ) < 1 / δ), one_div, inv_inv] - omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -lemma streamMeasure_sum_sub_mean_le_le {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) +lemma streamMeasure_sum_sub_mean_le_le {σ2 : ℝ≥0} (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - (a : α) (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - streamMeasure ν {ω | (∑ m ∈ range k, ω m a) + - √(2 * k * ↑σ2 * Real.log (1 / δ)) ≤ k * (ν a)[id]} ≤ - ENNReal.ofReal δ := by - calc - streamMeasure ν {ω | (∑ m ∈ range k, ω m a) + - √(2 * k * ↑σ2 * Real.log (1 / δ)) ≤ k * (ν a)[id]} - _ = streamMeasure ν - {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ - -√(2 * k * ↑σ2 * Real.log (1 / δ))} := by - congr with ω - rw [Finset.sum_sub_distrib, Finset.sum_const, Finset.card_range, nsmul_eq_mul] - constructor <;> intro h <;> linarith - _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / - (2 * k * ↑σ2))) := by - rw [← ofReal_measureReal] - gcongr - refine HasSubgaussianMGF.measure_sum_range_le_le_of_iIndepFun (c := σ2) ?_ ?_ - (by positivity) - · exact (iIndepFun_eval_streamMeasure'' ν a).comp - (fun i ω ↦ ω - (ν a)[id]) (fun _ ↦ by fun_prop) - · intro i _; exact (hν a).congr_identDistrib - ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) - _ = ENNReal.ofReal δ := exp_neg_sq_div_eq_delta hσ2 k hk δ hδ hδ1 + (a : α) (k : ℕ) {ε : ℝ} (hε : 0 ≤ ε) : + streamMeasure ν {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ -ε} ≤ + ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * k * σ2))) := by + rw [← ofReal_measureReal] + gcongr + refine HasSubgaussianMGF.measure_sum_range_le_le_of_iIndepFun (c := σ2) ?_ ?_ hε + · exact (iIndepFun_eval_streamMeasure'' ν a).comp + (fun i ω ↦ ω - (ν a)[id]) (fun _ ↦ by fun_prop) + · intro i _; exact (hν a).congr_identDistrib + ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -lemma streamMeasure_sum_sub_mean_ge_le {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) +lemma streamMeasure_sum_sub_mean_ge_le {σ2 : ℝ≥0} (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - (a : α) (k : ℕ) (hk : k ≠ 0) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - streamMeasure ν {ω | k * (ν a)[id] + - √(2 * k * ↑σ2 * Real.log (1 / δ)) ≤ (∑ m ∈ range k, ω m a)} ≤ - ENNReal.ofReal δ := by - calc - streamMeasure ν {ω | k * (ν a)[id] + - √(2 * k * ↑σ2 * Real.log (1 / δ)) ≤ (∑ m ∈ range k, ω m a)} - _ = streamMeasure ν - {ω | √(2 * k * ↑σ2 * Real.log (1 / δ)) ≤ - (∑ s ∈ range k, (ω s a - (ν a)[id]))} := by - congr with ω - rw [Finset.sum_sub_distrib, Finset.sum_const, Finset.card_range, nsmul_eq_mul] - constructor <;> intro h <;> linarith - _ ≤ ENNReal.ofReal (Real.exp (-(√(2 * k * ↑σ2 * Real.log (1 / δ)))^2 / - (2 * k * ↑σ2))) := by - rw [← ofReal_measureReal] - gcongr - refine HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := σ2) ?_ ?_ - (by positivity) - · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) - (fun _ ↦ by fun_prop) - · intro i _; exact (hν a).congr_identDistrib - ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) - _ = ENNReal.ofReal δ := exp_neg_sq_div_eq_delta hσ2 k hk δ hδ hδ1 + (a : α) (k : ℕ) {ε : ℝ} (hε : 0 ≤ ε) : + streamMeasure ν {ω | ε ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} ≤ + ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * k * σ2))) := by + rw [← ofReal_measureReal] + gcongr + refine HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := σ2) ?_ ?_ hε + · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) + (fun _ ↦ by fun_prop) + · intro i _; exact (hν a).congr_identDistrib + ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -lemma streamMeasure_sum_sub_mean_mem_le {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) +lemma streamMeasure_sum_sub_mean_mem_le {σ2 : ℝ≥0} (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - (a : α) {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) (m : ℕ) (hm : m ≠ 0) : - streamMeasure ν {ω : ℕ → α → ℝ | ∑ i ∈ range m, ω i a ∈ - {x | x + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ m * (ν a)[id]} ∪ - {x | m * (ν a)[id] + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ x}} ≤ - ENNReal.ofReal (2 * δ) := - calc streamMeasure ν {ω : ℕ → α → ℝ | ∑ i ∈ range m, ω i a ∈ - {x | x + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ m * (ν a)[id]} ∪ - {x | m * (ν a)[id] + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ x}} - ≤ streamMeasure ν {ω | (∑ i ∈ range m, ω i a) + - √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ m * (ν a)[id]} + - streamMeasure ν {ω | m * (ν a)[id] + - √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ (∑ i ∈ range m, ω i a)} := by + (a : α) (m : ℕ) {ε : ℝ} (hε : 0 ≤ ε) : + streamMeasure ν {ω | ε ≤ |(∑ i ∈ range m, ω i a) - m * (ν a)[id]|} ≤ + ENNReal.ofReal (2 * Real.exp (-ε ^ 2 / (2 * m * σ2))) := by + have h_eq : {ω : ℕ → α → ℝ | ε ≤ |(∑ i ∈ range m, ω i a) - ↑m * (ν a)[id]|} = + {ω | ε ≤ |∑ s ∈ range m, (ω s a - (ν a)[id])|} := by + congr with ω + rw [Finset.sum_sub_distrib, Finset.sum_const, Finset.card_range, nsmul_eq_mul] + rw [h_eq] + calc streamMeasure ν {ω | ε ≤ |∑ s ∈ range m, (ω s a - (ν a)[id])|} + ≤ streamMeasure ν {ω | (∑ s ∈ range m, (ω s a - (ν a)[id])) ≤ -ε} + + streamMeasure ν {ω | ε ≤ (∑ s ∈ range m, (ω s a - (ν a)[id]))} := by apply (measure_mono (fun ω hω ↦ ?_)).trans (measure_union_le _ _) - simp only [Set.mem_setOf_eq, Set.mem_union] at hω ⊢; exact hω - _ ≤ ENNReal.ofReal δ + ENNReal.ofReal δ := by + simp only [Set.mem_setOf_eq, Set.mem_union] at hω ⊢ + exact (le_abs.mp hω).symm.imp le_neg.mp id + _ ≤ ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * m * σ2))) + + ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * m * σ2))) := by gcongr - · exact streamMeasure_sum_sub_mean_le_le hσ2 hν a m hm δ hδ hδ1 - · exact streamMeasure_sum_sub_mean_ge_le hσ2 hν a m hm δ hδ hδ1 - _ = ENNReal.ofReal (2 * δ) := by + · exact streamMeasure_sum_sub_mean_le_le hν a m hε + · exact streamMeasure_sum_sub_mean_ge_le hν a m hε + _ = ENNReal.ofReal (2 * Real.exp (-ε ^ 2 / (2 * m * σ2))) := by rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf +private lemma exp_neg_sqrt_sq_div_le {σ2 : ℝ≥0} (hσ2 : 0 < σ2) {m : ℕ} (hm : 0 < m) + {δ : ℝ} (hδ : 0 < δ) : + Real.exp (-√(2 * ↑m * ↑σ2 * Real.log (1 / δ)) ^ 2 / (2 * ↑m * ↑σ2)) ≤ δ := by + by_cases hδ1 : δ < 1 + · have : 0 < Real.log (1 / δ) := Real.log_pos ((one_lt_div hδ).2 hδ1) + have : Real.exp (-√(2 * ↑m * ↑σ2 * Real.log (1 / δ)) ^ 2 / (2 * ↑m * ↑σ2)) = δ := by + rw [Real.sq_sqrt (by positivity), neg_div] + field_simp + simp [Real.exp_log (by positivity)] + linarith + · calc Real.exp _ ≤ Real.exp 0 := by + gcongr + simp only [neg_div, neg_nonpos] + positivity + _ ≤ δ := by + simp [Real.exp_zero] + linarith + +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in +lemma streamMeasure_sum_sub_mean_mem_le' {σ2 : ℝ≥0} (hσ2 : 0 < σ2) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) + (a : α) (m : ℕ) (hm : 0 < m) {δ : ℝ} (hδ : 0 < δ) : + streamMeasure ν {ω | √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ + |(∑ i ∈ range m, ω i a) - m * (ν a)[id]|} ≤ + ENNReal.ofReal (2 * δ) := + (streamMeasure_sum_sub_mean_mem_le hν a m (by positivity)).trans + (by gcongr; exact exp_neg_sqrt_sq_div_le hσ2 hm hδ) + lemma prob_abs_sumRewards_sub_mean_ge_le [Countable α] - {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) - {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) : + {δ : ℝ} (hδ : 0 < δ) : P (⋃ s ∈ Finset.range n, {ω | pullCount A a s ω ≠ 0 ∧ √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|}) ≤ @@ -465,16 +450,10 @@ lemma prob_abs_sumRewards_sub_mean_ge_le [Countable α] by_cases hn : n = 0 · simp [hn] have hn : 0 < n := Nat.pos_of_ne_zero hn - let B := fun m : ℕ ↦ - {x : ℝ | x + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ m * (ν a)[id]} ∪ - {x : ℝ | m * (ν a)[id] + √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ x} - have h_stream_bound : ∀ m : ℕ, m ≠ 0 → - streamMeasure ν {ω : ℕ → α → ℝ | ∑ i ∈ range m, ω i a ∈ B m} ≤ - ENNReal.ofReal (2 * δ) := - fun m hm0 ↦ streamMeasure_sum_sub_mean_mem_le hσ2 hν a hδ hδ1 m hm0 - have hB_meas : ∀ m, MeasurableSet (B m) := fun m ↦ - MeasurableSet.union (measurableSet_le (by fun_prop) (by fun_prop)) - (measurableSet_le (by fun_prop) (by fun_prop)) + let B := fun m : ℕ ↦ {x : ℝ | √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ |x - m * (ν a)[id]|} + have hB_meas : ∀ m, MeasurableSet (B m) := fun m ↦ by + simp only [B] + measurability let S := Finset.Icc 1 (n - 1) have hS_card : S.card = n - 1 := by simp only [Nat.card_Icc, S]; omega have h_decomp : ⋃ s ∈ Finset.range n, {ω | pullCount A a s ω ≠ 0 ∧ @@ -492,16 +471,10 @@ lemma prob_abs_sumRewards_sub_mean_ge_le [Countable α] have hm_le : m ≤ n - 1 := by have h1 : m ≤ s := pullCount_le (A := A) a s ω omega - refine ⟨m, ⟨hm_pos, hm_le⟩, s, hs, rfl, ?_⟩ - simp only [Set.mem_union, B, Set.mem_setOf_eq] - rcases le_abs'.mp hbad.2 with h | h <;> [left; right] <;> linarith + exact ⟨m, ⟨hm_pos, hm_le⟩, s, hs, rfl, hbad.2⟩ · rintro ⟨m, ⟨hm_pos, hm_le⟩, s, hs, hpc, hB⟩ - refine ⟨s, hs, ?_, ?_⟩ - · rw [hpc]; omega - · simp only [hpc, Set.mem_union, B, Set.mem_setOf_eq] at hB ⊢ - rcases hB with h | h - · exact le_abs.mpr (.inr (by linarith)) - · exact le_abs.mpr (.inl (by linarith)) + subst hpc + exact ⟨s, hs, by omega, hB⟩ rw [h_decomp] calc P (⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B m}) @@ -520,7 +493,7 @@ lemma prob_abs_sumRewards_sub_mean_ge_le [Countable α] prob_exists_pullCount_eq_and_sumRewards_mem_le h a m (hB_meas m) _ ≤ ∑ _m ∈ S, ENNReal.ofReal (2 * δ) := Finset.sum_le_sum fun m hm ↦ - h_stream_bound m (Nat.one_le_iff_ne_zero.mp (Finset.mem_Icc.mp hm).1) + streamMeasure_sum_sub_mean_mem_le' hσ2 hν a m (Finset.mem_Icc.mp hm).1 hδ _ = (n - 1) • ENNReal.ofReal (2 * δ) := by simp only [Finset.sum_const, hS_card] _ ≤ ENNReal.ofReal (2 * n * δ) := by @@ -530,10 +503,10 @@ lemma prob_abs_sumRewards_sub_mean_ge_le [Countable α] nlinarith [(Nat.cast_le (α := ℝ)).mpr (Nat.sub_le n 1), hδ.le]) lemma prob_abs_sumRewards_sub_mean_ge_fintype_le [Fintype α] - {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) - {δ : ℝ} (hδ : 0 < δ) (hδ1 : δ < 1) : + {δ : ℝ} (hδ : 0 < δ) : P {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} ≤ @@ -552,7 +525,7 @@ lemma prob_abs_sumRewards_sub_mean_ge_fintype_le [Fintype α] have h_arm_bound : ∀ a : α, P (⋃ s ∈ Finset.range n, badSet a s) ≤ ENNReal.ofReal (2 * n * δ) := by intro a - exact prob_abs_sumRewards_sub_mean_ge_le hσ2 hν h hδ hδ1 + exact prob_abs_sumRewards_sub_mean_ge_le hσ2 hν h hδ calc P (⋃ a : α, ⋃ s ∈ Finset.range n, badSet a s) ≤ ∑ a : α, P (⋃ s ∈ Finset.range n, badSet a s) := measure_iUnion_fintype_le _ _ diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 3223da5c..f611b6b4 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -200,7 +200,7 @@ variable [IsMarkovKernel κ] lemma prob_concentration_fail_delta [Nonempty (Fin K)] {alg : Algorithm (Fin K) ℝ} (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : P {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ @@ -261,8 +261,8 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] |empMean IT.action IT.reward a s ω - ((κ.sectR e) a)[id]|} := by simp only [badSetIT, Kernel.sectR_apply] rw [this] - have h_cf := prob_abs_sumRewards_sub_mean_ge_fintype_le (n := n) hσ2 - (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) h_isAlgEnvSeq hδ hδ1 + have h_cf := prob_abs_sumRewards_sub_mean_ge_fintype_le (n := n) (hσ2) + (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) h_isAlgEnvSeq hδ simp only [Fintype.card_fin] at h_cf refine le_trans (measure_mono fun ω hω ↦ ?_) h_cf simp only [Set.mem_setOf_eq, empMean] at hω @@ -307,7 +307,7 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] {alg : Algorithm (Fin K) ℝ} (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : P {ω | ∃ s < n, pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) s ω ≠ 0 ∧ @@ -362,9 +362,10 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] intro a; simp only [badSetIT, Kernel.sectR_apply] rw [h_eq] set ba := IsBayesAlgEnvSeq.bestAction κ id e - have h_ccb := prob_abs_sumRewards_sub_mean_ge_le (a := ba) (n := n) hσ2 + have h_ccb := prob_abs_sumRewards_sub_mean_ge_le (a := ba) (n := n) + (hσ2) (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) - h_isAlgEnvSeq hδ hδ1 + h_isAlgEnvSeq hδ refine le_trans (measure_mono fun ω hω ↦ ?_) h_ccb simp only [Set.mem_iUnion, Finset.mem_range, Set.mem_setOf_eq] at hω ⊢ obtain ⟨s, hs, hpc, hle⟩ := hω @@ -485,7 +486,7 @@ lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) - {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {l u : ℝ} (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : @@ -731,7 +732,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) - {σ2 : ℝ≥0} (hσ2 : σ2 ≠ 0) + {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {lo hi : ℝ} (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (t : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A t] From 23b3ccfcde38769703d8e8c1fde51e89188d0b48 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 27 Mar 2026 14:43:19 +0000 Subject: [PATCH 088/155] Refactor SumRewards.lean (in progress) --- LeanBandits/Bandit/SumRewards.lean | 133 +++++++++++++-------------- LeanBandits/BanditAlgorithms/TS.lean | 4 +- 2 files changed, 66 insertions(+), 71 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 97a1a9c6..0f338af9 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -357,65 +357,56 @@ lemma probReal_sumRewards_le_sumRewards_le [Fintype α] (h : IsAlgEnvSeq A R alg section Subgaussian -omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -lemma streamMeasure_sum_sub_mean_le_le {σ2 : ℝ≥0} - (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - (a : α) (k : ℕ) {ε : ℝ} (hε : 0 ≤ ε) : - streamMeasure ν {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ -ε} ≤ - ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * k * σ2))) := by +namespace StreamMeasure + +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] + +lemma prob_sum_range_sub_ge_le_of_HasSubgaussianMGF {σ2 : ℝ≥0} + (h : HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) {ε : ℝ} (hε : 0 ≤ ε) (n : ℕ) : + streamMeasure ν {ω | ε ≤ ∑ k ∈ range n, (ω k a - (ν a)[id])} ≤ + ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * n * σ2))) := by rw [← ofReal_measureReal] gcongr - refine HasSubgaussianMGF.measure_sum_range_le_le_of_iIndepFun (c := σ2) ?_ ?_ hε - · exact (iIndepFun_eval_streamMeasure'' ν a).comp - (fun i ω ↦ ω - (ν a)[id]) (fun _ ↦ by fun_prop) - · intro i _; exact (hν a).congr_identDistrib - ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) - -omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -lemma streamMeasure_sum_sub_mean_ge_le {σ2 : ℝ≥0} - (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - (a : α) (k : ℕ) {ε : ℝ} (hε : 0 ≤ ε) : - streamMeasure ν {ω | ε ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} ≤ - ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * k * σ2))) := by + apply HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun _ _ hε + · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun _ x ↦ x - (ν a)[id]) (by fun_prop) + · intro _ _ + exact h.congr_identDistrib ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) + +lemma prob_sum_range_sub_le_le_of_HasSubgaussianMGF {σ2 : ℝ≥0} + (h : HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) {ε : ℝ} (hε : 0 ≤ ε) (n : ℕ) : + streamMeasure ν {ω | ∑ k ∈ range n, (ω k a - (ν a)[id]) ≤ -ε} ≤ + ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * n * σ2))) := by rw [← ofReal_measureReal] gcongr - refine HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := σ2) ?_ ?_ hε - · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) - (fun _ ↦ by fun_prop) - · intro i _; exact (hν a).congr_identDistrib - ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) - -omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -lemma streamMeasure_sum_sub_mean_mem_le {σ2 : ℝ≥0} - (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - (a : α) (m : ℕ) {ε : ℝ} (hε : 0 ≤ ε) : - streamMeasure ν {ω | ε ≤ |(∑ i ∈ range m, ω i a) - m * (ν a)[id]|} ≤ - ENNReal.ofReal (2 * Real.exp (-ε ^ 2 / (2 * m * σ2))) := by - have h_eq : {ω : ℕ → α → ℝ | ε ≤ |(∑ i ∈ range m, ω i a) - ↑m * (ν a)[id]|} = - {ω | ε ≤ |∑ s ∈ range m, (ω s a - (ν a)[id])|} := by - congr with ω - rw [Finset.sum_sub_distrib, Finset.sum_const, Finset.card_range, nsmul_eq_mul] - rw [h_eq] - calc streamMeasure ν {ω | ε ≤ |∑ s ∈ range m, (ω s a - (ν a)[id])|} - ≤ streamMeasure ν {ω | (∑ s ∈ range m, (ω s a - (ν a)[id])) ≤ -ε} + - streamMeasure ν {ω | ε ≤ (∑ s ∈ range m, (ω s a - (ν a)[id]))} := by - apply (measure_mono (fun ω hω ↦ ?_)).trans (measure_union_le _ _) - simp only [Set.mem_setOf_eq, Set.mem_union] at hω ⊢ - exact (le_abs.mp hω).symm.imp le_neg.mp id - _ ≤ ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * m * σ2))) + - ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * m * σ2))) := by - gcongr - · exact streamMeasure_sum_sub_mean_le_le hν a m hε - · exact streamMeasure_sum_sub_mean_ge_le hν a m hε - _ = ENNReal.ofReal (2 * Real.exp (-ε ^ 2 / (2 * m * σ2))) := by - rw [← ENNReal.ofReal_add (by positivity) (by positivity)]; ring_nf - -private lemma exp_neg_sqrt_sq_div_le {σ2 : ℝ≥0} (hσ2 : 0 < σ2) {m : ℕ} (hm : 0 < m) - {δ : ℝ} (hδ : 0 < δ) : - Real.exp (-√(2 * ↑m * ↑σ2 * Real.log (1 / δ)) ^ 2 / (2 * ↑m * ↑σ2)) ≤ δ := by + apply HasSubgaussianMGF.measure_sum_range_le_le_of_iIndepFun _ _ hε + · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun _ x ↦ x - (ν a)[id]) (by fun_prop) + · intro _ _ + exact h.congr_identDistrib ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) + +lemma prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF {σ2 : ℝ≥0} + (h : HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) {ε : ℝ} (hε : 0 ≤ ε) (n : ℕ) : + streamMeasure ν {ω | ε ≤ |∑ k ∈ range n, (ω k a - (ν a)[id])|} ≤ + ENNReal.ofReal (2 * Real.exp (-ε ^ 2 / (2 * n * σ2))) := by + calc streamMeasure ν {ω | ε ≤ |∑ k ∈ range n, (ω k a - (ν a)[id])|} + _ = streamMeasure ν ({ω | ε ≤ ∑ k ∈ range n, (ω k a - (ν a)[id])} ∪ + {ω | ∑ k ∈ range n, (ω k a - (ν a)[id]) ≤ -ε}) := by + simp_rw [le_abs, le_neg] + rfl + _ ≤ streamMeasure ν {ω | ε ≤ ∑ k ∈ range n, (ω k a - (ν a)[id])} + + streamMeasure ν {ω | ∑ k ∈ range n, (ω k a - (ν a)[id]) ≤ -ε} := + measure_union_le _ _ + _ ≤ ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * n * σ2))) + + ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * n * σ2))) := + add_le_add (prob_sum_range_sub_ge_le_of_HasSubgaussianMGF h hε n) + (prob_sum_range_sub_le_le_of_HasSubgaussianMGF h hε n) + _ = ENNReal.ofReal (2 * Real.exp (-ε ^ 2 / (2 * n * σ2))) := by + rw [← ENNReal.ofReal_add (by positivity) (by positivity), ← two_mul] + +private lemma exp_neg_sqrt_sq_div_le {σ2 : ℝ≥0} (hσ2 : 0 < σ2) {δ : ℝ} (hδ : 0 < δ) (hn : 0 < n) : + Real.exp (-√(2 * n * σ2 * Real.log (1 / δ)) ^ 2 / (2 * n * σ2)) ≤ δ := by by_cases hδ1 : δ < 1 · have : 0 < Real.log (1 / δ) := Real.log_pos ((one_lt_div hδ).2 hδ1) - have : Real.exp (-√(2 * ↑m * ↑σ2 * Real.log (1 / δ)) ^ 2 / (2 * ↑m * ↑σ2)) = δ := by + have : Real.exp (-√(2 * n * σ2 * Real.log (1 / δ)) ^ 2 / (2 * n * σ2)) = δ := by rw [Real.sq_sqrt (by positivity), neg_div] field_simp simp [Real.exp_log (by positivity)] @@ -428,17 +419,17 @@ private lemma exp_neg_sqrt_sq_div_le {σ2 : ℝ≥0} (hσ2 : 0 < σ2) {m : ℕ} simp [Real.exp_zero] linarith -omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in -lemma streamMeasure_sum_sub_mean_mem_le' {σ2 : ℝ≥0} (hσ2 : 0 < σ2) - (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - (a : α) (m : ℕ) (hm : 0 < m) {δ : ℝ} (hδ : 0 < δ) : - streamMeasure ν {ω | √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ - |(∑ i ∈ range m, ω i a) - m * (ν a)[id]|} ≤ - ENNReal.ofReal (2 * δ) := - (streamMeasure_sum_sub_mean_mem_le hν a m (by positivity)).trans - (by gcongr; exact exp_neg_sqrt_sq_div_le hσ2 hm hδ) - -lemma prob_abs_sumRewards_sub_mean_ge_le [Countable α] +lemma prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF' {σ2 : ℝ≥0} (hσ2 : 0 < σ2) + (h : HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) {δ : ℝ} (hδ : 0 < δ) (hn : 0 < n) : + streamMeasure ν {ω | √(2 * n * σ2 * Real.log (1 / δ)) ≤ + |∑ k ∈ range n, (ω k a - (ν a)[id])|} ≤ ENNReal.ofReal (2 * δ) := by + apply (prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF h (by positivity) n).trans + gcongr + exact exp_neg_sqrt_sq_div_le hσ2 hδ hn + +end StreamMeasure + +lemma prob_abs_sumRewards_sub_ge_le [Countable α] {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) @@ -491,9 +482,13 @@ lemma prob_abs_sumRewards_sub_mean_ge_le [Countable α] measure_mono fun ω ⟨s, _, hpc, hB⟩ ↦ ⟨s, hpc, hB⟩ _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := prob_exists_pullCount_eq_and_sumRewards_mem_le h a m (hB_meas m) - _ ≤ ∑ _m ∈ S, ENNReal.ofReal (2 * δ) := - Finset.sum_le_sum fun m hm ↦ - streamMeasure_sum_sub_mean_mem_le' hσ2 hν a m (Finset.mem_Icc.mp hm).1 hδ + _ ≤ ∑ _m ∈ S, ENNReal.ofReal (2 * δ) := by + apply sum_le_sum + intro m hm + convert StreamMeasure.prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF' + hσ2 (hν a) hδ (Finset.mem_Icc.mp hm).1 using 2 + simp_rw [B, Set.mem_setOf_eq, Finset.sum_sub_distrib, Finset.sum_const, + Finset.card_range, nsmul_eq_mul] _ = (n - 1) • ENNReal.ofReal (2 * δ) := by simp only [Finset.sum_const, hS_card] _ ≤ ENNReal.ofReal (2 * n * δ) := by @@ -502,7 +497,7 @@ lemma prob_abs_sumRewards_sub_mean_ge_le [Countable α] exact ENNReal.ofReal_le_ofReal (by nlinarith [(Nat.cast_le (α := ℝ)).mpr (Nat.sub_le n 1), hδ.le]) -lemma prob_abs_sumRewards_sub_mean_ge_fintype_le [Fintype α] +lemma prob_abs_sumRewards_sub_ge_fintype_le [Fintype α] {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) @@ -525,7 +520,7 @@ lemma prob_abs_sumRewards_sub_mean_ge_fintype_le [Fintype α] have h_arm_bound : ∀ a : α, P (⋃ s ∈ Finset.range n, badSet a s) ≤ ENNReal.ofReal (2 * n * δ) := by intro a - exact prob_abs_sumRewards_sub_mean_ge_le hσ2 hν h hδ + exact prob_abs_sumRewards_sub_ge_le hσ2 hν h hδ calc P (⋃ a : α, ⋃ s ∈ Finset.range n, badSet a s) ≤ ∑ a : α, P (⋃ s ∈ Finset.range n, badSet a s) := measure_iUnion_fintype_le _ _ diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index f611b6b4..48594c46 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -261,7 +261,7 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] |empMean IT.action IT.reward a s ω - ((κ.sectR e) a)[id]|} := by simp only [badSetIT, Kernel.sectR_apply] rw [this] - have h_cf := prob_abs_sumRewards_sub_mean_ge_fintype_le (n := n) (hσ2) + have h_cf := prob_abs_sumRewards_sub_ge_fintype_le (n := n) (hσ2) (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) h_isAlgEnvSeq hδ simp only [Fintype.card_fin] at h_cf refine le_trans (measure_mono fun ω hω ↦ ?_) h_cf @@ -362,7 +362,7 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] intro a; simp only [badSetIT, Kernel.sectR_apply] rw [h_eq] set ba := IsBayesAlgEnvSeq.bestAction κ id e - have h_ccb := prob_abs_sumRewards_sub_mean_ge_le (a := ba) (n := n) + have h_ccb := prob_abs_sumRewards_sub_ge_le (a := ba) (n := n) (hσ2) (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) h_isAlgEnvSeq hδ From 7a4451817815e1b2bd4f7590c41d93e1d46858fe Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 30 Mar 2026 11:55:37 +0100 Subject: [PATCH 089/155] Refactor SumRewards.lean (in progress) --- LeanBandits/Bandit/SumRewards.lean | 35 +++++++++++++++--------------- 1 file changed, 17 insertions(+), 18 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 0f338af9..3938f128 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -402,30 +402,29 @@ lemma prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF {σ2 : ℝ≥0} _ = ENNReal.ofReal (2 * Real.exp (-ε ^ 2 / (2 * n * σ2))) := by rw [← ENNReal.ofReal_add (by positivity) (by positivity), ← two_mul] +/-- Auxiliary lemma for `prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF'`. -/ private lemma exp_neg_sqrt_sq_div_le {σ2 : ℝ≥0} (hσ2 : 0 < σ2) {δ : ℝ} (hδ : 0 < δ) (hn : 0 < n) : Real.exp (-√(2 * n * σ2 * Real.log (1 / δ)) ^ 2 / (2 * n * σ2)) ≤ δ := by - by_cases hδ1 : δ < 1 - · have : 0 < Real.log (1 / δ) := Real.log_pos ((one_lt_div hδ).2 hδ1) - have : Real.exp (-√(2 * n * σ2 * Real.log (1 / δ)) ^ 2 / (2 * n * σ2)) = δ := by - rw [Real.sq_sqrt (by positivity), neg_div] - field_simp - simp [Real.exp_log (by positivity)] - linarith - · calc Real.exp _ ≤ Real.exp 0 := by - gcongr - simp only [neg_div, neg_nonpos] - positivity - _ ≤ δ := by - simp [Real.exp_zero] - linarith + by_cases hd : δ < 1 + · have hl : 0 < Real.log (1 / δ) := Real.log_pos ((one_lt_div hδ).2 hd) + rw [Real.sq_sqrt (by positivity)] + field_simp + simp [Real.exp_log hδ] + · push_neg at hd + have hl : Real.log (1 / δ) ≤ 0 := Real.log_nonpos (by positivity) (div_le_one_of_le₀ hd (hδ.le)) + rw [Real.sqrt_eq_zero_of_nonpos (mul_nonpos_of_nonneg_of_nonpos (by positivity) hl)] + simp [hd] lemma prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF' {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (h : HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) {δ : ℝ} (hδ : 0 < δ) (hn : 0 < n) : streamMeasure ν {ω | √(2 * n * σ2 * Real.log (1 / δ)) ≤ - |∑ k ∈ range n, (ω k a - (ν a)[id])|} ≤ ENNReal.ofReal (2 * δ) := by - apply (prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF h (by positivity) n).trans - gcongr - exact exp_neg_sqrt_sq_div_le hσ2 hδ hn + |∑ k ∈ range n, (ω k a - (ν a)[id])|} ≤ ENNReal.ofReal (2 * δ) := + calc + _ ≤ ENNReal.ofReal (2 * Real.exp (-√(2 * n * σ2 * Real.log (1 / δ)) ^ 2 / (2 * n * σ2))) := + prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF h (by positivity) n + _ ≤ ENNReal.ofReal (2 * δ) := by + gcongr + exact exp_neg_sqrt_sq_div_le hσ2 hδ hn end StreamMeasure From 54e75f2b8569cbbe1fe9344304275aded6f1ba74 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 30 Mar 2026 12:56:38 +0100 Subject: [PATCH 090/155] Refactor SumRewards.lean (in progress) --- LeanBandits/Bandit/SumRewards.lean | 117 ++++++++++----------------- LeanBandits/BanditAlgorithms/TS.lean | 6 +- 2 files changed, 46 insertions(+), 77 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 3938f128..965d029e 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -428,75 +428,51 @@ lemma prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF' {σ2 : ℝ≥0} (hσ2 : end StreamMeasure -lemma prob_abs_sumRewards_sub_ge_le [Countable α] - {σ2 : ℝ≥0} (hσ2 : 0 < σ2) - (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) - {δ : ℝ} (hδ : 0 < δ) : - P (⋃ s ∈ Finset.range n, {ω | pullCount A a s ω ≠ 0 ∧ - √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|}) ≤ - ENNReal.ofReal (2 * n * δ) := by - by_cases hn : n = 0 - · simp [hn] - have hn : 0 < n := Nat.pos_of_ne_zero hn - let B := fun m : ℕ ↦ {x : ℝ | √(2 * m * ↑σ2 * Real.log (1 / δ)) ≤ |x - m * (ν a)[id]|} - have hB_meas : ∀ m, MeasurableSet (B m) := fun m ↦ by - simp only [B] - measurability - let S := Finset.Icc 1 (n - 1) - have hS_card : S.card = n - 1 := by simp only [Nat.card_Icc, S]; omega - have h_decomp : ⋃ s ∈ Finset.range n, {ω | pullCount A a s ω ≠ 0 ∧ - √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} = - ⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ - sumRewards A R a s ω ∈ B m} := by - ext ω - simp only [Set.mem_iUnion, Finset.mem_range, exists_prop, Set.mem_setOf_eq, - Finset.mem_Icc, S] - constructor - · rintro ⟨s, hs, hbad⟩ - let m := pullCount A a s ω - have hm_pos : 0 < m := Nat.pos_of_ne_zero hbad.1 - have hm_le : m ≤ n - 1 := by - have h1 : m ≤ s := pullCount_le (A := A) a s ω - omega - exact ⟨m, ⟨hm_pos, hm_le⟩, s, hs, rfl, hbad.2⟩ - · rintro ⟨m, ⟨hm_pos, hm_le⟩, s, hs, hpc, hB⟩ - subst hpc - exact ⟨s, hs, by omega, hB⟩ - rw [h_decomp] - calc P (⋃ m ∈ S, {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ - sumRewards A R a s ω ∈ B m}) - ≤ ∑ m ∈ S, P {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ +lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le [Countable α] {σ2 : ℝ≥0} (hσ2 : 0 < σ2) + (ha : HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {δ : ℝ} (hδ : 0 < δ) : + P {ω | ∃ t < n, pullCount A a t ω ≠ 0 ∧ + √(2 * (pullCount A a t ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R a t ω - (pullCount A a t ω : ℝ) * (ν a)[id]|} ≤ + ENNReal.ofReal (2 * n * δ) := + let B (m : ℕ) := {x : ℝ | √(2 * m * σ2 * Real.log (1 / δ)) ≤ |x - m * (ν a)[id]|} + calc + _ ≤ P (⋃ m ∈ Finset.Icc 1 (n - 1), {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ + sumRewards A R a s ω ∈ B m}) := by + apply measure_mono + intro ω ⟨s, hs, hne, hbad⟩ + simp only [Set.mem_iUnion, exists_prop, Finset.mem_Icc] + exact ⟨_, ⟨Nat.pos_of_ne_zero hne, + (pullCount_le (A := A) a s ω).trans (by omega)⟩, s, hs, rfl, hbad⟩ + _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), P {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ sumRewards A R a s ω ∈ B m} := - measure_biUnion_finset_le S _ - _ ≤ ∑ m ∈ S, streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := by + measure_biUnion_finset_le _ _ + _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), P {ω | ∃ s, pullCount A a s ω = m ∧ + sumRewards A R a s ω ∈ B m} := by apply Finset.sum_le_sum - intro m hm - calc P {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ - sumRewards A R a s ω ∈ B m} - ≤ P {ω | ∃ s, pullCount A a s ω = m ∧ - sumRewards A R a s ω ∈ B m} := - measure_mono fun ω ⟨s, _, hpc, hB⟩ ↦ ⟨s, hpc, hB⟩ - _ ≤ streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := - prob_exists_pullCount_eq_and_sumRewards_mem_le h a m (hB_meas m) - _ ≤ ∑ _m ∈ S, ENNReal.ofReal (2 * δ) := by + intro m _ + apply measure_mono + intro ω ⟨s, _, h⟩ + exact ⟨s, h⟩ + _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := by + gcongr with m hm + exact prob_exists_pullCount_eq_and_sumRewards_mem_le h a m (by measurability) + _ ≤ ∑ _m ∈ (Finset.Icc 1 (n - 1)), ENNReal.ofReal (2 * δ) := by apply sum_le_sum intro m hm convert StreamMeasure.prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF' - hσ2 (hν a) hδ (Finset.mem_Icc.mp hm).1 using 2 + hσ2 ha hδ (Finset.mem_Icc.mp hm).1 using 2 simp_rw [B, Set.mem_setOf_eq, Finset.sum_sub_distrib, Finset.sum_const, Finset.card_range, nsmul_eq_mul] _ = (n - 1) • ENNReal.ofReal (2 * δ) := by - simp only [Finset.sum_const, hS_card] + simp [Finset.sum_const, Nat.card_Icc] _ ≤ ENNReal.ofReal (2 * n * δ) := by rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast (n - 1), ← ENNReal.ofReal_mul (Nat.cast_nonneg (n - 1))] exact ENNReal.ofReal_le_ofReal (by nlinarith [(Nat.cast_le (α := ℝ)).mpr (Nat.sub_le n 1), hδ.le]) -lemma prob_abs_sumRewards_sub_ge_fintype_le [Fintype α] +lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype [Fintype α] {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) @@ -504,27 +480,20 @@ lemma prob_abs_sumRewards_sub_ge_fintype_le [Fintype α] P {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} ≤ - ENNReal.ofReal (2 * Fintype.card α * n * δ) := by - let badSet := fun (a : α) (s : ℕ) ↦ {ω : Ω | - pullCount A a s ω ≠ 0 ∧ - √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} - have h_set_eq : {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ - √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} = - ⋃ a : α, ⋃ s ∈ Finset.range n, badSet a s := by - ext ω; simp only [Set.mem_setOf_eq, Set.mem_iUnion, Finset.mem_range, badSet, exists_prop] - exact ⟨fun ⟨s, hs, a, ha⟩ ↦ ⟨a, s, hs, ha⟩, fun ⟨a, s, hs, ha⟩ ↦ ⟨s, hs, a, ha⟩⟩ - rw [h_set_eq] - have h_arm_bound : ∀ a : α, - P (⋃ s ∈ Finset.range n, badSet a s) ≤ ENNReal.ofReal (2 * n * δ) := by - intro a - exact prob_abs_sumRewards_sub_ge_le hσ2 hν h hδ - calc P (⋃ a : α, ⋃ s ∈ Finset.range n, badSet a s) - ≤ ∑ a : α, P (⋃ s ∈ Finset.range n, badSet a s) := + ENNReal.ofReal (2 * Fintype.card α * n * δ) := + calc + _ ≤ P (⋃ a : α, {ω | ∃ s < n, pullCount A a s ω ≠ 0 ∧ + √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|}) := by + apply measure_mono + intro ω ⟨s, hs, a, ha⟩ + exact Set.mem_iUnion.mpr ⟨a, s, hs, ha⟩ + _ ≤ ∑ a : α, P {ω | ∃ s < n, pullCount A a s ω ≠ 0 ∧ + √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} := measure_iUnion_fintype_le _ _ _ ≤ ∑ _a : α, ENNReal.ofReal (2 * n * δ) := - Finset.sum_le_sum fun a _ ↦ h_arm_bound a + Finset.sum_le_sum fun a _ ↦ prob_abs_sumRewards_sub_pullCount_mul_ge_le hσ2 (hν a) h hδ _ = Fintype.card α • ENNReal.ofReal (2 * n * δ) := by simp [Finset.sum_const] _ = ENNReal.ofReal (2 * Fintype.card α * n * δ) := by diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 48594c46..dfa10764 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -261,7 +261,7 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] |empMean IT.action IT.reward a s ω - ((κ.sectR e) a)[id]|} := by simp only [badSetIT, Kernel.sectR_apply] rw [this] - have h_cf := prob_abs_sumRewards_sub_ge_fintype_le (n := n) (hσ2) + have h_cf := prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype (n := n) (hσ2) (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) h_isAlgEnvSeq hδ simp only [Fintype.card_fin] at h_cf refine le_trans (measure_mono fun ω hω ↦ ?_) h_cf @@ -362,9 +362,9 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] intro a; simp only [badSetIT, Kernel.sectR_apply] rw [h_eq] set ba := IsBayesAlgEnvSeq.bestAction κ id e - have h_ccb := prob_abs_sumRewards_sub_ge_le (a := ba) (n := n) + have h_ccb := prob_abs_sumRewards_sub_pullCount_mul_ge_le (a := ba) (n := n) (hσ2) - (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) + (by simp only [Kernel.sectR_apply]; exact hs ba e) h_isAlgEnvSeq hδ refine le_trans (measure_mono fun ω hω ↦ ?_) h_ccb simp only [Set.mem_iUnion, Finset.mem_range, Set.mem_setOf_eq] at hω ⊢ From b804a9f20c680d670d3e6876ad9ab667a2c8469f Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 30 Mar 2026 15:47:24 +0100 Subject: [PATCH 091/155] Refactor SumRewards.lean (in progress) --- LeanBandits/Bandit/SumRewards.lean | 71 ++++++++++++++-------------- LeanBandits/BanditAlgorithms/TS.lean | 11 +++-- 2 files changed, 42 insertions(+), 40 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 965d029e..a9ced059 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -431,33 +431,28 @@ end StreamMeasure lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le [Countable α] {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (ha : HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {δ : ℝ} (hδ : 0 < δ) : - P {ω | ∃ t < n, pullCount A a t ω ≠ 0 ∧ - √(2 * (pullCount A a t ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R a t ω - (pullCount A a t ω : ℝ) * (ν a)[id]|} ≤ - ENNReal.ofReal (2 * n * δ) := + P {ω | ∃ t < n, pullCount A a t ω ≠ 0 ∧ √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R a t ω - pullCount A a t ω * (ν a)[id]|} ≤ + ENNReal.ofReal (2 * (n - 1) * δ) := let B (m : ℕ) := {x : ℝ | √(2 * m * σ2 * Real.log (1 / δ)) ≤ |x - m * (ν a)[id]|} calc - _ ≤ P (⋃ m ∈ Finset.Icc 1 (n - 1), {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ - sumRewards A R a s ω ∈ B m}) := by + _ ≤ P (⋃ m ∈ Finset.Icc 1 (n - 1), {ω | ∃ t, t < n ∧ pullCount A a t ω = m ∧ + sumRewards A R a t ω ∈ B m}) := by apply measure_mono - intro ω ⟨s, hs, hne, hbad⟩ - simp only [Set.mem_iUnion, exists_prop, Finset.mem_Icc] - exact ⟨_, ⟨Nat.pos_of_ne_zero hne, - (pullCount_le (A := A) a s ω).trans (by omega)⟩, s, hs, rfl, hbad⟩ - _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), P {ω | ∃ s, s < n ∧ pullCount A a s ω = m ∧ - sumRewards A R a s ω ∈ B m} := + intro ω ⟨t, ht, hp, hb⟩ + have hm : pullCount A a t ω ∈ Finset.Icc 1 (n - 1) := + Finset.mem_Icc.mpr ⟨Nat.pos_of_ne_zero hp, (pullCount_le a t ω).trans (by omega)⟩ + exact Set.mem_biUnion hm ⟨t, ht, rfl, hb⟩ + _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), P {ω | ∃ t, t < n ∧ pullCount A a t ω = m ∧ + sumRewards A R a t ω ∈ B m} := measure_biUnion_finset_le _ _ - _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), P {ω | ∃ s, pullCount A a s ω = m ∧ - sumRewards A R a s ω ∈ B m} := by - apply Finset.sum_le_sum - intro m _ - apply measure_mono - intro ω ⟨s, _, h⟩ - exact ⟨s, h⟩ - _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := by - gcongr with m hm - exact prob_exists_pullCount_eq_and_sumRewards_mem_le h a m (by measurability) - _ ≤ ∑ _m ∈ (Finset.Icc 1 (n - 1)), ENNReal.ofReal (2 * δ) := by + _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), P {ω | ∃ t, pullCount A a t ω = m ∧ + sumRewards A R a t ω ∈ B m} := + Finset.sum_le_sum (fun _ _ ↦ measure_mono (fun _ ⟨s, _, h⟩ ↦ ⟨s, h⟩)) + _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := + Finset.sum_le_sum + (fun m _ ↦ prob_exists_pullCount_eq_and_sumRewards_mem_le h a m (by measurability)) + _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), ENNReal.ofReal (2 * δ) := by apply sum_le_sum intro m hm convert StreamMeasure.prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF' @@ -465,12 +460,16 @@ lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le [Countable α] {σ2 : ℝ≥0} simp_rw [B, Set.mem_setOf_eq, Finset.sum_sub_distrib, Finset.sum_const, Finset.card_range, nsmul_eq_mul] _ = (n - 1) • ENNReal.ofReal (2 * δ) := by - simp [Finset.sum_const, Nat.card_Icc] - _ ≤ ENNReal.ofReal (2 * n * δ) := by + rw [Finset.sum_const] + congr 1 + simp [Nat.card_Icc] + _ = ENNReal.ofReal (↑(n - 1) * (2 * δ)) := by rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast (n - 1), - ← ENNReal.ofReal_mul (Nat.cast_nonneg (n - 1))] - exact ENNReal.ofReal_le_ofReal (by - nlinarith [(Nat.cast_le (α := ℝ)).mpr (Nat.sub_le n 1), hδ.le]) + ← ENNReal.ofReal_mul (Nat.cast_nonneg _)] + _ = ENNReal.ofReal (2 * (n - 1) * δ) := by + by_cases hn : n = 0 + · simp [hn, hδ.le] + · congr 1; rw [Nat.cast_sub (by omega : 1 ≤ n)]; ring lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype [Fintype α] {σ2 : ℝ≥0} (hσ2 : 0 < σ2) @@ -480,7 +479,7 @@ lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype [Fintype α] P {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} ≤ - ENNReal.ofReal (2 * Fintype.card α * n * δ) := + ENNReal.ofReal (2 * Fintype.card α * (n - 1) * δ) := calc _ ≤ P (⋃ a : α, {ω | ∃ s < n, pullCount A a s ω ≠ 0 ∧ √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ @@ -492,15 +491,15 @@ lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype [Fintype α] √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} := measure_iUnion_fintype_le _ _ - _ ≤ ∑ _a : α, ENNReal.ofReal (2 * n * δ) := - Finset.sum_le_sum fun a _ ↦ prob_abs_sumRewards_sub_pullCount_mul_ge_le hσ2 (hν a) h hδ - _ = Fintype.card α • ENNReal.ofReal (2 * n * δ) := by + _ ≤ ∑ _a : α, ENNReal.ofReal (2 * (n - 1) * δ) := + Finset.sum_le_sum fun a _ ↦ + prob_abs_sumRewards_sub_pullCount_mul_ge_le hσ2 (hν a) h hδ + _ = Fintype.card α • ENNReal.ofReal (2 * (n - 1) * δ) := by simp [Finset.sum_const] - _ = ENNReal.ofReal (2 * Fintype.card α * n * δ) := by - simp only [nsmul_eq_mul] - rw [← ENNReal.ofReal_natCast (Fintype.card α), + _ = ENNReal.ofReal (2 * Fintype.card α * (n - 1) * δ) := by + rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast (Fintype.card α), ← ENNReal.ofReal_mul (Nat.cast_nonneg (Fintype.card α))] - congr 1; ring + ring_nf omit [DecidableEq α] [StandardBorelSpace α] in lemma probReal_sum_le_sum_streamMeasure [Fintype α] {c : ℝ≥0} diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index dfa10764..b1a6e3a1 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -114,6 +114,7 @@ private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 have : √n * √n = n := Real.mul_self_sqrt (by positivity) nlinarith +/-- This bound could be improved. -/ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, μ a ∈ Set.Icc l u) (hi : l ≤ u) (hc : ∀ s < n, pullCount A (A s ω) s ω ≠ 0 → |empMean A R' (A s ω) s ω - μ (A s ω)| < √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω))) : @@ -264,6 +265,8 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] have h_cf := prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype (n := n) (hσ2) (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) h_isAlgEnvSeq hδ simp only [Fintype.card_fin] at h_cf + replace h_cf : _ ≤ ENNReal.ofReal (2 * K * n * δ) := + h_cf.trans (ENNReal.ofReal_le_ofReal (by nlinarith [Nat.cast_nonneg (α := ℝ) K])) refine le_trans (measure_mono fun ω hω ↦ ?_) h_cf simp only [Set.mem_setOf_eq, empMean] at hω obtain ⟨s, hs, a, hpc, hle⟩ := hω @@ -362,10 +365,10 @@ lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] intro a; simp only [badSetIT, Kernel.sectR_apply] rw [h_eq] set ba := IsBayesAlgEnvSeq.bestAction κ id e - have h_ccb := prob_abs_sumRewards_sub_pullCount_mul_ge_le (a := ba) (n := n) - (hσ2) - (by simp only [Kernel.sectR_apply]; exact hs ba e) - h_isAlgEnvSeq hδ + have h_ccb : _ ≤ ENNReal.ofReal (2 * n * δ) := + (prob_abs_sumRewards_sub_pullCount_mul_ge_le (a := ba) (n := n) hσ2 + (by simp only [Kernel.sectR_apply]; exact hs ba e) + h_isAlgEnvSeq hδ).trans (ENNReal.ofReal_le_ofReal (by nlinarith)) refine le_trans (measure_mono fun ω hω ↦ ?_) h_ccb simp only [Set.mem_iUnion, Finset.mem_range, Set.mem_setOf_eq] at hω ⊢ obtain ⟨s, hs, hpc, hle⟩ := hω From d31396be17bfcc69ec61ebb24a53a4bccc3f1b87 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 31 Mar 2026 11:45:31 +0100 Subject: [PATCH 092/155] Refactor SumRewards.lean (in progress) --- LeanBandits/Bandit/SumRewards.lean | 48 +++++++++++++----------------- 1 file changed, 20 insertions(+), 28 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index a9ced059..66eef485 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -230,11 +230,11 @@ lemma _root_.Learning.IsAlgEnvSeq.identDistrib_pullCount_sumRewards have hc1 : (fun ω n a ↦ (pullCount A a n ω, sumRewards A R a n ω)) = f ∘ (fun ω n ↦ (A n ω, R n ω)) := by ext ω n a : 3 - simp_rw [Function.comp, f, pullCount, Finset.card_filter, sumRewards] + simp_rw [Function.comp, f, pullCount, card_filter, sumRewards] have hc2 : (fun ω' n a ↦ (pullCount A₂ a n ω', sumRewards A₂ R₂ a n ω')) = f ∘ (fun ω' n ↦ (A₂ n ω', R₂ n ω')) := by ext ω' n a : 3 - simp_rw [Function.comp, f, pullCount, Finset.card_filter, sumRewards] + simp_rw [Function.comp, f, pullCount, card_filter, sumRewards] have hf : Measurable f := by simp_rw [f, measurable_pi_iff] intro n a @@ -432,44 +432,36 @@ lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le [Countable α] {σ2 : ℝ≥0} (ha : HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {δ : ℝ} (hδ : 0 < δ) : P {ω | ∃ t < n, pullCount A a t ω ≠ 0 ∧ √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R a t ω - pullCount A a t ω * (ν a)[id]|} ≤ - ENNReal.ofReal (2 * (n - 1) * δ) := + |sumRewards A R a t ω - pullCount A a t ω * (ν a)[id]|} ≤ ENNReal.ofReal (2 * (n - 1) * δ) := let B (m : ℕ) := {x : ℝ | √(2 * m * σ2 * Real.log (1 / δ)) ≤ |x - m * (ν a)[id]|} calc - _ ≤ P (⋃ m ∈ Finset.Icc 1 (n - 1), {ω | ∃ t, t < n ∧ pullCount A a t ω = m ∧ + _ ≤ P (⋃ m ∈ Icc 1 (n - 1), {ω | ∃ t, t < n ∧ pullCount A a t ω = m ∧ sumRewards A R a t ω ∈ B m}) := by apply measure_mono intro ω ⟨t, ht, hp, hb⟩ - have hm : pullCount A a t ω ∈ Finset.Icc 1 (n - 1) := - Finset.mem_Icc.mpr ⟨Nat.pos_of_ne_zero hp, (pullCount_le a t ω).trans (by omega)⟩ + have hm : pullCount A a t ω ∈ Icc 1 (n - 1) := mem_Icc.mpr ⟨Nat.one_le_iff_ne_zero.mpr hp, + (pullCount_le a t ω).trans (Nat.le_sub_one_of_lt ht)⟩ exact Set.mem_biUnion hm ⟨t, ht, rfl, hb⟩ - _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), P {ω | ∃ t, t < n ∧ pullCount A a t ω = m ∧ + _ ≤ ∑ m ∈ Icc 1 (n - 1), P {ω | ∃ t, t < n ∧ pullCount A a t ω = m ∧ sumRewards A R a t ω ∈ B m} := measure_biUnion_finset_le _ _ - _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), P {ω | ∃ t, pullCount A a t ω = m ∧ - sumRewards A R a t ω ∈ B m} := - Finset.sum_le_sum (fun _ _ ↦ measure_mono (fun _ ⟨s, _, h⟩ ↦ ⟨s, h⟩)) - _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := - Finset.sum_le_sum - (fun m _ ↦ prob_exists_pullCount_eq_and_sumRewards_mem_le h a m (by measurability)) - _ ≤ ∑ m ∈ Finset.Icc 1 (n - 1), ENNReal.ofReal (2 * δ) := by + _ ≤ ∑ m ∈ Icc 1 (n - 1), P {ω | ∃ t, pullCount A a t ω = m ∧ sumRewards A R a t ω ∈ B m} := + sum_le_sum (fun _ _ ↦ measure_mono (fun _ ⟨t, _, hps⟩ ↦ ⟨t, hps⟩)) + _ ≤ ∑ m ∈ Icc 1 (n - 1), streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := by + apply sum_le_sum + exact (fun m _ ↦ prob_exists_pullCount_eq_and_sumRewards_mem_le h a m (by measurability)) + _ ≤ ∑ m ∈ Icc 1 (n - 1), ENNReal.ofReal (2 * δ) := by apply sum_le_sum intro m hm convert StreamMeasure.prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF' - hσ2 ha hδ (Finset.mem_Icc.mp hm).1 using 2 - simp_rw [B, Set.mem_setOf_eq, Finset.sum_sub_distrib, Finset.sum_const, - Finset.card_range, nsmul_eq_mul] - _ = (n - 1) • ENNReal.ofReal (2 * δ) := by - rw [Finset.sum_const] - congr 1 - simp [Nat.card_Icc] - _ = ENNReal.ofReal (↑(n - 1) * (2 * δ)) := by - rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast (n - 1), - ← ENNReal.ofReal_mul (Nat.cast_nonneg _)] + hσ2 ha hδ (mem_Icc.mp hm).1 using 2 + simp [B] _ = ENNReal.ofReal (2 * (n - 1) * δ) := by by_cases hn : n = 0 · simp [hn, hδ.le] - · congr 1; rw [Nat.cast_sub (by omega : 1 ≤ n)]; ring + · rw [sum_const, Nat.card_Icc, add_tsub_cancel_right, ← ENNReal.ofReal_nsmul, nsmul_eq_mul, + Nat.cast_sub (Nat.one_le_iff_ne_zero.mpr hn)] + ring_nf lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype [Fintype α] {σ2 : ℝ≥0} (hσ2 : 0 < σ2) @@ -492,10 +484,10 @@ lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype [Fintype α] |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} := measure_iUnion_fintype_le _ _ _ ≤ ∑ _a : α, ENNReal.ofReal (2 * (n - 1) * δ) := - Finset.sum_le_sum fun a _ ↦ + sum_le_sum fun a _ ↦ prob_abs_sumRewards_sub_pullCount_mul_ge_le hσ2 (hν a) h hδ _ = Fintype.card α • ENNReal.ofReal (2 * (n - 1) * δ) := by - simp [Finset.sum_const] + simp [sum_const] _ = ENNReal.ofReal (2 * Fintype.card α * (n - 1) * δ) := by rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast (Fintype.card α), ← ENNReal.ofReal_mul (Nat.cast_nonneg (Fintype.card α))] From b872d1590fd9597f8c96fbdf2e700d70ced3a40e Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 31 Mar 2026 13:36:50 +0100 Subject: [PATCH 093/155] Refactor SumRewards.lean --- LeanBandits/Bandit/SumRewards.lean | 39 ++++++++++------------------ LeanBandits/BanditAlgorithms/TS.lean | 8 +++--- 2 files changed, 19 insertions(+), 28 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 66eef485..b1fa1125 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -463,34 +463,23 @@ lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le [Countable α] {σ2 : ℝ≥0} Nat.cast_sub (Nat.one_le_iff_ne_zero.mpr hn)] ring_nf -lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype [Fintype α] - {σ2 : ℝ≥0} (hσ2 : 0 < σ2) +lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype [Fintype α] {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) - (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) - {δ : ℝ} (hδ : 0 < δ) : - P {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ - √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} ≤ - ENNReal.ofReal (2 * Fintype.card α * (n - 1) * δ) := + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {δ : ℝ} (hδ : 0 < δ) : + P {ω | ∃ a, ∃ t < n, pullCount A a t ω ≠ 0 ∧ + √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R a t ω - pullCount A a t ω * (ν a)[id]|} ≤ + ENNReal.ofReal (2 * Fintype.card α * (n - 1) * δ) := calc - _ ≤ P (⋃ a : α, {ω | ∃ s < n, pullCount A a s ω ≠ 0 ∧ - √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|}) := by - apply measure_mono - intro ω ⟨s, hs, a, ha⟩ - exact Set.mem_iUnion.mpr ⟨a, s, hs, ha⟩ - _ ≤ ∑ a : α, P {ω | ∃ s < n, pullCount A a s ω ≠ 0 ∧ - √(2 * (pullCount A a s ω : ℝ) * ↑σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R a s ω - (pullCount A a s ω : ℝ) * (ν a)[id]|} := - measure_iUnion_fintype_le _ _ - _ ≤ ∑ _a : α, ENNReal.ofReal (2 * (n - 1) * δ) := - sum_le_sum fun a _ ↦ - prob_abs_sumRewards_sub_pullCount_mul_ge_le hσ2 (hν a) h hδ - _ = Fintype.card α • ENNReal.ofReal (2 * (n - 1) * δ) := by - simp [sum_const] + _ ≤ ∑ a, P {ω | ∃ t < n, pullCount A a t ω ≠ 0 ∧ + √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R a t ω - pullCount A a t ω * (ν a)[id]|} := by + rw [Set.setOf_exists] + exact measure_iUnion_fintype_le _ _ + _ ≤ ∑ a, ENNReal.ofReal (2 * (n - 1) * δ) := + sum_le_sum fun a _ ↦ prob_abs_sumRewards_sub_pullCount_mul_ge_le hσ2 (hν a) h hδ _ = ENNReal.ofReal (2 * Fintype.card α * (n - 1) * δ) := by - rw [nsmul_eq_mul, ← ENNReal.ofReal_natCast (Fintype.card α), - ← ENNReal.ofReal_mul (Nat.cast_nonneg (Fintype.card α))] + rw [sum_const, Finset.card_univ, ← ENNReal.ofReal_nsmul, nsmul_eq_mul] ring_nf omit [DecidableEq α] [StandardBorelSpace α] in diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index b1a6e3a1..47538828 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -257,10 +257,12 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) := by rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h filter_upwards [h_cond_ae] with e h_isAlgEnvSeq - have : badSetIT e = {ω | ∃ s < n, ∃ a, pullCount IT.action a s ω ≠ 0 ∧ + have : badSetIT e = {ω | ∃ a, ∃ s < n, pullCount IT.action a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ |empMean IT.action IT.reward a s ω - ((κ.sectR e) a)[id]|} := by simp only [badSetIT, Kernel.sectR_apply] + ext ω; simp only [Set.mem_setOf_eq] + exact ⟨fun ⟨s, hs, a, ha⟩ ↦ ⟨a, s, hs, ha⟩, fun ⟨a, s, hs, ha⟩ ↦ ⟨s, hs, a, ha⟩⟩ rw [this] have h_cf := prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype (n := n) (hσ2) (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) h_isAlgEnvSeq hδ @@ -269,11 +271,11 @@ lemma prob_concentration_fail_delta [Nonempty (Fin K)] h_cf.trans (ENNReal.ofReal_le_ofReal (by nlinarith [Nat.cast_nonneg (α := ℝ) K])) refine le_trans (measure_mono fun ω hω ↦ ?_) h_cf simp only [Set.mem_setOf_eq, empMean] at hω - obtain ⟨s, hs, a, hpc, hle⟩ := hω + obtain ⟨a, s, hs, hpc, hle⟩ := hω simp only [Set.mem_setOf_eq] have hk : (0 : ℝ) < pullCount IT.action a s ω := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hpc) - refine ⟨s, hs, a, hpc, ?_⟩ + refine ⟨a, s, hs, hpc, ?_⟩ rw [show sumRewards IT.action IT.reward a s ω / ↑(pullCount IT.action a s ω) - ((κ.sectR e) a)[id] = (sumRewards IT.action IT.reward a s ω - ↑(pullCount IT.action a s ω) * ((κ.sectR e) a)[id]) / From 7d0742f559d7b0472622a92db8e4646b8e4ffae1 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 7 Apr 2026 15:45:48 +0100 Subject: [PATCH 094/155] Refactor TS.lean (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 466 +++++++++++---------------- 1 file changed, 194 insertions(+), 272 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 47538828..2107bcb2 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -114,7 +114,7 @@ private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 have : √n * √n = n := Real.mul_self_sqrt (by positivity) nlinarith -/-- This bound could be improved. -/ +/-- This bound could be improved slightly. -/ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, μ a ∈ Set.Icc l u) (hi : l ≤ u) (hc : ∀ s < n, pullCount A (A s ω) s ω ≠ 0 → |empMean A R' (A s ω) s ω - μ (A s ω)| < √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω))) : @@ -185,273 +185,200 @@ end Bandits open Bandits -/-! ### Algorithm-generic Bayesian lemmas -/ - -section BayesianConcentration - -variable {K : ℕ} {𝓔 : Type*} [MeasurableSpace 𝓔] {Ω : Type*} [MeasurableSpace Ω] -variable (E : Ω → 𝓔) (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → ℝ) -variable (Q : Measure 𝓔) (κ : Kernel (𝓔 × Fin K) ℝ) -variable (P : Measure Ω) [IsProbabilityMeasure P] - namespace Learning.IsBayesAlgEnvSeq -variable [IsMarkovKernel κ] - -lemma prob_concentration_fail_delta [Nonempty (Fin K)] - {alg : Algorithm (Fin K) ℝ} - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - {σ2 : ℝ≥0} (hσ2 : 0 < σ2) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - P {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} - ≤ ENNReal.ofReal (2 * K * n * δ) := by - let badSetIT := fun (e : 𝓔) ↦ {ω : ℕ → (Fin K) × ℝ | - ∃ s < n, ∃ a, pullCount IT.action a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} - have h_set_eq : {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} = - (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' - {p | p.2 ∈ badSetIT p.1} := by - ext ω - simp only [Set.mem_setOf_eq, Set.mem_preimage, badSetIT, IsBayesAlgEnvSeq.actionMean] - rfl - rw [h_set_eq] - have h_meas_pair : - Measurable (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) := - h.measurable_E.prodMk (IsBayesAlgEnvSeq.measurable_trajectory h.measurable_A h.measurable_R) - have h_disint : P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) = - P.map E ⊗ₘ condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P := - (compProd_map_condDistrib - (IsBayesAlgEnvSeq.measurable_trajectory - h.measurable_A h.measurable_R).aemeasurable).symm - have h_kernel : ∀ a, Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := - fun a ↦ stronglyMeasurable_id.integral_kernel.measurable.comp - (measurable_fst.prodMk measurable_const) - have h_meas_set : MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ badSetIT p.1} := by - have h_eq : {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ badSetIT p.1} = - ⋃ s ∈ Finset.range n, ⋃ a : Fin K, {p | - pullCount IT.action a s p.2 ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s p.2 : ℝ)) ≤ - |empMean IT.action IT.reward a s p.2 - (κ (p.1, a))[id]|} := by - ext p; simp only [badSetIT, Set.mem_setOf_eq, Set.mem_iUnion, Finset.mem_range] - exact ⟨fun ⟨s, hs, a, ha⟩ ↦ ⟨s, hs, a, ha⟩, fun ⟨s, hs, a, ha⟩ ↦ ⟨s, hs, a, ha⟩⟩ - rw [h_eq] - exact .biUnion (Finset.range n).countable_toSet fun s _ ↦ - .iUnion fun a ↦ - MeasurableSet.inter - (((measurable_pullCount IT.measurable_action a s).comp measurable_snd) - (measurableSet_singleton (0 : ℕ)).compl) - (measurableSet_le (by fun_prop) - (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp - measurable_snd).sub (h_kernel a)).abs) - have h_cond_bound : ∀ᵐ e ∂(P.map E), - (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) (badSetIT e) ≤ - ENNReal.ofReal (2 * K * n * δ) := by - have h_cond_ae : ∀ᵐ e ∂(P.map E), IsAlgEnvSeq IT.action IT.reward - alg (stationaryEnv (κ.sectR e)) - (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) := by - rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h - filter_upwards [h_cond_ae] with e h_isAlgEnvSeq - have : badSetIT e = {ω | ∃ a, ∃ s < n, pullCount IT.action a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - ((κ.sectR e) a)[id]|} := by - simp only [badSetIT, Kernel.sectR_apply] - ext ω; simp only [Set.mem_setOf_eq] - exact ⟨fun ⟨s, hs, a, ha⟩ ↦ ⟨a, s, hs, ha⟩, fun ⟨a, s, hs, ha⟩ ↦ ⟨s, hs, a, ha⟩⟩ - rw [this] - have h_cf := prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype (n := n) (hσ2) - (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs a e) h_isAlgEnvSeq hδ - simp only [Fintype.card_fin] at h_cf - replace h_cf : _ ≤ ENNReal.ofReal (2 * K * n * δ) := - h_cf.trans (ENNReal.ofReal_le_ofReal (by nlinarith [Nat.cast_nonneg (α := ℝ) K])) - refine le_trans (measure_mono fun ω hω ↦ ?_) h_cf - simp only [Set.mem_setOf_eq, empMean] at hω - obtain ⟨a, s, hs, hpc, hle⟩ := hω - simp only [Set.mem_setOf_eq] - have hk : (0 : ℝ) < pullCount IT.action a s ω := - Nat.cast_pos.mpr (Nat.pos_of_ne_zero hpc) - refine ⟨a, s, hs, hpc, ?_⟩ - rw [show sumRewards IT.action IT.reward a s ω / ↑(pullCount IT.action a s ω) - - ((κ.sectR e) a)[id] = (sumRewards IT.action IT.reward a s ω - - ↑(pullCount IT.action a s ω) * ((κ.sectR e) a)[id]) / - ↑(pullCount IT.action a s ω) from by field_simp, - abs_div, abs_of_pos hk, le_div_iff₀ hk] at hle - have hlog : (0 : ℝ) < Real.log (1 / δ) := Real.log_pos (by rw [lt_div_iff₀ hδ]; linarith) - rwa [show √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount IT.action a s ω)) * - ↑(pullCount IT.action a s ω) = - √(2 * ↑(pullCount IT.action a s ω) * ↑σ2 * Real.log (1 / δ)) from by - rw [show √_ * ↑(pullCount IT.action a s ω) = - √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount IT.action a s ω) * - ↑(pullCount IT.action a s ω) ^ 2) from by - rw [Real.sqrt_mul (by positivity), Real.sqrt_sq hk.le] - ] - congr 1; field_simp] at hle - calc P ((fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' - {p | p.2 ∈ badSetIT p.1}) - = (P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω))) - {p | p.2 ∈ badSetIT p.1} := by - rw [Measure.map_apply h_meas_pair h_meas_set] - _ = (P.map E ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P) - {p | p.2 ∈ badSetIT p.1} := by - rw [h_disint] - _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) - (badSetIT e) ∂(P.map E) := by - rw [Measure.compProd_apply h_meas_set]; rfl - _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * K * n * δ) ∂(P.map E) := by - apply lintegral_mono_ae h_cond_bound - _ = ENNReal.ofReal (2 * K * n * δ) := by - rw [lintegral_const, Measure.map_apply h.measurable_E MeasurableSet.univ] +variable {K : ℕ} [Nonempty (Fin K)] +variable {𝓔 Ω : Type*} [MeasurableSpace 𝓔] [MeasurableSpace Ω] +variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] {alg : Algorithm (Fin K) ℝ} +variable {E : Ω → 𝓔} {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → ℝ} +variable {P : Measure Ω} [IsProbabilityMeasure P] + +lemma prob_abs_sumRewards_sub_pullCount_mul_actionMean_ge_le + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} + (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : + P {ω | ∃ a, ∃ t < n, pullCount A a t ω ≠ 0 ∧ + √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R' a t ω - pullCount A a t ω * actionMean κ E a ω|} + ≤ ENNReal.ofReal (2 * K * (n - 1) * δ) := by + let s (e : 𝓔) := {ω | ∃ a, ∃ s < n, pullCount IT.action a s ω ≠ 0 ∧ + √(2 * pullCount IT.action a s ω * σ2 * Real.log (1 / δ)) ≤ + |sumRewards IT.action IT.reward a s ω - pullCount IT.action a s ω * (κ (e, a))[id]|} + have : {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ s p.1} = + ⋃ a, ⋃ t ∈ Finset.range n, + (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (pullCount IT.action a t p.2, + sumRewards IT.action IT.reward a t p.2, (κ (p.1, a))[id])) ⁻¹' + {t : ℕ × ℝ × ℝ | t.1 ≠ 0 ∧ + √(2 * t.1 * σ2 * Real.log (1 / δ)) ≤ |t.2.1 - t.1 * t.2.2|} := by + ext p + simp only [s, Set.mem_setOf_eq, Set.mem_iUnion, Set.mem_preimage, Finset.mem_range] + tauto + have (a : Fin K) : Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := + measurable_actionMean measurable_fst + calc P {ω | ∃ a, ∃ s < n, pullCount A a s ω ≠ 0 ∧ + √(2 * pullCount A a s ω * σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R' a s ω - pullCount A a s ω * actionMean κ E a ω|} + _ = P ((fun ω ↦ (E ω, trajectory A R' ω)) ⁻¹' {p | p.2 ∈ s p.1}) := by + congr 1 + _ = (P.map (fun ω ↦ (E ω, trajectory A R' ω))) {p | p.2 ∈ s p.1} := + (Measure.map_apply (h.measurable_E.prodMk + (measurable_trajectory h.measurable_A h.measurable_R)) (by measurability)).symm + _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) {p | p.2 ∈ s p.1} := by + rw [(compProd_map_condDistrib + (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable).symm] + _ = ∫⁻ e, (condDistrib (trajectory A R') E P e) (s e) ∂(P.map E) := by + rw [Measure.compProd_apply (by measurability)] + rfl + _ ≤ ∫⁻ _, ENNReal.ofReal (2 * K * (n - 1) * δ) ∂(P.map E) := by + apply lintegral_mono_ae + rw [h.hasLaw_env.map_eq] + filter_upwards [ae_IsAlgEnvSeq h] with e he + simpa [Fintype.card_fin, s, Kernel.sectR_apply] using + prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 + (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs e a) he hδ + _ = ENNReal.ofReal (2 * K * (n - 1) * δ) := by + rw [lintegral_const, Measure.map_apply h.measurable_E .univ] simp [measure_univ] -lemma prob_concentration_bestArm_fail_delta [Nonempty (Fin K)] - {alg : Algorithm (Fin K) ℝ} - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - {σ2 : ℝ≥0} (hσ2 : 0 < σ2) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : - P {ω | ∃ s < n, pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / - (pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) s ω : ℝ)) ≤ - |empMean A R' (IsBayesAlgEnvSeq.bestAction κ E ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω|} - ≤ ENNReal.ofReal (2 * n * δ) := by +lemma prob_abs_sumRewards_bestAction_sub_pullCount_mul_actionMean_ge_le + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : + P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ + √(2 * pullCount A (bestAction κ E ω) t ω * σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R' (bestAction κ E ω) t ω - + pullCount A (bestAction κ E ω) t ω * actionMean κ E (bestAction κ E ω) ω|} + ≤ ENNReal.ofReal (2 * (n - 1) * δ) := by by_cases hn : n = 0 · simp [hn] - have hn' : 0 < n := Nat.pos_of_ne_zero hn - rw [show IsBayesAlgEnvSeq.bestAction κ E = IsBayesAlgEnvSeq.bestAction κ id ∘ E from - rfl] - let badSetIT := fun (a : Fin K) (s : ℕ) (e : 𝓔) ↦ {ω : ℕ → (Fin K) × ℝ | - pullCount IT.action a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - (κ (e, a))[id]|} - have h_set_eq : {ω | ∃ s < n, pullCount A ((IsBayesAlgEnvSeq.bestAction κ id ∘ E) ω) s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / - (pullCount A ((IsBayesAlgEnvSeq.bestAction κ id ∘ E) ω) s ω : ℝ)) ≤ - |empMean A R' ((IsBayesAlgEnvSeq.bestAction κ id ∘ E) ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E ((IsBayesAlgEnvSeq.bestAction κ id ∘ E) ω) ω|} = - (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' - {p | p.2 ∈ ⋃ s ∈ Finset.range n, - badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} := by - ext ω - simp only [Set.mem_setOf_eq, Finset.mem_range, Set.mem_preimage, Set.mem_iUnion, - badSetIT, IsBayesAlgEnvSeq.actionMean, Function.comp_apply, exists_prop] - rfl - rw [h_set_eq] - have h_meas_pair : - Measurable (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) := - h.measurable_E.prodMk (IsBayesAlgEnvSeq.measurable_trajectory h.measurable_A h.measurable_R) - have h_disint : P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) = - P.map E ⊗ₘ condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P := - (compProd_map_condDistrib - (IsBayesAlgEnvSeq.measurable_trajectory - h.measurable_A h.measurable_R).aemeasurable).symm - have h_cond_best : ∀ᵐ e ∂(P.map E), - (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) - (⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id e) s e) ≤ - ENNReal.ofReal (2 * n * δ) := by - have h_cond_ae : ∀ᵐ e ∂(P.map E), IsAlgEnvSeq IT.action IT.reward - alg (stationaryEnv (κ.sectR e)) - (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) := by - rw [h.hasLaw_env.map_eq]; exact IsBayesAlgEnvSeq.ae_IsAlgEnvSeq h - filter_upwards [h_cond_ae] with e h_isAlgEnvSeq - have h_eq : ∀ a, ⋃ s ∈ Finset.range n, badSetIT a s e = - ⋃ s ∈ Finset.range n, {ω | pullCount IT.action a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s ω : ℝ)) ≤ - |empMean IT.action IT.reward a s ω - ((κ.sectR e) a)[id]|} := by - intro a; simp only [badSetIT, Kernel.sectR_apply] - rw [h_eq] - set ba := IsBayesAlgEnvSeq.bestAction κ id e - have h_ccb : _ ≤ ENNReal.ofReal (2 * n * δ) := - (prob_abs_sumRewards_sub_pullCount_mul_ge_le (a := ba) (n := n) hσ2 - (by simp only [Kernel.sectR_apply]; exact hs ba e) - h_isAlgEnvSeq hδ).trans (ENNReal.ofReal_le_ofReal (by nlinarith)) - refine le_trans (measure_mono fun ω hω ↦ ?_) h_ccb - simp only [Set.mem_iUnion, Finset.mem_range, Set.mem_setOf_eq] at hω ⊢ - obtain ⟨s, hs, hpc, hle⟩ := hω - simp only [empMean] at hle - have hk : (0 : ℝ) < pullCount IT.action ba s ω := - Nat.cast_pos.mpr (Nat.pos_of_ne_zero hpc) - refine ⟨s, hs, hpc, ?_⟩ - rw [show sumRewards IT.action IT.reward ba s ω / ↑(pullCount IT.action ba s ω) - - ((κ.sectR e) ba)[id] = (sumRewards IT.action IT.reward ba s ω - - ↑(pullCount IT.action ba s ω) * ((κ.sectR e) ba)[id]) / - ↑(pullCount IT.action ba s ω) from by field_simp, - abs_div, abs_of_pos hk, le_div_iff₀ hk] at hle - have hlog : (0 : ℝ) < Real.log (1 / δ) := Real.log_pos (by rw [lt_div_iff₀ hδ]; linarith) - rwa [show √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount IT.action ba s ω)) * - ↑(pullCount IT.action ba s ω) = - √(2 * ↑(pullCount IT.action ba s ω) * ↑σ2 * Real.log (1 / δ)) from by - rw [show √_ * ↑(pullCount IT.action ba s ω) = - √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount IT.action ba s ω) * - ↑(pullCount IT.action ba s ω) ^ 2) from by - rw [Real.sqrt_mul (by positivity), Real.sqrt_sq hk.le]] - congr 1; field_simp] at hle - have h_kernel : ∀ a, Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := - fun a ↦ stronglyMeasurable_id.integral_kernel.measurable.comp - (measurable_fst.prodMk measurable_const) - have h_meas_badSetIT : ∀ a s, MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | - p.2 ∈ badSetIT a s p.1} := by - intro a s - simp only [badSetIT] - change MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | - pullCount IT.action a s p.2 ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount IT.action a s p.2 : ℝ)) ≤ - |empMean IT.action IT.reward a s p.2 - (κ (p.1, a))[id]|} - exact MeasurableSet.inter - (((measurable_pullCount IT.measurable_action a s).comp measurable_snd) - (measurableSet_singleton (0 : ℕ)).compl) - (measurableSet_le (by fun_prop) - (((measurable_empMean IT.measurable_action IT.measurable_reward a s).comp - measurable_snd).sub (h_kernel a)).abs) - have h_meas_set : MeasurableSet {p : 𝓔 × (ℕ → (Fin K) × ℝ) | - p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} := by - have h_eq : {p : 𝓔 × (ℕ → (Fin K) × ℝ) | - p.2 ∈ ⋃ s ∈ Finset.range n, badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} = - ⋃ a : Fin K, ((IsBayesAlgEnvSeq.bestAction κ id ∘ Prod.fst) ⁻¹' {a} ∩ - ⋃ s ∈ Finset.range n, {p | p.2 ∈ badSetIT a s p.1}) := by - ext p; simp only [Set.mem_setOf_eq, Set.mem_iUnion, Set.mem_inter_iff, - Set.mem_preimage, Function.comp_apply, Set.mem_singleton_iff, Finset.mem_range] - constructor - · intro ⟨s, hs, hm⟩; exact ⟨IsBayesAlgEnvSeq.bestAction κ id p.1, rfl, s, hs, hm⟩ - · rintro ⟨a, ha, s, hs, hm⟩; exact ⟨s, hs, ha ▸ hm⟩ - rw [h_eq] - exact .iUnion fun a ↦ .inter - ((IsBayesAlgEnvSeq.measurable_bestAction (κ := κ) measurable_id |>.comp - measurable_fst) (measurableSet_singleton a)) - (.biUnion (Finset.range n).countable_toSet fun s _ ↦ h_meas_badSetIT a s) - calc P ((fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω)) ⁻¹' - {p | p.2 ∈ ⋃ s ∈ Finset.range n, - badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1}) - = (P.map (fun ω ↦ (E ω, IsBayesAlgEnvSeq.trajectory A R' ω))) - {p | p.2 ∈ ⋃ s ∈ Finset.range n, - badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} := by - rw [Measure.map_apply h_meas_pair h_meas_set] - _ = (P.map E ⊗ₘ - condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P) - {p | p.2 ∈ ⋃ s ∈ Finset.range n, - badSetIT (IsBayesAlgEnvSeq.bestAction κ id p.1) s p.1} := by - rw [h_disint] - _ = ∫⁻ e, (condDistrib (IsBayesAlgEnvSeq.trajectory A R') E P e) - (⋃ s ∈ Finset.range n, - badSetIT (IsBayesAlgEnvSeq.bestAction κ id e) s e) ∂(P.map E) := by - rw [Measure.compProd_apply h_meas_set]; rfl - _ ≤ ∫⁻ _e, ENNReal.ofReal (2 * n * δ) ∂(P.map E) := by - apply lintegral_mono_ae h_cond_best - _ = ENNReal.ofReal (2 * n * δ) := by - rw [lintegral_const, Measure.map_apply h.measurable_E MeasurableSet.univ] + let s (a : Fin K) (t : ℕ) (e : 𝓔) := {ω : ℕ → (Fin K) × ℝ | + pullCount IT.action a t ω ≠ 0 ∧ + √(2 * pullCount IT.action a t ω * σ2 * Real.log (1 / δ)) ≤ + |sumRewards IT.action IT.reward a t ω - + pullCount IT.action a t ω * (κ (e, a))[id]|} + have (a : Fin K) : Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := + measurable_actionMean measurable_fst + have : {p : 𝓔 × (ℕ → (Fin K) × ℝ) | + p.2 ∈ ⋃ t ∈ Finset.range n, s (bestAction κ id p.1) t p.1} = + ⋃ a : Fin K, ((bestAction κ id ∘ Prod.fst) ⁻¹' {a} ∩ + ⋃ t ∈ Finset.range n, + (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (pullCount IT.action a t p.2, + sumRewards IT.action IT.reward a t p.2, (κ (p.1, a))[id])) ⁻¹' + {u : ℕ × ℝ × ℝ | u.1 ≠ 0 ∧ + √(2 * u.1 * σ2 * Real.log (1 / δ)) ≤ |u.2.1 - u.1 * u.2.2|}) := by + ext p + simp only [s, Set.mem_setOf_eq, Set.mem_iUnion, Set.mem_inter_iff, + Set.mem_preimage, Function.comp_apply, Set.mem_singleton_iff, Finset.mem_range] + constructor + · intro ⟨t, ht, hm⟩ + exact ⟨bestAction κ id p.1, rfl, t, ht, hm⟩ + · rintro ⟨a, ha, t, ht, hm⟩ + exact ⟨t, ht, ha ▸ hm⟩ + calc P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ + √(2 * pullCount A (bestAction κ E ω) t ω * σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R' (bestAction κ E ω) t ω - + pullCount A (bestAction κ E ω) t ω * actionMean κ E (bestAction κ E ω) ω|} + _ = P ((fun ω ↦ (E ω, trajectory A R' ω)) ⁻¹' + {p | p.2 ∈ ⋃ t ∈ Finset.range n, s (bestAction κ id p.1) t p.1}) := by + congr 1 + ext ω + simp only [Set.mem_setOf_eq, Finset.mem_range, Set.mem_preimage, + Set.mem_iUnion, s, actionMean, exists_prop] + rfl + _ = (P.map (fun ω ↦ (E ω, trajectory A R' ω))) + {p | p.2 ∈ ⋃ t ∈ Finset.range n, s (bestAction κ id p.1) t p.1} := + (Measure.map_apply (h.measurable_E.prodMk + (measurable_trajectory h.measurable_A h.measurable_R)) (by measurability)).symm + _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) + {p | p.2 ∈ ⋃ t ∈ Finset.range n, s (bestAction κ id p.1) t p.1} := by + rw [(compProd_map_condDistrib + (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable).symm] + _ = ∫⁻ e, (condDistrib (trajectory A R') E P e) + (⋃ t ∈ Finset.range n, s (bestAction κ id e) t e) ∂(P.map E) := by + rw [Measure.compProd_apply (by measurability)] + rfl + _ ≤ ∫⁻ _, ENNReal.ofReal (2 * (n - 1) * δ) ∂(P.map E) := by + apply lintegral_mono_ae + rw [h.hasLaw_env.map_eq] + filter_upwards [ae_IsAlgEnvSeq h] with e he + have : {ω | ∃ t < n, ω ∈ s (bestAction κ id e) t e} = + ⋃ t ∈ Finset.range n, s (bestAction κ id e) t e := by + ext ω + simp [Finset.mem_range] + simp only [← this, s, Set.mem_setOf_eq] + exact prob_abs_sumRewards_sub_pullCount_mul_ge_le hσ2 + (by simp only [Kernel.sectR_apply]; exact hs e _) he hδ + _ = ENNReal.ofReal (2 * (n - 1) * δ) := by + rw [lintegral_const, Measure.map_apply h.measurable_E .univ] simp [measure_univ] -end Learning.IsBayesAlgEnvSeq - -end BayesianConcentration +omit [Nonempty (Fin K)] [MeasurableSpace 𝓔] [MeasurableSpace Ω] in +private lemma abs_sumRewards_sub_pullCount_mul_ge {a : Fin K} {n : ℕ} {ω : Ω} + {μ σ2 δ : ℝ} (hpc : pullCount A a n ω ≠ 0) + (h : √(2 * σ2 * Real.log (1 / δ) / pullCount A a n ω) ≤ + |empMean A R' a n ω - μ|) : + √(2 * pullCount A a n ω * σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R' a n ω - pullCount A a n ω * μ| := by + have hk : (0 : ℝ) < pullCount A a n ω := + Nat.cast_pos.mpr (Nat.pos_of_ne_zero hpc) + simp only [empMean] at h + calc √(2 * pullCount A a n ω * σ2 * Real.log (1 / δ)) + _ ≤ √(2 * σ2 * Real.log (1 / δ) / pullCount A a n ω) * pullCount A a n ω := by + by_cases hc : 0 ≤ 2 * σ2 * Real.log (1 / δ) + · apply le_of_eq + have : 2 * pullCount A a n ω * σ2 * Real.log (1 / δ) = + 2 * σ2 * Real.log (1 / δ) / pullCount A a n ω * pullCount A a n ω ^ 2 := by + field_simp + rw [this, Real.sqrt_mul (div_nonneg hc hk.le), Real.sqrt_sq hk.le] + · rw [Real.sqrt_eq_zero_of_nonpos (by push_neg at hc; nlinarith)] + exact mul_nonneg (Real.sqrt_nonneg _) hk.le + _ ≤ |sumRewards A R' a n ω / pullCount A a n ω - μ| * pullCount A a n ω := + mul_le_mul_of_nonneg_right h hk.le + _ = |sumRewards A R' a n ω - pullCount A a n ω * μ| := by + have h_div : sumRewards A R' a n ω / ↑(pullCount A a n ω) - μ = + (sumRewards A R' a n ω - pullCount A a n ω * μ) / pullCount A a n ω := by + field_simp + rw [h_div, abs_div, abs_of_pos hk, div_mul_cancel₀ _ (ne_of_gt hk)] + +lemma prob_abs_empMean_sub_actionMean_ge_le + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} + (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : + P {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ + √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ + |empMean A R' a t ω - actionMean κ E a ω|} + ≤ ENNReal.ofReal (2 * K * (n - 1) * δ) := + calc _ + _ ≤ P {ω | ∃ a, ∃ t < n, pullCount A a t ω ≠ 0 ∧ + √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R' a t ω - pullCount A a t ω * actionMean κ E a ω|} := by + apply measure_mono + intro ω ⟨t, ht, a, hpc, hle⟩ + exact ⟨a, t, ht, hpc, abs_sumRewards_sub_pullCount_mul_ge hpc hle⟩ + _ ≤ _ := h.prob_abs_sumRewards_sub_pullCount_mul_actionMean_ge_le hσ2 hs hδ n + +lemma prob_abs_empMean_bestAction_sub_actionMean_ge_le + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : + P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ + √(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω : ℝ)) ≤ + |empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω|} + ≤ ENNReal.ofReal (2 * (n - 1) * δ) := + calc _ + _ ≤ P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ + √(2 * pullCount A (bestAction κ E ω) t ω * σ2 * Real.log (1 / δ)) ≤ + |sumRewards A R' (bestAction κ E ω) t ω - + pullCount A (bestAction κ E ω) t ω * + actionMean κ E (bestAction κ E ω) ω|} := by + apply measure_mono + intro ω ⟨t, ht, hpc, hle⟩ + exact ⟨t, ht, hpc, abs_sumRewards_sub_pullCount_mul_ge hpc hle⟩ + _ ≤ _ := + h.prob_abs_sumRewards_bestAction_sub_pullCount_mul_actionMean_ge_le + hσ2 hs hδ n -/-! ### TS-specific regret bounds -/ +end Learning.IsBayesAlgEnvSeq namespace Bandits @@ -492,9 +419,9 @@ lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {l u : ℝ} (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) - (n : ℕ) (δ : ℝ) (hδ : 0 < δ) (hδ1 : δ < 1) : + (n : ℕ) (δ : ℝ) (hδ : 0 < δ) : P[IsBayesAlgEnvSeq.regret κ E A n] ≤ (u - l) * ↑K + 2 * (↑K + 1) * (u - l) * n ^ 2 * δ + 4 * √(2 * ↑σ2 * Real.log (1 / δ) * ↑K * ↑n) := by @@ -661,8 +588,8 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp |empMean A R' a s ω - armMean a ω|} := by ext ω; simp only [Eδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push_neg; rfl rw [this] - exact IsBayesAlgEnvSeq.prob_concentration_fail_delta (E := E) (A := A) (R' := R') - (Q := Q) (κ := κ) (P := P) h hσ2 hs n δ hδ hδ1 + exact (h.prob_abs_empMean_sub_actionMean_ge_le hσ2 hs hδ n).trans + (ENNReal.ofReal_le_ofReal (by nlinarith [hδ.le, Nat.cast_nonneg (α := ℝ) K])) have hm_emp : ∀ a s, Measurable (fun ω ↦ empMean A R' a s ω) := fun a s ↦ measurable_empMean (fun n ↦ h.measurable_A n) (fun n ↦ h.measurable_R n) a s have hm_pc : ∀ a s, Measurable (fun ω ↦ (pullCount A a s ω : ℝ)) := @@ -698,8 +625,8 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω|} := by ext ω; simp only [Fδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push_neg; rfl rw [this] - exact IsBayesAlgEnvSeq.prob_concentration_bestArm_fail_delta (E := E) (A := A) (R' := R') - (Q := Q) (κ := κ) (P := P) h hσ2 hs n δ hδ hδ1 + exact (h.prob_abs_empMean_bestAction_sub_actionMean_ge_le hσ2 hs hδ n).trans + (ENNReal.ofReal_le_ofReal (by nlinarith [hδ.le])) rw [h_swap] set f1 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, (armMean (bestArm ω) ω - uc (bestArm ω) s ω) @@ -738,7 +665,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) - (hs : ∀ a e, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {lo hi : ℝ} (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (t : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A t] ≤ (3 * K + 2) * (hi - lo) + 8 * √(σ2 * K * t * Real.log t) := by @@ -767,11 +694,6 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] have htpos : (0 : ℝ) < t := by positivity have _ht1 : (1 : ℝ) ≤ t := by exact_mod_cast Nat.pos_of_ne_zero ht have hδ : (0 : ℝ) < 1 / (t : ℝ) ^ 2 := by positivity - have hδ1 : 1 / (t : ℝ) ^ 2 < 1 := by - rw [div_lt_one (pow_pos htpos 2)] - have ht2_real : (2 : ℝ) ≤ t := Nat.ofNat_le_cast.mpr ht2 - calc (1 : ℝ) < 2 ^ 2 := by norm_num - _ ≤ (t : ℝ) ^ 2 := by gcongr -- First term: (hi-lo)*K + 2*(K+1)*(hi-lo)*t²*(1/t²) = (3K+2)*(hi-lo) have h_first : (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * ↑t ^ 2 * (1 / (↑t) ^ 2) = (3 * ↑K + 2) * (hi - lo) := by @@ -783,7 +705,7 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] ≤ (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * ↑t ^ 2 * (1 / (↑t) ^ 2) + 4 * √(2 * ↑σ2 * Real.log (1 / (1 / (↑t) ^ 2)) * ↑K * ↑t) := bayesRegret_le_of_delta (hK := hK) (E := E) (A := A) (R' := R') (Q := Q) - (κ := κ) (P := P) h hσ2 hs hm t (1 / (↑t) ^ 2) hδ hδ1 + (κ := κ) (P := P) h hσ2 hs hm t (1 / (↑t) ^ 2) hδ _ = (3 * ↑K + 2) * (hi - lo) + 8 * √(↑σ2 * ↑K * ↑t * Real.log ↑t) := by rw [h_first, h_log]; congr 1 rw [show (2 : ℝ) * ↑σ2 * (2 * Real.log ↑t) * ↑K * ↑t = From 382b0aa37181e990990be913a8cd52783da6d668 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 8 Apr 2026 12:14:15 +0100 Subject: [PATCH 095/155] Refactor TS.lean (in progress) --- LeanBandits/BanditAlgorithms/TS.lean | 93 ++++++++++------------------ 1 file changed, 32 insertions(+), 61 deletions(-) diff --git a/LeanBandits/BanditAlgorithms/TS.lean b/LeanBandits/BanditAlgorithms/TS.lean index 2107bcb2..afdc2e01 100644 --- a/LeanBandits/BanditAlgorithms/TS.lean +++ b/LeanBandits/BanditAlgorithms/TS.lean @@ -183,8 +183,6 @@ end TS end Bandits -open Bandits - namespace Learning.IsBayesAlgEnvSeq variable {K : ℕ} [Nonempty (Fin K)] @@ -201,44 +199,27 @@ lemma prob_abs_sumRewards_sub_pullCount_mul_actionMean_ge_le √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ |sumRewards A R' a t ω - pullCount A a t ω * actionMean κ E a ω|} ≤ ENNReal.ofReal (2 * K * (n - 1) * δ) := by - let s (e : 𝓔) := {ω | ∃ a, ∃ s < n, pullCount IT.action a s ω ≠ 0 ∧ - √(2 * pullCount IT.action a s ω * σ2 * Real.log (1 / δ)) ≤ - |sumRewards IT.action IT.reward a s ω - pullCount IT.action a s ω * (κ (e, a))[id]|} - have : {p : 𝓔 × (ℕ → (Fin K) × ℝ) | p.2 ∈ s p.1} = - ⋃ a, ⋃ t ∈ Finset.range n, - (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (pullCount IT.action a t p.2, - sumRewards IT.action IT.reward a t p.2, (κ (p.1, a))[id])) ⁻¹' - {t : ℕ × ℝ × ℝ | t.1 ≠ 0 ∧ - √(2 * t.1 * σ2 * Real.log (1 / δ)) ≤ |t.2.1 - t.1 * t.2.2|} := by - ext p - simp only [s, Set.mem_setOf_eq, Set.mem_iUnion, Set.mem_preimage, Finset.mem_range] - tauto - have (a : Fin K) : Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := - measurable_actionMean measurable_fst - calc P {ω | ∃ a, ∃ s < n, pullCount A a s ω ≠ 0 ∧ - √(2 * pullCount A a s ω * σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R' a s ω - pullCount A a s ω * actionMean κ E a ω|} - _ = P ((fun ω ↦ (E ω, trajectory A R' ω)) ⁻¹' {p | p.2 ∈ s p.1}) := by - congr 1 - _ = (P.map (fun ω ↦ (E ω, trajectory A R' ω))) {p | p.2 ∈ s p.1} := - (Measure.map_apply (h.measurable_E.prodMk - (measurable_trajectory h.measurable_A h.measurable_R)) (by measurability)).symm - _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) {p | p.2 ∈ s p.1} := by - rw [(compProd_map_condDistrib - (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable).symm] - _ = ∫⁻ e, (condDistrib (trajectory A R') E P e) (s e) ∂(P.map E) := by + have := h.measurable_E + have := h.measurable_A + have := h.measurable_R + let B e := {τ | ∃ a, ∃ t < n, pullCount IT.action a t τ ≠ 0 ∧ + √(2 * pullCount IT.action a t τ * σ2 * Real.log (1 / δ)) ≤ + |sumRewards IT.action IT.reward a t τ - pullCount IT.action a t τ * actionMean κ id a e|} + calc P ((fun ω ↦ (E ω, trajectory A R' ω)) ⁻¹' {(e, τ) : 𝓔 × (ℕ → Fin K × ℝ) | τ ∈ B e}) + _ = (P.map (fun ω ↦ (E ω, trajectory A R' ω))) {(e, τ) : 𝓔 × (ℕ → Fin K × ℝ) | τ ∈ B e} := + (Measure.map_apply (by fun_prop) (by measurability)).symm + _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) {(e, τ) : 𝓔 × _ | τ ∈ B e} := by + rw [← compProd_map_condDistrib (by fun_prop)] + _ = ∫⁻ e, (condDistrib (trajectory A R') E P e) (B e) ∂(P.map E) := by rw [Measure.compProd_apply (by measurability)] rfl - _ ≤ ∫⁻ _, ENNReal.ofReal (2 * K * (n - 1) * δ) ∂(P.map E) := by + _ ≤ ∫⁻ e, ENNReal.ofReal (2 * (Fintype.card (Fin K)) * (n - 1) * δ) ∂(P.map E) := by apply lintegral_mono_ae rw [h.hasLaw_env.map_eq] - filter_upwards [ae_IsAlgEnvSeq h] with e he - simpa [Fintype.card_fin, s, Kernel.sectR_apply] using - prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 - (fun a ↦ by simp only [Kernel.sectR_apply]; exact hs e a) he hδ + filter_upwards [h.ae_IsAlgEnvSeq] with e he + exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 (hs e) he hδ _ = ENNReal.ofReal (2 * K * (n - 1) * δ) := by - rw [lintegral_const, Measure.map_apply h.measurable_E .univ] - simp [measure_univ] + simp [lintegral_const, Measure.map_apply h.measurable_E] lemma prob_abs_sumRewards_bestAction_sub_pullCount_mul_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) @@ -251,23 +232,21 @@ lemma prob_abs_sumRewards_bestAction_sub_pullCount_mul_actionMean_ge_le ≤ ENNReal.ofReal (2 * (n - 1) * δ) := by by_cases hn : n = 0 · simp [hn] - let s (a : Fin K) (t : ℕ) (e : 𝓔) := {ω : ℕ → (Fin K) × ℝ | + let B (a : Fin K) (t : ℕ) (e : 𝓔) := {ω : ℕ → (Fin K) × ℝ | pullCount IT.action a t ω ≠ 0 ∧ √(2 * pullCount IT.action a t ω * σ2 * Real.log (1 / δ)) ≤ |sumRewards IT.action IT.reward a t ω - - pullCount IT.action a t ω * (κ (e, a))[id]|} - have (a : Fin K) : Measurable (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (κ (p.1, a))[id]) := - measurable_actionMean measurable_fst + pullCount IT.action a t ω * actionMean κ id a e|} have : {p : 𝓔 × (ℕ → (Fin K) × ℝ) | - p.2 ∈ ⋃ t ∈ Finset.range n, s (bestAction κ id p.1) t p.1} = + p.2 ∈ ⋃ t ∈ Finset.range n, B (bestAction κ id p.1) t p.1} = ⋃ a : Fin K, ((bestAction κ id ∘ Prod.fst) ⁻¹' {a} ∩ ⋃ t ∈ Finset.range n, (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (pullCount IT.action a t p.2, - sumRewards IT.action IT.reward a t p.2, (κ (p.1, a))[id])) ⁻¹' + sumRewards IT.action IT.reward a t p.2, actionMean κ Prod.fst a p)) ⁻¹' {u : ℕ × ℝ × ℝ | u.1 ≠ 0 ∧ √(2 * u.1 * σ2 * Real.log (1 / δ)) ≤ |u.2.1 - u.1 * u.2.2|}) := by ext p - simp only [s, Set.mem_setOf_eq, Set.mem_iUnion, Set.mem_inter_iff, + simp only [B, actionMean, Set.mem_setOf_eq, Set.mem_iUnion, Set.mem_inter_iff, Set.mem_preimage, Function.comp_apply, Set.mem_singleton_iff, Finset.mem_range] constructor · intro ⟨t, ht, hm⟩ @@ -279,34 +258,34 @@ lemma prob_abs_sumRewards_bestAction_sub_pullCount_mul_actionMean_ge_le |sumRewards A R' (bestAction κ E ω) t ω - pullCount A (bestAction κ E ω) t ω * actionMean κ E (bestAction κ E ω) ω|} _ = P ((fun ω ↦ (E ω, trajectory A R' ω)) ⁻¹' - {p | p.2 ∈ ⋃ t ∈ Finset.range n, s (bestAction κ id p.1) t p.1}) := by + {p | p.2 ∈ ⋃ t ∈ Finset.range n, B (bestAction κ id p.1) t p.1}) := by congr 1 ext ω simp only [Set.mem_setOf_eq, Finset.mem_range, Set.mem_preimage, - Set.mem_iUnion, s, actionMean, exists_prop] + Set.mem_iUnion, B, actionMean, exists_prop] rfl _ = (P.map (fun ω ↦ (E ω, trajectory A R' ω))) - {p | p.2 ∈ ⋃ t ∈ Finset.range n, s (bestAction κ id p.1) t p.1} := + {p | p.2 ∈ ⋃ t ∈ Finset.range n, B (bestAction κ id p.1) t p.1} := (Measure.map_apply (h.measurable_E.prodMk (measurable_trajectory h.measurable_A h.measurable_R)) (by measurability)).symm _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) - {p | p.2 ∈ ⋃ t ∈ Finset.range n, s (bestAction κ id p.1) t p.1} := by + {p | p.2 ∈ ⋃ t ∈ Finset.range n, B (bestAction κ id p.1) t p.1} := by rw [(compProd_map_condDistrib (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable).symm] _ = ∫⁻ e, (condDistrib (trajectory A R') E P e) - (⋃ t ∈ Finset.range n, s (bestAction κ id e) t e) ∂(P.map E) := by + (⋃ t ∈ Finset.range n, B (bestAction κ id e) t e) ∂(P.map E) := by rw [Measure.compProd_apply (by measurability)] rfl _ ≤ ∫⁻ _, ENNReal.ofReal (2 * (n - 1) * δ) ∂(P.map E) := by apply lintegral_mono_ae rw [h.hasLaw_env.map_eq] filter_upwards [ae_IsAlgEnvSeq h] with e he - have : {ω | ∃ t < n, ω ∈ s (bestAction κ id e) t e} = - ⋃ t ∈ Finset.range n, s (bestAction κ id e) t e := by + have : {ω | ∃ t < n, ω ∈ B (bestAction κ id e) t e} = + ⋃ t ∈ Finset.range n, B (bestAction κ id e) t e := by ext ω simp [Finset.mem_range] - simp only [← this, s, Set.mem_setOf_eq] - exact prob_abs_sumRewards_sub_pullCount_mul_ge_le hσ2 + simp only [← this, B, Set.mem_setOf_eq] + exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le hσ2 (by simp only [Kernel.sectR_apply]; exact hs e _) he hδ _ = ENNReal.ofReal (2 * (n - 1) * δ) := by rw [lintegral_const, Measure.map_apply h.measurable_E .univ] @@ -380,9 +359,7 @@ lemma prob_abs_empMean_bestAction_sub_actionMean_ge_le end Learning.IsBayesAlgEnvSeq -namespace Bandits - -section TSRegret +namespace Bandits.TS variable {K : ℕ} variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] @@ -392,8 +369,6 @@ variable (E : Ω → 𝓔) (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → ℝ variable (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] variable (P : Measure Ω) [IsProbabilityMeasure P] -namespace TS - lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (t : ℕ) : condDistrib (A (t + 1)) (IsAlgEnvSeq.hist A R' t) P @@ -714,8 +689,4 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 2)] ring -end TS - -end TSRegret - -end Bandits +end Bandits.TS From 8a729aa47349ee80b50705e7207c5a744a433f29 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Thu, 9 Apr 2026 08:45:38 +0200 Subject: [PATCH 096/155] fix --- LeanMachineLearning.lean | 2 ++ LeanMachineLearning/Bandit/SumRewards.lean | 2 +- LeanMachineLearning/BanditAlgorithms/TS.lean | 16 ++++++++++------ .../BanditAlgorithms/Uniform.lean | 8 ++++++-- LeanMachineLearning/ForMathlib/FullSupport.lean | 6 +++++- LeanMachineLearning/ForMathlib/WithDensity.lean | 9 +++++++-- .../SequentialLearning/AlgorithmDensity.lean | 10 +++++++--- .../SequentialLearning/BayesStationaryEnv.lean | 10 +++++++--- 8 files changed, 45 insertions(+), 18 deletions(-) diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index ea398ef1..6fd85fdf 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -12,6 +12,7 @@ public import LeanMachineLearning.BanditAlgorithms.UCB public import LeanMachineLearning.BanditAlgorithms.Uniform public import LeanMachineLearning.ForMathlib.CondDistrib public import LeanMachineLearning.ForMathlib.CondIndepFun +public import LeanMachineLearning.ForMathlib.FullSupport public import LeanMachineLearning.ForMathlib.HasCondDistrib public import LeanMachineLearning.ForMathlib.IndepFun public import LeanMachineLearning.ForMathlib.IndepInfinitePi @@ -22,6 +23,7 @@ public import LeanMachineLearning.ForMathlib.MeasurableArgMax public import LeanMachineLearning.ForMathlib.StandardBorel public import LeanMachineLearning.ForMathlib.SubGaussian public import LeanMachineLearning.ForMathlib.Traj +public import LeanMachineLearning.ForMathlib.WithDensity public import LeanMachineLearning.SequentialLearning.Algorithm public import LeanMachineLearning.SequentialLearning.AlgorithmDensity public import LeanMachineLearning.SequentialLearning.BayesStationaryEnv diff --git a/LeanMachineLearning/Bandit/SumRewards.lean b/LeanMachineLearning/Bandit/SumRewards.lean index 031d8c62..6f2b1b0d 100644 --- a/LeanMachineLearning/Bandit/SumRewards.lean +++ b/LeanMachineLearning/Bandit/SumRewards.lean @@ -414,7 +414,7 @@ private lemma exp_neg_sqrt_sq_div_le {σ2 : ℝ≥0} (hσ2 : 0 < σ2) {δ : ℝ} rw [Real.sq_sqrt (by positivity)] field_simp simp [Real.exp_log hδ] - · push_neg at hd + · push Not at hd have hl : Real.log (1 / δ) ≤ 0 := Real.log_nonpos (by positivity) (div_le_one_of_le₀ hd (hδ.le)) rw [Real.sqrt_eq_zero_of_nonpos (mul_nonpos_of_nonneg_of_nonpos (by positivity) hl)] simp [hd] diff --git a/LeanMachineLearning/BanditAlgorithms/TS.lean b/LeanMachineLearning/BanditAlgorithms/TS.lean index afdc2e01..99b6ea4f 100644 --- a/LeanMachineLearning/BanditAlgorithms/TS.lean +++ b/LeanMachineLearning/BanditAlgorithms/TS.lean @@ -3,12 +3,16 @@ 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.Bandit.SumRewards -import LeanBandits.BanditAlgorithms.Uniform -import LeanBandits.SequentialLearning.AlgorithmDensity +module + +public import LeanMachineLearning.Bandit.SumRewards +public import LeanMachineLearning.BanditAlgorithms.Uniform +public import LeanMachineLearning.SequentialLearning.AlgorithmDensity /-! # The Thompson Sampling Algorithm -/ +@[expose] public section + open MeasureTheory ProbabilityTheory Finset Learning open scoped NNReal @@ -309,7 +313,7 @@ private lemma abs_sumRewards_sub_pullCount_mul_ge {a : Fin K} {n : ℕ} {ω : Ω 2 * σ2 * Real.log (1 / δ) / pullCount A a n ω * pullCount A a n ω ^ 2 := by field_simp rw [this, Real.sqrt_mul (div_nonneg hc hk.le), Real.sqrt_sq hk.le] - · rw [Real.sqrt_eq_zero_of_nonpos (by push_neg at hc; nlinarith)] + · rw [Real.sqrt_eq_zero_of_nonpos (by push Not at hc; nlinarith)] exact mul_nonneg (Real.sqrt_nonneg _) hk.le _ ≤ |sumRewards A R' a n ω / pullCount A a n ω - μ| * pullCount A a n ω := mul_le_mul_of_nonneg_right h hk.le @@ -561,7 +565,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp have : Eδᶜ = {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ |empMean A R' a s ω - armMean a ω|} := by - ext ω; simp only [Eδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push_neg; rfl + ext ω; simp only [Eδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push Not; rfl rw [this] exact (h.prob_abs_empMean_sub_actionMean_ge_le hσ2 hs hδ n).trans (ENNReal.ofReal_le_ofReal (by nlinarith [hδ.le, Nat.cast_nonneg (α := ℝ) K])) @@ -598,7 +602,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp have : Fδᶜ = {ω | ∃ s < n, pullCount A (bestArm ω) s ω ≠ 0 ∧ √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A (bestArm ω) s ω : ℝ)) ≤ |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω|} := by - ext ω; simp only [Fδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push_neg; rfl + ext ω; simp only [Fδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push Not; rfl rw [this] exact (h.prob_abs_empMean_bestAction_sub_actionMean_ge_le hσ2 hs hδ n).trans (ENNReal.ofReal_le_ofReal (by nlinarith [hδ.le])) diff --git a/LeanMachineLearning/BanditAlgorithms/Uniform.lean b/LeanMachineLearning/BanditAlgorithms/Uniform.lean index f385a454..b47ab953 100644 --- a/LeanMachineLearning/BanditAlgorithms/Uniform.lean +++ b/LeanMachineLearning/BanditAlgorithms/Uniform.lean @@ -3,11 +3,15 @@ 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.FullSupport -import LeanBandits.SequentialLearning.Algorithm +module + +public import LeanMachineLearning.ForMathlib.FullSupport +public import LeanMachineLearning.SequentialLearning.Algorithm /-! # The Uniform Algorithm -/ +@[expose] public section + open MeasureTheory ProbabilityTheory Learning namespace Bandits diff --git a/LeanMachineLearning/ForMathlib/FullSupport.lean b/LeanMachineLearning/ForMathlib/FullSupport.lean index 51bd3bc7..aa0f07e3 100644 --- a/LeanMachineLearning/ForMathlib/FullSupport.lean +++ b/LeanMachineLearning/ForMathlib/FullSupport.lean @@ -3,7 +3,11 @@ 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 Mathlib.Probability.Kernel.Composition.MeasureCompProd +module + +public import Mathlib.Probability.Kernel.Composition.MeasureCompProd + +@[expose] public section open MeasureTheory ProbabilityTheory diff --git a/LeanMachineLearning/ForMathlib/WithDensity.lean b/LeanMachineLearning/ForMathlib/WithDensity.lean index b806f5aa..47ba5f6e 100644 --- a/LeanMachineLearning/ForMathlib/WithDensity.lean +++ b/LeanMachineLearning/ForMathlib/WithDensity.lean @@ -3,8 +3,11 @@ 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 Mathlib.Probability.Kernel.CompProdEqIff -import Mathlib.Probability.Kernel.Composition.MeasureComp +module + +public import Mathlib.Probability.Kernel.CompProdEqIff +public import Mathlib.Probability.Kernel.Composition.MeasureComp + /-! # Interactions of `withDensity` with `compProd`, `map`, and `swap` @@ -12,6 +15,8 @@ Lemmas for pushing `Measure.withDensity` and `Kernel.withDensity` through `compProd`, `MeasurableEquiv.map`, `Prod.swap`, and composition. -/ +@[expose] public section + open MeasureTheory ProbabilityTheory open scoped ENNReal diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean index 54cb6049..c7e9564c 100644 --- a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean @@ -3,9 +3,13 @@ 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.FullSupport -import LeanBandits.ForMathlib.WithDensity -import LeanBandits.SequentialLearning.BayesStationaryEnv +module + +public import LeanMachineLearning.ForMathlib.FullSupport +public import LeanMachineLearning.ForMathlib.WithDensity +public import LeanMachineLearning.SequentialLearning.BayesStationaryEnv + +@[expose] public section open MeasureTheory ProbabilityTheory Finset diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 16dff851..37708f8a 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -3,12 +3,16 @@ 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.Bandit.Regret -import LeanBandits.ForMathlib.MeasurableArgMax -import LeanBandits.SequentialLearning.StationaryEnv +module + +public import LeanMachineLearning.Bandit.Regret +public import LeanMachineLearning.ForMathlib.MeasurableArgMax +public import LeanMachineLearning.SequentialLearning.StationaryEnv /-! # Bayesian stationary environments -/ +@[expose] public section + open MeasureTheory ProbabilityTheory Finset namespace Learning From 636c66af1d22002604c676e96a3aa4e99806b432 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 13 Apr 2026 10:44:46 +0100 Subject: [PATCH 097/155] Refactor TS.lean (in progress) --- LeanMachineLearning/BanditAlgorithms/TS.lean | 141 ++++++++---------- .../ForMathlib/Measurable.lean | 14 -- .../BayesStationaryEnv.lean | 8 + .../SequentialLearning/FiniteActions.lean | 25 ++++ 4 files changed, 92 insertions(+), 96 deletions(-) diff --git a/LeanMachineLearning/BanditAlgorithms/TS.lean b/LeanMachineLearning/BanditAlgorithms/TS.lean index 99b6ea4f..f6b30696 100644 --- a/LeanMachineLearning/BanditAlgorithms/TS.lean +++ b/LeanMachineLearning/BanditAlgorithms/TS.lean @@ -70,6 +70,18 @@ lemma measurable_ucb [MeasurableSpace Ω] {a : Fin K} {n : ℕ} (hA : ∀ t, Mea (hR : ∀ t, Measurable (R' t)) : Measurable (ucb A R' l u σ2 δ a n) := Measurable.ite (by measurability) (by fun_prop) (by fun_prop) +@[fun_prop] +lemma measurable_uncurry_ucb_comp [MeasurableSpace Ω] (hA : ∀ t, Measurable (A t)) + (hR : ∀ t, Measurable (R' t)) {f : Ω → Fin K} (hf : Measurable f) {g : Ω → ℕ} + (hg : Measurable g) : Measurable (fun ω ↦ ucb A R' l u σ2 δ (f ω) (g ω) ω) := by + change Measurable ((fun aω ↦ ucb A R' l u σ2 δ aω.1 (g aω.2) aω.2) ∘ fun ω ↦ (f ω, ω)) + apply Measurable.comp _ (by fun_prop) + apply measurable_from_prod_countable_right + intro a + change Measurable ((fun tω ↦ ucb A R' l u σ2 δ a tω.1 tω.2) ∘ fun ω ↦ (g ω, ω)) + apply Measurable.comp _ (by fun_prop) + exact measurable_from_prod_countable_right (fun _ ↦ measurable_ucb hA hR) + lemma ucb_mem_Icc (h : l ≤ u) {a : Fin K} {n : ℕ} {ω : Ω} : ucb A R' l u σ2 δ a n ω ∈ Set.Icc l u := by unfold ucb @@ -234,66 +246,29 @@ lemma prob_abs_sumRewards_bestAction_sub_pullCount_mul_actionMean_ge_le |sumRewards A R' (bestAction κ E ω) t ω - pullCount A (bestAction κ E ω) t ω * actionMean κ E (bestAction κ E ω) ω|} ≤ ENNReal.ofReal (2 * (n - 1) * δ) := by - by_cases hn : n = 0 - · simp [hn] - let B (a : Fin K) (t : ℕ) (e : 𝓔) := {ω : ℕ → (Fin K) × ℝ | - pullCount IT.action a t ω ≠ 0 ∧ - √(2 * pullCount IT.action a t ω * σ2 * Real.log (1 / δ)) ≤ - |sumRewards IT.action IT.reward a t ω - - pullCount IT.action a t ω * actionMean κ id a e|} - have : {p : 𝓔 × (ℕ → (Fin K) × ℝ) | - p.2 ∈ ⋃ t ∈ Finset.range n, B (bestAction κ id p.1) t p.1} = - ⋃ a : Fin K, ((bestAction κ id ∘ Prod.fst) ⁻¹' {a} ∩ - ⋃ t ∈ Finset.range n, - (fun p : 𝓔 × (ℕ → (Fin K) × ℝ) ↦ (pullCount IT.action a t p.2, - sumRewards IT.action IT.reward a t p.2, actionMean κ Prod.fst a p)) ⁻¹' - {u : ℕ × ℝ × ℝ | u.1 ≠ 0 ∧ - √(2 * u.1 * σ2 * Real.log (1 / δ)) ≤ |u.2.1 - u.1 * u.2.2|}) := by - ext p - simp only [B, actionMean, Set.mem_setOf_eq, Set.mem_iUnion, Set.mem_inter_iff, - Set.mem_preimage, Function.comp_apply, Set.mem_singleton_iff, Finset.mem_range] - constructor - · intro ⟨t, ht, hm⟩ - exact ⟨bestAction κ id p.1, rfl, t, ht, hm⟩ - · rintro ⟨a, ha, t, ht, hm⟩ - exact ⟨t, ht, ha ▸ hm⟩ - calc P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ - √(2 * pullCount A (bestAction κ E ω) t ω * σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R' (bestAction κ E ω) t ω - - pullCount A (bestAction κ E ω) t ω * actionMean κ E (bestAction κ E ω) ω|} - _ = P ((fun ω ↦ (E ω, trajectory A R' ω)) ⁻¹' - {p | p.2 ∈ ⋃ t ∈ Finset.range n, B (bestAction κ id p.1) t p.1}) := by - congr 1 - ext ω - simp only [Set.mem_setOf_eq, Finset.mem_range, Set.mem_preimage, - Set.mem_iUnion, B, actionMean, exists_prop] - rfl - _ = (P.map (fun ω ↦ (E ω, trajectory A R' ω))) - {p | p.2 ∈ ⋃ t ∈ Finset.range n, B (bestAction κ id p.1) t p.1} := - (Measure.map_apply (h.measurable_E.prodMk - (measurable_trajectory h.measurable_A h.measurable_R)) (by measurability)).symm - _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) - {p | p.2 ∈ ⋃ t ∈ Finset.range n, B (bestAction κ id p.1) t p.1} := by - rw [(compProd_map_condDistrib - (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable).symm] - _ = ∫⁻ e, (condDistrib (trajectory A R') E P e) - (⋃ t ∈ Finset.range n, B (bestAction κ id e) t e) ∂(P.map E) := by + have := h.measurable_E + have := h.measurable_A + have := h.measurable_R + let B e := {τ | ∃ t < n, pullCount IT.action (bestAction κ id e) t τ ≠ 0 ∧ + √(2 * pullCount IT.action (bestAction κ id e) t τ * σ2 * Real.log (1 / δ)) ≤ + |sumRewards IT.action IT.reward (bestAction κ id e) t τ - + pullCount IT.action (bestAction κ id e) t τ * actionMean κ id (bestAction κ id e) e|} + calc P ((fun ω ↦ (E ω, trajectory A R' ω)) ⁻¹' {(e, τ) : 𝓔 × (ℕ → Fin K × ℝ) | τ ∈ B e}) + _ = (P.map (fun ω ↦ (E ω, trajectory A R' ω))) {(e, τ) : 𝓔 × (ℕ → Fin K × ℝ) | τ ∈ B e} := + (Measure.map_apply (by fun_prop) (by measurability)).symm + _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) {(e, τ) : 𝓔 × _ | τ ∈ B e} := by + rw [← compProd_map_condDistrib (by fun_prop)] + _ = ∫⁻ e, (condDistrib (trajectory A R') E P e) (B e) ∂(P.map E) := by rw [Measure.compProd_apply (by measurability)] rfl - _ ≤ ∫⁻ _, ENNReal.ofReal (2 * (n - 1) * δ) ∂(P.map E) := by + _ ≤ ∫⁻ e, ENNReal.ofReal (2 * (n - 1) * δ) ∂(P.map E) := by apply lintegral_mono_ae rw [h.hasLaw_env.map_eq] - filter_upwards [ae_IsAlgEnvSeq h] with e he - have : {ω | ∃ t < n, ω ∈ B (bestAction κ id e) t e} = - ⋃ t ∈ Finset.range n, B (bestAction κ id e) t e := by - ext ω - simp [Finset.mem_range] - simp only [← this, B, Set.mem_setOf_eq] - exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le hσ2 - (by simp only [Kernel.sectR_apply]; exact hs e _) he hδ + filter_upwards [h.ae_IsAlgEnvSeq] with e he + exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le (ν := κ.sectR e) + hσ2 (hs e (bestAction κ id e)) he hδ _ = ENNReal.ofReal (2 * (n - 1) * δ) := by - rw [lintegral_const, Measure.map_apply h.measurable_E .univ] - simp [measure_univ] + simp [lintegral_const, Measure.map_apply h.measurable_E] omit [Nonempty (Fin K)] [MeasurableSpace 𝓔] [MeasurableSpace Ω] in private lemma abs_sumRewards_sub_pullCount_mul_ge {a : Fin K} {n : ℕ} {ω : Ω} @@ -302,26 +277,24 @@ private lemma abs_sumRewards_sub_pullCount_mul_ge {a : Fin K} {n : ℕ} {ω : Ω |empMean A R' a n ω - μ|) : √(2 * pullCount A a n ω * σ2 * Real.log (1 / δ)) ≤ |sumRewards A R' a n ω - pullCount A a n ω * μ| := by - have hk : (0 : ℝ) < pullCount A a n ω := - Nat.cast_pos.mpr (Nat.pos_of_ne_zero hpc) - simp only [empMean] at h - calc √(2 * pullCount A a n ω * σ2 * Real.log (1 / δ)) - _ ≤ √(2 * σ2 * Real.log (1 / δ) / pullCount A a n ω) * pullCount A a n ω := by - by_cases hc : 0 ≤ 2 * σ2 * Real.log (1 / δ) - · apply le_of_eq - have : 2 * pullCount A a n ω * σ2 * Real.log (1 / δ) = - 2 * σ2 * Real.log (1 / δ) / pullCount A a n ω * pullCount A a n ω ^ 2 := by - field_simp - rw [this, Real.sqrt_mul (div_nonneg hc hk.le), Real.sqrt_sq hk.le] - · rw [Real.sqrt_eq_zero_of_nonpos (by push Not at hc; nlinarith)] - exact mul_nonneg (Real.sqrt_nonneg _) hk.le - _ ≤ |sumRewards A R' a n ω / pullCount A a n ω - μ| * pullCount A a n ω := - mul_le_mul_of_nonneg_right h hk.le - _ = |sumRewards A R' a n ω - pullCount A a n ω * μ| := by - have h_div : sumRewards A R' a n ω / ↑(pullCount A a n ω) - μ = - (sumRewards A R' a n ω - pullCount A a n ω * μ) / pullCount A a n ω := by - field_simp - rw [h_div, abs_div, abs_of_pos hk, div_mul_cancel₀ _ (ne_of_gt hk)] + have hk : (0 : ℝ) < pullCount A a n ω := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hpc) + by_cases hc : 0 ≤ 2 * σ2 * Real.log (1 / δ) + · calc + _ = √(2 * σ2 * Real.log (1 / δ) / pullCount A a n ω) * pullCount A a n ω := by + have : 2 * pullCount A a n ω * σ2 * Real.log (1 / δ) = + 2 * σ2 * Real.log (1 / δ) / pullCount A a n ω * pullCount A a n ω ^ 2 := by + field_simp + rw [this, Real.sqrt_mul (div_nonneg hc hk.le), Real.sqrt_sq hk.le] + _ ≤ |sumRewards A R' a n ω / pullCount A a n ω - μ| * pullCount A a n ω := + mul_le_mul_of_nonneg_right h hk.le + _ = |sumRewards A R' a n ω - pullCount A a n ω * μ| := by + have : sumRewards A R' a n ω / ↑(pullCount A a n ω) - μ = + (sumRewards A R' a n ω - pullCount A a n ω * μ) / pullCount A a n ω := by + field_simp + rw [this, abs_div, abs_of_pos hk, div_mul_cancel₀ _ (ne_of_gt hk)] + · calc + _ = 0 := Real.sqrt_eq_zero_of_nonpos (by push Not at hc; nlinarith) + _ ≤ |sumRewards A R' a n ω - pullCount A a n ω * μ| := abs_nonneg _ lemma prob_abs_empMean_sub_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} @@ -331,7 +304,7 @@ lemma prob_abs_empMean_sub_actionMean_ge_le √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ |empMean A R' a t ω - actionMean κ E a ω|} ≤ ENNReal.ofReal (2 * K * (n - 1) * δ) := - calc _ + calc _ ≤ P {ω | ∃ a, ∃ t < n, pullCount A a t ω ≠ 0 ∧ √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ |sumRewards A R' a t ω - pullCount A a t ω * actionMean κ E a ω|} := by @@ -348,7 +321,7 @@ lemma prob_abs_empMean_bestAction_sub_actionMean_ge_le √(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω : ℝ)) ≤ |empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω|} ≤ ENNReal.ofReal (2 * (n - 1) * δ) := - calc _ + calc _ ≤ P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ √(2 * pullCount A (bestAction κ E ω) t ω * σ2 * Real.log (1 / δ)) ≤ |sumRewards A R' (bestAction κ E ω) t ω - @@ -450,15 +423,18 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (armMean (bestArm ω) ω - uc (bestArm ω) s ω)) P := by apply Integrable.of_bound (C := ↑n * (u - l)) · exact (Finset.measurable_fun_sum _ fun s _ ↦ - (measurable_apply_fin hm_arm hm_best).sub - (measurable_apply_fin (fun a ↦ hm_ucb a s) hm_best)).aestronglyMeasurable + (IsBayesAlgEnvSeq.measurable_uncurry_actionMean_comp h.measurable_E hm_best).sub + (measurable_uncurry_ucb_comp h.measurable_A h.measurable_R hm_best + measurable_const)).aestronglyMeasurable · filter_upwards with ω; rw [Real.norm_eq_abs]; exact h_first_bound ω have h_int_sum2 : Integrable (fun ω ↦ ∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω)) P := by apply Integrable.of_bound (C := ↑n * (u - l)) · exact (Finset.measurable_fun_sum _ fun s _ ↦ - (measurable_apply_fin (fun a ↦ hm_ucb a s) (h.measurable_A s)).sub - (measurable_apply_fin hm_arm (h.measurable_A s))).aestronglyMeasurable + (measurable_uncurry_ucb_comp h.measurable_A h.measurable_R + (h.measurable_A s) measurable_const).sub + (IsBayesAlgEnvSeq.measurable_uncurry_actionMean_comp h.measurable_E (h.measurable_A s)) + ).aestronglyMeasurable · filter_upwards with ω; rw [Real.norm_eq_abs]; exact h_second_bound ω have h_swap : P[IsBayesAlgEnvSeq.regret κ E A n] = @@ -468,7 +444,8 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp (uc (A s ω) s ω - armMean (A s ω) ω)] := by have h_int_ucb : ∀ s {f : Ω → Fin K}, Measurable f → Integrable (fun ω ↦ uc (f ω) s ω) P := fun s {_} hf ↦ - ⟨(measurable_apply_fin (fun a ↦ hm_ucb a s) hf).aestronglyMeasurable, + ⟨(measurable_uncurry_ucb_comp + h.measurable_A h.measurable_R hf measurable_const).aestronglyMeasurable, HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ by rw [Real.norm_eq_abs] exact abs_le_max_abs_abs (ucb_mem_Icc hlo).1 diff --git a/LeanMachineLearning/ForMathlib/Measurable.lean b/LeanMachineLearning/ForMathlib/Measurable.lean index 1768dedc..70a2f36e 100644 --- a/LeanMachineLearning/ForMathlib/Measurable.lean +++ b/LeanMachineLearning/ForMathlib/Measurable.lean @@ -60,18 +60,4 @@ lemma measurable_sum_Icc_of_le {f : ℕ → α → ℝ} {g : α → ℕ} {n : refine Measurable.ite ?_ (by fun_prop) (by fun_prop) exact (measurableSet_singleton _).preimage (by fun_prop) -lemma measurable_apply_fin {α' : Type*} [MeasurableSpace α'] [Finite α'] - [MeasurableSingletonClass α'] - {f : α' → α → ℝ} {g : α → α'} - (hf : ∀ a, Measurable (f a)) (hg : Measurable g) : - Measurable (fun ω ↦ f (g ω) ω) := by - classical - have := Fintype.ofFinite α' - have : (fun ω ↦ f (g ω) ω) = fun ω ↦ ∑ a : α', if g ω = a then f a ω else 0 := by - ext ω; simp [Finset.sum_ite_eq] - rw [this] - apply Finset.measurable_fun_sum - intro a _ - exact Measurable.ite (hg (measurableSet_singleton a)) (hf a) measurable_const - end MeasureTheory diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 37708f8a..77f6dbb0 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -58,6 +58,14 @@ lemma measurable_actionMean {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} {a Measurable (actionMean κ E a) := stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop) +@[fun_prop] +lemma measurable_uncurry_actionMean_comp [Countable α] [MeasurableSingletonClass α] + {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} (hE : Measurable E) {f : Ω → α} (hf : Measurable f) : + Measurable (fun ω ↦ actionMean κ E (f ω) ω) := by + change Measurable ((fun aω ↦ actionMean κ E aω.1 aω.2) ∘ fun ω ↦ (f ω, ω)) + apply Measurable.comp _ (by fun_prop) + exact measurable_from_prod_countable_right (fun _ ↦ measurable_actionMean hE) + noncomputable def bestAction [Nonempty α] [Fintype α] [Encodable α] [MeasurableSingletonClass α] (κ : Kernel (𝓔 × α) ℝ) (E : Ω → 𝓔) (ω : Ω) : α := diff --git a/LeanMachineLearning/SequentialLearning/FiniteActions.lean b/LeanMachineLearning/SequentialLearning/FiniteActions.lean index f4258cb4..9734d5ad 100644 --- a/LeanMachineLearning/SequentialLearning/FiniteActions.lean +++ b/LeanMachineLearning/SequentialLearning/FiniteActions.lean @@ -191,6 +191,18 @@ lemma measurable_uncurry_pullCount [MeasurableEq α] exact measurableSet_eq_fun (by fun_prop) (by fun_prop) fun_prop +@[fun_prop] +lemma measurable_uncurry_pullCount_comp [Countable α] [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) {f : Ω → α} (hf : Measurable f) {g : Ω → ℕ} (hg : Measurable g) : + Measurable (fun ω ↦ pullCount A (f ω) (g ω) ω) := by + change Measurable ((fun aω ↦ pullCount A aω.1 (g aω.2) aω.2) ∘ fun ω ↦ (f ω, ω)) + apply Measurable.comp _ (by fun_prop) + apply measurable_from_prod_countable_right + intro a + change Measurable ((fun tω ↦ pullCount A a tω.1 tω.2) ∘ fun ω ↦ (g ω, ω)) + apply Measurable.comp _ (by fun_prop) + exact measurable_from_prod_countable_right (fun t ↦ measurable_pullCount hA a t) + @[fun_prop] lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : Measurable (fun h : Iic n → α × R ↦ pullCount' n h a) := by @@ -866,6 +878,19 @@ lemma measurable_sumRewards [MeasurableSingletonClass α] {R' : ℕ → Ω → exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop +@[fun_prop] +lemma measurable_uncurry_sumRewards_comp [Countable α] [MeasurableSingletonClass α] + {R' : ℕ → Ω → ℝ} (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) {f : Ω → α} + (hf : Measurable f) {g : Ω → ℕ} (hg : Measurable g) : + Measurable (fun ω ↦ sumRewards A R' (f ω) (g ω) ω) := by + change Measurable ((fun aω ↦ sumRewards A R' aω.1 (g aω.2) aω.2) ∘ fun ω ↦ (f ω, ω)) + apply Measurable.comp _ (by fun_prop) + apply measurable_from_prod_countable_right + intro a + change Measurable ((fun tω ↦ sumRewards A R' a tω.1 tω.2) ∘ fun ω ↦ (g ω, ω)) + apply Measurable.comp _ (by fun_prop) + exact measurable_from_prod_countable_right (fun t ↦ measurable_sumRewards hA hR' a t) + @[fun_prop] lemma measurable_empMean [MeasurableSingletonClass α] {R' : ℕ → Ω → ℝ} (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (n : ℕ) : From 8dc181ceb29b0bedd8a47666b579d3049fc4b178 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 13 Apr 2026 14:24:51 +0100 Subject: [PATCH 098/155] Refactor TS.lean (in progress) --- LeanMachineLearning/BanditAlgorithms/TS.lean | 148 +----------------- .../BayesStationaryEnv.lean | 97 +++++++++++- 2 files changed, 98 insertions(+), 147 deletions(-) diff --git a/LeanMachineLearning/BanditAlgorithms/TS.lean b/LeanMachineLearning/BanditAlgorithms/TS.lean index f6b30696..fae5595d 100644 --- a/LeanMachineLearning/BanditAlgorithms/TS.lean +++ b/LeanMachineLearning/BanditAlgorithms/TS.lean @@ -5,7 +5,6 @@ Authors: Rémy Degenne, Paulo Rauber -/ module -public import LeanMachineLearning.Bandit.SumRewards public import LeanMachineLearning.BanditAlgorithms.Uniform public import LeanMachineLearning.SequentialLearning.AlgorithmDensity @@ -195,149 +194,6 @@ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, end UCB -end TS - -end Bandits - -namespace Learning.IsBayesAlgEnvSeq - -variable {K : ℕ} [Nonempty (Fin K)] -variable {𝓔 Ω : Type*} [MeasurableSpace 𝓔] [MeasurableSpace Ω] -variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] {alg : Algorithm (Fin K) ℝ} -variable {E : Ω → 𝓔} {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → ℝ} -variable {P : Measure Ω} [IsProbabilityMeasure P] - -lemma prob_abs_sumRewards_sub_pullCount_mul_actionMean_ge_le - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} - (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : - P {ω | ∃ a, ∃ t < n, pullCount A a t ω ≠ 0 ∧ - √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R' a t ω - pullCount A a t ω * actionMean κ E a ω|} - ≤ ENNReal.ofReal (2 * K * (n - 1) * δ) := by - have := h.measurable_E - have := h.measurable_A - have := h.measurable_R - let B e := {τ | ∃ a, ∃ t < n, pullCount IT.action a t τ ≠ 0 ∧ - √(2 * pullCount IT.action a t τ * σ2 * Real.log (1 / δ)) ≤ - |sumRewards IT.action IT.reward a t τ - pullCount IT.action a t τ * actionMean κ id a e|} - calc P ((fun ω ↦ (E ω, trajectory A R' ω)) ⁻¹' {(e, τ) : 𝓔 × (ℕ → Fin K × ℝ) | τ ∈ B e}) - _ = (P.map (fun ω ↦ (E ω, trajectory A R' ω))) {(e, τ) : 𝓔 × (ℕ → Fin K × ℝ) | τ ∈ B e} := - (Measure.map_apply (by fun_prop) (by measurability)).symm - _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) {(e, τ) : 𝓔 × _ | τ ∈ B e} := by - rw [← compProd_map_condDistrib (by fun_prop)] - _ = ∫⁻ e, (condDistrib (trajectory A R') E P e) (B e) ∂(P.map E) := by - rw [Measure.compProd_apply (by measurability)] - rfl - _ ≤ ∫⁻ e, ENNReal.ofReal (2 * (Fintype.card (Fin K)) * (n - 1) * δ) ∂(P.map E) := by - apply lintegral_mono_ae - rw [h.hasLaw_env.map_eq] - filter_upwards [h.ae_IsAlgEnvSeq] with e he - exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 (hs e) he hδ - _ = ENNReal.ofReal (2 * K * (n - 1) * δ) := by - simp [lintegral_const, Measure.map_apply h.measurable_E] - -lemma prob_abs_sumRewards_bestAction_sub_pullCount_mul_actionMean_ge_le - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) - (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : - P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ - √(2 * pullCount A (bestAction κ E ω) t ω * σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R' (bestAction κ E ω) t ω - - pullCount A (bestAction κ E ω) t ω * actionMean κ E (bestAction κ E ω) ω|} - ≤ ENNReal.ofReal (2 * (n - 1) * δ) := by - have := h.measurable_E - have := h.measurable_A - have := h.measurable_R - let B e := {τ | ∃ t < n, pullCount IT.action (bestAction κ id e) t τ ≠ 0 ∧ - √(2 * pullCount IT.action (bestAction κ id e) t τ * σ2 * Real.log (1 / δ)) ≤ - |sumRewards IT.action IT.reward (bestAction κ id e) t τ - - pullCount IT.action (bestAction κ id e) t τ * actionMean κ id (bestAction κ id e) e|} - calc P ((fun ω ↦ (E ω, trajectory A R' ω)) ⁻¹' {(e, τ) : 𝓔 × (ℕ → Fin K × ℝ) | τ ∈ B e}) - _ = (P.map (fun ω ↦ (E ω, trajectory A R' ω))) {(e, τ) : 𝓔 × (ℕ → Fin K × ℝ) | τ ∈ B e} := - (Measure.map_apply (by fun_prop) (by measurability)).symm - _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) {(e, τ) : 𝓔 × _ | τ ∈ B e} := by - rw [← compProd_map_condDistrib (by fun_prop)] - _ = ∫⁻ e, (condDistrib (trajectory A R') E P e) (B e) ∂(P.map E) := by - rw [Measure.compProd_apply (by measurability)] - rfl - _ ≤ ∫⁻ e, ENNReal.ofReal (2 * (n - 1) * δ) ∂(P.map E) := by - apply lintegral_mono_ae - rw [h.hasLaw_env.map_eq] - filter_upwards [h.ae_IsAlgEnvSeq] with e he - exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le (ν := κ.sectR e) - hσ2 (hs e (bestAction κ id e)) he hδ - _ = ENNReal.ofReal (2 * (n - 1) * δ) := by - simp [lintegral_const, Measure.map_apply h.measurable_E] - -omit [Nonempty (Fin K)] [MeasurableSpace 𝓔] [MeasurableSpace Ω] in -private lemma abs_sumRewards_sub_pullCount_mul_ge {a : Fin K} {n : ℕ} {ω : Ω} - {μ σ2 δ : ℝ} (hpc : pullCount A a n ω ≠ 0) - (h : √(2 * σ2 * Real.log (1 / δ) / pullCount A a n ω) ≤ - |empMean A R' a n ω - μ|) : - √(2 * pullCount A a n ω * σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R' a n ω - pullCount A a n ω * μ| := by - have hk : (0 : ℝ) < pullCount A a n ω := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hpc) - by_cases hc : 0 ≤ 2 * σ2 * Real.log (1 / δ) - · calc - _ = √(2 * σ2 * Real.log (1 / δ) / pullCount A a n ω) * pullCount A a n ω := by - have : 2 * pullCount A a n ω * σ2 * Real.log (1 / δ) = - 2 * σ2 * Real.log (1 / δ) / pullCount A a n ω * pullCount A a n ω ^ 2 := by - field_simp - rw [this, Real.sqrt_mul (div_nonneg hc hk.le), Real.sqrt_sq hk.le] - _ ≤ |sumRewards A R' a n ω / pullCount A a n ω - μ| * pullCount A a n ω := - mul_le_mul_of_nonneg_right h hk.le - _ = |sumRewards A R' a n ω - pullCount A a n ω * μ| := by - have : sumRewards A R' a n ω / ↑(pullCount A a n ω) - μ = - (sumRewards A R' a n ω - pullCount A a n ω * μ) / pullCount A a n ω := by - field_simp - rw [this, abs_div, abs_of_pos hk, div_mul_cancel₀ _ (ne_of_gt hk)] - · calc - _ = 0 := Real.sqrt_eq_zero_of_nonpos (by push Not at hc; nlinarith) - _ ≤ |sumRewards A R' a n ω - pullCount A a n ω * μ| := abs_nonneg _ - -lemma prob_abs_empMean_sub_actionMean_ge_le - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} - (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : - P {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ - |empMean A R' a t ω - actionMean κ E a ω|} - ≤ ENNReal.ofReal (2 * K * (n - 1) * δ) := - calc - _ ≤ P {ω | ∃ a, ∃ t < n, pullCount A a t ω ≠ 0 ∧ - √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R' a t ω - pullCount A a t ω * actionMean κ E a ω|} := by - apply measure_mono - intro ω ⟨t, ht, a, hpc, hle⟩ - exact ⟨a, t, ht, hpc, abs_sumRewards_sub_pullCount_mul_ge hpc hle⟩ - _ ≤ _ := h.prob_abs_sumRewards_sub_pullCount_mul_actionMean_ge_le hσ2 hs hδ n - -lemma prob_abs_empMean_bestAction_sub_actionMean_ge_le - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) - (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : - P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω : ℝ)) ≤ - |empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω|} - ≤ ENNReal.ofReal (2 * (n - 1) * δ) := - calc - _ ≤ P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ - √(2 * pullCount A (bestAction κ E ω) t ω * σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R' (bestAction κ E ω) t ω - - pullCount A (bestAction κ E ω) t ω * - actionMean κ E (bestAction κ E ω) ω|} := by - apply measure_mono - intro ω ⟨t, ht, hpc, hle⟩ - exact ⟨t, ht, hpc, abs_sumRewards_sub_pullCount_mul_ge hpc hle⟩ - _ ≤ _ := - h.prob_abs_sumRewards_bestAction_sub_pullCount_mul_actionMean_ge_le - hσ2 hs hδ n - -end Learning.IsBayesAlgEnvSeq - -namespace Bandits.TS - variable {K : ℕ} variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] variable (hK : 0 < K) @@ -670,4 +526,6 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 2)] ring -end Bandits.TS +end TS + +end Bandits diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 77f6dbb0..dc9ff0b6 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -5,15 +5,15 @@ Authors: Rémy Degenne, Paulo Rauber -/ module -public import LeanMachineLearning.Bandit.Regret +public import LeanMachineLearning.Bandit.SumRewards public import LeanMachineLearning.ForMathlib.MeasurableArgMax -public import LeanMachineLearning.SequentialLearning.StationaryEnv /-! # Bayesian stationary environments -/ @[expose] public section open MeasureTheory ProbabilityTheory Finset +open scoped ENNReal NNReal namespace Learning @@ -232,6 +232,99 @@ lemma ae_IsAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P end CondDistribIsAlgEnvSeq +section HasSubgaussianMGF + +variable {K : ℕ} [Nonempty (Fin K)] +variable {κ : Kernel (𝓔 × Fin K) ℝ} {alg : Algorithm (Fin K) ℝ} +variable {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → ℝ} +variable [IsProbabilityMeasure P] + +private lemma sqrt_two_mul_le_abs_sub_of_sqrt_div_le {s μ σ L : ℝ} {k : ℕ} (hk : k ≠ 0) + (h : √(2 * σ * L / k) ≤ |s / k - μ|) : √(2 * k * σ * L) ≤ |s - k * μ| := by + have hkp : (0 : ℝ) < k := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hk) + have hks : s / (k : ℝ) - μ = (s - k * μ) / k := by field_simp + by_cases hc : 0 ≤ 2 * σ * L + · have hkc : (2 * k * σ * L : ℝ) = 2 * σ * L / k * k ^ 2 := by field_simp + calc √(2 * k * σ * L) + _ = √(2 * σ * L / k) * k := by + rw [hkc, Real.sqrt_mul (div_nonneg hc hkp.le), Real.sqrt_sq hkp.le] + _ ≤ |s / k - μ| * k := mul_le_mul_of_nonneg_right h hkp.le + _ = |s - k * μ| := by rw [hks, abs_div, abs_of_pos hkp, div_mul_cancel₀ _ (ne_of_gt hkp)] + · push Not at hc + calc √(2 * k * σ * L) + _ = 0 := Real.sqrt_eq_zero_of_nonpos (by nlinarith) + _ ≤ |s - k * μ| := abs_nonneg _ + +lemma prob_abs_empMean_sub_actionMean_ge_le [IsMarkovKernel κ] + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} + (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : + P {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ + √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ + |empMean A R' a t ω - actionMean κ E a ω|} + ≤ ENNReal.ofReal (2 * K * (n - 1) * δ) := by + have := h.measurable_E + have := h.measurable_A + have := h.measurable_R + let S : Set (𝓔 × (ℕ → Fin K × ℝ)) := {(e, τ) | ∃ a, ∃ t < n, pullCount IT.action a t τ ≠ 0 ∧ + √(2 * pullCount IT.action a t τ * σ2 * Real.log (1 / δ)) ≤ + |sumRewards IT.action IT.reward a t τ - pullCount IT.action a t τ * actionMean κ id a e|} + calc _ + _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by + rw [Measure.map_apply (by fun_prop) (by measurability)] + apply measure_mono + intro ω ⟨t, ht, a, hpc, hle⟩ + exact ⟨a, t, ht, hpc, + sqrt_two_mul_le_abs_sub_of_sqrt_div_le hpc (by simpa [empMean] using hle)⟩ + _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by + rw [← compProd_map_condDistrib (by fun_prop)] + _ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := + Measure.compProd_apply (by measurability) + _ ≤ ∫⁻ _, ENNReal.ofReal (2 * K * (n - 1) * δ) ∂(P.map E) := by + apply lintegral_mono_ae + rw [h.hasLaw_env.map_eq] + filter_upwards [h.ae_IsAlgEnvSeq] with e he + convert Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 (hs e) he hδ + exact (Fintype.card_fin K).symm + _ = _ := by simp [Measure.map_apply h.measurable_E] + +lemma prob_abs_empMean_bestAction_sub_actionMean_ge_le [IsMarkovKernel κ] + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : + P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ + √(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω : ℝ)) ≤ + |empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω|} + ≤ ENNReal.ofReal (2 * (n - 1) * δ) := by + have := h.measurable_E + have := h.measurable_A + have := h.measurable_R + let S : Set (𝓔 × (ℕ → Fin K × ℝ)) := + {(e, τ) | ∃ t < n, pullCount IT.action (bestAction κ id e) t τ ≠ 0 ∧ + √(2 * pullCount IT.action (bestAction κ id e) t τ * σ2 * Real.log (1 / δ)) ≤ + |sumRewards IT.action IT.reward (bestAction κ id e) t τ - + pullCount IT.action (bestAction κ id e) t τ * actionMean κ id (bestAction κ id e) e|} + calc _ + _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by + rw [Measure.map_apply (by fun_prop) (by measurability)] + apply measure_mono + intro ω ⟨t, ht, hpc, hle⟩ + exact ⟨t, ht, hpc, + sqrt_two_mul_le_abs_sub_of_sqrt_div_le hpc (by simpa [empMean] using hle)⟩ + _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by + rw [← compProd_map_condDistrib (by fun_prop)] + _ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := + Measure.compProd_apply (by measurability) + _ ≤ ∫⁻ _, ENNReal.ofReal (2 * (n - 1) * δ) ∂(P.map E) := by + apply lintegral_mono_ae + rw [h.hasLaw_env.map_eq] + filter_upwards [h.ae_IsAlgEnvSeq] with e he + exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le (ν := κ.sectR e) hσ2 + (hs e (bestAction κ id e)) he hδ + _ = _ := by simp [Measure.map_apply h.measurable_E] + +end HasSubgaussianMGF + end IsBayesAlgEnvSeq section IsAlgEnvSeq From 5bd9f642dd79cc09f41ad0175506f6b52c79bada Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 14 Apr 2026 10:48:05 +0100 Subject: [PATCH 099/155] Refactor TS.lean (in progress) --- .../BayesStationaryEnv.lean | 84 +++++++++---------- 1 file changed, 40 insertions(+), 44 deletions(-) diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index dc9ff0b6..afa9870b 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -234,94 +234,90 @@ end CondDistribIsAlgEnvSeq section HasSubgaussianMGF +private lemma sqrt_two_mul_le {k : ℕ} (hk : k ≠ 0) {s μ σ l : ℝ} + (h : √(2 * σ * l / k) ≤ |s / k - μ|) : √(2 * k * σ * l) ≤ |s - k * μ| := by + have hkp : (0 : ℝ) < k := by positivity + calc √(2 * k * σ * l) + _ = √(2 * σ * l / k * k ^ 2) := by + field_simp + _ = √(2 * σ * l / k) * k := by + rw [Real.sqrt_mul' _ (sq_nonneg _), Real.sqrt_sq hkp.le] + _ ≤ |s / k - μ| * k := by + nlinarith + _ = |s - k * μ| := by + field_simp + grind + variable {K : ℕ} [Nonempty (Fin K)] -variable {κ : Kernel (𝓔 × Fin K) ℝ} {alg : Algorithm (Fin K) ℝ} +variable {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] {alg : Algorithm (Fin K) ℝ} variable {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → ℝ} variable [IsProbabilityMeasure P] -private lemma sqrt_two_mul_le_abs_sub_of_sqrt_div_le {s μ σ L : ℝ} {k : ℕ} (hk : k ≠ 0) - (h : √(2 * σ * L / k) ≤ |s / k - μ|) : √(2 * k * σ * L) ≤ |s - k * μ| := by - have hkp : (0 : ℝ) < k := Nat.cast_pos.mpr (Nat.pos_of_ne_zero hk) - have hks : s / (k : ℝ) - μ = (s - k * μ) / k := by field_simp - by_cases hc : 0 ≤ 2 * σ * L - · have hkc : (2 * k * σ * L : ℝ) = 2 * σ * L / k * k ^ 2 := by field_simp - calc √(2 * k * σ * L) - _ = √(2 * σ * L / k) * k := by - rw [hkc, Real.sqrt_mul (div_nonneg hc hkp.le), Real.sqrt_sq hkp.le] - _ ≤ |s / k - μ| * k := mul_le_mul_of_nonneg_right h hkp.le - _ = |s - k * μ| := by rw [hks, abs_div, abs_of_pos hkp, div_mul_cancel₀ _ (ne_of_gt hkp)] - · push Not at hc - calc √(2 * k * σ * L) - _ = 0 := Real.sqrt_eq_zero_of_nonpos (by nlinarith) - _ ≤ |s - k * μ| := abs_nonneg _ - -lemma prob_abs_empMean_sub_actionMean_ge_le [IsMarkovKernel κ] - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} +lemma prob_abs_empMean_sub_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : P {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ - |empMean A R' a t ω - actionMean κ E a ω|} + √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ |empMean A R' a t ω - actionMean κ E a ω|} ≤ ENNReal.ofReal (2 * K * (n - 1) * δ) := by have := h.measurable_E have := h.measurable_A have := h.measurable_R - let S : Set (𝓔 × (ℕ → Fin K × ℝ)) := {(e, τ) | ∃ a, ∃ t < n, pullCount IT.action a t τ ≠ 0 ∧ + let S := {(e, τ) | ∃ a, ∃ t < n, pullCount IT.action a t τ ≠ 0 ∧ √(2 * pullCount IT.action a t τ * σ2 * Real.log (1 / δ)) ≤ |sumRewards IT.action IT.reward a t τ - pullCount IT.action a t τ * actionMean κ id a e|} - calc _ + calc _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by rw [Measure.map_apply (by fun_prop) (by measurability)] apply measure_mono intro ω ⟨t, ht, a, hpc, hle⟩ - exact ⟨a, t, ht, hpc, - sqrt_two_mul_le_abs_sub_of_sqrt_div_le hpc (by simpa [empMean] using hle)⟩ + rw [empMean] at hle + exact ⟨a, t, ht, hpc, sqrt_two_mul_le hpc hle⟩ _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by rw [← compProd_map_condDistrib (by fun_prop)] _ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := Measure.compProd_apply (by measurability) - _ ≤ ∫⁻ _, ENNReal.ofReal (2 * K * (n - 1) * δ) ∂(P.map E) := by + _ ≤ ∫⁻ e, ENNReal.ofReal (2 * Fintype.card (Fin K) * (n - 1) * δ) ∂(P.map E) := by apply lintegral_mono_ae rw [h.hasLaw_env.map_eq] filter_upwards [h.ae_IsAlgEnvSeq] with e he - convert Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 (hs e) he hδ - exact (Fintype.card_fin K).symm - _ = _ := by simp [Measure.map_apply h.measurable_E] + exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 (hs e) he hδ + _ = ENNReal.ofReal (2 * K * (n - 1) * δ) := by + simp [Measure.map_apply h.measurable_E] -lemma prob_abs_empMean_bestAction_sub_actionMean_ge_le [IsMarkovKernel κ] - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) +lemma prob_abs_empMean_bestAction_sub_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) + {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω : ℝ)) ≤ + √(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω)) ≤ |empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω|} ≤ ENNReal.ofReal (2 * (n - 1) * δ) := by have := h.measurable_E have := h.measurable_A have := h.measurable_R - let S : Set (𝓔 × (ℕ → Fin K × ℝ)) := - {(e, τ) | ∃ t < n, pullCount IT.action (bestAction κ id e) t τ ≠ 0 ∧ - √(2 * pullCount IT.action (bestAction κ id e) t τ * σ2 * Real.log (1 / δ)) ≤ - |sumRewards IT.action IT.reward (bestAction κ id e) t τ - - pullCount IT.action (bestAction κ id e) t τ * actionMean κ id (bestAction κ id e) e|} - calc _ + let S := {(e, τ) | ∃ t < n, pullCount IT.action (bestAction κ id e) t τ ≠ 0 ∧ + √(2 * pullCount IT.action (bestAction κ id e) t τ * σ2 * Real.log (1 / δ)) ≤ + |sumRewards IT.action IT.reward (bestAction κ id e) t τ - + pullCount IT.action (bestAction κ id e) t τ * actionMean κ id (bestAction κ id e) e|} + calc _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by rw [Measure.map_apply (by fun_prop) (by measurability)] apply measure_mono intro ω ⟨t, ht, hpc, hle⟩ - exact ⟨t, ht, hpc, - sqrt_two_mul_le_abs_sub_of_sqrt_div_le hpc (by simpa [empMean] using hle)⟩ + rw [empMean] at hle + exact ⟨t, ht, hpc, sqrt_two_mul_le hpc hle⟩ _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by rw [← compProd_map_condDistrib (by fun_prop)] _ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := Measure.compProd_apply (by measurability) - _ ≤ ∫⁻ _, ENNReal.ofReal (2 * (n - 1) * δ) ∂(P.map E) := by + _ ≤ ∫⁻ e, ENNReal.ofReal (2 * (n - 1) * δ) ∂(P.map E) := by apply lintegral_mono_ae rw [h.hasLaw_env.map_eq] filter_upwards [h.ae_IsAlgEnvSeq] with e he - exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le (ν := κ.sectR e) hσ2 - (hs e (bestAction κ id e)) he hδ - _ = _ := by simp [Measure.map_apply h.measurable_E] + exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le (ν := κ.sectR e) hσ2 (hs e _) he + hδ + _ = ENNReal.ofReal (2 * (n - 1) * δ) := by + simp [Measure.map_apply h.measurable_E] end HasSubgaussianMGF From ea80b8b56248beb477e8273a82351bd59939a53c Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 14 Apr 2026 12:00:26 +0100 Subject: [PATCH 100/155] Refactor TS.lean (in progress) --- LeanMachineLearning/BanditAlgorithms/TS.lean | 103 +++++++++---------- 1 file changed, 49 insertions(+), 54 deletions(-) diff --git a/LeanMachineLearning/BanditAlgorithms/TS.lean b/LeanMachineLearning/BanditAlgorithms/TS.lean index fae5595d..ae93fb0d 100644 --- a/LeanMachineLearning/BanditAlgorithms/TS.lean +++ b/LeanMachineLearning/BanditAlgorithms/TS.lean @@ -53,12 +53,37 @@ def tsAlgorithm (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 namespace TS -section UCB - +variable (hK : 0 < K) variable {Ω : Type*} -variable {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} +variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} +variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] + +lemma condDistrib_action_ae_eq_condDistrib_bestAction [Nonempty (Fin K)] [MeasurableSpace Ω] + {P : Measure Ω} [IsProbabilityMeasure P] + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (t : ℕ) : + condDistrib (A (t + 1)) (IsAlgEnvSeq.hist A R' t) P =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] + condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' t) P := by + have hm : Measurable (IsBayesAlgEnvSeq.bestAction κ id) := by fun_prop + calc ↑(condDistrib (A (t + 1)) (IsAlgEnvSeq.hist A R' t) P) + _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] + (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) t).map + (IsBayesAlgEnvSeq.bestAction κ id) := + (h.hasCondDistrib_action' t).condDistrib_eq + _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] + (condDistrib E (IsAlgEnvSeq.hist A R' t) P).map (IsBayesAlgEnvSeq.bestAction κ id) := by + filter_upwards [(h.hasCondDistrib_env_hist + (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) + (absolutelyContinuous_uniformAlgorithm hK _) t).condDistrib_eq] with x hx + simp [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hx] + _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] + condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' t) P := + (condDistrib_comp (mβ := MeasurableSpace.pi) (μ := P) + (IsAlgEnvSeq.hist A R' t) h.measurable_E.aemeasurable hm).symm + variable {l u σ2 δ : ℝ} +section UCB + noncomputable def ucb (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ := if pullCount A a n ω = 0 then u @@ -194,42 +219,12 @@ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, end UCB -variable {K : ℕ} -variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] -variable (hK : 0 < K) -variable {Ω : Type*} [MeasurableSpace Ω] -variable (E : Ω → 𝓔) (A : ℕ → Ω → (Fin K)) (R' : ℕ → Ω → ℝ) -variable (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] -variable (P : Measure Ω) [IsProbabilityMeasure P] +section BayesRegret -lemma ts_identity [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (t : ℕ) : - condDistrib (A (t + 1)) (IsAlgEnvSeq.hist A R' t) P - =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' t) P := by - have h_ba_comp : IsBayesAlgEnvSeq.bestAction κ E - = IsBayesAlgEnvSeq.bestAction κ id ∘ E := rfl - rw [h_ba_comp] - have hm := IsBayesAlgEnvSeq.measurable_bestAction (κ := κ) measurable_id - have h_comp := condDistrib_comp (mβ := MeasurableSpace.pi) (μ := P) - (IsAlgEnvSeq.hist A R' t) h.measurable_E.aemeasurable hm - have h_map : (condDistrib E (IsAlgEnvSeq.hist A R' t) P).map - (IsBayesAlgEnvSeq.bestAction κ id) =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) t).map - (IsBayesAlgEnvSeq.bestAction κ id) := by - filter_upwards [(h.hasCondDistrib_env_hist - (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) - (absolutelyContinuous_uniformAlgorithm hK _) t).condDistrib_eq] - with x hx - simp only [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hx] - exact (h.hasCondDistrib_action' t).condDistrib_eq.trans (h_comp.trans h_map).symm - -lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) - {σ2 : ℝ≥0} (hσ2 : 0 < σ2) - (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - {l u : ℝ} (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) - (n : ℕ) (δ : ℝ) (hδ : 0 < δ) : +lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure Ω} + [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hσ2 : 0 < σ2) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) + (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hδ : 0 < δ) (n : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A n] ≤ (u - l) * ↑K + 2 * (↑K + 1) * (u - l) * n ^ 2 * δ + 4 * √(2 * ↑σ2 * Real.log (1 / δ) * ↑K * ↑n) := by @@ -319,7 +314,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp simp [h_ucb_zero] exact (integral_congr_ae (ae_of_all _ this)).trans (integral_zero _ _) | succ t => - have hts := ts_identity hK E A R' Q κ P h t + have hts := condDistrib_action_ae_eq_condDistrib_bestAction hK h t have h_map_eq : P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω)) = P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E ω)) := by rw [← compProd_map_condDistrib (hY := (h.measurable_A (t + 1)).aemeasurable), @@ -474,15 +469,14 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonemp measureReal_nonneg (μ := P) (s := Fδᶜ), measureReal_nonneg (μ := P) (s := Eδᶜ)] -lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) - {σ2 : ℝ≥0} (hσ2 : 0 < σ2) - (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - {lo hi : ℝ} (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc lo hi)) (t : ℕ) : +lemma bayesRegret_le [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hσ2 : 0 < σ2) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) + (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (t : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A t] - ≤ (3 * K + 2) * (hi - lo) + 8 * √(σ2 * K * t * Real.log t) := by + ≤ (3 * K + 2) * (u - l) + 8 * √(σ2 * K * t * Real.log t) := by have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _) - have hlo : lo ≤ hi := h1.trans h2 + have hlo : l ≤ u := h1.trans h2 by_cases ht : t = 0 · simp [ht, IsBayesAlgEnvSeq.regret, Bandits.regret] nlinarith [sub_nonneg.mpr hlo, Nat.cast_pos (α := ℝ).mpr hK, @@ -491,14 +485,14 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] · subst ht1_eq simp only [Nat.cast_one, Real.log_one, mul_zero, Real.sqrt_zero, mul_zero, add_zero] calc P[IsBayesAlgEnvSeq.regret κ E A 1] - ≤ hi - lo := by + ≤ u - l := by rw [IsBayesAlgEnvSeq.regret_eq_sum_gap'] simp only [Finset.range_one, Finset.sum_singleton] exact (integral_mono_of_nonneg (ae_of_all _ fun ω ↦ IsBayesAlgEnvSeq.gap_nonneg_of_le (fun e a ↦ (hm e a).2)) (integrable_const _) (ae_of_all _ fun ω ↦ IsBayesAlgEnvSeq.gap_le_of_mem_Icc hm)).trans (by simp) - _ ≤ (3 * ↑K + 2) * (hi - lo) := by + _ ≤ (3 * ↑K + 2) * (u - l) := by nlinarith [Nat.one_le_cast (α := ℝ).mpr (Nat.one_le_of_lt hK), sub_nonneg.mpr hlo] -- For t ≥ 2, we have δ = 1/t² < 1 @@ -507,18 +501,17 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] have _ht1 : (1 : ℝ) ≤ t := by exact_mod_cast Nat.pos_of_ne_zero ht have hδ : (0 : ℝ) < 1 / (t : ℝ) ^ 2 := by positivity -- First term: (hi-lo)*K + 2*(K+1)*(hi-lo)*t²*(1/t²) = (3K+2)*(hi-lo) - have h_first : (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * ↑t ^ 2 * (1 / (↑t) ^ 2) - = (3 * ↑K + 2) * (hi - lo) := by + have h_first : (u - l) * ↑K + 2 * (↑K + 1) * (u - l) * ↑t ^ 2 * (1 / (↑t) ^ 2) + = (3 * ↑K + 2) * (u - l) := by field_simp; ring -- Second term simplification: log(1/(1/t²)) = log(t²) = 2 log(t) have h_log : Real.log (1 / (1 / (↑t : ℝ) ^ 2)) = 2 * Real.log ↑t := by rw [one_div_one_div, Real.log_pow]; norm_cast calc P[IsBayesAlgEnvSeq.regret κ E A t] - ≤ (hi - lo) * ↑K + 2 * (↑K + 1) * (hi - lo) * ↑t ^ 2 * (1 / (↑t) ^ 2) + ≤ (u - l) * ↑K + 2 * (↑K + 1) * (u - l) * ↑t ^ 2 * (1 / (↑t) ^ 2) + 4 * √(2 * ↑σ2 * Real.log (1 / (1 / (↑t) ^ 2)) * ↑K * ↑t) := - bayesRegret_le_of_delta (hK := hK) (E := E) (A := A) (R' := R') (Q := Q) - (κ := κ) (P := P) h hσ2 hs hm t (1 / (↑t) ^ 2) hδ - _ = (3 * ↑K + 2) * (hi - lo) + 8 * √(↑σ2 * ↑K * ↑t * Real.log ↑t) := by + bayesRegret_le_of_delta (δ := 1 / (↑t) ^ 2) hK h hσ2 hs hm hδ t + _ = (3 * ↑K + 2) * (u - l) + 8 * √(↑σ2 * ↑K * ↑t * Real.log ↑t) := by rw [h_first, h_log]; congr 1 rw [show (2 : ℝ) * ↑σ2 * (2 * Real.log ↑t) * ↑K * ↑t = (2 : ℝ) ^ 2 * (↑σ2 * ↑K * ↑t * Real.log ↑t) from by ring, @@ -526,6 +519,8 @@ lemma bayesRegret_le [Nonempty (Fin K)] [StandardBorelSpace Ω] [Nonempty Ω] Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 2)] ring +end BayesRegret + end TS end Bandits From 5ea738bc16fb0ab7273cfe095e0829ca13f78b29 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 14 Apr 2026 14:25:07 +0100 Subject: [PATCH 101/155] Refactor TS.lean (in progress) --- LeanMachineLearning/BanditAlgorithms/TS.lean | 46 ++++++++++---------- 1 file changed, 24 insertions(+), 22 deletions(-) diff --git a/LeanMachineLearning/BanditAlgorithms/TS.lean b/LeanMachineLearning/BanditAlgorithms/TS.lean index ae93fb0d..96b23b17 100644 --- a/LeanMachineLearning/BanditAlgorithms/TS.lean +++ b/LeanMachineLearning/BanditAlgorithms/TS.lean @@ -58,27 +58,29 @@ variable {Ω : Type*} variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] -lemma condDistrib_action_ae_eq_condDistrib_bestAction [Nonempty (Fin K)] [MeasurableSpace Ω] - {P : Measure Ω} [IsProbabilityMeasure P] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (t : ℕ) : - condDistrib (A (t + 1)) (IsAlgEnvSeq.hist A R' t) P =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' t) P := by - have hm : Measurable (IsBayesAlgEnvSeq.bestAction κ id) := by fun_prop - calc ↑(condDistrib (A (t + 1)) (IsAlgEnvSeq.hist A R' t) P) - _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) t).map - (IsBayesAlgEnvSeq.bestAction κ id) := - (h.hasCondDistrib_action' t).condDistrib_eq - _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - (condDistrib E (IsAlgEnvSeq.hist A R' t) P).map (IsBayesAlgEnvSeq.bestAction κ id) := by - filter_upwards [(h.hasCondDistrib_env_hist - (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) - (absolutelyContinuous_uniformAlgorithm hK _) t).condDistrib_eq] with x hx - simp [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hx] - _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' t)] - condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' t) P := - (condDistrib_comp (mβ := MeasurableSpace.pi) (μ := P) - (IsAlgEnvSeq.hist A R' t) h.measurable_E.aemeasurable hm).symm +lemma hasCondDistrib_action [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure Ω} + [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (n : ℕ) : + HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) + (condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' n) P) P where + aemeasurable_fst := (h.measurable_A (n + 1)).aemeasurable + aemeasurable_snd := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable + condDistrib_eq := by + have hm : Measurable (IsBayesAlgEnvSeq.bestAction κ id) := by fun_prop + calc + _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' n)] + (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) n).map + (IsBayesAlgEnvSeq.bestAction κ id) := + (h.hasCondDistrib_action' n).condDistrib_eq + _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' n)] + (condDistrib E (IsAlgEnvSeq.hist A R' n) P).map + (IsBayesAlgEnvSeq.bestAction κ id) := by + filter_upwards [(h.hasCondDistrib_env_hist + (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) + (absolutelyContinuous_uniformAlgorithm hK _) n).condDistrib_eq] with _ hc + simp_rw [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hc] + _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' n)] + condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' n) P := + (condDistrib_comp (IsAlgEnvSeq.hist A R' n) h.measurable_E.aemeasurable hm).symm variable {l u σ2 δ : ℝ} @@ -314,7 +316,7 @@ lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measu simp [h_ucb_zero] exact (integral_congr_ae (ae_of_all _ this)).trans (integral_zero _ _) | succ t => - have hts := condDistrib_action_ae_eq_condDistrib_bestAction hK h t + have hts := (hasCondDistrib_action hK h t).condDistrib_eq have h_map_eq : P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω)) = P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E ω)) := by rw [← compProd_map_condDistrib (hY := (h.measurable_A (t + 1)).aemeasurable), From f2f91ad3937ac9987e030320e0c0c1448eec7500 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 15 Apr 2026 14:34:21 +0100 Subject: [PATCH 102/155] Refactor TS.lean (in progress) --- LeanMachineLearning/BanditAlgorithms/TS.lean | 639 ++++++++++-------- .../BayesStationaryEnv.lean | 8 + .../SequentialLearning/FiniteActions.lean | 7 + 3 files changed, 366 insertions(+), 288 deletions(-) diff --git a/LeanMachineLearning/BanditAlgorithms/TS.lean b/LeanMachineLearning/BanditAlgorithms/TS.lean index 96b23b17..3e5f920c 100644 --- a/LeanMachineLearning/BanditAlgorithms/TS.lean +++ b/LeanMachineLearning/BanditAlgorithms/TS.lean @@ -91,6 +91,11 @@ def ucb (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) if pullCount A a n ω = 0 then u else max l (min u (empMean A R' a n ω + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a n ω)))) +lemma ucb_mem_Icc (h : l ≤ u) {a : Fin K} {n : ℕ} {ω : Ω} : + ucb A R' l u σ2 δ a n ω ∈ Set.Icc l u := by + unfold ucb + grind + @[fun_prop] lemma measurable_ucb [MeasurableSpace Ω] {a : Fin K} {n : ℕ} (hA : ∀ t, Measurable (A t)) (hR : ∀ t, Measurable (R' t)) : Measurable (ucb A R' l u σ2 δ a n) := @@ -108,10 +113,13 @@ lemma measurable_uncurry_ucb_comp [MeasurableSpace Ω] (hA : ∀ t, Measurable ( apply Measurable.comp _ (by fun_prop) exact measurable_from_prod_countable_right (fun _ ↦ measurable_ucb hA hR) -lemma ucb_mem_Icc (h : l ≤ u) {a : Fin K} {n : ℕ} {ω : Ω} : - ucb A R' l u σ2 δ a n ω ∈ Set.Icc l u := by - unfold ucb - grind +lemma integrable_uncurry_ucb_comp [MeasurableSpace Ω] (hA : ∀ t, Measurable (A t)) + (hR : ∀ t, Measurable (R' t)) {f : Ω → Fin K} (hf : Measurable f) {g : Ω → ℕ} + (hg : Measurable g) {P : Measure Ω} [IsFiniteMeasure P] (hlu : l ≤ u) : + Integrable (fun ω ↦ ucb A R' l u σ2 δ (f ω) (g ω) ω) P := by + refine ⟨(measurable_uncurry_ucb_comp hA hR hf hg).aestronglyMeasurable, ?_⟩ + apply HasFiniteIntegral.of_bounded + filter_upwards with ω using abs_le_max_abs_abs (ucb_mem_Icc hlu).1 (ucb_mem_Icc hlu).2 noncomputable def ucb' (n : ℕ) (h : Iic n → Fin K × ℝ) (l u σ2 δ : ℝ) (a : Fin K) : ℝ := @@ -221,307 +229,362 @@ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, end UCB -section BayesRegret - -lemma bayesRegret_le_of_delta [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure Ω} - [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hσ2 : 0 < σ2) - (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) - (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hδ : 0 < δ) (n : ℕ) : - P[IsBayesAlgEnvSeq.regret κ E A n] - ≤ (u - l) * ↑K + 2 * (↑K + 1) * (u - l) * n ^ 2 * δ + - 4 * √(2 * ↑σ2 * Real.log (1 / δ) * ↑K * ↑n) := by - have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _) - have hlo : l ≤ u := h1.trans h2 - let bestArm := IsBayesAlgEnvSeq.bestAction κ E - let armMean := IsBayesAlgEnvSeq.actionMean κ E - let uc := ucb A R' l u (↑σ2) δ - set Eδ : Set Ω := {ω | ∀ s < n, ∀ a, pullCount A a s ω ≠ 0 → - |empMean A R' a s ω - armMean a ω| - < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ))} - set Fδ : Set Ω := {ω | ∀ s < n, pullCount A (bestArm ω) s ω ≠ 0 → - |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω| - < √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A (bestArm ω) s ω : ℝ))} - have hm_ucb : ∀ a t, Measurable (ucb A R' l u (↑σ2) δ a t) := - fun _ _ ↦ measurable_ucb h.measurable_A h.measurable_R - have hm_arm : ∀ a, Measurable (IsBayesAlgEnvSeq.actionMean κ E a) := - fun a ↦ IsBayesAlgEnvSeq.measurable_actionMean (a := a) h.measurable_E - have hm_best : Measurable (IsBayesAlgEnvSeq.bestAction κ E) := - IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E - have h_first_bound : ∀ ω, - |∑ s ∈ range n, (armMean (bestArm ω) ω - uc (bestArm ω) s ω)| - ≤ n * (u - l) := fun ω ↦ - calc |∑ s ∈ range n, (armMean (bestArm ω) ω - uc (bestArm ω) s ω)| - ≤ ∑ s ∈ range n, |armMean (bestArm ω) ω - uc (bestArm ω) s ω| := - Finset.abs_sum_le_sum_abs _ _ - _ ≤ ∑ s ∈ range n, (u - l) := by - gcongr with s _ - exact abs_sub_le_of_le_of_le (hm _ _).1 (hm _ _).2 - ((ucb_mem_Icc hlo).1) - (ucb_mem_Icc hlo).2 - _ = ↑n * (u - l) := by - rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] - have h_second_bound : ∀ ω, - |∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω)| - ≤ n * (u - l) := fun ω ↦ - calc |∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω)| - ≤ ∑ s ∈ range n, |uc (A s ω) s ω - armMean (A s ω) ω| := - Finset.abs_sum_le_sum_abs _ _ - _ ≤ ∑ s ∈ range n, (u - l) := by - gcongr with s _ - exact abs_sub_le_of_le_of_le (ucb_mem_Icc hlo).1 - (ucb_mem_Icc hlo).2 (hm _ _).1 (hm _ _).2 - _ = ↑n * (u - l) := by - rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul] - have h_int_sum1 : Integrable (fun ω ↦ ∑ s ∈ range n, - (armMean (bestArm ω) ω - uc (bestArm ω) s ω)) P := by - apply Integrable.of_bound (C := ↑n * (u - l)) - · exact (Finset.measurable_fun_sum _ fun s _ ↦ - (IsBayesAlgEnvSeq.measurable_uncurry_actionMean_comp h.measurable_E hm_best).sub - (measurable_uncurry_ucb_comp h.measurable_A h.measurable_R hm_best - measurable_const)).aestronglyMeasurable - · filter_upwards with ω; rw [Real.norm_eq_abs]; exact h_first_bound ω - have h_int_sum2 : Integrable (fun ω ↦ ∑ s ∈ range n, - (uc (A s ω) s ω - armMean (A s ω) ω)) P := by - apply Integrable.of_bound (C := ↑n * (u - l)) - · exact (Finset.measurable_fun_sum _ fun s _ ↦ - (measurable_uncurry_ucb_comp h.measurable_A h.measurable_R - (h.measurable_A s) measurable_const).sub - (IsBayesAlgEnvSeq.measurable_uncurry_actionMean_comp h.measurable_E (h.measurable_A s)) - ).aestronglyMeasurable - · filter_upwards with ω; rw [Real.norm_eq_abs]; exact h_second_bound ω - have h_swap : - P[IsBayesAlgEnvSeq.regret κ E A n] = +section IntegralRegret + +lemma integral_ucb_action_eq_integral_ucb_bestAction [Nonempty (Fin K)] + [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (n : ℕ) : + P[fun ω ↦ ucb A R' l u σ2 δ (A n ω) n ω] = + P[fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) n ω] := by + have := h.measurable_A + have := h.measurable_E + have := h.measurable_R + by_cases hs : n = 0 + · simp [hs, ucb, pullCount_zero] + obtain ⟨t, rfl⟩ := Nat.exists_eq_succ_of_ne_zero hs + let ucb'' : (Iic t → Fin K × ℝ) × Fin K → ℝ := fun p ↦ ucb' t p.1 l u σ2 δ p.2 + calc P[fun ω ↦ ucb A R' l u σ2 δ (A (t + 1) ω) (t + 1) ω] + = P[fun ω ↦ ucb'' (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω)] := by + simp_rw [ucb'', ucb_succ_eq_ucb'] + _ = ∫ p, ucb'' p ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω)) := by + rw [← integral_map (by fun_prop) (by fun_prop)] + _ = ∫ p, ucb'' p ∂P.map + (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E ω)) := by + rw [← compProd_map_condDistrib (hY := (h.measurable_A (t + 1)).aemeasurable), + ← compProd_map_condDistrib + (hY := (by fun_prop : AEMeasurable (IsBayesAlgEnvSeq.bestAction κ E) _)), + Measure.compProd_congr (hasCondDistrib_action hK h t).condDistrib_eq] + _ = P[fun ω ↦ ucb'' (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E ω)] := by + rw [integral_map (by fun_prop) (by fun_prop)] + _ = P[fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) (t + 1) ω] := by + simp_rw [ucb'', ucb_succ_eq_ucb'] + +lemma integral_regret_eq_add [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure Ω} + [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) + (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) : + P[IsBayesAlgEnvSeq.regret κ E A n] = P[fun ω ↦ ∑ s ∈ range n, - (armMean (bestArm ω) ω - uc (bestArm ω) s ω)] + + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)] + P[fun ω ↦ ∑ s ∈ range n, - (uc (A s ω) s ω - armMean (A s ω) ω)] := by - have h_int_ucb : ∀ s {f : Ω → Fin K}, Measurable f → - Integrable (fun ω ↦ uc (f ω) s ω) P := fun s {_} hf ↦ - ⟨(measurable_uncurry_ucb_comp - h.measurable_A h.measurable_R hf measurable_const).aestronglyMeasurable, - HasFiniteIntegral.of_bounded (ae_of_all _ fun ω ↦ by - rw [Real.norm_eq_abs] - exact abs_le_max_abs_abs (ucb_mem_Icc hlo).1 - (ucb_mem_Icc hlo).2)⟩ - have h_int_ucb_sub : ∀ s, Integrable (fun ω ↦ uc (A s ω) s ω - uc (bestArm ω) s ω) P := - fun s ↦ (h_int_ucb s (h.measurable_A s)).sub (h_int_ucb s hm_best) - have h_ucb_zero : ∀ a (ω : Ω), ucb A R' l u (↑σ2) δ a 0 ω = u := by - intro a ω; unfold ucb; simp [pullCount_zero] - have h_ucb_swap : ∀ s, ∫ ω, (uc (A s ω) s ω - uc (bestArm ω) s ω) ∂P = 0 := by - intro s - cases s with - | zero => - have : ∀ ω, uc (A 0 ω) 0 ω - uc (bestArm ω) 0 ω = 0 := fun ω ↦ by - change ucb A R' l u (↑σ2) δ _ 0 ω - ucb A R' l u (↑σ2) δ _ 0 ω = 0 - simp [h_ucb_zero] - exact (integral_congr_ae (ae_of_all _ this)).trans (integral_zero _ _) - | succ t => - have hts := (hasCondDistrib_action hK h t).condDistrib_eq - have h_map_eq : P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω)) = - P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E ω)) := by - rw [← compProd_map_condDistrib (hY := (h.measurable_A (t + 1)).aemeasurable), - ← compProd_map_condDistrib (hY := hm_best.aemeasurable)] - exact Measure.compProd_congr hts - have h_int_eq : ∀ (f : (Iic t → Fin K × ℝ) × Fin K → ℝ), Measurable f → - ∫ ω, f (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω) ∂P = - ∫ ω, f (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E ω) ∂P := by - intro f hf - have hm_hist := IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R t - rw [← integral_map - (hm_hist.prodMk (h.measurable_A (t + 1))).aemeasurable - hf.aestronglyMeasurable, - ← integral_map - (hm_hist.prodMk hm_best).aemeasurable - hf.aestronglyMeasurable, - h_map_eq] - let g : (Iic t → Fin K × ℝ) × Fin K → ℝ := - fun p ↦ ucb' t p.1 l u (↑σ2) δ p.2 - have hg_eq : ∀ a (ω : Ω), ucb A R' l u (↑σ2) δ a (t + 1) ω = - g (IsAlgEnvSeq.hist A R' t ω, a) := fun _ _ ↦ ucb_succ_eq_ucb' - have hg_meas : Measurable g := measurable_uncurry_ucb' - rw [show (fun ω ↦ uc (A (t + 1) ω) (t + 1) ω - - uc (bestArm ω) (t + 1) ω) = - fun ω ↦ (fun ω ↦ uc (A (t + 1) ω) (t + 1) ω) ω - - (fun ω ↦ uc (bestArm ω) (t + 1) ω) ω from rfl, - integral_sub (h_int_ucb (t + 1) (h.measurable_A (t + 1))) - (h_int_ucb (t + 1) hm_best), - funext fun ω ↦ hg_eq _ _, funext fun ω ↦ hg_eq _ _, - h_int_eq g hg_meas, sub_self] - have h_ucb_sum_zero : ∫ ω, ∑ s ∈ range n, - (uc (A s ω) s ω - uc (bestArm ω) s ω) ∂P = 0 := by - rw [integral_finset_sum _ (fun s _ ↦ h_int_ucb_sub s)] - exact Finset.sum_eq_zero fun s _ ↦ h_ucb_swap s - have h_int_gap : Integrable (fun ω ↦ IsBayesAlgEnvSeq.regret κ E A n ω) P := - IsBayesAlgEnvSeq.integrable_regret h.measurable_E (h.measurable_A) hm - simp_rw [IsBayesAlgEnvSeq.regret_eq_sum_gap, IsBayesAlgEnvSeq.gap_eq_sub] at h_int_gap ⊢ - have h_pw : ∀ ω, (∑ s ∈ range n, (armMean (bestArm ω) ω - uc (bestArm ω) s ω)) + - (∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω)) = - (∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω)) + - (∑ s ∈ range n, (uc (A s ω) s ω - uc (bestArm ω) s ω)) := by - intro ω - simp only [← Finset.sum_add_distrib] - apply Finset.sum_congr rfl; intros; ring - have h_int_ucb_swap : Integrable - (fun ω ↦ ∑ s ∈ range n, (uc (A s ω) s ω - uc (bestArm ω) s ω)) P := - integrable_finset_sum _ fun s _ ↦ h_int_ucb_sub s - calc ∫ ω, ∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω) ∂P - = ∫ ω, ((∑ s ∈ range n, (armMean (bestArm ω) ω - armMean (A s ω) ω)) + - (∑ s ∈ range n, (uc (A s ω) s ω - uc (bestArm ω) s ω))) ∂P := by - rw [integral_add h_int_gap h_int_ucb_swap, h_ucb_sum_zero, add_zero] - _ = ∫ ω, ((∑ s ∈ range n, (armMean (bestArm ω) ω - uc (bestArm ω) s ω)) + - (∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω))) ∂P := by - congr 1; ext ω; linarith [h_pw ω] - _ = _ := integral_add h_int_sum1 h_int_sum2 - have h_first_Fδ : ∀ ω ∈ Fδ, - ∑ s ∈ range n, (armMean (bestArm ω) ω - uc (bestArm ω) s ω) - ≤ 0 := by - intro ω hω - apply Finset.sum_nonpos - intro s hs - have : armMean (bestArm ω) ω ≤ uc (bestArm ω) s ω := by - simp only [armMean, uc]; unfold ucb - split_ifs with h0 - · exact (hm (E ω) (bestArm ω)).2 - · have := abs_lt.mp ((hω s (mem_range.mp hs)) h0) - exact le_max_of_le_right (le_min (hm (E ω) (bestArm ω)).2 (by linarith)) - linarith - have h_second_Eδ : ∀ ω ∈ Eδ, - ∑ s ∈ range n, (uc (A s ω) s ω - armMean (A s ω) ω) - ≤ (u - l) * ↑K + 4 * √(2 * ↑σ2 * Real.log (1 / δ) * ↑K * ↑n) := by - intro ω hω - exact sum_ucb_sub_mean_le (fun a ↦ armMean a ω) (hm (E ω)) hlo - (fun s hs hpc => hω s hs (A s ω) hpc) - have h_prob : P Eδᶜ ≤ ENNReal.ofReal (2 * K * n * δ) := by - have : Eδᶜ = {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ - |empMean A R' a s ω - armMean a ω|} := by - ext ω; simp only [Eδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push Not; rfl - rw [this] - exact (h.prob_abs_empMean_sub_actionMean_ge_le hσ2 hs hδ n).trans - (ENNReal.ofReal_le_ofReal (by nlinarith [hδ.le, Nat.cast_nonneg (α := ℝ) K])) - have hm_emp : ∀ a s, Measurable (fun ω ↦ empMean A R' a s ω) := - fun a s ↦ measurable_empMean (fun n ↦ h.measurable_A n) (fun n ↦ h.measurable_R n) a s - have hm_pc : ∀ a s, Measurable (fun ω ↦ (pullCount A a s ω : ℝ)) := - fun a s ↦ measurable_from_top.comp (measurable_pullCount (fun n ↦ h.measurable_A n) a s) - have h_arm_meas : ∀ s a, MeasurableSet {ω : Ω | pullCount A a s ω ≠ 0 → - |empMean A R' a s ω - armMean a ω| - < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} := by - intro s a - have : {ω : Ω | pullCount A a s ω ≠ 0 → - |empMean A R' a s ω - armMean a ω| - < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} = - {ω | (pullCount A a s ω : ℝ) = 0} ∪ {ω | - |empMean A R' a s ω - armMean a ω| - < √(2 * ↑σ2 * Real.log (1 / δ) / ↑(pullCount A a s ω))} := by - ext ω; simp only [Set.mem_setOf_eq, Set.mem_union, Nat.cast_eq_zero]; tauto - rw [this] - exact .union (hm_pc a s (measurableSet_singleton _)) - (measurableSet_lt (by fun_prop) (by fun_prop)) - have hEδ_meas : MeasurableSet Eδ := by - simp only [Eδ, Set.setOf_forall] - exact .iInter fun s ↦ .iInter fun _ ↦ .iInter fun a ↦ h_arm_meas s a - have hFδ_meas : MeasurableSet Fδ := by - simp only [Fδ, Set.setOf_forall] - refine .iInter fun s ↦ .iInter fun _ ↦ ?_ - convert MeasurableSet.iUnion fun a ↦ - (hm_best (measurableSet_singleton a)).inter (h_arm_meas s a) using 1 - ext ω; simp only [Set.mem_iUnion, Set.mem_inter_iff, Set.mem_preimage, - Set.mem_singleton_iff, Set.mem_setOf_eq] - exact ⟨fun h => ⟨_, rfl, h⟩, fun ⟨_, rfl, h⟩ => h⟩ - have h_prob_F : P Fδᶜ ≤ ENNReal.ofReal (2 * ↑n * δ) := by - have : Fδᶜ = {ω | ∃ s < n, pullCount A (bestArm ω) s ω ≠ 0 ∧ - √(2 * ↑σ2 * Real.log (1 / δ) / (pullCount A (bestArm ω) s ω : ℝ)) ≤ - |empMean A R' (bestArm ω) s ω - armMean (bestArm ω) ω|} := by - ext ω; simp only [Fδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push Not; rfl - rw [this] - exact (h.prob_abs_empMean_bestAction_sub_actionMean_ge_le hσ2 hs hδ n).trans - (ENNReal.ofReal_le_ofReal (by nlinarith [hδ.le])) - rw [h_swap] - set f1 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, - (armMean (bestArm ω) ω - uc (bestArm ω) s ω) - set f2 : Ω → ℝ := fun ω ↦ ∑ s ∈ range n, - (uc (A s ω) s ω - armMean (A s ω) ω) - set B := (u - l) * ↑K + 4 * √(2 * ↑σ2 * Real.log (1 / δ) * ↑K * ↑n) - have h1g : ∫ ω in Fδ, f1 ω ∂P ≤ 0 := - setIntegral_nonpos hFδ_meas fun ω hω ↦ h_first_Fδ ω hω - have h1b : ∫ ω in Fδᶜ, f1 ω ∂P ≤ ↑n * (u - l) * P.real Fδᶜ := by - have := setIntegral_mono_on (hf := h_int_sum1.integrableOn) (hg := integrableOn_const) - hFδ_meas.compl fun ω _ ↦ (abs_le.mp (h_first_bound ω)).2 - rwa [setIntegral_const, smul_eq_mul, mul_comm] at this - have h2g : ∫ ω in Eδ, f2 ω ∂P ≤ B := by - have hB : 0 ≤ B := by have : 0 ≤ u - l := sub_nonneg.mpr hlo; positivity - have := setIntegral_mono_on (hf := h_int_sum2.integrableOn) - (hg := integrableOn_const) hEδ_meas - fun ω hω ↦ h_second_Eδ ω hω - rw [setIntegral_const, smul_eq_mul, mul_comm] at this - exact le_trans this (mul_le_of_le_one_right hB measureReal_le_one) - have h2b : ∫ ω in Eδᶜ, f2 ω ∂P ≤ ↑n * (u - l) * P.real Eδᶜ := by - have := setIntegral_mono_on (hf := h_int_sum2.integrableOn) (hg := integrableOn_const) - hEδ_meas.compl fun ω _ ↦ (abs_le.mp (h_second_bound ω)).2 - rwa [setIntegral_const, smul_eq_mul, mul_comm] at this - have hPF : P.real Fδᶜ ≤ 2 * ↑n * δ := - ENNReal.toReal_le_of_le_ofReal (by positivity) h_prob_F - have hPE : P.real Eδᶜ ≤ 2 * ↑K * ↑n * δ := - ENNReal.toReal_le_of_le_ofReal (by positivity) h_prob - rw [(integral_add_compl hFδ_meas h_int_sum1).symm, - (integral_add_compl hEδ_meas h_int_sum2).symm] - have : 0 ≤ ↑n * (u - l) := by nlinarith - nlinarith [mul_le_mul_of_nonneg_left hPF this, - mul_le_mul_of_nonneg_left hPE this, - measureReal_nonneg (μ := P) (s := Fδᶜ), - measureReal_nonneg (μ := P) (s := Eδᶜ)] - -lemma bayesRegret_le [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)] := by + have := h.measurable_A + have := h.measurable_E + have := h.measurable_R + calc P[IsBayesAlgEnvSeq.regret κ E A n] + = ∫ ω, ∑ s ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P := by + simp_rw [IsBayesAlgEnvSeq.regret_eq_sum_gap, IsBayesAlgEnvSeq.gap_eq_sub] + _ = (∫ ω, ∑ s ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P) + 0 := + (add_zero _).symm + _ = (∫ ω, ∑ s ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P) + + ∫ ω, ∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω) ∂P := by + congr 1 + symm + rw [integral_finset_sum _ ?_] + · apply Finset.sum_eq_zero + intro s _ + rw [integral_sub ?_ ?_, + integral_ucb_action_eq_integral_ucb_bestAction (hK := hK) h s, + sub_self] + · exact integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A s) + measurable_const hlu + · exact integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (by fun_prop) + measurable_const hlu + · intro s _ + exact (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A s) + measurable_const hlu).sub + (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (by fun_prop) + measurable_const hlu) + _ = ∫ ω, + ((∑ s ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)) + + (∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω))) ∂P := by + rw [← integral_add] + · apply integrable_finset_sum + intro s _ + exact (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E + (by fun_prop) hm).sub + (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E + (h.measurable_A s) hm) + · apply integrable_finset_sum + intro s _ + exact (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A s) + measurable_const hlu).sub + (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (by fun_prop) + measurable_const hlu) + _ = ∫ ω, + ((∑ s ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)) + + (∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω))) ∂P := by + congr 1 + ext ω + simp only [← Finset.sum_add_distrib] + apply Finset.sum_congr rfl + intros + ring + _ = (∫ ω, ∑ s ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω) ∂P) + + ∫ ω, ∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P := by + apply integral_add + · apply integrable_finset_sum + intro s _ + exact (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E + (by fun_prop) hm).sub + (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (by fun_prop) + measurable_const hlu) + · apply integrable_finset_sum + intro s _ + exact (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A s) + measurable_const hlu).sub + (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E + (h.measurable_A s) hm) + +lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le [Nonempty (Fin K)] + [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hσ2 : 0 < σ2) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) + (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hδ : 0 < δ) (n : ℕ) : + P[fun ω ↦ ∑ s ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)] ≤ + 2 * (u - l) * n ^ 2 * δ := by + have := h.measurable_A + have := h.measurable_E + have := h.measurable_R + set Fδ := {ω | ∀ t < n, pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) t ω ≠ 0 → + |empMean A R' (IsBayesAlgEnvSeq.bestAction κ E ω) t ω - + IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω| + < √(2 * σ2 * Real.log (1 / δ) / (pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) t ω))} + have hFδ_meas : MeasurableSet Fδ := by measurability + have h_int : Integrable (fun ω ↦ ∑ s ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)) P := + integrable_finset_sum _ fun s _ ↦ + (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E (by fun_prop) hm).sub + (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (by fun_prop) + measurable_const hlu) + calc ∫ ω, ∑ s ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω) ∂P + = (∫ ω in Fδ, ∑ s ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω) ∂P) + + ∫ ω in Fδᶜ, ∑ s ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω) ∂P := + (integral_add_compl hFδ_meas h_int).symm + _ ≤ 0 + ∫ ω in Fδᶜ, ∑ s ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω) ∂P := by + gcongr + apply setIntegral_nonpos hFδ_meas + intro ω hω + apply Finset.sum_nonpos + intro s hs + have : IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ≤ + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω := by + unfold ucb + split_ifs with h0 + · exact (hm (E ω) (IsBayesAlgEnvSeq.bestAction κ E ω)).2 + · have := abs_lt.mp ((hω s (mem_range.mp hs)) h0) + exact le_max_of_le_right + (le_min (hm (E ω) (IsBayesAlgEnvSeq.bestAction κ E ω)).2 (by linarith)) + linarith + _ ≤ 0 + ∫ _ω in Fδᶜ, n * (u - l) ∂P := by + apply add_le_add le_rfl + apply setIntegral_mono_on h_int.integrableOn integrableOn_const hFδ_meas.compl + intro ω _ + refine le_of_le_of_eq (Finset.sum_le_card_nsmul (range n) _ (u - l) fun s _ ↦ ?_) ?_ + · have h1 : IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ≤ u := + (hm _ _).2 + have h2 : l ≤ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω := + (ucb_mem_Icc hlu).1 + linarith + · rw [Finset.card_range, nsmul_eq_mul] + _ = 0 + n * (u - l) * P.real Fδᶜ := by + rw [setIntegral_const, smul_eq_mul, mul_comm] + _ ≤ 0 + n * (u - l) * (2 * n * δ) := by + gcongr + · nlinarith + · apply ENNReal.toReal_le_of_le_ofReal (by positivity) + have : Fδᶜ = {ω | ∃ s < n, pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) s ω ≠ 0 ∧ + √(2 * σ2 * Real.log (1 / δ) / + (pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) s ω : ℝ)) ≤ + |empMean A R' (IsBayesAlgEnvSeq.bestAction κ E ω) s ω - + IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω|} := by + ext ω; simp only [Fδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push Not; rfl + rw [this] + exact (h.prob_abs_empMean_bestAction_sub_actionMean_ge_le hσ2 hs hδ n).trans + (ENNReal.ofReal_le_ofReal (by nlinarith [hδ.le])) + _ = 2 * (u - l) * n ^ 2 * δ := by ring + +lemma integral_sum_range_ucb_action_sub_actionMean_action_le [Nonempty (Fin K)] + [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) - (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (t : ℕ) : + (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hδ : 0 < δ) (n : ℕ) : + P[fun ω ↦ ∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)] ≤ + (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + 2 * K * (u - l) * n ^ 2 * δ := by + have := h.measurable_A + have := h.measurable_E + have := h.measurable_R + set Eδ := {ω | ∀ t < n, ∀ a, pullCount A a t ω ≠ 0 → + |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E a ω| + < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω))} + have hEδ_meas : MeasurableSet Eδ := by measurability + have h_int : Integrable (fun ω ↦ ∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)) P := + integrable_finset_sum _ fun s _ ↦ + (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A s) + measurable_const hlu).sub + (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E (h.measurable_A s) hm) + calc ∫ ω, ∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P + = (∫ ω in Eδ, ∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P) + + ∫ ω in Eδᶜ, ∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P := + (integral_add_compl hEδ_meas h_int).symm + _ ≤ (∫ _ω in Eδ, (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) ∂P) + + ∫ ω in Eδᶜ, ∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P := by + apply add_le_add _ le_rfl + apply setIntegral_mono_on h_int.integrableOn integrableOn_const hEδ_meas + intro ω hω + exact sum_ucb_sub_mean_le (fun a ↦ IsBayesAlgEnvSeq.actionMean κ E a ω) (hm (E ω)) hlu + fun s hs hpc ↦ hω s hs (A s ω) hpc + _ = ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) * P.real Eδ + + ∫ ω in Eδᶜ, ∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P := by + rw [setIntegral_const, smul_eq_mul, mul_comm] + _ ≤ ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) + + ∫ ω in Eδᶜ, ∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - + IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P := by + apply add_le_add _ le_rfl + exact mul_le_of_le_one_right + (by have : 0 ≤ u - l := sub_nonneg.mpr hlu; positivity) measureReal_le_one + _ ≤ ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) + + ∫ _ω in Eδᶜ, n * (u - l) ∂P := by + apply add_le_add le_rfl + apply setIntegral_mono_on h_int.integrableOn integrableOn_const hEδ_meas.compl + intro ω _ + refine le_of_le_of_eq (Finset.sum_le_card_nsmul (range n) _ (u - l) fun s _ ↦ ?_) ?_ + · have h1 : ucb A R' l u σ2 δ (A s ω) s ω ≤ u := (ucb_mem_Icc hlu).2 + have h2 : l ≤ IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω := (hm _ _).1 + linarith + · rw [Finset.card_range, nsmul_eq_mul] + _ = ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) + + n * (u - l) * P.real Eδᶜ := by + rw [setIntegral_const, smul_eq_mul]; ring + _ ≤ ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) + + n * (u - l) * (2 * K * n * δ) := by + gcongr + · nlinarith + · apply ENNReal.toReal_le_of_le_ofReal (by positivity) + have : Eδᶜ = {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ + |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} := by + ext ω; simp only [Eδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push Not; rfl + rw [this] + exact (h.prob_abs_empMean_sub_actionMean_ge_le hσ2 hs hδ n).trans + (ENNReal.ofReal_le_ofReal + (by nlinarith [hδ.le, Nat.cast_nonneg (α := ℝ) K])) + _ = (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + + 2 * K * (u - l) * n ^ 2 * δ := by ring + +lemma integral_regret_le_of_delta_pos [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure Ω} + [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hσ2 : 0 < σ2) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) + (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hδ : 0 < δ) (n : ℕ) : + P[IsBayesAlgEnvSeq.regret κ E A n] ≤ + (u - l) * K + 2 * (K + 1) * (u - l) * n ^ 2 * δ + + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by + calc P[IsBayesAlgEnvSeq.regret κ E A n] + = P[fun ω ↦ ∑ s ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)] + + P[fun ω ↦ ∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)] := + integral_regret_eq_add (hK := hK) (σ2 := σ2) (δ := δ) h hlu hm n + _ ≤ 2 * (u - l) * n ^ 2 * δ + + ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + + 2 * K * (u - l) * n ^ 2 * δ) := + add_le_add + (integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le + (hK := hK) (σ2 := σ2) (δ := δ) h hσ2 hs hlu hm hδ n) + (integral_sum_range_ucb_action_sub_actionMean_action_le + (hK := hK) (σ2 := σ2) (δ := δ) h hσ2 hs hlu hm hδ n) + _ = (u - l) * K + 2 * (K + 1) * (u - l) * n ^ 2 * δ + + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by ring + +lemma integral_regret_le [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure Ω} + [IsProbabilityMeasure P] + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hσ2 : 0 < σ2) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) + (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (t : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A t] ≤ (3 * K + 2) * (u - l) + 8 * √(σ2 * K * t * Real.log t) := by - have ⟨h1, h2⟩ := hm (Classical.arbitrary _) (Classical.arbitrary _) - have hlo : l ≤ u := h1.trans h2 by_cases ht : t = 0 · simp [ht, IsBayesAlgEnvSeq.regret, Bandits.regret] - nlinarith [sub_nonneg.mpr hlo, Nat.cast_pos (α := ℝ).mpr hK, - Real.sqrt_nonneg (↑σ2 * ↑K * (0 : ℝ) * Real.log (0 : ℝ))] + nlinarith by_cases ht1_eq : t = 1 - · subst ht1_eq - simp only [Nat.cast_one, Real.log_one, mul_zero, Real.sqrt_zero, mul_zero, add_zero] - calc P[IsBayesAlgEnvSeq.regret κ E A 1] - ≤ u - l := by + · calc P[IsBayesAlgEnvSeq.regret κ E A t] + = P[IsBayesAlgEnvSeq.regret κ E A 1] := by rw [ht1_eq] + _ ≤ u - l := by rw [IsBayesAlgEnvSeq.regret_eq_sum_gap'] simp only [Finset.range_one, Finset.sum_singleton] exact (integral_mono_of_nonneg (ae_of_all _ fun ω ↦ IsBayesAlgEnvSeq.gap_nonneg_of_le (fun e a ↦ (hm e a).2)) (integrable_const _) (ae_of_all _ fun ω ↦ IsBayesAlgEnvSeq.gap_le_of_mem_Icc hm)).trans (by simp) - _ ≤ (3 * ↑K + 2) * (u - l) := by - nlinarith [Nat.one_le_cast (α := ℝ).mpr (Nat.one_le_of_lt hK), - sub_nonneg.mpr hlo] - -- For t ≥ 2, we have δ = 1/t² < 1 - · have ht2 : 2 ≤ t := by omega - have htpos : (0 : ℝ) < t := by positivity - have _ht1 : (1 : ℝ) ≤ t := by exact_mod_cast Nat.pos_of_ne_zero ht - have hδ : (0 : ℝ) < 1 / (t : ℝ) ^ 2 := by positivity - -- First term: (hi-lo)*K + 2*(K+1)*(hi-lo)*t²*(1/t²) = (3K+2)*(hi-lo) - have h_first : (u - l) * ↑K + 2 * (↑K + 1) * (u - l) * ↑t ^ 2 * (1 / (↑t) ^ 2) - = (3 * ↑K + 2) * (u - l) := by - field_simp; ring - -- Second term simplification: log(1/(1/t²)) = log(t²) = 2 log(t) - have h_log : Real.log (1 / (1 / (↑t : ℝ) ^ 2)) = 2 * Real.log ↑t := by - rw [one_div_one_div, Real.log_pow]; norm_cast - calc P[IsBayesAlgEnvSeq.regret κ E A t] - ≤ (u - l) * ↑K + 2 * (↑K + 1) * (u - l) * ↑t ^ 2 * (1 / (↑t) ^ 2) - + 4 * √(2 * ↑σ2 * Real.log (1 / (1 / (↑t) ^ 2)) * ↑K * ↑t) := - bayesRegret_le_of_delta (δ := 1 / (↑t) ^ 2) hK h hσ2 hs hm hδ t - _ = (3 * ↑K + 2) * (u - l) + 8 * √(↑σ2 * ↑K * ↑t * Real.log ↑t) := by - rw [h_first, h_log]; congr 1 - rw [show (2 : ℝ) * ↑σ2 * (2 * Real.log ↑t) * ↑K * ↑t = - (2 : ℝ) ^ 2 * (↑σ2 * ↑K * ↑t * Real.log ↑t) from by ring, - Real.sqrt_mul (by positivity : (0 : ℝ) ≤ 2 ^ 2), - Real.sqrt_sq (by norm_num : (0 : ℝ) ≤ 2)] + _ ≤ (3 * K + 2) * (u - l) + 8 * √(σ2 * K * t * Real.log t) := by + rw [ht1_eq] + simp only [Nat.cast_one, Real.log_one, mul_zero, Real.sqrt_zero, mul_zero, add_zero] + nlinarith + · calc P[IsBayesAlgEnvSeq.regret κ E A t] + ≤ (u - l) * K + 2 * (K + 1) * (u - l) * t ^ 2 * (1 / (t : ℝ) ^ 2) + + 4 * √(2 * σ2 * Real.log (1 / (1 / (t : ℝ) ^ 2)) * K * t) := + integral_regret_le_of_delta_pos (δ := 1 / t ^ 2) hK h hσ2 hs hlu hm (by positivity) t + _ = (3 * K + 2) * (u - l) + + 4 * √(2 * σ2 * Real.log (1 / (1 / (t : ℝ) ^ 2)) * K * t) := by + congr 1 + field_simp + ring + _ = (3 * K + 2) * (u - l) + 4 * √(2 * σ2 * (2 * Real.log t) * K * t) := by + rw [one_div_one_div, Real.log_pow] + norm_cast + _ = (3 * K + 2) * (u - l) + 4 * √((2 : ℝ) ^ 2 * (σ2 * K * t * Real.log t)) := by + congr 2 + ring_nf + _ = (3 * K + 2) * (u - l) + 4 * (2 * √(σ2 * K * t * Real.log t)) := by + rw [Real.sqrt_mul (by positivity), Real.sqrt_sq (by norm_num)] + _ = (3 * K + 2) * (u - l) + 8 * √(σ2 * K * t * Real.log t) := by ring -end BayesRegret +end IntegralRegret end TS diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index afa9870b..19bb9c94 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -66,6 +66,14 @@ lemma measurable_uncurry_actionMean_comp [Countable α] [MeasurableSingletonClas apply Measurable.comp _ (by fun_prop) exact measurable_from_prod_countable_right (fun _ ↦ measurable_actionMean hE) +lemma integrable_uncurry_actionMean_comp [Countable α] [MeasurableSingletonClass α] + {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} (hE : Measurable E) {f : Ω → α} (hf : Measurable f) + {P : Measure Ω} [IsFiniteMeasure P] {l u : ℝ} (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) : + Integrable (fun ω ↦ actionMean κ E (f ω) ω) P := by + refine ⟨(measurable_uncurry_actionMean_comp hE hf).aestronglyMeasurable, ?_⟩ + apply HasFiniteIntegral.of_bounded + filter_upwards with ω using abs_le_max_abs_abs (hm (E ω) (f ω)).1 (hm (E ω) (f ω)).2 + noncomputable def bestAction [Nonempty α] [Fintype α] [Encodable α] [MeasurableSingletonClass α] (κ : Kernel (𝓔 × α) ℝ) (E : Ω → 𝓔) (ω : Ω) : α := diff --git a/LeanMachineLearning/SequentialLearning/FiniteActions.lean b/LeanMachineLearning/SequentialLearning/FiniteActions.lean index 9734d5ad..0c3a22fb 100644 --- a/LeanMachineLearning/SequentialLearning/FiniteActions.lean +++ b/LeanMachineLearning/SequentialLearning/FiniteActions.lean @@ -898,6 +898,13 @@ lemma measurable_empMean [MeasurableSingletonClass α] {R' : ℕ → Ω → ℝ} unfold empMean fun_prop +@[fun_prop] +lemma measurable_uncurry_empMean_comp [Countable α] [MeasurableSingletonClass α] {R' : ℕ → Ω → ℝ} + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) {f : Ω → α} (hf : Measurable f) + {g : Ω → ℕ} (hg : Measurable g) : Measurable (fun ω ↦ empMean A R' (f ω) (g ω) ω) := by + unfold empMean + fun_prop + @[fun_prop] lemma measurable_sumRewards' [MeasurableSingletonClass α] (n : ℕ) (a : α) : Measurable (fun h ↦ sumRewards' n h a) := by From 498094b90b3c0f2cc66b6f00fae7940c6378c948 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 16 Apr 2026 08:53:14 +0100 Subject: [PATCH 103/155] Refactor TS.lean (in progress) --- LeanMachineLearning/BanditAlgorithms/TS.lean | 72 ++++++++++---------- 1 file changed, 36 insertions(+), 36 deletions(-) diff --git a/LeanMachineLearning/BanditAlgorithms/TS.lean b/LeanMachineLearning/BanditAlgorithms/TS.lean index 3e5f920c..b1843ddb 100644 --- a/LeanMachineLearning/BanditAlgorithms/TS.lean +++ b/LeanMachineLearning/BanditAlgorithms/TS.lean @@ -231,8 +231,10 @@ end UCB section IntegralRegret -lemma integral_ucb_action_eq_integral_ucb_bestAction [Nonempty (Fin K)] - [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] +variable [Nonempty (Fin K)] +variable [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + +lemma integral_ucb_action_eq_integral_ucb_bestAction (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (n : ℕ) : P[fun ω ↦ ucb A R' l u σ2 δ (A n ω) n ω] = P[fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) n ω] := by @@ -259,8 +261,7 @@ lemma integral_ucb_action_eq_integral_ucb_bestAction [Nonempty (Fin K)] _ = P[fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) (t + 1) ω] := by simp_rw [ucb'', ucb_succ_eq_ucb'] -lemma integral_regret_eq_add [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure Ω} - [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) +lemma integral_regret_eq_add (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A n] = P[fun ω ↦ ∑ s ∈ range n, @@ -356,11 +357,11 @@ lemma integral_regret_eq_add [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measur (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E (h.measurable_A s) hm) -lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le [Nonempty (Fin K)] - [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hσ2 : 0 < σ2) +lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) + (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) - (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hδ : 0 < δ) (n : ℕ) : + (hδ : 0 < δ) (n : ℕ) : P[fun ω ↦ ∑ s ∈ range n, (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)] ≤ @@ -435,11 +436,11 @@ lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le [Nonempty ( (ENNReal.ofReal_le_ofReal (by nlinarith [hδ.le])) _ = 2 * (u - l) * n ^ 2 * δ := by ring -lemma integral_sum_range_ucb_action_sub_actionMean_action_le [Nonempty (Fin K)] - [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hσ2 : 0 < σ2) +lemma integral_sum_range_ucb_action_sub_actionMean_action_le + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) + (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) - (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hδ : 0 < δ) (n : ℕ) : + (hδ : 0 < δ) (n : ℕ) : P[fun ω ↦ ∑ s ∈ range n, (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)] ≤ (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + 2 * K * (u - l) * n ^ 2 * δ := by @@ -515,8 +516,8 @@ lemma integral_sum_range_ucb_action_sub_actionMean_action_le [Nonempty (Fin K)] _ = (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + 2 * K * (u - l) * n ^ 2 * δ := by ring -lemma integral_regret_le_of_delta_pos [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure Ω} - [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hσ2 : 0 < σ2) +lemma integral_regret_le_of_delta_pos + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hδ : 0 < δ) (n : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A n] ≤ @@ -534,24 +535,23 @@ lemma integral_regret_le_of_delta_pos [Nonempty (Fin K)] [MeasurableSpace Ω] {P 2 * K * (u - l) * n ^ 2 * δ) := add_le_add (integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le - (hK := hK) (σ2 := σ2) (δ := δ) h hσ2 hs hlu hm hδ n) + (hK := hK) (σ2 := σ2) (δ := δ) h hlu hm hσ2 hs hδ n) (integral_sum_range_ucb_action_sub_actionMean_action_le - (hK := hK) (σ2 := σ2) (δ := δ) h hσ2 hs hlu hm hδ n) + (hK := hK) (σ2 := σ2) (δ := δ) h hlu hm hσ2 hs hδ n) _ = (u - l) * K + 2 * (K + 1) * (u - l) * n ^ 2 * δ + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by ring -lemma integral_regret_le [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure Ω} - [IsProbabilityMeasure P] - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hσ2 : 0 < σ2) - (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) - (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (t : ℕ) : - P[IsBayesAlgEnvSeq.regret κ E A t] - ≤ (3 * K + 2) * (u - l) + 8 * √(σ2 * K * t * Real.log t) := by - by_cases ht : t = 0 +lemma integral_regret_le + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) + (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (n : ℕ) : + P[IsBayesAlgEnvSeq.regret κ E A n] + ≤ (3 * K + 2) * (u - l) + 8 * √(σ2 * K * n * Real.log n) := by + by_cases ht : n = 0 · simp [ht, IsBayesAlgEnvSeq.regret, Bandits.regret] nlinarith - by_cases ht1_eq : t = 1 - · calc P[IsBayesAlgEnvSeq.regret κ E A t] + by_cases ht1_eq : n = 1 + · calc P[IsBayesAlgEnvSeq.regret κ E A n] = P[IsBayesAlgEnvSeq.regret κ E A 1] := by rw [ht1_eq] _ ≤ u - l := by rw [IsBayesAlgEnvSeq.regret_eq_sum_gap'] @@ -560,28 +560,28 @@ lemma integral_regret_le [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure Ω (ae_of_all _ fun ω ↦ IsBayesAlgEnvSeq.gap_nonneg_of_le (fun e a ↦ (hm e a).2)) (integrable_const _) (ae_of_all _ fun ω ↦ IsBayesAlgEnvSeq.gap_le_of_mem_Icc hm)).trans (by simp) - _ ≤ (3 * K + 2) * (u - l) + 8 * √(σ2 * K * t * Real.log t) := by + _ ≤ (3 * K + 2) * (u - l) + 8 * √(σ2 * K * n * Real.log n) := by rw [ht1_eq] simp only [Nat.cast_one, Real.log_one, mul_zero, Real.sqrt_zero, mul_zero, add_zero] nlinarith - · calc P[IsBayesAlgEnvSeq.regret κ E A t] - ≤ (u - l) * K + 2 * (K + 1) * (u - l) * t ^ 2 * (1 / (t : ℝ) ^ 2) - + 4 * √(2 * σ2 * Real.log (1 / (1 / (t : ℝ) ^ 2)) * K * t) := - integral_regret_le_of_delta_pos (δ := 1 / t ^ 2) hK h hσ2 hs hlu hm (by positivity) t + · calc P[IsBayesAlgEnvSeq.regret κ E A n] + ≤ (u - l) * K + 2 * (K + 1) * (u - l) * n ^ 2 * (1 / (n : ℝ) ^ 2) + + 4 * √(2 * σ2 * Real.log (1 / (1 / (n : ℝ) ^ 2)) * K * n) := + integral_regret_le_of_delta_pos (δ := 1 / n ^ 2) hK h hσ2 hs hlu hm (by positivity) n _ = (3 * K + 2) * (u - l) - + 4 * √(2 * σ2 * Real.log (1 / (1 / (t : ℝ) ^ 2)) * K * t) := by + + 4 * √(2 * σ2 * Real.log (1 / (1 / (n : ℝ) ^ 2)) * K * n) := by congr 1 field_simp ring - _ = (3 * K + 2) * (u - l) + 4 * √(2 * σ2 * (2 * Real.log t) * K * t) := by + _ = (3 * K + 2) * (u - l) + 4 * √(2 * σ2 * (2 * Real.log n) * K * n) := by rw [one_div_one_div, Real.log_pow] norm_cast - _ = (3 * K + 2) * (u - l) + 4 * √((2 : ℝ) ^ 2 * (σ2 * K * t * Real.log t)) := by + _ = (3 * K + 2) * (u - l) + 4 * √((2 : ℝ) ^ 2 * (σ2 * K * n * Real.log n)) := by congr 2 ring_nf - _ = (3 * K + 2) * (u - l) + 4 * (2 * √(σ2 * K * t * Real.log t)) := by + _ = (3 * K + 2) * (u - l) + 4 * (2 * √(σ2 * K * n * Real.log n)) := by rw [Real.sqrt_mul (by positivity), Real.sqrt_sq (by norm_num)] - _ = (3 * K + 2) * (u - l) + 8 * √(σ2 * K * t * Real.log t) := by + _ = (3 * K + 2) * (u - l) + 8 * √(σ2 * K * n * Real.log n) := by ring end IntegralRegret From 4402052545d250750409174bd2c4aaedd7821c06 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 16 Apr 2026 10:00:40 +0100 Subject: [PATCH 104/155] Refactor TS.lean (in progress) --- LeanMachineLearning/BanditAlgorithms/TS.lean | 35 ++++++++++---------- 1 file changed, 18 insertions(+), 17 deletions(-) diff --git a/LeanMachineLearning/BanditAlgorithms/TS.lean b/LeanMachineLearning/BanditAlgorithms/TS.lean index b1843ddb..3e4e2c68 100644 --- a/LeanMachineLearning/BanditAlgorithms/TS.lean +++ b/LeanMachineLearning/BanditAlgorithms/TS.lean @@ -91,6 +91,10 @@ def ucb (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) if pullCount A a n ω = 0 then u else max l (min u (empMean A R' a n ω + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a n ω)))) +@[simp] +lemma ucb_zero {a : Fin K} {ω : Ω} : ucb A R' l u σ2 δ a 0 ω = u := by + simp [ucb] + lemma ucb_mem_Icc (h : l ≤ u) {a : Fin K} {n : ℕ} {ω : Ω} : ucb A R' l u σ2 δ a n ω ∈ Set.Icc l u := by unfold ucb @@ -241,25 +245,22 @@ lemma integral_ucb_action_eq_integral_ucb_bestAction have := h.measurable_A have := h.measurable_E have := h.measurable_R - by_cases hs : n = 0 - · simp [hs, ucb, pullCount_zero] - obtain ⟨t, rfl⟩ := Nat.exists_eq_succ_of_ne_zero hs - let ucb'' : (Iic t → Fin K × ℝ) × Fin K → ℝ := fun p ↦ ucb' t p.1 l u σ2 δ p.2 - calc P[fun ω ↦ ucb A R' l u σ2 δ (A (t + 1) ω) (t + 1) ω] - = P[fun ω ↦ ucb'' (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω)] := by - simp_rw [ucb'', ucb_succ_eq_ucb'] - _ = ∫ p, ucb'' p ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, A (t + 1) ω)) := by + by_cases hn : n = 0 + · simp [hn] + obtain ⟨n, rfl⟩ := Nat.exists_eq_succ_of_ne_zero hn + let u' (ha : (Iic n → Fin K × ℝ) × Fin K) := ucb' n ha.1 l u σ2 δ ha.2 + calc + _ = P[fun ω ↦ u' (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)] := by + simp_rw [u', ucb_succ_eq_ucb'] + _ = ∫ ha, u' ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) := by rw [← integral_map (by fun_prop) (by fun_prop)] - _ = ∫ p, ucb'' p ∂P.map - (fun ω ↦ (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E ω)) := by - rw [← compProd_map_condDistrib (hY := (h.measurable_A (t + 1)).aemeasurable), - ← compProd_map_condDistrib - (hY := (by fun_prop : AEMeasurable (IsBayesAlgEnvSeq.bestAction κ E) _)), - Measure.compProd_congr (hasCondDistrib_action hK h t).condDistrib_eq] - _ = P[fun ω ↦ ucb'' (IsAlgEnvSeq.hist A R' t ω, IsBayesAlgEnvSeq.bestAction κ E ω)] := by + _ = ∫ ha, u' ha ∂P.map + (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, IsBayesAlgEnvSeq.bestAction κ E ω)) := by + rw [← compProd_map_condDistrib (by fun_prop), ← compProd_map_condDistrib (by fun_prop), + Measure.compProd_congr (hasCondDistrib_action hK h n).condDistrib_eq] + _ = P[fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) (n + 1) ω] := by rw [integral_map (by fun_prop) (by fun_prop)] - _ = P[fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) (t + 1) ω] := by - simp_rw [ucb'', ucb_succ_eq_ucb'] + simp_rw [u', ucb_succ_eq_ucb'] lemma integral_regret_eq_add (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) : From 0e6780803b760c2463a8c682a595d0cb8131f2b9 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 16 Apr 2026 16:42:13 +0100 Subject: [PATCH 105/155] Refactor TS.lean (in progress) --- LeanMachineLearning/BanditAlgorithms/TS.lean | 154 +++++++------------ 1 file changed, 53 insertions(+), 101 deletions(-) diff --git a/LeanMachineLearning/BanditAlgorithms/TS.lean b/LeanMachineLearning/BanditAlgorithms/TS.lean index 3e4e2c68..6cd5e5f3 100644 --- a/LeanMachineLearning/BanditAlgorithms/TS.lean +++ b/LeanMachineLearning/BanditAlgorithms/TS.lean @@ -117,13 +117,17 @@ lemma measurable_uncurry_ucb_comp [MeasurableSpace Ω] (hA : ∀ t, Measurable ( apply Measurable.comp _ (by fun_prop) exact measurable_from_prod_countable_right (fun _ ↦ measurable_ucb hA hR) +@[fun_prop] lemma integrable_uncurry_ucb_comp [MeasurableSpace Ω] (hA : ∀ t, Measurable (A t)) (hR : ∀ t, Measurable (R' t)) {f : Ω → Fin K} (hf : Measurable f) {g : Ω → ℕ} - (hg : Measurable g) {P : Measure Ω} [IsFiniteMeasure P] (hlu : l ≤ u) : + (hg : Measurable g) {P : Measure Ω} [IsFiniteMeasure P] : Integrable (fun ω ↦ ucb A R' l u σ2 δ (f ω) (g ω) ω) P := by refine ⟨(measurable_uncurry_ucb_comp hA hR hf hg).aestronglyMeasurable, ?_⟩ - apply HasFiniteIntegral.of_bounded - filter_upwards with ω using abs_le_max_abs_abs (ucb_mem_Icc hlu).1 (ucb_mem_Icc hlu).2 + apply HasFiniteIntegral.of_bounded (C := max |l| |u|) + filter_upwards with ω + rw [Real.norm_eq_abs] + unfold ucb + grind noncomputable def ucb' (n : ℕ) (h : Iic n → Fin K × ℝ) (l u σ2 δ : ℝ) (a : Fin K) : ℝ := @@ -263,100 +267,48 @@ lemma integral_ucb_action_eq_integral_ucb_bestAction simp_rw [u', ucb_succ_eq_ucb'] lemma integral_regret_eq_add (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) - (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) : + (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A n] = - P[fun ω ↦ ∑ s ∈ range n, + P[fun ω ↦ ∑ t ∈ range n, (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)] + - P[fun ω ↦ ∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)] := by - have := h.measurable_A - have := h.measurable_E - have := h.measurable_R - calc P[IsBayesAlgEnvSeq.regret κ E A n] - = ∫ ω, ∑ s ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P := by + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω)] + + P[fun ω ↦ ∑ t ∈ range n, + (ucb A R' l u σ2 δ (A t ω) t ω - IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω)] := by + have hua (t : ℕ) : Integrable (fun ω ↦ ucb A R' l u σ2 δ (A t ω) t ω) P := + integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A t) measurable_const + have hub (t : ℕ) : + Integrable (fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω) P := + integrable_uncurry_ucb_comp h.measurable_A h.measurable_R + (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E) measurable_const + have haa (t : ℕ) : Integrable (fun ω ↦ IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω) P := + IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E (h.measurable_A t) hm + have hab : + Integrable (fun ω ↦ IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω) P := + IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E + (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E) hm + calc + _ = (∑ t ∈ range n, + ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ∂P) - + ∑ t ∈ range n, ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω ∂P := by simp_rw [IsBayesAlgEnvSeq.regret_eq_sum_gap, IsBayesAlgEnvSeq.gap_eq_sub] - _ = (∫ ω, ∑ s ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P) + 0 := - (add_zero _).symm - _ = (∫ ω, ∑ s ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P) + - ∫ ω, ∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω) ∂P := by - congr 1 - symm - rw [integral_finset_sum _ ?_] - · apply Finset.sum_eq_zero - intro s _ - rw [integral_sub ?_ ?_, - integral_ucb_action_eq_integral_ucb_bestAction (hK := hK) h s, - sub_self] - · exact integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A s) - measurable_const hlu - · exact integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (by fun_prop) - measurable_const hlu - · intro s _ - exact (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A s) - measurable_const hlu).sub - (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (by fun_prop) - measurable_const hlu) - _ = ∫ ω, - ((∑ s ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)) + - (∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω))) ∂P := by - rw [← integral_add] - · apply integrable_finset_sum - intro s _ - exact (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E - (by fun_prop) hm).sub - (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E - (h.measurable_A s) hm) - · apply integrable_finset_sum - intro s _ - exact (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A s) - measurable_const hlu).sub - (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (by fun_prop) - measurable_const hlu) - _ = ∫ ω, - ((∑ s ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)) + - (∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω))) ∂P := by - congr 1 - ext ω - simp only [← Finset.sum_add_distrib] - apply Finset.sum_congr rfl - intros - ring - _ = (∫ ω, ∑ s ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω) ∂P) + - ∫ ω, ∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P := by - apply integral_add - · apply integrable_finset_sum - intro s _ - exact (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E - (by fun_prop) hm).sub - (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (by fun_prop) - measurable_const hlu) - · apply integrable_finset_sum - intro s _ - exact (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A s) - measurable_const hlu).sub - (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E - (h.measurable_A s) hm) + rw [integral_finset_sum _ (by fun_prop), ← Finset.sum_sub_distrib] + simp_rw [integral_sub hab (haa _)] + _ = ((∑ t ∈ range n, + ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ∂P) - + ∑ t ∈ range n, + ∫ ω, ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω ∂P) + + ((∑ t ∈ range n, ∫ ω, ucb A R' l u σ2 δ (A t ω) t ω ∂P) - + ∑ t ∈ range n, ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω ∂P) := by + simp [integral_ucb_action_eq_integral_ucb_bestAction hK h] + _ = (∑ t ∈ range n, + ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω ∂P) + + ∑ t ∈ range n, ∫ ω, ucb A R' l u σ2 δ (A t ω) t ω - + IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω ∂P := by + rw [← Finset.sum_sub_distrib, ← Finset.sum_sub_distrib] + simp_rw [← integral_sub hab (hub _), ← integral_sub (hua _) (haa _)] + _ = _ := by + rw [← integral_finset_sum _ (by fun_prop), ← integral_finset_sum _ (by fun_prop)] lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) @@ -381,10 +333,10 @@ lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le integrable_finset_sum _ fun s _ ↦ (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E (by fun_prop) hm).sub (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (by fun_prop) - measurable_const hlu) - calc ∫ ω, ∑ s ∈ range n, + measurable_const) + calc P[fun ω ↦ ∑ s ∈ range n, (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω) ∂P + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)] = (∫ ω in Fδ, ∑ s ∈ range n, (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω) ∂P) + @@ -456,10 +408,10 @@ lemma integral_sum_range_ucb_action_sub_actionMean_action_le (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)) P := integrable_finset_sum _ fun s _ ↦ (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A s) - measurable_const hlu).sub + measurable_const).sub (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E (h.measurable_A s) hm) - calc ∫ ω, ∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P + calc P[fun ω ↦ ∑ s ∈ range n, + (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)] = (∫ ω in Eδ, ∑ s ∈ range n, (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P) + @@ -530,7 +482,7 @@ lemma integral_regret_le_of_delta_pos ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)] + P[fun ω ↦ ∑ s ∈ range n, (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)] := - integral_regret_eq_add (hK := hK) (σ2 := σ2) (δ := δ) h hlu hm n + integral_regret_eq_add (hK := hK) (σ2 := σ2) (δ := δ) h hm n _ ≤ 2 * (u - l) * n ^ 2 * δ + ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + 2 * K * (u - l) * n ^ 2 * δ) := From a7e4a7b82e4b20dd351f4834148884ebece59d0b Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 21 Apr 2026 16:58:29 +0100 Subject: [PATCH 106/155] Refactor TS.lean (in progress) --- .../Online/Bandit/Algorithms/TS.lean | 337 +++++++----------- 1 file changed, 126 insertions(+), 211 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index c652d165..fe25ef5f 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -172,7 +172,7 @@ private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 have : √n * √n = n := Real.mul_self_sqrt (by positivity) nlinarith -/-- This bound could be improved slightly. -/ +/-- This bound could be slightly improved. -/ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, μ a ∈ Set.Icc l u) (hi : l ≤ u) (hc : ∀ s < n, pullCount A (A s ω) s ω ≠ 0 → |empMean A R' (A s ω) s ω - μ (A s ω)| < √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω))) : @@ -213,7 +213,7 @@ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, _ = 2 * √(2 * σ2 * Real.log (1 / δ)) * ∑ a, ∑ j ∈ range (pullCount A a n ω), (1 / √j) := by rw [sum_comp_pullCount (fun j => 1 / √j)] - _ ≤ 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * ∑ a, √(pullCount A a n ω)) := by + _ ≤ 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * ∑ a, √(pullCount A a n ω)) := by -- loose rw [mul_sum _ _ 2] gcongr with a by_cases ha : pullCount A a n ω = 0 @@ -252,19 +252,19 @@ lemma integral_ucb_action_eq_integral_ucb_bestAction by_cases hn : n = 0 · simp [hn] obtain ⟨n, rfl⟩ := Nat.exists_eq_succ_of_ne_zero hn - let u' (ha : (Iic n → Fin K × ℝ) × Fin K) := ucb' n ha.1 l u σ2 δ ha.2 + let uc (ha : (Iic n → Fin K × ℝ) × Fin K) := ucb' n ha.1 l u σ2 δ ha.2 calc - _ = P[fun ω ↦ u' (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)] := by - simp_rw [u', ucb_succ_eq_ucb'] - _ = ∫ ha, u' ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) := by + _ = P[fun ω ↦ uc (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)] := by + simp_rw [uc, ucb_succ_eq_ucb'] + _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) := by rw [← integral_map (by fun_prop) (by fun_prop)] - _ = ∫ ha, u' ha ∂P.map + _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, IsBayesAlgEnvSeq.bestAction κ E ω)) := by rw [← compProd_map_condDistrib (by fun_prop), ← compProd_map_condDistrib (by fun_prop), Measure.compProd_congr (hasCondDistrib_action hK h n).condDistrib_eq] _ = P[fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) (n + 1) ω] := by rw [integral_map (by fun_prop) (by fun_prop)] - simp_rw [u', ucb_succ_eq_ucb'] + simp_rw [uc, ucb_succ_eq_ucb'] lemma integral_regret_eq_add (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) : @@ -310,232 +310,147 @@ lemma integral_regret_eq_add (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E _ = _ := by rw [← integral_finset_sum _ (by fun_prop), ← integral_finset_sum _ (by fun_prop)] -lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) +omit [Nonempty 𝓔] [StandardBorelSpace 𝓔] [IsProbabilityMeasure Q] in +/-- This bound could be improved by using a one-sided concentration inequality. -/ +lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algorithm (Fin K) ℝ} + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (hδ : 0 < δ) (n : ℕ) : - P[fun ω ↦ ∑ s ∈ range n, + P[fun ω ↦ ∑ t ∈ range n, (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)] ≤ - 2 * (u - l) * n ^ 2 * δ := by + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω)] ≤ + 2 * (u - l) * (n - 1) * n * δ := by + by_cases hn : n = 0 + · simp [hn] + let F := {ω | ∃ t < n, pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) t ω ≠ 0 ∧ + √(2 * σ2 * Real.log (1 / δ) / (pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) t ω)) ≤ + |empMean A R' (IsBayesAlgEnvSeq.bestAction κ E ω) t ω - + IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω|} have := h.measurable_A have := h.measurable_E have := h.measurable_R - set Fδ := {ω | ∀ t < n, pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) t ω ≠ 0 → - |empMean A R' (IsBayesAlgEnvSeq.bestAction κ E ω) t ω - - IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω| - < √(2 * σ2 * Real.log (1 / δ) / (pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) t ω))} - have hFδ_meas : MeasurableSet Fδ := by measurability - have h_int : Integrable (fun ω ↦ ∑ s ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)) P := - integrable_finset_sum _ fun s _ ↦ - (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E (by fun_prop) hm).sub - (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (by fun_prop) - measurable_const) - calc P[fun ω ↦ ∑ s ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)] - = (∫ ω in Fδ, ∑ s ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω) ∂P) + - ∫ ω in Fδᶜ, ∑ s ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω) ∂P := - (integral_add_compl hFδ_meas h_int).symm - _ ≤ 0 + ∫ ω in Fδᶜ, ∑ s ∈ range n, + have hF : MeasurableSet F := by measurability + have : + Integrable (fun ω ↦ IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω) P := + IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp (by fun_prop) (by fun_prop) hm + calc + _ ≤ ∫ ω in F, ∑ t ∈ range n, (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω) ∂P := by - gcongr - apply setIntegral_nonpos hFδ_meas + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω) ∂P := by + rw [← integral_add_compl hF (by fun_prop)] + apply add_le_of_nonpos_right + apply setIntegral_nonpos hF.compl intro ω hω - apply Finset.sum_nonpos - intro s hs - have : IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ≤ - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω := by - unfold ucb - split_ifs with h0 - · exact (hm (E ω) (IsBayesAlgEnvSeq.bestAction κ E ω)).2 - · have := abs_lt.mp ((hω s (mem_range.mp hs)) h0) - exact le_max_of_le_right - (le_min (hm (E ω) (IsBayesAlgEnvSeq.bestAction κ E ω)).2 (by linarith)) - linarith - _ ≤ 0 + ∫ _ω in Fδᶜ, n * (u - l) ∂P := by - apply add_le_add le_rfl - apply setIntegral_mono_on h_int.integrableOn integrableOn_const hFδ_meas.compl - intro ω _ - refine le_of_le_of_eq (Finset.sum_le_card_nsmul (range n) _ (u - l) fun s _ ↦ ?_) ?_ - · have h1 : IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ≤ u := - (hm _ _).2 - have h2 : l ≤ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω := - (ucb_mem_Icc hlu).1 - linarith - · rw [Finset.card_range, nsmul_eq_mul] - _ = 0 + n * (u - l) * P.real Fδᶜ := by - rw [setIntegral_const, smul_eq_mul, mul_comm] - _ ≤ 0 + n * (u - l) * (2 * n * δ) := by + apply sum_nonpos + intro t ht + rw [Set.mem_compl_iff, Set.mem_setOf_eq] at hω + push Not at hω + grind [hω t (mem_range.mp ht), ucb, IsBayesAlgEnvSeq.actionMean] + _ ≤ ∫ ω in F, ∑ t ∈ range n, (u - l) ∂P := by + apply setIntegral_mono_on (Integrable.integrableOn (by fun_prop)) + (Integrable.integrableOn (by fun_prop)) hF + intro ω hω + apply sum_le_sum + intro t ht + grind [IsBayesAlgEnvSeq.actionMean, ucb] + _ = P.real F * (n * (u - l)) := by + simp_rw [sum_const, card_range, nsmul_eq_mul, setIntegral_const, smul_eq_mul] + _ ≤ (2 * (n - 1) * δ) * (n * (u - l)) := by gcongr · nlinarith - · apply ENNReal.toReal_le_of_le_ofReal (by positivity) - have : Fδᶜ = {ω | ∃ s < n, pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) s ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / - (pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) s ω : ℝ)) ≤ - |empMean A R' (IsBayesAlgEnvSeq.bestAction κ E ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω|} := by - ext ω; simp only [Fδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push Not; rfl - rw [this] - exact (h.prob_abs_empMean_bestAction_sub_actionMean_ge_le hσ2 hs hδ n).trans - (ENNReal.ofReal_le_ofReal (by nlinarith [hδ.le])) - _ = 2 * (u - l) * n ^ 2 * δ := by ring - -lemma integral_sum_range_ucb_action_sub_actionMean_action_le - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) + · have : (1 : ℝ) ≤ n := by simp [Nat.one_le_iff_ne_zero, hn] + apply ENNReal.toReal_le_of_le_ofReal (by nlinarith) + exact h.prob_abs_empMean_bestAction_sub_actionMean_ge_le hσ2 hs hδ n + _ = _ := by + ring + +omit [Nonempty 𝓔] [StandardBorelSpace 𝓔] [IsProbabilityMeasure Q] in +/-- This bound could be improved by using a one-sided concentration inequality. -/ +lemma integral_sum_range_ucb_action_sub_actionMean_action_le {alg : Algorithm (Fin K) ℝ} + (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (hδ : 0 < δ) (n : ℕ) : - P[fun ω ↦ ∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)] ≤ - (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + 2 * K * (u - l) * n ^ 2 * δ := by + P[fun ω ↦ ∑ t ∈ range n, + (ucb A R' l u σ2 δ (A t ω) t ω - IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω)] ≤ + (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + + 2 * (u - l) * K * (n - 1) * n * δ := by + by_cases hn : n = 0 + · simp [hn, hlu, mul_nonneg] + let F := {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ + √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ + |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} have := h.measurable_A have := h.measurable_E have := h.measurable_R - set Eδ := {ω | ∀ t < n, ∀ a, pullCount A a t ω ≠ 0 → - |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E a ω| - < √(2 * σ2 * Real.log (1 / δ) / (pullCount A a t ω))} - have hEδ_meas : MeasurableSet Eδ := by measurability - have h_int : Integrable (fun ω ↦ ∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)) P := - integrable_finset_sum _ fun s _ ↦ - (integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A s) - measurable_const).sub - (IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E (h.measurable_A s) hm) - calc P[fun ω ↦ ∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)] - = (∫ ω in Eδ, ∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P) + - ∫ ω in Eδᶜ, ∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P := - (integral_add_compl hEδ_meas h_int).symm - _ ≤ (∫ _ω in Eδ, (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) ∂P) + - ∫ ω in Eδᶜ, ∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P := by - apply add_le_add _ le_rfl - apply setIntegral_mono_on h_int.integrableOn integrableOn_const hEδ_meas - intro ω hω - exact sum_ucb_sub_mean_le (fun a ↦ IsBayesAlgEnvSeq.actionMean κ E a ω) (hm (E ω)) hlu - fun s hs hpc ↦ hω s hs (A s ω) hpc - _ = ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) * P.real Eδ + - ∫ ω in Eδᶜ, ∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P := by - rw [setIntegral_const, smul_eq_mul, mul_comm] - _ ≤ ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) + - ∫ ω in Eδᶜ, ∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω) ∂P := by - apply add_le_add _ le_rfl - exact mul_le_of_le_one_right - (by have : 0 ≤ u - l := sub_nonneg.mpr hlu; positivity) measureReal_le_one - _ ≤ ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) + - ∫ _ω in Eδᶜ, n * (u - l) ∂P := by - apply add_le_add le_rfl - apply setIntegral_mono_on h_int.integrableOn integrableOn_const hEδ_meas.compl - intro ω _ - refine le_of_le_of_eq (Finset.sum_le_card_nsmul (range n) _ (u - l) fun s _ ↦ ?_) ?_ - · have h1 : ucb A R' l u σ2 δ (A s ω) s ω ≤ u := (ucb_mem_Icc hlu).2 - have h2 : l ≤ IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω := (hm _ _).1 - linarith - · rw [Finset.card_range, nsmul_eq_mul] - _ = ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) + - n * (u - l) * P.real Eδᶜ := by - rw [setIntegral_const, smul_eq_mul]; ring - _ ≤ ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) + - n * (u - l) * (2 * K * n * δ) := by + have hF : MeasurableSet F := by measurability + have : ∀ t, Integrable (fun ω ↦ IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω) P := + fun t ↦ IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp (by fun_prop) (by fun_prop) hm + calc + _ ≤ (∫ ω in F, ∑ t ∈ range n, (u - l) ∂P) + + ∫ ω in Fᶜ, (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) ∂P := by + rw [← integral_add_compl hF (by fun_prop)] + apply add_le_add + · apply setIntegral_mono_on (Integrable.integrableOn (by fun_prop)) + (Integrable.integrableOn (by fun_prop)) hF + intro ω hω + apply sum_le_sum + intro t ht + grind [ucb, IsBayesAlgEnvSeq.actionMean] + · apply setIntegral_mono_on (Integrable.integrableOn (by fun_prop)) + (Integrable.integrableOn (by fun_prop)) hF.compl + intro ω hω + rw [Set.mem_compl_iff, Set.mem_setOf_eq] at hω + push Not at hω + exact sum_ucb_sub_mean_le (fun a ↦ (κ (E ω, a))[id]) (hm (E ω)) hlu + (fun t ht hpc ↦ hω t ht (A t ω) hpc) + _ = P.real F * (n * (u - l)) + + P.real Fᶜ * ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) := by + simp_rw [sum_const, card_range, nsmul_eq_mul, setIntegral_const, smul_eq_mul] + _ ≤ (2 * K * (n - 1) * δ) * (n * (u - l)) + + 1 * ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) := by + have : 0 ≤ u - l := sub_nonneg.2 hlu gcongr - · nlinarith - · apply ENNReal.toReal_le_of_le_ofReal (by positivity) - have : Eδᶜ = {ω | ∃ s < n, ∃ a, pullCount A a s ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / (pullCount A a s ω : ℝ)) ≤ - |empMean A R' a s ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} := by - ext ω; simp only [Eδ, Set.mem_compl_iff, Set.mem_setOf_eq]; push Not; rfl - rw [this] - exact (h.prob_abs_empMean_sub_actionMean_ge_le hσ2 hs hδ n).trans - (ENNReal.ofReal_le_ofReal - (by nlinarith [hδ.le, Nat.cast_nonneg (α := ℝ) K])) - _ = (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + - 2 * K * (u - l) * n ^ 2 * δ := by ring - -lemma integral_regret_le_of_delta_pos - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hσ2 : 0 < σ2) - (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) - (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hδ : 0 < δ) (n : ℕ) : - P[IsBayesAlgEnvSeq.regret κ E A n] ≤ - (u - l) * K + 2 * (K + 1) * (u - l) * n ^ 2 * δ + - 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by - calc P[IsBayesAlgEnvSeq.regret κ E A n] - = P[fun ω ↦ ∑ s ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) s ω)] + - P[fun ω ↦ ∑ s ∈ range n, - (ucb A R' l u σ2 δ (A s ω) s ω - IsBayesAlgEnvSeq.actionMean κ E (A s ω) ω)] := - integral_regret_eq_add (hK := hK) (σ2 := σ2) (δ := δ) h hm n - _ ≤ 2 * (u - l) * n ^ 2 * δ + - ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + - 2 * K * (u - l) * n ^ 2 * δ) := - add_le_add - (integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le - (hK := hK) (σ2 := σ2) (δ := δ) h hlu hm hσ2 hs hδ n) - (integral_sum_range_ucb_action_sub_actionMean_action_le - (hK := hK) (σ2 := σ2) (δ := δ) h hlu hm hσ2 hs hδ n) - _ = (u - l) * K + 2 * (K + 1) * (u - l) * n ^ 2 * δ + - 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by ring - -lemma integral_regret_le - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) + · have : (0 : ℝ) ≤ n - 1 := by simp [Nat.one_le_iff_ne_zero, hn] + apply ENNReal.toReal_le_of_le_ofReal (by positivity) + exact h.prob_abs_empMean_sub_actionMean_ge_le hσ2 hs hδ n + · exact measureReal_le_one + _ = _ := by + ring + +/-- This bound could be improved. -/ +lemma integral_regret_le (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (n : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A n] ≤ (3 * K + 2) * (u - l) + 8 * √(σ2 * K * n * Real.log n) := by - by_cases ht : n = 0 - · simp [ht, IsBayesAlgEnvSeq.regret, Bandits.regret] + by_cases hn : n = 0 + · simp [hn, IsBayesAlgEnvSeq.regret, Bandits.regret] nlinarith - by_cases ht1_eq : n = 1 - · calc P[IsBayesAlgEnvSeq.regret κ E A n] - = P[IsBayesAlgEnvSeq.regret κ E A 1] := by rw [ht1_eq] - _ ≤ u - l := by - rw [IsBayesAlgEnvSeq.regret_eq_sum_gap'] - simp only [Finset.range_one, Finset.sum_singleton] - exact (integral_mono_of_nonneg - (ae_of_all _ fun ω ↦ IsBayesAlgEnvSeq.gap_nonneg_of_le (fun e a ↦ (hm e a).2)) - (integrable_const _) - (ae_of_all _ fun ω ↦ IsBayesAlgEnvSeq.gap_le_of_mem_Icc hm)).trans (by simp) - _ ≤ (3 * K + 2) * (u - l) + 8 * √(σ2 * K * n * Real.log n) := by - rw [ht1_eq] - simp only [Nat.cast_one, Real.log_one, mul_zero, Real.sqrt_zero, mul_zero, add_zero] - nlinarith - · calc P[IsBayesAlgEnvSeq.regret κ E A n] - ≤ (u - l) * K + 2 * (K + 1) * (u - l) * n ^ 2 * (1 / (n : ℝ) ^ 2) - + 4 * √(2 * σ2 * Real.log (1 / (1 / (n : ℝ) ^ 2)) * K * n) := - integral_regret_le_of_delta_pos (δ := 1 / n ^ 2) hK h hσ2 hs hlu hm (by positivity) n - _ = (3 * K + 2) * (u - l) - + 4 * √(2 * σ2 * Real.log (1 / (1 / (n : ℝ) ^ 2)) * K * n) := by - congr 1 - field_simp - ring - _ = (3 * K + 2) * (u - l) + 4 * √(2 * σ2 * (2 * Real.log n) * K * n) := by - rw [one_div_one_div, Real.log_pow] - norm_cast - _ = (3 * K + 2) * (u - l) + 4 * √((2 : ℝ) ^ 2 * (σ2 * K * n * Real.log n)) := by - congr 2 - ring_nf - _ = (3 * K + 2) * (u - l) + 4 * (2 * √(σ2 * K * n * Real.log n)) := by - rw [Real.sqrt_mul (by positivity), Real.sqrt_sq (by norm_num)] - _ = (3 * K + 2) * (u - l) + 8 * √(σ2 * K * n * Real.log n) := by - ring + have hδ : (0 : ℝ) < 1 / n ^ 2 := by positivity + calc P[IsBayesAlgEnvSeq.regret κ E A n] + = _ := + integral_regret_eq_add hK h hm n + _ ≤ _ := + add_le_add + (integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le h hlu hm hσ2 hs hδ n) + (integral_sum_range_ucb_action_sub_actionMean_action_le h hlu hm hσ2 hs hδ n) + _ = K * (u - l) + 2 * (K + 1) * (u - l) * ((n - 1) / n) + + 4 * √((2 : ℝ) ^ 2 * (σ2 * K * n * Real.log n)) := by + field_simp + rw [Real.log_pow] + ring_nf + _ = K * (u - l) + 2 * (K + 1) * (u - l) * ((n - 1) / n) + 8 * √(σ2 * K * n * Real.log n) := by + rw [Real.sqrt_mul (by positivity), Real.sqrt_sq (by norm_num)] + ring + _ ≤ K * (u - l) + 2 * (K + 1) * (u - l) * 1 + 8 * √(σ2 * K * n * Real.log n) := by -- loose + have : 0 ≤ u - l := sub_nonneg.2 hlu + gcongr + rw [div_le_one (by positivity)] + linarith + _ = _ := by + ring end IntegralRegret From 2bbd57eef36c0bf3b2534f4b4e7dac83ecdcae87 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 23 Apr 2026 11:16:27 +0100 Subject: [PATCH 107/155] Refactor TS.lean --- .../Online/Bandit/Algorithms/TS.lean | 218 ++++++++++-------- 1 file changed, 118 insertions(+), 100 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index fe25ef5f..2c5a2553 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -5,8 +5,8 @@ Authors: Rémy Degenne, Paulo Rauber -/ module -public import LeanMachineLearning.SequentialLearning.Algorithms.Uniform public import LeanMachineLearning.SequentialLearning.AlgorithmDensity +public import LeanMachineLearning.SequentialLearning.Algorithms.Uniform /-! # The Thompson Sampling Algorithm -/ @@ -18,49 +18,50 @@ open scoped NNReal namespace Bandits +section Algorithm + variable {K : ℕ} variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] -namespace TS - noncomputable -def policy (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] - (hK : 0 < K) (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := +def TS.policy (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) + [IsMarkovKernel κ] (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) n).map (IsBayesAlgEnvSeq.bestAction κ id) -instance {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] - {hK : 0 < K} (n : ℕ) : IsMarkovKernel (policy Q κ hK n) := +instance {hK : 0 < K} {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} + [IsMarkovKernel κ] {n : ℕ} : IsMarkovKernel (TS.policy hK Q κ n) := Kernel.IsMarkovKernel.map _ (by fun_prop) noncomputable -def initialPolicy (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) - [IsMarkovKernel κ] (hK : 0 < K) : Measure (Fin K) := +def TS.initialPolicy (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] + (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] : Measure (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK Q.map (IsBayesAlgEnvSeq.bestAction κ id) -instance {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] - {hK : 0 < K} : IsProbabilityMeasure (initialPolicy Q κ hK) := +instance {hK : 0 < K} {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} + [IsMarkovKernel κ] : IsProbabilityMeasure (TS.initialPolicy hK Q κ) := Measure.isProbabilityMeasure_map (by fun_prop) -end TS - noncomputable -def tsAlgorithm (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) - [IsMarkovKernel κ] (hK : 0 < K) : Algorithm (Fin K) ℝ where - policy := TS.policy Q κ hK - p0 := TS.initialPolicy Q κ hK +def tsAlgorithm (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) + [IsMarkovKernel κ] : Algorithm (Fin K) ℝ where + policy := TS.policy hK Q κ + p0 := TS.initialPolicy hK Q κ + +end Algorithm namespace TS -variable (hK : 0 < K) -variable {Ω : Type*} +variable {K : ℕ} [Nonempty (Fin K)] +variable {Ω : Type*} [MeasurableSpace Ω] +variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] +variable {P : Measure Ω} [IsProbabilityMeasure P] -lemma hasCondDistrib_action [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure Ω} - [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (n : ℕ) : - HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) +lemma hasCondDistrib_action (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) + (n : ℕ) : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) (condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' n) P) P where aemeasurable_fst := (h.measurable_A (n + 1)).aemeasurable aemeasurable_snd := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable @@ -82,9 +83,12 @@ lemma hasCondDistrib_action [Nonempty (Fin K)] [MeasurableSpace Ω] {P : Measure condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' n) P := (condDistrib_comp (IsAlgEnvSeq.hist A R' n) h.measurable_E.aemeasurable hm).symm -variable {l u σ2 δ : ℝ} +end TS + +namespace ClippedUCB -section UCB +variable {K : ℕ} {l u σ2 δ : ℝ} +variable {Ω : Type*} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} noncomputable def ucb (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ := @@ -235,82 +239,13 @@ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, rw [← Real.sqrt_mul' _ (by positivity)] ring_nf -end UCB - -section IntegralRegret - variable [Nonempty (Fin K)] -variable [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] - -lemma integral_ucb_action_eq_integral_ucb_bestAction - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) (n : ℕ) : - P[fun ω ↦ ucb A R' l u σ2 δ (A n ω) n ω] = - P[fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) n ω] := by - have := h.measurable_A - have := h.measurable_E - have := h.measurable_R - by_cases hn : n = 0 - · simp [hn] - obtain ⟨n, rfl⟩ := Nat.exists_eq_succ_of_ne_zero hn - let uc (ha : (Iic n → Fin K × ℝ) × Fin K) := ucb' n ha.1 l u σ2 δ ha.2 - calc - _ = P[fun ω ↦ uc (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)] := by - simp_rw [uc, ucb_succ_eq_ucb'] - _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) := by - rw [← integral_map (by fun_prop) (by fun_prop)] - _ = ∫ ha, uc ha ∂P.map - (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, IsBayesAlgEnvSeq.bestAction κ E ω)) := by - rw [← compProd_map_condDistrib (by fun_prop), ← compProd_map_condDistrib (by fun_prop), - Measure.compProd_congr (hasCondDistrib_action hK h n).condDistrib_eq] - _ = P[fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) (n + 1) ω] := by - rw [integral_map (by fun_prop) (by fun_prop)] - simp_rw [uc, ucb_succ_eq_ucb'] - -lemma integral_regret_eq_add (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) - (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) : - P[IsBayesAlgEnvSeq.regret κ E A n] = - P[fun ω ↦ ∑ t ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω)] + - P[fun ω ↦ ∑ t ∈ range n, - (ucb A R' l u σ2 δ (A t ω) t ω - IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω)] := by - have hua (t : ℕ) : Integrable (fun ω ↦ ucb A R' l u σ2 δ (A t ω) t ω) P := - integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A t) measurable_const - have hub (t : ℕ) : - Integrable (fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω) P := - integrable_uncurry_ucb_comp h.measurable_A h.measurable_R - (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E) measurable_const - have haa (t : ℕ) : Integrable (fun ω ↦ IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω) P := - IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E (h.measurable_A t) hm - have hab : - Integrable (fun ω ↦ IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω) P := - IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E - (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E) hm - calc - _ = (∑ t ∈ range n, - ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ∂P) - - ∑ t ∈ range n, ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω ∂P := by - simp_rw [IsBayesAlgEnvSeq.regret_eq_sum_gap, IsBayesAlgEnvSeq.gap_eq_sub] - rw [integral_finset_sum _ (by fun_prop), ← Finset.sum_sub_distrib] - simp_rw [integral_sub hab (haa _)] - _ = ((∑ t ∈ range n, - ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ∂P) - - ∑ t ∈ range n, - ∫ ω, ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω ∂P) + - ((∑ t ∈ range n, ∫ ω, ucb A R' l u σ2 δ (A t ω) t ω ∂P) - - ∑ t ∈ range n, ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω ∂P) := by - simp [integral_ucb_action_eq_integral_ucb_bestAction hK h] - _ = (∑ t ∈ range n, - ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω ∂P) + - ∑ t ∈ range n, ∫ ω, ucb A R' l u σ2 δ (A t ω) t ω - - IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω ∂P := by - rw [← Finset.sum_sub_distrib, ← Finset.sum_sub_distrib] - simp_rw [← integral_sub hab (hub _), ← integral_sub (hua _) (haa _)] - _ = _ := by - rw [← integral_finset_sum _ (by fun_prop), ← integral_finset_sum _ (by fun_prop)] +variable [MeasurableSpace Ω] +variable {𝓔 : Type*} [MeasurableSpace 𝓔] +variable {E : Ω → 𝓔} +variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] +variable {P : Measure Ω} [IsProbabilityMeasure P] -omit [Nonempty 𝓔] [StandardBorelSpace 𝓔] [IsProbabilityMeasure Q] in /-- This bound could be improved by using a one-sided concentration inequality. -/ lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algorithm (Fin K) ℝ} (h : IsBayesAlgEnvSeq Q κ alg E A R' P) @@ -365,7 +300,6 @@ lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algo _ = _ := by ring -omit [Nonempty 𝓔] [StandardBorelSpace 𝓔] [IsProbabilityMeasure Q] in /-- This bound could be improved by using a one-sided concentration inequality. -/ lemma integral_sum_range_ucb_action_sub_actionMean_action_le {alg : Algorithm (Fin K) ℝ} (h : IsBayesAlgEnvSeq Q κ alg E A R' P) @@ -419,8 +353,92 @@ lemma integral_sum_range_ucb_action_sub_actionMean_action_le {alg : Algorithm (F _ = _ := by ring +end ClippedUCB + +namespace TS + +section IntegralRegret + +open ClippedUCB + +variable {K : ℕ} [Nonempty (Fin K)] +variable {l u σ2 δ : ℝ} +variable {Ω : Type*} [MeasurableSpace Ω] +variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] +variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] +variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} +variable {P : Measure Ω} [IsProbabilityMeasure P] + +lemma integral_ucb_action_eq_integral_ucb_bestAction (hK : 0 < K) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) (n : ℕ) : + P[fun ω ↦ ucb A R' l u σ2 δ (A n ω) n ω] = + P[fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) n ω] := by + have := h.measurable_A + have := h.measurable_E + have := h.measurable_R + by_cases hn : n = 0 + · simp [hn] + obtain ⟨n, rfl⟩ := Nat.exists_eq_succ_of_ne_zero hn + let uc (ha : (Iic n → Fin K × ℝ) × Fin K) := ucb' n ha.1 l u σ2 δ ha.2 + calc + _ = P[fun ω ↦ uc (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)] := by + simp_rw [uc, ucb_succ_eq_ucb'] + _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) := by + rw [← integral_map (by fun_prop) (by fun_prop)] + _ = ∫ ha, uc ha ∂P.map + (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, IsBayesAlgEnvSeq.bestAction κ E ω)) := by + rw [← compProd_map_condDistrib (by fun_prop), ← compProd_map_condDistrib (by fun_prop), + Measure.compProd_congr (hasCondDistrib_action hK h n).condDistrib_eq] + _ = P[fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) (n + 1) ω] := by + rw [integral_map (by fun_prop) (by fun_prop)] + simp_rw [uc, ucb_succ_eq_ucb'] + +lemma integral_regret_eq_add (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) + (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) : + P[IsBayesAlgEnvSeq.regret κ E A n] = + P[fun ω ↦ ∑ t ∈ range n, + (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω)] + + P[fun ω ↦ ∑ t ∈ range n, + (ucb A R' l u σ2 δ (A t ω) t ω - IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω)] := by + have hua (t : ℕ) : Integrable (fun ω ↦ ucb A R' l u σ2 δ (A t ω) t ω) P := + integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A t) measurable_const + have hub (t : ℕ) : + Integrable (fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω) P := + integrable_uncurry_ucb_comp h.measurable_A h.measurable_R + (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E) measurable_const + have haa (t : ℕ) : Integrable (fun ω ↦ IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω) P := + IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E (h.measurable_A t) hm + have hab : + Integrable (fun ω ↦ IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω) P := + IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E + (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E) hm + calc + _ = (∑ t ∈ range n, + ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ∂P) - + ∑ t ∈ range n, ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω ∂P := by + simp_rw [IsBayesAlgEnvSeq.regret_eq_sum_gap, IsBayesAlgEnvSeq.gap_eq_sub] + rw [integral_finset_sum _ (by fun_prop), ← Finset.sum_sub_distrib] + simp_rw [integral_sub hab (haa _)] + _ = ((∑ t ∈ range n, + ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ∂P) - + ∑ t ∈ range n, + ∫ ω, ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω ∂P) + + ((∑ t ∈ range n, ∫ ω, ucb A R' l u σ2 δ (A t ω) t ω ∂P) - + ∑ t ∈ range n, ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω ∂P) := by + simp [integral_ucb_action_eq_integral_ucb_bestAction hK h] + _ = (∑ t ∈ range n, + ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - + ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω ∂P) + + ∑ t ∈ range n, ∫ ω, ucb A R' l u σ2 δ (A t ω) t ω - + IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω ∂P := by + rw [← Finset.sum_sub_distrib, ← Finset.sum_sub_distrib] + simp_rw [← integral_sub hab (hub _), ← integral_sub (hua _) (haa _)] + _ = _ := by + rw [← integral_finset_sum _ (by fun_prop), ← integral_finset_sum _ (by fun_prop)] + /-- This bound could be improved. -/ -lemma integral_regret_le (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm Q κ hK) E A R' P) +lemma integral_regret_le (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (n : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A n] From e06ab095654fcd55ae12be66e520534c16df9207 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 23 Apr 2026 11:26:46 +0100 Subject: [PATCH 108/155] Add Uniform --- LeanMachineLearning.lean | 1 + 1 file changed, 1 insertion(+) diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index 25f6bc55..e1b208a9 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -25,6 +25,7 @@ public import LeanMachineLearning.SequentialLearning.Algorithm public import LeanMachineLearning.SequentialLearning.AlgorithmDensity public import LeanMachineLearning.SequentialLearning.Algorithms.AuxSums public import LeanMachineLearning.SequentialLearning.Algorithms.RoundRobin +public import LeanMachineLearning.SequentialLearning.Algorithms.Uniform public import LeanMachineLearning.SequentialLearning.BayesStationaryEnv public import LeanMachineLearning.SequentialLearning.Deterministic public import LeanMachineLearning.SequentialLearning.FiniteActions From 5268980ae03669871c51c010cb0121504849e4fe Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 23 Apr 2026 11:52:15 +0100 Subject: [PATCH 109/155] Remove map_bind_condDistrib --- .../Probability/Independence/CondDistrib.lean | 5 ----- LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean | 4 ++-- 2 files changed, 2 insertions(+), 7 deletions(-) diff --git a/LeanMachineLearning/Probability/Independence/CondDistrib.lean b/LeanMachineLearning/Probability/Independence/CondDistrib.lean index ceb6e19e..b8bfac33 100644 --- a/LeanMachineLearning/Probability/Independence/CondDistrib.lean +++ b/LeanMachineLearning/Probability/Independence/CondDistrib.lean @@ -396,11 +396,6 @@ lemma ae_eq_of_condDistrib_eq_deterministic {f : β → Ω} (hf : Measurable f) rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h hfX exact ae_eq_of_map_prodMk_eq hf hX hY (hfX ▸ h) -/-- The marginal law of `Y` is obtained by integrating `condDistrib Y X μ` against `μ.map X`. -/ -lemma map_bind_condDistrib (hX : Measurable X) (hY : AEMeasurable Y μ) : - (μ.map X).bind (condDistrib Y X μ) = μ.map Y := by - rw [← Measure.snd_compProd, compProd_map_condDistrib hY, Measure.snd_map_prodMk hX] - /-- The joint law of `(Y, X)` equals the compProd of `μ.map X` and `condDistrib Y X μ`, swapped. -/ lemma compProd_map_condDistrib_swap (hX : Measurable X) (hY : Measurable Y) : (μ.map X ⊗ₘ condDistrib Y X μ).map Prod.swap = μ.map (fun ω ↦ (Y ω, X ω)) := by diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean index 4011bd2d..a65df041 100644 --- a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean @@ -167,10 +167,10 @@ lemma hasLaw_hist_withDensity (h : IsBayesAlgEnvSeq Q κ alg E A R' P) have hR₀ := h₀.measurable_R have hE := h.measurable_E have hE₀ := h₀.measurable_E - rw [← map_bind_condDistrib hE (by fun_prop), h.hasLaw_env.map_eq, + rw [← condDistrib_comp_map hE.aemeasurable (by fun_prop), h.hasLaw_env.map_eq, Measure.bind_congr_right (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), Kernel.comp_withDensity_const (by fun_prop), - ← h₀.hasLaw_env.map_eq, map_bind_condDistrib hE₀ (by fun_prop)] + ← h₀.hasLaw_env.map_eq, condDistrib_comp_map hE₀.aemeasurable (by fun_prop)] variable [StandardBorelSpace 𝓔] [Nonempty 𝓔] variable [IsProbabilityMeasure Q] From c319f3719e88ad62df90626760f8462100e17e59 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 24 Apr 2026 16:20:03 +0100 Subject: [PATCH 110/155] Improve regret bound --- .../Online/Bandit/Algorithms/TS.lean | 33 +++-- .../Online/Bandit/SumRewards.lean | 134 +++++++++++------- .../BayesStationaryEnv.lean | 64 +++++---- 3 files changed, 132 insertions(+), 99 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index 2c5a2553..f1908845 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -178,7 +178,7 @@ private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 /-- This bound could be slightly improved. -/ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, μ a ∈ Set.Icc l u) (hi : l ≤ u) - (hc : ∀ s < n, pullCount A (A s ω) s ω ≠ 0 → |empMean A R' (A s ω) s ω - μ (A s ω)| + (hc : ∀ s < n, pullCount A (A s ω) s ω ≠ 0 → empMean A R' (A s ω) s ω - μ (A s ω) < √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω))) : ∑ s ∈ range n, (ucb A R' l u σ2 δ (A s ω) s ω - μ (A s ω)) ≤ (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by @@ -206,6 +206,7 @@ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, ≤ ∑ s ∈ S₁, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) := by gcongr with s hs unfold ucb + have : 0 ≤ √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) := by positivity grind _ ≤ ∑ s ∈ range n, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) := sum_le_sum_of_subset_of_nonneg (filter_subset _ _) (fun _ _ _ => by positivity) @@ -246,7 +247,6 @@ variable {E : Ω → 𝓔} variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] variable {P : Measure Ω} [IsProbabilityMeasure P] -/-- This bound could be improved by using a one-sided concentration inequality. -/ lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algorithm (Fin K) ℝ} (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) @@ -255,13 +255,13 @@ lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algo P[fun ω ↦ ∑ t ∈ range n, (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω)] ≤ - 2 * (u - l) * (n - 1) * n * δ := by + (u - l) * (n - 1) * n * δ := by by_cases hn : n = 0 · simp [hn] let F := {ω | ∃ t < n, pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) t ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / (pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) t ω)) ≤ - |empMean A R' (IsBayesAlgEnvSeq.bestAction κ E ω) t ω - - IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω|} + empMean A R' (IsBayesAlgEnvSeq.bestAction κ E ω) t ω - + IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ≤ + -√(2 * σ2 * Real.log (1 / δ) / (pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) t ω))} have := h.measurable_A have := h.measurable_E have := h.measurable_R @@ -291,16 +291,15 @@ lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algo grind [IsBayesAlgEnvSeq.actionMean, ucb] _ = P.real F * (n * (u - l)) := by simp_rw [sum_const, card_range, nsmul_eq_mul, setIntegral_const, smul_eq_mul] - _ ≤ (2 * (n - 1) * δ) * (n * (u - l)) := by + _ ≤ ((n - 1) * δ) * (n * (u - l)) := by gcongr · nlinarith · have : (1 : ℝ) ≤ n := by simp [Nat.one_le_iff_ne_zero, hn] apply ENNReal.toReal_le_of_le_ofReal (by nlinarith) - exact h.prob_abs_empMean_bestAction_sub_actionMean_ge_le hσ2 hs hδ n + exact h.prob_empMean_bestAction_sub_actionMean_le_le hσ2 hs hδ n _ = _ := by ring -/-- This bound could be improved by using a one-sided concentration inequality. -/ lemma integral_sum_range_ucb_action_sub_actionMean_action_le {alg : Algorithm (Fin K) ℝ} (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) @@ -309,12 +308,12 @@ lemma integral_sum_range_ucb_action_sub_actionMean_action_le {alg : Algorithm (F P[fun ω ↦ ∑ t ∈ range n, (ucb A R' l u σ2 δ (A t ω) t ω - IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω)] ≤ (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + - 2 * (u - l) * K * (n - 1) * n * δ := by + (u - l) * K * (n - 1) * n * δ := by by_cases hn : n = 0 · simp [hn, hlu, mul_nonneg] let F := {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ - |empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E a ω|} + empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E a ω} have := h.measurable_A have := h.measurable_E have := h.measurable_R @@ -342,13 +341,13 @@ lemma integral_sum_range_ucb_action_sub_actionMean_action_le {alg : Algorithm (F _ = P.real F * (n * (u - l)) + P.real Fᶜ * ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) := by simp_rw [sum_const, card_range, nsmul_eq_mul, setIntegral_const, smul_eq_mul] - _ ≤ (2 * K * (n - 1) * δ) * (n * (u - l)) + + _ ≤ (K * (n - 1) * δ) * (n * (u - l)) + 1 * ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) := by have : 0 ≤ u - l := sub_nonneg.2 hlu gcongr · have : (0 : ℝ) ≤ n - 1 := by simp [Nat.one_le_iff_ne_zero, hn] apply ENNReal.toReal_le_of_le_ofReal (by positivity) - exact h.prob_abs_empMean_sub_actionMean_ge_le hσ2 hs hδ n + exact h.prob_empMean_sub_actionMean_ge_le hσ2 hs hδ n · exact measureReal_le_one _ = _ := by ring @@ -442,7 +441,7 @@ lemma integral_regret_le (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (n : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A n] - ≤ (3 * K + 2) * (u - l) + 8 * √(σ2 * K * n * Real.log n) := by + ≤ (2 * K + 1) * (u - l) + 8 * √(σ2 * K * n * Real.log n) := by by_cases hn : n = 0 · simp [hn, IsBayesAlgEnvSeq.regret, Bandits.regret] nlinarith @@ -454,15 +453,15 @@ lemma integral_regret_le (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK add_le_add (integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le h hlu hm hσ2 hs hδ n) (integral_sum_range_ucb_action_sub_actionMean_action_le h hlu hm hσ2 hs hδ n) - _ = K * (u - l) + 2 * (K + 1) * (u - l) * ((n - 1) / n) + _ = K * (u - l) + (K + 1) * (u - l) * ((n - 1) / n) + 4 * √((2 : ℝ) ^ 2 * (σ2 * K * n * Real.log n)) := by field_simp rw [Real.log_pow] ring_nf - _ = K * (u - l) + 2 * (K + 1) * (u - l) * ((n - 1) / n) + 8 * √(σ2 * K * n * Real.log n) := by + _ = K * (u - l) + (K + 1) * (u - l) * ((n - 1) / n) + 8 * √(σ2 * K * n * Real.log n) := by rw [Real.sqrt_mul (by positivity), Real.sqrt_sq (by norm_num)] ring - _ ≤ K * (u - l) + 2 * (K + 1) * (u - l) * 1 + 8 * √(σ2 * K * n * Real.log n) := by -- loose + _ ≤ K * (u - l) + (K + 1) * (u - l) * 1 + 8 * √(σ2 * K * n * Real.log n) := by -- loose have : 0 ≤ u - l := sub_nonneg.2 hlu gcongr rw [div_le_one (by positivity)] diff --git a/LeanMachineLearning/Online/Bandit/SumRewards.lean b/LeanMachineLearning/Online/Bandit/SumRewards.lean index 7705ef91..e545feb7 100644 --- a/LeanMachineLearning/Online/Bandit/SumRewards.lean +++ b/LeanMachineLearning/Online/Bandit/SumRewards.lean @@ -387,26 +387,7 @@ lemma prob_sum_range_sub_le_le_of_HasSubgaussianMGF {σ2 : ℝ≥0} · intro _ _ exact h.congr_identDistrib ((identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _) -lemma prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF {σ2 : ℝ≥0} - (h : HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) {ε : ℝ} (hε : 0 ≤ ε) (n : ℕ) : - streamMeasure ν {ω | ε ≤ |∑ k ∈ range n, (ω k a - (ν a)[id])|} ≤ - ENNReal.ofReal (2 * Real.exp (-ε ^ 2 / (2 * n * σ2))) := by - calc streamMeasure ν {ω | ε ≤ |∑ k ∈ range n, (ω k a - (ν a)[id])|} - _ = streamMeasure ν ({ω | ε ≤ ∑ k ∈ range n, (ω k a - (ν a)[id])} ∪ - {ω | ∑ k ∈ range n, (ω k a - (ν a)[id]) ≤ -ε}) := by - simp_rw [le_abs, le_neg] - rfl - _ ≤ streamMeasure ν {ω | ε ≤ ∑ k ∈ range n, (ω k a - (ν a)[id])} + - streamMeasure ν {ω | ∑ k ∈ range n, (ω k a - (ν a)[id]) ≤ -ε} := - measure_union_le _ _ - _ ≤ ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * n * σ2))) + - ENNReal.ofReal (Real.exp (-ε ^ 2 / (2 * n * σ2))) := - add_le_add (prob_sum_range_sub_ge_le_of_HasSubgaussianMGF h hε n) - (prob_sum_range_sub_le_le_of_HasSubgaussianMGF h hε n) - _ = ENNReal.ofReal (2 * Real.exp (-ε ^ 2 / (2 * n * σ2))) := by - rw [← ENNReal.ofReal_add (by positivity) (by positivity), ← two_mul] - -/-- Auxiliary lemma for `prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF'`. -/ +/-- Auxiliary lemma for `prob_sum_range_sub_*_le_of_HasSubgaussianMGF'`. -/ private lemma exp_neg_sqrt_sq_div_le {σ2 : ℝ≥0} (hσ2 : 0 < σ2) {δ : ℝ} (hδ : 0 < δ) (hn : 0 < n) : Real.exp (-√(2 * n * σ2 * Real.log (1 / δ)) ^ 2 / (2 * n * σ2)) ≤ δ := by by_cases hd : δ < 1 @@ -419,28 +400,39 @@ private lemma exp_neg_sqrt_sq_div_le {σ2 : ℝ≥0} (hσ2 : 0 < σ2) {δ : ℝ} rw [Real.sqrt_eq_zero_of_nonpos (mul_nonpos_of_nonneg_of_nonpos (by positivity) hl)] simp [hd] -lemma prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF' {σ2 : ℝ≥0} (hσ2 : 0 < σ2) +lemma prob_sum_range_sub_ge_le_of_HasSubgaussianMGF' {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (h : HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) {δ : ℝ} (hδ : 0 < δ) (hn : 0 < n) : streamMeasure ν {ω | √(2 * n * σ2 * Real.log (1 / δ)) ≤ - |∑ k ∈ range n, (ω k a - (ν a)[id])|} ≤ ENNReal.ofReal (2 * δ) := + ∑ k ∈ range n, (ω k a - (ν a)[id])} ≤ ENNReal.ofReal δ := calc - _ ≤ ENNReal.ofReal (2 * Real.exp (-√(2 * n * σ2 * Real.log (1 / δ)) ^ 2 / (2 * n * σ2))) := - prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF h (by positivity) n - _ ≤ ENNReal.ofReal (2 * δ) := by - gcongr - exact exp_neg_sqrt_sq_div_le hσ2 hδ hn + _ ≤ ENNReal.ofReal (Real.exp (-√(2 * n * σ2 * Real.log (1 / δ)) ^ 2 / (2 * n * σ2))) := + prob_sum_range_sub_ge_le_of_HasSubgaussianMGF h (by positivity) n + _ ≤ ENNReal.ofReal δ := by + gcongr + exact exp_neg_sqrt_sq_div_le hσ2 hδ hn + +lemma prob_sum_range_sub_le_le_of_HasSubgaussianMGF' {σ2 : ℝ≥0} (hσ2 : 0 < σ2) + (h : HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) {δ : ℝ} (hδ : 0 < δ) (hn : 0 < n) : + streamMeasure ν {ω | ∑ k ∈ range n, (ω k a - (ν a)[id]) ≤ + -√(2 * n * σ2 * Real.log (1 / δ))} ≤ ENNReal.ofReal δ := + calc + _ ≤ ENNReal.ofReal (Real.exp (-√(2 * n * σ2 * Real.log (1 / δ)) ^ 2 / (2 * n * σ2))) := + prob_sum_range_sub_le_le_of_HasSubgaussianMGF h (by positivity) n + _ ≤ ENNReal.ofReal δ := by + gcongr + exact exp_neg_sqrt_sq_div_le hσ2 hδ hn end StreamMeasure -lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le [Countable α] {σ2 : ℝ≥0} (hσ2 : 0 < σ2) +lemma prob_sumRewards_sub_pullCount_mul_ge_le [Countable α] {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (ha : HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {δ : ℝ} (hδ : 0 < δ) : P {ω | ∃ t < n, pullCount A a t ω ≠ 0 ∧ √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R a t ω - pullCount A a t ω * (ν a)[id]|} ≤ ENNReal.ofReal (2 * (n - 1) * δ) := - let B (m : ℕ) := {x : ℝ | √(2 * m * σ2 * Real.log (1 / δ)) ≤ |x - m * (ν a)[id]|} + sumRewards A R a t ω - pullCount A a t ω * (ν a)[id]} ≤ ENNReal.ofReal ((n - 1) * δ) := + let B (m : ℕ) := {x : ℝ | √(2 * m * σ2 * Real.log (1 / δ)) ≤ x - m * (ν a)[id]} calc _ ≤ P (⋃ m ∈ Icc 1 (n - 1), {ω | ∃ t, t < n ∧ pullCount A a t ω = m ∧ - sumRewards A R a t ω ∈ B m}) := by + sumRewards A R a t ω ∈ B m}) := by apply measure_mono intro ω ⟨t, ht, hp, hb⟩ have hm : pullCount A a t ω ∈ Icc 1 (n - 1) := mem_Icc.mpr ⟨Nat.one_le_iff_ne_zero.mpr hp, @@ -454,37 +446,73 @@ lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le [Countable α] {σ2 : ℝ≥0} _ ≤ ∑ m ∈ Icc 1 (n - 1), streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := by apply sum_le_sum exact (fun m _ ↦ prob_exists_pullCount_eq_and_sumRewards_mem_le h a m (by measurability)) - _ ≤ ∑ m ∈ Icc 1 (n - 1), ENNReal.ofReal (2 * δ) := by + _ ≤ ∑ m ∈ Icc 1 (n - 1), ENNReal.ofReal δ := by + apply sum_le_sum + intro m hm + convert StreamMeasure.prob_sum_range_sub_ge_le_of_HasSubgaussianMGF' hσ2 ha hδ + (mem_Icc.mp hm).1 using 2 + simp [B] + _ = ENNReal.ofReal ((n - 1) * δ) := by + by_cases hn : n = 0 + · simp [hn, hδ.le] + · rw [sum_const, Nat.card_Icc, add_tsub_cancel_right, ← ENNReal.ofReal_nsmul, nsmul_eq_mul, + Nat.cast_sub (Nat.one_le_iff_ne_zero.mpr hn)] + ring_nf + +lemma prob_sumRewards_sub_pullCount_mul_le_le [Countable α] {σ2 : ℝ≥0} (hσ2 : 0 < σ2) + (ha : HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {δ : ℝ} (hδ : 0 < δ) : + P {ω | ∃ t < n, pullCount A a t ω ≠ 0 ∧ + sumRewards A R a t ω - pullCount A a t ω * (ν a)[id] ≤ + -√(2 * pullCount A a t ω * σ2 * Real.log (1 / δ))} ≤ ENNReal.ofReal ((n - 1) * δ) := + let B (m : ℕ) := {x : ℝ | x - m * (ν a)[id] ≤ -√(2 * m * σ2 * Real.log (1 / δ))} + calc + _ ≤ P (⋃ m ∈ Icc 1 (n - 1), {ω | ∃ t, t < n ∧ pullCount A a t ω = m ∧ + sumRewards A R a t ω ∈ B m}) := by + apply measure_mono + intro ω ⟨t, ht, hp, hb⟩ + have hm : pullCount A a t ω ∈ Icc 1 (n - 1) := mem_Icc.mpr ⟨Nat.one_le_iff_ne_zero.mpr hp, + (pullCount_le a t ω).trans (Nat.le_sub_one_of_lt ht)⟩ + exact Set.mem_biUnion hm ⟨t, ht, rfl, hb⟩ + _ ≤ ∑ m ∈ Icc 1 (n - 1), P {ω | ∃ t, t < n ∧ pullCount A a t ω = m ∧ + sumRewards A R a t ω ∈ B m} := + measure_biUnion_finset_le _ _ + _ ≤ ∑ m ∈ Icc 1 (n - 1), P {ω | ∃ t, pullCount A a t ω = m ∧ sumRewards A R a t ω ∈ B m} := + sum_le_sum (fun _ _ ↦ measure_mono (fun _ ⟨t, _, hps⟩ ↦ ⟨t, hps⟩)) + _ ≤ ∑ m ∈ Icc 1 (n - 1), streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B m} := by apply sum_le_sum - intro m hm - convert StreamMeasure.prob_abs_sum_range_sub_ge_le_of_HasSubgaussianMGF' - hσ2 ha hδ (mem_Icc.mp hm).1 using 2 - simp [B] - _ = ENNReal.ofReal (2 * (n - 1) * δ) := by - by_cases hn : n = 0 - · simp [hn, hδ.le] - · rw [sum_const, Nat.card_Icc, add_tsub_cancel_right, ← ENNReal.ofReal_nsmul, nsmul_eq_mul, - Nat.cast_sub (Nat.one_le_iff_ne_zero.mpr hn)] - ring_nf - -lemma prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype [Fintype α] {σ2 : ℝ≥0} (hσ2 : 0 < σ2) + exact (fun m _ ↦ prob_exists_pullCount_eq_and_sumRewards_mem_le h a m (by measurability)) + _ ≤ ∑ m ∈ Icc 1 (n - 1), ENNReal.ofReal δ := by + apply sum_le_sum + intro m hm + convert StreamMeasure.prob_sum_range_sub_le_le_of_HasSubgaussianMGF' hσ2 ha hδ + (mem_Icc.mp hm).1 using 2 + simp [B] + _ = ENNReal.ofReal ((n - 1) * δ) := by + by_cases hn : n = 0 + · simp [hn, hδ.le] + · rw [sum_const, Nat.card_Icc, add_tsub_cancel_right, ← ENNReal.ofReal_nsmul, nsmul_eq_mul, + Nat.cast_sub (Nat.one_le_iff_ne_zero.mpr hn)] + ring_nf + +lemma prob_sumRewards_sub_pullCount_mul_ge_le_of_Fintype [Fintype α] {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {δ : ℝ} (hδ : 0 < δ) : P {ω | ∃ a, ∃ t < n, pullCount A a t ω ≠ 0 ∧ √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R a t ω - pullCount A a t ω * (ν a)[id]|} ≤ - ENNReal.ofReal (2 * Fintype.card α * (n - 1) * δ) := + sumRewards A R a t ω - pullCount A a t ω * (ν a)[id]} ≤ + ENNReal.ofReal (Fintype.card α * (n - 1) * δ) := calc _ ≤ ∑ a, P {ω | ∃ t < n, pullCount A a t ω ≠ 0 ∧ - √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ - |sumRewards A R a t ω - pullCount A a t ω * (ν a)[id]|} := by + √(2 * pullCount A a t ω * σ2 * Real.log (1 / δ)) ≤ + sumRewards A R a t ω - pullCount A a t ω * (ν a)[id]} := by rw [Set.setOf_exists] exact measure_iUnion_fintype_le _ _ - _ ≤ ∑ a, ENNReal.ofReal (2 * (n - 1) * δ) := - sum_le_sum fun a _ ↦ prob_abs_sumRewards_sub_pullCount_mul_ge_le hσ2 (hν a) h hδ - _ = ENNReal.ofReal (2 * Fintype.card α * (n - 1) * δ) := by - rw [sum_const, Finset.card_univ, ← ENNReal.ofReal_nsmul, nsmul_eq_mul] - ring_nf + _ ≤ ∑ a, ENNReal.ofReal ((n - 1) * δ) := + sum_le_sum fun a _ ↦ prob_sumRewards_sub_pullCount_mul_ge_le hσ2 (hν a) h hδ + _ = ENNReal.ofReal (Fintype.card α * (n - 1) * δ) := by + rw [sum_const, Finset.card_univ, ← ENNReal.ofReal_nsmul, nsmul_eq_mul] + ring_nf omit [DecidableEq α] [StandardBorelSpace α] in lemma probReal_sum_le_sum_streamMeasure [Fintype α] {c : ℝ≥0} diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 146e1fca..5ec77dbf 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -242,89 +242,95 @@ end CondDistribIsAlgEnvSeq section HasSubgaussianMGF -private lemma sqrt_two_mul_le {k : ℕ} (hk : k ≠ 0) {s μ σ l : ℝ} - (h : √(2 * σ * l / k) ≤ |s / k - μ|) : √(2 * k * σ * l) ≤ |s - k * μ| := by +variable {K : ℕ} [Nonempty (Fin K)] +variable {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] {alg : Algorithm (Fin K) ℝ} +variable {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → ℝ} +variable [IsProbabilityMeasure P] + +/-- Auxiliary lemma for `prob_empMean_sub_actionMean_ge_le`. -/ +private lemma sqrt_two_mul_le_sub {k : ℕ} (hk : k ≠ 0) {s μ σ l : ℝ} + (h : √(2 * σ * l / k) ≤ s / k - μ) : √(2 * k * σ * l) ≤ s - k * μ := by have hkp : (0 : ℝ) < k := by positivity calc √(2 * k * σ * l) _ = √(2 * σ * l / k * k ^ 2) := by field_simp _ = √(2 * σ * l / k) * k := by rw [Real.sqrt_mul' _ (sq_nonneg _), Real.sqrt_sq hkp.le] - _ ≤ |s / k - μ| * k := by + _ ≤ (s / k - μ) * k := by nlinarith - _ = |s - k * μ| := by + _ = s - k * μ := by field_simp - grind -variable {K : ℕ} [Nonempty (Fin K)] -variable {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] {alg : Algorithm (Fin K) ℝ} -variable {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → ℝ} -variable [IsProbabilityMeasure P] - -lemma prob_abs_empMean_sub_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} +lemma prob_empMean_sub_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : P {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ |empMean A R' a t ω - actionMean κ E a ω|} - ≤ ENNReal.ofReal (2 * K * (n - 1) * δ) := by + √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ empMean A R' a t ω - actionMean κ E a ω} + ≤ ENNReal.ofReal (K * (n - 1) * δ) := by have := h.measurable_E have := h.measurable_A have := h.measurable_R let S := {(e, τ) | ∃ a, ∃ t < n, pullCount IT.action a t τ ≠ 0 ∧ √(2 * pullCount IT.action a t τ * σ2 * Real.log (1 / δ)) ≤ - |sumRewards IT.action IT.reward a t τ - pullCount IT.action a t τ * actionMean κ id a e|} + sumRewards IT.action IT.reward a t τ - pullCount IT.action a t τ * actionMean κ id a e} calc _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by rw [Measure.map_apply (by fun_prop) (by measurability)] apply measure_mono intro ω ⟨t, ht, a, hpc, hle⟩ rw [empMean] at hle - exact ⟨a, t, ht, hpc, sqrt_two_mul_le hpc hle⟩ + exact ⟨a, t, ht, hpc, sqrt_two_mul_le_sub hpc hle⟩ _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by rw [← compProd_map_condDistrib (by fun_prop)] _ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := Measure.compProd_apply (by measurability) - _ ≤ ∫⁻ e, ENNReal.ofReal (2 * Fintype.card (Fin K) * (n - 1) * δ) ∂(P.map E) := by + _ ≤ ∫⁻ e, ENNReal.ofReal (Fintype.card (Fin K) * (n - 1) * δ) ∂(P.map E) := by apply lintegral_mono_ae rw [h.hasLaw_env.map_eq] filter_upwards [h.ae_IsAlgEnvSeq] with e he - exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 (hs e) he hδ - _ = ENNReal.ofReal (2 * K * (n - 1) * δ) := by + exact Bandits.prob_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 (hs e) he hδ + _ = ENNReal.ofReal (K * (n - 1) * δ) := by simp [Measure.map_apply h.measurable_E] -lemma prob_abs_empMean_bestAction_sub_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) +/-- Auxiliary lemma for `prob_empMean_bestAction_sub_actionMean_le_le`. -/ +private lemma sub_le_neg_sqrt_two_mul {k : ℕ} (hk : k ≠ 0) {s μ σ l : ℝ} + (h : s / k - μ ≤ -√(2 * σ * l / k)) : s - k * μ ≤ -√(2 * k * σ * l) := by + have : √(2 * k * σ * l) ≤ -s - k * -μ := sqrt_two_mul_le_sub hk (by grind) + linarith + +lemma prob_empMean_bestAction_sub_actionMean_le_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω)) ≤ - |empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω|} - ≤ ENNReal.ofReal (2 * (n - 1) * δ) := by + empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω ≤ + -√(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω))} + ≤ ENNReal.ofReal ((n - 1) * δ) := by have := h.measurable_E have := h.measurable_A have := h.measurable_R let S := {(e, τ) | ∃ t < n, pullCount IT.action (bestAction κ id e) t τ ≠ 0 ∧ - √(2 * pullCount IT.action (bestAction κ id e) t τ * σ2 * Real.log (1 / δ)) ≤ - |sumRewards IT.action IT.reward (bestAction κ id e) t τ - - pullCount IT.action (bestAction κ id e) t τ * actionMean κ id (bestAction κ id e) e|} + sumRewards IT.action IT.reward (bestAction κ id e) t τ - + pullCount IT.action (bestAction κ id e) t τ * actionMean κ id (bestAction κ id e) e ≤ + -√(2 * pullCount IT.action (bestAction κ id e) t τ * σ2 * Real.log (1 / δ))} calc _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by rw [Measure.map_apply (by fun_prop) (by measurability)] apply measure_mono intro ω ⟨t, ht, hpc, hle⟩ rw [empMean] at hle - exact ⟨t, ht, hpc, sqrt_two_mul_le hpc hle⟩ + exact ⟨t, ht, hpc, sub_le_neg_sqrt_two_mul hpc hle⟩ _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by rw [← compProd_map_condDistrib (by fun_prop)] _ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := Measure.compProd_apply (by measurability) - _ ≤ ∫⁻ e, ENNReal.ofReal (2 * (n - 1) * δ) ∂(P.map E) := by + _ ≤ ∫⁻ e, ENNReal.ofReal ((n - 1) * δ) ∂(P.map E) := by apply lintegral_mono_ae rw [h.hasLaw_env.map_eq] filter_upwards [h.ae_IsAlgEnvSeq] with e he - exact Bandits.prob_abs_sumRewards_sub_pullCount_mul_ge_le (ν := κ.sectR e) hσ2 (hs e _) he + exact Bandits.prob_sumRewards_sub_pullCount_mul_le_le (ν := κ.sectR e) hσ2 (hs e _) he hδ - _ = ENNReal.ofReal (2 * (n - 1) * δ) := by + _ = ENNReal.ofReal ((n - 1) * δ) := by simp [Measure.map_apply h.measurable_E] end HasSubgaussianMGF From 47c57dee8cf8130ee11b8bb47d6d8281a45ead87 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 27 Apr 2026 11:09:59 +0100 Subject: [PATCH 111/155] Finish refactor (excluding Probability/ and MeasureTheory/) --- LeanMachineLearning/Online/Bandit/Algorithms/TS.lean | 2 -- 1 file changed, 2 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index f1908845..de789586 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -176,7 +176,6 @@ private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 have : √n * √n = n := Real.mul_self_sqrt (by positivity) nlinarith -/-- This bound could be slightly improved. -/ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, μ a ∈ Set.Icc l u) (hi : l ≤ u) (hc : ∀ s < n, pullCount A (A s ω) s ω ≠ 0 → empMean A R' (A s ω) s ω - μ (A s ω) < √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω))) : @@ -436,7 +435,6 @@ lemma integral_regret_eq_add (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorith _ = _ := by rw [← integral_finset_sum _ (by fun_prop), ← integral_finset_sum _ (by fun_prop)] -/-- This bound could be improved. -/ lemma integral_regret_le (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (n : ℕ) : From 7f300db86bb2a797f4bef53c403311219ebcc573 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 4 May 2026 15:25:03 +0100 Subject: [PATCH 112/155] Refactor FullSupport.lean --- LeanMachineLearning.lean | 4 +- .../MeasureTheory/FullSupport.lean | 37 ------------------- .../Measure/AbsolutelyContinuous.lean | 26 +++++++++++++ .../MeasureTheory/OuterMeasure/Basic.lean | 29 +++++++++++++++ .../Kernel/Composition/MeasureCompProd.lean | 31 ++++++++++++++++ .../SequentialLearning/AlgorithmDensity.lean | 4 +- .../Algorithms/Uniform.lean | 10 ++--- 7 files changed, 96 insertions(+), 45 deletions(-) delete mode 100644 LeanMachineLearning/MeasureTheory/FullSupport.lean create mode 100644 LeanMachineLearning/MeasureTheory/Measure/AbsolutelyContinuous.lean create mode 100644 LeanMachineLearning/MeasureTheory/OuterMeasure/Basic.lean create mode 100644 LeanMachineLearning/Probability/Kernel/Composition/MeasureCompProd.lean diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index 66a14971..b8f59c93 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -2,8 +2,9 @@ module -- shake: keep-all public import LeanMachineLearning.MeasureTheory.Constructions.BorelSpace.MeasurableArgMax public import LeanMachineLearning.MeasureTheory.Constructions.Polish.StandardBorel -public import LeanMachineLearning.MeasureTheory.FullSupport public import LeanMachineLearning.MeasureTheory.Measurable +public import LeanMachineLearning.MeasureTheory.Measure.AbsolutelyContinuous +public import LeanMachineLearning.MeasureTheory.OuterMeasure.Basic public import LeanMachineLearning.Online.Bandit.Algorithms.ETC public import LeanMachineLearning.Online.Bandit.Algorithms.UCB public import LeanMachineLearning.Online.Bandit.Algorithms.TS @@ -17,6 +18,7 @@ public import LeanMachineLearning.Probability.Independence.CondIndepFun public import LeanMachineLearning.Probability.Independence.IndepFun public import LeanMachineLearning.Probability.Independence.IndepInfinitePi public import LeanMachineLearning.Probability.Integrable +public import LeanMachineLearning.Probability.Kernel.Composition.MeasureCompProd public import LeanMachineLearning.Probability.Kernel.IonescuTulcea.Traj public import LeanMachineLearning.Probability.Kernel.KernelSub public import LeanMachineLearning.Probability.Moments.SubGaussian diff --git a/LeanMachineLearning/MeasureTheory/FullSupport.lean b/LeanMachineLearning/MeasureTheory/FullSupport.lean deleted file mode 100644 index aa0f07e3..00000000 --- a/LeanMachineLearning/MeasureTheory/FullSupport.lean +++ /dev/null @@ -1,37 +0,0 @@ -/- -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 --/ -module - -public import Mathlib.Probability.Kernel.Composition.MeasureCompProd - -@[expose] public section - -open MeasureTheory ProbabilityTheory - -variable {α β : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {μ ν : Measure α} - -namespace Measure - -/-- Any measure is absolutely continuous wrt any measure giving positive mass to all singletons. -/ -lemma absolutelyContinuous_of_forall_singleton_pos (hν : ∀ a : α, ν {a} > 0) : μ ≪ ν := by - intro s hs - rcases s.eq_empty_or_nonempty with rfl | ⟨a, ha⟩ - · exact measure_empty - · exact absurd (measure_mono_null (Set.singleton_subset_iff.mpr ha) hs) (hν a).ne' - -end Measure - -variable {γ : Type*} {mγ : MeasurableSpace γ} - -namespace Measure.AbsolutelyContinuous - -/-- If `κ a` is absolutely continuous wrt `η a`, then so is the kernel compProd at `a`. -/ -lemma kernel_compProd_left {κ η : Kernel α β} [IsSFiniteKernel κ] [IsSFiniteKernel η] - {ξ : Kernel (α × β) γ} [IsSFiniteKernel ξ] {a : α} (hac : κ a ≪ η a) : - (κ ⊗ₖ ξ) a ≪ (η ⊗ₖ ξ) a := by - simp_rw [Kernel.compProd_apply_eq_compProd_sectR, hac.compProd_left _] - -end Measure.AbsolutelyContinuous diff --git a/LeanMachineLearning/MeasureTheory/Measure/AbsolutelyContinuous.lean b/LeanMachineLearning/MeasureTheory/Measure/AbsolutelyContinuous.lean new file mode 100644 index 00000000..a4770f4f --- /dev/null +++ b/LeanMachineLearning/MeasureTheory/Measure/AbsolutelyContinuous.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 +-/ +module + +public import Mathlib.MeasureTheory.Measure.AbsolutelyContinuous +public import LeanMachineLearning.MeasureTheory.OuterMeasure.Basic + +@[expose] public section + +variable {α : Type*} + +namespace MeasureTheory + +variable {mα : MeasurableSpace α} {μ ν : Measure α} + +namespace Measure + +lemma absolutelyContinuous_of_measure_singleton_ne_zero (h : ∀ a, ν {a} ≠ 0) : μ ≪ ν := + fun s hs ↦ by simp [(measure_null_iff_eq_empty_of_measure_singleton_ne_zero h).1 hs] + +end Measure + +end MeasureTheory diff --git a/LeanMachineLearning/MeasureTheory/OuterMeasure/Basic.lean b/LeanMachineLearning/MeasureTheory/OuterMeasure/Basic.lean new file mode 100644 index 00000000..8b74ec2b --- /dev/null +++ b/LeanMachineLearning/MeasureTheory/OuterMeasure/Basic.lean @@ -0,0 +1,29 @@ +/- +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 +-/ +module + +public import Mathlib.MeasureTheory.OuterMeasure.Basic + +@[expose] public section + +open scoped ENNReal + +namespace MeasureTheory + +section OuterMeasureClass + +variable {α F : Type*} [FunLike F (Set α) ℝ≥0∞] [OuterMeasureClass F α] + {μ : F} {s : Set α} + +lemma measure_null_iff_eq_empty_of_measure_singleton_ne_zero (h : ∀ a, μ {a} ≠ 0) : + μ s = 0 ↔ s = ∅ := by + refine ⟨fun hs ↦ ?_, fun he ↦ by simp [he]⟩ + apply Set.eq_empty_of_forall_notMem + exact fun a ha ↦ h a (measure_mono_null (Set.singleton_subset_iff.mpr ha) hs) + +end OuterMeasureClass + +end MeasureTheory diff --git a/LeanMachineLearning/Probability/Kernel/Composition/MeasureCompProd.lean b/LeanMachineLearning/Probability/Kernel/Composition/MeasureCompProd.lean new file mode 100644 index 00000000..74690c0f --- /dev/null +++ b/LeanMachineLearning/Probability/Kernel/Composition/MeasureCompProd.lean @@ -0,0 +1,31 @@ +/- +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 +-/ +module + +public import Mathlib.Probability.Kernel.Composition.MeasureCompProd + +@[expose] public section + +open ProbabilityTheory + +namespace MeasureTheory.Measure + +variable {α β : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {κ η : Kernel α β} + +section AbsolutelyContinuous + +lemma AbsolutelyContinuous.compProd_left_apply {γ : Type*} {mγ : MeasurableSpace γ} + [IsSFiniteKernel η] {a : α} (hac : κ a ≪ η a) (ξ : Kernel (α × β) γ) : + (κ ⊗ₖ ξ) a ≪ (η ⊗ₖ ξ) a := by + by_cases hκ : IsSFiniteKernel κ + · by_cases hξ : IsSFiniteKernel ξ + · simp_rw [Kernel.compProd_apply_eq_compProd_sectR, hac.compProd_left _] + · simp [Kernel.compProd_of_not_isSFiniteKernel_right _ _ hξ] + · simp [Kernel.compProd_of_not_isSFiniteKernel_left _ _ hκ] + +end AbsolutelyContinuous + +end MeasureTheory.Measure diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean index a65df041..8810af75 100644 --- a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean @@ -5,7 +5,7 @@ Authors: Rémy Degenne, Paulo Rauber -/ module -public import LeanMachineLearning.MeasureTheory.FullSupport +public import LeanMachineLearning.Probability.Kernel.Composition.MeasureCompProd public import LeanMachineLearning.Probability.WithDensity public import LeanMachineLearning.SequentialLearning.BayesStationaryEnv @@ -98,7 +98,7 @@ lemma absolutelyContinuous_map_hist (h : IsAlgEnvSeq A R' alg env P) rw [Measure.compProd_congr (h.hasCondDistrib_step n).condDistrib_eq, Measure.compProd_congr (h₀.hasCondDistrib_step n).condDistrib_eq] apply Measure.AbsolutelyContinuous.compProd ih - filter_upwards with h' using Measure.AbsolutelyContinuous.kernel_compProd_left (hc.policy n h') + filter_upwards with h' using Measure.AbsolutelyContinuous.compProd_left_apply (hc.policy n h') _ lemma hasLaw_hist_withDensity (h : IsAlgEnvSeq A R' alg env P) (h₀ : IsAlgEnvSeq A₀ R₀ alg₀ env P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : HasLaw (IsAlgEnvSeq.hist A R' n) diff --git a/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean b/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean index 4e2a95e1..30a0c5fe 100644 --- a/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean +++ b/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean @@ -5,7 +5,7 @@ Authors: Rémy Degenne, Paulo Rauber -/ module -public import LeanMachineLearning.MeasureTheory.FullSupport +public import LeanMachineLearning.MeasureTheory.Measure.AbsolutelyContinuous public import LeanMachineLearning.SequentialLearning.Algorithm /-! # The Uniform Algorithm -/ @@ -29,9 +29,9 @@ def uniformAlgorithm (hK : 0 < K) : Algorithm (Fin K) ℝ := lemma absolutelyContinuous_uniformAlgorithm (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) : alg ≪ₐ uniformAlgorithm hK where - p0 := Measure.absolutelyContinuous_of_forall_singleton_pos - (by simp [uniformAlgorithm, uniformOn, cond_pos_of_inter_ne_zero]) - policy n h := Measure.absolutelyContinuous_of_forall_singleton_pos - (by simp [uniformAlgorithm, uniformOn, cond_pos_of_inter_ne_zero]) + p0 := Measure.absolutelyContinuous_of_measure_singleton_ne_zero + (by simp [uniformAlgorithm, uniformOn, ← pos_iff_ne_zero, cond_pos_of_inter_ne_zero]) + policy n h := Measure.absolutelyContinuous_of_measure_singleton_ne_zero + (by simp [uniformAlgorithm, uniformOn, ← pos_iff_ne_zero, cond_pos_of_inter_ne_zero]) end Bandits From 263c89cb8d71c1c0f157939a3e039e3c8312e211 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 5 May 2026 11:34:10 +0100 Subject: [PATCH 113/155] Refactor CondDistrib.lean --- .../Probability/HasCondDistrib.lean | 4 ++- .../Probability/Independence/CondDistrib.lean | 27 ++++++++----------- .../SequentialLearning/AlgorithmDensity.lean | 4 +-- 3 files changed, 16 insertions(+), 19 deletions(-) diff --git a/LeanMachineLearning/Probability/HasCondDistrib.lean b/LeanMachineLearning/Probability/HasCondDistrib.lean index f3158eae..084c5590 100644 --- a/LeanMachineLearning/Probability/HasCondDistrib.lean +++ b/LeanMachineLearning/Probability/HasCondDistrib.lean @@ -322,7 +322,9 @@ lemma HasCondDistrib.ae_hasCondDistrib_sectR [IsFiniteMeasure μ] have h_comp_fst : (condDistrib (f ∘ W) Z μ) =ᵐ[μ.map Z] (condDistrib W Z μ).map f := condDistrib_comp Z hW hf - have h_nested := (Kernel.ae_eq_map_prod_iff_ae_condDistrib hfW).mp hcd.condDistrib_eq + have h_eq := hcd.condDistrib_eq + rw [(compProd_map_condDistrib hfW).symm] at h_eq + have h_nested := Measure.ae_ae_of_ae_compProd h_eq filter_upwards [h_prod, h_comp_pair, h_comp_fst, h_nested] with z h_prod_z h_pair_z h_fst_z h_nested_z refine ⟨hg.aemeasurable, hf.aemeasurable, ?_⟩ diff --git a/LeanMachineLearning/Probability/Independence/CondDistrib.lean b/LeanMachineLearning/Probability/Independence/CondDistrib.lean index b8bfac33..10980c37 100644 --- a/LeanMachineLearning/Probability/Independence/CondDistrib.lean +++ b/LeanMachineLearning/Probability/Independence/CondDistrib.lean @@ -28,15 +28,15 @@ section CondDistrib variable [IsFiniteMeasure μ] -/-- An ae-equality of kernels wrt the joint law `μ.map (X, Y)` is equivalent to an ae-equality -fiberwise via the conditional distribution of `Y` given `X`. -/ -lemma Kernel.ae_eq_map_prod_iff_ae_condDistrib - [MeasurableSpace.CountableOrCountablyGenerated (β × Ω) δ] - (hY : AEMeasurable Y μ) {f g : Kernel (β × Ω) δ} [IsFiniteKernel f] [IsFiniteKernel g] : - f =ᵐ[μ.map (fun ω ↦ (X ω, Y ω))] g ↔ - ∀ᵐ x ∂(μ.map X), ∀ᵐ y ∂(condDistrib Y X μ x), f (x, y) = g (x, y) := by - rw [← compProd_map_condDistrib hY] - exact Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _) +lemma map_swap_compProd_map_condDistrib (hY : AEMeasurable Y μ) : + (μ.map X ⊗ₘ condDistrib Y X μ).map Prod.swap = μ.map (fun a ↦ (Y a, X a)) := by + by_cases hX : AEMeasurable X μ + · rw [compProd_map_condDistrib hY, + AEMeasurable.map_map_of_aemeasurable measurable_swap.aemeasurable (hX.prodMk hY)] + rfl + · have hYX : ¬ AEMeasurable (fun a ↦ (Y a, X a)) μ := + fun h ↦ hX (measurable_snd.comp_aemeasurable h) + simp [hX, hYX] lemma condDistrib_prod_left [StandardBorelSpace β] [Nonempty β] (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) (hT : AEMeasurable T μ) : @@ -53,7 +53,8 @@ lemma condDistrib_prod_self_left [StandardBorelSpace β] [Nonempty β] [Standard 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) - have h_fst' := (Kernel.ae_eq_map_prod_iff_ae_condDistrib hX).mp h_fst + 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] simp only [Kernel.deterministic_apply] at hz2 @@ -396,12 +397,6 @@ lemma ae_eq_of_condDistrib_eq_deterministic {f : β → Ω} (hf : Measurable f) rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h hfX exact ae_eq_of_map_prodMk_eq hf hX hY (hfX ▸ h) -/-- The joint law of `(Y, X)` equals the compProd of `μ.map X` and `condDistrib Y X μ`, swapped. -/ -lemma compProd_map_condDistrib_swap (hX : Measurable X) (hY : Measurable Y) : - (μ.map X ⊗ₘ condDistrib Y X μ).map Prod.swap = μ.map (fun ω ↦ (Y ω, X ω)) := by - rw [compProd_map_condDistrib hY.aemeasurable, Measure.map_map measurable_swap (hX.prodMk hY)] - rfl - end CondDistrib section Cond diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean index 8810af75..b323813d 100644 --- a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean @@ -189,11 +189,11 @@ lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) have hE := h.measurable_E have hE₀ := h₀.measurable_E rw [condDistrib_ae_eq_iff_measure_eq_compProd _ h.measurable_E.aemeasurable, - ← compProd_map_condDistrib_swap hE (by fun_prop), h.hasLaw_env.map_eq, + ← map_swap_compProd_map_condDistrib (by fun_prop), h.hasLaw_env.map_eq, Measure.compProd_eq_compProd_withDensity (by fun_prop) (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), Measure.map_swap_withDensity_fst (by fun_prop), - ← h₀.hasLaw_env.map_eq, compProd_map_condDistrib_swap hE₀ (by fun_prop), + ← h₀.hasLaw_env.map_eq, map_swap_compProd_map_condDistrib (by fun_prop), ← compProd_map_condDistrib (by fun_prop), ← Measure.withDensity_compProd_left (by fun_prop), ← (hasLaw_hist_withDensity h h₀ hc n).map_eq] From 91385bb57ee65e0bcba7e160c85f32c4b49d83aa Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 5 May 2026 11:54:35 +0100 Subject: [PATCH 114/155] Lint --- LeanMachineLearning/Online/Bandit/Algorithms/TS.lean | 3 +-- LeanMachineLearning/Probability/WithDensity.lean | 2 +- .../SequentialLearning/BayesStationaryEnv.lean | 4 ++-- 3 files changed, 4 insertions(+), 5 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index de789586..0cb4b20b 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -34,8 +34,7 @@ instance {hK : 0 < K} {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel ( Kernel.IsMarkovKernel.map _ (by fun_prop) noncomputable -def TS.initialPolicy (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] - (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] : Measure (Fin K) := +def TS.initialPolicy (hK : 0 < K) (Q : Measure 𝓔) (κ : Kernel (𝓔 × Fin K) ℝ) : Measure (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK Q.map (IsBayesAlgEnvSeq.bestAction κ id) diff --git a/LeanMachineLearning/Probability/WithDensity.lean b/LeanMachineLearning/Probability/WithDensity.lean index 47ba5f6e..b7563aef 100644 --- a/LeanMachineLearning/Probability/WithDensity.lean +++ b/LeanMachineLearning/Probability/WithDensity.lean @@ -59,7 +59,7 @@ lemma withDensity_map_equiv /-- Mapping a `withDensity` through a `MeasurableEquiv` from the snd component. -/ lemma map_swap_withDensity_fst - {μ : Measure (α × β)} [SFinite μ] + {μ : Measure (α × β)} {f : β → ℝ≥0∞} (hf : Measurable f) : (μ.withDensity (f ∘ Prod.snd)).map Prod.swap = (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := by diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 5ec77dbf..99b5bd23 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -91,8 +91,8 @@ def gap (κ : Kernel (𝓔 × α) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → α) omit [MeasurableSpace Ω] in /-- The gap is non-negative if the means are bounded by `u : ℝ` (even if `α` is not `Finite`). -/ -lemma gap_nonneg_of_le [Nonempty α] {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → α} {n : ℕ} - {ω : Ω} {u : ℝ} (h : ∀ e a, (κ (e, a))[id] ≤ u) : 0 ≤ gap κ E A n ω := by +lemma gap_nonneg_of_le {κ : Kernel (𝓔 × α) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → α} {n : ℕ} {ω : Ω} {u : ℝ} + (h : ∀ e a, (κ (e, a))[id] ≤ u) : 0 ≤ gap κ E A n ω := by simp_rw [gap, Bandits.gap, Kernel.sectR_apply] linarith [le_ciSup ⟨u, Set.forall_mem_range.2 fun a ↦ (h (E ω) a)⟩ (A n ω)] From eb2bf83206f8d4cdfc115a5ac585ca027791cd19 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 5 May 2026 12:32:48 +0100 Subject: [PATCH 115/155] Lint --- LeanMachineLearning/Online/Bandit/Algorithms/TS.lean | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index 0cb4b20b..96a5df8b 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -38,8 +38,8 @@ def TS.initialPolicy (hK : 0 < K) (Q : Measure 𝓔) (κ : Kernel (𝓔 × Fin K have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK Q.map (IsBayesAlgEnvSeq.bestAction κ id) -instance {hK : 0 < K} {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} - [IsMarkovKernel κ] : IsProbabilityMeasure (TS.initialPolicy hK Q κ) := +instance {hK : 0 < K} {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} : + IsProbabilityMeasure (TS.initialPolicy hK Q κ) := Measure.isProbabilityMeasure_map (by fun_prop) noncomputable From 03d1d336b566d0c391cf60dd5aee1e6c75c36b63 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 7 May 2026 14:06:58 +0100 Subject: [PATCH 116/155] Refactor HasCondDistrib.lean (in progress) --- .../Probability/HasCondDistrib.lean | 32 ++++++++----------- .../Algorithms/RandomSampling.lean | 2 +- .../BayesStationaryEnv.lean | 2 +- 3 files changed, 16 insertions(+), 20 deletions(-) diff --git a/LeanMachineLearning/Probability/HasCondDistrib.lean b/LeanMachineLearning/Probability/HasCondDistrib.lean index 084c5590..ecc43593 100644 --- a/LeanMachineLearning/Probability/HasCondDistrib.lean +++ b/LeanMachineLearning/Probability/HasCondDistrib.lean @@ -186,8 +186,8 @@ lemma hasCondDistrib_prod_right_iff [IsFiniteMeasure μ] [IsFiniteKernel κ] (X rw [← Measure.map_prod_map _ _ (by fun_prop) (by fun_prop), Measure.map_id, Measure.map_dirac' (by fun_prop)] -lemma hasLaw_of_hasCondDistrib_const [IsProbabilityMeasure μ] {Q : Measure Ω} [SFinite Q] - (h : HasCondDistrib Y X (Kernel.const _ Q) μ) : HasLaw Y Q μ := by +lemma HasCondDistrib.hasLaw_of_const [IsProbabilityMeasure μ] {Q : Measure Ω} [SFinite Q] + (h : HasCondDistrib Y X (Kernel.const β Q) μ) : HasLaw Y Q μ := by obtain ⟨hY, hX, h⟩ := h refine ⟨hY, ?_⟩ have h_snd : (μ.map (fun ω => (X ω, Y ω))).snd = Q := by @@ -199,23 +199,19 @@ lemma hasLaw_of_hasCondDistrib_const [IsProbabilityMeasure μ] {Q : Measure Ω} simp [MeasureTheory.Measure.map_apply_of_aemeasurable hX] rwa [Measure.snd_map_prodMk₀ hX] at h_snd --- Replace `hasLaw_of_hasCondDistrib_const`? -lemma HasCondDistrib.hasLaw_of_const {Q : Measure Ω} - [IsProbabilityMeasure μ] [SFinite Q] - (h : HasCondDistrib Y X (Kernel.const β Q) μ) : HasLaw Y Q μ := - hasLaw_of_hasCondDistrib_const h +lemma HasCondDistrib.indepFun_of_const [IsProbabilityMeasure μ] {Q : Measure Ω} [SFinite Q] + (h : HasCondDistrib Y X (Kernel.const β Q) μ) : IndepFun X Y μ := by + rw [indepFun_iff_condDistrib_eq_const h.aemeasurable_snd h.aemeasurable_fst, + h.hasLaw_of_const.map_eq] + exact h.condDistrib_eq -lemma HasCondDistrib.swap_const {Q : Measure Ω} - [StandardBorelSpace β] [Nonempty β] - [IsProbabilityMeasure μ] [IsFiniteMeasure Q] - (h : HasCondDistrib Y X (Kernel.const β Q) μ) : - HasCondDistrib X Y (Kernel.const Ω (μ.map X)) μ := by - have h_indep : IndepFun X Y μ := by - rw [indepFun_iff_condDistrib_eq_const h.aemeasurable_snd h.aemeasurable_fst, - h.hasLaw_of_const.map_eq] - exact h.condDistrib_eq - exact ⟨h.aemeasurable_snd, h.aemeasurable_fst, - condDistrib_of_indepFun h_indep.symm h.aemeasurable_fst h.aemeasurable_snd⟩ +lemma HasCondDistrib.const_map_of_const [IsProbabilityMeasure μ] {Q : Measure Ω} [SFinite Q] + (h : HasCondDistrib Y X (Kernel.const β Q) μ) [StandardBorelSpace β] [Nonempty β] : + HasCondDistrib X Y (Kernel.const Ω (μ.map X)) μ where + aemeasurable_fst := h.aemeasurable_snd + aemeasurable_snd := h.aemeasurable_fst + condDistrib_eq := + condDistrib_of_indepFun h.indepFun_of_const.symm h.aemeasurable_fst h.aemeasurable_snd lemma HasLaw.prod_of_hasCondDistrib {P : Measure β} [IsFiniteMeasure μ] [IsSFiniteKernel κ] (h1 : HasLaw X P μ) (h2 : HasCondDistrib Y X κ μ) : diff --git a/LeanMachineLearning/SequentialLearning/Algorithms/RandomSampling.lean b/LeanMachineLearning/SequentialLearning/Algorithms/RandomSampling.lean index f6a35718..01ef6c7b 100644 --- a/LeanMachineLearning/SequentialLearning/Algorithms/RandomSampling.lean +++ b/LeanMachineLearning/SequentialLearning/Algorithms/RandomSampling.lean @@ -57,7 +57,7 @@ lemma hasLaw_action (h : IsAlgEnvSeq A R (randomSampling μ) env P) (n : ℕ) : exact h.hasLaw_action_zero · push Not at hn obtain ⟨k, rfl⟩ := Nat.exists_eq_succ_of_ne_zero hn - exact hasLaw_of_hasCondDistrib_const <| h.hasCondDistrib_action k + exact (h.hasCondDistrib_action k).hasLaw_of_const /-- Actions are mutually independent. -/ lemma iIndep_action (h : IsAlgEnvSeq A R (randomSampling μ) env P) : diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 99b5bd23..73bcc699 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -365,7 +365,7 @@ lemma IsAlgEnvSeq.isBayesAlgEnvSeq hasCondDistrib_action_zero := by have hc : HasCondDistrib (fun ω ↦ (R' 0 ω).1) (A 0) (Kernel.const _ Q) P := by simpa [bayesStationaryEnv] using h.hasCondDistrib_reward_zero.fst - simpa [h.hasLaw_action_zero.map_eq, Algorithm.prod_left] using hc.swap_const + simpa [h.hasLaw_action_zero.map_eq, Algorithm.prod_left] using hc.const_map_of_const hasCondDistrib_reward_zero := h.hasCondDistrib_reward_zero.of_compProd.comp_right MeasurableEquiv.prodComm hasCondDistrib_action n := by From 074321985000ed7bde74780826048b18591b26b5 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 7 May 2026 14:45:58 +0100 Subject: [PATCH 117/155] Open defininitions from IsBayesAlgEnvSeq --- .../Online/Bandit/Algorithms/TS.lean | 92 ++++++++----------- 1 file changed, 38 insertions(+), 54 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index 96a5df8b..f22a8f31 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -13,6 +13,7 @@ public import LeanMachineLearning.SequentialLearning.Algorithms.Uniform @[expose] public section open MeasureTheory ProbabilityTheory Finset Learning +open IsBayesAlgEnvSeq (bestAction actionMean) open scoped NNReal @@ -27,7 +28,7 @@ noncomputable def TS.policy (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) n).map (IsBayesAlgEnvSeq.bestAction κ id) + (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) n).map (bestAction κ id) instance {hK : 0 < K} {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] {n : ℕ} : IsMarkovKernel (TS.policy hK Q κ n) := @@ -36,7 +37,7 @@ instance {hK : 0 < K} {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel ( noncomputable def TS.initialPolicy (hK : 0 < K) (Q : Measure 𝓔) (κ : Kernel (𝓔 × Fin K) ℝ) : Measure (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - Q.map (IsBayesAlgEnvSeq.bestAction κ id) + Q.map (bestAction κ id) instance {hK : 0 < K} {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} : IsProbabilityMeasure (TS.initialPolicy hK Q κ) := @@ -61,25 +62,23 @@ variable {P : Measure Ω} [IsProbabilityMeasure P] lemma hasCondDistrib_action (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) (n : ℕ) : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) - (condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' n) P) P where + (condDistrib (bestAction κ E) (IsAlgEnvSeq.hist A R' n) P) P where aemeasurable_fst := (h.measurable_A (n + 1)).aemeasurable aemeasurable_snd := (IsAlgEnvSeq.measurable_hist h.measurable_A h.measurable_R n).aemeasurable condDistrib_eq := by - have hm : Measurable (IsBayesAlgEnvSeq.bestAction κ id) := by fun_prop + have hm : Measurable (bestAction κ id) := by fun_prop calc _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' n)] - (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) n).map - (IsBayesAlgEnvSeq.bestAction κ id) := + (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) n).map (bestAction κ id) := (h.hasCondDistrib_action' n).condDistrib_eq _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' n)] - (condDistrib E (IsAlgEnvSeq.hist A R' n) P).map - (IsBayesAlgEnvSeq.bestAction κ id) := by + (condDistrib E (IsAlgEnvSeq.hist A R' n) P).map (bestAction κ id) := by filter_upwards [(h.hasCondDistrib_env_hist (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) (absolutelyContinuous_uniformAlgorithm hK _) n).condDistrib_eq] with _ hc simp_rw [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hc] _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' n)] - condDistrib (IsBayesAlgEnvSeq.bestAction κ E) (IsAlgEnvSeq.hist A R' n) P := + condDistrib (bestAction κ E) (IsAlgEnvSeq.hist A R' n) P := (condDistrib_comp (IsAlgEnvSeq.hist A R' n) h.measurable_E.aemeasurable hm).symm end TS @@ -251,26 +250,22 @@ lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algo (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (hδ : 0 < δ) (n : ℕ) : P[fun ω ↦ ∑ t ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω)] ≤ - (u - l) * (n - 1) * n * δ := by + (actionMean κ E (bestAction κ E ω) ω - ucb A R' l u σ2 δ (bestAction κ E ω) t ω)] ≤ + (u - l) * (n - 1) * n * δ := by by_cases hn : n = 0 · simp [hn] - let F := {ω | ∃ t < n, pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) t ω ≠ 0 ∧ - empMean A R' (IsBayesAlgEnvSeq.bestAction κ E ω) t ω - - IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ≤ - -√(2 * σ2 * Real.log (1 / δ) / (pullCount A (IsBayesAlgEnvSeq.bestAction κ E ω) t ω))} + let F := {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ + empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω ≤ + -√(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω))} have := h.measurable_A have := h.measurable_E have := h.measurable_R have hF : MeasurableSet F := by measurability - have : - Integrable (fun ω ↦ IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω) P := + have : Integrable (fun ω ↦ actionMean κ E (bestAction κ E ω) ω) P := IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp (by fun_prop) (by fun_prop) hm calc _ ≤ ∫ ω in F, ∑ t ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω) ∂P := by + (actionMean κ E (bestAction κ E ω) ω - ucb A R' l u σ2 δ (bestAction κ E ω) t ω) ∂P := by rw [← integral_add_compl hF (by fun_prop)] apply add_le_of_nonpos_right apply setIntegral_nonpos hF.compl @@ -279,14 +274,14 @@ lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algo intro t ht rw [Set.mem_compl_iff, Set.mem_setOf_eq] at hω push Not at hω - grind [hω t (mem_range.mp ht), ucb, IsBayesAlgEnvSeq.actionMean] + grind [hω t (mem_range.mp ht), ucb, actionMean] _ ≤ ∫ ω in F, ∑ t ∈ range n, (u - l) ∂P := by apply setIntegral_mono_on (Integrable.integrableOn (by fun_prop)) (Integrable.integrableOn (by fun_prop)) hF intro ω hω apply sum_le_sum intro t ht - grind [IsBayesAlgEnvSeq.actionMean, ucb] + grind [actionMean, ucb] _ = P.real F * (n * (u - l)) := by simp_rw [sum_const, card_range, nsmul_eq_mul, setIntegral_const, smul_eq_mul] _ ≤ ((n - 1) * δ) * (n * (u - l)) := by @@ -303,20 +298,17 @@ lemma integral_sum_range_ucb_action_sub_actionMean_action_le {alg : Algorithm (F (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (hδ : 0 < δ) (n : ℕ) : - P[fun ω ↦ ∑ t ∈ range n, - (ucb A R' l u σ2 δ (A t ω) t ω - IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω)] ≤ - (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + - (u - l) * K * (n - 1) * n * δ := by + P[fun ω ↦ ∑ t ∈ range n, (ucb A R' l u σ2 δ (A t ω) t ω - actionMean κ E (A t ω) ω)] ≤ + (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + (u - l) * K * (n - 1) * n * δ := by by_cases hn : n = 0 · simp [hn, hlu, mul_nonneg] let F := {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ - empMean A R' a t ω - IsBayesAlgEnvSeq.actionMean κ E a ω} + √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ empMean A R' a t ω - actionMean κ E a ω} have := h.measurable_A have := h.measurable_E have := h.measurable_R have hF : MeasurableSet F := by measurability - have : ∀ t, Integrable (fun ω ↦ IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω) P := + have : ∀ t, Integrable (fun ω ↦ actionMean κ E (A t ω) ω) P := fun t ↦ IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp (by fun_prop) (by fun_prop) hm calc _ ≤ (∫ ω in F, ∑ t ∈ range n, (u - l) ∂P) + @@ -328,7 +320,7 @@ lemma integral_sum_range_ucb_action_sub_actionMean_action_le {alg : Algorithm (F intro ω hω apply sum_le_sum intro t ht - grind [ucb, IsBayesAlgEnvSeq.actionMean] + grind [ucb, actionMean] · apply setIntegral_mono_on (Integrable.integrableOn (by fun_prop)) (Integrable.integrableOn (by fun_prop)) hF.compl intro ω hω @@ -369,7 +361,7 @@ variable {P : Measure Ω} [IsProbabilityMeasure P] lemma integral_ucb_action_eq_integral_ucb_bestAction (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) (n : ℕ) : P[fun ω ↦ ucb A R' l u σ2 δ (A n ω) n ω] = - P[fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) n ω] := by + P[fun ω ↦ ucb A R' l u σ2 δ (bestAction κ E ω) n ω] := by have := h.measurable_A have := h.measurable_E have := h.measurable_R @@ -382,11 +374,10 @@ lemma integral_ucb_action_eq_integral_ucb_bestAction (hK : 0 < K) simp_rw [uc, ucb_succ_eq_ucb'] _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) := by rw [← integral_map (by fun_prop) (by fun_prop)] - _ = ∫ ha, uc ha ∂P.map - (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, IsBayesAlgEnvSeq.bestAction κ E ω)) := by + _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, bestAction κ E ω)) := by rw [← compProd_map_condDistrib (by fun_prop), ← compProd_map_condDistrib (by fun_prop), Measure.compProd_congr (hasCondDistrib_action hK h n).condDistrib_eq] - _ = P[fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) (n + 1) ω] := by + _ = P[fun ω ↦ ucb A R' l u σ2 δ (bestAction κ E ω) (n + 1) ω] := by rw [integral_map (by fun_prop) (by fun_prop)] simp_rw [uc, ucb_succ_eq_ucb'] @@ -394,41 +385,34 @@ lemma integral_regret_eq_add (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorith (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A n] = P[fun ω ↦ ∑ t ∈ range n, - (IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω)] + + (actionMean κ E (bestAction κ E ω) ω - ucb A R' l u σ2 δ (bestAction κ E ω) t ω)] + P[fun ω ↦ ∑ t ∈ range n, - (ucb A R' l u σ2 δ (A t ω) t ω - IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω)] := by + (ucb A R' l u σ2 δ (A t ω) t ω - actionMean κ E (A t ω) ω)] := by have hua (t : ℕ) : Integrable (fun ω ↦ ucb A R' l u σ2 δ (A t ω) t ω) P := integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (h.measurable_A t) measurable_const - have hub (t : ℕ) : - Integrable (fun ω ↦ ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω) P := + have hub (t : ℕ) : Integrable (fun ω ↦ ucb A R' l u σ2 δ (bestAction κ E ω) t ω) P := integrable_uncurry_ucb_comp h.measurable_A h.measurable_R (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E) measurable_const - have haa (t : ℕ) : Integrable (fun ω ↦ IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω) P := + have haa (t : ℕ) : Integrable (fun ω ↦ actionMean κ E (A t ω) ω) P := IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E (h.measurable_A t) hm - have hab : - Integrable (fun ω ↦ IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω) P := + have hab : Integrable (fun ω ↦ actionMean κ E (bestAction κ E ω) ω) P := IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E) hm calc - _ = (∑ t ∈ range n, - ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ∂P) - - ∑ t ∈ range n, ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω ∂P := by + _ = (∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω ∂P) - + ∑ t ∈ range n, ∫ ω, actionMean κ E (A t ω) ω ∂P := by simp_rw [IsBayesAlgEnvSeq.regret_eq_sum_gap, IsBayesAlgEnvSeq.gap_eq_sub] rw [integral_finset_sum _ (by fun_prop), ← Finset.sum_sub_distrib] simp_rw [integral_sub hab (haa _)] - _ = ((∑ t ∈ range n, - ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω ∂P) - - ∑ t ∈ range n, - ∫ ω, ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω ∂P) + + _ = ((∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω ∂P) - + ∑ t ∈ range n, ∫ ω, ucb A R' l u σ2 δ (bestAction κ E ω) t ω ∂P) + ((∑ t ∈ range n, ∫ ω, ucb A R' l u σ2 δ (A t ω) t ω ∂P) - - ∑ t ∈ range n, ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω ∂P) := by + ∑ t ∈ range n, ∫ ω, actionMean κ E (A t ω) ω ∂P) := by simp [integral_ucb_action_eq_integral_ucb_bestAction hK h] - _ = (∑ t ∈ range n, - ∫ ω, IsBayesAlgEnvSeq.actionMean κ E (IsBayesAlgEnvSeq.bestAction κ E ω) ω - - ucb A R' l u σ2 δ (IsBayesAlgEnvSeq.bestAction κ E ω) t ω ∂P) + + _ = (∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω - + ucb A R' l u σ2 δ (bestAction κ E ω) t ω ∂P) + ∑ t ∈ range n, ∫ ω, ucb A R' l u σ2 δ (A t ω) t ω - - IsBayesAlgEnvSeq.actionMean κ E (A t ω) ω ∂P := by + actionMean κ E (A t ω) ω ∂P := by rw [← Finset.sum_sub_distrib, ← Finset.sum_sub_distrib] simp_rw [← integral_sub hab (hub _), ← integral_sub (hua _) (haa _)] _ = _ := by From bfa848101b17eb35de7f7b49082df77745621816 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 7 May 2026 15:51:18 +0100 Subject: [PATCH 118/155] Move concentration inequalities out of BayesStationaryEnv --- .../Online/Bandit/Algorithms/TS.lean | 1 + .../Online/Bandit/SumRewards.lean | 99 ++++++++++++++++++- .../BayesStationaryEnv.lean | 98 +----------------- 3 files changed, 101 insertions(+), 97 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index f22a8f31..cddbfa87 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -5,6 +5,7 @@ Authors: Rémy Degenne, Paulo Rauber -/ module +public import LeanMachineLearning.Online.Bandit.SumRewards public import LeanMachineLearning.SequentialLearning.AlgorithmDensity public import LeanMachineLearning.SequentialLearning.Algorithms.Uniform diff --git a/LeanMachineLearning/Online/Bandit/SumRewards.lean b/LeanMachineLearning/Online/Bandit/SumRewards.lean index e545feb7..4a4b3dda 100644 --- a/LeanMachineLearning/Online/Bandit/SumRewards.lean +++ b/LeanMachineLearning/Online/Bandit/SumRewards.lean @@ -6,8 +6,8 @@ Authors: Rémy Degenne module public import LeanMachineLearning.Online.Bandit.ArrayProbSpace -public import LeanMachineLearning.Online.Bandit.Regret public import LeanMachineLearning.Probability.Moments.SubGaussian +public import LeanMachineLearning.SequentialLearning.BayesStationaryEnv /-! # Law of the sum of rewards -/ @@ -658,3 +658,100 @@ lemma todo' {σ2 : ℝ≥0} {c : ℝ} end Subgaussian end Bandits + +namespace Learning.IsBayesAlgEnvSeq + +variable {𝓔 Ω : Type*} [MeasurableSpace 𝓔] [MeasurableSpace Ω] +variable {K : ℕ} [Nonempty (Fin K)] +variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] +variable {alg : Algorithm (Fin K) ℝ} +variable {E : Ω → 𝓔} {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → ℝ} +variable {P : Measure Ω} [IsProbabilityMeasure P] + +/-- Auxiliary lemma for `prob_empMean_sub_actionMean_ge_le`. -/ +private lemma sqrt_two_mul_le_sub {k : ℕ} (hk : k ≠ 0) {s μ σ l : ℝ} + (h : √(2 * σ * l / k) ≤ s / k - μ) : √(2 * k * σ * l) ≤ s - k * μ := by + have hkp : (0 : ℝ) < k := by positivity + calc √(2 * k * σ * l) + _ = √(2 * σ * l / k * k ^ 2) := by + field_simp + _ = √(2 * σ * l / k) * k := by + rw [Real.sqrt_mul' _ (sq_nonneg _), Real.sqrt_sq hkp.le] + _ ≤ (s / k - μ) * k := by + nlinarith + _ = s - k * μ := by + field_simp + +lemma prob_empMean_sub_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} + (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : + P {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ + √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ empMean A R' a t ω - actionMean κ E a ω} + ≤ ENNReal.ofReal (K * (n - 1) * δ) := by + have := h.measurable_E + have := h.measurable_A + have := h.measurable_R + let S := {(e, τ) | ∃ a, ∃ t < n, pullCount IT.action a t τ ≠ 0 ∧ + √(2 * pullCount IT.action a t τ * σ2 * Real.log (1 / δ)) ≤ + sumRewards IT.action IT.reward a t τ - pullCount IT.action a t τ * actionMean κ id a e} + calc + _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by + rw [Measure.map_apply (by fun_prop) (by measurability)] + apply measure_mono + intro ω ⟨t, ht, a, hpc, hle⟩ + rw [empMean] at hle + exact ⟨a, t, ht, hpc, sqrt_two_mul_le_sub hpc hle⟩ + _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by + rw [← compProd_map_condDistrib (by fun_prop)] + _ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := + Measure.compProd_apply (by measurability) + _ ≤ ∫⁻ e, ENNReal.ofReal (Fintype.card (Fin K) * (n - 1) * δ) ∂(P.map E) := by + apply lintegral_mono_ae + rw [h.hasLaw_env.map_eq] + filter_upwards [h.ae_IsAlgEnvSeq] with e he + exact Bandits.prob_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 (hs e) he hδ + _ = ENNReal.ofReal (K * (n - 1) * δ) := by + simp [Measure.map_apply h.measurable_E] + +/-- Auxiliary lemma for `prob_empMean_bestAction_sub_actionMean_le_le`. -/ +private lemma sub_le_neg_sqrt_two_mul {k : ℕ} (hk : k ≠ 0) {s μ σ l : ℝ} + (h : s / k - μ ≤ -√(2 * σ * l / k)) : s - k * μ ≤ -√(2 * k * σ * l) := by + have : √(2 * k * σ * l) ≤ -s - k * -μ := sqrt_two_mul_le_sub hk (by grind) + linarith + +lemma prob_empMean_bestAction_sub_actionMean_le_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) + {σ2 : ℝ≥0} (hσ2 : 0 < σ2) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) + {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : + P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ + empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω ≤ + -√(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω))} + ≤ ENNReal.ofReal ((n - 1) * δ) := by + have := h.measurable_E + have := h.measurable_A + have := h.measurable_R + let S := {(e, τ) | ∃ t < n, pullCount IT.action (bestAction κ id e) t τ ≠ 0 ∧ + sumRewards IT.action IT.reward (bestAction κ id e) t τ - + pullCount IT.action (bestAction κ id e) t τ * actionMean κ id (bestAction κ id e) e ≤ + -√(2 * pullCount IT.action (bestAction κ id e) t τ * σ2 * Real.log (1 / δ))} + calc + _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by + rw [Measure.map_apply (by fun_prop) (by measurability)] + apply measure_mono + intro ω ⟨t, ht, hpc, hle⟩ + rw [empMean] at hle + exact ⟨t, ht, hpc, sub_le_neg_sqrt_two_mul hpc hle⟩ + _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by + rw [← compProd_map_condDistrib (by fun_prop)] + _ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := + Measure.compProd_apply (by measurability) + _ ≤ ∫⁻ e, ENNReal.ofReal ((n - 1) * δ) ∂(P.map E) := by + apply lintegral_mono_ae + rw [h.hasLaw_env.map_eq] + filter_upwards [h.ae_IsAlgEnvSeq] with e he + exact Bandits.prob_sumRewards_sub_pullCount_mul_le_le (ν := κ.sectR e) hσ2 (hs e _) he + hδ + _ = ENNReal.ofReal ((n - 1) * δ) := by + simp [Measure.map_apply h.measurable_E] + +end Learning.IsBayesAlgEnvSeq diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 73bcc699..88f2819c 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -5,8 +5,9 @@ Authors: Rémy Degenne, Paulo Rauber -/ module -public import LeanMachineLearning.Online.Bandit.SumRewards public import LeanMachineLearning.MeasureTheory.Constructions.BorelSpace.MeasurableArgMax +public import LeanMachineLearning.Online.Bandit.Regret +public import LeanMachineLearning.SequentialLearning.StationaryEnv /-! # Bayesian stationary environments -/ @@ -240,101 +241,6 @@ lemma ae_IsAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P end CondDistribIsAlgEnvSeq -section HasSubgaussianMGF - -variable {K : ℕ} [Nonempty (Fin K)] -variable {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] {alg : Algorithm (Fin K) ℝ} -variable {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → ℝ} -variable [IsProbabilityMeasure P] - -/-- Auxiliary lemma for `prob_empMean_sub_actionMean_ge_le`. -/ -private lemma sqrt_two_mul_le_sub {k : ℕ} (hk : k ≠ 0) {s μ σ l : ℝ} - (h : √(2 * σ * l / k) ≤ s / k - μ) : √(2 * k * σ * l) ≤ s - k * μ := by - have hkp : (0 : ℝ) < k := by positivity - calc √(2 * k * σ * l) - _ = √(2 * σ * l / k * k ^ 2) := by - field_simp - _ = √(2 * σ * l / k) * k := by - rw [Real.sqrt_mul' _ (sq_nonneg _), Real.sqrt_sq hkp.le] - _ ≤ (s / k - μ) * k := by - nlinarith - _ = s - k * μ := by - field_simp - -lemma prob_empMean_sub_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} - (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : - P {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ empMean A R' a t ω - actionMean κ E a ω} - ≤ ENNReal.ofReal (K * (n - 1) * δ) := by - have := h.measurable_E - have := h.measurable_A - have := h.measurable_R - let S := {(e, τ) | ∃ a, ∃ t < n, pullCount IT.action a t τ ≠ 0 ∧ - √(2 * pullCount IT.action a t τ * σ2 * Real.log (1 / δ)) ≤ - sumRewards IT.action IT.reward a t τ - pullCount IT.action a t τ * actionMean κ id a e} - calc - _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by - rw [Measure.map_apply (by fun_prop) (by measurability)] - apply measure_mono - intro ω ⟨t, ht, a, hpc, hle⟩ - rw [empMean] at hle - exact ⟨a, t, ht, hpc, sqrt_two_mul_le_sub hpc hle⟩ - _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by - rw [← compProd_map_condDistrib (by fun_prop)] - _ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := - Measure.compProd_apply (by measurability) - _ ≤ ∫⁻ e, ENNReal.ofReal (Fintype.card (Fin K) * (n - 1) * δ) ∂(P.map E) := by - apply lintegral_mono_ae - rw [h.hasLaw_env.map_eq] - filter_upwards [h.ae_IsAlgEnvSeq] with e he - exact Bandits.prob_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 (hs e) he hδ - _ = ENNReal.ofReal (K * (n - 1) * δ) := by - simp [Measure.map_apply h.measurable_E] - -/-- Auxiliary lemma for `prob_empMean_bestAction_sub_actionMean_le_le`. -/ -private lemma sub_le_neg_sqrt_two_mul {k : ℕ} (hk : k ≠ 0) {s μ σ l : ℝ} - (h : s / k - μ ≤ -√(2 * σ * l / k)) : s - k * μ ≤ -√(2 * k * σ * l) := by - have : √(2 * k * σ * l) ≤ -s - k * -μ := sqrt_two_mul_le_sub hk (by grind) - linarith - -lemma prob_empMean_bestAction_sub_actionMean_le_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) - {σ2 : ℝ≥0} (hσ2 : 0 < σ2) - (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) - {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : - P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ - empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω ≤ - -√(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω))} - ≤ ENNReal.ofReal ((n - 1) * δ) := by - have := h.measurable_E - have := h.measurable_A - have := h.measurable_R - let S := {(e, τ) | ∃ t < n, pullCount IT.action (bestAction κ id e) t τ ≠ 0 ∧ - sumRewards IT.action IT.reward (bestAction κ id e) t τ - - pullCount IT.action (bestAction κ id e) t τ * actionMean κ id (bestAction κ id e) e ≤ - -√(2 * pullCount IT.action (bestAction κ id e) t τ * σ2 * Real.log (1 / δ))} - calc - _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by - rw [Measure.map_apply (by fun_prop) (by measurability)] - apply measure_mono - intro ω ⟨t, ht, hpc, hle⟩ - rw [empMean] at hle - exact ⟨t, ht, hpc, sub_le_neg_sqrt_two_mul hpc hle⟩ - _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by - rw [← compProd_map_condDistrib (by fun_prop)] - _ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := - Measure.compProd_apply (by measurability) - _ ≤ ∫⁻ e, ENNReal.ofReal ((n - 1) * δ) ∂(P.map E) := by - apply lintegral_mono_ae - rw [h.hasLaw_env.map_eq] - filter_upwards [h.ae_IsAlgEnvSeq] with e he - exact Bandits.prob_sumRewards_sub_pullCount_mul_le_le (ν := κ.sectR e) hσ2 (hs e _) he - hδ - _ = ENNReal.ofReal ((n - 1) * δ) := by - simp [Measure.map_apply h.measurable_E] - -end HasSubgaussianMGF - end IsBayesAlgEnvSeq section IsAlgEnvSeq From 3b5e07599bc74d9b3003b0f5cdddff3e5ff1b91a Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 8 May 2026 10:32:10 +0100 Subject: [PATCH 119/155] Refactor HasCondDistrib (in progress) --- .../Probability/HasCondDistrib.lean | 40 ++++++++----------- 1 file changed, 17 insertions(+), 23 deletions(-) diff --git a/LeanMachineLearning/Probability/HasCondDistrib.lean b/LeanMachineLearning/Probability/HasCondDistrib.lean index ecc43593..60f786bb 100644 --- a/LeanMachineLearning/Probability/HasCondDistrib.lean +++ b/LeanMachineLearning/Probability/HasCondDistrib.lean @@ -224,31 +224,25 @@ lemma HasLaw.prod_of_hasCondDistrib {P : Measure β} [IsFiniteMeasure μ] [IsSFi rw [← h1.map_eq] exact h2.condDistrib_eq --- Claude -lemma HasCondDistrib.of_compProd [IsFiniteMeasure μ] [IsFiniteKernel κ] - {Z : α → Ω'} {η : Kernel (β × Ω) Ω'} [IsMarkovKernel η] - (h : HasCondDistrib (fun ω ↦ (Y ω, Z ω)) X (κ ⊗ₖ η) μ) : - HasCondDistrib Z (fun ω ↦ (X ω, Y ω)) η μ := by - have hY : AEMeasurable Y μ := h.aemeasurable_fst.fst +lemma HasCondDistrib.of_compProd [IsFiniteMeasure μ] [IsFiniteKernel κ] {Z : α → Ω'} + {η : Kernel (β × Ω) Ω'} [IsMarkovKernel η] + (h : HasCondDistrib (fun a ↦ (Y a, Z a)) X (κ ⊗ₖ η) μ) : + HasCondDistrib Z (fun a ↦ (X a, Y a)) η μ := by have hZ : AEMeasurable Z μ := h.aemeasurable_fst.snd have hX : AEMeasurable X μ := h.aemeasurable_snd - refine ⟨hZ, by fun_prop, ?_⟩ - have h_eq := h.condDistrib_eq - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop) _] at h_eq ⊢ - have h_assoc : (μ.map X ⊗ₘ κ ⊗ₘ η).map MeasurableEquiv.prodAssoc = μ.map X ⊗ₘ (κ ⊗ₖ η) := - Measure.compProd_assoc' - calc μ.map (fun ω ↦ ((X ω, Y ω), Z ω)) - _ = (μ.map (fun ω ↦ (X ω, Y ω, Z ω))).map MeasurableEquiv.prodAssoc.symm := by - rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)]; rfl - _ = (μ.map X ⊗ₘ (κ ⊗ₖ η)).map MeasurableEquiv.prodAssoc.symm := by rw [h_eq] - _ = μ.map X ⊗ₘ κ ⊗ₘ η := by - rw [← h_assoc, Measure.map_map (by fun_prop) (by fun_prop)] - simp only [MeasurableEquiv.symm_comp_self, Measure.map_id] - _ = μ.map (fun ω ↦ (X ω, Y ω)) ⊗ₘ η := by - have h_fst := h.fst - rw [Kernel.fst_compProd] at h_fst - have h_fst_eq := (condDistrib_ae_eq_iff_measure_eq_compProd X hY _).mp h_fst.condDistrib_eq - rw [h_fst_eq] + have hY : AEMeasurable Y μ := h.aemeasurable_fst.fst + refine ⟨hZ, (hX.prodMk hY), ?_⟩ + have hc := h.condDistrib_eq + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at hc ⊢ + calc μ.map (fun a ↦ ((X a, Y a), Z a)) + _ = (μ.map X ⊗ₘ (κ ⊗ₖ η)).map MeasurableEquiv.prodAssoc.symm := by + rw [← hc, AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + rfl + _ = μ.map X ⊗ₘ κ ⊗ₘ η := + Measure.compProd_assoc + _ = μ.map (fun a ↦ (X a, Y a)) ⊗ₘ η := by + rw [← (condDistrib_ae_eq_iff_measure_eq_compProd X hY κ).1] + simpa using h.fst.condDistrib_eq lemma HasCondDistrib.prod [IsFiniteMeasure μ] [IsFiniteKernel κ] {Z : α → Ω'} {η : Kernel (β × Ω) Ω'} [IsFiniteKernel η] From a26079f7511f485493ce614be17e4f5746dc422b Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 8 May 2026 14:30:16 +0100 Subject: [PATCH 120/155] Refactor HasCondDistrib (in progress) --- .../Probability/HasCondDistrib.lean | 90 ++++++------------- .../BayesStationaryEnv.lean | 8 +- 2 files changed, 29 insertions(+), 69 deletions(-) diff --git a/LeanMachineLearning/Probability/HasCondDistrib.lean b/LeanMachineLearning/Probability/HasCondDistrib.lean index 60f786bb..ba14bb45 100644 --- a/LeanMachineLearning/Probability/HasCondDistrib.lean +++ b/LeanMachineLearning/Probability/HasCondDistrib.lean @@ -75,48 +75,34 @@ lemma HasCondDistrib.snd {Y : α → Ω × Ω'} {κ : Kernel β (Ω × Ω')} [Is rw [Kernel.snd_eq] exact HasCondDistrib.comp h measurable_snd +/-- Rename to `HasCondDistrib.comp_right`? -/ +lemma HasCondDistrib.comp_right' [IsFiniteMeasure μ] [IsFiniteKernel κ] {f : γ → β} + (hf : Measurable f) {Z : α → γ} (h : HasCondDistrib Y Z (κ.comap f hf) μ) : + HasCondDistrib Y (f ∘ Z) κ μ := by + have hY : AEMeasurable Y μ := h.aemeasurable_fst + have hZ : AEMeasurable Z μ := h.aemeasurable_snd + have hfZ : AEMeasurable (f ∘ Z) μ := hf.comp_aemeasurable hZ + refine ⟨hY, hfZ, ?_⟩ + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ hY] + calc μ.map (fun a ↦ ((f ∘ Z) a, Y a)) + _ = (μ.map (fun a ↦ (Z a, Y a))).map (Prod.map f id) := by + rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (hZ.prodMk hY)] + rfl + _ = (μ.map Z ⊗ₘ κ.comap f hf).map (Prod.map f id) := by + rw [(condDistrib_ae_eq_iff_measure_eq_compProd Z hY _).mp h.condDistrib_eq] + _ = (μ.map Z).map f ⊗ₘ κ := by + ext s hs + rw [Measure.map_apply (by fun_prop) hs, Measure.compProd_apply (by measurability), + Measure.compProd_apply hs, lintegral_map (Kernel.measurable_kernel_prodMk_left hs) hf] + rfl + _ = μ.map (f ∘ Z) ⊗ₘ κ := by + rw [AEMeasurable.map_map_of_aemeasurable hf.aemeasurable hZ] + lemma HasCondDistrib.comp_right [IsFiniteMeasure μ] [IsFiniteKernel κ] (h : HasCondDistrib Y X κ μ) (f : β ≃ᵐ γ) : HasCondDistrib Y (f ∘ X) (κ.comap f.symm (by fun_prop) : Kernel γ Ω) μ := 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) + apply HasCondDistrib.comp_right' f.measurable + simpa [← Kernel.comap_comp_right] lemma HasCondDistrib.prod_right [IsFiniteMeasure μ] [IsFiniteKernel κ] (h : HasCondDistrib Y X κ μ) {f : β → γ} (hf : Measurable f) : @@ -267,30 +253,6 @@ lemma HasCondDistrib.prod [IsFiniteMeasure μ] [IsFiniteKernel κ] AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] rfl --- Claude -lemma HasCondDistrib.comp_left [IsFiniteMeasure μ] [IsFiniteKernel κ] {f : γ → β} - (hf : Measurable f) {Z : α → γ} (h : HasCondDistrib Y Z (κ.comap f hf) μ) : - HasCondDistrib Y (f ∘ Z) κ μ where - aemeasurable_fst := h.aemeasurable_fst - aemeasurable_snd := hf.comp_aemeasurable h.aemeasurable_snd - condDistrib_eq := by - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ h.aemeasurable_fst] - calc μ.map (fun ω ↦ ((f ∘ Z) ω, Y ω)) - _ = (μ.map (fun ω ↦ (Z ω, Y ω))).map (Prod.map f id) := by - rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) - (h.aemeasurable_snd.prodMk h.aemeasurable_fst)]; rfl - _ = (μ.map Z ⊗ₘ κ.comap f hf).map (Prod.map f id) := by - rw [(condDistrib_ae_eq_iff_measure_eq_compProd Z h.aemeasurable_fst _).mp h.condDistrib_eq] - _ = μ.map (f ∘ Z) ⊗ₘ κ := by - rw [← AEMeasurable.map_map_of_aemeasurable hf.aemeasurable h.aemeasurable_snd] - ext s hs - rw [Measure.map_apply (by fun_prop) hs, Measure.compProd_apply hs, - Measure.compProd_apply (hs.preimage (by fun_prop))] - rw [lintegral_map (Kernel.measurable_kernel_prodMk_left hs) hf] - refine lintegral_congr fun x ↦ ?_ - rw [Kernel.comap_apply] - congr 1 - /-- Transfer a `HasCondDistrib` from the outer probability space to the conditional distribution of `W` given `Z`. If `g ∘ W` is conditionally distributed as `η` given `(Z, f ∘ W)`, then in the conditional space given `Z = z`, `g` is conditionally distributed as `η.sectR z` given `f`. -/ @@ -340,6 +302,4 @@ lemma HasCondDistrib.ae_hasCondDistrib_sectL [IsFiniteMeasure μ] ∀ᵐ z ∂(μ.map Z), HasCondDistrib g f (η.sectL z) (condDistrib W Z μ z) := (hcd.comp_right .prodComm).ae_hasCondDistrib_sectR hf hg hW hZ - - end ProbabilityTheory diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 88f2819c..af8c2039 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -168,11 +168,11 @@ lemma hasLaw_action_zero [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ alg lemma hasCondDistrib_action' (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) (alg.policy n) P := - (h.hasCondDistrib_action n).comp_left (by fun_prop) + (h.hasCondDistrib_action n).comp_right' (by fun_prop) lemma hasCondDistrib_reward' [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : HasCondDistrib (R' (n + 1)) (fun ω ↦ (E ω, A (n + 1) ω)) κ P := - (h.hasCondDistrib_reward n).comp_left (by fun_prop) + (h.hasCondDistrib_reward n).comp_right' (by fun_prop) end Laws @@ -280,7 +280,7 @@ lemma IsAlgEnvSeq.isBayesAlgEnvSeq have hc : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) (((alg.policy n).comap Prod.snd (by fun_prop)).comap f (by fun_prop)) P := h.hasCondDistrib_action n - exact hc.comp_left (f := f) + exact hc.comp_right' (f := f) hasCondDistrib_reward n := by let f : (Iic n → α × 𝓔 × R) × α → (Iic n → α × R) × 𝓔 × α := fun p ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), (p.1 ⟨0, by simp⟩).2.1, p.2) @@ -288,7 +288,7 @@ lemma IsAlgEnvSeq.isBayesAlgEnvSeq (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) ((Kernel.prodMkLeft ((Iic n) → α × R) κ).comap f (by fun_prop)) P := by simpa [bayesStationaryEnv, Kernel.snd_prod] using (h.hasCondDistrib_reward n).snd - exact hc.comp_left (by fun_prop) + exact hc.comp_right' (by fun_prop) end IsAlgEnvSeq From 30fa0052538696d50c8f2c3988fa5cbc07e52944 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sat, 9 May 2026 10:41:38 +0200 Subject: [PATCH 121/155] fix --- LeanMachineLearning/Online/Bandit/SumRewards.lean | 1 - LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean | 1 + 2 files changed, 1 insertion(+), 1 deletion(-) diff --git a/LeanMachineLearning/Online/Bandit/SumRewards.lean b/LeanMachineLearning/Online/Bandit/SumRewards.lean index 18b98bfb..4a4b3dda 100644 --- a/LeanMachineLearning/Online/Bandit/SumRewards.lean +++ b/LeanMachineLearning/Online/Bandit/SumRewards.lean @@ -8,7 +8,6 @@ module public import LeanMachineLearning.Online.Bandit.ArrayProbSpace public import LeanMachineLearning.Probability.Moments.SubGaussian public import LeanMachineLearning.SequentialLearning.BayesStationaryEnv -public import LeanMachineLearning.SequentialLearning.IonescuTulceaSpace /-! # Law of the sum of rewards -/ diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index af8c2039..f8e84f74 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -7,6 +7,7 @@ module public import LeanMachineLearning.MeasureTheory.Constructions.BorelSpace.MeasurableArgMax public import LeanMachineLearning.Online.Bandit.Regret +public import LeanMachineLearning.SequentialLearning.IonescuTulceaSpace public import LeanMachineLearning.SequentialLearning.StationaryEnv /-! # Bayesian stationary environments -/ From 71c3ba1d27662d62a686cf828ed0b06a5d15b78c Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 14 May 2026 10:58:11 +0100 Subject: [PATCH 122/155] Refactor HasCondDistrib.lean (in progress) --- .../Probability/HasCondDistrib.lean | 58 ++++++------------- .../Probability/Independence/CondDistrib.lean | 29 ++++++++++ .../BayesStationaryEnv.lean | 6 +- 3 files changed, 48 insertions(+), 45 deletions(-) diff --git a/LeanMachineLearning/Probability/HasCondDistrib.lean b/LeanMachineLearning/Probability/HasCondDistrib.lean index ba14bb45..16baeeb6 100644 --- a/LeanMachineLearning/Probability/HasCondDistrib.lean +++ b/LeanMachineLearning/Probability/HasCondDistrib.lean @@ -253,53 +253,29 @@ lemma HasCondDistrib.prod [IsFiniteMeasure μ] [IsFiniteKernel κ] AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] rfl -/-- Transfer a `HasCondDistrib` from the outer probability space to the conditional distribution -of `W` given `Z`. If `g ∘ W` is conditionally distributed as `η` given `(Z, f ∘ W)`, then in the -conditional space given `Z = z`, `g` is conditionally distributed as `η.sectR z` given `f`. -/ lemma HasCondDistrib.ae_hasCondDistrib_sectR [IsFiniteMeasure μ] [StandardBorelSpace β] [Nonempty β] - {δ : Type*} [MeasurableSpace δ] [StandardBorelSpace δ] [Nonempty δ] - {W : α → δ} {Z : α → γ} - {f : δ → β} {g : δ → Ω} + {W : α → Ω'} {Z : α → γ} + {f : Ω' → β} {g : Ω' → Ω} {η : Kernel (γ × β) Ω} [IsFiniteKernel η] (hf : Measurable f) (hg : Measurable g) - (hW : AEMeasurable W μ) (hZ : AEMeasurable Z μ) - (hcd : HasCondDistrib (g ∘ W) (fun ω ↦ (Z ω, (f ∘ W) ω)) η μ) : + (hW : AEMeasurable W μ) + (hcd : HasCondDistrib (g ∘ W) (fun a ↦ (Z a, (f (W a)))) η μ) : ∀ᵐ z ∂(μ.map Z), HasCondDistrib g f (η.sectR z) (condDistrib W Z μ z) := by have hfW := hf.comp_aemeasurable hW - have h_prod := condDistrib_prod_left hfW (hg.comp_aemeasurable hW) hZ (μ := μ) - have h_comp_pair : (condDistrib (fun ω ↦ ((f ∘ W) ω, (g ∘ W) ω)) Z μ) - =ᵐ[μ.map Z] (condDistrib W Z μ).map (fun w ↦ (f w, g w)) := - condDistrib_comp Z hW (hf.prodMk hg) - have h_comp_fst : (condDistrib (f ∘ W) Z μ) - =ᵐ[μ.map Z] (condDistrib W Z μ).map f := - condDistrib_comp Z hW hf - have h_eq := hcd.condDistrib_eq - rw [(compProd_map_condDistrib hfW).symm] at h_eq - have h_nested := Measure.ae_ae_of_ae_compProd h_eq - filter_upwards [h_prod, h_comp_pair, h_comp_fst, h_nested] - with z h_prod_z h_pair_z h_fst_z h_nested_z + have h_eq : (condDistrib (g ∘ W) (fun ω ↦ (Z ω, (f ∘ W) ω)) μ) + =ᵐ[μ.map Z ⊗ₘ condDistrib (f ∘ W) Z μ] η := by + rw [compProd_map_condDistrib (X := Z) hfW] + exact hcd.condDistrib_eq + filter_upwards [ + condDistrib_condDistrib_ae_eq_sectR_condDistrib hf hg hW hcd.aemeasurable_snd.fst, + condDistrib_comp Z hW hf, + Measure.ae_ae_of_ae_compProd h_eq] with z h_tower h_fst h_nested refine ⟨hg.aemeasurable, hf.aemeasurable, ?_⟩ - rw [condDistrib_ae_eq_iff_measure_eq_compProd f hg.aemeasurable, - ← Kernel.map_apply _ (hf.prodMk hg), ← h_pair_z, - ← Kernel.map_apply _ hf, ← h_fst_z, - h_prod_z, Kernel.compProd_apply_eq_compProd_sectR] - exact Measure.compProd_congr (h_nested_z.mono fun a ha ↦ by - simp only [Kernel.sectR_apply]; exact ha) - -/-- Variant of `ae_hasCondDistrib_sectR` where `Z` appears second in the conditioning pair. -If `g ∘ W` is conditionally distributed as `η` given `(f ∘ W, Z)`, then in the conditional space -given `Z = z`, `g` is conditionally distributed as `η.sectL z` given `f`. -/ -lemma HasCondDistrib.ae_hasCondDistrib_sectL [IsFiniteMeasure μ] - [StandardBorelSpace β] [Nonempty β] - {δ : Type*} [MeasurableSpace δ] [StandardBorelSpace δ] [Nonempty δ] - {W : α → δ} {Z : α → γ} - {f : δ → β} {g : δ → Ω} - {η : Kernel (β × γ) Ω} [IsFiniteKernel η] - (hf : Measurable f) (hg : Measurable g) - (hW : AEMeasurable W μ) (hZ : AEMeasurable Z μ) - (hcd : HasCondDistrib (g ∘ W) (fun ω ↦ ((f ∘ W) ω, Z ω)) η μ) : - ∀ᵐ z ∂(μ.map Z), HasCondDistrib g f (η.sectL z) (condDistrib W Z μ z) := - (hcd.comp_right .prodComm).ae_hasCondDistrib_sectR hf hg hW hZ + refine h_tower.trans ?_ + have h_meas : (condDistrib W Z μ z).map f = condDistrib (f ∘ W) Z μ z := by + rw [← Kernel.map_apply _ hf, ← h_fst] + rw [h_meas] + exact h_nested.mono fun b hb ↦ by simp only [Kernel.sectR_apply]; exact hb end ProbabilityTheory diff --git a/LeanMachineLearning/Probability/Independence/CondDistrib.lean b/LeanMachineLearning/Probability/Independence/CondDistrib.lean index 10980c37..0513fca9 100644 --- a/LeanMachineLearning/Probability/Independence/CondDistrib.lean +++ b/LeanMachineLearning/Probability/Independence/CondDistrib.lean @@ -47,6 +47,35 @@ lemma condDistrib_prod_left [StandardBorelSpace β] [Nonempty β] AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] rfl +lemma condDistrib_condDistrib_ae_eq_sectR_condDistrib [StandardBorelSpace β] [Nonempty β] + {f : Ω' → β} {g : Ω' → Ω} + (hf : Measurable f) (hg : Measurable g) + (hZ : AEMeasurable Z μ) (hT : AEMeasurable T μ) : + ∀ᵐ t ∂(μ.map T), + condDistrib g f (condDistrib Z T μ t) + =ᵐ[(condDistrib Z T μ t).map f] + (condDistrib (g ∘ Z) (fun a ↦ (T a, f (Z a))) μ).sectR t := by + have hfZ := hf.comp_aemeasurable hZ + have hgZ := hg.comp_aemeasurable hZ + filter_upwards [ + condDistrib_prod_left hfZ hgZ hT, + condDistrib_comp T hZ (hf.prodMk hg), + condDistrib_comp T hZ hf] with t h_prod h_pair h_fst + rw [condDistrib_ae_eq_iff_measure_eq_compProd f hg.aemeasurable] + calc (condDistrib Z T μ t).map (fun w ↦ (f w, g w)) + _ = ((condDistrib Z T μ).map (fun w ↦ (f w, g w))) t := + (Kernel.map_apply _ (hf.prodMk hg) t).symm + _ = condDistrib (fun ω ↦ ((f ∘ Z) ω, (g ∘ Z) ω)) T μ t := h_pair.symm + _ = (condDistrib (f ∘ Z) T μ ⊗ₖ condDistrib (g ∘ Z) (fun ω ↦ (T ω, (f ∘ Z) ω)) μ) t := h_prod + _ = condDistrib (f ∘ Z) T μ t + ⊗ₘ (condDistrib (g ∘ Z) (fun ω ↦ (T ω, (f ∘ Z) ω)) μ).sectR t := + Kernel.compProd_apply_eq_compProd_sectR _ _ t + _ = ((condDistrib Z T μ).map f) t + ⊗ₘ (condDistrib (g ∘ Z) (fun ω ↦ (T ω, (f ∘ Z) ω)) μ).sectR t := by rw [h_fst] + _ = (condDistrib Z T μ t).map f + ⊗ₘ (condDistrib (g ∘ Z) (fun ω ↦ (T ω, (f ∘ Z) ω)) μ).sectR t := by + rw [Kernel.map_apply _ hf] + lemma condDistrib_prod_self_left [StandardBorelSpace β] [Nonempty β] [StandardBorelSpace γ] [Nonempty γ] (hX : AEMeasurable X μ) (hT : AEMeasurable T μ) : diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index af8c2039..10345891 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -196,7 +196,6 @@ lemma hasCondDistrib_IT_reward_zero [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q exact h.hasCondDistrib_reward_zero.ae_hasCondDistrib_sectR (IT.measurable_action 0) (IT.measurable_reward 0) (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable - h.measurable_E.aemeasurable lemma hasCondDistrib_IT_action (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : ∀ᵐ e ∂Q, HasCondDistrib (IT.action (n + 1)) (IT.hist n) (alg.policy n) @@ -204,8 +203,7 @@ lemma hasCondDistrib_IT_action (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ rw [← h.hasLaw_env.map_eq] filter_upwards [(h.hasCondDistrib_action n).ae_hasCondDistrib_sectR (IT.measurable_hist n) (IT.measurable_action (n + 1)) - (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable - h.measurable_E.aemeasurable] with _ he + (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable] with _ he rwa [Kernel.sectR_prodMkLeft] at he lemma hasCondDistrib_IT_reward [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : @@ -219,7 +217,7 @@ lemma hasCondDistrib_IT_reward [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ al ((MeasurableEquiv.prodCongr .prodComm (.refl _)).trans .prodAssoc)) exact hc.ae_hasCondDistrib_sectR ((IT.measurable_hist n).prodMk (IT.measurable_action (n + 1))) (IT.measurable_reward (n + 1)) - (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable h.measurable_E.aemeasurable + (measurable_trajectory h.measurable_A h.measurable_R).aemeasurable lemma hasLaw_IT_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : ∀ᵐ e ∂Q, HasLaw (IT.hist n) (condDistrib (IsAlgEnvSeq.hist A R' n) E P e) From 4457e9ebfb01082e35148bfed05a2cb2b0aea9bf Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 15 May 2026 15:56:33 +0100 Subject: [PATCH 123/155] Refactor HasCondDistrib.lean (in progress) --- .../Probability/HasCondDistrib.lean | 33 ++++++++----------- .../Probability/Independence/CondDistrib.lean | 31 +++++++---------- 2 files changed, 25 insertions(+), 39 deletions(-) diff --git a/LeanMachineLearning/Probability/HasCondDistrib.lean b/LeanMachineLearning/Probability/HasCondDistrib.lean index bdc8c0a8..ee50efe5 100644 --- a/LeanMachineLearning/Probability/HasCondDistrib.lean +++ b/LeanMachineLearning/Probability/HasCondDistrib.lean @@ -253,29 +253,22 @@ lemma HasCondDistrib.prod [IsFiniteMeasure μ] [IsFiniteKernel κ] AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] rfl -lemma HasCondDistrib.ae_hasCondDistrib_sectR [IsFiniteMeasure μ] - [StandardBorelSpace β] [Nonempty β] - {W : α → Ω'} {Z : α → γ} - {f : Ω' → β} {g : Ω' → Ω} - {η : Kernel (γ × β) Ω} [IsFiniteKernel η] - (hf : Measurable f) (hg : Measurable g) - (hW : AEMeasurable W μ) - (hcd : HasCondDistrib (g ∘ W) (fun a ↦ (Z a, (f (W a)))) η μ) : +lemma HasCondDistrib.ae_hasCondDistrib_sectR [IsFiniteMeasure μ] [StandardBorelSpace β] [Nonempty β] + {W : α → Ω'} {Z : α → γ} {f : Ω' → β} {g : Ω' → Ω} {η : Kernel (γ × β) Ω} [IsFiniteKernel η] + (hf : Measurable f) (hg : Measurable g) (hW : AEMeasurable W μ) + (hcd : HasCondDistrib (g ∘ W) (fun a ↦ (Z a, f (W a))) η μ) : ∀ᵐ z ∂(μ.map Z), HasCondDistrib g f (η.sectR z) (condDistrib W Z μ z) := by - have hfW := hf.comp_aemeasurable hW - have h_eq : (condDistrib (g ∘ W) (fun ω ↦ (Z ω, (f ∘ W) ω)) μ) - =ᵐ[μ.map Z ⊗ₘ condDistrib (f ∘ W) Z μ] η := by - rw [compProd_map_condDistrib (X := Z) hfW] + have h_eq : (condDistrib (g ∘ W) (fun a ↦ (Z a, f (W a))) μ) + =ᵐ[μ.map Z ⊗ₘ (condDistrib W Z μ).map f] η := by + rw [← Measure.compProd_congr (condDistrib_comp Z hW hf), + compProd_map_condDistrib (X := Z) (hf.comp_aemeasurable hW)] exact hcd.condDistrib_eq filter_upwards [ - condDistrib_condDistrib_ae_eq_sectR_condDistrib hf hg hW hcd.aemeasurable_snd.fst, - condDistrib_comp Z hW hf, - Measure.ae_ae_of_ae_compProd h_eq] with z h_tower h_fst h_nested + ae_condDistrib_condDistrib_ae_eq_sectR_condDistrib hf hg hW hcd.aemeasurable_snd.fst, + Measure.ae_ae_of_ae_compProd h_eq] with z ht hn refine ⟨hg.aemeasurable, hf.aemeasurable, ?_⟩ - refine h_tower.trans ?_ - have h_meas : (condDistrib W Z μ z).map f = condDistrib (f ∘ W) Z μ z := by - rw [← Kernel.map_apply _ hf, ← h_fst] - rw [h_meas] - exact h_nested.mono fun b hb ↦ by simp only [Kernel.sectR_apply]; exact hb + apply ht.trans + rw [← Kernel.map_apply _ hf] + exact hn.mono (fun _ hb ↦ hb) end ProbabilityTheory diff --git a/LeanMachineLearning/Probability/Independence/CondDistrib.lean b/LeanMachineLearning/Probability/Independence/CondDistrib.lean index 0513fca9..cb343e7d 100644 --- a/LeanMachineLearning/Probability/Independence/CondDistrib.lean +++ b/LeanMachineLearning/Probability/Independence/CondDistrib.lean @@ -47,34 +47,27 @@ lemma condDistrib_prod_left [StandardBorelSpace β] [Nonempty β] AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] rfl -lemma condDistrib_condDistrib_ae_eq_sectR_condDistrib [StandardBorelSpace β] [Nonempty β] - {f : Ω' → β} {g : Ω' → Ω} - (hf : Measurable f) (hg : Measurable g) - (hZ : AEMeasurable Z μ) (hT : AEMeasurable T μ) : +lemma ae_condDistrib_condDistrib_ae_eq_sectR_condDistrib [StandardBorelSpace β] [Nonempty β] + {f : Ω' → β} {g : Ω' → Ω} (hf : Measurable f) (hg : Measurable g) (hZ : AEMeasurable Z μ) + (hT : AEMeasurable T μ) : ∀ᵐ t ∂(μ.map T), - condDistrib g f (condDistrib Z T μ t) - =ᵐ[(condDistrib Z T μ t).map f] - (condDistrib (g ∘ Z) (fun a ↦ (T a, f (Z a))) μ).sectR t := by - have hfZ := hf.comp_aemeasurable hZ - have hgZ := hg.comp_aemeasurable hZ + condDistrib g f (condDistrib Z T μ t) =ᵐ[(condDistrib Z T μ t).map f] + (condDistrib (g ∘ Z) (fun a ↦ (T a, f (Z a))) μ).sectR t := by filter_upwards [ - condDistrib_prod_left hfZ hgZ hT, + condDistrib_prod_left (hf.comp_aemeasurable hZ) (hg.comp_aemeasurable hZ) hT, condDistrib_comp T hZ (hf.prodMk hg), condDistrib_comp T hZ hf] with t h_prod h_pair h_fst rw [condDistrib_ae_eq_iff_measure_eq_compProd f hg.aemeasurable] calc (condDistrib Z T μ t).map (fun w ↦ (f w, g w)) - _ = ((condDistrib Z T μ).map (fun w ↦ (f w, g w))) t := - (Kernel.map_apply _ (hf.prodMk hg) t).symm - _ = condDistrib (fun ω ↦ ((f ∘ Z) ω, (g ∘ Z) ω)) T μ t := h_pair.symm - _ = (condDistrib (f ∘ Z) T μ ⊗ₖ condDistrib (g ∘ Z) (fun ω ↦ (T ω, (f ∘ Z) ω)) μ) t := h_prod + _ = condDistrib (fun ω ↦ ((f ∘ Z) ω, (g ∘ Z) ω)) T μ t := by + rw [← Kernel.map_apply _ (hf.prodMk hg)] + exact h_pair.symm _ = condDistrib (f ∘ Z) T μ t - ⊗ₘ (condDistrib (g ∘ Z) (fun ω ↦ (T ω, (f ∘ Z) ω)) μ).sectR t := - Kernel.compProd_apply_eq_compProd_sectR _ _ t - _ = ((condDistrib Z T μ).map f) t - ⊗ₘ (condDistrib (g ∘ Z) (fun ω ↦ (T ω, (f ∘ Z) ω)) μ).sectR t := by rw [h_fst] + ⊗ₘ (condDistrib (g ∘ Z) (fun ω ↦ (T ω, (f ∘ Z) ω)) μ).sectR t := by + rw [h_prod, Kernel.compProd_apply_eq_compProd_sectR] _ = (condDistrib Z T μ t).map f ⊗ₘ (condDistrib (g ∘ Z) (fun ω ↦ (T ω, (f ∘ Z) ω)) μ).sectR t := by - rw [Kernel.map_apply _ hf] + rw [h_fst, Kernel.map_apply _ hf] lemma condDistrib_prod_self_left [StandardBorelSpace β] [Nonempty β] [StandardBorelSpace γ] [Nonempty γ] From 7dbacffeba70f772738d306ce12b234b12ef2863 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 18 May 2026 15:13:42 +0100 Subject: [PATCH 124/155] Refactor HasCondDistrib.lean --- .../Probability/HasCondDistrib.lean | 17 ++++++++--------- .../Probability/Independence/CondDistrib.lean | 18 +++++++----------- .../SequentialLearning/BayesStationaryEnv.lean | 6 +++--- 3 files changed, 18 insertions(+), 23 deletions(-) diff --git a/LeanMachineLearning/Probability/HasCondDistrib.lean b/LeanMachineLearning/Probability/HasCondDistrib.lean index ee50efe5..ec8a8789 100644 --- a/LeanMachineLearning/Probability/HasCondDistrib.lean +++ b/LeanMachineLearning/Probability/HasCondDistrib.lean @@ -253,22 +253,21 @@ lemma HasCondDistrib.prod [IsFiniteMeasure μ] [IsFiniteKernel κ] AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] rfl -lemma HasCondDistrib.ae_hasCondDistrib_sectR [IsFiniteMeasure μ] [StandardBorelSpace β] [Nonempty β] +lemma HasCondDistrib.hasCondDistrib_sectR [IsFiniteMeasure μ] [StandardBorelSpace β] [Nonempty β] {W : α → Ω'} {Z : α → γ} {f : Ω' → β} {g : Ω' → Ω} {η : Kernel (γ × β) Ω} [IsFiniteKernel η] (hf : Measurable f) (hg : Measurable g) (hW : AEMeasurable W μ) - (hcd : HasCondDistrib (g ∘ W) (fun a ↦ (Z a, f (W a))) η μ) : + (hcd : HasCondDistrib (g ∘ W) (fun a ↦ (Z a, (f ∘ W) a)) η μ) : ∀ᵐ z ∂(μ.map Z), HasCondDistrib g f (η.sectR z) (condDistrib W Z μ z) := by - have h_eq : (condDistrib (g ∘ W) (fun a ↦ (Z a, f (W a))) μ) + have h_eq : condDistrib (g ∘ W) (fun a ↦ (Z a, (f ∘ W) a)) μ =ᵐ[μ.map Z ⊗ₘ (condDistrib W Z μ).map f] η := by rw [← Measure.compProd_congr (condDistrib_comp Z hW hf), - compProd_map_condDistrib (X := Z) (hf.comp_aemeasurable hW)] + compProd_map_condDistrib (hf.comp_aemeasurable hW)] exact hcd.condDistrib_eq filter_upwards [ - ae_condDistrib_condDistrib_ae_eq_sectR_condDistrib hf hg hW hcd.aemeasurable_snd.fst, - Measure.ae_ae_of_ae_compProd h_eq] with z ht hn + condDistrib_condDistrib_ae_eq_sectR_condDistrib hf hg hW hcd.aemeasurable_snd.fst, + Measure.ae_ae_of_ae_compProd h_eq] with z hc ha refine ⟨hg.aemeasurable, hf.aemeasurable, ?_⟩ - apply ht.trans - rw [← Kernel.map_apply _ hf] - exact hn.mono (fun _ hb ↦ hb) + rw [Kernel.map_apply _ hf] at ha + filter_upwards [hc, ha] with b hcb hab using hcb.trans hab end ProbabilityTheory diff --git a/LeanMachineLearning/Probability/Independence/CondDistrib.lean b/LeanMachineLearning/Probability/Independence/CondDistrib.lean index cb343e7d..31e2c93e 100644 --- a/LeanMachineLearning/Probability/Independence/CondDistrib.lean +++ b/LeanMachineLearning/Probability/Independence/CondDistrib.lean @@ -47,27 +47,23 @@ lemma condDistrib_prod_left [StandardBorelSpace β] [Nonempty β] AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] rfl -lemma ae_condDistrib_condDistrib_ae_eq_sectR_condDistrib [StandardBorelSpace β] [Nonempty β] +lemma condDistrib_condDistrib_ae_eq_sectR_condDistrib [StandardBorelSpace β] [Nonempty β] {f : Ω' → β} {g : Ω' → Ω} (hf : Measurable f) (hg : Measurable g) (hZ : AEMeasurable Z μ) (hT : AEMeasurable T μ) : ∀ᵐ t ∂(μ.map T), condDistrib g f (condDistrib Z T μ t) =ᵐ[(condDistrib Z T μ t).map f] - (condDistrib (g ∘ Z) (fun a ↦ (T a, f (Z a))) μ).sectR t := by + (condDistrib (g ∘ Z) (fun a ↦ (T a, (f ∘ Z) a)) μ).sectR t := by filter_upwards [ condDistrib_prod_left (hf.comp_aemeasurable hZ) (hg.comp_aemeasurable hZ) hT, - condDistrib_comp T hZ (hf.prodMk hg), - condDistrib_comp T hZ hf] with t h_prod h_pair h_fst + condDistrib_comp T hZ (hf.prodMk hg), condDistrib_comp T hZ hf] with t h_prod h_pair h_fst rw [condDistrib_ae_eq_iff_measure_eq_compProd f hg.aemeasurable] - calc (condDistrib Z T μ t).map (fun w ↦ (f w, g w)) - _ = condDistrib (fun ω ↦ ((f ∘ Z) ω, (g ∘ Z) ω)) T μ t := by + calc (condDistrib Z T μ t).map (fun ω' ↦ (f ω', g ω')) + _ = condDistrib (fun a ↦ ((f ∘ Z) a, (g ∘ Z) a)) T μ t := by rw [← Kernel.map_apply _ (hf.prodMk hg)] exact h_pair.symm - _ = condDistrib (f ∘ Z) T μ t - ⊗ₘ (condDistrib (g ∘ Z) (fun ω ↦ (T ω, (f ∘ Z) ω)) μ).sectR t := by - rw [h_prod, Kernel.compProd_apply_eq_compProd_sectR] _ = (condDistrib Z T μ t).map f - ⊗ₘ (condDistrib (g ∘ Z) (fun ω ↦ (T ω, (f ∘ Z) ω)) μ).sectR t := by - rw [h_fst, Kernel.map_apply _ hf] + ⊗ₘ (condDistrib (g ∘ Z) (fun a ↦ (T a, (f ∘ Z) a)) μ).sectR t := by + rw [h_prod, Kernel.compProd_apply_eq_compProd_sectR, h_fst, Kernel.map_apply _ hf] lemma condDistrib_prod_self_left [StandardBorelSpace β] [Nonempty β] [StandardBorelSpace γ] [Nonempty γ] diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 81f80f96..c97941e3 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -194,7 +194,7 @@ lemma hasCondDistrib_IT_feedback_zero [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq ∀ᵐ e ∂Q, HasCondDistrib (IT.feedback 0) (IT.action 0) (κ.sectR e) (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] - exact h.hasCondDistrib_feedback_zero.ae_hasCondDistrib_sectR + exact h.hasCondDistrib_feedback_zero.hasCondDistrib_sectR (IT.measurable_action 0) (IT.measurable_feedback 0) (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable @@ -202,7 +202,7 @@ lemma hasCondDistrib_IT_action (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ ∀ᵐ e ∂Q, HasCondDistrib (IT.action (n + 1)) (IT.hist n) (alg.policy n) (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] - filter_upwards [(h.hasCondDistrib_action n).ae_hasCondDistrib_sectR + filter_upwards [(h.hasCondDistrib_action n).hasCondDistrib_sectR (IT.measurable_hist n) (IT.measurable_action (n + 1)) (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable] with _ he rwa [Kernel.sectR_prodMkLeft] at he @@ -217,7 +217,7 @@ lemma hasCondDistrib_IT_feedback [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ (κ.comap (fun (e, _, a) ↦ (e, a)) (by fun_prop)) P := (h.hasCondDistrib_feedback n).comp_right (MeasurableEquiv.prodAssoc.symm.trans ((MeasurableEquiv.prodCongr .prodComm (.refl _)).trans .prodAssoc)) - exact hc.ae_hasCondDistrib_sectR ((IT.measurable_hist n).prodMk + exact hc.hasCondDistrib_sectR ((IT.measurable_hist n).prodMk (IT.measurable_action (n + 1))) (IT.measurable_feedback (n + 1)) (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable From 6b4b6560e35d46dd600a69cf0821fdffe0e8b2a8 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 18 May 2026 17:07:08 +0100 Subject: [PATCH 125/155] Refactor WithDensity.lean (in progress) --- .../Probability/HasCondDistrib.lean | 4 +- .../Probability/WithDensity.lean | 197 +++++++----------- .../BayesStationaryEnv.lean | 2 +- 3 files changed, 84 insertions(+), 119 deletions(-) diff --git a/LeanMachineLearning/Probability/HasCondDistrib.lean b/LeanMachineLearning/Probability/HasCondDistrib.lean index ec8a8789..ab497f48 100644 --- a/LeanMachineLearning/Probability/HasCondDistrib.lean +++ b/LeanMachineLearning/Probability/HasCondDistrib.lean @@ -254,8 +254,8 @@ lemma HasCondDistrib.prod [IsFiniteMeasure μ] [IsFiniteKernel κ] rfl lemma HasCondDistrib.hasCondDistrib_sectR [IsFiniteMeasure μ] [StandardBorelSpace β] [Nonempty β] - {W : α → Ω'} {Z : α → γ} {f : Ω' → β} {g : Ω' → Ω} {η : Kernel (γ × β) Ω} [IsFiniteKernel η] - (hf : Measurable f) (hg : Measurable g) (hW : AEMeasurable W μ) + {W : α → Ω'} {Z : α → γ} {f : Ω' → β} {g : Ω' → Ω} {η : Kernel (γ × β) Ω} (hf : Measurable f) + (hg : Measurable g) (hW : AEMeasurable W μ) (hcd : HasCondDistrib (g ∘ W) (fun a ↦ (Z a, (f ∘ W) a)) η μ) : ∀ᵐ z ∂(μ.map Z), HasCondDistrib g f (η.sectR z) (condDistrib W Z μ z) := by have h_eq : condDistrib (g ∘ W) (fun a ↦ (Z a, (f ∘ W) a)) μ diff --git a/LeanMachineLearning/Probability/WithDensity.lean b/LeanMachineLearning/Probability/WithDensity.lean index b7563aef..a437de4b 100644 --- a/LeanMachineLearning/Probability/WithDensity.lean +++ b/LeanMachineLearning/Probability/WithDensity.lean @@ -8,142 +8,110 @@ module public import Mathlib.Probability.Kernel.CompProdEqIff public import Mathlib.Probability.Kernel.Composition.MeasureComp -/-! -# Interactions of `withDensity` with `compProd`, `map`, and `swap` - -Lemmas for pushing `Measure.withDensity` and `Kernel.withDensity` through -`compProd`, `MeasurableEquiv.map`, `Prod.swap`, and composition. --/ - @[expose] public section open MeasureTheory ProbabilityTheory open scoped ENNReal -variable {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} - {mγ : MeasurableSpace γ} {μ : Measure α} +variable {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} +variable {μ : Measure α} namespace Measure -/-- Composing `withDensity` on the measure side of a `compProd`: -`(μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst)`. -/ -lemma withDensity_compProd_left [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] - {f : α → ℝ≥0∞} (hf : Measurable f) : - (μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst) := by - ext s hs - rw [Measure.compProd_apply hs, withDensity_apply _ hs, - lintegral_withDensity_eq_lintegral_mul₀ hf.aemeasurable - (Kernel.measurable_kernel_prodMk_left hs).aemeasurable, - ← lintegral_indicator hs, - Measure.lintegral_compProd ((hf.comp measurable_fst).indicator hs)] - congr 1 - ext a - simp_rw [Pi.mul_apply] - have : (fun b ↦ s.indicator (f ∘ Prod.fst) (a, b)) = - fun b ↦ (Prod.mk a ⁻¹' s).indicator (fun _ ↦ f a) b := by - ext b; simp only [Function.comp, Set.indicator, Set.mem_preimage]; rfl - rw [this, lintegral_indicator_const (hs.preimage (by fun_prop))] - -/-- Pushing a `withDensity` through a `MeasurableEquiv`: -`(μ.withDensity f).map e = (μ.map e).withDensity (f ∘ e.symm)`. -/ -lemma withDensity_map_equiv - {e : α ≃ᵐ β} {f : α → ℝ≥0∞} (hf : Measurable f) : - (μ.withDensity f).map e = (μ.map e).withDensity (f ∘ e.symm) := by - ext s hs - rw [Measure.map_apply e.measurable hs, - withDensity_apply _ (e.measurable hs), - withDensity_apply _ hs, Measure.restrict_map e.measurable hs, - lintegral_map (hf.comp e.symm.measurable) e.measurable] - simp_rw [Function.comp_apply, e.symm_apply_apply] - -/-- Mapping a `withDensity` through a `MeasurableEquiv` from the snd component. -/ -lemma map_swap_withDensity_fst - {μ : Measure (α × β)} - {f : β → ℝ≥0∞} (hf : Measurable f) : - (μ.withDensity (f ∘ Prod.snd)).map Prod.swap - = (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := by - ext s hs - rw [Measure.map_apply measurable_swap hs, withDensity_apply _ (measurable_swap hs), - withDensity_apply _ hs, Measure.restrict_map measurable_swap hs] - exact (lintegral_map (hf.comp measurable_fst) measurable_swap).symm - -/-- `(μ.withDensity (f ∘ g)).map g = (μ.map g).withDensity f`. -/ -lemma map_withDensity_comp - {g : α → γ} {f : γ → ℝ≥0∞} - (hg : Measurable g) (hf : Measurable f) : +lemma withDensity_compProd_left [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] {f : α → ℝ≥0∞} + (hf : Measurable f) : (μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst) := by + refine Measure.ext_of_lintegral _ fun g hg ↦ ?_ + calc ∫⁻ p, g p ∂((μ.withDensity f) ⊗ₘ κ) + = ∫⁻ a, ∫⁻ b, g (a, b) ∂κ a ∂(μ.withDensity f) := + Measure.lintegral_compProd hg + _ = ∫⁻ a, f a * ∫⁻ b, g (a, b) ∂κ a ∂μ := + lintegral_withDensity_eq_lintegral_mul _ hf hg.lintegral_kernel_prod_right' + _ = ∫⁻ a, ∫⁻ b, f a * g (a, b) ∂κ a ∂μ := by + refine lintegral_congr fun a ↦ ?_ + rw [← lintegral_const_mul _ (by fun_prop)] + _ = ∫⁻ p, (f ∘ Prod.fst) p * g p ∂(μ ⊗ₘ κ) := + (Measure.lintegral_compProd ((hf.comp measurable_fst).mul hg)).symm + _ = ∫⁻ p, g p ∂((μ ⊗ₘ κ).withDensity (f ∘ Prod.fst)) := + (lintegral_withDensity_eq_lintegral_mul _ (hf.comp measurable_fst) hg).symm + +lemma map_withDensity_comp {g : α → γ} {f : γ → ℝ≥0∞} (hg : Measurable g) (hf : Measurable f) : (μ.withDensity (f ∘ g)).map g = (μ.map g).withDensity f := by ext s hs simp only [Measure.map_apply hg hs, withDensity_apply _ (hg hs), withDensity_apply _ hs, setLIntegral_map hs hf hg, Function.comp] -/-- `(f · μ) ⊗ₘ (g · κ) = ((a, c) ↦ f a * g a c) · (μ ⊗ₘ κ)`. -/ -lemma withDensity_compProd_withDensity [SFinite μ] - {κ : Kernel α γ} [IsSFiniteKernel κ] - {f : α → ℝ≥0∞} {g : α → γ → ℝ≥0∞} - (hf : Measurable f) (hg : Measurable (Function.uncurry g)) +lemma withDensity_map_equiv {e : α ≃ᵐ β} {f : α → ℝ≥0∞} (hf : Measurable f) : + (μ.withDensity f).map e = (μ.map e).withDensity (f ∘ e.symm) := + calc (μ.withDensity f).map e + = (μ.withDensity ((f ∘ e.symm) ∘ e)).map e := by + congr + funext x + simp + _ = (μ.map e).withDensity (f ∘ e.symm) := + map_withDensity_comp e.measurable (hf.comp e.symm.measurable) + +lemma map_swap_withDensity_fst {μ : Measure (α × β)} {f : β → ℝ≥0∞} (hf : Measurable f) : + (μ.withDensity (f ∘ Prod.snd)).map Prod.swap = (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := + calc (μ.withDensity (f ∘ Prod.snd)).map Prod.swap + _ = (μ.withDensity ((f ∘ Prod.fst) ∘ Prod.swap)).map Prod.swap := + rfl + _ = (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := + map_withDensity_comp measurable_swap (hf.comp measurable_fst) + +lemma withDensity_compProd_withDensity [SFinite μ] {κ : Kernel α γ} [IsSFiniteKernel κ] + {f : α → ℝ≥0∞} {g : α → γ → ℝ≥0∞} (hf : Measurable f) (hg : Measurable (Function.uncurry g)) [IsSFiniteKernel (κ.withDensity g)] : - (μ.withDensity f) ⊗ₘ (κ.withDensity g) = - (μ ⊗ₘ κ).withDensity (fun (a, c) => f a * g a c) := by + (μ.withDensity f) ⊗ₘ (κ.withDensity g) = (μ ⊗ₘ κ).withDensity (fun (a, c) => f a * g a c) := by rw [Measure.compProd_withDensity hg, withDensity_compProd_left hf] exact (withDensity_mul _ (hf.comp measurable_fst) hg).symm -/-- If `κ =ᵐ[μ] η.withDensity (fun _ => f)`, then `μ ⊗ₘ κ = (μ ⊗ₘ η).withDensity (f ∘ Prod.snd)`. - Unlike `compProd_congr` + `compProd_withDensity`, does not need - `IsSFiniteKernel (η.withDensity (fun _ => f))`. -/ -lemma compProd_eq_compProd_withDensity [SFinite μ] - {κ η : Kernel α β} [IsSFiniteKernel κ] [IsSFiniteKernel η] - {f : β → ℝ≥0∞} (hf : Measurable f) - (h : κ =ᵐ[μ] η.withDensity (fun _ b ↦ f b)) : - μ ⊗ₘ κ = (μ ⊗ₘ η).withDensity (f ∘ Prod.snd) := by - have hf_uncurry : Measurable (Function.uncurry (fun (_ : α) => f)) := - hf.comp measurable_snd - ext s hs - have lhs : (μ ⊗ₘ κ) s = ∫⁻ a, (κ a) (Prod.mk a ⁻¹' s) ∂μ := - Measure.compProd_apply hs - have rhs : ((μ ⊗ₘ η).withDensity (f ∘ Prod.snd)) s = - ∫⁻ a, ∫⁻ b in Prod.mk a ⁻¹' s, f b ∂(η a) ∂μ := by - rw [withDensity_apply _ hs, ← lintegral_indicator hs, - Measure.lintegral_compProd ((hf.comp measurable_snd).indicator hs)] - congr 1; ext a - have : (fun b => s.indicator (f ∘ Prod.snd) (a, b)) = (Prod.mk a ⁻¹' s).indicator f := by - ext b; simp only [Set.indicator, Set.mem_preimage]; rfl - rw [this, lintegral_indicator (hs.preimage measurable_prodMk_left)] - rw [lhs, rhs] - apply lintegral_congr_ae - filter_upwards [h] with a ha - rw [ha, Kernel.withDensity_apply _ hf_uncurry, withDensity_apply _ (hs.preimage (by fun_prop))] +lemma compProd_eq_compProd_withDensity [SFinite μ] {κ η : Kernel α β} [IsSFiniteKernel κ] + [IsSFiniteKernel η] {f : β → ℝ≥0∞} (hf : Measurable f) + (h : κ =ᵐ[μ] η.withDensity (fun _ b ↦ f b)) : μ ⊗ₘ κ = (μ ⊗ₘ η).withDensity (f ∘ Prod.snd) := by + refine Measure.ext_of_lintegral _ fun g hg ↦ ?_ + calc ∫⁻ p, g p ∂(μ ⊗ₘ κ) + = ∫⁻ a, ∫⁻ b, g (a, b) ∂κ a ∂μ := + Measure.lintegral_compProd hg + _ = ∫⁻ a, ∫⁻ b, g (a, b) ∂(η.withDensity (fun _ b ↦ f b)) a ∂μ := by + refine lintegral_congr_ae ?_ + filter_upwards [h] with a ha; rw [ha] + _ = ∫⁻ a, ∫⁻ b, g (a, b) ∂((η a).withDensity f) ∂μ := by + refine lintegral_congr fun a ↦ ?_ + rw [Kernel.withDensity_apply _ (by fun_prop)] + _ = ∫⁻ a, ∫⁻ b, f b * g (a, b) ∂η a ∂μ := by + refine lintegral_congr fun a ↦ ?_ + exact lintegral_withDensity_eq_lintegral_mul _ hf (by fun_prop) + _ = ∫⁻ p, (f ∘ Prod.snd) p * g p ∂(μ ⊗ₘ η) := + (Measure.lintegral_compProd ((hf.comp measurable_snd).mul hg)).symm + _ = ∫⁻ p, g p ∂((μ ⊗ₘ η).withDensity (f ∘ Prod.snd)) := + (lintegral_withDensity_eq_lintegral_mul _ (hf.comp measurable_snd) hg).symm end Measure namespace ProbabilityTheory.Kernel -/-- `(κ.withDensity (fun _ => f)) ∘ₘ μ = (κ ∘ₘ μ).withDensity f`. -/ -lemma comp_withDensity_const - {κ : Kernel α γ} [IsSFiniteKernel κ] - {f : γ → ℝ≥0∞} (hf : Measurable f) : - (κ.withDensity (fun _ c ↦ f c)) ∘ₘ μ = (κ ∘ₘ μ).withDensity f := by - have hf_uncurry : Measurable (Function.uncurry (fun (_ : α) => f)) := - hf.comp measurable_snd - ext s hs - have lhs : ((κ.withDensity (fun _ => f)) ∘ₘ μ) s = ∫⁻ a, ∫⁻ x in s, f x ∂(κ a) ∂μ := by - rw [Measure.bind_apply hs (Kernel.measurable _).aemeasurable] - congr 1; ext a - rw [Kernel.withDensity_apply _ hf_uncurry, withDensity_apply _ hs] - have rhs : ((κ ∘ₘ μ).withDensity f) s = ∫⁻ a, ∫⁻ x in s, f x ∂(κ a) ∂μ := by - rw [withDensity_apply _ hs, ← lintegral_indicator hs f, - Measure.lintegral_bind (Kernel.measurable _).aemeasurable ((hf.indicator hs).aemeasurable)] - congr 1; ext a; rw [lintegral_indicator hs f] - rw [lhs, rhs] - -/-- Composing `Kernel.withDensity` on the left kernel of `Kernel.compProd`: -`(κ.withDensity f) ⊗ₖ η = (κ ⊗ₖ η).withDensity (fun a (b, _) => f a b)`. -/ -lemma withDensity_compProd_left - {κ : Kernel α β} {η : Kernel (α × β) γ} {f : α → β → ℝ≥0∞} +lemma comp_withDensity_const {κ : Kernel α γ} [IsSFiniteKernel κ] {f : γ → ℝ≥0∞} + (hf : Measurable f) : (κ.withDensity (fun _ c ↦ f c)) ∘ₘ μ = (κ ∘ₘ μ).withDensity f := by + refine Measure.ext_of_lintegral _ fun g hg ↦ ?_ + calc ∫⁻ x, g x ∂((κ.withDensity (fun _ c ↦ f c)) ∘ₘ μ) + = ∫⁻ a, ∫⁻ x, g x ∂(κ.withDensity (fun _ c ↦ f c)) a ∂μ := + Measure.lintegral_bind (Kernel.measurable _).aemeasurable hg.aemeasurable + _ = ∫⁻ a, ∫⁻ x, g x ∂((κ a).withDensity f) ∂μ := by + refine lintegral_congr fun a ↦ ?_ + rw [Kernel.withDensity_apply _ (by fun_prop)] + _ = ∫⁻ a, ∫⁻ x, f x * g x ∂κ a ∂μ := by + refine lintegral_congr fun a ↦ ?_ + exact lintegral_withDensity_eq_lintegral_mul _ hf hg + _ = ∫⁻ x, f x * g x ∂(κ ∘ₘ μ) := + (Measure.lintegral_bind (Kernel.measurable _).aemeasurable (hf.mul hg).aemeasurable).symm + _ = ∫⁻ x, g x ∂((κ ∘ₘ μ).withDensity f) := + (lintegral_withDensity_eq_lintegral_mul _ hf hg).symm + +lemma withDensity_compProd_left {κ : Kernel α β} {η : Kernel (α × β) γ} {f : α → β → ℝ≥0∞} [IsSFiniteKernel κ] [IsSFiniteKernel η] [IsSFiniteKernel (κ.withDensity f)] (hf : Measurable (Function.uncurry f)) : - (κ.withDensity f) ⊗ₖ η = - (κ ⊗ₖ η).withDensity (fun a (b, _) ↦ f a b) := by + (κ.withDensity f) ⊗ₖ η = (κ ⊗ₖ η).withDensity (fun a (b, _) ↦ f a b) := by have hg : Measurable (Function.uncurry (fun a (bc : β × γ) => f a bc.1)) := hf.comp (measurable_fst.prodMk (measurable_fst.comp measurable_snd)) ext x : 1 @@ -153,11 +121,8 @@ lemma withDensity_compProd_left Kernel.withDensity_apply _ hg] exact Measure.withDensity_compProd_left hf.of_uncurry_left -/-- If `κ a ≪ η a` for all `a`, then `η.withDensity (κ.rnDeriv η) = κ`. -/ -lemma withDensity_rnDeriv_eq' {κ η : Kernel α β} - [MeasurableSpace.CountableOrCountablyGenerated α β] - [IsFiniteKernel κ] [IsFiniteKernel η] - (h : ∀ a, κ a ≪ η a) : +lemma withDensity_rnDeriv_eq' {κ η : Kernel α β} [MeasurableSpace.CountableOrCountablyGenerated α β] + [IsFiniteKernel κ] [IsFiniteKernel η] (h : ∀ a, κ a ≪ η a) : η.withDensity (κ.rnDeriv η) = κ := by ext a : 1 exact Kernel.withDensity_rnDeriv_eq (h a) diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index c97941e3..36dfb8f0 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -190,7 +190,7 @@ lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : rw [← Kernel.map_apply _ (IT.measurable_action 0), ← hc, show IT.action 0 ∘ trajectory A R' = A 0 from rfl, hcd, Kernel.const_apply]⟩ -lemma hasCondDistrib_IT_feedback_zero [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : +lemma hasCondDistrib_IT_feedback_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : ∀ᵐ e ∂Q, HasCondDistrib (IT.feedback 0) (IT.action 0) (κ.sectR e) (condDistrib (trajectory A R') E P e) := by rw [← h.hasLaw_env.map_eq] From bb4ba54b5839add7310e2bc6a9c6f94ebcb7f2ae Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Mon, 18 May 2026 17:12:18 +0100 Subject: [PATCH 126/155] Refactor WithDensity.lean (in progress) --- LeanMachineLearning/Probability/WithDensity.lean | 12 ++++++------ .../SequentialLearning/AlgorithmDensity.lean | 12 ++++++------ 2 files changed, 12 insertions(+), 12 deletions(-) diff --git a/LeanMachineLearning/Probability/WithDensity.lean b/LeanMachineLearning/Probability/WithDensity.lean index a437de4b..1845feee 100644 --- a/LeanMachineLearning/Probability/WithDensity.lean +++ b/LeanMachineLearning/Probability/WithDensity.lean @@ -19,7 +19,7 @@ variable {μ : Measure α} namespace Measure -lemma withDensity_compProd_left [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] {f : α → ℝ≥0∞} +lemma compProd_withDensity_left [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] {f : α → ℝ≥0∞} (hf : Measurable f) : (μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst) := by refine Measure.ext_of_lintegral _ fun g hg ↦ ?_ calc ∫⁻ p, g p ∂((μ.withDensity f) ⊗ₘ κ) @@ -41,7 +41,7 @@ lemma map_withDensity_comp {g : α → γ} {f : γ → ℝ≥0∞} (hg : Measura simp only [Measure.map_apply hg hs, withDensity_apply _ (hg hs), withDensity_apply _ hs, setLIntegral_map hs hf hg, Function.comp] -lemma withDensity_map_equiv {e : α ≃ᵐ β} {f : α → ℝ≥0∞} (hf : Measurable f) : +lemma map_withDensity_equiv {e : α ≃ᵐ β} {f : α → ℝ≥0∞} (hf : Measurable f) : (μ.withDensity f).map e = (μ.map e).withDensity (f ∘ e.symm) := calc (μ.withDensity f).map e = (μ.withDensity ((f ∘ e.symm) ∘ e)).map e := by @@ -59,11 +59,11 @@ lemma map_swap_withDensity_fst {μ : Measure (α × β)} {f : β → ℝ≥0∞} _ = (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := map_withDensity_comp measurable_swap (hf.comp measurable_fst) -lemma withDensity_compProd_withDensity [SFinite μ] {κ : Kernel α γ} [IsSFiniteKernel κ] +lemma compProd_withDensity_withDensity [SFinite μ] {κ : Kernel α γ} [IsSFiniteKernel κ] {f : α → ℝ≥0∞} {g : α → γ → ℝ≥0∞} (hf : Measurable f) (hg : Measurable (Function.uncurry g)) [IsSFiniteKernel (κ.withDensity g)] : (μ.withDensity f) ⊗ₘ (κ.withDensity g) = (μ ⊗ₘ κ).withDensity (fun (a, c) => f a * g a c) := by - rw [Measure.compProd_withDensity hg, withDensity_compProd_left hf] + rw [Measure.compProd_withDensity hg, compProd_withDensity_left hf] exact (withDensity_mul _ (hf.comp measurable_fst) hg).symm lemma compProd_eq_compProd_withDensity [SFinite μ] {κ η : Kernel α β} [IsSFiniteKernel κ] @@ -108,7 +108,7 @@ lemma comp_withDensity_const {κ : Kernel α γ} [IsSFiniteKernel κ] {f : γ _ = ∫⁻ x, g x ∂((κ ∘ₘ μ).withDensity f) := (lintegral_withDensity_eq_lintegral_mul _ hf hg).symm -lemma withDensity_compProd_left {κ : Kernel α β} {η : Kernel (α × β) γ} {f : α → β → ℝ≥0∞} +lemma compProd_withDensity_left {κ : Kernel α β} {η : Kernel (α × β) γ} {f : α → β → ℝ≥0∞} [IsSFiniteKernel κ] [IsSFiniteKernel η] [IsSFiniteKernel (κ.withDensity f)] (hf : Measurable (Function.uncurry f)) : (κ.withDensity f) ⊗ₖ η = (κ ⊗ₖ η).withDensity (fun a (b, _) ↦ f a b) := by @@ -119,7 +119,7 @@ lemma withDensity_compProd_left {κ : Kernel α β} {η : Kernel (α × β) γ} rw [← Kernel.withDensity_apply _ hf]; infer_instance simp only [Kernel.compProd_apply_eq_compProd_sectR, Kernel.withDensity_apply _ hf, Kernel.withDensity_apply _ hg] - exact Measure.withDensity_compProd_left hf.of_uncurry_left + exact Measure.compProd_withDensity_left hf.of_uncurry_left lemma withDensity_rnDeriv_eq' {κ η : Kernel α β} [MeasurableSpace.CountableOrCountablyGenerated α β] [IsFiniteKernel κ] [IsFiniteKernel η] (h : ∀ a, κ a ≪ η a) : diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean index b1cf827a..195d0311 100644 --- a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean @@ -111,21 +111,21 @@ lemma hasLaw_hist_withDensity (h : IsAlgEnvSeq A R' alg env P) (h₀ : IsAlgEnvS rw [h.hasLaw_hist_zero.map_eq, h₀.hasLaw_hist_zero.map_eq, h.hasLaw_step_zero.map_eq, h₀.hasLaw_step_zero.map_eq] rw [← Measure.withDensity_rnDeriv_eq _ _ hc.p0, - Measure.withDensity_compProd_left (by fun_prop)] - exact Measure.withDensity_map_equiv (by fun_prop) + Measure.compProd_withDensity_left (by fun_prop)] + exact Measure.map_withDensity_equiv (by fun_prop) | succ n ih => let ρ h' (ar : α × R) := Kernel.rnDeriv (alg.policy n) (alg₀.policy n) h' ar.1 have hs : stepKernel alg env n = (stepKernel alg₀ env n).withDensity ρ := by rw [stepKernel, ← Kernel.withDensity_rnDeriv_eq' (hc.policy n)] - exact Kernel.withDensity_compProd_left (Kernel.measurable_rnDeriv _ _) + exact Kernel.compProd_withDensity_left (Kernel.measurable_rnDeriv _ _) have : IsMarkovKernel ((stepKernel alg₀ env n).withDensity ρ) := by rw [← hs] infer_instance rw [(h.hasLaw_hist_succ n).map_eq, (h₀.hasLaw_hist_succ n).map_eq, Measure.compProd_congr (h.hasCondDistrib_step n).condDistrib_eq, Measure.compProd_congr (h₀.hasCondDistrib_step n).condDistrib_eq, ih, hs, - Measure.withDensity_compProd_withDensity (by fun_prop) (by fun_prop)] - exact Measure.withDensity_map_equiv (by fun_prop) + Measure.compProd_withDensity_withDensity (by fun_prop) (by fun_prop)] + exact Measure.map_withDensity_equiv (by fun_prop) end IsAlgEnvSeq @@ -198,7 +198,7 @@ lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) Measure.map_swap_withDensity_fst (by fun_prop), ← h₀.hasLaw_env.map_eq, map_swap_compProd_map_condDistrib (by fun_prop), ← compProd_map_condDistrib (by fun_prop), - ← Measure.withDensity_compProd_left (by fun_prop), + ← Measure.compProd_withDensity_left (by fun_prop), ← (hasLaw_hist_withDensity h h₀ hc n).map_eq] end IsBayesAlgEnvSeq From a8cfc8707a40f5e4e6c52cc878cfbc69d11f7a27 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 19 May 2026 11:35:17 +0100 Subject: [PATCH 127/155] Refactor WithDensity.lean (in progress) --- .../Probability/WithDensity.lean | 40 ++++++++++--------- .../SequentialLearning/AlgorithmDensity.lean | 6 +-- 2 files changed, 25 insertions(+), 21 deletions(-) diff --git a/LeanMachineLearning/Probability/WithDensity.lean b/LeanMachineLearning/Probability/WithDensity.lean index 1845feee..4121e440 100644 --- a/LeanMachineLearning/Probability/WithDensity.lean +++ b/LeanMachineLearning/Probability/WithDensity.lean @@ -17,23 +17,7 @@ open scoped ENNReal variable {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} variable {μ : Measure α} -namespace Measure - -lemma compProd_withDensity_left [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] {f : α → ℝ≥0∞} - (hf : Measurable f) : (μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst) := by - refine Measure.ext_of_lintegral _ fun g hg ↦ ?_ - calc ∫⁻ p, g p ∂((μ.withDensity f) ⊗ₘ κ) - = ∫⁻ a, ∫⁻ b, g (a, b) ∂κ a ∂(μ.withDensity f) := - Measure.lintegral_compProd hg - _ = ∫⁻ a, f a * ∫⁻ b, g (a, b) ∂κ a ∂μ := - lintegral_withDensity_eq_lintegral_mul _ hf hg.lintegral_kernel_prod_right' - _ = ∫⁻ a, ∫⁻ b, f a * g (a, b) ∂κ a ∂μ := by - refine lintegral_congr fun a ↦ ?_ - rw [← lintegral_const_mul _ (by fun_prop)] - _ = ∫⁻ p, (f ∘ Prod.fst) p * g p ∂(μ ⊗ₘ κ) := - (Measure.lintegral_compProd ((hf.comp measurable_fst).mul hg)).symm - _ = ∫⁻ p, g p ∂((μ ⊗ₘ κ).withDensity (f ∘ Prod.fst)) := - (lintegral_withDensity_eq_lintegral_mul _ (hf.comp measurable_fst) hg).symm +namespace MeasureTheory lemma map_withDensity_comp {g : α → γ} {f : γ → ℝ≥0∞} (hg : Measurable g) (hf : Measurable f) : (μ.withDensity (f ∘ g)).map g = (μ.map g).withDensity f := by @@ -59,6 +43,26 @@ lemma map_swap_withDensity_fst {μ : Measure (α × β)} {f : β → ℝ≥0∞} _ = (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := map_withDensity_comp measurable_swap (hf.comp measurable_fst) +end MeasureTheory + +namespace MeasureTheory.Measure + +lemma compProd_withDensity_left [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] {f : α → ℝ≥0∞} + (hf : Measurable f) : (μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst) := by + refine Measure.ext_of_lintegral _ fun g hg ↦ ?_ + calc ∫⁻ p, g p ∂((μ.withDensity f) ⊗ₘ κ) + = ∫⁻ a, ∫⁻ b, g (a, b) ∂κ a ∂(μ.withDensity f) := + Measure.lintegral_compProd hg + _ = ∫⁻ a, f a * ∫⁻ b, g (a, b) ∂κ a ∂μ := + lintegral_withDensity_eq_lintegral_mul _ hf hg.lintegral_kernel_prod_right' + _ = ∫⁻ a, ∫⁻ b, f a * g (a, b) ∂κ a ∂μ := by + refine lintegral_congr fun a ↦ ?_ + rw [← lintegral_const_mul _ (by fun_prop)] + _ = ∫⁻ p, (f ∘ Prod.fst) p * g p ∂(μ ⊗ₘ κ) := + (Measure.lintegral_compProd ((hf.comp measurable_fst).mul hg)).symm + _ = ∫⁻ p, g p ∂((μ ⊗ₘ κ).withDensity (f ∘ Prod.fst)) := + (lintegral_withDensity_eq_lintegral_mul _ (hf.comp measurable_fst) hg).symm + lemma compProd_withDensity_withDensity [SFinite μ] {κ : Kernel α γ} [IsSFiniteKernel κ] {f : α → ℝ≥0∞} {g : α → γ → ℝ≥0∞} (hf : Measurable f) (hg : Measurable (Function.uncurry g)) [IsSFiniteKernel (κ.withDensity g)] : @@ -87,7 +91,7 @@ lemma compProd_eq_compProd_withDensity [SFinite μ] {κ η : Kernel α β} [IsSF _ = ∫⁻ p, g p ∂((μ ⊗ₘ η).withDensity (f ∘ Prod.snd)) := (lintegral_withDensity_eq_lintegral_mul _ (hf.comp measurable_snd) hg).symm -end Measure +end MeasureTheory.Measure namespace ProbabilityTheory.Kernel diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean index 195d0311..28445859 100644 --- a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean @@ -112,7 +112,7 @@ lemma hasLaw_hist_withDensity (h : IsAlgEnvSeq A R' alg env P) (h₀ : IsAlgEnvS h₀.hasLaw_step_zero.map_eq] rw [← Measure.withDensity_rnDeriv_eq _ _ hc.p0, Measure.compProd_withDensity_left (by fun_prop)] - exact Measure.map_withDensity_equiv (by fun_prop) + exact map_withDensity_equiv (by fun_prop) | succ n ih => let ρ h' (ar : α × R) := Kernel.rnDeriv (alg.policy n) (alg₀.policy n) h' ar.1 have hs : stepKernel alg env n = (stepKernel alg₀ env n).withDensity ρ := by @@ -125,7 +125,7 @@ lemma hasLaw_hist_withDensity (h : IsAlgEnvSeq A R' alg env P) (h₀ : IsAlgEnvS Measure.compProd_congr (h.hasCondDistrib_step n).condDistrib_eq, Measure.compProd_congr (h₀.hasCondDistrib_step n).condDistrib_eq, ih, hs, Measure.compProd_withDensity_withDensity (by fun_prop) (by fun_prop)] - exact Measure.map_withDensity_equiv (by fun_prop) + exact map_withDensity_equiv (by fun_prop) end IsAlgEnvSeq @@ -195,7 +195,7 @@ lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) ← map_swap_compProd_map_condDistrib (by fun_prop), h.hasLaw_env.map_eq, Measure.compProd_eq_compProd_withDensity (by fun_prop) (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), - Measure.map_swap_withDensity_fst (by fun_prop), + map_swap_withDensity_fst (by fun_prop), ← h₀.hasLaw_env.map_eq, map_swap_compProd_map_condDistrib (by fun_prop), ← compProd_map_condDistrib (by fun_prop), ← Measure.compProd_withDensity_left (by fun_prop), From c72c0b79ba8465ea93f72af37f1bdfd805241f98 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 19 May 2026 19:22:55 +0100 Subject: [PATCH 128/155] Refactor WithDensity.lean (in progress) --- .../Probability/WithDensity.lean | 32 ++++++++----------- .../SequentialLearning/AlgorithmDensity.lean | 4 +-- 2 files changed, 16 insertions(+), 20 deletions(-) diff --git a/LeanMachineLearning/Probability/WithDensity.lean b/LeanMachineLearning/Probability/WithDensity.lean index 4121e440..d35c220a 100644 --- a/LeanMachineLearning/Probability/WithDensity.lean +++ b/LeanMachineLearning/Probability/WithDensity.lean @@ -22,26 +22,22 @@ namespace MeasureTheory lemma map_withDensity_comp {g : α → γ} {f : γ → ℝ≥0∞} (hg : Measurable g) (hf : Measurable f) : (μ.withDensity (f ∘ g)).map g = (μ.map g).withDensity f := by ext s hs - simp only [Measure.map_apply hg hs, withDensity_apply _ (hg hs), withDensity_apply _ hs, - setLIntegral_map hs hf hg, Function.comp] - -lemma map_withDensity_equiv {e : α ≃ᵐ β} {f : α → ℝ≥0∞} (hf : Measurable f) : - (μ.withDensity f).map e = (μ.map e).withDensity (f ∘ e.symm) := - calc (μ.withDensity f).map e - = (μ.withDensity ((f ∘ e.symm) ∘ e)).map e := by - congr - funext x - simp - _ = (μ.map e).withDensity (f ∘ e.symm) := - map_withDensity_comp e.measurable (hf.comp e.symm.measurable) + rw [Measure.map_apply hg hs, withDensity_apply _ (hg hs), withDensity_apply _ hs, + setLIntegral_map hs hf hg] + rfl + +lemma map_equiv_withDensity {e : α ≃ᵐ β} {f : α → ℝ≥0∞} (hf : Measurable f) : + (μ.withDensity f).map e = (μ.map e).withDensity (f ∘ e.symm) := by + rw [← map_withDensity_comp e.measurable (hf.comp e.symm.measurable)] + congr + ext a + simp lemma map_swap_withDensity_fst {μ : Measure (α × β)} {f : β → ℝ≥0∞} (hf : Measurable f) : - (μ.withDensity (f ∘ Prod.snd)).map Prod.swap = (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := - calc (μ.withDensity (f ∘ Prod.snd)).map Prod.swap - _ = (μ.withDensity ((f ∘ Prod.fst) ∘ Prod.swap)).map Prod.swap := - rfl - _ = (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := - map_withDensity_comp measurable_swap (hf.comp measurable_fst) + (μ.withDensity (f ∘ Prod.snd)).map Prod.swap = + (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := by + rw [← map_withDensity_comp measurable_swap (hf.comp measurable_fst)] + congr end MeasureTheory diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean index 28445859..8ca4bbaa 100644 --- a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean @@ -112,7 +112,7 @@ lemma hasLaw_hist_withDensity (h : IsAlgEnvSeq A R' alg env P) (h₀ : IsAlgEnvS h₀.hasLaw_step_zero.map_eq] rw [← Measure.withDensity_rnDeriv_eq _ _ hc.p0, Measure.compProd_withDensity_left (by fun_prop)] - exact map_withDensity_equiv (by fun_prop) + exact map_equiv_withDensity (by fun_prop) | succ n ih => let ρ h' (ar : α × R) := Kernel.rnDeriv (alg.policy n) (alg₀.policy n) h' ar.1 have hs : stepKernel alg env n = (stepKernel alg₀ env n).withDensity ρ := by @@ -125,7 +125,7 @@ lemma hasLaw_hist_withDensity (h : IsAlgEnvSeq A R' alg env P) (h₀ : IsAlgEnvS Measure.compProd_congr (h.hasCondDistrib_step n).condDistrib_eq, Measure.compProd_congr (h₀.hasCondDistrib_step n).condDistrib_eq, ih, hs, Measure.compProd_withDensity_withDensity (by fun_prop) (by fun_prop)] - exact map_withDensity_equiv (by fun_prop) + exact map_equiv_withDensity (by fun_prop) end IsAlgEnvSeq From 7ef662114aa1790f1df61bdf1b8118eb8655099e Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 20 May 2026 09:32:07 +0100 Subject: [PATCH 129/155] Refactor WithDensity.lean (in progress) --- LeanMachineLearning/Probability/WithDensity.lean | 11 +++++------ .../SequentialLearning/AlgorithmDensity.lean | 2 +- 2 files changed, 6 insertions(+), 7 deletions(-) diff --git a/LeanMachineLearning/Probability/WithDensity.lean b/LeanMachineLearning/Probability/WithDensity.lean index d35c220a..1df8a1f1 100644 --- a/LeanMachineLearning/Probability/WithDensity.lean +++ b/LeanMachineLearning/Probability/WithDensity.lean @@ -28,16 +28,15 @@ lemma map_withDensity_comp {g : α → γ} {f : γ → ℝ≥0∞} (hg : Measura lemma map_equiv_withDensity {e : α ≃ᵐ β} {f : α → ℝ≥0∞} (hf : Measurable f) : (μ.withDensity f).map e = (μ.map e).withDensity (f ∘ e.symm) := by - rw [← map_withDensity_comp e.measurable (hf.comp e.symm.measurable)] - congr - ext a - simp + simp_rw [← map_withDensity_comp e.measurable (hf.comp e.symm.measurable), + Function.comp_assoc, MeasurableEquiv.symm_comp_self] + rfl -lemma map_swap_withDensity_fst {μ : Measure (α × β)} {f : β → ℝ≥0∞} (hf : Measurable f) : +lemma map_swap_withDensity_comp_snd {μ : Measure (α × β)} {f : β → ℝ≥0∞} (hf : Measurable f) : (μ.withDensity (f ∘ Prod.snd)).map Prod.swap = (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := by rw [← map_withDensity_comp measurable_swap (hf.comp measurable_fst)] - congr + rfl end MeasureTheory diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean index 8ca4bbaa..007f34fc 100644 --- a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean @@ -195,7 +195,7 @@ lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) ← map_swap_compProd_map_condDistrib (by fun_prop), h.hasLaw_env.map_eq, Measure.compProd_eq_compProd_withDensity (by fun_prop) (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), - map_swap_withDensity_fst (by fun_prop), + map_swap_withDensity_comp_snd (by fun_prop), ← h₀.hasLaw_env.map_eq, map_swap_compProd_map_condDistrib (by fun_prop), ← compProd_map_condDistrib (by fun_prop), ← Measure.compProd_withDensity_left (by fun_prop), From 6b59c70aef21106694b7161a53f4c18384f55da1 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 20 May 2026 10:32:20 +0100 Subject: [PATCH 130/155] Refactor WithDensity.lean (in progress) --- .../Probability/WithDensity.lean | 22 +++++++++---------- 1 file changed, 11 insertions(+), 11 deletions(-) diff --git a/LeanMachineLearning/Probability/WithDensity.lean b/LeanMachineLearning/Probability/WithDensity.lean index 1df8a1f1..e566a41e 100644 --- a/LeanMachineLearning/Probability/WithDensity.lean +++ b/LeanMachineLearning/Probability/WithDensity.lean @@ -33,9 +33,9 @@ lemma map_equiv_withDensity {e : α ≃ᵐ β} {f : α → ℝ≥0∞} (hf : Mea rfl lemma map_swap_withDensity_comp_snd {μ : Measure (α × β)} {f : β → ℝ≥0∞} (hf : Measurable f) : - (μ.withDensity (f ∘ Prod.snd)).map Prod.swap = - (μ.map Prod.swap).withDensity (f ∘ Prod.fst) := by - rw [← map_withDensity_comp measurable_swap (hf.comp measurable_fst)] + (μ.withDensity (fun ab ↦ f ab.2)).map Prod.swap = + (μ.map Prod.swap).withDensity (fun ba ↦ f ba.1) := by + rw [← map_withDensity_comp measurable_swap (by fun_prop)] rfl end MeasureTheory @@ -43,19 +43,18 @@ end MeasureTheory namespace MeasureTheory.Measure lemma compProd_withDensity_left [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] {f : α → ℝ≥0∞} - (hf : Measurable f) : (μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (f ∘ Prod.fst) := by + (hf : Measurable f) : (μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (fun ab ↦ f ab.1) := by refine Measure.ext_of_lintegral _ fun g hg ↦ ?_ - calc ∫⁻ p, g p ∂((μ.withDensity f) ⊗ₘ κ) + calc ∫⁻ ab, g ab ∂((μ.withDensity f) ⊗ₘ κ) = ∫⁻ a, ∫⁻ b, g (a, b) ∂κ a ∂(μ.withDensity f) := Measure.lintegral_compProd hg _ = ∫⁻ a, f a * ∫⁻ b, g (a, b) ∂κ a ∂μ := lintegral_withDensity_eq_lintegral_mul _ hf hg.lintegral_kernel_prod_right' - _ = ∫⁻ a, ∫⁻ b, f a * g (a, b) ∂κ a ∂μ := by - refine lintegral_congr fun a ↦ ?_ - rw [← lintegral_const_mul _ (by fun_prop)] - _ = ∫⁻ p, (f ∘ Prod.fst) p * g p ∂(μ ⊗ₘ κ) := + _ = ∫⁻ a, ∫⁻ b, f a * g (a, b) ∂κ a ∂μ := + lintegral_congr fun a ↦ (lintegral_const_mul _ (by fun_prop)).symm + _ = ∫⁻ ab, (fun ab ↦ f ab.1) ab * g ab ∂(μ ⊗ₘ κ) := (Measure.lintegral_compProd ((hf.comp measurable_fst).mul hg)).symm - _ = ∫⁻ p, g p ∂((μ ⊗ₘ κ).withDensity (f ∘ Prod.fst)) := + _ = ∫⁻ ab, g ab ∂((μ ⊗ₘ κ).withDensity (fun ab ↦ f ab.1)) := (lintegral_withDensity_eq_lintegral_mul _ (hf.comp measurable_fst) hg).symm lemma compProd_withDensity_withDensity [SFinite μ] {κ : Kernel α γ} [IsSFiniteKernel κ] @@ -67,7 +66,8 @@ lemma compProd_withDensity_withDensity [SFinite μ] {κ : Kernel α γ} [IsSFini lemma compProd_eq_compProd_withDensity [SFinite μ] {κ η : Kernel α β} [IsSFiniteKernel κ] [IsSFiniteKernel η] {f : β → ℝ≥0∞} (hf : Measurable f) - (h : κ =ᵐ[μ] η.withDensity (fun _ b ↦ f b)) : μ ⊗ₘ κ = (μ ⊗ₘ η).withDensity (f ∘ Prod.snd) := by + (h : κ =ᵐ[μ] η.withDensity (fun _ b ↦ f b)) : + μ ⊗ₘ κ = (μ ⊗ₘ η).withDensity (fun ab ↦ f ab.2) := by refine Measure.ext_of_lintegral _ fun g hg ↦ ?_ calc ∫⁻ p, g p ∂(μ ⊗ₘ κ) = ∫⁻ a, ∫⁻ b, g (a, b) ∂κ a ∂μ := From d540261e5eaded7c837d3b41f957c3b5c32766e2 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 20 May 2026 11:15:59 +0100 Subject: [PATCH 131/155] Refactor WithDensity.lean (in progress) --- .../Probability/WithDensity.lean | 21 ++++++++++--------- 1 file changed, 11 insertions(+), 10 deletions(-) diff --git a/LeanMachineLearning/Probability/WithDensity.lean b/LeanMachineLearning/Probability/WithDensity.lean index e566a41e..b817771c 100644 --- a/LeanMachineLearning/Probability/WithDensity.lean +++ b/LeanMachineLearning/Probability/WithDensity.lean @@ -44,34 +44,35 @@ namespace MeasureTheory.Measure lemma compProd_withDensity_left [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] {f : α → ℝ≥0∞} (hf : Measurable f) : (μ.withDensity f) ⊗ₘ κ = (μ ⊗ₘ κ).withDensity (fun ab ↦ f ab.1) := by - refine Measure.ext_of_lintegral _ fun g hg ↦ ?_ + refine ext_of_lintegral _ fun g hg ↦ ?_ calc ∫⁻ ab, g ab ∂((μ.withDensity f) ⊗ₘ κ) = ∫⁻ a, ∫⁻ b, g (a, b) ∂κ a ∂(μ.withDensity f) := - Measure.lintegral_compProd hg + lintegral_compProd hg _ = ∫⁻ a, f a * ∫⁻ b, g (a, b) ∂κ a ∂μ := lintegral_withDensity_eq_lintegral_mul _ hf hg.lintegral_kernel_prod_right' _ = ∫⁻ a, ∫⁻ b, f a * g (a, b) ∂κ a ∂μ := lintegral_congr fun a ↦ (lintegral_const_mul _ (by fun_prop)).symm _ = ∫⁻ ab, (fun ab ↦ f ab.1) ab * g ab ∂(μ ⊗ₘ κ) := - (Measure.lintegral_compProd ((hf.comp measurable_fst).mul hg)).symm + (lintegral_compProd ((hf.comp measurable_fst).mul hg)).symm _ = ∫⁻ ab, g ab ∂((μ ⊗ₘ κ).withDensity (fun ab ↦ f ab.1)) := (lintegral_withDensity_eq_lintegral_mul _ (hf.comp measurable_fst) hg).symm -lemma compProd_withDensity_withDensity [SFinite μ] {κ : Kernel α γ} [IsSFiniteKernel κ] - {f : α → ℝ≥0∞} {g : α → γ → ℝ≥0∞} (hf : Measurable f) (hg : Measurable (Function.uncurry g)) +lemma compProd_withDensity_withDensity [SFinite μ] {κ : Kernel α β} [IsSFiniteKernel κ] + {f : α → ℝ≥0∞} {g : α → β → ℝ≥0∞} (hf : Measurable f) (hg : Measurable (Function.uncurry g)) [IsSFiniteKernel (κ.withDensity g)] : - (μ.withDensity f) ⊗ₘ (κ.withDensity g) = (μ ⊗ₘ κ).withDensity (fun (a, c) => f a * g a c) := by - rw [Measure.compProd_withDensity hg, compProd_withDensity_left hf] + (μ.withDensity f) ⊗ₘ (κ.withDensity g) = + (μ ⊗ₘ κ).withDensity (fun ac ↦ f ac.1 * g ac.1 ac.2) := by + rw [compProd_withDensity hg, compProd_withDensity_left hf] exact (withDensity_mul _ (hf.comp measurable_fst) hg).symm lemma compProd_eq_compProd_withDensity [SFinite μ] {κ η : Kernel α β} [IsSFiniteKernel κ] [IsSFiniteKernel η] {f : β → ℝ≥0∞} (hf : Measurable f) (h : κ =ᵐ[μ] η.withDensity (fun _ b ↦ f b)) : μ ⊗ₘ κ = (μ ⊗ₘ η).withDensity (fun ab ↦ f ab.2) := by - refine Measure.ext_of_lintegral _ fun g hg ↦ ?_ + refine ext_of_lintegral _ fun g hg ↦ ?_ calc ∫⁻ p, g p ∂(μ ⊗ₘ κ) = ∫⁻ a, ∫⁻ b, g (a, b) ∂κ a ∂μ := - Measure.lintegral_compProd hg + lintegral_compProd hg _ = ∫⁻ a, ∫⁻ b, g (a, b) ∂(η.withDensity (fun _ b ↦ f b)) a ∂μ := by refine lintegral_congr_ae ?_ filter_upwards [h] with a ha; rw [ha] @@ -82,7 +83,7 @@ lemma compProd_eq_compProd_withDensity [SFinite μ] {κ η : Kernel α β} [IsSF refine lintegral_congr fun a ↦ ?_ exact lintegral_withDensity_eq_lintegral_mul _ hf (by fun_prop) _ = ∫⁻ p, (f ∘ Prod.snd) p * g p ∂(μ ⊗ₘ η) := - (Measure.lintegral_compProd ((hf.comp measurable_snd).mul hg)).symm + (lintegral_compProd ((hf.comp measurable_snd).mul hg)).symm _ = ∫⁻ p, g p ∂((μ ⊗ₘ η).withDensity (f ∘ Prod.snd)) := (lintegral_withDensity_eq_lintegral_mul _ (hf.comp measurable_snd) hg).symm From efc791497b65026c54ad6b083c18ccce13a9c290 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 22 May 2026 11:49:26 +0100 Subject: [PATCH 132/155] Refactor WithDensity.lean (in progress) --- .../Probability/WithDensity.lean | 50 +++++++++---------- .../SequentialLearning/AlgorithmDensity.lean | 4 +- 2 files changed, 26 insertions(+), 28 deletions(-) diff --git a/LeanMachineLearning/Probability/WithDensity.lean b/LeanMachineLearning/Probability/WithDensity.lean index b817771c..f3ddcce1 100644 --- a/LeanMachineLearning/Probability/WithDensity.lean +++ b/LeanMachineLearning/Probability/WithDensity.lean @@ -65,47 +65,45 @@ lemma compProd_withDensity_withDensity [SFinite μ] {κ : Kernel α β} [IsSFini rw [compProd_withDensity hg, compProd_withDensity_left hf] exact (withDensity_mul _ (hf.comp measurable_fst) hg).symm -lemma compProd_eq_compProd_withDensity [SFinite μ] {κ η : Kernel α β} [IsSFiniteKernel κ] +/-- A proof based on `compProd_congr` requires `IsSFiniteKernel (η.withDensity fun _ b ↦ f b)`. -/ +lemma compProd_eq_compProd_withDensity_comp_snd [SFinite μ] {κ η : Kernel α β} [IsSFiniteKernel κ] [IsSFiniteKernel η] {f : β → ℝ≥0∞} (hf : Measurable f) (h : κ =ᵐ[μ] η.withDensity (fun _ b ↦ f b)) : μ ⊗ₘ κ = (μ ⊗ₘ η).withDensity (fun ab ↦ f ab.2) := by refine ext_of_lintegral _ fun g hg ↦ ?_ - calc ∫⁻ p, g p ∂(μ ⊗ₘ κ) + calc ∫⁻ ab, g ab ∂(μ ⊗ₘ κ) = ∫⁻ a, ∫⁻ b, g (a, b) ∂κ a ∂μ := lintegral_compProd hg - _ = ∫⁻ a, ∫⁻ b, g (a, b) ∂(η.withDensity (fun _ b ↦ f b)) a ∂μ := by - refine lintegral_congr_ae ?_ - filter_upwards [h] with a ha; rw [ha] _ = ∫⁻ a, ∫⁻ b, g (a, b) ∂((η a).withDensity f) ∂μ := by - refine lintegral_congr fun a ↦ ?_ - rw [Kernel.withDensity_apply _ (by fun_prop)] + apply lintegral_congr_ae + filter_upwards [h] with a ha + rw [ha, Kernel.withDensity_apply _ (by fun_prop)] _ = ∫⁻ a, ∫⁻ b, f b * g (a, b) ∂η a ∂μ := by - refine lintegral_congr fun a ↦ ?_ + congr + ext a exact lintegral_withDensity_eq_lintegral_mul _ hf (by fun_prop) - _ = ∫⁻ p, (f ∘ Prod.snd) p * g p ∂(μ ⊗ₘ η) := + _ = ∫⁻ ab, f ab.2 * g ab ∂(μ ⊗ₘ η) := (lintegral_compProd ((hf.comp measurable_snd).mul hg)).symm - _ = ∫⁻ p, g p ∂((μ ⊗ₘ η).withDensity (f ∘ Prod.snd)) := + _ = ∫⁻ ab, g ab ∂((μ ⊗ₘ η).withDensity (fun ab ↦ f ab.2)) := (lintegral_withDensity_eq_lintegral_mul _ (hf.comp measurable_snd) hg).symm end MeasureTheory.Measure namespace ProbabilityTheory.Kernel -lemma comp_withDensity_const {κ : Kernel α γ} [IsSFiniteKernel κ] {f : γ → ℝ≥0∞} - (hf : Measurable f) : (κ.withDensity (fun _ c ↦ f c)) ∘ₘ μ = (κ ∘ₘ μ).withDensity f := by +lemma bind_withDensity_eq_withDensity_bind {κ : Kernel α β} [IsSFiniteKernel κ] {f : β → ℝ≥0∞} + (hf : Measurable f) : (κ.withDensity (fun _ b ↦ f b)) ∘ₘ μ = (κ ∘ₘ μ).withDensity f := by refine Measure.ext_of_lintegral _ fun g hg ↦ ?_ - calc ∫⁻ x, g x ∂((κ.withDensity (fun _ c ↦ f c)) ∘ₘ μ) - = ∫⁻ a, ∫⁻ x, g x ∂(κ.withDensity (fun _ c ↦ f c)) a ∂μ := - Measure.lintegral_bind (Kernel.measurable _).aemeasurable hg.aemeasurable - _ = ∫⁻ a, ∫⁻ x, g x ∂((κ a).withDensity f) ∂μ := by - refine lintegral_congr fun a ↦ ?_ - rw [Kernel.withDensity_apply _ (by fun_prop)] - _ = ∫⁻ a, ∫⁻ x, f x * g x ∂κ a ∂μ := by - refine lintegral_congr fun a ↦ ?_ - exact lintegral_withDensity_eq_lintegral_mul _ hf hg - _ = ∫⁻ x, f x * g x ∂(κ ∘ₘ μ) := - (Measure.lintegral_bind (Kernel.measurable _).aemeasurable (hf.mul hg).aemeasurable).symm - _ = ∫⁻ x, g x ∂((κ ∘ₘ μ).withDensity f) := + calc ∫⁻ b, g b ∂((κ.withDensity (fun _ b ↦ f b)) ∘ₘ μ) + = ∫⁻ a, ∫⁻ b, g b ∂(κ.withDensity (fun _ b ↦ f b)) a ∂μ := + Measure.lintegral_bind (measurable _).aemeasurable hg.aemeasurable + _ = ∫⁻ a, ∫⁻ b, f b * g b ∂κ a ∂μ := by + congr + ext a + exact lintegral_withDensity _ (by fun_prop) _ hg + _ = ∫⁻ b, f b * g b ∂(κ ∘ₘ μ) := + (Measure.lintegral_bind (measurable _).aemeasurable (hf.mul hg).aemeasurable).symm + _ = ∫⁻ b, g b ∂((κ ∘ₘ μ).withDensity f) := (lintegral_withDensity_eq_lintegral_mul _ hf hg).symm lemma compProd_withDensity_left {κ : Kernel α β} {η : Kernel (α × β) γ} {f : α → β → ℝ≥0∞} @@ -117,7 +115,7 @@ lemma compProd_withDensity_left {κ : Kernel α β} {η : Kernel (α × β) γ} ext x : 1 haveI : SFinite ((κ x).withDensity (f x)) := by rw [← Kernel.withDensity_apply _ hf]; infer_instance - simp only [Kernel.compProd_apply_eq_compProd_sectR, Kernel.withDensity_apply _ hf, + simp only [compProd_apply_eq_compProd_sectR, Kernel.withDensity_apply _ hf, Kernel.withDensity_apply _ hg] exact Measure.compProd_withDensity_left hf.of_uncurry_left @@ -125,6 +123,6 @@ lemma withDensity_rnDeriv_eq' {κ η : Kernel α β} [MeasurableSpace.CountableO [IsFiniteKernel κ] [IsFiniteKernel η] (h : ∀ a, κ a ≪ η a) : η.withDensity (κ.rnDeriv η) = κ := by ext a : 1 - exact Kernel.withDensity_rnDeriv_eq (h a) + exact withDensity_rnDeriv_eq (h a) end ProbabilityTheory.Kernel diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean index 007f34fc..8cd9f9f9 100644 --- a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean @@ -171,7 +171,7 @@ lemma hasLaw_hist_withDensity (h : IsBayesAlgEnvSeq Q κ alg E A R' P) have hE₀ := h₀.measurable_E rw [← condDistrib_comp_map hE.aemeasurable (by fun_prop), h.hasLaw_env.map_eq, Measure.bind_congr_right (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), - Kernel.comp_withDensity_const (by fun_prop), + Kernel.bind_withDensity_eq_withDensity_bind (by fun_prop), ← h₀.hasLaw_env.map_eq, condDistrib_comp_map hE₀.aemeasurable (by fun_prop)] variable [StandardBorelSpace 𝓔] [Nonempty 𝓔] @@ -193,7 +193,7 @@ lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) have hE₀ := h₀.measurable_E rw [condDistrib_ae_eq_iff_measure_eq_compProd _ h.measurable_E.aemeasurable, ← map_swap_compProd_map_condDistrib (by fun_prop), h.hasLaw_env.map_eq, - Measure.compProd_eq_compProd_withDensity (by fun_prop) + Measure.compProd_eq_compProd_withDensity_comp_snd (by fun_prop) (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), map_swap_withDensity_comp_snd (by fun_prop), ← h₀.hasLaw_env.map_eq, map_swap_compProd_map_condDistrib (by fun_prop), From ebeb5f07a1f34a89830c7cf03a671c31e757eef1 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 22 May 2026 13:40:30 +0100 Subject: [PATCH 133/155] Refactor WithDensity.lean --- .../Probability/WithDensity.lean | 23 +++++++++---------- 1 file changed, 11 insertions(+), 12 deletions(-) diff --git a/LeanMachineLearning/Probability/WithDensity.lean b/LeanMachineLearning/Probability/WithDensity.lean index f3ddcce1..8869a422 100644 --- a/LeanMachineLearning/Probability/WithDensity.lean +++ b/LeanMachineLearning/Probability/WithDensity.lean @@ -109,20 +109,19 @@ lemma bind_withDensity_eq_withDensity_bind {κ : Kernel α β} [IsSFiniteKernel lemma compProd_withDensity_left {κ : Kernel α β} {η : Kernel (α × β) γ} {f : α → β → ℝ≥0∞} [IsSFiniteKernel κ] [IsSFiniteKernel η] [IsSFiniteKernel (κ.withDensity f)] (hf : Measurable (Function.uncurry f)) : - (κ.withDensity f) ⊗ₖ η = (κ ⊗ₖ η).withDensity (fun a (b, _) ↦ f a b) := by - have hg : Measurable (Function.uncurry (fun a (bc : β × γ) => f a bc.1)) := - hf.comp (measurable_fst.prodMk (measurable_fst.comp measurable_snd)) - ext x : 1 - haveI : SFinite ((κ x).withDensity (f x)) := by - rw [← Kernel.withDensity_apply _ hf]; infer_instance - simp only [compProd_apply_eq_compProd_sectR, Kernel.withDensity_apply _ hf, - Kernel.withDensity_apply _ hg] - exact Measure.compProd_withDensity_left hf.of_uncurry_left + (κ.withDensity f) ⊗ₖ η = (κ ⊗ₖ η).withDensity (fun a bc ↦ f a bc.1) := by + ext a : 1 + calc ((κ.withDensity f) ⊗ₖ η) a + = (κ a).withDensity (f a) ⊗ₘ η.sectR a := by + rw [compProd_apply_eq_compProd_sectR, Kernel.withDensity_apply _ hf] + _ = ((κ a) ⊗ₘ (η.sectR a)).withDensity (fun bc ↦ f a bc.1) := + Measure.compProd_withDensity_left (by fun_prop) + _ = ((κ ⊗ₖ η).withDensity (fun a bc ↦ f a bc.1)) a := by + rw [← compProd_apply_eq_compProd_sectR, Kernel.withDensity_apply _ (by fun_prop)] lemma withDensity_rnDeriv_eq' {κ η : Kernel α β} [MeasurableSpace.CountableOrCountablyGenerated α β] [IsFiniteKernel κ] [IsFiniteKernel η] (h : ∀ a, κ a ≪ η a) : - η.withDensity (κ.rnDeriv η) = κ := by - ext a : 1 - exact withDensity_rnDeriv_eq (h a) + η.withDensity (κ.rnDeriv η) = κ := + Kernel.ext fun a ↦ withDensity_rnDeriv_eq (h a) end ProbabilityTheory.Kernel From c24b45cfc416f63485d4b52a5ac118f269d527e0 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 22 May 2026 16:57:33 +0100 Subject: [PATCH 134/155] Add documentation --- .../Online/Bandit/Algorithms/TS.lean | 8 ++ .../SequentialLearning/Algorithm.lean | 11 ++- .../SequentialLearning/AlgorithmDensity.lean | 3 + .../BayesStationaryEnv.lean | 92 +++++++++++-------- 4 files changed, 71 insertions(+), 43 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index 6e90355d..87689b9b 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -25,6 +25,9 @@ section Algorithm variable {K : ℕ} variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] +/-- The Thompson sampling policy samples an action according to its probability of being optimal +under the posterior over "environments" given the history so far. +The posterior under a uniform algorithm is used to avoid a circular definition. -/ noncomputable def TS.policy (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := @@ -35,6 +38,8 @@ instance {hK : 0 < K} {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel ( [IsMarkovKernel κ] {n : ℕ} : IsMarkovKernel (TS.policy hK Q κ n) := Kernel.IsMarkovKernel.map _ (by fun_prop) +/-- The initial action is sampled according to its probability of being optimal under the prior over +"environments". -/ noncomputable def TS.initialPolicy (hK : 0 < K) (Q : Measure 𝓔) (κ : Kernel (𝓔 × Fin K) ℝ) : Measure (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK @@ -44,6 +49,7 @@ instance {hK : 0 < K} {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel ( IsProbabilityMeasure (TS.initialPolicy hK Q κ) := Measure.isProbabilityMeasure_map (by fun_prop) +/-- The Thompson sampling algorithm. -/ noncomputable def tsAlgorithm (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] : Algorithm (Fin K) ℝ where @@ -90,6 +96,7 @@ namespace ClippedUCB variable {K : ℕ} {l u σ2 δ : ℝ} variable {Ω : Type*} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} +/-- Clipped upper confidence bound used in the regret analysis of Thompson sampling. -/ noncomputable def ucb (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ := if pullCount A a n ω = 0 then u @@ -133,6 +140,7 @@ lemma integrable_uncurry_ucb_comp [MeasurableSpace Ω] (hA : ∀ t, Measurable ( unfold ucb grind +/-- Clipped upper confidence bound (history-based version). -/ noncomputable def ucb' (n : ℕ) (h : Iic n → Fin K × ℝ) (l u σ2 δ : ℝ) (a : Fin K) : ℝ := if pullCount' n h a = 0 then u diff --git a/LeanMachineLearning/SequentialLearning/Algorithm.lean b/LeanMachineLearning/SequentialLearning/Algorithm.lean index 8822bf4d..31495e2c 100644 --- a/LeanMachineLearning/SequentialLearning/Algorithm.lean +++ b/LeanMachineLearning/SequentialLearning/Algorithm.lean @@ -55,16 +55,19 @@ structure Algorithm (𝓐 𝓨 : Type*) [MeasurableSpace 𝓐] [MeasurableSpace instance (alg : Algorithm 𝓐 𝓨) (n : ℕ) : IsMarkovKernel (alg.policy n) := alg.h_policy n instance (alg : Algorithm 𝓐 𝓨) : IsProbabilityMeasure alg.p0 := alg.hp0 +/-- For every time and history, the distribution over actions according to `alg` is absolutely +continuous with respect to the distribution over actions according to `alg₀`. -/ structure Algorithm.AbsolutelyContinuous (alg alg₀ : Algorithm 𝓐 𝓨) : Prop where p0 : alg.p0 ≪ alg₀.p0 policy n h : alg.policy n h ≪ alg₀.policy n h +@[inherit_doc Algorithm.AbsolutelyContinuous] scoped notation:50 alg " ≪ₐ " alg₀ => Algorithm.AbsolutelyContinuous alg alg₀ -/-- An algorithm that receives observations in `E × R` created form an algorithm that receives -observations in `R` by ignoring the additional information. -/ -def Algorithm.prod_left (E : Type*) [MeasurableSpace E] (alg : Algorithm 𝓐 𝓨) : - Algorithm 𝓐 (E × 𝓨) where +/-- An algorithm with observations in `𝓧 × 𝓨` obtained from an algorithm with observations in `𝓨` +by ignoring the `𝓧` component of each observation. -/ +def Algorithm.prod_left (𝓧 : Type*) [MeasurableSpace 𝓧] (alg : Algorithm 𝓐 𝓨) : + Algorithm 𝓐 (𝓧 × 𝓨) where policy n := (alg.policy n).comap (fun h i ↦ ((h i).1, (h i).2.2)) (by fun_prop) p0 := alg.p0 diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean index 8cd9f9f9..98aa4653 100644 --- a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean @@ -21,6 +21,9 @@ variable {α R : Type*} [MeasurableSpace α] [MeasurableSpace R] namespace Algorithm +/-- If the algorithm `alg` is absolutely continuous with respect to the algorithm `alg₀` and they +are both interacting with the same environment, then the law of the history at time `n` under `alg` +is the law of the history at time `n` under `alg₀` with density `alg.density alg₀ n`. -/ noncomputable def density [MeasurableSpace.CountablyGenerated α] (alg alg₀ : Algorithm α R) : (n : ℕ) → (Iic n → α × R) → ℝ≥0∞ diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 36dfb8f0..d46475a5 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -22,36 +22,42 @@ namespace Learning variable {𝓔 𝓐 𝓨 Ω : Type*} variable [MeasurableSpace 𝓔] [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] [MeasurableSpace Ω] +/-- `IsBayesAlgEnvSeq Q κ alg E A Y P` states that the random variable `E` has law `Q` under `P`. +It also states that, under `P`, the sequences of actions `A` and feedbacks `Y` are generated by the +algorithm `alg` interacting with the environment `stationaryEnv (κ.sectR (E ω))`. -/ structure IsBayesAlgEnvSeq [StandardBorelSpace 𝓐] [Nonempty 𝓐] [StandardBorelSpace 𝓨] [Nonempty 𝓨] (Q : Measure 𝓔) (κ : Kernel (𝓔 × 𝓐) 𝓨) (alg : Algorithm 𝓐 𝓨) - (E : Ω → 𝓔) (A : ℕ → Ω → 𝓐) (R' : ℕ → Ω → 𝓨) + (E : Ω → 𝓔) (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → 𝓨) (P : Measure Ω) [IsFiniteMeasure P] : Prop where measurable_E : Measurable E := by fun_prop -- todo rename measurable_action n : Measurable (A n) := by fun_prop - measurable_feedback n : Measurable (R' n) := by fun_prop + measurable_feedback n : Measurable (Y n) := by fun_prop hasLaw_env : HasLaw E Q P hasCondDistrib_action_zero : HasCondDistrib (A 0) E (Kernel.const _ alg.p0) P - hasCondDistrib_feedback_zero : HasCondDistrib (R' 0) (fun ω ↦ (E ω, A 0 ω)) κ P + hasCondDistrib_feedback_zero : HasCondDistrib (Y 0) (fun ω ↦ (E ω, A 0 ω)) κ P hasCondDistrib_action n : - HasCondDistrib (A (n + 1)) (fun ω ↦ (E ω, IsAlgEnvSeq.hist A R' n ω)) + HasCondDistrib (A (n + 1)) (fun ω ↦ (E ω, IsAlgEnvSeq.hist A Y n ω)) ((alg.policy n).prodMkLeft _) P hasCondDistrib_feedback n : - HasCondDistrib (R' (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, E ω, A (n + 1) ω)) + HasCondDistrib (Y (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A Y n ω, E ω, A (n + 1) ω)) (κ.prodMkLeft _) P namespace IsBayesAlgEnvSeq -def trajectory (A : ℕ → Ω → 𝓐) (R' : ℕ → Ω → 𝓨) (ω : Ω) : ℕ → 𝓐 × 𝓨 := fun n ↦ (A n ω, R' n ω) +/-- A random variable that gives the sequence of pairs of actions and feedbacks +(cf. `IsAlgEnvSeq.step`). -/ +def trajectory (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → 𝓨) : Ω → ℕ → 𝓐 × 𝓨 := fun ω n ↦ (A n ω, Y n ω) @[fun_prop] -lemma measurable_trajectory {A : ℕ → Ω → 𝓐} {R' : ℕ → Ω → 𝓨} (hA : ∀ n, Measurable (A n)) - (hR : ∀ n, Measurable (R' n)) : Measurable (trajectory A R') := by +lemma measurable_trajectory {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} (hA : ∀ n, Measurable (A n)) + (hR : ∀ n, Measurable (Y n)) : Measurable (trajectory A Y) := by unfold trajectory fun_prop section Real +/-- A random variable that gives the mean feedback of action `a`. -/ noncomputable def actionMean (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (a : 𝓐) (ω : Ω) : ℝ := (κ (E ω, a))[id] @@ -76,6 +82,7 @@ lemma integrable_uncurry_actionMean_comp [Countable 𝓐] [MeasurableSingletonCl apply HasFiniteIntegral.of_bounded filter_upwards with ω using abs_le_max_abs_abs (hm (E ω) (f ω)).1 (hm (E ω) (f ω)).2 +/-- A random variable that gives the action with the highest mean feedback. -/ noncomputable def bestAction [Nonempty 𝓐] [Fintype 𝓐] [Encodable 𝓐] [MeasurableSingletonClass 𝓐] (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (ω : Ω) : 𝓐 := @@ -86,7 +93,7 @@ lemma measurable_bestAction [Nonempty 𝓐] [Fintype 𝓐] [Encodable 𝓐] [Mea {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} (hE : Measurable E) : Measurable (bestAction κ E) := measurable_measurableArgmax (by fun_prop) -/-- The gap at time `n`. -/ +/-- A random variable that gives the gap at time `n`. -/ noncomputable def gap (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → 𝓐) (n : ℕ) (ω : Ω) : ℝ := Bandits.gap (κ.sectR (E ω)) (A n ω) @@ -129,6 +136,7 @@ lemma integrable_gap [Countable 𝓐] [Nonempty 𝓐] {κ : Kernel (𝓔 × 𝓐 rw [Real.norm_eq_abs, abs_of_nonneg (gap_nonneg_of_le (fun e a ↦ (h e a).2))] exact gap_le_of_mem_Icc h +/-- A random variable that gives the regret at time `n`. -/ noncomputable def regret (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → 𝓐) (n : ℕ) (ω : Ω) : ℝ := Bandits.regret (κ.sectR (E ω)) A n ω @@ -159,28 +167,28 @@ end Real variable [StandardBorelSpace 𝓐] [Nonempty 𝓐] [StandardBorelSpace 𝓨] [Nonempty 𝓨] variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × 𝓐) 𝓨} {alg : Algorithm 𝓐 𝓨} -variable {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {R' : ℕ → Ω → 𝓨} +variable {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} variable {P : Measure Ω} [IsFiniteMeasure P] section Laws -lemma hasLaw_action_zero [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : +lemma hasLaw_action_zero [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ alg E A Y P) : HasLaw (A 0) alg.p0 P := h.hasCondDistrib_action_zero.hasLaw_of_const -lemma hasCondDistrib_action' (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : - HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) (alg.policy n) P := +lemma hasCondDistrib_action' (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (n : ℕ) : + HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A Y n) (alg.policy n) P := (h.hasCondDistrib_action n).comp_right' (by fun_prop) -lemma hasCondDistrib_feedback' [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : - HasCondDistrib (R' (n + 1)) (fun ω ↦ (E ω, A (n + 1) ω)) κ P := +lemma hasCondDistrib_feedback' [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (n : ℕ) : + HasCondDistrib (Y (n + 1)) (fun ω ↦ (E ω, A (n + 1) ω)) κ P := (h.hasCondDistrib_feedback n).comp_right' (by fun_prop) end Laws section CondDistribIsAlgEnvSeq -lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : - ∀ᵐ e ∂Q, HasLaw (IT.action 0) alg.p0 (condDistrib (trajectory A R') E P e) := by +lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A Y P) : + ∀ᵐ e ∂Q, HasLaw (IT.action 0) alg.p0 (condDistrib (trajectory A Y) E P e) := by rw [← h.hasLaw_env.map_eq] filter_upwards [condDistrib_comp E ((measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable) @@ -188,32 +196,32 @@ lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : h.hasCondDistrib_action_zero.condDistrib_eq] with _ hc hcd exact ⟨(IT.measurable_action 0).aemeasurable, by rw [← Kernel.map_apply _ (IT.measurable_action 0), ← hc, - show IT.action 0 ∘ trajectory A R' = A 0 from rfl, hcd, Kernel.const_apply]⟩ + show IT.action 0 ∘ trajectory A Y = A 0 from rfl, hcd, Kernel.const_apply]⟩ -lemma hasCondDistrib_IT_feedback_zero (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : +lemma hasCondDistrib_IT_feedback_zero (h : IsBayesAlgEnvSeq Q κ alg E A Y P) : ∀ᵐ e ∂Q, HasCondDistrib (IT.feedback 0) (IT.action 0) (κ.sectR e) - (condDistrib (trajectory A R') E P e) := by + (condDistrib (trajectory A Y) E P e) := by rw [← h.hasLaw_env.map_eq] exact h.hasCondDistrib_feedback_zero.hasCondDistrib_sectR (IT.measurable_action 0) (IT.measurable_feedback 0) (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable -lemma hasCondDistrib_IT_action (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : +lemma hasCondDistrib_IT_action (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (n : ℕ) : ∀ᵐ e ∂Q, HasCondDistrib (IT.action (n + 1)) (IT.hist n) (alg.policy n) - (condDistrib (trajectory A R') E P e) := by + (condDistrib (trajectory A Y) E P e) := by rw [← h.hasLaw_env.map_eq] filter_upwards [(h.hasCondDistrib_action n).hasCondDistrib_sectR (IT.measurable_hist n) (IT.measurable_action (n + 1)) (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable] with _ he rwa [Kernel.sectR_prodMkLeft] at he -lemma hasCondDistrib_IT_feedback [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) +lemma hasCondDistrib_IT_feedback [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (n : ℕ) : ∀ᵐ e ∂Q, HasCondDistrib (IT.feedback (n + 1)) (fun τ ↦ (IT.hist n τ, IT.action (n + 1) τ)) - ((κ.sectR e).prodMkLeft _) (condDistrib (trajectory A R') E P e) := by + ((κ.sectR e).prodMkLeft _) (condDistrib (trajectory A Y) E P e) := by rw [← h.hasLaw_env.map_eq] - have hc : HasCondDistrib (R' (n + 1)) - (fun ω ↦ (E ω, IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) + have hc : HasCondDistrib (Y (n + 1)) + (fun ω ↦ (E ω, IsAlgEnvSeq.hist A Y n ω, A (n + 1) ω)) (κ.comap (fun (e, _, a) ↦ (e, a)) (by fun_prop)) P := (h.hasCondDistrib_feedback n).comp_right (MeasurableEquiv.prodAssoc.symm.trans ((MeasurableEquiv.prodCongr .prodComm (.refl _)).trans .prodAssoc)) @@ -221,19 +229,19 @@ lemma hasCondDistrib_IT_feedback [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ (IT.measurable_action (n + 1))) (IT.measurable_feedback (n + 1)) (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable -lemma hasLaw_IT_hist (h : IsBayesAlgEnvSeq Q κ alg E A R' P) (n : ℕ) : - ∀ᵐ e ∂Q, HasLaw (IT.hist n) (condDistrib (IsAlgEnvSeq.hist A R' n) E P e) - (condDistrib (trajectory A R') E P e) := by - rw [← h.hasLaw_env.map_eq, show IsAlgEnvSeq.hist A R' n = IT.hist n ∘ trajectory A R' from rfl] +lemma hasLaw_IT_hist (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (n : ℕ) : + ∀ᵐ e ∂Q, HasLaw (IT.hist n) (condDistrib (IsAlgEnvSeq.hist A Y n) E P e) + (condDistrib (trajectory A Y) E P e) := by + rw [← h.hasLaw_env.map_eq, show IsAlgEnvSeq.hist A Y n = IT.hist n ∘ trajectory A Y from rfl] filter_upwards [condDistrib_comp E (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable (IT.measurable_hist n)] with _ he exact ⟨(IT.measurable_hist n).aemeasurable, by rw [← Kernel.map_apply _ (IT.measurable_hist n), he]⟩ -lemma ae_IsAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A R' P) : +lemma ae_IsAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A Y P) : ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.feedback alg (stationaryEnv (κ.sectR e)) - (condDistrib (trajectory A R') E P e) := by + (condDistrib (trajectory A Y) E P e) := by filter_upwards [hasLaw_IT_action_zero h, hasCondDistrib_IT_feedback_zero h, ae_all_iff.2 (hasCondDistrib_IT_action h), ae_all_iff.2 (hasCondDistrib_IT_feedback h)] with _ ha0 hr0 hA hR @@ -245,6 +253,9 @@ end IsBayesAlgEnvSeq section IsAlgEnvSeq +/-- An environment with observations in `𝓔 × 𝓨`. The first element `e` of an observation is +sampled from `Q` once and remains constant. The second element of an observation is sampled from +`κ (e, a)`, where `a` is the corresponding action. -/ noncomputable def bayesStationaryEnv (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × 𝓐) 𝓨) [IsMarkovKernel κ] : Environment 𝓐 (𝓔 × 𝓨) where @@ -256,12 +267,12 @@ def bayesStationaryEnv (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel variable [Nonempty 𝓐] [Nonempty 𝓔] [Nonempty 𝓨] variable [StandardBorelSpace 𝓐] [StandardBorelSpace 𝓔] [StandardBorelSpace 𝓨] variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × 𝓐) 𝓨} [IsMarkovKernel κ] -variable {alg : Algorithm 𝓐 𝓨} {A : ℕ → Ω → 𝓐} {R' : ℕ → Ω → 𝓔 × 𝓨} +variable {alg : Algorithm 𝓐 𝓨} {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓔 × 𝓨} variable {P : Measure Ω} [IsProbabilityMeasure P] lemma IsAlgEnvSeq.isBayesAlgEnvSeq - (h : IsAlgEnvSeq A R' (alg.prod_left 𝓔) (bayesStationaryEnv Q κ) P) : - IsBayesAlgEnvSeq Q κ alg (fun ω ↦ (R' 0 ω).1) A (fun n ω ↦ (R' n ω).2) P where + (h : IsAlgEnvSeq A Y (alg.prod_left 𝓔) (bayesStationaryEnv Q κ) P) : + IsBayesAlgEnvSeq Q κ alg (fun ω ↦ (Y 0 ω).1) A (fun n ω ↦ (Y n ω).2) P where measurable_E := (h.measurable_feedback 0).fst measurable_action := h.measurable_action measurable_feedback n := (h.measurable_feedback n).snd @@ -269,7 +280,7 @@ lemma IsAlgEnvSeq.isBayesAlgEnvSeq apply HasCondDistrib.hasLaw_of_const simpa [bayesStationaryEnv] using h.hasCondDistrib_feedback_zero.fst hasCondDistrib_action_zero := by - have hc : HasCondDistrib (fun ω ↦ (R' 0 ω).1) (A 0) (Kernel.const _ Q) P := by + have hc : HasCondDistrib (fun ω ↦ (Y 0 ω).1) (A 0) (Kernel.const _ Q) P := by simpa [bayesStationaryEnv] using h.hasCondDistrib_feedback_zero.fst simpa [h.hasLaw_action_zero.map_eq, Algorithm.prod_left] using hc.const_map_of_const hasCondDistrib_feedback_zero := @@ -277,15 +288,15 @@ lemma IsAlgEnvSeq.isBayesAlgEnvSeq hasCondDistrib_action n := by let f : (Iic n → 𝓐 × 𝓔 × 𝓨) → 𝓔 × (Iic n → 𝓐 × 𝓨) := fun h ↦ ((h ⟨0, by simp⟩).2.1, fun i ↦ ((h i).1, (h i).2.2)) - have hc : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) + have hc : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A Y n) (((alg.policy n).comap Prod.snd (by fun_prop)).comap f (by fun_prop)) P := h.hasCondDistrib_action n exact hc.comp_right' (f := f) hasCondDistrib_feedback n := by let f : (Iic n → 𝓐 × 𝓔 × 𝓨) × 𝓐 → (Iic n → 𝓐 × 𝓨) × 𝓔 × 𝓐 := fun p ↦ ((fun i ↦ ((p.1 i).1, (p.1 i).2.2)), (p.1 ⟨0, by simp⟩).2.1, p.2) - have hc : HasCondDistrib (fun ω ↦ (R' (n + 1) ω).2) - (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) + have hc : HasCondDistrib (fun ω ↦ (Y (n + 1) ω).2) + (fun ω ↦ (IsAlgEnvSeq.hist A Y n ω, A (n + 1) ω)) ((Kernel.prodMkLeft ((Iic n) → 𝓐 × 𝓨) κ).comap f (by fun_prop)) P := by simpa [bayesStationaryEnv, Kernel.snd_prod] using (h.hasCondDistrib_feedback n).snd exact hc.comp_right' (by fun_prop) @@ -294,6 +305,8 @@ end IsAlgEnvSeq namespace IT +/-- A measure `P` on a measurable space that carries random variables `E`, `A`, and `Y` such that +`IsBayesAlgEnvSeq Q κ alg E A Y P`. -/ noncomputable def bayesTrajMeasure (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × 𝓐) 𝓨) [IsMarkovKernel κ] (alg : Algorithm 𝓐 𝓨) : Measure (ℕ → 𝓐 × 𝓔 × 𝓨) := @@ -309,6 +322,7 @@ lemma isBayesAlgEnvSeq_bayesTrajMeasure IsBayesAlgEnvSeq Q κ alg (fun ω ↦ (ω 0).2.1) action (fun n ω ↦ (ω n).2.2) (bayesTrajMeasure Q κ alg) := (isAlgEnvSeq_trajMeasure _ _).isBayesAlgEnvSeq +/-- A kernel that represents the posterior over `E` given the history up to time `n`. -/ noncomputable def bayesTrajMeasurePosterior [StandardBorelSpace 𝓔] [Nonempty 𝓔] (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × 𝓐) 𝓨) [IsMarkovKernel κ] From 372ab166932e096a203dd4ca0f2b75630eb52345 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 26 May 2026 15:59:15 +0100 Subject: [PATCH 135/155] Minor changes --- .../Measure/AbsolutelyContinuous.lean | 4 +- .../MeasureTheory/OuterMeasure/Basic.lean | 4 +- .../Online/Bandit/Algorithms/TS.lean | 115 +++++++++--------- .../Online/Bandit/SumRewards.lean | 24 ++-- .../Algorithms/Uniform.lean | 4 +- .../BayesStationaryEnv.lean | 4 +- 6 files changed, 77 insertions(+), 78 deletions(-) diff --git a/LeanMachineLearning/MeasureTheory/Measure/AbsolutelyContinuous.lean b/LeanMachineLearning/MeasureTheory/Measure/AbsolutelyContinuous.lean index a4770f4f..09d90000 100644 --- a/LeanMachineLearning/MeasureTheory/Measure/AbsolutelyContinuous.lean +++ b/LeanMachineLearning/MeasureTheory/Measure/AbsolutelyContinuous.lean @@ -1,7 +1,7 @@ /- -Copyright (c) 2026 Rémy Degenne. All rights reserved. +Copyright (c) 2026 Paulo Rauber. All rights reserved. Released under Apache 2.0 license as described in the file LICENSE. -Authors: Rémy Degenne, Paulo Rauber +Authors: Paulo Rauber -/ module diff --git a/LeanMachineLearning/MeasureTheory/OuterMeasure/Basic.lean b/LeanMachineLearning/MeasureTheory/OuterMeasure/Basic.lean index 8b74ec2b..4f804bd4 100644 --- a/LeanMachineLearning/MeasureTheory/OuterMeasure/Basic.lean +++ b/LeanMachineLearning/MeasureTheory/OuterMeasure/Basic.lean @@ -1,7 +1,7 @@ /- -Copyright (c) 2026 Rémy Degenne. All rights reserved. +Copyright (c) 2026 Paulo Rauber. All rights reserved. Released under Apache 2.0 license as described in the file LICENSE. -Authors: Rémy Degenne, Paulo Rauber +Authors: Paulo Rauber -/ module diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index 87689b9b..708d765d 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -1,12 +1,11 @@ /- -Copyright (c) 2026 Rémy Degenne. All rights reserved. +Copyright (c) 2026 Paulo Rauber. All rights reserved. Released under Apache 2.0 license as described in the file LICENSE. -Authors: Rémy Degenne, Paulo Rauber +Authors: Paulo Rauber -/ module public import LeanMachineLearning.Online.Bandit.SumRewards -public import LeanMachineLearning.SequentialLearning.AlgorithmDensity public import LeanMachineLearning.SequentialLearning.Algorithms.Uniform /-! # The Thompson Sampling Algorithm -/ @@ -63,76 +62,76 @@ namespace TS variable {K : ℕ} [Nonempty (Fin K)] variable {Ω : Type*} [MeasurableSpace Ω] variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] -variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} +variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] variable {P : Measure Ω} [IsProbabilityMeasure P] -lemma hasCondDistrib_action (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) - (n : ℕ) : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) - (condDistrib (bestAction κ E) (IsAlgEnvSeq.hist A R' n) P) P where +lemma hasCondDistrib_action (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) + (n : ℕ) : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R n) + (condDistrib (bestAction κ E) (IsAlgEnvSeq.hist A R n) P) P where aemeasurable_fst := (h.measurable_action (n + 1)).aemeasurable aemeasurable_snd := (IsAlgEnvSeq.measurable_hist h.measurable_action h.measurable_feedback n).aemeasurable condDistrib_eq := by have hm : Measurable (bestAction κ id) := by fun_prop calc - _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' n)] + _ =ᵐ[P.map (IsAlgEnvSeq.hist A R n)] (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) n).map (bestAction κ id) := (h.hasCondDistrib_action' n).condDistrib_eq - _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' n)] - (condDistrib E (IsAlgEnvSeq.hist A R' n) P).map (bestAction κ id) := by + _ =ᵐ[P.map (IsAlgEnvSeq.hist A R n)] + (condDistrib E (IsAlgEnvSeq.hist A R n) P).map (bestAction κ id) := by filter_upwards [(h.hasCondDistrib_env_hist (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) (absolutelyContinuous_uniformAlgorithm hK _) n).condDistrib_eq] with _ hc simp_rw [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hc] - _ =ᵐ[P.map (IsAlgEnvSeq.hist A R' n)] - condDistrib (bestAction κ E) (IsAlgEnvSeq.hist A R' n) P := - (condDistrib_comp (IsAlgEnvSeq.hist A R' n) h.measurable_E.aemeasurable hm).symm + _ =ᵐ[P.map (IsAlgEnvSeq.hist A R n)] + condDistrib (bestAction κ E) (IsAlgEnvSeq.hist A R n) P := + (condDistrib_comp (IsAlgEnvSeq.hist A R n) h.measurable_E.aemeasurable hm).symm end TS namespace ClippedUCB variable {K : ℕ} {l u σ2 δ : ℝ} -variable {Ω : Type*} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} +variable {Ω : Type*} {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} /-- Clipped upper confidence bound used in the regret analysis of Thompson sampling. -/ noncomputable -def ucb (A : ℕ → Ω → Fin K) (R' : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ := +def ucb (A : ℕ → Ω → Fin K) (R : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ := if pullCount A a n ω = 0 then u - else max l (min u (empMean A R' a n ω + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a n ω)))) + else max l (min u (empMean A R a n ω + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a n ω)))) @[simp] -lemma ucb_zero {a : Fin K} {ω : Ω} : ucb A R' l u σ2 δ a 0 ω = u := by +lemma ucb_zero {a : Fin K} {ω : Ω} : ucb A R l u σ2 δ a 0 ω = u := by simp [ucb] lemma ucb_mem_Icc (h : l ≤ u) {a : Fin K} {n : ℕ} {ω : Ω} : - ucb A R' l u σ2 δ a n ω ∈ Set.Icc l u := by + ucb A R l u σ2 δ a n ω ∈ Set.Icc l u := by unfold ucb grind @[fun_prop] lemma measurable_ucb [MeasurableSpace Ω] {a : Fin K} {n : ℕ} (hA : ∀ t, Measurable (A t)) - (hR : ∀ t, Measurable (R' t)) : Measurable (ucb A R' l u σ2 δ a n) := + (hR : ∀ t, Measurable (R t)) : Measurable (ucb A R l u σ2 δ a n) := Measurable.ite (by measurability) (by fun_prop) (by fun_prop) @[fun_prop] lemma measurable_uncurry_ucb_comp [MeasurableSpace Ω] (hA : ∀ t, Measurable (A t)) - (hR : ∀ t, Measurable (R' t)) {f : Ω → Fin K} (hf : Measurable f) {g : Ω → ℕ} - (hg : Measurable g) : Measurable (fun ω ↦ ucb A R' l u σ2 δ (f ω) (g ω) ω) := by - change Measurable ((fun aω ↦ ucb A R' l u σ2 δ aω.1 (g aω.2) aω.2) ∘ fun ω ↦ (f ω, ω)) + (hR : ∀ t, Measurable (R t)) {f : Ω → Fin K} (hf : Measurable f) {g : Ω → ℕ} + (hg : Measurable g) : Measurable (fun ω ↦ ucb A R l u σ2 δ (f ω) (g ω) ω) := by + change Measurable ((fun aω ↦ ucb A R l u σ2 δ aω.1 (g aω.2) aω.2) ∘ fun ω ↦ (f ω, ω)) apply Measurable.comp _ (by fun_prop) apply measurable_from_prod_countable_right intro a - change Measurable ((fun tω ↦ ucb A R' l u σ2 δ a tω.1 tω.2) ∘ fun ω ↦ (g ω, ω)) + change Measurable ((fun tω ↦ ucb A R l u σ2 δ a tω.1 tω.2) ∘ fun ω ↦ (g ω, ω)) apply Measurable.comp _ (by fun_prop) exact measurable_from_prod_countable_right (fun _ ↦ measurable_ucb hA hR) @[fun_prop] lemma integrable_uncurry_ucb_comp [MeasurableSpace Ω] (hA : ∀ t, Measurable (A t)) - (hR : ∀ t, Measurable (R' t)) {f : Ω → Fin K} (hf : Measurable f) {g : Ω → ℕ} + (hR : ∀ t, Measurable (R t)) {f : Ω → Fin K} (hf : Measurable f) {g : Ω → ℕ} (hg : Measurable g) {P : Measure Ω} [IsFiniteMeasure P] : - Integrable (fun ω ↦ ucb A R' l u σ2 δ (f ω) (g ω) ω) P := by + Integrable (fun ω ↦ ucb A R l u σ2 δ (f ω) (g ω) ω) P := by refine ⟨(measurable_uncurry_ucb_comp hA hR hf hg).aestronglyMeasurable, ?_⟩ apply HasFiniteIntegral.of_bounded (C := max |l| |u|) filter_upwards with ω @@ -152,10 +151,10 @@ lemma measurable_uncurry_ucb' {n : ℕ} : Measurable.ite (by measurability) (by fun_prop) (by fun_prop) lemma ucb_succ_eq_ucb' {a : Fin K} {n : ℕ} {ω : Ω} : - ucb A R' l u σ2 δ a (n + 1) ω = ucb' n (IsAlgEnvSeq.hist A R' n ω) l u σ2 δ a := by - have hp : pullCount A a (n + 1) ω = pullCount' n (IsAlgEnvSeq.hist A R' n ω) a := + ucb A R l u σ2 δ a (n + 1) ω = ucb' n (IsAlgEnvSeq.hist A R n ω) l u σ2 δ a := by + have hp : pullCount A a (n + 1) ω = pullCount' n (IsAlgEnvSeq.hist A R n ω) a := pullCount_add_one_eq_pullCount' - have he : empMean A R' a (n + 1) ω = empMean' n (IsAlgEnvSeq.hist A R' n ω) a := + have he : empMean A R a (n + 1) ω = empMean' n (IsAlgEnvSeq.hist A R n ω) a := empMean_add_one_eq_empMean' rw [ucb, ucb', hp, he] @@ -185,9 +184,9 @@ private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 nlinarith lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, μ a ∈ Set.Icc l u) (hi : l ≤ u) - (hc : ∀ s < n, pullCount A (A s ω) s ω ≠ 0 → empMean A R' (A s ω) s ω - μ (A s ω) + (hc : ∀ s < n, pullCount A (A s ω) s ω ≠ 0 → empMean A R (A s ω) s ω - μ (A s ω) < √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω))) : - ∑ s ∈ range n, (ucb A R' l u σ2 δ (A s ω) s ω - μ (A s ω)) + ∑ s ∈ range n, (ucb A R l u σ2 δ (A s ω) s ω - μ (A s ω)) ≤ (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by let S₀ := {s ∈ range n | pullCount A (A s ω) s ω = 0} let S₁ := {s ∈ range n | pullCount A (A s ω) s ω ≠ 0} @@ -195,9 +194,9 @@ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, have hd : Disjoint S₀ S₁ := disjoint_filter_filter_not _ _ _ rw [← hu, sum_union hd] gcongr - · calc ∑ s ∈ S₀, (ucb A R' l u σ2 δ (A s ω) s ω - μ (A s ω)) + · calc ∑ s ∈ S₀, (ucb A R l u σ2 δ (A s ω) s ω - μ (A s ω)) ≤ ∑ s ∈ S₀, (u - l) := - have (s : ℕ) : ucb A R' l u σ2 δ (A s ω) s ω ∈ Set.Icc l u := ucb_mem_Icc hi + have (s : ℕ) : ucb A R l u σ2 δ (A s ω) s ω ∈ Set.Icc l u := ucb_mem_Icc hi sum_le_sum (by grind) _ = ∑ s ∈ range n, if pullCount A (A s ω) s ω = 0 then (u - l) else 0 := by rw [sum_filter] @@ -209,7 +208,7 @@ lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, grind _ = (u - l) * K := by rw [Fin.sum_const, nsmul_eq_mul, mul_comm] - · calc ∑ s ∈ S₁, (ucb A R' l u σ2 δ (A s ω) s ω - μ (A s ω)) + · calc ∑ s ∈ S₁, (ucb A R l u σ2 δ (A s ω) s ω - μ (A s ω)) ≤ ∑ s ∈ S₁, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) := by gcongr with s hs unfold ucb @@ -255,17 +254,17 @@ variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ variable {P : Measure Ω} [IsProbabilityMeasure P] lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algorithm (Fin K) ℝ} - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) + (h : IsBayesAlgEnvSeq Q κ alg E A R P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (hδ : 0 < δ) (n : ℕ) : P[fun ω ↦ ∑ t ∈ range n, - (actionMean κ E (bestAction κ E ω) ω - ucb A R' l u σ2 δ (bestAction κ E ω) t ω)] ≤ + (actionMean κ E (bestAction κ E ω) ω - ucb A R l u σ2 δ (bestAction κ E ω) t ω)] ≤ (u - l) * (n - 1) * n * δ := by by_cases hn : n = 0 · simp [hn] let F := {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ - empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω ≤ + empMean A R (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω ≤ -√(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω))} have := h.measurable_action have := h.measurable_E @@ -275,7 +274,7 @@ lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algo IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp (by fun_prop) (by fun_prop) hm calc _ ≤ ∫ ω in F, ∑ t ∈ range n, - (actionMean κ E (bestAction κ E ω) ω - ucb A R' l u σ2 δ (bestAction κ E ω) t ω) ∂P := by + (actionMean κ E (bestAction κ E ω) ω - ucb A R l u σ2 δ (bestAction κ E ω) t ω) ∂P := by rw [← integral_add_compl hF (by fun_prop)] apply add_le_of_nonpos_right apply setIntegral_nonpos hF.compl @@ -304,16 +303,16 @@ lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algo ring lemma integral_sum_range_ucb_action_sub_actionMean_action_le {alg : Algorithm (Fin K) ℝ} - (h : IsBayesAlgEnvSeq Q κ alg E A R' P) + (h : IsBayesAlgEnvSeq Q κ alg E A R P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (hδ : 0 < δ) (n : ℕ) : - P[fun ω ↦ ∑ t ∈ range n, (ucb A R' l u σ2 δ (A t ω) t ω - actionMean κ E (A t ω) ω)] ≤ + P[fun ω ↦ ∑ t ∈ range n, (ucb A R l u σ2 δ (A t ω) t ω - actionMean κ E (A t ω) ω)] ≤ (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + (u - l) * K * (n - 1) * n * δ := by by_cases hn : n = 0 · simp [hn, hlu, mul_nonneg] let F := {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ empMean A R' a t ω - actionMean κ E a ω} + √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ empMean A R a t ω - actionMean κ E a ω} have := h.measurable_action have := h.measurable_E have := h.measurable_feedback @@ -365,13 +364,13 @@ variable {l u σ2 δ : ℝ} variable {Ω : Type*} [MeasurableSpace Ω] variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] -variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R' : ℕ → Ω → ℝ} +variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} variable {P : Measure Ω} [IsProbabilityMeasure P] lemma integral_ucb_action_eq_integral_ucb_bestAction (hK : 0 < K) - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) (n : ℕ) : - P[fun ω ↦ ucb A R' l u σ2 δ (A n ω) n ω] = - P[fun ω ↦ ucb A R' l u σ2 δ (bestAction κ E ω) n ω] := by + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) (n : ℕ) : + P[fun ω ↦ ucb A R l u σ2 δ (A n ω) n ω] = + P[fun ω ↦ ucb A R l u σ2 δ (bestAction κ E ω) n ω] := by have := h.measurable_action have := h.measurable_E have := h.measurable_feedback @@ -380,28 +379,28 @@ lemma integral_ucb_action_eq_integral_ucb_bestAction (hK : 0 < K) obtain ⟨n, rfl⟩ := Nat.exists_eq_succ_of_ne_zero hn let uc (ha : (Iic n → Fin K × ℝ) × Fin K) := ucb' n ha.1 l u σ2 δ ha.2 calc - _ = P[fun ω ↦ uc (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)] := by + _ = P[fun ω ↦ uc (IsAlgEnvSeq.hist A R n ω, A (n + 1) ω)] := by simp_rw [uc, ucb_succ_eq_ucb'] - _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) := by + _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R n ω, A (n + 1) ω)) := by rw [← integral_map (by fun_prop) (by fun_prop)] - _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, bestAction κ E ω)) := by + _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R n ω, bestAction κ E ω)) := by rw [← compProd_map_condDistrib (by fun_prop), ← compProd_map_condDistrib (by fun_prop), Measure.compProd_congr (hasCondDistrib_action hK h n).condDistrib_eq] - _ = P[fun ω ↦ ucb A R' l u σ2 δ (bestAction κ E ω) (n + 1) ω] := by + _ = P[fun ω ↦ ucb A R l u σ2 δ (bestAction κ E ω) (n + 1) ω] := by rw [integral_map (by fun_prop) (by fun_prop)] simp_rw [uc, ucb_succ_eq_ucb'] -lemma integral_regret_eq_add (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) +lemma integral_regret_eq_add (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A n] = P[fun ω ↦ ∑ t ∈ range n, - (actionMean κ E (bestAction κ E ω) ω - ucb A R' l u σ2 δ (bestAction κ E ω) t ω)] + + (actionMean κ E (bestAction κ E ω) ω - ucb A R l u σ2 δ (bestAction κ E ω) t ω)] + P[fun ω ↦ ∑ t ∈ range n, - (ucb A R' l u σ2 δ (A t ω) t ω - actionMean κ E (A t ω) ω)] := by - have hua (t : ℕ) : Integrable (fun ω ↦ ucb A R' l u σ2 δ (A t ω) t ω) P := + (ucb A R l u σ2 δ (A t ω) t ω - actionMean κ E (A t ω) ω)] := by + have hua (t : ℕ) : Integrable (fun ω ↦ ucb A R l u σ2 δ (A t ω) t ω) P := integrable_uncurry_ucb_comp h.measurable_action h.measurable_feedback (h.measurable_action t) measurable_const - have hub (t : ℕ) : Integrable (fun ω ↦ ucb A R' l u σ2 δ (bestAction κ E ω) t ω) P := + have hub (t : ℕ) : Integrable (fun ω ↦ ucb A R l u σ2 δ (bestAction κ E ω) t ω) P := integrable_uncurry_ucb_comp h.measurable_action h.measurable_feedback (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E) measurable_const have haa (t : ℕ) : Integrable (fun ω ↦ actionMean κ E (A t ω) ω) P := @@ -416,20 +415,20 @@ lemma integral_regret_eq_add (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorith rw [integral_finset_sum _ (by fun_prop), ← Finset.sum_sub_distrib] simp_rw [integral_sub hab (haa _)] _ = ((∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω ∂P) - - ∑ t ∈ range n, ∫ ω, ucb A R' l u σ2 δ (bestAction κ E ω) t ω ∂P) + - ((∑ t ∈ range n, ∫ ω, ucb A R' l u σ2 δ (A t ω) t ω ∂P) - + ∑ t ∈ range n, ∫ ω, ucb A R l u σ2 δ (bestAction κ E ω) t ω ∂P) + + ((∑ t ∈ range n, ∫ ω, ucb A R l u σ2 δ (A t ω) t ω ∂P) - ∑ t ∈ range n, ∫ ω, actionMean κ E (A t ω) ω ∂P) := by simp [integral_ucb_action_eq_integral_ucb_bestAction hK h] _ = (∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω - - ucb A R' l u σ2 δ (bestAction κ E ω) t ω ∂P) + - ∑ t ∈ range n, ∫ ω, ucb A R' l u σ2 δ (A t ω) t ω - + ucb A R l u σ2 δ (bestAction κ E ω) t ω ∂P) + + ∑ t ∈ range n, ∫ ω, ucb A R l u σ2 δ (A t ω) t ω - actionMean κ E (A t ω) ω ∂P := by rw [← Finset.sum_sub_distrib, ← Finset.sum_sub_distrib] simp_rw [← integral_sub hab (hub _), ← integral_sub (hua _) (haa _)] _ = _ := by rw [← integral_finset_sum _ (by fun_prop), ← integral_finset_sum _ (by fun_prop)] -lemma integral_regret_le (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R' P) +lemma integral_regret_le (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (n : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A n] diff --git a/LeanMachineLearning/Online/Bandit/SumRewards.lean b/LeanMachineLearning/Online/Bandit/SumRewards.lean index c6cc8036..cf2461e4 100644 --- a/LeanMachineLearning/Online/Bandit/SumRewards.lean +++ b/LeanMachineLearning/Online/Bandit/SumRewards.lean @@ -1,7 +1,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 +Authors: Rémy Degenne, Paulo Rauber -/ module @@ -665,7 +665,7 @@ variable {𝓔 Ω : Type*} [MeasurableSpace 𝓔] [MeasurableSpace Ω] variable {K : ℕ} [Nonempty (Fin K)] variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] variable {alg : Algorithm (Fin K) ℝ} -variable {E : Ω → 𝓔} {A : ℕ → Ω → (Fin K)} {R' : ℕ → Ω → ℝ} +variable {E : Ω → 𝓔} {A : ℕ → Ω → (Fin K)} {R : ℕ → Ω → ℝ} variable {P : Measure Ω} [IsProbabilityMeasure P] /-- Auxiliary lemma for `prob_empMean_sub_actionMean_ge_le`. -/ @@ -682,11 +682,11 @@ private lemma sqrt_two_mul_le_sub {k : ℕ} (hk : k ≠ 0) {s μ σ l : ℝ} _ = s - k * μ := by field_simp -lemma prob_empMean_sub_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) {σ2 : ℝ≥0} +lemma prob_empMean_sub_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : P {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ empMean A R' a t ω - actionMean κ E a ω} + √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ empMean A R a t ω - actionMean κ E a ω} ≤ ENNReal.ofReal (K * (n - 1) * δ) := by have := h.measurable_E have := h.measurable_action @@ -695,15 +695,15 @@ lemma prob_empMean_sub_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) √(2 * pullCount IT.action a t τ * σ2 * Real.log (1 / δ)) ≤ sumRewards IT.action IT.feedback a t τ - pullCount IT.action a t τ * actionMean κ id a e} calc - _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by + _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R ω))) S := by rw [Measure.map_apply (by fun_prop) (by measurability)] apply measure_mono intro ω ⟨t, ht, a, hpc, hle⟩ rw [empMean] at hle exact ⟨a, t, ht, hpc, sqrt_two_mul_le_sub hpc hle⟩ - _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by + _ = (P.map E ⊗ₘ condDistrib (trajectory A R) E P) S := by rw [← compProd_map_condDistrib (by fun_prop)] - _ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := + _ = ∫⁻ e, condDistrib (trajectory A R) E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := Measure.compProd_apply (by measurability) _ ≤ ∫⁻ e, ENNReal.ofReal (Fintype.card (Fin K) * (n - 1) * δ) ∂(P.map E) := by apply lintegral_mono_ae @@ -719,12 +719,12 @@ private lemma sub_le_neg_sqrt_two_mul {k : ℕ} (hk : k ≠ 0) {s μ σ l : ℝ} have : √(2 * k * σ * l) ≤ -s - k * -μ := sqrt_two_mul_le_sub hk (by grind) linarith -lemma prob_empMean_bestAction_sub_actionMean_le_le (h : IsBayesAlgEnvSeq Q κ alg E A R' P) +lemma prob_empMean_bestAction_sub_actionMean_le_le (h : IsBayesAlgEnvSeq Q κ alg E A R P) {σ2 : ℝ≥0} (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) σ2 (κ (e, a))) {δ : ℝ} (hδ : 0 < δ) (n : ℕ) : P {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ - empMean A R' (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω ≤ + empMean A R (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω ≤ -√(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω))} ≤ ENNReal.ofReal ((n - 1) * δ) := by have := h.measurable_E @@ -735,15 +735,15 @@ lemma prob_empMean_bestAction_sub_actionMean_le_le (h : IsBayesAlgEnvSeq Q κ al pullCount IT.action (bestAction κ id e) t τ * actionMean κ id (bestAction κ id e) e ≤ -√(2 * pullCount IT.action (bestAction κ id e) t τ * σ2 * Real.log (1 / δ))} calc - _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R' ω))) S := by + _ ≤ (P.map (fun ω ↦ (E ω, trajectory A R ω))) S := by rw [Measure.map_apply (by fun_prop) (by measurability)] apply measure_mono intro ω ⟨t, ht, hpc, hle⟩ rw [empMean] at hle exact ⟨t, ht, hpc, sub_le_neg_sqrt_two_mul hpc hle⟩ - _ = (P.map E ⊗ₘ condDistrib (trajectory A R') E P) S := by + _ = (P.map E ⊗ₘ condDistrib (trajectory A R) E P) S := by rw [← compProd_map_condDistrib (by fun_prop)] - _ = ∫⁻ e, condDistrib (trajectory A R') E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := + _ = ∫⁻ e, condDistrib (trajectory A R) E P e (Prod.mk e ⁻¹' S) ∂(P.map E) := Measure.compProd_apply (by measurability) _ ≤ ∫⁻ e, ENNReal.ofReal ((n - 1) * δ) ∂(P.map E) := by apply lintegral_mono_ae diff --git a/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean b/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean index f968f1f3..6ee4dd09 100644 --- a/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean +++ b/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean @@ -1,7 +1,7 @@ /- -Copyright (c) 2026 Rémy Degenne. All rights reserved. +Copyright (c) 2026 Paulo Rauber. All rights reserved. Released under Apache 2.0 license as described in the file LICENSE. -Authors: Rémy Degenne, Paulo Rauber +Authors: Paulo Rauber, Rémy Degenne -/ module diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index d46475a5..38611ae3 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -1,7 +1,7 @@ /- -Copyright (c) 2026 Rémy Degenne. All rights reserved. +Copyright (c) 2026 Paulo Rauber. All rights reserved. Released under Apache 2.0 license as described in the file LICENSE. -Authors: Rémy Degenne, Paulo Rauber +Authors: Paulo Rauber, Rémy Degenne -/ module From 7ef5b112f02e0f72acd5c536697b77acbc8560a5 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 26 May 2026 16:15:48 +0100 Subject: [PATCH 136/155] Move trajectory --- LeanMachineLearning/SequentialLearning/Algorithm.lean | 9 +++++++++ .../SequentialLearning/BayesStationaryEnv.lean | 10 ---------- 2 files changed, 9 insertions(+), 10 deletions(-) diff --git a/LeanMachineLearning/SequentialLearning/Algorithm.lean b/LeanMachineLearning/SequentialLearning/Algorithm.lean index f2d67863..db32ba6c 100644 --- a/LeanMachineLearning/SequentialLearning/Algorithm.lean +++ b/LeanMachineLearning/SequentialLearning/Algorithm.lean @@ -107,6 +107,15 @@ lemma IsAlgEnvSeq.measurable_step (n : ℕ) (hA : Measurable (A n)) unfold IsAlgEnvSeq.step fun_prop +/-- A random variable that gives the sequence of action-feedback pairs. -/ +def trajectory (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → 𝓨) (ω : Ω) : ℕ → 𝓐 × 𝓨 := fun n ↦ (A n ω, Y n ω) + +@[fun_prop] +lemma measurable_trajectory {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} (hA : ∀ n, Measurable (A n)) + (hR : ∀ n, Measurable (Y n)) : Measurable (trajectory A Y) := by + unfold trajectory + fun_prop + /-- History of the algorithm-environment sequence up to time `n`. -/ def IsAlgEnvSeq.hist (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → 𝓨) (n : ℕ) (ω : Ω) : Iic n → 𝓐 × 𝓨 := fun i ↦ (A i ω, Y i ω) diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 38611ae3..60241be5 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -45,16 +45,6 @@ structure IsBayesAlgEnvSeq namespace IsBayesAlgEnvSeq -/-- A random variable that gives the sequence of pairs of actions and feedbacks -(cf. `IsAlgEnvSeq.step`). -/ -def trajectory (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → 𝓨) : Ω → ℕ → 𝓐 × 𝓨 := fun ω n ↦ (A n ω, Y n ω) - -@[fun_prop] -lemma measurable_trajectory {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} (hA : ∀ n, Measurable (A n)) - (hR : ∀ n, Measurable (Y n)) : Measurable (trajectory A Y) := by - unfold trajectory - fun_prop - section Real /-- A random variable that gives the mean feedback of action `a`. -/ From 313f8781a4791495eefd2a2d68c4045a575cd4d9 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 26 May 2026 16:36:21 +0100 Subject: [PATCH 137/155] Add AlgorithmDensityBayes.lean --- LeanMachineLearning.lean | 1 + .../Online/Bandit/Algorithms/TS.lean | 1 + .../SequentialLearning/AlgorithmDensity.lean | 76 +------------- .../AlgorithmDensityBayes.lean | 99 +++++++++++++++++++ 4 files changed, 102 insertions(+), 75 deletions(-) create mode 100644 LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index 2d639e31..6d7c17f5 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -27,6 +27,7 @@ public import LeanMachineLearning.Probability.Moments.SubGaussian public import LeanMachineLearning.Probability.WithDensity public import LeanMachineLearning.SequentialLearning.Algorithm public import LeanMachineLearning.SequentialLearning.AlgorithmDensity +public import LeanMachineLearning.SequentialLearning.AlgorithmDensityBayes public import LeanMachineLearning.SequentialLearning.Algorithms.RandomSampling public import LeanMachineLearning.SequentialLearning.Algorithms.RoundRobin public import LeanMachineLearning.SequentialLearning.Algorithms.Uniform diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index 708d765d..90cc3a7d 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -6,6 +6,7 @@ Authors: Paulo Rauber module public import LeanMachineLearning.Online.Bandit.SumRewards +public import LeanMachineLearning.SequentialLearning.AlgorithmDensityBayes public import LeanMachineLearning.SequentialLearning.Algorithms.Uniform /-! # The Thompson Sampling Algorithm -/ diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean index 580c9225..4812c94a 100644 --- a/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensity.lean @@ -7,7 +7,7 @@ module public import LeanMachineLearning.Probability.Kernel.Composition.MeasureCompProd public import LeanMachineLearning.Probability.WithDensity -public import LeanMachineLearning.SequentialLearning.BayesStationaryEnv +public import LeanMachineLearning.SequentialLearning.Algorithm /-! # Algorithm density @@ -140,78 +140,4 @@ lemma hasLaw_hist_withDensity (h : IsAlgEnvSeq A Y alg env P) (h₀ : IsAlgEnvSe end IsAlgEnvSeq -namespace IsBayesAlgEnvSeq - -variable {𝓔 : Type*} [MeasurableSpace 𝓔] -variable [StandardBorelSpace 𝓐] [Nonempty 𝓐] [StandardBorelSpace 𝓨] [Nonempty 𝓨] -variable {Q : Measure 𝓔} -variable {κ : Kernel (𝓔 × 𝓐) 𝓨} [IsMarkovKernel κ] - -variable {Ω : Type*} [MeasurableSpace Ω] -variable {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} -variable {alg : Algorithm 𝓐 𝓨} -variable {P : Measure Ω} [IsProbabilityMeasure P] - -variable {Ω₀ : Type*} [MeasurableSpace Ω₀] -variable {E₀ : Ω₀ → 𝓔} {A₀ : ℕ → Ω₀ → 𝓐} {Y₀ : ℕ → Ω₀ → 𝓨} -variable {alg₀ : Algorithm 𝓐 𝓨} -variable {P₀ : Measure Ω₀} [IsProbabilityMeasure P₀] - -lemma condDistrib_hist_eq_condDistrib_hist_withDensity (h : IsBayesAlgEnvSeq Q κ alg E A Y P) - (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ Y₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : - condDistrib (IsAlgEnvSeq.hist A Y n) E P =ᵐ[Q] - ((condDistrib (IsAlgEnvSeq.hist A₀ Y₀ n) E₀ P₀).withDensity - (fun _ ↦ alg.density alg₀ n)) := by - filter_upwards [h.ae_IsAlgEnvSeq, h₀.ae_IsAlgEnvSeq, h.hasLaw_IT_hist n, h₀.hasLaw_IT_hist n] - with _ hae hae₀ he he₀ - rw [Kernel.withDensity_apply _ (by fun_prop), ← he.map_eq, ← he₀.map_eq] - exact (hae.hasLaw_hist_withDensity hae₀ hc n).map_eq - -lemma hasLaw_hist_withDensity (h : IsBayesAlgEnvSeq Q κ alg E A Y P) - (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ Y₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : - HasLaw (IsAlgEnvSeq.hist A Y n) - ((P₀.map (IsAlgEnvSeq.hist A₀ Y₀ n)).withDensity (alg.density alg₀ n)) P where - aemeasurable := - (IsAlgEnvSeq.measurable_hist h.measurable_action h.measurable_feedback n).aemeasurable - map_eq := by - have hA := h.measurable_action - have hY := h.measurable_feedback - have hA₀ := h₀.measurable_action - have hY₀ := h₀.measurable_feedback - have hE := h.measurable_E - have hE₀ := h₀.measurable_E - rw [← condDistrib_comp_map hE.aemeasurable (by fun_prop), h.hasLaw_env.map_eq, - Measure.bind_congr_right (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), - Kernel.comp_withDensity_eq_withDensity_comp (by fun_prop), - ← h₀.hasLaw_env.map_eq, condDistrib_comp_map hE₀.aemeasurable (by fun_prop)] - -variable [StandardBorelSpace 𝓔] [Nonempty 𝓔] -variable [IsProbabilityMeasure Q] - -lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A Y P) - (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ Y₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : - HasCondDistrib E (IsAlgEnvSeq.hist A Y n) - (condDistrib E₀ (IsAlgEnvSeq.hist A₀ Y₀ n) P₀) P where - aemeasurable_fst := h.measurable_E.aemeasurable - aemeasurable_snd := - (IsAlgEnvSeq.measurable_hist h.measurable_action h.measurable_feedback n).aemeasurable - condDistrib_eq := by - have hA := h.measurable_action - have hY := h.measurable_feedback - have hA₀ := h₀.measurable_action - have hY₀ := h₀.measurable_feedback - have hE := h.measurable_E - have hE₀ := h₀.measurable_E - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ h.measurable_E.aemeasurable, - ← map_swap_compProd_map_condDistrib (by fun_prop), h.hasLaw_env.map_eq, - Measure.compProd_eq_compProd_withDensity_comp_snd (by fun_prop) - (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), - map_swap_withDensity_comp_snd (by fun_prop), - ← h₀.hasLaw_env.map_eq, map_swap_compProd_map_condDistrib (by fun_prop), - ← compProd_map_condDistrib (by fun_prop), - ← Measure.compProd_withDensity_left (by fun_prop), - ← (hasLaw_hist_withDensity h h₀ hc n).map_eq] - -end IsBayesAlgEnvSeq - end Learning diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean new file mode 100644 index 00000000..59fa62a1 --- /dev/null +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean @@ -0,0 +1,99 @@ +/- +Copyright (c) 2026 Paulo Rauber. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Paulo Rauber +-/ +module + +public import LeanMachineLearning.SequentialLearning.AlgorithmDensity +public import LeanMachineLearning.SequentialLearning.BayesStationaryEnv + +/-! +# Algorithm density under `IsBayesAlgEnvSeq` + +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Finset + +namespace Learning + +open scoped Algorithm + +namespace IsBayesAlgEnvSeq + +variable {𝓐 𝓨 : Type*} [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] +variable {𝓔 : Type*} [MeasurableSpace 𝓔] +variable [StandardBorelSpace 𝓐] [Nonempty 𝓐] [StandardBorelSpace 𝓨] [Nonempty 𝓨] +variable {Q : Measure 𝓔} +variable {κ : Kernel (𝓔 × 𝓐) 𝓨} [IsMarkovKernel κ] + +variable {Ω : Type*} [MeasurableSpace Ω] +variable {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} +variable {alg : Algorithm 𝓐 𝓨} +variable {P : Measure Ω} [IsProbabilityMeasure P] + +variable {Ω₀ : Type*} [MeasurableSpace Ω₀] +variable {E₀ : Ω₀ → 𝓔} {A₀ : ℕ → Ω₀ → 𝓐} {Y₀ : ℕ → Ω₀ → 𝓨} +variable {alg₀ : Algorithm 𝓐 𝓨} +variable {P₀ : Measure Ω₀} [IsProbabilityMeasure P₀] + +lemma condDistrib_hist_eq_condDistrib_hist_withDensity (h : IsBayesAlgEnvSeq Q κ alg E A Y P) + (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ Y₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : + condDistrib (IsAlgEnvSeq.hist A Y n) E P =ᵐ[Q] + ((condDistrib (IsAlgEnvSeq.hist A₀ Y₀ n) E₀ P₀).withDensity + (fun _ ↦ alg.density alg₀ n)) := by + filter_upwards [h.ae_IsAlgEnvSeq, h₀.ae_IsAlgEnvSeq, h.hasLaw_IT_hist n, h₀.hasLaw_IT_hist n] + with _ hae hae₀ he he₀ + rw [Kernel.withDensity_apply _ (by fun_prop), ← he.map_eq, ← he₀.map_eq] + exact (hae.hasLaw_hist_withDensity hae₀ hc n).map_eq + +lemma hasLaw_hist_withDensity (h : IsBayesAlgEnvSeq Q κ alg E A Y P) + (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ Y₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : + HasLaw (IsAlgEnvSeq.hist A Y n) + ((P₀.map (IsAlgEnvSeq.hist A₀ Y₀ n)).withDensity (alg.density alg₀ n)) P where + aemeasurable := + (IsAlgEnvSeq.measurable_hist h.measurable_action h.measurable_feedback n).aemeasurable + map_eq := by + have hA := h.measurable_action + have hY := h.measurable_feedback + have hA₀ := h₀.measurable_action + have hY₀ := h₀.measurable_feedback + have hE := h.measurable_E + have hE₀ := h₀.measurable_E + rw [← condDistrib_comp_map hE.aemeasurable (by fun_prop), h.hasLaw_env.map_eq, + Measure.bind_congr_right (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), + Kernel.comp_withDensity_eq_withDensity_comp (by fun_prop), + ← h₀.hasLaw_env.map_eq, condDistrib_comp_map hE₀.aemeasurable (by fun_prop)] + +variable [StandardBorelSpace 𝓔] [Nonempty 𝓔] +variable [IsProbabilityMeasure Q] + +lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A Y P) + (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ Y₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : + HasCondDistrib E (IsAlgEnvSeq.hist A Y n) + (condDistrib E₀ (IsAlgEnvSeq.hist A₀ Y₀ n) P₀) P where + aemeasurable_fst := h.measurable_E.aemeasurable + aemeasurable_snd := + (IsAlgEnvSeq.measurable_hist h.measurable_action h.measurable_feedback n).aemeasurable + condDistrib_eq := by + have hA := h.measurable_action + have hY := h.measurable_feedback + have hA₀ := h₀.measurable_action + have hY₀ := h₀.measurable_feedback + have hE := h.measurable_E + have hE₀ := h₀.measurable_E + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ h.measurable_E.aemeasurable, + ← map_swap_compProd_map_condDistrib (by fun_prop), h.hasLaw_env.map_eq, + Measure.compProd_eq_compProd_withDensity_comp_snd (by fun_prop) + (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), + map_swap_withDensity_comp_snd (by fun_prop), + ← h₀.hasLaw_env.map_eq, map_swap_compProd_map_condDistrib (by fun_prop), + ← compProd_map_condDistrib (by fun_prop), + ← Measure.compProd_withDensity_left (by fun_prop), + ← (hasLaw_hist_withDensity h h₀ hc n).map_eq] + +end IsBayesAlgEnvSeq + +end Learning From cbbb4f85a5ebb4c844213af26b5a2bf078c37a49 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 26 May 2026 16:53:02 +0100 Subject: [PATCH 138/155] Generalize and document uniformAlgorithm --- .../Algorithms/Uniform.lean | 20 +++++++++++++++---- 1 file changed, 16 insertions(+), 4 deletions(-) diff --git a/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean b/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean index 6ee4dd09..3aaddfbb 100644 --- a/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean +++ b/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean @@ -8,7 +8,19 @@ module public import LeanMachineLearning.MeasureTheory.Measure.AbsolutelyContinuous public import LeanMachineLearning.SequentialLearning.AlgorithmDensity -/-! # The Uniform Algorithm -/ +/-! # The Uniform algorithm + +We introduce an algorithm that chooses actions uniformly at random. + +## Main definitions + +* `uniformAlgorithm hK`: a uniform algorithm with actions in `Fin K` given `hK : 0 < K`. + +## Main results + +* `absolutelyContinuous_uniformAlgorithm`: every algorithm with actions in `Fin K` is absolutely + continuous with respect to the uniform algorithm with the same type of feedback. +-/ @[expose] public section @@ -18,18 +30,18 @@ open scoped Algorithm namespace Bandits -variable {K : ℕ} +variable {𝓨 : Type*} {K : ℕ} /-- The Uniform algorithm: actions are chosen uniformly at random. -/ noncomputable -def uniformAlgorithm (hK : 0 < K) : Algorithm (Fin K) ℝ := +def uniformAlgorithm (hK : 0 < K) : Algorithm (Fin K) 𝓨 := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK have : IsProbabilityMeasure (uniformOn (Set.univ : Set (Fin K))) := isProbabilityMeasure_uniformOn Set.finite_univ Set.univ_nonempty { policy _ := Kernel.const _ (uniformOn Set.univ) p0 := uniformOn Set.univ } -lemma absolutelyContinuous_uniformAlgorithm (hK : 0 < K) (alg : Algorithm (Fin K) ℝ) : +lemma absolutelyContinuous_uniformAlgorithm (hK : 0 < K) (alg : Algorithm (Fin K) 𝓨) : alg ≪ₐ uniformAlgorithm hK where p0 := Measure.absolutelyContinuous_of_measure_singleton_ne_zero (by simp [uniformAlgorithm, uniformOn, ← pos_iff_ne_zero, cond_pos_of_inter_ne_zero]) From 28f7229f33b68335fb85fad7564ee6712610ff4a Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 26 May 2026 16:59:09 +0100 Subject: [PATCH 139/155] Fix --- LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean b/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean index 3aaddfbb..b989c8fe 100644 --- a/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean +++ b/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean @@ -30,7 +30,7 @@ open scoped Algorithm namespace Bandits -variable {𝓨 : Type*} {K : ℕ} +variable {𝓨 : Type*} {m𝓨 : MeasurableSpace 𝓨} {K : ℕ} /-- The Uniform algorithm: actions are chosen uniformly at random. -/ noncomputable From 9d4d7841a1756fdd3451c66700470f2f2c2ffc88 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 26 May 2026 17:05:03 +0100 Subject: [PATCH 140/155] Minor --- LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean b/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean index b989c8fe..6123e5d0 100644 --- a/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean +++ b/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean @@ -10,7 +10,7 @@ public import LeanMachineLearning.SequentialLearning.AlgorithmDensity /-! # The Uniform algorithm -We introduce an algorithm that chooses actions uniformly at random. +An algorithm that chooses actions uniformly at random in every situation. ## Main definitions From cedee49210d5ca5b05b8663596fa0b468590563b Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 29 May 2026 13:09:30 +0100 Subject: [PATCH 141/155] Document TS --- .../Online/Bandit/Algorithms/TS.lean | 38 ++++++++++++++++++- 1 file changed, 37 insertions(+), 1 deletion(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index 90cc3a7d..8cb7774c 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -9,7 +9,43 @@ public import LeanMachineLearning.Online.Bandit.SumRewards public import LeanMachineLearning.SequentialLearning.AlgorithmDensityBayes public import LeanMachineLearning.SequentialLearning.Algorithms.Uniform -/-! # The Thompson Sampling Algorithm -/ +/-! +# Thompson Sampling + +This file defines the Thompson sampling algorithm. This algorithm samples an action according to its +probability of being optimal under the posterior over environments given the history so far. + +We also provide a Bayesian regret upper bound (`integral_regret_le`) for this algorithm under the +assumption (among others) that it has the correct prior over environments. + +The Bayesian regret upper bound relies on a clipped upper confidence bound whose definition +and properties are also given in this file. + +## Main definitions + +* `tsAlgorithm hK Q κ`: a Thompson sampling algorithm with actions in `Fin K` given `hK : 0 < K`, + a prior distribution over "environments" `Q : Measure 𝓔`, and a Markov kernel + `κ : Kernel (𝓔 × Fin K) ℝ`. This kernel defines how an "environment" `e : 𝓔` gives rise to + an actual (stationary) environment `stationaryEnv (κ.sectR e) : Environment (Fin K) ℝ`. +* `ucb A R l u σ2 δ a n` : clipped upper confidence bound used in the regret analysis of Thompson + sampling for a sequence of actions `A : ℕ → Ω → Fin K`, rewards `R : ℕ → Ω → ℝ`, reward lower + bound `l : ℝ`, reward upper bound `u : ℝ`, sub-Gaussian variance proxy `σ2 : ℝ`, confidence + parameter `δ : ℝ`, action `a : Fin K`, and time `n : ℕ`. +* `ucb' n h l u σ2 δ a`: clipped upper confidence bound for action `a : Fin K` at time `n : ℕ` given + the history `h : Iic n → Fin K × ℝ` (rather than the entire sequences of actions and rewards). + +## Main results + +* `hasCondDistrib_action` : if Thompson sampling has the correct prior over environments, then + the conditional distribution of the next action given the history so far is equal to the + conditional distribution of the best action given the history so far. + +* `integral_regret_le`: if Thompson sampling has the correct prior over environments and every + "environment" has `K` actions, each of which has a corresponding reward between `l` and `u` that + is sub-Gaussian with variance proxy `σ2` after its mean is subtracted, then the Bayesian regret at + time `n` is at most `(2 * K + 1) * (u - l) + 8 * √(σ2 * K * n * Real.log n)`. + +-/ @[expose] public section From 206bc85aa56ca0161a71bafc9942630a2fed779e Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 29 May 2026 16:11:52 +0100 Subject: [PATCH 142/155] Document BayesStationaryEnv --- .../BayesStationaryEnv.lean | 50 +++++++++++++++++-- 1 file changed, 46 insertions(+), 4 deletions(-) diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 60241be5..4aec194f 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -10,7 +10,48 @@ public import LeanMachineLearning.Online.Bandit.Regret public import LeanMachineLearning.SequentialLearning.IonescuTulceaSpace public import LeanMachineLearning.SequentialLearning.StationaryEnv -/-! # Bayesian stationary environments -/ +/-! +# Bayesian stationary environments + +This file defines the structure `IsBayesAlgEnvSeq` and provides its basic properties. + +## Main definitions + +* `IsBayesAlgEnvSeq Q κ alg E A Y P`: states that there is a measure `P : Measure Ω` such + that the random variable `E : Ω → 𝓔` has law `Q` and that the sequences of actions `A : ℕ → Ω → 𝓐` + and feedbacks `Y : ℕ → Ω → 𝓨` are generated by the algorithm `alg : Algorithm 𝓐 𝓨` interacting + with an underlying environment that depends on `E` and `κ` (`stationaryEnv (κ.sectR (E ω))`). +* `bayesTrajMeasure Q κ alg`: a probability measure `P : Measure (ℕ → 𝓐 × 𝓔 × 𝓨)` on a space that + carries `E`, `A`, and `Y` such that `IsBayesAlgEnvSeq Q κ alg E A Y P` for any choice of + probability measure `Q : Measure 𝓔`, Markov kernel `κ : Kernel (𝓔 × 𝓐) 𝓨`, and + algorithm `alg : Algorithm 𝓐 𝓨`. +* `bayesTrajMeasurePosterior Q κ alg n`: a `Kernel (Iic n → 𝓐 × 𝓨) 𝓔` that represents the posterior + over `E` given the history up to time `n` under the prior `Q` and the algorithm `alg`, assuming + that the kernel `κ` controls how `E` gives rise to the underlying (stationary) environment. + See also `LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean`. + +The following definitions require feedback in `ℝ`: +* `actionMean κ E a`: the mean feedback associated with action `a : 𝓐` based on the random variable + `E`, which defines the underlying stationary environment together with the kernel `κ`. +* `bestAction κ E`: (one of) the action(s) with the highest associated mean feedback based on `E`. +* `gap κ E A n`: the difference between the highest mean feedback associated with an action and the + mean feedback associated with the action at time `n` based on `E` and the sequence of actions `A`. +* `regret κ E A n`: the regret at time `n` based on `E` and the sequence of actions `A`. If + `IsBayesAlgEnvSeq Q κ alg E A Y P`, then `P[regret κ E A n]` is the so-called Bayesian regret of + algorithm `alg` under the prior `Q`. + +## Main results + +* `ae_IsAlgEnvSeq h`: if `h : IsBayesAlgEnvSeq Q κ alg E A Y P`, for `Q`-almost every `e : 𝓔`, + `IsAlgEnvSeq A' Y' alg (stationaryEnv (κ.sectR e)) (condDistrib (trajectory A Y) E P e)` for some + sequence of actions `A' : ℕ → (ℕ → 𝓐 × 𝓨) → 𝓐` and sequence of feedbacks + `Y' : ℕ → (ℕ → 𝓐 × 𝓨) → 𝓨`. Intuitively, once the observable trajectory is conditioned on an + "environment" `e : 𝓔`, the measure that carries the `IsBayesAlgEnvSeq` structure reveals a measure + that carries an `IsAlgEnvSeq` structure under the environment `stationaryEnv (κ.sectR e)` + and the same algorithm. This allows transferring results from the `IsAlgEnvSeq` structure to the + `IsBayesAlgEnvSeq` structure. + +-/ @[expose] public section @@ -22,9 +63,10 @@ namespace Learning variable {𝓔 𝓐 𝓨 Ω : Type*} variable [MeasurableSpace 𝓔] [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] [MeasurableSpace Ω] -/-- `IsBayesAlgEnvSeq Q κ alg E A Y P` states that the random variable `E` has law `Q` under `P`. -It also states that, under `P`, the sequences of actions `A` and feedbacks `Y` are generated by the -algorithm `alg` interacting with the environment `stationaryEnv (κ.sectR (E ω))`. -/ +/-- `IsBayesAlgEnvSeq Q κ alg E A Y P` states that there is a measure `P : Measure Ω` such + that the random variable `E : Ω → 𝓔` has law `Q` and that the sequences of actions `A : ℕ → Ω → 𝓐` + and feedbacks `Y : ℕ → Ω → 𝓨` are generated by the algorithm `alg : Algorithm 𝓐 𝓨` interacting + with an underlying environment that depends on `E` and `κ` (`stationaryEnv (κ.sectR (E ω))`). -/ structure IsBayesAlgEnvSeq [StandardBorelSpace 𝓐] [Nonempty 𝓐] [StandardBorelSpace 𝓨] [Nonempty 𝓨] (Q : Measure 𝓔) (κ : Kernel (𝓔 × 𝓐) 𝓨) (alg : Algorithm 𝓐 𝓨) From e79304d6e465a6582385a6e564b738cf3a0f6d33 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Fri, 29 May 2026 17:04:27 +0100 Subject: [PATCH 143/155] Document AlgorithmDensityBayes --- .../AlgorithmDensityBayes.lean | 19 ++++++++++++++++++- 1 file changed, 18 insertions(+), 1 deletion(-) diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean index 59fa62a1..f1e9d3f9 100644 --- a/LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean @@ -9,7 +9,24 @@ public import LeanMachineLearning.SequentialLearning.AlgorithmDensity public import LeanMachineLearning.SequentialLearning.BayesStationaryEnv /-! -# Algorithm density under `IsBayesAlgEnvSeq` +# Algorithm density under Bayesian stationary environments + +This file provides results about `Algorithm.density` for the Bayesian stationary environment +setting. + +## Main results + +Let `h : IsBayesAlgEnvSeq Q κ alg E A Y P`, `h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ Y₀ P₀`, and +`hc : alg ≪ₐ alg₀`. + +* `hasLaw_hist_withDensity h h₀ hc n`: the law of the history at time `n` under `P` is the law of + the history at time `n` under `P₀` with density `alg.density alg₀ n`. Intuitively, the law of the + history under `alg` can be obtained from the law of the history under `alg₀` when they are + interacting with underlying stationary environments drawn from the same distribution. +* `hasCondDistrib_env_hist h h₀ hc n`: the conditional distribution of `E` given the history at time + `n` under `P` is almost everywhere equal to the conditional distribution of `E₀` given the history + at time `n` under `P₀`. Intuitively, the posterior is independent of the algorithm used to observe + the history. -/ From dac797c383206736bd1f13513f9c1db89659d795 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 16 Jun 2026 13:26:43 +0100 Subject: [PATCH 144/155] Fix merge issues --- .../MeasureTheory/Measure/AbsolutelyContinuous.lean | 4 ++++ .../MeasureTheory/OuterMeasure/Basic.lean | 5 +++++ LeanMachineLearning/Online/Bandit/Algorithms/TS.lean | 11 +++++------ LeanMachineLearning/Online/Bandit/SumRewards.lean | 10 ++++------ .../SequentialLearning/BayesStationaryEnv.lean | 11 ++++++----- 5 files changed, 24 insertions(+), 17 deletions(-) diff --git a/LeanMachineLearning/MeasureTheory/Measure/AbsolutelyContinuous.lean b/LeanMachineLearning/MeasureTheory/Measure/AbsolutelyContinuous.lean index 09d90000..41b0522d 100644 --- a/LeanMachineLearning/MeasureTheory/Measure/AbsolutelyContinuous.lean +++ b/LeanMachineLearning/MeasureTheory/Measure/AbsolutelyContinuous.lean @@ -8,6 +8,10 @@ module public import Mathlib.MeasureTheory.Measure.AbsolutelyContinuous public import LeanMachineLearning.MeasureTheory.OuterMeasure.Basic +/-! +# Lemma about measures that assign non-zero probability to every singleton. +-/ + @[expose] public section variable {α : Type*} diff --git a/LeanMachineLearning/MeasureTheory/OuterMeasure/Basic.lean b/LeanMachineLearning/MeasureTheory/OuterMeasure/Basic.lean index 4f804bd4..6cb4f21a 100644 --- a/LeanMachineLearning/MeasureTheory/OuterMeasure/Basic.lean +++ b/LeanMachineLearning/MeasureTheory/OuterMeasure/Basic.lean @@ -7,6 +7,11 @@ module public import Mathlib.MeasureTheory.OuterMeasure.Basic +/-! +# Lemma about measures that assign non-zero probability to every singleton. + +-/ + @[expose] public section open scoped ENNReal diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index 8cb7774c..3dffe1f9 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -332,10 +332,9 @@ lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algo simp_rw [sum_const, card_range, nsmul_eq_mul, setIntegral_const, smul_eq_mul] _ ≤ ((n - 1) * δ) * (n * (u - l)) := by gcongr - · nlinarith - · have : (1 : ℝ) ≤ n := by simp [Nat.one_le_iff_ne_zero, hn] - apply ENNReal.toReal_le_of_le_ofReal (by nlinarith) - exact h.prob_empMean_bestAction_sub_actionMean_le_le hσ2 hs hδ n + have : (1 : ℝ) ≤ n := by simp [Nat.one_le_iff_ne_zero, hn] + apply ENNReal.toReal_le_of_le_ofReal (by nlinarith) + exact h.prob_empMean_bestAction_sub_actionMean_le_le hσ2 hs hδ n _ = _ := by ring @@ -449,7 +448,7 @@ lemma integral_regret_eq_add (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorith _ = (∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω ∂P) - ∑ t ∈ range n, ∫ ω, actionMean κ E (A t ω) ω ∂P := by simp_rw [IsBayesAlgEnvSeq.regret_eq_sum_gap, IsBayesAlgEnvSeq.gap_eq_sub] - rw [integral_finset_sum _ (by fun_prop), ← Finset.sum_sub_distrib] + rw [integral_finsetSum _ (by fun_prop), ← Finset.sum_sub_distrib] simp_rw [integral_sub hab (haa _)] _ = ((∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω ∂P) - ∑ t ∈ range n, ∫ ω, ucb A R l u σ2 δ (bestAction κ E ω) t ω ∂P) + @@ -463,7 +462,7 @@ lemma integral_regret_eq_add (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorith rw [← Finset.sum_sub_distrib, ← Finset.sum_sub_distrib] simp_rw [← integral_sub hab (hub _), ← integral_sub (hua _) (haa _)] _ = _ := by - rw [← integral_finset_sum _ (by fun_prop), ← integral_finset_sum _ (by fun_prop)] + rw [← integral_finsetSum _ (by fun_prop), ← integral_finsetSum _ (by fun_prop)] lemma integral_regret_le (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) diff --git a/LeanMachineLearning/Online/Bandit/SumRewards.lean b/LeanMachineLearning/Online/Bandit/SumRewards.lean index d4a29b27..e218ff0c 100644 --- a/LeanMachineLearning/Online/Bandit/SumRewards.lean +++ b/LeanMachineLearning/Online/Bandit/SumRewards.lean @@ -454,9 +454,8 @@ lemma prob_sumRewards_sub_pullCount_mul_ge_le [Countable 𝓐] {σ2 : ℝ≥0} ( _ ≤ ∑ m ∈ Icc 1 (n - 1), ENNReal.ofReal δ := by apply sum_le_sum intro m hm - convert StreamMeasure.prob_sum_range_sub_ge_le_of_HasSubgaussianMGF' hσ2 ha hδ - (mem_Icc.mp hm).1 using 2 - simp [B] + exact le_of_eq_of_le (by simp [B]) + (StreamMeasure.prob_sum_range_sub_ge_le_of_HasSubgaussianMGF' hσ2 ha hδ (mem_Icc.mp hm).1) _ = ENNReal.ofReal ((n - 1) * δ) := by by_cases hn : n = 0 · simp [hn, hδ.le] @@ -490,9 +489,8 @@ lemma prob_sumRewards_sub_pullCount_mul_le_le [Countable 𝓐] {σ2 : ℝ≥0} ( _ ≤ ∑ m ∈ Icc 1 (n - 1), ENNReal.ofReal δ := by apply sum_le_sum intro m hm - convert StreamMeasure.prob_sum_range_sub_le_le_of_HasSubgaussianMGF' hσ2 ha hδ - (mem_Icc.mp hm).1 using 2 - simp [B] + exact le_of_eq_of_le (by simp [B]) + (StreamMeasure.prob_sum_range_sub_le_le_of_HasSubgaussianMGF' hσ2 ha hδ (mem_Icc.mp hm).1) _ = ENNReal.ofReal ((n - 1) * δ) := by by_cases hn : n = 0 · simp [hn, hδ.le] diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 4aec194f..0bc0c6cf 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -193,7 +193,7 @@ lemma integrable_regret [Countable 𝓐] [Nonempty 𝓐] {κ : Kernel (𝓔 × (hA : ∀ t, Measurable (A t)) {l u : ℝ} (h : ∀ e a, (κ (e, a))[id] ∈ Set.Icc l u) : Integrable (regret κ E A n) P := by rw [regret_eq_sum_gap'] - exact integrable_finset_sum _ (fun _ _ ↦ integrable_gap hE hA h) + exact integrable_finsetSum _ (fun _ _ ↦ integrable_gap hE hA h) end Real @@ -303,7 +303,7 @@ variable {alg : Algorithm 𝓐 𝓨} {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω variable {P : Measure Ω} [IsProbabilityMeasure P] lemma IsAlgEnvSeq.isBayesAlgEnvSeq - (h : IsAlgEnvSeq A Y (alg.prod_left 𝓔) (bayesStationaryEnv Q κ) P) : + (h : IsAlgEnvSeq A Y (alg.prodLeft 𝓔) (bayesStationaryEnv Q κ) P) : IsBayesAlgEnvSeq Q κ alg (fun ω ↦ (Y 0 ω).1) A (fun n ω ↦ (Y n ω).2) P where measurable_E := (h.measurable_feedback 0).fst measurable_action := h.measurable_action @@ -314,7 +314,7 @@ lemma IsAlgEnvSeq.isBayesAlgEnvSeq hasCondDistrib_action_zero := by have hc : HasCondDistrib (fun ω ↦ (Y 0 ω).1) (A 0) (Kernel.const _ Q) P := by simpa [bayesStationaryEnv] using h.hasCondDistrib_feedback_zero.fst - simpa [h.hasLaw_action_zero.map_eq, Algorithm.prod_left] using hc.const_map_of_const + simpa [h.hasLaw_action_zero.map_eq, Algorithm.prodLeft] using hc.const_map_of_const hasCondDistrib_feedback_zero := h.hasCondDistrib_feedback_zero.of_compProd.comp_right MeasurableEquiv.prodComm hasCondDistrib_action n := by @@ -330,7 +330,8 @@ lemma IsAlgEnvSeq.isBayesAlgEnvSeq have hc : HasCondDistrib (fun ω ↦ (Y (n + 1) ω).2) (fun ω ↦ (IsAlgEnvSeq.hist A Y n ω, A (n + 1) ω)) ((Kernel.prodMkLeft ((Iic n) → 𝓐 × 𝓨) κ).comap f (by fun_prop)) P := by - simpa [bayesStationaryEnv, Kernel.snd_prod] using (h.hasCondDistrib_feedback n).snd + simpa [bayesStationaryEnv, Kernel.prodMkLeft, ← Kernel.comap_comp_right, Function.comp_def] + using (h.hasCondDistrib_feedback n).snd exact hc.comp_right' (by fun_prop) end IsAlgEnvSeq @@ -342,7 +343,7 @@ namespace IT noncomputable def bayesTrajMeasure (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × 𝓐) 𝓨) [IsMarkovKernel κ] (alg : Algorithm 𝓐 𝓨) : Measure (ℕ → 𝓐 × 𝓔 × 𝓨) := - trajMeasure (alg.prod_left 𝓔) (bayesStationaryEnv Q κ) + trajMeasure (alg.prodLeft 𝓔) (bayesStationaryEnv Q κ) deriving IsProbabilityMeasure lemma isBayesAlgEnvSeq_bayesTrajMeasure From f011a47d11b3b2ee088344b7dec4d69c775fa454 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 16 Jun 2026 13:32:27 +0100 Subject: [PATCH 145/155] Fix --- LeanMachineLearning.lean | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index f634f7e8..f187c68b 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -6,8 +6,8 @@ public import LeanMachineLearning.MeasureTheory.Measurable public import LeanMachineLearning.MeasureTheory.Measure.AbsolutelyContinuous public import LeanMachineLearning.MeasureTheory.OuterMeasure.Basic public import LeanMachineLearning.Online.Bandit.Algorithms.ETC -public import LeanMachineLearning.Online.Bandit.Algorithms.UCB public import LeanMachineLearning.Online.Bandit.Algorithms.TS +public import LeanMachineLearning.Online.Bandit.Algorithms.UCB public import LeanMachineLearning.Online.Bandit.ArrayProbSpace public import LeanMachineLearning.Online.Bandit.Regret public import LeanMachineLearning.Online.Bandit.RewardByCountMeasure From fe49e6b4ec0b05e1ea91d1a831ca002a855fce0a Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Tue, 16 Jun 2026 15:09:52 +0100 Subject: [PATCH 146/155] Generalize uniformAlgorithm --- .../Online/Bandit/Algorithms/TS.lean | 8 +++---- .../Algorithms/Uniform.lean | 24 ++++++++----------- 2 files changed, 14 insertions(+), 18 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index 3dffe1f9..5f256e02 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -68,7 +68,7 @@ noncomputable def TS.policy (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] (n : ℕ) : Kernel (Iic n → (Fin K) × ℝ) (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) n).map (bestAction κ id) + (IT.bayesTrajMeasurePosterior Q κ uniformAlgorithm n).map (bestAction κ id) instance {hK : 0 < K} {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] {n : ℕ} : IsMarkovKernel (TS.policy hK Q κ n) := @@ -113,13 +113,13 @@ lemma hasCondDistrib_action (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm have hm : Measurable (bestAction κ id) := by fun_prop calc _ =ᵐ[P.map (IsAlgEnvSeq.hist A R n)] - (IT.bayesTrajMeasurePosterior Q κ (uniformAlgorithm hK) n).map (bestAction κ id) := + (IT.bayesTrajMeasurePosterior Q κ uniformAlgorithm n).map (bestAction κ id) := (h.hasCondDistrib_action' n).condDistrib_eq _ =ᵐ[P.map (IsAlgEnvSeq.hist A R n)] (condDistrib E (IsAlgEnvSeq.hist A R n) P).map (bestAction κ id) := by filter_upwards [(h.hasCondDistrib_env_hist - (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ (uniformAlgorithm hK)) - (absolutelyContinuous_uniformAlgorithm hK _) n).condDistrib_eq] with _ hc + (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ uniformAlgorithm) + absolutelyContinuous_uniformAlgorithm n).condDistrib_eq] with _ hc simp_rw [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hc] _ =ᵐ[P.map (IsAlgEnvSeq.hist A R n)] condDistrib (bestAction κ E) (IsAlgEnvSeq.hist A R n) P := diff --git a/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean b/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean index 6123e5d0..e9b60f3f 100644 --- a/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean +++ b/LeanMachineLearning/SequentialLearning/Algorithms/Uniform.lean @@ -7,6 +7,7 @@ module public import LeanMachineLearning.MeasureTheory.Measure.AbsolutelyContinuous public import LeanMachineLearning.SequentialLearning.AlgorithmDensity +public import LeanMachineLearning.SequentialLearning.Algorithms.RandomSampling /-! # The Uniform algorithm @@ -14,11 +15,11 @@ An algorithm that chooses actions uniformly at random in every situation. ## Main definitions -* `uniformAlgorithm hK`: a uniform algorithm with actions in `Fin K` given `hK : 0 < K`. +* `uniformAlgorithm`: a uniform algorithm with actions in a finite non-empty type `𝓐`. ## Main results -* `absolutelyContinuous_uniformAlgorithm`: every algorithm with actions in `Fin K` is absolutely +* `absolutelyContinuous_uniformAlgorithm`: every algorithm with actions in `𝓐` is absolutely continuous with respect to the uniform algorithm with the same type of feedback. -/ @@ -28,24 +29,19 @@ open MeasureTheory ProbabilityTheory Learning open scoped Algorithm -namespace Bandits +namespace Learning -variable {𝓨 : Type*} {m𝓨 : MeasurableSpace 𝓨} {K : ℕ} +variable {𝓐 𝓨 : Type*} {m𝓐 : MeasurableSpace 𝓐} {m𝓨 : MeasurableSpace 𝓨} /-- The Uniform algorithm: actions are chosen uniformly at random. -/ noncomputable -def uniformAlgorithm (hK : 0 < K) : Algorithm (Fin K) 𝓨 := - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - have : IsProbabilityMeasure (uniformOn (Set.univ : Set (Fin K))) := - isProbabilityMeasure_uniformOn Set.finite_univ Set.univ_nonempty - { policy _ := Kernel.const _ (uniformOn Set.univ) - p0 := uniformOn Set.univ } - -lemma absolutelyContinuous_uniformAlgorithm (hK : 0 < K) (alg : Algorithm (Fin K) 𝓨) : - alg ≪ₐ uniformAlgorithm hK where +def uniformAlgorithm [Finite 𝓐] [Nonempty 𝓐] : Algorithm 𝓐 𝓨 := randomSampling (uniformOn Set.univ) + +lemma absolutelyContinuous_uniformAlgorithm [Finite 𝓐] [Nonempty 𝓐] {alg : Algorithm 𝓐 𝓨} : + alg ≪ₐ uniformAlgorithm where p0 := Measure.absolutelyContinuous_of_measure_singleton_ne_zero (by simp [uniformAlgorithm, uniformOn, ← pos_iff_ne_zero, cond_pos_of_inter_ne_zero]) policy n h := Measure.absolutelyContinuous_of_measure_singleton_ne_zero (by simp [uniformAlgorithm, uniformOn, ← pos_iff_ne_zero, cond_pos_of_inter_ne_zero]) -end Bandits +end Learning From d37bbb982ea0cdafd97646eefd1e13f92d1d9152 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 17 Jun 2026 10:40:18 +0100 Subject: [PATCH 147/155] Change field name --- .../Online/Bandit/Algorithms/TS.lean | 17 +++++++++-------- .../Online/Bandit/SumRewards.lean | 8 ++++---- .../AlgorithmDensityBayes.lean | 12 ++++++------ .../SequentialLearning/BayesStationaryEnv.lean | 4 ++-- 4 files changed, 21 insertions(+), 20 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index 5f256e02..dda7fd07 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -123,7 +123,7 @@ lemma hasCondDistrib_action (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm simp_rw [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hc] _ =ᵐ[P.map (IsAlgEnvSeq.hist A R n)] condDistrib (bestAction κ E) (IsAlgEnvSeq.hist A R n) P := - (condDistrib_comp (IsAlgEnvSeq.hist A R n) h.measurable_E.aemeasurable hm).symm + (condDistrib_comp (IsAlgEnvSeq.hist A R n) h.measurable_param.aemeasurable hm).symm end TS @@ -304,7 +304,7 @@ lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algo empMean A R (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω ≤ -√(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω))} have := h.measurable_action - have := h.measurable_E + have := h.measurable_param have := h.measurable_feedback have hF : MeasurableSet F := by measurability have : Integrable (fun ω ↦ actionMean κ E (bestAction κ E ω) ω) P := @@ -350,7 +350,7 @@ lemma integral_sum_range_ucb_action_sub_actionMean_action_le {alg : Algorithm (F let F := {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ empMean A R a t ω - actionMean κ E a ω} have := h.measurable_action - have := h.measurable_E + have := h.measurable_param have := h.measurable_feedback have hF : MeasurableSet F := by measurability have : ∀ t, Integrable (fun ω ↦ actionMean κ E (A t ω) ω) P := @@ -408,7 +408,7 @@ lemma integral_ucb_action_eq_integral_ucb_bestAction (hK : 0 < K) P[fun ω ↦ ucb A R l u σ2 δ (A n ω) n ω] = P[fun ω ↦ ucb A R l u σ2 δ (bestAction κ E ω) n ω] := by have := h.measurable_action - have := h.measurable_E + have := h.measurable_param have := h.measurable_feedback by_cases hn : n = 0 · simp [hn] @@ -438,12 +438,13 @@ lemma integral_regret_eq_add (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorith measurable_const have hub (t : ℕ) : Integrable (fun ω ↦ ucb A R l u σ2 δ (bestAction κ E ω) t ω) P := integrable_uncurry_ucb_comp h.measurable_action h.measurable_feedback - (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E) measurable_const + (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_param) measurable_const have haa (t : ℕ) : Integrable (fun ω ↦ actionMean κ E (A t ω) ω) P := - IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E (h.measurable_action t) hm + IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_param + (h.measurable_action t) hm have hab : Integrable (fun ω ↦ actionMean κ E (bestAction κ E ω) ω) P := - IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_E - (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_E) hm + IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_param + (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_param) hm calc _ = (∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω ∂P) - ∑ t ∈ range n, ∫ ω, actionMean κ E (A t ω) ω ∂P := by diff --git a/LeanMachineLearning/Online/Bandit/SumRewards.lean b/LeanMachineLearning/Online/Bandit/SumRewards.lean index e68519ee..e3f8f0c7 100644 --- a/LeanMachineLearning/Online/Bandit/SumRewards.lean +++ b/LeanMachineLearning/Online/Bandit/SumRewards.lean @@ -691,7 +691,7 @@ lemma prob_empMean_sub_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R P) P {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ empMean A R a t ω - actionMean κ E a ω} ≤ ENNReal.ofReal (K * (n - 1) * δ) := by - have := h.measurable_E + have := h.measurable_param have := h.measurable_action have := h.measurable_feedback let S := {(e, τ) | ∃ a, ∃ t < n, pullCount IT.action a t τ ≠ 0 ∧ @@ -714,7 +714,7 @@ lemma prob_empMean_sub_actionMean_ge_le (h : IsBayesAlgEnvSeq Q κ alg E A R P) filter_upwards [h.ae_IsAlgEnvSeq] with e he exact Bandits.prob_sumRewards_sub_pullCount_mul_ge_le_of_Fintype hσ2 (hs e) he hδ _ = ENNReal.ofReal (K * (n - 1) * δ) := by - simp [Measure.map_apply h.measurable_E] + simp [Measure.map_apply h.measurable_param] /-- Auxiliary lemma for `prob_empMean_bestAction_sub_actionMean_le_le`. -/ private lemma sub_le_neg_sqrt_two_mul {k : ℕ} (hk : k ≠ 0) {s μ σ l : ℝ} @@ -730,7 +730,7 @@ lemma prob_empMean_bestAction_sub_actionMean_le_le (h : IsBayesAlgEnvSeq Q κ al empMean A R (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω ≤ -√(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω))} ≤ ENNReal.ofReal ((n - 1) * δ) := by - have := h.measurable_E + have := h.measurable_param have := h.measurable_action have := h.measurable_feedback let S := {(e, τ) | ∃ t < n, pullCount IT.action (bestAction κ id e) t τ ≠ 0 ∧ @@ -755,6 +755,6 @@ lemma prob_empMean_bestAction_sub_actionMean_le_le (h : IsBayesAlgEnvSeq Q κ al exact Bandits.prob_sumRewards_sub_pullCount_mul_le_le (ν := κ.sectR e) hσ2 (hs e _) he hδ _ = ENNReal.ofReal ((n - 1) * δ) := by - simp [Measure.map_apply h.measurable_E] + simp [Measure.map_apply h.measurable_param] end Learning.IsBayesAlgEnvSeq diff --git a/LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean b/LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean index f1e9d3f9..42f381b7 100644 --- a/LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean +++ b/LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean @@ -77,8 +77,8 @@ lemma hasLaw_hist_withDensity (h : IsBayesAlgEnvSeq Q κ alg E A Y P) have hY := h.measurable_feedback have hA₀ := h₀.measurable_action have hY₀ := h₀.measurable_feedback - have hE := h.measurable_E - have hE₀ := h₀.measurable_E + have hE := h.measurable_param + have hE₀ := h₀.measurable_param rw [← condDistrib_comp_map hE.aemeasurable (by fun_prop), h.hasLaw_env.map_eq, Measure.bind_congr_right (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), Kernel.comp_withDensity_eq_withDensity_comp (by fun_prop), @@ -91,7 +91,7 @@ lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (h₀ : IsBayesAlgEnvSeq Q κ alg₀ E₀ A₀ Y₀ P₀) (hc : alg ≪ₐ alg₀) (n : ℕ) : HasCondDistrib E (IsAlgEnvSeq.hist A Y n) (condDistrib E₀ (IsAlgEnvSeq.hist A₀ Y₀ n) P₀) P where - aemeasurable_fst := h.measurable_E.aemeasurable + aemeasurable_fst := h.measurable_param.aemeasurable aemeasurable_snd := (IsAlgEnvSeq.measurable_hist h.measurable_action h.measurable_feedback n).aemeasurable condDistrib_eq := by @@ -99,9 +99,9 @@ lemma hasCondDistrib_env_hist (h : IsBayesAlgEnvSeq Q κ alg E A Y P) have hY := h.measurable_feedback have hA₀ := h₀.measurable_action have hY₀ := h₀.measurable_feedback - have hE := h.measurable_E - have hE₀ := h₀.measurable_E - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ h.measurable_E.aemeasurable, + have hE := h.measurable_param + have hE₀ := h₀.measurable_param + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ h.measurable_param.aemeasurable, ← map_swap_compProd_map_condDistrib (by fun_prop), h.hasLaw_env.map_eq, Measure.compProd_eq_compProd_withDensity_comp_snd (by fun_prop) (h.condDistrib_hist_eq_condDistrib_hist_withDensity h₀ hc n), diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 7c1c90f0..c9dc3806 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -72,7 +72,7 @@ structure IsBayesAlgEnvSeq (Q : Measure 𝓔) (κ : Kernel (𝓔 × 𝓐) 𝓨) (alg : Algorithm 𝓐 𝓨) (E : Ω → 𝓔) (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → 𝓨) (P : Measure Ω) [IsFiniteMeasure P] : Prop where - measurable_E : Measurable E := by fun_prop -- todo rename + measurable_param : Measurable E := by fun_prop measurable_action n : Measurable (A n) := by fun_prop measurable_feedback n : Measurable (Y n) := by fun_prop hasLaw_env : HasLaw E Q P @@ -305,7 +305,7 @@ variable {P : Measure Ω} [IsProbabilityMeasure P] lemma IsAlgEnvSeq.isBayesAlgEnvSeq (h : IsAlgEnvSeq A Y (alg.prodLeft 𝓔) (bayesStationaryEnv Q κ) P) : IsBayesAlgEnvSeq Q κ alg (fun ω ↦ (Y 0 ω).1) A (fun n ω ↦ (Y n ω).2) P where - measurable_E := (h.measurable_feedback 0).fst + measurable_param := (h.measurable_feedback 0).fst measurable_action := h.measurable_action measurable_feedback n := (h.measurable_feedback n).snd hasLaw_env := by From 76ad93694cb1eaefaf3c882ddcf363514ebf8c8c Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 17 Jun 2026 11:07:39 +0100 Subject: [PATCH 148/155] Update documentation --- .../Online/Bandit/Algorithms/TS.lean | 12 +++++----- .../BayesStationaryEnv.lean | 24 +++++++++---------- 2 files changed, 18 insertions(+), 18 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index dda7fd07..30fcbc2a 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -24,9 +24,9 @@ and properties are also given in this file. ## Main definitions * `tsAlgorithm hK Q κ`: a Thompson sampling algorithm with actions in `Fin K` given `hK : 0 < K`, - a prior distribution over "environments" `Q : Measure 𝓔`, and a Markov kernel - `κ : Kernel (𝓔 × Fin K) ℝ`. This kernel defines how an "environment" `e : 𝓔` gives rise to - an actual (stationary) environment `stationaryEnv (κ.sectR e) : Environment (Fin K) ℝ`. + a prior distribution over parameters `Q : Measure 𝓔`, and a Markov kernel + `κ : Kernel (𝓔 × Fin K) ℝ`. This kernel defines how a parameter `e : 𝓔` gives rise to + a stationary environment: `stationaryEnv (κ.sectR e) : Environment (Fin K) ℝ`. * `ucb A R l u σ2 δ a n` : clipped upper confidence bound used in the regret analysis of Thompson sampling for a sequence of actions `A : ℕ → Ω → Fin K`, rewards `R : ℕ → Ω → ℝ`, reward lower bound `l : ℝ`, reward upper bound `u : ℝ`, sub-Gaussian variance proxy `σ2 : ℝ`, confidence @@ -41,7 +41,7 @@ and properties are also given in this file. conditional distribution of the best action given the history so far. * `integral_regret_le`: if Thompson sampling has the correct prior over environments and every - "environment" has `K` actions, each of which has a corresponding reward between `l` and `u` that + environment has `K` actions, each of which has a corresponding reward between `l` and `u` that is sub-Gaussian with variance proxy `σ2` after its mean is subtracted, then the Bayesian regret at time `n` is at most `(2 * K + 1) * (u - l) + 8 * √(σ2 * K * n * Real.log n)`. @@ -62,7 +62,7 @@ variable {K : ℕ} variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] /-- The Thompson sampling policy samples an action according to its probability of being optimal -under the posterior over "environments" given the history so far. +under the posterior over environments given the history so far. The posterior under a uniform algorithm is used to avoid a circular definition. -/ noncomputable def TS.policy (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) @@ -75,7 +75,7 @@ instance {hK : 0 < K} {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel ( Kernel.IsMarkovKernel.map _ (by fun_prop) /-- The initial action is sampled according to its probability of being optimal under the prior over -"environments". -/ +environments. -/ noncomputable def TS.initialPolicy (hK : 0 < K) (Q : Measure 𝓔) (κ : Kernel (𝓔 × Fin K) ℝ) : Measure (Fin K) := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index c9dc3806..2381ead7 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -18,21 +18,21 @@ This file defines the structure `IsBayesAlgEnvSeq` and provides its basic proper ## Main definitions * `IsBayesAlgEnvSeq Q κ alg E A Y P`: states that there is a measure `P : Measure Ω` such - that the random variable `E : Ω → 𝓔` has law `Q` and that the sequences of actions `A : ℕ → Ω → 𝓐` + that the parameter `E : Ω → 𝓔` has law `Q` and that the sequences of actions `A : ℕ → Ω → 𝓐` and feedbacks `Y : ℕ → Ω → 𝓨` are generated by the algorithm `alg : Algorithm 𝓐 𝓨` interacting with an underlying environment that depends on `E` and `κ` (`stationaryEnv (κ.sectR (E ω))`). -* `bayesTrajMeasure Q κ alg`: a probability measure `P : Measure (ℕ → 𝓐 × 𝓔 × 𝓨)` on a space that - carries `E`, `A`, and `Y` such that `IsBayesAlgEnvSeq Q κ alg E A Y P` for any choice of - probability measure `Q : Measure 𝓔`, Markov kernel `κ : Kernel (𝓔 × 𝓐) 𝓨`, and - algorithm `alg : Algorithm 𝓐 𝓨`. +* `bayesTrajMeasure Q κ alg`: for any choice of probability measure `Q : Measure 𝓔`, Markov kernel + `κ : Kernel (𝓔 × 𝓐) 𝓨`, and algorithm `alg : Algorithm 𝓐 𝓨`, provides a probability measure + `P : Measure (ℕ → 𝓐 × 𝓔 × 𝓨)` on a space that carries `E`, `A`, and `Y` such that + `IsBayesAlgEnvSeq Q κ alg E A Y P`. * `bayesTrajMeasurePosterior Q κ alg n`: a `Kernel (Iic n → 𝓐 × 𝓨) 𝓔` that represents the posterior over `E` given the history up to time `n` under the prior `Q` and the algorithm `alg`, assuming - that the kernel `κ` controls how `E` gives rise to the underlying (stationary) environment. + that the kernel `κ` specifies how `E` gives rise to the underlying (stationary) environment. See also `LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean`. The following definitions require feedback in `ℝ`: -* `actionMean κ E a`: the mean feedback associated with action `a : 𝓐` based on the random variable - `E`, which defines the underlying stationary environment together with the kernel `κ`. +* `actionMean κ E a`: the mean feedback associated with action `a : 𝓐` based on the parameter `E`, + which defines the underlying stationary environment together with the kernel `κ`. * `bestAction κ E`: (one of) the action(s) with the highest associated mean feedback based on `E`. * `gap κ E A n`: the difference between the highest mean feedback associated with an action and the mean feedback associated with the action at time `n` based on `E` and the sequence of actions `A`. @@ -45,9 +45,9 @@ The following definitions require feedback in `ℝ`: * `ae_IsAlgEnvSeq h`: if `h : IsBayesAlgEnvSeq Q κ alg E A Y P`, for `Q`-almost every `e : 𝓔`, `IsAlgEnvSeq A' Y' alg (stationaryEnv (κ.sectR e)) (condDistrib (trajectory A Y) E P e)` for some sequence of actions `A' : ℕ → (ℕ → 𝓐 × 𝓨) → 𝓐` and sequence of feedbacks - `Y' : ℕ → (ℕ → 𝓐 × 𝓨) → 𝓨`. Intuitively, once the observable trajectory is conditioned on an - "environment" `e : 𝓔`, the measure that carries the `IsBayesAlgEnvSeq` structure reveals a measure - that carries an `IsAlgEnvSeq` structure under the environment `stationaryEnv (κ.sectR e)` + `Y' : ℕ → (ℕ → 𝓐 × 𝓨) → 𝓨`. Intuitively, if the observable trajectory is generated by an + underlying parameter `e : 𝓔`, the measure that carries the `IsBayesAlgEnvSeq` structure reveals a + measure that carries an `IsAlgEnvSeq` structure under the environment `stationaryEnv (κ.sectR e)` and the same algorithm. This allows transferring results from the `IsAlgEnvSeq` structure to the `IsBayesAlgEnvSeq` structure. @@ -64,7 +64,7 @@ variable {𝓔 𝓐 𝓨 Ω : Type*} variable [MeasurableSpace 𝓔] [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] [MeasurableSpace Ω] /-- `IsBayesAlgEnvSeq Q κ alg E A Y P` states that there is a measure `P : Measure Ω` such - that the random variable `E : Ω → 𝓔` has law `Q` and that the sequences of actions `A : ℕ → Ω → 𝓐` + that the parameter `E : Ω → 𝓔` has law `Q` and that the sequences of actions `A : ℕ → Ω → 𝓐` and feedbacks `Y : ℕ → Ω → 𝓨` are generated by the algorithm `alg : Algorithm 𝓐 𝓨` interacting with an underlying environment that depends on `E` and `κ` (`stationaryEnv (κ.sectR (E ω))`). -/ structure IsBayesAlgEnvSeq From f9aa7afb50d79e1f3fc0c7e6f72aef9428371c7c Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 17 Jun 2026 11:16:11 +0100 Subject: [PATCH 149/155] Reorder sections --- .../BayesStationaryEnv.lean | 164 +++++++++--------- 1 file changed, 80 insertions(+), 84 deletions(-) diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 2381ead7..0c889b8e 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -87,6 +87,86 @@ structure IsBayesAlgEnvSeq namespace IsBayesAlgEnvSeq +section Laws + +variable [StandardBorelSpace 𝓐] [Nonempty 𝓐] [StandardBorelSpace 𝓨] [Nonempty 𝓨] +variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × 𝓐) 𝓨} {alg : Algorithm 𝓐 𝓨} +variable {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} +variable {P : Measure Ω} [IsFiniteMeasure P] + +lemma hasLaw_action_zero [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ alg E A Y P) : + HasLaw (A 0) alg.p0 P := h.hasCondDistrib_action_zero.hasLaw_of_const + +lemma hasCondDistrib_action' (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (n : ℕ) : + HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A Y n) (alg.policy n) P := + (h.hasCondDistrib_action n).comp_right' (by fun_prop) + +lemma hasCondDistrib_feedback' [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (n : ℕ) : + HasCondDistrib (Y (n + 1)) (fun ω ↦ (E ω, A (n + 1) ω)) κ P := + (h.hasCondDistrib_feedback n).comp_right' (by fun_prop) + +lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A Y P) : + ∀ᵐ e ∂Q, HasLaw (IT.action 0) alg.p0 (condDistrib (trajectory A Y) E P e) := by + rw [← h.hasLaw_env.map_eq] + filter_upwards [condDistrib_comp E + ((measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable) + (IT.measurable_action (𝓐 := 𝓐) (𝓨 := 𝓨) 0), + h.hasCondDistrib_action_zero.condDistrib_eq] with _ hc hcd + exact ⟨(IT.measurable_action 0).aemeasurable, by + rw [← Kernel.map_apply _ (IT.measurable_action 0), ← hc, + show IT.action 0 ∘ trajectory A Y = A 0 from rfl, hcd, Kernel.const_apply]⟩ + +lemma hasCondDistrib_IT_feedback_zero (h : IsBayesAlgEnvSeq Q κ alg E A Y P) : + ∀ᵐ e ∂Q, HasCondDistrib (IT.feedback 0) (IT.action 0) (κ.sectR e) + (condDistrib (trajectory A Y) E P e) := by + rw [← h.hasLaw_env.map_eq] + exact h.hasCondDistrib_feedback_zero.hasCondDistrib_sectR + (IT.measurable_action 0) (IT.measurable_feedback 0) + (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable + +lemma hasCondDistrib_IT_action (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (n : ℕ) : + ∀ᵐ e ∂Q, HasCondDistrib (IT.action (n + 1)) (IT.hist n) (alg.policy n) + (condDistrib (trajectory A Y) E P e) := by + rw [← h.hasLaw_env.map_eq] + filter_upwards [(h.hasCondDistrib_action n).hasCondDistrib_sectR + (IT.measurable_hist n) (IT.measurable_action (n + 1)) + (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable] with _ he + rwa [Kernel.sectR_prodMkLeft] at he + +lemma hasCondDistrib_IT_feedback [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A Y P) + (n : ℕ) : + ∀ᵐ e ∂Q, HasCondDistrib (IT.feedback (n + 1)) (fun τ ↦ (IT.hist n τ, IT.action (n + 1) τ)) + ((κ.sectR e).prodMkLeft _) (condDistrib (trajectory A Y) E P e) := by + rw [← h.hasLaw_env.map_eq] + have hc : HasCondDistrib (Y (n + 1)) + (fun ω ↦ (E ω, IsAlgEnvSeq.hist A Y n ω, A (n + 1) ω)) + (κ.comap (fun (e, _, a) ↦ (e, a)) (by fun_prop)) P := + (h.hasCondDistrib_feedback n).comp_right (MeasurableEquiv.prodAssoc.symm.trans + ((MeasurableEquiv.prodCongr .prodComm (.refl _)).trans .prodAssoc)) + exact hc.hasCondDistrib_sectR ((IT.measurable_hist n).prodMk + (IT.measurable_action (n + 1))) (IT.measurable_feedback (n + 1)) + (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable + +lemma hasLaw_IT_hist (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (n : ℕ) : + ∀ᵐ e ∂Q, HasLaw (IT.hist n) (condDistrib (IsAlgEnvSeq.hist A Y n) E P e) + (condDistrib (trajectory A Y) E P e) := by + rw [← h.hasLaw_env.map_eq, show IsAlgEnvSeq.hist A Y n = IT.hist n ∘ trajectory A Y from rfl] + filter_upwards [condDistrib_comp E + (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable + (IT.measurable_hist n)] with _ he + exact ⟨(IT.measurable_hist n).aemeasurable, by + rw [← Kernel.map_apply _ (IT.measurable_hist n), he]⟩ + +lemma ae_IsAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A Y P) : + ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.feedback alg (stationaryEnv (κ.sectR e)) + (condDistrib (trajectory A Y) E P e) := by + filter_upwards [hasLaw_IT_action_zero h, hasCondDistrib_IT_feedback_zero h, + ae_all_iff.2 (hasCondDistrib_IT_action h), ae_all_iff.2 (hasCondDistrib_IT_feedback h)] + with _ ha0 hr0 hA hR + exact ⟨IT.measurable_action, IT.measurable_feedback, ha0, hr0, hA, hR⟩ + +end Laws + section Real /-- A random variable that gives the mean feedback of action `a`. -/ @@ -197,90 +277,6 @@ lemma integrable_regret [Countable 𝓐] [Nonempty 𝓐] {κ : Kernel (𝓔 × end Real -variable [StandardBorelSpace 𝓐] [Nonempty 𝓐] [StandardBorelSpace 𝓨] [Nonempty 𝓨] -variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × 𝓐) 𝓨} {alg : Algorithm 𝓐 𝓨} -variable {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} -variable {P : Measure Ω} [IsFiniteMeasure P] - -section Laws - -lemma hasLaw_action_zero [IsProbabilityMeasure P] (h : IsBayesAlgEnvSeq Q κ alg E A Y P) : - HasLaw (A 0) alg.p0 P := h.hasCondDistrib_action_zero.hasLaw_of_const - -lemma hasCondDistrib_action' (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (n : ℕ) : - HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A Y n) (alg.policy n) P := - (h.hasCondDistrib_action n).comp_right' (by fun_prop) - -lemma hasCondDistrib_feedback' [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (n : ℕ) : - HasCondDistrib (Y (n + 1)) (fun ω ↦ (E ω, A (n + 1) ω)) κ P := - (h.hasCondDistrib_feedback n).comp_right' (by fun_prop) - -end Laws - -section CondDistribIsAlgEnvSeq - -lemma hasLaw_IT_action_zero (h : IsBayesAlgEnvSeq Q κ alg E A Y P) : - ∀ᵐ e ∂Q, HasLaw (IT.action 0) alg.p0 (condDistrib (trajectory A Y) E P e) := by - rw [← h.hasLaw_env.map_eq] - filter_upwards [condDistrib_comp E - ((measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable) - (IT.measurable_action (𝓐 := 𝓐) (𝓨 := 𝓨) 0), - h.hasCondDistrib_action_zero.condDistrib_eq] with _ hc hcd - exact ⟨(IT.measurable_action 0).aemeasurable, by - rw [← Kernel.map_apply _ (IT.measurable_action 0), ← hc, - show IT.action 0 ∘ trajectory A Y = A 0 from rfl, hcd, Kernel.const_apply]⟩ - -lemma hasCondDistrib_IT_feedback_zero (h : IsBayesAlgEnvSeq Q κ alg E A Y P) : - ∀ᵐ e ∂Q, HasCondDistrib (IT.feedback 0) (IT.action 0) (κ.sectR e) - (condDistrib (trajectory A Y) E P e) := by - rw [← h.hasLaw_env.map_eq] - exact h.hasCondDistrib_feedback_zero.hasCondDistrib_sectR - (IT.measurable_action 0) (IT.measurable_feedback 0) - (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable - -lemma hasCondDistrib_IT_action (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (n : ℕ) : - ∀ᵐ e ∂Q, HasCondDistrib (IT.action (n + 1)) (IT.hist n) (alg.policy n) - (condDistrib (trajectory A Y) E P e) := by - rw [← h.hasLaw_env.map_eq] - filter_upwards [(h.hasCondDistrib_action n).hasCondDistrib_sectR - (IT.measurable_hist n) (IT.measurable_action (n + 1)) - (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable] with _ he - rwa [Kernel.sectR_prodMkLeft] at he - -lemma hasCondDistrib_IT_feedback [IsFiniteKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A Y P) - (n : ℕ) : - ∀ᵐ e ∂Q, HasCondDistrib (IT.feedback (n + 1)) (fun τ ↦ (IT.hist n τ, IT.action (n + 1) τ)) - ((κ.sectR e).prodMkLeft _) (condDistrib (trajectory A Y) E P e) := by - rw [← h.hasLaw_env.map_eq] - have hc : HasCondDistrib (Y (n + 1)) - (fun ω ↦ (E ω, IsAlgEnvSeq.hist A Y n ω, A (n + 1) ω)) - (κ.comap (fun (e, _, a) ↦ (e, a)) (by fun_prop)) P := - (h.hasCondDistrib_feedback n).comp_right (MeasurableEquiv.prodAssoc.symm.trans - ((MeasurableEquiv.prodCongr .prodComm (.refl _)).trans .prodAssoc)) - exact hc.hasCondDistrib_sectR ((IT.measurable_hist n).prodMk - (IT.measurable_action (n + 1))) (IT.measurable_feedback (n + 1)) - (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable - -lemma hasLaw_IT_hist (h : IsBayesAlgEnvSeq Q κ alg E A Y P) (n : ℕ) : - ∀ᵐ e ∂Q, HasLaw (IT.hist n) (condDistrib (IsAlgEnvSeq.hist A Y n) E P e) - (condDistrib (trajectory A Y) E P e) := by - rw [← h.hasLaw_env.map_eq, show IsAlgEnvSeq.hist A Y n = IT.hist n ∘ trajectory A Y from rfl] - filter_upwards [condDistrib_comp E - (measurable_trajectory h.measurable_action h.measurable_feedback).aemeasurable - (IT.measurable_hist n)] with _ he - exact ⟨(IT.measurable_hist n).aemeasurable, by - rw [← Kernel.map_apply _ (IT.measurable_hist n), he]⟩ - -lemma ae_IsAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A Y P) : - ∀ᵐ e ∂Q, IsAlgEnvSeq IT.action IT.feedback alg (stationaryEnv (κ.sectR e)) - (condDistrib (trajectory A Y) E P e) := by - filter_upwards [hasLaw_IT_action_zero h, hasCondDistrib_IT_feedback_zero h, - ae_all_iff.2 (hasCondDistrib_IT_action h), ae_all_iff.2 (hasCondDistrib_IT_feedback h)] - with _ ha0 hr0 hA hR - exact ⟨IT.measurable_action, IT.measurable_feedback, ha0, hr0, hA, hR⟩ - -end CondDistribIsAlgEnvSeq - end IsBayesAlgEnvSeq section IsAlgEnvSeq From 1ece892ec12ff3365a90e9e9d0335f66ec4beeb9 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 17 Jun 2026 13:17:24 +0100 Subject: [PATCH 150/155] Add gap properties to Regret.lean --- LeanMachineLearning/Online/Bandit/Regret.lean | 11 +++++++++++ .../SequentialLearning/BayesStationaryEnv.lean | 10 ++++------ 2 files changed, 15 insertions(+), 6 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Regret.lean b/LeanMachineLearning/Online/Bandit/Regret.lean index 20d68f74..c5b60e4f 100644 --- a/LeanMachineLearning/Online/Bandit/Regret.lean +++ b/LeanMachineLearning/Online/Bandit/Regret.lean @@ -42,6 +42,17 @@ lemma gap_nonneg [Finite 𝓐] : 0 ≤ gap ν a := by rw [gap, sub_nonneg] exact le_ciSup (f := fun i ↦ (ν i)[id]) (by simp) a +omit [DecidableEq 𝓐] in +/-- The gap is non-negative if the means are bounded by `u : ℝ` (even if `𝓐` is not `Finite`). -/ +lemma gap_nonneg_of_le {u : ℝ} (h : ∀ a, (ν a)[id] ≤ u) : 0 ≤ gap ν a := by + rw [gap, sub_nonneg] + exact le_ciSup ⟨u, Set.forall_mem_range.2 h⟩ a + +omit [DecidableEq 𝓐] in +lemma gap_le_of_mem_Icc [Nonempty 𝓐] {l u : ℝ} (h : ∀ a, (ν a)[id] ∈ Set.Icc l u) : + gap ν a ≤ u - l := by + grind [gap, ciSup_le (fun i ↦ (h i).2)] + /-- Regret of a sequence of pulls `k : ℕ → 𝓐` at time `t` for the reward kernel `ν ; Kernel 𝓐 ℝ`. -/ noncomputable def regret (ν : Kernel 𝓐 ℝ) (A : ℕ → Ω → 𝓐) (t : ℕ) (ω : Ω) : ℝ := diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index 0c889b8e..ca3a1c11 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -213,15 +213,13 @@ def gap (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → omit [MeasurableSpace Ω] in /-- The gap is non-negative if the means are bounded by `u : ℝ` (even if `𝓐` is not `Finite`). -/ lemma gap_nonneg_of_le {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} {ω : Ω} {u : ℝ} - (h : ∀ e a, (κ (e, a))[id] ≤ u) : 0 ≤ gap κ E A n ω := by - simp_rw [gap, Bandits.gap, Kernel.sectR_apply] - linarith [le_ciSup ⟨u, Set.forall_mem_range.2 fun a ↦ (h (E ω) a)⟩ (A n ω)] + (h : ∀ e a, (κ (e, a))[id] ≤ u) : 0 ≤ gap κ E A n ω := + Bandits.gap_nonneg_of_le (h (E ω)) omit [MeasurableSpace Ω] in lemma gap_le_of_mem_Icc [Nonempty 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} - {ω : Ω} {l u : ℝ} (h : ∀ e a, (κ (e, a))[id] ∈ Set.Icc l u) : gap κ E A n ω ≤ u - l := by - simp_rw [gap, Bandits.gap, Kernel.sectR_apply] - grind [ciSup_le (fun a ↦ (h (E ω) a).2)] + {ω : Ω} {l u : ℝ} (h : ∀ e a, (κ (e, a))[id] ∈ Set.Icc l u) : gap κ E A n ω ≤ u - l := + Bandits.gap_le_of_mem_Icc (h (E ω)) omit [MeasurableSpace Ω] in lemma gap_eq_sub [Nonempty 𝓐] [Fintype 𝓐] [Encodable 𝓐] [MeasurableSingletonClass 𝓐] From 23012e0d0e95370783a53e4cd523ac381506163e Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 17 Jun 2026 13:55:54 +0100 Subject: [PATCH 151/155] Add BayesRegret --- LeanMachineLearning.lean | 1 + .../Online/Bandit/BayesRegret.lean | 148 ++++++++++++++++++ .../Online/Bandit/SumRewards.lean | 1 + .../BayesStationaryEnv.lean | 124 --------------- 4 files changed, 150 insertions(+), 124 deletions(-) create mode 100644 LeanMachineLearning/Online/Bandit/BayesRegret.lean diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index 88212f14..fa5372c1 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -22,6 +22,7 @@ public import LeanMachineLearning.Online.Bandit.Algorithms.ETC public import LeanMachineLearning.Online.Bandit.Algorithms.TS public import LeanMachineLearning.Online.Bandit.Algorithms.UCB public import LeanMachineLearning.Online.Bandit.ArrayProbSpace +public import LeanMachineLearning.Online.Bandit.BayesRegret public import LeanMachineLearning.Online.Bandit.Regret public import LeanMachineLearning.Online.Bandit.RewardByCountMeasure public import LeanMachineLearning.Online.Bandit.SumRewards diff --git a/LeanMachineLearning/Online/Bandit/BayesRegret.lean b/LeanMachineLearning/Online/Bandit/BayesRegret.lean new file mode 100644 index 00000000..0bed476d --- /dev/null +++ b/LeanMachineLearning/Online/Bandit/BayesRegret.lean @@ -0,0 +1,148 @@ +/- +Copyright (c) 2026 Paulo Rauber. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Paulo Rauber, Rémy Degenne +-/ +module + +public import LeanMachineLearning.ForMathlib.MeasureTheory.Constructions.BorelSpace.MeasurableArgMax +public import LeanMachineLearning.Online.Bandit.Regret + +/-! +# Bayesian regret + +This file defines `actionMean`, `bestAction`, `gap`, and `regret` as random variables in measurable +space `Ω`. These definitions are useful when `IsBayesAlgEnvSeq Q κ alg E A Y P`. + +Recall that `IsBayesAlgEnvSeq Q κ alg E A Y P` states that there is a measure `P : Measure Ω` such +that the parameter `E : Ω → 𝓔` has law `Q` and that the sequences of actions `A : ℕ → Ω → 𝓐` and +feedbacks `Y : ℕ → Ω → 𝓨` are generated by the algorithm `alg : Algorithm 𝓐 𝓨` interacting with an +underlying environment that depends on `E` and `κ` (`stationaryEnv (κ.sectR (E ω))`) + +## Main definitions + +* `actionMean κ E a`: the mean feedback associated with action `a : 𝓐` based on the parameter `E`, + which defines the underlying stationary environment together with the kernel `κ`. +* `bestAction κ E`: (one of) the action(s) with the highest associated mean feedback based on `E`. +* `gap κ E A n`: the difference between the highest mean feedback associated with an action and the + mean feedback associated with the action at time `n` based on `E` and the sequence of actions `A`. +* `regret κ E A n`: the regret at time `n` based on `E` and the sequence of actions `A`. If + `IsBayesAlgEnvSeq Q κ alg E A Y P`, then `P[regret κ E A n]` is the so-called Bayesian regret of + algorithm `alg` under the prior `Q`. + +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Finset + +namespace Learning.IsBayesAlgEnvSeq + +variable {𝓔 𝓐 𝓨 Ω : Type*} +variable [MeasurableSpace 𝓔] [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] [MeasurableSpace Ω] + +/-- A random variable that gives the mean feedback of action `a`. -/ +noncomputable +def actionMean (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (a : 𝓐) (ω : Ω) : ℝ := (κ (E ω, a))[id] + +@[fun_prop] +lemma measurable_actionMean {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {a : 𝓐} (hE : Measurable E) : + Measurable (actionMean κ E a) := + stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop) + +@[fun_prop] +lemma measurable_uncurry_actionMean_comp [Countable 𝓐] [MeasurableSingletonClass 𝓐] + {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} (hE : Measurable E) {f : Ω → 𝓐} (hf : Measurable f) : + Measurable (fun ω ↦ actionMean κ E (f ω) ω) := by + change Measurable ((fun aω ↦ actionMean κ E aω.1 aω.2) ∘ fun ω ↦ (f ω, ω)) + apply Measurable.comp _ (by fun_prop) + exact measurable_from_prod_countable_right (fun _ ↦ measurable_actionMean hE) + +lemma integrable_uncurry_actionMean_comp [Countable 𝓐] [MeasurableSingletonClass 𝓐] + {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} (hE : Measurable E) {f : Ω → 𝓐} (hf : Measurable f) + {P : Measure Ω} [IsFiniteMeasure P] {l u : ℝ} (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) : + Integrable (fun ω ↦ actionMean κ E (f ω) ω) P := by + refine ⟨(measurable_uncurry_actionMean_comp hE hf).aestronglyMeasurable, ?_⟩ + apply HasFiniteIntegral.of_bounded + filter_upwards with ω using abs_le_max_abs_abs (hm (E ω) (f ω)).1 (hm (E ω) (f ω)).2 + +/-- A random variable that gives the action with the highest mean feedback. -/ +noncomputable +def bestAction [Nonempty 𝓐] [Fintype 𝓐] [Encodable 𝓐] [MeasurableSingletonClass 𝓐] + (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (ω : Ω) : 𝓐 := + measurableArgmax (fun ω' a ↦ actionMean κ E a ω') ω + +@[fun_prop] +lemma measurable_bestAction [Nonempty 𝓐] [Fintype 𝓐] [Encodable 𝓐] [MeasurableSingletonClass 𝓐] + {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} (hE : Measurable E) : Measurable (bestAction κ E) := + measurable_measurableArgmax (by fun_prop) + +/-- A random variable that gives the gap at time `n`. -/ +noncomputable +def gap (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → 𝓐) (n : ℕ) (ω : Ω) : ℝ := + Bandits.gap (κ.sectR (E ω)) (A n ω) + +omit [MeasurableSpace Ω] in +/-- The gap is non-negative if the means are bounded by `u : ℝ` (even if `𝓐` is not `Finite`). -/ +lemma gap_nonneg_of_le {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} {ω : Ω} {u : ℝ} + (h : ∀ e a, (κ (e, a))[id] ≤ u) : 0 ≤ gap κ E A n ω := + Bandits.gap_nonneg_of_le (h (E ω)) + +omit [MeasurableSpace Ω] in +lemma gap_le_of_mem_Icc [Nonempty 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} + {ω : Ω} {l u : ℝ} (h : ∀ e a, (κ (e, a))[id] ∈ Set.Icc l u) : gap κ E A n ω ≤ u - l := + Bandits.gap_le_of_mem_Icc (h (E ω)) + +omit [MeasurableSpace Ω] in +lemma gap_eq_sub [Nonempty 𝓐] [Fintype 𝓐] [Encodable 𝓐] [MeasurableSingletonClass 𝓐] + {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} {ω : Ω} : + gap κ E A n ω = actionMean κ E (bestAction κ E ω) ω - actionMean κ E (A n ω) ω := by + rw [gap, Bandits.gap] + congr + apply le_antisymm + · exact ciSup_le (isMaxOn_measurableArgmax (fun ω' a ↦ actionMean κ E a ω') ω) + · exact Finite.le_ciSup (fun a ↦ actionMean κ E a ω) _ + +@[fun_prop] +lemma measurable_gap [Countable 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} + (hE : Measurable E) (hA : ∀ t, Measurable (A t)) : Measurable (gap κ E A n) := + (Measurable.iSup fun _ ↦ stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop)).sub + (stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop)) + +lemma integrable_gap [Countable 𝓐] [Nonempty 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} + {A : ℕ → Ω → 𝓐} {n : ℕ} {P : Measure Ω} [IsFiniteMeasure P] (hE : Measurable E) + (hA : ∀ t, Measurable (A t)) {l u : ℝ} (h : ∀ e a, (κ (e, a))[id] ∈ Set.Icc l u) : + Integrable (gap κ E A n) P := by + apply Integrable.of_bound (by fun_prop) (u - l) + filter_upwards with ω + rw [Real.norm_eq_abs, abs_of_nonneg (gap_nonneg_of_le (fun e a ↦ (h e a).2))] + exact gap_le_of_mem_Icc h + +/-- A random variable that gives the regret at time `n`. -/ +noncomputable +def regret (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → 𝓐) (n : ℕ) (ω : Ω) : ℝ := + Bandits.regret (κ.sectR (E ω)) A n ω + +omit [MeasurableSpace Ω] in +lemma regret_eq_sum_gap {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} {ω : Ω} : + regret κ E A n ω = ∑ s ∈ range n, gap κ E A s ω := by + simp [regret, Bandits.regret, gap, Bandits.gap] + +omit [MeasurableSpace Ω] in +lemma regret_eq_sum_gap' {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} : + regret κ E A n = fun ω ↦ ∑ s ∈ range n, gap κ E A s ω := funext fun _ ↦ regret_eq_sum_gap + +@[fun_prop] +lemma measurable_regret [Countable 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} + (hE : Measurable E) (hA : ∀ t, Measurable (A t)) : Measurable (regret κ E A n) := by + rw [regret_eq_sum_gap'] + fun_prop + +lemma integrable_regret [Countable 𝓐] [Nonempty 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} + {A : ℕ → Ω → 𝓐} {n : ℕ} {P : Measure Ω} [IsFiniteMeasure P] (hE : Measurable E) + (hA : ∀ t, Measurable (A t)) {l u : ℝ} (h : ∀ e a, (κ (e, a))[id] ∈ Set.Icc l u) : + Integrable (regret κ E A n) P := by + rw [regret_eq_sum_gap'] + exact integrable_finsetSum _ (fun _ _ ↦ integrable_gap hE hA h) + +end Learning.IsBayesAlgEnvSeq diff --git a/LeanMachineLearning/Online/Bandit/SumRewards.lean b/LeanMachineLearning/Online/Bandit/SumRewards.lean index e3f8f0c7..7a455c64 100644 --- a/LeanMachineLearning/Online/Bandit/SumRewards.lean +++ b/LeanMachineLearning/Online/Bandit/SumRewards.lean @@ -7,6 +7,7 @@ module public import LeanMachineLearning.ForMathlib.Probability.Moments.SubGaussian public import LeanMachineLearning.Online.Bandit.ArrayProbSpace +public import LeanMachineLearning.Online.Bandit.BayesRegret public import LeanMachineLearning.SequentialLearning.BayesStationaryEnv /-! # Law of the sum of rewards diff --git a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean index ca3a1c11..00fb411a 100644 --- a/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/BayesStationaryEnv.lean @@ -5,8 +5,6 @@ Authors: Paulo Rauber, Rémy Degenne -/ module -public import LeanMachineLearning.ForMathlib.MeasureTheory.Constructions.BorelSpace.MeasurableArgMax -public import LeanMachineLearning.Online.Bandit.Regret public import LeanMachineLearning.SequentialLearning.IonescuTulceaSpace public import LeanMachineLearning.SequentialLearning.StationaryEnv @@ -30,16 +28,6 @@ This file defines the structure `IsBayesAlgEnvSeq` and provides its basic proper that the kernel `κ` specifies how `E` gives rise to the underlying (stationary) environment. See also `LeanMachineLearning/SequentialLearning/AlgorithmDensityBayes.lean`. -The following definitions require feedback in `ℝ`: -* `actionMean κ E a`: the mean feedback associated with action `a : 𝓐` based on the parameter `E`, - which defines the underlying stationary environment together with the kernel `κ`. -* `bestAction κ E`: (one of) the action(s) with the highest associated mean feedback based on `E`. -* `gap κ E A n`: the difference between the highest mean feedback associated with an action and the - mean feedback associated with the action at time `n` based on `E` and the sequence of actions `A`. -* `regret κ E A n`: the regret at time `n` based on `E` and the sequence of actions `A`. If - `IsBayesAlgEnvSeq Q κ alg E A Y P`, then `P[regret κ E A n]` is the so-called Bayesian regret of - algorithm `alg` under the prior `Q`. - ## Main results * `ae_IsAlgEnvSeq h`: if `h : IsBayesAlgEnvSeq Q κ alg E A Y P`, for `Q`-almost every `e : 𝓔`, @@ -87,8 +75,6 @@ structure IsBayesAlgEnvSeq namespace IsBayesAlgEnvSeq -section Laws - variable [StandardBorelSpace 𝓐] [Nonempty 𝓐] [StandardBorelSpace 𝓨] [Nonempty 𝓨] variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × 𝓐) 𝓨} {alg : Algorithm 𝓐 𝓨} variable {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} @@ -165,116 +151,6 @@ lemma ae_IsAlgEnvSeq [IsMarkovKernel κ] (h : IsBayesAlgEnvSeq Q κ alg E A Y P) with _ ha0 hr0 hA hR exact ⟨IT.measurable_action, IT.measurable_feedback, ha0, hr0, hA, hR⟩ -end Laws - -section Real - -/-- A random variable that gives the mean feedback of action `a`. -/ -noncomputable -def actionMean (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (a : 𝓐) (ω : Ω) : ℝ := (κ (E ω, a))[id] - -@[fun_prop] -lemma measurable_actionMean {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {a : 𝓐} (hE : Measurable E) : - Measurable (actionMean κ E a) := - stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop) - -@[fun_prop] -lemma measurable_uncurry_actionMean_comp [Countable 𝓐] [MeasurableSingletonClass 𝓐] - {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} (hE : Measurable E) {f : Ω → 𝓐} (hf : Measurable f) : - Measurable (fun ω ↦ actionMean κ E (f ω) ω) := by - change Measurable ((fun aω ↦ actionMean κ E aω.1 aω.2) ∘ fun ω ↦ (f ω, ω)) - apply Measurable.comp _ (by fun_prop) - exact measurable_from_prod_countable_right (fun _ ↦ measurable_actionMean hE) - -lemma integrable_uncurry_actionMean_comp [Countable 𝓐] [MeasurableSingletonClass 𝓐] - {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} (hE : Measurable E) {f : Ω → 𝓐} (hf : Measurable f) - {P : Measure Ω} [IsFiniteMeasure P] {l u : ℝ} (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) : - Integrable (fun ω ↦ actionMean κ E (f ω) ω) P := by - refine ⟨(measurable_uncurry_actionMean_comp hE hf).aestronglyMeasurable, ?_⟩ - apply HasFiniteIntegral.of_bounded - filter_upwards with ω using abs_le_max_abs_abs (hm (E ω) (f ω)).1 (hm (E ω) (f ω)).2 - -/-- A random variable that gives the action with the highest mean feedback. -/ -noncomputable -def bestAction [Nonempty 𝓐] [Fintype 𝓐] [Encodable 𝓐] [MeasurableSingletonClass 𝓐] - (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (ω : Ω) : 𝓐 := - measurableArgmax (fun ω' a ↦ actionMean κ E a ω') ω - -@[fun_prop] -lemma measurable_bestAction [Nonempty 𝓐] [Fintype 𝓐] [Encodable 𝓐] [MeasurableSingletonClass 𝓐] - {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} (hE : Measurable E) : Measurable (bestAction κ E) := - measurable_measurableArgmax (by fun_prop) - -/-- A random variable that gives the gap at time `n`. -/ -noncomputable -def gap (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → 𝓐) (n : ℕ) (ω : Ω) : ℝ := - Bandits.gap (κ.sectR (E ω)) (A n ω) - -omit [MeasurableSpace Ω] in -/-- The gap is non-negative if the means are bounded by `u : ℝ` (even if `𝓐` is not `Finite`). -/ -lemma gap_nonneg_of_le {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} {ω : Ω} {u : ℝ} - (h : ∀ e a, (κ (e, a))[id] ≤ u) : 0 ≤ gap κ E A n ω := - Bandits.gap_nonneg_of_le (h (E ω)) - -omit [MeasurableSpace Ω] in -lemma gap_le_of_mem_Icc [Nonempty 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} - {ω : Ω} {l u : ℝ} (h : ∀ e a, (κ (e, a))[id] ∈ Set.Icc l u) : gap κ E A n ω ≤ u - l := - Bandits.gap_le_of_mem_Icc (h (E ω)) - -omit [MeasurableSpace Ω] in -lemma gap_eq_sub [Nonempty 𝓐] [Fintype 𝓐] [Encodable 𝓐] [MeasurableSingletonClass 𝓐] - {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} {ω : Ω} : - gap κ E A n ω = actionMean κ E (bestAction κ E ω) ω - actionMean κ E (A n ω) ω := by - rw [gap, Bandits.gap] - congr - apply le_antisymm - · exact ciSup_le (isMaxOn_measurableArgmax (fun ω' a ↦ actionMean κ E a ω') ω) - · exact Finite.le_ciSup (fun a ↦ actionMean κ E a ω) _ - -@[fun_prop] -lemma measurable_gap [Countable 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} - (hE : Measurable E) (hA : ∀ t, Measurable (A t)) : Measurable (gap κ E A n) := - (Measurable.iSup fun _ ↦ stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop)).sub - (stronglyMeasurable_id.integral_kernel.measurable.comp (by fun_prop)) - -lemma integrable_gap [Countable 𝓐] [Nonempty 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} - {A : ℕ → Ω → 𝓐} {n : ℕ} {P : Measure Ω} [IsFiniteMeasure P] (hE : Measurable E) - (hA : ∀ t, Measurable (A t)) {l u : ℝ} (h : ∀ e a, (κ (e, a))[id] ∈ Set.Icc l u) : - Integrable (gap κ E A n) P := by - apply Integrable.of_bound (by fun_prop) (u - l) - filter_upwards with ω - rw [Real.norm_eq_abs, abs_of_nonneg (gap_nonneg_of_le (fun e a ↦ (h e a).2))] - exact gap_le_of_mem_Icc h - -/-- A random variable that gives the regret at time `n`. -/ -noncomputable -def regret (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (A : ℕ → Ω → 𝓐) (n : ℕ) (ω : Ω) : ℝ := - Bandits.regret (κ.sectR (E ω)) A n ω - -omit [MeasurableSpace Ω] in -lemma regret_eq_sum_gap {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} {ω : Ω} : - regret κ E A n ω = ∑ s ∈ range n, gap κ E A s ω := by - simp [regret, Bandits.regret, gap, Bandits.gap] - -omit [MeasurableSpace Ω] in -lemma regret_eq_sum_gap' {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} : - regret κ E A n = fun ω ↦ ∑ s ∈ range n, gap κ E A s ω := funext fun _ ↦ regret_eq_sum_gap - -@[fun_prop] -lemma measurable_regret [Countable 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} - (hE : Measurable E) (hA : ∀ t, Measurable (A t)) : Measurable (regret κ E A n) := by - rw [regret_eq_sum_gap'] - fun_prop - -lemma integrable_regret [Countable 𝓐] [Nonempty 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} - {A : ℕ → Ω → 𝓐} {n : ℕ} {P : Measure Ω} [IsFiniteMeasure P] (hE : Measurable E) - (hA : ∀ t, Measurable (A t)) {l u : ℝ} (h : ∀ e a, (κ (e, a))[id] ∈ Set.Icc l u) : - Integrable (regret κ E A n) P := by - rw [regret_eq_sum_gap'] - exact integrable_finsetSum _ (fun _ _ ↦ integrable_gap hE hA h) - -end Real - end IsBayesAlgEnvSeq section IsAlgEnvSeq From d67251dcb1d91f9ba0aee449afbe0ff5b8bacf47 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 17 Jun 2026 14:01:55 +0100 Subject: [PATCH 152/155] Documentation fix --- LeanMachineLearning/Online/Bandit/BayesRegret.lean | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/BayesRegret.lean b/LeanMachineLearning/Online/Bandit/BayesRegret.lean index 0bed476d..981f1c70 100644 --- a/LeanMachineLearning/Online/Bandit/BayesRegret.lean +++ b/LeanMachineLearning/Online/Bandit/BayesRegret.lean @@ -11,8 +11,8 @@ public import LeanMachineLearning.Online.Bandit.Regret /-! # Bayesian regret -This file defines `actionMean`, `bestAction`, `gap`, and `regret` as random variables in measurable -space `Ω`. These definitions are useful when `IsBayesAlgEnvSeq Q κ alg E A Y P`. +This file defines `actionMean`, `bestAction`, `gap`, and `regret` as random variables in a +measurable space `Ω`. These definitions are useful when `IsBayesAlgEnvSeq Q κ alg E A Y P`. Recall that `IsBayesAlgEnvSeq Q κ alg E A Y P` states that there is a measure `P : Measure Ω` such that the parameter `E : Ω → 𝓔` has law `Q` and that the sequences of actions `A : ℕ → Ω → 𝓐` and From 79167e3abb16ff50f5b8029bfc3714248f0b6b91 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 18 Jun 2026 10:15:32 +0100 Subject: [PATCH 153/155] Split definition and regret analysis --- LeanMachineLearning.lean | 1 + .../Algorithms/Regret/BayesRegretTS.lean | 415 ++++++++++++++++++ .../Online/Bandit/Algorithms/TS.lean | 403 +---------------- 3 files changed, 419 insertions(+), 400 deletions(-) create mode 100644 LeanMachineLearning/Online/Bandit/Algorithms/Regret/BayesRegretTS.lean diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index fa5372c1..12be32f8 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -19,6 +19,7 @@ public import LeanMachineLearning.ForMathlib.Probability.Kernel.KernelSub public import LeanMachineLearning.ForMathlib.Probability.Moments.SubGaussian public import LeanMachineLearning.ForMathlib.Probability.WithDensity public import LeanMachineLearning.Online.Bandit.Algorithms.ETC +public import LeanMachineLearning.Online.Bandit.Algorithms.Regret.BayesRegretTS public import LeanMachineLearning.Online.Bandit.Algorithms.TS public import LeanMachineLearning.Online.Bandit.Algorithms.UCB public import LeanMachineLearning.Online.Bandit.ArrayProbSpace diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/Regret/BayesRegretTS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/Regret/BayesRegretTS.lean new file mode 100644 index 00000000..e0bc2e75 --- /dev/null +++ b/LeanMachineLearning/Online/Bandit/Algorithms/Regret/BayesRegretTS.lean @@ -0,0 +1,415 @@ +/- +Copyright (c) 2026 Paulo Rauber. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Paulo Rauber +-/ +module + +public import LeanMachineLearning.Online.Bandit.Algorithms.TS +public import LeanMachineLearning.Online.Bandit.SumRewards + +/-! +# Bayesian regret of Thompson sampling + +This file provides a Bayesian regret upper bound (`integral_regret_le`) for Thompson sampling under +the assumption (among others) that it has the correct prior over environments. + +The Bayesian regret upper bound relies on a clipped upper confidence bound whose definition +and properties are also given in this file. + +## Main definitions + +* `ucb A R l u σ2 δ a n` : clipped upper confidence bound used in the regret analysis of Thompson + sampling for a sequence of actions `A : ℕ → Ω → Fin K`, rewards `R : ℕ → Ω → ℝ`, reward lower + bound `l : ℝ`, reward upper bound `u : ℝ`, sub-Gaussian variance proxy `σ2 : ℝ`, confidence + parameter `δ : ℝ`, action `a : Fin K`, and time `n : ℕ`. +* `ucb' n h l u σ2 δ a`: clipped upper confidence bound for action `a : Fin K` at time `n : ℕ` given + the history `h : Iic n → Fin K × ℝ` (rather than the entire sequences of actions and rewards). + +## Main results + +* `integral_regret_le`: if Thompson sampling has the correct prior over environments and every + environment has `K` actions, each of which has a corresponding reward between `l` and `u` that + is sub-Gaussian with variance proxy `σ2` after its mean is subtracted, then the Bayesian regret at + time `n` is at most `(2 * K + 1) * (u - l) + 8 * √(σ2 * K * n * Real.log n)`. + +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Finset Learning +open IsBayesAlgEnvSeq (bestAction actionMean) + +namespace Bandits + +namespace ClippedUCB + +variable {K : ℕ} {l u σ2 δ : ℝ} +variable {Ω : Type*} {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} + +/-- Clipped upper confidence bound used in the regret analysis of Thompson sampling. -/ +noncomputable +def ucb (A : ℕ → Ω → Fin K) (R : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ := + if pullCount A a n ω = 0 then u + else max l (min u (empMean A R a n ω + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a n ω)))) + +@[simp] +lemma ucb_zero {a : Fin K} {ω : Ω} : ucb A R l u σ2 δ a 0 ω = u := by + simp [ucb] + +lemma ucb_mem_Icc (h : l ≤ u) {a : Fin K} {n : ℕ} {ω : Ω} : + ucb A R l u σ2 δ a n ω ∈ Set.Icc l u := by + unfold ucb + grind + +@[fun_prop] +lemma measurable_ucb [MeasurableSpace Ω] {a : Fin K} {n : ℕ} (hA : ∀ t, Measurable (A t)) + (hR : ∀ t, Measurable (R t)) : Measurable (ucb A R l u σ2 δ a n) := + Measurable.ite (by measurability) (by fun_prop) (by fun_prop) + +@[fun_prop] +lemma measurable_uncurry_ucb_comp [MeasurableSpace Ω] (hA : ∀ t, Measurable (A t)) + (hR : ∀ t, Measurable (R t)) {f : Ω → Fin K} (hf : Measurable f) {g : Ω → ℕ} + (hg : Measurable g) : Measurable (fun ω ↦ ucb A R l u σ2 δ (f ω) (g ω) ω) := by + change Measurable ((fun aω ↦ ucb A R l u σ2 δ aω.1 (g aω.2) aω.2) ∘ fun ω ↦ (f ω, ω)) + apply Measurable.comp _ (by fun_prop) + apply measurable_from_prod_countable_right + intro a + change Measurable ((fun tω ↦ ucb A R l u σ2 δ a tω.1 tω.2) ∘ fun ω ↦ (g ω, ω)) + apply Measurable.comp _ (by fun_prop) + exact measurable_from_prod_countable_right (fun _ ↦ measurable_ucb hA hR) + +@[fun_prop] +lemma integrable_uncurry_ucb_comp [MeasurableSpace Ω] (hA : ∀ t, Measurable (A t)) + (hR : ∀ t, Measurable (R t)) {f : Ω → Fin K} (hf : Measurable f) {g : Ω → ℕ} + (hg : Measurable g) {P : Measure Ω} [IsFiniteMeasure P] : + Integrable (fun ω ↦ ucb A R l u σ2 δ (f ω) (g ω) ω) P := by + refine ⟨(measurable_uncurry_ucb_comp hA hR hf hg).aestronglyMeasurable, ?_⟩ + apply HasFiniteIntegral.of_bounded (C := max |l| |u|) + filter_upwards with ω + rw [Real.norm_eq_abs] + unfold ucb + grind + +/-- Clipped upper confidence bound (history-based version). -/ +noncomputable +def ucb' (n : ℕ) (h : Iic n → Fin K × ℝ) (l u σ2 δ : ℝ) (a : Fin K) : ℝ := + if pullCount' n h a = 0 then u + else max l (min u (empMean' n h a + √(2 * σ2 * Real.log (1 / δ) / (pullCount' n h a)))) + +@[fun_prop] +lemma measurable_uncurry_ucb' {n : ℕ} : + Measurable (fun p : (Iic n → Fin K × ℝ) × Fin K ↦ ucb' n p.1 l u σ2 δ p.2) := + Measurable.ite (by measurability) (by fun_prop) (by fun_prop) + +lemma ucb_succ_eq_ucb' {a : Fin K} {n : ℕ} {ω : Ω} : + ucb A R l u σ2 δ a (n + 1) ω = ucb' n (IsAlgEnvSeq.hist A R n ω) l u σ2 δ a := by + have hp : pullCount A a (n + 1) ω = pullCount' n (IsAlgEnvSeq.hist A R n ω) a := + pullCount_add_one_eq_pullCount' + have he : empMean A R a (n + 1) ω = empMean' n (IsAlgEnvSeq.hist A R n ω) a := + empMean_add_one_eq_empMean' + rw [ucb, ucb', hp, he] + +/-- Helper for `sum_ucb_sub_mean_le`. -/ +private lemma sum_sqrt_le {ι : Type*} {c : ι → ℝ} (s : Finset ι) (hc : ∀ i, 0 ≤ c i) : + ∑ i ∈ s, √(c i) ≤ √(#s * ∑ i ∈ s, c i) := by + have h := Real.sum_sqrt_mul_sqrt_le s hc (fun _ => zero_le_one) + simp only [Real.sqrt_one, mul_one, sum_const, nsmul_eq_mul] at h + rwa [Real.sqrt_mul (by positivity), mul_comm] + +/-- Helper for `sum_ucb_sub_mean_le`. -/ +private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 / √k ≤ 2 * √n - 1 := by + induction n with + | zero => simp at h + | succ n ih => + rw [sum_range_succ] + by_cases hn : n = 0 + · rw [hn] + simp + norm_num + · have hi := ih (Nat.pos_of_ne_zero hn) + suffices 1 / √↑(n + 1) ≤ 2 * (√↑(n + 1) - √n) by linarith + push_cast + field_simp + have : √(n + 1) * √(n + 1) = (n + 1) := Real.mul_self_sqrt (by positivity) + have : √n * √n = n := Real.mul_self_sqrt (by positivity) + nlinarith + +lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, μ a ∈ Set.Icc l u) (hi : l ≤ u) + (hc : ∀ s < n, pullCount A (A s ω) s ω ≠ 0 → empMean A R (A s ω) s ω - μ (A s ω) + < √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω))) : + ∑ s ∈ range n, (ucb A R l u σ2 δ (A s ω) s ω - μ (A s ω)) + ≤ (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by + let S₀ := {s ∈ range n | pullCount A (A s ω) s ω = 0} + let S₁ := {s ∈ range n | pullCount A (A s ω) s ω ≠ 0} + have hu : S₀ ∪ S₁ = range n := filter_union_filter_not_eq _ _ + have hd : Disjoint S₀ S₁ := disjoint_filter_filter_not _ _ _ + rw [← hu, sum_union hd] + gcongr + · calc ∑ s ∈ S₀, (ucb A R l u σ2 δ (A s ω) s ω - μ (A s ω)) + ≤ ∑ s ∈ S₀, (u - l) := + have (s : ℕ) : ucb A R l u σ2 δ (A s ω) s ω ∈ Set.Icc l u := ucb_mem_Icc hi + sum_le_sum (by grind) + _ = ∑ s ∈ range n, if pullCount A (A s ω) s ω = 0 then (u - l) else 0 := by + rw [sum_filter] + _ = ∑ a, ∑ j ∈ range (pullCount A a n ω), if j = 0 then (u - l) else 0 := + sum_comp_pullCount (fun j => if j = 0 then (u - l) else 0) n ω + _ ≤ ∑ a, (u - l) := by + gcongr + rw [sum_ite_eq'] + grind + _ = (u - l) * K := by + rw [Fin.sum_const, nsmul_eq_mul, mul_comm] + · calc ∑ s ∈ S₁, (ucb A R l u σ2 δ (A s ω) s ω - μ (A s ω)) + ≤ ∑ s ∈ S₁, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) := by + gcongr with s hs + unfold ucb + have : 0 ≤ √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) := by positivity + grind + _ ≤ ∑ s ∈ range n, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) := + sum_le_sum_of_subset_of_nonneg (filter_subset _ _) (fun _ _ _ => by positivity) + _ = 2 * √(2 * σ2 * Real.log (1 / δ)) * ∑ s ∈ range n, (1 / √(pullCount A (A s ω) s ω)) := by + rw [mul_sum] + congr with s + rw [Real.sqrt_div' _ (by positivity)] + ring + _ = 2 * √(2 * σ2 * Real.log (1 / δ)) * + ∑ a, ∑ j ∈ range (pullCount A a n ω), (1 / √j) := by + rw [sum_comp_pullCount (fun j => 1 / √j)] + _ ≤ 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * ∑ a, √(pullCount A a n ω)) := by -- loose + rw [mul_sum _ _ 2] + gcongr with a + by_cases ha : pullCount A a n ω = 0 + · simp [ha] + · have hi := sum_inv_sqrt_le (Nat.pos_of_ne_zero ha) + rw [sum_range_succ] at hi + have : 0 ≤ 1 / √(pullCount A a n ω) := by positivity + linarith + _ ≤ 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * √(K * ∑ a, (pullCount A a n ω))) := by + gcongr + have h := sum_sqrt_le Finset.univ (fun a => Nat.cast_nonneg (pullCount A a n ω)) + rw [Finset.card_fin] at h + exact_mod_cast h + _ = 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * √(K * n)) := by + congr + exact sum_pullCount (ω := ω) + _ = 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by + ring_nf + rw [← Real.sqrt_mul' _ (by positivity)] + ring_nf + +variable [Nonempty (Fin K)] +variable [MeasurableSpace Ω] +variable {𝓔 : Type*} [MeasurableSpace 𝓔] +variable {E : Ω → 𝓔} +variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] +variable {P : Measure Ω} [IsProbabilityMeasure P] + +lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algorithm (Fin K) ℝ} + (h : IsBayesAlgEnvSeq Q κ alg E A R P) + (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) + (hδ : 0 < δ) (n : ℕ) : + P[fun ω ↦ ∑ t ∈ range n, + (actionMean κ E (bestAction κ E ω) ω - ucb A R l u σ2 δ (bestAction κ E ω) t ω)] ≤ + (u - l) * (n - 1) * n * δ := by + by_cases hn : n = 0 + · simp [hn] + let F := {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ + empMean A R (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω ≤ + -√(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω))} + have := h.measurable_action + have := h.measurable_param + have := h.measurable_feedback + have hF : MeasurableSet F := by measurability + have : Integrable (fun ω ↦ actionMean κ E (bestAction κ E ω) ω) P := + IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp (by fun_prop) (by fun_prop) hm + calc + _ ≤ ∫ ω in F, ∑ t ∈ range n, + (actionMean κ E (bestAction κ E ω) ω - ucb A R l u σ2 δ (bestAction κ E ω) t ω) ∂P := by + rw [← integral_add_compl hF (by fun_prop)] + apply add_le_of_nonpos_right + apply setIntegral_nonpos hF.compl + intro ω hω + apply sum_nonpos + intro t ht + rw [Set.mem_compl_iff, Set.mem_setOf_eq] at hω + push Not at hω + grind [hω t (mem_range.mp ht), ucb, actionMean] + _ ≤ ∫ ω in F, ∑ t ∈ range n, (u - l) ∂P := by + apply setIntegral_mono_on (Integrable.integrableOn (by fun_prop)) + (Integrable.integrableOn (by fun_prop)) hF + intro ω hω + apply sum_le_sum + intro t ht + grind [actionMean, ucb] + _ = P.real F * (n * (u - l)) := by + simp_rw [sum_const, card_range, nsmul_eq_mul, setIntegral_const, smul_eq_mul] + _ ≤ ((n - 1) * δ) * (n * (u - l)) := by + gcongr + have : (1 : ℝ) ≤ n := by simp [Nat.one_le_iff_ne_zero, hn] + apply ENNReal.toReal_le_of_le_ofReal (by nlinarith) + exact h.prob_empMean_bestAction_sub_actionMean_le_le hσ2 hs hδ n + _ = _ := by + ring + +lemma integral_sum_range_ucb_action_sub_actionMean_action_le {alg : Algorithm (Fin K) ℝ} + (h : IsBayesAlgEnvSeq Q κ alg E A R P) + (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) + (hδ : 0 < δ) (n : ℕ) : + P[fun ω ↦ ∑ t ∈ range n, (ucb A R l u σ2 δ (A t ω) t ω - actionMean κ E (A t ω) ω)] ≤ + (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + (u - l) * K * (n - 1) * n * δ := by + by_cases hn : n = 0 + · simp [hn, hlu, mul_nonneg] + let F := {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ + √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ empMean A R a t ω - actionMean κ E a ω} + have := h.measurable_action + have := h.measurable_param + have := h.measurable_feedback + have hF : MeasurableSet F := by measurability + have : ∀ t, Integrable (fun ω ↦ actionMean κ E (A t ω) ω) P := + fun t ↦ IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp (by fun_prop) (by fun_prop) hm + calc + _ ≤ (∫ ω in F, ∑ t ∈ range n, (u - l) ∂P) + + ∫ ω in Fᶜ, (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) ∂P := by + rw [← integral_add_compl hF (by fun_prop)] + apply add_le_add + · apply setIntegral_mono_on (Integrable.integrableOn (by fun_prop)) + (Integrable.integrableOn (by fun_prop)) hF + intro ω hω + apply sum_le_sum + intro t ht + grind [ucb, actionMean] + · apply setIntegral_mono_on (Integrable.integrableOn (by fun_prop)) + (Integrable.integrableOn (by fun_prop)) hF.compl + intro ω hω + rw [Set.mem_compl_iff, Set.mem_setOf_eq] at hω + push Not at hω + exact sum_ucb_sub_mean_le (fun a ↦ (κ (E ω, a))[id]) (hm (E ω)) hlu + (fun t ht hpc ↦ hω t ht (A t ω) hpc) + _ = P.real F * (n * (u - l)) + + P.real Fᶜ * ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) := by + simp_rw [sum_const, card_range, nsmul_eq_mul, setIntegral_const, smul_eq_mul] + _ ≤ (K * (n - 1) * δ) * (n * (u - l)) + + 1 * ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) := by + have : 0 ≤ u - l := sub_nonneg.2 hlu + gcongr + · have : (0 : ℝ) ≤ n - 1 := by simp [Nat.one_le_iff_ne_zero, hn] + apply ENNReal.toReal_le_of_le_ofReal (by positivity) + exact h.prob_empMean_sub_actionMean_ge_le hσ2 hs hδ n + · exact measureReal_le_one + _ = _ := by + ring + +end ClippedUCB + +namespace TS + +open ClippedUCB + +variable {K : ℕ} [Nonempty (Fin K)] +variable {l u σ2 δ : ℝ} +variable {Ω : Type*} [MeasurableSpace Ω] +variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] +variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] +variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} +variable {P : Measure Ω} [IsProbabilityMeasure P] + +lemma integral_ucb_action_eq_integral_ucb_bestAction (hK : 0 < K) + (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) (n : ℕ) : + P[fun ω ↦ ucb A R l u σ2 δ (A n ω) n ω] = + P[fun ω ↦ ucb A R l u σ2 δ (bestAction κ E ω) n ω] := by + have := h.measurable_action + have := h.measurable_param + have := h.measurable_feedback + by_cases hn : n = 0 + · simp [hn] + obtain ⟨n, rfl⟩ := Nat.exists_eq_succ_of_ne_zero hn + let uc (ha : (Iic n → Fin K × ℝ) × Fin K) := ucb' n ha.1 l u σ2 δ ha.2 + calc + _ = P[fun ω ↦ uc (IsAlgEnvSeq.hist A R n ω, A (n + 1) ω)] := by + simp_rw [uc, ucb_succ_eq_ucb'] + _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R n ω, A (n + 1) ω)) := by + rw [← integral_map (by fun_prop) (by fun_prop)] + _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R n ω, bestAction κ E ω)) := by + rw [← compProd_map_condDistrib (by fun_prop), ← compProd_map_condDistrib (by fun_prop), + Measure.compProd_congr (hasCondDistrib_action hK h n).condDistrib_eq] + _ = P[fun ω ↦ ucb A R l u σ2 δ (bestAction κ E ω) (n + 1) ω] := by + rw [integral_map (by fun_prop) (by fun_prop)] + simp_rw [uc, ucb_succ_eq_ucb'] + +lemma integral_regret_eq_add (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) + (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) : + P[IsBayesAlgEnvSeq.regret κ E A n] = + P[fun ω ↦ ∑ t ∈ range n, + (actionMean κ E (bestAction κ E ω) ω - ucb A R l u σ2 δ (bestAction κ E ω) t ω)] + + P[fun ω ↦ ∑ t ∈ range n, + (ucb A R l u σ2 δ (A t ω) t ω - actionMean κ E (A t ω) ω)] := by + have hua (t : ℕ) : Integrable (fun ω ↦ ucb A R l u σ2 δ (A t ω) t ω) P := + integrable_uncurry_ucb_comp h.measurable_action h.measurable_feedback (h.measurable_action t) + measurable_const + have hub (t : ℕ) : Integrable (fun ω ↦ ucb A R l u σ2 δ (bestAction κ E ω) t ω) P := + integrable_uncurry_ucb_comp h.measurable_action h.measurable_feedback + (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_param) measurable_const + have haa (t : ℕ) : Integrable (fun ω ↦ actionMean κ E (A t ω) ω) P := + IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_param + (h.measurable_action t) hm + have hab : Integrable (fun ω ↦ actionMean κ E (bestAction κ E ω) ω) P := + IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_param + (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_param) hm + calc + _ = (∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω ∂P) - + ∑ t ∈ range n, ∫ ω, actionMean κ E (A t ω) ω ∂P := by + simp_rw [IsBayesAlgEnvSeq.regret_eq_sum_gap, IsBayesAlgEnvSeq.gap_eq_sub] + rw [integral_finsetSum _ (by fun_prop), ← Finset.sum_sub_distrib] + simp_rw [integral_sub hab (haa _)] + _ = ((∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω ∂P) - + ∑ t ∈ range n, ∫ ω, ucb A R l u σ2 δ (bestAction κ E ω) t ω ∂P) + + ((∑ t ∈ range n, ∫ ω, ucb A R l u σ2 δ (A t ω) t ω ∂P) - + ∑ t ∈ range n, ∫ ω, actionMean κ E (A t ω) ω ∂P) := by + simp [integral_ucb_action_eq_integral_ucb_bestAction hK h] + _ = (∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω - + ucb A R l u σ2 δ (bestAction κ E ω) t ω ∂P) + + ∑ t ∈ range n, ∫ ω, ucb A R l u σ2 δ (A t ω) t ω - + actionMean κ E (A t ω) ω ∂P := by + rw [← Finset.sum_sub_distrib, ← Finset.sum_sub_distrib] + simp_rw [← integral_sub hab (hub _), ← integral_sub (hua _) (haa _)] + _ = _ := by + rw [← integral_finsetSum _ (by fun_prop), ← integral_finsetSum _ (by fun_prop)] + +lemma integral_regret_le (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) + (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) + (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (n : ℕ) : + P[IsBayesAlgEnvSeq.regret κ E A n] + ≤ (2 * K + 1) * (u - l) + 8 * √(σ2 * K * n * Real.log n) := by + by_cases hn : n = 0 + · simp [hn, IsBayesAlgEnvSeq.regret, Bandits.regret] + nlinarith + have hδ : (0 : ℝ) < 1 / n ^ 2 := by positivity + calc P[IsBayesAlgEnvSeq.regret κ E A n] + = _ := + integral_regret_eq_add hK h hm n + _ ≤ _ := + add_le_add + (integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le h hlu hm hσ2 hs hδ n) + (integral_sum_range_ucb_action_sub_actionMean_action_le h hlu hm hσ2 hs hδ n) + _ = K * (u - l) + (K + 1) * (u - l) * ((n - 1) / n) + + 4 * √((2 : ℝ) ^ 2 * (σ2 * K * n * Real.log n)) := by + field_simp + rw [Real.log_pow] + ring_nf + _ = K * (u - l) + (K + 1) * (u - l) * ((n - 1) / n) + 8 * √(σ2 * K * n * Real.log n) := by + rw [Real.sqrt_mul (by positivity), Real.sqrt_sq (by norm_num)] + ring + _ ≤ K * (u - l) + (K + 1) * (u - l) * 1 + 8 * √(σ2 * K * n * Real.log n) := by -- loose + have : 0 ≤ u - l := sub_nonneg.2 hlu + gcongr + rw [div_le_one (by positivity)] + linarith + _ = _ := by + ring + +end TS + +end Bandits diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index 30fcbc2a..de67b3ea 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -5,7 +5,7 @@ Authors: Paulo Rauber -/ module -public import LeanMachineLearning.Online.Bandit.SumRewards +public import LeanMachineLearning.Online.Bandit.BayesRegret public import LeanMachineLearning.SequentialLearning.AlgorithmDensityBayes public import LeanMachineLearning.SequentialLearning.Algorithms.Uniform @@ -15,24 +15,12 @@ public import LeanMachineLearning.SequentialLearning.Algorithms.Uniform This file defines the Thompson sampling algorithm. This algorithm samples an action according to its probability of being optimal under the posterior over environments given the history so far. -We also provide a Bayesian regret upper bound (`integral_regret_le`) for this algorithm under the -assumption (among others) that it has the correct prior over environments. - -The Bayesian regret upper bound relies on a clipped upper confidence bound whose definition -and properties are also given in this file. - ## Main definitions * `tsAlgorithm hK Q κ`: a Thompson sampling algorithm with actions in `Fin K` given `hK : 0 < K`, a prior distribution over parameters `Q : Measure 𝓔`, and a Markov kernel `κ : Kernel (𝓔 × Fin K) ℝ`. This kernel defines how a parameter `e : 𝓔` gives rise to a stationary environment: `stationaryEnv (κ.sectR e) : Environment (Fin K) ℝ`. -* `ucb A R l u σ2 δ a n` : clipped upper confidence bound used in the regret analysis of Thompson - sampling for a sequence of actions `A : ℕ → Ω → Fin K`, rewards `R : ℕ → Ω → ℝ`, reward lower - bound `l : ℝ`, reward upper bound `u : ℝ`, sub-Gaussian variance proxy `σ2 : ℝ`, confidence - parameter `δ : ℝ`, action `a : Fin K`, and time `n : ℕ`. -* `ucb' n h l u σ2 δ a`: clipped upper confidence bound for action `a : Fin K` at time `n : ℕ` given - the history `h : Iic n → Fin K × ℝ` (rather than the entire sequences of actions and rewards). ## Main results @@ -40,19 +28,12 @@ and properties are also given in this file. the conditional distribution of the next action given the history so far is equal to the conditional distribution of the best action given the history so far. -* `integral_regret_le`: if Thompson sampling has the correct prior over environments and every - environment has `K` actions, each of which has a corresponding reward between `l` and `u` that - is sub-Gaussian with variance proxy `σ2` after its mean is subtracted, then the Bayesian regret at - time `n` is at most `(2 * K + 1) * (u - l) + 8 * √(σ2 * K * n * Real.log n)`. - -/ @[expose] public section open MeasureTheory ProbabilityTheory Finset Learning -open IsBayesAlgEnvSeq (bestAction actionMean) - -open scoped NNReal +open IsBayesAlgEnvSeq (bestAction) namespace Bandits @@ -94,8 +75,6 @@ def tsAlgorithm (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : K end Algorithm -namespace TS - variable {K : ℕ} [Nonempty (Fin K)] variable {Ω : Type*} [MeasurableSpace Ω] variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] @@ -103,7 +82,7 @@ variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] variable {P : Measure Ω} [IsProbabilityMeasure P] -lemma hasCondDistrib_action (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) +lemma TS.hasCondDistrib_action (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) (n : ℕ) : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R n) (condDistrib (bestAction κ E) (IsAlgEnvSeq.hist A R n) P) P where aemeasurable_fst := (h.measurable_action (n + 1)).aemeasurable @@ -125,380 +104,4 @@ lemma hasCondDistrib_action (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm condDistrib (bestAction κ E) (IsAlgEnvSeq.hist A R n) P := (condDistrib_comp (IsAlgEnvSeq.hist A R n) h.measurable_param.aemeasurable hm).symm -end TS - -namespace ClippedUCB - -variable {K : ℕ} {l u σ2 δ : ℝ} -variable {Ω : Type*} {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} - -/-- Clipped upper confidence bound used in the regret analysis of Thompson sampling. -/ -noncomputable -def ucb (A : ℕ → Ω → Fin K) (R : ℕ → Ω → ℝ) (l u σ2 δ : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ := - if pullCount A a n ω = 0 then u - else max l (min u (empMean A R a n ω + √(2 * σ2 * Real.log (1 / δ) / (pullCount A a n ω)))) - -@[simp] -lemma ucb_zero {a : Fin K} {ω : Ω} : ucb A R l u σ2 δ a 0 ω = u := by - simp [ucb] - -lemma ucb_mem_Icc (h : l ≤ u) {a : Fin K} {n : ℕ} {ω : Ω} : - ucb A R l u σ2 δ a n ω ∈ Set.Icc l u := by - unfold ucb - grind - -@[fun_prop] -lemma measurable_ucb [MeasurableSpace Ω] {a : Fin K} {n : ℕ} (hA : ∀ t, Measurable (A t)) - (hR : ∀ t, Measurable (R t)) : Measurable (ucb A R l u σ2 δ a n) := - Measurable.ite (by measurability) (by fun_prop) (by fun_prop) - -@[fun_prop] -lemma measurable_uncurry_ucb_comp [MeasurableSpace Ω] (hA : ∀ t, Measurable (A t)) - (hR : ∀ t, Measurable (R t)) {f : Ω → Fin K} (hf : Measurable f) {g : Ω → ℕ} - (hg : Measurable g) : Measurable (fun ω ↦ ucb A R l u σ2 δ (f ω) (g ω) ω) := by - change Measurable ((fun aω ↦ ucb A R l u σ2 δ aω.1 (g aω.2) aω.2) ∘ fun ω ↦ (f ω, ω)) - apply Measurable.comp _ (by fun_prop) - apply measurable_from_prod_countable_right - intro a - change Measurable ((fun tω ↦ ucb A R l u σ2 δ a tω.1 tω.2) ∘ fun ω ↦ (g ω, ω)) - apply Measurable.comp _ (by fun_prop) - exact measurable_from_prod_countable_right (fun _ ↦ measurable_ucb hA hR) - -@[fun_prop] -lemma integrable_uncurry_ucb_comp [MeasurableSpace Ω] (hA : ∀ t, Measurable (A t)) - (hR : ∀ t, Measurable (R t)) {f : Ω → Fin K} (hf : Measurable f) {g : Ω → ℕ} - (hg : Measurable g) {P : Measure Ω} [IsFiniteMeasure P] : - Integrable (fun ω ↦ ucb A R l u σ2 δ (f ω) (g ω) ω) P := by - refine ⟨(measurable_uncurry_ucb_comp hA hR hf hg).aestronglyMeasurable, ?_⟩ - apply HasFiniteIntegral.of_bounded (C := max |l| |u|) - filter_upwards with ω - rw [Real.norm_eq_abs] - unfold ucb - grind - -/-- Clipped upper confidence bound (history-based version). -/ -noncomputable -def ucb' (n : ℕ) (h : Iic n → Fin K × ℝ) (l u σ2 δ : ℝ) (a : Fin K) : ℝ := - if pullCount' n h a = 0 then u - else max l (min u (empMean' n h a + √(2 * σ2 * Real.log (1 / δ) / (pullCount' n h a)))) - -@[fun_prop] -lemma measurable_uncurry_ucb' {n : ℕ} : - Measurable (fun p : (Iic n → Fin K × ℝ) × Fin K ↦ ucb' n p.1 l u σ2 δ p.2) := - Measurable.ite (by measurability) (by fun_prop) (by fun_prop) - -lemma ucb_succ_eq_ucb' {a : Fin K} {n : ℕ} {ω : Ω} : - ucb A R l u σ2 δ a (n + 1) ω = ucb' n (IsAlgEnvSeq.hist A R n ω) l u σ2 δ a := by - have hp : pullCount A a (n + 1) ω = pullCount' n (IsAlgEnvSeq.hist A R n ω) a := - pullCount_add_one_eq_pullCount' - have he : empMean A R a (n + 1) ω = empMean' n (IsAlgEnvSeq.hist A R n ω) a := - empMean_add_one_eq_empMean' - rw [ucb, ucb', hp, he] - -/-- Helper for `sum_ucb_sub_mean_le`. -/ -private lemma sum_sqrt_le {ι : Type*} {c : ι → ℝ} (s : Finset ι) (hc : ∀ i, 0 ≤ c i) : - ∑ i ∈ s, √(c i) ≤ √(#s * ∑ i ∈ s, c i) := by - have h := Real.sum_sqrt_mul_sqrt_le s hc (fun _ => zero_le_one) - simp only [Real.sqrt_one, mul_one, sum_const, nsmul_eq_mul] at h - rwa [Real.sqrt_mul (by positivity), mul_comm] - -/-- Helper for `sum_ucb_sub_mean_le`. -/ -private lemma sum_inv_sqrt_le {n : ℕ} (h : 0 < n) : ∑ k ∈ range (n + 1), 1 / √k ≤ 2 * √n - 1 := by - induction n with - | zero => simp at h - | succ n ih => - rw [sum_range_succ] - by_cases hn : n = 0 - · rw [hn] - simp - norm_num - · have hi := ih (Nat.pos_of_ne_zero hn) - suffices 1 / √↑(n + 1) ≤ 2 * (√↑(n + 1) - √n) by linarith - push_cast - field_simp - have : √(n + 1) * √(n + 1) = (n + 1) := Real.mul_self_sqrt (by positivity) - have : √n * √n = n := Real.mul_self_sqrt (by positivity) - nlinarith - -lemma sum_ucb_sub_mean_le {n : ℕ} {ω : Ω} (μ : Fin K → ℝ) (hμ : ∀ a, μ a ∈ Set.Icc l u) (hi : l ≤ u) - (hc : ∀ s < n, pullCount A (A s ω) s ω ≠ 0 → empMean A R (A s ω) s ω - μ (A s ω) - < √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω))) : - ∑ s ∈ range n, (ucb A R l u σ2 δ (A s ω) s ω - μ (A s ω)) - ≤ (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by - let S₀ := {s ∈ range n | pullCount A (A s ω) s ω = 0} - let S₁ := {s ∈ range n | pullCount A (A s ω) s ω ≠ 0} - have hu : S₀ ∪ S₁ = range n := filter_union_filter_not_eq _ _ - have hd : Disjoint S₀ S₁ := disjoint_filter_filter_not _ _ _ - rw [← hu, sum_union hd] - gcongr - · calc ∑ s ∈ S₀, (ucb A R l u σ2 δ (A s ω) s ω - μ (A s ω)) - ≤ ∑ s ∈ S₀, (u - l) := - have (s : ℕ) : ucb A R l u σ2 δ (A s ω) s ω ∈ Set.Icc l u := ucb_mem_Icc hi - sum_le_sum (by grind) - _ = ∑ s ∈ range n, if pullCount A (A s ω) s ω = 0 then (u - l) else 0 := by - rw [sum_filter] - _ = ∑ a, ∑ j ∈ range (pullCount A a n ω), if j = 0 then (u - l) else 0 := - sum_comp_pullCount (fun j => if j = 0 then (u - l) else 0) n ω - _ ≤ ∑ a, (u - l) := by - gcongr - rw [sum_ite_eq'] - grind - _ = (u - l) * K := by - rw [Fin.sum_const, nsmul_eq_mul, mul_comm] - · calc ∑ s ∈ S₁, (ucb A R l u σ2 δ (A s ω) s ω - μ (A s ω)) - ≤ ∑ s ∈ S₁, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) := by - gcongr with s hs - unfold ucb - have : 0 ≤ √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) := by positivity - grind - _ ≤ ∑ s ∈ range n, 2 * √(2 * σ2 * Real.log (1 / δ) / (pullCount A (A s ω) s ω)) := - sum_le_sum_of_subset_of_nonneg (filter_subset _ _) (fun _ _ _ => by positivity) - _ = 2 * √(2 * σ2 * Real.log (1 / δ)) * ∑ s ∈ range n, (1 / √(pullCount A (A s ω) s ω)) := by - rw [mul_sum] - congr with s - rw [Real.sqrt_div' _ (by positivity)] - ring - _ = 2 * √(2 * σ2 * Real.log (1 / δ)) * - ∑ a, ∑ j ∈ range (pullCount A a n ω), (1 / √j) := by - rw [sum_comp_pullCount (fun j => 1 / √j)] - _ ≤ 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * ∑ a, √(pullCount A a n ω)) := by -- loose - rw [mul_sum _ _ 2] - gcongr with a - by_cases ha : pullCount A a n ω = 0 - · simp [ha] - · have hi := sum_inv_sqrt_le (Nat.pos_of_ne_zero ha) - rw [sum_range_succ] at hi - have : 0 ≤ 1 / √(pullCount A a n ω) := by positivity - linarith - _ ≤ 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * √(K * ∑ a, (pullCount A a n ω))) := by - gcongr - have h := sum_sqrt_le Finset.univ (fun a => Nat.cast_nonneg (pullCount A a n ω)) - rw [Finset.card_fin] at h - exact_mod_cast h - _ = 2 * √(2 * σ2 * Real.log (1 / δ)) * (2 * √(K * n)) := by - congr - exact sum_pullCount (ω := ω) - _ = 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) := by - ring_nf - rw [← Real.sqrt_mul' _ (by positivity)] - ring_nf - -variable [Nonempty (Fin K)] -variable [MeasurableSpace Ω] -variable {𝓔 : Type*} [MeasurableSpace 𝓔] -variable {E : Ω → 𝓔} -variable {Q : Measure 𝓔} {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] -variable {P : Measure Ω} [IsProbabilityMeasure P] - -lemma integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le {alg : Algorithm (Fin K) ℝ} - (h : IsBayesAlgEnvSeq Q κ alg E A R P) - (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) - (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) - (hδ : 0 < δ) (n : ℕ) : - P[fun ω ↦ ∑ t ∈ range n, - (actionMean κ E (bestAction κ E ω) ω - ucb A R l u σ2 δ (bestAction κ E ω) t ω)] ≤ - (u - l) * (n - 1) * n * δ := by - by_cases hn : n = 0 - · simp [hn] - let F := {ω | ∃ t < n, pullCount A (bestAction κ E ω) t ω ≠ 0 ∧ - empMean A R (bestAction κ E ω) t ω - actionMean κ E (bestAction κ E ω) ω ≤ - -√(2 * σ2 * Real.log (1 / δ) / (pullCount A (bestAction κ E ω) t ω))} - have := h.measurable_action - have := h.measurable_param - have := h.measurable_feedback - have hF : MeasurableSet F := by measurability - have : Integrable (fun ω ↦ actionMean κ E (bestAction κ E ω) ω) P := - IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp (by fun_prop) (by fun_prop) hm - calc - _ ≤ ∫ ω in F, ∑ t ∈ range n, - (actionMean κ E (bestAction κ E ω) ω - ucb A R l u σ2 δ (bestAction κ E ω) t ω) ∂P := by - rw [← integral_add_compl hF (by fun_prop)] - apply add_le_of_nonpos_right - apply setIntegral_nonpos hF.compl - intro ω hω - apply sum_nonpos - intro t ht - rw [Set.mem_compl_iff, Set.mem_setOf_eq] at hω - push Not at hω - grind [hω t (mem_range.mp ht), ucb, actionMean] - _ ≤ ∫ ω in F, ∑ t ∈ range n, (u - l) ∂P := by - apply setIntegral_mono_on (Integrable.integrableOn (by fun_prop)) - (Integrable.integrableOn (by fun_prop)) hF - intro ω hω - apply sum_le_sum - intro t ht - grind [actionMean, ucb] - _ = P.real F * (n * (u - l)) := by - simp_rw [sum_const, card_range, nsmul_eq_mul, setIntegral_const, smul_eq_mul] - _ ≤ ((n - 1) * δ) * (n * (u - l)) := by - gcongr - have : (1 : ℝ) ≤ n := by simp [Nat.one_le_iff_ne_zero, hn] - apply ENNReal.toReal_le_of_le_ofReal (by nlinarith) - exact h.prob_empMean_bestAction_sub_actionMean_le_le hσ2 hs hδ n - _ = _ := by - ring - -lemma integral_sum_range_ucb_action_sub_actionMean_action_le {alg : Algorithm (Fin K) ℝ} - (h : IsBayesAlgEnvSeq Q κ alg E A R P) - (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) - (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) - (hδ : 0 < δ) (n : ℕ) : - P[fun ω ↦ ∑ t ∈ range n, (ucb A R l u σ2 δ (A t ω) t ω - actionMean κ E (A t ω) ω)] ≤ - (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) + (u - l) * K * (n - 1) * n * δ := by - by_cases hn : n = 0 - · simp [hn, hlu, mul_nonneg] - let F := {ω | ∃ t < n, ∃ a, pullCount A a t ω ≠ 0 ∧ - √(2 * σ2 * Real.log (1 / δ) / pullCount A a t ω) ≤ empMean A R a t ω - actionMean κ E a ω} - have := h.measurable_action - have := h.measurable_param - have := h.measurable_feedback - have hF : MeasurableSet F := by measurability - have : ∀ t, Integrable (fun ω ↦ actionMean κ E (A t ω) ω) P := - fun t ↦ IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp (by fun_prop) (by fun_prop) hm - calc - _ ≤ (∫ ω in F, ∑ t ∈ range n, (u - l) ∂P) + - ∫ ω in Fᶜ, (u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n) ∂P := by - rw [← integral_add_compl hF (by fun_prop)] - apply add_le_add - · apply setIntegral_mono_on (Integrable.integrableOn (by fun_prop)) - (Integrable.integrableOn (by fun_prop)) hF - intro ω hω - apply sum_le_sum - intro t ht - grind [ucb, actionMean] - · apply setIntegral_mono_on (Integrable.integrableOn (by fun_prop)) - (Integrable.integrableOn (by fun_prop)) hF.compl - intro ω hω - rw [Set.mem_compl_iff, Set.mem_setOf_eq] at hω - push Not at hω - exact sum_ucb_sub_mean_le (fun a ↦ (κ (E ω, a))[id]) (hm (E ω)) hlu - (fun t ht hpc ↦ hω t ht (A t ω) hpc) - _ = P.real F * (n * (u - l)) + - P.real Fᶜ * ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) := by - simp_rw [sum_const, card_range, nsmul_eq_mul, setIntegral_const, smul_eq_mul] - _ ≤ (K * (n - 1) * δ) * (n * (u - l)) + - 1 * ((u - l) * K + 4 * √(2 * σ2 * Real.log (1 / δ) * K * n)) := by - have : 0 ≤ u - l := sub_nonneg.2 hlu - gcongr - · have : (0 : ℝ) ≤ n - 1 := by simp [Nat.one_le_iff_ne_zero, hn] - apply ENNReal.toReal_le_of_le_ofReal (by positivity) - exact h.prob_empMean_sub_actionMean_ge_le hσ2 hs hδ n - · exact measureReal_le_one - _ = _ := by - ring - -end ClippedUCB - -namespace TS - -section IntegralRegret - -open ClippedUCB - -variable {K : ℕ} [Nonempty (Fin K)] -variable {l u σ2 δ : ℝ} -variable {Ω : Type*} [MeasurableSpace Ω] -variable {𝓔 : Type*} [MeasurableSpace 𝓔] [StandardBorelSpace 𝓔] [Nonempty 𝓔] -variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] -variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} -variable {P : Measure Ω} [IsProbabilityMeasure P] - -lemma integral_ucb_action_eq_integral_ucb_bestAction (hK : 0 < K) - (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) (n : ℕ) : - P[fun ω ↦ ucb A R l u σ2 δ (A n ω) n ω] = - P[fun ω ↦ ucb A R l u σ2 δ (bestAction κ E ω) n ω] := by - have := h.measurable_action - have := h.measurable_param - have := h.measurable_feedback - by_cases hn : n = 0 - · simp [hn] - obtain ⟨n, rfl⟩ := Nat.exists_eq_succ_of_ne_zero hn - let uc (ha : (Iic n → Fin K × ℝ) × Fin K) := ucb' n ha.1 l u σ2 δ ha.2 - calc - _ = P[fun ω ↦ uc (IsAlgEnvSeq.hist A R n ω, A (n + 1) ω)] := by - simp_rw [uc, ucb_succ_eq_ucb'] - _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R n ω, A (n + 1) ω)) := by - rw [← integral_map (by fun_prop) (by fun_prop)] - _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R n ω, bestAction κ E ω)) := by - rw [← compProd_map_condDistrib (by fun_prop), ← compProd_map_condDistrib (by fun_prop), - Measure.compProd_congr (hasCondDistrib_action hK h n).condDistrib_eq] - _ = P[fun ω ↦ ucb A R l u σ2 δ (bestAction κ E ω) (n + 1) ω] := by - rw [integral_map (by fun_prop) (by fun_prop)] - simp_rw [uc, ucb_succ_eq_ucb'] - -lemma integral_regret_eq_add (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) - (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (n : ℕ) : - P[IsBayesAlgEnvSeq.regret κ E A n] = - P[fun ω ↦ ∑ t ∈ range n, - (actionMean κ E (bestAction κ E ω) ω - ucb A R l u σ2 δ (bestAction κ E ω) t ω)] + - P[fun ω ↦ ∑ t ∈ range n, - (ucb A R l u σ2 δ (A t ω) t ω - actionMean κ E (A t ω) ω)] := by - have hua (t : ℕ) : Integrable (fun ω ↦ ucb A R l u σ2 δ (A t ω) t ω) P := - integrable_uncurry_ucb_comp h.measurable_action h.measurable_feedback (h.measurable_action t) - measurable_const - have hub (t : ℕ) : Integrable (fun ω ↦ ucb A R l u σ2 δ (bestAction κ E ω) t ω) P := - integrable_uncurry_ucb_comp h.measurable_action h.measurable_feedback - (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_param) measurable_const - have haa (t : ℕ) : Integrable (fun ω ↦ actionMean κ E (A t ω) ω) P := - IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_param - (h.measurable_action t) hm - have hab : Integrable (fun ω ↦ actionMean κ E (bestAction κ E ω) ω) P := - IsBayesAlgEnvSeq.integrable_uncurry_actionMean_comp h.measurable_param - (IsBayesAlgEnvSeq.measurable_bestAction h.measurable_param) hm - calc - _ = (∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω ∂P) - - ∑ t ∈ range n, ∫ ω, actionMean κ E (A t ω) ω ∂P := by - simp_rw [IsBayesAlgEnvSeq.regret_eq_sum_gap, IsBayesAlgEnvSeq.gap_eq_sub] - rw [integral_finsetSum _ (by fun_prop), ← Finset.sum_sub_distrib] - simp_rw [integral_sub hab (haa _)] - _ = ((∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω ∂P) - - ∑ t ∈ range n, ∫ ω, ucb A R l u σ2 δ (bestAction κ E ω) t ω ∂P) + - ((∑ t ∈ range n, ∫ ω, ucb A R l u σ2 δ (A t ω) t ω ∂P) - - ∑ t ∈ range n, ∫ ω, actionMean κ E (A t ω) ω ∂P) := by - simp [integral_ucb_action_eq_integral_ucb_bestAction hK h] - _ = (∑ t ∈ range n, ∫ ω, actionMean κ E (bestAction κ E ω) ω - - ucb A R l u σ2 δ (bestAction κ E ω) t ω ∂P) + - ∑ t ∈ range n, ∫ ω, ucb A R l u σ2 δ (A t ω) t ω - - actionMean κ E (A t ω) ω ∂P := by - rw [← Finset.sum_sub_distrib, ← Finset.sum_sub_distrib] - simp_rw [← integral_sub hab (hub _), ← integral_sub (hua _) (haa _)] - _ = _ := by - rw [← integral_finsetSum _ (by fun_prop), ← integral_finsetSum _ (by fun_prop)] - -lemma integral_regret_le (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) - (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) - (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (n : ℕ) : - P[IsBayesAlgEnvSeq.regret κ E A n] - ≤ (2 * K + 1) * (u - l) + 8 * √(σ2 * K * n * Real.log n) := by - by_cases hn : n = 0 - · simp [hn, IsBayesAlgEnvSeq.regret, Bandits.regret] - nlinarith - have hδ : (0 : ℝ) < 1 / n ^ 2 := by positivity - calc P[IsBayesAlgEnvSeq.regret κ E A n] - = _ := - integral_regret_eq_add hK h hm n - _ ≤ _ := - add_le_add - (integral_sum_range_actionMean_bestAction_sub_ucb_bestAction_le h hlu hm hσ2 hs hδ n) - (integral_sum_range_ucb_action_sub_actionMean_action_le h hlu hm hσ2 hs hδ n) - _ = K * (u - l) + (K + 1) * (u - l) * ((n - 1) / n) - + 4 * √((2 : ℝ) ^ 2 * (σ2 * K * n * Real.log n)) := by - field_simp - rw [Real.log_pow] - ring_nf - _ = K * (u - l) + (K + 1) * (u - l) * ((n - 1) / n) + 8 * √(σ2 * K * n * Real.log n) := by - rw [Real.sqrt_mul (by positivity), Real.sqrt_sq (by norm_num)] - ring - _ ≤ K * (u - l) + (K + 1) * (u - l) * 1 + 8 * √(σ2 * K * n * Real.log n) := by -- loose - have : 0 ≤ u - l := sub_nonneg.2 hlu - gcongr - rw [div_le_one (by positivity)] - linarith - _ = _ := by - ring - -end IntegralRegret - -end TS - end Bandits From 2739c4a6a8e9b2717e11a7bc852a73923141bf88 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 18 Jun 2026 10:27:50 +0100 Subject: [PATCH 154/155] Fix hist --- .../Algorithms/Regret/BayesRegretTS.lean | 12 +++++------ .../Online/Bandit/Algorithms/TS.lean | 20 +++++++++---------- 2 files changed, 16 insertions(+), 16 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/Regret/BayesRegretTS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/Regret/BayesRegretTS.lean index e0bc2e75..22e6c5ba 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/Regret/BayesRegretTS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/Regret/BayesRegretTS.lean @@ -103,10 +103,10 @@ lemma measurable_uncurry_ucb' {n : ℕ} : Measurable.ite (by measurability) (by fun_prop) (by fun_prop) lemma ucb_succ_eq_ucb' {a : Fin K} {n : ℕ} {ω : Ω} : - ucb A R l u σ2 δ a (n + 1) ω = ucb' n (IsAlgEnvSeq.hist A R n ω) l u σ2 δ a := by - have hp : pullCount A a (n + 1) ω = pullCount' n (IsAlgEnvSeq.hist A R n ω) a := + ucb A R l u σ2 δ a (n + 1) ω = ucb' n (history A R n ω) l u σ2 δ a := by + have hp : pullCount A a (n + 1) ω = pullCount' n (history A R n ω) a := pullCount_add_one_eq_pullCount' - have he : empMean A R a (n + 1) ω = empMean' n (IsAlgEnvSeq.hist A R n ω) a := + have he : empMean A R a (n + 1) ω = empMean' n (history A R n ω) a := empMean_add_one_eq_empMean' rw [ucb, ucb', hp, he] @@ -328,11 +328,11 @@ lemma integral_ucb_action_eq_integral_ucb_bestAction (hK : 0 < K) obtain ⟨n, rfl⟩ := Nat.exists_eq_succ_of_ne_zero hn let uc (ha : (Iic n → Fin K × ℝ) × Fin K) := ucb' n ha.1 l u σ2 δ ha.2 calc - _ = P[fun ω ↦ uc (IsAlgEnvSeq.hist A R n ω, A (n + 1) ω)] := by + _ = P[fun ω ↦ uc (history A R n ω, A (n + 1) ω)] := by simp_rw [uc, ucb_succ_eq_ucb'] - _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R n ω, A (n + 1) ω)) := by + _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (history A R n ω, A (n + 1) ω)) := by rw [← integral_map (by fun_prop) (by fun_prop)] - _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (IsAlgEnvSeq.hist A R n ω, bestAction κ E ω)) := by + _ = ∫ ha, uc ha ∂P.map (fun ω ↦ (history A R n ω, bestAction κ E ω)) := by rw [← compProd_map_condDistrib (by fun_prop), ← compProd_map_condDistrib (by fun_prop), Measure.compProd_congr (hasCondDistrib_action hK h n).condDistrib_eq] _ = P[fun ω ↦ ucb A R l u σ2 δ (bestAction κ E ω) (n + 1) ω] := by diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index de67b3ea..44e155ce 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -83,25 +83,25 @@ variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K variable {P : Measure Ω} [IsProbabilityMeasure P] lemma TS.hasCondDistrib_action (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) - (n : ℕ) : HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R n) - (condDistrib (bestAction κ E) (IsAlgEnvSeq.hist A R n) P) P where + (n : ℕ) : HasCondDistrib (A (n + 1)) (history A R n) + (condDistrib (bestAction κ E) (history A R n) P) P where aemeasurable_fst := (h.measurable_action (n + 1)).aemeasurable aemeasurable_snd := - (IsAlgEnvSeq.measurable_hist h.measurable_action h.measurable_feedback n).aemeasurable + (measurable_history h.measurable_action h.measurable_feedback n).aemeasurable condDistrib_eq := by have hm : Measurable (bestAction κ id) := by fun_prop calc - _ =ᵐ[P.map (IsAlgEnvSeq.hist A R n)] + _ =ᵐ[P.map (history A R n)] (IT.bayesTrajMeasurePosterior Q κ uniformAlgorithm n).map (bestAction κ id) := (h.hasCondDistrib_action' n).condDistrib_eq - _ =ᵐ[P.map (IsAlgEnvSeq.hist A R n)] - (condDistrib E (IsAlgEnvSeq.hist A R n) P).map (bestAction κ id) := by - filter_upwards [(h.hasCondDistrib_env_hist + _ =ᵐ[P.map (history A R n)] + (condDistrib E (history A R n) P).map (bestAction κ id) := by + filter_upwards [(h.hasCondDistrib_env_history (IT.isBayesAlgEnvSeq_bayesTrajMeasure Q κ uniformAlgorithm) absolutelyContinuous_uniformAlgorithm n).condDistrib_eq] with _ hc simp_rw [Kernel.map_apply _ hm, IT.bayesTrajMeasurePosterior, hc] - _ =ᵐ[P.map (IsAlgEnvSeq.hist A R n)] - condDistrib (bestAction κ E) (IsAlgEnvSeq.hist A R n) P := - (condDistrib_comp (IsAlgEnvSeq.hist A R n) h.measurable_param.aemeasurable hm).symm + _ =ᵐ[P.map (history A R n)] + condDistrib (bestAction κ E) (history A R n) P := + (condDistrib_comp (history A R n) h.measurable_param.aemeasurable hm).symm end Bandits From 8aca9c355204f900df581147dd7d649a8fe2d434 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Thu, 18 Jun 2026 12:22:09 +0100 Subject: [PATCH 155/155] Add documentation --- .../Bandit/Algorithms/Regret/BayesRegretTS.lean | 6 +++++- LeanMachineLearning/Online/Bandit/Algorithms/TS.lean | 11 ++++++++++- 2 files changed, 15 insertions(+), 2 deletions(-) diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/Regret/BayesRegretTS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/Regret/BayesRegretTS.lean index 22e6c5ba..444f290d 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/Regret/BayesRegretTS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/Regret/BayesRegretTS.lean @@ -378,7 +378,11 @@ lemma integral_regret_eq_add (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorith _ = _ := by rw [← integral_finsetSum _ (by fun_prop), ← integral_finsetSum _ (by fun_prop)] -lemma integral_regret_le (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) +/-- If Thompson sampling has the correct prior over environments and every environment has `K` +actions, each of which has a corresponding reward between `l` and `u` that is sub-Gaussian with +variance proxy `σ2` after its mean is subtracted, then the Bayesian regret at time `n` is at most +`(2 * K + 1) * (u - l) + 8 * √(σ2 * K * n * Real.log n)`. -/ +theorem integral_regret_le (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) (hlu : l ≤ u) (hm : ∀ e a, (κ (e, a))[id] ∈ (Set.Icc l u)) (hσ2 : 0 < σ2) (hs : ∀ e a, HasSubgaussianMGF (fun x ↦ x - (κ (e, a))[id]) ⟨σ2, hσ2.le⟩ (κ (e, a))) (n : ℕ) : P[IsBayesAlgEnvSeq.regret κ E A n] diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean index 44e155ce..1310275b 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/TS.lean @@ -66,7 +66,13 @@ instance {hK : 0 < K} {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel ( IsProbabilityMeasure (TS.initialPolicy hK Q κ) := Measure.isProbabilityMeasure_map (by fun_prop) -/-- The Thompson sampling algorithm. -/ +/-- The Thompson sampling algorithm with actions in `Fin K`, where `Q : Measure 𝓔` is a prior + distribution over parameters, and `κ : Kernel (𝓔 × Fin K) ℝ` is a Markov kernel that defines the + stationary environment `stationaryEnv (κ.sectR e)` that corresponds to a parameter `e : 𝓔`. + + At every time `n`, the Thompson sampling policy uses the posterior over the parameters given the + history up to time `n` to derive the probability of each action being optimal. The action for time + `n` is sampled according to these probabilities. -/ noncomputable def tsAlgorithm (hK : 0 < K) (Q : Measure 𝓔) [IsProbabilityMeasure Q] (κ : Kernel (𝓔 × Fin K) ℝ) [IsMarkovKernel κ] : Algorithm (Fin K) ℝ where @@ -82,6 +88,9 @@ variable {E : Ω → 𝓔} {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} variable {Q : Measure 𝓔} [IsProbabilityMeasure Q] {κ : Kernel (𝓔 × Fin K) ℝ} [IsMarkovKernel κ] variable {P : Measure Ω} [IsProbabilityMeasure P] +/-- If Thompson sampling has the correct prior over environments, then the conditional distribution +of the next action given the history so far is equal to the conditional distribution of the best +action given the history so far. -/ lemma TS.hasCondDistrib_action (hK : 0 < K) (h : IsBayesAlgEnvSeq Q κ (tsAlgorithm hK Q κ) E A R P) (n : ℕ) : HasCondDistrib (A (n + 1)) (history A R n) (condDistrib (bestAction κ E) (history A R n) P) P where