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 @@ -5,6 +5,7 @@ import LeanBandits.Bandit
import LeanBandits.ETC
import LeanBandits.ForMathlib.CondDistrib
import LeanBandits.ForMathlib.IdentDistrib
import LeanBandits.ForMathlib.IndepFun
import LeanBandits.ForMathlib.KernelSub
import LeanBandits.ForMathlib.Measurable
import LeanBandits.ForMathlib.SubGaussian
Expand Down
8 changes: 6 additions & 2 deletions LeanBandits/Algorithm.lean
Original file line number Diff line number Diff line change
Expand Up @@ -156,8 +156,12 @@ lemma measurable_hist_filtration (n : ℕ) : Measurable[Learning.filtration α R
simp [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe,
measurable_iff_comap_le]

-- todo: due to the type of `Adapted` and the fact that `Iic n → α × R` depends on `n`, we cannot
-- state that `hist` is adapted.
lemma adapted_hist [TopologicalSpace α] [TopologicalSpace.PseudoMetrizableSpace α]
[SecondCountableTopology α] [OpensMeasurableSpace α]
[TopologicalSpace R] [TopologicalSpace.PseudoMetrizableSpace R]
[SecondCountableTopology R] [OpensMeasurableSpace R] :
Adapted (Learning.filtration α R) hist :=
fun n ↦ (measurable_hist_filtration n).stronglyMeasurable

lemma measurable_action_filtration (n : ℕ) : Measurable[Learning.filtration α R n] (action n) := by
simp only [Learning.filtration, Filtration.piLE_eq_comap_frestrictLe, ← hist_eq_frestrictLe]
Expand Down
48 changes: 36 additions & 12 deletions LeanBandits/ETC.lean
Original file line number Diff line number Diff line change
Expand Up @@ -118,9 +118,42 @@ lemma arm_of_ge {n : ℕ} (hm : m ≠ 0) (hn : K * m ≤ n) : arm n =ᵐ[𝔓t]
| base => rfl
| succ n hmn h_ind => rw [h_ae n hmn, h_ind]

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
sorry
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]

