From ff83243363729a4cd30d2fbb6ba80cd9e8466528 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sun, 15 Feb 2026 20:59:45 +0100 Subject: [PATCH] pick a few results --- LeanBandits/Bandit/SumRewards.lean | 57 ++++--------------- .../SequentialLearning/FiniteActions.lean | 44 +++++++------- 2 files changed, 34 insertions(+), 67 deletions(-) diff --git a/LeanBandits/Bandit/SumRewards.lean b/LeanBandits/Bandit/SumRewards.lean index 81ec80a3..64d071f9 100644 --- a/LeanBandits/Bandit/SumRewards.lean +++ b/LeanBandits/Bandit/SumRewards.lean @@ -17,6 +17,15 @@ namespace Bandits namespace ArrayModel +lemma sum_Icc_one_eq_sum_range {m : ℕ} {f : ℕ → ℝ} : + ∑ i ∈ Icc 1 m, f (i - 1) = ∑ i ∈ range m, f i := by + have h : Icc 1 m = (range m).image (· + 1) := by + ext x; simp only [mem_Icc, mem_image, mem_range]; constructor + · intro ⟨h1, h2⟩; exact ⟨x - 1, by omega, by omega⟩ + · rintro ⟨a, ha, rfl⟩; omega + rw [h, Finset.sum_image (fun _ _ _ _ h => by omega)] + simp + variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [Countable α] [StandardBorelSpace α] [Nonempty α] {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] @@ -88,19 +97,7 @@ lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount (n : ℕ) : · infer_instance ext a : 1 congr 1 - let e : Icc 1 (pullCount A a n ω) ≃ range (pullCount A a n ω) := - { toFun x := ⟨x - 1, by have h := x.2; simp only [mem_Icc] at h; simp; grind⟩ - invFun x := ⟨x + 1, by - have h := x.2 - simp only [mem_Icc, le_add_iff_nonneg_left, zero_le, true_and, ge_iff_le] - simp only [mem_range] at h - grind⟩ - left_inv x := by have h := x.2; simp only [mem_Icc] at h; grind - right_inv x := by have h := x.2; grind } - rw [← sum_coe_sort (Icc 1 (pullCount A a n ω)), ← sum_coe_sort (range (pullCount A a n ω)), - sum_equiv e] - · simp - · simp [e] + exact sum_Icc_one_eq_sum_range.symm lemma identDistrib_pullCount_prod_sumRewards (n : ℕ) : IdentDistrib (fun ω a ↦ (pullCount A a n ω, sumRewards A R a n ω)) @@ -363,38 +360,8 @@ lemma _root_.Learning.IsAlgEnvSeq.law_pullCount_sumRewards_unique (h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') : P.map (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω)) = - P'.map (fun ω ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) := by - have hA := h1.measurable_A - have hR := h1.measurable_R - have hA2 := h2.measurable_A - have hR2 := h2.measurable_R - have h_unique := isAlgEnvSeq_unique h1 h2 - let f := fun p : ℕ → α × ℝ ↦ (∑ i ∈ range n, if (p i).1 = a then 1 else 0, - ∑ i ∈ range n, if (p i).1 = a then (p i).2 else 0) - have hf : Measurable f := by - refine Measurable.prod ?_ ?_ - · simp only [f] - refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact (measurableSet_singleton _).preimage (by fun_prop) - · simp only [f] - refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact (measurableSet_singleton _).preimage (by fun_prop) - have h_eq_comp : (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω)) - = f ∘ (fun ω n ↦ (A n ω, R n ω)) := by - ext ω : 1 - rw [pullCount_eq_comp (R := R), sumRewards_eq_comp] - grind - have h_eq_comp2 : (fun ω ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) - = f ∘ (fun ω n ↦ (A₂ n ω, R₂ n ω)) := by - ext ω : 1 - rw [pullCount_eq_comp (R := R₂), sumRewards_eq_comp] - grind - rw [h_eq_comp, h_eq_comp2, ← Measure.map_map hf, h_unique, Measure.map_map hf, - ← h_eq_comp2] - · rw [measurable_pi_iff] - exact fun n ↦ Measurable.prodMk (hA2 n) (hR2 n) - · rw [measurable_pi_iff] - exact fun n ↦ Measurable.prodMk (hA n) (hR n) + P'.map (fun ω ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) := + ((h1.law_pullCount_sumRewards_unique' h2 (n := n)).comp (u := fun f ↦ f a) (by fun_prop)).map_eq -- this is what we will use for UCB lemma prob_pullCount_prod_sumRewards_mem_le [Countable α] diff --git a/LeanBandits/SequentialLearning/FiniteActions.lean b/LeanBandits/SequentialLearning/FiniteActions.lean index ef3a113f..b46415be 100644 --- a/LeanBandits/SequentialLearning/FiniteActions.lean +++ b/LeanBandits/SequentialLearning/FiniteActions.lean @@ -95,10 +95,7 @@ lemma pullCount_eq_pullCount' {n : ℕ} {ω : Ω} (hn : n ≠ 0) : pullCount A a n ω = pullCount' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by cases n with | zero => exact absurd rfl hn - | succ n => - rw [pullCount_add_one_eq_pullCount' (R' := R')] - have : n + 1 - 1 = n := by simp - exact this ▸ rfl + | succ n => simp [pullCount_add_one_eq_pullCount' (R' := R')] lemma pullCount'_mono {n m : ℕ} (hnm : n ≤ m) : pullCount' n (fun i ↦ (A i ω, R' i ω)) a ≤ pullCount' m (fun i ↦ (A i ω, R' i ω)) a := by @@ -269,10 +266,7 @@ lemma stepsUntil_zero_of_ne (hka : A 0 ω ≠ a) : stepsUntil A a 0 ω = 0 := by lemma stepsUntil_zero_of_eq (hka : A 0 ω = a) : stepsUntil A a 0 ω = ⊤ := by rw [stepsUntil_eq_top_iff] - suffices 0 < pullCount A a 1 ω by - intro n hn - refine lt_irrefl 0 ?_ - exact this.trans_le (le_trans (monotone_pullCount _ _ (by omega)) hn.le) + suffices 0 < pullCount A a 1 ω from fun _ ↦ (this.trans_le (monotone_pullCount _ _ (by lia))).ne' rw [← hka, ← zero_add 1, pullCount_action_eq_pullCount_add_one] simp @@ -362,11 +356,7 @@ lemma stepsUntil_eq_zero_iff : simp only [stepsUntil_eq_dite, h_exists, ↓reduceDIte, Nat.cast_eq_zero, Nat.find_eq_zero, zero_add] at h' rw [pullCount_one] at h' - by_cases hka : A 0 ω = a - · simp only [hka, ↓reduceIte] at h' - simp [h'.symm, hka] - · simp only [hka, ↓reduceIte] at h' - simp [h'.symm, hka] + by_cases hka : A 0 ω = a <;> simp_all · cases h' with | inl h => rw [h.1, stepsUntil_zero_of_ne h.2] @@ -641,11 +631,6 @@ theorem isStoppingTime_stepsUntil_filtrationAction [MeasurableSingletonClass α] · rw [IsAlgEnvSeq.filtrationAction_eq_comap _ hn] exact measurableSet_stepsUntil_eq hA hR' a m n --- /-- Sigma-algebra generated by the stopping time `stepsUntil a m`. -/ --- def stepsUntilMeasurableSpace [Nonempty R] [MeasurableSingletonClass α] (a : α) (m : ℕ) : --- MeasurableSpace (ℕ → α × R) := --- (isStoppingTime_stepsUntil_filtrationAction a m (mR := mR)).measurableSpace - end Measurability end StepsUntil @@ -772,6 +757,24 @@ lemma sumRewards_add_one {R' : ℕ → Ω → ℝ} : unfold sumRewards rw [sum_range_succ] +lemma sumRewards_eq_of_pullCount_eq {R' : ℕ → Ω → ℝ} {s t : ℕ} + (h_eq : pullCount A a s ω = pullCount A a t ω) : + sumRewards A R' a s ω = sumRewards A R' a t ω := by + wlog hst : s ≤ t + · have hts : t ≤ s := by lia + exact (this h_eq.symm hts).symm + induction t, hst using Nat.le_induction with + | base => rfl + | succ t hst' ih => + have h_mono' : pullCount A a t ω ≤ pullCount A a (t + 1) ω := pullCount_mono a (Nat.le_succ t) ω + have h_eq_t : pullCount A a s ω = pullCount A a t ω := + le_antisymm (pullCount_mono a hst' ω) (h_eq ▸ h_mono') + have hne : A t ω ≠ a := by + intro ha + have h1 := ha ▸ pullCount_action_eq_pullCount_add_one (A := A) t ω + lia + rw [sumRewards_add_one, if_neg hne, add_zero, ih h_eq_t] + lemma sumRewards_eq_pullCount_mul_empMean {R' : ℕ → Ω → ℝ} {ω : Ω} (h_pull : pullCount A a t ω ≠ 0) : sumRewards A R' a t ω = pullCount A a t ω * empMean A R' a t ω := by unfold empMean; field_simp @@ -801,10 +804,7 @@ lemma sumRewards_eq_sumRewards' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (h sumRewards A R' a n ω = sumRewards' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by cases n with | zero => exact absurd rfl hn - | succ n => - rw [sumRewards_add_one_eq_sumRewards'] - have : n + 1 - 1 = n := by simp - exact this ▸ rfl + | succ n => simp [sumRewards_add_one_eq_sumRewards'] lemma empMean_add_one_eq_empMean' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} : empMean A R' a (n + 1) ω = empMean' n (fun i ↦ (A i ω, R' i ω)) a := by