|
| 1 | +/- |
| 2 | +Copyright (c) 2026 Gaëtan Serré. All rights reserved. |
| 3 | +Released under Apache 2.0 license as described in the file LICENSE. |
| 4 | +Authors: Gaëtan Serré |
| 5 | +-/ |
| 6 | +module |
| 7 | + |
| 8 | +public import LeanMachineLearning.Tactic.EqLift.ForMathlib.Kernel |
| 9 | +public import LeanMachineLearning.Tactic.EqLift.ForMathlib.MeasurableEquiv |
| 10 | +public import Mathlib.Probability.Kernel.Composition.CompProd |
| 11 | + |
| 12 | +/-! |
| 13 | +# Kernel Lift |
| 14 | +
|
| 15 | +This file defines the `lift` operation on kernels, which allows to cast kernels to different types |
| 16 | +in the same universe level, as long as there are measurable equivalences between the types. |
| 17 | +
|
| 18 | +## Main declarations |
| 19 | +* `Kernel.lift`: the main definition of the lift operation. |
| 20 | +* `Kernel.isSFinite_lift`: a kernel is s-finite if and only if its lift is s-finite. |
| 21 | +* `Kernel.lift_congr`: two kernels are equal if and only if their lifts are equal. |
| 22 | +* `Kernel.lift_comp`: the lift of a composition is the composition of the lifts. |
| 23 | +* `Kernel.parallelComp_lift`: the lift of a parallel composition is the parallel composition of the |
| 24 | +lifts. |
| 25 | +* `Kernel.prod_lift`: the lift of a product is the product of the lifts. |
| 26 | +-/ |
| 27 | + |
| 28 | +@[expose] public section |
| 29 | + |
| 30 | +open MeasureTheory ProbabilityTheory MeasurableEquiv |
| 31 | + |
| 32 | +namespace ProbabilityTheory.Kernel |
| 33 | + |
| 34 | +universe x y z w t |
| 35 | + |
| 36 | +variable {X : Type x} [MeasurableSpace X] {Y : Type y} [MeasurableSpace Y] |
| 37 | + {X' : Type w} [MeasurableSpace X'] {Y' : Type w} [MeasurableSpace Y'] |
| 38 | + |
| 39 | +/-- Cast a kernel to different types in the same universe level, using measurable equivalences. -/ |
| 40 | +noncomputable def lift {ex : X' ≃ᵐ X} {ey : Y' ≃ᵐ Y} (κ : Kernel X Y) : Kernel X' Y' := |
| 41 | + (κ.map ey.symm).comap ex ex.measurable |
| 42 | + |
| 43 | +variable (ex : X' ≃ᵐ X) (ey : Y' ≃ᵐ Y) |
| 44 | + |
| 45 | +lemma lift_apply (κ : Kernel X Y) (a : X') : |
| 46 | + κ.lift (ex := ex) (ey := ey) a = (κ.map ey.symm) (ex a) := rfl |
| 47 | + |
| 48 | +lemma lift_apply' (κ : Kernel X Y) (a : X') {s : Set Y'} (hs : MeasurableSet s) : |
| 49 | + κ.lift (ex := ex) (ey := ey) a s = κ (ex a) (ey '' s) := by |
| 50 | + simp only [lift, coe_comap, Function.comp_apply] |
| 51 | + rw [map_apply' _ ey.symm.measurable _ hs, preimage_symm] |
| 52 | + |
| 53 | +lemma isSFinite_lift (κ : Kernel X Y) : |
| 54 | + IsSFiniteKernel κ ↔ IsSFiniteKernel (κ.lift (ex := ex) (ey := ey)) := by |
| 55 | + constructor |
| 56 | + · intro h |
| 57 | + simp only [lift] |
| 58 | + infer_instance |
| 59 | + · rintro ⟨κs, hfinite_κs, h⟩ |
| 60 | + constructor |
| 61 | + let κs' (i : ℕ) := ((κs i).map ey).comap ex.symm ex.symm.measurable |
| 62 | + refine ⟨κs', ⟨fun i ↦ ?_, ?_⟩⟩ |
| 63 | + · exact IsFiniteKernel.comap ((κs i).map ey) ex.symm.measurable |
| 64 | + · simp only [κs'] |
| 65 | + ext a s hs |
| 66 | + replace h := DFunLike.congr (x := ey.symm '' s) (DFunLike.congr (x := ex.symm a) h rfl) rfl |
| 67 | + rw [sum_apply, Measure.sum_apply] at h ⊢ |
| 68 | + · rw [lift_apply'] at h |
| 69 | + · convert h with x |
| 70 | + · simp |
| 71 | + · rw [image_symm] |
| 72 | + simp |
| 73 | + · simp only [coe_comap, Function.comp_apply] |
| 74 | + rw [map_apply' _ ey.measurable _ hs, image_symm] |
| 75 | + all_goals measurability |
| 76 | + all_goals measurability |
| 77 | + |
| 78 | +instance (κ : Kernel X Y) [IsSFiniteKernel κ] : IsSFiniteKernel (lift (ex := ex) (ey := ey) κ) := |
| 79 | + (isSFinite_lift ex ey κ).mp ‹_› |
| 80 | + |
| 81 | +instance (κ : Kernel X Y) [IsMarkovKernel κ] : IsMarkovKernel (lift (ex := ex) (ey := ey) κ) := by |
| 82 | + simp only [lift] |
| 83 | + have := IsMarkovKernel.map κ ey.symm.measurable |
| 84 | + exact IsMarkovKernel.comap _ ex.measurable |
| 85 | + |
| 86 | +lemma lift_congr (κ η : Kernel X Y) : |
| 87 | + κ = η ↔ κ.lift (ex := ex) (ey := ey) = η.lift (ex := ex) (ey := ey) := by |
| 88 | + constructor |
| 89 | + · grind |
| 90 | + · intro h |
| 91 | + ext a s hs |
| 92 | + replace h := DFunLike.congr (x := ey.symm '' s) (DFunLike.congr (x := ex.symm a) h rfl) rfl |
| 93 | + rw [lift_apply', lift_apply'] at h |
| 94 | + · simp only [apply_symm_apply] at h |
| 95 | + rwa [image_symm, image_preimage] at h |
| 96 | + · measurability |
| 97 | + · measurability |
| 98 | + |
| 99 | +variable {Z : Type z} [MeasurableSpace Z] {T : Type t} [MeasurableSpace T] |
| 100 | + {Z' : Type w} [MeasurableSpace Z'] {T' : Type w} [MeasurableSpace T'] |
| 101 | + (ez : Z' ≃ᵐ Z) (et : T' ≃ᵐ T) |
| 102 | + |
| 103 | +lemma comp_lift (η : Kernel X Y) (κ : Kernel Z X) : |
| 104 | + η.lift (ex := ex) (ey := ey) ∘ₖ κ.lift (ex := ez) (ey := ex) = |
| 105 | + (η ∘ₖ κ).lift (ex := ez) (ey := ey) := by |
| 106 | + ext _ _ hs |
| 107 | + rw [lift_apply', comp_apply', comp_apply', lift_apply, lintegral_map] |
| 108 | + · congr with y |
| 109 | + simp [lift_apply' _ _ _ _ hs] |
| 110 | + all_goals try fun_prop |
| 111 | + all_goals try measurability |
| 112 | + · exact Kernel.measurable_coe _ hs |
| 113 | + |
| 114 | +lemma parallelComp_lift (κ : Kernel X Y) (η : Kernel Z T) : |
| 115 | + κ.lift (ex := ex) (ey := ey) ∥ₖ η.lift (ex := ez) (ey := et) = |
| 116 | + lift (ex := ex.prodCongr ez) (ey := ey.prodCongr et) (κ ∥ₖ η) := by |
| 117 | + by_cases hκ : IsSFiniteKernel <| lift (ex := ex) (ey := ey) κ |
| 118 | + swap |
| 119 | + · simp only [hκ, not_false_eq_true, parallelComp_of_not_isSFiniteKernel_left, |
| 120 | + (isSFinite_lift ex ey κ).not.mpr hκ] |
| 121 | + simp [lift] |
| 122 | + by_cases hη : IsSFiniteKernel <| lift (ex := ez) (ey := et) η |
| 123 | + swap |
| 124 | + · simp only [hη, not_false_eq_true, parallelComp_of_not_isSFiniteKernel_right, |
| 125 | + (isSFinite_lift ez et η).not.mpr hη] |
| 126 | + simp [lift] |
| 127 | + simp only [lift] |
| 128 | + replace hκ := (isSFinite_lift ex ey κ).mpr hκ |
| 129 | + replace hη := (isSFinite_lift ez et η).mpr hη |
| 130 | + rw [comap_parallelComp_comap, map_parallelComp_map] |
| 131 | + · rfl |
| 132 | + all_goals fun_prop |
| 133 | + |
| 134 | +lemma id_lift : Kernel.id (α := X') = Kernel.id.lift (ex := ex) (ey := ex) := by |
| 135 | + ext _ _ hs |
| 136 | + rw [lift_apply' _ _ _ _ hs] |
| 137 | + simp only [id_apply] |
| 138 | + rw [Measure.dirac_apply' _ hs, Measure.dirac_apply'] |
| 139 | + · exact Set.indicator_eq_indicator (by simp) rfl |
| 140 | + all_goals measurability |
| 141 | + |
| 142 | +lemma discard_lift : discard.{_, w} X' = (discard X).lift (ex := ex) (ey := punit) := by |
| 143 | + ext _ _ hs |
| 144 | + rw [lift_apply' _ _ _ _ hs] |
| 145 | + simp only [discard_apply, MeasurableSpace.measurableSet_top, Measure.dirac_apply'] |
| 146 | + exact Set.indicator_eq_indicator (by grind) rfl |
| 147 | + |
| 148 | + |
| 149 | +lemma copy_lift : copy X' = (copy X).lift (ex := ex) (ey := ex.prodCongr ex) := by |
| 150 | + ext _ _ hs |
| 151 | + rw [lift_apply' _ _ _ _ hs] |
| 152 | + simp only [copy_apply] |
| 153 | + rw [Measure.dirac_apply' _ hs, Measure.dirac_apply'] |
| 154 | + · refine Set.indicator_eq_indicator ?_ rfl |
| 155 | + simp [MeasurableEquiv.prodCongr] |
| 156 | + · measurability |
| 157 | + |
| 158 | +lemma swap_lift : swap X' Y' = (swap X Y).lift (ex := ex.prodCongr ey) (ey := ey.prodCongr ex) := by |
| 159 | + ext a s hs |
| 160 | + rw [lift_apply' _ _ _ _ hs] |
| 161 | + simp only [swap_apply] |
| 162 | + rw [Measure.dirac_apply' _ hs, Measure.dirac_apply'] |
| 163 | + · refine Set.indicator_eq_indicator ?_ rfl |
| 164 | + simp only [prodCongr, MeasurableEquiv.coe_mk, Equiv.prodCongr_apply, Prod.map, coe_toEquiv, |
| 165 | + Prod.swap_prod_mk, Set.mem_image, Prod.mk.injEq, EmbeddingLike.apply_eq_iff_eq, Prod.exists, |
| 166 | + exists_eq_right_right, exists_eq_right] |
| 167 | + grind |
| 168 | + · measurability |
| 169 | + |
| 170 | +lemma prod_lift (κ : Kernel X Y) (η : Kernel X Z) : |
| 171 | + κ.lift (ex := ex) (ey := ey) ×ₖ η.lift (ex := ex) (ey := ez) = |
| 172 | + lift (ex := ex) (ey := ey.prodCongr ez) (κ ×ₖ η) := by |
| 173 | + by_cases hκ : IsSFiniteKernel <| lift (ex := ex) (ey := ey) κ |
| 174 | + swap |
| 175 | + · simp only [hκ, not_false_eq_true, prod_of_not_isSFiniteKernel_left, |
| 176 | + (isSFinite_lift ex ey κ).not.mpr hκ] |
| 177 | + simp [lift] |
| 178 | + by_cases hη : IsSFiniteKernel <| lift (ex := ex) (ey := ez) η |
| 179 | + swap |
| 180 | + · simp only [hη, not_false_eq_true, prod_of_not_isSFiniteKernel_right, |
| 181 | + (isSFinite_lift ex ez η).not.mpr hη] |
| 182 | + simp [lift] |
| 183 | + simp only [prod] |
| 184 | + rw [← comp_lift (ex := ex.prodCongr ex), ← parallelComp_lift, ← copy_lift] |
| 185 | + |
| 186 | +lemma compProd_lift (κ : Kernel X Y) (η : Kernel (X × Y) Z) : |
| 187 | + κ.lift (ex := ex) (ey := ey) ⊗ₖ η.lift (ex := ex.prodCongr ey) (ey := ez) = |
| 188 | + lift (ex := ex) (ey := ey.prodCongr ez) (κ ⊗ₖ η) := by |
| 189 | + by_cases hκ : IsSFiniteKernel <| lift (ex := ex) (ey := ey) κ |
| 190 | + swap |
| 191 | + · simp only [hκ, not_false_eq_true, compProd_of_not_isSFiniteKernel_left, |
| 192 | + (isSFinite_lift ex ey κ).not.mpr hκ] |
| 193 | + simp [lift] |
| 194 | + · by_cases hη : IsSFiniteKernel <| lift (ex := ex.prodCongr ey) (ey := ez) η |
| 195 | + swap |
| 196 | + · simp [hη, (isSFinite_lift (ex.prodCongr ey) ez η).not.mpr hη] |
| 197 | + simp [lift] |
| 198 | + · simp only [compProd] |
| 199 | + rw [← comp_lift (ex := ex.prodCongr ex) (ey := ey.prodCongr ez), ← copy_lift, |
| 200 | + ← comp_lift (ex := ex.prodCongr ey), ← parallelComp_lift, ← id_lift, |
| 201 | + ← comp_lift (ex := ex.prodCongr (ey.prodCongr ey)), ← parallelComp_lift, ← id_lift, |
| 202 | + ← copy_lift, |
| 203 | + ← comp_lift (ex := (ex.prodCongr ey).prodCongr ey), |
| 204 | + ← comp_lift (ex := ez.prodCongr ey), ← parallelComp_lift, ← id_lift, ← swap_lift] |
| 205 | + congr |
| 206 | + simp only [lift] |
| 207 | + rw [deterministic_map (MeasurableEquiv.measurable _) (MeasurableEquiv.measurable _)] |
| 208 | + ext _ : 1 |
| 209 | + simp [comap_apply, deterministic_apply, MeasurableEquiv.prodCongr, prodAssoc] |
| 210 | + |
| 211 | + |
| 212 | +instance {κ : Kernel X Y} [IsDeterministic κ] : |
| 213 | + IsDeterministic (κ.lift (ex := ex) (ey := ey)) where |
| 214 | + parallelComp_self_comp_copy' := by |
| 215 | + rw [parallelComp_lift, copy_lift (ex := ex), copy_lift (ex := ey), comp_lift, comp_lift, |
| 216 | + ← lift_congr, κ.parallelComp_self_comp_copy] |
| 217 | + |
| 218 | +end ProbabilityTheory.Kernel |
0 commit comments