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 @@ -6,6 +6,7 @@ import LeanBandits.ForMathlib.CondDistrib
import LeanBandits.ForMathlib.KernelCompositionLemmas
import LeanBandits.ForMathlib.KernelCompositionParallelComp
import LeanBandits.ForMathlib.KernelSub
import LeanBandits.ForMathlib.SubGaussian
import LeanBandits.ForMathlib.Traj
import LeanBandits.Regret
import LeanBandits.RewardByCountMeasure
Expand Down
27 changes: 27 additions & 0 deletions LeanBandits/Bandit.lean
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,33 @@ lemma snd_measure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν]

end Bandit

section StreamMeasure

lemma _root_.hasLaw_eval_infinitePi {ι : Type*} {X : ι → Type*} {mX : ∀ i, MeasurableSpace (X i)}
(μ : (i : ι) → Measure (X i)) [hμ : ∀ i, IsProbabilityMeasure (μ i)] (i : ι) :
HasLaw (Function.eval i) (μ i) (Measure.infinitePi μ) where
aemeasurable := Measurable.aemeasurable (by fun_prop)
map_eq := by exact (measurePreserving_eval_infinitePi μ i).map_eq

lemma hasLaw_eval_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
HasLaw (fun h : ℕ → α → R ↦ h n) (Measure.infinitePi ν) (Bandit.streamMeasure ν) :=
hasLaw_eval_infinitePi (fun _ ↦ Measure.infinitePi ν) n

lemma hasLaw_eval_eval_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) (a : α) :
HasLaw (fun h : ℕ → α → R ↦ h n a) (ν a) (Bandit.streamMeasure ν) :=
(hasLaw_eval_infinitePi ν a).comp (hasLaw_eval_streamMeasure ν n)

lemma identDistrib_eval_eval_id_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) (a : α) :
IdentDistrib (fun h : ℕ → α → R ↦ h n a) id (Bandit.streamMeasure ν) (ν a) where
aemeasurable_fst := Measurable.aemeasurable (by fun_prop)
aemeasurable_snd := Measurable.aemeasurable (by fun_prop)
map_eq := by
rw [← (hasLaw_eval_eval_streamMeasure ν n a).map_eq,
Measure.map_map (by fun_prop) (by fun_prop)]
simp

end StreamMeasure

/-- `arm n` is the arm pulled at time `n`. This is a random variable on the measurable space
`ℕ → α × ℝ`. -/
def arm (n : ℕ) (h : ℕ → α × R) : α := (h n).1
Expand Down
253 changes: 213 additions & 40 deletions LeanBandits/ETC.lean

Large diffs are not rendered by default.

91 changes: 91 additions & 0 deletions LeanBandits/ForMathlib/SubGaussian.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,91 @@
/-
Copyright (c) 2025 Rémy Degenne. All rights reserved.
Released under Apache 2.0 license as described in the file LICENSE.
Authors: Rémy Degenne
-/
import Mathlib.Probability.Moments.SubGaussian

open MeasureTheory
open scoped ENNReal NNReal

namespace ProbabilityTheory

theorem mgf_const_mul {Ω : Type*} {m : MeasurableSpace Ω} {X : Ω → ℝ} {μ : Measure Ω}
{t : ℝ} (α : ℝ) : mgf (fun ω ↦ α * X ω) μ t = mgf X μ (α * t) := by
rw [← mgf_smul_left]
rfl

namespace Kernel.HasSubgaussianMGF

variable {Ω Ω' : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'}
{ν : Measure Ω'} {κ : Kernel Ω' Ω} {X : Ω → ℝ} {c : ℝ≥0}

lemma id_map_iff (hX : Measurable X) :
HasSubgaussianMGF X c κ ν ↔ HasSubgaussianMGF id c (κ.map X) ν := by
refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩
· constructor
· intro t
rw [← Kernel.deterministic_comp_eq_map hX, ← Measure.comp_assoc,
Measure.deterministic_comp_eq_map]
rw [integrable_map_measure (by fun_prop) hX.aemeasurable]
exact h.integrable_exp_mul t
· simp_rw [Kernel.map_apply _ hX, mgf_id_map hX.aemeasurable]
exact h.mgf_le
· have : X = id ∘ X := rfl
rw [this]
exact .of_map hX h

