diff --git a/LeanBandits.lean b/LeanBandits.lean index aaea891f..c1a8f31c 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -1,17 +1,24 @@ import LeanBandits.Bandit.Bandit import LeanBandits.Bandit.Regret +import LeanBandits.Bandit.RewardByCountMeasure +import LeanBandits.Bandit.SumRewards +import LeanBandits.BanditAlgorithms.AuxSums import LeanBandits.BanditAlgorithms.ETC import LeanBandits.BanditAlgorithms.UCB import LeanBandits.ForMathlib.CondDistrib +import LeanBandits.ForMathlib.CondIndepFun +import LeanBandits.ForMathlib.HasCondDistrib import LeanBandits.ForMathlib.IndepFun import LeanBandits.ForMathlib.IndepInfinitePi +import LeanBandits.ForMathlib.KernelRepresentation import LeanBandits.ForMathlib.KernelSub import LeanBandits.ForMathlib.Measurable import LeanBandits.ForMathlib.MeasurableArgMax +import LeanBandits.ForMathlib.StandardBorel import LeanBandits.ForMathlib.SubGaussian import LeanBandits.ForMathlib.Traj -import LeanBandits.RewardByCountMeasure import LeanBandits.SequentialLearning.Algorithm import LeanBandits.SequentialLearning.Deterministic import LeanBandits.SequentialLearning.FiniteActions +import LeanBandits.SequentialLearning.IonescuTulceaSpace import LeanBandits.SequentialLearning.StationaryEnv diff --git a/LeanBandits/Bandit/Bandit.lean b/LeanBandits/Bandit/Bandit.lean index 5e047baa..8c81862f 100644 --- a/LeanBandits/Bandit/Bandit.lean +++ b/LeanBandits/Bandit/Bandit.lean @@ -3,10 +3,14 @@ Copyright (c) 2025 Rémy Degenne. All rights reserved. Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne, Paulo Rauber -/ +import LeanBandits.ForMathlib.CondIndepFun +import LeanBandits.ForMathlib.IndepFun import LeanBandits.ForMathlib.IndepInfinitePi +import LeanBandits.ForMathlib.KernelRepresentation +import LeanBandits.ForMathlib.StandardBorel import LeanBandits.SequentialLearning.Deterministic +import LeanBandits.SequentialLearning.FiniteActions import LeanBandits.SequentialLearning.StationaryEnv -import Mathlib.Probability.IdentDistrib /-! # Bandit @@ -24,7 +28,8 @@ section MeasureSpace namespace Bandit -/-- Kernel describing the distribution of the next arm-reward pair given the history up to `n`. -/ +/-- Kernel describing the distribution of the next action-reward pair given the history up to +time `n`. -/ noncomputable def stepKernel (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : Kernel (Iic n → α × R) (α × R) := @@ -41,13 +46,13 @@ lemma snd_stepKernel (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel (stepKernel alg ν n).snd = ν ∘ₖ alg.policy n := by rw [stepKernel, Learning.stepKernel, stationaryEnv_feedback, Kernel.snd_compProd_prodMkLeft] -/-- Measure on the sequence of arms pulled and rewards observed generated by the bandit. -/ +/-- Measure on the sequence of actions pulled and rewards observed generated by the bandit. -/ noncomputable def trajMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : Measure (ℕ → α × R) := Learning.trajMeasure alg (stationaryEnv ν) deriving IsProbabilityMeasure -/-- Measure of an infinite stream of rewards from each arm. -/ +/-- Measure of an infinite stream of rewards from each action. -/ noncomputable def streamMeasure (ν : Kernel α R) : Measure (ℕ → α → R) := Measure.infinitePi fun _ ↦ Measure.infinitePi ν @@ -56,8 +61,8 @@ instance (ν : Kernel α R) [IsMarkovKernel ν] : IsProbabilityMeasure (streamMe unfold streamMeasure infer_instance -/-- Joint distribution of the sequence of arm pulled and rewards, and a stream of independent -rewards from all arms. -/ +/-- Joint distribution of the sequence of action pulled and rewards, and a stream of independent +rewards from all actions. -/ noncomputable def measure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : Measure ((ℕ → α × R) × (ℕ → α → R)) := @@ -156,117 +161,1106 @@ lemma indepFun_eval_snd_measure (alg : Algorithm α R) (ν : Kernel α R) [IsMar end StreamMeasure -/-- `arm n` is the arm pulled at time `n`. This is a random variable on the measurable space -`ℕ → α × ℝ`. -/ -def arm (n : ℕ) (h : ℕ → α × R) : α := (h n).1 +namespace ArrayModel -/-- `reward n` is the reward at time `n`. This is a random variable on the measurable space -`ℕ → α × R`. -/ -def reward (n : ℕ) (h : ℕ → α × R) : R := (h n).2 +open unitInterval -/-- `hist n` is the history up to time `n`. This is a random variable on the measurable space -`ℕ → α × R`. -/ -def hist (n : ℕ) (h : ℕ → α × R) : Iic n → α × R := fun i ↦ h i +section ProbabilitySpace + +variable (α R) in +/-- Probability space for the array model of stochastic bandits. -/ +def probSpace : Type _ := (ℕ → I) × (ℕ → α → R) + +instance {α R : Type*} [MeasurableSpace R] : MeasurableSpace (probSpace α R) := + inferInstanceAs (MeasurableSpace ((ℕ → I) × (ℕ → α → R))) + +instance {α R : Type*} [Countable α] [MeasurableSpace R] [StandardBorelSpace R] : + StandardBorelSpace (probSpace α R) := + inferInstanceAs (StandardBorelSpace ((ℕ → I) × (ℕ → α → R))) + +/-- Probability measure for the array model of stochastic bandits. -/ +noncomputable +def arrayMeasure (ν : Kernel α R) : Measure (probSpace α R) := + (Measure.infinitePi fun _ ↦ volume).prod (Bandit.streamMeasure ν) + +instance (ν : Kernel α R) [IsMarkovKernel ν] : IsProbabilityMeasure (arrayMeasure ν) := + Measure.prod.instIsProbabilityMeasure _ _ + +variable [Nonempty α] [StandardBorelSpace α] + +/-- The initial action is the image of a uniform random variable by this function. -/ +noncomputable +def initAlgFunction (alg : Algorithm α R) : I → α := + (representation_measure alg.p0).choose + +lemma initAlgFunction_map (alg : Algorithm α R) : volume.map (initAlgFunction alg) = alg.p0 := + (representation_measure alg.p0).choose_spec.2 @[fun_prop] -lemma measurable_arm (n : ℕ) : Measurable (arm n (α := α) (R := R)) := measurable_action n +lemma measurable_initAlgFunction (alg : Algorithm α R) : + Measurable (initAlgFunction alg) := (representation_measure alg.p0).choose_spec.1 + +/-- The next action is the image of the history and a uniform random variable by this function. -/ +noncomputable +def algFunction (alg : Algorithm α R) (n : ℕ) : + (Iic n → α × R) → I → α := + (Kernel.representation (alg.policy n)).choose + +lemma algFunction_map (alg : Algorithm α R) (n : ℕ) (h : Iic n → α × R) : + volume.map (algFunction alg n h) = alg.policy n h := + (Kernel.representation (alg.policy n)).choose_spec.2 h @[fun_prop] -lemma measurable_arm_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ arm p.1 p.2) := - measurable_action_prod +lemma measurable_algFunction (alg : Algorithm α R) (n : ℕ) : + Measurable (Function.uncurry (algFunction alg n)) := + (Kernel.representation (alg.policy n)).choose_spec.1 + +end ProbabilitySpace + +variable [Nonempty α] [StandardBorelSpace α] + +section HistoryActionReward + +/-- History of actions and rewards up to time `n` in the array model. -/ +noncomputable +def hist [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) : (n : ℕ) → Iic n → α × R +| 0 => fun _ ↦ (initAlgFunction alg (ω.1 0), ω.2 0 (initAlgFunction alg (ω.1 0))) +| n + 1 => + let hn : Iic n → α × R := hist alg ω n + let a : α := algFunction alg n hn (ω.1 (n + 1)) + fun i ↦ if hin : i ≤ n then hn ⟨i, by simp [hin]⟩ else (a, ω.2 (pullCount' n hn a) a) + +@[simp] +lemma hist_zero [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) : + hist alg ω 0 = fun _ ↦ (initAlgFunction alg (ω.1 0), ω.2 0 (initAlgFunction alg (ω.1 0))) := + rfl + +lemma hist_add_one [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) (n : ℕ) : + let a : α := algFunction alg n (hist alg ω n) (ω.1 (n + 1)) + hist alg ω (n + 1) = + fun (i : Iic (n + 1)) ↦ if hin : i ≤ n then hist alg ω n ⟨i, by simp [hin]⟩ + else (a, ω.2 (pullCount' n (hist alg ω n) a) a) := rfl + +lemma hist_eq [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) (n : ℕ) : + hist alg ω n = fun i : Iic n ↦ hist alg ω i ⟨i.1, by simp⟩ := by + induction n with + | zero => + ext i : 1 + simp only [hist] + rw [Unique.eq_default i] + simp [coe_default_Iic_zero] + | succ n hn => + ext i : 1 + by_cases hin : i ≤ n + · rw [hist_add_one] + simp only [hin, ↓reduceDIte] + rw [funext_iff] at hn + simp_rw [hn] + · grind + +lemma hist_add_one_eq_IicSuccProd' [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) + (n : ℕ) : + let a : α := algFunction alg n (hist alg ω n) (ω.1 (n + 1)) + hist alg ω (n + 1) = + (MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n).symm + (hist alg ω n, (a, ω.2 (pullCount' n (hist alg ω n) a) a)) := by + intro a + rw [hist_add_one] + ext i : 1 + simp only [Kernel.symm_IicSuccProd, MeasurableEquiv.prodCongr, MeasurableEquiv.refl_toEquiv, + MeasurableEquiv.piSingleton, eq_rec_constant, MeasurableEquiv.IicProdIoc, + MeasurableEquiv.trans_apply, MeasurableEquiv.coe_mk, Equiv.prodCongr_apply, Equiv.coe_refl, + Equiv.coe_fn_mk, Prod.map_apply, id_eq] + rfl + +lemma measurable_action_add_one' [DecidableEq α] {alg : Algorithm α R} + (n : ℕ) (h : Measurable (hist alg · n)) : + Measurable (fun x ↦ algFunction alg n (hist alg x n) (x.1 (n + 1))) := by fun_prop + +lemma measurable_pullCount'_action_add_one [DecidableEq α] {alg : Algorithm α R} + (n : ℕ) (h_hist : Measurable (hist alg · n)) : + Measurable (fun x ↦ + pullCount' n (hist alg x n) (algFunction alg n (hist alg x n) (x.1 (n + 1)))) := by + have h_alg_meas : Measurable (fun x ↦ algFunction alg n (hist alg x n) (x.1 (n + 1))) := + measurable_action_add_one' n h_hist + exact (measurable_uncurry_pullCount' (α := α) n).comp (h_hist.prodMk h_alg_meas) @[fun_prop] -lemma measurable_reward (n : ℕ) : Measurable (reward n (α := α) (R := R)) := - Learning.measurable_reward n +lemma measurable_hist [DecidableEq α] [Countable α] (alg : Algorithm α R) (n : ℕ) : + Measurable (fun ω ↦ hist alg ω n) := by + induction n with + | zero => + simp_rw [hist_zero, measurable_pi_iff] + refine fun _ ↦ Measurable.prodMk (by fun_prop) ?_ + change Measurable ((fun x : α × ((ℕ → I) × (ℕ → α → R)) ↦ x.2.2 0 x.1) ∘ + (fun x : (ℕ → I) × (ℕ → α → R) ↦ (initAlgFunction alg (x.1 0), x))) + have : Measurable (fun x : α × ((ℕ → I) × (ℕ → α → R)) ↦ x.2.2 0 x.1) := + measurable_from_prod_countable_right fun p ↦ by simp only; fun_prop + exact Measurable.comp (by fun_prop) (Measurable.prodMk (by fun_prop) (by fun_prop)) + | succ n hn => + refine measurable_pi_iff.mpr fun i ↦ ?_ + by_cases hin : i ≤ n + · simp only [hist, hin, ↓reduceDIte] + rw [measurable_pi_iff] at hn + exact hn ⟨i.1, by simp [hin]⟩ + · simp only [hist, hin, ↓reduceDIte] + refine Measurable.prodMk (by fun_prop) ?_ + change Measurable ((fun (x : (ℕ → α → R) × ℕ × α) ↦ x.1 x.2.1 x.2.2) ∘ + (fun x ↦ (x.2, pullCount' n (hist alg x n) (algFunction alg n (hist alg x n) (x.1 (n + 1))), + (algFunction alg n (hist alg x n) (x.1 (n + 1)))))) + have h1 : Measurable (fun (x : (ℕ → α → R) × ℕ × α) ↦ x.1 x.2.1 x.2.2) := + measurable_from_prod_countable_left fun p : ℕ × α ↦ (by simp only; fun_prop) + refine Measurable.comp (by fun_prop) (Measurable.prodMk (by fun_prop) ?_) + refine Measurable.prodMk ?_ (by fun_prop) + exact measurable_pullCount'_action_add_one n hn + +/-- Action taken at time `n` in the array model. -/ +noncomputable +def action [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : α := + (hist alg ω n ⟨n, by simp⟩).1 + +lemma action_zero [DecidableEq α] (alg : Algorithm α R) : + action alg 0 = fun ω ↦ initAlgFunction alg (ω.1 0) := by + ext + simp [action, hist_zero] + +lemma action_add_one_eq [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : + action alg (n + 1) = fun ω ↦ algFunction alg n (hist alg ω n) (ω.1 (n + 1)) := by + ext ω + rw [action, hist_add_one] + simp only [add_le_iff_nonpos_right, nonpos_iff_eq_zero, one_ne_zero, ↓reduceDIte] @[fun_prop] -lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ reward p.1 p.2) := - Learning.measurable_reward_prod +lemma measurable_action [DecidableEq α] [Countable α] (alg : Algorithm α R) (n : ℕ) : + Measurable (action alg n) := by unfold action; fun_prop + +/-- Reward received at time `n` in the array model. -/ +noncomputable +def reward [DecidableEq α] (alg : Algorithm α R) (n : ℕ) (ω : probSpace α R) : R := + (hist alg ω n ⟨n, by simp⟩).2 + +lemma reward_zero [DecidableEq α] (alg : Algorithm α R) : + reward alg 0 = fun ω ↦ ω.2 0 (action alg 0 ω) := by + ext + simp [reward, hist_zero, action_zero] + +lemma reward_add_one [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : + reward alg (n + 1) = + fun ω ↦ ω.2 (pullCount' n (hist alg ω n) (action alg (n + 1) ω)) (action alg (n + 1) ω) := by + ext ω + simp [reward, hist_add_one, action_add_one_eq] + +lemma reward_eq [DecidableEq α] (alg : Algorithm α R) (n : ℕ) : + reward alg n = fun ω ↦ ω.2 (pullCount (action alg) (action alg n ω) n ω) (action alg n ω) := by + cases n with + | zero => ext; simp [reward_zero, action_zero] + | succ n => + ext ω + rw [reward, hist_add_one] + simp only [add_le_iff_nonpos_right, nonpos_iff_eq_zero, one_ne_zero, ↓reduceDIte] + rw [action_add_one_eq, pullCount_eq_pullCount' (R' := reward alg) (by simp)] + simp only [Nat.add_one_sub_one] + rw [hist_eq] + rfl @[fun_prop] -lemma measurable_hist (n : ℕ) : Measurable (hist n (α := α) (R := R)) := - Learning.measurable_hist n +lemma measurable_reward [DecidableEq α] [Countable α] (alg : Algorithm α R) (n : ℕ) : + Measurable (reward alg n) := by unfold reward; fun_prop + +lemma hist_add_one_eq_IicSuccProd [DecidableEq α] (alg : Algorithm α R) (ω : probSpace α R) + (n : ℕ) : + hist alg ω (n + 1) = + (MeasurableEquiv.IicSuccProd (fun _ ↦ α × R) n).symm + (hist alg ω n, (action alg (n + 1) ω, reward alg (n + 1) ω)) := by + rw [hist_add_one_eq_IicSuccProd', reward_add_one, action_add_one_eq] -lemma hist_eq_frestrictLe : - hist = Preorder.frestrictLe («π» := fun _ ↦ α × R) := by - ext n h i : 3 - simp [hist, Preorder.frestrictLe] +end HistoryActionReward -/-- Filtration of the bandit process. -/ -protected def filtration (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : - Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) := - MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R) +variable [DecidableEq α] + +section Congruence + +-- very useful to prove measurability +lemma hist_congr (alg : Algorithm α R) (n : ℕ) {ω ω' : probSpace α R} + (hω1 : ∀ i ≤ n, ω.1 i = ω'.1 i) + (hω2 : ∀ i a, i < pullCount (action alg) a (n + 1) ω → ω.2 i a = ω'.2 i a) : + hist alg ω n = hist alg ω' n := by + induction n with + | zero => + simp only [zero_add, pullCount_one] at hω2 + simp_rw [hist_zero] + ext i : 1 + simp only [le_refl, hω1, Prod.mk.injEq, true_and] + refine hω2 0 _ ?_ + simp [action, hω1] + | succ n hn => + simp_rw [hist_add_one_eq_IicSuccProd] + specialize hn fun i hin ↦ hω1 i (by grind) + have h_hist : hist alg ω n = hist alg ω' n := by + refine hn fun i a hi ↦ hω2 i a (hi.trans_le ?_) + exact pullCount_mono _ (by lia) _ + have h_action : action alg (n + 1) ω = action alg (n + 1) ω' := by + simp_rw [action_add_one_eq] + rw [h_hist, hω1 _ le_rfl] + congr 3 + simp only [reward_add_one, h_hist, h_action] + refine hω2 _ _ ?_ + rw [pullCount_add_one, h_action] + simp only [↓reduceIte] + rw [pullCount_eq_pullCount' (R' := reward alg) (by simp)] + simp only [Nat.add_one_sub_one] + rw [← h_hist, hist_eq] + change pullCount' n (fun i ↦ (action alg i ω, reward alg i ω)) (action alg (n + 1) ω') < + pullCount' n (fun i ↦ (action alg i ω, reward alg i ω)) (action alg (n + 1) ω') + 1 + grind + +lemma stepsUntil_congr_aux (alg : Algorithm α R) + (a : α) (m n : ℕ) {ω ω' : probSpace α R} + (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) + (hω2_eq : ∀ i, i + 1 ≤ m → ω.2 i a = ω'.2 i a) + (h_eq : action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m) : + action alg (n + 1) ω' = a ∧ pullCount (action alg) a (n + 1) ω' = m := by + obtain ⟨h_action, h_pc⟩ := h_eq + have h_hist := hist_congr alg n (ω := ω) (ω' := ω') (by grind) fun i b hi ↦ ?_ + swap + · rcases eq_or_ne b a with (rfl | hba) + · refine hω2_eq i ?_ + rw [h_pc] at hi + grind + · grind + constructor + · rw [← h_action, action_add_one_eq] + simp [h_hist, hω1] + · simp_rw [← h_pc, pullCount_eq_sum] + refine Finset.sum_congr rfl fun i hi ↦ ?_ + congr 2 + rw [hist_eq _ _ n, hist_eq _ _ n, funext_iff] at h_hist + unfold action + specialize h_hist ⟨i, by grind⟩ + simp only at h_hist + rw [h_hist] + +lemma stepsUntil_congr (alg : Algorithm α R) (a : α) (m n : ℕ) {ω ω' : probSpace α R} + (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) + (hω2_eq : ∀ i, i + 1 ≤ m → ω.2 i a = ω'.2 i a) : + (action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m) ↔ + (action alg (n + 1) ω' = a ∧ pullCount (action alg) a (n + 1) ω' = m) := + ⟨stepsUntil_congr_aux alg a m n hω1 hω2_ne hω2_eq, + stepsUntil_congr_aux alg a m n (by grind) (by grind) (by grind)⟩ + +lemma stepsUntil_indicator_congr (alg : Algorithm α R) (a : α) (m n : ℕ) {ω ω' : probSpace α R} + (hω1 : ∀ i, ω.1 i = ω'.1 i) (hω2_ne : ∀ i b, b ≠ a → ω.2 i b = ω'.2 i b) + (hω2_eq : ∀ i, i + 1 ≤ m → ω.2 i a = ω'.2 i a) : + {ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}.indicator (fun _ ↦ 1) + ω = + {ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}.indicator + (fun _ ↦ 1) ω' := by + simp only [Set.indicator_apply, Set.mem_setOf_eq] + simp_rw [stepsUntil_congr alg a m n hω1 hω2_ne hω2_eq] + +end Congruence section Laws -lemma hasLaw_step_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : - HasLaw (fun h : ℕ → α × R ↦ h 0) (alg.p0 ⊗ₘ ν) (Bandit.trajMeasure alg ν) := - Learning.hasLaw_step_zero alg (stationaryEnv ν) +variable [Countable α] -lemma hasLaw_arm_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : - HasLaw (arm 0) alg.p0 (Bandit.trajMeasure alg ν) := - Learning.hasLaw_action_zero alg (stationaryEnv ν) +lemma hasLaw_action_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : + HasLaw (action alg 0) alg.p0 (arrayMeasure ν) where + map_eq := by + calc (arrayMeasure ν).map (fun ω ↦ initAlgFunction alg (ω.1 0)) + _ = ((arrayMeasure ν).fst.map (Function.eval 0)).map (initAlgFunction alg) := by + rw [Measure.fst, Measure.map_map (by fun_prop) (by fun_prop), + Measure.map_map (by fun_prop) (by fun_prop)] + rfl + _ = (volume : Measure I).map (initAlgFunction alg) := by + simp only [arrayMeasure, Measure.fst_prod] + rw [(measurePreserving_eval_infinitePi (fun _ ↦ volume) 0).map_eq] + _ = alg.p0 := initAlgFunction_map alg -lemma condDistrib_arm_reward [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - condDistrib (fun h ↦ (arm (n + 1) h, reward (n + 1) h)) (hist n) (Bandit.trajMeasure alg ν) - =ᵐ[(Bandit.trajMeasure alg ν).map (hist n)] Bandit.stepKernel alg ν n := - Learning.condDistrib_step alg (stationaryEnv ν) n +omit [Nonempty α] [StandardBorelSpace α] [DecidableEq α] [Countable α] in +lemma indepFun_fst_snd (ν : Kernel α R) [IsMarkovKernel ν] : + IndepFun Prod.fst Prod.snd (arrayMeasure ν) := + indepFun_prod measurable_id measurable_id -lemma condDistrib_reward' [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - condDistrib (reward (n + 1)) (fun ω ↦ (hist n ω, arm (n + 1) ω)) (Bandit.trajMeasure alg ν) - =ᵐ[(Bandit.trajMeasure alg ν).map (fun ω ↦ (hist n ω, arm (n + 1) ω))] ν.prodMkLeft _ := - Learning.condDistrib_reward alg (stationaryEnv ν) n +omit [Nonempty α] [StandardBorelSpace α] [DecidableEq α] [Countable α] in +lemma indepFun_fst_zero_snd_zero_action (ν : Kernel α R) [IsMarkovKernel ν] (a : α) : + IndepFun (fun ω ↦ ω.1 0) (fun ω ↦ ω.2 0 a) (arrayMeasure ν) := + indepFun_prod (X := fun ω : ℕ → I ↦ ω 0) (Y := fun ω : ℕ → α → R ↦ ω 0 a) + (by fun_prop) (by fun_prop) -lemma condDistrib_reward [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - condDistrib (reward n) (arm n) (Bandit.trajMeasure alg ν) - =ᵐ[(Bandit.trajMeasure alg ν).map (arm n)] ν := - Learning.condDistrib_reward_stationaryEnv n +omit [Nonempty α] [StandardBorelSpace α] [DecidableEq α] [Countable α] in +lemma map_snd_apply_arrayMeasure {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) (a : α) : + (arrayMeasure ν).map (fun ω ↦ ω.2 n a) = ν a := by + calc (arrayMeasure ν).map (fun ω ↦ ω.2 n a) + _ = (arrayMeasure ν).snd.map (fun ω ↦ ω n a) := by + rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)] + rfl + _ = ν a := by + rw [arrayMeasure, Measure.snd_prod, Bandit.streamMeasure] + have : (fun ω ↦ ω n a) = (fun h : α → R ↦ h a) ∘ (fun ω : ℕ → α → R ↦ ω n) := rfl + rw [this, ← Measure.map_map (by fun_prop) (by fun_prop), Measure.infinitePi_map_eval, + Measure.infinitePi_map_eval] + +variable [StandardBorelSpace R] [Nonempty R] + +lemma hasCondDistrib_reward_zero (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : + HasCondDistrib (reward alg 0) (action alg 0) ν (arrayMeasure ν) where + condDistrib_eq := by + refine (condDistrib_ae_eq_cond (by fun_prop) (by fun_prop)).trans ?_ + rw [Filter.EventuallyEq, ae_iff_of_countable] + intro a ha + simp only [reward_zero] + calc ((arrayMeasure ν)[|action alg 0 ⁻¹' {a}]).map (fun ω ↦ ω.2 0 (action alg 0 ω)) + _ = ((arrayMeasure ν)[|action alg 0 ⁻¹' {a}]).map (fun ω ↦ ω.2 0 a) := by + refine Measure.map_congr + (ae_cond_of_forall_mem ((measurableSet_singleton _).preimage (by fun_prop)) ?_) + intro x hx + simp only [Set.mem_preimage, Set.mem_singleton_iff] at hx + simp [hx] + _ = ν a := by + rw [cond_of_indepFun] + · exact map_snd_apply_arrayMeasure 0 a + · have : (fun ω ↦ ω.1 0) ⟂ᵢ[arrayMeasure ν] fun ω ↦ ω.2 0 a := + indepFun_fst_zero_snd_zero_action ν a + rw [action_zero] + exact this.comp (φ := initAlgFunction alg) (by fun_prop) measurable_id + · fun_prop + · fun_prop + · simp + · rwa [Measure.map_apply (by fun_prop) (by simp)] at ha + +-- proved by Claude, then slightly golfed +omit [DecidableEq α] [Nonempty α] [StandardBorelSpace α] [Countable α] [StandardBorelSpace R] + [Nonempty R] in +lemma indepFun_fst_add_one_aux (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + (fun ω ↦ ω.1 (n + 1)) ⟂ᵢ[arrayMeasure ν] (fun ω ↦ (fun (i : Iic n) ↦ ω.1 i, ω.2)) := by + let μ₁ : Measure (ℕ → I) := Measure.infinitePi fun _ ↦ volume + let μ₂ : Measure (ℕ → α → R) := Bandit.streamMeasure ν + -- Coordinates of μ₁ are independent + have h_indep : iIndepFun (fun i (ω : ℕ → I) ↦ ω i) μ₁ := + iIndepFun_infinitePi (fun _ ↦ measurable_id) + have h_indep_n : IndepFun (fun ω ↦ ω (n + 1)) (fun ω ↦ fun i : Iic n ↦ ω i) μ₁ := by + have h := h_indep.indepFun_finset₀ {n + 1} (Iic n) (by simp) + (fun i ↦ (measurable_pi_apply i).aemeasurable) + convert h.comp (measurable_pi_apply ⟨n + 1, by simp⟩) measurable_id using 1 + rw [indepFun_iff_measure_inter_preimage_eq_mul] + intro s t hs ht + let X : (ℕ → I) × (ℕ → α → R) → I := fun ω ↦ ω.1 (n + 1) + let Y : (ℕ → I) × (ℕ → α → R) → (Iic n → I) × (ℕ → α → R) := fun ω ↦ (fun i ↦ ω.1 i, ω.2) + change (μ₁.prod μ₂) (X ⁻¹' s ∩ Y ⁻¹' t) = (μ₁.prod μ₂) (X ⁻¹' s) * (μ₁.prod μ₂) (Y ⁻¹' t) + -- Rewrite using Fubini + rw [Measure.prod_apply (hs.preimage (by fun_prop : Measurable X)), + Measure.prod_apply (ht.preimage (by fun_prop : Measurable Y)), + Measure.prod_apply ((hs.preimage (by fun_prop : Measurable X)).inter + (ht.preimage (by fun_prop : Measurable Y)))] + -- Compute fibers + have hX_fst ω₁ : μ₂ (Prod.mk ω₁ ⁻¹' (X ⁻¹' s)) = s.indicator 1 (ω₁ (n + 1)) := by + simp only [X, Set.preimage_preimage] + by_cases h : ω₁ (n + 1) ∈ s <;> simp [h] + have hY_fst ω₁ : μ₂ (Prod.mk ω₁ ⁻¹' (Y ⁻¹' t)) = μ₂ {y | ((fun i : Iic n ↦ ω₁ i), y) ∈ t} := rfl + have hXY ω₁ : μ₂ (Prod.mk ω₁ ⁻¹' (X ⁻¹' s ∩ Y ⁻¹' t)) = + s.indicator 1 (ω₁ (n + 1)) * μ₂ {y | ((fun i : Iic n ↦ ω₁ i), y) ∈ t} := by + simp only [X, Y, Set.preimage_inter, Set.preimage_preimage] + by_cases h : ω₁ (n + 1) ∈ s + · simp [h] + grind + · simp [h] + simp_rw [hY_fst, hX_fst, hXY] + -- Factor the integral using independence + let g : (Iic n → I) → ENNReal := fun x ↦ μ₂ {y | (x, y) ∈ t} + have hg_meas : Measurable g := measurable_measure_prodMk_left ht + have hf_meas : Measurable (fun ω₁ : ℕ → I ↦ s.indicator (1 : I → ENNReal) (ω₁ (n + 1))) := + (measurable_one.indicator hs).comp (measurable_pi_apply _) + have hindep_fg : IndepFun (fun ω₁ ↦ s.indicator (1 : I → ENNReal) (ω₁ (n + 1))) + (fun ω₁ ↦ g (fun i ↦ ω₁ i)) μ₁ := + h_indep_n.comp (measurable_one.indicator hs) hg_meas + have h_eq (ω₁ : ℕ → I) : μ₂ {y | ((fun i : Iic n ↦ ω₁ i), y) ∈ t} = g (fun i ↦ ω₁ i) := rfl + simp_rw [h_eq] + exact lintegral_mul_eq_lintegral_mul_lintegral_of_indepFun hf_meas (by fun_prop) hindep_fg + +omit [StandardBorelSpace R] [Nonempty R] in +lemma measurable_hist_todo (alg : Algorithm α R) (n : ℕ) : + Measurable[MeasurableSpace.comap (fun ω ↦ (fun (i : Iic n) ↦ ω.1 i, ω.2)) inferInstance] + (hist alg · n) := by + have h_eq : (hist alg · n) = + ((hist alg · n) ∘ (fun p ↦ (fun i : ℕ ↦ p.1 ⟨min i n, by grind⟩, p.2))) ∘ + (fun ω ↦ (fun (i : Iic n) ↦ ω.1 i, ω.2)) := by + ext ω : 1 + exact hist_congr alg n (by grind) (by simp) + rw [h_eq] + refine measurable_comp_comap _ (Measurable.comp (by fun_prop) ?_) + refine Measurable.prodMk ?_ (by fun_prop) + rw [measurable_pi_iff] + intro i + change Measurable ((fun p ↦ p ⟨min i n, by simp⟩) ∘ (fun x : (Iic n → I) × (ℕ → α → R) ↦ x.1)) + exact Measurable.comp (by fun_prop) measurable_fst + +lemma indepFun_fst_add_one_hist (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + IndepFun (fun ω ↦ ω.1 (n + 1)) (hist alg · n) (arrayMeasure ν) := + (indepFun_fst_add_one_aux ν n).of_measurable_right (measurable_hist_todo alg n) + +lemma hasCondDistrib_action' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + HasCondDistrib (action alg (n + 1)) (hist alg · n) (alg.policy n) (arrayMeasure ν) := by + rw [action_add_one_eq] + have h_fun ω := algFunction_map alg n (hist alg ω n) + refine ⟨by fun_prop, by fun_prop, ?_⟩ + refine condDistrib_ae_eq_of_measure_eq_compProd _ (by fun_prop) ?_ + have h_indep : (arrayMeasure ν).map (fun ω ↦ (ω.1 (n + 1), hist alg ω n)) = + (ℙ).prod ((arrayMeasure ν).map (hist alg · n)) := by + have h_indep' := indepFun_fst_add_one_hist alg ν n + rw [indepFun_iff_map_prod_eq_prod_map_map (by fun_prop) (by fun_prop)] at h_indep' + rw [h_indep'] + congr + simp only [arrayMeasure] + calc ((Measure.infinitePi fun x ↦ ℙ).prod (Bandit.streamMeasure ν)).map (fun ω ↦ ω.1 (n + 1)) + _ = (Measure.infinitePi fun x ↦ ℙ).map (Function.eval (n + 1)) := by + nth_rw 2 [← Measure.fst_prod (μ := Measure.infinitePi fun x ↦ ℙ) + (ν := Bandit.streamMeasure ν)] + rw [Measure.fst, Measure.map_map (by fun_prop) (by fun_prop)] + rfl + _ = ℙ := by rw [Measure.infinitePi_map_eval] + have : (fun x ↦ (hist alg x n, algFunction alg n (hist alg x n) (x.1 (n + 1)))) = + (fun p ↦ (p.2, algFunction alg n (p.2) (p.1))) ∘ (fun x ↦ (x.1 (n + 1), hist alg x n)) := rfl + rw [this, ← Measure.map_map (by fun_prop) (by fun_prop), h_indep] + have : (ℙ : Measure I).prod ((arrayMeasure ν).map (hist alg · n)) = + ((Kernel.const _ ℙ) ×ₖ Kernel.id) ∘ₘ ((arrayMeasure ν).map (hist alg · n)) := by + have h := Measure.compProd_const (μ := (arrayMeasure ν).map (hist alg · n)) + (ν := (ℙ : Measure I)) + rw [Measure.compProd_eq_comp_prod] at h + rw [← Measure.prod_swap, ← h, ← Measure.deterministic_comp_eq_map (by fun_prop), + Measure.comp_assoc, ← Kernel.swap, Kernel.swap_prod] + rw [this, ← Measure.deterministic_comp_eq_map (by fun_prop), + ← Measure.deterministic_comp_eq_map (by fun_prop), Measure.compProd_eq_comp_prod, + Measure.comp_assoc, Measure.comp_assoc, Measure.comp_assoc] + congr 2 + ext ω : 1 + simp only [Kernel.deterministic_comp_eq_map, Kernel.comp_deterministic_eq_comap, Kernel.coe_comap, + Function.comp_apply] + rw [Kernel.map_apply _ (by fun_prop), Kernel.prod_apply, Kernel.const_apply, Kernel.id_apply, + Kernel.prod_apply, Kernel.id_apply, ← h_fun] + calc (((ℙ).prod (Measure.dirac (hist alg ω n)))).map (fun p ↦ (p.2, algFunction alg n p.2 p.1)) + _ = (((ℙ).prod (Measure.dirac (hist alg ω n))).map Prod.swap).map + (fun p ↦ (p.1, algFunction alg n p.1 p.2)) := by + rw [Measure.map_map (by fun_prop) (by fun_prop)] + rfl + _ = ((Measure.dirac (hist alg ω n)).prod ℙ).map (fun p ↦ (p.1, algFunction alg n p.1 p.2)) := by + rw [Measure.prod_swap] + _ = (Measure.dirac (hist alg ω n)).prod ((ℙ).map (algFunction alg n (hist alg ω n))) := by + ext s hs + rw [Measure.map_apply (by fun_prop) hs, Measure.prod_apply, lintegral_dirac, Measure.prod_apply, + lintegral_dirac, Measure.map_apply (by fun_prop)] + · congr + · exact hs.preimage (by fun_prop) + · exact hs + · exact hs.preimage (by fun_prop) + +-- very bad name +/-- All random variables in the space, except for the unseen rewards for action `a` after +time `n`. -/ +noncomputable +def truePast (alg : Algorithm α R) (a : α) (n : ℕ) (ω : probSpace α R) : + probSpace α R := + (ω.1, fun i b ↦ if b = a then if pullCount (action alg) a (n + 1) ω ≠ 0 then + ω.2 (min i ((pullCount (action alg) a (n + 1) ω) - 1)) a else Nonempty.some inferInstance + else ω.2 i b) -lemma condDistrib_arm [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] +omit [Countable α] [StandardBorelSpace R] in +lemma truePast_eq_of_pullCount_eq (alg : Algorithm α R) + (a : α) (n m : ℕ) (ω : probSpace α R) + (h_pc : pullCount (action alg) a (n + 1) ω = m) : + truePast alg a n ω = (ω.1, fun i b ↦ if b = a then if m ≠ 0 then + ω.2 (min i (m - 1)) a else Nonempty.some inferInstance else ω.2 i b) := by + simp [truePast, h_pc] + +omit [Countable α] [StandardBorelSpace R] in +lemma truePast_eq_of_pullCount_eq_of_ne_zero (alg : Algorithm α R) + (a : α) (n m : ℕ) (ω : probSpace α R) + (h_pc : pullCount (action alg) a (n + 1) ω = m) (hm : m ≠ 0) : + truePast alg a n ω = (ω.1, fun i b ↦ if b = a then + ω.2 (min i (m - 1)) a else ω.2 i b) := by + simp [truePast, h_pc, hm] + +omit [StandardBorelSpace R] in +lemma measurable_hist_truePast (alg : Algorithm α R) + (a : α) (n : ℕ) : + Measurable[MeasurableSpace.comap (truePast alg a n) inferInstance] (hist alg · n) := by + have h_eq : (hist alg · n) = (hist alg · n) ∘ (truePast alg a n) := by + ext ω : 1 + refine hist_congr alg n (fun _ _ ↦ rfl) fun i b hi ↦ ?_ + by_cases hb : b = a + · subst hb + simp only [truePast, ↓reduceIte] + rw [min_eq_left, if_pos (by grind)] + grind + · simp [truePast, hb] + rw [h_eq] + refine Measurable.comp ?_ (Measurable.of_comap_le le_rfl) + fun_prop + +omit [StandardBorelSpace R] in +lemma measurable_action_add_one_truePast (alg : Algorithm α R) + (a : α) (n : ℕ) : + Measurable[MeasurableSpace.comap (truePast alg a n) inferInstance] + (action alg (n + 1)) := by + rw [action_add_one_eq] + change Measurable[MeasurableSpace.comap (truePast alg a n) inferInstance] + ((fun p ↦ algFunction alg n p.1 p.2) ∘ (fun ω ↦ (hist alg ω n, ω.1 (n + 1)))) + refine (measurable_algFunction alg n).comp (Measurable.prodMk ?_ ?_) + · exact measurable_hist_truePast alg a n + · have : (fun ω ↦ ω.1 (n + 1)) = + (fun (p : probSpace α R) ↦ p.1 (n + 1)) ∘ (truePast alg a n) := rfl + rw [this] + exact Measurable.comp (by fun_prop) (Measurable.of_comap_le le_rfl) + +omit [StandardBorelSpace R] in +lemma measurable_pullCount_add_one_truePast (alg : Algorithm α R) (a : α) (n : ℕ) : + Measurable[MeasurableSpace.comap (truePast alg a n) inferInstance] + (pullCount (action alg) a (n + 1)) := by + change Measurable[MeasurableSpace.comap (truePast alg a n) inferInstance] + (fun ω ↦ pullCount (action alg) a (n + 1) ω) + simp_rw [pullCount_eq_sum] + refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + refine (measurableSet_singleton _).preimage ?_ + have h_meas := measurable_hist_truePast alg a n + simp_rw [hist_eq _ _ n, @measurable_pi_iff] at h_meas + exact (h_meas ⟨i, by grind⟩).fst + +omit [Nonempty α] [StandardBorelSpace α] [Countable α] [StandardBorelSpace R] in +lemma indepFun_snd_apply_aux (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m : ℕ) : + (fun ω ↦ ω.2 m a) ⟂ᵢ[arrayMeasure ν] + (fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b)) := by + unfold arrayMeasure + let μ₁ : Measure (ℕ → I) := Measure.infinitePi fun _ ↦ volume + let μ₂ : Measure (ℕ → α → R) := Measure.infinitePi fun _ ↦ Measure.infinitePi ν + -- Independence within μ₂: coordinates ω i are independent + have h_indep₂ : iIndepFun (fun i (ω : ℕ → α → R) ↦ ω i) μ₂ := + iIndepFun_infinitePi (fun _ ↦ measurable_id) + -- Independence within each infinitePi ν: coordinates f b are independent + have h_indep_inner : iIndepFun (fun (b : α) (f : α → R) ↦ f b) (Measure.infinitePi ν) := + iIndepFun_infinitePi (fun _ ↦ measurable_id) + rw [indepFun_iff_measure_inter_preimage_eq_mul] + intro s t hs ht + let X : (ℕ → I) × (ℕ → α → R) → R := fun ω ↦ ω.2 m a + let Y : (ℕ → I) × (ℕ → α → R) → (ℕ → I) × (ℕ → α → R) := + fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b) + have hX_meas : Measurable X := + (measurable_pi_apply a).comp ((measurable_pi_apply m).comp measurable_snd) + have hY_meas : Measurable Y := by + change Measurable (fun ω : (ℕ → I) × (ℕ → α → R) ↦ + (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b)) + refine Measurable.prod measurable_fst ?_ + refine measurable_pi_lambda _ (fun k ↦ ?_) + refine measurable_pi_lambda _ (fun b ↦ ?_) + by_cases hb : b = a + · simp only [hb, ↓reduceIte] + by_cases hm : m ≠ 0 + · simp only [ne_eq, hm, not_false_eq_true, ↓reduceIte] + exact (measurable_pi_apply a).comp + ((measurable_pi_apply (min k (m - 1))).comp measurable_snd) + · simp only [hm, ↓reduceIte] + exact measurable_const + · simp only [hb, ↓reduceIte] + exact (measurable_pi_apply b).comp ((measurable_pi_apply k).comp measurable_snd) + change (μ₁.prod μ₂) (X ⁻¹' s ∩ Y ⁻¹' t) = (μ₁.prod μ₂) (X ⁻¹' s) * (μ₁.prod μ₂) (Y ⁻¹' t) + -- Use Fubini on μ₁.prod μ₂ + rw [Measure.prod_apply (hs.preimage hX_meas), + Measure.prod_apply (ht.preimage hY_meas), + Measure.prod_apply ((hs.preimage hX_meas).inter (ht.preimage hY_meas))] + -- X only depends on ω₂, so its fiber is constant in ω₁ + have hX_fst : ∀ ω₁, μ₂ (Prod.mk ω₁ ⁻¹' (X ⁻¹' s)) = μ₂ ((fun ω₂ ↦ ω₂ m a) ⁻¹' s) := fun _ ↦ rfl + simp_rw [hX_fst] + -- The LHS integral: fiber of X ∩ Y at ω₁ + -- Key: X depends only on ω₂ m a, while Y's dependence on ω₂ avoids (m, a) + -- Define the "truncation" map on ω₂ + let trunc : (ℕ → α → R) → (ℕ → α → R) := + fun ω₂ k b ↦ if b = a then if m ≠ 0 then ω₂ (min k (m - 1)) b + else Nonempty.some inferInstance else ω₂ k b + -- The fiber of Y at ω₁ only depends on trunc(ω₂) + have hY_fiber : ∀ ω₁, Prod.mk ω₁ ⁻¹' (Y ⁻¹' t) = (fun ω₂ ↦ (ω₁, trunc ω₂)) ⁻¹' t := fun _ ↦ rfl + -- The fiber of X ∩ Y factors + have hXY_fiber : ∀ ω₁, Prod.mk ω₁ ⁻¹' (X ⁻¹' s ∩ Y ⁻¹' t) = + ((fun ω₂ ↦ ω₂ m a) ⁻¹' s) ∩ ((fun ω₂ ↦ (ω₁, trunc ω₂)) ⁻¹' t) := fun _ ↦ rfl + simp_rw [hXY_fiber, hY_fiber] + -- Now we use independence in μ₂: (ω₂ m a) is independent of (trunc ω₂) + -- because trunc only uses indices (k, a) with k < m, and (k, b) with b ≠ a + have h_trunc_meas : Measurable trunc := by + refine measurable_pi_lambda _ (fun k ↦ ?_) + refine measurable_pi_lambda _ (fun b ↦ ?_) + simp only [trunc] + by_cases hb : b = a + · simp only [hb, ↓reduceIte] + by_cases hm : m = 0 + · simp only [hm] + exact measurable_const + · simp only [ne_eq, hm, not_false_eq_true, ↓reduceIte] + exact (measurable_pi_apply a).comp (measurable_pi_apply (min k (m - 1))) + · simp only [hb, ↓reduceIte] + exact (measurable_pi_apply b).comp (measurable_pi_apply k) + -- Key independence: (ω₂ m a) ⟂ trunc because trunc only uses coordinates ≠ (m, a) + have h_indep_trunc : IndepFun (fun ω₂ ↦ ω₂ m a) trunc μ₂ := by + -- Factor trunc through proj which extracts the relevant coordinates + let proj : (ℕ → α → R) → ((ℕ → R) × (ℕ → {b : α // b ≠ a} → R)) := fun ω₂ ↦ + (fun k ↦ if m ≠ 0 then ω₂ (min k (m - 1)) a else Nonempty.some inferInstance, + fun k ⟨b, _⟩ ↦ ω₂ k b) + have h_trunc_proj : ∀ ω₂, trunc ω₂ = (fun p k b ↦ + if h : b = a then if m ≠ 0 then p.1 k + else Nonempty.some inferInstance else p.2 k ⟨b, h⟩) (proj ω₂) := by + intro ω₂; ext k b; simp only [trunc, proj]; by_cases hb : b = a <;> simp [hb]; grind + have h_proj_meas : Measurable proj := by + refine Measurable.prod ?_ ?_ + · refine measurable_pi_lambda _ fun k ↦ ?_ + by_cases hm : m ≠ 0 + · simp only [proj, ne_eq, hm, not_false_eq_true, ↓reduceIte] + exact (measurable_pi_apply a).comp (measurable_pi_apply (min k (m - 1))) + · simp [proj, hm] + · exact measurable_pi_lambda _ (fun k ↦ measurable_pi_lambda _ (fun ⟨b, _⟩ ↦ + (measurable_pi_apply b).comp (measurable_pi_apply k))) + have h_g_meas : Measurable (fun p : (ℕ → R) × (ℕ → {b : α // b ≠ a} → R) ↦ + (fun k b ↦ if h : b = a then if m ≠ 0 then p.1 k else Nonempty.some inferInstance + else p.2 k ⟨b, h⟩)) := by + refine measurable_pi_lambda _ (fun k ↦ measurable_pi_lambda _ (fun b ↦ ?_)) + by_cases hb : b = a + · simp only [hb, ↓reduceDIte] + by_cases hm : m ≠ 0 + · simp only [ne_eq, hm, not_false_eq_true] + exact (measurable_pi_apply k).comp measurable_fst + · simp [hm] + · simp only [hb, ↓reduceDIte] + exact (measurable_pi_apply (⟨b, hb⟩ : {b : α // b ≠ a})).comp + ((measurable_pi_apply k).comp measurable_snd) + -- Show (ω₂ m a) ⟂ proj: proj uses coordinates disjoint from (m, a) + have h_indep_proj : IndepFun (fun ω₂ ↦ ω₂ m a) proj μ₂ := by + have h_row_bound (hm : m ≠ 0) : ∀ k, min k (m - 1) < m := by + intro k + calc min k (m - 1) ≤ m - 1 := Nat.min_le_right k (m - 1) + _ < m := Nat.sub_lt (by grind) Nat.one_pos + rw [indepFun_iff_measure_inter_preimage_eq_mul] + intro s t' hs ht' + -- rows_lt_m extracts column a at rows < m, other_cols extracts columns ≠ a + let rows_lt_m : (ℕ → α → R) → (Iio m → R) := fun ω₂ ⟨j, _⟩ ↦ ω₂ j a + let other_cols : (ℕ → α → R) → (ℕ → {b : α // b ≠ a} → R) := fun ω₂ k ⟨b, _⟩ ↦ ω₂ k b + have h_proj_factor : ∀ ω₂, proj ω₂ = + ((fun r k ↦ if hm : m ≠ 0 then r ⟨min k (m - 1), Finset.mem_Iio.mpr (h_row_bound hm k)⟩ + else Nonempty.some inferInstance) (rows_lt_m ω₂), + other_cols ω₂) := by + intro ω₂; ext1 + · ext k + by_cases hm : m ≠ 0 + · simp [proj, rows_lt_m, hm] + · simp [proj, hm] + · rfl + -- Use iIndepFun structure of the doubly-indexed infinite product + have h_iindep : iIndepFun (fun (p : ℕ × α) ω ↦ ω p.1 p.2) μ₂ := + iIndepFun_uncurry_infinitePi' (X := fun _ _ ↦ id) (fun _ ↦ ν) (by fun_prop) + have h_rows_meas : Measurable rows_lt_m := + measurable_pi_lambda _ (fun ⟨j, _⟩ ↦ (measurable_pi_apply a).comp (measurable_pi_apply j)) + have h_other_meas : Measurable other_cols := + measurable_pi_lambda _ (fun k ↦ measurable_pi_lambda _ (fun ⟨b, _⟩ ↦ + (measurable_pi_apply b).comp (measurable_pi_apply k))) + -- Show (ω₂ m a) ⟂ (rows_lt_m, other_cols) via indep_iSup_of_disjoint + have h_indep_combined : IndepFun (fun ω₂ ↦ ω₂ m a) + (fun ω₂ ↦ (rows_lt_m ω₂, other_cols ω₂)) μ₂ := by + rw [IndepFun_iff_Indep] + have h_comap_le : (MeasurableSpace.pi.prod MeasurableSpace.pi).comap + (fun ω₂ ↦ (rows_lt_m ω₂, other_cols ω₂)) ≤ + ⨆ (p : {p : ℕ × α // p ≠ (m, a)}), mR.comap (fun ω ↦ ω p.val.1 p.val.2) := by + rw [MeasurableSpace.comap_prodMk] + refine sup_le ?_ ?_ + · rw [MeasurableSpace.comap_pi] + refine iSup_le (fun ⟨j, hj⟩ ↦ ?_) + have h_ne : (j, a) ≠ (m, a) := fun h ↦ (Finset.mem_Iio.mp hj).ne (Prod.mk.inj h).1 + exact le_iSup_of_le ⟨(j, a), h_ne⟩ le_rfl + · rw [MeasurableSpace.comap_pi] + refine iSup_le (fun k ↦ ?_) + rw [MeasurableSpace.comap_pi] + refine iSup_le (fun ⟨b, hb⟩ ↦ ?_) + have h_ne : (k, b) ≠ (m, a) := fun h ↦ hb (Prod.mk.inj h).2 + exact le_iSup_of_le ⟨(k, b), h_ne⟩ le_rfl + refine indep_of_indep_of_le_right ?_ h_comap_le + have h_disjoint : Disjoint ({(m, a)} : Set (ℕ × α)) {p | p ≠ (m, a)} := by simp + have h_le : ∀ p : ℕ × α, mR.comap (fun ω : ℕ → α → R ↦ ω p.1 p.2) ≤ + MeasurableSpace.pi (m := fun _ ↦ MeasurableSpace.pi) := fun p ↦ + Measurable.comap_le ((measurable_pi_apply p.2).comp (measurable_pi_apply p.1)) + have h_iindep' : iIndep (fun p : ℕ × α ↦ mR.comap (fun ω : ℕ → α → R ↦ ω p.1 p.2)) μ₂ := + h_iindep.iIndep + have h_indep := indep_iSup_of_disjoint h_le h_iindep' h_disjoint + convert h_indep using 2 + · simp only [Set.mem_singleton_iff, iSup_iSup_eq_left] + · simp only [ne_eq, Set.mem_setOf_eq, iSup_subtype'] + have h_proj_preimage : proj ⁻¹' t' = (fun ω₂ ↦ (rows_lt_m ω₂, other_cols ω₂)) ⁻¹' + {p | ((fun r k ↦ if hm : m ≠ 0 then + r ⟨min k (m - 1), Finset.mem_Iio.mpr (h_row_bound hm k)⟩ + else Nonempty.some inferInstance) p.1, p.2) ∈ t'} + := by ext ω₂; simp only [Set.mem_preimage, Set.mem_setOf_eq, h_proj_factor] + rw [indepFun_iff_measure_inter_preimage_eq_mul] at h_indep_combined + rw [h_proj_preimage] + let T : Set ((Iio m → R) × (ℕ → {b : α // b ≠ a} → R)) := + {p | ((fun r k ↦ if hm : m ≠ 0 then + r ⟨min k (m - 1), Finset.mem_Iio.mpr (h_row_bound hm k)⟩ + else Nonempty.some inferInstance) p.1, p.2) ∈ t'} + have hT_meas : MeasurableSet T := by + refine ht'.preimage (Measurable.prod ?_ measurable_snd) + refine measurable_pi_lambda _ (fun k ↦ ?_) + by_cases hm : m ≠ 0 + · simp only [ne_eq, hm, not_false_eq_true, ↓reduceDIte] + exact (measurable_pi_apply (⟨min k (m - 1), Finset.mem_Iio.mpr (h_row_bound hm k)⟩ : + Iio m)).comp measurable_fst + · simp [hm] + change μ₂ ((fun ω₂ ↦ ω₂ m a) ⁻¹' s ∩ (fun ω₂ ↦ (rows_lt_m ω₂, other_cols ω₂)) ⁻¹' T) = + μ₂ ((fun ω₂ ↦ ω₂ m a) ⁻¹' s) * + μ₂ ((fun ω₂ ↦ (rows_lt_m ω₂, other_cols ω₂)) ⁻¹' T) + exact h_indep_combined s T hs hT_meas + have h_eq : trunc = (fun p k b ↦ if h : b = a then if m ≠ 0 then p.1 k + else Nonempty.some inferInstance else p.2 k ⟨b, h⟩) ∘ proj := by + funext ω₂; exact h_trunc_proj ω₂ + rw [h_eq] + exact h_indep_proj.comp measurable_id h_g_meas + rw [indepFun_iff_measure_inter_preimage_eq_mul] at h_indep_trunc + have h_const : ∀ ω₁, μ₂ (((fun ω₂ ↦ ω₂ m a) ⁻¹' s) ∩ ((fun ω₂ ↦ (ω₁, trunc ω₂)) ⁻¹' t)) = + μ₂ ((fun ω₂ ↦ ω₂ m a) ⁻¹' s) * μ₂ ((fun ω₂ ↦ (ω₁, trunc ω₂)) ⁻¹' t) := fun ω₁ ↦ + h_indep_trunc s _ hs (ht.preimage (by fun_prop)) + simp_rw [h_const] + let c := μ₂ ((fun ω₂ ↦ ω₂ m a) ⁻¹' s) + change ∫⁻ x, c * μ₂ ((fun ω₂ ↦ (x, trunc ω₂)) ⁻¹' t) ∂μ₁ = + (∫⁻ _, c ∂μ₁) * ∫⁻ x, μ₂ ((fun ω₂ ↦ (x, trunc ω₂)) ⁻¹' t) ∂μ₁ + have h_preimage : ∀ x, (fun ω₂ ↦ (x, trunc ω₂)) ⁻¹' t = trunc ⁻¹' (Prod.mk x ⁻¹' t) := fun _ ↦ rfl + simp_rw [h_preimage] + have h_map : ∀ x, μ₂ (trunc ⁻¹' (Prod.mk x ⁻¹' t)) = (μ₂.map trunc) (Prod.mk x ⁻¹' t) := by + intro x; rw [Measure.map_apply h_trunc_meas (ht.preimage (by fun_prop))] + simp_rw [h_map] + rw [lintegral_const_mul _ (measurable_measure_prodMk_left_finite ht), + lintegral_const, measure_univ, mul_one] + + +omit [StandardBorelSpace R] in +lemma measurable_stepsUntil (alg : Algorithm α R) (a : α) (m n : ℕ) : + Measurable[MeasurableSpace.comap + (fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b)) inferInstance] + (({ω | action alg (n + 1) ω = a ∧ + pullCount (action alg) a (n + 1) ω = m}).indicator (fun _ ↦ 1)) := by + let f := ({ω | action alg (n + 1) ω = a ∧ pullCount (action alg) a (n + 1) ω = m}).indicator + (fun _ ↦ 1) + have h_eq : f = f ∘ + fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b) := by + ext ω + exact stepsUntil_indicator_congr alg a m n (by grind) (by grind) (by grind) + change Measurable[MeasurableSpace.comap + (fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b)) inferInstance] f + rw [h_eq] + refine Measurable.comp ?_ (Measurable.of_comap_le le_rfl) + refine Measurable.indicator (by fun_prop) ?_ + exact MeasurableSet.inter ((measurableSet_singleton _).preimage (by fun_prop)) + ((measurableSet_singleton _).preimage (by fun_prop)) + +omit [StandardBorelSpace R] in +lemma indepFun_snd_apply_pullCount_action (alg : Algorithm α R) + (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (m n : ℕ) : + (fun ω ↦ ω.2 m a) ⟂ᵢ[arrayMeasure ν] + ({ω | action alg (n + 1) ω = a ∧ + pullCount (action alg) a (n + 1) ω = m}).indicator (fun _ ↦ 1) := + (indepFun_snd_apply_aux ν a m).of_measurable_right (measurable_stepsUntil alg a m n) + +omit [StandardBorelSpace R] [Nonempty R] in +@[fun_prop] +lemma measurable_pullCount_action_add_one (alg : Algorithm α R) (n : ℕ) : + Measurable (fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω) := by + change Measurable ((fun p : (probSpace α R) × α ↦ pullCount (action alg) p.2 (n + 1) p.1) ∘ + (fun ω : probSpace α R ↦ (ω, action alg (n + 1) ω))) + exact (measurable_uncurry_pullCount (by fun_prop) _).comp (by fun_prop) + +/-- The conditional distribution of the reward at time `n + 1`, given the action at time `n + 1` +and the number of times that action has been pulled before time `n + 1`, is equal to +the kernel `ν`. -/ +lemma hasCondDistrib_reward_pullCount_action (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : - condDistrib (arm (n + 1)) (hist n) (Bandit.trajMeasure alg ν) - =ᵐ[(Bandit.trajMeasure alg ν).map (hist n)] alg.policy n := - Learning.condDistrib_action alg (stationaryEnv ν) n - -/-- The reward at time `n+1` is independent of the history up to time `n` given the arm at `n+1`. -/ -lemma condIndepFun_reward_hist_arm [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] - {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) : - CondIndepFun (MeasurableSpace.comap (arm (n + 1)) inferInstance) - (measurable_arm _).comap_le (reward (n + 1)) (hist n) (Bandit.trajMeasure alg ν) := - Learning.condIndepFun_reward_hist_action n + HasCondDistrib (reward alg (n + 1)) + (fun ω ↦ (action alg (n + 1) ω, pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) + (ν.prodMkRight _) (arrayMeasure ν) := by + have h_meas : Measurable fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω := by + change Measurable ((fun p : (probSpace α R) × α ↦ pullCount (action alg) p.2 (n + 1) p.1) ∘ + (fun ω : probSpace α R ↦ (ω, action alg (n + 1) ω))) + exact (measurable_uncurry_pullCount (by fun_prop) _).comp (by fun_prop) + refine ⟨by fun_prop, by fun_prop, ?_⟩ + refine (condDistrib_ae_eq_cond + (Measurable.prodMk (by fun_prop) (by fun_prop)) (by fun_prop)).trans ?_ + rw [Filter.EventuallyEq, ae_iff_of_countable] + intro ⟨a, m⟩ ham + simp only [Kernel.prodMkRight_apply] + calc + Measure.map (reward alg (n + 1)) + (arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}] + _ = Measure.map (fun ω ↦ ω.2 m a) + (arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}] := by + rw [reward_eq] + refine Measure.map_congr + (ae_cond_of_forall_mem ((measurableSet_singleton _).preimage (by fun_prop)) (fun x hx ↦ ?_)) + simp only [Set.mem_preimage, Set.mem_singleton_iff, Prod.mk.injEq] at hx + simp only [hx.1] at hx ⊢ + simp [hx.2] + _ = Measure.map (fun ω ↦ ω.2 m a) + (arrayMeasure ν)[|({ω | action alg (n + 1) ω = a ∧ + pullCount (action alg) a (n + 1) ω = m}).indicator 1 ⁻¹' {1}] := by + congr with ω + simp only [Set.mem_preimage, Set.mem_singleton_iff, Prod.mk.injEq, Set.indicator_apply, + Set.mem_setOf_eq, Pi.one_apply, ite_eq_left_iff, not_and, zero_ne_one, imp_false, + Classical.not_imp, Decidable.not_not, and_congr_right_iff] + intro ha + simp [ha] + _ = ν a := by + rw [cond_of_indepFun, map_snd_apply_arrayMeasure m a] + · exact (indepFun_snd_apply_pullCount_action alg ν a m n).symm + · refine Measurable.indicator (by fun_prop) ?_ + exact MeasurableSet.inter ((measurableSet_singleton _).preimage (by fun_prop)) + ((measurableSet_singleton _).preimage (by fun_prop)) + · fun_prop + · simp + · rw [Measure.map_apply (by fun_prop) (by simp)] at ham + convert ham + ext ω + simp only [Set.mem_preimage, Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply, + Set.mem_singleton_iff, ite_eq_left_iff, not_and, zero_ne_one, imp_false, Classical.not_imp, + Decidable.not_not, Prod.mk.injEq, and_congr_right_iff] + intro ha + simp [ha] -end Laws +omit [StandardBorelSpace R] [Nonempty R] in +lemma reward_ae_eq_cond (alg : Algorithm α R) (ν : Kernel α R) (a : α) (n m : ℕ) : + reward alg (n + 1) =ᵐ[(arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}]] + (fun ω ↦ ω.2 m a) := by + rw [reward_eq] + refine ae_cond_of_forall_mem ?_ ?_ + · have : Measurable fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω := by + change Measurable ((fun p : (probSpace α R) × α ↦ pullCount (action alg) p.2 (n + 1) p.1) ∘ + (fun ω : probSpace α R ↦ (ω, action alg (n + 1) ω))) + exact (measurable_uncurry_pullCount (by fun_prop) _).comp (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + intro ω hω + simp only [Set.mem_preimage, Set.mem_singleton_iff, Prod.mk.injEq] at hω + simp only [hω.2] + simp [hω.1] + +lemma indepFun_todo {α β γ δ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {mδ : MeasurableSpace δ} [MeasurableSingletonClass δ] {μ : Measure α} + {X : α → β} {Y : α → γ} (hXY : X ⟂ᵢ[μ] Y) (hY : Measurable Y) + {Z : γ → δ} (hZ : Measurable Z) (z : δ) : + X ⟂ᵢ[μ[|(Z ∘ Y) ⁻¹' {z}]] Y := by + have h_preim : (Z ∘ Y) ⁻¹' {z} = Y ⁻¹' (Z ⁻¹' {z}) := by grind + simp_rw [h_preim] + exact indepFun_cond_of_indepFun hXY hY (hZ (measurableSet_singleton z)) -section DetAlgorithm +lemma indepFun_snd_hist_cond (alg : Algorithm α R) + (ν : Kernel α R) [IsMarkovKernel ν] (a : α) (n m : ℕ) : + (fun ω ↦ ω.2 m a) ⟂ᵢ[(arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}]] + (hist alg · n) := by + have h_meas := measurable_hist_truePast alg a n + refine IndepFun.of_measurable_right ?_ h_meas + have h_ae_eq : truePast alg a n =ᵐ[(arrayMeasure ν)[|(fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)}]] + (fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b)) := by + refine ae_cond_of_forall_mem ?_ fun x hx ↦ ?_ + · refine (measurableSet_singleton _).preimage ?_ + have h_meas_pc : Measurable fun ω ↦ + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω := by + change Measurable ((fun p : (probSpace α R) × α ↦ pullCount (action alg) p.2 (n + 1) p.1) ∘ + (fun ω : probSpace α R ↦ (ω, action alg (n + 1) ω))) + exact (measurable_uncurry_pullCount (by fun_prop) _).comp (by fun_prop) + fun_prop + simp only [Set.mem_preimage, Set.mem_singleton_iff, Prod.mk.injEq] at hx + simp only [truePast] + congr with i b + by_cases hb : b = a + · simp only [hb, ↓reduceIte] + simp only [hx.1, true_and] at hx + congr! + · simp [hb] + refine IndepFun.congr ?_ EventuallyEq.rfl h_ae_eq.symm + suffices (fun ω ↦ ω.2 m a) ⟂ᵢ[(arrayMeasure ν)[|(({ω | action alg (n + 1) ω = a ∧ + pullCount (action alg) a (n + 1) ω = m}).indicator (fun _ ↦ 1)) ⁻¹' {1}]] + fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b + else Nonempty.some inferInstance else ω.2 k b) by + convert this + ext ω + simp only [Set.mem_preimage, Set.mem_singleton_iff, Prod.mk.injEq, Set.indicator_apply, + Set.mem_setOf_eq, ite_eq_left_iff, not_and, zero_ne_one, imp_false, + Classical.not_imp, Decidable.not_not, and_congr_right_iff] + intro ha + simp [ha] + have h_meas := measurable_stepsUntil alg a m n + obtain ⟨f, hf, hf_eq⟩ := h_meas.exists_eq_measurable_comp + simp_rw [hf_eq] + refine indepFun_todo (Z := f) (z := 1) ?_ ?_ hf + · exact indepFun_snd_apply_aux ν a m + · refine Measurable.prodMk (by fun_prop) ?_ + simp_rw [measurable_pi_iff] + intro i b + refine Measurable.ite (MeasurableSet.const _) ?_ (by fun_prop) + refine Measurable.ite (MeasurableSet.const _) (by fun_prop) (by fun_prop) -variable {nextArm : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextArm n)} - {arm0 : α} {ν : Kernel α R} [IsMarkovKernel ν] +/-- The conditional distribution of the reward at time `n + 1`, given the history up to time `n`, +the action at time `n + 1`, and the number of times that action has been pulled before time `n + 1`, +is equal to the kernel `ν`. -/ +lemma hasCondDistrib_reward_hist_action_pullCount + (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + HasCondDistrib (reward alg (n + 1)) + (fun ω ↦ (hist alg ω n, action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) + ((ν.prodMkRight _).prodMkLeft _) (arrayMeasure ν) := by + have h_meas : Measurable fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω := by + change Measurable ((fun p : (probSpace α R) × α ↦ pullCount (action alg) p.2 (n + 1) p.1) ∘ + (fun ω : probSpace α R ↦ (ω, action alg (n + 1) ω))) + exact (measurable_uncurry_pullCount (by fun_prop) _).comp (by fun_prop) + refine ⟨by fun_prop, by fun_prop, ?_⟩ + refine condDistrib_prod_of_forall_condDistrib_cond (by fun_prop) (by fun_prop) (by fun_prop) _ ?_ + intro (a, m) ham + have h_eq : ((ν.prodMkRight _).prodMkLeft _).comap (fun ω : (Iic n → α × R) ↦ (ω, a, m)) + (by fun_prop) = + Kernel.const _ (ν a) := by ext; simp + rw [h_eq, condDistrib_congr_left (reward_ae_eq_cond alg ν a n m)] + refine (condDistrib_of_indepFun ?_ (by fun_prop) (by fun_prop)).trans (ae_of_all _ fun ω ↦ ?_) + · exact (indepFun_snd_hist_cond alg ν a n m).symm + · simp only [Kernel.const_apply] + have : (fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)) ⁻¹' {(a, m)} = + ({ω | action alg (n + 1) ω = a ∧ + pullCount (action alg) a (n + 1) ω = m}).indicator 1 ⁻¹' {1} := by + ext ω + simp [Set.indicator_apply] + grind + rw [this, cond_of_indepFun, map_snd_apply_arrayMeasure m a] + · exact (indepFun_snd_apply_pullCount_action alg ν a m n).symm + · refine Measurable.indicator (by fun_prop) ?_ + exact MeasurableSet.inter ((measurableSet_singleton _).preimage (by fun_prop)) + ((measurableSet_singleton _).preimage (by fun_prop)) + · fun_prop + · simp + · convert ham + ext ω + simp only [Set.mem_preimage, Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply, + Set.mem_singleton_iff, ite_eq_left_iff, not_and, zero_ne_one, imp_false, Classical.not_imp, + Decidable.not_not, Prod.mk.injEq, and_congr_right_iff] + intro ha + simp [ha] + +/-- The reward at time `n + 1` is conditionally independent of the history up to time `n`, +given the action at time `n + 1` and the number of times that action has been pulled before +time `n + 1`. -/ +lemma condIndepFun_reward_hist (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + (reward alg (n + 1)) ⟂ᵢ[(fun ω ↦ (action alg (n + 1) ω, + pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω)), + Measurable.prodMk (by fun_prop) (measurable_pullCount_action_add_one alg n); + arrayMeasure ν] + (hist alg · n) := by + have h_cond := hasCondDistrib_reward_hist_action_pullCount alg ν n + refine condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft (by fun_prop) (by fun_prop) ?_ + h_cond.condDistrib_eq + exact Measurable.prodMk (by fun_prop) (measurable_pullCount_action_add_one alg n) -local notation "𝔓t" => Bandit.trajMeasure (detAlgorithm nextArm h_next arm0) ν +omit [Countable α] [StandardBorelSpace R] [Nonempty R] in +lemma measurable_pullCount_action_add_one_hist (alg : Algorithm α R) (n : ℕ) : + Measurable[MeasurableSpace.comap (fun ω ↦ (action alg (n + 1) ω, hist alg ω n)) inferInstance] + (fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω) := by + simp_rw [pullCount_eq_sum] + refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + refine measurableSet_eq_fun ?_ (measurable_comp_comap _ measurable_fst) + simp_rw [hist_eq _ _ n] + unfold action + refine Measurable.fst (mγ := inferInstance) ?_ + have : (hist alg · i ⟨i, by grind⟩) = + (fun ω : α × (Iic n → α × R) ↦ ω.2 ⟨i, by grind⟩) ∘ + (fun ω ↦ (action alg (n + 1) ω, fun i : Iic n ↦ hist alg ω i ⟨i, by grind⟩)) := rfl + rw [this] + exact measurable_comp_comap _ (Measurable.prodMk (by fun_prop) (by fun_prop)) -lemma HasLaw_arm_zero_detAlgorithm : HasLaw (arm 0) (Measure.dirac arm0) 𝔓t where - map_eq := (hasLaw_arm_zero _ _).map_eq +/-- The conditional distribution of the reward at time `n + 1`, given the history up to time `n` +and the action at time `n + 1`, is equal to the kernel `ν`. -/ +lemma hasCondDistrib_reward' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + HasCondDistrib (reward alg (n + 1)) (fun ω ↦ (hist alg ω n, action alg (n + 1) ω)) + (ν.prodMkLeft _) (arrayMeasure ν) := by + let R' := reward alg (n + 1) + let H := (hist alg · n) + let A := action alg (n + 1) + let P := fun ω ↦ pullCount (action alg) (action alg (n + 1) ω) (n + 1) ω + have hP : Measurable P := measurable_pullCount_action_add_one alg n + change HasCondDistrib R' (fun ω ↦ (H ω, A ω)) (ν.prodMkLeft _) _ + suffices HasCondDistrib R' (fun ω ↦ (A ω, H ω)) (ν.prodMkRight _) (arrayMeasure ν) by + have h_eq : (fun ω ↦ (H ω, A ω)) = MeasurableEquiv.prodComm ∘ (fun ω ↦ (A ω, H ω)) := rfl + rw [h_eq] + exact this.comp_right (κ := ν.prodMkRight _) _ + suffices HasCondDistrib R' (fun ω ↦ ((A ω, H ω), P ω)) + ((ν.prodMkRight _).prodMkRight _) (arrayMeasure ν) by + -- use that `P` is measurable wrt `(A, H)` to drop it from the conditioning + have hP_meas : + Measurable[MeasurableSpace.comap (fun ω ↦ (A ω, H ω)) inferInstance] P := + measurable_pullCount_action_add_one_hist alg n + obtain ⟨f, hf_meas, hf_eq⟩ := hP_meas.exists_eq_measurable_comp + simp only [hf_eq, Function.comp_apply] at this + rwa [hasCondDistrib_prod_right_iff _ _ hf_meas] at this + suffices HasCondDistrib R' (fun ω ↦ ((A ω, P ω), H ω)) + ((ν.prodMkRight _).prodMkRight _) (arrayMeasure ν) by + let e : ((α × ℕ) × (Iic n → α × R)) ≃ᵐ ((α × (Iic n → α × R)) × ℕ) := + { toFun := fun x ↦ ((x.1.1, x.2), x.1.2) + invFun := fun x ↦ ((x.1.1, x.2), x.1.2) + measurable_toFun := by fun_prop + measurable_invFun := by fun_prop } + exact this.comp_right e + suffices HasCondDistrib R' (fun ω ↦ (A ω, P ω)) (ν.prodMkRight _) (arrayMeasure ν) by + have h_indep : H ⟂ᵢ[(fun ω ↦ (A ω, P ω)), (by fun_prop); arrayMeasure ν] R' := + (condIndepFun_reward_hist alg ν n).symm + have h_condDistrib := this.condDistrib_eq + rw [condIndepFun_iff_condDistrib_prod_ae_eq_prodMkRight (by fun_prop) (by fun_prop) + (by fun_prop)] at h_indep + refine ⟨by fun_prop, by fun_prop, ?_⟩ + refine h_indep.trans ?_ + rw [Filter.EventuallyEq, ae_map_iff] at h_condDistrib ⊢ + · simpa only [Kernel.prodMkRight_apply] + · fun_prop + · exact Kernel.measurableSet_eq _ _ + · fun_prop + · exact Kernel.measurableSet_eq _ _ + exact hasCondDistrib_reward_pullCount_action alg ν n -lemma arm_zero_detAlgorithm [MeasurableSingletonClass α] : - arm 0 =ᵐ[𝔓t] fun _ ↦ arm0 := - Learning.action_zero_detAlgorithm +lemma hasCondDistrib_action (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + HasCondDistrib (action alg (n + 1)) + (fun ω (i : Iic n) ↦ (action alg i ω, reward alg i ω)) + (alg.policy n) (arrayMeasure ν) := by + convert hasCondDistrib_action' alg ν n with ω i + · simp only [action] + rw [hist_eq _ _ n] + · simp only [reward] + rw [hist_eq _ _ n] -lemma arm_detAlgorithm_ae_eq [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] (n : ℕ) : - arm (n + 1) =ᵐ[𝔓t] fun h ↦ nextArm n (fun i ↦ h i) := - Learning.action_detAlgorithm_ae_eq n +lemma hasCondDistrib_reward (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] + (n : ℕ) : + HasCondDistrib (reward alg (n + 1)) + (fun ω ↦ (fun (i : Iic n) ↦ (action alg i ω, reward alg i ω), action alg (n + 1) ω)) + ((stationaryEnv ν).feedback n) (arrayMeasure ν) := by + convert hasCondDistrib_reward' alg ν n with ω i + · simp only [action] + rw [hist_eq _ _ n] + · simp only [reward] + rw [hist_eq _ _ n] -example [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] : - ∀ᵐ h ∂(𝔓t), arm 0 h = arm0 ∧ ∀ n, arm (n + 1) h = nextArm n (fun i ↦ h i) := by - rw [eventually_and, ae_all_iff] - exact ⟨arm_zero_detAlgorithm, arm_detAlgorithm_ae_eq⟩ +lemma isAlgEnvSeq_arrayMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : + IsAlgEnvSeq (action alg) (reward alg) alg (stationaryEnv ν) (arrayMeasure ν) where + hasLaw_action_zero := hasLaw_action_zero alg ν + hasCondDistrib_reward_zero := hasCondDistrib_reward_zero alg ν + hasCondDistrib_action := hasCondDistrib_action alg ν + hasCondDistrib_reward := hasCondDistrib_reward alg ν + +end Laws -end DetAlgorithm +end ArrayModel end MeasureSpace diff --git a/LeanBandits/Bandit/Regret.lean b/LeanBandits/Bandit/Regret.lean index acabf6c8..fdbfb8da 100644 --- a/LeanBandits/Bandit/Regret.lean +++ b/LeanBandits/Bandit/Regret.lean @@ -3,11 +3,10 @@ Copyright (c) 2025 Rémy Degenne. All rights reserved. Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne, Paulo Rauber -/ -import LeanBandits.Bandit.Bandit import LeanBandits.SequentialLearning.FiniteActions /-! -# Regret +# Regret, gap, best arm -/ @@ -17,17 +16,12 @@ open scoped ENNReal NNReal namespace Bandits -variable {α : Type*} [DecidableEq α] {mα : MeasurableSpace α} {ν : Kernel α ℝ} - {h : ℕ → α × ℝ} {m n t : ℕ} {a : α} +variable {α Ω : Type*} [DecidableEq α] {mα : MeasurableSpace α} {mΩ : MeasurableSpace Ω} + {ν : Kernel α ℝ} + {A : ℕ → Ω → α} {R : ℕ → Ω → ℝ} + {ω : Ω} {m n t : ℕ} {a : α} -/-! ### Definitions of regret, gaps, pull counts -/ - -/-- Regret of a sequence of pulls `k : ℕ → α` at time `t` for the reward kernel `ν ; Kernel α ℝ`. -/ -noncomputable -def regret (ν : Kernel α ℝ) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := - t * (⨆ a, (ν a)[id]) - ∑ s ∈ range t, (ν (arm s h))[id] - -/-- Gap of an arm `a`: difference between the highest mean of the arms and the mean of `a`. -/ +/-- Gap of an action `a`: difference between the highest mean of the actions and the mean of `a`. -/ noncomputable def gap (ν : Kernel α ℝ) (a : α) : ℝ := (⨆ i, (ν i)[id]) - (ν a)[id] @@ -36,28 +30,35 @@ lemma gap_nonneg [Fintype α] : 0 ≤ gap ν a := by rw [gap, sub_nonneg] exact le_ciSup (f := fun i ↦ (ν i)[id]) (by simp) a -lemma arm_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) h = m) : - arm (stepsUntil a m h).toNat h = a := by - exact action_stepsUntil hm h_exists +/-- Regret of a sequence of pulls `k : ℕ → α` at time `t` for the reward kernel `ν ; Kernel α ℝ`. -/ +noncomputable +def regret (ν : Kernel α ℝ) (A : ℕ → Ω → α) (t : ℕ) (ω : Ω) : ℝ := + t * (⨆ a, (ν a)[id]) - ∑ s ∈ range t, (ν (A s ω))[id] + +omit [DecidableEq α] in +lemma regret_eq_sum_gap : regret ν A t ω = ∑ s ∈ range t, gap ν (A s ω) := by + simp [regret, gap] -lemma arm_eq_of_stepsUntil_eq_coe {ω : ℕ → α × ℝ} (hm : m ≠ 0) - (h : stepsUntil a m ω = n) : - arm n ω = a := by - exact action_eq_of_stepsUntil_eq_coe hm h +omit [DecidableEq α] in +lemma regret_nonneg [Fintype α] : 0 ≤ regret ν A t ω := by + rw [regret_eq_sum_gap] + exact sum_nonneg (fun _ _ ↦ gap_nonneg) -section RewardByCount +omit [DecidableEq α] in +lemma gap_eq_zero_of_regret_eq_zero [Fintype α] (hr : regret ν A t ω = 0) {s : ℕ} (hs : s < t) : + gap ν (A s ω) = 0 := by + rw [regret_eq_sum_gap] at hr + exact (sum_eq_zero_iff_of_nonneg fun _ _ ↦ gap_nonneg).1 hr s (mem_range.2 hs) lemma regret_eq_sum_pullCount_mul_gap [Fintype α] : - regret ν t h = ∑ a, pullCount a t h * gap ν a := by - simp [sum_pullCount_mul, regret, gap, sum_sub_distrib, arm, action] + regret ν A t ω = ∑ a, pullCount A a t ω * gap ν a := by + simp_rw [regret_eq_sum_gap, sum_pullCount_mul] -end RewardByCount - -section BestArm +section bestArm variable [Fintype α] [Nonempty α] -/-- Arm with the highest mean. -/ +/-- action with the highest mean. -/ noncomputable def bestArm (ν : Kernel α ℝ) : α := (exists_max_image univ (fun a ↦ (ν a)[id]) (univ_nonempty_iff.mpr inferInstance)).choose @@ -78,6 +79,47 @@ omit [DecidableEq α] in lemma gap_bestArm : gap ν (bestArm ν) = 0 := by rw [gap_eq_bestArm_sub, sub_self] -end BestArm +omit [DecidableEq α] in +lemma integral_eq_of_gap_eq_zero (hg : gap ν a = 0) : (ν (bestArm ν))[id] = (ν a)[id] := by + rwa [← sub_eq_zero, ← gap_eq_bestArm_sub] + +end bestArm + +section Asymptotics + +omit [DecidableEq α] in +/-- If the regret is sublinear, the average mean reward tends to the highest mean of the arms. -/ +lemma avg_mean_reward_tendsto_of_sublinear_regret + (hr : (regret ν A · ω) =o[atTop] fun t ↦ (t : ℝ)) : + Tendsto (fun t ↦ (∑ s ∈ range t, (ν (A s ω))[id]) / (t : ℝ)) + atTop (nhds (⨆ a, (ν a)[id])) := by + have ht : Tendsto (fun t ↦ (⨆ a, (ν a)[id]) - regret ν A t ω / t) + atTop (nhds (⨆ a, (ν a)[id])) := by + simpa using tendsto_const_nhds.sub hr.tendsto_div_nhds_zero + apply ht.congr' + filter_upwards [eventually_ne_atTop 0] with t ht + rw [regret] + field_simp + ring + +/-- If the regret is sublinear, the rate of suboptimal arm pulls tends to zero. -/ +lemma pullCount_rate_tendsto_of_sublinear_regret [Fintype α] + (hr : (regret ν A · ω) =o[atTop] fun t ↦ (t : ℝ)) (hg : 0 < gap ν a) : + Tendsto (fun t ↦ (pullCount A a t ω : ℝ) / t) atTop (nhds 0) := by + have hb (t : ℕ) : (pullCount A a t ω : ℝ) * gap ν a ≤ regret ν A t ω := by + rw [regret_eq_sum_pullCount_mul_gap] + exact single_le_sum (f := fun a ↦ pullCount A a t ω * gap ν a) + (fun _ _ ↦ mul_nonneg (Nat.cast_nonneg _) gap_nonneg) (mem_univ a) + have hb' (t : ℕ) : (pullCount A a t ω : ℝ) / t ≤ regret ν A t ω / t / gap ν a := by + obtain ht | ht := eq_or_ne t 0 + · simp [ht] + · calc (pullCount A a t ω : ℝ) / t + = pullCount A a t ω * gap ν a / gap ν a / t := by field_simp + _ ≤ regret ν A t ω / gap ν a / t := by gcongr; exact hb t + _ = regret ν A t ω / t / gap ν a := by ring + apply squeeze_zero' (Eventually.of_forall fun _ ↦ by positivity) (Eventually.of_forall hb') + simpa using hr.tendsto_div_nhds_zero.div_const (gap ν a) + +end Asymptotics end Bandits diff --git a/LeanBandits/Bandit/RewardByCountMeasure.lean b/LeanBandits/Bandit/RewardByCountMeasure.lean new file mode 100644 index 00000000..85b8aa46 --- /dev/null +++ b/LeanBandits/Bandit/RewardByCountMeasure.lean @@ -0,0 +1,325 @@ +/- +Copyright (c) 2025 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +import LeanBandits.Bandit.Bandit +import Mathlib.Probability.IdentDistribIndep + +/-! # Laws of `stepsUntil` and `rewardByCount` +-/ + +open MeasureTheory ProbabilityTheory Finset Learning +open scoped ENNReal NNReal + +namespace Bandits + +variable {α Ω : Type*} {mα : MeasurableSpace α} {mΩ : MeasurableSpace Ω} [DecidableEq α] + [StandardBorelSpace α] [Nonempty α] + {A : ℕ → Ω → α} {R : ℕ → Ω → ℝ} {P : Measure Ω} [IsProbabilityMeasure P] + {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] + {h_inter : IsAlgEnvSeq A R alg (stationaryEnv ν) P} + +local notation "𝔓'" => P.prod (Bandit.streamMeasure ν) + +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in +lemma hasLaw_Z (a : α) (m : ℕ) : + HasLaw (fun ω ↦ ω.2 m a) (ν a) 𝔓' where + map_eq := by + calc (𝔓').map (fun ω ↦ ω.2 m a) + _ = ((𝔓').snd).map (fun ω ↦ ω m a) := by + rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)] + rfl + _ = (Bandit.streamMeasure ν).map (fun ω ↦ ω m a) := by simp + _ = ((Measure.infinitePi fun _ ↦ Measure.infinitePi ν).map (fun ω ↦ ω m)).map + (fun ω ↦ ω a) := by + rw [Bandit.streamMeasure, Measure.map_map (by fun_prop) (by fun_prop)] + rfl + _ = ν a := by simp_rw [(measurePreserving_eval_infinitePi _ _).map_eq] + +/-- Law of `Y` conditioned on the event `s`.-/ +notation "𝓛[" Y " | " s "; " μ "]" => Measure.map Y (μ[|s]) +/-- Law of `Y` conditioned on the event that `X` is in `s`. -/ +notation "𝓛[" Y " | " X " in " s "; " μ "]" => Measure.map Y (μ[|X ⁻¹' s]) +/-- Law of `Y` conditioned on the event that `X` equals `x`. -/ +notation "𝓛[" Y " | " X " ← " x "; " μ "]" => Measure.map Y (μ[|X ⁻¹' {x}]) + +local notation "𝔓t" => Bandit.trajMeasure alg ν +local notation "𝔓" => Bandit.measure alg ν + +omit [DecidableEq α] in +lemma condDistrib_reward'' [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (n : ℕ) : + 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1; 𝔓'] =ᵐ[(𝔓').map (fun ω ↦ A n ω.1)] ν := by + have hA := h.measurable_A + have hR := h.measurable_R + have h_ra' : 𝓛[R n | A n; P] =ᵐ[P.map (A n)] ν := h.condDistrib_reward_stationaryEnv n + have h_law : (𝔓').map (fun ω ↦ A n ω.1) = P.map (A n) := by + change ((𝔓').map (A n ∘ Prod.fst)) = _ + rw [← Measure.map_map (by fun_prop) (by fun_prop), ← Measure.fst, Measure.fst_prod] + rw [h_law] + have h_prod : 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1; 𝔓'] + =ᵐ[P.map (A n)] 𝓛[R n | A n; P] := + condDistrib_fst_prod _ (by fun_prop) _ + filter_upwards [h_ra', h_prod] with ω h_eq h_prod + rw [h_prod, h_eq] + +omit [DecidableEq α] in +lemma reward_cond_action [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n : ℕ) + (hμa : (𝔓').map (fun ω ↦ A n ω.1) {a} ≠ 0) : + 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1 ← a; 𝔓'] = ν a := by + have hA := h.measurable_A + have hR := h.measurable_R + have h_ra : 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1; 𝔓'] =ᵐ[(𝔓').map (fun ω ↦ A n ω.1)] ν := + condDistrib_reward'' h n + have h_eq := condDistrib_ae_eq_cond (μ := 𝔓') + (X := fun ω ↦ A n ω.1) (Y := fun ω ↦ R n ω.1) (by fun_prop) (by fun_prop) + rw [Filter.EventuallyEq, ae_iff_of_countable] at h_ra h_eq + specialize h_ra a hμa + specialize h_eq a hμa + rw [h_ra] at h_eq + exact h_eq.symm + +lemma condIndepFun_reward_stepsUntil_action' [StandardBorelSpace Ω] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m n : ℕ) : + R n ⟂ᵢ[A n, h.measurable_A n; P] {ω | stepsUntil A a m ω = ↑n}.indicator (fun _ ↦ 1) := by + -- the indicator of `stepsUntil ... = n` is a function of `hist (n-1)` and `action n`. + -- It thus suffices to use the independence of `reward n` and `hist (n-1)` conditionally + -- on `action n`. + have hA := h.measurable_A + have hR := h.measurable_R + by_cases hn : n = 0 + · have h_indep : R 0 ⟂ᵢ[A 0, hA 0; P] A 0 := + condIndepFun_self_right (by fun_prop) (by fun_prop) + simp only [hn, CharP.cast_eq_zero] + refine h_indep.of_measurable_right (hX := hA 0) ?_ + exact measurable_comap_indicator_stepsUntil_eq_zero a m + · have h_indep : R n ⟂ᵢ[A n, hA n; P] fun ω ↦ (IsAlgEnvSeq.hist A R (n - 1) ω, A n ω) := + IsAlgEnvSeq.condIndepFun_reward_hist_action_action' h n (by grind) + refine h_indep.of_measurable_right (hX := hA n) ?_ + exact measurable_comap_indicator_stepsUntil_eq hA hR a m n + +lemma condIndepFun_reward_stepsUntil_action [StandardBorelSpace Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + (a : α) (m n : ℕ) : + CondIndepFun (mα.comap (fun ω ↦ A n ω.1)) ((h.measurable_A n).comp measurable_fst).comap_le + (fun ω ↦ R n ω.1) ({ω | stepsUntil A a m ω.1 = ↑n}.indicator (fun _ ↦ 1)) 𝔓' := by + have hA := h.measurable_A + have hR := h.measurable_R + exact condIndepFun_fst_prod (ν := Bandit.streamMeasure ν) + (measurable_indicator_stepsUntil_eq hA hR a m n) (by fun_prop) (by fun_prop) + (condIndepFun_reward_stepsUntil_action' h a m n) + +lemma reward_cond_stepsUntil [StandardBorelSpace Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m n : ℕ) + (hm : m ≠ 0) (hμn : 𝔓' ((fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n}) ≠ 0) : + 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ stepsUntil A a m ω.1 ← ↑n; 𝔓'] = ν a := by + have hA := h.measurable_A + have hR := h.measurable_R + have hμna : + 𝔓' ((fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ A n ω.1) ⁻¹' {a}) ≠ 0 := by + suffices ((fun ω : Ω × (ℕ → α → ℝ) ↦ + stepsUntil A a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ A n ω.1) ⁻¹' {a}) + = (fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n} by simpa [this] using hμn + ext ω + simp only [Set.mem_inter_iff, Set.mem_preimage, Set.mem_singleton_iff, and_iff_left_iff_imp] + exact action_eq_of_stepsUntil_eq_coe hm + have hμa : (𝔓').map (fun ω ↦ A n ω.1) {a} ≠ 0 := by + rw [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] + refine fun h_zero ↦ hμn (measure_mono_null (fun ω ↦ ?_) h_zero) + simp only [Set.mem_preimage, Set.mem_singleton_iff] + exact action_eq_of_stepsUntil_eq_coe hm + calc 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ stepsUntil A a m ω.1 ← (n : ℕ∞); 𝔓'] + _ = (𝔓'[|(fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ A n ω.1) ⁻¹' {a}]).map + (fun ω ↦ R n ω.1) := by + congr with ω + simp only [Set.mem_preimage, Set.mem_singleton_iff, Set.mem_inter_iff, iff_self_and] + exact action_eq_of_stepsUntil_eq_coe hm + _ = (𝔓'[|(fun ω ↦ A n ω.1) ⁻¹' {a} + ∩ {ω : Ω × (ℕ → α → ℝ) | stepsUntil A a m ω.1 = ↑n}.indicator 1 ⁻¹' {1} ]).map + (fun ω ↦ R n ω.1) := by + congr 2 with ω + simp only [Set.mem_inter_iff, Set.mem_preimage, Set.mem_singleton_iff, Set.indicator_apply, + Set.mem_setOf_eq, Pi.one_apply, ite_eq_left_iff, zero_ne_one, imp_false, Decidable.not_not] + rw [and_comm] + _ = 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1 ← a; 𝔓'] := by + rw [cond_of_condIndepFun (by fun_prop)] + · exact condIndepFun_reward_stepsUntil_action h a m n + · refine measurable_one.indicator ?_ + exact measurableSet_eq_fun (by fun_prop) (by fun_prop) + · fun_prop + · convert hμna using 2 + rw [Set.inter_comm] + congr 1 with ω + simp [Set.indicator_apply] + _ = ν a := reward_cond_action h a n hμa + +/-- The conditional distribution of the reward received at the `m`-th pull of action `a` +given the time at which number of pulls is `m` is the constant kernel with value `ν a`. -/ +theorem condDistrib_rewardByCount_stepsUntil [StandardBorelSpace Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m : ℕ) (hm : m ≠ 0) : + condDistrib (rewardByCount A R a m) (fun ω ↦ stepsUntil A a m ω.1) 𝔓' + =ᵐ[(𝔓').map (fun ω ↦ stepsUntil A a m ω.1)] Kernel.const _ (ν a) := by + have hA := h.measurable_A + have hR := h.measurable_R + refine (condDistrib_ae_eq_cond (μ := 𝔓') + (X := fun ω ↦ stepsUntil A a m ω.1) (by fun_prop) (by fun_prop)).trans ?_ + rw [Filter.EventuallyEq, ae_iff_of_countable] + intro n hn + simp only [Kernel.const_apply] + cases n with + | top => + rw [Measure.map_congr (g := fun ω ↦ ω.2 m a)] + swap + · refine ae_cond_of_forall_mem ((measurableSet_singleton _).preimage (by fun_prop)) ?_ + simp only [Set.mem_preimage, Set.mem_singleton_iff] + exact fun ω ↦ rewardByCount_of_stepsUntil_eq_top + rw [cond_of_indepFun _ (by fun_prop) (by fun_prop) (measurableSet_singleton _)] + · exact (hasLaw_Z a m).map_eq + · rwa [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] at hn + · exact indepFun_prod (X := fun ω : Ω ↦ stepsUntil A a m ω) + (Y := fun ω : ℕ → α → ℝ ↦ ω m a) (by fun_prop) (by fun_prop) + | coe n => + rw [Measure.map_congr (g := fun ω ↦ R n ω.1)] + swap + · refine ae_cond_of_forall_mem ((measurableSet_singleton _).preimage (by fun_prop)) ?_ + simp only [Set.mem_preimage, Set.mem_singleton_iff] + exact fun ω ↦ rewardByCount_of_stepsUntil_eq_coe + refine reward_cond_stepsUntil h a m n hm ?_ + rwa [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] at hn + +/-- The reward received at the `m`-th pull of action `a` has law `ν a`. -/ +lemma hasLaw_rewardByCount [StandardBorelSpace Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m : ℕ) (hm : m ≠ 0) : + HasLaw (rewardByCount A R a m) (ν a) 𝔓' where + aemeasurable := (measurable_rewardByCount h.measurable_A h.measurable_R a m).aemeasurable + map_eq := by + have hA := h.measurable_A + have hR := h.measurable_R + have h_condDistrib : + condDistrib (rewardByCount A R a m) (fun ω ↦ stepsUntil A a m ω.1) 𝔓' + =ᵐ[(𝔓').map (fun ω ↦ stepsUntil A a m ω.1)] + Kernel.const _ (ν a) := condDistrib_rewardByCount_stepsUntil h a m hm + calc (𝔓').map (rewardByCount A R a m) + _ = (condDistrib (rewardByCount A R a m) (fun ω ↦ stepsUntil A a m ω.1) 𝔓') + ∘ₘ ((𝔓').map (fun ω ↦ stepsUntil A a m ω.1)) := by + rw [condDistrib_comp_map (by fun_prop) (by fun_prop)] + _ = (Kernel.const _ (ν a)) ∘ₘ ((𝔓').map (fun ω ↦ stepsUntil A a m ω.1)) := + Measure.comp_congr h_condDistrib + _ = ν a := by + have : IsProbabilityMeasure ((𝔓').map (fun ω ↦ stepsUntil A a m ω.1)) := + Measure.isProbabilityMeasure_map (by fun_prop) + simp + +lemma identDistrib_rewardByCount [StandardBorelSpace Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n m : ℕ) + (hn : n ≠ 0) (hm : m ≠ 0) : + IdentDistrib (rewardByCount A R a n) (rewardByCount A R a m) 𝔓' 𝔓' where + aemeasurable_fst := (measurable_rewardByCount h.measurable_A h.measurable_R a n).aemeasurable + aemeasurable_snd := (measurable_rewardByCount h.measurable_A h.measurable_R a m).aemeasurable + map_eq := by rw [(hasLaw_rewardByCount h a n hn).map_eq, (hasLaw_rewardByCount h a m hm).map_eq] + +lemma identDistrib_rewardByCount_id [StandardBorelSpace Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n : ℕ) (hn : n ≠ 0) : + IdentDistrib (rewardByCount A R a n) id 𝔓' (ν a) where + aemeasurable_fst := (measurable_rewardByCount h.measurable_A h.measurable_R a n).aemeasurable + aemeasurable_snd := Measurable.aemeasurable <| by fun_prop + map_eq := by rw [(hasLaw_rewardByCount h a n hn).map_eq, Measure.map_id] + +lemma identDistrib_rewardByCount_eval [StandardBorelSpace Ω] [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n m : ℕ) (hn : n ≠ 0) : + IdentDistrib (rewardByCount A R a n) (fun ω ↦ ω m a) 𝔓' (Bandit.streamMeasure ν) := + (identDistrib_rewardByCount_id h a n hn).trans + (identDistrib_eval_eval_id_streamMeasure ν m a).symm + +-- lemma indepFun_rewardByCount_Iic [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) +-- (n : ℕ) : +-- (rewardByCount A R a (n + 1)) ⟂ᵢ[𝔓'] fun ω (i : Iic n) ↦ rewardByCount A R a i ω := by +-- sorry + +-- lemma iIndepFun_rewardByCount' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) : +-- iIndepFun (rewardByCount A R a) 𝔓' := by +-- have hA := h.measurable_A +-- have hR := h.measurable_R +-- rw [iIndepFun_nat_iff_forall_indepFun (by fun_prop)] +-- exact indepFun_rewardByCount_Iic h a + +-- lemma iIndepFun_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) : +-- iIndepFun (fun (p : α × ℕ) ↦ rewardByCount A R p.1 (p.2 + 1)) 𝔓' := by +-- sorry + +-- lemma identDistrib_rewardByCount_stream_all [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) : +-- IdentDistrib (fun ω (p : α × ℕ) ↦ rewardByCount A R p.1 (p.2 + 1) ω) +-- (fun ω p ↦ ω p.2 p.1) 𝔓' (Bandit.streamMeasure ν) := by +-- refine IdentDistrib.pi (fun p ↦ ?_) ?_ ?_ +-- · refine identDistrib_rewardByCount_eval h p.1 (p.2 + 1) p.2 (by simp) (ν := ν) +-- · exact iIndepFun_rewardByCount h +-- · sorry + +-- lemma identDistrib_rewardByCount_stream' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) : +-- IdentDistrib (fun ω n ↦ rewardByCount A R a (n + 1) ω) (fun ω n ↦ ω n a) +-- 𝔓' (Bandit.streamMeasure ν) := by +-- refine IdentDistrib.pi (fun n ↦ ?_) ?_ ?_ +-- · refine identDistrib_rewardByCount_eval h a (n + 1) n (by simp) (ν := ν) +-- · have h_indep := iIndepFun_rewardByCount' h a +-- exact iIndepFun.precomp (g := fun n ↦ n + 1) (fun i j hij ↦ by grind) h_indep +-- · exact iIndepFun_eval_streamMeasure'' ν a + +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in +lemma identDistrib_eval_streamMeasure_measure (a : α) : + IdentDistrib (fun ω n ↦ ω n a) (fun ω n ↦ ω.2 n a) + (Bandit.streamMeasure ν) 𝔓 := by + refine IdentDistrib.pi (fun n ↦ ?_) ?_ ?_ + · rw [← Bandit.snd_measure alg ν, Measure.snd, + identDistrib_map_left_iff (by fun_prop) (by fun_prop) + (Measurable.aemeasurable <| by fun_prop)] + exact IdentDistrib.refl (by fun_prop) + · exact iIndepFun_eval_streamMeasure'' ν a + · change iIndepFun (fun n ↦ ((fun ω ↦ ω n a) ∘ Prod.snd)) 𝔓 + rw [← iIndepFun_map_iff (by fun_prop) (fun _ ↦ Measurable.aemeasurable (by fun_prop))] + rw [← Measure.snd, Bandit.snd_measure] + exact iIndepFun_eval_streamMeasure'' ν a + +-- lemma identDistrib_rewardByCount_stream [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) : +-- IdentDistrib (fun ω n ↦ rewardByCount A R a (n + 1) ω) (fun ω n ↦ ω.2 n a) 𝔓' 𝔓 := +-- (identDistrib_rewardByCount_stream' h a).trans (identDistrib_eval_streamMeasure_measure a) + +-- lemma indepFun_rewardByCount_of_ne [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {a b : α} (hab : a ≠ b) : +-- IndepFun (fun ω s ↦ rewardByCount A R a s ω) (fun ω s ↦ rewardByCount A R b s ω) 𝔓' := by +-- sorry + +-- lemma identDistrib_sum_Icc_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α] +-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (m : ℕ) (a : α) : +-- IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount A R a s ω) +-- (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓' 𝔓 := by +-- have h1 (a : α) : +-- IdentDistrib (fun ω s ↦ rewardByCount A R a (s + 1) ω) (fun ω s ↦ ω.2 s a) 𝔓' 𝔓 := +-- identDistrib_rewardByCount_stream h a +-- have h_eq (ω : Ω × (ℕ → α → ℝ)) : ∑ s ∈ Icc 1 m, rewardByCount A R a s ω +-- = ∑ s ∈ range m, rewardByCount A R a (s + 1) ω := by +-- let e : Icc 1 m ≃ range m := +-- { toFun x := ⟨x - 1, by have h := x.2; simp only [mem_Icc] at h; simp; grind⟩ +-- invFun x := ⟨x + 1, by +-- have h := x.2 +-- simp only [mem_Icc, le_add_iff_nonneg_left, zero_le, true_and, ge_iff_le] +-- simp only [mem_range] at h +-- grind⟩ +-- left_inv x := by have h := x.2; simp only [mem_Icc] at h; grind +-- right_inv x := by have h := x.2; grind } +-- rw [← sum_coe_sort (Icc 1 m), ← sum_coe_sort (range m), sum_equiv e] +-- · simp +-- · simp only [univ_eq_attach, mem_attach, forall_const, Subtype.forall, mem_Icc, +-- forall_and_index] +-- grind +-- simp_rw [h_eq] +-- exact IdentDistrib.comp (h1 a) (u := fun p ↦ ∑ s ∈ range m, p s) (by fun_prop) + +end Bandits diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean new file mode 100644 index 00000000..65ecf7e8 --- /dev/null +++ b/LeanBandits/Bandit/SumRewards.lean @@ -0,0 +1,612 @@ +/- +Copyright (c) 2025 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +import LeanBandits.Bandit.Bandit +import LeanBandits.Bandit.Regret +import LeanBandits.ForMathlib.SubGaussian + +/-! # Law of the sum of rewards +-/ + +open MeasureTheory ProbabilityTheory Finset Learning +open scoped ENNReal NNReal + +lemma measurable_sum_range_of_le {α : Type*} {mα : MeasurableSpace α} + {f : ℕ → α → ℝ} {g : α → ℕ} {n : ℕ} (hg_le : ∀ a, g a ≤ n) (hf : ∀ i, Measurable (f i)) + (hg : Measurable g) : + Measurable (fun a ↦ ∑ i ∈ range (g a), f i a) := by + have h_eq : (fun a ↦ ∑ i ∈ range (g a), f i a) + = fun a ↦ ∑ i ∈ range (n + 1), if g a = i then ∑ j ∈ range i, f j a else 0 := by + ext ω + rw [sum_ite_eq_of_mem] + grind + rw [h_eq] + refine measurable_sum _ fun n hn ↦ ?_ + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + +lemma measurable_sum_Icc_of_le {α : Type*} {mα : MeasurableSpace α} + {f : ℕ → α → ℝ} {g : α → ℕ} {n : ℕ} (hg_le : ∀ a, g a ≤ n) (hf : ∀ i, Measurable (f i)) + (hg : Measurable g) : + Measurable (fun a ↦ ∑ i ∈ Icc 1 (g a), f i a) := by + have h_eq : (fun a ↦ ∑ i ∈ Icc 1 (g a), f i a) + = fun a ↦ ∑ i ∈ range (n + 1), if g a = i then ∑ j ∈ Icc 1 i, f j a else 0 := by + ext ω + rw [sum_ite_eq_of_mem] + grind + rw [h_eq] + refine measurable_sum _ fun n hn ↦ ?_ + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + +namespace Bandits + +namespace ArrayModel + +variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [Countable α] + [StandardBorelSpace α] [Nonempty α] + {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] + +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 (Bandit.streamMeasure ν)) 𝔓 where + aemeasurable_fst := by + refine Measurable.aemeasurable ?_ + rw [measurable_pi_iff] + refine fun a ↦ Measurable.prod (by fun_prop) ?_ + exact measurable_sum_Icc_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) + aemeasurable_snd := by + refine Measurable.aemeasurable ?_ + rw [measurable_pi_iff] + refine fun a ↦ Measurable.prod (by fun_prop) ?_ + exact measurable_sum_Icc_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) + map_eq := by + by_cases hn : n = 0 + · simp [hn] + have h_eq (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 (μ := 𝔓) (ν := Bandit.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 (Bandit.streamMeasure ν)) 𝔓 := by + convert identDistrib_pullCount_prod_sum_Icc_rewardByCount' n using 2 with ω + rotate_left + · infer_instance + · infer_instance + ext a : 1 + congr 1 + let e : Icc 1 (pullCount A a n ω) ≃ range (pullCount A a n ω) := + { toFun x := ⟨x - 1, by have h := x.2; simp only [mem_Icc] at h; simp; grind⟩ + invFun x := ⟨x + 1, by + have h := x.2 + simp only [mem_Icc, le_add_iff_nonneg_left, zero_le, true_and, ge_iff_le] + simp only [mem_range] at h + grind⟩ + left_inv x := by have h := x.2; simp only [mem_Icc] at h; grind + right_inv x := by have h := x.2; grind } + rw [← sum_coe_sort (Icc 1 (pullCount A a n ω)), ← sum_coe_sort (range (pullCount A a n ω)), + sum_equiv e] + · simp + · simp [e] + +lemma identDistrib_pullCount_prod_sumRewards (n : ℕ) : + IdentDistrib (fun ω a ↦ (pullCount A a n ω, sumRewards A R a n ω)) + (fun ω a ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) 𝔓 𝔓 := by + suffices IdentDistrib (fun ω a ↦ (pullCount A a n ω.1, sumRewards A R a n ω.1)) + (fun ω a ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) + ((𝔓).prod (Bandit.streamMeasure ν)) 𝔓 by + -- todo: missing lemma about IdentDistrib? + constructor + · refine Measurable.aemeasurable ?_ + fun_prop + · refine Measurable.aemeasurable ?_ + rw [measurable_pi_iff] + refine fun a ↦ Measurable.prod (by fun_prop) ?_ + exact measurable_sum_range_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) + have h_eq := this.map_eq + nth_rw 1 [← Measure.fst_prod (μ := 𝔓) (ν := Bandit.streamMeasure ν), Measure.fst, + Measure.map_map (by fun_prop) (by fun_prop)] + exact h_eq + simp_rw [← sum_rewardByCount_eq_sumRewards] + exact identDistrib_pullCount_prod_sum_Icc_rewardByCount n + +lemma identDistrib_pullCount_prod_sumRewards_arm (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) + 𝔓 (Bandit.streamMeasure ν) where + aemeasurable_fst := by fun_prop + aemeasurable_snd := (measurable_sum _ fun i _ ↦ by fun_prop).aemeasurable + map_eq := by + rw [← Measure.snd_prod (μ := (Measure.infinitePi fun (_ : ℕ) ↦ (volume : Measure unitInterval))) + (ν := Bandit.streamMeasure ν), Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)] + rfl + +lemma prob_pullCount_prod_sumRewards_mem_le (a : α) (n : ℕ) + {s : Set (ℕ × ℝ)} [DecidablePred (· ∈ Prod.fst '' s)] (hs : MeasurableSet s) : + 𝔓 {ω | (pullCount A a n ω, sumRewards A R a n ω) ∈ s} ≤ + ∑ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), + Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := by + have h_ident := identDistrib_pullCount_prod_sumRewards_arm a n (ν := ν) (alg := alg) + have : 𝔓 {ω | (pullCount A a n ω, 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) + calc 𝔓 ((fun ω ↦ (pullCount A a n ω, ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) ⁻¹' s) + _ ≤ 𝔓 {ω | ∃ k ≤ n, (k, ∑ i ∈ range k, ω.2 i a) ∈ s} := by + refine measure_mono fun ω hω ↦ ?_ + simp only [Set.mem_setOf_eq] at hω ⊢ + exact ⟨pullCount A a n ω, pullCount_le _ _ _, hω⟩ + _ = 𝔓 (⋃ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), + {ω | (k, ∑ i ∈ range k, ω.2 i a) ∈ s}) := by congr 1; ext; simp; grind + _ ≤ ∑ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), + 𝔓 {ω | ∑ i ∈ range k, ω.2 i a ∈ Prod.mk k ⁻¹' s} := measure_biUnion_finset_le _ _ + _ = ∑ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), + (Bandit.streamMeasure ν) {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := by + congr with k + have : (𝔓).map (fun ω ↦ ∑ i ∈ range k, ω.2 i a) = + (Bandit.streamMeasure ν).map (fun ω ↦ ∑ i ∈ range k, ω i a) := + (identDistrib_sum_range_snd a k).map_eq + rw [Measure.ext_iff] at this + specialize this (Prod.mk k ⁻¹' s) (hs.preimage (by fun_prop)) + rwa [Measure.map_apply (by fun_prop) (hs.preimage (by fun_prop)), + Measure.map_apply (by fun_prop) (hs.preimage (by fun_prop))] at this + +lemma prob_pullCount_mem_and_sumRewards_mem_le (a : α) (n : ℕ) + {s : Set ℕ} [DecidablePred (· ∈ s)] (hs : MeasurableSet s) {B : Set ℝ} (hB : MeasurableSet B) : + 𝔓 {ω | pullCount A a n ω ∈ s ∧ sumRewards A R a n ω ∈ B} ≤ + ∑ k ∈ (range (n + 1)).filter (· ∈ s), + Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := by + classical + rcases Set.eq_empty_or_nonempty B with h_empty | h_nonempty + · simp [h_empty] + convert prob_pullCount_prod_sumRewards_mem_le a n (hs.prod hB) (ν := ν) (alg := alg) with _ _ k hk + · ext n + have : ∃ x, x ∈ B := h_nonempty + simp [this] + · ext x + simp only [Set.mem_image, Set.mem_prod, Prod.exists, exists_and_right, exists_and_left, + exists_eq_right, mem_filter, mem_range] at hk + simp [hk.2.1] + +lemma prob_sumRewards_le_sumRewards_le [Fintype α] (a : α) (n m₁ m₂ : ℕ) : + (𝔓) {ω | pullCount A (bestArm ν) n ω = m₁ ∧ pullCount A a n ω = m₂ ∧ + sumRewards A R (bestArm ν) n ω ≤ sumRewards A R a n ω} ≤ + Bandit.streamMeasure ν + {ω | ∑ i ∈ range m₁, ω i (bestArm ν) ≤ ∑ i ∈ range m₂, ω i a} := by + have h_ident := identDistrib_pullCount_prod_sumRewards_two_arms (bestArm ν) a n + (ν := ν) (alg := alg) + let s := {p : ℕ × ℕ × ℝ × ℝ | p.1 = m₁ ∧ p.2.1 = m₂ ∧ p.2.2.1 ≤ p.2.2.2} + have hs : MeasurableSet s := by simp only [measurableSet_setOf, s]; fun_prop + calc 𝔓 {ω | pullCount A (bestArm ν) n ω = m₁ ∧ pullCount A a n ω = m₂ ∧ + sumRewards A R (bestArm ν) n ω ≤ sumRewards A R a n ω} + _ = 𝔓 ((fun ω ↦ (pullCount A (bestArm ν) n ω, pullCount A a n ω, + sumRewards A R (bestArm ν) n ω, sumRewards A R a n ω)) ⁻¹' + {p | p.1 = m₁ ∧ p.2.1 = m₂ ∧ p.2.2.1 ≤ p.2.2.2}) := rfl + _ = 𝔓 ((fun ω ↦ (pullCount A (bestArm ν) n ω, pullCount A a n ω, + ∑ i ∈ range (pullCount A (bestArm ν) n ω), ω.2 i (bestArm ν), + ∑ i ∈ range (pullCount A a n ω), ω.2 i a)) ⁻¹' + {p | p.1 = m₁ ∧ p.2.1 = m₂ ∧ p.2.2.1 ≤ p.2.2.2}) := by + rw [← Measure.map_apply (by fun_prop) hs, h_ident.map_eq, + Measure.map_apply _ hs] + refine Measurable.prod (by fun_prop) (Measurable.prod (by fun_prop) ?_) + refine Measurable.prod ?_ ?_ + · exact measurable_sum_range_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) + · exact measurable_sum_range_of_le (n := n) (pullCount_le _ _) (by fun_prop) (by fun_prop) + _ ≤ 𝔓 ((fun ω ↦ (∑ i ∈ range m₁, ω.2 i (bestArm ν), ∑ i ∈ range m₂, ω.2 i a)) ⁻¹' + {p | p.1 ≤ p.2}) := by + refine measure_mono fun ω hω ↦ ?_ + simp only [Set.preimage_setOf_eq, Set.mem_setOf_eq] at hω ⊢ + grind + _ = Bandit.streamMeasure ν + {ω | ∑ i ∈ range m₁, ω i (bestArm ν) ≤ ∑ i ∈ range m₂, ω i a} := by + rw [← Measure.snd_prod (μ := (Measure.infinitePi fun (_ : ℕ) ↦ (volume : Measure unitInterval))) + (ν := Bandit.streamMeasure ν), Measure.snd, Measure.map_apply (by fun_prop)] + · rfl + simp only [measurableSet_setOf] + fun_prop + +lemma probReal_sumRewards_le_sumRewards_le [Fintype α] (a : α) (n m₁ m₂ : ℕ) : + (𝔓).real {ω | pullCount A (bestArm ν) n ω = m₁ ∧ pullCount A a n ω = m₂ ∧ + sumRewards A R (bestArm ν) n ω ≤ sumRewards A R a n ω} ≤ + (Bandit.streamMeasure ν).real + {ω | ∑ i ∈ range m₁, ω i (bestArm ν) ≤ ∑ i ∈ range m₂, ω i a} := by + simp_rw [measureReal_def] + gcongr + · finiteness + · exact prob_sumRewards_le_sumRewards_le a n m₁ m₂ + +end ArrayModel + +variable {α Ω Ω' : Type*} [DecidableEq α] {mα : MeasurableSpace α} {mΩ : MeasurableSpace Ω} + {mΩ' : MeasurableSpace Ω'} + {P : Measure Ω} [IsProbabilityMeasure P] {P' : Measure Ω'} [IsProbabilityMeasure P'] + {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] + {A : ℕ → Ω → α} {R : ℕ → Ω → ℝ} {A₂ : ℕ → Ω' → α} {R₂ : ℕ → Ω' → ℝ} + {ω : Ω} {m n t : ℕ} {a : α} + +lemma sumRewards_eq_comp : + sumRewards A R a n = + (fun p ↦ ∑ i ∈ range n, if (p i).1 = a then (p i).2 else 0) ∘ (fun ω n ↦ (A n ω, R n ω)) := by + ext + simp [sumRewards] + +lemma pullCount_eq_comp : + pullCount A a n = + (fun p ↦ ∑ i ∈ range n, if (p i).1 = a then 1 else 0) ∘ (fun ω n ↦ (A n ω, R n ω)) := by + ext + simp [pullCount] + +variable [StandardBorelSpace α] [Nonempty α] + +-- todo: write those lemmas with IdentDistrib instead of equality of maps +lemma _root_.Learning.IsAlgEnvSeq.law_sumRewards_unique + (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + (h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') : + P.map (sumRewards A R a n) = P'.map (sumRewards A₂ R₂ a n) := by + have hA := h1.measurable_A + have hR := h1.measurable_R + have hA2 := h2.measurable_A + have hR2 := h2.measurable_R + have h_unique := isAlgEnvSeq_unique h1 h2 + rw [sumRewards_eq_comp, sumRewards_eq_comp, ← Measure.map_map, h_unique, Measure.map_map, + ← sumRewards_eq_comp] + · refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + · rw [measurable_pi_iff] + exact fun n ↦ Measurable.prodMk (hA2 n) (hR2 n) + · refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + · rw [measurable_pi_iff] + exact fun n ↦ Measurable.prodMk (hA n) (hR n) + +lemma _root_.Learning.IsAlgEnvSeq.law_pullCount_sumRewards_unique' + (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + (h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') : + IdentDistrib (fun ω a ↦ (pullCount A a n ω, sumRewards A R a n ω)) + (fun ω a ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) P P' := by + have hA := h1.measurable_A + have hR := h1.measurable_R + have hA2 := h2.measurable_A + have hR2 := h2.measurable_R + constructor + · refine Measurable.aemeasurable ?_ + rw [measurable_pi_iff] + exact fun a ↦ Measurable.prod (by fun_prop) (measurable_sumRewards hA hR _ _) + · refine Measurable.aemeasurable ?_ + rw [measurable_pi_iff] + exact fun a ↦ Measurable.prod (by fun_prop) (measurable_sumRewards hA2 hR2 _ _) + have h_unique := isAlgEnvSeq_unique h1 h2 + let f := fun (p : ℕ → α × ℝ ) (a : α) ↦ (∑ i ∈ range n, if (p i).1 = a then 1 else 0, + ∑ i ∈ range n, if (p i).1 = a then (p i).2 else 0) + have hf : Measurable f := by + rw [measurable_pi_iff] + intro a + refine Measurable.prod ?_ ?_ + · simp only [f] + refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + · simp only [f] + refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + have h_eq_comp : (fun ω a ↦ (pullCount A a n ω, sumRewards A R a n ω)) + = f ∘ (fun ω n ↦ (A n ω, R n ω)) := by + ext ω a : 2 + rw [pullCount_eq_comp (R := R), sumRewards_eq_comp] + grind + have h_eq_comp2 : (fun ω a ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) + = f ∘ (fun ω n ↦ (A₂ n ω, R₂ n ω)) := by + ext ω a : 2 + rw [pullCount_eq_comp (R := R₂), sumRewards_eq_comp] + grind + rw [h_eq_comp, h_eq_comp2, ← Measure.map_map hf, h_unique, Measure.map_map hf, + ← h_eq_comp2] + · rw [measurable_pi_iff] + exact fun n ↦ Measurable.prodMk (hA2 n) (hR2 n) + · rw [measurable_pi_iff] + exact fun n ↦ Measurable.prodMk (hA n) (hR n) + +lemma _root_.Learning.IsAlgEnvSeq.law_pullCount_sumRewards_unique + (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + (h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') : + 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) + +-- this is what we will use for UCB +lemma prob_pullCount_prod_sumRewards_mem_le [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + {s : Set (ℕ × ℝ)} [DecidablePred (· ∈ Prod.fst '' s)] (hs : MeasurableSet s) : + P {ω | (pullCount A a n ω, sumRewards A R a n ω) ∈ s} ≤ + ∑ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), + Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := by + have hA := h.measurable_A + have hR := h.measurable_R + calc P {ω | (pullCount A a n ω, sumRewards A R a n ω) ∈ s} + _ = (P.map (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω))) s := by + rw [Measure.map_apply (by fun_prop) hs]; rfl + _ = ((ArrayModel.arrayMeasure ν).map + (fun ω ↦ (pullCount (ArrayModel.action alg) a n ω, + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a n ω))) s := by + rw [h.law_pullCount_sumRewards_unique (ArrayModel.isAlgEnvSeq_arrayMeasure alg ν)] + _ = (ArrayModel.arrayMeasure ν) {ω | (pullCount (ArrayModel.action alg) a n ω, + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a n ω) ∈ s} := by + rw [Measure.map_apply (by fun_prop) hs]; rfl + _ ≤ ∑ k ∈ (range (n + 1)).filter (· ∈ Prod.fst '' s), + Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := + ArrayModel.prob_pullCount_prod_sumRewards_mem_le a n hs + +lemma prob_pullCount_mem_and_sumRewards_mem_le [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + {s : Set ℕ} [DecidablePred (· ∈ s)] (hs : MeasurableSet s) {B : Set ℝ} (hB : MeasurableSet B) : + P {ω | pullCount A a n ω ∈ s ∧ sumRewards A R a n ω ∈ B} ≤ + ∑ k ∈ (range (n + 1)).filter (· ∈ s), + Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := by + classical + rcases Set.eq_empty_or_nonempty B with h_empty | h_nonempty + · simp [h_empty] + convert prob_pullCount_prod_sumRewards_mem_le h (hs.prod hB) (ν := ν) (alg := alg) with _ _ k hk + · ext n + have : ∃ x, x ∈ B := h_nonempty + simp [this] + · ext x + simp only [Set.mem_image, Set.mem_prod, Prod.exists, exists_and_right, exists_and_left, + exists_eq_right, mem_filter, mem_range] at hk + simp [hk.2.1] + +lemma prob_sumRewards_mem_le [Countable α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + {B : Set ℝ} (hB : MeasurableSet B) : + P (sumRewards A R a n ⁻¹' B) ≤ + ∑ k ∈ range (n + 1), Bandit.streamMeasure ν {ω | ∑ i ∈ range k, ω i a ∈ B} := by + classical + have h_le := prob_pullCount_mem_and_sumRewards_mem_le h .univ hB (a := a) (n := n) + simpa using h_le + +lemma prob_pullCount_eq_and_sumRewards_mem_le [Countable α] + (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + {m : ℕ} (hm : m ≤ n) {B : Set ℝ} (hB : MeasurableSet B) : + P {ω | pullCount A a n ω = m ∧ sumRewards A R a n ω ∈ B} ≤ + Bandit.streamMeasure ν {ω | ∑ i ∈ range m, ω i a ∈ B} := by + have h_le := prob_pullCount_mem_and_sumRewards_mem_le h (s := {m}) (by simp) hB (a := a) (n := n) + have hm' : m < n + 1 := by lia + simpa [hm'] using h_le + +lemma probReal_sumRewards_le_sumRewards_le [Fintype α] (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) + (a : α) (n m₁ m₂ : ℕ) : + P.real {ω | pullCount A (bestArm ν) n ω = m₁ ∧ pullCount A a n ω = m₂ ∧ + sumRewards A R (bestArm ν) n ω ≤ sumRewards A R a n ω} ≤ + (Bandit.streamMeasure ν).real + {ω | ∑ i ∈ range m₁, ω i (bestArm ν) ≤ ∑ i ∈ range m₂, ω i a} := by + have hA := h.measurable_A + have hR := h.measurable_R + refine le_trans (le_of_eq ?_) + (ArrayModel.probReal_sumRewards_le_sumRewards_le (alg := alg) a n m₁ m₂) + let s := {p : ℕ × ℕ × ℝ × ℝ | p.1 = m₁ ∧ p.2.1 = m₂ ∧ p.2.2.1 ≤ p.2.2.2} + have hs : MeasurableSet s := by simp only [measurableSet_setOf, s]; fun_prop + change P.real ((fun ω ↦ (pullCount A (bestArm ν) n ω, + pullCount A a n ω, sumRewards A R (bestArm ν) n ω, sumRewards A R a n ω)) ⁻¹' s) = + (ArrayModel.arrayMeasure ν).real + ((fun ω ↦ (pullCount (ArrayModel.action alg) (bestArm ν) n ω, + pullCount (ArrayModel.action alg) a n ω, + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) (bestArm ν) n ω, + sumRewards (ArrayModel.action alg) (ArrayModel.reward alg) a n ω)) ⁻¹' s) + simp_rw [measureReal_def] + congr 1 + rw [← Measure.map_apply ?_ hs, ← Measure.map_apply (by fun_prop) hs] + swap + · refine Measurable.prod (by fun_prop) (Measurable.prod (by fun_prop) ?_) + exact (measurable_sumRewards hA hR _ _).prod (measurable_sumRewards hA hR _ _) + congr 1 + refine IdentDistrib.map_eq ?_ + have h_eq := h.law_pullCount_sumRewards_unique' (ArrayModel.isAlgEnvSeq_arrayMeasure alg ν) + (n := n) + exact h_eq.comp (u := fun p ↦ ((p (bestArm ν)).1, (p a).1, (p (bestArm ν)).2, (p a).2)) + (by fun_prop) + +section Subgaussian + +omit [DecidableEq α] [StandardBorelSpace α] in +lemma probReal_sum_le_sum_streamMeasure [Fintype α] + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : α) (m : ℕ) : + (Bandit.streamMeasure ν).real + {ω | ∑ s ∈ range m, ω s (bestArm ν) ≤ ∑ s ∈ range m, ω s a} ≤ + Real.exp (-↑m * gap ν a ^ 2 / 4) := by + by_cases ha : a = bestArm ν + · simp [ha] + refine (HasSubgaussianMGF.measure_sum_le_sum_le' (cX := fun _ ↦ 1) (cY := fun _ ↦ 1) + ?_ ?_ ?_ ?_ ?_ ?_).trans_eq ?_ + · exact iIndepFun_eval_streamMeasure'' ν (bestArm ν) + · exact iIndepFun_eval_streamMeasure'' ν a + · intro i him + simp_rw [integral_eval_streamMeasure] + refine (hν (bestArm ν)).congr_identDistrib ?_ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ + · intro i him + simp_rw [integral_eval_streamMeasure] + refine (hν a).congr_identDistrib ?_ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ + · exact indepFun_eval_streamMeasure' ν (Ne.symm ha) + · gcongr 1 with i him + simp_rw [integral_eval_streamMeasure] + exact le_bestArm a + · congr 1 + simp_rw [integral_eval_streamMeasure] + simp only [id_eq, sum_const, card_range, nsmul_eq_mul, mul_one, NNReal.coe_natCast, + gap_eq_bestArm_sub, neg_mul] + field_simp + ring + +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in +lemma prob_sum_le_sqrt_log + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) {c : ℝ} (hc : 0 ≤ c) + (a : α) (k : ℕ) (hk : k ≠ 0) : + Bandit.streamMeasure ν + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ - √(c * k * Real.log (n + 1))} ≤ + 1 / (n + 1) ^ (c / 2) := by + calc + Bandit.streamMeasure ν + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ - √(c * k * Real.log (n + 1))} + _ ≤ ENNReal.ofReal (Real.exp (-(√(c * k * Real.log (n + 1))) ^ 2 / (2 * k * 1))) := by + rw [← ofReal_measureReal] + gcongr + refine (HasSubgaussianMGF.measure_sum_range_le_le_of_iIndepFun (c := 1) ?_ ?_ (by positivity)) + · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) + (fun _ ↦ by fun_prop) + · intro i him + refine (hν a).congr_identDistrib ?_ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ + _ = 1 / (n + 1) ^ (c / 2) := by + rw [Real.sq_sqrt] + swap; · exact mul_nonneg (by positivity) (Real.log_nonneg (by simp)) + field_simp + rw [div_eq_inv_mul, ← mul_assoc, ← Real.log_rpow (by positivity), ← Real.log_inv, + Real.exp_log (by positivity), one_div, ENNReal.ofReal_inv_of_pos (by positivity), + ← ENNReal.ofReal_rpow_of_nonneg (by positivity) (by positivity)] + congr 2 + · norm_cast + · field + +omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in +lemma prob_sum_ge_sqrt_log + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) {c : ℝ} (hc : 0 ≤ c) + (a : α) (k : ℕ) (hk : k ≠ 0) : + Bandit.streamMeasure ν + {ω | √(c * k * Real.log (n + 1)) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} ≤ + 1 / (n + 1) ^ (c / 2) := by + calc + Bandit.streamMeasure ν + {ω | √(c * k * Real.log (n + 1)) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} + _ ≤ ENNReal.ofReal (Real.exp (-(√(c * k * Real.log (n + 1))) ^ 2 / (2 * k * 1))) := by + rw [← ofReal_measureReal] + gcongr + refine (HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := 1) ?_ ?_ (by positivity)) + · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) + (fun _ ↦ by fun_prop) + · intro i him + refine (hν a).congr_identDistrib ?_ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ + _ = 1 / (n + 1) ^ (c / 2) := by + rw [Real.sq_sqrt] + swap; · exact mul_nonneg (by positivity) (Real.log_nonneg (by simp)) + field_simp + rw [div_eq_inv_mul, ← mul_assoc, ← Real.log_rpow (by positivity), ← Real.log_inv, + Real.exp_log (by positivity), one_div, ENNReal.ofReal_inv_of_pos (by positivity), + ← ENNReal.ofReal_rpow_of_nonneg (by positivity) (by positivity)] + congr 2 + · norm_cast + · field + +end Subgaussian + +end Bandits diff --git a/LeanBandits/BanditAlgorithms/AuxSums.lean b/LeanBandits/BanditAlgorithms/AuxSums.lean new file mode 100644 index 00000000..27bdf41f --- /dev/null +++ b/LeanBandits/BanditAlgorithms/AuxSums.lean @@ -0,0 +1,47 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +import Mathlib.Algebra.BigOperators.Intervals +import Mathlib.Algebra.BigOperators.Ring.Finset +import Mathlib.Tactic.Ring.RingNF + +open Finset + +lemma sum_mod_range {K : ℕ} (hK : 0 < K) (a : Fin K) : + (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) = 1 := by + have h_iff (s : ℕ) (hs : s < K) : ⟨s % K, Nat.mod_lt _ hK⟩ = a ↔ s = a := by + simp only [Nat.mod_eq_of_lt hs, Fin.ext_iff] + calc (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) + _ = ∑ s ∈ range K, if s = a then 1 else 0 := sum_congr rfl fun s hs ↦ by grind + _ = _ := by + rw [sum_ite_eq'] + simp + +lemma sum_mod_range_mul {K : ℕ} (hK : 0 < K) (m : ℕ) (a : Fin K) : + (∑ s ∈ range (K * m), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) = m := by + induction m with + | zero => simp + | succ n hn => + calc (∑ s ∈ range (K * (n + 1)), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) + _ = (∑ s ∈ range (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by ring_nf + _ = (∑ s ∈ range (K * n), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) + + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by + rw [sum_range_add_sum_Ico] + grind + _ = n + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by + rw [hn] + _ = n + (∑ s ∈ range K, if ⟨(s + K * n) % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by + congr 1 + let e : ℕ ↪ ℕ := ⟨fun i : ℕ ↦ i + K * n, fun i j hij ↦ by grind⟩ + have : Finset.map e (range K) = Ico (K * n) (K * n + K) := by + ext x + simp only [mem_map, mem_range, Function.Embedding.coeFn_mk, mem_Ico, e] + refine ⟨fun h ↦ by grind, fun h ↦ ?_⟩ + use x - K * n + grind + rw [← this, Finset.sum_map] + congr + _ = n + (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by simp + _ = n + 1 := by rw [sum_mod_range hK] diff --git a/LeanBandits/BanditAlgorithms/ETC.lean b/LeanBandits/BanditAlgorithms/ETC.lean index 048cb41c..14445cd2 100644 --- a/LeanBandits/BanditAlgorithms/ETC.lean +++ b/LeanBandits/BanditAlgorithms/ETC.lean @@ -3,9 +3,9 @@ Copyright (c) 2025 Rémy Degenne. All rights reserved. Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne -/ +import LeanBandits.Bandit.SumRewards +import LeanBandits.BanditAlgorithms.AuxSums import LeanBandits.ForMathlib.MeasurableArgMax -import LeanBandits.ForMathlib.SubGaussian -import LeanBandits.RewardByCountMeasure /-! # The Explore-Then-Commit Algorithm @@ -22,58 +22,6 @@ lemma ae_eq_set_iff {α : Type*} {mα : MeasurableSpace α} {μ : Measure α} {s simp only [eq_iff_iff] congr! ---todo: generalize Icc -lemma measurable_sum_of_le {α : Type*} {mα : MeasurableSpace α} - {f : ℕ → α → ℝ} {g : α → ℕ} {n : ℕ} (hg_le : ∀ a, g a ≤ n) (hf : ∀ i, Measurable (f i)) - (hg : Measurable g) : - Measurable (fun a ↦ ∑ i ∈ Icc 1 (g a), f i a) := by - have h_eq : (fun a ↦ ∑ i ∈ Icc 1 (g a), f i a) - = fun a ↦ ∑ i ∈ range (n + 1), if g a = i then ∑ j ∈ Icc 1 i, f j a else 0 := by - ext ω - rw [sum_ite_eq_of_mem] - grind - rw [h_eq] - refine measurable_sum _ fun n hn ↦ ?_ - refine Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact (measurableSet_singleton _).preimage (by fun_prop) - -lemma sum_mod_range {K : ℕ} (hK : 0 < K) (a : Fin K) : - (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) = 1 := by - have h_iff (s : ℕ) (hs : s < K) : ⟨s % K, Nat.mod_lt _ hK⟩ = a ↔ s = a := by - simp only [Nat.mod_eq_of_lt hs, Fin.ext_iff] - calc (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) - _ = ∑ s ∈ range K, if s = a then 1 else 0 := sum_congr rfl fun s hs ↦ by grind - _ = _ := by - rw [sum_ite_eq'] - simp - -lemma sum_mod_range_mul {K : ℕ} (hK : 0 < K) (m : ℕ) (a : Fin K) : - (∑ s ∈ range (K * m), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) = m := by - induction m with - | zero => simp - | succ n hn => - calc (∑ s ∈ range (K * (n + 1)), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) - _ = (∑ s ∈ range (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by ring_nf - _ = (∑ s ∈ range (K * n), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) - + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by - rw [sum_range_add_sum_Ico] - grind - _ = n + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by - rw [hn] - _ = n + (∑ s ∈ range K, if ⟨(s + K * n) % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by - congr 1 - let e : ℕ ↪ ℕ := ⟨fun i : ℕ ↦ i + K * n, fun i j hij ↦ by grind⟩ - have : Finset.map e (range K) = Ico (K * n) (K * n + K) := by - ext x - simp only [mem_map, mem_range, Function.Embedding.coeFn_mk, mem_Ico, e] - refine ⟨fun h ↦ by grind, fun h ↦ ?_⟩ - use x - K * n - grind - rw [← this, Finset.sum_map] - congr - _ = n + (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by simp - _ = n + 1 := by rw [sum_mod_range hK] - end Aux namespace Bandits @@ -116,55 +64,64 @@ end AlgorithmDefinition namespace ETC variable {hK : 0 < K} {m : ℕ} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] + {Ω : Type*} {mΩ : MeasurableSpace Ω} + {P : Measure Ω} [IsProbabilityMeasure P] + {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} -local notation "𝔓t" => Bandit.trajMeasure (etcAlgorithm hK m) ν -local notation "𝔓" => Bandit.measure (etcAlgorithm hK m) ν - -lemma arm_zero : arm 0 =ᵐ[𝔓t] fun _ ↦ ⟨0, hK⟩ := by +lemma arm_zero [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) : + A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - exact arm_zero_detAlgorithm + exact h.action_zero_detAlgorithm -lemma arm_ae_eq_etcNextArm (n : ℕ) : - arm (n + 1) =ᵐ[𝔓t] fun h ↦ nextArm hK m n (fun i ↦ h i) := by +lemma arm_ae_eq_etcNextArm [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (n : ℕ) : + A (n + 1) =ᵐ[P] fun ω ↦ nextArm hK m n (IsAlgEnvSeq.hist A R n ω) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - exact arm_detAlgorithm_ae_eq n + exact h.action_detAlgorithm_ae_eq n /-- For `n < K * m`, the arm pulled at time `n` is the arm `n % K`. -/ -lemma arm_of_lt {n : ℕ} (hn : n < K * m) : - arm n =ᵐ[𝔓t] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := by +lemma arm_of_lt [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) {n : ℕ} (hn : n < K * m) : + A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := by cases n with - | zero => exact arm_zero + | zero => exact arm_zero h | succ n => - filter_upwards [arm_ae_eq_etcNextArm n] with h hn_eq + filter_upwards [arm_ae_eq_etcNextArm h n] with h hn_eq rw [hn_eq, nextArm, dif_pos] grind /-- The arm pulled at time `K * m` is the arm with the highest empirical mean after the exploration phase. -/ -lemma arm_mul (hm : m ≠ 0) : - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - arm (K * m) =ᵐ[𝔓t] fun h ↦ measurableArgmax (empMean' (K * m - 1)) (fun i ↦ h i) := by +lemma arm_mul [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (hm : m ≠ 0) : + A (K * m) =ᵐ[P] fun ω ↦ measurableArgmax (empMean' (K * m - 1)) + (IsAlgEnvSeq.hist A R (K * m - 1) ω) := by have : K * m = (K * m - 1) + 1 := by have : 0 < K * m := Nat.mul_pos hK hm.bot_lt grind rw [this] - filter_upwards [arm_ae_eq_etcNextArm (K * m - 1)] with h hn_eq + filter_upwards [arm_ae_eq_etcNextArm h (K * m - 1)] with ω hn_eq rw [hn_eq, nextArm, dif_neg (by simp), dif_pos rfl] exact this ▸ rfl /-- For `n ≥ K * m`, the arm pulled at time `n + 1` is the same as the arm pulled at time `n`. -/ -lemma arm_add_one_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : - arm (n + 1) =ᵐ[𝔓t] fun ω ↦ arm n ω := by - filter_upwards [arm_ae_eq_etcNextArm n] with ω hn_eq +lemma arm_add_one_of_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : + A (n + 1) =ᵐ[P] fun ω ↦ A n ω := by + filter_upwards [arm_ae_eq_etcNextArm h n] with ω hn_eq rw [hn_eq, nextArm, dif_neg (by grind), dif_neg] · rfl · have : 0 < K * m := Nat.mul_pos hK hm.bot_lt grind /-- For `n ≥ K * m`, the arm pulled at time `n` is the same as the arm pulled at time `K * m`. -/ -lemma arm_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : - arm n =ᵐ[𝔓t] arm (K * m) := by - have h_ae n : K * m ≤ n → arm (n + 1) =ᵐ[𝔓t] fun ω ↦ arm n ω := arm_add_one_of_ge hm +lemma arm_of_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : + A n =ᵐ[P] A (K * m) := by + have h_ae n : K * m ≤ n → A (n + 1) =ᵐ[P] fun ω ↦ A n ω := arm_add_one_of_ge h hm simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_ae filter_upwards [h_ae] with ω h_ae induction n, hn using Nat.le_induction with @@ -172,36 +129,42 @@ lemma arm_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : | succ n hmn h_ind => rw [h_ae n hmn, h_ind] /-- At time `K * m`, the number of pulls of each arm is equal to `m`. -/ -lemma pullCount_mul (a : Fin K) : pullCount a (K * m) =ᵐ[𝔓t] fun _ ↦ m := by +lemma pullCount_mul [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (a : Fin K) : + pullCount A a (K * m) =ᵐ[P] fun _ ↦ m := by rw [Filter.EventuallyEq] simp_rw [pullCount_eq_sum] - have h_arm (n : range (K * m)) : arm n =ᵐ[𝔓t] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := - arm_of_lt (mem_range.mp n.2) + have h_arm (n : range (K * m)) : A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := + arm_of_lt h (mem_range.mp n.2) simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_arm filter_upwards [h_arm] with ω h_arm - have h_arm' {i : ℕ} (hi : i ∈ range (K * m)) : arm i ω = ⟨i % K, Nat.mod_lt _ hK⟩ := h_arm ⟨i, hi⟩ - calc (∑ s ∈ range (K * m), if arm s ω = a then 1 else 0) + have h_arm' {i : ℕ} (hi : i ∈ range (K * m)) : A i ω = ⟨i % K, Nat.mod_lt _ hK⟩ := h_arm ⟨i, hi⟩ + calc (∑ s ∈ range (K * m), if A s ω = a then 1 else 0) _ = (∑ s ∈ range (K * m), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := sum_congr rfl fun s hs ↦ by rw [h_arm' hs] _ = m := sum_mod_range_mul hK m a -lemma pullCount_add_one_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : - pullCount a (n + 1) - =ᵐ[𝔓t] fun ω ↦ pullCount a n ω + {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by +lemma pullCount_add_one_of_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : + pullCount A a (n + 1) + =ᵐ[P] fun ω ↦ pullCount A a n ω + {ω' | A (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by simp_rw [Filter.EventuallyEq, pullCount_add_one] - filter_upwards [arm_of_ge hm hn] with ω h_arm + filter_upwards [arm_of_ge h hm hn] with ω h_arm congr 3 /-- For `n ≥ K * m`, the number of pulls of each arm `a` at time `n` is equal to `m` plus `n - K * m` if arm `a` is the best arm after the exploration phase. -/ -lemma pullCount_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : - pullCount a n - =ᵐ[𝔓t] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by - have h_ae n : K * m ≤ n → pullCount a (n + 1) - =ᵐ[𝔓t] fun ω ↦ pullCount a n ω + {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := - pullCount_add_one_of_ge a hm +lemma pullCount_of_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : + pullCount A a n + =ᵐ[P] fun ω ↦ m + (n - K * m) * {ω' | A (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by + have h_ae n : K * m ≤ n → pullCount A a (n + 1) + =ᵐ[P] fun ω ↦ pullCount A a n ω + {ω' | A (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := + pullCount_add_one_of_ge h a hm simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_ae - have h_ae_Km : pullCount a (K * m) =ᵐ[𝔓t] fun _ ↦ m := pullCount_mul a + have h_ae_Km : pullCount A a (K * m) =ᵐ[P] fun _ ↦ m := pullCount_mul h a filter_upwards [h_ae_Km, h_ae] with ω h_Km h_ae induction n, hn using Nat.le_induction with | base => simp [h_Km] @@ -212,13 +175,14 @@ lemma pullCount_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : /-- If at time `K * m` the algorithm chooses arm `a`, then the total reward obtained by pulling arm `a` is at least the total reward obtained by pulling the best arm. -/ -lemma sumRewards_bestArm_le_of_arm_mul_eq (a : Fin K) (hm : m ≠ 0) : - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - ∀ᵐ h ∂𝔓t, arm (K * m) h = a → sumRewards (bestArm ν) (K * m) h ≤ sumRewards a (K * m) h := by - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - filter_upwards [arm_mul hm, pullCount_mul a, pullCount_mul (bestArm ν)] with h h_arm ha h_best - h_eq - have h_max := isMaxOn_measurableArgmax (empMean' (K * m - 1)) (fun i ↦ h i) (bestArm ν) +lemma sumRewards_bestArm_le_of_arm_mul_eq [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (a : Fin K) (hm : m ≠ 0) : + ∀ᵐ h ∂P, A (K * m) h = a → sumRewards A R (bestArm ν) (K * m) h ≤ + sumRewards A R a (K * m) h := by + filter_upwards [arm_mul h hm, pullCount_mul h a, pullCount_mul h (bestArm ν)] + with h h_arm ha h_best h_eq + have h_max := isMaxOn_measurableArgmax (empMean' (K * m - 1)) (IsAlgEnvSeq.hist A R (K * m - 1) h) + (bestArm ν) rw [← h_arm, h_eq] at h_max rw [sumRewards_eq_pullCount_mul_empMean, sumRewards_eq_pullCount_mul_empMean, ha, h_best] · gcongr @@ -227,150 +191,52 @@ lemma sumRewards_bestArm_le_of_arm_mul_eq (a : Fin K) (hm : m ≠ 0) : · simp [ha, hm] · simp [h_best, hm] -lemma identDistrib_aux (m : ℕ) (a b : Fin K) : - IdentDistrib - (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount a s ω, ∑ s ∈ Icc 1 m, rewardByCount b s ω)) - (fun ω ↦ (∑ s ∈ range m, ω.2 s a, ∑ s ∈ range m, ω.2 s b)) 𝔓 𝔓 := by - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - have h2 (a : Fin K) : IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount a s ω) - (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓 𝔓 := identDistrib_sum_Icc_rewardByCount m a - by_cases hab : a = b - · simp only [hab] - exact (h2 b).comp (u := fun p ↦ (p, p)) (by fun_prop) - refine (h2 a).prodMk (h2 b) ?_ ?_ - · suffices IndepFun (fun ω s ↦ rewardByCount a s ω) (fun ω s ↦ rewardByCount b s ω) - 𝔓 by - exact this.comp (φ := fun p ↦ ∑ i ∈ Icc 1 m, p i) (ψ := fun p ↦ ∑ j ∈ Icc 1 m, p j) - (by fun_prop) (by fun_prop) - exact indepFun_rewardByCount_of_ne hab - · suffices IndepFun (fun ω s ↦ ω.2 s a) (fun ω s ↦ ω.2 s b) 𝔓 by - exact this.comp (φ := fun p ↦ ∑ i ∈ range m, p i) (ψ := fun p ↦ ∑ j ∈ range m, p j) - (by fun_prop) (by fun_prop) - exact indepFun_eval_snd_measure _ ν hab +lemma probReal_sumRewards_le_sumRewards_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) : + P.real {ω | sumRewards A R (bestArm ν) (K * m) ω ≤ sumRewards A R a (K * m) ω} ≤ + Real.exp (-↑m * gap ν a ^ 2 / 4) := by + have hA := h.measurable_A + have hR := h.measurable_R + have h1 := Bandits.probReal_sumRewards_le_sumRewards_le h a (K * m) m m + have h2 := probReal_sum_le_sum_streamMeasure hν a m + refine le_trans (le_of_eq ?_) (h1.trans h2) + simp_rw [measureReal_def] + congr 1 + refine measure_congr ?_ + rw [ae_eq_set_iff] + filter_upwards [pullCount_mul h a, pullCount_mul h (bestArm ν)] with ω ha h_best + simp [ha, h_best] /-- The probability that at time `K * m` the ETC algorithm chooses arm `a` is at most `exp(- m * Δ_a^2 / 4)`. -/ -lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) +lemma prob_arm_mul_eq_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) (hm : m ≠ 0) : - (𝔓t).real {ω | arm (K * m) ω = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + P.real {ω | A (K * m) ω = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by have h_pos : 0 < K * m := Nat.mul_pos hK hm.bot_lt - have h_le : (𝔓t).real {ω | arm (K * m) ω = a} - ≤ (𝔓t).real {ω | sumRewards (bestArm ν) (K * m) ω ≤ sumRewards a (K * m) ω} := by + have h_le : P.real {ω | A (K * m) ω = a} + ≤ P.real {ω | sumRewards A R (bestArm ν) (K * m) ω ≤ sumRewards A R a (K * m) ω} := by simp_rw [measureReal_def] gcongr 1 · simp refine measure_mono_ae ?_ - exact sumRewards_bestArm_le_of_arm_mul_eq a hm - refine h_le.trans ?_ - -- extend the probability space to include the stream of independent rewards - suffices (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} - ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) by - suffices (𝔓t).real {ω | sumRewards (bestArm ν) (K * m) ω ≤ sumRewards a (K * m) ω} - = (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} by rwa [this] - calc (𝔓t).real {ω | sumRewards (bestArm ν) (K * m) ω ≤ sumRewards a (K * m) ω} - _ = ((𝔓).fst).real {ω | sumRewards (bestArm ν) (K * m) ω ≤ sumRewards a (K * m) ω} := by simp - _ = (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} := by - rw [Measure.fst, map_measureReal_apply (by fun_prop)] - · rfl - · exact measurableSet_le (by fun_prop) (by fun_prop) - calc (𝔓).real {ω | sumRewards (bestArm ν) (K * m) ω.1 ≤ sumRewards a (K * m) ω.1} - _ = (𝔓).real {ω | ∑ s ∈ Icc 1 (pullCount (bestArm ν) (K * m) ω.1), rewardByCount (bestArm ν) s ω - ≤ ∑ s ∈ Icc 1 (pullCount a (K * m) ω.1), rewardByCount a s ω} := by - congr with ω - congr! 1 <;> rw [sum_rewardByCount_eq_sumRewards] - _ = (𝔓).real {ω | ∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω - ≤ ∑ s ∈ Icc 1 m, rewardByCount a s ω} := by - simp_rw [measureReal_def] - congr 1 - refine measure_congr ?_ - have ha := pullCount_mul a (hK := hK) (ν := ν) (m := m) - have h_best := pullCount_mul (bestArm ν) (hK := hK) (ν := ν) (m := m) - rw [ae_eq_set_iff] - change ∀ᵐ ω ∂((𝔓t).prod _), _ - rw [Measure.ae_prod_iff_ae_ae] - · filter_upwards [ha, h_best] with ω ha h_best - refine ae_of_all _ fun ω' ↦ ?_ - rw [ha, h_best] - · simp only [Set.mem_setOf_eq] - let f₁ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ - ∑ s ∈ Icc 1 (pullCount (bestArm ν) (K * m) ω.1), rewardByCount (bestArm ν) s ω - let g₁ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ - ∑ s ∈ Icc 1 (pullCount a (K * m) ω.1), rewardByCount a s ω - let f₂ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ - ∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω - let g₂ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ ∑ s ∈ Icc 1 m, rewardByCount a s ω - change MeasurableSet {x | f₁ x ≤ g₁ x ↔ f₂ x ≤ g₂ x} - have hf₁ : Measurable f₁ := by - refine measurable_sum_of_le (n := K * m + 1) - (g := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ pullCount (bestArm ν) (K * m) ω.1) - (f := rewardByCount (bestArm ν)) (fun ω ↦ ?_) - (by fun_prop) (by fun_prop) - have h_le := pullCount_le (bestArm ν) (K * m) ω.1 - grind - have hg₁ : Measurable g₁ := by - refine measurable_sum_of_le (n := K * m + 1) - (g := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ pullCount a (K * m) ω.1) - (f := rewardByCount a) (fun ω ↦ ?_) (by fun_prop) (by fun_prop) - have h_le := pullCount_le a (K * m) ω.1 - grind - refine MeasurableSet.iff ?_ ?_ - · exact measurableSet_le (by fun_prop) (by fun_prop) - · exact measurableSet_le (by fun_prop) (by fun_prop) - _ = (𝔓).real {ω | ∑ s ∈ range m, ω.2 s (bestArm ν) ≤ ∑ s ∈ range m, ω.2 s a} := by - simp_rw [measureReal_def] - congr 1 - have : (𝔓).map (fun ω ↦ (∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω, - ∑ s ∈ Icc 1 m, rewardByCount a s ω)) - = (𝔓).map (fun ω ↦ (∑ s ∈ range m, ω.2 s (bestArm ν), ∑ s ∈ range m, ω.2 s a)) := - (identDistrib_aux m (bestArm ν) a).map_eq - rw [Measure.ext_iff] at this - have h_meas : MeasurableSet {x : ℝ × ℝ | x.1 ≤ x.2} := - measurableSet_le (by fun_prop) (by fun_prop) - specialize this {x | x.1 ≤ x.2} h_meas - rwa [Measure.map_apply (by fun_prop) h_meas, Measure.map_apply (by fun_prop) h_meas] at this - _ = (Bandit.streamMeasure ν).real - {ω | ∑ s ∈ range m, ω s (bestArm ν) ≤ ∑ s ∈ range m, ω s a} := by - simp_rw [measureReal_def] - congr 1 - rw [← Bandit.snd_measure (etcAlgorithm hK m), Measure.snd_apply] - · rfl - · exact measurableSet_le (by fun_prop) (by fun_prop) - _ ≤ Real.exp (-↑m * gap ν a ^ 2 / 4) := by - by_cases ha : a = bestArm ν - · simp [ha] - -- Apply a sub-Gaussian concentration inequality - refine (HasSubgaussianMGF.measure_sum_le_sum_le' (cX := fun _ ↦ 1) (cY := fun _ ↦ 1) - ?_ ?_ ?_ ?_ ?_ ?_).trans_eq ?_ - · exact iIndepFun_eval_streamMeasure'' ν (bestArm ν) - · exact iIndepFun_eval_streamMeasure'' ν a - · intro i him - simp_rw [integral_eval_streamMeasure] - refine (hν (bestArm ν)).congr_identDistrib ?_ - exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ - · intro i him - simp_rw [integral_eval_streamMeasure] - refine (hν a).congr_identDistrib ?_ - exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ - · exact indepFun_eval_streamMeasure' ν (Ne.symm ha) - · gcongr 1 with i him - simp_rw [integral_eval_streamMeasure] - exact le_bestArm a - · congr 1 - simp_rw [integral_eval_streamMeasure] - simp only [id_eq, sum_const, card_range, nsmul_eq_mul, mul_one, NNReal.coe_natCast, - gap_eq_bestArm_sub, neg_mul] - field_simp - ring + exact sumRewards_bestArm_le_of_arm_mul_eq h a hm + exact h_le.trans (probReal_sumRewards_le_sumRewards_le h hν a) /-- Bound on the expectation of the number of pulls of each arm by the ETC algorithm. -/ -lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma expectation_pullCount_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : - 𝔓t[fun ω ↦ (pullCount a n ω : ℝ)] + P[fun ω ↦ (pullCount A a n ω : ℝ)] ≤ m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by - have : (fun ω ↦ (pullCount a n ω : ℝ)) - =ᵐ[𝔓t] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by - filter_upwards [pullCount_of_ge a hm hn] with ω h + have hA := h.measurable_A + have hR := h.measurable_R + have : (fun ω ↦ (pullCount A a n ω : ℝ)) + =ᵐ[P] fun ω ↦ m + (n - K * m) * {ω' | A (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by + filter_upwards [pullCount_of_ge h a hm hn] with ω h simp only [h, Set.indicator_apply, Set.mem_setOf_eq, mul_ite, mul_one, mul_zero, Nat.cast_add, Nat.cast_ite, CharP.cast_eq_zero, add_right_inj] norm_cast @@ -386,21 +252,25 @@ lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - ( simp rw [integral_indicator_const, smul_eq_mul, mul_one] · rw [← neg_mul] - exact prob_arm_mul_eq_le hν a hm + exact prob_arm_mul_eq_le h hν a hm · exact (measurableSet_singleton _).preimage (by fun_prop) /-- Regret bound for the ETC algorithm. -/ -lemma regret_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hm : m ≠ 0) +lemma regret_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hm : m ≠ 0) (n : ℕ) (hn : K * m ≤ n) : - 𝔓t[regret ν n] ≤ ∑ a, gap ν a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4)) := by + P[regret ν A n] ≤ + ∑ a, gap ν a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4)) := by + have hA := h.measurable_A simp_rw [regret_eq_sum_pullCount_mul_gap] rw [integral_finset_sum] - swap; · exact fun i _ ↦ (integrable_pullCount i n).mul_const _ + swap; · exact fun i _ ↦ (integrable_pullCount hA i n).mul_const _ gcongr with a rw [mul_comm (gap _ _), integral_mul_const] gcongr · exact gap_nonneg - · exact expectation_pullCount_le hν a hm hn + · exact expectation_pullCount_le h hν a hm hn end ETC diff --git a/LeanBandits/BanditAlgorithms/UCB.lean b/LeanBandits/BanditAlgorithms/UCB.lean index 0fd5462b..121bb0e8 100644 --- a/LeanBandits/BanditAlgorithms/UCB.lean +++ b/LeanBandits/BanditAlgorithms/UCB.lean @@ -3,7 +3,9 @@ Copyright (c) 2025 Rémy Degenne. All rights reserved. Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne -/ -import LeanBandits.BanditAlgorithms.ETC +import LeanBandits.Bandit.SumRewards +import LeanBandits.BanditAlgorithms.AuxSums +import LeanBandits.ForMathlib.MeasurableArgMax /-! # UCB algorithm @@ -50,143 +52,159 @@ end Algorithm namespace UCB -variable {hK : 0 < K} {c : ℝ} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] {n : ℕ} {h : ℕ → Fin K × ℝ} +variable {hK : 0 < K} {c : ℝ} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] + {Ω : Type*} {mΩ : MeasurableSpace Ω} + {P : Measure Ω} [IsProbabilityMeasure P] + {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} + {n : ℕ} {ω : Ω} /-- The exploration bonus of the UCB algorithm, which corresponds to the width of a confidence interval. -/ -noncomputable def ucbWidth (c : ℝ) (a : Fin K) (n : ℕ) (h : ℕ → Fin K × ℝ) : ℝ := - √(c * log (n + 1) / pullCount a n h) +noncomputable def ucbWidth (A : ℕ → Ω → Fin K) (c : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ := + √(c * log (n + 1) / pullCount A a n ω) @[fun_prop] -lemma measurable_ucbWidth (c : ℝ) (a : Fin K) : Measurable (ucbWidth c a n) := by +lemma measurable_ucbWidth (hA : ∀ n, Measurable (A n)) (c : ℝ) (a : Fin K) : + Measurable (ucbWidth A c a n) := by unfold ucbWidth fun_prop -lemma ucbWidth_eq_ucbWidth' (c : ℝ) (a : Fin K) (n : ℕ) (h : ℕ → Fin K × ℝ) (hn : n ≠ 0) : - ucbWidth c a n h = ucbWidth' c (n - 1) (fun i ↦ h i) a := by - simp only [ucbWidth, pullCount_eq_pullCount' hn, Nat.cast_nonneg, sqrt_div', ucbWidth'] +lemma ucbWidth_eq_ucbWidth' (c : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) (hn : n ≠ 0) : + ucbWidth A c a n ω = ucbWidth' c (n - 1) (IsAlgEnvSeq.hist A R (n - 1) ω) a := by + simp only [ucbWidth, pullCount_eq_pullCount' (A := A) (R' := R) hn, Nat.cast_nonneg, sqrt_div', + ucbWidth'] congr 4 norm_cast grind -local notation "𝔓t" => Bandit.trajMeasure (ucbAlgorithm hK c) ν -local notation "𝔓" => Bandit.measure (ucbAlgorithm hK c) ν - -lemma arm_zero : arm 0 =ᵐ[𝔓t] fun _ ↦ ⟨0, hK⟩ := by +lemma arm_zero [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) : + A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - exact arm_zero_detAlgorithm + exact h.action_zero_detAlgorithm -lemma arm_ae_eq_ucbNextArm (n : ℕ) : - arm (n + 1) =ᵐ[𝔓t] fun h ↦ nextArm hK c n (fun i ↦ h i) := by +lemma arm_ae_eq_ucbNextArm [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (n : ℕ) : + A (n + 1) =ᵐ[P] fun ω ↦ nextArm hK c n (IsAlgEnvSeq.hist A R n ω) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - exact arm_detAlgorithm_ae_eq n + exact h.action_detAlgorithm_ae_eq n -lemma arm_ae_all_eq : - ∀ᵐ h ∂𝔓t, arm 0 h = ⟨0, hK⟩ ∧ ∀ n, arm (n + 1) h = nextArm hK c n (fun i ↦ h i) := by +lemma arm_ae_all_eq [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) : + ∀ᵐ h ∂P, A 0 h = ⟨0, hK⟩ ∧ ∀ n, A (n + 1) h = nextArm hK c n (IsAlgEnvSeq.hist A R n h) := by rw [eventually_and, ae_all_iff] - exact ⟨arm_zero, arm_ae_eq_ucbNextArm⟩ + exact ⟨arm_zero h, arm_ae_eq_ucbNextArm h⟩ -lemma ucbIndex_le_ucbIndex_arm (a : Fin K) (hn : K ≤ n) : - ∀ᵐ h ∂𝔓t, empMean a n h + ucbWidth c a n h ≤ - empMean (arm n h) n h + ucbWidth c (arm n h) n h := by - filter_upwards [arm_ae_eq_ucbNextArm (n - 1)] with h h_arm +lemma ucbIndex_le_ucbIndex_arm [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) (hn : K ≤ n) : + ∀ᵐ h ∂P, empMean A R a n h + ucbWidth A c a n h ≤ + empMean A R (A n h) n h + ucbWidth A c (A n h) n h := by + filter_upwards [arm_ae_eq_ucbNextArm h (n - 1)] with h h_arm have : n - 1 + 1 = n := by grind have h_not_lt : ¬ n - 1 < K - 1 := by grind simp only [this, nextArm, h_not_lt, ↓reduceIte] at h_arm have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK simp_rw [h_arm, empMean_eq_empMean' (by grind : n ≠ 0), - ucbWidth_eq_ucbWidth' _ _ _ _ (by grind : n ≠ 0)] + ucbWidth_eq_ucbWidth' (A := A) (R := R) _ _ _ _ (by grind : n ≠ 0)] exact isMaxOn_measurableArgmax (fun h a ↦ empMean' (n - 1) h a + ucbWidth' c (n - 1) h a) - (fun i ↦ h i) a + (IsAlgEnvSeq.hist A R (n - 1) h) a -lemma forall_arm_eq_mod_of_lt : - ∀ᵐ h ∂𝔓t, ∀ n < K, arm n h = ⟨n % K, Nat.mod_lt _ hK⟩ := by +lemma forall_arm_eq_mod_of_lt [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) : + ∀ᵐ h ∂P, ∀ n < K, A n h = ⟨n % K, Nat.mod_lt _ hK⟩ := by simp_rw [ae_all_iff] intro n hn induction n with - | zero => exact arm_zero + | zero => exact arm_zero h | succ n _ => - filter_upwards [arm_ae_eq_ucbNextArm n] with h h_eq + filter_upwards [arm_ae_eq_ucbNextArm h n] with h h_eq rw [h_eq, nextArm, if_pos] grind -lemma forall_ucbIndex_le_ucbIndex_arm (a : Fin K) : - ∀ᵐ h ∂𝔓t, ∀ n, K ≤ n → - empMean a n h + ucbWidth c a n h ≤ empMean (arm n h) n h + ucbWidth c (arm n h) n h := by +lemma forall_ucbIndex_le_ucbIndex_arm [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) : + ∀ᵐ h ∂P, ∀ n, K ≤ n → + empMean A R a n h + ucbWidth A c a n h ≤ + empMean A R (A n h) n h + ucbWidth A c (A n h) n h := by simp_rw [ae_all_iff] - exact fun _ ↦ ucbIndex_le_ucbIndex_arm a - -lemma forall_arm_prop : - ∀ᵐ h ∂𝔓t, - (∀ n < K, arm n h = ⟨n % K, Nat.mod_lt _ hK⟩) ∧ - (∀ n, K ≤ n → ∀ a, empMean a n h + ucbWidth c a n h ≤ - empMean (arm n h) n h + ucbWidth c (arm n h) n h) := by + exact fun _ ↦ ucbIndex_le_ucbIndex_arm h a + +lemma forall_arm_prop [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) : + ∀ᵐ h ∂P, + (∀ n < K, A n h = ⟨n % K, Nat.mod_lt _ hK⟩) ∧ + (∀ n, K ≤ n → ∀ a, empMean A R a n h + ucbWidth A c a n h ≤ + empMean A R (A n h) n h + ucbWidth A c (A n h) n h) := by simp only [eventually_and] constructor - · exact forall_arm_eq_mod_of_lt + · exact forall_arm_eq_mod_of_lt h · simp_rw [ae_all_iff] intro n hn a - have h_ae := forall_ucbIndex_le_ucbIndex_arm (ν := ν) (c := c) (hK := hK) a + have h_ae := forall_ucbIndex_le_ucbIndex_arm h a simp_rw [ae_all_iff] at h_ae exact h_ae n hn -lemma pullCount_eq_of_time_eq (a : Fin K) : - ∀ᵐ ω ∂𝔓t, pullCount a K ω = 1 := by - filter_upwards [forall_arm_eq_mod_of_lt] with h h_eq +lemma pullCount_eq_of_time_eq [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) : + ∀ᵐ ω ∂P, pullCount A a K ω = 1 := by + filter_upwards [forall_arm_eq_mod_of_lt h] with h h_eq rw [pullCount_eq_sum] conv_rhs => rw [← sum_mod_range hK a] refine Finset.sum_congr rfl fun s hs ↦ ?_ congr exact h_eq s (by grind) -lemma time_gt_of_pullCount_gt_one (a : Fin K) : - ∀ᵐ ω ∂𝔓t, ∀ n, 1 < pullCount a n ω → K < n := by - filter_upwards [pullCount_eq_of_time_eq a] with h h_eq n hn +lemma time_gt_of_pullCount_gt_one [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) : + ∀ᵐ ω ∂P, ∀ n, 1 < pullCount A a n ω → K < n := by + filter_upwards [pullCount_eq_of_time_eq h a] with h h_eq n hn rw [← h_eq] at hn by_contra! h_lt exact hn.not_ge (pullCount_mono _ h_lt _) -lemma pullCount_pos_of_time_ge : - ∀ᵐ ω ∂𝔓t, ∀ n, K ≤ n → ∀ b : Fin K, 0 < pullCount b n ω := by - have h_ae a := pullCount_eq_of_time_eq (ν := ν) (c := c) (hK := hK) a +lemma pullCount_pos_of_time_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) : + ∀ᵐ ω ∂P, ∀ n, K ≤ n → ∀ b : Fin K, 0 < pullCount A b n ω := by + have h_ae a := pullCount_eq_of_time_eq h a rw [← ae_all_iff] at h_ae filter_upwards [h_ae] with ω hω n hn a refine Nat.one_pos.trans_le ?_ rw [← hω a] exact pullCount_mono _ hn _ -lemma pullCount_pos_of_pullCount_gt_one (a : Fin K) : - ∀ᵐ ω ∂𝔓t, ∀ n, 1 < pullCount a n ω → ∀ b : Fin K, 0 < pullCount b n ω := by - filter_upwards [time_gt_of_pullCount_gt_one a, pullCount_pos_of_time_ge] with ω h1 h2 n h_gt a +lemma pullCount_pos_of_pullCount_gt_one [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) : + ∀ᵐ ω ∂P, ∀ n, 1 < pullCount A a n ω → ∀ b : Fin K, 0 < pullCount A b n ω := by + filter_upwards [time_gt_of_pullCount_gt_one h a, pullCount_pos_of_time_ge h] with ω h1 h2 n h_gt a exact h2 n (h1 n h_gt).le a omit [IsMarkovKernel ν] in lemma gap_arm_le_two_mul_ucbWidth [Nonempty (Fin K)] - (h_best : (ν (bestArm ν))[id] ≤ empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h) - (h_arm : empMean (arm n h) n h - ucbWidth c (arm n h) n h ≤ (ν (arm n h))[id]) - (h_le : empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h ≤ - empMean (arm n h) n h + ucbWidth c (arm n h) n h) : - gap ν (arm n h) ≤ 2 * ucbWidth c (arm n h) n h := by + (h_best : (ν (bestArm ν))[id] ≤ empMean A R (bestArm ν) n ω + ucbWidth A c (bestArm ν) n ω) + (h_arm : empMean A R (A n ω) n ω - ucbWidth A c (A n ω) n ω ≤ (ν (A n ω))[id]) + (h_le : empMean A R (bestArm ν) n ω + ucbWidth A c (bestArm ν) n ω ≤ + empMean A R (A n ω) n ω + ucbWidth A c (A n ω) n ω) : + gap ν (A n ω) ≤ 2 * ucbWidth A c (A n ω) n ω := by rw [gap_eq_bestArm_sub, sub_le_iff_le_add'] calc (ν (bestArm ν))[id] - _ ≤ empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h := h_best - _ ≤ empMean (arm n h) n h + ucbWidth c (arm n h) n h := h_le - _ ≤ (ν (arm n h))[id] + 2 * ucbWidth c (arm n h) n h := by + _ ≤ empMean A R (bestArm ν) n ω + ucbWidth A c (bestArm ν) n ω := h_best + _ ≤ empMean A R (A n ω) n ω + ucbWidth A c (A n ω) n ω := h_le + _ ≤ (ν (A n ω))[id] + 2 * ucbWidth A c (A n ω) n ω := by rw [two_mul, ← add_assoc] gcongr rwa [sub_le_iff_le_add] at h_arm omit [IsMarkovKernel ν] in lemma pullCount_arm_le [Nonempty (Fin K)] (hc : 0 ≤ c) - (h_best : (ν (bestArm ν))[id] ≤ empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h) - (h_arm : empMean (arm n h) n h - ucbWidth c (arm n h) n h ≤ (ν (arm n h))[id]) - (h_le : empMean (bestArm ν) n h + ucbWidth c (bestArm ν) n h ≤ - empMean (arm n h) n h + ucbWidth c (arm n h) n h) - (h_gap_pos : 0 < gap ν (arm n h)) (h_pull_pos : 0 < pullCount (arm n h) n h) : - pullCount (arm n h) n h ≤ 4 * c * log (n + 1) / gap ν (arm n h) ^ 2 := by + (h_best : (ν (bestArm ν))[id] ≤ empMean A R (bestArm ν) n ω + ucbWidth A c (bestArm ν) n ω) + (h_arm : empMean A R (A n ω) n ω - ucbWidth A c (A n ω) n ω ≤ (ν (A n ω))[id]) + (h_le : empMean A R (bestArm ν) n ω + ucbWidth A c (bestArm ν) n ω ≤ + empMean A R (A n ω) n ω + ucbWidth A c (A n ω) n ω) + (h_gap_pos : 0 < gap ν (A n ω)) (h_pull_pos : 0 < pullCount A (A n ω) n ω) : + pullCount A (A n ω) n ω ≤ 4 * c * log (n + 1) / gap ν (A n ω) ^ 2 := by have h_gap_le := gap_arm_le_two_mul_ucbWidth h_best h_arm h_le rw [ucbWidth] at h_gap_le - have h2 : (gap ν (arm n h)) ^ 2 ≤ (2 * √(c * log (n + 1) / pullCount (arm n h) n h)) ^ 2 := by + have h2 : (gap ν (A n ω)) ^ 2 ≤ (2 * √(c * log (n + 1) / pullCount A (A n ω) n ω)) ^ 2 := by gcongr rw [mul_pow, sq_sqrt] at h2 · have : (2 : ℝ) ^ 2 = 4 := by norm_num @@ -198,154 +216,76 @@ lemma pullCount_arm_le [Nonempty (Fin K)] (hc : 0 ≤ c) lemma todo (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n k : ℕ) (hk : k ≠ 0) : - 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} ≤ + Bandit.streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2) := by - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - have h_meas : MeasurableSet {ω | ω / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := - measurableSet_le (by fun_prop) measurable_const have h_log_nonneg : 0 ≤ log (n + 1) := log_nonneg (by simp) - calc - 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} - _ = ((𝔓).map (fun ω ↦ ∑ m ∈ Icc 1 k, rewardByCount a m ω)) - {ω | ω / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := by - rw [Measure.map_apply (by fun_prop) h_meas] - rfl - _ = ((𝔓).map (fun ω ↦ ∑ s ∈ range k, ω.2 s a)) - {ω | ω / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := by - rw [IdentDistrib.map_eq (identDistrib_sum_Icc_rewardByCount k a)] - _ = 𝔓 {ω | (∑ s ∈ range k, ω.2 s a) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} := by - rw [Measure.map_apply (by fun_prop) h_meas] - rfl - _ = 𝔓 {ω | (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) / k ≤ - √(c * log (n + 1) / k)} := by + calc Bandit.streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(c * log (n + 1) / k) ≤ (ν a)[id]} + _ = Bandit.streamMeasure ν + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) / k ≤ - √(c * log (n + 1) / k)} := by congr with ω field_simp rw [Finset.sum_sub_distrib] simp grind - _ = 𝔓 {ω | (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) ≤ - √(c * k * log (n + 1))} := by + _ = Bandit.streamMeasure ν + {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ - √(c * k * log (n + 1))} := by congr with ω field_simp congr! 2 rw [sqrt_div (by positivity), ← mul_div_assoc, mul_comm, mul_div_assoc, div_sqrt, mul_assoc (k : ℝ), sqrt_mul (x := (k : ℝ)) (by positivity), mul_comm] - _ = Bandit.streamMeasure ν - {ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ - √(c * k * log (n + 1))} := by - rw [← Bandit.snd_measure (ucbAlgorithm hK c), Measure.snd_apply] - · rfl - · exact measurableSet_le (by fun_prop) (by fun_prop) - _ ≤ ENNReal.ofReal (exp (-(√(c * k * log (n + 1))) ^ 2 / (2 * k * 1))) := by - rw [← ofReal_measureReal] - gcongr - refine (HasSubgaussianMGF.measure_sum_range_le_le_of_iIndepFun (c := 1) ?_ ?_ (by positivity)) - · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) - (fun _ ↦ by fun_prop) - · intro i him - refine (hν a).congr_identDistrib ?_ - exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ - _ = 1 / (n + 1) ^ (c / 2) := by - rw [sq_sqrt] - swap; · exact mul_nonneg (by positivity) (log_nonneg (by simp)) - field_simp - rw [div_eq_inv_mul, ← mul_assoc, ← Real.log_rpow (by positivity), ← Real.log_inv, - Real.exp_log (by positivity), one_div, ENNReal.ofReal_inv_of_pos (by positivity), - ← ENNReal.ofReal_rpow_of_nonneg (by positivity) (by positivity)] - congr 2 - · norm_cast - · field + _ ≤ 1 / (n + 1) ^ (c / 2) := prob_sum_le_sqrt_log hν hc a k hk lemma todo' (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n k : ℕ) (hk : k ≠ 0) : - 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - √(c * log (n + 1) / k)} ≤ + Bandit.streamMeasure ν + {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - √(c * log (n + 1) / k)} ≤ 1 / (n + 1) ^ (c / 2) := by - have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - have h_meas : MeasurableSet {ω |(ν a)[id] ≤ ω / k - √(c * log (n + 1) / k)} := - measurableSet_le (by fun_prop) (by fun_prop) have h_log_nonneg : 0 ≤ log (n + 1) := log_nonneg (by simp) - calc - 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - √(c * log (n + 1) / k)} - _ = ((𝔓).map (fun ω ↦ ∑ m ∈ Icc 1 k, rewardByCount a m ω)) - {ω | (ν a)[id] ≤ ω / k - √(c * log (n + 1) / k)} := by - rw [Measure.map_apply (by fun_prop) h_meas] - rfl - _ = ((𝔓).map (fun ω ↦ ∑ s ∈ range k, ω.2 s a)) - {ω | (ν a)[id] ≤ ω / k - √(c * log (n + 1) / k)} := by - rw [IdentDistrib.map_eq (identDistrib_sum_Icc_rewardByCount k a)] - _ = 𝔓 {ω | (ν a)[id] ≤ (∑ s ∈ range k, ω.2 s a) / k - √(c * log (n + 1) / k)} := by - rw [Measure.map_apply (by fun_prop) h_meas] - rfl - _ = 𝔓 {ω | √(c * log (n + 1) / k) ≤ (∑ s ∈ range k, (ω.2 s a - (ν a)[id])) / k} := by + calc Bandit.streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - √(c * log (n + 1) / k)} + _ = Bandit.streamMeasure ν + {ω | √(c * log (n + 1) / k) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id])) / k} := by congr with ω field_simp rw [Finset.sum_sub_distrib] simp grind - _ = 𝔓 {ω | √(c * k * log (n + 1)) ≤ (∑ s ∈ range k, (ω.2 s a - (ν a)[id]))} := by + _ = Bandit.streamMeasure ν + {ω | √(c * k * log (n + 1)) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} := by congr with ω field_simp congr! 1 rw [sqrt_div (by positivity), ← mul_div_assoc, mul_comm, mul_div_assoc, div_sqrt, mul_comm _ (k : ℝ), sqrt_mul (x := (k : ℝ)) (by positivity), mul_comm] - _ = Bandit.streamMeasure ν - {ω | √(c * k * log (n + 1)) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} := by - rw [← Bandit.snd_measure (ucbAlgorithm hK c), Measure.snd_apply] - · rfl - · exact measurableSet_le (by fun_prop) (by fun_prop) - _ ≤ ENNReal.ofReal (exp (-(√(c * k * log (n + 1))) ^ 2 / (2 * k * 1))) := by - rw [← ofReal_measureReal] - gcongr - refine (HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := 1) ?_ ?_ (by positivity)) - · exact (iIndepFun_eval_streamMeasure'' ν a).comp (fun i ω ↦ ω - (ν a)[id]) - (fun _ ↦ by fun_prop) - · intro i him - refine (hν a).congr_identDistrib ?_ - exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ - _ = 1 / (n + 1) ^ (c / 2) := by - rw [sq_sqrt] - swap; · exact mul_nonneg (by positivity) (log_nonneg (by simp)) - field_simp - rw [div_eq_inv_mul, ← mul_assoc, ← Real.log_rpow (by positivity), ← Real.log_inv, - Real.exp_log (by positivity), one_div, ENNReal.ofReal_inv_of_pos (by positivity), - ← ENNReal.ofReal_rpow_of_nonneg (by positivity) (by positivity)] - congr 2 - · norm_cast - · field - -lemma prob_ucbIndex_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) + _ ≤ 1 / (n + 1) ^ (c / 2) := prob_sum_ge_sqrt_log hν hc a k hk + +lemma prob_ucbIndex_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : - 𝔓t {h | 0 < pullCount a n h ∧ empMean a n h + ucbWidth c a n h ≤ (ν a)[id]} ≤ + P {h | 0 < pullCount A a n h ∧ empMean A R a n h + ucbWidth A c a n h ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by - -- extend the probability space - suffices 𝔓 {ω | 0 < pullCount a n ω.1 ∧ - empMean a n ω.1 + ucbWidth c a n ω.1 ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2 - 1) by - rwa [← Bandit.fst_measure (ucbAlgorithm hK c) ν, Measure.fst_apply] - change MeasurableSet ({h | 0 < pullCount a n h} - ∩ {h | empMean a n h + ucbWidth c a n h ≤ ∫ (x : ℝ), id x ∂ν a}) - refine MeasurableSet.inter ?_ ?_ - · exact measurableSet_lt (by fun_prop) (by fun_prop) - · exact measurableSet_le (by fun_prop) (by fun_prop) - -- express with `rewardByCount` and `pullCount` - unfold empMean ucbWidth - simp_rw [← sum_rewardByCount_eq_sumRewards] - calc - 𝔓 {ω | 0 < pullCount a n ω.1 ∧ - (∑ m ∈ Icc 1 (pullCount a n ω.1), rewardByCount a m ω) / pullCount a n ω.1 + - √(c * log (↑n + 1) / pullCount a n ω.1) ≤ (ν a)[id]} - -- list the possible values of `pullCount a n ω.1` - _ ≤ 𝔓 {ω | ∃ k ≤ n, 0 < k ∧ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + - √(c * log (↑n + 1) / k) ≤ (ν a)[id]} := by - refine measure_mono fun ω hω ↦ ?_ - simp only [Nat.cast_nonneg, sqrt_div', id_eq, Set.mem_setOf_eq] at hω ⊢ - exact ⟨pullCount a n ω.1, pullCount_le _ _ _, hω⟩ - _ = 𝔓 (⋃ k ∈ Icc 1 n, {ω |(∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + - √(c * log (↑n + 1) / k) ≤ (ν a)[id]}) := by - congr 1 - ext ω - simp - grind - -- Union bound over the possible values of `pullCount a n ω.1` + let s : Set (ℕ × ℝ) := {(m, x) | 0 < m ∧ x / m + √(c * log (↑n + 1) / m) ≤ (ν a)[id]} + have hs : MeasurableSet s := by + simp only [Nat.cast_nonneg, sqrt_div', id_eq, measurableSet_setOf, s] + fun_prop + classical + calc P {h | 0 < pullCount A a n h ∧ empMean A R a n h + ucbWidth A c a n h ≤ (ν a)[id]} + _ ≤ ∑ k ∈ range (n + 1) with k ∈ Prod.fst '' s, + (Bandit.streamMeasure ν) {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := + prob_pullCount_prod_sumRewards_mem_le h hs _ ≤ ∑ k ∈ Icc 1 n, - 𝔓 {ω | (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k + √(c * log (↑n + 1) / k) ≤ (ν a)[id]} := - measure_biUnion_finset_le _ _ + (Bandit.streamMeasure ν) {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := by + refine Finset.sum_le_sum_of_subset_of_nonneg (fun m ↦ ?_) fun _ _ _ ↦ by positivity + simp [s] + grind + _ = ∑ k ∈ Icc 1 n, + (Bandit.streamMeasure ν) {ω | (∑ i ∈ range k, ω i a) / k + √(c * log (↑n + 1) / k) ≤ + (ν a)[id]} := by + refine Finset.sum_congr rfl fun k hk ↦ ?_ + congr with ω + have hk : 0 < k := by grind + simp [s, hk] _ ≤ ∑ k ∈ Icc 1 n, (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by gcongr with k hk exact todo hν hc a n k (by grind) @@ -359,42 +299,33 @@ lemma prob_ucbIndex_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id] rw [ENNReal.rpow_sub _ _ (by simp) (by finiteness), ENNReal.rpow_one, div_eq_mul_inv, ENNReal.div_eq_inv_mul, ENNReal.mul_inv (by simp) (by simp), inv_inv] -lemma prob_ucbIndex_ge (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma prob_ucbIndex_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : - 𝔓t {h | 0 < pullCount a n h ∧ - (ν a)[id] ≤ empMean a n h - ucbWidth c a n h} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by - -- extend the probability space - suffices 𝔓 {ω | 0 < pullCount a n ω.1 ∧ - (ν a)[id] ≤ empMean a n ω.1 - ucbWidth c a n ω.1} ≤ 1 / (n + 1) ^ (c / 2 - 1) by - rwa [← Bandit.fst_measure (ucbAlgorithm hK c) ν, Measure.fst_apply] - change MeasurableSet ({h | 0 < pullCount a n h} - ∩ {h | (ν a)[id] ≤ empMean a n h - ucbWidth c a n h}) - refine MeasurableSet.inter ?_ ?_ - · exact measurableSet_lt (by fun_prop) (by fun_prop) - · exact measurableSet_le (by fun_prop) (by fun_prop) - -- express with `rewardByCount` and `pullCount` - unfold empMean ucbWidth - simp_rw [← sum_rewardByCount_eq_sumRewards] - calc - 𝔓 {ω | 0 < pullCount a n ω.1 ∧ - (ν a)[id] ≤ (∑ m ∈ Icc 1 (pullCount a n ω.1), rewardByCount a m ω) / pullCount a n ω.1 - - √(c * log (↑n + 1) / pullCount a n ω.1)} - -- list the possible values of `pullCount a n ω.1` - _ ≤ 𝔓 {ω | ∃ k ≤ n, 0 < k ∧ (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - - √(c * log (↑n + 1) / k)} := by - refine measure_mono fun ω hω ↦ ?_ - simp only [Nat.cast_nonneg, sqrt_div', id_eq, Set.mem_setOf_eq] at hω ⊢ - exact ⟨pullCount a n ω.1, pullCount_le _ _ _, hω⟩ - _ = 𝔓 (⋃ k ∈ Icc 1 n, {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - - √(c * log (↑n + 1) / k)}) := by - congr 1 - ext ω - simp - grind - -- Union bound over the possible values of `pullCount a n ω.1` + P {h | 0 < pullCount A a n h ∧ + (ν a)[id] ≤ empMean A R a n h - ucbWidth A c a n h} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by + let s : Set (ℕ × ℝ) := {(m, x) | 0 < m ∧ (ν a)[id] ≤ x / m - √(c * log (↑n + 1) / m)} + have hs : MeasurableSet s := by + simp only [Nat.cast_nonneg, sqrt_div', id_eq, measurableSet_setOf, s] + fun_prop + classical + calc P {h | 0 < pullCount A a n h ∧ (ν a)[id] ≤ empMean A R a n h - ucbWidth A c a n h} + _ ≤ ∑ k ∈ range (n + 1) with k ∈ Prod.fst '' s, + (Bandit.streamMeasure ν) {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := + prob_pullCount_prod_sumRewards_mem_le h hs _ ≤ ∑ k ∈ Icc 1 n, - 𝔓 {ω | (ν a)[id] ≤ (∑ m ∈ Icc 1 k, rewardByCount a m ω) / k - √(c * log (↑n + 1) / k)} := - measure_biUnion_finset_le _ _ + (Bandit.streamMeasure ν) {ω | ∑ i ∈ range k, ω i a ∈ Prod.mk k ⁻¹' s} := by + refine Finset.sum_le_sum_of_subset_of_nonneg (fun m ↦ ?_) fun _ _ _ ↦ by positivity + simp [s] + grind + _ = ∑ k ∈ Icc 1 n, + (Bandit.streamMeasure ν) + {ω | (ν a)[id] ≤ (∑ i ∈ range k, ω i a) / k - √(c * log (↑n + 1) / k)} := by + refine Finset.sum_congr rfl fun k hk ↦ ?_ + congr with ω + have hk : 0 < k := by grind + simp [s, hk] _ ≤ ∑ k ∈ Icc 1 n, (1 : ℝ≥0∞) / (n + 1) ^ (c / 2) := by gcongr with k hk exact todo' hν hc a n k (by grind) @@ -408,96 +339,64 @@ lemma prob_ucbIndex_ge (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id] rw [ENNReal.rpow_sub _ _ (by simp) (by finiteness), ENNReal.rpow_one, div_eq_mul_inv, ENNReal.div_eq_inv_mul, ENNReal.mul_inv (by simp) (by simp), inv_inv] -lemma probReal_ucbIndex_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma probReal_ucbIndex_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : - (𝔓t).real {h | 0 < pullCount a n h ∧ empMean a n h + ucbWidth c a n h ≤ (ν a)[id]} ≤ + P.real {h | 0 < pullCount A a n h ∧ empMean A R a n h + ucbWidth A c a n h ≤ (ν a)[id]} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by rw [measureReal_def] - grw [prob_ucbIndex_le hν hc a n] + grw [prob_ucbIndex_le h hν hc a n] swap; · finiteness simp only [one_div, ENNReal.toReal_inv] rw [← ENNReal.toReal_rpow] norm_cast -lemma probReal_ucbIndex_ge (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma probReal_ucbIndex_ge [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 ≤ c) (a : Fin K) (n : ℕ) : - (𝔓t).real {h | 0 < pullCount a n h ∧ - (ν a)[id] ≤ empMean a n h - ucbWidth c a n h} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by + P.real {h | 0 < pullCount A a n h ∧ + (ν a)[id] ≤ empMean A R a n h - ucbWidth A c a n h} ≤ 1 / (n + 1) ^ (c / 2 - 1) := by rw [measureReal_def] - grw [prob_ucbIndex_ge hν hc a n] + grw [prob_ucbIndex_ge h hν hc a n] swap; · finiteness simp only [one_div, ENNReal.toReal_inv] rw [← ENNReal.toReal_rpow] norm_cast -lemma pullCount_le_add (a : Fin K) (n C : ℕ) (ω : ℕ → Fin K × ℝ) : - pullCount a n ω ≤ C + 1 + - ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω}.indicator 1 s := by - rw [pullCount_eq_sum] - calc ∑ s ∈ range n, if arm s ω = a then 1 else 0 - _ ≤ ∑ s ∈ range n, ({s | arm s ω = a ∧ pullCount a s ω ≤ C}.indicator 1 s + - {s | arm s ω = a ∧ C < pullCount a s ω}.indicator 1 s) := by - gcongr with s hs - simp [Set.indicator_apply] - grind - _ = ∑ s ∈ range n, {s | arm s ω = a ∧ pullCount a s ω ≤ C}.indicator 1 s + - ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω}.indicator 1 s := by - rw [Finset.sum_add_distrib] - _ ≤ C + 1 + ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω}.indicator 1 s := by - gcongr - have h_le n : ∑ s ∈ range n, {s | arm s ω = a ∧ pullCount a s ω ≤ C}.indicator 1 s ≤ - pullCount a n ω := by - rw [pullCount_eq_sum] - gcongr with s hs - simp only [Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply, arm, action] - grind - induction n with - | zero => simp - | succ n hn => - rw [Finset.sum_range_succ] - rcases le_or_gt (pullCount a n ω) C with h_pc | h_pc - · have hn' : ∑ s ∈ range n, {s | arm s ω = a ∧ pullCount a s ω ≤ C}.indicator 1 s ≤ C := - (h_le n).trans h_pc - grw [hn'] - gcongr - simp only [Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] - grind - · refine le_trans ?_ hn - simp [h_pc] - omit [IsMarkovKernel ν] in -lemma pullCount_le_add_three [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (ω : ℕ → Fin K × ℝ) : - pullCount a n ω ≤ C + 1 + - ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω ∧ - (ν (bestArm ν))[id] ≤ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω ∧ - empMean (arm s ω) s ω - ucbWidth c (arm s ω) s ω ≤ (ν (arm s ω))[id]}.indicator 1 s + +lemma pullCount_le_add_three [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (ω : Ω) : + pullCount A a n ω ≤ C + 1 + + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω ∧ + (ν (bestArm ν))[id] ≤ empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω ∧ + empMean A R (A s ω) s ω - ucbWidth A c (A s ω) s ω ≤ (ν (A s ω))[id]}.indicator 1 s + ∑ s ∈ range n, - {s | C < pullCount a s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < + {s | C < pullCount A a s ω ∧ empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω < (ν (bestArm ν))[id]}.indicator 1 s + ∑ s ∈ range n, - {s | C < pullCount a s ω ∧ (ν a)[id] < - empMean a s ω - ucbWidth c a s ω}.indicator 1 s := by + {s | C < pullCount A a s ω ∧ (ν a)[id] < + empMean A R a s ω - ucbWidth A c a s ω}.indicator 1 s := by refine (pullCount_le_add a n C ω).trans ?_ simp_rw [add_assoc] gcongr simp_rw [← add_assoc] - let A := {s | arm s ω = a ∧ C < pullCount a s ω} - let B := {s | arm s ω = a ∧ C < pullCount a s ω ∧ - (ν (bestArm ν))[id] ≤ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω ∧ - empMean (arm s ω) s ω - ucbWidth c (arm s ω) s ω ≤ (ν (arm s ω))[id]} - let C' := {s | C < pullCount a s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < - (ν (bestArm ν))[id]} - let D := {s | C < pullCount a s ω ∧ (ν a)[id] < - empMean a s ω - ucbWidth c a s ω} - change ∑ s ∈ range n, A.indicator 1 s ≤ + let A' := {s | A s ω = a ∧ C < pullCount A a s ω} + let B := {s | A s ω = a ∧ C < pullCount A a s ω ∧ + (ν (bestArm ν))[id] ≤ empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω ∧ + empMean A R (A s ω) s ω - ucbWidth A c (A s ω) s ω ≤ (ν (A s ω))[id]} + let C' := {s | C < pullCount A a s ω ∧ + empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω < (ν (bestArm ν))[id]} + let D := {s | C < pullCount A a s ω ∧ (ν a)[id] < empMean A R a s ω - ucbWidth A c a s ω} + change ∑ s ∈ range n, A'.indicator 1 s ≤ ∑ s ∈ range n, B.indicator 1 s + ∑ s ∈ range n, C'.indicator 1 s + ∑ s ∈ range n, D.indicator 1 s - have h_union : A ⊆ B ∪ C' ∪ D := by simp [A, B, C', D]; grind + have h_union : A' ⊆ B ∪ C' ∪ D := by simp [A', B, C', D]; grind calc - (∑ s ∈ range n, A.indicator 1 s) + (∑ s ∈ range n, A'.indicator 1 s) _ ≤ (∑ s ∈ range n, (B ∪ C' ∪ D).indicator (fun _ ↦ (1 : ℕ)) s) := by gcongr with n hn - by_cases h : n ∈ A + by_cases h : n ∈ A' · have : n ∈ B ∪ C' ∪ D := h_union h simp [h, this] · simp [h] @@ -509,33 +408,38 @@ lemma pullCount_le_add_three [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (ω : ∑ s ∈ range n, D.indicator 1 s := by rw [Finset.sum_add_distrib, Finset.sum_add_distrib] -lemma pullCount_le_add_three_ae [Nonempty (Fin K)] (a : Fin K) (n C : ℕ) (hC : C ≠ 0) : - ∀ᵐ ω ∂𝔓t, - pullCount a n ω ≤ C + 1 + - ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω ∧ - (ν (bestArm ν))[id] ≤ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω ∧ - empMean (arm s ω) s ω - ucbWidth c (arm s ω) s ω ≤ (ν (arm s ω))[id]}.indicator 1 s + +lemma pullCount_le_add_three_ae [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (a : Fin K) (n C : ℕ) (hC : C ≠ 0) : + ∀ᵐ ω ∂P, + pullCount A a n ω ≤ C + 1 + + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω ∧ + (ν (bestArm ν))[id] ≤ empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω ∧ + empMean A R (A s ω) s ω - ucbWidth A c (A s ω) s ω ≤ (ν (A s ω))[id]}.indicator 1 s + ∑ s ∈ range n, - {s | 0 < pullCount (bestArm ν) s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < - (ν (bestArm ν))[id]}.indicator 1 s + + {s | 0 < pullCount A (bestArm ν) s ω ∧ + empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω < + (ν (bestArm ν))[id]}.indicator 1 s + ∑ s ∈ range n, - {s | 0 < pullCount a s ω ∧ (ν a)[id] < - empMean a s ω - ucbWidth c a s ω}.indicator 1 s := by - filter_upwards [pullCount_pos_of_pullCount_gt_one a] with ω hω - refine (pullCount_le_add_three a n C ω (ν := ν) (c := c)).trans ?_ + {s | 0 < pullCount A a s ω ∧ (ν a)[id] < + empMean A R a s ω - ucbWidth A c a s ω}.indicator 1 s := by + filter_upwards [pullCount_pos_of_pullCount_gt_one h a] with ω hω + refine (pullCount_le_add_three (R := R) a n C ω (ν := ν) (c := c)).trans ?_ gcongr 5 with k hk j k hk j · gcongr 1 exact fun h_gt ↦ hω _ (lt_of_le_of_lt (by grind) h_gt) _ · exact fun h_gt ↦ hω _ (lt_of_le_of_lt (by grind) h_gt) _ -lemma some_sum_eq_zero [Nonempty (Fin K)] (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 < gap ν a) (n C : ℕ) +lemma some_sum_eq_zero [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 < gap ν a) (n C : ℕ) (hC : C ≠ 0) (hC' : 4 * c * log (n + 1) / gap ν a ^ 2 ≤ C) : - ∀ᵐ ω ∂𝔓t, - ∑ s ∈ range n, {s | arm s ω = a ∧ C < pullCount a s ω ∧ - (ν (bestArm ν))[id] ≤ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω ∧ - empMean (arm s ω) s ω - ucbWidth c (arm s ω) s ω ≤ (ν (arm s ω))[id]}.indicator 1 s = 0 := by - have h_ae := forall_ucbIndex_le_ucbIndex_arm (bestArm ν) (ν := ν) (c := c) (hK := hK) - have h_gt := time_gt_of_pullCount_gt_one a (ν := ν) (c := c) (hK := hK) + ∀ᵐ ω ∂P, + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω ∧ + (ν (bestArm ν))[id] ≤ empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω ∧ + empMean A R (A s ω) s ω - ucbWidth A c (A s ω) s ω ≤ (ν (A s ω))[id]}.indicator 1 s = 0 := by + have h_ae := forall_ucbIndex_le_ucbIndex_arm h (bestArm ν) (ν := ν) (c := c) (hK := hK) + have h_gt := time_gt_of_pullCount_gt_one h a (ν := ν) (c := c) (hK := hK) filter_upwards [h_ae, h_gt] with ω h_le h_time_ge simp only [id_eq, tsub_le_iff_right, sum_eq_zero_iff, mem_range, Set.indicator_apply_eq_zero, Set.mem_setOf_eq, Pi.one_apply, one_ne_zero, imp_false, not_and, not_le] @@ -557,18 +461,21 @@ lemma some_sum_eq_zero [Nonempty (Fin K)] (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 · rw [h_arm] gcongr -lemma pullCount_ae_le_add_two [Nonempty (Fin K)] (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 < gap ν a) +lemma pullCount_ae_le_add_two [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hc : 0 ≤ c) (a : Fin K) (h_gap : 0 < gap ν a) (n C : ℕ) (hC : C ≠ 0) (hC' : 4 * c * log (n + 1) / gap ν a ^ 2 ≤ C) : - ∀ᵐ ω ∂𝔓t, - pullCount a n ω ≤ C + 1 + + ∀ᵐ ω ∂P, + pullCount A a n ω ≤ C + 1 + ∑ s ∈ range n, - {s | 0 < pullCount (bestArm ν) s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < - (ν (bestArm ν))[id]}.indicator 1 s + + {s | 0 < pullCount A (bestArm ν) s ω ∧ + empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω < + (ν (bestArm ν))[id]}.indicator 1 s + ∑ s ∈ range n, - {s | 0 < pullCount a s ω ∧ (ν a)[id] < - empMean a s ω - ucbWidth c a s ω}.indicator 1 s := by - filter_upwards [some_sum_eq_zero hc a h_gap n C hC hC', - pullCount_le_add_three_ae a n C hC] with ω hω_zero hω_le + {s | 0 < pullCount A a s ω ∧ (ν a)[id] < + empMean A R a s ω - ucbWidth A c a s ω}.indicator 1 s := by + filter_upwards [some_sum_eq_zero h hc a h_gap n C hC hC', + pullCount_le_add_three_ae h a n C hC] with ω hω_zero hω_le refine (hω_le).trans_eq ?_ rw [hω_zero] @@ -583,57 +490,62 @@ lemma constSum_lt_top (c : ℝ) (n : ℕ) : constSum c n < ∞ := by positivity /-- Bound on the expectation of the number of pulls of each arm by the UCB algorithm. -/ -lemma expectation_pullCount_le' (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma expectation_pullCount_le' [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 < c) (a : Fin K) (h_gap : 0 < gap ν a) (n : ℕ) : - ∫⁻ ω, pullCount a n ω ∂𝔓t ≤ + ∫⁻ ω, pullCount A a n ω ∂P ≤ ENNReal.ofReal (4 * c * log (n + 1) / gap ν a ^ 2 + 1) + 1 + 2 * constSum c n := by + have hA := h.measurable_A + have hR := h.measurable_R by_cases hn_zero : n = 0 · simp [hn_zero] let C a : ℕ := ⌈4 * c * log (n + 1) / gap ν a ^ 2⌉₊ have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - have h_set_1 b : MeasurableSet {a_1 | 0 < pullCount a b a_1 ∧ - (ν a)[id] < empMean a b a_1 - ucbWidth c a b a_1} := by - change MeasurableSet ({a_1 | 0 < pullCount a b a_1} ∩ - {a_1 | (ν a)[id] < empMean a b a_1 - ucbWidth c a b a_1}) + have h_set_1 b : MeasurableSet {a_1 | 0 < pullCount A a b a_1 ∧ + (ν a)[id] < empMean A R a b a_1 - ucbWidth A c a b a_1} := by + change MeasurableSet ({a_1 | 0 < pullCount A a b a_1} ∩ + {a_1 | (ν a)[id] < empMean A R a b a_1 - ucbWidth A c a b a_1}) exact (measurableSet_lt (by fun_prop) (by fun_prop)).inter (measurableSet_lt (by fun_prop) (by fun_prop)) - have h_set_2 b : MeasurableSet {a | 0 < pullCount (bestArm ν) b a ∧ - empMean (bestArm ν) b a + ucbWidth c (bestArm ν) b a < (ν (bestArm ν))[id]} := by - change MeasurableSet ({a | 0 < pullCount (bestArm ν) b a} ∩ - {a | empMean (bestArm ν) b a + ucbWidth c (bestArm ν) b a < (ν (bestArm ν))[id]}) + have h_set_2 b : MeasurableSet {a | 0 < pullCount A (bestArm ν) b a ∧ + empMean A R (bestArm ν) b a + ucbWidth A c (bestArm ν) b a < (ν (bestArm ν))[id]} := by + change MeasurableSet ({a | 0 < pullCount A (bestArm ν) b a} ∩ + {a | empMean A R (bestArm ν) b a + ucbWidth A c (bestArm ν) b a < (ν (bestArm ν))[id]}) exact (measurableSet_lt (by fun_prop) (by fun_prop)).inter (measurableSet_lt (by fun_prop) (by fun_prop)) - have h_meas_1 b : Measurable fun h ↦ {s | 0 < pullCount a s h ∧ (ν a)[id] < - empMean a s h - ucbWidth c a s h}.indicator (1 : ℕ → ℕ) b := by + have h_meas_1 b : Measurable fun h ↦ {s | 0 < pullCount A a s h ∧ (ν a)[id] < + empMean A R a s h - ucbWidth A c a s h}.indicator (1 : ℕ → ℕ) b := by simp only [id_eq, Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] exact Measurable.ite (h_set_1 _) (by fun_prop) (by fun_prop) - have h_meas_2 b : Measurable fun h ↦ {s | 0 < pullCount (bestArm ν) s h ∧ - empMean (bestArm ν) s h + ucbWidth c (bestArm ν) s h < + have h_meas_2 b : Measurable fun h ↦ {s | 0 < pullCount A (bestArm ν) s h ∧ + empMean A R (bestArm ν) s h + ucbWidth A c (bestArm ν) s h < (ν (bestArm ν))[id]}.indicator (1 : ℕ → ℕ) b := by simp only [id_eq, Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] exact Measurable.ite (h_set_2 _) (by fun_prop) (by fun_prop) - calc ∫⁻ ω, pullCount a n ω ∂𝔓t + calc ∫⁻ ω, pullCount A a n ω ∂P _ ≤ ∫⁻ ω, C a + 1 + ∑ s ∈ range n, - {s | 0 < pullCount (bestArm ν) s ω ∧ empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < - (ν (bestArm ν))[id]}.indicator (1 : ℕ → ℕ) s + + {s | 0 < pullCount A (bestArm ν) s ω ∧ + empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω < + (ν (bestArm ν))[id]}.indicator (1 : ℕ → ℕ) s + ∑ s ∈ range n, - {s | 0 < pullCount a s ω ∧ (ν a)[id] < - empMean a s ω - ucbWidth c a s ω}.indicator (1 : ℕ → ℕ) s ∂𝔓t := by + {s | 0 < pullCount A a s ω ∧ (ν a)[id] < + empMean A R a s ω - ucbWidth A c a s ω}.indicator (1 : ℕ → ℕ) s ∂P := by refine lintegral_mono_ae ?_ have hCa : C a ≠ 0 := by simp only [ne_eq, Nat.ceil_eq_zero, not_le, C] have : 0 < log (n + 1) := log_pos (by simp; grind) positivity - filter_upwards [pullCount_ae_le_add_two hc.le a h_gap n (C a) hCa (Nat.le_ceil _)] with ω hω + filter_upwards [pullCount_ae_le_add_two h hc.le a h_gap n (C a) hCa (Nat.le_ceil _)] with ω hω simp only [id_eq, Nat.cast_sum] norm_cast _ ≤ (C a : ℝ≥0∞) + 1 + ∑ s ∈ range n, - 𝔓t {ω | 0 < pullCount (bestArm ν) s ω ∧ - empMean (bestArm ν) s ω + ucbWidth c (bestArm ν) s ω < (ν (bestArm ν))[id]} + + P {ω | 0 < pullCount A (bestArm ν) s ω ∧ + empMean A R (bestArm ν) s ω + ucbWidth A c (bestArm ν) s ω < (ν (bestArm ν))[id]} + ∑ s ∈ range n, - 𝔓t {ω | 0 < pullCount a s ω ∧ (ν a)[id] < empMean a s ω - ucbWidth c a s ω} := by + P {ω | 0 < pullCount A a s ω ∧ (ν a)[id] < empMean A R a s ω - ucbWidth A c a s ω} := by simp only [id_eq, Nat.cast_sum] rw [lintegral_add_left (by fun_prop), lintegral_add_left (by fun_prop)] simp only [lintegral_const, measure_univ, mul_one] @@ -651,9 +563,9 @@ lemma expectation_pullCount_le' (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - ( ∑ s ∈ range n, 1 / ((s : ℝ≥0∞) + 1) ^ (c / 2 - 1) + ∑ s ∈ range n, 1 / ((s : ℝ≥0∞) + 1) ^ (c / 2 - 1) := by gcongr with s hs s hs - · refine (measure_mono ?_).trans (prob_ucbIndex_le hν hc.le (bestArm ν) s) + · refine (measure_mono ?_).trans (prob_ucbIndex_le h hν hc.le (bestArm ν) s) grind - · refine (measure_mono ?_).trans (prob_ucbIndex_ge hν hc.le a s) + · refine (measure_mono ?_).trans (prob_ucbIndex_ge h hν hc.le a s) grind _ ≤ ENNReal.ofReal (4 * c * log (n + 1) / gap ν a ^ 2 + 1) + 1 + 2 * constSum c n := by rw [two_mul, add_assoc, constSum] @@ -666,15 +578,18 @@ lemma expectation_pullCount_le' (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - ( positivity /-- Bound on the expectation of the number of pulls of each arm by the UCB algorithm. -/ -lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) +lemma expectation_pullCount_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 < c) (a : Fin K) (h_gap : 0 < gap ν a) (n : ℕ) : - 𝔓t[fun ω ↦ (pullCount a n ω : ℝ)] ≤ + P[fun ω ↦ (pullCount A a n ω : ℝ)] ≤ 4 * c * log (n + 1) / gap ν a ^ 2 + 2 + 2 * (constSum c n).toReal := by - have h := expectation_pullCount_le' hν hc a h_gap n (hK := hK) + have hA := h.measurable_A + have h := expectation_pullCount_le' h hν hc a h_gap n (hK := hK) simp_rw [← ENNReal.ofReal_natCast] at h rw [← ofReal_integral_eq_lintegral_ofReal] at h rotate_left - · exact integrable_pullCount _ _ + · exact integrable_pullCount hA _ _ · exact ae_of_all _ fun _ ↦ by simp simp only have : 0 ≤ log (n + 1) := log_nonneg (by simp) @@ -694,18 +609,21 @@ lemma expectation_pullCount_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - ( ring /-- Regret bound for the UCB algorithm. -/ -lemma regret_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 < c) (n : ℕ) : - 𝔓t[regret ν n] ≤ +lemma regret_le [Nonempty (Fin K)] + (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) + (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) 1 (ν a)) (hc : 0 < c) (n : ℕ) : + P[regret ν A n] ≤ ∑ a, (4 * c * log (n + 1) / gap ν a + gap ν a * (2 + 2 * (constSum c n).toReal)) := by + have hA := h.measurable_A simp_rw [regret_eq_sum_pullCount_mul_gap] rw [integral_finset_sum] - swap; · exact fun i _ ↦ (integrable_pullCount i n).mul_const _ + swap; · exact fun i _ ↦ (integrable_pullCount hA i n).mul_const _ gcongr with a rw [integral_mul_const] by_cases h_gap : gap ν a = 0 · simp [h_gap] replace h_gap : 0 < gap ν a := lt_of_le_of_ne gap_nonneg (Ne.symm h_gap) - grw [expectation_pullCount_le hν hc a h_gap n] + grw [expectation_pullCount_le h hν hc a h_gap n] refine le_of_eq ?_ rw [mul_add] field diff --git a/LeanBandits/ForMathlib/CondDistrib.lean b/LeanBandits/ForMathlib/CondDistrib.lean index 9aa67375..030e1bd6 100644 --- a/LeanBandits/ForMathlib/CondDistrib.lean +++ b/LeanBandits/ForMathlib/CondDistrib.lean @@ -33,14 +33,6 @@ lemma Measurable.toNat {f : α → ℕ∞} (hf : Measurable f) : Measurable (fun namespace MeasureTheory.Measure -lemma trim_eq_map {hm : m ≤ mα} : μ.trim hm = @Measure.map _ _ mα m id μ := by - refine @Measure.ext _ m _ _ fun s hs ↦ ?_ - rw [trim_measurableSet_eq _ hs, Measure.map_apply _ hs] - · simp - · intro t ht - simp only [Set.preimage_id_eq, id_eq] - exact hm _ ht - lemma trim_comap_apply (hX : Measurable X) {s : Set β} (hs : MeasurableSet s) : μ.trim hX.comap_le (X ⁻¹' s) = μ.map X s := by rw [trim_measurableSet_eq, Measure.map_apply (by fun_prop) hs] @@ -52,13 +44,6 @@ namespace ProbabilityTheory section IndepFun --- fix the lemma in mathlib to allow different types for the functions -theorem CondIndepFun.symm' - [StandardBorelSpace α] {hm : m ≤ mα} [IsFiniteMeasure μ] {f : α → β} {g : α → γ} - (hfg : CondIndepFun m hm f g μ) : - CondIndepFun m hm g f μ := - Kernel.IndepFun.symm hfg - lemma Kernel.IndepFun.of_prod_right {ε Ω : Type*} {mΩ : MeasurableSpace Ω} {mε : MeasurableSpace ε} {μ : Measure Ω} {κ : Kernel Ω α} {X : α → β} {Y : α → γ} {T : α → ε} (h : IndepFun X (fun ω ↦ (Y ω, T ω)) κ μ) : @@ -87,14 +72,6 @@ lemma CondIndepFun.of_prod_left {ε : Type*} {mε : MeasurableSpace ε} X ⟂ᵢ[Z, hZ; μ] Y := Kernel.IndepFun.of_prod_left h -lemma CondIndepFun.prod_right [StandardBorelSpace α] [StandardBorelSpace β] [Nonempty β] - [StandardBorelSpace γ] [Nonempty γ] [StandardBorelSpace δ] [Nonempty δ] [IsFiniteMeasure μ] - {X : α → β} {Y : α → γ} {Z : α → δ} - (hX : Measurable X) (hY : Measurable Y) (hZ : Measurable Z) - (h : X ⟂ᵢ[Z, hZ; μ] Y) : - X ⟂ᵢ[Z, hZ; μ] (fun ω ↦ (Y ω, Z ω)) := by - sorry - end IndepFun section CondDistrib @@ -110,6 +87,77 @@ lemma condDistrib_prod_left [StandardBorelSpace β] [Nonempty β] AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] rfl +lemma condDistrib_prod_self_left [StandardBorelSpace β] [Nonempty β] [StandardBorelSpace γ] + [Nonempty γ] + (hX : AEMeasurable X μ) (hT : AEMeasurable T μ) : + condDistrib (fun ω ↦ (X ω, T ω)) T μ =ᵐ[μ.map T] condDistrib X T μ ×ₖ Kernel.id := by + have h_prod := condDistrib_prod_left hX hT hT (μ := μ) + have h_fst := condDistrib_comp_self (μ := μ) (fun ω ↦ (T ω, X ω)) (f := Prod.fst) (by fun_prop) + rw [(compProd_map_condDistrib hX).symm] at h_fst + have h_fst' := (Measure.ae_compProd_iff (Kernel.measurableSet_eq _ _)).mp h_fst + filter_upwards [h_prod, h_fst'] with z hz1 hz2 + rw [hz1] + simp only [Kernel.deterministic_apply] at hz2 + change ∀ᵐ y ∂(condDistrib X T μ z), condDistrib T (fun ω ↦ (T ω, X ω)) μ (z, y) = Measure.dirac z + at hz2 + ext t ht + rw [Kernel.compProd_apply ht] + calc ∫⁻ y, condDistrib T (fun ω ↦ (T ω, X ω)) μ (z, y) (Prod.mk y ⁻¹' t) ∂condDistrib X T μ z + _ = ∫⁻ y, (Measure.dirac z) (Prod.mk y ⁻¹' t) ∂condDistrib X T μ z := + lintegral_congr_ae (hz2.mono fun y hy ↦ by simp only [hy]) + _ = ∫⁻ y, (Prod.mk y ⁻¹' t).indicator 1 z ∂condDistrib X T μ z := + lintegral_congr fun y ↦ Measure.dirac_apply' _ (ht.preimage (by fun_prop)) + _ = (condDistrib X T μ z) ((fun y ↦ (y, z)) ⁻¹' t) := by + rw [← lintegral_indicator_one (ht.preimage (by fun_prop : Measurable fun y ↦ (y, z)))] + exact lintegral_congr fun _ ↦ rfl + _ = ((condDistrib X T μ ×ₖ Kernel.id) z) t := by + rw [Kernel.prod_apply, Kernel.id_apply, Measure.prod_apply_symm ht, lintegral_dirac] + +-- proved by Claude, then modified +lemma CondIndepFun.prod_right [StandardBorelSpace α] [StandardBorelSpace β] [Nonempty β] + [StandardBorelSpace γ] [Nonempty γ] [StandardBorelSpace δ] [Nonempty δ] + {X : α → β} {Y : α → γ} {Z : α → δ} + (hX : Measurable X) (hY : Measurable Y) (hZ : Measurable Z) + (h : X ⟂ᵢ[Z, hZ; μ] Y) : + X ⟂ᵢ[Z, hZ; μ] (fun ω ↦ (Y ω, Z ω)) := by + rw [condIndepFun_iff_condDistrib_prod_ae_eq_prodMkRight hY hX hZ, + condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h + rw [condIndepFun_iff_condDistrib_prod_ae_eq_prodMkRight (by fun_prop) hX hZ, + condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] + -- Key: condDistrib (Y, Z) Z μ z = (condDistrib Y Z μ z).map (y ↦ (y, z)) + have h_cond : condDistrib (fun ω ↦ (Y ω, Z ω)) Z μ =ᵐ[μ.map Z] + fun z ↦ (condDistrib Y Z μ z).map (fun y ↦ (y, z)) := by + suffices condDistrib (fun ω ↦ (Y ω, Z ω)) Z μ =ᵐ[μ.map Z] + (condDistrib Y Z μ) ×ₖ Kernel.id by + refine this.trans (ae_of_all _ fun x ↦ ?_) + simp only + rw [Kernel.prod_apply, Kernel.id_apply] + ext s hs + rw [Measure.map_apply (by fun_prop) hs, Measure.prod_apply_symm hs, lintegral_dirac] + exact condDistrib_prod_self_left hY.aemeasurable hZ.aemeasurable + -- Main calculation + calc μ.map (fun x ↦ ((Z x, X x), (Y x, Z x))) + _ = (μ.map (fun x ↦ ((Z x, X x), Y x))).map (fun p ↦ (p.1, (p.2, p.1.1))) := by + rw [Measure.map_map (by fun_prop) (by fun_prop)]; rfl + _ = (μ.map (fun ω ↦ (Z ω, X ω)) ⊗ₘ (condDistrib Y Z μ).prodMkRight β).map + (fun p ↦ (p.1, (p.2, p.1.1))) := by rw [h] + _ = μ.map (fun ω ↦ (Z ω, X ω)) ⊗ₘ (condDistrib (fun ω ↦ (Y ω, Z ω)) Z μ).prodMkRight β := by + ext s hs + rw [Measure.map_apply (by fun_prop) hs, + Measure.compProd_apply (hs.preimage (by fun_prop)), Measure.compProd_apply hs] + have h_cond' : ∀ᵐ p ∂(μ.map (fun ω ↦ (Z ω, X ω))), + condDistrib (fun ω ↦ (Y ω, Z ω)) Z μ p.1 = + (condDistrib Y Z μ p.1).map (fun y ↦ (y, p.1)) := by + have h_fst : (μ.map (fun ω ↦ (Z ω, X ω))).map Prod.fst = μ.map Z := by + rw [Measure.map_map (by fun_prop) (by fun_prop)]; rfl + rw [← h_fst] at h_cond + exact mem_ae_of_mem_ae_map (by fun_prop) h_cond + refine lintegral_congr_ae (h_cond'.mono fun ⟨z, x⟩ hzx ↦ ?_) + simp only [Kernel.prodMkRight_apply, hzx, + Measure.map_apply (by fun_prop : Measurable fun y ↦ (y, z)) + (hs.preimage (by fun_prop : Measurable (Prod.mk (z, x))))] + congr 1 + lemma fst_condDistrib_prod [StandardBorelSpace β] [Nonempty β] (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) (hT : AEMeasurable T μ) : (condDistrib (fun ω ↦ (X ω, Y ω)) T μ).fst =ᵐ[μ.map T] condDistrib X T μ := by @@ -407,6 +455,120 @@ lemma condDistrib_ae_eq_cond [Countable β] [MeasurableSingletonClass β] · congr · exact hb +lemma lintegral_cond {μ : Measure α} (s : Set α) (f : α → ℝ≥0∞) : + ∫⁻ x, f x ∂μ[|s] = (μ s)⁻¹ * ∫⁻ (a : α) in s, f a ∂μ := by + unfold cond + simp [lintegral_smul_measure] + +omit [Nonempty Ω'] in +lemma condDistrib_prod_of_forall_condDistrib_cond [Countable Ω'] [IsFiniteMeasure μ] + (hX : Measurable X) (hY : Measurable Y) (hZ : Measurable Z) + (κ : Kernel (β × Ω') Ω) [IsFiniteKernel κ] + (h_cond : ∀ b, μ (Z ⁻¹' {b}) ≠ 0 → condDistrib Y X μ[|Z ⁻¹' {b}] =ᵐ[μ[|Z ⁻¹' {b}].map X] + (κ.comap (fun ω ↦ (ω, b)) (by fun_prop))) : + condDistrib Y (fun ω ↦ (X ω, Z ω)) μ =ᵐ[μ.map (fun ω ↦ (X ω, Z ω))] κ := by + refine condDistrib_ae_eq_of_measure_eq_compProd _ (by fun_prop) ?_ + ext s hs + suffices ∀ b, (Measure.map (fun x ↦ ((X x, Z x), Y x)) μ) (s ∩ {p | p.1.2 = b}) = + (Measure.map (fun ω ↦ (X ω, Z ω)) μ ⊗ₘ κ) (s ∩ {p | p.1.2 = b}) by + have hs_iUnion : s = ⋃ b, s ∩ {p | p.1.2 = b} := by + ext p + simp only [Set.mem_iUnion, Set.mem_inter_iff, Set.mem_setOf_eq] + grind + have h_disj : Pairwise (Function.onFun Disjoint fun b ↦ s ∩ {p | p.1.2 = b}) := by + intro i j hij + simp only [Set.disjoint_iff_inter_eq_empty] + ext + grind + have h_meas (b : Ω') : MeasurableSet (s ∩ {p | p.1.2 = b}) := + hs.inter ((measurableSet_singleton _).preimage (by fun_prop)) + rw [hs_iUnion, measure_iUnion h_disj h_meas, measure_iUnion h_disj h_meas] + congr with b + exact this b + intro b + by_cases hb : μ (Z ⁻¹' {b}) = 0 + · have h_left : (Measure.map (fun x ↦ ((X x, Z x), Y x)) μ) (s ∩ {p | p.1.2 = b}) = 0 := by + suffices (Measure.map (fun x ↦ ((X x, Z x), Y x)) μ) {p | p.1.2 = b} = 0 from + measure_mono_null Set.inter_subset_right this + rw [Measure.map_apply (by fun_prop)] + · simpa + · exact (measurableSet_singleton _).preimage (by fun_prop) + have h_right : (Measure.map (fun ω ↦ (X ω, Z ω)) μ ⊗ₘ κ) (s ∩ {p | p.1.2 = b}) = 0 := by + suffices (Measure.map (fun ω ↦ (X ω, Z ω)) μ ⊗ₘ κ) {p | p.1.2 = b} = 0 from + measure_mono_null Set.inter_subset_right this + rw [Measure.compProd_apply, lintegral_map] + rotate_left + · exact Kernel.measurable_kernel_prodMk_left + ((measurableSet_singleton _).preimage (by fun_prop)) + · fun_prop + · exact (measurableSet_singleton _).preimage (by fun_prop) + simp only [Set.preimage_setOf_eq] + classical + have h_le : ∫⁻ a, (κ (X a, Z a)) {a_1 | Z a = b} ∂μ ≤ + ∫⁻ a, {a' | Z a' = b}.indicator (fun _ ↦ κ.bound) a ∂μ := by + gcongr with a + by_cases hZ : Z a = b + · simp only [hZ, Set.setOf_true, Set.mem_setOf_eq, Set.indicator_of_mem] + exact κ.measure_le_bound _ _ + · simp [hZ] + refine le_antisymm (h_le.trans ?_) zero_le' + rw [lintegral_indicator] + swap; · exact (measurableSet_singleton _).preimage (by fun_prop) + simp only [lintegral_const, MeasurableSet.univ, Measure.restrict_apply, Set.univ_inter, + nonpos_iff_eq_zero, mul_eq_zero] + exact .inr hb + rw [h_left, h_right] + specialize h_cond b hb + rw [condDistrib_ae_eq_iff_measure_eq_compProd] at h_cond + swap; · fun_prop + rw [Measure.ext_iff] at h_cond + have hs' : MeasurableSet {p : β × Ω | ((p.1, b), p.2) ∈ s} := hs.preimage (by fun_prop) + have h1 := h_cond {p | ((p.1, b), p.2) ∈ s} hs' + have h_indicator : Measurable ({ω' | Z ω' = b}.indicator (fun x ↦ 1)) := + Measurable.indicator (by fun_prop) ((measurableSet_singleton _).preimage (by fun_prop)) + rw [Measure.map_apply] at h1 ⊢ + rotate_left + · fun_prop + · exact hs.inter ((measurableSet_singleton _).preimage (by fun_prop)) + · fun_prop + · exact hs' + rw [cond_apply] at h1 + swap; · exact (measurableSet_singleton _).preimage (by fun_prop) + have h1' : μ (Z ⁻¹' {b} ∩ (fun x ↦ (X x, Y x)) ⁻¹' {p | ((p.1, b), p.2) ∈ s}) = + (μ (Z ⁻¹' {b})) * + (Measure.map X μ[|Z ⁻¹' {b}] ⊗ₘ κ.comap (fun ω ↦ (ω, b)) (by fun_prop)) + {p | ((p.1, b), p.2) ∈ s} := by + rw [← h1, ← mul_assoc, ENNReal.mul_inv_cancel hb (by simp), one_mul] + convert h1' + · ext x + simp only [Set.preimage_inter, Set.preimage_setOf_eq, Set.mem_inter_iff, Set.mem_preimage, + Set.mem_setOf_eq] + grind + · rw [Measure.compProd_apply, Measure.compProd_apply, lintegral_map, lintegral_map] + rotate_left + · exact Kernel.measurable_kernel_prodMk_left hs' + · fun_prop + · apply Kernel.measurable_kernel_prodMk_left + exact hs.inter ((measurableSet_singleton _).preimage (by fun_prop)) + · fun_prop + · exact hs' + · exact hs.inter ((measurableSet_singleton _).preimage (by fun_prop)) + rw [lintegral_cond, ← mul_assoc, ENNReal.mul_inv_cancel hb (by simp), one_mul] + simp only [Set.preimage_inter, Set.preimage_setOf_eq, Kernel.coe_comap, Function.comp_apply] + classical + have h_eq : (fun a ↦ κ (X a, Z a) (Prod.mk (X a, Z a) ⁻¹' s ∩ {a_1 | Z a = b})) = + {a | Z a = b}.indicator + (fun a ↦ κ (X a, b) (Prod.mk (X a, b) ⁻¹' s ∩ {a_1 | Z a = b})) := by + ext a + by_cases hZ : Z a = b <;> simp [hZ] + simp_rw [h_eq] + rw [lintegral_indicator] + swap; · exact (measurableSet_singleton _).preimage (by fun_prop) + refine setLIntegral_congr_fun ((measurableSet_singleton _).preimage (by fun_prop)) fun a ha ↦ ?_ + congr 1 with ω + simp only [Set.mem_inter_iff, Set.mem_preimage, Set.mem_setOf_eq, and_iff_left_iff_imp] + grind + lemma cond_of_indepFun [IsZeroOrProbabilityMeasure μ] (h : IndepFun X T μ) (hX : Measurable X) (hT : Measurable T) {s : Set β} (hs : MeasurableSet s) (hμs : μ (X ⁻¹' s) ≠ 0) : diff --git a/LeanBandits/ForMathlib/CondIndepFun.lean b/LeanBandits/ForMathlib/CondIndepFun.lean new file mode 100644 index 00000000..d361705a --- /dev/null +++ b/LeanBandits/ForMathlib/CondIndepFun.lean @@ -0,0 +1,71 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +import Mathlib.MeasureTheory.Function.FactorsThrough +import Mathlib.Probability.Independence.Basic +import Mathlib.Probability.Independence.Conditional + +/-! # Laws of `stepsUntil` and `rewardByCount` +-/ + +open MeasureTheory ProbabilityTheory Finset +open scoped ENNReal NNReal + +namespace ProbabilityTheory + +variable {α β γ δ γ' δ' : Type*} + {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} + {mδ : MeasurableSpace δ} {mγ' : MeasurableSpace γ'} {mδ' : MeasurableSpace δ'} + [StandardBorelSpace δ'] [Nonempty δ'] [StandardBorelSpace γ'] [Nonempty γ'] + {μ : Measure α} + {X : α → β} {hX : Measurable X} {Y : α → γ} {Z : α → δ} {Y' : α → γ'} {Z' : α → δ'} + +lemma IndepFun.of_measurable (h_indep : Y ⟂ᵢ[μ] Z) + (hY_meas : Measurable[mγ.comap Y] Y') (hZ_meas : Measurable[mδ.comap Z] Z') : + Y' ⟂ᵢ[μ] Z' := by + obtain ⟨φ, hφ_meas, h_eqY⟩ : ∃ φ, Measurable φ ∧ Y' = φ ∘ Y := hY_meas.exists_eq_measurable_comp + obtain ⟨ψ, hψ_meas, h_eqZ⟩ : ∃ ψ, Measurable ψ ∧ Z' = ψ ∘ Z := hZ_meas.exists_eq_measurable_comp + rw [h_eqY, h_eqZ] + exact h_indep.comp hφ_meas hψ_meas + +lemma IndepFun.of_measurable_left + (h_indep : Y ⟂ᵢ[μ] Z) (hY_meas : Measurable[mγ.comap Y] Y') : + Y' ⟂ᵢ[μ] Z := by + obtain ⟨φ, hφ_meas, h_eqY⟩ : ∃ φ, Measurable φ ∧ Y' = φ ∘ Y := hY_meas.exists_eq_measurable_comp + rw [h_eqY] + exact h_indep.comp hφ_meas measurable_id + +lemma IndepFun.of_measurable_right + (h_indep : Y ⟂ᵢ[μ] Z) (hZ_meas : Measurable[mδ.comap Z] Z') : + Y ⟂ᵢ[μ] Z' := by + obtain ⟨ψ, hψ_meas, h_eqZ⟩ : ∃ ψ, Measurable ψ ∧ Z' = ψ ∘ Z := hZ_meas.exists_eq_measurable_comp + rw [h_eqZ] + exact h_indep.comp measurable_id hψ_meas + +variable [StandardBorelSpace α] [IsFiniteMeasure μ] + +lemma CondIndepFun.of_measurable (h_indep : Y ⟂ᵢ[X, hX; μ] Z) + (hY_meas : Measurable[mγ.comap Y] Y') (hZ_meas : Measurable[mδ.comap Z] Z') : + Y' ⟂ᵢ[X, hX; μ] Z' := by + obtain ⟨φ, hφ_meas, h_eqY⟩ : ∃ φ, Measurable φ ∧ Y' = φ ∘ Y := hY_meas.exists_eq_measurable_comp + obtain ⟨ψ, hψ_meas, h_eqZ⟩ : ∃ ψ, Measurable ψ ∧ Z' = ψ ∘ Z := hZ_meas.exists_eq_measurable_comp + rw [h_eqY, h_eqZ] + exact h_indep.comp hφ_meas hψ_meas + +lemma CondIndepFun.of_measurable_left + (h_indep : Y ⟂ᵢ[X, hX; μ] Z) (hY_meas : Measurable[mγ.comap Y] Y') : + Y' ⟂ᵢ[X, hX; μ] Z := by + obtain ⟨φ, hφ_meas, h_eqY⟩ : ∃ φ, Measurable φ ∧ Y' = φ ∘ Y := hY_meas.exists_eq_measurable_comp + rw [h_eqY] + exact h_indep.comp hφ_meas measurable_id + +lemma CondIndepFun.of_measurable_right + (h_indep : Y ⟂ᵢ[X, hX; μ] Z) (hZ_meas : Measurable[mδ.comap Z] Z') : + Y ⟂ᵢ[X, hX; μ] Z' := by + obtain ⟨ψ, hψ_meas, h_eqZ⟩ : ∃ ψ, Measurable ψ ∧ Z' = ψ ∘ Z := hZ_meas.exists_eq_measurable_comp + rw [h_eqZ] + exact h_indep.comp measurable_id hψ_meas + +end ProbabilityTheory diff --git a/LeanBandits/ForMathlib/HasCondDistrib.lean b/LeanBandits/ForMathlib/HasCondDistrib.lean new file mode 100644 index 00000000..a151124e --- /dev/null +++ b/LeanBandits/ForMathlib/HasCondDistrib.lean @@ -0,0 +1,219 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne, Paulo Rauber +-/ +import LeanBandits.ForMathlib.CondDistrib +import Mathlib.Probability.HasLaw + +/-! +# A predicate for having a specified conditional distribution +-/ + +open MeasureTheory + +namespace ProbabilityTheory + +variable {α β γ Ω Ω' : Type*} + {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ} + {mΩ : MeasurableSpace Ω} [StandardBorelSpace Ω] [Nonempty Ω] + {mΩ' : MeasurableSpace Ω'} [StandardBorelSpace Ω'] [Nonempty Ω'] + {μ : Measure α} {X : α → β} {Y : α → Ω} {κ : Kernel β Ω} + +/-- Predicate stating that the conditional distribution of `Y` given `X` under the measure `μ` +is equal to the kernel `κ`. -/ +structure HasCondDistrib (Y : α → Ω) (X : α → β) (κ : Kernel β Ω) + (μ : Measure α) [IsFiniteMeasure μ] : Prop where + aemeasurable_fst : AEMeasurable Y μ := by fun_prop + aemeasurable_snd : AEMeasurable X μ := by fun_prop + condDistrib_eq : condDistrib Y X μ =ᵐ[μ.map X] κ + +attribute [fun_prop] HasCondDistrib.aemeasurable_fst HasCondDistrib.aemeasurable_snd + +lemma hasCondDistrib_fst_prod {Y : α → Ω} {X : α → β} + {κ : Kernel β Ω} + {μ : Measure α} [IsFiniteMeasure μ] {ν : Measure γ} [IsProbabilityMeasure ν] + (h : HasCondDistrib Y X κ μ) : + HasCondDistrib (fun ω ↦ Y ω.1) (fun ω ↦ X ω.1) κ (μ.prod ν) where + aemeasurable_fst := by have := h.aemeasurable_fst; fun_prop + aemeasurable_snd := by have := h.aemeasurable_snd; fun_prop + condDistrib_eq := by + have : ((μ.prod ν).map (fun ω ↦ X ω.1)) = μ.map X := by + conv_rhs => rw [← Measure.fst_prod (μ := μ) (ν := ν), Measure.fst] + rw [AEMeasurable.map_map_of_aemeasurable _ (by fun_prop)] + · rfl + · have := h.aemeasurable_snd + simpa + rw [this] + exact (condDistrib_fst_prod X h.aemeasurable_fst ν).trans h.condDistrib_eq + +lemma HasCondDistrib.comp [IsFiniteMeasure μ] + (h : HasCondDistrib Y X κ μ) {f : Ω → Ω'} (hf : Measurable f) : + HasCondDistrib (fun ω ↦ f (Y ω)) X (κ.map f) μ where + aemeasurable_fst := by have := h.aemeasurable_fst; fun_prop + aemeasurable_snd := by have := h.aemeasurable_snd; fun_prop + condDistrib_eq := by + have h_comp := condDistrib_comp X (Y := Y) (f := f) (mβ := mβ) h.aemeasurable_fst hf + refine h_comp.trans ?_ + have h' := h.condDistrib_eq + filter_upwards [h'] with ω hω + rw [Kernel.map_apply _ hf, hω, Kernel.map_apply _ hf] + +lemma HasCondDistrib.fst {Y : α → Ω × Ω'} {κ : Kernel β (Ω × Ω')} [IsFiniteMeasure μ] + (h : HasCondDistrib Y X κ μ) : + HasCondDistrib (fun ω ↦ (Y ω).1) X κ.fst μ := by + rw [Kernel.fst_eq] + exact HasCondDistrib.comp h measurable_fst + +lemma HasCondDistrib.snd {Y : α → Ω × Ω'} {κ : Kernel β (Ω × Ω')} [IsFiniteMeasure μ] + (h : HasCondDistrib Y X κ μ) : + HasCondDistrib (fun ω ↦ (Y ω).2) X κ.snd μ := by + rw [Kernel.snd_eq] + exact HasCondDistrib.comp h measurable_snd + +lemma HasCondDistrib.comp_right [IsFiniteMeasure μ] [IsFiniteKernel κ] (h : HasCondDistrib Y X κ μ) + (f : β ≃ᵐ γ) : + HasCondDistrib Y (f ∘ X) (κ.comap f.symm (by fun_prop)) μ := by + have hY := h.aemeasurable_fst + have hX := h.aemeasurable_snd + refine ⟨h.aemeasurable_fst, by fun_prop, ?_⟩ + have h_eq := h.condDistrib_eq + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop) _] at h_eq ⊢ + calc μ.map (fun ω ↦ ((f ∘ X) ω, Y ω)) + _ = μ.map ((fun p ↦ (f p.1, p.2)) ∘ fun ω ↦ (X ω, Y ω)) := by congr + _ = (μ.map (fun ω ↦ (X ω, Y ω))).map (fun p ↦ (f p.1, p.2)) := by + rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + _ = (μ.map X ⊗ₘ κ).map (fun p ↦ (f p.1, p.2)) := by rw [h_eq] + _ = μ.map (f ∘ X) ⊗ₘ (κ.comap f.symm (by fun_prop)) := by + -- this is probably very inefficient. + have hX_eq : X = f.symm ∘ (f ∘ X) := by ext; simp + conv_lhs => rw [hX_eq] + rw [← AEMeasurable.map_map_of_aemeasurable, Measure.compProd_eq_comp_prod, + ← Measure.deterministic_comp_eq_map (f := f.symm), ← Measure.deterministic_comp_eq_map] + rotate_left + · fun_prop + · fun_prop + · fun_prop + · fun_prop + rw [← Kernel.comp_deterministic_eq_comap, Measure.compProd_eq_comp_prod] + simp_rw [Measure.comp_assoc] + congr 1 + ext c : 1 + rw [Kernel.comp_apply, Kernel.comp_apply, Kernel.prod_apply, Kernel.comp_apply] + simp only [Kernel.deterministic_apply, Kernel.id_apply, Measure.dirac_bind κ.measurable, + Measure.dirac_bind (Kernel.id ×ₖ κ).measurable, Kernel.prod_apply, + Measure.deterministic_comp_eq_map] + ext s hs + rw [Measure.map_apply (by fun_prop) hs, Measure.prod_apply, Measure.prod_apply, + lintegral_dirac', lintegral_dirac'] + · congr + ext + simp + · exact measurable_measure_prodMk_left hs + · exact measurable_measure_prodMk_left (hs.preimage (by fun_prop)) + · exact hs + · exact hs.preimage (by fun_prop) + +lemma HasCondDistrib.prod_right [IsFiniteMeasure μ] [IsFiniteKernel κ] (h : HasCondDistrib Y X κ μ) + {f : β → γ} (hf : Measurable f) : + HasCondDistrib Y (fun a ↦ (X a, f (X a))) (κ.prodMkRight _) μ := by + have hY := h.aemeasurable_fst + have hX := h.aemeasurable_snd + refine ⟨h.aemeasurable_fst, by fun_prop, ?_⟩ + have h_eq := h.condDistrib_eq + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop) _] at h_eq ⊢ + calc μ.map (fun x ↦ ((X x, f (X x)), Y x)) + _ = (μ.map (fun ω ↦ (X ω, Y ω))).map (fun p ↦ ((p.1, f p.1), p.2)) := by + rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + congr + _ = (μ.map X ⊗ₘ κ).map (fun p ↦ ((p.1, f p.1), p.2)) := by rw [h_eq] + _ = (μ.map X).map (fun a ↦ (a, f a)) ⊗ₘ κ.prodMkRight γ := by + rw [Measure.compProd_eq_comp_prod, Measure.compProd_eq_comp_prod, + ← Measure.deterministic_comp_eq_map (f := fun a ↦ (a, f a)), + ← Measure.deterministic_comp_eq_map, Measure.comp_assoc, Measure.comp_assoc] + swap; · fun_prop + swap; · fun_prop + congr 1 + ext b : 1 + rw [Kernel.comp_apply, Kernel.comp_apply, Kernel.prod_apply, Kernel.deterministic_apply, + Kernel.id_apply, Measure.dirac_bind (Kernel.measurable _), Kernel.prod_apply, + Measure.deterministic_comp_eq_map, Kernel.prodMkRight_apply, Kernel.id_apply] + change Measure.map (Prod.map (fun x ↦ (x, f x)) id) ((Measure.dirac b).prod (κ b)) = + (Measure.dirac (b, f b)).prod (κ b) + rw [← Measure.map_prod_map _ _ (by fun_prop) (by fun_prop), Measure.map_id, + Measure.map_dirac (by fun_prop)] + _ = μ.map (fun a ↦ (X a, f (X a))) ⊗ₘ κ.prodMkRight γ := by + rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + congr + +lemma hasCondDistrib_prod_right_iff [IsFiniteMeasure μ] [IsFiniteKernel κ] (X : α → β) (Y : α → Ω) + {f : β → γ} (hf : Measurable f) : + HasCondDistrib Y (fun a ↦ (X a, f (X a))) (κ.prodMkRight _) μ ↔ HasCondDistrib Y X κ μ := by + refine ⟨fun h ↦ ?_, fun h ↦ h.prod_right hf⟩ + have hX : AEMeasurable X μ := by + have := h.aemeasurable_snd + have h_eq : X = (fun p ↦ p.1) ∘ (fun a ↦ (X a, f (X a))) := by ext; simp + rw [h_eq] + exact Measurable.comp_aemeasurable (by fun_prop) (by fun_prop) + have hY := h.aemeasurable_fst + refine ⟨by fun_prop, by fun_prop, ?_⟩ + have h_eq := h.condDistrib_eq + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop) _] at h_eq ⊢ + calc μ.map (fun x ↦ (X x, Y x)) + _ = (μ.map (fun ω ↦ ((X ω, f (X ω)), Y ω))).map (fun p ↦ (p.1.1, p.2)) := by + rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + congr + _ = (μ.map (fun a ↦ (X a, f (X a))) ⊗ₘ κ.prodMkRight γ).map (fun p ↦ (p.1.1, p.2)) := by rw [h_eq] + _ = ((μ.map X).map (fun a ↦ (a, f a)) ⊗ₘ κ.prodMkRight γ).map (fun p ↦ (p.1.1, p.2)) := by + rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + congr + _ = μ.map X ⊗ₘ κ := by + simp_rw [Measure.compProd_eq_comp_prod, + ← Measure.deterministic_comp_eq_map (f := fun a ↦ (a, f a)) (by fun_prop), + ← Measure.deterministic_comp_eq_map (f := fun p : (β × γ) × Ω ↦ (p.1.1, p.2)) (by fun_prop), + Measure.comp_assoc] + congr 1 + ext b : 1 + rw [Kernel.comp_apply, Kernel.comp_apply, Kernel.prod_apply, Kernel.id_apply, + Kernel.deterministic_apply, Measure.dirac_bind (Kernel.measurable _), + Kernel.prod_apply, Measure.deterministic_comp_eq_map, Kernel.prodMkRight_apply, + Kernel.id_apply] + change Measure.map (Prod.map (fun x ↦ x.1) id) ((Measure.dirac (b, f b)).prod (κ b)) = _ + rw [← Measure.map_prod_map _ _ (by fun_prop) (by fun_prop), Measure.map_id, + Measure.map_dirac (by fun_prop)] + +lemma HasLaw.prod_of_hasCondDistrib {P : Measure β} [IsFiniteMeasure μ] [IsSFiniteKernel κ] + (h1 : HasLaw X P μ) (h2 : HasCondDistrib Y X κ μ) : + HasLaw (fun ω ↦ (X ω, Y ω)) (P ⊗ₘ κ) μ := by + have hX := h1.aemeasurable + have hY := h2.aemeasurable_fst + refine ⟨by fun_prop, ?_⟩ + rw [← compProd_map_condDistrib (by fun_prop), h1.map_eq] + refine Measure.compProd_congr ?_ + rw [← h1.map_eq] + exact h2.condDistrib_eq + +lemma HasCondDistrib.prod [IsFiniteMeasure μ] [IsFiniteKernel κ] + {Z : α → Ω'} {η : Kernel (β × Ω) Ω'} [IsFiniteKernel η] + (h1 : HasCondDistrib Y X κ μ) (h2 : HasCondDistrib Z (fun ω ↦ (X ω, Y ω)) η μ) : + HasCondDistrib (fun ω ↦ (Y ω, Z ω)) X (κ ⊗ₖ η) μ := by + have hX := h1.aemeasurable_snd + have hY := h1.aemeasurable_fst + have hZ := h2.aemeasurable_fst + refine ⟨by fun_prop, by fun_prop, ?_⟩ + have h_condDistrib_Y := h1.condDistrib_eq + have h_condDistrib_Z := h2.condDistrib_eq + have h_prod := condDistrib_prod_left hY hZ hX + have h_prod' : 𝓛[fun ω ↦ (Y ω, Z ω) | X; μ] =ᵐ[μ.map X] (κ ⊗ₖ 𝓛[Z | fun ω ↦ (X ω, Y ω); μ]) := by + filter_upwards [h_condDistrib_Y, h_prod] with ω hω₁ hω₂ + rw [hω₂] + ext s hs + rw [Kernel.compProd_apply hs, Kernel.compProd_apply hs] + simp [hω₁] + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] + at h_condDistrib_Z h_condDistrib_Y ⊢ + rw [← Measure.compProd_assoc', ← h_condDistrib_Y, ← h_condDistrib_Z, + AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + rfl + +end ProbabilityTheory diff --git a/LeanBandits/ForMathlib/IndepFun.lean b/LeanBandits/ForMathlib/IndepFun.lean index 15901bd4..d0834089 100644 --- a/LeanBandits/ForMathlib/IndepFun.lean +++ b/LeanBandits/ForMathlib/IndepFun.lean @@ -9,6 +9,35 @@ variable {α Ω Ω' E ι : Type*} [Countable ι] {mα : MeasurableSpace α} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} {mE : MeasurableSpace E} {μ ν : Measure Ω} +@[simp] +lemma indepFun_zero_measure {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} (X : α → β) (Y : α → γ) : + X ⟂ᵢ[(0 : Measure α)] Y := by + simp [indepFun_iff_measure_inter_preimage_eq_mul] + +lemma indepFun_cond_of_indepFun {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {μ : Measure α} + {X : α → β} {Y : α → γ} (hXY : X ⟂ᵢ[μ] Y) (hY : Measurable Y) {s : Set γ} + (hs : MeasurableSet s) : + X ⟂ᵢ[μ[|Y ⁻¹' s]] Y := by + by_cases h_zero : μ[|Y ⁻¹' s] = 0 + · simp [h_zero] + rw [cond_eq_zero] at h_zero + push_neg at h_zero -- `h_zero : μ (Y ⁻¹' s) ≠ ⊤ ∧ μ (Y ⁻¹' s) ≠ 0` + rw [indepFun_iff_measure_inter_preimage_eq_mul] at hXY ⊢ + intro u t hu ht + rw [cond_apply (hs.preimage hY), cond_apply (hs.preimage hY), cond_apply (hs.preimage hY)] + have h_eq : Y ⁻¹' s ∩ (X ⁻¹' u ∩ Y ⁻¹' t) = X ⁻¹' u ∩ Y ⁻¹' (s ∩ t) := by grind + have hsu : μ (X ⁻¹' u ∩ Y ⁻¹' s) = μ (X ⁻¹' u) * μ (Y ⁻¹' s) := hXY u s hu hs + rw [Set.inter_comm] at hsu + have hust : μ (X ⁻¹' u ∩ Y ⁻¹' (s ∩ t)) = μ (X ⁻¹' u) * μ (Y ⁻¹' (s ∩ t)) := + hXY u (s ∩ t) hu (hs.inter ht) + rw [hsu, h_eq, hust] + simp_rw [mul_assoc] + congr 1 + rw [← mul_assoc (μ (Y ⁻¹' s)), ENNReal.mul_inv_cancel h_zero.2 h_zero.1, one_mul] + congr + lemma iIndepFun_nat_iff_forall_indepFun [IsProbabilityMeasure μ] {X : ℕ → Ω → E} (hX : ∀ n, AEMeasurable (X n) μ) : iIndepFun X μ ↔ ∀ n, X (n + 1) ⟂ᵢ[μ] fun ω (i : Iic n) ↦ X i ω := by diff --git a/LeanBandits/ForMathlib/KernelRepresentation.lean b/LeanBandits/ForMathlib/KernelRepresentation.lean new file mode 100644 index 00000000..aebf144b --- /dev/null +++ b/LeanBandits/ForMathlib/KernelRepresentation.lean @@ -0,0 +1,154 @@ +/- +Copyright (c) 2025 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré, Rémy Degenne +-/ + +import Mathlib.Analysis.SpecialFunctions.Sigmoid +import Mathlib.MeasureTheory.Constructions.UnitInterval +import Mathlib.Order.CompletePartialOrder +import Mathlib.Probability.CDF + +-- copied from PR #30112 + +/-! +# Representation of kernels + +This file contains results about isolation of kernels randomness. In particular, it shows that, +when the target space is a standard Borel space, any Markov kernel can be represented as the image +of the uniform measure on `[0,1]` by a deterministic map. It corresponds to Lemma 4.22 in +"Foundations of Modern Probability" by Olav Kallenberg, 2021. + +## Statements + +* `ProbabilityTheory.Kernel.unitInterval_representation`: + for a Markov kernel `κ : Kernel α I`, there exists a jointly measurable function + `f : α → I → I` such that for all `a : α`, `volume.map (f a) = κ a`. + +* `ProbabilityTheory.Kernel.embedding_representation`: + for a measurable embedding `g : β → I` and a Markov kernel `κ : Kernel α β`, + there exists a jointly measurable function `f : α → I → β` such that for all `a : α`, + `volume.map (f a) = κ a`. + +* `ProbabilityTheory.Kernel.representation`: + for a Markov kernel `κ : Kernel α β` with `β` a standard Borel space, + there exists a jointly measurable function `f : α → I → β` such that for all `a : α`, + `volume.map (f a) = κ a`. + This is a consequence of `ProbabilityTheory.Kernel.embedding_representation` and the fact that + any standard Borel space can be embedded in `ℝ`, and then composed with `unitInterval.sigmoid`. +-/ + +open MeasureTheory ProbabilityTheory Set ENNReal unitInterval Filter Topology Function + +namespace ProbabilityTheory.Kernel + +variable {α : Type*} [MeasurableSpace α] + +lemma unitInterval_representation (κ : Kernel α I) [IsMarkovKernel κ] : + ∃ (f : α → I → I), Measurable (uncurry f) ∧ ∀ a, volume.map (f a) = κ a := by + let f := fun s (t : I) ↦ sSup {x | (κ s).real (Icc 0 x) < t} + have measurable_f : Measurable (uncurry f) := by + refine measurable_of_Ioi fun a ↦ ?_ + simp only [preimage, uncurry, mem_Ioi] + have h_monotone s : Monotone (fun x ↦ (κ s).real (Icc 0 x)) := by + intro x y hxy + suffices h : Icc 0 x ⊆ Icc 0 y from measureReal_mono h + exact Icc_subset_Icc_right hxy + have sSup_eq_iUnion_rat : {x : α × I | a < f x.1 x.2} = ⋃ (q : ℚ), ⋃ (hqI : ↑q ∈ I), + ⋃ (_ : a < (q : ℝ)), {e | (κ e.1).real (Icc 0 ⟨q, hqI⟩) < e.2} := by + ext e + simp only [f] + constructor + · intro (he : a < sSup {x | (κ e.1).real (Icc 0 x) < e.2}) + simp_rw [Set.mem_iUnion] + rw [lt_sSup_iff] at he + obtain ⟨y, y_mem, (hy : a.1 < y.1)⟩ := he + obtain ⟨q, hqa, hqy⟩ := exists_rat_btwn hy + have q_in_I : (q : ℝ) ∈ I := ⟨a.2.1.trans hqa.le, hqy.le.trans y.2.2⟩ + refine ⟨q, q_in_I, hqa, ?_⟩ + exact lt_of_lt_of_le' y_mem (h_monotone e.1 hqy.le) + · intro he + simp_all only [lt_sSup_iff, Set.mem_iUnion] + obtain ⟨q, q_in_I, hqa, h⟩ := he + exact ⟨⟨q, q_in_I⟩, h, hqa⟩ + rw [sSup_eq_iUnion_rat] + refine MeasurableSet.iUnion (fun b ↦ MeasurableSet.iUnion + (fun bI ↦ MeasurableSet.iUnion (fun _ ↦ ?_))) + refine measurableSet_lt ?_ measurable_snd.subtype_val + simp_rw [measureReal_def] + have hκ := κ.measurable_coe (s := Icc 0 ⟨b, bI⟩) measurableSet_Icc + fun_prop + refine ⟨f, measurable_f, fun a ↦ (volume.map (f a)).ext_of_Iic (κ a) fun x ↦ ?_⟩ + rw [volume.map_apply measurable_f.of_uncurry_left measurableSet_Iic, preimage] + simp only [mem_Iic] + have Iic_to_Icc : Iic x = Icc 0 x := by ext; simp + rw [Iic_to_Icc] + clear Iic_to_Icc + rw [← ofReal_measureReal (measure_ne_top (κ a) _)] + have κ_in_I : ((κ a).real (Icc 0 x)) ∈ I := ⟨measureReal_nonneg, measureReal_le_one⟩ + rw [← volume_Iic ⟨_, κ_in_I⟩] + congr with ξ + constructor + swap + · intro (hξ : ξ ≤ (κ a).real (Icc 0 x)) + simp only [sSup_le_iff, f] + intro c hc + have le1 := lt_of_le_of_lt' hξ hc + by_contra h + push_neg at h + have le2 : (κ a).real (Icc 0 x) ≤ (κ a).real (Icc 0 c) := by + suffices h : Icc 0 x ⊆ Icc 0 c from measureReal_mono h + refine (Icc_subset_Icc_iff unitInterval.nonneg').mpr ?_ + exact ⟨nonneg', h.le⟩ + linarith + · intro (hξ : f a ξ ≤ x) + change ξ ≤ (κ a).real (Icc 0 x) + by_cases hx : x = 1 + · simp [hx, ← univ_eq_Icc, ξ.2.2] + let g := fun y ↦ (κ a).real (Icc 0 y) + letI nebot : NeBot (𝓝[>] x) := by + refine nhdsGT_neBot_of_exists_gt ?_ + use 1 + exact lt_of_le_of_ne x.2.2 hx + refine le_of_tendsto_of_tendsto (b := 𝓝[>] x) (g := g) continuousWithinAt_const ?_ ?_ + · let h := cdf ((κ a).map Subtype.val) + have h_continuousWithinAt := continuousWithinAt_Ioi_iff_Ici.mpr (h.right_continuous x) + simp_rw [g, ← unitInterval.cdf_eq_real (κ a)] + exact h_continuousWithinAt.comp (Continuous.continuousWithinAt (by fun_prop)) (fun y hy ↦ hy) + · apply eventually_nhdsWithin_of_forall + intro y hy + by_contra h + push_neg at h + simp only [sSup_le_iff, f] at hξ + specialize hξ y h + replace hξ : y.1 ≤ x.1 := hξ + have : y.1 > x.1 := hy + linarith + +lemma embedding_representation {β : Type*} [Nonempty β] [MeasurableSpace β] {g : β → I} + (hg : MeasurableEmbedding g) (κ : Kernel α β) [IsMarkovKernel κ] : + ∃ (f : α → I → β), Measurable (uncurry f) ∧ ∀ a, volume.map (f a) = κ a := by + have hκg : IsMarkovKernel (κ.map g) := Kernel.IsMarkovKernel.map κ hg.measurable + classical + have hg'κ : κ = (κ.map g).map hg.invFun := by + rw [← Kernel.map_comp_right _ hg.measurable (by fun_prop), LeftInverse.id hg.leftInverse_invFun, + Kernel.map_id] + obtain ⟨f', hf', hf'κ⟩ := (κ.map g).unitInterval_representation + refine ⟨fun a u ↦ hg.invFun (f' a u), by fun_prop, fun a ↦ ?_⟩ + rw [hg'κ, Kernel.map_apply _ (by fun_prop), ← hf'κ, Measure.map_map (by fun_prop) (by fun_prop)] + rfl + +theorem representation {β : Type*} [Nonempty β] [MeasurableSpace β] [StandardBorelSpace β] + (κ : Kernel α β) [IsMarkovKernel κ] : + ∃ (f : α → I → β), Measurable (uncurry f) ∧ ∀ a, volume.map (f a) = κ a := + κ.embedding_representation (measurableEmbedding_sigmoid_comp_embeddingReal β) + +end ProbabilityTheory.Kernel + +theorem ProbabilityTheory.representation_measure {β : Type*} {mβ : MeasurableSpace β} + [Nonempty β] [StandardBorelSpace β] + (μ : Measure β) [IsProbabilityMeasure μ] : + ∃ (f : I → β), Measurable f ∧ volume.map f = μ := by + obtain ⟨f, hf_meas, hf_map⟩ := Kernel.representation (Kernel.const Unit μ) + specialize hf_map ⟨⟩ + exact ⟨f ⟨⟩, by fun_prop, by simpa⟩ diff --git a/LeanBandits/ForMathlib/StandardBorel.lean b/LeanBandits/ForMathlib/StandardBorel.lean new file mode 100644 index 00000000..e31407b7 --- /dev/null +++ b/LeanBandits/ForMathlib/StandardBorel.lean @@ -0,0 +1,18 @@ +/- +Copyright (c) 2025 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +import Mathlib.MeasureTheory.Constructions.Polish.Basic + +/-! +# Properties of standard Borel spaces +-/ + +open MeasureTheory + +variable {Ω : Type*} {mΩ : MeasurableSpace Ω} + +instance [StandardBorelSpace Ω] : MeasurableEq Ω := by + letI := upgradeStandardBorel Ω + infer_instance diff --git a/LeanBandits/ForMathlib/Traj.lean b/LeanBandits/ForMathlib/Traj.lean index 8df53a22..d9fdf7e2 100644 --- a/LeanBandits/ForMathlib/Traj.lean +++ b/LeanBandits/ForMathlib/Traj.lean @@ -1,12 +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 LeanBandits.ForMathlib.HasCondDistrib import Mathlib.Probability.Kernel.IonescuTulcea.Traj -import Mathlib.Probability.Kernel.CondDistrib -import LeanBandits.ForMathlib.CondDistrib +import Mathlib.Probability.Process.FiniteDimensionalLaws open Filter Finset Function MeasurableEquiv MeasurableSpace MeasureTheory Preorder ProbabilityTheory -variable {X : ℕ → Type*} [∀ n, MeasurableSpace (X n)] -variable {κ : (n : ℕ) → Kernel (Π i : Iic n, X i) (X (n + 1))} [∀ n, IsMarkovKernel (κ n)] -variable {μ₀ : Measure (X 0)} [IsProbabilityMeasure μ₀] +variable {Ω : Type*} {mΩ : MeasurableSpace Ω} {P : Measure Ω} [IsFiniteMeasure P] + {X : ℕ → Type*} [∀ n, MeasurableSpace (X n)] + {κ : (n : ℕ) → Kernel (Π i : Iic n, X i) (X (n + 1))} [∀ n, IsMarkovKernel (κ n)] + {μ₀ : Measure (X 0)} [IsProbabilityMeasure μ₀] section MeasurableEquiv @@ -27,4 +33,102 @@ lemma traj_zero_map_eval_zero : rw [← Kernel.traj_map_frestrictLe, ← Kernel.map_comp_right _ (by fun_prop) (by fun_prop)] rfl +/-- Measurable equivalence between a product up to `n + 1` and the pair of the product up to `n` and +the space at `n + 1`. -/ +def _root_.MeasurableEquiv.IicSuccProd (X : ℕ → Type*) [∀ n, MeasurableSpace (X n)] (n : ℕ) : + MeasurableEquiv (Π i : Iic (n + 1), X i) ((Π i : Iic n, X i) × X (n + 1)) := + (MeasurableEquiv.IicProdIoc (Nat.le_succ n)).symm.trans + (MeasurableEquiv.prodCongr (MeasurableEquiv.refl _) (MeasurableEquiv.piSingleton n).symm) + +lemma symm_IicSuccProd (n : ℕ) : + (MeasurableEquiv.IicSuccProd X n).symm = + (MeasurableEquiv.prodCongr (MeasurableEquiv.refl _) (MeasurableEquiv.piSingleton n)).trans + (MeasurableEquiv.IicProdIoc (Nat.le_succ n)) := rfl + +@[simp] +lemma MeasurableEquiv.IicSuccProd_apply (n : ℕ) (h : Π i : Iic (n + 1), X i) : + MeasurableEquiv.IicSuccProd X n h = (fun i : Iic n ↦ h ⟨i.1, by grind⟩, h ⟨n + 1, by simp⟩) := + rfl + +lemma MeasurableEquiv.coe_prodCongr {α β γ δ : Type*} + {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {mδ : MeasurableSpace δ} + (e₁ : MeasurableEquiv α β) (e₂ : MeasurableEquiv γ δ) : + (MeasurableEquiv.prodCongr e₁ e₂ : (α × γ) → (β × δ)) = Prod.map e₁ e₂ := rfl + +lemma MeasurableEquiv.coe_refl {α : Type*} {mα : MeasurableSpace α} : + (MeasurableEquiv.refl α : α → α) = id := rfl + +theorem hasLaw_Iic_of_forall_hasCondDistrib [∀ n, StandardBorelSpace (X n)] [∀ n, Nonempty (X n)] + {Y : (n : ℕ) → Ω → X n} (h0 : HasLaw (Y 0) μ₀ P) + (h_condDistrib : ∀ n, HasCondDistrib (Y (n + 1)) (fun ω ↦ fun i : Iic n ↦ Y i ω) (κ n) P) + (n : ℕ) : + HasLaw (fun ω (i : Iic n) ↦ Y i ω) + ((partialTraj κ 0 n) ∘ₘ (μ₀.map (MeasurableEquiv.piUnique _).symm)) P := by + induction n with + | zero => + simp only [piUnique_symm_apply, partialTraj_self, Measure.id_comp] + rw [← h0.map_eq, AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + constructor + · have h_meas := h0.aemeasurable + have : (fun ω (i : Iic 0) ↦ Y i ω) = (MeasurableEquiv.piUnique _).symm ∘ (Y 0) := by + ext ω i + simp only [piUnique_symm_apply, Function.comp_apply] + rw [Unique.eq_default i] + simp [coe_default_Iic_zero] + rw [this] + exact AEMeasurable.comp_aemeasurable (by fun_prop) h_meas + · congr + ext ω i + simp only [Function.comp_apply] + rw [Unique.eq_default i] + simp [coe_default_Iic_zero] + | succ n hn => + specialize h_condDistrib n + have h_law := hn.prod_of_hasCondDistrib h_condDistrib + have : (fun ω (i : Iic (n + 1)) ↦ Y i ω) = + (MeasurableEquiv.IicSuccProd X n).symm ∘ + (fun ω ↦ (fun i : Iic n ↦ Y i ω, Y (n + 1) ω)) := by + suffices (MeasurableEquiv.IicSuccProd X n) ∘ (fun ω (i : Iic (n + 1)) ↦ Y i ω) = + (fun ω ↦ (fun i : Iic n ↦ Y i ω, Y (n + 1) ω)) by + rw [← this, ← Function.comp_assoc, MeasurableEquiv.symm_comp_self] + simp + ext ω : 1 + simp + rw [this] + refine HasLaw.comp ⟨by fun_prop, ?_⟩ h_law + rw [Measure.compProd_eq_comp_prod, partialTraj_succ_eq_comp (by simp), Measure.comp_assoc, + ← Measure.deterministic_comp_eq_map (by fun_prop), Measure.comp_assoc] + congr 1 + rw [← Kernel.comp_assoc] + congr + rw [Kernel.deterministic_comp_eq_map, partialTraj_succ_self, symm_IicSuccProd] + rw [MeasurableEquiv.coe_trans, MeasurableEquiv.coe_prodCongr] + rw [Kernel.map_comp_right _ (by fun_prop) (by fun_prop), + ← Kernel.map_prod_map _ _ (by fun_prop) (by fun_prop)] + congr + simp [MeasurableEquiv.coe_refl] + +omit [IsProbabilityMeasure μ₀] in +lemma trajMeasure_map_frestrictLe (n : ℕ) : + (trajMeasure μ₀ κ).map (frestrictLe n) = + (partialTraj κ 0 n) ∘ₘ (μ₀.map (MeasurableEquiv.piUnique _).symm) := by + rw [trajMeasure, ← Measure.deterministic_comp_eq_map (by fun_prop), Measure.comp_assoc, + Kernel.deterministic_comp_eq_map, traj_map_frestrictLe] + +-- todo: switch to `HasLaw` +/-- Uniqueness of `trajMeasure`. -/ +theorem eq_trajMeasure [∀ n, StandardBorelSpace (X n)] [∀ n, Nonempty (X n)] + {Y : (n : ℕ) → Ω → X n} (hY_meas : ∀ n, Measurable (Y n)) + (h0 : HasLaw (Y 0) μ₀ P) + (h_condDistrib : ∀ n, HasCondDistrib (Y (n + 1)) (fun ω ↦ fun i : Iic n ↦ Y i ω) (κ n) P) : + P.map (fun ω n ↦ Y n ω) = trajMeasure μ₀ κ := by + refine IsProjectiveLimit.unique (P := fun (J : Finset ℕ) ↦ P.map (fun ω (i : J) ↦ Y i ω)) ?_ ?_ + · exact isProjectiveLimit_map (by fun_prop) + rw [isProjectiveLimit_nat_iff] + swap; · exact isProjectiveMeasureFamily_map_restrict (by fun_prop) + intro n + rw [(hasLaw_Iic_of_forall_hasCondDistrib h0 h_condDistrib n).map_eq, + trajMeasure_map_frestrictLe] + end ProbabilityTheory.Kernel diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean deleted file mode 100644 index baac2523..00000000 --- a/LeanBandits/RewardByCountMeasure.lean +++ /dev/null @@ -1,347 +0,0 @@ -/- -Copyright (c) 2025 Rémy Degenne. All rights reserved. -Released under Apache 2.0 license as described in the file LICENSE. -Authors: Rémy Degenne --/ -import LeanBandits.Bandit.Regret -import LeanBandits.ForMathlib.IndepFun -import Mathlib.Probability.IdentDistribIndep - -/-! # Laws of `stepsUntil` and `rewardByCount` --/ - -open MeasureTheory ProbabilityTheory Finset Learning -open scoped ENNReal NNReal - -namespace Bandits - -variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [MeasurableSingletonClass α] - -lemma integrable_pullCount {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] - (a : α) (n : ℕ) : - Integrable (fun ω ↦ (pullCount a n ω : ℝ)) (Bandit.trajMeasure alg ν) := by - refine integrable_of_le_of_le (g₁ := 0) (g₂ := fun _ ↦ n) (by fun_prop) - (ae_of_all _ fun ω ↦ by simp) (ae_of_all _ fun ω ↦ ?_) (integrable_const _) (integrable_const _) - simp only [Nat.cast_le] - exact pullCount_le a n ω - -variable {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] - -omit [DecidableEq α] [MeasurableSingletonClass α] in -lemma hasLaw_Z (a : α) (m : ℕ) : - HasLaw (fun ω ↦ ω.2 m a) (ν a) (Bandit.measure alg ν) where - map_eq := by - calc (Bandit.measure alg ν).map (fun ω ↦ ω.2 m a) - _ = ((Bandit.measure alg ν).snd).map (fun ω ↦ ω m a) := by - rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)] - rfl - _ = (Bandit.streamMeasure ν).map (fun ω ↦ ω m a) := by simp - _ = ((Measure.infinitePi fun _ ↦ Measure.infinitePi ν).map (fun ω ↦ ω m)).map - (fun ω ↦ ω a) := by - rw [Bandit.streamMeasure, Measure.map_map (by fun_prop) (by fun_prop)] - rfl - _ = ν a := by simp_rw [(measurePreserving_eval_infinitePi _ _).map_eq] - -/-- Law of `Y` conditioned on the event `s`.-/ -notation "𝓛[" Y " | " s "; " μ "]" => Measure.map Y (μ[|s]) -/-- Law of `Y` conditioned on the event that `X` is in `s`. -/ -notation "𝓛[" Y " | " X " in " s "; " μ "]" => Measure.map Y (μ[|X ⁻¹' s]) -/-- Law of `Y` conditioned on the event that `X` equals `x`. -/ -notation "𝓛[" Y " | " X " ← " x "; " μ "]" => Measure.map Y (μ[|X ⁻¹' {x}]) - -local notation "𝔓t" => Bandit.trajMeasure alg ν -local notation "𝔓" => Bandit.measure alg ν - -omit [DecidableEq α] [MeasurableSingletonClass α] in -lemma condDistrib_reward'' [StandardBorelSpace α] [Nonempty α] (n : ℕ) : - 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1; 𝔓] - =ᵐ[(𝔓).map (fun ω ↦ arm n ω.1)] ν := by - have h_ra' : 𝓛[reward n | arm n; 𝔓t] =ᵐ[(𝔓t).map (arm n)] ν := condDistrib_reward alg ν n - have h_law : (𝔓).map (fun ω ↦ arm n ω.1) = (𝔓t).map (arm n) := by - rw [← Bandit.fst_measure, Measure.fst, Measure.map_map (by fun_prop) (by fun_prop)] - rfl - rw [h_law] - have h_prod : 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1; 𝔓] - =ᵐ[(𝔓t).map (arm n)] 𝓛[reward n | arm n; 𝔓t] := - condDistrib_fst_prod _ (by fun_prop) _ - filter_upwards [h_ra', h_prod] with ω h_eq h_prod - rw [h_prod, h_eq] - -omit [DecidableEq α] in -lemma reward_cond_arm [StandardBorelSpace α] [Nonempty α] [Countable α] (a : α) (n : ℕ) - (hμa : (𝔓).map (fun ω ↦ arm n ω.1) {a} ≠ 0) : - 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1 ← a; 𝔓] = ν a := by - have h_ra : 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1; 𝔓] =ᵐ[(𝔓).map (fun ω ↦ arm n ω.1)] ν := - condDistrib_reward'' n - have h_eq := condDistrib_ae_eq_cond (μ := 𝔓) - (X := fun ω ↦ arm n ω.1) (Y := fun ω ↦ reward n ω.1) (by fun_prop) (by fun_prop) - rw [Filter.EventuallyEq, ae_iff_of_countable] at h_ra h_eq - specialize h_ra a hμa - specialize h_eq a hμa - rw [h_ra] at h_eq - exact h_eq.symm - --- after the Mathlib stopping time refactor, we will be able to prove that stepsUntil is a --- stopping time -lemma measurable_comap_indicator_stepsUntil_eq (a : α) (m n : ℕ) : - Measurable[MeasurableSpace.comap (fun ω : ℕ → α × ℝ ↦ (hist (n-1) ω, arm n ω)) inferInstance] - ({ω | stepsUntil a m ω = ↑n}.indicator fun _ ↦ 1) := by - let k : ((Iic (n - 1) → α × ℝ) × α) → (ℕ → α × ℝ) := fun x i ↦ - if hi : i ∈ Iic (n - 1) then (x.1 ⟨i, hi⟩) else if i = n then (x.2, 0) else (a, 0) - have hk : Measurable k := by - unfold k - rw [measurable_pi_iff] - intro i - split_ifs <;> fun_prop - let φ : ((Iic (n - 1) → α × ℝ) × α) → ℕ := fun x ↦ if stepsUntil a m (k x) = ↑n then 1 else 0 - have hφ : Measurable φ := - Measurable.ite ((measurableSet_singleton _).preimage (by fun_prop)) (by fun_prop) (by fun_prop) - suffices {ω | stepsUntil a m ω = ↑n}.indicator (fun x ↦ 1) - = φ ∘ fun ω ↦ (hist (n - 1) ω, arm n ω) from this ▸ measurable_comp_comap _ hφ - ext ω - classical - simp only [Set.indicator_apply, Set.mem_setOf_eq, Function.comp_apply, φ] - congr 1 - rw [stepsUntil_eq_congr] - intro i hin - simp only [arm, mem_Iic, hist, dite_eq_ite, k, action] - grind - -lemma measurable_indicator_stepsUntil_eq (a : α) (m n : ℕ) : - Measurable ({ω : ℕ → α × ℝ | stepsUntil a m ω = ↑n}.indicator fun _ ↦ 1) := by - refine (measurable_comap_indicator_stepsUntil_eq a m n).mono ?_ le_rfl - refine Measurable.comap_le ?_ - fun_prop - -lemma measurableSet_stepsUntil_eq (a : α) (m n : ℕ) : - MeasurableSet[MeasurableSpace.comap (fun ω : ℕ → α × ℝ ↦ (hist (n-1) ω, arm n ω)) inferInstance] - {ω : ℕ → α × ℝ | stepsUntil a m ω = ↑n} := by - let mProd := MeasurableSpace.comap (fun ω : ℕ → α × ℝ ↦ (hist (n-1) ω, arm n ω)) inferInstance - suffices Measurable[mProd] ({ω | stepsUntil a m ω = ↑n}.indicator fun x ↦ 1) by - rwa [measurable_indicator_const_iff] at this - exact measurable_comap_indicator_stepsUntil_eq a m n - -lemma condIndepFun_reward_stepsUntil_arm' [StandardBorelSpace α] [Countable α] [Nonempty α] - (a : α) (m n : ℕ) (hm : m ≠ 0) : - reward n ⟂ᵢ[arm n, measurable_arm n; 𝔓t] {ω | stepsUntil a m ω = ↑n}.indicator (fun _ ↦ 1) := by - -- the indicator of `stepsUntil ... = n` is a function of - -- `hist (n-1)` and `arm n`. - -- It thus suffices to prove the independence of `reward n` and `hist (n-1)` conditionally - -- on `arm n`. - by_cases hn : n = 0 - · simp only [hn, CharP.cast_eq_zero] - simp only [stepsUntil_eq_zero_iff, hm, ne_eq, false_and, false_or] - by_cases hm1 : m = 1 - · simp only [hm1, true_and] - have h_indep := condIndepFun_self_right (X := reward 0) (Z := arm 0) - (mβ := inferInstance) (mβ' := inferInstance) (μ := 𝔓t) - (by fun_prop) (by fun_prop) - have : {ω : ℕ → α × ℝ | action 0 ω = a}.indicator (fun x ↦ 1) - = {b | b = a}.indicator (fun _ ↦ 1) ∘ action 0 := by ext; simp [Set.indicator] - rw [this] - exact h_indep.comp measurable_id (by fun_prop) - · simp only [hm1, false_and, Set.setOf_false, Set.indicator_empty] - exact condIndepFun_const_right (reward 0) 0 - have h_indep : reward n ⟂ᵢ[arm n, measurable_arm n; 𝔓t] hist (n - 1) := by - convert condIndepFun_reward_hist_arm (alg := alg) (ν := ν) (n - 1) - <;> rw [Nat.sub_add_cancel (by grind)] - have h_indep' : reward n ⟂ᵢ[arm n, measurable_arm n; 𝔓t] fun ω ↦ (hist (n - 1) ω, arm n ω) := - h_indep.prod_right (by fun_prop) (by fun_prop) (by fun_prop) - obtain ⟨φ, hφ_meas, h_eq⟩ : ∃ φ : ((Iic (n - 1) → α × ℝ) × α) → ℕ, Measurable φ ∧ - {ω | stepsUntil a m ω = ↑n}.indicator (fun _ ↦ 1) = φ ∘ (fun ω ↦ (hist (n - 1) ω, arm n ω)) := - (measurable_comap_indicator_stepsUntil_eq a m n).exists_eq_measurable_comp - rw [h_eq] - exact h_indep'.comp measurable_id hφ_meas - -lemma condIndepFun_reward_stepsUntil_arm [StandardBorelSpace α] [Countable α] [Nonempty α] - (a : α) (m n : ℕ) (hm : m ≠ 0) : - CondIndepFun (mα.comap (fun ω ↦ arm n ω.1)) ((measurable_arm n).comp measurable_fst).comap_le - (fun ω ↦ reward n ω.1) ({ω | stepsUntil a m ω.1 = ↑n}.indicator (fun _ ↦ 1)) 𝔓 := - condIndepFun_fst_prod (ν := Bandit.streamMeasure ν) - (measurable_indicator_stepsUntil_eq a m n) (by fun_prop) (by fun_prop) - (condIndepFun_reward_stepsUntil_arm' a m n hm) - -lemma reward_cond_stepsUntil [StandardBorelSpace α] [Countable α] [Nonempty α] (a : α) (m n : ℕ) - (hm : m ≠ 0) (hμn : 𝔓 ((fun ω ↦ stepsUntil a m ω.1) ⁻¹' {↑n}) ≠ 0) : - 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ stepsUntil a m ω.1 ← ↑n; 𝔓] = ν a := by - have hμna : - 𝔓 ((fun ω ↦ stepsUntil a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}) ≠ 0 := by - suffices ((fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ - stepsUntil a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}) - = (fun ω ↦ stepsUntil a m ω.1) ⁻¹' {↑n} by simpa [this] using hμn - ext ω - simp only [Set.mem_inter_iff, Set.mem_preimage, Set.mem_singleton_iff, and_iff_left_iff_imp] - exact arm_eq_of_stepsUntil_eq_coe hm - have hμa : (𝔓).map (fun ω ↦ arm n ω.1) {a} ≠ 0 := by - rw [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] - refine fun h_zero ↦ hμn (measure_mono_null (fun ω ↦ ?_) h_zero) - simp only [Set.mem_preimage, Set.mem_singleton_iff] - exact arm_eq_of_stepsUntil_eq_coe hm - calc 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ stepsUntil a m ω.1 ← (n : ℕ∞); 𝔓] - _ = (𝔓[|(fun ω ↦ stepsUntil a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}]).map - (fun ω ↦ reward n ω.1) := by - congr with ω - simp only [Set.mem_preimage, Set.mem_singleton_iff, Set.mem_inter_iff, iff_self_and] - exact arm_eq_of_stepsUntil_eq_coe hm - _ = (𝔓[|(fun ω ↦ arm n ω.1) ⁻¹' {a} - ∩ {ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) | stepsUntil a m ω.1 = ↑n}.indicator 1 ⁻¹' {1} ]).map - (fun ω ↦ reward n ω.1) := by - congr 2 with ω - simp only [Set.mem_inter_iff, Set.mem_preimage, Set.mem_singleton_iff, Set.indicator_apply, - Set.mem_setOf_eq, Pi.one_apply, ite_eq_left_iff, zero_ne_one, imp_false, Decidable.not_not] - rw [and_comm] - _ = 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1 ← a; 𝔓] := by - rw [cond_of_condIndepFun (by fun_prop)] - · exact condIndepFun_reward_stepsUntil_arm a m n hm - · refine measurable_one.indicator ?_ - exact measurableSet_eq_fun (by fun_prop) (by fun_prop) - · fun_prop - · convert hμna using 2 - rw [Set.inter_comm] - congr 1 with ω - simp [Set.indicator_apply] - _ = ν a := reward_cond_arm a n hμa - -lemma condDistrib_rewardByCount_stepsUntil [Countable α] [StandardBorelSpace α] [Nonempty α] - (a : α) (m : ℕ) (hm : m ≠ 0) : - condDistrib (rewardByCount a m) (fun ω ↦ stepsUntil a m ω.1) 𝔓 - =ᵐ[(𝔓).map (fun ω ↦ stepsUntil a m ω.1)] Kernel.const _ (ν a) := by - refine (condDistrib_ae_eq_cond (μ := 𝔓) - (X := fun ω ↦ stepsUntil a m ω.1) (by fun_prop) (by fun_prop)).trans ?_ - rw [Filter.EventuallyEq, ae_iff_of_countable] - intro n hn - simp only [Kernel.const_apply] - cases n with - | top => - rw [Measure.map_congr (g := fun ω ↦ ω.2 m a)] - swap - · refine ae_cond_of_forall_mem ((measurableSet_singleton _).preimage (by fun_prop)) ?_ - simp only [Set.mem_preimage, Set.mem_singleton_iff] - exact fun ω ↦ rewardByCount_of_stepsUntil_eq_top - rw [cond_of_indepFun _ (by fun_prop) (by fun_prop) (measurableSet_singleton _)] - · exact (hasLaw_Z a m).map_eq - · rwa [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] at hn - · exact indepFun_prod (X := fun ω : ℕ → α × ℝ ↦ stepsUntil a m ω) - (Y := fun ω : ℕ → α → ℝ ↦ ω m a) (by fun_prop) (by fun_prop) - | coe n => - rw [Measure.map_congr (g := fun ω ↦ reward n ω.1)] - swap - · refine ae_cond_of_forall_mem ((measurableSet_singleton _).preimage (by fun_prop)) ?_ - simp only [Set.mem_preimage, Set.mem_singleton_iff] - exact fun ω ↦ rewardByCount_of_stepsUntil_eq_coe - refine reward_cond_stepsUntil a m n hm ?_ - rwa [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] at hn - -/-- The reward received at the `m`-th pull of arm `a` has law `ν a`. -/ -lemma hasLaw_rewardByCount [Countable α] [StandardBorelSpace α] [Nonempty α] - (a : α) (m : ℕ) (hm : m ≠ 0) : - HasLaw (rewardByCount a m) (ν a) 𝔓 where - map_eq := by - have h_condDistrib : - condDistrib (rewardByCount a m) (fun ω ↦ stepsUntil a m ω.1) 𝔓 - =ᵐ[(𝔓).map (fun ω ↦ stepsUntil a m ω.1)] - Kernel.const _ (ν a) := condDistrib_rewardByCount_stepsUntil a m hm - calc (𝔓).map (rewardByCount a m) - _ = (condDistrib (rewardByCount a m) (fun ω ↦ stepsUntil a m ω.1) 𝔓) - ∘ₘ ((𝔓).map (fun ω ↦ stepsUntil a m ω.1)) := by - rw [condDistrib_comp_map (by fun_prop) (by fun_prop)] - _ = (Kernel.const _ (ν a)) ∘ₘ ((𝔓).map (fun ω ↦ stepsUntil a m ω.1)) := - Measure.comp_congr h_condDistrib - _ = ν a := by - have : IsProbabilityMeasure ((𝔓).map (fun ω ↦ stepsUntil a m ω.1)) := - Measure.isProbabilityMeasure_map (by fun_prop) - simp - -lemma identDistrib_rewardByCount [Countable α] [StandardBorelSpace α] [Nonempty α] (a : α) (n m : ℕ) - (hn : n ≠ 0) (hm : m ≠ 0) : - IdentDistrib (rewardByCount a n) (rewardByCount a m) 𝔓 𝔓 where - aemeasurable_fst := by fun_prop - aemeasurable_snd := by fun_prop - map_eq := by rw [(hasLaw_rewardByCount a n hn).map_eq, (hasLaw_rewardByCount a m hm).map_eq] - -lemma identDistrib_rewardByCount_id [Countable α] [StandardBorelSpace α] [Nonempty α] - (a : α) (n : ℕ) (hn : n ≠ 0) : - IdentDistrib (rewardByCount a n) id 𝔓 (ν a) where - aemeasurable_fst := by fun_prop - aemeasurable_snd := Measurable.aemeasurable <| by fun_prop - map_eq := by rw [(hasLaw_rewardByCount a n hn).map_eq, Measure.map_id] - -lemma identDistrib_rewardByCount_eval [Countable α] [StandardBorelSpace α] [Nonempty α] - (a : α) (n m : ℕ) (hn : n ≠ 0) : - IdentDistrib (rewardByCount a n) (fun ω ↦ ω m a) 𝔓 (Bandit.streamMeasure ν) := - (identDistrib_rewardByCount_id a n hn).trans (identDistrib_eval_eval_id_streamMeasure ν m a).symm - -lemma indepFun_rewardByCount_Iic (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (a : α) - (n : ℕ) : - (rewardByCount a (n + 1)) ⟂ᵢ[𝔓] fun ω (i : Iic n) ↦ rewardByCount a i ω := by - sorry - -lemma iIndepFun_rewardByCount' (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (a : α) : - iIndepFun (rewardByCount a) (Bandit.measure alg ν) := by - rw [iIndepFun_nat_iff_forall_indepFun (by fun_prop)] - exact indepFun_rewardByCount_Iic alg ν a - -lemma iIndepFun_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] : - iIndepFun (fun (p : α × ℕ) ↦ rewardByCount p.1 p.2) 𝔓 := by - sorry - -lemma identDistrib_rewardByCount_stream' [Countable α] [StandardBorelSpace α] [Nonempty α] - (a : α) : - IdentDistrib (fun ω n ↦ rewardByCount a (n + 1) ω) (fun ω n ↦ ω n a) - 𝔓 (Bandit.streamMeasure ν) := by - refine IdentDistrib.pi (fun n ↦ ?_) ?_ ?_ - · refine identDistrib_rewardByCount_eval a (n + 1) n (by simp) (ν := ν) - · have h_indep := iIndepFun_rewardByCount' alg ν a - exact iIndepFun.precomp (g := fun n ↦ n + 1) (fun i j hij ↦ by grind) h_indep - · exact iIndepFun_eval_streamMeasure'' ν a - -omit [DecidableEq α] [MeasurableSingletonClass α] in -lemma identDistrib_eval_streamMeasure_measure (a : α) : - IdentDistrib (fun ω n ↦ ω n a) (fun ω n ↦ ω.2 n a) - (Bandit.streamMeasure ν) 𝔓 := by - refine IdentDistrib.pi (fun n ↦ ?_) ?_ ?_ - · rw [← Bandit.snd_measure alg ν, Measure.snd, - identDistrib_map_left_iff (by fun_prop) (by fun_prop) - (Measurable.aemeasurable <| by fun_prop)] - exact IdentDistrib.refl (by fun_prop) - · exact iIndepFun_eval_streamMeasure'' ν a - · change iIndepFun (fun n ↦ ((fun ω ↦ ω n a) ∘ Prod.snd)) 𝔓 - rw [← iIndepFun_map_iff (by fun_prop) (fun _ ↦ Measurable.aemeasurable (by fun_prop))] - rw [← Measure.snd, Bandit.snd_measure] - exact iIndepFun_eval_streamMeasure'' ν a - -lemma identDistrib_rewardByCount_stream [Countable α] [StandardBorelSpace α] [Nonempty α] - (a : α) : - IdentDistrib (fun ω n ↦ rewardByCount a (n + 1) ω) (fun ω n ↦ ω.2 n a) 𝔓 𝔓 := - (identDistrib_rewardByCount_stream' a).trans (identDistrib_eval_streamMeasure_measure a) - -lemma indepFun_rewardByCount_of_ne {a b : α} (hab : a ≠ b) : - IndepFun (fun ω s ↦ rewardByCount a s ω) (fun ω s ↦ rewardByCount b s ω) 𝔓 := by - sorry - -lemma identDistrib_sum_Icc_rewardByCount [Nonempty α] [Countable α] (m : ℕ) (a : α) : - IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount a s ω) - (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓 𝔓 := by - have h1 (a : α) : - IdentDistrib (fun ω s ↦ rewardByCount a (s + 1) ω) (fun ω s ↦ ω.2 s a) 𝔓 𝔓 := - identDistrib_rewardByCount_stream a - have h_eq (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) : ∑ s ∈ Icc 1 m, rewardByCount a s ω - = ∑ s ∈ range m, rewardByCount a (s + 1) ω := by - let e : Icc 1 m ≃ range m := - { toFun x := ⟨x - 1, by have h := x.2; simp only [mem_Icc] at h; simp; grind⟩ - invFun x := ⟨x + 1, by - have h := x.2 - simp only [mem_Icc, le_add_iff_nonneg_left, zero_le, true_and, ge_iff_le] - simp only [mem_range] at h - grind⟩ - left_inv x := by have h := x.2; simp only [mem_Icc] at h; grind - right_inv x := by have h := x.2; grind } - rw [← sum_coe_sort (Icc 1 m), ← sum_coe_sort (range m), sum_equiv e] - · simp - · simp only [univ_eq_attach, mem_attach, forall_const, Subtype.forall, mem_Icc, - forall_and_index] - grind - simp_rw [h_eq] - exact IdentDistrib.comp (h1 a) (u := fun p ↦ ∑ s ∈ range m, p s) (by fun_prop) - -end Bandits diff --git a/LeanBandits/SequentialLearning/Algorithm.lean b/LeanBandits/SequentialLearning/Algorithm.lean index c9736eaf..df3c528b 100644 --- a/LeanBandits/SequentialLearning/Algorithm.lean +++ b/LeanBandits/SequentialLearning/Algorithm.lean @@ -5,7 +5,6 @@ Authors: Rémy Degenne, Paulo Rauber -/ import LeanBandits.ForMathlib.Measurable import LeanBandits.ForMathlib.Traj -import Mathlib.Probability.HasLaw /-! # Algorithms @@ -17,7 +16,7 @@ open scoped ENNReal NNReal namespace Learning -variable {α R : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} +variable {α R Ω : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} {mΩ : MeasurableSpace Ω} /-- A stochastic, sequential algorithm. -/ structure Algorithm (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] where @@ -56,183 +55,190 @@ lemma fst_stepKernel (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : (stepKernel alg env n).fst = alg.policy n := by rw [stepKernel, Kernel.fst_compProd] -/-- Kernel sending a partial trajectory of the bandit interaction `Iic n → α × ℝ` to a measure -on `ℕ → α × ℝ`, supported on full trajectories that start with the partial one. -/ -noncomputable def traj (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : - Kernel (Iic n → α × R) (ℕ → α × R) := - Kernel.traj (X := fun _ ↦ α × R) (stepKernel alg env) n -deriving IsMarkovKernel - -/-- Measure on the sequence of actions and observations generated by the algorithm/environment. -/ -noncomputable -def trajMeasure (alg : Algorithm α R) (env : Environment α R) : - Measure (ℕ → α × R) := - Kernel.trajMeasure (alg.p0 ⊗ₘ env.ν0) (stepKernel alg env) -deriving IsProbabilityMeasure +section IsAlgEnvSeq -/-- Action and reward at step `n`. -/ -def step (n : ℕ) (h : ℕ → α × R) : α × R := h n +variable {A : ℕ → Ω → α} {R' : ℕ → Ω → R} {alg : Algorithm α R} {env : Environment α R} + {P : Measure Ω} [IsFiniteMeasure P] -/-- `action n` is the action pulled at time `n`. This is a random variable on the measurable space -`ℕ → α × ℝ`. -/ -def action (n : ℕ) (h : ℕ → α × R) : α := (h n).1 - -/-- `reward n` is the reward at time `n`. This is a random variable on the measurable space -`ℕ → α × R`. -/ -def reward (n : ℕ) (h : ℕ → α × R) : R := (h n).2 - -/-- `hist n` is the history up to time `n`. This is a random variable on the measurable space -`ℕ → α × R`. -/ -def hist (n : ℕ) (h : ℕ → α × R) : Iic n → α × R := fun i ↦ h i - -lemma fst_comp_step (n : ℕ) : Prod.fst ∘ step (α := α) (R := R) n = action n := rfl +/-- Step of the algorithm-environment sequence: the action-reward pair at time `n`. -/ +def IsAlgEnvSeq.step (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (n : ℕ) (ω : Ω) : α × R := + (A n ω, R' n ω) @[fun_prop] -lemma measurable_step (n : ℕ) : Measurable (step n (α := α) (R := R)) := by - unfold step; fun_prop - -@[fun_prop] -lemma measurable_step_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ step p.1 p.2) := by - refine measurable_from_prod_countable_right fun n ↦ ?_ - simp only +lemma IsAlgEnvSeq.measurable_step (n : ℕ) (hA : Measurable (A n)) + (hR' : Measurable (R' n)) : + Measurable (IsAlgEnvSeq.step A R' n) := by + unfold IsAlgEnvSeq.step fun_prop -@[fun_prop] -lemma measurable_action (n : ℕ) : Measurable (action n (α := α) (R := R)) := by - unfold action; fun_prop - -@[fun_prop] -lemma measurable_action_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ action p.1 p.2) := by - refine measurable_from_prod_countable_right fun n ↦ ?_ - simp only - fun_prop - -@[fun_prop] -lemma measurable_reward (n : ℕ) : Measurable (reward n (α := α) (R := R)) := by - unfold reward; fun_prop +/-- History of the algorithm-environment sequence up to time `n`. -/ +def IsAlgEnvSeq.hist (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (n : ℕ) (ω : Ω) : Iic n → α × R := + fun i ↦ (A i ω, R' i ω) @[fun_prop] -lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ reward p.1 p.2) := by - refine measurable_from_prod_countable_right fun n ↦ ?_ - simp only +lemma IsAlgEnvSeq.measurable_hist (hA : ∀ n, Measurable (A n)) + (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : + Measurable (IsAlgEnvSeq.hist A R' n) := by + unfold IsAlgEnvSeq.hist fun_prop -@[fun_prop] -lemma measurable_hist (n : ℕ) : Measurable (hist n (α := α) (R := R)) := by unfold hist; fun_prop - -lemma hist_eq_frestrictLe : - hist = Preorder.frestrictLe («π» := fun _ ↦ α × R) := by - ext n h i : 3 - simp [hist, Preorder.frestrictLe] - -/-- Filtration of the algorithm interaction. -/ -protected def filtration (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : - Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) := - MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R) - -lemma step_eq_eval_comp_hist (n : ℕ) : - step (α := α) (R := R) n = (fun x ↦ x ⟨n, by simp⟩) ∘ (hist n) := rfl - -lemma action_eq_eval_comp_hist (n : ℕ) : - action (α := α) (R := R) n = (fun x ↦ (x ⟨n, by simp⟩).1) ∘ (hist n) := rfl - -lemma reward_eq_eval_comp_hist (n : ℕ) : - reward (α := α) (R := R) n = (fun x ↦ (x ⟨n, by simp⟩).2) ∘ (hist n) := rfl - -lemma measurable_step_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (step n) := by - simp only [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe] - rw [step_eq_eval_comp_hist] +lemma IsAlgEnvSeq.eval_comp_hist (n : ℕ) : + (fun x ↦ x ⟨n, by simp⟩) ∘ (hist A R' n) = step A R' n := rfl + +lemma IsAlgEnvSeq.fst_eval_comp_hist (n : ℕ) : + (fun x ↦ (x ⟨n, by simp⟩).1) ∘ (hist A R' n) = A n := rfl + +lemma IsAlgEnvSeq.snd_eval_comp_hist (n : ℕ) : + (fun x ↦ (x ⟨n, by simp⟩).2) ∘ (hist A R' n) = R' n := rfl + +/-- An algorithm-environment sequence: a sequence of actions and rewards generated +by an algorithm interacting with an environment. -/ +structure IsAlgEnvSeq + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (alg : Algorithm α R) (env : Environment α R) + (P : Measure Ω) [IsFiniteMeasure P] : Prop where + measurable_A n : Measurable (A n) := by fun_prop + measurable_R n : Measurable (R' n) := by fun_prop + hasLaw_action_zero : HasLaw (fun ω ↦ (A 0 ω)) alg.p0 P + hasCondDistrib_reward_zero : HasCondDistrib (R' 0) (A 0) env.ν0 P + hasCondDistrib_action n : + HasCondDistrib (A (n + 1)) (IsAlgEnvSeq.hist A R' n) (alg.policy n) P + hasCondDistrib_reward n : + HasCondDistrib (R' (n + 1)) (fun ω ↦ (IsAlgEnvSeq.hist A R' n ω, A (n + 1) ω)) + (env.feedback n) P + +lemma IsAlgEnvSeq.hasLaw_step_zero + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + (h : IsAlgEnvSeq A R' alg env P) : + HasLaw (step A R' 0) (alg.p0 ⊗ₘ env.ν0) P := + HasLaw.prod_of_hasCondDistrib h.hasLaw_action_zero h.hasCondDistrib_reward_zero + +lemma IsAlgEnvSeq.hasCondDistrib_step + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + (h : IsAlgEnvSeq A R' alg env P) (n : ℕ) : + HasCondDistrib (step A R' (n + 1)) (hist A R' n) (stepKernel alg env n) P := + HasCondDistrib.prod (h.hasCondDistrib_action n) (h.hasCondDistrib_reward n) + +/-- Filtration generated by the history up to time `n`. -/ +def IsAlgEnvSeq.filtration (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) : + Filtration ℕ mΩ where + seq i := MeasurableSpace.comap (hist A R' i) inferInstance + mono' i j hij := by + simp only + rw [← measurable_iff_comap_le] + have : hist A R' i = (fun h k ↦ h ⟨k.1, by grind⟩) ∘ hist A R' j := rfl + rw [this] + exact measurable_comp_comap _ (by fun_prop) + le' i := by + rw [← measurable_iff_comap_le] + exact measurable_hist hA hR' i + +lemma IsAlgEnvSeq.measurable_action_filtration + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (n : ℕ) : + Measurable[IsAlgEnvSeq.filtration hA hR' n] (A n) := by + have : A n = (fun h ↦ (h ⟨n, by simp⟩).1) ∘ (hist A R' n) := by + ext ω : 1 + simp [IsAlgEnvSeq.hist] + rw [this] exact measurable_comp_comap _ (by fun_prop) -lemma adapted_step [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] - [SecondCountableTopology α] [OpensMeasurableSpace α] - [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] - [SecondCountableTopology R] [OpensMeasurableSpace R] : - Adapted (Learning.filtration α R) (step (α := α) (R := R)) := - fun n ↦ (measurable_step_filtration n).stronglyMeasurable - -lemma measurable_hist_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (hist n) := by - simp [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe, - measurable_iff_comap_le] - -lemma adapted_hist [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] - [SecondCountableTopology α] [OpensMeasurableSpace α] - [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] - [SecondCountableTopology R] [OpensMeasurableSpace R] : - Adapted (Learning.filtration α R) hist := - fun n ↦ (measurable_hist_filtration n).stronglyMeasurable - -lemma measurable_action_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (action n) := by - simp only [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe] - rw [action_eq_eval_comp_hist] - exact measurable_comp_comap _ (by fun_prop) +/-- Filtration generated by the history at time `n-1` together with the action at time `n`. -/ +def IsAlgEnvSeq.filtrationAction + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) : + Filtration ℕ mΩ where + seq n := if n = 0 then MeasurableSpace.comap (A 0) inferInstance + else IsAlgEnvSeq.filtration hA hR' (n - 1) ⊔ MeasurableSpace.comap (A n) inferInstance + mono' n m hnm := by + simp only + by_cases hn : n = 0 + · by_cases hm : m = 0 + · simp [hn, hm] + · simp only [hn, ↓reduceIte, hm] + refine le_sup_of_le_left ?_ + rw [← measurable_iff_comap_le] + suffices Measurable[IsAlgEnvSeq.filtration hA hR' 0] (A 0) from + this.mono ((IsAlgEnvSeq.filtration hA hR').mono zero_le') le_rfl + exact measurable_action_filtration hA hR' 0 + have hm : m ≠ 0 := by grind + simp only [hn, hm, ↓reduceIte] + have hnm' : n - 1 ≤ m - 1 := by grind + simp only [sup_le_iff] + constructor + · refine le_sup_of_le_left ?_ + exact (IsAlgEnvSeq.filtration hA hR').mono hnm' + · rcases eq_or_lt_of_le hnm with rfl | hlt + · exact le_sup_of_le_right le_rfl + refine le_sup_of_le_left ?_ + rw [← measurable_iff_comap_le] + have h_le : n ≤ m - 1 := by grind + suffices Measurable[IsAlgEnvSeq.filtration hA hR' n] (A n) from + this.mono ((IsAlgEnvSeq.filtration hA hR').mono h_le) le_rfl + exact measurable_action_filtration hA hR' n + le' n := by + by_cases hn : n = 0 + · simp only [hn, ↓reduceIte] + rw [← measurable_iff_comap_le] + fun_prop + simp only [hn, ↓reduceIte, sup_le_iff] + constructor + · exact (IsAlgEnvSeq.filtration hA hR').le _ + · rw [← measurable_iff_comap_le] + fun_prop + +lemma IsAlgEnvSeq.filtrationAction_zero_eq_comap + {hA : ∀ n, Measurable (A n)} {hR' : ∀ n, Measurable (R' n)} : + filtrationAction hA hR' 0 = MeasurableSpace.comap (A 0) inferInstance := by + simp [filtrationAction] + +lemma IsAlgEnvSeq.filtrationAction_eq_comap + {hA : ∀ n, Measurable (A n)} {hR' : ∀ n, Measurable (R' n)} (n : ℕ) (hn : n ≠ 0) : + filtrationAction hA hR' n = + MeasurableSpace.comap (fun ω ↦ (hist A R' (n - 1) ω, A n ω)) inferInstance := by + simp only [filtrationAction, filtration, ← MeasurableSpace.comap_prodMk, hn, ↓reduceIte] + rfl -lemma adapted_action [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] - [SecondCountableTopology α] [OpensMeasurableSpace α] : - Adapted (Learning.filtration α R) action := - fun n ↦ (measurable_action_filtration n).stronglyMeasurable +end IsAlgEnvSeq -lemma measurable_reward_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (reward n) := by - simp only [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe] - rw [reward_eq_eval_comp_hist] - exact measurable_comp_comap _ (by fun_prop) +/-- Kernel sending a partial trajectory of the bandit Seq `Iic n → α × ℝ` to a measure +on `ℕ → α × ℝ`, supported on full trajectories that start with the partial one. -/ +noncomputable def traj (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : + Kernel (Iic n → α × R) (ℕ → α × R) := + Kernel.traj (X := fun _ ↦ α × R) (stepKernel alg env) n +deriving IsMarkovKernel -lemma adapted_reward [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] - [SecondCountableTopology R] [OpensMeasurableSpace R] : - Adapted (Learning.filtration α R) reward := - fun n ↦ (measurable_reward_filtration n).stronglyMeasurable - -lemma condDistrib_step [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : - condDistrib (step (n + 1)) (hist n) (trajMeasure alg env) - =ᵐ[(trajMeasure alg env).map (hist n)] stepKernel alg env n := - Kernel.condDistrib_trajMeasure - -lemma condDistrib_action [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : - condDistrib (action (n + 1)) (hist n) (trajMeasure alg env) - =ᵐ[(trajMeasure alg env).map (hist n)] alg.policy n := by - rw [← fst_comp_step] - refine (condDistrib_comp _ (by fun_prop) (by fun_prop)).trans ?_ - filter_upwards [condDistrib_step alg env n] with h h_eq - rw [Kernel.map_apply _ (by fun_prop), h_eq, ← Kernel.map_apply _ (by fun_prop), ← Kernel.fst_eq, - fst_stepKernel] - -lemma condDistrib_reward [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : - condDistrib (reward (n + 1)) (fun ω ↦ (hist n ω, action (n + 1) ω)) (trajMeasure alg env) - =ᵐ[(trajMeasure alg env).map (fun ω ↦ (hist n ω, action (n + 1) ω))] env.feedback n := by - have h_step := condDistrib_step alg env n - have h_action := condDistrib_action alg env n - rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h_step h_action ⊢ - rw [h_action, ← Measure.compProd_assoc, ← stepKernel, ← h_step, - Measure.map_map (by fun_prop) (by fun_prop)] - rfl +/-- Measure on the sequence of actions and observations generated by the algorithm/environment. -/ +noncomputable +def trajMeasure (alg : Algorithm α R) (env : Environment α R) : + Measure (ℕ → α × R) := + Kernel.trajMeasure (alg.p0 ⊗ₘ env.ν0) (stepKernel alg env) +deriving IsProbabilityMeasure -lemma hasLaw_step_zero (alg : Algorithm α R) (env : Environment α R) : - HasLaw (step 0) (alg.p0 ⊗ₘ env.ν0) (trajMeasure alg env) where - aemeasurable := Measurable.aemeasurable (by fun_prop) - map_eq := by - unfold step - rw [← coe_default_Iic_zero] - simp only [trajMeasure, Kernel.trajMeasure] - rw [← Measure.deterministic_comp_eq_map (by fun_prop), Measure.comp_assoc, - Kernel.deterministic_comp_eq_map, Kernel.traj_zero_map_eval_zero, - Measure.deterministic_comp_eq_map, Measure.map_map (by fun_prop) (by fun_prop)] - exact Measure.map_id - -lemma hasLaw_action_zero (alg : Algorithm α R) (env : Environment α R) : - HasLaw (action 0) alg.p0 (trajMeasure alg env) where - map_eq := by - rw [← fst_comp_step, ← Measure.map_map (by fun_prop) (by fun_prop), - (hasLaw_step_zero alg env).map_eq, ← Measure.fst, Measure.fst_compProd] - -lemma condDistrib_reward_zero [StandardBorelSpace R] [Nonempty R] - (alg : Algorithm α R) (env : Environment α R) : - condDistrib (reward 0) (action 0) (trajMeasure alg env) - =ᵐ[(trajMeasure alg env).map (action 0)] env.ν0 := by - have h_step := (hasLaw_step_zero alg env).map_eq - have h_action := (hasLaw_action_zero alg env).map_eq - rwa [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop), h_action] +section ModelEquivalence + +variable {Ω Ω' : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + {alg : Algorithm α R} {env : Environment α R} + {P : Measure Ω} [IsProbabilityMeasure P] {P' : Measure Ω'} [IsProbabilityMeasure P'] + {A₁ : ℕ → Ω → α} {R₁ : ℕ → Ω → R} {A₂ : ℕ → Ω' → α} {R₂ : ℕ → Ω' → R} + +theorem eq_trajMeasure_of_isAlgEnvSeq (h : IsAlgEnvSeq A₁ R₁ alg env P) : + P.map (fun ω n ↦ (A₁ n ω, R₁ n ω)) = trajMeasure alg env := by + rw [trajMeasure] + have h := Kernel.eq_trajMeasure (Y := fun n ω ↦ (A₁ n ω, R₁ n ω)) (P := P) + (μ₀ := alg.p0 ⊗ₘ env.ν0) (κ := stepKernel alg env) (fun n ↦ ?_) ?_ (fun n ↦ ?_) + · exact h + · have hA := h.measurable_A n + have hR := h.measurable_R n + fun_prop + · simp only + exact h.hasLaw_step_zero + · exact h.hasCondDistrib_step n + +theorem isAlgEnvSeq_unique (h1 : IsAlgEnvSeq A₁ R₁ alg env P) + (h2 : IsAlgEnvSeq A₂ R₂ alg env P') : + P.map (fun ω n ↦ (A₁ n ω, R₁ n ω)) = P'.map (fun ω n ↦ (A₂ n ω, R₂ n ω)) := by + rw [eq_trajMeasure_of_isAlgEnvSeq h1, eq_trajMeasure_of_isAlgEnvSeq h2] + +end ModelEquivalence end Learning diff --git a/LeanBandits/SequentialLearning/Deterministic.lean b/LeanBandits/SequentialLearning/Deterministic.lean index f2240745..5ab42a05 100644 --- a/LeanBandits/SequentialLearning/Deterministic.lean +++ b/LeanBandits/SequentialLearning/Deterministic.lean @@ -3,7 +3,7 @@ Copyright (c) 2025 Rémy Degenne. All rights reserved. Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne -/ -import LeanBandits.SequentialLearning.Algorithm +import LeanBandits.SequentialLearning.IonescuTulceaSpace /-! # Deterministic algorithms @@ -17,38 +17,81 @@ namespace Learning variable {α R : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} -/-- A deterministic algorithm. -/ +/-- A deterministic algorithm, which chooses the action given by the function `nextAction`. -/ @[simps] noncomputable -def detAlgorithm (nextaction : (n : ℕ) → (Iic n → α × R) → α) - (h_next : ∀ n, Measurable (nextaction n)) (action0 : α) : +def detAlgorithm (nextAction : (n : ℕ) → (Iic n → α × R) → α) + (h_next : ∀ n, Measurable (nextAction n)) (action0 : α) : Algorithm α R where - policy n := Kernel.deterministic (nextaction n) (h_next n) + policy n := Kernel.deterministic (nextAction n) (h_next n) p0 := Measure.dirac action0 -variable {nextaction : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextaction n)} +variable {nextAction : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextAction n)} {action0 : α} {env : Environment α R} -local notation "𝔓" => trajMeasure (detAlgorithm nextaction h_next action0) env +namespace IsAlgEnvSeq -lemma HasLaw_action_zero_detAlgorithm : HasLaw (action 0) (Measure.dirac action0) 𝔓 where - map_eq := (hasLaw_action_zero _ _).map_eq +variable {Ω : Type*} {mΩ : MeasurableSpace Ω} + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] + {P : Measure Ω} [IsProbabilityMeasure P] {A : ℕ → Ω → α} {R' : ℕ → Ω → R} -lemma action_zero_detAlgorithm [MeasurableSingletonClass α] : action 0 =ᵐ[𝔓] fun _ ↦ action0 := by - have h_eq : ∀ᵐ x ∂((𝔓).map (action 0)), x = action0 := by - rw [(hasLaw_action_zero _ _).map_eq] +lemma HasLaw_action_zero_detAlgorithm + (h : IsAlgEnvSeq A R' (detAlgorithm nextAction h_next action0) env P) : + HasLaw (A 0) (Measure.dirac action0) P where + aemeasurable := have hA := h.measurable_A; by fun_prop + map_eq := (hasLaw_action_zero h).map_eq + +lemma action_zero_detAlgorithm + (h : IsAlgEnvSeq A R' (detAlgorithm nextAction h_next action0) env P) : + A 0 =ᵐ[P] fun _ ↦ action0 := by + have h_eq : ∀ᵐ x ∂(P.map (A 0)), x = action0 := by + rw [(hasLaw_action_zero h).map_eq] + simp [detAlgorithm] + have hA := h.measurable_A + exact ae_of_ae_map (by fun_prop) h_eq + +lemma action_detAlgorithm_ae_eq + (h : IsAlgEnvSeq A R' (detAlgorithm nextAction h_next action0) env P) (n : ℕ) : + A (n + 1) =ᵐ[P] fun ω ↦ nextAction n (hist A R' n ω) := by + have hA := h.measurable_A + have hR' := h.measurable_R + exact ae_eq_of_condDistrib_eq_deterministic (by fun_prop) (by fun_prop) (by fun_prop) + (h.hasCondDistrib_action n).condDistrib_eq + +lemma action_detAlgorithm_ae_all_eq + (h : IsAlgEnvSeq A R' (detAlgorithm nextAction h_next action0) env P) : + ∀ᵐ ω ∂P, A 0 ω = action0 ∧ ∀ n, A (n + 1) ω = nextAction n (hist A R' n ω) := by + rw [eventually_and, ae_all_iff] + exact ⟨action_zero_detAlgorithm h, action_detAlgorithm_ae_eq h⟩ + +end IsAlgEnvSeq + +namespace IT + +local notation "𝔓" => trajMeasure (detAlgorithm nextAction h_next action0) env + +lemma HasLaw_action_zero_detAlgorithm : HasLaw (IT.action 0) (Measure.dirac action0) 𝔓 where + map_eq := (IT.hasLaw_action_zero _ _).map_eq + +lemma action_zero_detAlgorithm [MeasurableSingletonClass α] : + IT.action 0 =ᵐ[𝔓] fun _ ↦ action0 := by + have h_eq : ∀ᵐ x ∂((𝔓).map (IT.action 0)), x = action0 := by + rw [(IT.hasLaw_action_zero _ _).map_eq] simp [detAlgorithm] exact ae_of_ae_map (by fun_prop) h_eq lemma action_detAlgorithm_ae_eq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] - [Nonempty R] (n : ℕ) : action (n + 1) =ᵐ[𝔓] fun h ↦ nextaction n (hist n h) := + [Nonempty R] (n : ℕ) : IT.action (n + 1) =ᵐ[𝔓] fun h ↦ nextAction n (IT.hist n h) := ae_eq_of_condDistrib_eq_deterministic (by fun_prop) (by fun_prop) (by fun_prop) - (condDistrib_action (detAlgorithm nextaction h_next action0) env n) + (IT.condDistrib_action (detAlgorithm nextAction h_next action0) env n) lemma action_detAlgorithm_ae_all_eq [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] : - ∀ᵐ h ∂𝔓, action 0 h = action0 ∧ ∀ n, action (n + 1) h = nextaction n (hist n h) := by + ∀ᵐ h ∂𝔓, IT.action 0 h = action0 ∧ ∀ n, IT.action (n + 1) h = nextAction n (IT.hist n h) := by rw [eventually_and, ae_all_iff] exact ⟨action_zero_detAlgorithm, action_detAlgorithm_ae_eq⟩ +end IT + end Learning diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index e4f6d6b5..527db1f6 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -20,19 +20,23 @@ be seen as a stochastic process indexed by time `t` on the measurable space `ℕ -/ -open MeasureTheory Finset +open MeasureTheory Finset Learning namespace Learning -variable {α R : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} [DecidableEq α] - {a : α} {m n t : ℕ} {h : ℕ → α × R} +variable {α R Ω : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} {mΩ : MeasurableSpace Ω} + [DecidableEq α] + {alg : Algorithm α R} {env : Environment α R} + {P : Measure Ω} [IsProbabilityMeasure P] + {A : ℕ → Ω → α} {R' : ℕ → Ω → R} + {a : α} {m n t : ℕ} {ω : Ω} section PullCount /-- Number of times action `a` was chosen up to time `t` (excluding `t`). -/ noncomputable -def pullCount (a : α) (t : ℕ) (h : ℕ → α × R) : ℕ := - #(filter (fun s ↦ action s h = a) (range t)) +def pullCount (A : ℕ → Ω → α) (a : α) (t : ℕ) (ω : Ω) : ℕ := + #(filter (fun s ↦ A s ω = a) (range t)) /-- Number of pulls of arm `a` up to (and including) time `n`. This is the number of entries in `h` in which the arm is `a`. -/ @@ -40,68 +44,72 @@ noncomputable def pullCount' (n : ℕ) (h : Iic n → α × R) (a : α) := #{s | (h s).1 = a} @[simp] -lemma pullCount_zero (a : α) : pullCount a 0 (R := R) = 0 := by ext; simp [pullCount] +lemma pullCount_zero (a : α) : pullCount A a 0 = 0 := by ext; simp [pullCount] -lemma pullCount_zero_apply (a : α) (h : ℕ → α × R) : pullCount a 0 h = 0 := by simp +lemma pullCount_zero_apply (a : α) (ω : Ω) : pullCount A a 0 ω = 0 := by simp -lemma pullCount_one : pullCount a 1 h = if action 0 h = a then 1 else 0 := by +lemma pullCount_one : pullCount A a 1 ω = if A 0 ω = a then 1 else 0 := by simp only [pullCount, range_one] split_ifs with h · rw [card_eq_one] refine ⟨0, by simp [h]⟩ · simp [h] -lemma monotone_pullCount (a : α) (h : ℕ → α × R) : Monotone (pullCount a · h) := +lemma monotone_pullCount (a : α) (ω : Ω) : Monotone (pullCount A a · ω) := fun _ _ _ ↦ card_le_card (filter_subset_filter _ (by simpa)) @[mono, gcongr] -lemma pullCount_mono (a : α) {n m : ℕ} (hnm : n ≤ m) (h : ℕ → α × R) : - pullCount a n h ≤ pullCount a m h := - monotone_pullCount a h hnm +lemma pullCount_mono (a : α) {n m : ℕ} (hnm : n ≤ m) (ω : Ω) : + pullCount A a n ω ≤ pullCount A a m ω := + monotone_pullCount a ω hnm -lemma pullCount_action_eq_pullCount_add_one (t : ℕ) (h : ℕ → α × R) : - pullCount (action t h) (t + 1) h = pullCount (action t h) t h + 1 := by +lemma pullCount_action_eq_pullCount_add_one (t : ℕ) (ω : Ω) : + pullCount A (A t ω) (t + 1) ω = pullCount A (A t ω) t ω + 1 := by simp [pullCount, range_add_one, filter_insert] -lemma pullCount_eq_pullCount_of_action_ne (ha : action t h ≠ a) : - pullCount a (t + 1) h = pullCount a t h := by +lemma pullCount_eq_pullCount_of_action_ne (ha : A t ω ≠ a) : + pullCount A a (t + 1) ω = pullCount A a t ω := by simp [pullCount, range_add_one, filter_insert, ha] lemma pullCount_add_one : - pullCount a (t + 1) h = pullCount a t h + if action t h = a then 1 else 0 := by + pullCount A a (t + 1) ω = pullCount A a t ω + if A t ω = a then 1 else 0 := by split_ifs with h · rw [← h, pullCount_action_eq_pullCount_add_one] · rw [pullCount_eq_pullCount_of_action_ne h, add_zero] -lemma pullCount_eq_sum (a : α) (t : ℕ) (h : ℕ → α × R) : - pullCount a t h = ∑ s ∈ range t, if action s h = a then 1 else 0 := by simp [pullCount] +lemma pullCount_eq_sum (a : α) (t : ℕ) (ω : Ω) : + pullCount A a t ω = ∑ s ∈ range t, if A s ω = a then 1 else 0 := by simp [pullCount] lemma pullCount'_eq_sum (n : ℕ) (h : Iic n → α × R) (a : α) : pullCount' n h a = ∑ s : Iic n, if (h s).1 = a then 1 else 0 := by simp [pullCount'] -lemma pullCount_add_one_eq_pullCount' {n : ℕ} {h : ℕ → α × R} : - pullCount a (n + 1) h = pullCount' n (fun i ↦ h i) a := by +lemma pullCount_add_one_eq_pullCount' {n : ℕ} {ω : Ω} : + pullCount A a (n + 1) ω = pullCount' n (fun i ↦ (A i ω, R' i ω)) a := by rw [pullCount_eq_sum, pullCount'_eq_sum] - unfold action - rw [Finset.sum_coe_sort (f := fun s ↦ if (h s).1 = a then 1 else 0) (Iic n)] + rw [Finset.sum_coe_sort (f := fun s ↦ if A s ω = a then 1 else 0) (Iic n)] congr with m simp only [mem_range, mem_Iic] grind -lemma pullCount_eq_pullCount' {n : ℕ} {h : ℕ → α × R} (hn : n ≠ 0) : - pullCount a n h = pullCount' (n - 1) (fun i ↦ h i) a := by +lemma pullCount_eq_pullCount' {n : ℕ} {ω : Ω} (hn : n ≠ 0) : + pullCount A a n ω = pullCount' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by cases n with | zero => exact absurd rfl hn | succ n => - rw [pullCount_add_one_eq_pullCount'] + rw [pullCount_add_one_eq_pullCount' (R' := R')] have : n + 1 - 1 = n := by simp exact this ▸ rfl -lemma pullCount_le (a : α) (t : ℕ) (h : ℕ → α × R) : pullCount a t h ≤ t := +lemma pullCount'_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 + rw [← pullCount_add_one_eq_pullCount', ← pullCount_add_one_eq_pullCount'] + exact pullCount_mono a (by lia) _ + +lemma pullCount_le (a : α) (t : ℕ) (ω : Ω) : pullCount A a t ω ≤ t := (card_filter_le _ _).trans_eq (by simp) -lemma pullCount_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h = action i h') : - pullCount a (n + 1) h = pullCount a (n + 1) h' := by +lemma pullCount_congr {ω' : Ω} (h_eq : ∀ i ≤ n, A i ω = A i ω') : + pullCount A a (n + 1) ω = pullCount A a (n + 1) ω' := by unfold pullCount congr 1 with s simp only [mem_filter, mem_range, and_congr_right_iff] @@ -109,8 +117,8 @@ lemma pullCount_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h = a rw [Nat.lt_add_one_iff] at hs rw [h_eq s hs] -lemma pullCount_lt_of_forall_ne (h_lt : ∀ s, pullCount a (s + 1) h ≠ t) (ht : t ≠ 0) : - pullCount a n h < t := by +lemma pullCount_lt_of_forall_ne (h_lt : ∀ s, pullCount A a (s + 1) ω ≠ t) (ht : t ≠ 0) : + pullCount A a n ω < t := by induction n with | zero => simpa using ht.bot_lt | succ n hn => @@ -118,24 +126,70 @@ lemma pullCount_lt_of_forall_ne (h_lt : ∀ s, pullCount a (s + 1) h ≠ t) (ht rw [pullCount_add_one] at h_lt ⊢ grind -lemma exists_pullCount_eq_of_le (hnm : t ≤ pullCount a (n + 1) h) (ht : t ≠ 0) : - ∃ s, pullCount a (s + 1) h = t := by +lemma exists_pullCount_eq_of_le (hnm : t ≤ pullCount A a (n + 1) ω) (ht : t ≠ 0) : + ∃ s, pullCount A a (s + 1) ω = t := by by_contra! h_contra - refine lt_irrefl (pullCount a (n + 1) h) ?_ + refine lt_irrefl (pullCount A a (n + 1) ω) ?_ refine lt_of_lt_of_le ?_ hnm exact pullCount_lt_of_forall_ne h_contra ht +lemma pullCount_le_add (a : α) (n C : ℕ) (ω : Ω) : + pullCount A a n ω ≤ C + 1 + + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s := by + rw [pullCount_eq_sum] + calc ∑ s ∈ range n, if A s ω = a then 1 else 0 + _ ≤ ∑ s ∈ range n, ({s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s + + {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s) := by + gcongr with s hs + simp [Set.indicator_apply] + grind + _ = ∑ s ∈ range n, {s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s + + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s := by + rw [Finset.sum_add_distrib] + _ ≤ C + 1 + ∑ s ∈ range n, {s | A s ω = a ∧ C < pullCount A a s ω}.indicator 1 s := by + gcongr + have h_le n : ∑ s ∈ range n, {s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s ≤ + pullCount A a n ω := by + rw [pullCount_eq_sum] + gcongr with s hs + simp only [Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] + grind + induction n with + | zero => simp + | succ n hn => + rw [Finset.sum_range_succ] + rcases le_or_gt (pullCount A a n ω) C with h_pc | h_pc + · have hn' : ∑ s ∈ range n, {s | A s ω = a ∧ pullCount A a s ω ≤ C}.indicator 1 s ≤ C := + (h_le n).trans h_pc + grw [hn'] + gcongr + simp only [Set.indicator_apply, Set.mem_setOf_eq, Pi.one_apply] + grind + · refine le_trans ?_ hn + simp [h_pc] + section Measurability @[fun_prop] -lemma measurable_pullCount [MeasurableSingletonClass α] (a : α) (t : ℕ) : - Measurable (fun h : ℕ → α × R ↦ pullCount a t h) := by +lemma measurable_pullCount [MeasurableSingletonClass α] (hA : ∀ n, Measurable (A n)) + (a : α) (t : ℕ) : + Measurable (fun ω : Ω ↦ pullCount A a t ω) := by simp_rw [pullCount_eq_sum] - have h_meas s : Measurable (fun h : ℕ → α × R ↦ if action s h = a then 1 else 0) := by + have h_meas s : Measurable (fun ω : Ω ↦ if A s ω = a then 1 else 0) := by refine Measurable.ite ?_ (by fun_prop) (by fun_prop) exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop +@[fun_prop] +lemma measurable_uncurry_pullCount [MeasurableEq α] + (hA : ∀ n, Measurable (A n)) (t : ℕ) : + Measurable (fun p : Ω × α ↦ pullCount A p.2 t p.1) := by + simp_rw [pullCount_eq_sum] + have h_meas s : Measurable (fun h : Ω × α ↦ if A s h.1 = h.2 then 1 else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact measurableSet_eq_fun (by fun_prop) (by fun_prop) + fun_prop + @[fun_prop] lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : Measurable (fun h : Iic n → α × R ↦ pullCount' n h a) := by @@ -145,23 +199,45 @@ lemma measurable_pullCount' [MeasurableSingletonClass α] (n : ℕ) (a : α) : exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop -lemma adapted_pullCount_add_one [MeasurableSingletonClass α] (a : α) : - Adapted (Learning.filtration α R) (fun n ↦ pullCount a (n + 1)) := by - refine fun n ↦ Measurable.stronglyMeasurable ?_ - simp only - have : pullCount a (n + 1) = (fun h : Iic n → α × R ↦ pullCount' n h a) ∘ (hist n) := by +lemma measurable_uncurry_pullCount' [MeasurableEq α] (n : ℕ) : + Measurable (fun p : (Iic n → α × R) × α ↦ pullCount' n p.1 p.2) := by + simp_rw [pullCount'_eq_sum] + have h_meas s : Measurable (fun h : (Iic n → α × R) × α ↦ if (h.1 s).1 = h.2 then 1 else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact measurableSet_eq_fun (by fun_prop) (by fun_prop) + fun_prop + +lemma adapted_pullCount_add_one' [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (n : ℕ) : + Measurable[IsAlgEnvSeq.filtration hA hR' n] (pullCount A a (n + 1)) := by + have : pullCount A a (n + 1) = (fun h : Iic n → α × R ↦ pullCount' n h a) ∘ + (IsAlgEnvSeq.hist A R' n) := by ext exact pullCount_add_one_eq_pullCount' - rw [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe, this] - exact measurable_comp_comap (hist n) (measurable_pullCount' n a) + rw [IsAlgEnvSeq.filtration, this] + exact measurable_comp_comap _ (measurable_pullCount' n a) + +lemma adapted_pullCount_add_one [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) : + Adapted (IsAlgEnvSeq.filtration hA hR') (fun n ↦ pullCount A a (n + 1)) := + fun n ↦ Measurable.stronglyMeasurable <| adapted_pullCount_add_one' hA hR' a n -lemma isPredictable_pullCount [MeasurableSingletonClass α] (a : α) : - IsPredictable (Learning.filtration α R) (pullCount a) := by +lemma isPredictable_pullCount [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) : + IsPredictable (IsAlgEnvSeq.filtration hA hR') (pullCount A a) := by rw [isPredictable_iff_measurable_add_one] - refine ⟨?_, fun n ↦ (adapted_pullCount_add_one a n).measurable⟩ + refine ⟨?_, fun n ↦ (adapted_pullCount_add_one hA hR' a n).measurable⟩ simp only [pullCount_zero] fun_prop +lemma integrable_pullCount [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (a : α) (n : ℕ) : + Integrable (fun ω ↦ (pullCount A a n ω : ℝ)) P := by + refine integrable_of_le_of_le (g₁ := 0) (g₂ := fun _ ↦ n) (by fun_prop) + (ae_of_all _ fun ω ↦ by simp) (ae_of_all _ fun ω ↦ ?_) (integrable_const _) (integrable_const _) + simp only [Nat.cast_le] + exact pullCount_le a n ω + end Measurability end PullCount @@ -171,21 +247,22 @@ section StepsUntil -- TODO: replace this by leastGE, once leastGE is generalized /-- Number of steps until action `a` was pulled exactly `m` times. -/ noncomputable -def stepsUntil (a : α) (m : ℕ) (h : ℕ → α × R) : ℕ∞ := sInf ((↑) '' {s | pullCount a (s + 1) h = m}) +def stepsUntil (A : ℕ → Ω → α) (a : α) (m : ℕ) (ω : Ω) : ℕ∞ := + sInf ((↑) '' {s | pullCount A a (s + 1) ω = m}) -lemma stepsUntil_eq_top_iff : stepsUntil a m h = ⊤ ↔ ∀ s, pullCount a (s + 1) h ≠ m := by +lemma stepsUntil_eq_top_iff : stepsUntil A a m ω = ⊤ ↔ ∀ s, pullCount A a (s + 1) ω ≠ m := by simp [stepsUntil, sInf_eq_top] -lemma stepsUntil_ne_top (h_exists : ∃ s, pullCount a (s + 1) h = m) : stepsUntil a m h ≠ ⊤ := by +lemma stepsUntil_ne_top (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : stepsUntil A a m ω ≠ ⊤ := by simpa [stepsUntil_eq_top_iff] -lemma exists_pullCount_eq (h' : stepsUntil a m h ≠ ⊤) : - ∃ s, pullCount a (s + 1) h = m := by +lemma exists_pullCount_eq (h' : stepsUntil A a m ω ≠ ⊤) : + ∃ s, pullCount A a (s + 1) ω = m := by by_contra! h_contra rw [← stepsUntil_eq_top_iff] at h_contra simp [h_contra] at h' -lemma stepsUntil_zero_of_ne (hka : action 0 h ≠ a) : stepsUntil a 0 h = 0 := by +lemma stepsUntil_zero_of_ne (hka : A 0 ω ≠ a) : stepsUntil A a 0 ω = 0 := by unfold stepsUntil simp_rw [← bot_eq_zero, sInf_eq_bot, bot_eq_zero] intro n hn @@ -194,19 +271,19 @@ lemma stepsUntil_zero_of_ne (hka : action 0 h ≠ a) : stepsUntil a 0 h = 0 := b rw [← zero_add 1, pullCount_eq_pullCount_of_action_ne hka] simp -lemma stepsUntil_zero_of_eq (hka : action 0 h = a) : stepsUntil a 0 h = ⊤ := by +lemma stepsUntil_zero_of_eq (hka : A 0 ω = a) : stepsUntil A a 0 ω = ⊤ := by rw [stepsUntil_eq_top_iff] - suffices 0 < pullCount a 1 h by + suffices 0 < pullCount A a 1 ω by intro n hn refine lt_irrefl 0 ?_ exact this.trans_le (le_trans (monotone_pullCount _ _ (by omega)) hn.le) rw [← hka, ← zero_add 1, pullCount_action_eq_pullCount_add_one] simp -lemma stepsUntil_eq_dite (a : α) (m : ℕ) (h : ℕ → α × R) - [Decidable (∃ s, pullCount a (s + 1) h = m)] : - stepsUntil a m h = - if h : ∃ s, pullCount a (s + 1) h = m then (Nat.find h : ℕ∞) else ⊤ := by +lemma stepsUntil_eq_dite (a : α) (m : ℕ) (ω : Ω) + [Decidable (∃ s, pullCount A a (s + 1) ω = m)] : + stepsUntil A a m ω = + if h : ∃ s, pullCount A a (s + 1) ω = m then (Nat.find h : ℕ∞) else ⊤ := by unfold stepsUntil split_ifs with h' · refine le_antisymm ?_ ?_ @@ -216,22 +293,22 @@ lemma stepsUntil_eq_dite (a : α) (m : ℕ) (h : ℕ → α × R) forall_apply_eq_imp_iff₂, Nat.cast_le, Nat.find_le_iff] exact fun n hn ↦ ⟨n, le_rfl, hn⟩ · push_neg at h' - suffices {s | pullCount a (s + 1) h = m} = ∅ by simp [this] + suffices {s | pullCount A a (s + 1) ω = m} = ∅ by simp [this] ext s simpa using (h' s) -- todo: this is in ℝ because of the limited def of leastGE lemma stepsUntil_eq_leastGE (a : α) (hm : m ≠ 0) : - stepsUntil a m = leastGE (fun n (h : ℕ → α × ℝ) ↦ pullCount a (n + 1) h) m := by + stepsUntil A a m = leastGE (fun n (ω : Ω) ↦ pullCount A a (n + 1) ω) m := by classical - ext h + ext ω rw [stepsUntil_eq_dite] unfold leastGE hittingAfter simp only [zero_le, Set.mem_Ici, Nat.cast_le, true_and, ENat.some_eq_coe] - have h_iff : (∃ s, pullCount a (s + 1) h = m) ↔ (∃ s, m ≤ pullCount a (s + 1) h) := by + have h_iff : (∃ s, pullCount A a (s + 1) ω = m) ↔ (∃ s, m ≤ pullCount A a (s + 1) ω) := by refine ⟨fun ⟨s, hs⟩ ↦ ⟨s, hs.ge⟩, fun ⟨s, hs⟩ ↦ ?_⟩ exact exists_pullCount_eq_of_le hs hm - by_cases h_exists : ∃ s, m ≤ pullCount a (s + 1) h + by_cases h_exists : ∃ s, m ≤ pullCount A a (s + 1) ω swap; · simp_rw [h_iff]; simp [h_exists] rw [if_pos h_exists, dif_pos] swap; · rwa [h_iff] @@ -240,45 +317,56 @@ lemma stepsUntil_eq_leastGE (a : α) (hm : m ≠ 0) : constructor · apply le_antisymm · by_contra! h_contra - obtain ⟨s, hs⟩ : ∃ s, pullCount a (s + 1) h = m := exists_pullCount_eq_of_le h_contra.le hm + obtain ⟨s, hs⟩ : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq_of_le h_contra.le hm rw [← hs] at h_contra refine h_contra.not_ge ?_ gcongr exact csInf_le (by simp) (by simp) - · exact Nat.sInf_mem (s := {j | m ≤ pullCount a (j + 1) h}) h_exists + · exact Nat.sInf_mem (s := {j | m ≤ pullCount A a (j + 1) ω}) h_exists · intro n hn h_contra refine hn.not_ge ?_ exact csInf_le (by simp) (by simp [h_contra]) -lemma stepsUntil_pullCount_le (h : ℕ → α × R) (a : α) (t : ℕ) : - stepsUntil a (pullCount a (t + 1) h) h ≤ t := by +lemma stepsUntil_mono (a : α) (ω : Ω) {n m : ℕ} (hn : n ≠ 0) (hnm : n ≤ m) : + stepsUntil A a n ω ≤ stepsUntil A a m ω := by + rw [stepsUntil_eq_leastGE a hn, stepsUntil_eq_leastGE a (by lia)] + simp_rw [leastGE] + have h_Ici_subset : Set.Ici (m : ℝ) ⊆ Set.Ici (n : ℝ) := by + intro x hx + simp only [Set.mem_Ici] at hx ⊢ + refine le_trans ?_ hx + exact mod_cast hnm + exact hittingAfter_anti (fun n ω ↦ (pullCount A a (n + 1) ω : ℝ)) 0 h_Ici_subset ω + +lemma stepsUntil_pullCount_le (ω : Ω) (a : α) (t : ℕ) : + stepsUntil A a (pullCount A a (t + 1) ω) ω ≤ t := by rw [stepsUntil] exact csInf_le (OrderBot.bddBelow _) ⟨t, rfl, rfl⟩ -lemma stepsUntil_pullCount_eq (h : ℕ → α × R) (t : ℕ) : - stepsUntil (action t h) (pullCount (action t h) (t + 1) h) h = t := by - apply le_antisymm (stepsUntil_pullCount_le h (action t h) t) - suffices ∀ t', pullCount (action t h) (t' + 1) h = pullCount (action t h) t h + 1 → t ≤ t' by +lemma stepsUntil_pullCount_eq (ω : Ω) (t : ℕ) : + stepsUntil A (A t ω) (pullCount A (A t ω) (t + 1) ω) ω = t := by + apply le_antisymm (stepsUntil_pullCount_le ω (A t ω) t) + suffices ∀ t', pullCount A (A t ω) (t' + 1) ω = pullCount A (A t ω) t ω + 1 → t ≤ t' by simpa [stepsUntil, pullCount_action_eq_pullCount_add_one] - exact fun t' h' ↦ Nat.le_of_lt_succ ((monotone_pullCount (action t h) h).reflect_lt + exact fun t' h' ↦ Nat.le_of_lt_succ ((monotone_pullCount (A t ω) ω).reflect_lt (h' ▸ lt_add_one _)) /-- If we pull action `a` at time 0, the first time at which it is pulled once is 0. -/ -lemma stepsUntil_one_of_eq (hka : action 0 h = a) : stepsUntil a 1 h = 0 := by +lemma stepsUntil_one_of_eq (hka : A 0 ω = a) : stepsUntil A a 1 ω = 0 := by classical - have h_pull : pullCount a 1 h = 1 := by simp [pullCount_one, hka] - have h_le := stepsUntil_pullCount_le h a 0 + have h_pull : pullCount A a 1 ω = 1 := by simp [pullCount_one, hka] + have h_le := stepsUntil_pullCount_le (A := A) ω a 0 simpa [h_pull] using h_le lemma stepsUntil_eq_zero_iff : - stepsUntil a m h = 0 ↔ (m = 0 ∧ action 0 h ≠ a) ∨ (m = 1 ∧ action 0 h = a) := by + stepsUntil A a m ω = 0 ↔ (m = 0 ∧ A 0 ω ≠ a) ∨ (m = 1 ∧ A 0 ω = a) := by classical refine ⟨fun h' ↦ ?_, fun h' ↦ ?_⟩ - · have h_exists : ∃ s, pullCount a (s + 1) h = m := exists_pullCount_eq (by simp [h']) + · have h_exists : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq (by simp [h']) simp only [stepsUntil_eq_dite, h_exists, ↓reduceDIte, Nat.cast_eq_zero, Nat.find_eq_zero, zero_add] at h' rw [pullCount_one] at h' - by_cases hka : action 0 h = a + by_cases hka : A 0 ω = a · simp only [hka, ↓reduceIte] at h' simp [h'.symm, hka] · simp only [hka, ↓reduceIte] at h' @@ -290,8 +378,8 @@ lemma stepsUntil_eq_zero_iff : rw [h.1] exact stepsUntil_one_of_eq h.2 -lemma action_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) h = m) : - action (stepsUntil a m h).toNat h = a := by +lemma action_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : + A (stepsUntil A a m ω).toNat ω = a := by classical simp only [stepsUntil_eq_dite, h_exists, ↓reduceDIte, ENat.toNat_coe] have h_spec := Nat.find_spec h_exists @@ -310,17 +398,17 @@ lemma action_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) h rwa [← pullCount_eq_pullCount_of_action_ne] exact h_ne -lemma action_eq_of_stepsUntil_eq_coe {ω : ℕ → α × R} (hm : m ≠ 0) - (h : stepsUntil a m ω = n) : - action n ω = a := by - have : n = (stepsUntil a m ω).toNat := by simp [h] - rw [this, action_stepsUntil hm] - exact exists_pullCount_eq (by simp [h]) +lemma action_eq_of_stepsUntil_eq_coe (hm : m ≠ 0) (h : stepsUntil A a m ω = n) : + A n ω = a := by + have : n = (stepsUntil A a m ω).toNat := by simp [h] + rw [this] + have h_exists : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq (by simp [h]) + exact action_stepsUntil hm h_exists -lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount a (s + 1) h = m) : - pullCount a (stepsUntil a m h + 1).toNat h = m := by +lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : + pullCount A a (stepsUntil A a m ω + 1).toNat ω = m := by classical - have h_eq := stepsUntil_eq_dite a m h + have h_eq := stepsUntil_eq_dite (A := A) a m ω simp only [h_exists, ↓reduceDIte] at h_eq have h' := Nat.find_spec h_exists rw [h_eq] @@ -328,10 +416,10 @@ lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount a (s + 1) h = m) simp only [ENat.toNat_coe, ENat.toNat_one] exact h' -lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) h = m) : - pullCount a (stepsUntil a m h).toNat h = m - 1 := by - have h_action := action_eq_of_stepsUntil_eq_coe (n := (stepsUntil a m h).toNat) (a := a) (ω := h) - hm ?_ +lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount A a (s + 1) ω = m) : + pullCount A a (stepsUntil A a m ω).toNat ω = m - 1 := by + have h_action := action_eq_of_stepsUntil_eq_coe (A := A) (n := (stepsUntil A a m ω).toNat) + (a := a) (ω := ω) hm ?_ swap; · symm; simpa [stepsUntil_eq_top_iff] have h_add_one := pullCount_stepsUntil_add_one h_exists nth_rw 1 [← h_action] at h_add_one @@ -340,46 +428,46 @@ lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) swap; · simpa [stepsUntil_eq_top_iff] grind -lemma pullCount_lt_of_le_stepsUntil (a : α) {n m : ℕ} (h : ℕ → α × R) - (h_exists : ∃ s, pullCount a (s + 1) h = m) (hn : n < stepsUntil a m h) : - pullCount a (n + 1) h < m := by +lemma pullCount_lt_of_le_stepsUntil (a : α) {n m : ℕ} (ω : Ω) + (h_exists : ∃ s, pullCount A a (s + 1) ω = m) (hn : n < stepsUntil A a m ω) : + pullCount A a (n + 1) ω < m := by classical - have h_eq := stepsUntil_eq_dite a m h + have h_eq := stepsUntil_eq_dite (A := A) a m ω simp only [h_exists, ↓reduceDIte] at h_eq rw [← ENat.coe_toNat (stepsUntil_ne_top h_exists)] at hn refine lt_of_le_of_ne ?_ ?_ - · calc pullCount a (n + 1) h - _ ≤ pullCount a (stepsUntil a m h + 1).toNat h := by - refine monotone_pullCount a h ?_ + · calc pullCount A a (n + 1) ω + _ ≤ pullCount A a (stepsUntil A a m ω + 1).toNat ω := by + refine monotone_pullCount a ω ?_ rw [ENat.toNat_add (stepsUntil_ne_top h_exists) (by simp)] simp only [ENat.toNat_one, add_le_add_iff_right] exact mod_cast hn.le _ = m := pullCount_stepsUntil_add_one h_exists · refine Nat.find_min h_exists (m := n) ?_ - suffices n < (stepsUntil a m h).toNat by + suffices n < (stepsUntil A a m ω).toNat by rwa [h_eq, ENat.toNat_coe] at this exact mod_cast hn -lemma pullCount_eq_of_stepsUntil_eq_coe {ω : ℕ → α × R} (hm : m ≠ 0) - (h : stepsUntil a m ω = n) : - pullCount a n ω = m - 1 := by - have : n = (stepsUntil a m ω).toNat := by simp [h] +lemma pullCount_eq_of_stepsUntil_eq_coe {ω : Ω} (hm : m ≠ 0) + (h : stepsUntil A a m ω = n) : + pullCount A a n ω = m - 1 := by + have : n = (stepsUntil A a m ω).toNat := by simp [h] rw [this, pullCount_stepsUntil hm] exact exists_pullCount_eq (by simp [h]) -lemma pullCount_add_one_eq_of_stepsUntil_eq_coe {ω : ℕ → α × R} - (h : stepsUntil a m ω = n) : - pullCount a (n + 1) ω = m := by - have : n + 1 = (stepsUntil a m ω + 1).toNat := by +lemma pullCount_add_one_eq_of_stepsUntil_eq_coe {ω : Ω} + (h : stepsUntil A a m ω = n) : + pullCount A a (n + 1) ω = m := by + have : n + 1 = (stepsUntil A a m ω + 1).toNat := by rw [ENat.toNat_add (by simp [h]) (by simp)]; simp [h] rw [this, pullCount_stepsUntil_add_one] exact exists_pullCount_eq (by simp [h]) -lemma stepsUntil_eq_iff {ω : ℕ → α × R} (n : ℕ) : - stepsUntil a m ω = n ↔ - pullCount a (n + 1) ω = m ∧ (∀ k < n, pullCount a (k + 1) ω < m) := by +lemma stepsUntil_eq_iff {ω : Ω} (n : ℕ) : + stepsUntil A a m ω = n ↔ + pullCount A a (n + 1) ω = m ∧ (∀ k < n, pullCount A a (k + 1) ω < m) := by refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ - · have h_exists : ∃ s, pullCount a (s + 1) ω = m := exists_pullCount_eq (by simp [h]) + · have h_exists : ∃ s, pullCount A a (s + 1) ω = m := exists_pullCount_eq (by simp [h]) refine ⟨pullCount_add_one_eq_of_stepsUntil_eq_coe h, fun k hk ↦ ?_⟩ exact pullCount_lt_of_le_stepsUntil a ω h_exists (by rw [h]; exact mod_cast hk) · classical @@ -388,8 +476,28 @@ lemma stepsUntil_eq_iff {ω : ℕ → α × R} (n : ℕ) : rw [Nat.find_eq_iff] exact ⟨h.1, fun k hk ↦ (h.2 k hk).ne⟩ -lemma stepsUntil_eq_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h = action i h') : - stepsUntil a m h = n ↔ stepsUntil a m h' = n := by +lemma stepsUntil_eq_iff' {ω : Ω} (hm : m ≠ 0) (n : ℕ) : + stepsUntil A a m ω = n ↔ A n ω = a ∧ pullCount A a n ω = m - 1 := by + by_cases hn : n = 0 + · simp [hn, stepsUntil_eq_zero_iff, hm] + grind + rw [stepsUntil_eq_iff n] + refine ⟨fun ⟨h1, h2⟩ ↦ ⟨?_, ?_⟩, fun ⟨h1, h2⟩ ↦ ⟨?_, fun k hk ↦ ?_⟩⟩ + · rw [pullCount_add_one] at h1 + specialize h2 (n - 1) (by lia) + grind + · rw [pullCount_add_one] at h1 + specialize h2 (n - 1) (by lia) + grind + · rw [pullCount_add_one, h1, h2] + grind + · rw [Nat.lt_iff_le_pred (by grind)] + rw [← h2] + refine monotone_pullCount a ω ?_ + grind + +lemma stepsUntil_eq_congr {ω' : Ω} (h_eq : ∀ i ≤ n, A i ω = A i ω') : + stepsUntil A a m ω = n ↔ stepsUntil A a m ω' = n := by simp_rw [stepsUntil_eq_iff n] congr! 1 · rw [pullCount_congr h_eq] @@ -397,34 +505,42 @@ lemma stepsUntil_eq_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h rw [pullCount_congr] grind -lemma isStoppingTime_stepsUntil [MeasurableSingletonClass α] (a : α) (hm : m ≠ 0) : - IsStoppingTime (Learning.filtration α ℝ) (stepsUntil a m) := by +section Measurability + +lemma isStoppingTime_stepsUntil [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (hm : m ≠ 0) : + IsStoppingTime (IsAlgEnvSeq.filtration hA hR') (stepsUntil A a m) := by rw [stepsUntil_eq_leastGE _ hm] refine Adapted.isStoppingTime_leastGE _ fun n ↦ ?_ - suffices StronglyMeasurable[Learning.filtration α ℝ n] (pullCount a (n + 1)) by fun_prop - exact adapted_pullCount_add_one a n + suffices StronglyMeasurable[IsAlgEnvSeq.filtration hA hR' n] (pullCount A a (n + 1)) by + fun_prop + exact adapted_pullCount_add_one hA hR' a n -- todo: get this from the stopping time property? @[fun_prop] -lemma measurable_stepsUntil [MeasurableSingletonClass α] (a : α) (m : ℕ) : - Measurable (fun h : ℕ → α × R ↦ stepsUntil a m h) := by +lemma measurable_stepsUntil [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (a : α) (m : ℕ) : + Measurable (stepsUntil A a m) := by classical - have h_union : {h' : ℕ → α × R | ∃ s, pullCount a (s + 1) h' = m} - = ⋃ s : ℕ, {h' | pullCount a (s + 1) h' = m} := by ext; simp - have h_meas_set : MeasurableSet {h' : ℕ → α × R | ∃ s, pullCount a (s + 1) h' = m} := by + have h_union : {h' : Ω | ∃ s, pullCount A a (s + 1) h' = m} + = ⋃ s : ℕ, {h' | pullCount A a (s + 1) h' = m} := by ext; simp + have h_meas_set : MeasurableSet {h' : Ω | ∃ s, pullCount A a (s + 1) h' = m} := by rw [h_union] - exact MeasurableSet.iUnion fun s ↦ (measurableSet_singleton _).preimage (by fun_prop) - simp_rw [stepsUntil_eq_dite] - suffices Measurable fun k ↦ if h : k ∈ {k' | ∃ s, pullCount a (s + 1) k' = m} - then (Nat.find h : ℕ∞) else ⊤ by convert this - refine Measurable.dite (s := {k' : ℕ → α × R | ∃ s, pullCount a (s + 1) k' = m}) + refine MeasurableSet.iUnion fun s ↦ (measurableSet_singleton _).preimage ?_ + exact measurable_pullCount hA a (s + 1) + suffices Measurable fun k ↦ if h : k ∈ {k' | ∃ s, pullCount A a (s + 1) k' = m} + then (Nat.find h : ℕ∞) else ⊤ by + convert this with ω + rw [stepsUntil_eq_dite a m ω] + rfl + refine Measurable.dite (s := {k' : Ω | ∃ s, pullCount A a (s + 1) k' = m}) (f := fun x ↦ (Nat.find x.2 : ℕ∞)) (g := fun _ ↦ ⊤) ?_ (by fun_prop) h_meas_set refine Measurable.coe_nat_enat ?_ refine measurable_find _ fun k ↦ ?_ - suffices MeasurableSet {x : ℕ → α × R | pullCount a (k + 1) x = m} by - have : Subtype.val '' {x : {k' : ℕ → α × R | - ∃ s, pullCount a (s + 1) k' = m} | pullCount a (k + 1) (x : ℕ → α × R) = m} - = {x : ℕ → α × R | pullCount a (k + 1) x = m} := by + suffices MeasurableSet {x : Ω | pullCount A a (k + 1) x = m} by + have : Subtype.val '' {x : {k' : Ω | + ∃ s, pullCount A a (s + 1) k' = m} | pullCount A a (k + 1) (x : Ω) = m} + = {x : Ω | pullCount A a (k + 1) x = m} := by ext x simp only [Set.mem_setOf_eq, Set.coe_setOf, Set.mem_image, Subtype.exists, exists_and_left, exists_prop, exists_eq_right_right, and_iff_left_iff_imp] @@ -434,9 +550,106 @@ lemma measurable_stepsUntil [MeasurableSingletonClass α] (a : α) (m : ℕ) : exact (measurableSet_singleton _).preimage (by fun_prop) exact (measurableSet_singleton _).preimage (by fun_prop) -lemma measurable_stepsUntil' [MeasurableSingletonClass α] (a : α) (m : ℕ) : - Measurable (fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ stepsUntil a m ω.1) := - (measurable_stepsUntil a m).comp measurable_fst +lemma measurable_stepsUntil' [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (a : α) (m : ℕ) : + Measurable (fun ω : Ω × (ℕ → α → R) ↦ stepsUntil A a m ω.1) := + (measurable_stepsUntil hA a m).comp measurable_fst + +lemma measurable_comap_indicator_stepsUntil_eq [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : + Measurable[MeasurableSpace.comap + (fun ω : Ω ↦ (IsAlgEnvSeq.hist A R' (n-1) ω, A n ω)) inferInstance] + ({ω | stepsUntil A a m ω = ↑n}.indicator fun _ ↦ 1) := by + by_cases hm : m = 0 + · simp only [hm] + by_cases hn : n = 0 + · simp only [hn, CharP.cast_eq_zero, stepsUntil_eq_zero_iff, ne_eq, true_and, zero_ne_one, + false_and, or_false] + refine Measurable.indicator measurable_const ?_ + refine (measurableSet_singleton _).compl.preimage ?_ + rw [measurable_iff_comap_le] + rw [Prod.instMeasurableSpace, MeasurableSpace.comap_prodMk] + exact le_sup_of_le_right le_rfl + · have : {ω | stepsUntil A a 0 ω = n} = ∅ := by + ext ω + by_cases ha : A 0 ω = a + · simp [stepsUntil_zero_of_eq ha] + · simp only [Set.mem_setOf_eq, stepsUntil_zero_of_ne ha, Set.mem_empty_iff_false, + iff_false] + norm_cast + exact Ne.symm hn + simp [this] + simp_rw [stepsUntil_eq_iff' hm] + refine Measurable.indicator measurable_const ?_ + refine ((measurableSet_singleton _).preimage ?_).inter ((measurableSet_singleton _).preimage ?_) + · rw [measurable_iff_comap_le] + rw [Prod.instMeasurableSpace, MeasurableSpace.comap_prodMk] + exact le_sup_of_le_right le_rfl + · rw [measurable_iff_comap_le] + rw [Prod.instMeasurableSpace, MeasurableSpace.comap_prodMk] + refine le_sup_of_le_left ?_ + rw [← measurable_iff_comap_le] + by_cases hn : n = 0 + · simp only [hn, pullCount_zero] + exact measurable_const + have h_meas := adapted_pullCount_add_one' hA hR' a (n - 1) + rwa [Nat.sub_add_cancel (by lia)] at h_meas + +lemma measurable_indicator_stepsUntil_eq [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : + Measurable ({ω : Ω | stepsUntil A a m ω = ↑n}.indicator fun _ ↦ 1) := by + refine (measurable_comap_indicator_stepsUntil_eq hA hR' a m n).mono ?_ le_rfl + refine Measurable.comap_le ?_ + fun_prop + +lemma measurableSet_stepsUntil_eq_zero [MeasurableSingletonClass α] (a : α) (m : ℕ) : + MeasurableSet[MeasurableSpace.comap (A 0) inferInstance] + {ω : Ω | stepsUntil A a m ω = 0} := by + simp only [stepsUntil_eq_zero_iff (a := a) (m := m), ne_eq] + by_cases hm : m = 0 + · simp only [hm, true_and, zero_ne_one, false_and, or_false] + refine (measurableSet_singleton _).compl.preimage ?_ + rw [measurable_iff_comap_le] + by_cases hm1 : m = 1 + swap; · simp [hm, hm1] + simp only [hm1, one_ne_zero, false_and, true_and, false_or] + refine (measurableSet_singleton _).preimage ?_ + rw [measurable_iff_comap_le] + +lemma measurable_comap_indicator_stepsUntil_eq_zero [MeasurableSingletonClass α] (a : α) (m : ℕ) : + Measurable[MeasurableSpace.comap (A 0) inferInstance] + ({ω | stepsUntil A a m ω = 0}.indicator fun _ ↦ 1) := by + rw [measurable_indicator_const_iff] + exact measurableSet_stepsUntil_eq_zero a m + +lemma measurableSet_stepsUntil_eq [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m n : ℕ) : + MeasurableSet[MeasurableSpace.comap (fun ω : Ω ↦ (IsAlgEnvSeq.hist A R' (n-1) ω, A n ω)) + inferInstance] + {ω : Ω | stepsUntil A a m ω = ↑n} := by + let mProd := MeasurableSpace.comap + (fun ω : Ω ↦ (IsAlgEnvSeq.hist A R' (n-1) ω, A n ω)) inferInstance + suffices Measurable[mProd] ({ω | stepsUntil A a m ω = ↑n}.indicator fun x ↦ 1) by + rwa [measurable_indicator_const_iff] at this + exact measurable_comap_indicator_stepsUntil_eq hA hR' a m n + +/-- `stepsUntil a m` is a stopping time with respect to the filtration `filtrationAction`. -/ +theorem isStoppingTime_stepsUntil_filtrationAction [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m : ℕ) : + IsStoppingTime (IsAlgEnvSeq.filtrationAction hA hR') (stepsUntil A a m) := by + refine isStoppingTime_of_measurableSet_eq fun n ↦ ?_ + by_cases hn : n = 0 + · simp only [hn, IsAlgEnvSeq.filtrationAction_zero_eq_comap, WithTop.coe_zero] + exact measurableSet_stepsUntil_eq_zero a m + · 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 @@ -446,64 +659,98 @@ section RewardByCount If it is never pulled `m` times, the reward is given by the second component of `ω`, which in applications will be indepedent with same law. -/ noncomputable -def rewardByCount (a : α) (m : ℕ) (ω : (ℕ → α × R) × (ℕ → α → R)) : R := - match (stepsUntil a m ω.1) with +def rewardByCount (A : ℕ → Ω → α) (R' : ℕ → Ω → R) (a : α) (m : ℕ) (ω : Ω × (ℕ → α → R)) : R := + match (stepsUntil A a m ω.1) with | ⊤ => ω.2 m a - | (n : ℕ) => reward n ω.1 + | (n : ℕ) => R' n ω.1 -lemma rewardByCount_eq_ite (a : α) (m : ℕ) (ω : (ℕ → α × R) × (ℕ → α → R)) : - rewardByCount a m ω = - if (stepsUntil a m ω.1) = ⊤ then ω.2 m a else reward (stepsUntil a m ω.1).toNat ω.1 := by +variable {ω : Ω × (ℕ → α → R)} + +lemma rewardByCount_eq_ite (a : α) (m : ℕ) (ω : Ω × (ℕ → α → R)) : + rewardByCount A R' a m ω = + if (stepsUntil A a m ω.1) = ⊤ then ω.2 m a else R' (stepsUntil A a m ω.1).toNat ω.1 := by unfold rewardByCount - cases stepsUntil a m ω.1 <;> simp + cases stepsUntil A a m ω.1 <;> simp + +lemma rewardByCount_eq_add [AddMonoid R] (a : α) (m : ℕ) : + rewardByCount A R' a m = + {ω : Ω × (ℕ → α → R) | stepsUntil A a m ω.1 ≠ ⊤}.indicator + (fun ω ↦ R' (stepsUntil A a m ω.1).toNat ω.1) + + {ω | stepsUntil A a m ω.1 = ⊤}.indicator (fun ω ↦ ω.2 m a) := by + ext ω + simp only [rewardByCount_eq_ite, ne_eq, Pi.add_apply, Set.indicator_apply, Set.mem_setOf_eq, + ite_not] + grind -lemma rewardByCount_of_stepsUntil_eq_top {ω : (ℕ → α × R) × (ℕ → α → R)} - (h : stepsUntil a m ω.1 = ⊤) : - rewardByCount a m ω = ω.2 m a := by simp [rewardByCount_eq_ite, h] +lemma rewardByCount_of_stepsUntil_eq_top (h : stepsUntil A a m ω.1 = ⊤) : + rewardByCount A R' a m ω = ω.2 m a := by simp [rewardByCount_eq_ite, h] -lemma rewardByCount_of_stepsUntil_eq_coe {ω : (ℕ → α × R) × (ℕ → α → R)} - (h : stepsUntil a m ω.1 = n) : - rewardByCount a m ω = reward n ω.1 := by simp [rewardByCount_eq_ite, h] +lemma rewardByCount_of_stepsUntil_ne_top (h : stepsUntil A a m ω.1 ≠ ⊤) : + rewardByCount A R' a m ω = R' (stepsUntil A a m ω.1).toNat ω.1 := by + simp [rewardByCount_eq_ite, h] -lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (ω : (ℕ → α × R) × (ℕ → α → R)) : - rewardByCount (action t ω.1) (pullCount (action t ω.1) t ω.1 + 1) ω = reward t ω.1 := by +lemma rewardByCount_eq_stoppedValue (h : stepsUntil A a m ω.1 ≠ ⊤) : + rewardByCount A R' a m ω = stoppedValue R' (stepsUntil A a m) ω.1 := by + rw [rewardByCount_of_stepsUntil_ne_top h, stoppedValue] + lift stepsUntil A a m ω.1 to ℕ using h with n + simp + +lemma rewardByCount_of_stepsUntil_eq_coe (h : stepsUntil A a m ω.1 = n) : + rewardByCount A R' a m ω = R' n ω.1 := by simp [rewardByCount_eq_ite, h] + +/-- The value at 0 does not matter (it would be the "zeroth" reward). +It should be considered a junk value. -/ +@[simp] +lemma rewardByCount_zero (a : α) (ω : Ω × (ℕ → α → R)) : + rewardByCount A R' a 0 ω = if A 0 ω.1 = a then ω.2 0 a else R' 0 ω.1 := by + rw [rewardByCount_eq_ite] + by_cases ha : A 0 ω.1 = a + · simp [ha, stepsUntil_zero_of_eq] + · simp [stepsUntil_zero_of_ne, ha] + +lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (ω : Ω × (ℕ → α → R)) : + rewardByCount A R' (A t ω.1) (pullCount A (A t ω.1) t ω.1 + 1) ω = R' t ω.1 := by rw [rewardByCount, ← pullCount_action_eq_pullCount_add_one, stepsUntil_pullCount_eq] @[fun_prop] -lemma measurable_rewardByCount [MeasurableSingletonClass α] (a : α) (m : ℕ) : - Measurable (fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ rewardByCount a m ω) := by +lemma measurable_rewardByCount [MeasurableSingletonClass α] + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (m : ℕ) : + Measurable (fun ω : Ω × (ℕ → α → R) ↦ rewardByCount A R' a m ω) := by simp_rw [rewardByCount_eq_ite] refine Measurable.ite ?_ ?_ ?_ - · exact (measurableSet_singleton _).preimage <| measurable_stepsUntil' a m + · exact (measurableSet_singleton _).preimage <| measurable_stepsUntil' hA a m · fun_prop - · change Measurable ((fun p : ℕ × (ℕ → α × R) ↦ reward p.1 p.2) - ∘ (fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ ((stepsUntil a m ω.1).toNat, ω.1))) - have : Measurable fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ ((stepsUntil a m ω.1).toNat, ω.1) := - (measurable_stepsUntil' a m).toNat.prodMk (by fun_prop) - exact Measurable.comp (by fun_prop) this + · change Measurable ((fun p : ℕ × Ω ↦ R' p.1 p.2) + ∘ (fun ω : Ω × (ℕ → α → R) ↦ ((stepsUntil A a m ω.1).toNat, ω.1))) + have : Measurable fun ω : Ω × (ℕ → α → R) ↦ ((stepsUntil A a m ω.1).toNat, ω.1) := + (measurable_stepsUntil' hA a m).toNat.prodMk (by fun_prop) + refine Measurable.comp ?_ this + refine measurable_from_prod_countable_right fun n ↦ ?_ + simp only + fun_prop end RewardByCount -lemma sum_pullCount_mul [Fintype α] [Semiring R] (h : ℕ → α × R) (f : α → R) (t : ℕ) : - ∑ a, pullCount a t h * f a = ∑ s ∈ range t, f (action s h) := by +lemma sum_pullCount_mul [Fintype α] [Semiring R] (ω : Ω) (f : α → R) (t : ℕ) : + ∑ a, pullCount A a t ω * f a = ∑ s ∈ range t, f (A s ω) := by unfold pullCount classical simp_rw [card_eq_sum_ones] push_cast simp_rw [sum_mul, one_mul] - exact sum_fiberwise' (range t) (action · h) f + exact sum_fiberwise' (range t) (A · ω) f -- todo: only in ℝ for now -lemma sum_pullCount [Fintype α] {h : ℕ → α × ℝ} : ∑ a, pullCount a t h = t := by - suffices ∑ a, pullCount a t h * (1 : ℝ) = t by norm_cast at this; simpa +lemma sum_pullCount [Fintype α] {ω : Ω} : ∑ a, pullCount A a t ω = t := by + suffices ∑ a, pullCount A a t ω * (1 : ℝ) = t by norm_cast at this; simpa rw [sum_pullCount_mul] simp section SumRewards /-- Sum of rewards obtained when pulling action `a` up to time `t` (exclusive). -/ -def sumRewards (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := - ∑ s ∈ range t, if action s h = a then reward s h else 0 +def sumRewards (A : ℕ → Ω → α) (R' : ℕ → Ω → ℝ) (a : α) (t : ℕ) (ω : Ω) : ℝ := + ∑ s ∈ range t, if A s ω = a then R' s ω else 0 /-- Sum of rewards of arm `a` up to (and including) time `n`. -/ noncomputable @@ -512,22 +759,24 @@ def sumRewards' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := /-- Empirical mean reward obtained when pulling action `a` up to time `t` (exclusive). -/ noncomputable -def empMean (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := sumRewards a t h / pullCount a t h +def empMean (A : ℕ → Ω → α) (R' : ℕ → Ω → ℝ) (a : α) (t : ℕ) (ω : Ω) : ℝ := + sumRewards A R' a t ω / pullCount A a t ω /-- Empirical mean of arm `a` at time `n`. -/ noncomputable def empMean' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := (sumRewards' n h a) / (pullCount' n h a) -lemma sumRewards_eq_pullCount_mul_empMean {h : ℕ → α × ℝ} (h_pull : pullCount a t h ≠ 0) : - sumRewards a t h = pullCount a t h * empMean a t h := by unfold empMean; field_simp +lemma sumRewards_eq_pullCount_mul_empMean {R' : ℕ → Ω → ℝ} {ω : Ω} + (h_pull : pullCount A a t ω ≠ 0) : + sumRewards A R' a t ω = pullCount A a t ω * empMean A R' a t ω := by unfold empMean; field_simp -lemma sum_rewardByCount_eq_sumRewards (a : α) (t : ℕ) (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) : - ∑ m ∈ Icc 1 (pullCount a t ω.1), rewardByCount a m ω = sumRewards a t ω.1 := by +lemma sum_rewardByCount_eq_sumRewards {R' : ℕ → Ω → ℝ} (a : α) (t : ℕ) (ω : Ω × (ℕ → α → ℝ)) : + ∑ m ∈ Icc 1 (pullCount A a t ω.1), rewardByCount A R' a m ω = sumRewards A R' a t ω.1 := by induction t with | zero => simp [pullCount, sumRewards] | succ t ht => - by_cases hta : action t ω.1 = a + by_cases hta : A t ω.1 = a · rw [← hta] at ht ⊢ rw [pullCount_action_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht] unfold sumRewards @@ -535,16 +784,16 @@ lemma sum_rewardByCount_eq_sumRewards (a : α) (t : ℕ) (ω : (ℕ → α × · unfold sumRewards rwa [pullCount_eq_pullCount_of_action_ne hta, sum_range_succ, if_neg hta, add_zero] -lemma sumRewards_add_one_eq_sumRewards' {n : ℕ} {h : ℕ → α × ℝ} : - sumRewards a (n + 1) h = sumRewards' n (fun i ↦ h i) a := by - unfold sumRewards sumRewards' action Learning.reward - rw [Finset.sum_coe_sort (f := fun s ↦ if (h s).1 = a then (h s).2 else 0) (Iic n)] +lemma sumRewards_add_one_eq_sumRewards' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} : + sumRewards A R' a (n + 1) ω = sumRewards' n (fun i ↦ (A i ω, R' i ω)) a := by + unfold sumRewards sumRewards' + rw [Finset.sum_coe_sort (f := fun s ↦ if A s ω = a then R' s ω else 0) (Iic n)] congr with m simp only [mem_range, mem_Iic] grind -lemma sumRewards_eq_sumRewards' {n : ℕ} {h : ℕ → α × ℝ} (hn : n ≠ 0) : - sumRewards a n h = sumRewards' (n - 1) (fun i ↦ h i) a := by +lemma sumRewards_eq_sumRewards' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (hn : n ≠ 0) : + sumRewards A R' a n ω = sumRewards' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by cases n with | zero => exact absurd rfl hn | succ n => @@ -552,28 +801,30 @@ lemma sumRewards_eq_sumRewards' {n : ℕ} {h : ℕ → α × ℝ} (hn : n ≠ 0) have : n + 1 - 1 = n := by simp exact this ▸ rfl -lemma empMean_add_one_eq_empMean' {n : ℕ} {h : ℕ → α × ℝ} : - empMean a (n + 1) h = empMean' n (fun i ↦ h i) a := by +lemma empMean_add_one_eq_empMean' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} : + empMean A R' a (n + 1) ω = empMean' n (fun i ↦ (A i ω, R' i ω)) a := by unfold empMean empMean' rw [sumRewards_add_one_eq_sumRewards', pullCount_add_one_eq_pullCount'] -lemma empMean_eq_empMean' {n : ℕ} {h : ℕ → α × ℝ} (hn : n ≠ 0) : - empMean a n h = empMean' (n - 1) (fun i ↦ h i) a := by +lemma empMean_eq_empMean' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (hn : n ≠ 0) : + empMean A R' a n ω = empMean' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by unfold empMean empMean' rw [sumRewards_eq_sumRewards' hn, pullCount_eq_pullCount' hn] @[fun_prop] -lemma measurable_sumRewards [MeasurableSingletonClass α] (a : α) (t : ℕ) : - Measurable (sumRewards a t) := by +lemma measurable_sumRewards [MeasurableSingletonClass α] {R' : ℕ → Ω → ℝ} + (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : α) (t : ℕ) : + Measurable (sumRewards A R' a t) := by unfold sumRewards - have h_meas s : Measurable (fun h : ℕ → α × ℝ ↦ if action s h = a then reward s h else 0) := by + have h_meas s : Measurable (fun h : Ω ↦ if A s h = a then R' s h else 0) := by refine Measurable.ite ?_ (by fun_prop) (by fun_prop) exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop @[fun_prop] -lemma measurable_empMean [MeasurableSingletonClass α] (a : α) (n : ℕ) : - Measurable (empMean a n) := by +lemma measurable_empMean [MeasurableSingletonClass α] {R' : ℕ → Ω → ℝ} (hA : ∀ n, Measurable (A n)) + (hR' : ∀ n, Measurable (R' n)) (a : α) (n : ℕ) : + Measurable (empMean A R' a n) := by unfold empMean fun_prop diff --git a/LeanBandits/SequentialLearning/IonescuTulceaSpace.lean b/LeanBandits/SequentialLearning/IonescuTulceaSpace.lean new file mode 100644 index 00000000..9057372f --- /dev/null +++ b/LeanBandits/SequentialLearning/IonescuTulceaSpace.lean @@ -0,0 +1,282 @@ +/- +Copyright (c) 2025 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne, Paulo Rauber +-/ +import LeanBandits.SequentialLearning.Algorithm + +/-! +# Algorithms +-/ + +open MeasureTheory ProbabilityTheory Filter Real Finset + +open scoped ENNReal NNReal + +namespace Learning + +variable {α R Ω : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} {mΩ : MeasurableSpace Ω} + +namespace IT + +/-- Action and reward at step `n`. -/ +def step (n : ℕ) (h : ℕ → α × R) : α × R := h n + +/-- `action n` is the action pulled at time `n`. This is a random variable on the measurable space +`ℕ → α × ℝ`. -/ +def action (n : ℕ) (h : ℕ → α × R) : α := (h n).1 + +/-- `reward n` is the reward at time `n`. This is a random variable on the measurable space +`ℕ → α × R`. -/ +def reward (n : ℕ) (h : ℕ → α × R) : R := (h n).2 + +/-- `hist n` is the history up to time `n`. This is a random variable on the measurable space +`ℕ → α × R`. -/ +def hist (n : ℕ) (h : ℕ → α × R) : Iic n → α × R := fun i ↦ h i + +lemma fst_comp_step (n : ℕ) : Prod.fst ∘ step (α := α) (R := R) n = action n := rfl + +@[fun_prop] +lemma measurable_step (n : ℕ) : Measurable (step n (α := α) (R := R)) := by + unfold step; fun_prop + +@[fun_prop] +lemma measurable_step_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ step p.1 p.2) := + measurable_from_prod_countable_right fun n ↦ (by simp only; fun_prop) + +@[fun_prop] +lemma measurable_action (n : ℕ) : Measurable (action n (α := α) (R := R)) := by + unfold action; fun_prop + +@[fun_prop] +lemma measurable_action_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ action p.1 p.2) := + measurable_from_prod_countable_right fun n ↦ (by simp only; fun_prop) + +@[fun_prop] +lemma measurable_reward (n : ℕ) : Measurable (reward n (α := α) (R := R)) := by + unfold reward; fun_prop + +@[fun_prop] +lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦ reward p.1 p.2) := + measurable_from_prod_countable_right fun n ↦ (by simp only; fun_prop) + +@[fun_prop] +lemma measurable_hist (n : ℕ) : Measurable (hist n (α := α) (R := R)) := by unfold hist; fun_prop + +lemma hist_eq_frestrictLe : + hist = Preorder.frestrictLe («π» := fun _ ↦ α × R) := by + ext n h i : 3 + simp [hist, Preorder.frestrictLe] + +/-- Filtration of the algorithm Seq. -/ +protected def filtration (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : + Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) := + MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R) + +lemma filtration_eq_comap (n : ℕ) : + IT.filtration α R n = MeasurableSpace.comap (hist n) inferInstance := by + simp [IT.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe] + +lemma step_eq_eval_comp_hist (n : ℕ) : + step (α := α) (R := R) n = (fun x ↦ x ⟨n, by simp⟩) ∘ (hist n) := rfl + +lemma action_eq_eval_comp_hist (n : ℕ) : + action (α := α) (R := R) n = (fun x ↦ (x ⟨n, by simp⟩).1) ∘ (hist n) := rfl + +lemma reward_eq_eval_comp_hist (n : ℕ) : + reward (α := α) (R := R) n = (fun x ↦ (x ⟨n, by simp⟩).2) ∘ (hist n) := rfl + +lemma measurable_step_filtration (n : ℕ) : Measurable[IT.filtration α R n] (step n) := by + rw [filtration_eq_comap, step_eq_eval_comp_hist] + exact measurable_comp_comap _ (by fun_prop) + +lemma adapted_step [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] + [SecondCountableTopology α] [OpensMeasurableSpace α] + [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] + [SecondCountableTopology R] [OpensMeasurableSpace R] : + Adapted (IT.filtration α R) (step (α := α) (R := R)) := + fun n ↦ (measurable_step_filtration n).stronglyMeasurable + +lemma measurable_hist_filtration (n : ℕ) : Measurable[IT.filtration α R n] (hist n) := by + simp [filtration_eq_comap, measurable_iff_comap_le] + +lemma adapted_hist [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] + [SecondCountableTopology α] [OpensMeasurableSpace α] + [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] + [SecondCountableTopology R] [OpensMeasurableSpace R] : + Adapted (IT.filtration α R) hist := + fun n ↦ (measurable_hist_filtration n).stronglyMeasurable + +lemma measurable_action_filtration (n : ℕ) : Measurable[IT.filtration α R n] (action n) := by + rw [filtration_eq_comap, action_eq_eval_comp_hist] + exact measurable_comp_comap _ (by fun_prop) + +lemma adapted_action [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α] + [SecondCountableTopology α] [OpensMeasurableSpace α] : + Adapted (IT.filtration α R) action := + fun n ↦ (measurable_action_filtration n).stronglyMeasurable + +lemma measurable_reward_filtration (n : ℕ) : Measurable[IT.filtration α R n] (reward n) := by + rw [filtration_eq_comap, reward_eq_eval_comp_hist] + exact measurable_comp_comap _ (by fun_prop) + +lemma adapted_reward [TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R] + [SecondCountableTopology R] [OpensMeasurableSpace R] : + Adapted (IT.filtration α R) reward := + fun n ↦ (measurable_reward_filtration n).stronglyMeasurable + +section FiltrationAction + +/-- Filtration generated by the history at time `n-1` together with the action at time `n`. -/ +def filtrationAction (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] : + Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) where + seq n := if n = 0 then MeasurableSpace.comap (action 0) inferInstance + else IT.filtration α R (n - 1) ⊔ MeasurableSpace.comap (action n) inferInstance + mono' n m hnm := by + simp only + by_cases hn : n = 0 + · by_cases hm : m = 0 + · simp [hn, hm] + · simp only [hn, ↓reduceIte, hm] + refine le_sup_of_le_left ?_ + rw [← measurable_iff_comap_le] + suffices Measurable[IT.filtration α R 0] (action 0) from + this.mono ((IT.filtration α R).mono zero_le') le_rfl + exact measurable_action_filtration 0 + have hm : m ≠ 0 := by grind + simp only [hn, hm, ↓reduceIte] + have hnm' : n - 1 ≤ m - 1 := by grind + simp only [sup_le_iff] + constructor + · refine le_sup_of_le_left ?_ + exact (IT.filtration α R).mono hnm' + · rcases eq_or_lt_of_le hnm with rfl | hlt + · exact le_sup_of_le_right le_rfl + refine le_sup_of_le_left ?_ + rw [← measurable_iff_comap_le] + have h_le : n ≤ m - 1 := by grind + suffices Measurable[IT.filtration α R n] (action n) from + this.mono ((IT.filtration α R).mono h_le) le_rfl + exact measurable_action_filtration n + le' n := by + by_cases hn : n = 0 + · simp only [hn, ↓reduceIte] + rw [← measurable_iff_comap_le] + fun_prop + simp only [hn, ↓reduceIte, sup_le_iff] + constructor + · exact (IT.filtration α R).le _ + · rw [← measurable_iff_comap_le] + fun_prop + +lemma filtrationAction_zero_eq_comap : + filtrationAction α R 0 = MeasurableSpace.comap (action 0) inferInstance := by + simp [filtrationAction] + +lemma filtrationAction_eq_comap (n : ℕ) (hn : n ≠ 0) : + filtrationAction α R n = + MeasurableSpace.comap (fun ω ↦ (hist (n - 1) ω, action n ω)) inferInstance := by + simp only [filtrationAction, filtration_eq_comap, ← MeasurableSpace.comap_prodMk, hn, ↓reduceIte] + rfl + +lemma filtration_le_filtrationAction_add_one (n : ℕ) : + IT.filtration α R n ≤ filtrationAction α R (n + 1) := le_sup_of_le_left le_rfl + +lemma filtration_le_filtrationAction {m n : ℕ} (h : n < m) : + IT.filtration α R n ≤ filtrationAction α R m := by + have h' : n + 1 ≤ m := by grind + exact (filtration_le_filtrationAction_add_one n).trans ((filtrationAction α R).mono h') + +lemma filtrationAction_le_filtration_self (n : ℕ) : + filtrationAction α R n ≤ IT.filtration α R n := by + by_cases hn : n = 0 + · simp only [hn, filtrationAction_zero_eq_comap] + rw [← measurable_iff_comap_le] + exact measurable_action_filtration 0 + simp only [filtrationAction, hn, ↓reduceIte, sup_le_iff] + constructor + · exact (IT.filtration α R).mono (by grind) + · rw [← measurable_iff_comap_le] + exact measurable_action_filtration _ + +lemma filtrationAction_le_filtration {m n : ℕ} (h : m ≤ n) : + filtrationAction α R m ≤ IT.filtration α R n := + (filtrationAction_le_filtration_self m).trans ((IT.filtration α R).mono h) + +lemma measurable_action_filtrationAction (n : ℕ) : + Measurable[filtrationAction α R n] (action n) := by + simp only [filtrationAction] + rw [measurable_iff_comap_le] + split_ifs with hn + · simp [hn] + · exact le_sup_of_le_right le_rfl + +end FiltrationAction + +section Laws + +lemma hasLaw_step_zero (alg : Algorithm α R) (env : Environment α R) : + HasLaw (step 0) (alg.p0 ⊗ₘ env.ν0) (trajMeasure alg env) where + aemeasurable := Measurable.aemeasurable (by fun_prop) + map_eq := by + unfold step + rw [← coe_default_Iic_zero] + simp only [trajMeasure, Kernel.trajMeasure] + rw [← Measure.deterministic_comp_eq_map (by fun_prop), Measure.comp_assoc, + Kernel.deterministic_comp_eq_map, Kernel.traj_zero_map_eval_zero, + Measure.deterministic_comp_eq_map, Measure.map_map (by fun_prop) (by fun_prop)] + exact Measure.map_id + +lemma hasLaw_action_zero (alg : Algorithm α R) (env : Environment α R) : + HasLaw (action 0) alg.p0 (trajMeasure alg env) where + map_eq := by + rw [← fst_comp_step, ← Measure.map_map (by fun_prop) (by fun_prop), + (hasLaw_step_zero alg env).map_eq, ← Measure.fst, Measure.fst_compProd] + +variable [StandardBorelSpace R] [Nonempty R] + +lemma condDistrib_reward_zero (alg : Algorithm α R) (env : Environment α R) : + condDistrib (reward 0) (action 0) (trajMeasure alg env) + =ᵐ[(trajMeasure alg env).map (action 0)] env.ν0 := by + have h_step := (hasLaw_step_zero alg env).map_eq + have h_action := (hasLaw_action_zero alg env).map_eq + rwa [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop), h_action] + +variable [StandardBorelSpace α] [Nonempty α] + +lemma condDistrib_step (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : + condDistrib (step (n + 1)) (hist n) (trajMeasure alg env) + =ᵐ[(trajMeasure alg env).map (hist n)] stepKernel alg env n := + Kernel.condDistrib_trajMeasure + +lemma condDistrib_action (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : + condDistrib (action (n + 1)) (hist n) (trajMeasure alg env) + =ᵐ[(trajMeasure alg env).map (hist n)] alg.policy n := by + rw [← fst_comp_step] + refine (condDistrib_comp _ (by fun_prop) (by fun_prop)).trans ?_ + filter_upwards [condDistrib_step alg env n] with h h_eq + rw [Kernel.map_apply _ (by fun_prop), h_eq, ← Kernel.map_apply _ (by fun_prop), ← Kernel.fst_eq, + fst_stepKernel] + +lemma condDistrib_reward (alg : Algorithm α R) (env : Environment α R) (n : ℕ) : + condDistrib (reward (n + 1)) (fun ω ↦ (hist n ω, action (n + 1) ω)) (trajMeasure alg env) + =ᵐ[(trajMeasure alg env).map (fun ω ↦ (hist n ω, action (n + 1) ω))] env.feedback n := by + have h_step := condDistrib_step alg env n + have h_action := condDistrib_action alg env n + rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h_step h_action ⊢ + rw [h_action, ← Measure.compProd_assoc, ← stepKernel, ← h_step, + Measure.map_map (by fun_prop) (by fun_prop)] + rfl + +lemma isAlgEnvSeq_trajMeasure (alg : Algorithm α R) (env : Environment α R) : + IsAlgEnvSeq action reward alg env (trajMeasure alg env) where + hasLaw_action_zero := hasLaw_action_zero alg env + hasCondDistrib_reward_zero := ⟨by fun_prop, by fun_prop, condDistrib_reward_zero alg env⟩ + hasCondDistrib_action n := ⟨by fun_prop, by fun_prop, condDistrib_action alg env n⟩ + hasCondDistrib_reward n := ⟨by fun_prop, by fun_prop, condDistrib_reward alg env n⟩ + +end Laws + +end IT + +end Learning diff --git a/LeanBandits/SequentialLearning/StationaryEnv.lean b/LeanBandits/SequentialLearning/StationaryEnv.lean index 995b7f48..a52b4a0c 100644 --- a/LeanBandits/SequentialLearning/StationaryEnv.lean +++ b/LeanBandits/SequentialLearning/StationaryEnv.lean @@ -3,7 +3,7 @@ Copyright (c) 2025 Rémy Degenne. All rights reserved. Released under Apache 2.0 license as described in the file LICENSE. Authors: Rémy Degenne, Paulo Rauber -/ -import LeanBandits.SequentialLearning.Algorithm +import LeanBandits.SequentialLearning.IonescuTulceaSpace /-! # Stationary environments @@ -24,24 +24,29 @@ def stationaryEnv (ν : Kernel α R) [IsMarkovKernel ν] : Environment α R wher feedback _ := ν.prodMkLeft _ ν0 := ν -variable {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] +variable {Ω : Type*} {mΩ : MeasurableSpace Ω} + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] + {P : Measure Ω} [IsProbabilityMeasure P] {A : ℕ → Ω → α} {R' : ℕ → Ω → R} -local notation "𝔓" => trajMeasure alg (stationaryEnv ν) +namespace IsAlgEnvSeq /-- The conditional distribution of the reward at time `n` given the action at time `n` is `ν`. -/ -lemma condDistrib_reward_stationaryEnv [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] (n : ℕ) : - condDistrib (reward n) (action n) 𝔓 =ᵐ[(𝔓).map (action n)] ν := by +lemma condDistrib_reward_stationaryEnv + (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) : + condDistrib (R' n) (A n) P =ᵐ[P.map (A n)] ν := by + have hA := h.measurable_A + have hR' := h.measurable_R cases n with | zero => rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] - change (𝔓).map (step 0) = (𝔓).map (action 0) ⊗ₘ ν - rw [(hasLaw_action_zero alg (stationaryEnv ν)).map_eq, - (hasLaw_step_zero alg (stationaryEnv ν)).map_eq, stationaryEnv_ν0] + change P.map (step A R' 0) = P.map (A 0) ⊗ₘ ν + rw [(hasLaw_action_zero h).map_eq, (hasLaw_step_zero h).map_eq, stationaryEnv_ν0] | succ n => - have h_eq := condDistrib_reward alg (stationaryEnv ν) n + have h_eq := (h.hasCondDistrib_reward n).condDistrib_eq rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at h_eq ⊢ - have : (𝔓).map (action (n + 1)) = ((𝔓).map (fun x ↦ (hist n x, action (n + 1) x))).snd := by + have : P.map (A (n + 1)) = + (P.map (fun x ↦ (hist A R' n x, A (n + 1) x))).snd := by rw [Measure.snd_map_prodMk (by fun_prop)] simp only [stationaryEnv_feedback] at h_eq rw [this, ← Measure.snd_prodAssoc_compProd_prodMkLeft, ← h_eq, @@ -50,10 +55,63 @@ lemma condDistrib_reward_stationaryEnv [StandardBorelSpace α] [Nonempty α] /-- The reward at time `n + 1` is conditionally independent of the history up to time `n` given the action at time `n + 1`. -/ -lemma condIndepFun_reward_hist_action [StandardBorelSpace α] [Nonempty α] - [StandardBorelSpace R] [Nonempty R] (n : ℕ) : - reward (n + 1) ⟂ᵢ[action (n + 1), measurable_action _ ; 𝔓] hist n := - condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft - (by fun_prop) (by fun_prop) (by fun_prop) (condDistrib_reward alg (stationaryEnv ν) n) +lemma condIndepFun_reward_hist_action [StandardBorelSpace Ω] + (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) : + R' (n + 1) ⟂ᵢ[A (n + 1), h.measurable_A _ ; P] hist A R' n := by + have hA := h.measurable_A + have hR' := h.measurable_R + exact condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft + (by fun_prop) (by fun_prop) (by fun_prop) (h.hasCondDistrib_reward n).condDistrib_eq + +lemma condIndepFun_reward_hist_action_action [StandardBorelSpace Ω] + (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) : + R' (n + 1) ⟂ᵢ[A (n + 1), h.measurable_A (n + 1); P] + (fun ω ↦ (hist A R' n ω, A (n + 1) ω)) := by + have h_indep : R' (n + 1) ⟂ᵢ[A (n + 1), h.measurable_A (n + 1); P] hist A R' n := by + convert h.condIndepFun_reward_hist_action n + have hA := h.measurable_A + have hR' := h.measurable_R + exact h_indep.prod_right (by fun_prop) (by fun_prop) (by fun_prop) + +lemma condIndepFun_reward_hist_action_action' [StandardBorelSpace Ω] + (h : IsAlgEnvSeq A R' alg (stationaryEnv ν) P) (n : ℕ) (hn : n ≠ 0) : + R' n ⟂ᵢ[A n, h.measurable_A n; P] (fun ω ↦ (hist A R' (n - 1) ω, A n ω)) := by + have := h.condIndepFun_reward_hist_action_action (n - 1) + grind + +end IsAlgEnvSeq + +namespace IT + +local notation "𝔓" => trajMeasure alg (stationaryEnv ν) + +/-- The conditional distribution of the reward at time `n` given the action at time `n` is `ν`. -/ +lemma condDistrib_reward_stationaryEnv (n : ℕ) : + condDistrib (IT.reward n) (IT.action n) 𝔓 =ᵐ[(𝔓).map (IT.action n)] ν := + IsAlgEnvSeq.condDistrib_reward_stationaryEnv + (IT.isAlgEnvSeq_trajMeasure alg (stationaryEnv ν)) n + +/-- The reward at time `n + 1` is conditionally independent of the history up to time `n` +given the action at time `n + 1`. -/ +lemma condIndepFun_reward_hist_action (n : ℕ) : + IT.reward (n + 1) ⟂ᵢ[IT.action (n + 1), IT.measurable_action _ ; 𝔓] IT.hist n := + IsAlgEnvSeq.condIndepFun_reward_hist_action + (IT.isAlgEnvSeq_trajMeasure alg (stationaryEnv ν)) n + +lemma condIndepFun_reward_hist_action_action + {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) : + reward (n + 1) ⟂ᵢ[action (n + 1), measurable_action (n + 1); trajMeasure alg (stationaryEnv ν)] + (fun ω ↦ (hist n ω, action (n + 1) ω)) := + IsAlgEnvSeq.condIndepFun_reward_hist_action_action + (IT.isAlgEnvSeq_trajMeasure alg (stationaryEnv ν)) n + +lemma condIndepFun_reward_hist_action_action' + {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) (hn : n ≠ 0) : + reward n ⟂ᵢ[action n, measurable_action n; trajMeasure alg (stationaryEnv ν)] + (fun ω ↦ (hist (n - 1) ω, action n ω)) := + IsAlgEnvSeq.condIndepFun_reward_hist_action_action' + (IT.isAlgEnvSeq_trajMeasure alg (stationaryEnv ν)) n hn + +end IT end Learning diff --git a/blueprint/lean_decls b/blueprint/lean_decls index a677429e..54612c19 100644 --- a/blueprint/lean_decls +++ b/blueprint/lean_decls @@ -2,25 +2,34 @@ Learning.Algorithm Learning.Environment Learning.detAlgorithm Learning.stationaryEnv +Learning.IsAlgEnvSeq +Learning.IsAlgEnvSeq.hist +Learning.IsAlgEnvSeq.step +Learning.IsAlgEnvSeq.hasLaw_step_zero +Learning.IsAlgEnvSeq.hasCondDistrib_step +Learning.IsAlgEnvSeq.filtration +Learning.IsAlgEnvSeq.filtrationAction +Learning.isAlgEnvSeq_unique +Learning.IsAlgEnvSeq.condDistrib_reward_stationaryEnv +Learning.IsAlgEnvSeq.condIndepFun_reward_hist_action ProbabilityTheory.Kernel.traj ProbabilityTheory.Kernel.trajMeasure -Learning.step -Learning.hist -Learning.filtration -Learning.adapted_step -Learning.adapted_hist +Learning.IT.step +Learning.IT.hist +Learning.IT.filtration +Learning.IT.adapted_step +Learning.IT.adapted_hist ProbabilityTheory.Kernel.condDistrib_trajMeasure -Learning.hasLaw_step_zero -Learning.action -Learning.reward -Learning.adapted_action -Learning.adapted_reward -Learning.condDistrib_action -Learning.condDistrib_reward -Learning.hasLaw_action_zero -Learning.condDistrib_reward_zero -Learning.condDistrib_reward_stationaryEnv -Learning.condIndepFun_reward_hist_action +Learning.IsAlgEnvSeq.hasLaw_step_zero +Learning.IT.action +Learning.IT.reward +Learning.IT.adapted_action +Learning.IT.adapted_reward +Learning.IT.condDistrib_action +Learning.IT.condDistrib_reward +Learning.IT.hasLaw_action_zero +Learning.IT.condDistrib_reward_zero +Learning.IT.isAlgEnvSeq_trajMeasure Learning.pullCount Learning.pullCount_zero Learning.pullCount_mono @@ -44,17 +53,17 @@ Learning.empMean Learning.sum_rewardByCount_eq_sumRewards Bandits.Bandit.trajMeasure Bandits.Bandit.measure -Bandits.measurable_comap_indicator_stepsUntil_eq -ProbabilityTheory.CondIndepFun.prod_right -Bandits.condIndepFun_reward_stepsUntil_arm +Bandits.ArrayModel.probSpace +Bandits.ArrayModel.arrayMeasure +Bandits.ArrayModel.algFunction +Bandits.ArrayModel.initAlgFunction +Bandits.ArrayModel.isAlgEnvSeq_arrayMeasure +Learning.measurable_comap_indicator_stepsUntil_eq +Bandits.condIndepFun_reward_stepsUntil_action Bandits.reward_cond_stepsUntil ProbabilityTheory.condDistrib_ae_eq_cond Bandits.condDistrib_rewardByCount_stepsUntil Bandits.hasLaw_rewardByCount -ProbabilityTheory.iIndepFun_nat_iff_forall_indepFun -Bandits.iIndepFun_rewardByCount' -Bandits.identDistrib_rewardByCount_stream -Bandits.identDistrib_sum_Icc_rewardByCount Bandits.regret Bandits.gap Learning.sum_pullCount_mul @@ -83,4 +92,6 @@ Bandits.UCB.pullCount_le_add_three Bandits.UCB.pullCount_le_add_three_ae Bandits.UCB.some_sum_eq_zero Bandits.UCB.expectation_pullCount_le -Bandits.UCB.regret_le \ No newline at end of file +Bandits.UCB.regret_le +ProbabilityTheory.CondIndepFun.prod_right +ProbabilityTheory.iIndepFun_nat_iff_forall_indepFun \ No newline at end of file diff --git a/blueprint/src/biblio.bib b/blueprint/src/biblio.bib index 82ef848f..f6e1818e 100644 --- a/blueprint/src/biblio.bib +++ b/blueprint/src/biblio.bib @@ -269,3 +269,10 @@ @inproceedings{hirata2023semantic year={2023}, organization={Schloss Dagstuhl--Leibniz-Zentrum f{\"u}r Informatik} } + +@book{lattimore2020bandit, + title={Bandit algorithms}, + author={Lattimore, Tor and Szepesv{\'a}ri, Csaba}, + year={2020}, + publisher={Cambridge University Press} +} diff --git a/blueprint/src/chapters/algorithm.tex b/blueprint/src/chapters/algorithm.tex index 7b3de6d7..044be591 100644 --- a/blueprint/src/chapters/algorithm.tex +++ b/blueprint/src/chapters/algorithm.tex @@ -1,12 +1,19 @@ \chapter{Iterative stochastic algorithms} +Warning: all times start at zero. + +TODO: notations + +All measurable spaces are assumed to be standard Borel. + + \begin{definition}[Algorithm]\label{def:algorithm} \leanok \lean{Learning.Algorithm} A sequential, stochastic algorithm with actions in a measurable space $\mathcal{A}$ and observations in a measurable space $\mathcal{R}$ is described by the following data: \begin{itemize} - \item for all $t \in \mathbb{N}$, a policy $\pi_t : (\mathcal{A} \times \mathcal{R})^{t+1} \rightsquigarrow \mathcal{A}$, a Markov kernel which gives the distribution of the action of the algorithm at time $t+1$ given the history of previous pulls and observations, + \item for all $t \in \mathbb{N}$, a policy $\pi_t : (\mathcal{A} \times \mathcal{R})^{t+1} \rightsquigarrow \mathcal{A}$, a Markov kernel which gives the distribution of the action of the algorithm at time $t+1$ given the history of previous actions and observations, \item $P_0 \in \mathcal{P}(\mathcal{A})$, a probability measure that gives the distribution of the first action. \end{itemize} \end{definition} @@ -40,10 +47,7 @@ \chapter{Iterative stochastic algorithms} An environment is stationary if there exists a Markov kernel $\nu : \mathcal{A} \rightsquigarrow \mathcal{R}$ such that $\nu'_0 = \nu$ and for all $t \in \mathbb{N}$, for all $h_t \in (\mathcal{A} \times \mathcal{R})^{t+1}$, for all $a \in \mathcal{A}$, $\nu_t(h_t, a) = \nu(a)$. \end{definition} - -\begin{remark}[Lean remark: properties vs constructors] -There are several ways to implement the last two definitions in Lean. We could write them as properties of algorithms and environments, or we can implement constructors that create algorithms and environments from the data in the definitions. We chose the latter option. Time will tell if it was a good choice. -\end{remark} +TODO: possibly change the ``stationary'' name. Let's detail four examples of interactions between an algorithm and an environment. @@ -64,18 +68,122 @@ \chapter{Iterative stochastic algorithms} \end{enumerate} +We will want to make global probabilistic statements about the whole sequence of actions and observations. +For example, we may want to prove that an optimization algorithm converges to the minimum of a function almost surely. +For such a statement to make sense, we need a probability space on which the whole sequence of actions and observations is defined as a random variable. + +We denote by $P[X \mid Y]$ the conditional distribution of a random variable $X$ given another random variable $Y$ under a probability measure $P$. + + +\begin{definition}[Algorithm-environment interaction]\label{def:IsAlgEnvSeq} + \uses{def:algorithm, def:environment} + \leanok + \lean{Learning.IsAlgEnvSeq} +Let $\mathfrak{A}$ be an algorithm as in Definition~\ref{def:algorithm} and $\mathfrak{E}$ be an environment as in Definition~\ref{def:environment}. +A probability space $(\Omega, P)$ and two sequences of random variables $A : \mathbb{N} \to \Omega \to \mathcal{A}$ and $R : \mathbb{N} \to \Omega \to \mathcal{R}$ form an algorithm-environment interaction for $\mathfrak{A}$ and $\mathfrak{E}$ if the following conditions hold: +\begin{enumerate} + \item The law of $A_0$ is $P_0$. + \item $P \left[ R_0 \mid A_0 \right] = \nu'_0$. + \item For all $t \in \mathbb{N}$, $P\left[A_{t+1} \mid A_0, R_0, \ldots, A_t, R_t \right] = \pi_t$. + \item For all $t \in \mathbb{N}$, $P\left[R_{t+1} \mid A_0, R_0, \ldots, A_t, R_t, A_{t+1}\right] = \nu_t$. +\end{enumerate} +\end{definition} + + +\begin{definition}[History]\label{def:history} + \leanok + \lean{Learning.IsAlgEnvSeq.hist, Learning.IsAlgEnvSeq.step} +For two sequences of random variables $A : \mathbb{N} \to \Omega \to \mathcal{A}$ and $R : \mathbb{N} \to \Omega \to \mathcal{R}$ (actions and observations), we call step of the interaction at time $t$ the random variable $X_t : \Omega \to \mathcal{A} \times \mathcal{R}$ defined by $X_t(\omega) = (A_t(\omega), R_t(\omega))$. +We call history up to time $t$ the random variable $H_t : \Omega \to (\mathcal{A} \times \mathcal{R})^{t+1}$ defined by $H_t(\omega) = (X_0(\omega), \ldots, X_t(\omega))$. +\end{definition} + + +\begin{lemma}\label{lem:law_step} + \uses{def:IsAlgEnvSeq, def:history} + \leanok + \lean{Learning.IsAlgEnvSeq.hasLaw_step_zero, Learning.IsAlgEnvSeq.hasCondDistrib_step} +In an algorithm-environment interaction $(A, R, P)$ as in Definition~\ref{def:IsAlgEnvSeq}, +\begin{itemize} + \item the law of the initial step $X_0$ is $P_0 \otimes \nu'_0$, + \item for all $t \in \mathbb{N}$, $P \left[ X_{t+1} \mid H_t \right] = \pi_t \otimes \nu_t$. +\end{itemize} +\end{lemma} + +\begin{proof}\leanok +Immediate from the properties of an algorithm-environment interaction. +\end{proof} + + +\begin{definition}\label{def:IsAlgEnvSeq.filtration} + \uses{def:IsAlgEnvSeq, def:history} + \leanok + \lean{Learning.IsAlgEnvSeq.filtration, Learning.IsAlgEnvSeq.filtrationAction} +For an algorithm-environment interaction $(A, R, P)$ as in Definition~\ref{def:IsAlgEnvSeq}, we denote by $\mathcal{F}_t$ the sigma-algebra generated by the history up to time $t$: $\mathcal{F}_t = \sigma(H_t)$. +We denote by $\mathcal{F}^A_t$ the sigma-algebra generated by the history up to time $t-1$ and the action at time $t$: $\mathcal{F}^A_t = \sigma(H_{t-1}, A_t)$. +\end{definition} + + +\begin{theorem}[\cite{lattimore2020bandit}, Proposition 4.8]\label{thm:isAlgEnvSeq_unique} + \uses{def:IsAlgEnvSeq} + \leanok + \lean{Learning.isAlgEnvSeq_unique} +If $(A, R, P)$ and $(A', R', P')$ are two algorithm-environment interactions for the same algorithm $\mathfrak{A}$ and environment $\mathfrak{E}$, then the joint distributions of the sequences of actions and observations are equal: the law of $(A_i, R_i)_{i \in \mathbb{N}}$ under $P$ is equal to the law of $(A'_i, R'_i)_{i \in \mathbb{N}}$ under $P'$. +\end{theorem} + +\begin{proof}\leanok + +\end{proof} + + + +\section{Stationary environment} + +Recall that in a stationary environment, there exists a Markov kernel $\nu : \mathcal{A} \rightsquigarrow \mathcal{R}$ such that $\nu'_0 = \nu$ and for all $t \in \mathbb{N}$, for all $h_t \in (\mathcal{A} \times \mathcal{R})^{t+1}$, for all $a \in \mathcal{A}$, $\nu_t(h_t, a) = \nu(a)$. + +Let $(A, R, P)$ be an algorithm-environment interaction in a stationary environment with kernel $\nu$. + +\begin{lemma}\label{lem:condDistrib_reward_stationaryEnv} + \uses{def:IsAlgEnvSeq, def:stationaryEnv} + \leanok + \lean{Learning.IsAlgEnvSeq.condDistrib_reward_stationaryEnv} +In a stationary environment, for any $t \in \mathbb{N}$, the conditional distribution $P\left[R_t \mid A_t\right]$ is $(A_{t*} P_{\mathcal{T}})$-almost surely equal to $\nu$. +\end{lemma} + +\begin{proof}\leanok + \uses{lem:law_step, def:stationaryEnv} + +\end{proof} + + +\begin{lemma}\label{lem:condIndepFun_reward_hist_action} + \uses{def:IsAlgEnvSeq, def:stationaryEnv} + \leanok + \lean{Learning.IsAlgEnvSeq.condIndepFun_reward_hist_action} +In a stationary environment, for any $t \in \mathbb{N}$, the reward $R_{t+1}$ is conditionally independent of the history $H_t$ given the action $A_{t+1}$ (more succinctly, $R_{t+1} \ind H_t \mid A_{t+1}$). +\end{lemma} + +\begin{proof}\leanok + +\end{proof} + \section{Probability space: Ionescu-Tulcea theorem} +In Theorem~\ref{thm:isAlgEnvSeq_unique}, we saw that the distribution of the sequence of actions and observations in a suitable probability space is uniquely determined by the algorithm and the environment. +We now show that such a probability space actually exists: for any algorithm and environment, we build an algorithm-environment interaction as in Definition~\ref{def:IsAlgEnvSeq}. + + + +\subsection{Ionescu-Tulcea theorem} + If we group together the policy of the algorithm and the kernel of the environment at each time step, we get a sequence of Markov kernels $(\kappa_t)_{t \in \mathbb{N}}$, with $\kappa_t : (\mathcal{A} \times \mathcal{R})^{t+1} \rightsquigarrow (\mathcal{A} \times \mathcal{R})$. -We will want to make global probabilistic statements about the whole sequence of actions and observations. -For example, we may want to prove that an optimization algorithm converges to the minimum of a function almost surely. -For such a statement to make sense, we need a probability space on which the whole sequence of actions and observations is defined as a random variable. + We now abstract that situation and consider a sequence of measurable spaces $(\Omega_t)_{t \in \mathbb{N}}$, a probability measure $\mu$ on $\Omega_0$ and a sequence of Markov kernels $\kappa_t : \prod_{s=0}^t \Omega_s \rightsquigarrow \Omega_{t+1}$. The Ionescu-Tulcea theorem builds a probability space from the sequence of kernels and the initial measure. + \begin{theorem}[Ionescu-Tulcea]\label{thm:ionescu-tulcea} \mathlibok \lean{ProbabilityTheory.Kernel.traj} @@ -100,9 +208,9 @@ \section{Probability space: Ionescu-Tulcea theorem} \end{definition} -\begin{definition}[Step and history]\label{def:history} +\begin{definition}[Step and history]\label{def:IT.history} \leanok - \lean{Learning.step, Learning.hist} + \lean{Learning.IT.step, Learning.IT.hist} For $t \in \mathbb{N}$, we denote by $X_t \in \Omega_t$ the random variable describing the time step $t$, and by $H_t \in \prod_{s=0}^t \Omega_s$ the history up to time $t$. Formally, these are measurable functions on $\Omega_{\mathcal{T}}$, defined by $X_t(\omega) = \omega_t$ and $H_t(\omega) = (\omega_1, \ldots, \omega_t)$. \end{definition} @@ -110,10 +218,10 @@ \section{Probability space: Ionescu-Tulcea theorem} Note: $(X_t)_{t \in \mathbb{N}}$ is the canonical process on $\Omega_{\mathcal{T}}$. $H_t$ is equal to $\pi_{[0,t]}$. -\begin{definition}[Filtration]\label{def:filtration} - \uses{def:history} +\begin{definition}[Filtration]\label{def:IT.filtration} + \uses{def:IT.history} \leanok - \lean{Learning.filtration} + \lean{Learning.IT.filtration} For $t \in \mathbb{N}$, we denote by $\mathcal{F}_t$ the sigma-algebra generated by the history up to time $t$: $\mathcal{F}_t = \sigma(H_t)$. The family $(\mathcal{F}_t)_{t \in \mathbb{N}}$ is a filtration on $\Omega_{\mathcal{T}}$. \end{definition} @@ -122,9 +230,9 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:adapted_history} - \uses{def:history, def:filtration} + \uses{def:IT.history, def:IT.filtration} \leanok - \lean{Learning.adapted_step, Learning.adapted_hist} + \lean{Learning.IT.adapted_step, Learning.IT.adapted_hist} The random variables $X_t$ and $H_t$ are $\mathcal{F}_t$-measurable. Said differently, the processes $(X_t)_{t \in \mathbb{N}}$ and $(H_t)_{t \in \mathbb{N}}$ are adapted to the filtration $(\mathcal{F}_t)_{t \in \mathbb{N}}$. \end{lemma} @@ -134,11 +242,8 @@ \section{Probability space: Ionescu-Tulcea theorem} \end{proof} -We now list properties of those random variables that follow from the construction of the trajectory measure. -We write $P[X \mid Y]$ for the conditional distribution of a random variable $X$ given another random variable $Y$ under a probability measure $P$. - \begin{lemma}\label{lem:condDistrib_X_add_one} - \uses{def:history, def:trajMeasure} + \uses{def:IT.history, def:trajMeasure} \leanok \lean{ProbabilityTheory.Kernel.condDistrib_trajMeasure} For any $t \in \mathbb{N}$, the conditional distribution $P_{\mathcal{T}}\left[X_{t+1} \mid H_t\right]$ is $((H_t)_* P_{\mathcal{T}})$-almost surely equal to $\kappa_t$. @@ -152,9 +257,9 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:law_X_zero} - \uses{def:history, def:trajMeasure} + \uses{def:IT.history, def:trajMeasure} \leanok - \lean{Learning.hasLaw_step_zero} + \lean{Learning.IsAlgEnvSeq.hasLaw_step_zero} The law of $X_0$ under $P_{\mathcal{T}}$ is $\mu$. \end{lemma} @@ -163,28 +268,28 @@ \section{Probability space: Ionescu-Tulcea theorem} \end{proof} -\paragraph{Case of an algorithm-environment interaction.} -We suppose now that, as in the algorithm-environment interaction, $\Omega_t = \mathcal{A}_t \times \mathcal{R}_t$ for some measurable spaces $\mathcal{A}_t$ and $\mathcal{R}_t$, and that for all $t \in \mathbb{N}$, $\kappa_t = \pi_t \otimes \nu_t$ for policy kernels $\pi_t : \prod_{s=0}^t(\mathcal{A}_s \times \mathcal{R}_s) \rightsquigarrow \mathcal{A}$ and feedback kernels $\nu_t : \prod_{s=0}^t(\mathcal{A}_s \times \mathcal{R}_s) \times \mathcal{A} \rightsquigarrow \mathcal{R}$. -Likewise, $\mu = \alpha_0 \otimes \nu'_0$ for a probability measure $\alpha_0$ on $\mathcal{A}_0$ and a Markov kernel $\nu'_0 : \mathcal{A}_0 \rightsquigarrow \mathcal{R}_0$. -The step random variable $X_t$ takes values in $\mathcal{A}_t \times \mathcal{R}_t$. -TODO: the code does not have $\mathcal{A}_t$ but a unique $\mathcal{A}$, same for $\mathcal{R}$. +\subsection{Case of an algorithm-environment interaction} -\begin{definition}\label{def:actionReward} - \uses{def:history} +We now go back to the setting of an algorithm interacting with an environment and suppose that $\Omega_t = \mathcal{A} \times \mathcal{R}$ for some measurable spaces $\mathcal{A}$ and $\mathcal{R}$, and that for all $t \in \mathbb{N}$, $\kappa_t = \pi_t \otimes \nu_t$ for policy kernels $\pi_t : (\mathcal{A} \times \mathcal{R})^{t+1} \rightsquigarrow \mathcal{A}$ and feedback kernels $\nu_t : (\mathcal{A} \times \mathcal{R})^{t+1} \times \mathcal{A} \rightsquigarrow \mathcal{R}$. +Likewise, $\mu = P_0 \otimes \nu'_0$ for a probability measure $P_0$ on $\mathcal{A}$ and a Markov kernel $\nu'_0 : \mathcal{A}_\rightsquigarrow \mathcal{R}$. +The step random variable $X_t$ takes values in $\mathcal{A} \times \mathcal{R}$. + +\begin{definition}\label{def:IT.actionReward} + \uses{def:IT.history} \leanok - \lean{Learning.action, Learning.reward} -We write $A_t$ and $R_t$ for the projections of $X_t$ on $\mathcal{A}_t$ and $\mathcal{R}_t$ respectively. + \lean{Learning.IT.action, Learning.IT.reward} +We write $A_t$ and $R_t$ for the projections of $X_t$ on $\mathcal{A}$ and $\mathcal{R}$ respectively. $A_t$ is the action taken at time $t$ and $R_t$ is the reward received at time $t$. -Formally, $A_t(\omega) = \omega_{t,1}$ and $R_t(\omega) = \omega_{t,2}$ for $\omega = \prod_{t=0}^{+\infty}(\omega_{t,1}, \omega_{t,2}) \in \prod_{t=0}^{+\infty} \mathcal{A}_t \times \mathcal{R}_t$. +Formally, $A_t(\omega) = \omega_{t,1}$ and $R_t(\omega) = \omega_{t,2}$ for $\omega = \prod_{t=0}^{+\infty}(\omega_{t,1}, \omega_{t,2}) \in \Omega_{\mathcal{T}} = \prod_{t=0}^{+\infty} \mathcal{A} \times \mathcal{R}$. \end{definition} \begin{lemma}\label{lem:adapted_action_reward} - \uses{def:actionReward, def:filtration} + \uses{def:IT.actionReward, def:IT.filtration} \leanok - \lean{Learning.adapted_action, Learning.adapted_reward} + \lean{Learning.IT.adapted_action, Learning.IT.adapted_reward} The random variables $A_t$ and $R_t$ are $\mathcal{F}_t$-measurable. Said differently, the processes $(A_t)_{t \in \mathbb{N}}$ and $(R_t)_{t \in \mathbb{N}}$ are adapted to the filtration $(\mathcal{F}_t)_{t \in \mathbb{N}}$. \end{lemma} @@ -198,9 +303,9 @@ \section{Probability space: Ionescu-Tulcea theorem} We need to check that the random variables $A_t$ and $R_t$ have the expected conditional distributions. \begin{lemma}\label{lem:condDistrib_A_add_one} - \uses{def:actionReward, def:trajMeasure, def:algorithm, def:environment} + \uses{def:IT.actionReward, def:trajMeasure, def:algorithm, def:environment} \leanok - \lean{Learning.condDistrib_action} + \lean{Learning.IT.condDistrib_action} For any $t \in \mathbb{N}$, the conditional distribution $P_{\mathcal{T}}\left[A_{t+1} \mid H_t\right]$ is $((H_t)_* P_{\mathcal{T}})$-almost surely equal to $\pi_t$. \end{lemma} @@ -212,9 +317,9 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:condDistrib_R_add_one} - \uses{def:actionReward, def:trajMeasure, def:algorithm, def:environment} + \uses{def:IT.actionReward, def:trajMeasure, def:algorithm, def:environment} \leanok - \lean{Learning.condDistrib_reward} + \lean{Learning.IT.condDistrib_reward} For any $t \in \mathbb{N}$, the conditional distribution $P_{\mathcal{T}}\left[R_{t+1} \mid H_t, A_{t+1}\right]$ is $((H_t, A_{t+1})_* P_{\mathcal{T}})$-almost surely equal to $\nu_t$. \end{lemma} @@ -231,9 +336,9 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:law_A_zero} - \uses{def:actionReward, def:trajMeasure, def:algorithm, def:environment} + \uses{def:IT.actionReward, def:trajMeasure, def:algorithm, def:environment} \leanok - \lean{Learning.hasLaw_action_zero} + \lean{Learning.IT.hasLaw_action_zero} The law of $A_0$ under $P_{\mathcal{T}}$ is $\alpha_0$. \end{lemma} @@ -244,9 +349,9 @@ \section{Probability space: Ionescu-Tulcea theorem} \begin{lemma}\label{lem:condDistrib_R_zero} - \uses{def:actionReward, def:trajMeasure, def:algorithm, def:environment} + \uses{def:IT.actionReward, def:trajMeasure, def:algorithm, def:environment} \leanok - \lean{Learning.condDistrib_reward_zero} + \lean{Learning.IT.condDistrib_reward_zero} The conditional distribution $P_{\mathcal{T}}\left[R_0 \mid A_0\right]$ is $(A_{0*} P_{\mathcal{T}})$-almost surely equal to $\nu'_0$. \end{lemma} @@ -260,34 +365,16 @@ \section{Probability space: Ionescu-Tulcea theorem} \end{proof} - -\section{Stationary environment} - -Recall that in a stationary environment, there exists a Markov kernel $\nu : \mathcal{A} \rightsquigarrow \mathcal{R}$ such that $\nu'_0 = \nu$ and for all $t \in \mathbb{N}$, for all $h_t \in (\mathcal{A} \times \mathcal{R})^{t+1}$, for all $a \in \mathcal{A}$, $\nu_t(h_t, a) = \nu(a)$. - - -\begin{lemma}\label{lem:condDistrib_reward_stationaryEnv} - \uses{def:actionReward, def:trajMeasure, def:algorithm, def:stationaryEnv} - \leanok - \lean{Learning.condDistrib_reward_stationaryEnv} -In a stationary environment, for any $t \in \mathbb{N}$, the conditional distribution $P_{\mathcal{T}}\left[R_t \mid A_t\right]$ is $(A_{t*} P_{\mathcal{T}})$-almost surely equal to $\nu$. -\end{lemma} - -\begin{proof}\leanok - \uses{lem:condDistrib_R_add_one, def:stationaryEnv} - -\end{proof} - - -\begin{lemma}\label{lem:condIndepFun_reward_hist_action} - \uses{def:actionReward, def:trajMeasure, def:algorithm, def:stationaryEnv} +\begin{theorem}\label{thm:isAlgEnvSeq_trajMeasure} + \uses{def:IsAlgEnvSeq, def:trajMeasure} \leanok - \lean{Learning.condIndepFun_reward_hist_action} -In a stationary environment, for any $t \in \mathbb{N}$, the reward $R_{t+1}$ is conditionally independent of the history $H_t$ given the action $A_{t+1}$ (more succinctly, $R_{t+1} \ind H_t \mid A_{t+1}$). -\end{lemma} + \lean{Learning.IT.isAlgEnvSeq_trajMeasure} +In the probability space $(\Omega_{\mathcal{T}}, P_{\mathcal{T}})$ constructed from an algorithm $\mathfrak{A}$ and an environment $\mathfrak{E}$ as above, the sequences of random variables $A : \mathbb{N} \to \Omega_{\mathcal{T}} \to \mathcal{A}$ and $R : \mathbb{N} \to \Omega_{\mathcal{T}} \to \mathcal{R}$ form an algorithm-environment interaction for $\mathfrak{A}$ and $\mathfrak{E}$. +\end{theorem} \begin{proof}\leanok - + \uses{lem:law_A_zero, lem:condDistrib_R_zero, lem:condDistrib_A_add_one, lem:condDistrib_R_add_one} +The four conditions of Definition~\ref{def:IsAlgEnvSeq} are exactly the statements of Lemmas~\ref{lem:law_A_zero}, \ref{lem:condDistrib_R_zero}, \ref{lem:condDistrib_A_add_one} and \ref{lem:condDistrib_R_add_one}. \end{proof} @@ -298,7 +385,7 @@ \section{Finitely many actions} We can also define the time step at which an action was chosen a certain number of times, and the value of the reward obtained when pulling an action for the $m$-th time. \begin{definition}[Pull counts]\label{def:pullCount} - \uses{def:actionReward} + \uses{def:IT.actionReward} \leanok \lean{Learning.pullCount} For an action $a \in \mathcal{A}$ and a time $t \in \mathbb{N}$, we denote by $N_{t,a}$ the number of times that action $a$ has been chosen before time $t$, that is $N_{t,a} = \sum_{s=0}^{t-1} \mathbb{I}\{A_s = a\}$. @@ -312,15 +399,15 @@ \section{Finitely many actions} That means that any tool used to define a policy must be a function defined on $(\mathcal{A} \times \mathcal{R})^{t+1}$. For example a definition of the empirical mean of an action must be a function $t : \mathbb{N} \to (\mathcal{A} \times \mathcal{R})^{t+1} \to \mathbb{R}$. -When we analyze an algorithm, we work on the other hand on the bandit probability space $(\Omega, \mathbb{P})$, in which $\Omega = (\mathcal{A} \times \mathcal{R})^{\mathbb{N}}$ is the full history, which describes the whole sequence of actions and rewards. +When we analyze an algorithm, we work on the other hand on a probability space $(\Omega, P)$, in which $\Omega$ could be for example $(\mathcal{A} \times \mathcal{R})^{\mathbb{N}}$, the full history, which describes the whole sequence of actions and rewards. As a stochastic process, the empirical mean of an action is a function $\mathbb{N} \to (\mathcal{A} \times \mathcal{R})^{\mathbb{N}} \to \mathbb{R}$. -Thus there are two similar but still distinct types of objects: those defined on the partial history, which are used to build algorithms, and those defined on the full history, which are used to analyze algorithms. +Thus there are two similar but still distinct types of objects: those defined on the partial history, which are used to build algorithms, and those defined on a generic probability space (the full history in the Ionescu-Tulcea construction), which are used to analyze algorithms. \end{remark} \begin{lemma}\label{lem:pullCount_basic} - \uses{def:pullCount, def:actionReward} + \uses{def:pullCount, def:IT.actionReward} \leanok \lean{Learning.pullCount_zero, Learning.pullCount_mono, Learning.pullCount_add_one, Learning.pullCount_le, Learning.pullCount_congr} We note the following basic properties of $N_{t,a}$: @@ -339,10 +426,10 @@ \section{Finitely many actions} \begin{lemma}\label{lem:predictable_pullCount} - \uses{def:filtration, def:pullCount} + \uses{def:IsAlgEnvSeq.filtration, def:pullCount} \leanok \lean{Learning.isPredictable_pullCount} -Let $a \in \mathcal{A}$. The process $(N_{t,a})_{t \in \mathbb{N}}$ is predictable with respect to the filtration $\mathcal{F}$. +Let $a \in \mathcal{A}$. The process $(N_{t,a})_{t \in \mathbb{N}}$ is predictable with respect to the filtration $\mathcal{F}$ of the algorithm-environment interaction. \end{lemma} \begin{proof}\leanok @@ -384,7 +471,7 @@ \section{Finitely many actions} \begin{lemma}\label{lem:isStoppingTime_stepsUntil} - \uses{def:filtration, def:stepsUntil} + \uses{def:IsAlgEnvSeq.filtration, def:stepsUntil} \leanok \lean{Learning.isStoppingTime_stepsUntil} Let $a \in \mathcal{A}$. For any $n > 0$, the random variable $T_{n,a}$ is a stopping time with respect to the filtration $\mathcal{F}$. @@ -441,7 +528,7 @@ \section{Scalar rewards} \begin{definition}[Sum of rewards]\label{def:sumRewards} - \uses{def:actionReward} + \uses{def:IT.actionReward} \leanok \lean{Learning.sumRewards} Let $S_{t, a} = \sum_{s=0}^{t-1} R_s \mathbb{I}\{A_s = a\}$ be the sum of the rewards obtained by chosing action $a$ before time $t$. diff --git a/blueprint/src/chapters/bandit.tex b/blueprint/src/chapters/bandit.tex index a3f3af6a..e89628b6 100644 --- a/blueprint/src/chapters/bandit.tex +++ b/blueprint/src/chapters/bandit.tex @@ -34,6 +34,54 @@ \section{Algorithm, bandit and probability space} \end{definition} +\section{The array model of rewards} + +We previously built a probability space on which we can define the sequence of arms and rewards generated by the interaction between the algorithm and the bandit, using the Ionescu-Tulcea theorem. +From Theorem~\ref{thm:isAlgEnvSeq_unique}, we know that the law of the sequence of arms and rewards is independent of the probability space used to define them. +Nonetheless, we now build an alternative model of the rewards, on which it will be easier to prove concentration inequalities. +By the uniqueness of the law, these statements will then transfer to any algorithm-environment interaction. + + +\begin{definition}\label{def:arrayMeasure} + \leanok + \lean{Bandits.ArrayModel.probSpace, Bandits.ArrayModel.arrayMeasure} +Let $I = [0,1]$ and let $P_U$ be the uniform distribution on $I$. We define the probability space $(\Omega_{\mathcal{A}}, P_{\mathcal{A}})$, where +\begin{align*} + \Omega_{\mathcal{A}} &:= I^{\mathbb{N}} \times \mathcal{R}^{\mathbb{N} \times \mathcal{A}} + \: , \\ + P_{\mathcal{A}} &:= \left( \bigotimes_{n \in \mathbb{N}} P_U \right) \otimes \left( \bigotimes_{n \in \mathbb{N}, a \in \mathcal{A}} \nu(a) \right) + \: . +\end{align*} +\end{definition} + + +\begin{definition}\label{def:algFunction} + \uses{def:algorithm} + \leanok + \lean{Bandits.ArrayModel.algFunction, Bandits.ArrayModel.initAlgFunction} +Since $\mathcal{A}$ and $\mathcal{R}$ are standard Borel spaces, there exists jointly measurable functions $f'_0 : I \to \mathcal{A}$ and $f_t : (\mathcal{A} \times \mathcal{R})^{t+1} \times I \to \mathcal{A}$ such that +\begin{itemize} + \item the law of $f'_0$ is $P_0$, + \item for all history $h_t \in (\mathcal{A} \times \mathcal{R})^{t+1}$, the law of $f_t(h_t, \cdot)$ is $\pi_t(h_t)$. +\end{itemize} +\end{definition} + + +TODO: lots of results + + +\begin{theorem}\label{thm:isAlgEnvSeq_arrayMeasure} + \uses{def:bandit, def:arrayMeasure} + \leanok + \lean{Bandits.ArrayModel.isAlgEnvSeq_arrayMeasure} +TODO +\end{theorem} + +\begin{proof} + +\end{proof} + + \section{Alternative models: rewards indexed by time or pull count}\label{sec:alt_model} The description of the bandit model above considers that at time $t$, a reward $R_t$ is generated, depending on the arm $A_t$ pulled at that time. @@ -51,7 +99,7 @@ \section{Alternative models: rewards indexed by time or pull count}\label{sec:al \begin{lemma}\label{lem:measurable_comap_indicator_stepsUntil_eq} \uses{def:stepsUntil} \leanok - \lean{Bandits.measurable_comap_indicator_stepsUntil_eq} + \lean{Learning.measurable_comap_indicator_stepsUntil_eq} The function $\mathbb{I}\{T_{n,a} = t\} : \Omega \to \{0, 1\}$ is measurable with respect to the sigma-algebra generated by $(H_{t-1}, A_t)$. \end{lemma} @@ -61,9 +109,9 @@ \section{Alternative models: rewards indexed by time or pull count}\label{sec:al \begin{lemma}\label{lem:condIndepFun_reward_stepsUntil_arm} - \uses{def:stepsUntil, def:actionReward, def:Bandit.measure} + \uses{def:stepsUntil, def:IT.actionReward, def:Bandit.measure} \leanok - \lean{Bandits.condIndepFun_reward_stepsUntil_arm} + \lean{Bandits.condIndepFun_reward_stepsUntil_action} For $t > 0$, $R_t \ind \mathbb{I}\{T_{n, a} = t\} \mid A_t$. \end{lemma} @@ -205,8 +253,6 @@ \section{Alternative models: rewards indexed by time or pull count}\label{sec:al \begin{lemma}\label{lem:iIndepFun_rewardByCount} \uses{def:rewardByCount} - \leanok - \lean{Bandits.iIndepFun_rewardByCount'} The rewards $(Y_{n,a})_{n \in \mathbb{N}}$ are independent. \end{lemma} @@ -235,12 +281,10 @@ \section{Alternative models: rewards indexed by time or pull count}\label{sec:al \begin{lemma}\label{lem:identDistrib_rewardByCount_stream} \uses{def:rewardByCount} - \leanok - \lean{Bandits.identDistrib_rewardByCount_stream} The random sequences $(Y_{n+1,a})_{n \in \mathbb{N}}$ and $(Z_{n,a})_{n \in \mathbb{N}}$ are identically distributed. \end{lemma} -\begin{proof}\leanok +\begin{proof} \uses{lem:hasLaw_rewardByCount, lem:iIndepFun_rewardByCount} \end{proof} @@ -248,12 +292,10 @@ \section{Alternative models: rewards indexed by time or pull count}\label{sec:al \begin{lemma}\label{lem:identDistrib_sum_Icc_rewardByCount} \uses{def:rewardByCount} - \leanok - \lean{Bandits.identDistrib_sum_Icc_rewardByCount} The random variables $\sum_{i=1}^n Y_{i,a}$ and $\sum_{i=0}^{n-1} Z_{i,a}$ are identically distributed. \end{lemma} -\begin{proof}\leanok +\begin{proof} \uses{lem:identDistrib_rewardByCount_stream} Immediate consequence of Lemma~\ref{lem:identDistrib_rewardByCount_stream}. \end{proof} @@ -270,7 +312,7 @@ \section{Regret and other bandit quantities} \begin{definition}[Regret]\label{def:regret} - \uses{def:armMean, def:actionReward} + \uses{def:armMean, def:IT.actionReward} \leanok \lean{Bandits.regret} The regret $R_T$ of a sequence of arms $A_0, \ldots, A_{T-1}$ after $T$ pulls is the difference between the cumulative reward of always playing the best arm and the cumulative reward of the sequence: diff --git a/blueprint/src/chapters/ucb.tex b/blueprint/src/chapters/ucb.tex index 1cfb5177..54ffc48a 100644 --- a/blueprint/src/chapters/ucb.tex +++ b/blueprint/src/chapters/ucb.tex @@ -1,7 +1,7 @@ \section{UCB} \begin{definition}[UCB algorithm]\label{def:ucbAlgorithm} - \uses{def:actionReward, def:pullCount, def:empMean} + \uses{def:IT.actionReward, def:pullCount, def:empMean} \leanok \lean{Bandits.UCB.nextArm, Bandits.ucbAlgorithm} The UCB algorithm with parameter $c \in \mathbb{R}_+$ is defined as follows: diff --git a/blueprint/src/macros/print.tex b/blueprint/src/macros/print.tex index 668fec13..78708343 100644 --- a/blueprint/src/macros/print.tex +++ b/blueprint/src/macros/print.tex @@ -26,4 +26,4 @@ \NewDocumentCommand{\proves}{m} {\clist_map_inline:nn{#1}{\vphantom{\ref{##1}}}% \ignorespaces} -\ExplSyntaxOff \ No newline at end of file +\ExplSyntaxOff diff --git a/blueprint/src/print.tex b/blueprint/src/print.tex index bf1c8685..4445c547 100644 --- a/blueprint/src/print.tex +++ b/blueprint/src/print.tex @@ -25,7 +25,7 @@ \input{macros/print} \title{LeanBandits\\ \Large{A Lean package for bandit algorithms}} -\author{Rémy Degenne} +\author{Rémy Degenne, Paulo Rauber} \begin{document} \maketitle diff --git a/blueprint/src/web.tex b/blueprint/src/web.tex index d7590660..8f3ac506 100644 --- a/blueprint/src/web.tex +++ b/blueprint/src/web.tex @@ -21,7 +21,7 @@ \dochome{https://RemyDegenne.github.io/lean-bandits/docs} \title{LeanBandits} -\author{Rémy Degenne} +\author{Rémy Degenne, Paulo Rauber} \begin{document} \maketitle