@@ -15,11 +15,16 @@ open MeasureTheory ProbabilityTheory Filter Real Finset
1515
1616open scoped ENNReal NNReal
1717
18- /-- Measurable equivalence between `Iic 0 → α` and `α`. -/
19- def MeasurableEquiv.piIicZero (α : Type *) [MeasurableSpace α] :
20- (Iic 0 → α) ≃ᵐ α :=
21- have : Unique (Iic 0 ) := by simp only [mem_Iic, nonpos_iff_eq_zero]; exact Unique.subtypeEq 0
22- MeasurableEquiv.funUnique _ _
18+ instance : Unique (Iic 0 ) := by simp only [mem_Iic, nonpos_iff_eq_zero]; exact Unique.subtypeEq 0
19+
20+ lemma coe_default_Iic_zero : ((default : Iic 0 ) : ℕ) = 0 := by
21+ calc _ = ((⟨0 , by simp⟩ : Iic 0 ) : ℕ) := by congr; exact (Unique.eq_default _).symm
22+ _ = _ := by simp
23+
24+ /-- Measurable equivalence between `Iic 0 → X i` and `X 0`. -/
25+ def MeasurableEquiv.piIicZero (X : ℕ → Type *) [∀ n, MeasurableSpace (X n)] :
26+ ((i : Iic 0 ) → X i) ≃ᵐ X 0 :=
27+ (MeasurableEquiv.piUnique _).trans (coe_default_Iic_zero.symm ▸ MeasurableEquiv.refl _)
2328
2429namespace Bandits
2530
@@ -39,6 +44,14 @@ structure Algorithm (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] wher
3944instance (alg : Algorithm α R) (n : ℕ) : IsMarkovKernel (alg.policy n) := alg.h_policy n
4045instance (alg : Algorithm α R) : IsProbabilityMeasure alg.p0 := alg.hp0
4146
47+ /-- A deterministic algorithm. -/
48+ noncomputable
49+ def detAlgorithm (nextArm : (n : ℕ) → (Iic n → α × R) → α) (h_next : ∀ n, Measurable (nextArm n))
50+ (arm0 : α) :
51+ Algorithm α R where
52+ policy n := Kernel.deterministic (nextArm n) (h_next n)
53+ p0 := Measure.dirac arm0
54+
4255namespace Bandit
4356
4457/-- Kernel describing the distribution of the next arm-reward pair given the history up to `n`. -/
@@ -71,7 +84,7 @@ deriving IsMarkovKernel
7184/-- Measure on the sequence of arms pulled and rewards observed generated by the bandit. -/
7285noncomputable
7386def trajMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : Measure (ℕ → α × R) :=
74- (traj alg ν 0 ) ∘ₘ ((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero _ ).symm)
87+ (traj alg ν 0 ) ∘ₘ ((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero ( fun _ ↦ α × R) ).symm)
7588
7689/-- Measure of an infinite stream of rewards from each arm. -/
7790noncomputable
@@ -82,7 +95,8 @@ deriving IsProbabilityMeasure
8295instance (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
8396 IsProbabilityMeasure (trajMeasure alg ν) := by
8497 rw [trajMeasure]
85- have : IsProbabilityMeasure ((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero _).symm) :=
98+ have : IsProbabilityMeasure
99+ ((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero (fun _ ↦ α × R)).symm) :=
86100 isProbabilityMeasure_map <| by fun_prop
87101 infer_instance
88102
@@ -130,29 +144,125 @@ lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦
130144@[fun_prop]
131145lemma measurable_hist (n : ℕ) : Measurable (hist n (α := α) (R := R)) := by unfold hist; fun_prop
132146
147+ lemma hist_eq_frestrictLe :
148+ hist = Preorder.frestrictLe («π » := fun _ ↦ α × R) := by
149+ ext n h i : 3
150+ simp [hist, Preorder.frestrictLe]
151+
133152/-- Filtration of the bandit process. -/
134153protected def filtration (α R : Type *) [MeasurableSpace α] [MeasurableSpace R] :
135154 Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) :=
136155 MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R)
137156
157+ section Traj
158+
159+ open Kernel Preorder
160+
161+ variable {X : ℕ → Type *} [∀ n, MeasurableSpace (X n)]
162+ {κ : (n : ℕ) → Kernel ((i : { x // x ∈ Iic n }) → X i) (X (n + 1 ))} [∀ n, IsMarkovKernel (κ n)]
163+
164+ lemma Measure.compProd_map {X Y Z : Type *} {mX : MeasurableSpace X} {mY : MeasurableSpace Y}
165+ {mZ : MeasurableSpace Z} {μ : Measure X} {κ : Kernel X Y} [SFinite μ] [IsSFiniteKernel κ]
166+ {f : Y → Z} (hf : Measurable f) :
167+ μ ⊗ₘ (κ.map f) = (μ ⊗ₘ κ).map (Prod.map id f) := by
168+ calc μ ⊗ₘ (κ.map f)
169+ _ = (Kernel.id ∥ₖ Kernel.deterministic f hf) ∘ₘ (Kernel.id ×ₖ κ) ∘ₘ μ := by
170+ rw [Measure.comp_assoc, Kernel.parallelComp_comp_prod, Measure.compProd_eq_comp_prod,
171+ Kernel.id_comp, Kernel.deterministic_comp_eq_map]
172+ _ = (Kernel.id ∥ₖ Kernel.deterministic f hf) ∘ₘ (μ ⊗ₘ κ) := by rw [Measure.compProd_eq_comp_prod]
173+ _ = (μ ⊗ₘ κ).map (Prod.map id f) := by
174+ rw [Kernel.id, Kernel.deterministic_parallelComp_deterministic,
175+ Measure.deterministic_comp_eq_map]
176+
177+ lemma partialTraj_compProd_eq_traj_map_frestrictLe (a : ℕ) (x₀ : (i : Iic 0 ) → X i) :
178+ (partialTraj κ 0 a x₀) ⊗ₘ (κ a) =
179+ (traj κ 0 x₀).map (fun x ↦ (frestrictLe a x, x (a + 1 ))) := by
180+ have h1 := partialTraj_compProd_traj (κ := κ) (zero_le a) x₀
181+ have h2 : (fun x : Π n, X n ↦ (frestrictLe a x, x (a + 1 ))) =
182+ (Prod.map id (fun x ↦ x (a + 1 ))) ∘ (fun x ↦ (frestrictLe a x, x)) := by ext <;> simp
183+ rw [h2, ← Measure.map_map (by fun_prop) (by fun_prop), ← h1, ← Measure.compProd_map (by fun_prop)]
184+ congr
185+ have : (fun x : Π n, X n ↦ x (a + 1 )) =
186+ (fun x : Π i : Iic (a + 1 ), X i ↦ x ⟨a+1 , by simp⟩) ∘ (frestrictLe (a + 1 )) := by ext; simp
187+ rw [this, map_comp_right _ (by fun_prop) (by fun_prop), traj_map_frestrictLe,
188+ partialTraj_succ_self, ← map_comp_right _ (by fun_prop) (by fun_prop)]
189+ have : (fun x : Π i : Iic (a + 1 ), X i ↦ x ⟨a+1 , by simp⟩) ∘ IicProdIoc a (a + 1 )
190+ = (MeasurableEquiv.piSingleton a).symm ∘ Prod.snd := by
191+ ext; simp [_root_.IicProdIoc, MeasurableEquiv.piSingleton]
192+ rw [this, map_comp_right _ (by fun_prop) (by fun_prop), ← snd_eq, snd_prod,
193+ ← map_comp_right _ (by fun_prop) (by fun_prop)]
194+ simp
195+
196+ lemma traj_cond_lemma1 {a : ℕ} (μ₀ : Measure ((i : Iic 0 ) → X i)) [IsFiniteMeasure μ₀] :
197+ (traj κ 0 ∘ₘ μ₀).map (fun x ↦ (frestrictLe a x, x (a + 1 )))
198+ = (traj κ 0 ∘ₘ μ₀).map (frestrictLe a) ⊗ₘ κ a := by
199+ rw [Measure.compProd_eq_comp_prod, Measure.map_comp _ _ (by fun_prop),
200+ Measure.map_comp _ _ (by fun_prop), Measure.comp_assoc, traj_map_frestrictLe]
201+ congr
202+ ext x₀ : 1
203+ rw [ProbabilityTheory.Kernel.comp_apply, ← Measure.compProd_eq_comp_prod]
204+ symm
205+ rw [Kernel.map_apply _ (by fun_prop)]
206+ exact partialTraj_compProd_eq_traj_map_frestrictLe a x₀
207+
208+ lemma condDistrib_lemma (μ₀ : Measure ((i : Iic 0 ) → X i)) [IsFiniteMeasure μ₀] (a : ℕ)
209+ [Nonempty (X (a + 1 ))] [StandardBorelSpace (X (a + 1 ))] :
210+ condDistrib (fun x ↦ x (a + 1 )) (frestrictLe a) (traj κ 0 ∘ₘ μ₀)
211+ =ᵐ[(traj κ 0 ∘ₘ μ₀).map (frestrictLe a)] κ a := by
212+ symm
213+ exact condDistrib_ae_eq_of_measure_eq_compProd (by fun_prop) (by fun_prop) _ (traj_cond_lemma1 μ₀)
214+
215+ lemma traj_zero_map_eval_zero :
216+ (Kernel.traj κ 0 ).map (fun h ↦ h 0 )
217+ = Kernel.deterministic (MeasurableEquiv.piIicZero X)
218+ (MeasurableEquiv.piIicZero X).measurable := by
219+ suffices (Kernel.traj κ 0 ).map (fun h ↦ h 0 ) = (Kernel.partialTraj κ 0 0 ).map
220+ (MeasurableEquiv.piIicZero X) by
221+ rwa [Kernel.partialTraj_zero,
222+ Kernel.deterministic_map _ (MeasurableEquiv.piIicZero X).measurable] at this
223+ rw [← Kernel.traj_map_frestrictLe, ← Kernel.map_comp_right _ (by fun_prop) (by fun_prop)]
224+ congr with h
225+ sorry
226+
227+ end Traj
228+
138229lemma condDistrib_arm_reward [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
139230 (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
140231 condDistrib (fun h ↦ (arm (n + 1 ) h, reward (n + 1 ) h)) (hist n) (Bandit.trajMeasure alg ν)
141- =ᵐ[(Bandit.trajMeasure alg ν).map (hist n)] Bandit.stepKernel alg ν n := by
142- sorry
232+ =ᵐ[(Bandit.trajMeasure alg ν).map (hist n)] Bandit.stepKernel alg ν n :=
233+ condDistrib_lemma (X := fun _ ↦ α × R)
234+ ((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero (fun _ ↦ α × R)).symm)
235+ (κ := Bandit.stepKernel alg ν) n
143236
144- lemma condDistrib_reward [StandardBorelSpace R ] [Nonempty R] (alg : Algorithm α R) (ν : Kernel α R)
145- [IsMarkovKernel ν] (n : ℕ) :
237+ lemma condDistrib_reward [StandardBorelSpace α ] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
238+ (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
146239 condDistrib (reward n) (arm n) (Bandit.trajMeasure alg ν)
147240 =ᵐ[(Bandit.trajMeasure alg ν).map (arm n)] ν := by
148- sorry
241+ cases n with
242+ | zero => sorry
243+ | succ n =>
244+ have h_ar := condDistrib_arm_reward alg ν n
245+ have h_prod := condDistrib_prod_left (X := arm (n + 1 )) (Y := reward (n + 1 ))
246+ (T := hist n) (μ := Bandit.trajMeasure alg ν) (by fun_prop) (by fun_prop) (by fun_prop)
247+ sorry
149248
150249lemma condDistrib_arm [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
151250 (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
152251 condDistrib (arm (n + 1 )) (hist n) (Bandit.trajMeasure alg ν)
153252 =ᵐ[(Bandit.trajMeasure alg ν).map (hist n)] alg.policy n := by
154253 sorry
155254
255+ lemma hasLaw_step_zero
256+ (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
257+ HasLaw (fun h : ℕ → α × R ↦ h 0 ) (alg.p0 ⊗ₘ ν) (Bandit.trajMeasure alg ν) where
258+ aemeasurable := Measurable.aemeasurable (by fun_prop)
259+ map_eq := by
260+ simp only [Bandit.trajMeasure]
261+ rw [← Measure.deterministic_comp_eq_map (by fun_prop), Measure.comp_assoc,
262+ Kernel.deterministic_comp_eq_map, Bandit.traj, traj_zero_map_eval_zero,
263+ Measure.deterministic_comp_eq_map, Measure.map_map (by fun_prop) (by fun_prop)]
264+ simp
265+
156266lemma hasLaw_arm_zero [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
157267 (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
158268 HasLaw (arm 0 ) alg.p0 (Bandit.trajMeasure alg ν) where
@@ -168,6 +278,37 @@ lemma condIndepFun_reward_hist_arm [StandardBorelSpace α] [Nonempty α]
168278 rw [condIndepFun_iff_condDistrib_prod_ae_eq_prodMkLeft (by fun_prop) (by fun_prop) (by fun_prop)]
169279 sorry
170280
281+ section DetAlgorithm
282+
283+ variable [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
284+ {nextArm : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextArm n)}
285+ {arm0 : α} {ν : Kernel α R} [IsMarkovKernel ν]
286+
287+ lemma HasLaw_arm_zero_detAlgorithm :
288+ HasLaw (arm 0 ) (Measure.dirac arm0)
289+ (Bandit.trajMeasure (detAlgorithm nextArm h_next arm0) ν) where
290+ map_eq := (hasLaw_arm_zero _ _).map_eq
291+
292+ lemma arm_zero_detAlgorithm :
293+ arm 0 =ᵐ[Bandit.trajMeasure (detAlgorithm nextArm h_next arm0) ν] fun _ ↦ arm0 := by
294+ have h_eq : ∀ᵐ x ∂((Bandit.trajMeasure (detAlgorithm nextArm h_next arm0) ν).map (arm 0 )), x
295+ = arm0 := by
296+ rw [(hasLaw_arm_zero _ _).map_eq]
297+ simp [detAlgorithm]
298+ exact ae_of_ae_map (by fun_prop) h_eq
299+
300+ lemma arm_detAlgorithm_ae_eq (n : ℕ) :
301+ arm (n + 1 ) =ᵐ[Bandit.trajMeasure (detAlgorithm nextArm h_next arm0) ν]
302+ fun h ↦ nextArm n (fun i ↦ h i) := by
303+ sorry
304+
305+ example : ∀ᵐ h ∂(Bandit.trajMeasure (detAlgorithm nextArm h_next arm0) ν),
306+ arm 0 h = arm0 ∧ ∀ n, arm (n + 1 ) h = nextArm n (fun i ↦ h i) := by
307+ rw [eventually_and, ae_all_iff]
308+ exact ⟨arm_zero_detAlgorithm, arm_detAlgorithm_ae_eq⟩
309+
310+ end DetAlgorithm
311+
171312end MeasureSpace
172313
173314end Bandits
0 commit comments