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
57 changes: 12 additions & 45 deletions LeanBandits/Bandit/SumRewards.lean
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,15 @@ namespace Bandits

namespace ArrayModel

lemma sum_Icc_one_eq_sum_range {m : ℕ} {f : ℕ → ℝ} :
∑ i ∈ Icc 1 m, f (i - 1) = ∑ i ∈ range m, f i := by
have h : Icc 1 m = (range m).image (· + 1) := by
ext x; simp only [mem_Icc, mem_image, mem_range]; constructor
· intro ⟨h1, h2⟩; exact ⟨x - 1, by omega, by omega⟩
· rintro ⟨a, ha, rfl⟩; omega
rw [h, Finset.sum_image (fun _ _ _ _ h => by omega)]
simp

variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [Countable α]
[StandardBorelSpace α] [Nonempty α]
{alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν]
Expand Down Expand Up @@ -88,19 +97,7 @@ lemma identDistrib_pullCount_prod_sum_Icc_rewardByCount (n : ℕ) :
· infer_instance
ext a : 1
congr 1
let e : Icc 1 (pullCount A a n ω) ≃ range (pullCount A a n ω) :=
{ toFun x := ⟨x - 1, by have h := x.2; simp only [mem_Icc] at h; simp; grind⟩
invFun x := ⟨x + 1, by
have h := x.2
simp only [mem_Icc, le_add_iff_nonneg_left, zero_le, true_and, ge_iff_le]
simp only [mem_range] at h
grind⟩
left_inv x := by have h := x.2; simp only [mem_Icc] at h; grind
right_inv x := by have h := x.2; grind }
rw [← sum_coe_sort (Icc 1 (pullCount A a n ω)), ← sum_coe_sort (range (pullCount A a n ω)),
sum_equiv e]
· simp
· simp [e]
exact sum_Icc_one_eq_sum_range.symm

lemma identDistrib_pullCount_prod_sumRewards (n : ℕ) :
IdentDistrib (fun ω a ↦ (pullCount A a n ω, sumRewards A R a n ω))
Expand Down Expand Up @@ -363,38 +360,8 @@ lemma _root_.Learning.IsAlgEnvSeq.law_pullCount_sumRewards_unique
(h1 : IsAlgEnvSeq A R alg (stationaryEnv ν) P)
(h2 : IsAlgEnvSeq A₂ R₂ alg (stationaryEnv ν) P') :
P.map (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω)) =
P'.map (fun ω ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) := by
have hA := h1.measurable_A
have hR := h1.measurable_R
have hA2 := h2.measurable_A
have hR2 := h2.measurable_R
have h_unique := isAlgEnvSeq_unique h1 h2
let f := fun p : ℕ → α × ℝ ↦ (∑ i ∈ range n, if (p i).1 = a then 1 else 0,
∑ i ∈ range n, if (p i).1 = a then (p i).2 else 0)
have hf : Measurable f := by
refine Measurable.prod ?_ ?_
· simp only [f]
refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop)
exact (measurableSet_singleton _).preimage (by fun_prop)
· simp only [f]
refine measurable_sum _ fun i hi ↦ Measurable.ite ?_ (by fun_prop) (by fun_prop)
exact (measurableSet_singleton _).preimage (by fun_prop)
have h_eq_comp : (fun ω ↦ (pullCount A a n ω, sumRewards A R a n ω))
= f ∘ (fun ω n ↦ (A n ω, R n ω)) := by
ext ω : 1
rw [pullCount_eq_comp (R := R), sumRewards_eq_comp]
grind
have h_eq_comp2 : (fun ω ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω))
= f ∘ (fun ω n ↦ (A₂ n ω, R₂ n ω)) := by
ext ω : 1
rw [pullCount_eq_comp (R := R₂), sumRewards_eq_comp]
grind
rw [h_eq_comp, h_eq_comp2, ← Measure.map_map hf, h_unique, Measure.map_map hf,
← h_eq_comp2]
· rw [measurable_pi_iff]
exact fun n ↦ Measurable.prodMk (hA2 n) (hR2 n)
· rw [measurable_pi_iff]
exact fun n ↦ Measurable.prodMk (hA n) (hR n)
P'.map (fun ω ↦ (pullCount A₂ a n ω, sumRewards A₂ R₂ a n ω)) :=
((h1.law_pullCount_sumRewards_unique' h2 (n := n)).comp (u := fun f ↦ f a) (by fun_prop)).map_eq

-- this is what we will use for UCB
lemma prob_pullCount_prod_sumRewards_mem_le [Countable α]
Expand Down
44 changes: 22 additions & 22 deletions LeanBandits/SequentialLearning/FiniteActions.lean
Original file line number Diff line number Diff line change
Expand Up @@ -95,10 +95,7 @@ lemma pullCount_eq_pullCount' {n : ℕ} {ω : Ω} (hn : n ≠ 0) :
pullCount A a n ω = pullCount' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by
cases n with
| zero => exact absurd rfl hn
| succ n =>
rw [pullCount_add_one_eq_pullCount' (R' := R')]
have : n + 1 - 1 = n := by simp
exact this ▸ rfl
| succ n => simp [pullCount_add_one_eq_pullCount' (R' := R')]

lemma pullCount'_mono {n m : ℕ} (hnm : n ≤ m) :
pullCount' n (fun i ↦ (A i ω, R' i ω)) a ≤ pullCount' m (fun i ↦ (A i ω, R' i ω)) a := by
Expand Down Expand Up @@ -269,10 +266,7 @@ lemma stepsUntil_zero_of_ne (hka : A 0 ω ≠ a) : stepsUntil A a 0 ω = 0 := by

