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
124 changes: 72 additions & 52 deletions LeanMachineLearning/Probability/HasCondDistrib.lean
Original file line number Diff line number Diff line change
Expand Up @@ -75,48 +75,34 @@ lemma HasCondDistrib.snd {Y : α → Ω × Ω'} {κ : Kernel β (Ω × Ω')} [Is
rw [Kernel.snd_eq]
exact HasCondDistrib.comp h measurable_snd

-- TODO: Rename to `HasCondDistrib.comp_right`?
lemma HasCondDistrib.comp_right' [IsFiniteMeasure μ] [IsFiniteKernel κ] {f : γ → β}
(hf : Measurable f) {Z : α → γ} (h : HasCondDistrib Y Z (κ.comap f hf) μ) :
HasCondDistrib Y (f ∘ Z) κ μ := by
have hY : AEMeasurable Y μ := h.aemeasurable_fst
have hZ : AEMeasurable Z μ := h.aemeasurable_snd
have hfZ : AEMeasurable (f ∘ Z) μ := hf.comp_aemeasurable hZ
refine ⟨hY, hfZ, ?_⟩
rw [condDistrib_ae_eq_iff_measure_eq_compProd _ hY]
calc μ.map (fun a ↦ ((f ∘ Z) a, Y a))
_ = (μ.map (fun a ↦ (Z a, Y a))).map (Prod.map f id) := by
rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (hZ.prodMk hY)]
rfl
_ = (μ.map Z ⊗ₘ κ.comap f hf).map (Prod.map f id) := by
rw [(condDistrib_ae_eq_iff_measure_eq_compProd Z hY _).mp h.condDistrib_eq]
_ = (μ.map Z).map f ⊗ₘ κ := by
ext s hs
rw [Measure.map_apply (by fun_prop) hs, Measure.compProd_apply (by measurability),
Measure.compProd_apply hs, lintegral_map (Kernel.measurable_kernel_prodMk_left hs) hf]
rfl
_ = μ.map (f ∘ Z) ⊗ₘ κ := by
rw [AEMeasurable.map_map_of_aemeasurable hf.aemeasurable hZ]

lemma HasCondDistrib.comp_right [IsFiniteMeasure μ] [IsFiniteKernel κ] (h : HasCondDistrib Y X κ μ)
(f : β ≃ᵐ γ) :
HasCondDistrib Y (f ∘ X) (κ.comap f.symm (by fun_prop) : Kernel γ Ω) μ := by
have hY := h.aemeasurable_fst
have hX := h.aemeasurable_snd
refine ⟨h.aemeasurable_fst, by fun_prop, ?_⟩
have h_eq := h.condDistrib_eq
rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop) _] at h_eq ⊢
calc μ.map (fun ω ↦ ((f ∘ X) ω, Y ω))
_ = μ.map ((fun p ↦ (f p.1, p.2)) ∘ fun ω ↦ (X ω, Y ω)) := by congr
_ = (μ.map (fun ω ↦ (X ω, Y ω))).map (fun p ↦ (f p.1, p.2)) := by
rw [AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)]
_ = (μ.map X ⊗ₘ κ).map (fun p ↦ (f p.1, p.2)) := by rw [h_eq]
_ = μ.map (f ∘ X) ⊗ₘ (κ.comap f.symm (by fun_prop)) := by
-- this is probably very inefficient.
have hX_eq : X = f.symm ∘ (f ∘ X) := by ext; simp
conv_lhs => rw [hX_eq]
rw [← AEMeasurable.map_map_of_aemeasurable, Measure.compProd_eq_comp_prod,
← Measure.deterministic_comp_eq_map (f := f.symm), ← Measure.deterministic_comp_eq_map]
rotate_left
· fun_prop
· fun_prop
· fun_prop
· fun_prop
rw [← Kernel.comp_deterministic_eq_comap, Measure.compProd_eq_comp_prod]
simp_rw [Measure.comp_assoc]
congr 1
ext c : 1
rw [Kernel.comp_apply, Kernel.comp_apply, Kernel.prod_apply, Kernel.comp_apply]
simp only [Kernel.deterministic_apply, Kernel.id_apply, Measure.dirac_bind κ.measurable,
Measure.dirac_bind (Kernel.id ×ₖ κ).measurable, Kernel.prod_apply,
Measure.deterministic_comp_eq_map]
ext s hs
rw [Measure.map_apply (by fun_prop) hs, Measure.prod_apply, Measure.prod_apply,
lintegral_dirac', lintegral_dirac']
· congr
ext
simp
· exact measurable_measure_prodMk_left hs
· exact measurable_measure_prodMk_left (hs.preimage (by fun_prop))
· exact hs
· exact hs.preimage (by fun_prop)
apply HasCondDistrib.comp_right' f.measurable
simpa [← Kernel.comap_comp_right]

