diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index c8f890c9..731b3ebc 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -19,7 +19,6 @@ public import LeanMachineLearning.MeasureTheory.Constructions.Polish.StandardBor public import LeanMachineLearning.Probability.Moments.SubGaussian public import LeanMachineLearning.Probability.Kernel.IonescuTulcea.Traj public import LeanMachineLearning.SequentialLearning.Algorithm -public import LeanMachineLearning.SequentialLearning.Algorithms.AuxSums public import LeanMachineLearning.SequentialLearning.Algorithms.RoundRobin public import LeanMachineLearning.SequentialLearning.Deterministic public import LeanMachineLearning.SequentialLearning.FiniteActions diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/ETC.lean b/LeanMachineLearning/Online/Bandit/Algorithms/ETC.lean index 5d52c849..c59cb570 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/ETC.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/ETC.lean @@ -25,14 +25,13 @@ variable {K : ℕ} section AlgorithmDefinition /-- Arm pulled by the ETC algorithm at time `n + 1`. -For `n < K * m - 1`, this is arm `n % K`. +For `n < K * m - 1`, this is arm `(n + 1) % K`. For `n = K * m - 1`, this is the arm with the highest empirical mean after the exploration phase. For `n ≥ K * m`, this is the same arm as at time `n`. -/ noncomputable def ETC.nextArm (hK : 0 < K) (m n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - if hn : n < K * m - 1 then - ⟨(n + 1) % K, Nat.mod_lt _ hK⟩ -- for `n = 0` we have pulled arm 0 already, and we pull arm 1 + if hn : n < K * m - 1 then RoundRobin.nextAction hK n else if hn_eq : n = K * m - 1 then measurableArgmax (empMean' n) h else (h ⟨n, by simp⟩).1 @@ -75,8 +74,7 @@ lemma isAlgEnvSeqUntil_roundRobinAlgorithm [Nonempty (Fin K)] convert h.hasCondDistrib_action n using 1 simp only [roundRobinAlgorithm, detAlgorithm_policy, etcAlgorithm] congr 1 with h - unfold ETC.nextArm RoundRobin.nextArm - simp [hn] + simp [ETC.nextArm, hn] hasCondDistrib_reward n _ := h.hasCondDistrib_reward n section AlgorithmBehavior @@ -84,7 +82,7 @@ section AlgorithmBehavior lemma arm_zero [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) : A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ := - RoundRobin.arm_zero ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono zero_le') + RoundRobin.action_zero ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono zero_le') lemma arm_ae_eq_etcNextArm [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (n : ℕ) : @@ -96,7 +94,7 @@ lemma arm_ae_eq_etcNextArm [Nonempty (Fin K)] 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⟩ := - RoundRobin.arm_ae_eq n ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono (by grind)) + RoundRobin.action_ae_eq n ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono (by grind)) /-- The arm pulled at time `K * m` is the arm with the highest empirical mean after the exploration phase. -/ diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/UCB.lean b/LeanMachineLearning/Online/Bandit/Algorithms/UCB.lean index e7eeb81d..4a4553cc 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/UCB.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/UCB.lean @@ -37,7 +37,7 @@ open Classical in noncomputable def UCB.nextArm (hK : 0 < K) (c : ℝ) (n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - if n < K - 1 then ⟨(n + 1) % K, Nat.mod_lt _ hK⟩ else + if n < K - 1 then RoundRobin.nextAction hK n else measurableArgmax (fun h a ↦ empMean' n h a + ucbWidth' c n h a) h @[fun_prop] @@ -75,8 +75,7 @@ lemma isAlgEnvSeqUntil_roundRobinAlgorithm [Nonempty (Fin K)] convert h.hasCondDistrib_action n using 1 simp only [roundRobinAlgorithm, detAlgorithm_policy, ucbAlgorithm] congr 1 with h - unfold UCB.nextArm RoundRobin.nextArm - simp [hn] + simp [UCB.nextArm, hn] hasCondDistrib_reward n _ := h.hasCondDistrib_reward n section AlgorithmBehavior @@ -103,7 +102,7 @@ lemma ucbWidth_eq_ucbWidth' (c : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) (hn : n lemma arm_zero [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) : A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ := - RoundRobin.arm_zero ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono zero_le') + RoundRobin.action_zero ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono zero_le') lemma arm_ae_eq_ucbNextArm [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (n : ℕ) : @@ -141,7 +140,8 @@ lemma forall_arm_eq_mod_of_lt [Nonempty (Fin K)] | succ n _ => filter_upwards [arm_ae_eq_ucbNextArm h n] with h h_eq rw [h_eq, nextArm, if_pos] - grind + · rfl + · grind lemma forall_ucbIndex_le_ucbIndex_arm [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) : diff --git a/LeanMachineLearning/SequentialLearning/Algorithms/AuxSums.lean b/LeanMachineLearning/SequentialLearning/Algorithms/AuxSums.lean deleted file mode 100644 index 69b104a0..00000000 --- a/LeanMachineLearning/SequentialLearning/Algorithms/AuxSums.lean +++ /dev/null @@ -1,56 +0,0 @@ -/- -Copyright (c) 2026 Rémy Degenne. All rights reserved. -Released under Apache 2.0 license as described in the file LICENSE. -Authors: Rémy Degenne --/ -module - -public import Mathlib.Algebra.BigOperators.Intervals -public import Mathlib.Algebra.BigOperators.Ring.Finset -public import Mathlib.Tactic.Ring.RingNF - -/-! -# Lemmas about sums of indicators - --/ - -@[expose] public section - -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/LeanMachineLearning/SequentialLearning/Algorithms/RoundRobin.lean b/LeanMachineLearning/SequentialLearning/Algorithms/RoundRobin.lean index 09d1694b..bfac2e07 100644 --- a/LeanMachineLearning/SequentialLearning/Algorithms/RoundRobin.lean +++ b/LeanMachineLearning/SequentialLearning/Algorithms/RoundRobin.lean @@ -5,15 +5,19 @@ Authors: Rémy Degenne -/ module -public import LeanMachineLearning.SequentialLearning.Algorithms.AuxSums public import LeanMachineLearning.SequentialLearning.Deterministic public import LeanMachineLearning.SequentialLearning.FiniteActions public import LeanMachineLearning.SequentialLearning.StationaryEnv /-! # Round-Robin algorithm -That algorithm pulls each action in a round-robin fashion. -That is, if there are `K` actions, then at time `n`, it pulls the action `n % K`. +That algorithm chooses each of finitely many actions in a round-robin fashion. +That is, if there are `K` actions numbered from 0 to `K - 1`, then at time `n` it chooses +he action `n % K`. + +## Main definitions + +* `roundRobinAlgorithm`: the Round-Robin algorithm. -/ @@ -22,21 +26,61 @@ That is, if there are `K` actions, then at time `n`, it pulls the action `n % K` open MeasureTheory ProbabilityTheory Finset Learning open scoped ENNReal NNReal -namespace Bandits +section Aux + +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 Learning variable {K : ℕ} section AlgorithmDefinition -/-- Arm pulled by the Round-Robin algorithm at time `n + 1`. This is arm `n % K`. -/ +/-- Action chosen by the Round-Robin algorithm at time `n + 1`. This is action `(n + 1) % K`. -/ noncomputable -def RoundRobin.nextArm (hK : 0 < K) (n : ℕ) : Fin K := ⟨(n + 1) % K, Nat.mod_lt _ hK⟩ +def RoundRobin.nextAction (hK : 0 < K) (n : ℕ) : Fin K := ⟨(n + 1) % K, Nat.mod_lt _ hK⟩ -/-- The Round-Robin algorithm: deterministic algorithm that chooses the next arm according -to `RoundRobin.nextArm`. -/ +/-- The Round-Robin algorithm: deterministic algorithm that chooses action `n % K` at time `n`. -/ noncomputable def roundRobinAlgorithm (hK : 0 < K) : Algorithm (Fin K) ℝ := - detAlgorithm (fun n _ ↦ RoundRobin.nextArm hK n) (by fun_prop) ⟨0, hK⟩ + detAlgorithm (fun n _ ↦ RoundRobin.nextAction hK n) (by fun_prop) ⟨0, hK⟩ end AlgorithmDefinition @@ -47,28 +91,28 @@ variable {hK : 0 < K} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν] {P : Measure Ω} [IsProbabilityMeasure P] {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ} -lemma arm_zero [Nonempty (Fin K)] +lemma action_zero [Nonempty (Fin K)] (h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P 0) : A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ := by have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK exact h.action_zero_detAlgorithm -lemma arm_ae_eq_roundRobinNextArm [Nonempty (Fin K)] (n : ℕ) +lemma action_ae_eq_roundRobinNextAction [Nonempty (Fin K)] (n : ℕ) (h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P (n + 1)) : - A (n + 1) =ᵐ[P] fun _ ↦ nextArm hK n := + A (n + 1) =ᵐ[P] fun _ ↦ nextAction hK n := h.action_detAlgorithm_ae_eq (by grind) -/-- The arm pulled at time `n` is the arm `n % K`. -/ -lemma arm_ae_eq [Nonempty (Fin K)] (n : ℕ) +/-- The action chosen at time `n` is the action `n % K`. -/ +lemma action_ae_eq [Nonempty (Fin K)] (n : ℕ) (h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P n) : A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := by cases n with - | zero => exact arm_zero h + | zero => exact action_zero h | succ n => - filter_upwards [arm_ae_eq_roundRobinNextArm n h] with h hn_eq - rw [hn_eq, nextArm] + filter_upwards [action_ae_eq_roundRobinNextAction n h] with h hn_eq + rw [hn_eq, nextAction] -/-- At time `K * m`, the number of pulls of each arm is equal to `m`. -/ +/-- At time `K * m`, the number of times each action is chosen is equal to `m`. -/ lemma pullCount_mul [Nonempty (Fin K)] (m : ℕ) (h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P (K * m - 1)) (a : Fin K) : @@ -76,7 +120,7 @@ lemma pullCount_mul [Nonempty (Fin K)] (m : ℕ) rw [Filter.EventuallyEq] simp_rw [pullCount_eq_sum] have h_arm (n : range (K * m)) : A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := - arm_ae_eq n (h.mono (by have := n.2; simp only [mem_range] at this; grind)) + action_ae_eq n (h.mono (by have := n.2; simp only [mem_range] at this; grind)) 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)) : A i ω = ⟨i % K, Nat.mod_lt _ hK⟩ := h_arm ⟨i, hi⟩ @@ -119,4 +163,4 @@ lemma pullCount_pos_of_pullCount_gt_one [Nonempty (Fin K)] end RoundRobin -end Bandits +end Learning diff --git a/blueprint/lean_decls b/blueprint/lean_decls index 266a4a7d..4a30eb78 100644 --- a/blueprint/lean_decls +++ b/blueprint/lean_decls @@ -114,9 +114,9 @@ Bandits.probReal_sumRewards_le_sumRewards_le Bandits.probReal_sum_le_sum_streamMeasure Bandits.prob_sum_le_sqrt_log Bandits.prob_sum_ge_sqrt_log -Bandits.RoundRobin.nextArm -Bandits.roundRobinAlgorithm -Bandits.RoundRobin.pullCount_mul +Learning.RoundRobin.nextAction +Learning.roundRobinAlgorithm +Learning.RoundRobin.pullCount_mul Bandits.ETC.nextArm Bandits.etcAlgorithm Bandits.ETC.isAlgEnvSeqUntil_roundRobinAlgorithm diff --git a/blueprint/src/chapters/etc.tex b/blueprint/src/chapters/etc.tex index 2fda856e..0525c179 100644 --- a/blueprint/src/chapters/etc.tex +++ b/blueprint/src/chapters/etc.tex @@ -8,7 +8,7 @@ \section{Round-Robin} \begin{definition}\label{def:roundRobinAlgorithm} \uses{def:detAlgorithm,def:algorithm} \leanok - \lean{Bandits.RoundRobin.nextArm, Bandits.roundRobinAlgorithm} + \lean{Learning.RoundRobin.nextAction, Learning.roundRobinAlgorithm} The Round-Robin algorithm is defined as follows: at time $t \in \mathbb{N}$, $A_t = t \mod K$. \end{definition} @@ -16,7 +16,7 @@ \section{Round-Robin} \begin{lemma}\label{lem:pullCount_roundRobinAlgorithm} \uses{def:stationaryEnv,def:IsAlgEnvSeq,def:pullCount,def:roundRobinAlgorithm} \leanok - \lean{Bandits.RoundRobin.pullCount_mul} + \lean{Learning.RoundRobin.pullCount_mul} For the Round-Robin algorithm, for any arm $a \in [K]$, at time $Km$ we have \begin{align*} N_{Km,a} diff --git a/tutorial/Manual/Pages/DefiningAlgorithm.lean b/tutorial/Manual/Pages/DefiningAlgorithm.lean index 76669dd5..916b6c5a 100644 --- a/tutorial/Manual/Pages/DefiningAlgorithm.lean +++ b/tutorial/Manual/Pages/DefiningAlgorithm.lean @@ -146,7 +146,7 @@ open Classical in noncomputable def UCB.nextArm (hK : 0 < K) (c : ℝ) (n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - if n < K - 1 then ⟨(n + 1) % K, Nat.mod_lt _ hK⟩ else + if n < K - 1 then RoundRobin.nextAction hK n else measurableArgmax (fun h a ↦ empMean' n h a + ucbWidth' c n h a) h @[fun_prop]