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 @@ -10,6 +10,7 @@ import LeanBandits.ForMathlib.CondIndepFun
import LeanBandits.ForMathlib.HasCondDistrib
import LeanBandits.ForMathlib.IndepFun
import LeanBandits.ForMathlib.IndepInfinitePi
import LeanBandits.ForMathlib.Integrable
import LeanBandits.ForMathlib.KernelRepresentation
import LeanBandits.ForMathlib.KernelSub
import LeanBandits.ForMathlib.Measurable
Expand Down
11 changes: 1 addition & 10 deletions LeanBandits/Bandit/Bandit.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.ForMathlib.CondIndepFun
import LeanBandits.ForMathlib.IndepFun
import LeanBandits.ForMathlib.IndepInfinitePi
import LeanBandits.ForMathlib.Integrable
import LeanBandits.ForMathlib.KernelRepresentation
import LeanBandits.ForMathlib.StandardBorel
import LeanBandits.SequentialLearning.FiniteActions
Expand Down Expand Up @@ -59,16 +60,6 @@ lemma identDistrib_eval_eval_id_streamMeasure (ν : Kernel α R) [IsMarkovKernel
Measure.map_map (by fun_prop) (by fun_prop)]
simp

lemma Integrable.congr_identDistrib {Ω Ω' : Type*}
{mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'}
{μ : Measure Ω} {μ' : Measure Ω'} {X : Ω → ℝ} {Y : Ω' → ℝ}
(hX : Integrable X μ) (hXY : IdentDistrib X Y μ μ') :
Integrable Y μ' := by
have hX' : Integrable id (μ.map X) := by
rwa [integrable_map_measure (by fun_prop) hXY.aemeasurable_fst]
rw [hXY.map_eq] at hX'
rwa [integrable_map_measure (by fun_prop) hXY.aemeasurable_snd] at hX'

lemma integrable_eval_streamMeasure (ν : Kernel α ℝ) [IsMarkovKernel ν] (n : ℕ) (a : α)
(h_int : Integrable id (ν a)) :
Integrable (fun h : ℕ → α → ℝ ↦ h n a) (streamMeasure ν) :=
Expand Down
27 changes: 27 additions & 0 deletions LeanBandits/Bandit/Regret.lean
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,33 @@ lemma regret_eq_sum_pullCount_mul_gap [Fintype α] :
regret ν A t ω = ∑ a, pullCount A a t ω * gap ν a := by
simp_rw [regret_eq_sum_gap, sum_pullCount_mul]

lemma integral_regret_eq_sum_gap_mul_integral_pullCount
[StandardBorelSpace α] [Fintype α] {P : Measure Ω} [IsProbabilityMeasure P]
(hA : ∀ n, Measurable (A n)) :
P[regret ν A n] = ∑ a, gap ν a * P[fun ω ↦ (pullCount A a n ω : ℝ)] := by
simp_rw [regret_eq_sum_pullCount_mul_gap]
rw [integral_finset_sum]
swap; · exact fun i _ ↦ (integrable_pullCount hA i n).mul_const _
congr with a
rw [integral_mul_const, mul_comm]

/-- To bound the expected regret, it suffices to bound the expected number of pulls for each action
with positive gap. -/
lemma integral_regret_le_of_forall_integral_pullCount_le
[Nonempty α] [StandardBorelSpace α] [Fintype α] {P : Measure Ω} [IsProbabilityMeasure P]
{alg : Algorithm α ℝ} {env : Environment α ℝ} {B : α → ℝ}
(h : IsAlgEnvSeq A R alg env P)
(h_le : ∀ a, gap ν a ≠ 0 → ∫ ω, (pullCount A a n ω : ℝ) ∂P ≤ B a) :
P[regret ν A n] ≤ ∑ a, gap ν a * B a := by
have hA := h.measurable_A
rw [integral_regret_eq_sum_gap_mul_integral_pullCount hA]
gcongr 1 with a
by_cases h_gap : gap ν a = 0
· simp [h_gap]
gcongr
· exact gap_nonneg
· exact h_le a h_gap

section bestArm

variable [Fintype α] [Nonempty α]
Expand Down
70 changes: 35 additions & 35 deletions LeanBandits/Bandit/RewardByCountMeasure.lean
Original file line number Diff line number Diff line change
Expand Up @@ -20,14 +20,14 @@ variable {α Ω : Type*} {mα : MeasurableSpace α} {mΩ : MeasurableSpace Ω} [
{alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν]
{h_inter : IsAlgEnvSeq A R alg (stationaryEnv ν) P}

local notation "𝔓'" => P.prod (streamMeasure ν)
local notation "𝔓" => P.prod (streamMeasure ν)

omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in
lemma hasLaw_Z (a : α) (m : ℕ) :
HasLaw (fun ω ↦ ω.2 m a) (ν a) 𝔓' where
HasLaw (fun ω ↦ ω.2 m a) (ν a) 𝔓 where
map_eq := by
calc (𝔓').map (fun ω ↦ ω.2 m a)
_ = ((𝔓').snd).map (fun ω ↦ ω m a) := by
calc (𝔓).map (fun ω ↦ ω.2 m a)
_ = ((𝔓).snd).map (fun ω ↦ ω m a) := by
rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)]
rfl
_ = (streamMeasure ν).map (fun ω ↦ ω m a) := by simp
Expand All @@ -47,15 +47,15 @@ notation "𝓛[" Y " | " X " ← " x "; " μ "]" => Measure.map Y (μ[|X ⁻¹'
omit [DecidableEq α] in
lemma condDistrib_reward'' [Countable α]
(h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (n : ℕ) :
𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1; 𝔓'] =ᵐ[(𝔓').map (fun ω ↦ A n ω.1)] ν := by
𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1; 𝔓] =ᵐ[(𝔓).map (fun ω ↦ A n ω.1)] ν := by
have hA := h.measurable_A
have hR := h.measurable_R
have h_ra' : 𝓛[R n | A n; P] =ᵐ[P.map (A n)] ν := h.condDistrib_reward_stationaryEnv n
have h_law : (𝔓').map (fun ω ↦ A n ω.1) = P.map (A n) := by
change ((𝔓').map (A n ∘ Prod.fst)) = _
have h_law : (𝔓).map (fun ω ↦ A n ω.1) = P.map (A n) := by
change ((𝔓).map (A n ∘ Prod.fst)) = _
rw [← Measure.map_map (by fun_prop) (by fun_prop), ← Measure.fst, Measure.fst_prod]
rw [h_law]
have h_prod : 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1; 𝔓']
have h_prod : 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1; 𝔓]
=ᵐ[P.map (A n)] 𝓛[R n | A n; P] :=
condDistrib_fst_prod _ (by fun_prop) _
filter_upwards [h_ra', h_prod] with ω h_eq h_prod
Expand All @@ -64,13 +64,13 @@ lemma condDistrib_reward'' [Countable α]
omit [DecidableEq α] in
lemma reward_cond_action [Countable α]
(h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n : ℕ)
(hμa : (𝔓').map (fun ω ↦ A n ω.1) {a} ≠ 0) :
𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1 ← a; 𝔓'] = ν a := by
(hμa : (𝔓).map (fun ω ↦ A n ω.1) {a} ≠ 0) :
𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1 ← a; 𝔓] = ν a := by
have hA := h.measurable_A
have hR := h.measurable_R
have h_ra : 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1; 𝔓'] =ᵐ[(𝔓').map (fun ω ↦ A n ω.1)] ν :=
have h_ra : 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1; 𝔓] =ᵐ[(𝔓).map (fun ω ↦ A n ω.1)] ν :=
condDistrib_reward'' h n
have h_eq := condDistrib_ae_eq_cond (μ := 𝔓')
have h_eq := condDistrib_ae_eq_cond (μ := 𝔓)
(X := fun ω ↦ A n ω.1) (Y := fun ω ↦ R n ω.1) (by fun_prop) (by fun_prop)
rw [Filter.EventuallyEq, ae_iff_of_countable] at h_ra h_eq
specialize h_ra a hμa
Expand Down Expand Up @@ -101,7 +101,7 @@ lemma condIndepFun_reward_stepsUntil_action [StandardBorelSpace Ω] [Countable
(h : IsAlgEnvSeq A R alg (stationaryEnv ν) P)
(a : α) (m n : ℕ) :
CondIndepFun (mα.comap (fun ω ↦ A n ω.1)) ((h.measurable_A n).comp measurable_fst).comap_le
(fun ω ↦ R n ω.1) ({ω | stepsUntil A a m ω.1 = ↑n}.indicator (fun _ ↦ 1)) 𝔓' := by
(fun ω ↦ R n ω.1) ({ω | stepsUntil A a m ω.1 = ↑n}.indicator (fun _ ↦ 1)) 𝔓 := by
have hA := h.measurable_A
have hR := h.measurable_R
exact condIndepFun_fst_prod (ν := streamMeasure ν)
Expand All @@ -110,37 +110,37 @@ lemma condIndepFun_reward_stepsUntil_action [StandardBorelSpace Ω] [Countable

lemma reward_cond_stepsUntil [StandardBorelSpace Ω] [Countable α]
(h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m n : ℕ)
(hm : m ≠ 0) (hμn : 𝔓' ((fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n}) ≠ 0) :
𝓛[fun ω ↦ R n ω.1 | fun ω ↦ stepsUntil A a m ω.1 ← ↑n; 𝔓'] = ν a := by
(hm : m ≠ 0) (hμn : 𝔓 ((fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n}) ≠ 0) :
𝓛[fun ω ↦ R n ω.1 | fun ω ↦ stepsUntil A a m ω.1 ← ↑n; 𝔓] = ν a := by
have hA := h.measurable_A
have hR := h.measurable_R
have hμna :
𝔓' ((fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ A n ω.1) ⁻¹' {a}) ≠ 0 := by
𝔓 ((fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ A n ω.1) ⁻¹' {a}) ≠ 0 := by
suffices ((fun ω : Ω × (ℕ → α → ℝ) ↦
stepsUntil A a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ A n ω.1) ⁻¹' {a})
= (fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n} by simpa [this] using hμn
ext ω
simp only [Set.mem_inter_iff, Set.mem_preimage, Set.mem_singleton_iff, and_iff_left_iff_imp]
exact action_eq_of_stepsUntil_eq_coe hm
have hμa : (𝔓').map (fun ω ↦ A n ω.1) {a} ≠ 0 := by
have hμa : (𝔓).map (fun ω ↦ A n ω.1) {a} ≠ 0 := by
rw [Measure.map_apply (by fun_prop) (measurableSet_singleton _)]
refine fun h_zero ↦ hμn (measure_mono_null (fun ω ↦ ?_) h_zero)
simp only [Set.mem_preimage, Set.mem_singleton_iff]
exact action_eq_of_stepsUntil_eq_coe hm
calc 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ stepsUntil A a m ω.1 ← (n : ℕ∞); 𝔓']
_ = (𝔓'[|(fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ A n ω.1) ⁻¹' {a}]).map
calc 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ stepsUntil A a m ω.1 ← (n : ℕ∞); 𝔓]
_ = (𝔓[|(fun ω ↦ stepsUntil A a m ω.1) ⁻¹' {↑n} ∩ (fun ω ↦ A n ω.1) ⁻¹' {a}]).map
(fun ω ↦ R n ω.1) := by
congr with ω
simp only [Set.mem_preimage, Set.mem_singleton_iff, Set.mem_inter_iff, iff_self_and]
exact action_eq_of_stepsUntil_eq_coe hm
_ = (𝔓'[|(fun ω ↦ A n ω.1) ⁻¹' {a}
_ = (𝔓[|(fun ω ↦ A n ω.1) ⁻¹' {a}
∩ {ω : Ω × (ℕ → α → ℝ) | stepsUntil A a m ω.1 = ↑n}.indicator 1 ⁻¹' {1} ]).map
(fun ω ↦ R n ω.1) := by
congr 2 with ω
simp only [Set.mem_inter_iff, Set.mem_preimage, Set.mem_singleton_iff, Set.indicator_apply,
Set.mem_setOf_eq, Pi.one_apply, ite_eq_left_iff, zero_ne_one, imp_false, Decidable.not_not]
rw [and_comm]
_ = 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1 ← a; 𝔓'] := by
_ = 𝓛[fun ω ↦ R n ω.1 | fun ω ↦ A n ω.1 ← a; 𝔓] := by
rw [cond_of_condIndepFun (by fun_prop)]
· exact condIndepFun_reward_stepsUntil_action h a m n
· refine measurable_one.indicator ?_
Expand All @@ -156,11 +156,11 @@ lemma reward_cond_stepsUntil [StandardBorelSpace Ω] [Countable α]
given the time at which number of pulls is `m` is the constant kernel with value `ν a`. -/
theorem condDistrib_rewardByCount_stepsUntil [StandardBorelSpace Ω] [Countable α]
(h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m : ℕ) (hm : m ≠ 0) :
condDistrib (rewardByCount A R a m) (fun ω ↦ stepsUntil A a m ω.1) 𝔓'
=ᵐ[(𝔓').map (fun ω ↦ stepsUntil A a m ω.1)] Kernel.const _ (ν a) := by
condDistrib (rewardByCount A R a m) (fun ω ↦ stepsUntil A a m ω.1) 𝔓
=ᵐ[(𝔓).map (fun ω ↦ stepsUntil A a m ω.1)] Kernel.const _ (ν a) := by
have hA := h.measurable_A
have hR := h.measurable_R
refine (condDistrib_ae_eq_cond (μ := 𝔓')
refine (condDistrib_ae_eq_cond (μ := 𝔓)
(X := fun ω ↦ stepsUntil A a m ω.1) (by fun_prop) (by fun_prop)).trans ?_
rw [Filter.EventuallyEq, ae_iff_of_countable]
intro n hn
Expand Down Expand Up @@ -189,44 +189,44 @@ theorem condDistrib_rewardByCount_stepsUntil [StandardBorelSpace Ω] [Countable
/-- The reward received at the `m`-th pull of action `a` has law `ν a`. -/
lemma hasLaw_rewardByCount [StandardBorelSpace Ω] [Countable α]
(h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (m : ℕ) (hm : m ≠ 0) :
HasLaw (rewardByCount A R a m) (ν a) 𝔓' where
HasLaw (rewardByCount A R a m) (ν a) 𝔓 where
aemeasurable := (measurable_rewardByCount h.measurable_A h.measurable_R a m).aemeasurable
map_eq := by
have hA := h.measurable_A
have hR := h.measurable_R
have h_condDistrib :
condDistrib (rewardByCount A R a m) (fun ω ↦ stepsUntil A a m ω.1) 𝔓'
=ᵐ[(𝔓').map (fun ω ↦ stepsUntil A a m ω.1)]
condDistrib (rewardByCount A R a m) (fun ω ↦ stepsUntil A a m ω.1) 𝔓
=ᵐ[(𝔓).map (fun ω ↦ stepsUntil A a m ω.1)]
Kernel.const _ (ν a) := condDistrib_rewardByCount_stepsUntil h a m hm
calc (𝔓').map (rewardByCount A R a m)
_ = (condDistrib (rewardByCount A R a m) (fun ω ↦ stepsUntil A a m ω.1) 𝔓')
∘ₘ ((𝔓').map (fun ω ↦ stepsUntil A a m ω.1)) := by
calc (𝔓).map (rewardByCount A R a m)
_ = (condDistrib (rewardByCount A R a m) (fun ω ↦ stepsUntil A a m ω.1) 𝔓)
∘ₘ ((𝔓).map (fun ω ↦ stepsUntil A a m ω.1)) := by
rw [condDistrib_comp_map (by fun_prop) (by fun_prop)]
_ = (Kernel.const _ (ν a)) ∘ₘ ((𝔓').map (fun ω ↦ stepsUntil A a m ω.1)) :=
_ = (Kernel.const _ (ν a)) ∘ₘ ((𝔓).map (fun ω ↦ stepsUntil A a m ω.1)) :=
Measure.comp_congr h_condDistrib
_ = ν a := by
have : IsProbabilityMeasure ((𝔓').map (fun ω ↦ stepsUntil A a m ω.1)) :=
have : IsProbabilityMeasure ((𝔓).map (fun ω ↦ stepsUntil A a m ω.1)) :=
Measure.isProbabilityMeasure_map (by fun_prop)
simp

lemma identDistrib_rewardByCount [StandardBorelSpace Ω] [Countable α]
(h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n m : ℕ)
(hn : n ≠ 0) (hm : m ≠ 0) :
IdentDistrib (rewardByCount A R a n) (rewardByCount A R a m) 𝔓' 𝔓' where
IdentDistrib (rewardByCount A R a n) (rewardByCount A R a m) 𝔓 𝔓 where
aemeasurable_fst := (measurable_rewardByCount h.measurable_A h.measurable_R a n).aemeasurable
aemeasurable_snd := (measurable_rewardByCount h.measurable_A h.measurable_R a m).aemeasurable
map_eq := by rw [(hasLaw_rewardByCount h a n hn).map_eq, (hasLaw_rewardByCount h a m hm).map_eq]

lemma identDistrib_rewardByCount_id [StandardBorelSpace Ω] [Countable α]
(h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n : ℕ) (hn : n ≠ 0) :
IdentDistrib (rewardByCount A R a n) id 𝔓' (ν a) where
IdentDistrib (rewardByCount A R a n) id 𝔓 (ν a) where
aemeasurable_fst := (measurable_rewardByCount h.measurable_A h.measurable_R a n).aemeasurable
aemeasurable_snd := Measurable.aemeasurable <| by fun_prop
map_eq := by rw [(hasLaw_rewardByCount h a n hn).map_eq, Measure.map_id]

lemma identDistrib_rewardByCount_eval [StandardBorelSpace Ω] [Countable α]
(h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n m : ℕ) (hn : n ≠ 0) :
IdentDistrib (rewardByCount A R a n) (fun ω ↦ ω m a) 𝔓' (streamMeasure ν) :=
IdentDistrib (rewardByCount A R a n) (fun ω ↦ ω m a) 𝔓 (streamMeasure ν) :=
(identDistrib_rewardByCount_id h a n hn).trans
(identDistrib_eval_eval_id_streamMeasure ν m a).symm

Expand Down
82 changes: 54 additions & 28 deletions LeanBandits/Bandit/SumRewards.lean
Original file line number Diff line number Diff line change
Expand Up @@ -13,34 +13,6 @@ import LeanBandits.ForMathlib.SubGaussian
open MeasureTheory ProbabilityTheory Finset Learning
open scoped ENNReal NNReal

lemma measurable_sum_range_of_le {α : Type*} {mα : MeasurableSpace α}
{f : ℕ → α → ℝ} {g : α → ℕ} {n : ℕ} (hg_le : ∀ a, g a ≤ n) (hf : ∀ i, Measurable (f i))
(hg : Measurable g) :
Measurable (fun a ↦ ∑ i ∈ range (g a), f i a) := by
have h_eq : (fun a ↦ ∑ i ∈ range (g a), f i a)
= fun a ↦ ∑ i ∈ range (n + 1), if g a = i then ∑ j ∈ range i, f j a else 0 := by
ext ω
rw [sum_ite_eq_of_mem]
grind
rw [h_eq]
refine measurable_sum _ fun n hn ↦ ?_
refine Measurable.ite ?_ (by fun_prop) (by fun_prop)
exact (measurableSet_singleton _).preimage (by fun_prop)

lemma measurable_sum_Icc_of_le {α : Type*} {mα : MeasurableSpace α}
{f : ℕ → α → ℝ} {g : α → ℕ} {n : ℕ} (hg_le : ∀ a, g a ≤ n) (hf : ∀ i, Measurable (f i))
(hg : Measurable g) :
Measurable (fun a ↦ ∑ i ∈ Icc 1 (g a), f i a) := by
have h_eq : (fun a ↦ ∑ i ∈ Icc 1 (g a), f i a)
= fun a ↦ ∑ i ∈ range (n + 1), if g a = i then ∑ j ∈ Icc 1 i, f j a else 0 := by
ext ω
rw [sum_ite_eq_of_mem]
grind
rw [h_eq]
refine measurable_sum _ fun n hn ↦ ?_
refine Measurable.ite ?_ (by fun_prop) (by fun_prop)
exact (measurableSet_singleton _).preimage (by fun_prop)

namespace Bandits

namespace ArrayModel
Expand Down Expand Up @@ -603,6 +575,60 @@ lemma prob_sum_ge_sqrt_log {σ2 : ℝ≥0}
← ENNReal.ofReal_rpow_of_nonneg (by positivity) (by positivity)]
norm_cast

open Real

omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in
lemma todo {σ2 : ℝ≥0} {c : ℝ}
(hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (hσ2 : σ2 ≠ 0)
(hc : 0 ≤ c) (a : α) (n k : ℕ) (hk : k ≠ 0) :
streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(2 * c * σ2 * log (n + 1) / k) ≤ (ν a)[id]} ≤
1 / (n + 1) ^ c := by
have h_log_nonneg : 0 ≤ log (n + 1) := log_nonneg (by simp)
calc
streamMeasure ν {ω | (∑ m ∈ range k, ω m a) / k + √(2 * c * σ2 * log (n + 1) / k) ≤ (ν a)[id]}
_ = streamMeasure ν
{ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) / k ≤ - √(2 * c * σ2 * log (n + 1) / k)} := by
congr with ω
field_simp
rw [Finset.sum_sub_distrib]
simp
grind
_ = streamMeasure ν
{ω | (∑ s ∈ range k, (ω s a - (ν a)[id])) ≤ - √(2 * c * k * σ2 * log (n + 1))} := by
congr with ω
field_simp
congr! 2
rw [sqrt_div (by positivity), ← mul_div_assoc, mul_comm, mul_div_assoc, div_sqrt,
mul_assoc (k : ℝ), mul_assoc (k : ℝ), mul_assoc (k : ℝ),
sqrt_mul (x := (k : ℝ)) (by positivity), mul_comm]
_ ≤ 1 / (n + 1) ^ c := prob_sum_le_sqrt_log hν hσ2 hc a k hk

omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in
lemma todo' {σ2 : ℝ≥0} {c : ℝ}
(hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[id]) σ2 (ν a)) (hσ2 : σ2 ≠ 0)
(hc : 0 ≤ c) (a : α) (n k : ℕ) (hk : k ≠ 0) :
streamMeasure ν
{ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - √(2 * c * σ2 *log (n + 1) / k)} ≤
1 / (n + 1) ^ c := by
have h_log_nonneg : 0 ≤ log (n + 1) := log_nonneg (by simp)
calc
streamMeasure ν {ω | (ν a)[id] ≤ (∑ m ∈ range k, ω m a) / k - √(2 * c * σ2 * log (n + 1) / k)}
_ = streamMeasure ν
{ω | √(2 * c * σ2 * log (n + 1) / k) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id])) / k} := by
congr with ω
field_simp
rw [Finset.sum_sub_distrib]
simp
grind
_ = streamMeasure ν
{ω | √(2 * c * k * σ2 * log (n + 1)) ≤ (∑ s ∈ range k, (ω s a - (ν a)[id]))} := by
congr with ω
field_simp
congr! 1
rw [sqrt_div (by positivity), ← mul_div_assoc, mul_comm, mul_div_assoc, div_sqrt,
mul_comm _ (k : ℝ), sqrt_mul (x := (k : ℝ)) (by positivity), mul_comm]
_ ≤ 1 / (n + 1) ^ c := prob_sum_ge_sqrt_log hν hσ2 hc a k hk

end Subgaussian

end Bandits
5 changes: 5 additions & 0 deletions LeanBandits/BanditAlgorithms/AuxSums.lean
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,11 @@ import Mathlib.Algebra.BigOperators.Intervals
import Mathlib.Algebra.BigOperators.Ring.Finset
import Mathlib.Tactic.Ring.RingNF

/-!
# Lemmas about sums of indicators

-/

open Finset

lemma sum_mod_range {K : ℕ} (hK : 0 < K) (a : Fin K) :
Expand Down
Loading