diff --git a/LeanBandits.lean b/LeanBandits.lean index a5d8fef3..050443eb 100644 --- a/LeanBandits.lean +++ b/LeanBandits.lean @@ -6,6 +6,7 @@ import LeanBandits.ForMathlib.CondDistrib import LeanBandits.ForMathlib.KernelCompositionLemmas import LeanBandits.ForMathlib.KernelCompositionParallelComp import LeanBandits.ForMathlib.KernelSub +import LeanBandits.ForMathlib.SubGaussian import LeanBandits.ForMathlib.Traj import LeanBandits.Regret import LeanBandits.RewardByCountMeasure diff --git a/LeanBandits/Bandit.lean b/LeanBandits/Bandit.lean index 18c12767..e27198f3 100644 --- a/LeanBandits/Bandit.lean +++ b/LeanBandits/Bandit.lean @@ -74,6 +74,33 @@ lemma snd_measure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] end Bandit +section StreamMeasure + +lemma _root_.hasLaw_eval_infinitePi {ι : Type*} {X : ι → Type*} {mX : ∀ i, MeasurableSpace (X i)} + (μ : (i : ι) → Measure (X i)) [hμ : ∀ i, IsProbabilityMeasure (μ i)] (i : ι) : + HasLaw (Function.eval i) (μ i) (Measure.infinitePi μ) where + aemeasurable := Measurable.aemeasurable (by fun_prop) + map_eq := by exact (measurePreserving_eval_infinitePi μ i).map_eq + +lemma hasLaw_eval_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : + HasLaw (fun h : ℕ → α → R ↦ h n) (Measure.infinitePi ν) (Bandit.streamMeasure ν) := + hasLaw_eval_infinitePi (fun _ ↦ Measure.infinitePi ν) n + +lemma hasLaw_eval_eval_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) (a : α) : + HasLaw (fun h : ℕ → α → R ↦ h n a) (ν a) (Bandit.streamMeasure ν) := + (hasLaw_eval_infinitePi ν a).comp (hasLaw_eval_streamMeasure ν n) + +lemma identDistrib_eval_eval_id_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) (a : α) : + IdentDistrib (fun h : ℕ → α → R ↦ h n a) id (Bandit.streamMeasure ν) (ν a) where + aemeasurable_fst := Measurable.aemeasurable (by fun_prop) + aemeasurable_snd := Measurable.aemeasurable (by fun_prop) + map_eq := by + rw [← (hasLaw_eval_eval_streamMeasure ν n a).map_eq, + Measure.map_map (by fun_prop) (by fun_prop)] + simp + +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 diff --git a/LeanBandits/ETC.lean b/LeanBandits/ETC.lean index e7f950ae..46f4b5b7 100644 --- a/LeanBandits/ETC.lean +++ b/LeanBandits/ETC.lean @@ -5,7 +5,9 @@ Authors: Rémy Degenne -/ import Mathlib.Probability.Moments.SubGaussian import LeanBandits.AlgorithmBuilding +import LeanBandits.ForMathlib.SubGaussian import LeanBandits.Regret +import LeanBandits.RewardByCountMeasure /-! # The Explore-Then-Commit Algorithm @@ -26,7 +28,7 @@ def ETC.nextArm (hK : 0 < K) (m n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K := ⟨(n + 1) % K, Nat.mod_lt _ hK⟩ -- for `n = 0` we have pulled arm 0 already, and we pull arm 1 else if hn_eq : n = K * m - 1 then measurableArgmax (empMean' n) h - else (h ⟨n - 1, by simp⟩).1 + else (h ⟨n, by simp⟩).1 @[fun_prop] lemma ETC.measurable_nextArm (hK : 0 < K) (m n : ℕ) : Measurable (nextArm hK m n) := by @@ -46,51 +48,202 @@ namespace ETC variable {hK : 0 < K} {m : ℕ} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] -local notation "𝔓b" => Bandit.trajMeasure (etcAlgorithm hK m) ν +local notation "𝔓t" => Bandit.trajMeasure (etcAlgorithm hK m) ν local notation "𝔓" => Bandit.measure (etcAlgorithm hK m) ν -lemma arm_zero : arm 0 =ᵐ[𝔓b] fun _ ↦ ⟨0, hK⟩ := by +lemma arm_zero : arm 0 =ᵐ[𝔓t] fun _ ↦ ⟨0, hK⟩ := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK exact arm_zero_detAlgorithm lemma arm_ae_eq_etcNextArm (n : ℕ) : - arm (n + 1) =ᵐ[𝔓b] fun h ↦ nextArm hK m n (fun i ↦ h i) := by + arm (n + 1) =ᵐ[𝔓t] fun h ↦ nextArm hK m n (fun i ↦ h i) := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK exact arm_detAlgorithm_ae_eq n -lemma pullCount_mul (a : Fin K) : - pullCount a (K * m) =ᵐ[𝔓b] fun _ ↦ m := by - sorry +lemma arm_of_lt {n : ℕ} (hn : n < K * m) : + arm n =ᵐ[𝔓t] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := by + cases n with + | zero => exact arm_zero + | succ n => + filter_upwards [arm_ae_eq_etcNextArm n] with h hn_eq + rw [hn_eq, nextArm, dif_pos] + grind -lemma pullCount_of_ge (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : +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 + 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 + rw [hn_eq, nextArm, dif_neg (by simp), dif_pos rfl] + exact this ▸ rfl + +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 + rw [hn_eq, nextArm, dif_neg (by grind), dif_neg] + · rfl + · have : 0 < K * m := Nat.mul_pos hK hm.bot_lt + grind + +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 + 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 + | base => rfl + | succ n hmn h_ind => rw [h_ae n hmn, h_ind] + +lemma pullCount_mul (a : Fin K) : pullCount a (K * m) =ᵐ[𝔓t] 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) + 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) + _ = (∑ 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 := by + sorry + +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 + simp_rw [Filter.EventuallyEq, pullCount_add_one] + filter_upwards [arm_of_ge hm hn] with ω h_arm + congr + +lemma pullCount_of_ge (a : Fin K) (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : pullCount a n - =ᵐ[𝔓b] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by - sorry + =ᵐ[𝔓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 + simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_ae + have h_ae_Km : pullCount a (K * m) =ᵐ[𝔓t] fun _ ↦ m := pullCount_mul 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] + | succ n hmn h_ind => + rw [h_ae n hmn, h_ind, add_assoc, ← add_one_mul] + congr + grind + +lemma pullCount_add_one_eq_pullCount' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} : + pullCount a (n + 1) h = pullCount' n (fun i ↦ h i) a := by + rw [pullCount_eq_sum, pullCount'_eq_sum] + unfold arm + rw [Finset.sum_coe_sort (f := fun s ↦ if (h s).1 = a then 1 else 0) (Iic n)] + congr with m + simp only [mem_range, mem_Iic] + grind -lemma prob_arm_mul_eq_le (a : Fin K) : - (𝔓b).real {ω | arm (K * m) ω = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by +lemma pullCount_eq_pullCount' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} (hn : n ≠ 0) : + pullCount a n h = pullCount' (n - 1) (fun i ↦ h i) a := by + cases n with + | zero => exact absurd rfl hn + | succ n => + rw [pullCount_add_one_eq_pullCount'] + have : n + 1 - 1 = n := by simp + exact this ▸ rfl + +lemma sumRewards_add_one_eq_sumRewards' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} : + sumRewards a (n + 1) h = sumRewards' n (fun i ↦ h i) a := by + unfold sumRewards sumRewards' arm reward + rw [Finset.sum_coe_sort (f := fun s ↦ if (h s).1 = a then (h s).2 else 0) (Iic n)] + congr with m + simp only [mem_range, mem_Iic] + grind + +lemma sumRewards_eq_sumRewards' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} (hn : n ≠ 0) : + sumRewards a n h = sumRewards' (n - 1) (fun i ↦ h i) a := by + cases n with + | zero => exact absurd rfl hn + | succ n => + rw [sumRewards_add_one_eq_sumRewards'] + have : n + 1 - 1 = n := by simp + exact this ▸ rfl + +lemma empMean_add_one_eq_empMean' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} : + empMean a (n + 1) h = empMean' n (fun i ↦ h i) a := by + unfold empMean empMean' + rw [sumRewards_add_one_eq_sumRewards', pullCount_add_one_eq_pullCount'] + +lemma empMean_eq_empMean' {a : Fin K} {n : ℕ} {h : ℕ → Fin K × ℝ} (hn : n ≠ 0) : + empMean a n h = empMean' (n - 1) (fun i ↦ h i) a := by + unfold empMean empMean' + rw [sumRewards_eq_sumRewards' hn, pullCount_eq_pullCount' hn] + +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 ν) + rw [← h_arm, h_eq] at h_max + rw [sumRewards_eq_pullCount_mul_empMean, sumRewards_eq_pullCount_mul_empMean, ha, h_best] + · gcongr + have : 0 < K * m := Nat.mul_pos hK hm.bot_lt + rwa [empMean_eq_empMean' this.ne', empMean_eq_empMean' this.ne'] + · simp [ha, hm] + · simp [h_best, hm] + +lemma ae_eq_set_iff {α : Type*} {mα : MeasurableSpace α} {μ : Measure α} {s t : Set α} : + s =ᵐ[μ] t ↔ ∀ᵐ a ∂μ, a ∈ s ↔ a ∈ t := by + rw [Filter.EventuallyEq] + simp only [eq_iff_iff] + congr! + +lemma prob_arm_mul_eq_le (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 + 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 + 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 {ω | arm (K * m) ω.1 = a} ≤ Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) by - suffices (𝔓b).real {ω | arm (K * m) ω = a} = (𝔓).real {ω | arm (K * m) ω.1 = a} by - rwa [this] - calc (𝔓b).real {ω | arm (K * m) ω = a} - _ = ((𝔓).fst).real {ω | arm (K * m) ω = a} := by simp - _ = (𝔓).real {ω | arm (K * m) ω.1 = a} := by + 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_singleton _).preimage (by fun_prop) - calc (𝔓).real {ω | arm (K * m) ω.1 = a} - _ ≤ (𝔓).real {ω | ∑ s ∈ range (K * m), (if (arm s ω.1) = bestArm ν then (reward s ω.1) else 0) - ≤ ∑ s ∈ range (K * m), if (arm s ω.1) = a then (reward s ω.1) else 0} := by - sorry + · 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 ω.1 ω.2 ≤ ∑ s ∈ Icc 1 (pullCount a (K * m) ω.1), rewardByCount a s ω.1 ω.2} := by - sorry + congr with ω + congr! 1 <;> rw [sum_rewardByCount_eq_sumRewards] _ = (𝔓).real {ω | ∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω.1 ω.2 ≤ ∑ s ∈ Icc 1 m, rewardByCount a s ω.1 ω.2} := by - sorry + 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] + sorry _ = (𝔓).real {ω | ∑ s ∈ range m, ω.2 s (bestArm ν) ≤ ∑ s ∈ range m, ω.2 s a} := by sorry _ = (𝔓).real {ω | m * gap ν a @@ -99,30 +252,45 @@ lemma prob_arm_mul_eq_le (a : Fin K) : simp only [gap_eq_bestArm_sub, id_eq, sum_sub_distrib, sum_const, card_range, nsmul_eq_mul] ring_nf simp + _ = (Bandit.streamMeasure ν).real {ω | m * gap ν a + ≤ ∑ s ∈ range m, ((ω s a - (ν a)[id]) - (ω s (bestArm ν) - (ν (bestArm ν))[id]))} := by + have : Bandit.streamMeasure ν = (𝔓).map Prod.snd := by rw [← Measure.snd, Bandit.snd_measure] + rw [this, measureReal_def, measureReal_def, Measure.map_apply (by fun_prop)] + · rfl + · exact measurableSet_le (by fun_prop) (by fun_prop) _ ≤ Real.exp (-↑m * gap ν a ^ 2 / 4) := by refine (HasSubgaussianMGF.measure_sum_range_ge_le_of_iIndepFun (c := 2) (ε := m * gap ν a) ?_ ?_ ?_).trans_eq ?_ - · suffices iIndepFun (fun s ω ↦ ω s a - (ν a)[id] - (ω s (bestArm ν) - (ν (bestArm ν))[id])) - (Bandit.streamMeasure ν) by + · suffices iIndepFun (fun s ω ↦ ω s a - ω s (bestArm ν)) (Bandit.streamMeasure ν) by sorry sorry · intro i him - sorry + rw [← one_add_one_eq_two] + refine HasSubgaussianMGF.sub_of_indepFun ?_ ?_ ?_ + · refine (hν a).congr_identDistrib ?_ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ + · refine (hν (bestArm ν)).congr_identDistrib ?_ + exact (identDistrib_eval_eval_id_streamMeasure _ _ _).symm.sub_const _ + · suffices IndepFun (fun ω ↦ ω i a) (fun ω ↦ ω i (bestArm ν)) (Bandit.streamMeasure ν) by + exact this.comp (φ := fun x ↦ x - (ν a)[id]) (ψ := fun x ↦ x - (ν (bestArm ν))[id]) + (by fun_prop) (by fun_prop) + sorry · have : 0 ≤ gap ν a := gap_nonneg positivity · congr 1 field_simp simp_rw [mul_assoc] simp only [NNReal.coe_ofNat, neg_inj, mul_eq_mul_left_iff, ne_eq, OfNat.ofNat_ne_zero, - not_false_eq_true, pow_eq_zero_iff, Nat.cast_eq_zero] + not_false_eq_true, pow_eq_zero_iff] norm_num -lemma expectation_pullCount_le (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : - 𝔓b[fun ω ↦ (pullCount a n ω : ℝ)] +lemma expectation_pullCount_le (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 ω : ℝ)] ≤ m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by have : (fun ω ↦ (pullCount a n ω : ℝ)) - =ᵐ[𝔓b] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by - filter_upwards [pullCount_of_ge a hn] with ω h + =ᵐ[𝔓t] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by + filter_upwards [pullCount_of_ge 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 @@ -139,21 +307,26 @@ lemma expectation_pullCount_le (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : simp rw [integral_indicator_const, smul_eq_mul, mul_one] · rw [← neg_mul] - exact prob_arm_mul_eq_le a + exact prob_arm_mul_eq_le hν a hm · exact (measurableSet_singleton _).preimage (by fun_prop) -lemma regret_le (n : ℕ) (hn : K * m ≤ n) : - 𝔓b[regret ν n] ≤ ∑ a, gap ν a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4)) := by +lemma integrable_pullCount (a : Fin K) (n : ℕ) : Integrable (fun ω ↦ (pullCount a n ω : ℝ)) 𝔓t := 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 ω + +lemma regret_le (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 simp_rw [regret_eq_sum_pullCount_mul_gap] rw [integral_finset_sum] - swap - · refine fun i _ ↦ Integrable.mul_const ?_ _ - sorry + swap; · exact fun i _ ↦ (integrable_pullCount i n).mul_const _ gcongr with a rw [mul_comm (gap _ _), integral_mul_const] gcongr · exact gap_nonneg - · exact expectation_pullCount_le a hn + · exact expectation_pullCount_le hν a hm hn end ETC diff --git a/LeanBandits/ForMathlib/SubGaussian.lean b/LeanBandits/ForMathlib/SubGaussian.lean new file mode 100644 index 00000000..ca7e09e3 --- /dev/null +++ b/LeanBandits/ForMathlib/SubGaussian.lean @@ -0,0 +1,91 @@ +/- +Copyright (c) 2025 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +import Mathlib.Probability.Moments.SubGaussian + +open MeasureTheory +open scoped ENNReal NNReal + +namespace ProbabilityTheory + +theorem mgf_const_mul {Ω : Type*} {m : MeasurableSpace Ω} {X : Ω → ℝ} {μ : Measure Ω} + {t : ℝ} (α : ℝ) : mgf (fun ω ↦ α * X ω) μ t = mgf X μ (α * t) := by + rw [← mgf_smul_left] + rfl + +namespace Kernel.HasSubgaussianMGF + +variable {Ω Ω' : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} + {ν : Measure Ω'} {κ : Kernel Ω' Ω} {X : Ω → ℝ} {c : ℝ≥0} + +lemma id_map_iff (hX : Measurable X) : + HasSubgaussianMGF X c κ ν ↔ HasSubgaussianMGF id c (κ.map X) ν := by + refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ + · constructor + · intro t + rw [← Kernel.deterministic_comp_eq_map hX, ← Measure.comp_assoc, + Measure.deterministic_comp_eq_map] + rw [integrable_map_measure (by fun_prop) hX.aemeasurable] + exact h.integrable_exp_mul t + · simp_rw [Kernel.map_apply _ hX, mgf_id_map hX.aemeasurable] + exact h.mgf_le + · have : X = id ∘ X := rfl + rw [this] + exact .of_map hX h + +protected lemma const_mul (h : HasSubgaussianMGF X c κ ν) (r : ℝ) : + HasSubgaussianMGF (fun ω ↦ r * X ω) (⟨r ^ 2, sq_nonneg r⟩ * c) κ ν where + integrable_exp_mul t := by + simp_rw [← mul_assoc] + exact h.integrable_exp_mul (t * r) + mgf_le := by + filter_upwards [h.mgf_le] with ω hω t + specialize hω (t * r) + rw [mgf_const_mul, mul_comm] + refine hω.trans_eq ?_ + congr 1 + simp only [NNReal.coe_mul, NNReal.coe_mk] + ring + +end Kernel.HasSubgaussianMGF + +namespace HasSubgaussianMGF + +variable {Ω : Type*} {m mΩ : MeasurableSpace Ω} {μ : Measure Ω} {X : Ω → ℝ} {c : ℝ≥0} + +lemma id_map_iff (hX : AEMeasurable X μ) : + HasSubgaussianMGF X c μ ↔ HasSubgaussianMGF id c (μ.map X) := by + refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ + · constructor + · intro t + rw [integrable_map_measure (by fun_prop) hX] + exact h.integrable_exp_mul t + · intro t + rw [mgf_id_map hX] + exact h.mgf_le t + · have : X = id ∘ X := rfl + rw [this] + exact .of_map hX h + +lemma congr_identDistrib {Ω' : Type*} {mΩ' : MeasurableSpace Ω'} {μ' : Measure Ω'} + {Y : Ω' → ℝ} (hX : HasSubgaussianMGF X c μ) (hXY : IdentDistrib X Y μ μ') : + HasSubgaussianMGF Y c μ' := by + rw [id_map_iff hXY.aemeasurable_fst] at hX + rwa [id_map_iff hXY.aemeasurable_snd, ← hXY.map_eq] + +protected lemma const_mul (h : HasSubgaussianMGF X c μ) (r : ℝ) : + HasSubgaussianMGF (fun ω ↦ r * X ω) (⟨r ^ 2, sq_nonneg r⟩ * c) μ := by + rw [HasSubgaussianMGF_iff_kernel] at h ⊢ + exact Kernel.HasSubgaussianMGF.const_mul h r + +lemma sub_of_indepFun {Y : Ω → ℝ} {cX cY : ℝ≥0} (hX : HasSubgaussianMGF X cX μ) + (hY : HasSubgaussianMGF Y cY μ) (hindep : IndepFun X Y μ) : + HasSubgaussianMGF (fun ω ↦ X ω - Y ω) (cX + cY) μ := by + simp_rw [sub_eq_add_neg] + exact hX.add_of_indepFun hY.neg hindep.neg_right + +end HasSubgaussianMGF + +end ProbabilityTheory diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index db6de8fc..ac4b8cb7 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -61,9 +61,18 @@ lemma pullCount_eq_pullCount_add_one (t : ℕ) (h : ℕ → α × ℝ) : lemma pullCount_eq_pullCount (ha : arm t h ≠ a) : pullCount a (t + 1) h = pullCount a t h := by simp [pullCount, range_succ, filter_insert, ha] +lemma pullCount_add_one : + pullCount a (t + 1) h = pullCount a t h + if arm t h = a then 1 else 0 := by + split_ifs with h + · rw [← h, pullCount_eq_pullCount_add_one] + · rw [pullCount_eq_pullCount h, add_zero] + lemma pullCount_eq_sum (a : α) (t : ℕ) (h : ℕ → α × ℝ) : pullCount a t h = ∑ s ∈ range t, if arm s h = a then 1 else 0 := by simp [pullCount] +lemma pullCount_le (a : α) (t : ℕ) (h : ℕ → α × ℝ) : pullCount a t h ≤ t := + (card_filter_le _ _).trans_eq (by simp) + /-- Number of steps until arm `a` was pulled exactly `m` times. -/ noncomputable def stepsUntil (a : α) (m : ℕ) (h : ℕ → α × ℝ) : ℕ∞ := sInf ((↑) '' {s | pullCount a (s + 1) h = m}) @@ -190,6 +199,23 @@ lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) pullCount a (stepsUntil a m h).toNat h = m - 1 := by sorry +section SumRewards + +/-- Sum of rewards obtained when pulling arm `a` up to time `t` (exclusive). -/ +def sumRewards (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := + ∑ s ∈ range t, if (arm s h) = a then (reward s h) else 0 + +/-- Empirical mean reward obtained when pulling arm `a` up to time `t` (exclusive). -/ +noncomputable +def empMean (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := sumRewards a t h / pullCount a t h + +lemma sumRewards_eq_pullCount_mul_empMean (h_pull : pullCount a t h ≠ 0) : + sumRewards a t h = pullCount a t h * empMean a t h := by unfold empMean; field_simp + +end SumRewards + +section RewardByCount + /-- Reward obtained when pulling arm `a` for the `m`-th time. -/ noncomputable def rewardByCount (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : ℝ := @@ -215,17 +241,18 @@ lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (h : ℕ → α × ℝ rewardByCount (arm t h) (pullCount (arm t h) t h + 1) h z = reward t h := by rw [rewardByCount, ← pullCount_eq_pullCount_add_one, stepsUntil_pullCount_eq] -lemma sum_rewardByCount_eq_sum_reward +lemma sum_rewardByCount_eq_sumRewards (a : α) (t : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : - ∑ m ∈ Icc 1 (pullCount a t h), rewardByCount a m h z = - ∑ s ∈ range t, if (arm s h) = a then (reward s h) else 0 := by + ∑ m ∈ Icc 1 (pullCount a t h), rewardByCount a m h z = sumRewards a t h := by induction' t with t ht - · simp [pullCount] + · simp [pullCount, sumRewards] by_cases hta : arm t h = a · rw [← hta] at ht ⊢ rw [pullCount_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht] + unfold sumRewards rw [sum_range_succ, if_pos rfl, rewardByCount_pullCount_add_one_eq_reward] - · rwa [pullCount_eq_pullCount hta, sum_range_succ, if_neg hta, add_zero] + · unfold sumRewards + rwa [pullCount_eq_pullCount hta, sum_range_succ, if_neg hta, add_zero] lemma sum_pullCount_mul [Fintype α] (h : ℕ → α × ℝ) (f : α → ℝ) (t : ℕ) : ∑ a, pullCount a t h * f a = ∑ s ∈ range t, f (arm s h) := by @@ -246,6 +273,8 @@ lemma regret_eq_sum_pullCount_mul_gap [Fintype α] : simp_rw [sum_pullCount_mul, regret, gap, sum_sub_distrib] simp +end RewardByCount + section BestArm variable [Fintype α] [Nonempty α] diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index 8acb3b6c..2ca511ba 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -24,6 +24,14 @@ lemma measurable_pullCount (a : α) (t : ℕ) : Measurable (fun h ↦ pullCount exact (measurableSet_singleton _).preimage (by fun_prop) fun_prop +@[fun_prop] +lemma measurable_sumRewards (a : α) (t : ℕ) : Measurable (sumRewards a t) := by + unfold sumRewards + have h_meas s : Measurable (fun h : ℕ → α × ℝ ↦ if arm s h = a then reward 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_stepsUntil (a : α) (m : ℕ) : Measurable (fun h ↦ stepsUntil a m h) := by classical diff --git a/blueprint/lean_decls b/blueprint/lean_decls index fa693245..d545e1de 100644 --- a/blueprint/lean_decls +++ b/blueprint/lean_decls @@ -23,7 +23,7 @@ Bandits.iIndepFun_rewardByCount Bandits.stepsUntil_pullCount_le Bandits.stepsUntil_pullCount_eq Bandits.rewardByCount_pullCount_add_one_eq_reward -Bandits.sum_rewardByCount_eq_sum_reward +Bandits.sum_rewardByCount_eq_sumRewards Bandits.regret Bandits.gap Bandits.regret_eq_sum_pullCount_mul_gap diff --git a/blueprint/src/chapters/bandit.tex b/blueprint/src/chapters/bandit.tex index abb30068..f767409a 100644 --- a/blueprint/src/chapters/bandit.tex +++ b/blueprint/src/chapters/bandit.tex @@ -267,7 +267,7 @@ \section{Alternative model}\label{sec:alt_model} \begin{lemma}\label{lem:sum_rewardByCount} \uses{def:rewardByCount,def:pullCount} \leanok - \lean{Bandits.sum_rewardByCount_eq_sum_reward} + \lean{Bandits.sum_rewardByCount_eq_sumRewards} \begin{align*} \sum_{n=1}^{N_{t, a}} Y_{n, a} = \sum_{s=0}^{t-1} \mathbb{I}\{A_s = a\} X_s \: .