lemma pullCount_mul (a : Fin K) : pullCount a (K * m) =ᵐ[𝔓t] fun _ ↦ m := by
rw [Filter.EventuallyEq]
Expand Down Expand Up @@ -262,6 +295,7 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i
let f₂ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦
∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω.1 ω.2
let g₂ := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ ∑ s ∈ Icc 1 m, rewardByCount a s ω.1 ω.2
change MeasurableSet {x | f₁ x ≤ g₁ x ↔ f₂ x ≤ g₂ x}
have hf₁ : Measurable f₁ := by
refine measurable_sum_of_le (n := K * m + 1)
(g := fun ω : (ℕ → Fin K × ℝ) × (ℕ → Fin K → ℝ) ↦ pullCount (bestArm ν) (K * m) ω.1)
Expand All @@ -275,18 +309,8 @@ lemma prob_arm_mul_eq_le (hν : ∀ a, HasSubgaussianMGF (fun x ↦ x - (ν a)[i
(f := fun s ω ↦ rewardByCount a s ω.1 ω.2) (fun ω ↦ ?_) (by fun_prop) (by fun_prop)
have h_le := pullCount_le a (K * m) ω.1
grind
change MeasurableSet {x | f₁ x ≤ g₁ x ↔ f₂ x ≤ g₂ x}
simp_rw [iff_def, imp_iff_not_or]
change MeasurableSet ({x | ¬f₁ x ≤ g₁ x ∨ f₂ x ≤ g₂ x} ∩ {x | ¬f₂ x ≤ g₂ x ∨ f₁ x ≤ g₁ x})
have h1 : {x | ¬f₁ x ≤ g₁ x ∨ f₂ x ≤ g₂ x} = {x | f₁ x ≤ g₁ x}ᶜ ∪ {x | f₂ x ≤ g₂ x} := by
ext; simp
have h2 : {x | ¬f₂ x ≤ g₂ x ∨ f₁ x ≤ g₁ x} = {x | f₂ x ≤ g₂ x}ᶜ ∪ {x | f₁ x ≤ g₁ x} := by
ext; simp
rw [h1, h2]
refine (MeasurableSet.union ?_ ?_).inter (MeasurableSet.union ?_ ?_)
· exact (measurableSet_le (by fun_prop) (by fun_prop)).compl
refine MeasurableSet.iff ?_ ?_
· exact measurableSet_le (by fun_prop) (by fun_prop)
· exact (measurableSet_le (by fun_prop) (by fun_prop)).compl
· exact measurableSet_le (by fun_prop) (by fun_prop)
_ = (𝔓).real {ω | ∑ s ∈ range m, ω.2 s (bestArm ν) ≤ ∑ s ∈ range m, ω.2 s a} := by
simp_rw [measureReal_def]
Expand Down
26 changes: 26 additions & 0 deletions LeanBandits/ForMathlib/IndepFun.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
import Mathlib

open MeasureTheory Finset

namespace ProbabilityTheory

variable {Ω E : Type*} {mΩ : MeasurableSpace Ω} {mE : MeasurableSpace E} {μ : Measure Ω}

lemma iIndepFun_nat_iff_forall_indepFun {X : ℕ → Ω → E} (hX : ∀ n, AEMeasurable (X n) μ) :
iIndepFun X μ ↔ ∀ n, X (n + 1) ⟂ᵢ[μ] fun ω (i : Iic n) ↦ X i ω := by
constructor
· intro h n
have h' := h.indepFun_finset₀ {n + 1} (Iic n) (by simp) hX
let f : (({n + 1} : Finset ℕ) → E) → E := fun x ↦ x ⟨n + 1, by simp⟩
have hf : Measurable f := by unfold f; fun_prop
have h_eq : X (n + 1) = f ∘ (fun ω (i : ({n + 1} : Finset ℕ)) ↦ X i ω) := rfl
exact h'.comp (ψ := id) hf measurable_id
· intro h
suffices ∀ n, iIndepFun ((Iic n).restrict X) μ by
rw [iIndepFun_iff_finset]
intro s
sorry
intro n
sorry

end ProbabilityTheory
19 changes: 19 additions & 0 deletions LeanBandits/ForMathlib/Measurable.lean
Original file line number Diff line number Diff line change
Expand Up @@ -19,4 +19,23 @@ lemma measurable_comp_comap {α β γ : Type*} {mβ : MeasurableSpace β} {mγ :
rw [← measurable_iff_comap_le]
exact hg

lemma MeasurableSet.imp {α : Type*} {mα : MeasurableSpace α} {p q : α → Prop}
(hs : MeasurableSet {x | p x}) (ht : MeasurableSet {x | q x}) :
MeasurableSet {x | p x → q x} := by
have h_eq : {x | p x → q x} = {x | p x}ᶜ ∪ {x | q x} := by
ext x
grind
rw [h_eq]
exact MeasurableSet.union hs.compl ht

lemma MeasurableSet.iff {α : Type*} {mα : MeasurableSpace α} {p q : α → Prop}
(hs : MeasurableSet {x | p x}) (ht : MeasurableSet {x | q x}) :
MeasurableSet {x | p x ↔ q x} := by
have h_eq : {x | p x ↔ q x} = ({x | p x}ᶜ ∪ {x | q x}) ∩ ({x | q x}ᶜ ∪ {x | p x}) := by
ext x
simp only [Set.mem_setOf_eq, Set.mem_inter_iff, Set.mem_union, Set.mem_compl_iff]
grind
rw [h_eq]
exact (MeasurableSet.union hs.compl ht).inter (MeasurableSet.union ht.compl hs)

end MeasureTheory
88 changes: 77 additions & 11 deletions LeanBandits/Regret.lean
Original file line number Diff line number Diff line change
Expand Up @@ -73,13 +73,31 @@ lemma pullCount_eq_sum (a : α) (t : ℕ) (h : ℕ → α × ℝ) :
lemma pullCount_le (a : α) (t : ℕ) (h : ℕ → α × ℝ) : pullCount a t h ≤ t :=
(card_filter_le _ _).trans_eq (by simp)

lemma pullCount_congr {h' : ℕ → α × ℝ} (h_eq : ∀ i ≤ n, arm i h = arm i h') :
pullCount a (n + 1) h = pullCount a (n + 1) h' := by
unfold pullCount
congr 1 with s
simp only [mem_filter, mem_range, and_congr_right_iff]
intro hs
rw [Nat.lt_add_one_iff] at hs
rw [h_eq s hs]

/-- 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})

lemma stepsUntil_eq_top_iff : stepsUntil a m h = ⊤ ↔ ∀ s, pullCount a (s + 1) h ≠ m := by
simp [stepsUntil, sInf_eq_top]

lemma stepsUntil_ne_top (h_exists : ∃ s, pullCount a (s + 1) h = m) : stepsUntil a m h ≠ ⊤ := by
simpa [stepsUntil_eq_top_iff]

lemma exists_pullCount_eq (h' : stepsUntil a m h ≠ ⊤) :
∃ s, pullCount a (s + 1) h = m := by
by_contra! h_contra
rw [← stepsUntil_eq_top_iff] at h_contra
simp [h_contra] at h'

lemma stepsUntil_zero_of_ne (hka : arm 0 h ≠ a) : stepsUntil a 0 h = 0 := by
unfold stepsUntil
simp_rw [← bot_eq_zero, sInf_eq_bot, bot_eq_zero]
Expand Down Expand Up @@ -139,10 +157,7 @@ lemma stepsUntil_eq_zero_iff :
stepsUntil a m h = 0 ↔ (m = 0 ∧ arm 0 h ≠ a) ∨ (m = 1 ∧ arm 0 h = a) := by
classical
refine ⟨fun h' ↦ ?_, fun h' ↦ ?_⟩
· have h_exists : ∃ s, pullCount a (s + 1) h = m := by
by_contra! h_contra
rw [← stepsUntil_eq_top_iff] at h_contra
simp [h_contra] at h'
· have h_exists : ∃ s, pullCount a (s + 1) h = m := exists_pullCount_eq (by simp [h'])
simp only [stepsUntil_eq_dite, h_exists, ↓reduceDIte, Nat.cast_eq_zero, Nat.find_eq_zero,
zero_add] at h'
rw [pullCount_one] at h'
Expand Down Expand Up @@ -183,13 +198,7 @@ lemma arm_eq_of_stepsUntil_eq_coe {ω : ℕ → α × ℝ} (hm : m ≠ 0)
arm n ω = a := by
have : n = (stepsUntil a m ω).toNat := by simp [h]
rw [this, arm_stepsUntil hm]
by_contra! h_contra
rw [← stepsUntil_eq_top_iff] at h_contra
simp [h_contra] at h

lemma stepsUntil_eq_congr {h' : ℕ → α × ℝ} (h_eq : ∀ i ≤ n, arm i h = arm i h') :
stepsUntil a m h = n ↔ stepsUntil a m h' = n := by
sorry
exact exists_pullCount_eq (by simp [h])

lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount a (s + 1) h = m) :
pullCount a (stepsUntil a m h + 1).toNat h = m := by
Expand All @@ -212,6 +221,63 @@ lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount a (s + 1)
swap; · simpa [stepsUntil_eq_top_iff]
grind

lemma pullCount_lt_of_le_stepsUntil (a : α) {n m : ℕ} (h : ℕ → α × ℝ)
(h_exists : ∃ s, pullCount a (s + 1) h = m) (hn : n < stepsUntil a m h) :
pullCount a (n + 1) h < m := by
classical
have h_eq := stepsUntil_eq_dite a m h
simp only [h_exists, ↓reduceDIte] at h_eq
rw [← ENat.coe_toNat (stepsUntil_ne_top h_exists)] at hn
refine lt_of_le_of_ne ?_ ?_
· calc pullCount a (n + 1) h
_ ≤ pullCount a (stepsUntil a m h + 1).toNat h := by
refine monotone_pullCount a h ?_
rw [ENat.toNat_add (stepsUntil_ne_top h_exists) (by simp)]
simp only [ENat.toNat_one, add_le_add_iff_right]
exact mod_cast hn.le
_ = m := pullCount_stepsUntil_add_one h_exists
· refine Nat.find_min h_exists (m := n) ?_
suffices n < (stepsUntil a m h).toNat by
rwa [h_eq, ENat.toNat_coe] at this
exact mod_cast hn

lemma pullCount_eq_of_stepsUntil_eq_coe {ω : ℕ → α × ℝ} (hm : m ≠ 0)
(h : stepsUntil a m ω = n) :
pullCount a n ω = m - 1 := by
have : n = (stepsUntil a m ω).toNat := by simp [h]
rw [this, pullCount_stepsUntil hm]
exact exists_pullCount_eq (by simp [h])

lemma pullCount_add_one_eq_of_stepsUntil_eq_coe {ω : ℕ → α × ℝ}
(h : stepsUntil a m ω = n) :
pullCount a (n + 1) ω = m := by
have : n + 1 = (stepsUntil a m ω + 1).toNat := by
rw [ENat.toNat_add (by simp [h]) (by simp)]; simp [h]
rw [this, pullCount_stepsUntil_add_one]
exact exists_pullCount_eq (by simp [h])

lemma stepsUntil_eq_iff {ω : ℕ → α × ℝ} (n : ℕ) :
stepsUntil a m ω = n ↔
pullCount a (n + 1) ω = m ∧ (∀ k < n, pullCount a (k + 1) ω < m) := by
refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩
· have h_exists : ∃ s, pullCount a (s + 1) ω = m := exists_pullCount_eq (by simp [h])
refine ⟨pullCount_add_one_eq_of_stepsUntil_eq_coe h, fun k hk ↦ ?_⟩
exact pullCount_lt_of_le_stepsUntil a ω h_exists (by rw [h]; exact mod_cast hk)
· classical
rw [stepsUntil_eq_dite a m ω, dif_pos ⟨n, h.1⟩]
simp only [Nat.cast_inj]
rw [Nat.find_eq_iff]
exact ⟨h.1, fun k hk ↦ (h.2 k hk).ne⟩

lemma stepsUntil_eq_congr {h' : ℕ → α × ℝ} (h_eq : ∀ i ≤ n, arm i h = arm i h') :
stepsUntil a m h = n ↔ stepsUntil a m h' = n := by
simp_rw [stepsUntil_eq_iff n]
congr! 1
· rw [pullCount_congr h_eq]
· congr! 3 with k hk
rw [pullCount_congr]
grind

section SumRewards

/-- Sum of rewards obtained when pulling arm `a` up to time `t` (exclusive). -/
Expand Down
10 changes: 9 additions & 1 deletion LeanBandits/RewardByCountMeasure.lean
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ Authors: Rémy Degenne
-/
import LeanBandits.Bandit
import LeanBandits.ForMathlib.IdentDistrib
import LeanBandits.ForMathlib.IndepFun
import LeanBandits.Regret

/-! # Laws of `stepsUntil` and `rewardByCount`
Expand Down Expand Up @@ -343,9 +344,16 @@ lemma identDistrib_rewardByCount_eval [Countable α] [StandardBorelSpace α] [No
(Bandit.measure alg ν) (Bandit.streamMeasure ν) :=
(identDistrib_rewardByCount_id a n hn).trans (identDistrib_eval_eval_id_streamMeasure ν m a).symm

lemma indepFun_rewardByCount_Iic (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (a : α)
(n : ℕ) :
(fun ω ↦ rewardByCount a (n + 1) ω.1 ω.2) ⟂ᵢ[Bandit.measure alg ν]
fun ω (i : Iic n) ↦ rewardByCount a i ω.1 ω.2 := by
sorry

lemma iIndepFun_rewardByCount' (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] (a : α) :
iIndepFun (fun n ω ↦ rewardByCount a n ω.1 ω.2) (Bandit.measure alg ν) := by
sorry
rw [iIndepFun_nat_iff_forall_indepFun (by fun_prop)]
exact indepFun_rewardByCount_Iic alg ν a

lemma iIndepFun_rewardByCount (alg : Algorithm α ℝ) (ν : Kernel α ℝ) [IsMarkovKernel ν] :
iIndepFun (fun (p : α × ℕ) ω ↦ rewardByCount p.1 p.2 ω.1 ω.2) (Bandit.measure alg ν) := by
Expand Down
16 changes: 8 additions & 8 deletions lake-manifest.json
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
"type": "git",
"subDir": null,
"scope": "",
"rev": "f07bd0325121718862be33645a23c4e55791271a",
"rev": "98b14016adbf3d90a2cc79399e49b9ee67b4155c",
"name": "mathlib",
"manifestFile": "lake-manifest.json",
"inputRev": null,
Expand All @@ -25,7 +25,7 @@
"type": "git",
"subDir": null,
"scope": "leanprover-community",
"rev": "7607162f5a1c1eb23c23027629a418b3a160670e",
"rev": "8864a73bf79aad549e34eff972c606343935106d",
"name": "plausible",
"manifestFile": "lake-manifest.json",
"inputRev": "main",
Expand All @@ -45,7 +45,7 @@
"type": "git",
"subDir": null,
"scope": "leanprover-community",
"rev": "e5c37730d22634ee0169c164f25dac49918ed951",
"rev": "451499ea6e97cee4c8979b507a9af5581a849161",
"name": "importGraph",
"manifestFile": "lake-manifest.json",
"inputRev": "main",
Expand All @@ -65,7 +65,7 @@
"type": "git",
"subDir": null,
"scope": "leanprover-community",
"rev": "cbe864cd5177966c9e005418cfdc1fb36db62e13",
"rev": "1fa48c6a63b4c4cda28be61e1037192776e77ac0",
"name": "aesop",
"manifestFile": "lake-manifest.json",
"inputRev": "master",
Expand All @@ -75,7 +75,7 @@
"type": "git",
"subDir": null,
"scope": "leanprover-community",
"rev": "593aa51c4aa07ee81e9233b53e1f61a5b4d9f761",
"rev": "95c2f8afe09d9e49d3cacca667261da04f7f93f7",
"name": "Qq",
"manifestFile": "lake-manifest.json",
"inputRev": "master",
Expand All @@ -85,7 +85,7 @@
"type": "git",
"subDir": null,
"scope": "leanprover-community",
"rev": "5bd478197f2e5d2a4fde527cf3581d83f49baa9b",
"rev": "c44068fa1b40041e6df42bd67639b690eb2764ca",
"name": "batteries",
"manifestFile": "lake-manifest.json",
"inputRev": "main",
Expand All @@ -95,10 +95,10 @@
"type": "git",
"subDir": null,
"scope": "leanprover",
"rev": "f75f4926aff7ba19949e16c19094d7298806b1a6",
"rev": "72ae7004d9f0ddb422aec5378204fdd7828c5672",
"name": "Cli",
"manifestFile": "lake-manifest.json",
"inputRev": "v4.25.0-rc1",
"inputRev": "v4.25.0-rc2",
"inherited": true,
"configFile": "lakefile.toml"}],
"name": "LeanBandits",
Expand Down
2 changes: 1 addition & 1 deletion lean-toolchain
Original file line number Diff line number Diff line change
@@ -1 +1 @@
leanprover/lean4:v4.25.0-rc1
leanprover/lean4:v4.25.0-rc2