diff --git a/LeanBandits/Algorithm.lean b/LeanBandits/Algorithm.lean index d886459d..49c5dc3c 100644 --- a/LeanBandits/Algorithm.lean +++ b/LeanBandits/Algorithm.lean @@ -44,21 +44,6 @@ structure Environment (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] wh instance (env : Environment α R) (n : ℕ) : IsMarkovKernel (env.feedback n) := env.h_feedback n instance (env : Environment α R) : IsMarkovKernel env.ν0 := env.hp0 -/-- A deterministic algorithm. -/ -noncomputable -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) - p0 := Measure.dirac action0 - -/-- A stationary environment, in which the distribution of the next reward depends only on the last -action. -/ -@[simps] -def stationaryEnv (ν : Kernel α R) [IsMarkovKernel ν] : Environment α R where - feedback _ := ν.prodMkLeft _ - ν0 := ν - /-- Kernel describing the distribution of the next action-reward pair given the history up to `n`. -/ noncomputable @@ -200,33 +185,83 @@ lemma condDistrib_reward_zero [StandardBorelSpace R] [Nonempty R] section DetAlgorithm +/-- A deterministic algorithm. -/ +@[simps] +noncomputable +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) + p0 := Measure.dirac action0 + variable {nextaction : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextaction n)} {action0 : α} {env : Environment α R} -lemma HasLaw_action_zero_detAlgorithm : - HasLaw (action 0) (Measure.dirac action0) - (trajMeasure (detAlgorithm nextaction h_next action0) env) where +local notation "𝔓" => trajMeasure (detAlgorithm nextaction h_next action0) env + +lemma HasLaw_action_zero_detAlgorithm : HasLaw (action 0) (Measure.dirac action0) 𝔓 where map_eq := (hasLaw_action_zero _ _).map_eq -lemma action_zero_detAlgorithm [MeasurableSingletonClass α] : - action 0 =ᵐ[trajMeasure (detAlgorithm nextaction h_next action0) env] fun _ ↦ action0 := by - have h_eq : ∀ᵐ x ∂((trajMeasure (detAlgorithm nextaction h_next action0) env).map (action 0)), x - = action0 := by +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] simp [detAlgorithm] exact ae_of_ae_map (by fun_prop) h_eq -lemma action_detAlgorithm_ae_eq (n : ℕ) : - action (n + 1) =ᵐ[trajMeasure (detAlgorithm nextaction h_next action0) env] - fun h ↦ nextaction n (fun i ↦ h i) := by +lemma action_detAlgorithm_ae_eq + [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] + (n : ℕ) : + action (n + 1) =ᵐ[𝔓] fun h ↦ nextaction n (fun i ↦ h i) := by + have h := condDistrib_action (detAlgorithm nextaction h_next action0) env n + simp only [detAlgorithm_policy] at h sorry -example [MeasurableSingletonClass α] : - ∀ᵐ h ∂(trajMeasure (detAlgorithm nextaction h_next action0) env), - action 0 h = action0 ∧ ∀ n, action (n + 1) h = nextaction n (fun i ↦ h i) := by +example [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] : + ∀ᵐ h ∂𝔓, action 0 h = action0 ∧ ∀ n, action (n + 1) h = nextaction n (fun i ↦ h i) := by rw [eventually_and, ae_all_iff] exact ⟨action_zero_detAlgorithm, action_detAlgorithm_ae_eq⟩ end DetAlgorithm +section stationaryEnv + +/-- A stationary environment, in which the distribution of the next reward depends only on the last +action. -/ +@[simps] +def stationaryEnv (ν : Kernel α R) [IsMarkovKernel ν] : Environment α R where + feedback _ := ν.prodMkLeft _ + ν0 := ν + +variable {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] + +local notation "𝔓" => trajMeasure alg (stationaryEnv ν) + +lemma condDistrib_reward_stationaryEnv [StandardBorelSpace α] [Nonempty α] + [StandardBorelSpace R] [Nonempty R] (n : ℕ) : + condDistrib (reward n) (action n) 𝔓 =ᵐ[(𝔓).map (action n)] ν := by + cases n with + | zero => + rw [condDistrib_ae_eq_iff_measure_eq_compProd₀ (by fun_prop) (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] + | succ n => + have h_eq := condDistrib_reward alg (stationaryEnv ν) n + rw [condDistrib_ae_eq_iff_measure_eq_compProd₀ (by fun_prop) (by fun_prop)] at h_eq ⊢ + have : (𝔓).map (action (n + 1)) = ((𝔓).map (fun x ↦ (hist n x, action (n + 1) x))).snd := by + rw [Measure.snd_map_prodMk (by fun_prop)] + simp only [stationaryEnv_feedback] at h_eq + rw [this, ← Measure.snd_prodAssoc_compProd_prodMkLeft, ← h_eq, + Measure.snd_map_prodMk (by fun_prop), Measure.map_map (by fun_prop) (by fun_prop)] + congr + +lemma condIndepFun_reward_hist_action [StandardBorelSpace α] [Nonempty α] + [StandardBorelSpace R] [Nonempty R] (n : ℕ) : + CondIndepFun (MeasurableSpace.comap (action (n + 1)) inferInstance) + (measurable_action _).comap_le (reward (n + 1)) (hist n) (𝔓) := + condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft + (by fun_prop) (by fun_prop) (by fun_prop) (condDistrib_reward alg (stationaryEnv ν) n) + +end stationaryEnv + end Learning diff --git a/LeanBandits/Bandit.lean b/LeanBandits/Bandit.lean index 06be7c03..18c12767 100644 --- a/LeanBandits/Bandit.lean +++ b/LeanBandits/Bandit.lean @@ -138,22 +138,8 @@ lemma condDistrib_reward' [StandardBorelSpace α] [Nonempty α] [StandardBorelSp 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)] ν := by - cases n with - | zero => - rw [condDistrib_ae_eq_iff_measure_eq_compProd₀ (by fun_prop) (by fun_prop)] - change (Bandit.trajMeasure alg ν).map (fun h ↦ h 0) - = (Bandit.trajMeasure alg ν).map (arm 0) ⊗ₘ ν - rw [(hasLaw_arm_zero alg ν).map_eq, (hasLaw_step_zero alg ν).map_eq] - | succ n => - have h_eq := condDistrib_reward' alg ν n - rw [condDistrib_ae_eq_iff_measure_eq_compProd₀ (by fun_prop) (by fun_prop)] at h_eq ⊢ - have : (Bandit.trajMeasure alg ν).map (arm (n + 1)) - = ((Bandit.trajMeasure alg ν).map (fun x ↦ (hist n x, arm (n + 1) x))).snd := by - rw [Measure.snd_map_prodMk (by fun_prop)] - rw [this, ← Measure.snd_prodAssoc_compProd_prodMkLeft, ← h_eq, - Measure.snd_map_prodMk (by fun_prop), Measure.map_map (by fun_prop) (by fun_prop)] - congr + =ᵐ[(Bandit.trajMeasure alg ν).map (arm n)] ν := + Learning.condDistrib_reward_stationaryEnv n lemma condDistrib_arm [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) : @@ -167,8 +153,7 @@ lemma condIndepFun_reward_hist_arm [StandardBorelSpace α] [Nonempty α] {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 ν) := - condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft - (by fun_prop) (by fun_prop) (by fun_prop) (condDistrib_reward' alg ν n) + Learning.condIndepFun_reward_hist_action n section DetAlgorithm diff --git a/LeanBandits/ETC.lean b/LeanBandits/ETC.lean index 52d1c850..e7f950ae 100644 --- a/LeanBandits/ETC.lean +++ b/LeanBandits/ETC.lean @@ -59,11 +59,11 @@ lemma arm_ae_eq_etcNextArm (n : ℕ) : exact arm_detAlgorithm_ae_eq n lemma pullCount_mul (a : Fin K) : - (fun ω ↦ pullCount (arm · ω) a (K * m)) =ᵐ[𝔓b] fun _ ↦ m := by + pullCount a (K * m) =ᵐ[𝔓b] fun _ ↦ m := by sorry lemma pullCount_of_ge (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : - (fun ω ↦ pullCount (arm · ω) a n) + pullCount a n =ᵐ[𝔓b] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by sorry @@ -84,9 +84,9 @@ lemma prob_arm_mul_eq_le (a : Fin K) : _ ≤ (𝔓).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 - _ = (𝔓).real {ω | ∑ s ∈ Icc 1 (pullCount (arm · ω.1) (bestArm ν) (K * m)), + _ = (𝔓).real {ω | ∑ s ∈ Icc 1 (pullCount (bestArm ν) (K * m) ω.1), rewardByCount (bestArm ν) s ω.1 ω.2 - ≤ ∑ s ∈ Icc 1 (pullCount (arm · ω.1) a (K * m)), rewardByCount a s ω.1 ω.2} := by + ≤ ∑ s ∈ Icc 1 (pullCount a (K * m) ω.1), rewardByCount a s ω.1 ω.2} := by sorry _ = (𝔓).real {ω | ∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω.1 ω.2 ≤ ∑ s ∈ Icc 1 m, rewardByCount a s ω.1 ω.2} := by @@ -118,9 +118,9 @@ lemma prob_arm_mul_eq_le (a : Fin K) : norm_num lemma expectation_pullCount_le (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : - 𝔓b[fun ω ↦ (pullCount (arm · ω) a n : ℝ)] + 𝔓b[fun ω ↦ (pullCount a n ω : ℝ)] ≤ m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by - have : (fun ω ↦ (pullCount (arm · ω) a n : ℝ)) + 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 simp only [h, Set.indicator_apply, Set.mem_setOf_eq, mul_ite, mul_one, mul_zero, Nat.cast_add, @@ -143,8 +143,7 @@ lemma expectation_pullCount_le (a : Fin K) {n : ℕ} (hn : K * m ≤ n) : · exact (measurableSet_singleton _).preimage (by fun_prop) lemma regret_le (n : ℕ) (hn : K * m ≤ n) : - 𝔓b[fun ω ↦ regret ν (arm · ω) n] - ≤ ∑ a, gap ν a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4)) := by + 𝔓b[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 diff --git a/LeanBandits/Regret.lean b/LeanBandits/Regret.lean index c28be981..db6de8fc 100644 --- a/LeanBandits/Regret.lean +++ b/LeanBandits/Regret.lean @@ -18,14 +18,14 @@ open scoped ENNReal NNReal namespace Bandits variable {α : Type*} [DecidableEq α] {mα : MeasurableSpace α} {ν : Kernel α ℝ} - {k : ℕ → α} {m n t : ℕ} {a : α} {h : ℕ → α × ℝ} + {h : ℕ → α × ℝ} {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 α ℝ) (k : ℕ → α) (t : ℕ) : ℝ := - t * (⨆ a, (ν a)[id]) - ∑ s ∈ range t, (ν (k s))[id] +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`. -/ noncomputable @@ -37,13 +37,13 @@ lemma gap_nonneg [Fintype α] : 0 ≤ gap ν a := by exact le_ciSup (f := fun i ↦ (ν i)[id]) (by simp) a /-- Number of times arm `a` was pulled up to time `t` (excluding `t`). -/ -noncomputable def pullCount [DecidableEq α] (k : ℕ → α) (a : α) (t : ℕ) : ℕ := - #(filter (fun s ↦ k s = a) (range t)) +noncomputable def pullCount [DecidableEq α] (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℕ := + #(filter (fun s ↦ arm s h = a) (range t)) @[simp] -lemma pullCount_zero (k : ℕ → α) (a : α) : pullCount k a 0 = 0 := by simp [pullCount] +lemma pullCount_zero (a : α) (h : ℕ → α × ℝ) : pullCount a 0 h = 0 := by simp [pullCount] -lemma pullCount_one : pullCount k a 1 = if k 0 = a then 1 else 0 := by +lemma pullCount_one : pullCount a 1 h = if arm 0 h = a then 1 else 0 := by simp only [pullCount, range_one] split_ifs with h · rw [card_eq_one] @@ -51,27 +51,27 @@ lemma pullCount_one : pullCount k a 1 = if k 0 = a then 1 else 0 := by · simp [h] open Classical in -lemma monotone_pullCount (k : ℕ → α) (a : α) : Monotone (pullCount k a) := +lemma monotone_pullCount (a : α) (h : ℕ → α × ℝ) : Monotone (pullCount a · h) := fun _ _ _ ↦ card_le_card (filter_subset_filter _ (by simpa)) -lemma pullCount_eq_pullCount_add_one (k : ℕ → α) (t : ℕ) : - pullCount k (k t) (t + 1) = pullCount k (k t) t + 1 := by +lemma pullCount_eq_pullCount_add_one (t : ℕ) (h : ℕ → α × ℝ) : + pullCount (arm t h) (t + 1) h = pullCount (arm t h) t h + 1 := by simp [pullCount, range_succ, filter_insert] -lemma pullCount_eq_pullCount (h : k t ≠ a) : pullCount k a (t + 1) = pullCount k a t := by - simp [pullCount, range_succ, filter_insert, 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_eq_sum (k : ℕ → α) (a : α) (t : ℕ) : - pullCount k a t = ∑ s ∈ range t, if k s = a then 1 else 0 := by simp [pullCount] +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] /-- Number of steps until arm `a` was pulled exactly `m` times. -/ noncomputable -def stepsUntil (k : ℕ → α) (a : α) (m : ℕ) : ℕ∞ := sInf ((↑) '' {s | pullCount k a (s + 1) = m}) +def stepsUntil (a : α) (m : ℕ) (h : ℕ → α × ℝ) : ℕ∞ := sInf ((↑) '' {s | pullCount a (s + 1) h = m}) -lemma stepsUntil_eq_top_iff : stepsUntil k a m = ⊤ ↔ ∀ s, pullCount k a (s + 1) ≠ m := by +lemma stepsUntil_eq_top_iff : stepsUntil a m h = ⊤ ↔ ∀ s, pullCount a (s + 1) h ≠ m := by simp [stepsUntil, sInf_eq_top] -lemma stepsUntil_zero_of_ne (hka : k 0 ≠ a) : stepsUntil k a 0 = 0 := by +lemma stepsUntil_zero_of_ne (hka : arm 0 h ≠ a) : stepsUntil a 0 h = 0 := by unfold stepsUntil simp_rw [← bot_eq_zero, sInf_eq_bot, bot_eq_zero] intro n hn @@ -80,75 +80,77 @@ lemma stepsUntil_zero_of_ne (hka : k 0 ≠ a) : stepsUntil k a 0 = 0 := by rw [← zero_add 1, pullCount_eq_pullCount hka] simp -lemma stepsUntil_zero_of_eq (hka : k 0 = a) : stepsUntil k a 0 = ⊤ := by +lemma stepsUntil_zero_of_eq (hka : arm 0 h = a) : stepsUntil a 0 h = ⊤ := by rw [stepsUntil_eq_top_iff] - suffices 0 < pullCount k a 1 by + suffices 0 < pullCount a 1 h 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_eq_pullCount_add_one] simp -lemma stepsUntil_eq_dite (k : ℕ → α) (a : α) (m : ℕ) [Decidable (∃ s, pullCount k a (s + 1) = m)] : - stepsUntil k a m = - if h : ∃ s, pullCount k a (s + 1) = m then (Nat.find h : ℕ∞) else ⊤ := by +lemma stepsUntil_eq_dite (a : α) (m : ℕ) (h : ℕ → α × ℝ) + [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 unfold stepsUntil - split_ifs with h + split_ifs with h' · refine le_antisymm ?_ ?_ · refine sInf_le ?_ - simpa using Nat.find_spec h + simpa using Nat.find_spec h' · simp only [le_sInf_iff, Set.mem_image, Set.mem_setOf_eq, forall_exists_index, and_imp, forall_apply_eq_imp_iff₂, Nat.cast_le, Nat.find_le_iff] exact fun n hn ↦ ⟨n, le_rfl, hn⟩ - · push_neg at h - suffices {s | pullCount k a (s + 1) = m} = ∅ by simp [this] + · push_neg at h' + suffices {s | pullCount a (s + 1) h = m} = ∅ by simp [this] ext s - simpa using (h s) + simpa using (h' s) -lemma stepsUntil_pullCount_le (k : ℕ → α) (a : α) (t : ℕ) : - stepsUntil k a (pullCount k a (t + 1)) ≤ t := by +lemma stepsUntil_pullCount_le (h : ℕ → α × ℝ) (a : α) (t : ℕ) : + stepsUntil a (pullCount a (t + 1) h) h ≤ t := by rw [stepsUntil] exact csInf_le (OrderBot.bddBelow _) ⟨t, rfl, rfl⟩ -lemma stepsUntil_pullCount_eq (k : ℕ → α) (t : ℕ) : - stepsUntil k (k t) (pullCount k (k t) (t + 1)) = t := by - apply le_antisymm (stepsUntil_pullCount_le k (k t) t) - suffices ∀ t', pullCount k (k t) (t' + 1) = pullCount k (k t) t + 1 → t ≤ t' by +lemma stepsUntil_pullCount_eq (h : ℕ → α × ℝ) (t : ℕ) : + stepsUntil (arm t h) (pullCount (arm t h) (t + 1) h) h = t := by + apply le_antisymm (stepsUntil_pullCount_le h (arm t h) t) + suffices ∀ t', pullCount (arm t h) (t' + 1) h = pullCount (arm t h) t h + 1 → t ≤ t' by simpa [stepsUntil, pullCount_eq_pullCount_add_one] - exact fun t' h ↦ Nat.le_of_lt_succ ((monotone_pullCount k (k t)).reflect_lt (h ▸ lt_add_one _)) + exact fun t' h' ↦ Nat.le_of_lt_succ ((monotone_pullCount (arm t h) h).reflect_lt + (h' ▸ lt_add_one _)) /-- If we pull arm `a` at time 0, the first time at which it is pulled once is 0. -/ -lemma stepsUntil_one_of_eq (hka : k 0 = a) : stepsUntil k a 1 = 0 := by +lemma stepsUntil_one_of_eq (hka : arm 0 h = a) : stepsUntil a 1 h = 0 := by classical - have h_pull : pullCount k a 1 = 1 := by simp [pullCount_one, hka] - have h_le := stepsUntil_pullCount_le k a 0 + have h_pull : pullCount a 1 h = 1 := by simp [pullCount_one, hka] + have h_le := stepsUntil_pullCount_le h a 0 simpa [h_pull] using h_le lemma stepsUntil_eq_zero_iff : - stepsUntil k a m = 0 ↔ (m = 0 ∧ k 0 ≠ a) ∨ (m = 1 ∧ k 0 = a) := by + stepsUntil a m h = 0 ↔ (m = 0 ∧ arm 0 h ≠ a) ∨ (m = 1 ∧ arm 0 h = a) := by classical - refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ - · have h_exists : ∃ s, pullCount k a (s + 1) = m := by + refine ⟨fun h' ↦ ?_, fun h' ↦ ?_⟩ + · have h_exists : ∃ s, pullCount a (s + 1) h = m := by by_contra! h_contra rw [← stepsUntil_eq_top_iff] at h_contra - simp [h_contra] at h + simp [h_contra] at 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 : k 0 = a - · simp only [hka, ↓reduceIte] at h - simp [h.symm, hka] - · simp only [hka, ↓reduceIte] at h - simp [h.symm, hka] - · cases h with + zero_add] at h' + rw [pullCount_one] at h' + by_cases hka : arm 0 h = a + · simp only [hka, ↓reduceIte] at h' + simp [h'.symm, hka] + · simp only [hka, ↓reduceIte] at h' + simp [h'.symm, hka] + · cases h' with | inl h => rw [h.1, stepsUntil_zero_of_ne h.2] | inr h => rw [h.1] exact stepsUntil_one_of_eq h.2 -lemma arm_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount (arm · h) a (s + 1) = m) : - arm (stepsUntil (arm · h) a m).toNat h = a := by +lemma arm_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1) h = m) : + arm (stepsUntil a m h).toNat h = a := by classical simp only [stepsUntil_eq_dite, h_exists, ↓reduceDIte, ENat.toNat_coe] have h_spec := Nat.find_spec h_exists @@ -168,55 +170,54 @@ lemma arm_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount (arm · h) a (s exact h_ne lemma arm_eq_of_stepsUntil_eq_coe {ω : ℕ → α × ℝ} (hm : m ≠ 0) - (h : stepsUntil (arm · ω) a m = n) : + (h : stepsUntil a m ω = n) : arm n ω = a := by - have : n = (stepsUntil (fun x ↦ arm x ω) a m).toNat := by simp [h] + have : n = (stepsUntil a m ω).toNat := by simp [h] rw [this, arm_stepsUntil hm] by_contra! h_contra rw [← stepsUntil_eq_top_iff] at h_contra simp [h_contra] at h -lemma stepsUntil_eq_congr {k' : ℕ → α} (h : ∀ i ≤ n, k i = k' i) : - stepsUntil k a m = n ↔ stepsUntil k' a m = n := by +lemma stepsUntil_eq_congr {h' : ℕ → α × ℝ} (h_eq : ∀ i ≤ n, arm i h = arm i h') : + stepsUntil a m h = n ↔ stepsUntil a m h' = n := by sorry -lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount k a (s + 1) = m) : - pullCount k a (stepsUntil k a m + 1).toNat = m := by +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 sorry -lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount k a (s + 1) = m) : - pullCount k a (stepsUntil k a m).toNat = m - 1 := by +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 sorry /-- Reward obtained when pulling arm `a` for the `m`-th time. -/ noncomputable def rewardByCount (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : ℝ := - match (stepsUntil (arm · h) a m) with + match (stepsUntil a m h) with | ⊤ => z m a | (n : ℕ) => reward n h lemma rewardByCount_eq_ite (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : rewardByCount a m h z = - if (stepsUntil (arm · h) a m) = ⊤ then z m a - else reward (stepsUntil (arm · h) a m).toNat h := by + if (stepsUntil a m h) = ⊤ then z m a else reward (stepsUntil a m h).toNat h := by unfold rewardByCount - cases stepsUntil (arm · h) a m <;> simp + cases stepsUntil a m h <;> simp lemma rewardByCount_of_stepsUntil_eq_top {ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)} - (h : stepsUntil (arm · ω.1) a m = ⊤) : + (h : stepsUntil a m ω.1 = ⊤) : rewardByCount a m ω.1 ω.2 = ω.2 m a := by simp [rewardByCount_eq_ite, h] lemma rewardByCount_of_stepsUntil_eq_coe {ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)} - (h : stepsUntil (arm · ω.1) a m = n) : + (h : stepsUntil a m ω.1 = n) : rewardByCount a m ω.1 ω.2 = reward n ω.1 := by simp [rewardByCount_eq_ite, h] lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : - rewardByCount (arm t h) (pullCount (arm · h) (arm t h) t + 1) h z = reward t h := by + 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 (a : α) (t : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : - ∑ m ∈ Icc 1 (pullCount (arm · h) a t), rewardByCount a m 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 induction' t with t ht · simp [pullCount] @@ -226,22 +227,22 @@ lemma sum_rewardByCount_eq_sum_reward 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] -lemma sum_pullCount_mul [Fintype α] (k : ℕ → α) (f : α → ℝ) (t : ℕ) : - ∑ a, pullCount k a t * f a = ∑ s ∈ range t, f (k s) := by +lemma sum_pullCount_mul [Fintype α] (h : ℕ → α × ℝ) (f : α → ℝ) (t : ℕ) : + ∑ a, pullCount a t h * f a = ∑ s ∈ range t, f (arm s h) := by unfold pullCount classical simp_rw [card_eq_sum_ones] push_cast simp_rw [sum_mul, one_mul] - exact sum_fiberwise' (range t) k f + exact sum_fiberwise' (range t) (arm · h) f -lemma sum_pullCount [Fintype α] : ∑ a, pullCount k a t = t := by - suffices ∑ a, pullCount k a t * (1 : ℝ) = t by norm_cast at this; simpa +lemma sum_pullCount [Fintype α] : ∑ a, pullCount a t h = t := by + suffices ∑ a, pullCount a t h * (1 : ℝ) = t by norm_cast at this; simpa rw [sum_pullCount_mul] simp lemma regret_eq_sum_pullCount_mul_gap [Fintype α] : - regret ν k t = ∑ a, pullCount k a t * gap ν a := by + regret ν t h = ∑ a, pullCount a t h * gap ν a := by simp_rw [sum_pullCount_mul, regret, gap, sum_sub_distrib] simp diff --git a/LeanBandits/RewardByCountMeasure.lean b/LeanBandits/RewardByCountMeasure.lean index a03533cd..8acb3b6c 100644 --- a/LeanBandits/RewardByCountMeasure.lean +++ b/LeanBandits/RewardByCountMeasure.lean @@ -17,32 +17,32 @@ namespace Bandits variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [MeasurableSingletonClass α] @[fun_prop] -lemma measurable_pullCount (a : α) (t : ℕ) : Measurable (fun k ↦ pullCount k a t) := by +lemma measurable_pullCount (a : α) (t : ℕ) : Measurable (fun h ↦ pullCount a t h) := by simp_rw [pullCount_eq_sum] - have h_meas s : Measurable (fun k : ℕ → α ↦ if k s = a then 1 else 0) := by + have h_meas s : Measurable (fun h : ℕ → α × ℝ ↦ if arm s h = 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_stepsUntil (a : α) (m : ℕ) : Measurable (fun k ↦ stepsUntil k a m) := by +lemma measurable_stepsUntil (a : α) (m : ℕ) : Measurable (fun h ↦ stepsUntil a m h) := by classical - have h_union : {k' | ∃ s, pullCount k' a (s + 1) = m} - = ⋃ s : ℕ, {k' | pullCount k' a (s + 1) = m} := by ext; simp - have h_meas_set : MeasurableSet {k' | ∃ s, pullCount k' a (s + 1) = m} := by + have h_union : {h' | ∃ s, pullCount a (s + 1) h' = m} + = ⋃ s : ℕ, {h' | pullCount a (s + 1) h' = m} := by ext; simp + have h_meas_set : MeasurableSet {h' | ∃ s, pullCount 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 k' a (s + 1) = m} + 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' : ℕ → α | ∃ s, pullCount k' a (s + 1) = m}) + refine Measurable.dite (s := {k' : ℕ → α × ℝ | ∃ s, pullCount a (s + 1) k' = m}) (f := fun x ↦ (Nat.find x.2 : ℕ∞)) (g := fun _ ↦ ⊤) ?_ (by fun_prop) h_meas_set refine Measurable.coe_nat_enat ?_ refine measurable_find _ fun k ↦ ?_ - suffices MeasurableSet {x : ℕ → α | pullCount x a (k + 1) = m} by + suffices MeasurableSet {x : ℕ → α × ℝ | pullCount a (k + 1) x = m} by have : Subtype.val '' - {x : {k' : ℕ → α | ∃ s, pullCount k' a (s + 1) = m} | pullCount x a (k + 1) = m} - = {x : ℕ → α | pullCount x a (k + 1) = m} := by + {x : {k' : ℕ → α × ℝ | ∃ s, pullCount a (s + 1) k' = m} | pullCount a (k + 1) x = m} + = {x : ℕ → α × ℝ | pullCount 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] @@ -52,13 +52,9 @@ lemma measurable_stepsUntil (a : α) (m : ℕ) : Measurable (fun k ↦ stepsUnti exact (measurableSet_singleton _).preimage (by fun_prop) exact (measurableSet_singleton _).preimage (by fun_prop) -lemma measurable_stepsUntil'' (a : α) (m : ℕ) : - Measurable (fun ω : (ℕ → α × ℝ) ↦ stepsUntil (arm · ω) a m) := - (measurable_stepsUntil a m).comp (by fun_prop) - lemma measurable_stepsUntil' (a : α) (m : ℕ) : - Measurable (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ stepsUntil (arm · ω.1) a m) := - (measurable_stepsUntil'' a m).comp measurable_fst + Measurable (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ stepsUntil a m ω.1) := + (measurable_stepsUntil a m).comp measurable_fst @[fun_prop] lemma measurable_rewardByCount (a : α) (m : ℕ) : @@ -68,9 +64,8 @@ lemma measurable_rewardByCount (a : α) (m : ℕ) : · exact (measurableSet_singleton _).preimage <| measurable_stepsUntil' a m · fun_prop · change Measurable ((fun p : ℕ × (ℕ → α × ℝ) ↦ reward p.1 p.2) - ∘ (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ ((stepsUntil (arm · ω.1) a m).toNat, ω.1))) - have : Measurable fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ - ((stepsUntil (arm · ω.1) a m).toNat, ω.1) := + ∘ (fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ ((stepsUntil a m ω.1).toNat, ω.1))) + have : Measurable fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ ((stepsUntil a m ω.1).toNat, ω.1) := (measurable_stepsUntil' a m).toNat.prodMk (by fun_prop) exact Measurable.comp (by fun_prop) this @@ -139,12 +134,12 @@ lemma reward_cond_arm [StandardBorelSpace α] [Nonempty α] [Countable α] (a : 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 (arm · ω.1) a m = ↑n}.indicator (fun _ ↦ 1)) + (fun ω ↦ reward n ω.1) ({ω | stepsUntil a m ω.1 = ↑n}.indicator (fun _ ↦ 1)) (Bandit.measure alg ν) := by -- first restrict to the `trajMeasure` side suffices h_indep : CondIndepFun (mα.comap (arm n)) (measurable_arm n).comap_le - (reward n) ({ω | stepsUntil (arm · ω) a m = ↑n}.indicator (fun _ ↦ 1)) + (reward n) ({ω | stepsUntil a m ω = ↑n}.indicator (fun _ ↦ 1)) (Bandit.trajMeasure alg ν) by sorry -- Now prove the independence : the indicator of `stepsUntil ... = n` is a function of @@ -173,16 +168,16 @@ lemma condIndepFun_reward_stepsUntil_arm [StandardBorelSpace α] [Countable α] (fun ω ↦ (hist (n - 1) ω, arm n ω)) (Bandit.trajMeasure alg ν) := h_indep.prod_right (by fun_prop) (by fun_prop) (by fun_prop) suffices ∃ φ : ((Iic (n - 1) → α × ℝ) × α) → ℕ, Measurable φ ∧ - ({ω : ℕ → α × ℝ | stepsUntil (arm · ω) a m = ↑n}.indicator (fun _ ↦ 1)) + ({ω : ℕ → α × ℝ | stepsUntil a m ω = ↑n}.indicator (fun _ ↦ 1)) = φ ∘ (fun ω : ℕ → α × ℝ ↦ (hist (n - 1) ω, arm n ω)) by obtain ⟨φ, hφ_meas, h_eq⟩ := this rw [h_eq] exact h_indep'.comp measurable_id hφ_meas -- it would follow from measurability wrt the sigma-algebra generated by -- `hist (n-1)` and `arm n`, but we can also give an explicit function - let k : ((Iic (n - 1) → α × ℝ) × α) → (ℕ → α) := fun x i ↦ - if hi : i ∈ Iic (n - 1) then (x.1 ⟨i, hi⟩).1 else if i = n then x.2 else a -- a is arbitrary - let φ : ((Iic (n - 1) → α × ℝ) × α) → ℕ := fun x ↦ if stepsUntil (k x) a m = ↑n then 1 else 0 + 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) + let φ : ((Iic (n - 1) → α × ℝ) × α) → ℕ := fun x ↦ if stepsUntil a m (k x) = ↑n then 1 else 0 classical have hφ_meas : Measurable φ := by refine Measurable.ite ?_ (by fun_prop) (by fun_prop) @@ -199,22 +194,24 @@ lemma condIndepFun_reward_stepsUntil_arm [StandardBorelSpace α] [Countable α] congr 1 rw [stepsUntil_eq_congr] intro i hin - simp only [arm, mem_Iic, hist, dite_eq_ite, left_eq_ite_iff, not_le, k] - intro hni - have : i = n := by grind - simp [this] + simp only [arm, mem_Iic, hist, dite_eq_ite, k] + split_ifs with h1 h2 + · rfl + · simp [h2] + · exfalso + grind lemma reward_cond_stepsUntil [StandardBorelSpace α] [Countable α] [Nonempty α] (a : α) (m n : ℕ) (hm : m ≠ 0) - (hμn : (Bandit.measure alg ν) ((fun ω ↦ stepsUntil (arm · ω.1) a m) ⁻¹' {↑n}) ≠ 0) : - 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ stepsUntil (fun x ↦ arm x ω.1) a m ← (n : ℕ∞); + (hμn : (Bandit.measure alg ν) ((fun ω ↦ stepsUntil a m ω.1) ⁻¹' {↑n}) ≠ 0) : + 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ stepsUntil a m ω.1 ← (n : ℕ∞); Bandit.measure alg ν] = ν a := by let μ := Bandit.measure alg ν have hμna : - μ ((fun ω ↦ stepsUntil (arm · ω.1) a m) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}) ≠ 0 := by + μ ((fun ω ↦ stepsUntil a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}) ≠ 0 := by suffices ((fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦ - stepsUntil (arm · ω.1) a m) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}) - = (fun ω ↦ stepsUntil (arm · ω.1) a m) ⁻¹' {↑n} by simpa [this] using hμn + 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 @@ -223,13 +220,13 @@ lemma reward_cond_stepsUntil [StandardBorelSpace α] [Countable α] [Nonempty α 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 (arm · ω.1) a m ← (n : ℕ∞); μ] - _ = (μ[|(fun ω ↦ stepsUntil (fun x ↦ arm x ω.1) a m) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}]).map + 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 - _ = (μ[|{ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) | stepsUntil (arm · ω.1) a m = ↑n}.indicator 1 ⁻¹' {1} + _ = (μ[|{ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) | stepsUntil a m ω.1 = ↑n}.indicator 1 ⁻¹' {1} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}]).map (fun ω ↦ reward n ω.1) := by congr 3 with ω simp [Set.indicator_apply] @@ -246,12 +243,12 @@ lemma reward_cond_stepsUntil [StandardBorelSpace α] [Countable α] [Nonempty α lemma condDistrib_rewardByCount_stepsUntil [Countable α] [StandardBorelSpace α] [Nonempty α] (a : α) (m : ℕ) (hm : m ≠ 0) : - condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil (arm · ω.1) a m) + condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil a m ω.1) (Bandit.measure alg ν) - =ᵐ[(Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)] Kernel.const _ (ν a) := by + =ᵐ[(Bandit.measure alg ν).map (fun ω ↦ stepsUntil a m ω.1)] Kernel.const _ (ν a) := by let μ := Bandit.measure alg ν refine (condDistrib_ae_eq_cond (μ := μ) - (X := fun ω ↦ stepsUntil (arm · ω.1) a m) (by fun_prop) (by fun_prop)).trans ?_ + (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] @@ -265,7 +262,7 @@ lemma condDistrib_rewardByCount_stepsUntil [Countable α] [StandardBorelSpace α rw [cond_of_indepFun _ (by fun_prop) (by fun_prop) (measurableSet_singleton _)] · exact (hasLaw_Z a m).map_eq · rwa [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] at hn - · exact indepFun_prod (X := fun ω : ℕ → α × ℝ ↦ stepsUntil (arm · ω) a m) + · 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)] @@ -282,21 +279,21 @@ lemma hasLaw_rewardByCount [Countable α] [StandardBorelSpace α] [Nonempty α] HasLaw (fun ω ↦ rewardByCount a m ω.1 ω.2) (ν a) (Bandit.measure alg ν) where map_eq := by have h_condDistrib : - condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil (arm · ω.1) a m) + condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil a m ω.1) (Bandit.measure alg ν) - =ᵐ[(Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)] + =ᵐ[(Bandit.measure alg ν).map (fun ω ↦ stepsUntil a m ω.1)] Kernel.const _ (ν a) := condDistrib_rewardByCount_stepsUntil a m hm calc (Bandit.measure alg ν).map (fun ω ↦ rewardByCount a m ω.1 ω.2) - _ = (condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil (arm · ω.1) a m) + _ = (condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil a m ω.1) (Bandit.measure alg ν)) - ∘ₘ ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)) := by + ∘ₘ ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil a m ω.1)) := by rw [condDistrib_comp_map (by fun_prop) (by fun_prop)] _ = (Kernel.const _ (ν a)) - ∘ₘ ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)) := + ∘ₘ ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil a m ω.1)) := Measure.comp_congr h_condDistrib _ = ν a := by have : IsProbabilityMeasure - ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)) := + ((Bandit.measure alg ν).map (fun ω ↦ stepsUntil a m ω.1)) := isProbabilityMeasure_map (by fun_prop) simp diff --git a/LeanBandits/UCB.lean b/LeanBandits/UCB.lean index b0db6429..29d2bc4b 100644 --- a/LeanBandits/UCB.lean +++ b/LeanBandits/UCB.lean @@ -17,7 +17,7 @@ open scoped ENNReal NNReal namespace Bandits -variable {α : Type*} {mα : MeasurableSpace α} {ν : Kernel α ℝ} {k : ℕ → α} {t : ℕ} {a : α} +variable {α : Type*} {mα : MeasurableSpace α} {ν : Kernel α ℝ} {t : ℕ} {a : α} section Algorithm