lemma HasCondDistrib.prod_right [IsFiniteMeasure μ] [IsFiniteKernel κ] (h : HasCondDistrib Y X κ μ)
{f : β → γ} (hf : Measurable f) :
Expand Down Expand Up @@ -186,6 +172,33 @@ lemma hasCondDistrib_prod_right_iff [IsFiniteMeasure μ] [IsFiniteKernel κ] (X
rw [← Measure.map_prod_map _ _ (by fun_prop) (by fun_prop), Measure.map_id,
Measure.map_dirac' (by fun_prop)]

lemma HasCondDistrib.hasLaw_of_const [IsProbabilityMeasure μ] {Q : Measure Ω} [SFinite Q]
(h : HasCondDistrib Y X (Kernel.const β Q) μ) : HasLaw Y Q μ := by
obtain ⟨hY, hX, h⟩ := h
refine ⟨hY, ?_⟩
have h_snd : (μ.map (fun ω => (X ω, Y ω))).snd = Q := by
have h_map : μ.map (fun ω => (X ω, Y ω)) = (μ.map X) ⊗ₘ (Kernel.const _ Q) :=
have h_map : μ.map (fun ω => (X ω, Y ω)) = (μ.map X) ⊗ₘ (condDistrib Y X μ) :=
(compProd_map_condDistrib hY).symm
h_map.trans (Measure.compProd_congr h)
rw [h_map, MeasureTheory.Measure.snd_compProd]
simp [MeasureTheory.Measure.map_apply_of_aemeasurable hX]
rwa [Measure.snd_map_prodMk₀ hX] at h_snd

lemma HasCondDistrib.indepFun_of_const [IsProbabilityMeasure μ] {Q : Measure Ω} [SFinite Q]
(h : HasCondDistrib Y X (Kernel.const β Q) μ) : IndepFun X Y μ := by
rw [indepFun_iff_condDistrib_eq_const h.aemeasurable_snd h.aemeasurable_fst,
h.hasLaw_of_const.map_eq]
exact h.condDistrib_eq

lemma HasCondDistrib.const_map_of_const [IsProbabilityMeasure μ] {Q : Measure Ω} [SFinite Q]
(h : HasCondDistrib Y X (Kernel.const β Q) μ) [StandardBorelSpace β] [Nonempty β] :
HasCondDistrib X Y (Kernel.const Ω (μ.map X)) μ where
aemeasurable_fst := h.aemeasurable_snd
aemeasurable_snd := h.aemeasurable_fst
condDistrib_eq :=
condDistrib_of_indepFun h.indepFun_of_const.symm h.aemeasurable_fst h.aemeasurable_snd

lemma HasLaw.prod_of_hasCondDistrib {P : Measure β} [IsFiniteMeasure μ] [IsSFiniteKernel κ]
(h1 : HasLaw X P μ) (h2 : HasCondDistrib Y X κ μ) :
HasLaw (fun ω ↦ (X ω, Y ω)) (P ⊗ₘ κ) μ := by
Expand All @@ -197,6 +210,26 @@ lemma HasLaw.prod_of_hasCondDistrib {P : Measure β} [IsFiniteMeasure μ] [IsSFi
rw [← h1.map_eq]
exact h2.condDistrib_eq

lemma HasCondDistrib.of_compProd [IsFiniteMeasure μ] [IsFiniteKernel κ] {Z : α → Ω'}
{η : Kernel (β × Ω) Ω'} [IsMarkovKernel η]
(h : HasCondDistrib (fun a ↦ (Y a, Z a)) X (κ ⊗ₖ η) μ) :
HasCondDistrib Z (fun a ↦ (X a, Y a)) η μ := by
have hZ : AEMeasurable Z μ := h.aemeasurable_fst.snd
have hX : AEMeasurable X μ := h.aemeasurable_snd
have hY : AEMeasurable Y μ := h.aemeasurable_fst.fst
refine ⟨hZ, (hX.prodMk hY), ?_⟩
have hc := h.condDistrib_eq
rw [condDistrib_ae_eq_iff_measure_eq_compProd _ (by fun_prop)] at hc ⊢
calc μ.map (fun a ↦ ((X a, Y a), Z a))
_ = (μ.map X ⊗ₘ (κ ⊗ₖ η)).map MeasurableEquiv.prodAssoc.symm := by
rw [← hc, AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)]
rfl
_ = μ.map X ⊗ₘ κ ⊗ₘ η :=
Measure.compProd_assoc
_ = μ.map (fun a ↦ (X a, Y a)) ⊗ₘ η := by
rw [← (condDistrib_ae_eq_iff_measure_eq_compProd X hY κ).1]
simpa using h.fst.condDistrib_eq

lemma HasCondDistrib.prod [IsFiniteMeasure μ] [IsFiniteKernel κ]
{Z : α → Ω'} {η : Kernel (β × Ω) Ω'} [IsFiniteKernel η]
(h1 : HasCondDistrib Y X κ μ) (h2 : HasCondDistrib Z (fun ω ↦ (X ω, Y ω)) η μ) :
Expand All @@ -220,17 +253,4 @@ lemma HasCondDistrib.prod [IsFiniteMeasure μ] [IsFiniteKernel κ]
AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)]
rfl

lemma hasLaw_of_hasCondDistrib_const [IsProbabilityMeasure μ] {Q : Measure Ω} [SFinite Q]
(h : HasCondDistrib Y X (Kernel.const _ Q) μ) : HasLaw Y Q μ := by
obtain ⟨hY, hX, h⟩ := h
refine ⟨hY, ?_⟩
have h_snd : (μ.map (fun ω => (X ω, Y ω))).snd = Q := by
have h_map : μ.map (fun ω => (X ω, Y ω)) = (μ.map X) ⊗ₘ (Kernel.const _ Q) :=
have h_map : μ.map (fun ω => (X ω, Y ω)) = (μ.map X) ⊗ₘ (condDistrib Y X μ) :=
(compProd_map_condDistrib hY).symm
h_map.trans (Measure.compProd_congr h)
rw [h_map, MeasureTheory.Measure.snd_compProd]
simp [MeasureTheory.Measure.map_apply_of_aemeasurable hX]
rwa [Measure.snd_map_prodMk₀ hX] at h_snd

end ProbabilityTheory
10 changes: 10 additions & 0 deletions LeanMachineLearning/Probability/Independence/CondDistrib.lean
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,16 @@ section CondDistrib

variable [IsFiniteMeasure μ]

lemma map_swap_compProd_map_condDistrib (hY : AEMeasurable Y μ) :
(μ.map X ⊗ₘ condDistrib Y X μ).map Prod.swap = μ.map (fun a ↦ (Y a, X a)) := by
by_cases hX : AEMeasurable X μ
· rw [compProd_map_condDistrib hY,
AEMeasurable.map_map_of_aemeasurable measurable_swap.aemeasurable (hX.prodMk hY)]
rfl
· have hYX : ¬ AEMeasurable (fun a ↦ (Y a, X a)) μ :=
fun h ↦ hX (measurable_snd.comp_aemeasurable h)
simp [hX, hYX]

lemma condDistrib_prod_left [StandardBorelSpace β] [Nonempty β]
(hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) (hT : AEMeasurable T μ) :
condDistrib (fun ω ↦ (X ω, Y ω)) T μ
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,7 @@ lemma hasLaw_action (h : IsAlgEnvSeq A Y (randomSampling μ) env P) (n : ℕ) :
exact h.hasLaw_action_zero
· push Not at hn
obtain ⟨k, rfl⟩ := Nat.exists_eq_succ_of_ne_zero hn
exact hasLaw_of_hasCondDistrib_const <| h.hasCondDistrib_action k
exact (h.hasCondDistrib_action k).hasLaw_of_const

/-- Actions are mutually independent. -/
lemma iIndep_action (h : IsAlgEnvSeq A Y (randomSampling μ) env P) :
Expand Down