protected lemma const_mul (h : HasSubgaussianMGF X c κ ν) (r : ℝ) :
HasSubgaussianMGF (fun ω ↦ r * X ω) (⟨r ^ 2, sq_nonneg r⟩ * c) κ ν where
integrable_exp_mul t := by
simp_rw [← mul_assoc]
exact h.integrable_exp_mul (t * r)
mgf_le := by
filter_upwards [h.mgf_le] with ω hω t
specialize hω (t * r)
rw [mgf_const_mul, mul_comm]
refine hω.trans_eq ?_
congr 1
simp only [NNReal.coe_mul, NNReal.coe_mk]
ring

end Kernel.HasSubgaussianMGF

namespace HasSubgaussianMGF

variable {Ω : Type*} {m mΩ : MeasurableSpace Ω} {μ : Measure Ω} {X : Ω → ℝ} {c : ℝ≥0}

lemma id_map_iff (hX : AEMeasurable X μ) :
HasSubgaussianMGF X c μ ↔ HasSubgaussianMGF id c (μ.map X) := by
refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩
· constructor
· intro t
rw [integrable_map_measure (by fun_prop) hX]
exact h.integrable_exp_mul t
· intro t
rw [mgf_id_map hX]
exact h.mgf_le t
· have : X = id ∘ X := rfl
rw [this]
exact .of_map hX h

lemma congr_identDistrib {Ω' : Type*} {mΩ' : MeasurableSpace Ω'} {μ' : Measure Ω'}
{Y : Ω' → ℝ} (hX : HasSubgaussianMGF X c μ) (hXY : IdentDistrib X Y μ μ') :
HasSubgaussianMGF Y c μ' := by
rw [id_map_iff hXY.aemeasurable_fst] at hX
rwa [id_map_iff hXY.aemeasurable_snd, ← hXY.map_eq]

protected lemma const_mul (h : HasSubgaussianMGF X c μ) (r : ℝ) :
HasSubgaussianMGF (fun ω ↦ r * X ω) (⟨r ^ 2, sq_nonneg r⟩ * c) μ := by
rw [HasSubgaussianMGF_iff_kernel] at h ⊢
exact Kernel.HasSubgaussianMGF.const_mul h r

lemma sub_of_indepFun {Y : Ω → ℝ} {cX cY : ℝ≥0} (hX : HasSubgaussianMGF X cX μ)
(hY : HasSubgaussianMGF Y cY μ) (hindep : IndepFun X Y μ) :
HasSubgaussianMGF (fun ω ↦ X ω - Y ω) (cX + cY) μ := by
simp_rw [sub_eq_add_neg]
exact hX.add_of_indepFun hY.neg hindep.neg_right

end HasSubgaussianMGF

end ProbabilityTheory
39 changes: 34 additions & 5 deletions LeanBandits/Regret.lean
Original file line number Diff line number Diff line change
Expand Up @@ -61,9 +61,18 @@ lemma pullCount_eq_pullCount_add_one (t : ℕ) (h : ℕ → α × ℝ) :
lemma pullCount_eq_pullCount (ha : arm t h ≠ a) : pullCount a (t + 1) h = pullCount a t h := by
simp [pullCount, range_succ, filter_insert, ha]

lemma pullCount_add_one :
pullCount a (t + 1) h = pullCount a t h + if arm t h = a then 1 else 0 := by
split_ifs with h
· rw [← h, pullCount_eq_pullCount_add_one]
· rw [pullCount_eq_pullCount h, add_zero]

lemma pullCount_eq_sum (a : α) (t : ℕ) (h : ℕ → α × ℝ) :
pullCount a t h = ∑ s ∈ range t, if arm s h = a then 1 else 0 := by simp [pullCount]

lemma pullCount_le (a : α) (t : ℕ) (h : ℕ → α × ℝ) : pullCount a t h ≤ t :=
(card_filter_le _ _).trans_eq (by simp)

/-- 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 Down Expand Up @@ -190,6 +199,23 @@ lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1)
pullCount a (stepsUntil a m h).toNat h = m - 1 := by
sorry

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

/-- Empirical mean reward obtained when pulling arm `a` up to time `t` (exclusive). -/
noncomputable
def empMean (a : α) (t : ℕ) (h : ℕ → α × ℝ) : ℝ := sumRewards a t h / pullCount a t h

