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
5 changes: 4 additions & 1 deletion LeanBandits/Bandit.lean
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ Released under Apache 2.0 license as described in the file LICENSE.
Authors: Rémy Degenne, Paulo Rauber
-/
import Mathlib
import LeanBandits.ForMathlib.CondDistrib

/-!
# Bandit
Expand Down Expand Up @@ -159,10 +160,12 @@ lemma hasLaw_arm_zero [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace
sorry

/-- The reward at time `n+1` is independent of the history up to time `n` given the arm at `n+1`. -/
lemma condIndepFun_reward_hist_arm [StandardBorelSpace α] [StandardBorelSpace R]
lemma condIndepFun_reward_hist_arm [StandardBorelSpace α] [Nonempty α]
[StandardBorelSpace R] [Nonempty R]
{alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) :
CondIndepFun (MeasurableSpace.comap (arm (n + 1)) inferInstance)
(measurable_arm _).comap_le (reward (n + 1)) (hist n) (Bandit.trajMeasure alg ν) := by
rw [condIndepFun_iff_condDistrib_prod_ae_eq_prodMkLeft (by fun_prop) (by fun_prop) (by fun_prop)]
sorry

end MeasureSpace
Expand Down
218 changes: 167 additions & 51 deletions LeanBandits/ForMathlib/CondDistrib.lean
Original file line number Diff line number Diff line change
Expand Up @@ -13,8 +13,10 @@ import Mathlib.Probability.Kernel.Condexp
open MeasureTheory ProbabilityTheory Finset
open scoped ENNReal NNReal

variable {α β γ Ω Ω' : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω]
{mα : MeasurableSpace α} {μ : Measure α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ}
variable {α β γ δ Ω Ω' : Type*}
{m mα : MeasurableSpace α} {μ : Measure α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ}
{mδ : MeasurableSpace δ}
[MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω]
[MeasurableSpace Ω'] [StandardBorelSpace Ω'] [Nonempty Ω']
{X : α → β} {Y : α → Ω} {Z : α → Ω'} {T : α → γ}

Expand Down Expand Up @@ -71,10 +73,96 @@ lemma Kernel.prod_apply_prod {κ : Kernel α β} {η : Kernel α γ}
(κ ×ₖ η) a (s ×ˢ t) = (κ a s) * (η a t) := by
rw [Kernel.prod_apply, Measure.prod_prod]

lemma CondIndepFun.prod_right
{α β γ δ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ}
{mδ : MeasurableSpace δ} [StandardBorelSpace α]
{μ : Measure α} [IsFiniteMeasure μ] {X : α → β} {Y : α → γ} {Z : α → δ}
-- fix the lemma in mathlib to allow different types for the functions
theorem CondIndepFun.symm'
[StandardBorelSpace α] {hm : m ≤ mα} [IsFiniteMeasure μ] {f : α → β} {g : α → γ}
(hfg : CondIndepFun m hm f g μ) :
CondIndepFun m hm g f μ :=
Kernel.IndepFun.symm hfg

lemma Kernel.indepFun_const_left {κ : Kernel α β} [IsZeroOrMarkovKernel κ] (c : δ) (X : β → γ) :
IndepFun (fun _ ↦ c) X κ μ := by
rw [IndepFun, MeasurableSpace.comap_const]
exact indep_bot_left _

lemma Kernel.indepFun_const_right {κ : Kernel α β} [IsZeroOrMarkovKernel κ] (X : β → γ) (c : δ) :
IndepFun X (fun _ ↦ c) κ μ :=
(Kernel.indepFun_const_left c X).symm

lemma condIndepFun_const_left [StandardBorelSpace α] {hm : m ≤ mα} [IsFiniteMeasure μ]
(c : γ) (X : α → β) :
CondIndepFun m hm (fun _ ↦ c) X μ :=
Kernel.indepFun_const_left c X

lemma condIndepFun_const_right [StandardBorelSpace α] {hm : m ≤ mα}
[IsFiniteMeasure μ] (X : α → β) (c : γ) :
CondIndepFun m hm X (fun _ ↦ c) μ :=
Kernel.indepFun_const_right X c

lemma condIndepFun_of_measurable_left [StandardBorelSpace α] {hm : m ≤ mα} [IsFiniteMeasure μ]
{X : α → β} {Y : α → γ} (hX : Measurable[m] X) (hY : Measurable Y) :
CondIndepFun m hm X Y μ := by
rw [condIndepFun_iff _ hm _ _ (hX.mono hm le_rfl) hY]
rintro _ _ ⟨s, hs, rfl⟩ ⟨t, ht, rfl⟩
have h_ind_eq ω : (X ⁻¹' s ∩ Y ⁻¹' t).indicator (fun _ ↦ (1 : ℝ)) ω
= (X ⁻¹' s).indicator (fun _ ↦ (1 : ℝ)) ω * (Y ⁻¹' t).indicator (fun _ ↦ (1 : ℝ)) ω := by
simp only [Set.indicator, Set.mem_inter_iff, Set.mem_preimage, mul_ite, mul_one, mul_zero]
split_ifs with h1 h2 h3 h4 h5
· rfl
· exfalso
rw [Set.mem_inter_iff] at h1
refine h3 h1.1
· exfalso
rw [Set.mem_inter_iff] at h1
exact h2 h1.2
· exfalso
rw [Set.mem_inter_iff] at h1
exact h1 ⟨h5, h4⟩
· rfl
· rfl
calc μ[(X ⁻¹' s ∩ Y ⁻¹' t).indicator fun ω ↦ (1 : ℝ)|m]
_ = μ[fun ω ↦ (X ⁻¹' s).indicator (fun _ ↦ 1) ω * (Y ⁻¹' t).indicator (fun _ ↦ 1) ω|m] := by
simp_rw [← h_ind_eq]
_ =ᵐ[μ] fun ω ↦ (X ⁻¹' s).indicator (fun _ ↦ 1) ω * μ[(Y ⁻¹' t).indicator (fun _ ↦ 1)|m] ω := by
refine condExp_mul_of_stronglyMeasurable_left ?_ ?_ ?_
· exact (Measurable.indicator (by fun_prop) (hX hs)).stronglyMeasurable
· have : ((X ⁻¹' s).indicator fun x ↦ (1 : ℝ)) * (Y ⁻¹' t).indicator (fun x ↦ 1)
= (X ⁻¹' s ∩ Y ⁻¹' t).indicator (fun _ ↦ (1 : ℝ)) := by ext; simp [h_ind_eq]
rw [this]
rw [integrable_indicator_iff]
· exact (integrable_const (1 : ℝ)).integrableOn
· exact (hm _ (hX hs)).inter (hY ht)
· rw [integrable_indicator_iff]
· exact (integrable_const (1 : ℝ)).integrableOn
· exact hY ht
_ =ᵐ[μ] μ[(X ⁻¹' s).indicator fun ω ↦ 1|m] * μ[(Y ⁻¹' t).indicator fun ω ↦ 1|m] := by
nth_rw 2 [condExp_of_stronglyMeasurable hm]
· rfl
· exact (Measurable.indicator (by fun_prop) (hX hs)).stronglyMeasurable
· rw [integrable_indicator_iff]
· exact (integrable_const (1 : ℝ)).integrableOn
· exact hm _ (hX hs)

lemma condIndepFun_of_measurable_right [StandardBorelSpace α] {hm : m ≤ mα} [IsFiniteMeasure μ]
{X : α → β} {Y : α → γ} (hX : Measurable X) (hY : Measurable[m] Y) :
CondIndepFun m hm X Y μ := by
refine CondIndepFun.symm' ?_
exact condIndepFun_of_measurable_left hY hX

lemma condIndepFun_self_left [StandardBorelSpace α] [IsFiniteMeasure μ]
{X : α → β} {Z : α → δ} (hX : Measurable X) (hZ : Measurable Z) :
CondIndepFun (mδ.comap Z) hZ.comap_le Z X μ := by
refine condIndepFun_of_measurable_left ?_ hX
rw [measurable_iff_comap_le]

lemma condIndepFun_self_right [StandardBorelSpace α] [IsFiniteMeasure μ]
{X : α → β} {Z : α → δ} (hX : Measurable X) (hZ : Measurable Z) :
CondIndepFun (mδ.comap Z) hZ.comap_le X Z μ := by
refine condIndepFun_of_measurable_right hX ?_
rw [measurable_iff_comap_le]

lemma CondIndepFun.prod_right [StandardBorelSpace α] [IsFiniteMeasure μ]
{X : α → β} {Y : α → γ} {Z : α → δ}
(hX : Measurable X) (hY : Measurable Y) (hZ : Measurable Z)
(h : CondIndepFun (mδ.comap Z) hZ.comap_le X Y μ) :
CondIndepFun (mδ.comap Z) hZ.comap_le X (fun ω ↦ (Y ω, Z ω)) μ := by
Expand Down Expand Up @@ -106,20 +194,26 @@ lemma condDistrib_congr_left {Y' : α → Ω} (hY : Y =ᵐ[μ] Y') :
lemma condDistrib_ae_eq_of_measure_eq_compProd₀
(hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) (κ : Kernel β Ω) [IsFiniteKernel κ]
(hκ : μ.map (fun x => (X x, Y x)) = μ.map X ⊗ₘ κ) :
∀ᵐ x ∂μ.map X, κ x = condDistrib Y X μ x := by
condDistrib Y X μ =ᵐ[μ.map X] κ := by
suffices ∀ᵐ x ∂μ.map (hX.mk X), κ x = condDistrib (hY.mk Y) (hX.mk X) μ x by
symm
rw [Measure.map_congr hX.ae_eq_mk]
convert this using 3 with b
rw [condDistrib_congr hY.ae_eq_mk hX.ae_eq_mk]
rw [condDistrib_congr hY.ae_eq_mk hX.ae_eq_mk, Filter.EventuallyEq]
refine condDistrib_ae_eq_of_measure_eq_compProd (μ := μ) hX.measurable_mk hY.measurable_mk κ
((Eq.trans ?_ hκ).trans ?_)
· refine Measure.map_congr ?_
filter_upwards [hX.ae_eq_mk, hY.ae_eq_mk] with a haX haY using by rw [haX, haY]
· rw [Measure.map_congr hX.ae_eq_mk]

lemma condDistrib_ae_eq_iff_measure_eq_compProd₀
(hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) (κ : Kernel β Ω) [IsFiniteKernel κ] :
(condDistrib Y X μ =ᵐ[μ.map X] κ) ↔ μ.map (fun x => (X x, Y x)) = μ.map X ⊗ₘ κ := by
refine ⟨fun h ↦ ?_, condDistrib_ae_eq_of_measure_eq_compProd₀ hX hY κ⟩
rw [Measure.compProd_congr h.symm, compProd_map_condDistrib hY]

lemma condDistrib_comp (hX : AEMeasurable X μ) {f : β → Ω} (hf : Measurable f) :
condDistrib (f ∘ X) X μ =ᵐ[μ.map X] Kernel.deterministic f hf := by
symm
refine condDistrib_ae_eq_of_measure_eq_compProd₀ hX (by fun_prop) _ ?_
rw [Measure.compProd_deterministic, AEMeasurable.map_map_of_aemeasurable (by fun_prop) hX]
rfl
Expand All @@ -133,7 +227,6 @@ lemma condDistrib_const (hX : AEMeasurable X μ) (c : Ω) :

lemma condDistrib_of_indepFun (h : IndepFun X Y μ) (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) :
condDistrib Y X μ =ᵐ[μ.map X] Kernel.const β (μ.map Y) := by
symm
refine condDistrib_ae_eq_of_measure_eq_compProd₀ (μ := μ) hX hY _ ?_
simp only [Measure.compProd_const]
exact (indepFun_iff_map_prod_eq_prod_map_map hX hY).mp h
Expand Down Expand Up @@ -333,51 +426,74 @@ theorem condIndepFun_comap_iff_map_prod_eq_prod_condDistrib_prod_condDistrib
· exact (h_left hs ht hu).symm
· exact (h_right hs ht hu).symm

-- todo: should be an iff
lemma condDistrib_prod_of_condIndepFun [StandardBorelSpace α] [StandardBorelSpace β] [Nonempty β]
(hX : Measurable X) (hY : Measurable Y) (hZ : Measurable Z)
(h : CondIndepFun (MeasurableSpace.comap Z inferInstance) hZ.comap_le Y X μ) :
condDistrib Y (fun ω ↦ (X ω, Z ω)) μ
=ᵐ[μ.map (fun ω ↦ (X ω, Z ω))] Kernel.prodMkLeft _ (condDistrib Y Z μ) := by
symm
refine condDistrib_ae_eq_of_measure_eq_compProd₀ (μ := μ) (hX.prodMk hZ).aemeasurable
hY.aemeasurable _ ?_
rw [condIndepFun_comap_iff_map_prod_eq_prod_condDistrib_prod_condDistrib hY hX hZ] at h
rw [Measure.compProd_eq_comp_prod]
calc μ.map (fun x ↦ ((X x, Z x), Y x))
_ = ((condDistrib X Z μ ×ₖ Kernel.id) ×ₖ condDistrib Y Z μ) ∘ₘ μ.map Z := by
-- up to shuffling, this is the previous lemma
sorry
_ = (Kernel.id ×ₖ Kernel.prodMkLeft β (condDistrib Y Z μ)) ∘ₘ Kernel.swap _ _
∘ₘ (μ.map Z ⊗ₘ condDistrib X Z μ) := by
rw [Measure.compProd_eq_comp_prod, Measure.comp_assoc, Measure.comp_assoc]
congr
rw [Kernel.comp_assoc, Kernel.swap_prod]
ext ω : 1
simp_rw [Kernel.prod_apply]
rw [Kernel.comp_apply, Kernel.prod_apply, Kernel.id_apply, ← Measure.compProd_eq_comp_prod]
ext s hs
rw [Measure.compProd_apply hs, Measure.prod_apply hs]
simp only [Kernel.prodMkLeft_apply]
rw [lintegral_prod, lintegral_prod]
· simp_rw [lintegral_dirac]
· refine Measurable.aemeasurable ?_
have : Measurable fun a ↦ (Kernel.prodMkLeft _ (condDistrib Y Z μ) a) (Prod.mk a ⁻¹' s) :=
Kernel.measurable_kernel_prodMk_left hs
exact this
· refine Measurable.aemeasurable ?_
have : Measurable fun x ↦ (Kernel.const _ ((condDistrib Y Z μ) ω) x) (Prod.mk x ⁻¹' s) :=
Kernel.measurable_kernel_prodMk_left hs
exact this
_ = (Kernel.id ×ₖ Kernel.prodMkLeft β (condDistrib Y Z μ)) ∘ₘ μ.map (fun a ↦ (X a, Z a)) := by
lemma condIndepFun_iff_condDistrib_prod_ae_eq_prodMkLeft
[StandardBorelSpace α] [StandardBorelSpace β] [Nonempty β]
(hX : Measurable X) (hY : Measurable Y) (hZ : Measurable Z) :
CondIndepFun (MeasurableSpace.comap Z inferInstance) hZ.comap_le Y X μ
↔ condDistrib Y (fun ω ↦ (X ω, Z ω)) μ
=ᵐ[μ.map (fun ω ↦ (X ω, Z ω))] Kernel.prodMkLeft _ (condDistrib Y Z μ) := by
rw [condDistrib_ae_eq_iff_measure_eq_compProd₀ (μ := μ) (hX.prodMk hZ).aemeasurable
hY.aemeasurable, condIndepFun_comap_iff_map_prod_eq_prod_condDistrib_prod_condDistrib hY hX hZ,
Measure.compProd_eq_comp_prod]
let e : Ω' × Ω × β ≃ᵐ (β × Ω') × Ω := {
toFun := fun p ↦ ((p.2.2, p.1), p.2.1)
invFun := fun p ↦ (p.1.2, p.2, p.1.1)
left_inv p := by simp
right_inv p := by simp
measurable_toFun := by simp only [Equiv.coe_fn_mk]; fun_prop
measurable_invFun := by simp only [Equiv.coe_fn_symm_mk]; fun_prop }
have h_eq : ((condDistrib X Z μ ×ₖ Kernel.id) ×ₖ condDistrib Y Z μ) ∘ₘ μ.map Z
= (Kernel.id ×ₖ Kernel.prodMkLeft β (condDistrib Y Z μ)) ∘ₘ μ.map (fun a ↦ (X a, Z a)) := by
calc ((condDistrib X Z μ ×ₖ Kernel.id) ×ₖ condDistrib Y Z μ) ∘ₘ μ.map Z
_ = (Kernel.id ×ₖ Kernel.prodMkLeft β (condDistrib Y Z μ)) ∘ₘ Kernel.swap _ _
∘ₘ (μ.map Z ⊗ₘ condDistrib X Z μ) := by
rw [Measure.compProd_eq_comp_prod, Measure.comp_assoc, Measure.comp_assoc]
congr
rw [Kernel.comp_assoc, Kernel.swap_prod]
ext ω : 1
simp_rw [Kernel.prod_apply]
rw [Kernel.comp_apply, Kernel.prod_apply, Kernel.id_apply, ← Measure.compProd_eq_comp_prod]
ext s hs
rw [Measure.compProd_apply hs, Measure.prod_apply hs]
simp only [Kernel.prodMkLeft_apply]
rw [lintegral_prod, lintegral_prod]
· simp_rw [lintegral_dirac]
· refine Measurable.aemeasurable ?_
have : Measurable fun a ↦ (Kernel.prodMkLeft _ (condDistrib Y Z μ) a) (Prod.mk a ⁻¹' s) :=
Kernel.measurable_kernel_prodMk_left hs
exact this
· refine Measurable.aemeasurable ?_
have : Measurable fun x ↦ (Kernel.const _ ((condDistrib Y Z μ) ω) x) (Prod.mk x ⁻¹' s) :=
Kernel.measurable_kernel_prodMk_left hs
exact this
_ = (Kernel.id ×ₖ Kernel.prodMkLeft β (condDistrib Y Z μ)) ∘ₘ μ.map (fun a ↦ (X a, Z a)) := by
congr
rw [compProd_map_condDistrib hX.aemeasurable, Measure.swap_comp,
Measure.map_map (by fun_prop) (by fun_prop)]
rfl
rw [← h_eq]
have h1 : μ.map (fun x ↦ ((X x, Z x), Y x)) = (μ.map (fun a ↦ (Z a , Y a, X a))).map e := by
rw [Measure.map_map (by fun_prop) (by fun_prop)]
congr
rw [compProd_map_condDistrib hX.aemeasurable, Measure.swap_comp,
Measure.map_map (by fun_prop) (by fun_prop)]
rfl
have h1_symm : μ.map (fun a ↦ (Z a , Y a, X a))
= (μ.map (fun x ↦ ((X x, Z x), Y x))).map e.symm := by
rw [h1, Measure.map_map (by fun_prop) (by fun_prop), MeasurableEquiv.symm_comp_self,
Measure.map_id]
have h2 : (condDistrib X Z μ ×ₖ Kernel.id ×ₖ condDistrib Y Z μ) ∘ₘ μ.map Z
= ((Kernel.id ×ₖ (condDistrib Y Z μ ×ₖ condDistrib X Z μ)) ∘ₘ μ.map Z).map e := by
rw [← Measure.deterministic_comp_eq_map e.measurable, Measure.comp_assoc]
sorry
have h2_symm : (Kernel.id ×ₖ (condDistrib Y Z μ ×ₖ condDistrib X Z μ)) ∘ₘ μ.map Z
= ((condDistrib X Z μ ×ₖ Kernel.id ×ₖ condDistrib Y Z μ) ∘ₘ μ.map Z).map e.symm := by
rw [h2, Measure.map_map (by fun_prop) (by fun_prop), MeasurableEquiv.symm_comp_self,
Measure.map_id]
rw [h1, h2]
exact ⟨fun h ↦ by rw [h], fun h ↦ by rw [h1_symm, h1, h2_symm, h2, h]⟩

lemma condDistrib_fst_prod (hX : AEMeasurable X μ) (hY : AEMeasurable Y μ)
(ν : Measure γ) [IsProbabilityMeasure ν] :
condDistrib (fun ω ↦ Y ω.1) (fun ω ↦ X ω.1) (μ.prod ν) =ᵐ[μ.map X] condDistrib Y X μ := by
symm
refine condDistrib_ae_eq_of_measure_eq_compProd₀ (μ := μ) hX hY _ ?_
have hX_map : (μ.prod ν).map (fun ω ↦ X ω.1) = μ.map X := by
calc (μ.prod ν).map (fun ω ↦ X ω.1)
Expand Down Expand Up @@ -439,7 +555,7 @@ lemma cond_of_condIndepFun [StandardBorelSpace α] [StandardBorelSpace β] [None
(hX : Measurable X) (hY : Measurable Y) {b : β} {ω : Ω'}
(hμ : μ (X ⁻¹' {b} ∩ Z ⁻¹' {ω}) ≠ 0) :
(μ[|X ⁻¹' {b} ∩ Z ⁻¹' {ω}]).map Y = (μ[|Z ⁻¹' {ω}]).map Y := by
have h := condDistrib_prod_of_condIndepFun hX hY hZ h
have h := (condIndepFun_iff_condDistrib_prod_ae_eq_prodMkLeft hX hY hZ).mp h
have h_left := condDistrib_ae_eq_cond (hX.prodMk hZ) hY (μ := μ)
have h_right := condDistrib_ae_eq_cond hZ hY (μ := μ)
rw [Filter.EventuallyEq, ae_iff_of_countable] at h h_left h_right
Expand Down
37 changes: 37 additions & 0 deletions LeanBandits/Regret.lean
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,13 @@ noncomputable def pullCount [DecidableEq α] (k : ℕ → α) (a : α) (t : ℕ)
@[simp]
lemma pullCount_zero (k : ℕ → α) (a : α) : pullCount k a 0 = 0 := by simp [pullCount]

lemma pullCount_one : pullCount k a 1 = if k 0 = a then 1 else 0 := by
simp only [pullCount, range_one]
split_ifs with h
· rw [card_eq_one]
refine ⟨0, by simp [h]⟩
· simp [h]

open Classical in
lemma monotone_pullCount (k : ℕ → α) (a : α) : Monotone (pullCount k a) :=
fun _ _ _ ↦ card_le_card (filter_subset_filter _ (by simpa))
Expand Down Expand Up @@ -110,6 +117,36 @@ lemma stepsUntil_pullCount_eq (k : ℕ → α) (t : ℕ) :
simpa [stepsUntil, pullCount_eq_pullCount_add_one]
exact fun t' h ↦ Nat.le_of_lt_succ ((monotone_pullCount k (k t)).reflect_lt (h ▸ lt_add_one _))

/-- If we pull arm `a` at time 0, the first time at which it is pulled once is 0. -/
lemma stepsUntil_one_of_eq (hka : k 0 = a) : stepsUntil k a 1 = 0 := by
classical
have h_pull : pullCount k a 1 = 1 := by simp [pullCount_one, hka]
have h_le := stepsUntil_pullCount_le k a 0
simpa [h_pull] using h_le

lemma stepsUntil_eq_zero_iff :
stepsUntil k a m = 0 ↔ (m = 0 ∧ k 0 ≠ a) ∨ (m = 1 ∧ k 0 = a) := by
classical
refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩
· have h_exists : ∃ s, pullCount k a (s + 1) = m := by
by_contra! h_contra
rw [← stepsUntil_eq_top_iff] at h_contra
simp [h_contra] at 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
by_cases hka : k 0 = a
· simp only [hka, ↓reduceIte] at h
simp [h.symm, hka]
· simp only [hka, ↓reduceIte] at h
simp [h.symm, hka]
· cases h with
| inl h =>
rw [h.1, stepsUntil_zero_of_ne h.2]
| inr h =>
rw [h.1]
exact stepsUntil_one_of_eq h.2

lemma arm_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount (arm · h) a (s + 1) = m) :
arm (stepsUntil (arm · h) a m).toNat h = a := by
classical
Expand Down
17 changes: 14 additions & 3 deletions LeanBandits/RewardByCountMeasure.lean
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,6 @@ Released under Apache 2.0 license as described in the file LICENSE.
Authors: Rémy Degenne
-/
import LeanBandits.Bandit
import LeanBandits.ForMathlib.CondDistrib
import LeanBandits.Regret

/-! # Laws of `stepsUntil` and `rewardByCount`
Expand Down Expand Up @@ -152,8 +151,20 @@ lemma condIndepFun_reward_stepsUntil_arm [StandardBorelSpace α] [Countable α]
-- `hist (n-1)` and `arm n`.
-- It thus suffices to prove the independence of `reward n` and `hist (n-1)` conditionally
-- on `arm n`.
have hn : n ≠ 0 := by
sorry -- assume it?
by_cases hn : n = 0
· simp only [hn, CharP.cast_eq_zero]
simp only [stepsUntil_eq_zero_iff, hm, ne_eq, false_and, false_or]
by_cases hm1 : m = 1
· simp only [hm1, true_and]
have h_indep := condIndepFun_self_right (X := reward 0) (Z := arm 0)
(mβ := inferInstance) (mδ := inferInstance) (μ := Bandit.trajMeasure alg ν)
(by fun_prop) (by fun_prop)
have : {ω : ℕ → α × ℝ | arm 0 ω = a}.indicator (fun x ↦ 1)
= {b | b = a}.indicator (fun _ ↦ 1) ∘ arm 0 := by ext; simp [Set.indicator]
rw [this]
exact h_indep.comp measurable_id (by fun_prop)
· simp only [hm1, false_and, Set.setOf_false, Set.indicator_empty]
exact condIndepFun_const_right (reward 0) 0
have h_indep : CondIndepFun (mα.comap (arm n)) (measurable_arm n).comap_le (reward n)
(hist (n - 1)) (Bandit.trajMeasure alg ν) := by
convert condIndepFun_reward_hist_arm (alg := alg) (ν := ν) (n - 1)
Expand Down