Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion LeanMachineLearning.lean
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
12 changes: 5 additions & 7 deletions LeanMachineLearning/Online/Bandit/Algorithms/ETC.lean
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -75,16 +74,15 @@ 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

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 : ℕ) :
Expand All @@ -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. -/
Expand Down
10 changes: 5 additions & 5 deletions LeanMachineLearning/Online/Bandit/Algorithms/UCB.lean
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down Expand Up @@ -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
Expand All @@ -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 : ℕ) :
Expand Down Expand Up @@ -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) :
Expand Down
56 changes: 0 additions & 56 deletions LeanMachineLearning/SequentialLearning/Algorithms/AuxSums.lean

This file was deleted.

84 changes: 64 additions & 20 deletions LeanMachineLearning/SequentialLearning/Algorithms/RoundRobin.lean
Original file line number Diff line number Diff line change
Expand Up @@ -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.

-/

Expand All @@ -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

Expand All @@ -47,36 +91,36 @@ 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) :
pullCount A a (K * m) =ᵐ[P] fun _ ↦ m := by
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⟩
Expand Down Expand Up @@ -119,4 +163,4 @@ lemma pullCount_pos_of_pullCount_gt_one [Nonempty (Fin K)]

end RoundRobin

end Bandits
end Learning
6 changes: 3 additions & 3 deletions blueprint/lean_decls
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
4 changes: 2 additions & 2 deletions blueprint/src/chapters/etc.tex
Original file line number Diff line number Diff line change
Expand Up @@ -8,15 +8,15 @@ \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}


\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}
Expand Down
2 changes: 1 addition & 1 deletion tutorial/Manual/Pages/DefiningAlgorithm.lean
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
Loading