lemma sumRewards_eq_pullCount_mul_empMean (h_pull : pullCount a t h ≠ 0) :
sumRewards a t h = pullCount a t h * empMean a t h := by unfold empMean; field_simp

end SumRewards

section RewardByCount

/-- Reward obtained when pulling arm `a` for the `m`-th time. -/
noncomputable
def rewardByCount (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : ℝ :=
Expand All @@ -215,17 +241,18 @@ lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (h : ℕ → α × ℝ
rewardByCount (arm t h) (pullCount (arm t h) t h + 1) h z = reward t h := by
rw [rewardByCount, ← pullCount_eq_pullCount_add_one, stepsUntil_pullCount_eq]

lemma sum_rewardByCount_eq_sum_reward
lemma sum_rewardByCount_eq_sumRewards
(a : α) (t : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) :
∑ m ∈ Icc 1 (pullCount a t h), rewardByCount a m h z =
∑ s ∈ range t, if (arm s h) = a then (reward s h) else 0 := by
∑ m ∈ Icc 1 (pullCount a t h), rewardByCount a m h z = sumRewards a t h := by
induction' t with t ht
· simp [pullCount]
· simp [pullCount, sumRewards]
by_cases hta : arm t h = a
· rw [← hta] at ht ⊢
rw [pullCount_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht]
unfold sumRewards
rw [sum_range_succ, if_pos rfl, rewardByCount_pullCount_add_one_eq_reward]
· rwa [pullCount_eq_pullCount hta, sum_range_succ, if_neg hta, add_zero]
· unfold sumRewards
rwa [pullCount_eq_pullCount hta, sum_range_succ, if_neg hta, add_zero]

lemma sum_pullCount_mul [Fintype α] (h : ℕ → α × ℝ) (f : α → ℝ) (t : ℕ) :
∑ a, pullCount a t h * f a = ∑ s ∈ range t, f (arm s h) := by
Expand All @@ -246,6 +273,8 @@ lemma regret_eq_sum_pullCount_mul_gap [Fintype α] :
simp_rw [sum_pullCount_mul, regret, gap, sum_sub_distrib]
simp

end RewardByCount

section BestArm

variable [Fintype α] [Nonempty α]
Expand Down
8 changes: 8 additions & 0 deletions LeanBandits/RewardByCountMeasure.lean
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,14 @@ lemma measurable_pullCount (a : α) (t : ℕ) : Measurable (fun h ↦ pullCount
exact (measurableSet_singleton _).preimage (by fun_prop)
fun_prop

@[fun_prop]
lemma measurable_sumRewards (a : α) (t : ℕ) : Measurable (sumRewards a t) := by
unfold sumRewards
have h_meas s : Measurable (fun h : ℕ → α × ℝ ↦ if arm s h = a then reward s h else 0) := by
refine Measurable.ite ?_ (by fun_prop) (by fun_prop)
exact (measurableSet_singleton _).preimage (by fun_prop)
fun_prop

@[fun_prop]
lemma measurable_stepsUntil (a : α) (m : ℕ) : Measurable (fun h ↦ stepsUntil a m h) := by
classical
Expand Down
2 changes: 1 addition & 1 deletion blueprint/lean_decls
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ Bandits.iIndepFun_rewardByCount
Bandits.stepsUntil_pullCount_le
Bandits.stepsUntil_pullCount_eq
Bandits.rewardByCount_pullCount_add_one_eq_reward
Bandits.sum_rewardByCount_eq_sum_reward
Bandits.sum_rewardByCount_eq_sumRewards
Bandits.regret
Bandits.gap
Bandits.regret_eq_sum_pullCount_mul_gap
Expand Down
2 changes: 1 addition & 1 deletion blueprint/src/chapters/bandit.tex
Original file line number Diff line number Diff line change
Expand Up @@ -267,7 +267,7 @@ \section{Alternative model}\label{sec:alt_model}
\begin{lemma}\label{lem:sum_rewardByCount}
\uses{def:rewardByCount,def:pullCount}
\leanok
\lean{Bandits.sum_rewardByCount_eq_sum_reward}
\lean{Bandits.sum_rewardByCount_eq_sumRewards}
\begin{align*}
\sum_{n=1}^{N_{t, a}} Y_{n, a} = \sum_{s=0}^{t-1} \mathbb{I}\{A_s = a\} X_s
\: .
Expand Down