Skip to content

Commit 30dc8bb

Browse files
authored
trajMeasure definition for Mathlib (#24)
2 parents 2b8092f + cda4b01 commit 30dc8bb

6 files changed

Lines changed: 145 additions & 86 deletions

File tree

‎LeanBandits.lean‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,9 @@ import LeanBandits.AlgorithmBuilding
22
import LeanBandits.Bandit
33
import LeanBandits.ETC
44
import LeanBandits.ForMathlib.CondDistrib
5+
import LeanBandits.ForMathlib.KernelCompositionLemmas
6+
import LeanBandits.ForMathlib.KernelCompositionParallelComp
7+
import LeanBandits.ForMathlib.Traj
58
import LeanBandits.Regret
69
import LeanBandits.RewardByCountMeasure
710
import LeanBandits.UCB

‎LeanBandits/Bandit.lean‎

Lines changed: 7 additions & 77 deletions
Original file line numberDiff line numberDiff line change
@@ -5,27 +5,17 @@ Authors: Rémy Degenne, Paulo Rauber
55
-/
66
import Mathlib
77
import LeanBandits.ForMathlib.CondDistrib
8+
import LeanBandits.ForMathlib.KernelCompositionLemmas
9+
import LeanBandits.ForMathlib.Traj
810

911
/-!
1012
# Bandit
11-
1213
-/
1314

1415
open MeasureTheory ProbabilityTheory Filter Real Finset
1516

1617
open scoped ENNReal NNReal
1718

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 _)
28-
2919
namespace Bandits
3020

3121
variable {α R : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R}
@@ -84,22 +74,15 @@ deriving IsMarkovKernel
8474
/-- Measure on the sequence of arms pulled and rewards observed generated by the bandit. -/
8575
noncomputable
8676
def trajMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : Measure (ℕ → α × R) :=
87-
(traj alg ν 0) ∘ₘ ((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero (fun _ ↦ α × R)).symm)
77+
Kernel.trajMeasure (alg.p0 ⊗ₘ ν) (stepKernel alg ν)
78+
deriving IsProbabilityMeasure
8879

8980
/-- Measure of an infinite stream of rewards from each arm. -/
9081
noncomputable
9182
def streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] : Measure (ℕ → α → R) :=
9283
Measure.infinitePi fun _ ↦ Measure.infinitePi ν
9384
deriving IsProbabilityMeasure
9485

