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
16 changes: 6 additions & 10 deletions LeanBandits/RewardByCountMeasure.lean
Original file line number Diff line number Diff line change
Expand Up @@ -31,12 +31,11 @@ omit [DecidableEq α] [MeasurableSingletonClass α] in
lemma hasLaw_Z (a : α) (m : ℕ) :
HasLaw (fun ω ↦ ω.2 m a) (ν a) (Bandit.measure alg ν) where
map_eq := by
calc ((Bandit.trajMeasure alg ν).prod (Bandit.streamMeasure ν)).map (fun ω ↦ ω.2 m a)
_ = (((Bandit.trajMeasure alg ν).prod (Bandit.streamMeasure ν)).map (fun ω ↦ ω.2)).map
(fun ω ↦ ω m a) := by
rw [Measure.map_map (by fun_prop) (by fun_prop)]
calc (Bandit.measure alg ν).map (fun ω ↦ ω.2 m a)
_ = ((Bandit.measure alg ν).snd).map (fun ω ↦ ω m a) := by
rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)]
rfl
_ = (Bandit.streamMeasure ν).map (fun ω ↦ ω m a) := by simp [Measure.map_snd_prod]
_ = (Bandit.streamMeasure ν).map (fun ω ↦ ω m a) := by simp
_ = ((Measure.infinitePi fun _ ↦ Measure.infinitePi ν).map (fun ω ↦ ω m)).map
(fun ω ↦ ω a) := by
rw [Bandit.streamMeasure, Measure.map_map (by fun_prop) (by fun_prop)]
Expand All @@ -59,11 +58,8 @@ lemma condDistrib_reward'' [StandardBorelSpace α] [Nonempty α] (n : ℕ) :
=ᵐ[(𝔓).map (fun ω ↦ arm n ω.1)] ν := by
have h_ra' : 𝓛[reward n | arm n; 𝔓t] =ᵐ[(𝔓t).map (arm n)] ν := condDistrib_reward alg ν n
have h_law : (𝔓).map (fun ω ↦ arm n ω.1) = (𝔓t).map (arm n) := by
calc (𝔓).map (fun ω ↦ arm n ω.1)
_ = ((𝔓).map (fun ω ↦ ω.1)).map (fun ω ↦ arm n ω) := by
rw [Measure.map_map (by fun_prop) (by fun_prop)]
rfl
_ = _ := by unfold Bandit.measure; simp [Measure.map_fst_prod]
rw [← Bandit.fst_measure, Measure.fst, Measure.map_map (by fun_prop) (by fun_prop)]
rfl
rw [h_law]
have h_prod : 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1; 𝔓]
=ᵐ[(𝔓t).map (arm n)] 𝓛[reward n | arm n; 𝔓t] :=
Expand Down
69 changes: 61 additions & 8 deletions LeanBandits/SequentialLearning/FiniteActions.lean
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,8 @@ namespace Learning
variable {α R : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R} [DecidableEq α]
{a : α} {m n t : ℕ} {h : ℕ → α × R}

section PullCount

/-- Number of times action `a` was chosen up to time `t` (excluding `t`). -/
noncomputable
def pullCount (a : α) (t : ℕ) (h : ℕ → α × R) : ℕ :=
Expand Down Expand Up @@ -107,6 +109,24 @@ lemma pullCount_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h = a
rw [Nat.lt_add_one_iff] at hs
rw [h_eq s hs]

lemma pullCount_lt_of_forall_ne (h_lt : ∀ s, pullCount a (s + 1) h ≠ t) (ht : t ≠ 0) :
pullCount a n h < t := by
induction n with
| zero => simpa using ht.bot_lt
| succ n hn =>
specialize h_lt n
rw [pullCount_add_one] at h_lt ⊢
grind

lemma exists_pullCount_eq_of_le (hnm : t ≤ pullCount a (n + 1) h) (ht : t ≠ 0) :
∃ s, pullCount a (s + 1) h = t := by
by_contra! h_contra
refine lt_irrefl (pullCount a (n + 1) h) ?_
refine lt_of_lt_of_le ?_ hnm
exact pullCount_lt_of_forall_ne h_contra ht

section Measurability

