@@ -5,15 +5,19 @@ Authors: Rémy Degenne
55-/
66module
77
8- public import LeanMachineLearning.SequentialLearning.Algorithms.AuxSums
98public import LeanMachineLearning.SequentialLearning.Deterministic
109public import LeanMachineLearning.SequentialLearning.FiniteActions
1110public import LeanMachineLearning.SequentialLearning.StationaryEnv
1211
1312/-! # Round-Robin algorithm
1413
15- That algorithm pulls each action in a round-robin fashion.
16- That is, if there are `K` actions, then at time `n`, it pulls the action `n % K`.
14+ That algorithm chooses each of finitely many actions in a round-robin fashion.
15+ That is, if there are `K` actions numbered from 0 to `K - 1`, then at time `n` it chooses
16+ he action `n % K`.
17+
18+ ## Main definitions
19+
20+ * `roundRobinAlgorithm`: the Round-Robin algorithm.
1721
1822 -/
1923
@@ -22,21 +26,61 @@ That is, if there are `K` actions, then at time `n`, it pulls the action `n % K`
2226open MeasureTheory ProbabilityTheory Finset Learning
2327open scoped ENNReal NNReal
2428
25- namespace Bandits
29+ section Aux
30+
31+ lemma sum_mod_range {K : ℕ} (hK : 0 < K) (a : Fin K) :
32+ (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0 ) = 1 := by
33+ have h_iff (s : ℕ) (hs : s < K) : ⟨s % K, Nat.mod_lt _ hK⟩ = a ↔ s = a := by
34+ simp only [Nat.mod_eq_of_lt hs, Fin.ext_iff]
35+ calc (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0 )
36+ _ = ∑ s ∈ range K, if s = a then 1 else 0 := sum_congr rfl fun s hs ↦ by grind
37+ _ = _ := by
38+ rw [sum_ite_eq']
39+ simp
40+
41+ lemma sum_mod_range_mul {K : ℕ} (hK : 0 < K) (m : ℕ) (a : Fin K) :
42+ (∑ s ∈ range (K * m), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0 ) = m := by
43+ induction m with
44+ | zero => simp
45+ | succ n hn =>
46+ calc (∑ s ∈ range (K * (n + 1 )), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0 )
47+ _ = (∑ s ∈ range (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0 ) := by ring_nf
48+ _ = (∑ s ∈ range (K * n), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0 )
49+ + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0 ) := by
50+ rw [sum_range_add_sum_Ico]
51+ grind
52+ _ = n + (∑ s ∈ Ico (K * n) (K * n + K), if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0 ) := by
53+ rw [hn]
54+ _ = n + (∑ s ∈ range K, if ⟨(s + K * n) % K, Nat.mod_lt _ hK⟩ = a then 1 else 0 ) := by
55+ congr 1
56+ let e : ℕ ↪ ℕ := ⟨fun i : ℕ ↦ i + K * n, fun i j hij ↦ by grind⟩
57+ have : Finset.map e (range K) = Ico (K * n) (K * n + K) := by
58+ ext x
59+ simp only [mem_map, mem_range, Function.Embedding.coeFn_mk, mem_Ico, e]
60+ refine ⟨fun h ↦ by grind, fun h ↦ ?_⟩
61+ use x - K * n
62+ grind
63+ rw [← this, Finset.sum_map]
64+ congr
65+ _ = n + (∑ s ∈ range K, if ⟨s % K, Nat.mod_lt _ hK⟩ = a then 1 else 0 ) := by simp
66+ _ = n + 1 := by rw [sum_mod_range hK]
67+
68+ end Aux
69+
70+ namespace Learning
2671
2772variable {K : ℕ}
2873
2974section AlgorithmDefinition
3075
31- /-- Arm pulled by the Round-Robin algorithm at time `n + 1`. This is arm `n % K`. -/
76+ /-- Action chosen by the Round-Robin algorithm at time `n + 1`. This is action `(n + 1) % K`. -/
3277noncomputable
33- def RoundRobin.nextArm (hK : 0 < K) (n : ℕ) : Fin K := ⟨(n + 1 ) % K, Nat.mod_lt _ hK⟩
78+ def RoundRobin.nextAction (hK : 0 < K) (n : ℕ) : Fin K := ⟨(n + 1 ) % K, Nat.mod_lt _ hK⟩
3479
35- /-- The Round-Robin algorithm: deterministic algorithm that chooses the next arm according
36- to `RoundRobin.nextArm`. -/
80+ /-- The Round-Robin algorithm: deterministic algorithm that chooses action `n % K` at time `n`. -/
3781noncomputable
3882def roundRobinAlgorithm (hK : 0 < K) : Algorithm (Fin K) ℝ :=
39- detAlgorithm (fun n _ ↦ RoundRobin.nextArm hK n) (by fun_prop) ⟨0 , hK⟩
83+ detAlgorithm (fun n _ ↦ RoundRobin.nextAction hK n) (by fun_prop) ⟨0 , hK⟩
4084
4185end AlgorithmDefinition
4286
@@ -47,36 +91,36 @@ variable {hK : 0 < K} {ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν]
4791 {P : Measure Ω} [IsProbabilityMeasure P]
4892 {A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ}
4993
50- lemma arm_zero [Nonempty (Fin K)]
94+ lemma action_zero [Nonempty (Fin K)]
5195 (h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P 0 ) :
5296 A 0 =ᵐ[P] fun _ ↦ ⟨0 , hK⟩ := by
5397 have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
5498 exact h.action_zero_detAlgorithm
5599
56- lemma arm_ae_eq_roundRobinNextArm [Nonempty (Fin K)] (n : ℕ)
100+ lemma action_ae_eq_roundRobinNextAction [Nonempty (Fin K)] (n : ℕ)
57101 (h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P (n + 1 )) :
58- A (n + 1 ) =ᵐ[P] fun _ ↦ nextArm hK n :=
102+ A (n + 1 ) =ᵐ[P] fun _ ↦ nextAction hK n :=
59103 h.action_detAlgorithm_ae_eq (by grind)
60104
61- /-- The arm pulled at time `n` is the arm `n % K`. -/
62- lemma arm_ae_eq [Nonempty (Fin K)] (n : ℕ)
105+ /-- The action chosen at time `n` is the action `n % K`. -/
106+ lemma action_ae_eq [Nonempty (Fin K)] (n : ℕ)
63107 (h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P n) :
64108 A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ := by
65109 cases n with
66- | zero => exact arm_zero h
110+ | zero => exact action_zero h
67111 | succ n =>
68- filter_upwards [arm_ae_eq_roundRobinNextArm n h] with h hn_eq
69- rw [hn_eq, nextArm ]
112+ filter_upwards [action_ae_eq_roundRobinNextAction n h] with h hn_eq
113+ rw [hn_eq, nextAction ]
70114
71- /-- At time `K * m`, the number of pulls of each arm is equal to `m`. -/
115+ /-- At time `K * m`, the number of times each action is chosen is equal to `m`. -/
72116lemma pullCount_mul [Nonempty (Fin K)] (m : ℕ)
73117 (h : IsAlgEnvSeqUntil A R (roundRobinAlgorithm hK) (stationaryEnv ν) P (K * m - 1 ))
74118 (a : Fin K) :
75119 pullCount A a (K * m) =ᵐ[P] fun _ ↦ m := by
76120 rw [Filter.EventuallyEq]
77121 simp_rw [pullCount_eq_sum]
78122 have h_arm (n : range (K * m)) : A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ :=
79- arm_ae_eq n (h.mono (by have := n.2 ; simp only [mem_range] at this; grind))
123+ action_ae_eq n (h.mono (by have := n.2 ; simp only [mem_range] at this; grind))
80124 simp_rw [Filter.EventuallyEq, ← ae_all_iff] at h_arm
81125 filter_upwards [h_arm] with ω h_arm
82126 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)]
119163
120164end RoundRobin
121165
122- end Bandits
166+ end Learning
0 commit comments