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: 1 addition & 0 deletions LeanBandits.lean
Original file line number Diff line number Diff line change
Expand Up @@ -16,3 +16,4 @@ import LeanBandits.ForMathlib.Traj
import LeanBandits.RewardByCountMeasure
import LeanBandits.SequentialLearning.Algorithm
import LeanBandits.SequentialLearning.Deterministic
import LeanBandits.SequentialLearning.StationaryEnv
3 changes: 2 additions & 1 deletion LeanBandits/AlgorithmBuilding.lean
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,8 @@ namespace Bandits

variable {α : Type*} [DecidableEq α] [MeasurableSpace α]

/-- Number of pulls of arm `a` up to (and including) time `n`. -/
/-- Number of pulls of arm `a` up to (and including) time `n`.
This is the number of entries in `h` in which the arm is `a`. -/
noncomputable
def pullCount' (n : ℕ) (h : Iic n → α × ℝ) (a : α) := #{s | (h s).1 = a}

Expand Down
1 change: 1 addition & 0 deletions LeanBandits/Bandit/Bandit.lean
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ Released under Apache 2.0 license as described in the file LICENSE.
Authors: Rémy Degenne, Paulo Rauber
-/
import LeanBandits.SequentialLearning.Deterministic
import LeanBandits.SequentialLearning.StationaryEnv
import LeanBandits.ForMathlib.IndepInfinitePi
import Mathlib.Probability.IdentDistrib

Expand Down
50 changes: 31 additions & 19 deletions LeanBandits/Bandit/Regret.lean
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ Authors: Rémy Degenne, Paulo Rauber
import LeanBandits.Bandit.Bandit
import Mathlib.Data.ENat.Lattice
import Mathlib.Order.CompletePartialOrder
import Mathlib.Probability.Martingale.BorelCantelli

/-!
# Regret
Expand Down Expand Up @@ -55,6 +56,11 @@ open Classical in
lemma monotone_pullCount (a : α) (h : ℕ → α × ℝ) : Monotone (pullCount a · h) :=
fun _ _ _ ↦ card_le_card (filter_subset_filter _ (by simpa))

@[mono, gcongr]
lemma pullCount_mono (a : α) {n m : ℕ} (hnm : n ≤ m) (h : ℕ → α × ℝ) :
pullCount a n h ≤ pullCount a m h :=
monotone_pullCount a h hnm

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_add_one, filter_insert]
Expand Down Expand Up @@ -83,6 +89,7 @@ lemma pullCount_congr {h' : ℕ → α × ℝ} (h_eq : ∀ i ≤ n, arm i h = ar
rw [Nat.lt_add_one_iff] at hs
rw [h_eq s hs]

-- TODO: replace this by leastGE?
/-- Number of steps until arm `a` was pulled exactly `m` times. -/
noncomputable
def stepsUntil (a : α) (m : ℕ) (h : ℕ → α × ℝ) : ℕ∞ := sInf ((↑) '' {s | pullCount a (s + 1) h = m})
Expand All @@ -93,6 +100,10 @@ lemma stepsUntil_eq_top_iff : stepsUntil a m h = ⊤ ↔ ∀ s, pullCount a (s +
lemma stepsUntil_ne_top (h_exists : ∃ s, pullCount a (s + 1) h = m) : stepsUntil a m h ≠ ⊤ := by
simpa [stepsUntil_eq_top_iff]

lemma stepsUntil_eq_leastGE (a : α) (m : ℕ) :
stepsUntil a m = leastGE (fun n h ↦ pullCount a (n + 1) h) m := by
sorry

lemma exists_pullCount_eq (h' : stepsUntil a m h ≠ ⊤) :
∃ s, pullCount a (s + 1) h = m := by
by_contra! h_contra
Expand Down Expand Up @@ -283,7 +294,7 @@ section SumRewards

/-- Sum of rewards obtained when pulling arm `a` up to time `t` (exclusive). -/
def sumRewards (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ :=
∑ s ∈ range t, if (arm s h) = a then (reward s h) else 0
∑ s ∈ range t, if arm s h = a then reward s h else 0

/-- Empirical mean reward obtained when pulling arm `a` up to time `t` (exclusive). -/
noncomputable
Expand All @@ -296,38 +307,39 @@ end SumRewards

section RewardByCount

/-- Reward obtained when pulling arm `a` for the `m`-th time. -/
/-- Reward obtained when pulling arm `a` for the `m`-th time.
If it is never pulled `m` times, the reward is given by the second component of `ω`, which in
applications will be indepedent with same law. -/
noncomputable
def rewardByCount (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : ℝ :=
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 a m h) = ⊤ then z m a else reward (stepsUntil a m h).toNat h := by
def rewardByCount (a : α) (m : ℕ) (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) : ℝ :=
match (stepsUntil a m ω.1) with
| ⊤ => ω.2 m a
| (n : ℕ) => reward n ω.1

lemma rewardByCount_eq_ite (a : α) (m : ℕ) (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) :
rewardByCount a m ω =
if (stepsUntil a m ω.1) = ⊤ then ω.2 m a else reward (stepsUntil a m ω.1).toNat ω.1 := by
unfold rewardByCount
cases stepsUntil a m h <;> simp
cases stepsUntil a m ω.1 <;> simp

lemma rewardByCount_of_stepsUntil_eq_top {ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)}
(h : stepsUntil a m ω.1 = ⊤) :
rewardByCount a m ω.1 ω.2 = ω.2 m a := by simp [rewardByCount_eq_ite, h]
rewardByCount a m ω = ω.2 m a := by simp [rewardByCount_eq_ite, h]

lemma rewardByCount_of_stepsUntil_eq_coe {ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)}
(h : stepsUntil a m ω.1 = n) :
rewardByCount a m ω.1 ω.2 = reward n ω.1 := by simp [rewardByCount_eq_ite, h]
rewardByCount a m ω = 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 t h) t h + 1) h z = reward t h := by
lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) :
rewardByCount (arm t ω.1) (pullCount (arm t ω.1) t ω.1 + 1) ω = reward t ω.1 := by
rw [rewardByCount, ← pullCount_eq_pullCount_add_one, stepsUntil_pullCount_eq]

lemma sum_rewardByCount_eq_sumRewards
(a : α) (t : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) :
∑ m ∈ Icc 1 (pullCount a t h), rewardByCount a m h z = sumRewards a t h := by
lemma sum_rewardByCount_eq_sumRewards (a : α) (t : ℕ) (ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)) :
∑ m ∈ Icc 1 (pullCount a t ω.1), rewardByCount a m ω = sumRewards a t ω.1 := by
induction t with
| zero => simp [pullCount, sumRewards]
| succ t ht =>
by_cases hta : arm t h = a
by_cases hta : arm t ω.1 = a
· rw [← hta] at ht ⊢
rw [pullCount_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht]
unfold sumRewards
Expand Down
Loading