95-
instance (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
96-
IsProbabilityMeasure (trajMeasure alg ν) := by
97-
rw [trajMeasure]
98-
have : IsProbabilityMeasure
99-
((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero (fun _ ↦ α × R)).symm) :=
100-
isProbabilityMeasure_map <| by fun_prop
101-
infer_instance
102-
10386
/-- Joint distribution of the sequence of arm pulled and rewards, and a stream of independent
10487
rewards from all arms. -/
10588
noncomputable
@@ -171,57 +154,6 @@ open Kernel Preorder
171154
variable {X : ℕ → Type*} [∀ n, MeasurableSpace (X n)]
172155
{κ : (n : ℕ) → Kernel ((i : { x // x ∈ Iic n }) → X i) (X (n + 1))} [∀ n, IsMarkovKernel (κ n)]
173156

174-
lemma Measure.compProd_map {X Y Z : Type*} {mX : MeasurableSpace X} {mY : MeasurableSpace Y}
175-
{mZ : MeasurableSpace Z} {μ : Measure X} {κ : Kernel X Y} [SFinite μ] [IsSFiniteKernel κ]
176-
{f : Y → Z} (hf : Measurable f) :
177-
μ ⊗ₘ (κ.map f) = (μ ⊗ₘ κ).map (Prod.map id f) := by
178-
calc μ ⊗ₘ (κ.map f)
179-
_ = (Kernel.id ∥ₖ Kernel.deterministic f hf) ∘ₘ (Kernel.id ×ₖ κ) ∘ₘ μ := by
180-
rw [Measure.comp_assoc, Kernel.parallelComp_comp_prod, Measure.compProd_eq_comp_prod,
181-
Kernel.id_comp, Kernel.deterministic_comp_eq_map]
182-
_ = (Kernel.id ∥ₖ Kernel.deterministic f hf) ∘ₘ (μ ⊗ₘ κ) := by rw [Measure.compProd_eq_comp_prod]
183-
_ = (μ ⊗ₘ κ).map (Prod.map id f) := by
184-
rw [Kernel.id, Kernel.deterministic_parallelComp_deterministic,
185-
Measure.deterministic_comp_eq_map]
186-
187-
lemma partialTraj_compProd_eq_traj_map_frestrictLe (a : ℕ) (x₀ : (i : Iic 0) → X i) :
188-
(partialTraj κ 0 a x₀) ⊗ₘ (κ a) =
189-
(traj κ 0 x₀).map (fun x ↦ (frestrictLe a x, x (a + 1))) := by
190-
have h1 := partialTraj_compProd_traj (κ := κ) (zero_le a) x₀
191-
have h2 : (fun x : Π n, X n ↦ (frestrictLe a x, x (a + 1))) =
192-
(Prod.map id (fun x ↦ x (a + 1))) ∘ (fun x ↦ (frestrictLe a x, x)) := by ext <;> simp
193-
rw [h2, ← Measure.map_map (by fun_prop) (by fun_prop), ← h1, ← Measure.compProd_map (by fun_prop)]
194-
congr
195-
have : (fun x : Π n, X n ↦ x (a + 1)) =
196-
(fun x : Π i : Iic (a + 1), X i ↦ x ⟨a+1, by simp⟩) ∘ (frestrictLe (a + 1)) := by ext; simp
197-
rw [this, map_comp_right _ (by fun_prop) (by fun_prop), traj_map_frestrictLe,
198-
partialTraj_succ_self, ← map_comp_right _ (by fun_prop) (by fun_prop)]
199-
have : (fun x : Π i : Iic (a + 1), X i ↦ x ⟨a+1, by simp⟩) ∘ IicProdIoc a (a + 1)
200-
= (MeasurableEquiv.piSingleton a).symm ∘ Prod.snd := by
201-
ext; simp [_root_.IicProdIoc, MeasurableEquiv.piSingleton]
202-
rw [this, map_comp_right _ (by fun_prop) (by fun_prop), ← snd_eq, snd_prod,
203-
← map_comp_right _ (by fun_prop) (by fun_prop)]
204-
simp
205-
206-
lemma traj_cond_lemma1 {a : ℕ} (μ₀ : Measure ((i : Iic 0) → X i)) [IsFiniteMeasure μ₀] :
207-
(traj κ 0 ∘ₘ μ₀).map (fun x ↦ (frestrictLe a x, x (a + 1)))
208-
= (traj κ 0 ∘ₘ μ₀).map (frestrictLe a) ⊗ₘ κ a := by
209-
rw [Measure.compProd_eq_comp_prod, Measure.map_comp _ _ (by fun_prop),
210-
Measure.map_comp _ _ (by fun_prop), Measure.comp_assoc, traj_map_frestrictLe]
211-
congr
212-
ext x₀ : 1
213-
rw [ProbabilityTheory.Kernel.comp_apply, ← Measure.compProd_eq_comp_prod]
214-
symm
215-
rw [Kernel.map_apply _ (by fun_prop)]
216-
exact partialTraj_compProd_eq_traj_map_frestrictLe a x₀
217-
218-
lemma condDistrib_lemma (μ₀ : Measure ((i : Iic 0) → X i)) [IsFiniteMeasure μ₀] (a : ℕ)
219-
[Nonempty (X (a + 1))] [StandardBorelSpace (X (a + 1))] :
220-
condDistrib (fun x ↦ x (a + 1)) (frestrictLe a) (traj κ 0 ∘ₘ μ₀)
221-
=ᵐ[(traj κ 0 ∘ₘ μ₀).map (frestrictLe a)] κ a := by
222-
symm
223-
exact condDistrib_ae_eq_of_measure_eq_compProd (by fun_prop) (by fun_prop) _ (traj_cond_lemma1 μ₀)
224-
225157
lemma traj_zero_map_eval_zero :
226158
(Kernel.traj κ 0).map (fun h ↦ h 0)
227159
= Kernel.deterministic (MeasurableEquiv.piIicZero X)
@@ -240,9 +172,7 @@ lemma condDistrib_arm_reward [StandardBorelSpace α] [Nonempty α] [StandardBore
240172
(alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
241173
condDistrib (fun h ↦ (arm (n + 1) h, reward (n + 1) h)) (hist n) (Bandit.trajMeasure alg ν)
242174
=ᵐ[(Bandit.trajMeasure alg ν).map (hist n)] Bandit.stepKernel alg ν n :=
243-
condDistrib_lemma (X := fun _ ↦ α × R)
244-
((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero (fun _ ↦ α × R)).symm)
245-
(κ := Bandit.stepKernel alg ν) n
175+
Kernel.condDistrib_trajMeasure_ae_eq_kernel
246176

247177
lemma condDistrib_reward [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
248178
(alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
@@ -267,9 +197,9 @@ lemma hasLaw_step_zero
267197
HasLaw (fun h : ℕ → α × R ↦ h 0) (alg.p0 ⊗ₘ ν) (Bandit.trajMeasure alg ν) where
268198
aemeasurable := Measurable.aemeasurable (by fun_prop)
269199
map_eq := by
270-
simp only [Bandit.trajMeasure]
200+
simp only [Bandit.trajMeasure, Kernel.trajMeasure]
271201
rw [← Measure.deterministic_comp_eq_map (by fun_prop), Measure.comp_assoc,
272-
Kernel.deterministic_comp_eq_map, Bandit.traj, traj_zero_map_eval_zero,
202+
Kernel.deterministic_comp_eq_map, traj_zero_map_eval_zero,
273203
Measure.deterministic_comp_eq_map, Measure.map_map (by fun_prop) (by fun_prop)]
274204
simp
275205

‎LeanBandits/ForMathlib/CondDistrib.lean‎

Lines changed: 1 addition & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@ import Mathlib.Probability.Independence.Conditional
88
import Mathlib.Probability.Kernel.Composition.Lemmas
99
import Mathlib.Probability.Kernel.CompProdEqIff
1010
import Mathlib.Probability.Kernel.Condexp
11-
11+
import LeanBandits.ForMathlib.KernelCompositionParallelComp
1212

1313
open MeasureTheory ProbabilityTheory Finset
1414
open scoped ENNReal NNReal
@@ -68,14 +68,6 @@ end MeasureTheory.Measure
6868

6969
namespace ProbabilityTheory
7070

71-
lemma Kernel.deterministic_parallelComp_deterministic
72-
{f : α → γ} {g : β → δ} (hf : Measurable f) (hg : Measurable g) :
73-
(deterministic f hf) ∥ₖ (deterministic g hg)
74-
= deterministic (Prod.map f g) (hf.prodMap hg) := by
75-
ext x : 1
76-
rw [parallelComp_apply, deterministic_apply, deterministic_apply, deterministic_apply, Prod.map,
77-
Measure.dirac_prod_dirac]
78-
7971
lemma Kernel.prod_apply_prod {κ : Kernel α β} {η : Kernel α γ}
8072
[IsSFiniteKernel κ] [IsSFiniteKernel η] {s : Set β} {t : Set γ} {a : α} :
8173
(κ ×ₖ η) a (s ×ˢ t) = (κ a s) * (η a t) := by
Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,21 @@
1+
import Mathlib.Probability.Kernel.Composition.Lemmas
2+
import LeanBandits.ForMathlib.KernelCompositionParallelComp
3+
4+
open MeasureTheory ProbabilityTheory
5+
open scoped ENNReal
6+
7+
variable {α β γ δ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β}
8+
{mγ : MeasurableSpace γ} {mδ : MeasurableSpace δ}
9+
{μ : Measure α} {ν : Measure β} {κ : Kernel α β}
10+
11+
-- PR: https://github.com/leanprover-community/mathlib4/pull/29555
12+
lemma MeasureTheory.Measure.compProd_map [SFinite μ] [IsSFiniteKernel κ]
13+
{f : β → γ} (hf : Measurable f) :
14+
μ ⊗ₘ (κ.map f) = (μ ⊗ₘ κ).map (Prod.map id f) := by
15+
calc μ ⊗ₘ (κ.map f)
16+
_ = (Kernel.id ∥ₖ Kernel.deterministic f hf) ∘ₘ (Kernel.id ×ₖ κ) ∘ₘ μ := by
17+
rw [comp_assoc, Kernel.parallelComp_comp_prod, compProd_eq_comp_prod,
18+
Kernel.id_comp, Kernel.deterministic_comp_eq_map]
19+
_ = (Kernel.id ∥ₖ Kernel.deterministic f hf) ∘ₘ (μ ⊗ₘ κ) := by rw [compProd_eq_comp_prod]
20+
_ = (μ ⊗ₘ κ).map (Prod.map id f) := by
21+
rw [Kernel.id, Kernel.deterministic_parallelComp_deterministic, deterministic_comp_eq_map]
Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,26 @@
1+
import Mathlib.Probability.Kernel.Composition.ParallelComp
2+
3+
open MeasureTheory
4+
open scoped ENNReal
5+
6+
namespace ProbabilityTheory.Kernel
7+
8+
variable {α β γ δ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β}
9+
{mγ : MeasurableSpace γ} {mδ : MeasurableSpace δ}
10+
{κ : Kernel α β} {η : Kernel γ δ} {x : α × γ}
11+
12+
-- PR: https://github.com/leanprover-community/mathlib4/pull/29555
13+
lemma parallelComp_apply_prod [IsSFiniteKernel κ] [IsSFiniteKernel η] (s : Set β) (t : Set δ) :
14+
(κ ∥ₖ η) x (s ×ˢ t) = (κ x.1 s) * (η x.2 t) := by
15+
rw [parallelComp_apply, Measure.prod_prod]
16+
17+
-- PR: https://github.com/leanprover-community/mathlib4/pull/29555
18+
lemma deterministic_parallelComp_deterministic
19+
{f : α → γ} {g : β → δ} (hf : Measurable f) (hg : Measurable g) :
20+
(deterministic f hf) ∥ₖ (deterministic g hg)
21+
= deterministic (Prod.map f g) (hf.prodMap hg) := by
22+
ext x : 1
23+
rw [parallelComp_apply, deterministic_apply, deterministic_apply, deterministic_apply, Prod.map,
24+
Measure.dirac_prod_dirac]
25+
26+
end ProbabilityTheory.Kernel

‎LeanBandits/ForMathlib/Traj.lean‎

Lines changed: 87 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
1+
import Mathlib.Probability.Kernel.IonescuTulcea.Traj
2+
import Mathlib.Probability.Kernel.CondDistrib
3+
import LeanBandits.ForMathlib.CondDistrib
4+
import LeanBandits.ForMathlib.KernelCompositionLemmas
5+
6+
open Filter Finset Function MeasurableEquiv MeasurableSpace MeasureTheory Preorder ProbabilityTheory
7+
8+
variable {X : ℕ → Type*} [∀ n, MeasurableSpace (X n)]
9+
variable {κ : (n : ℕ) → Kernel (Π i : Iic n, X i) (X (n + 1))} [∀ n, IsMarkovKernel (κ n)]
10+
variable {μ₀ : Measure (X 0)} [IsProbabilityMeasure μ₀]
11+
12+
section MeasurableEquiv
13+
14+
instance : Unique (Iic 0) := by simp only [mem_Iic, nonpos_iff_eq_zero]; exact Unique.subtypeEq 0
15+
16+
lemma coe_default_Iic_zero : ((default : Iic 0) : ℕ) = 0 := by
17+
calc _ = ((⟨0, by simp⟩ : Iic 0) : ℕ) := by congr; exact (Unique.eq_default _).symm
18+
_ = _ := by simp
19+
20+
/-- Measurable equivalence between `Iic 0 → X i` and `X 0`. -/
21+
def MeasurableEquiv.piIicZero (X : ℕ → Type*) [∀ n, MeasurableSpace (X n)] :
22+
((i : Iic 0) → X i) ≃ᵐ X 0 :=
23+
(MeasurableEquiv.piUnique _).trans (coe_default_Iic_zero.symm ▸ MeasurableEquiv.refl _)
24+
25+
end MeasurableEquiv
26+
27+
namespace ProbabilityTheory.Kernel
28+
29+
-- Probability/Kernel/IonescuTulcea/Traj.lean
30+
/-- Distribution of the infinite trajectory given the distribution of `X 0`. -/
31+
noncomputable
32+
def trajMeasure (μ₀ : Measure (X 0)) (κ : (n : ℕ) → Kernel (Π i : Iic n, X i) (X (n + 1)))
33+
[∀ n, IsMarkovKernel (κ n)] :
34+
Measure (Π n, X n) :=
35+
(traj κ 0) ∘ₘ (μ₀.map (MeasurableEquiv.piIicZero _).symm)
36+
37+
-- Probability/Kernel/IonescuTulcea/Traj.lean
38+
instance : IsProbabilityMeasure (trajMeasure μ₀ κ) := by
39+
rw [trajMeasure]
40+
have : IsProbabilityMeasure (μ₀.map (MeasurableEquiv.piIicZero _).symm) :=
41+
isProbabilityMeasure_map <| by fun_prop
42+
infer_instance
43+
44+
-- Probability/Kernel/IonescuTulcea/Traj.lean
45+
lemma traj_map_eq_kernel {a : ℕ} : (traj κ a).map (fun x ↦ x (a + 1)) = κ a := by
46+
set f : (Π n, X n) → X (a + 1) := fun x ↦ x (a + 1)
47+
set g : (Π n : Iic (a + 1), X n) → X (a + 1) := fun x ↦ x ⟨a + 1, by simp⟩
48+
have hf : f = g ∘ (frestrictLe (a + 1)) := by rfl
49+
have hp : g ∘ IicProdIoc a (a + 1) = (piSingleton a).symm ∘ Prod.snd := by
50+
ext
51+
simp [g, _root_.IicProdIoc, piSingleton]
52+
rw [hf, map_comp_right, traj_map_frestrictLe, partialTraj_succ_self, ← map_comp_right, hp,
53+
map_comp_right, ← snd_eq, snd_prod, ← map_comp_right]
54+
all_goals measurability
55+
56+
-- Probability/Kernel/IonescuTulcea/Traj.lean
57+
lemma partialTraj_compProd_kernel_eq_traj_map {a : ℕ} {x₀ : Π n : Iic 0, X n} :
58+
(partialTraj κ 0 a x₀) ⊗ₘ (κ a) = (traj κ 0 x₀).map (fun x ↦ (frestrictLe a x, x (a + 1))) := by
59+
set f := fun x ↦ (frestrictLe a x, x (a + 1))
60+
set g := fun x ↦ (frestrictLe a x, x)
61+
have hf : f = (Prod.map id (fun x ↦ x (a + 1))) ∘ g := rfl
62+
rw [hf, ← Measure.map_map, ← partialTraj_compProd_traj, ← MeasureTheory.Measure.compProd_map,
63+
traj_map_eq_kernel]
64+
all_goals measurability
65+
66+
-- (Extract kernel lemmas from rewrites?) Probability/Kernel/IonescuTulcea/Traj.lean
67+
lemma trajMeasure_map_frestrictLe_compProd_kernel_eq_trajMeasure_map {a : ℕ} :
68+
(trajMeasure μ₀ κ).map (frestrictLe a) ⊗ₘ κ a =
69+
(trajMeasure μ₀ κ).map (fun x ↦ (frestrictLe a x, x (a + 1))) := by
70+
rw [Measure.compProd_eq_comp_prod, trajMeasure, Measure.map_comp, traj_map_frestrictLe,
71+
Measure.comp_assoc, Measure.map_comp]
72+
any_goals fun_prop
73+
congr
74+
ext1 x₀
75+
rw [comp_apply, ← Measure.compProd_eq_comp_prod, map_apply,
76+
partialTraj_compProd_kernel_eq_traj_map]
77+
fun_prop
78+
79+
-- Probability/Kernel/IonescuTulcea/Traj.lean
80+
lemma condDistrib_trajMeasure_ae_eq_kernel {a : ℕ}
81+
[StandardBorelSpace (X (a + 1))] [Nonempty (X (a + 1))] :
82+
condDistrib (fun x ↦ x (a + 1)) (frestrictLe a) (trajMeasure μ₀ κ)
83+
=ᵐ[(trajMeasure μ₀ κ).map (frestrictLe a)] κ a := by
84+
apply condDistrib_ae_eq_of_measure_eq_compProd₀ (by measurability) (by measurability)
85+
exact trajMeasure_map_frestrictLe_compProd_kernel_eq_trajMeasure_map.symm
86+
87+
end ProbabilityTheory.Kernel

0 commit comments

Comments
 (0)