lemma stepsUntil_zero_of_eq (hka : A 0 ω = a) : stepsUntil A a 0 ω = ⊤ := by
rw [stepsUntil_eq_top_iff]
suffices 0 < pullCount A a 1 ω by
intro n hn
refine lt_irrefl 0 ?_
exact this.trans_le (le_trans (monotone_pullCount _ _ (by omega)) hn.le)
suffices 0 < pullCount A a 1 ω from fun _ ↦ (this.trans_le (monotone_pullCount _ _ (by lia))).ne'
rw [← hka, ← zero_add 1, pullCount_action_eq_pullCount_add_one]
simp

Expand Down Expand Up @@ -362,11 +356,7 @@ lemma stepsUntil_eq_zero_iff :
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'
by_cases hka : A 0 ω = a
· simp only [hka, ↓reduceIte] at h'
simp [h'.symm, hka]
· simp only [hka, ↓reduceIte] at h'
simp [h'.symm, hka]
by_cases hka : A 0 ω = a <;> simp_all
· cases h' with
| inl h =>
rw [h.1, stepsUntil_zero_of_ne h.2]
Expand Down Expand Up @@ -641,11 +631,6 @@ theorem isStoppingTime_stepsUntil_filtrationAction [MeasurableSingletonClass α]
· rw [IsAlgEnvSeq.filtrationAction_eq_comap _ hn]
exact measurableSet_stepsUntil_eq hA hR' a m n

-- /-- Sigma-algebra generated by the stopping time `stepsUntil a m`. -/
-- def stepsUntilMeasurableSpace [Nonempty R] [MeasurableSingletonClass α] (a : α) (m : ℕ) :
-- MeasurableSpace (ℕ → α × R) :=
-- (isStoppingTime_stepsUntil_filtrationAction a m (mR := mR)).measurableSpace

end Measurability

end StepsUntil
Expand Down Expand Up @@ -772,6 +757,24 @@ lemma sumRewards_add_one {R' : ℕ → Ω → ℝ} :
unfold sumRewards
rw [sum_range_succ]

lemma sumRewards_eq_of_pullCount_eq {R' : ℕ → Ω → ℝ} {s t : ℕ}
(h_eq : pullCount A a s ω = pullCount A a t ω) :
sumRewards A R' a s ω = sumRewards A R' a t ω := by
wlog hst : s ≤ t
· have hts : t ≤ s := by lia
exact (this h_eq.symm hts).symm
induction t, hst using Nat.le_induction with
| base => rfl
| succ t hst' ih =>
have h_mono' : pullCount A a t ω ≤ pullCount A a (t + 1) ω := pullCount_mono a (Nat.le_succ t) ω
have h_eq_t : pullCount A a s ω = pullCount A a t ω :=
le_antisymm (pullCount_mono a hst' ω) (h_eq ▸ h_mono')
have hne : A t ω ≠ a := by
intro ha
have h1 := ha ▸ pullCount_action_eq_pullCount_add_one (A := A) t ω
lia
rw [sumRewards_add_one, if_neg hne, add_zero, ih h_eq_t]

lemma sumRewards_eq_pullCount_mul_empMean {R' : ℕ → Ω → ℝ} {ω : Ω}
(h_pull : pullCount A a t ω ≠ 0) :
sumRewards A R' a t ω = pullCount A a t ω * empMean A R' a t ω := by unfold empMean; field_simp
Expand Down Expand Up @@ -801,10 +804,7 @@ lemma sumRewards_eq_sumRewards' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (h
sumRewards A R' a n ω = sumRewards' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by
cases n with
| zero => exact absurd rfl hn
| succ n =>
rw [sumRewards_add_one_eq_sumRewards']
have : n + 1 - 1 = n := by simp
exact this ▸ rfl
| succ n => simp [sumRewards_add_one_eq_sumRewards']

lemma empMean_add_one_eq_empMean' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} :
empMean A R' a (n + 1) ω = empMean' n (fun i ↦ (A i ω, R' i ω)) a := by
Expand Down