@[fun_prop]
lemma measurable_pullCount [MeasurableSingletonClass α] (a : α) (t : ℕ) :
Measurable (fun h : ℕ → α × R ↦ pullCount a t h) := by
Expand Down Expand Up @@ -142,7 +162,13 @@ lemma isPredictable_pullCount [MeasurableSingletonClass α] (a : α) :
simp only [pullCount_zero]
fun_prop

-- TODO: replace this by leastGE
end Measurability

end PullCount

section StepsUntil

-- TODO: replace this by leastGE, once leastGE is generalized
/-- Number of steps until action `a` was pulled exactly `m` times. -/
noncomputable
def stepsUntil (a : α) (m : ℕ) (h : ℕ → α × R) : ℕ∞ := sInf ((↑) '' {s | pullCount a (s + 1) h = m})
Expand All @@ -153,11 +179,6 @@ 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]

-- todo: this is in ℝ because of the limited def of leastGE
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 @@ -199,6 +220,36 @@ lemma stepsUntil_eq_dite (a : α) (m : ℕ) (h : ℕ → α × R)
ext s
simpa using (h' s)

-- todo: this is in ℝ because of the limited def of leastGE
lemma stepsUntil_eq_leastGE (a : α) (hm : m ≠ 0) :
stepsUntil a m = leastGE (fun n (h : ℕ → α × ℝ) ↦ pullCount a (n + 1) h) m := by
classical
ext h
rw [stepsUntil_eq_dite]
unfold leastGE hittingAfter
simp only [zero_le, Set.mem_Ici, Nat.cast_le, true_and, ENat.some_eq_coe]
have h_iff : (∃ s, pullCount a (s + 1) h = m) ↔ (∃ s, m ≤ pullCount a (s + 1) h) := by
refine ⟨fun ⟨s, hs⟩ ↦ ⟨s, hs.ge⟩, fun ⟨s, hs⟩ ↦ ?_⟩
exact exists_pullCount_eq_of_le hs hm
by_cases h_exists : ∃ s, m ≤ pullCount a (s + 1) h
swap; · simp_rw [h_iff]; simp [h_exists]
rw [if_pos h_exists, dif_pos]
swap; · rwa [h_iff]
norm_cast
rw [Nat.find_eq_iff]
constructor
· apply le_antisymm
· by_contra! h_contra
obtain ⟨s, hs⟩ : ∃ s, pullCount a (s + 1) h = m := exists_pullCount_eq_of_le h_contra.le hm
rw [← hs] at h_contra
refine h_contra.not_ge ?_
gcongr
exact csInf_le (by simp) (by simp)
· exact Nat.sInf_mem (s := {j | m ≤ pullCount a (j + 1) h}) h_exists
· intro n hn h_contra
refine hn.not_ge ?_
exact csInf_le (by simp) (by simp [h_contra])

lemma stepsUntil_pullCount_le (h : ℕ → α × R) (a : α) (t : ℕ) :
stepsUntil a (pullCount a (t + 1) h) h ≤ t := by
rw [stepsUntil]
Expand Down Expand Up @@ -346,9 +397,9 @@ lemma stepsUntil_eq_congr {h' : ℕ → α × R} (h_eq : ∀ i ≤ n, action i h
rw [pullCount_congr]
grind

lemma isStoppingTime_stepsUntil [MeasurableSingletonClass α] (a : α) (m : ℕ) :
lemma isStoppingTime_stepsUntil [MeasurableSingletonClass α] (a : α) (hm : m ≠ 0) :
IsStoppingTime (Learning.filtration α ℝ) (stepsUntil a m) := by
rw [stepsUntil_eq_leastGE]
rw [stepsUntil_eq_leastGE _ hm]
refine Adapted.isStoppingTime_leastGE _ fun n ↦ ?_
suffices StronglyMeasurable[Learning.filtration α ℝ n] (pullCount a (n + 1)) by fun_prop
exact adapted_pullCount_add_one a n
Expand Down Expand Up @@ -387,6 +438,8 @@ lemma measurable_stepsUntil' [MeasurableSingletonClass α] (a : α) (m : ℕ) :
Measurable (fun ω : (ℕ → α × R) × (ℕ → α → R) ↦ stepsUntil a m ω.1) :=
(measurable_stepsUntil a m).comp measurable_fst

end StepsUntil

section RewardByCount

/-- Reward obtained when pulling action `a` for the `m`-th time.
Expand Down
Loading