Skip to content

Commit cbf06d4

Browse files
authored
Doc and names for roundRobin (#89)
2 parents 7713ffb + 23ce9b9 commit cbf06d4

8 files changed

Lines changed: 80 additions & 95 deletions

File tree

‎LeanMachineLearning.lean‎

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,6 @@ public import LeanMachineLearning.MeasureTheory.Constructions.Polish.StandardBor
1919
public import LeanMachineLearning.Probability.Moments.SubGaussian
2020
public import LeanMachineLearning.Probability.Kernel.IonescuTulcea.Traj
2121
public import LeanMachineLearning.SequentialLearning.Algorithm
22-
public import LeanMachineLearning.SequentialLearning.Algorithms.AuxSums
2322
public import LeanMachineLearning.SequentialLearning.Algorithms.RoundRobin
2423
public import LeanMachineLearning.SequentialLearning.Deterministic
2524
public import LeanMachineLearning.SequentialLearning.FiniteActions

‎LeanMachineLearning/Online/Bandit/Algorithms/ETC.lean‎

Lines changed: 5 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -25,14 +25,13 @@ variable {K : ℕ}
2525
section AlgorithmDefinition
2626

2727
/-- Arm pulled by the ETC algorithm at time `n + 1`.
28-
For `n < K * m - 1`, this is arm `n % K`.
28+
For `n < K * m - 1`, this is arm `(n + 1) % K`.
2929
For `n = K * m - 1`, this is the arm with the highest empirical mean after the exploration phase.
3030
For `n ≥ K * m`, this is the same arm as at time `n`. -/
3131
noncomputable
3232
def ETC.nextArm (hK : 0 < K) (m n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K :=
3333
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
34-
if hn : n < K * m - 1 then
35-
⟨(n + 1) % K, Nat.mod_lt _ hK⟩ -- for `n = 0` we have pulled arm 0 already, and we pull arm 1
34+
if hn : n < K * m - 1 then RoundRobin.nextAction hK n
3635
else
3736
if hn_eq : n = K * m - 1 then measurableArgmax (empMean' n) h
3837
else (h ⟨n, by simp⟩).1
@@ -75,16 +74,15 @@ lemma isAlgEnvSeqUntil_roundRobinAlgorithm [Nonempty (Fin K)]
7574
convert h.hasCondDistrib_action n using 1
7675
simp only [roundRobinAlgorithm, detAlgorithm_policy, etcAlgorithm]
7776
congr 1 with h
78-
unfold ETC.nextArm RoundRobin.nextArm
79-
simp [hn]
77+
simp [ETC.nextArm, hn]
8078
hasCondDistrib_reward n _ := h.hasCondDistrib_reward n
8179

8280
section AlgorithmBehavior
8381

8482
lemma arm_zero [Nonempty (Fin K)]
8583
(h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) :
8684
A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ :=
87-
RoundRobin.arm_zero ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono zero_le')
85+
RoundRobin.action_zero ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono zero_le')
8886

8987
lemma arm_ae_eq_etcNextArm [Nonempty (Fin K)]
9088
(h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (n : ℕ) :
@@ -96,7 +94,7 @@ lemma arm_ae_eq_etcNextArm [Nonempty (Fin K)]
9694
lemma arm_of_lt [Nonempty (Fin K)]
9795
(h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) {n : ℕ} (hn : n < K * m) :
9896
A n =ᵐ[P] fun _ ↦ ⟨n % K, Nat.mod_lt _ hK⟩ :=
99-
RoundRobin.arm_ae_eq n ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono (by grind))
97+
RoundRobin.action_ae_eq n ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono (by grind))
10098

10199
/-- The arm pulled at time `K * m` is the arm with the highest empirical mean after the exploration
102100
phase. -/

‎LeanMachineLearning/Online/Bandit/Algorithms/UCB.lean‎

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -37,7 +37,7 @@ open Classical in
3737
noncomputable
3838
def UCB.nextArm (hK : 0 < K) (c : ℝ) (n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K :=
3939
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
40-
if n < K - 1 then ⟨(n + 1) % K, Nat.mod_lt _ hK⟩ else
40+
if n < K - 1 then RoundRobin.nextAction hK n else
4141
measurableArgmax (fun h a ↦ empMean' n h a + ucbWidth' c n h a) h
4242

4343
@[fun_prop]
@@ -75,8 +75,7 @@ lemma isAlgEnvSeqUntil_roundRobinAlgorithm [Nonempty (Fin K)]
7575
convert h.hasCondDistrib_action n using 1
7676
simp only [roundRobinAlgorithm, detAlgorithm_policy, ucbAlgorithm]
7777
congr 1 with h
78-
unfold UCB.nextArm RoundRobin.nextArm
79-
simp [hn]
78+
simp [UCB.nextArm, hn]
8079
hasCondDistrib_reward n _ := h.hasCondDistrib_reward n
8180

8281
section AlgorithmBehavior
@@ -103,7 +102,7 @@ lemma ucbWidth_eq_ucbWidth' (c : ℝ) (a : Fin K) (n : ℕ) (ω : Ω) (hn : n
103102
lemma arm_zero [Nonempty (Fin K)]
104103
(h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) :
105104
A 0 =ᵐ[P] fun _ ↦ ⟨0, hK⟩ :=
106-
RoundRobin.arm_zero ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono zero_le')
105+
RoundRobin.action_zero ((isAlgEnvSeqUntil_roundRobinAlgorithm h).mono zero_le')
107106

108107
lemma arm_ae_eq_ucbNextArm [Nonempty (Fin K)]
109108
(h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (n : ℕ) :
@@ -141,7 +140,8 @@ lemma forall_arm_eq_mod_of_lt [Nonempty (Fin K)]
141140
| succ n _ =>
142141
filter_upwards [arm_ae_eq_ucbNextArm h n] with h h_eq
143142
rw [h_eq, nextArm, if_pos]
144-
grind
143+
· rfl
144+
· grind
145145

146146
lemma forall_ucbIndex_le_ucbIndex_arm [Nonempty (Fin K)]
147147
(h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) (a : Fin K) :

‎LeanMachineLearning/SequentialLearning/Algorithms/AuxSums.lean‎

Lines changed: 0 additions & 56 deletions
This file was deleted.

‎LeanMachineLearning/SequentialLearning/Algorithms/RoundRobin.lean‎

Lines changed: 64 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -5,15 +5,19 @@ Authors: Rémy Degenne
55
-/
66
module
77

8-
public import LeanMachineLearning.SequentialLearning.Algorithms.AuxSums
98
public import LeanMachineLearning.SequentialLearning.Deterministic
109
public import LeanMachineLearning.SequentialLearning.FiniteActions
1110
public 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`
2226
open MeasureTheory ProbabilityTheory Finset Learning
2327
open 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

2772
variable {K : ℕ}
2873

2974
section 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`. -/
3277
noncomputable
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`. -/
3781
noncomputable
3882
def 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

4185
end 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`. -/
72116
lemma 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

120164
end RoundRobin
121165

122-
end Bandits
166+
end Learning

‎blueprint/lean_decls‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -114,9 +114,9 @@ Bandits.probReal_sumRewards_le_sumRewards_le
114114
Bandits.probReal_sum_le_sum_streamMeasure
115115
Bandits.prob_sum_le_sqrt_log
116116
Bandits.prob_sum_ge_sqrt_log
117-
Bandits.RoundRobin.nextArm
118-
Bandits.roundRobinAlgorithm
119-
Bandits.RoundRobin.pullCount_mul
117+
Learning.RoundRobin.nextAction
118+
Learning.roundRobinAlgorithm
119+
Learning.RoundRobin.pullCount_mul
120120
Bandits.ETC.nextArm
121121
Bandits.etcAlgorithm
122122
Bandits.ETC.isAlgEnvSeqUntil_roundRobinAlgorithm

‎blueprint/src/chapters/etc.tex‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,15 +8,15 @@ \section{Round-Robin}
88
\begin{definition}\label{def:roundRobinAlgorithm}
99
\uses{def:detAlgorithm,def:algorithm}
1010
\leanok
11-
\lean{Bandits.RoundRobin.nextArm, Bandits.roundRobinAlgorithm}
11+
\lean{Learning.RoundRobin.nextAction, Learning.roundRobinAlgorithm}
1212
The Round-Robin algorithm is defined as follows: at time $t \in \mathbb{N}$, $A_t = t \mod K$.
1313
\end{definition}
1414

1515

1616
\begin{lemma}\label{lem:pullCount_roundRobinAlgorithm}
1717
\uses{def:stationaryEnv,def:IsAlgEnvSeq,def:pullCount,def:roundRobinAlgorithm}
1818
\leanok
19-
\lean{Bandits.RoundRobin.pullCount_mul}
19+
\lean{Learning.RoundRobin.pullCount_mul}
2020
For the Round-Robin algorithm, for any arm $a \in [K]$, at time $Km$ we have
2121
\begin{align*}
2222
N_{Km,a}

‎tutorial/Manual/Pages/DefiningAlgorithm.lean‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -146,7 +146,7 @@ open Classical in
146146
noncomputable
147147
def UCB.nextArm (hK : 0 < K) (c : ℝ) (n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K :=
148148
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
149-
if n < K - 1 then ⟨(n + 1) % K, Nat.mod_lt _ hK⟩ else
149+
if n < K - 1 then RoundRobin.nextAction hK n else
150150
measurableArgmax (fun h a ↦ empMean' n h a + ucbWidth' c n h a) h
151151

152152
@[fun_prop]

0 commit comments

Comments
 (0)