From 531ef96061c99e00591017f4ccae45b1addc2db6 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Thu, 20 Aug 2026 15:19:03 +0200 Subject: [PATCH 1/4] Add the `kernel_hom` tactic --- LeanMachineLearning.lean | 25 ++ LeanMachineLearning/Tactic/EqLift.lean | 19 + .../Tactic/EqLift/ForMathlib/Kernel.lean | 71 +++ .../EqLift/ForMathlib/MeasurableEquiv.lean | 49 +++ .../Tactic/EqLift/Kernel/Lift.lean | 218 +++++++++ .../EqLift/Tactic/Kernel/KernelLift.lean | 186 ++++++++ .../EqLift/Tactic/Kernel/KernelUnlift.lean | 196 +++++++++ .../Tactic/EqLift/Tactic/Kernel/Utils.lean | 88 ++++ .../Tactic/EqLift/Tactic/Lift.lean | 78 ++++ .../Tactic/EqLift/Tactic/Location.lean | 75 ++++ .../Tactic/EqLift/Tactic/Universe.lean | 61 +++ .../Tactic/EqLift/Tactic/Unlift.lean | 66 +++ .../Tactic/EqLift/Tactic/Utils.lean | 91 ++++ LeanMachineLearning/Tactic/KernelHom.lean | 20 + .../Tactic/KernelHom/ForMathlib/Kernel.lean | 53 +++ .../KernelHom/ForMathlib/LIntegral.lean | 27 ++ .../KernelHom/ForMathlib/MeasurableEquiv.lean | 26 ++ .../Tactic/KernelHom/Kernel/Hom.lean | 267 +++++++++++ .../Tactic/KernelHom/Kernel/MonoidalComp.lean | 105 +++++ .../Tactic/KernelHom/Tactic/Delaborators.lean | 50 +++ .../Tactic/KernelHom/Tactic/HomKernel.lean | 275 ++++++++++++ .../Tactic/KernelHom/Tactic/KernelCat.lean | 55 +++ .../KernelHom/Tactic/KernelDiagram.lean | 206 +++++++++ .../Tactic/KernelHom/Tactic/KernelHom.lean | 415 ++++++++++++++++++ .../Tactic/KernelHom/Tactic/Reassoc.lean | 131 ++++++ .../Tactic/KernelHom/Tactic/Utils.lean | 37 ++ scripts/update_tactics.sh | 95 ++++ 27 files changed, 2985 insertions(+) create mode 100644 LeanMachineLearning/Tactic/EqLift.lean create mode 100644 LeanMachineLearning/Tactic/EqLift/ForMathlib/Kernel.lean create mode 100644 LeanMachineLearning/Tactic/EqLift/ForMathlib/MeasurableEquiv.lean create mode 100644 LeanMachineLearning/Tactic/EqLift/Kernel/Lift.lean create mode 100644 LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/KernelLift.lean create mode 100644 LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/KernelUnlift.lean create mode 100644 LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/Utils.lean create mode 100644 LeanMachineLearning/Tactic/EqLift/Tactic/Lift.lean create mode 100644 LeanMachineLearning/Tactic/EqLift/Tactic/Location.lean create mode 100644 LeanMachineLearning/Tactic/EqLift/Tactic/Universe.lean create mode 100644 LeanMachineLearning/Tactic/EqLift/Tactic/Unlift.lean create mode 100644 LeanMachineLearning/Tactic/EqLift/Tactic/Utils.lean create mode 100644 LeanMachineLearning/Tactic/KernelHom.lean create mode 100644 LeanMachineLearning/Tactic/KernelHom/ForMathlib/Kernel.lean create mode 100644 LeanMachineLearning/Tactic/KernelHom/ForMathlib/LIntegral.lean create mode 100644 LeanMachineLearning/Tactic/KernelHom/ForMathlib/MeasurableEquiv.lean create mode 100644 LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean create mode 100644 LeanMachineLearning/Tactic/KernelHom/Kernel/MonoidalComp.lean create mode 100644 LeanMachineLearning/Tactic/KernelHom/Tactic/Delaborators.lean create mode 100644 LeanMachineLearning/Tactic/KernelHom/Tactic/HomKernel.lean create mode 100644 LeanMachineLearning/Tactic/KernelHom/Tactic/KernelCat.lean create mode 100644 LeanMachineLearning/Tactic/KernelHom/Tactic/KernelDiagram.lean create mode 100644 LeanMachineLearning/Tactic/KernelHom/Tactic/KernelHom.lean create mode 100644 LeanMachineLearning/Tactic/KernelHom/Tactic/Reassoc.lean create mode 100644 LeanMachineLearning/Tactic/KernelHom/Tactic/Utils.lean create mode 100755 scripts/update_tactics.sh diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index 9bc9ac39..d6e92832 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -42,3 +42,28 @@ public import LeanMachineLearning.SequentialLearning.EvaluationEnv public import LeanMachineLearning.SequentialLearning.FiniteActions public import LeanMachineLearning.SequentialLearning.IonescuTulceaSpace public import LeanMachineLearning.SequentialLearning.StationaryEnv +public import LeanMachineLearning.Tactic.EqLift +public import LeanMachineLearning.Tactic.EqLift.ForMathlib.Kernel +public import LeanMachineLearning.Tactic.EqLift.ForMathlib.MeasurableEquiv +public import LeanMachineLearning.Tactic.EqLift.Kernel.Lift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Kernel.KernelLift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Kernel.KernelUnlift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Kernel.Utils +public import LeanMachineLearning.Tactic.EqLift.Tactic.Lift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Location +public import LeanMachineLearning.Tactic.EqLift.Tactic.Universe +public import LeanMachineLearning.Tactic.EqLift.Tactic.Unlift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Utils +public import LeanMachineLearning.Tactic.KernelHom +public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.Kernel +public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.LIntegral +public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.MeasurableEquiv +public import LeanMachineLearning.Tactic.KernelHom.Kernel.Hom +public import LeanMachineLearning.Tactic.KernelHom.Kernel.MonoidalComp +public import LeanMachineLearning.Tactic.KernelHom.Tactic.Delaborators +public import LeanMachineLearning.Tactic.KernelHom.Tactic.HomKernel +public import LeanMachineLearning.Tactic.KernelHom.Tactic.KernelCat +public import LeanMachineLearning.Tactic.KernelHom.Tactic.KernelDiagram +public import LeanMachineLearning.Tactic.KernelHom.Tactic.KernelHom +public import LeanMachineLearning.Tactic.KernelHom.Tactic.Reassoc +public import LeanMachineLearning.Tactic.KernelHom.Tactic.Utils diff --git a/LeanMachineLearning/Tactic/EqLift.lean b/LeanMachineLearning/Tactic/EqLift.lean new file mode 100644 index 00000000..2c4f0570 --- /dev/null +++ b/LeanMachineLearning/Tactic/EqLift.lean @@ -0,0 +1,19 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ + +module -- shake: keep-all --deprecated_module: ignore + +public import LeanMachineLearning.Tactic.EqLift.ForMathlib.Kernel +public import LeanMachineLearning.Tactic.EqLift.ForMathlib.MeasurableEquiv +public import LeanMachineLearning.Tactic.EqLift.Kernel.Lift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Kernel.KernelLift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Kernel.KernelUnlift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Kernel.Utils +public import LeanMachineLearning.Tactic.EqLift.Tactic.Lift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Location +public import LeanMachineLearning.Tactic.EqLift.Tactic.Universe +public import LeanMachineLearning.Tactic.EqLift.Tactic.Unlift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Utils diff --git a/LeanMachineLearning/Tactic/EqLift/ForMathlib/Kernel.lean b/LeanMachineLearning/Tactic/EqLift/ForMathlib/Kernel.lean new file mode 100644 index 00000000..52c0b616 --- /dev/null +++ b/LeanMachineLearning/Tactic/EqLift/ForMathlib/Kernel.lean @@ -0,0 +1,71 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import Mathlib.Probability.Kernel.Deterministic + +/-! +# Kernel utilities + +This file provides helper lemmas for working with kernels. + +## Main declarations + +* `comap_parallelComp_comap`: the comap of a parallel composition is the parallel composition of + the comaps. +* `map_parallelComp_map`: the map of a parallel composition is the parallel composition of the maps. +-/ + +@[expose] public section + +open ProbabilityTheory MeasureTheory ENNReal Set + +variable {α β γ ι : Type*} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace γ] + [MeasurableSpace ι] + +namespace ProbabilityTheory.Kernel + +lemma comap_parallelComp_comap {α₂ γ₂ : Type*} [MeasurableSpace α₂] [MeasurableSpace γ₂] + (κ : Kernel α β) (η : Kernel γ ι) [IsSFiniteKernel κ] [IsSFiniteKernel η] + {f : α₂ → α} {g : γ₂ → γ} (hf : Measurable f) (hg : Measurable g) : + κ.comap f hf ∥ₖ η.comap g hg = (κ ∥ₖ η).comap (fun a ↦ (f a.1, g a.2)) (by fun_prop) := by + ext : 1 + rw [Kernel.parallelComp_apply, Kernel.comap_apply, Kernel.comap_apply, Kernel.comap_apply, + Kernel.parallelComp_apply] + +lemma map_parallelComp_map {β₂ ι₂ : Type*} [MeasurableSpace β₂] [MeasurableSpace ι₂] + (κ : Kernel α β) (η : Kernel γ ι) [IsSFiniteKernel κ] [IsSFiniteKernel η] + {f : β → β₂} {g : ι → ι₂} (hf : Measurable f) (hg : Measurable g) : + κ.map f ∥ₖ η.map g = (κ ∥ₖ η).map (fun a ↦ (f a.1, g a.2)) := by + ext a s hs + rw [Kernel.parallelComp_apply', Kernel.lintegral_map, Kernel.map_apply', + Kernel.parallelComp_apply'] + · congr with x + rw [Kernel.map_apply' _ (by fun_prop) _ (by measurability)] + congr + all_goals try fun_prop + all_goals try measurability + exact measurable_measure_prodMk_left hs + +instance (κ : Kernel α β) [IsDeterministic κ] : IsSFiniteKernel κ := by + by_contra + have : ∀ C < ∞, ∃ a, C < (κ a) univ := by + by_contra! h + have : IsFiniteKernel κ := ⟨h⟩ + have : IsSFiniteKernel κ := inferInstance + contradiction + obtain ⟨a, ha⟩ := this 0 (by simp) + have h := DFunLike.congr_fun κ.parallelComp_self_comp_copy a + simp_all only [not_false_eq_true, parallelComp_of_not_isSFiniteKernel_left, zero_comp, zero_apply] + replace h := DFunLike.congr_fun h Set.univ + rw [comp_apply'] at h + · simp_rw [copy_apply, Measure.dirac_apply' _ MeasurableSet.univ, indicator_univ] at h + simp only [Measure.coe_zero, Pi.zero_apply, Pi.one_apply, MeasureTheory.lintegral_const, + one_mul] at h + exact ha.ne h + exact MeasurableSet.univ + +end ProbabilityTheory.Kernel diff --git a/LeanMachineLearning/Tactic/EqLift/ForMathlib/MeasurableEquiv.lean b/LeanMachineLearning/Tactic/EqLift/ForMathlib/MeasurableEquiv.lean new file mode 100644 index 00000000..dea6839f --- /dev/null +++ b/LeanMachineLearning/Tactic/EqLift/ForMathlib/MeasurableEquiv.lean @@ -0,0 +1,49 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import Mathlib.MeasureTheory.MeasurableSpace.Embedding + +/-! +# Measurable equivalences + +This file extends the theory of measurable equivalences, providing utilities for +working with products and unit types. + +## Main declarations + +* `MeasurableEquiv.prod`: product of measurable equivalences. +* `MeasurableEquiv.punit`: measurable equivalence between `PUnit`s. +-/ + +@[expose] public section + +namespace MeasurableEquiv + +universe w x y + +variable {X Y X' Y' : Type*} [MeasurableSpace X] [MeasurableSpace Y] [MeasurableSpace X'] + [MeasurableSpace Y'] (ex : X' ≃ᵐ X) (ey : Y' ≃ᵐ Y) + +/-- The product of two measurable equivalences is a measurable equivalence. -/ +def prod : X' × Y' ≃ᵐ X × Y where + toFun := fun (x', y') ↦ (ex x', ey y') + invFun := fun (x, y) ↦ (ex.symm x, ey.symm y) + left_inv := by simp [Function.LeftInverse] + right_inv := by simp [Function.RightInverse, Function.LeftInverse] + measurable_toFun := by simp only [Equiv.coe_fn_mk]; fun_prop + measurable_invFun := by simp only [Equiv.coe_fn_symm_mk]; fun_prop + +/-- The measurable equivalence between two `PUnit`s. -/ +def punit : PUnit.{w + 1} ≃ᵐ PUnit.{x + 1} where + toFun := fun _ ↦ PUnit.unit + invFun := fun _ ↦ PUnit.unit + left_inv := by grind + right_inv := by grind + measurable_toFun := measurable_id + measurable_invFun := measurable_id + +end MeasurableEquiv diff --git a/LeanMachineLearning/Tactic/EqLift/Kernel/Lift.lean b/LeanMachineLearning/Tactic/EqLift/Kernel/Lift.lean new file mode 100644 index 00000000..cae2baba --- /dev/null +++ b/LeanMachineLearning/Tactic/EqLift/Kernel/Lift.lean @@ -0,0 +1,218 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import LeanMachineLearning.Tactic.EqLift.ForMathlib.Kernel +public import LeanMachineLearning.Tactic.EqLift.ForMathlib.MeasurableEquiv +public import Mathlib.Probability.Kernel.Composition.CompProd + +/-! +# Kernel Lift + +This file defines the `lift` operation on kernels, which allows to cast kernels to different types +in the same universe level, as long as there are measurable equivalences between the types. + +## Main declarations +* `Kernel.lift`: the main definition of the lift operation. +* `Kernel.isSFinite_lift`: a kernel is s-finite if and only if its lift is s-finite. +* `Kernel.lift_congr`: two kernels are equal if and only if their lifts are equal. +* `Kernel.lift_comp`: the lift of a composition is the composition of the lifts. +* `Kernel.parallelComp_lift`: the lift of a parallel composition is the parallel composition of the +lifts. +* `Kernel.prod_lift`: the lift of a product is the product of the lifts. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory MeasurableEquiv + +namespace ProbabilityTheory.Kernel + +universe x y z w t + +variable {X : Type x} [MeasurableSpace X] {Y : Type y} [MeasurableSpace Y] + {X' : Type w} [MeasurableSpace X'] {Y' : Type w} [MeasurableSpace Y'] + +/-- Cast a kernel to different types in the same universe level, using measurable equivalences. -/ +noncomputable def lift {ex : X' ≃ᵐ X} {ey : Y' ≃ᵐ Y} (κ : Kernel X Y) : Kernel X' Y' := + (κ.map ey.symm).comap ex ex.measurable + +variable (ex : X' ≃ᵐ X) (ey : Y' ≃ᵐ Y) + +lemma lift_apply (κ : Kernel X Y) (a : X') : + κ.lift (ex := ex) (ey := ey) a = (κ.map ey.symm) (ex a) := rfl + +lemma lift_apply' (κ : Kernel X Y) (a : X') {s : Set Y'} (hs : MeasurableSet s) : + κ.lift (ex := ex) (ey := ey) a s = κ (ex a) (ey '' s) := by + simp only [lift, coe_comap, Function.comp_apply] + rw [map_apply' _ ey.symm.measurable _ hs, preimage_symm] + +lemma isSFinite_lift (κ : Kernel X Y) : + IsSFiniteKernel κ ↔ IsSFiniteKernel (κ.lift (ex := ex) (ey := ey)) := by + constructor + · intro h + simp only [lift] + infer_instance + · rintro ⟨κs, hfinite_κs, h⟩ + constructor + let κs' (i : ℕ) := ((κs i).map ey).comap ex.symm ex.symm.measurable + refine ⟨κs', ⟨fun i ↦ ?_, ?_⟩⟩ + · exact IsFiniteKernel.comap ((κs i).map ey) ex.symm.measurable + · simp only [κs'] + ext a s hs + replace h := DFunLike.congr (x := ey.symm '' s) (DFunLike.congr (x := ex.symm a) h rfl) rfl + rw [sum_apply, Measure.sum_apply] at h ⊢ + · rw [lift_apply'] at h + · convert h with x + · simp + · rw [image_symm] + simp + · simp only [coe_comap, Function.comp_apply] + rw [map_apply' _ ey.measurable _ hs, image_symm] + all_goals measurability + all_goals measurability + +instance (κ : Kernel X Y) [IsSFiniteKernel κ] : IsSFiniteKernel (lift (ex := ex) (ey := ey) κ) := + (isSFinite_lift ex ey κ).mp ‹_› + +instance (κ : Kernel X Y) [IsMarkovKernel κ] : IsMarkovKernel (lift (ex := ex) (ey := ey) κ) := by + simp only [lift] + have := IsMarkovKernel.map κ ey.symm.measurable + exact IsMarkovKernel.comap _ ex.measurable + +lemma lift_congr (κ η : Kernel X Y) : + κ = η ↔ κ.lift (ex := ex) (ey := ey) = η.lift (ex := ex) (ey := ey) := by + constructor + · grind + · intro h + ext a s hs + replace h := DFunLike.congr (x := ey.symm '' s) (DFunLike.congr (x := ex.symm a) h rfl) rfl + rw [lift_apply', lift_apply'] at h + · simp only [apply_symm_apply] at h + rwa [image_symm, image_preimage] at h + · measurability + · measurability + +variable {Z : Type z} [MeasurableSpace Z] {T : Type t} [MeasurableSpace T] + {Z' : Type w} [MeasurableSpace Z'] {T' : Type w} [MeasurableSpace T'] + (ez : Z' ≃ᵐ Z) (et : T' ≃ᵐ T) + +lemma comp_lift (η : Kernel X Y) (κ : Kernel Z X) : + η.lift (ex := ex) (ey := ey) ∘ₖ κ.lift (ex := ez) (ey := ex) = + (η ∘ₖ κ).lift (ex := ez) (ey := ey) := by + ext _ _ hs + rw [lift_apply', comp_apply', comp_apply', lift_apply, lintegral_map] + · congr with y + simp [lift_apply' _ _ _ _ hs] + all_goals try fun_prop + all_goals try measurability + · exact Kernel.measurable_coe _ hs + +lemma parallelComp_lift (κ : Kernel X Y) (η : Kernel Z T) : + κ.lift (ex := ex) (ey := ey) ∥ₖ η.lift (ex := ez) (ey := et) = + lift (ex := ex.prod ez) (ey := ey.prod et) (κ ∥ₖ η) := by + by_cases hκ : IsSFiniteKernel <| lift (ex := ex) (ey := ey) κ + swap + · simp only [hκ, not_false_eq_true, parallelComp_of_not_isSFiniteKernel_left, + (isSFinite_lift ex ey κ).not.mpr hκ] + simp [lift] + by_cases hη : IsSFiniteKernel <| lift (ex := ez) (ey := et) η + swap + · simp only [hη, not_false_eq_true, parallelComp_of_not_isSFiniteKernel_right, + (isSFinite_lift ez et η).not.mpr hη] + simp [lift] + simp only [lift] + replace hκ := (isSFinite_lift ex ey κ).mpr hκ + replace hη := (isSFinite_lift ez et η).mpr hη + rw [comap_parallelComp_comap, map_parallelComp_map] + · rfl + all_goals fun_prop + +lemma id_lift : Kernel.id (α := X') = Kernel.id.lift (ex := ex) (ey := ex) := by + ext _ _ hs + rw [lift_apply' _ _ _ _ hs] + simp only [id_apply] + rw [Measure.dirac_apply' _ hs, Measure.dirac_apply'] + · exact Set.indicator_eq_indicator (by simp) rfl + all_goals measurability + +lemma discard_lift : discard.{_, w} X' = (discard X).lift (ex := ex) (ey := punit) := by + ext _ _ hs + rw [lift_apply' _ _ _ _ hs] + simp only [discard_apply, MeasurableSpace.measurableSet_top, Measure.dirac_apply'] + exact Set.indicator_eq_indicator (by grind) rfl + + +lemma copy_lift : copy X' = (copy X).lift (ex := ex) (ey := ex.prod ex) := by + ext _ _ hs + rw [lift_apply' _ _ _ _ hs] + simp only [copy_apply] + rw [Measure.dirac_apply' _ hs, Measure.dirac_apply'] + · refine Set.indicator_eq_indicator ?_ rfl + simp [MeasurableEquiv.prod] + · measurability + +lemma swap_lift : swap X' Y' = (swap X Y).lift (ex := ex.prod ey) (ey := ey.prod ex) := by + ext a s hs + rw [lift_apply' _ _ _ _ hs] + simp only [swap_apply] + rw [Measure.dirac_apply' _ hs, Measure.dirac_apply'] + · refine Set.indicator_eq_indicator ?_ rfl + simp only [MeasurableEquiv.prod, MeasurableEquiv.coe_mk, Equiv.coe_fn_mk, Prod.swap_prod_mk, + Set.mem_image, Prod.mk.injEq, EmbeddingLike.apply_eq_iff_eq, Prod.exists, + exists_eq_right_right, exists_eq_right] + grind + · measurability + +lemma prod_lift (κ : Kernel X Y) (η : Kernel X Z) : + κ.lift (ex := ex) (ey := ey) ×ₖ η.lift (ex := ex) (ey := ez) = + lift (ex := ex) (ey := ey.prod ez) (κ ×ₖ η) := by + by_cases hκ : IsSFiniteKernel <| lift (ex := ex) (ey := ey) κ + swap + · simp only [hκ, not_false_eq_true, prod_of_not_isSFiniteKernel_left, + (isSFinite_lift ex ey κ).not.mpr hκ] + simp [lift] + by_cases hη : IsSFiniteKernel <| lift (ex := ex) (ey := ez) η + swap + · simp only [hη, not_false_eq_true, prod_of_not_isSFiniteKernel_right, + (isSFinite_lift ex ez η).not.mpr hη] + simp [lift] + simp only [prod] + rw [← comp_lift (ex := ex.prod ex), ← parallelComp_lift, ← copy_lift] + +lemma compProd_lift (κ : Kernel X Y) (η : Kernel (X × Y) Z) : + κ.lift (ex := ex) (ey := ey) ⊗ₖ η.lift (ex := ex.prod ey) (ey := ez) = + lift (ex := ex) (ey := ey.prod ez) (κ ⊗ₖ η) := by + by_cases hκ : IsSFiniteKernel <| lift (ex := ex) (ey := ey) κ + swap + · simp only [hκ, not_false_eq_true, compProd_of_not_isSFiniteKernel_left, + (isSFinite_lift ex ey κ).not.mpr hκ] + simp [lift] + by_cases hη : IsSFiniteKernel <| lift (ex := ex.prod ey) (ey := ez) η + swap + · simp only [hη, not_false_eq_true, compProd_of_not_isSFiniteKernel_right, + (isSFinite_lift (ex.prod ey) ez η).not.mpr hη] + simp [lift] + simp only [compProd] + rw [← comp_lift (ex := ex.prod ex) (ey := ey.prod ez), ← copy_lift, + ← comp_lift (ex := ex.prod ey), ← parallelComp_lift, ← id_lift, + ← comp_lift (ex := ex.prod (ey.prod ey)), ← parallelComp_lift, ← id_lift, ← copy_lift, + ← comp_lift (ex := (ex.prod ey).prod ey), + ← comp_lift (ex := ez.prod ey), ← parallelComp_lift, ← id_lift, ← swap_lift] + congr + simp only [lift] + rw [deterministic_map (MeasurableEquiv.measurable _) (MeasurableEquiv.measurable _)] + ext _ : 1 + simp [comap_apply, deterministic_apply, MeasurableEquiv.prod, prodAssoc] + + +instance {κ : Kernel X Y} [IsDeterministic κ] : + IsDeterministic (κ.lift (ex := ex) (ey := ey)) where + parallelComp_self_comp_copy' := by + rw [parallelComp_lift, copy_lift (ex := ex), copy_lift (ex := ey), comp_lift, comp_lift, + ← lift_congr, κ.parallelComp_self_comp_copy] + +end ProbabilityTheory.Kernel diff --git a/LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/KernelLift.lean b/LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/KernelLift.lean new file mode 100644 index 00000000..f2a8f240 --- /dev/null +++ b/LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/KernelLift.lean @@ -0,0 +1,186 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import LeanMachineLearning.Tactic.EqLift.Kernel.Lift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Lift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Kernel.Utils + +/-! +# Implementation of the `lift_eq` tactic for kernels. + +This file contains functions that propagate the lifting of kernel expressions through several +operators and primitives, and constructs the necessary proofs for the `lift_eq` tactic. +-/ + +public meta section + +open Lean Meta Parser.Tactic ProbabilityTheory ProbabilityTheory.Kernel + +/-- Lifts a composition of kernels by lifting the inner kernels. -/ +def liftComposition (e : Expr) (maxLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.comp do + throwError "Expected a composition of kernels, but got {e}." + let args := e.getAppArgs + let η := args[args.size - 2]! + let κ := args[args.size - 1]! + let (X, Y, xLvl, yLvl) ← getTypesFromKernel η + let (Z, _, tLvl, _) ← getTypesFromKernel κ + let ex ← constructMeasurableEquiv X xLvl maxLvl + let ey ← constructMeasurableEquiv Y yLvl maxLvl + let ez ← constructMeasurableEquiv Z tLvl maxLvl + let comp_lift_proof ← mkAppM ``comp_lift #[ex, ey, ez, η, κ] + let (η', proofs_η) ← liftExpr η maxLvl proofs + let (κ', proofs_κ) ← liftExpr κ maxLvl proofs_η + return (← mkAppM ``Kernel.comp #[η', κ'], comp_lift_proof :: proofs_κ) + +initialize registerLiftExpr liftComposition + +/-- Lifts a parallel composition of kernels by lifting the inner kernels. -/ +def liftParallelComp (e : Expr) (maxLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.parallelComp do + throwError "Expected a parallel composition of kernels, but got {e}." + let args := e.getAppArgs + let κ := args[args.size - 2]! + let η := args[args.size - 1]! + let (X, Y, xLvl, yLvl) ← getTypesFromKernel κ + let (Z, T, zLvl, tLvl) ← getTypesFromKernel η + let ex ← constructMeasurableEquiv X xLvl maxLvl + let ey ← constructMeasurableEquiv Y yLvl maxLvl + let ez ← constructMeasurableEquiv Z zLvl maxLvl + let et ← constructMeasurableEquiv T tLvl maxLvl + let parallelComp_lift_proof ← mkAppM ``parallelComp_lift #[ex, ey, ez, et, κ, η] + let (κ', proofs_κ) ← liftExpr κ maxLvl proofs + let (η', proofs_η) ← liftExpr η maxLvl proofs_κ + return (← mkAppM ``Kernel.parallelComp #[κ', η'], parallelComp_lift_proof :: proofs_η) + +initialize registerLiftExpr liftParallelComp + +/-- Lifts a product of kernels by lifting the inner kernels. -/ +def liftProd (e : Expr) (maxLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.prod do + throwError "Expected a product of kernels, but got {e}." + let args := e.getAppArgs + let κ := args[args.size - 2]! + let η := args[args.size - 1]! + let (X, Y, xLvl, yLvl) ← getTypesFromKernel κ + let (_, Z, _, zLvl) ← getTypesFromKernel η + let ex ← constructMeasurableEquiv X xLvl maxLvl + let ey ← constructMeasurableEquiv Y yLvl maxLvl + let ez ← constructMeasurableEquiv Z zLvl maxLvl + let prod_lift_proof ← mkAppM ``prod_lift #[ex, ey, ez, κ, η] + let (κ', proofs_κ) ← liftExpr κ maxLvl proofs + let (η', proofs_η) ← liftExpr η maxLvl proofs_κ + return (← mkAppM ``Kernel.prod #[κ', η'], prod_lift_proof :: proofs_η) + +initialize registerLiftExpr liftProd + +/-- Lifts a composition-product of kernels by lifting the inner kernels. -/ +def liftCompProd (e : Expr) (maxLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.compProd do + throwError "Expected a composition of product of kernels, but got {e}." + let args := e.getAppArgs + let κ := args[args.size - 2]! + let η := args[args.size - 1]! + let (X, Y, xLvl, yLvl) ← getTypesFromKernel κ + let (_, Z, _, zLvl) ← getTypesFromKernel η + let ex ← constructMeasurableEquiv X xLvl maxLvl + let ey ← constructMeasurableEquiv Y yLvl maxLvl + let ez ← constructMeasurableEquiv Z zLvl maxLvl + let compProd_lift_proof ← mkAppM ``compProd_lift #[ex, ey, ez, κ, η] + let (κ', proofs_κ) ← liftExpr κ maxLvl proofs + let (η', proofs_η) ← liftExpr η maxLvl proofs_κ + return (← mkAppM ``Kernel.compProd #[κ', η'], compProd_lift_proof :: proofs_η) + +initialize registerLiftExpr liftCompProd + +/-- Lifts a composition of kernels by lifting the carrier type. -/ +def liftId (e : Expr) (maxLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.id do + throwError "Expected the identity kernel, but got {e}." + let (X, _, xLvl, _) ← getTypesFromKernel e + let ex ← constructMeasurableEquiv X xLvl maxLvl + let (X', _) ← getTypesFromMeasurableEquiv ex + let id_lift_proof ← mkAppM ``id_lift #[ex] + let mX' ← synthInstance (mkApp (mkConst ``MeasurableSpace [maxLvl]) X') + return (← mkAppOptM ``Kernel.id #[X', mX'], id_lift_proof :: proofs) + +initialize registerLiftExpr liftId + +/-- Lifts a discard kernel by lifting the carrier type. -/ +def liftDiscard (e : Expr) (maxLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.discard do + throwError "Expected the discard kernel, but got {e}." + let (X, _, xLvl, punitLvl) ← getTypesFromKernel e + let ex ← constructMeasurableEquiv X xLvl maxLvl + let (X', _) ← getTypesFromMeasurableEquiv ex + let discard_const := mkConst ``discard_lift [xLvl, maxLvl, punitLvl] + let discard_lift_proof ← mkAppM' discard_const #[ex] + let discard_const := mkConst ``Kernel.discard [maxLvl, maxLvl] + return (← mkAppOptM' discard_const #[X', none], discard_lift_proof :: proofs) + +initialize registerLiftExpr liftDiscard + +/-- Lifts a copy kernel by lifting the carrier type. -/ +def liftCopy (e : Expr) (maxLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.copy do + throwError "Expected the copy kernel, but got {e}." + let (X, _, xLvl, _) ← getTypesFromKernel e + let ex ← constructMeasurableEquiv X xLvl maxLvl + let (X', _) ← getTypesFromMeasurableEquiv ex + let copy_lift_proof ← mkAppM ``copy_lift #[ex] + return (← mkAppOptM ``Kernel.copy #[X', none], copy_lift_proof :: proofs) + +initialize registerLiftExpr liftCopy + +/-- Lifts a swap kernel by lifting the carrier types. -/ +def liftSwap (e : Expr) (maxLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.swap do + throwError "Expected the swap kernel, but got {e}." + let args := e.getAppArgs + let X := args[0]! + let Y := args[1]! + let xLvl := (← getDecLevel X) + let yLvl := (← getDecLevel Y) + let ex ← constructMeasurableEquiv X xLvl maxLvl + let ey ← constructMeasurableEquiv Y yLvl maxLvl + let (X', _) ← getTypesFromMeasurableEquiv ex + let (Y', _) ← getTypesFromMeasurableEquiv ey + let swap_lift_proof ← mkAppM ``swap_lift #[ex, ey] + return (← mkAppOptM ``Kernel.swap #[X', Y', none, none], swap_lift_proof :: proofs) + +initialize registerLiftExpr liftSwap + +/-- Lifts a kernel using `Kernel.lift`. -/ +def liftKernel (e : Expr) (maxLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + let (X, Y, xLvl, yLvl) ← getTypesFromKernel e + let ex ← constructMeasurableEquiv X xLvl maxLvl + let ey ← constructMeasurableEquiv Y yLvl maxLvl + let expr ← mkAppOptM ``Kernel.lift + #[none, none, none, none, none, none, none, none, ex, ey, e] + return (expr, proofs) + +initialize registerLiftExpr liftKernel + +/-- Constructs the finisher proof that concludes the lifting after rewriting the equalities. -/ +def finisherKernelLift (lhs rhs _ _ : Expr) (maxLvl : Level) : MetaM Expr := do + let (X, Y, xLvl, yLvl) ← getTypesFromKernel lhs + let ex ← constructMeasurableEquiv X xLvl maxLvl + let ey ← constructMeasurableEquiv Y yLvl maxLvl + mkAppM ``lift_congr #[ex, ey, lhs, rhs] + +initialize registerLiftFinisher finisherKernelLift + +end diff --git a/LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/KernelUnlift.lean b/LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/KernelUnlift.lean new file mode 100644 index 00000000..0290d10a --- /dev/null +++ b/LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/KernelUnlift.lean @@ -0,0 +1,196 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import LeanMachineLearning.Tactic.EqLift.Kernel.Lift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Unlift +public import LeanMachineLearning.Tactic.EqLift.Tactic.Kernel.Utils + +/-! +# Implementation of the `unlift_eq` tactic for kernels. + +This file contains functions that propagate the unlifting of lifted kernel expressions through +several operators and primitives, and constructs the necessary proofs for the `unlift_eq` tactic. +-/ + +public meta section + +open Lean Meta Parser.Tactic ProbabilityTheory ProbabilityTheory.Kernel + +/-- Unlifts a composition of kernels by unlifting the inner kernels. -/ +def unliftComposition (e : Expr) (eLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.comp do + throwError "Expected a composition of kernels, but got {e}." + let args := e.getAppArgs + let η' := args[args.size - 2]! + let κ' := args[args.size - 1]! + let (η, proofs_η) ← unliftExpr η' eLvl proofs + let (κ, proofs_κ) ← unliftExpr κ' eLvl proofs_η + let (X, Y, xLvl, yLvl) ← getTypesFromKernel η + let (Z, _, tLvl, _) ← getTypesFromKernel κ + let ex ← constructMeasurableEquiv X xLvl eLvl + let ey ← constructMeasurableEquiv Y yLvl eLvl + let ez ← constructMeasurableEquiv Z tLvl eLvl + let comp_unlift_proof ← mkAppM ``comp_lift #[ex, ey, ez, η, κ] + return (← mkAppM ``Kernel.comp #[η, κ], comp_unlift_proof :: proofs_κ) + +initialize registerUnliftExpr unliftComposition + +/-- Unlifts a parallel composition of kernels by unlifting the inner kernels. -/ +def unliftParallelComp (e : Expr) (eLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.parallelComp do + throwError "Expected a parallel composition of kernels, but got {e}." + let args := e.getAppArgs + let κ' := args[args.size - 2]! + let η' := args[args.size - 1]! + let (κ, proofs_κ) ← unliftExpr κ' eLvl proofs + let (η, proofs_η) ← unliftExpr η' eLvl proofs_κ + let (X, Y, xLvl, yLvl) ← getTypesFromKernel κ + let (Z, T, zLvl, tLvl) ← getTypesFromKernel η + let ex ← constructMeasurableEquiv X xLvl eLvl + let ey ← constructMeasurableEquiv Y yLvl eLvl + let ez ← constructMeasurableEquiv Z zLvl eLvl + let et ← constructMeasurableEquiv T tLvl eLvl + let parallelComp_unlift_proof ← mkAppM ``parallelComp_lift #[ex, ey, ez, et, κ, η] + return (← mkAppM ``Kernel.parallelComp #[κ, η], parallelComp_unlift_proof :: proofs_η) + +initialize registerUnliftExpr unliftParallelComp + +/-- Unlifts a product of kernels by unlifting the inner kernels. -/ +def unliftProd (e : Expr) (eLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.prod do + throwError "Expected a product of kernels, but got {e}." + let args := e.getAppArgs + let κ' := args[args.size - 2]! + let η' := args[args.size - 1]! + let (κ, proofs_κ) ← unliftExpr κ' eLvl proofs + let (η, proofs_η) ← unliftExpr η' eLvl proofs_κ + let (X, Y, xLvl, yLvl) ← getTypesFromKernel κ + let (_, Z, _, zLvl) ← getTypesFromKernel η + let ex ← constructMeasurableEquiv X xLvl eLvl + let ey ← constructMeasurableEquiv Y yLvl eLvl + let ez ← constructMeasurableEquiv Z zLvl eLvl + let prod_unlift_proof ← mkAppM ``prod_lift #[ex, ey, ez, κ, η] + return (← mkAppM ``Kernel.prod #[κ, η], prod_unlift_proof :: proofs_η) + +initialize registerUnliftExpr unliftProd + +/-- Unlifts a composition-product of kernels by unlifting the inner kernels. -/ +def unliftCompProd (e : Expr) (eLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.compProd do + throwError "Expected a composition of product of kernels, but got {e}." + let args := e.getAppArgs + let κ' := args[args.size - 2]! + let η' := args[args.size - 1]! + let (κ, proofs_κ) ← unliftExpr κ' eLvl proofs + let (η, proofs_η) ← unliftExpr η' eLvl proofs_κ + let (X, Y, xLvl, yLvl) ← getTypesFromKernel κ + let (_, Z, _, zLvl) ← getTypesFromKernel η + let ex ← constructMeasurableEquiv X xLvl eLvl + let ey ← constructMeasurableEquiv Y yLvl eLvl + let ez ← constructMeasurableEquiv Z zLvl eLvl + let compProd_unlift_proof ← mkAppM ``compProd_lift #[ex, ey, ez, κ, η] + return (← mkAppM ``compProd #[κ, η], compProd_unlift_proof :: proofs_η) + +initialize registerUnliftExpr unliftCompProd + +/-- Unlifts the identity kernel by unlifting the carrier type. -/ +def unliftId (e : Expr) (eLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.id do + throwError "Expected the identity kernel, but got {e}." + let (X', _, _, _) ← getTypesFromKernel e + let (X, xLvl) ← getOriginalType X' + if X == X' then + return (e, proofs) + else + let ex ← constructMeasurableEquiv X xLvl eLvl + let id_unlift_proof ← mkAppM ``id_lift #[ex] + let mX ← synthInstance (mkApp (mkConst ``MeasurableSpace [xLvl]) X) + return (← mkAppOptM ``Kernel.id #[X, mX], id_unlift_proof :: proofs) + +initialize registerUnliftExpr unliftId + +/-- Unlifts the discard kernel by unlifting the carrier type. -/ +def unliftDiscard (e : Expr) (eLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.discard do + throwError "Expected the discard kernel, but got {e}." + let (X', _, _, _) ← getTypesFromKernel e + let (X, xLvl) ← getOriginalType X' + if X == X' then + return (e, proofs) + else + let ex ← constructMeasurableEquiv X xLvl eLvl + let discard_const := mkConst ``discard_lift [xLvl, eLvl, Level.zero] + let discard_unlift_proof ← mkAppM' discard_const #[ex] + let discard_const := mkConst ``Kernel.discard [xLvl, 0] + return (← mkAppOptM' discard_const #[X, none], discard_unlift_proof :: proofs) + +initialize registerUnliftExpr unliftDiscard + +/-- Unlifts the copy kernel by unlifting the carrier type. -/ +def unliftCopy (e : Expr) (eLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.copy do + throwError "Expected the copy kernel, but got {e}." + let (X', _, _, _) ← getTypesFromKernel e + let (X, xLvl) ← getOriginalType X' + if X == X' then + return (e, proofs) + else + let ex ← constructMeasurableEquiv X xLvl eLvl + let copy_unlift_proof ← mkAppM ``copy_lift #[ex] + return (← mkAppOptM ``Kernel.copy #[X, none], copy_unlift_proof :: proofs) + +initialize registerUnliftExpr unliftCopy + +/-- Unlifts the swap kernel by unlifting the carrier types. -/ +def unliftSwap (e : Expr) (eLvl : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.swap do + throwError "Expected the swap kernel, but got {e}." + let args := e.getAppArgs + let X' := args[0]! + let Y' := args[1]! + let (X, xLvl) ← getOriginalType X' + let (Y, yLvl) ← getOriginalType Y' + if X == X' && Y == Y' then + return (e, proofs) + else + let ex ← constructMeasurableEquiv X xLvl eLvl + let ey ← constructMeasurableEquiv Y yLvl eLvl + let swap_unlift_proof ← mkAppM ``swap_lift #[ex, ey] + return (← mkAppOptM ``Kernel.swap #[X, Y, none, none], swap_unlift_proof :: proofs) + +initialize registerUnliftExpr unliftSwap + +/-- Unlifts a lifted kernel by returning the inner kernel. -/ +def unliftKernel (e : Expr) (_ : Level) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + unless e.isAppOf ``Kernel.lift do + throwError "Expected a lifted kernel, but got {e}." + let args := e.getAppArgs + let κ := args[args.size - 1]! + return (κ, proofs) + +initialize registerUnliftExpr unliftKernel + +/-- Constructs the finisher proof that concludes the unlifting after rewriting the equalities. -/ +def finisherKernelUnlift (_ _ lhs rhs : Expr) (maxLvl : Level) : MetaM Expr := do + let (X, Y, xLvl, yLvl) ← getTypesFromKernel lhs + let ex ← constructMeasurableEquiv X xLvl maxLvl + let ey ← constructMeasurableEquiv Y yLvl maxLvl + let lift_congr_expr ← mkAppM ``lift_congr #[ex, ey, lhs, rhs] + mkAppM ``Iff.symm #[lift_congr_expr] + +initialize registerUnliftFinisher finisherKernelUnlift + +end diff --git a/LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/Utils.lean b/LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/Utils.lean new file mode 100644 index 00000000..5f89dd32 --- /dev/null +++ b/LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/Utils.lean @@ -0,0 +1,88 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import LeanMachineLearning.Tactic.EqLift.ForMathlib.MeasurableEquiv +public import Mathlib.Probability.Kernel.Composition.Prod +public import Mathlib.Probability.Kernel.Composition.CompProd + +/-! +# Kernel lifting utilities + +This file provides helper functions for lifting and unlifting kernel expressions, including type +extraction and equivalence construction. + +## Main declarations + +* `getTypesFromKernel`: extracts carrier types and universe levels from kernel expressions. +* `constructMeasurableEquiv`: recursively builds measurable equivalences. +* `getOriginalType`: retrieves the original type from a lifted type. +-/ + +public meta section + +open Lean Meta ProbabilityTheory Elab Term + +/-- Extract `(X, Y, u, v)` from an expression of type `Kernel X Y`. -/ +def getTypesFromKernel (κ : Expr) : MetaM (Expr × Expr × Level × Level) := do + let κType ← inferType κ + match κType.getAppFn with + | Expr.const ``Kernel univs => + let args := κType.getAppArgs + if args.size < 2 then + throwError "Kernel type with insufficient arguments: {κType}." + let X := args[0]! + let Y := args[1]! + let xLevel := univs[0]! + let yLevel := univs[1]! + return (X, Y, xLevel, yLevel) + | _ => throwError "Expected a kernel type, got: {κType}." + +/-- Build a measurable equivalence for `e` into universe `maxLvl` (recursive on products). -/ +partial def constructMeasurableEquiv (e : Expr) (eLevel maxLvl : Level) : MetaM Expr := do + let ewhnf ← whnf e + match ewhnf.getAppFn with + | Expr.const ``PUnit _ | Expr.const ``Unit _ => + mkAppOptM' (Expr.const `MeasurableEquiv.punit [maxLvl, eLevel]) #[] + | Expr.const ``Prod univs => + let args := ewhnf.getAppArgs + let X := args[0]! + let Y := args[1]! + let xLevel := univs[0]! + let yLevel := univs[1]! + let ex ← constructMeasurableEquiv X xLevel maxLvl + let ey ← constructMeasurableEquiv Y yLevel maxLvl + let res ← mkAppOptM' (Expr.const ``MeasurableEquiv.prod [xLevel, yLevel, maxLvl, maxLvl]) + #[none, none, none, none, none, none, none, none, ex, ey] + return res + | _ => mkAppOptM' (Expr.const ``MeasurableEquiv.ulift [eLevel, maxLvl]) #[e, none] + +/-- Get departure and target types from a `MeasurableEquiv` expression. -/ +def getTypesFromMeasurableEquiv (e : Expr) : MetaM (Expr × Expr) := do + let equivT ← (whnf (← inferType e)) + match equivT.getAppFn with + | Expr.const ``MeasurableEquiv _ => + let args := equivT.getAppArgs + return (args[0]!, args[1]!) + | _ => throwError "Expected a MeasurableEquiv, got: {e}." + +/-- Get the original type from a lifted type. -/ +partial def getOriginalType (t : Expr) : MetaM (Expr × Level) := do + let twhnf ← whnf t + match twhnf.getAppFn with + | Expr.const ``PUnit _ | Expr.const ``Unit _ => + return (mkConst ``Unit [], 0) + | Expr.const ``ULift univs => + return (twhnf.getAppArgs[0]!, univs[1]!) + | Expr.const ``Prod _ => + let args := twhnf.getAppArgs + let (X, xLvl) ← getOriginalType args[0]! + let (Y, yLvl) ← getOriginalType args[1]! + return (← mkAppM ``Prod #[X, Y], .max xLvl yLvl) + | _ => + return (t, ← getDecLevel (← inferType t)) + +end diff --git a/LeanMachineLearning/Tactic/EqLift/Tactic/Lift.lean b/LeanMachineLearning/Tactic/EqLift/Tactic/Lift.lean new file mode 100644 index 00000000..8acc2f47 --- /dev/null +++ b/LeanMachineLearning/Tactic/EqLift/Tactic/Lift.lean @@ -0,0 +1,78 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +import Lean.Elab.Tactic.Location +public import LeanMachineLearning.Tactic.EqLift.Tactic.Utils +public import LeanMachineLearning.Tactic.EqLift.Tactic.Universe +public import LeanMachineLearning.Tactic.EqLift.Tactic.Location + +/-! +# Lift tactic + +This file defines the `lift_eq` tactic, which lifts an equality to a common universe level. It +propagates the lifting through the structure of the equality. The tactic is implemented in a way +that it can be easily extended to support new types of expressions by registering new lifting +functions. +-/ + +public meta section + +open Lean Elab Tactic Meta Parser.Tactic + +private initialize liftImplRef : IO.Ref (Array liftMetadata) ← IO.mkRef #[] + +/-- Registers a new lifting function for expressions. The function should take an expression, a +universe level (most likely the common universe level where the equality is being lifted), and a +list of proofs, and return a new expression and an updated list of proofs. -/ +def registerLiftExpr (f : liftMetadata) : IO Unit := do liftImplRef.modify (·.push f) + +private initialize liftFinisherRef : IO.Ref (Array (finisherMetadata)) ← IO.mkRef #[] + +/-- Registers a new finisher function for constructing the final proof of equality after lifting +inner expressions. The function should take the original left-hand side and right-hand side, the +transformed left-hand side and right-hand side, the common universe level, and return a proof of +equality. -/ +def registerLiftFinisher (f : finisherMetadata) : IO Unit := do + liftFinisherRef.modify (·.push f) + +/-- Lifts an expression to a common universe level using the registered lifting functions. -/ +def liftExpr := fun a b c ↦ transformExpr a b c liftImplRef + +/-- Gets the maximum universe level from an equality expression by collecting all universe levels +from the left-hand side and right-hand side of the equality. -/ +def getMaxLvl (eq : Expr) : MetaM Level := do + let univs ← collectExprUniverses eq + computeMaxLevel univs + +/-- Lifts an equality expression to a common universe level using the registered lifting functions +and finisher functions. -/ +def liftEquality := transformEquality getMaxLvl liftImplRef liftFinisherRef + +/-- Same as `liftEquality`, but allows specifying a universe level that will be taken into account +when computing the maximum universe level. -/ +def liftEqualityWithLevel (Lvl : Level) (eq : Expr) : MetaM (Expr × Expr) := do + let getMaxLvl := fun e ↦ do + let univs ← collectExprUniverses e + computeMaxLevel <| Lvl :: univs + transformEquality getMaxLvl liftImplRef liftFinisherRef eq + +/-- Transforms an equality expression by lifting both sides to a common universe level. + +The tactic supports location specifiers like `rw` or `simp`: +* `lift_eq` — applies to the goal +* `lift_eq at h` — applies to hypothesis `h` +* `lift_eq at h₁ h₂` — applies to multiple hypotheses +* `lift_eq at h ⊢` — applies to hypothesis `h` and the goal +* `lift_eq at *` — applies to all hypotheses and the goal +-/ +syntax (name := EqLift) "lift_eq" (ppSpace location)? : tactic + +elab_rules : tactic + | `(tactic| lift_eq $[$loc]?) => + expandOptLocation (mkOptionalNode loc) |> applyLocTactic <| liftEquality + +end diff --git a/LeanMachineLearning/Tactic/EqLift/Tactic/Location.lean b/LeanMachineLearning/Tactic/EqLift/Tactic/Location.lean new file mode 100644 index 00000000..89b931c3 --- /dev/null +++ b/LeanMachineLearning/Tactic/EqLift/Tactic/Location.lean @@ -0,0 +1,75 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public meta import Lean.Elab.Tactic.Location + +/-! +# Tactic location support + +This module provides utilities for applying tactics to multiple goals and hypotheses +specified by location patterns, following the standard Lean syntax (like in `rw` or `simp`). + +## Main declarations + +* `applyLocTactic`: applies a tactic to goals and hypotheses at specified locations. +-/ + +public meta section + +open Lean Elab Tactic Meta + +/-- Replace an equality in a goal or hypothesis with a transformed expression, using a provided +transformation function. -/ +def replaceEquality (goal : MVarId) (fvarId : Option FVarId) + (transform : Expr → MetaM (Expr × Expr)) : TacticM MVarId := do + goal.withContext do + let expr ← match fvarId with + | some fid => do + let decl ← fid.getDecl + pure decl.type + | none => goal.getType + let (lift_expr, eq_proof) ← transform expr + match fvarId with + | some fid => do + let mvarId ← getMainGoal + let h_proof ← mkEqMP eq_proof (mkFVar fid) + let userName := (← fid.getDecl).userName + let mvarId ← mvarId.assert userName lift_expr h_proof + let mvarId ← mvarId.tryClear fid + let (_, mvarId) ← mvarId.intro1P + pure mvarId + | none => do + let mvarId ← getMainGoal + mvarId.replaceTargetEq lift_expr eq_proof + +/-- Apply a given transformation to all goals and/or hypotheses specified by a `Location`. -/ +def applyLocTactic (loc : Location) (transform : Expr → MetaM (Expr × Expr)) : + TacticM Unit := do + match loc with + | Location.targets hyps target => + for hyp in hyps do + let hFVarId ← getFVarId hyp + let newGoal ← replaceEquality (← getMainGoal) (some hFVarId) transform + replaceMainGoal [newGoal] + if target then + let newGoal ← replaceEquality (← getMainGoal) none transform + replaceMainGoal [newGoal] + | Location.wildcard => + let goal ← getMainGoal + goal.withContext do + let lctx ← getLCtx + let mut currentGoal := goal + for decl in lctx do + if decl.isImplementationDetail then continue + try + currentGoal ← replaceEquality currentGoal (some decl.fvarId) transform + replaceMainGoal [currentGoal] + catch _ => continue + try + currentGoal ← replaceEquality currentGoal none transform + replaceMainGoal [currentGoal] + catch _ => pure () diff --git a/LeanMachineLearning/Tactic/EqLift/Tactic/Universe.lean b/LeanMachineLearning/Tactic/EqLift/Tactic/Universe.lean new file mode 100644 index 00000000..11124bd5 --- /dev/null +++ b/LeanMachineLearning/Tactic/EqLift/Tactic/Universe.lean @@ -0,0 +1,61 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import Lean.Meta.DecLevel +public import Lean.Meta.Transform +public import Lean.Util.Recognizers +import Mathlib.Probability.Kernel.Composition.Prod + +/-! +# Universe level utilities + +This file provides utilities for working with universe levels in metaprograms. +It includes conversion functions between levels and syntax, and universe level collection. + +## Main declarations + +* `collectExprUniverses`: recursively collects universe levels from expressions. +* `getUniverseFromEq`: extracts the universe level from the left-hand side of an equality +expression. +-/ + +public meta section + +open Lean Meta ProbabilityTheory + +/-- Recursively traverses an expression and collects all universe levels. -/ +def collectExprUniverses.aux (e : Expr) : List Level := + match e with + | Expr.const _ univs => univs + | Expr.sort u => [u] + | Expr.app f a => aux f ++ aux a + | Expr.lam _ t b _ => aux t ++ aux b + | Expr.forallE _ t b _ => aux t ++ aux b + | Expr.letE _ t v b _ => aux t ++ aux v ++ aux b + | Expr.mdata _ b => aux b + | Expr.proj _ _ b => aux b + | Expr.bvar _ | Expr.fvar _ | Expr.mvar _ | Expr.lit _ => [] + +/-- Recursively traverse an expression and collect universe levels found. +Returns a list of all unique universe levels encountered. -/ +def collectExprUniverses (e : Expr) : MetaM (List Level) := do + let e ← instantiateMVars e + let e ← zetaReduce e + return (collectExprUniverses.aux e).eraseDups + +/-- Compute the maximum universe level from a list of levels. -/ +def computeMaxLevel (levels : List Level) : MetaM Level := + match levels with + | [] => throwError "Expected at least one universe level, got an empty list." + | head :: tail => pure (tail.foldl Level.max head) + +/-- Extract the universe level from the left side of an equality expression. -/ +def getLevelFromEq (eq : Expr) : MetaM Level := do + let eq ← whnf (← zetaReduce (← instantiateMVars eq)) + let eq := eq.consumeMData + let some (_, lhs, _) := eq.eq? | throwError "Expected an equality, got: {eq}." + getDecLevel (← inferType lhs) diff --git a/LeanMachineLearning/Tactic/EqLift/Tactic/Unlift.lean b/LeanMachineLearning/Tactic/EqLift/Tactic/Unlift.lean new file mode 100644 index 00000000..a6210b05 --- /dev/null +++ b/LeanMachineLearning/Tactic/EqLift/Tactic/Unlift.lean @@ -0,0 +1,66 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +import Lean.Elab.Tactic.Location +public import LeanMachineLearning.Tactic.EqLift.Tactic.Utils +public import LeanMachineLearning.Tactic.EqLift.Tactic.Universe +public import LeanMachineLearning.Tactic.EqLift.Tactic.Location + +/-! +# Unlift tactic + +This file defines the `unlift_eq` tactic, which performs the inverse operation of `lift_eq`. It +takes an equality that has been lifted to a common universe level and attempts to unlift it back to +its original form. The tactic is designed to work with various types of expressions, and can be +extended by registering new unlift functions. +-/ + +public meta section + +open Lean Elab Tactic Meta Parser.Tactic + +private initialize unliftImplRef : IO.Ref (Array liftMetadata) ← IO.mkRef #[] + +/-- Registers a new unlifting function for expressions. The function should take an expression, a +universe level (most likely the common universe level where the equality is being lifted), and a +list of proofs, and return a new expression and an updated list of proofs. -/ +def registerUnliftExpr (f : liftMetadata) : IO Unit := do unliftImplRef.modify (·.push f) + +private initialize unliftFinisherRef : IO.Ref (Array finisherMetadata) ← IO.mkRef #[] + +/-- Registers a new finisher function for constructing the final proof of equality after unlifting +inner expressions. The function should take the original left-hand side and right-hand side, the +transformed left-hand side and right-hand side, the common universe level, and return a proof of +equality. -/ +def registerUnliftFinisher (f : finisherMetadata) : IO Unit := do + unliftFinisherRef.modify (·.push f) + +/-- Unlifts an expression that has been lifted to a common universe level using the registered +unlifting functions. -/ +def unliftExpr := fun a b c ↦ transformExpr a b c unliftImplRef + +/-- Unlifts an equality expression that has been lifted to a common universe level using the +registered unlifting functions and finisher functions. -/ +def unliftEquality := transformEquality getLevelFromEq unliftImplRef unliftFinisherRef + +/-- Performs the inverse operation of `lift_eq`, transforming an equality that has been lifted to a +common universe level back to its original form. + +The tactic supports location specifiers like `rw` or `simp`: +* `unlift_eq` — applies to the goal +* `unlift_eq at h` — applies to hypothesis `h` +* `unlift_eq at h₁ h₂` — applies to multiple hypotheses +* `unlift_eq at h ⊢` — applies to hypothesis `h` and the goal +* `unlift_eq at *` — applies to all hypotheses and the goal +-/ +syntax (name := EqUnlift) "unlift_eq" (ppSpace location)? : tactic + +elab_rules : tactic + | `(tactic| unlift_eq $[$loc]?) => + expandOptLocation (mkOptionalNode loc) |> applyLocTactic <| unliftEquality + +end diff --git a/LeanMachineLearning/Tactic/EqLift/Tactic/Utils.lean b/LeanMachineLearning/Tactic/EqLift/Tactic/Utils.lean new file mode 100644 index 00000000..d93e7591 --- /dev/null +++ b/LeanMachineLearning/Tactic/EqLift/Tactic/Utils.lean @@ -0,0 +1,91 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public meta import Lean.Meta.Tactic.Replace +public meta import Lean.Meta.Tactic.Rewrite + +/-! + +# Lift and Unlift utilities + +This file provides utility functions for lifting and unlifting equalities. + +-/ + +public meta section + +open Lean Elab Tactic Meta Parser.Tactic + +/-- A type alias for lifting/unlifting functions. -/ +abbrev liftMetadata := Expr → Level → List Expr → MetaM (Expr × List Expr) + +/-- A type alias for finisher functions that construct the final proof of equality after lifting/ +unlifting inner expressions. -/ +abbrev finisherMetadata := Expr → Expr → Expr → Expr → Level → MetaM Expr + +/-- Transforms an expression using the registered lifting/unlifting functions given in `impl_ref`. Returns the first successful transformation along with the updated list of proofs. -/ +def transformExpr (e : Expr) (maxLvl : Level) (proofs : List Expr) + (impl_ref : IO.Ref (Array (liftMetadata))) : MetaM (Expr × List Expr) := do + let handlers ← impl_ref.get + let (lift_expr, proofs) ← handlers.firstM (fun h => h e maxLvl proofs) <|> do + throwError "No transform handler found for {e}." + return (lift_expr, proofs) + +/-- Rewrites the type of `mvarId` at the `n`-th occurrence using `heq`.-/ +def Lean.MVarId.nthRewrite (mvarId : MVarId) (n : Nat) (heq : Expr) : MetaM MVarId := do + let r ← mvarId.rewrite (← mvarId.getType) heq (config := { occs := .pos [n] }) + mvarId.replaceTargetEq r.eNew r.eqProof + +/-- Constructs a proof of equality between the original and transformed expressions using the provided proofs and finisher functions. -/ +def constructProof (eqProofType lhs rhs lhs_t rhs_t : Expr) (maxLvl : Level) (proofs : List Expr) + (finisher_ref : IO.Ref (Array finisherMetadata)) : MetaM Expr := do + let mvar ← mkFreshExprSyntheticOpaqueMVar eqProofType + let mvarId := mvar.mvarId! + let propext := mkConst ``propext + match ← mvarId.apply propext with + | [mvarId] => + let proofs := proofs.reverse + let mut mvarId := mvarId + for proof in proofs do + mvarId ← mvarId.nthRewrite 1 proof + let handlers ← finisher_ref.get + let e ← handlers.firstM (fun h => do + let finisher ← h lhs rhs lhs_t rhs_t maxLvl + unless ← isDefEq (← mvarId.getType) (← inferType finisher) do + throwError "Type mismatch: expected {← mvarId.getType}, got {← inferType finisher}." + mvarId.assign finisher + instantiateMVars mvar + ) <|> do + throwError m!"No finisher found for {eqProofType}." + return e + | _ => + throwError "Failed to apply propext while building kernel_lift equivalence proof for + {eqProofType}." + +/-- Lifts or unlifts an equality expression by transforming both sides using the registered lifting/ +unlifting functions. Returns the transformed equality and a proof of equality between the original +and transformed expressions. -/ +def transformEquality (getLvl : Expr → MetaM Level) (lift_ref : IO.Ref (Array liftMetadata)) + (finisher_ref : IO.Ref (Array finisherMetadata)) (eq : Expr) : MetaM (Expr × Expr) := do + let e ← whnfR <| ← zetaReduce <| ← instantiateMVars eq + let e := e.consumeMData + let lvl ← getLvl eq + let some (_, lhs, rhs) := e.eq? | throwError "Expected an equality, got: {e}." + let (lhs_transformed, proofs) ← transformExpr lhs lvl [] lift_ref + let (rhs_transformed, proofs) ← transformExpr rhs lvl proofs lift_ref + let eq_transformed ← mkEq lhs_transformed rhs_transformed + let eq_proof_type ← mkEq eq eq_transformed + let proof ← constructProof + eq_proof_type + lhs rhs + lhs_transformed rhs_transformed + lvl + proofs + finisher_ref + return (eq_transformed, proof) + +end diff --git a/LeanMachineLearning/Tactic/KernelHom.lean b/LeanMachineLearning/Tactic/KernelHom.lean new file mode 100644 index 00000000..efff87e0 --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom.lean @@ -0,0 +1,20 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ + +module -- shake: keep-all --deprecated_module: ignore + +public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.Kernel +public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.LIntegral +public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.MeasurableEquiv +public import LeanMachineLearning.Tactic.KernelHom.Kernel.Hom +public import LeanMachineLearning.Tactic.KernelHom.Kernel.MonoidalComp +public import LeanMachineLearning.Tactic.KernelHom.Tactic.Delaborators +public import LeanMachineLearning.Tactic.KernelHom.Tactic.HomKernel +public import LeanMachineLearning.Tactic.KernelHom.Tactic.KernelCat +public import LeanMachineLearning.Tactic.KernelHom.Tactic.KernelDiagram +public import LeanMachineLearning.Tactic.KernelHom.Tactic.KernelHom +public import LeanMachineLearning.Tactic.KernelHom.Tactic.Reassoc +public import LeanMachineLearning.Tactic.KernelHom.Tactic.Utils diff --git a/LeanMachineLearning/Tactic/KernelHom/ForMathlib/Kernel.lean b/LeanMachineLearning/Tactic/KernelHom/ForMathlib/Kernel.lean new file mode 100644 index 00000000..cb49f7ab --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom/ForMathlib/Kernel.lean @@ -0,0 +1,53 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import Mathlib.Probability.Kernel.Composition.ParallelComp + +/-! +# Kernel utilities + +This file provides helper lemmas for working with kernels. + +## Main declarations + +* `comap_parallelComp_comap`: the comap of a parallel composition is the parallel composition of + the comaps. +* `map_parallelComp_map`: the map of a parallel composition is the parallel composition of the maps. +-/ + +@[expose] public section + +open ProbabilityTheory MeasureTheory + +variable {α β γ ι : Type*} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace γ] + [MeasurableSpace ι] + +namespace ProbabilityTheory.Kernel + +lemma comap_parallelComp_comap {α₂ γ₂ : Type*} [MeasurableSpace α₂] [MeasurableSpace γ₂] + (κ : Kernel α β) (η : Kernel γ ι) [IsSFiniteKernel κ] [IsSFiniteKernel η] + {f : α₂ → α} {g : γ₂ → γ} (hf : Measurable f) (hg : Measurable g) : + κ.comap f hf ∥ₖ η.comap g hg = (κ ∥ₖ η).comap (fun a ↦ (f a.1, g a.2)) (by fun_prop) := by + ext : 1 + rw [Kernel.parallelComp_apply, Kernel.comap_apply, Kernel.comap_apply, Kernel.comap_apply, + Kernel.parallelComp_apply] + +lemma map_parallelComp_map {β₂ ι₂ : Type*} [MeasurableSpace β₂] [MeasurableSpace ι₂] + (κ : Kernel α β) (η : Kernel γ ι) [IsSFiniteKernel κ] [IsSFiniteKernel η] + {f : β → β₂} {g : ι → ι₂} (hf : Measurable f) (hg : Measurable g) : + κ.map f ∥ₖ η.map g = (κ ∥ₖ η).map (fun a ↦ (f a.1, g a.2)) := by + ext a s hs + rw [Kernel.parallelComp_apply', Kernel.lintegral_map, Kernel.map_apply', + Kernel.parallelComp_apply'] + · congr with x + rw [Kernel.map_apply' _ (by fun_prop) _ (by measurability)] + congr + all_goals try fun_prop + all_goals try measurability + exact measurable_measure_prodMk_left hs + +end ProbabilityTheory.Kernel diff --git a/LeanMachineLearning/Tactic/KernelHom/ForMathlib/LIntegral.lean b/LeanMachineLearning/Tactic/KernelHom/ForMathlib/LIntegral.lean new file mode 100644 index 00000000..7bb7b435 --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom/ForMathlib/LIntegral.lean @@ -0,0 +1,27 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import Mathlib.MeasureTheory.Integral.Lebesgue.Countable + +/-! +# Lebesgue integral utilities + +This file provides helper lemmas for the Lebesgue integral with Dirac measures. + +## Main declarations + +* `lintegral_lintegral_dirac`: computing nested integrals with Dirac measures. +-/ + +@[expose] public section + +open ENNReal MeasureTheory + +lemma lintegral_lintegral_dirac {α β : Type*} [MeasurableSpace α] [MeasurableSpace β] + {μ : Measure α} {f : β → ℝ≥0∞} {g : α → β} + (hf : Measurable f) : ∫⁻ a, ∫⁻ b, f b ∂(Measure.dirac (g a)) ∂μ = ∫⁻ a, f (g a) ∂μ := by + simp_rw [lintegral_dirac' _ hf] diff --git a/LeanMachineLearning/Tactic/KernelHom/ForMathlib/MeasurableEquiv.lean b/LeanMachineLearning/Tactic/KernelHom/ForMathlib/MeasurableEquiv.lean new file mode 100644 index 00000000..4a6b95f1 --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom/ForMathlib/MeasurableEquiv.lean @@ -0,0 +1,26 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import Mathlib.MeasureTheory.MeasurableSpace.Embedding + +/-! +# Measurable equivalences +-/ + +@[expose] public section + +namespace MeasurableEquiv + +variable (α : Type*) [MeasurableSpace α] + +/-- The identity measurable equivalence. -/ +def id : α ≃ᵐ α where + toEquiv := .refl α + measurable_toFun := measurable_id + measurable_invFun := measurable_id + +end MeasurableEquiv diff --git a/LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean b/LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean new file mode 100644 index 00000000..e4ff8685 --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean @@ -0,0 +1,267 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import LeanMachineLearning.Tactic.EqLift.Kernel.Lift +public import Mathlib.Combinatorics.Quiver.ReflQuiver +public import Mathlib.Probability.Kernel.Category.SFinKer + +/-! +# Kernel morphisms + +This file defines the transformation between categorical morphisms in `SFinKer` and kernel objects. + +## Main declarations + +* `fromHom`: transforms a categorical morphism in `SFinKer` to a `Kernel`. +* `hom`: transforms a `Kernel` to a categorical morphism in `SFinKer`. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory MeasurableEquiv CategoryTheory +open scoped SFinKer CategoryTheory CategoryTheory.MonoidalCategory + +namespace ProbabilityTheory.Kernel + +variable {X Y T Z : Type*} [MeasurableSpace X] [MeasurableSpace Y] [MeasurableSpace T] + [MeasurableSpace Z] + +section + +variable {SX SY ST SZ : SFinKer} {ex : SX ≃ᵐ X} {ey : SY ≃ᵐ Y} + +/-- Transform a morphism in `SFinKer` into a kernel. -/ +noncomputable def fromHom (κ : SX ⟶ SY) : Kernel X Y := (κ.1.comap ex.symm (by fun_prop)).map ey + +instance {κ : SX ⟶ SY} : IsSFiniteKernel (fromHom (ex := ex) (ey := ey) κ) := by + simp only [fromHom] + have := κ.2 + infer_instance + +/-- Transform a kernel into a morphism in `SFinKer`. -/ +noncomputable def hom (κ : Kernel X Y) [IsSFiniteKernel κ] : SX ⟶ SY := by + refine ⟨(κ.map ey.symm).comap ex (by fun_prop), ?_⟩ + have := κ.2 + infer_instance + +lemma hom_apply (κ : Kernel X Y) [IsSFiniteKernel κ] (a : SX) : + (κ.hom (ex := ex) (ey := ey)).1 a = (κ.map ey.symm) (ex a) := rfl + +lemma hom_apply' (κ : Kernel X Y) [IsSFiniteKernel κ] (a : SX) {s : Set SY} + (hs : MeasurableSet s) : + (κ.hom (ex := ex) (ey := ey)).1 a s = κ (ex a) (ey '' s) := by + simp only [hom, coe_comap, Function.comp_apply] + rw [map_apply' _ ey.symm.measurable _ hs, preimage_symm] + +instance {κ : Kernel X Y} [IsDeterministic κ] [IsMarkovKernel κ] : + Deterministic (hom (ex := ex) (ey := ey) κ) := by + set κ_hom := hom (ex := ex) (ey := ey) κ + have : IsDeterministic κ_hom.hom := by + refine ⟨?_⟩ + ext a s hs + simp only [hom, κ_hom] + have := κ.parallelComp_self_comp_copy + have := DFunLike.congr_fun (x := ex a) this + have := DFunLike.congr_fun (x := ey.prod ey '' s) this + rw [comap_parallelComp_comap, map_parallelComp_map, comp_apply', comp_apply', + copy, deterministic_apply, lintegral_dirac', comap_apply', map_apply', parallelComp_apply', + lintegral_comap, lintegral_map] + · rw [comp_apply', comp_apply', copy, deterministic_apply, lintegral_dirac', + parallelComp_apply'] at this + · convert this + all_goals try simp + · ext y + simp [MeasurableEquiv.prod] + aesop + · simp only [copy, deterministic_apply] + rw [Measure.dirac_apply', Measure.dirac_apply'] + · refine Set.indicator_eq_indicator ?_ rfl + simp [MeasurableEquiv.prod] + aesop + · exact (measurableSet_image (ey.prod ey)).mpr hs + · exact hs + all_goals try measurability + · exact Kernel.measurable_coe _ (by measurability) + all_goals try measurability + · exact Kernel.measurable_coe _ hs + · exact Kernel.measurable_coe _ hs + have : IsMarkovKernel κ_hom.hom := + have : IsMarkovKernel (κ.map ey.symm) := + IsMarkovKernel.map _ (by fun_prop) + IsMarkovKernel.comap _ (by fun_prop) + exact SX.deterministic_deterministic SY κ_hom.hom + +end + +lemma hom_congr (SX SY : SFinKer) (ex : SX ≃ᵐ X) (ey : SY ≃ᵐ Y) + (κ η : Kernel X Y) [IsSFiniteKernel κ] [IsSFiniteKernel η] : + κ = η ↔ κ.hom (ex := ex) (ey := ey) = η.hom (ex := ex) (ey := ey) := by + constructor + · grind + · intro h + ext a s hs + replace h := DFunLike.congr (x := ex.symm a) (congrArg SFinKer.Hom.hom h) rfl + replace h := DFunLike.congr (x := ey.symm '' s) h rfl + rw [hom_apply', hom_apply'] at h + · simp only [apply_symm_apply] at h + rwa [image_symm, image_preimage] at h + · measurability + · measurability + +section + +variable (SX SY SZ ST : SFinKer) (ex : SX ≃ᵐ X) (ey : SY ≃ᵐ Y) (ez : SZ ≃ᵐ Z) (et : ST ≃ᵐ T) + +lemma comp_hom (η : Kernel X Y) (κ : Kernel Z X) [IsSFiniteKernel η] [IsSFiniteKernel κ] : + κ.hom (ex := ez) (ey := ex) ≫ η.hom (ex := ex) (ey := ey) = + (η ∘ₖ κ).hom (ex := ez) (ey := ey) := by + ext a s hs + dsimp + rw [hom_apply', comp_apply', comp_apply', hom_apply, lintegral_map] + · congr with y + simp [hom_apply' _ _ hs] + all_goals try fun_prop + all_goals try measurability + · exact Kernel.measurable_coe η.hom.hom hs + +lemma parallelComp_hom (κ : Kernel X Y) (η : Kernel Z T) [IsSFiniteKernel η] [IsSFiniteKernel κ] : + κ.hom (ex := ex) (ey := ey) ⊗ₘ η.hom (ex := ez) (ey := et) = + hom (ex := ex.prod ez) (ey := ey.prod et) (κ ∥ₖ η) := by + ext : 1; dsimp + simp only [hom] + rw [id_parallelComp_comp_parallelComp_id, comap_parallelComp_comap, map_parallelComp_map] + · rfl + all_goals fun_prop + +lemma id_hom : 𝟙 SX = Kernel.id.hom (ex := ex) (ey := ex) := by + ext; dsimp + rw [hom_apply', id_apply, id_apply, Measure.dirac_apply', Measure.dirac_apply'] + · exact Set.indicator_eq_indicator (by simp) rfl + all_goals measurability + +lemma whiskerLeft (κ : Kernel X Y) [IsSFiniteKernel κ] : SZ ◁ κ.hom (ex := ex) (ey := ey) = + (Kernel.id (α := Z) ∥ₖ κ).hom (ex := ez.prod ex) (ey := ez.prod ey) := by + ext _ _ hs; dsimp + simp only [hom] + rw [parallelComp_apply, comap_apply, map_apply, id_apply, + comap_apply, map_apply, parallelComp_apply, id_apply] + · simp only [Measure.dirac_prod, MeasurableEquiv.prod] + rw [Measure.map_map, Measure.map_map, Measure.map_apply, Measure.map_apply] + · congr with y + · simp + · simp + all_goals try fun_prop + all_goals exact hs + all_goals fun_prop + +lemma whiskerRight (κ : Kernel X Y) [IsSFiniteKernel κ] : + κ.hom (ex := ex) (ey := ey) ▷ SZ = + (κ ∥ₖ Kernel.id (α := Z)).hom (ex := ex.prod ez) (ey := ey.prod ez) := by + ext _ _ hs; dsimp + simp only [hom] + rw [parallelComp_apply, comap_apply, map_apply, id_apply, comap_apply, map_apply, + parallelComp_apply, id_apply] + · simp only [Measure.prod_dirac, MeasurableEquiv.prod] + rw [Measure.map_map, Measure.map_map, Measure.map_apply, Measure.map_apply] + · congr with y + · simp + · simp + all_goals try fun_prop + all_goals exact hs + all_goals fun_prop + +open scoped ComonObj + +lemma counit : ε[SX] = (Kernel.discard X).hom (ex := ex) (ey := punit) := by + ext : 1; dsimp + simp only [hom, discard] + rw [deterministic_map (by fun_prop) (by fun_prop)] + rfl + +lemma comul : Δ[SX] = (Kernel.copy X).hom (ex := ex) (ey := ex.prod ex) := by + ext : 1; dsimp + simp only [hom, copy] + rw [deterministic_map (by fun_prop) (by fun_prop)] + congr with x + all_goals simp [MeasurableEquiv.prod] + +lemma braiding_hom : (β_ SX SY).hom = + (Kernel.swap X Y).hom (ex := ex.prod ey) (ey := ey.prod ex) := by + ext : 1; dsimp + simp only [hom, swap] + rw [deterministic_map (by fun_prop) (by fun_prop)] + congr with x + all_goals simp [MeasurableEquiv.prod] + +variable {X₀ Y₀ Z₀ : Type*} [MeasurableSpace X₀] [MeasurableSpace Y₀] [MeasurableSpace Z₀] + (ex₀ : X ≃ᵐ X₀) (ey₀ : Y ≃ᵐ Y₀) (ez₀ : Z ≃ᵐ Z₀) + +lemma leftUnitor_hom : (λ_ SX).hom = hom (ex := punit.prod ex) (ey := ex) + (lift (Kernel.id.map (Prod.snd : PUnit × X₀ → X₀)) (ex := punit.prod ex₀) (ey := ex₀)) := by + ext; dsimp + rw [hom_apply', lift_apply', id_map (by fun_prop), id_map (by fun_prop), deterministic_apply', + deterministic_apply', Set.image] + · refine Set.indicator_eq_indicator ?_ rfl + simp [MeasurableEquiv.prod] + all_goals measurability + +lemma leftUnitor_inv : (λ_ SX).inv = hom (ex := ex) (ey := punit.prod ex) + (lift (Kernel.id.map (fun x ↦ (PUnit.unit, x))) (ex := ex₀) (ey := punit.prod ex₀)) := by + ext; dsimp + rw [hom_apply', lift_apply', id_map (by fun_prop), id_map (by fun_prop), deterministic_apply', + deterministic_apply'] + · refine Set.indicator_eq_indicator ?_ rfl + simp [Set.image, MeasurableEquiv.prod] + constructor + all_goals simp_all + all_goals measurability + +lemma rightUnitor_hom : (ρ_ SX).hom = hom (ex := ex.prod punit) (ey := ex) + (lift (Kernel.id.map (Prod.fst : X₀ × PUnit → X₀)) (ex := ex₀.prod punit) (ey := ex₀)) := by + ext; dsimp + rw [hom_apply', lift_apply', id_map (by fun_prop), id_map (by fun_prop), deterministic_apply', + deterministic_apply'] + · refine Set.indicator_eq_indicator ?_ rfl + simp [MeasurableEquiv.prod] + all_goals measurability + +lemma rightUnitor_inv : (ρ_ SX).inv = hom (ex := ex) (ey := ex.prod punit) + (lift (Kernel.id.map (fun x ↦ (x, PUnit.unit))) (ex := ex₀) (ey := ex₀.prod punit)) := by + ext; dsimp + rw [hom_apply', lift_apply', id_map (by fun_prop), id_map (by fun_prop), deterministic_apply', + deterministic_apply'] + · refine Set.indicator_eq_indicator ?_ rfl + simp [Set.image, MeasurableEquiv.prod] + constructor + all_goals simp_all + all_goals measurability + +lemma associator_hom : (α_ SX SY SZ).hom = + hom (ex := (ex.prod ey).prod ez) (ey := ex.prod (ey.prod ez)) + (lift (Kernel.deterministic prodAssoc (by fun_prop)) (ex := (ex₀.prod ey₀).prod ez₀) + (ey := ex₀.prod (ey₀.prod ez₀))) := by + ext; dsimp + simp only [hom] + rw [comap_apply', map_apply', lift_apply', deterministic_apply', deterministic_apply'] + · refine Set.indicator_eq_indicator ?_ rfl + simp [MeasurableEquiv.prod, prodAssoc] + all_goals measurability + +lemma associator_inv : (α_ SX SY SZ).inv = + hom (ex := ex.prod (ey.prod ez)) (ey := (ex.prod ey).prod ez) + (lift (Kernel.deterministic prodAssoc.symm (by fun_prop)) (ex := ex₀.prod (ey₀.prod ez₀)) + (ey := (ex₀.prod ey₀).prod ez₀)) := by + ext; dsimp + simp only [hom] + rw [comap_apply', map_apply', lift_apply', deterministic_apply', deterministic_apply'] + · refine Set.indicator_eq_indicator ?_ rfl + simp [MeasurableEquiv.prod, prodAssoc] + all_goals measurability + +end + +end ProbabilityTheory.Kernel diff --git a/LeanMachineLearning/Tactic/KernelHom/Kernel/MonoidalComp.lean b/LeanMachineLearning/Tactic/KernelHom/Kernel/MonoidalComp.lean new file mode 100644 index 00000000..272f52ca --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom/Kernel/MonoidalComp.lean @@ -0,0 +1,105 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.LIntegral +public import LeanMachineLearning.Tactic.KernelHom.Kernel.Hom + +/-! +# Measurable coherence + +This file introduces the monoidal composition for s-finite kernels (noted `⊗≫ₖ`). +## Main declarations + +* `MeasurableCoherence`: class witnessing measurable equivalences between types. +* `monoComp`: monoidal composition of kernels using measurable equivalences to transport to + `SFinKer`. +* `hom_monoComp`: the `SFinKer` morphism of the kernelized monoidal composition is the monoidal + composition of the morphisms in `SFinKer`. +-/ + +@[expose] public section + +open CategoryTheory MeasureTheory ProbabilityTheory MeasurableEquiv + +open scoped MonoidalCategory SFinKer + +/-- A class witnessing the existence of a measurable equivalence between two measurable spaces. -/ +class MeasurableCoherence (X Y : Type*) [MeasurableSpace X] [MeasurableSpace Y] where + /-- A measurable equivalence between `X` and `Y`. -/ + miso : X ≃ᵐ Y + +namespace MeasurableCoherence + +variable {X Y : Type*} [MeasurableSpace X] [MeasurableSpace Y] [mXY : MeasurableCoherence X Y] + +instance : MeasurableCoherence X X where + miso := MeasurableEquiv.refl X + +/-- Given measurable equivalences `ex : X ≃ᵐ X'` and `ey : Y ≃ᵐ Y'`, we can transport the +`MeasurableCoherence` instance from `X` and `Y` to `X'` and `Y'`. -/ +@[reducible] +def TransEquiv {X' Y' : Type*} [MeasurableSpace X'] [MeasurableSpace Y'] + (ex : X' ≃ᵐ X) (ey : Y' ≃ᵐ Y) : MeasurableCoherence X' Y' where + miso := ex.trans <| mXY.miso.trans ey.symm + +/-- `MeasurableCoherence` gives an instance of `MonoidalCoherence` in the `SFinKer` category. -/ +@[reducible] +noncomputable def monoidalCoherence {SX SY : SFinKer} (ex : SX.carrier ≃ᵐ X) + (ey : SY.carrier ≃ᵐ Y) : MonoidalCoherence SX SY where + iso := by + let e := ex.trans <| mXY.miso.trans ey.symm + refine ⟨⟨Kernel.id.map e, inferInstance⟩, + ⟨Kernel.id.map e.symm, inferInstance⟩, ?_, ?_⟩ + all_goals ext; dsimp + · rw [Kernel.id_map (by fun_prop), Kernel.id_map (by fun_prop), + Kernel.deterministic_comp_deterministic, Kernel.id] + congr + simp + · rw [Kernel.id_map (by fun_prop), Kernel.id_map (by fun_prop), + Kernel.deterministic_comp_deterministic, Kernel.id] + congr + simp + +end MeasurableCoherence + +namespace ProbabilityTheory.Kernel + +open MeasurableCoherence + +variable {W X Y Z : Type*} [MeasurableSpace W] [MeasurableSpace X] [MeasurableSpace Y] + [MeasurableSpace Z] {SW SX SY SZ : SFinKer} (ew : SW ≃ᵐ W) (ex : SX ≃ᵐ X) + (ey : SY ≃ᵐ Y) (ez : SZ ≃ᵐ Z) [MeasurableCoherence X Y] (κ : Kernel W X) [IsSFiniteKernel κ] + (η : Kernel Y Z) [IsSFiniteKernel η] + +/-- The kernelized version of the monoidal composition of kernels using the `SFinKer` category. +It uses arbitrary measurable equivalences to transport the kernels to the `SFinKer` category. -/ +noncomputable def monoComp₀ : Kernel W Z := + have := monoidalCoherence ex ey + fromHom (ex := ew) (ey := ez) <| hom (ex := ew) (ey := ex) κ ⊗≫ + hom (ex := ey) (ey := ez) η + +instance monoComp'_sfinite : IsSFiniteKernel (monoComp₀ ew ex ey ez κ η) := by + simp only [monoComp₀] + infer_instance + +/-- The kernelized version of the monoidal composition of kernels using the `SFinKer` category. -/ +noncomputable abbrev monoComp : Kernel W Z := + monoComp₀ + (SW := SFinKer.of <| ULift W) + (SX := SFinKer.of <| ULift X) + (SY := SFinKer.of <| ULift Y) + (SZ := SFinKer.of <| ULift Z) + ulift.{_, max u_1 u_2 u_3 u_4} + ulift.{_, max u_1 u_2 u_3 u_4} + ulift.{_, max u_1 u_2 u_3 u_4} + ulift.{_, max u_1 u_2 u_3 u_4} + κ η + +@[inherit_doc Kernel.monoComp] +scoped[ProbabilityTheory] infixr:80 " ⊗≫ₖ " => Kernel.monoComp + +end ProbabilityTheory.Kernel diff --git a/LeanMachineLearning/Tactic/KernelHom/Tactic/Delaborators.lean b/LeanMachineLearning/Tactic/KernelHom/Tactic/Delaborators.lean new file mode 100644 index 00000000..e10178f1 --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom/Tactic/Delaborators.lean @@ -0,0 +1,50 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import LeanMachineLearning.Tactic.KernelHom.Kernel.Hom + +/-! +# Delaborators for simplified kernel presentations + +This file implements delaborators that provide simplified pretty-printing kernel-related +categorical operations that are generated using the `kernel_hom` tactic. + +To use these delaborators, simply open the `KernelHom` namespace. +-/ + +public meta section + +open Lean Meta Elab Command PrettyPrinter Delaborator +open Lean.PrettyPrinter.Delaborator.SubExpr + +namespace KernelHom + +/-- Removes the `ULift` wrapper for readability. -/ +@[scoped app_delab ULift] +meta def delabULift : Delab := do + let x ← withNaryArg 0 delab + `($x) + +/-- Only display the carrier space of `SFinKer.of` for readability. -/ +@[scoped app_delab SFinKer.of] +meta def delabSFinKerOf : Delab := do + let x ← withNaryArg 0 delab + `($x) + +/-- Only display the underlying kernel of `Kernel.hom` for readability. -/ +@[scoped app_delab ProbabilityTheory.Kernel.hom] +meta def delabKernelHom : Delab := do + let x ← withNaryArg 8 delab + `($x) + +/-- Only display the underlying kernel of `Kernel.lift` for readability. -/ +@[scoped app_delab ProbabilityTheory.Kernel.lift] +meta def delabKernelLift : Delab := do + let x ← withNaryArg 10 delab + `($x) + +end KernelHom diff --git a/LeanMachineLearning/Tactic/KernelHom/Tactic/HomKernel.lean b/LeanMachineLearning/Tactic/KernelHom/Tactic/HomKernel.lean new file mode 100644 index 00000000..38fcf9f1 --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom/Tactic/HomKernel.lean @@ -0,0 +1,275 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import LeanMachineLearning.Tactic.KernelHom.Tactic.KernelHom +public import LeanMachineLearning.Tactic.EqLift.Tactic.Kernel.KernelUnlift + +/-! +# `hom_kernel` tactic + +This file implements the `hom_kernel` tactic, the inverse of `kernel_hom`. +It transforms equalities written in the monoidal category back into +equivalent equalities of kernels. + +## Main declarations + +* `transformHomToKernel`: recursive translation from categorical morphism expressions to + kernel expressions. +* `applyHomKernel`: core implementation on goals and hypotheses. +* `hom_kernel`: user-facing tactic (with location support). +-/ + +public meta section + +open Lean Elab Tactic Meta CategoryTheory Parser.Tactic ProbabilityTheory MonoidalCategory +open ProbabilityTheory.Kernel + +/-- Get the original type and its universe from a `SFinKer.of` expression. -/ +partial def getTypeFromSFinKer (e : Expr) : MetaM Expr := do + match e.getAppFn with + | Expr.const ``tensorUnit [eLvl, _] => + return mkConst ``PUnit [eLvl.succ] + | Expr.const ``SFinKer.of _ => + let args := e.getAppArgs + return args[0]! + | Expr.const ``MonoidalCategory.tensorObj _ => + let args := e.getAppArgs + let SY := args[args.size - 1]! + let SX := args[args.size - 2]! + let Y ← getTypeFromSFinKer SY + let X ← getTypeFromSFinKer SX + mkAppOptM ``Prod #[X, Y] + | _ => throwError "Expected a SFinKer.of expression, got: {e}." + +/-- Deconstruct a left or right whisker. -/ +def deconstructWhiskersHomArgs (e : Expr) (eLvl : Level) (left : Bool) : + MetaM (Expr × Expr × Expr × Expr × Expr × Expr × Expr × Expr) := do + let args := e.getAppArgs + let SZ := if left then args[args.size - 4]! else args[args.size - 1]! + let SY := if left then args[args.size - 2]! else args[args.size - 3]! + let SX := if left then args[args.size - 3]! else args[args.size - 4]! + let κ := if left then args[args.size - 1]! else args[args.size - 2]! + let Z ← getTypeFromSFinKer SZ + let Y ← getTypeFromSFinKer SY + let X ← getTypeFromSFinKer SX + let mXUnit ← synthInstance (mkApp (mkConst ``MeasurableSpace [eLvl]) Z) + let kernel_id ← mkAppOptM ``Kernel.id #[Z, mXUnit] + return (κ, kernel_id, SX, SY, SZ, X, Y, Z) + +/-- Deconstruct a braiding morphism. -/ +def deconstructBraiding (e : Expr) : MetaM (Expr × Expr) := do + let args := e.getAppArgs + let SY := args[args.size - 1]! + let SX := args[args.size - 2]! + let Y ← getTypeFromSFinKer SY + let X ← getTypeFromSFinKer SX + let swap_hom_proof ← mkAppM ``braiding_hom #[SX, SY, ← idME X, ← idME Y] + return (← mkAppOptM ``Kernel.swap #[X, Y, none, none], swap_hom_proof) + +/-- Given an equality between a categorical morphism (left) and a "morphized" kernel (right), get +the kernel on the right side of the equality. -/ +def getKernelRHSEqProofType (e : Expr) : MetaM Expr := do + let some (_, _, hom_expr) := (← inferType e).eq? | throwError "Expected an equality, got: {e}." + match hom_expr.getAppFn with + | Expr.const ``Kernel.hom _ => + let args := hom_expr.getAppArgs + return args[args.size - 2]! + | _ => throwError "Expected a hom expression, got: {hom_expr}." + +/-- Deconstruct a left or right unitor [inverse] morphism. -/ +def deconstructUnitors (e : Expr) (eLvl : Level) (left hom : Bool) : + MetaM (Expr × Expr) := do + let args := e.getAppArgs + let SX := args[args.size - 1]! + let X ← getTypeFromSFinKer SX + let ex ← idME X + let (X₀, x₀Lvl) ← getOriginalType X + let ex₀ ← constructMeasurableEquiv X₀ x₀Lvl eLvl + let const_args := [eLvl, x₀Lvl, eLvl, Level.zero] + let const_name := + if left then + if hom then ``leftUnitor_hom + else ``leftUnitor_inv + else + if hom then ``rightUnitor_hom + else ``rightUnitor_inv + let const := mkConst const_name const_args + let unitor_proof_eq ← mkAppM' const #[SX, ex, ex₀] + return (← getKernelRHSEqProofType unitor_proof_eq, unitor_proof_eq) + +/-- Deconstruct an associator [inverse] morphism. -/ +def deconstructAssociator (e : Expr) (eLvl : Level) (hom : Bool) : MetaM (Expr × Expr) := do + let args := e.getAppArgs + let SZ := args[args.size - 1]! + let SY := args[args.size - 2]! + let SX := args[args.size - 3]! + let Z ← getTypeFromSFinKer SZ + let Y ← getTypeFromSFinKer SY + let X ← getTypeFromSFinKer SX + let (Z₀, z₀Lvl) ← getOriginalType Z + let (Y₀, y₀Lvl) ← getOriginalType Y + let (X₀, x₀Lvl) ← getOriginalType X + let ez₀ ← constructMeasurableEquiv Z₀ z₀Lvl eLvl + let ey₀ ← constructMeasurableEquiv Y₀ y₀Lvl eLvl + let ex₀ ← constructMeasurableEquiv X₀ x₀Lvl eLvl + let associator_const := mkConst + (if hom then ``Kernel.associator_hom else ``Kernel.associator_inv) + [eLvl, eLvl, eLvl, x₀Lvl, y₀Lvl, z₀Lvl, eLvl] + let associator_proof_eq ← mkAppM' associator_const + #[SX, SY, SZ, ← idME X, ← idME Y, ← idME Z, ex₀, ey₀, ez₀] + return (← getKernelRHSEqProofType associator_proof_eq, associator_proof_eq) + +/-- Recursive transformation from morphism expression in `SFinKer` to kernel expression. -/ +partial def transformHomToKernel (e : Expr) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + match e.getAppFn with + | Expr.const ``tensorHom _ => + let args := e.getAppArgs + let κ := args[args.size - 2]! + let η := args[args.size - 1]! + let ST := args[args.size - 3]! + let SZ := args[args.size - 4]! + let SY := args[args.size - 5]! + let SX := args[args.size - 6]! + let (κ', proofs_κ) ← transformHomToKernel κ proofs + let (η', proofs_η) ← transformHomToKernel η proofs_κ + let (X, Y, _, _) ← getTypesFromKernel κ' + let (Z, T, _, _) ← getTypesFromKernel η' + let parallelComp_hom_proof ← mkAppMInst ``parallelComp_hom + #[SX, SY, SZ, ST, ← idME X, ← idME Y, ← idME Z, ← idME T, κ', η'] 2 + return (← mkAppM ``Kernel.parallelComp #[κ', η'], parallelComp_hom_proof :: proofs_η) + | Expr.const ``CategoryStruct.comp _ => + let args := e.getAppArgs + let κ := args[args.size - 2]! + let η := args[args.size - 1]! + let SY := args[args.size - 3]! + let SX := args[args.size - 4]! + let SZ := args[args.size - 5]! + let (κ', proofs_κ) ← transformHomToKernel κ proofs + let (η', proofs_η) ← transformHomToKernel η proofs_κ + let (X, Y, _, _) ← getTypesFromKernel η' + let (Z, _, _, _) ← getTypesFromKernel κ' + let comp_hom_proof ← mkAppMInst ``comp_hom + #[SX, SY, SZ, ← idME X, ← idME Y, ← idME Z, η', κ'] 2 + return (← mkAppM ``Kernel.comp #[η', κ'], comp_hom_proof :: proofs_η) + | Expr.const ``CategoryStruct.id [xLvl, _] => + let args := e.getAppArgs + let SX := args[args.size - 1]! + let X ← getTypeFromSFinKer SX + let mX' ← synthInstance (mkApp (mkConst ``MeasurableSpace [xLvl]) X) + let id ← mkAppOptM ``Kernel.id #[X, mX'] + let id_hom_proof ← mkAppM ``id_hom #[SX, ← idME X] + return (id, id_hom_proof :: proofs) + | Expr.const ``ComonObj.counit [xLvl, _] => + let args := e.getAppArgs + let SX := args[args.size - 2]! + let X ← getTypeFromSFinKer SX + let discard_kernel_const := mkConst ``Kernel.discard [xLvl, xLvl] + let discard_const := mkConst ``counit [xLvl, xLvl, xLvl] + let discard_hom_proof ← mkAppM' discard_const #[SX, ← idME X] + return (← mkAppOptM' discard_kernel_const #[X, none], discard_hom_proof :: proofs) + | Expr.const ``ComonObj.comul [xLvl, _] => + let args := e.getAppArgs + let SX := args[args.size - 2]! + let X ← getTypeFromSFinKer SX + let copy_kernel_const := mkConst ``Kernel.copy [xLvl] + let copy_hom_proof ← mkAppM ``comul #[SX, ← idME X] + return (← mkAppOptM' copy_kernel_const #[X, none], copy_hom_proof :: proofs) + | Expr.const ``Kernel.hom _ => + let args := e.getAppArgs + let κ := args[args.size - 2]! + return (κ, proofs) + | Expr.const ``MonoidalCategory.whiskerLeft [eLvl, _] => + let (κ, kernel_id, SX, SY, SZ, X, Y, Z) ← deconstructWhiskersHomArgs e eLvl true + let (κ', proofs_κ) ← transformHomToKernel κ proofs + let whisker_left_hom_proof ← mkAppMInst ``Kernel.whiskerLeft + #[SX, SY, SZ, ← idME X, ← idME Y, ← idME Z, κ'] 1 + return (← mkAppM ``Kernel.parallelComp #[kernel_id, κ'], whisker_left_hom_proof :: proofs_κ) + | Expr.const ``MonoidalCategory.whiskerRight [eLvl, _] => + let (κ, kernel_id, SX, SY, SZ, X, Y, Z) ← deconstructWhiskersHomArgs e eLvl false + let (κ', proofs_κ) ← transformHomToKernel κ proofs + let whisker_right_hom_proof ← mkAppMInst ``Kernel.whiskerRight + #[SX, SY, SZ, ← idME X, ← idME Y, ← idME Z, κ'] 1 + return (← mkAppM ``Kernel.parallelComp #[κ', kernel_id], whisker_right_hom_proof :: proofs_κ) + | Expr.const ``Iso.hom _ => + let args := e.getAppArgs + let iso := args[args.size - 1]! + match iso.getAppFn with + | Expr.const ``BraidedCategory.braiding _ => + let (braiding_expr, swap_hom_proof) ← deconstructBraiding iso + return (braiding_expr, swap_hom_proof :: proofs) + | Expr.const ``leftUnitor [eLvl, _] => + let (left_unitor_expr, left_unitor_hom_proof) ← deconstructUnitors iso eLvl true true + return (left_unitor_expr, left_unitor_hom_proof :: proofs) + | Expr.const ``rightUnitor [eLvl, _] => + let (right_unitor_expr, right_unitor_hom_proof) ← deconstructUnitors iso eLvl false true + return (right_unitor_expr, right_unitor_hom_proof :: proofs) + | Expr.const ``MonoidalCategory.associator [eLvl, _] => + let (associator_expr, associator_hom_proof) ← deconstructAssociator iso eLvl true + return (associator_expr, associator_hom_proof :: proofs) + | _ => throwError "Unexpected isomorphism {iso}." + | Expr.const ``Iso.inv _ => + let args := e.getAppArgs + let iso := args[args.size - 1]! + match iso.getAppFn with + | Expr.const ``BraidedCategory.braiding _ => + let (braiding_expr, swap_hom_proof) ← deconstructBraiding iso + return (braiding_expr, swap_hom_proof :: proofs) + | Expr.const ``leftUnitor [eLvl, _] => + let (left_unitor_expr, left_unitor_inv_hom_proof) ← deconstructUnitors iso eLvl true false + return (left_unitor_expr, left_unitor_inv_hom_proof :: proofs) + | Expr.const ``rightUnitor [eLvl, _] => + let (right_unitor_expr, right_unitor_inv_hom_proof) ← deconstructUnitors iso eLvl false false + return (right_unitor_expr, right_unitor_inv_hom_proof :: proofs) + | Expr.const ``MonoidalCategory.associator [eLvl, _] => + let (associator_expr, associator_inv_hom_proof) ← deconstructAssociator iso eLvl false + return (associator_expr, associator_inv_hom_proof :: proofs) + | _ => throwError "Unexpected isomorphism {iso}." + | _ => throwError "Expected a hom expression, got: {e}." + +/-- Get the universe level from the left side of an equality expression. -/ +def getUniverseFromEq (eq : Expr) : MetaM Level := do + let eq ← instantiateMVars eq + let eq ← zetaReduce eq + let eq ← whnf eq + let eq := eq.consumeMData + let some (_, lhs, _) := eq.eq? | throwError "Expected an equality, got: {eq}." + let l ← getLevel (← inferType lhs) + match l with + | Level.succ l' => return l' + | _ => throwError "Expected a universe level ≥ 1, got: {l}" + +/-- Transform a `SFinKer` equality into an equivalent equality of kernels, along with a proof of +equivalence. -/ +def KernelEquality (eq : Expr) : MetaM (Expr × Expr) := do + let eq ← whnfR <| ← instantiateMVars eq + let some (_, lhs_hom, rhs_hom) := eq.eq? | throwError "Expected an equality, got: {eq}." + let (lhs, proofs) ← transformHomToKernel lhs_hom [] + let (rhs, proofs) ← transformHomToKernel rhs_hom proofs + let kernel_expr ← mkEq lhs rhs + let (unlifted_expr, unlifted_proof) ← unliftEquality kernel_expr + let kernel_eq_proof_type ← mkEq kernel_expr eq + let kernel_eq_proof ← mkAppM ``Eq.symm #[← mkKernelHomEqProof kernel_eq_proof_type lhs rhs proofs] + return (unlifted_expr, ← mkEqTrans kernel_eq_proof unlifted_proof) + +/-- The `hom_kernel` tactic is the inverse of `kernel_hom`: it transforms an +equality written in the monoidal category back to an equivalent equality of +s-finite kernels. + +The tactic supports location specifiers like `rw` or `simp`: +- `hom_kernel` — applies to the goal +- `hom_kernel at h` — applies to hypothesis `h` +- `hom_kernel at h₁ h₂` — applies to multiple hypotheses +- `hom_kernel at h ⊢` — applies to hypothesis `h` and the goal +- `hom_kernel at *` — applies to all hypotheses and the goal + +It is useful to switch back to kernel equations once categorical rewrites are done. -/ +syntax (name := homKernel) "hom_kernel" (ppSpace location)? : tactic + +elab_rules : tactic + | `(tactic| hom_kernel $[$loc]?) => + expandOptLocation (Lean.mkOptionalNode loc) |> applyLocTactic <| KernelEquality diff --git a/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelCat.lean b/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelCat.lean new file mode 100644 index 00000000..d6195e07 --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelCat.lean @@ -0,0 +1,55 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import LeanMachineLearning.Tactic.KernelHom.Tactic.KernelHom +public import Mathlib.Tactic.CategoryTheory.Coherence + +/-! +# Kernel category tactics + +This file implements the `kernel_coherence` and `kernel_monoidal` tactics, which apply the +`kernel_hom` transformation and then use categorical `coherence` or `monoidal` tactics to solve the +resulting goal. + +## Main declarations + +* `kernel_coherence`: tactic combining kernel_hom and categorical coherence. +* `kernel_monoidal`: tactic combining kernel_hom and categorical monoidal coherence. +-/ + +public meta section + +open Lean Elab Tactic CategoryTheory +open Lean Elab Tactic Meta CategoryTheory Parser.Tactic ProbabilityTheory MonoidalCategory + + +/-- The `kernel_monoidal` tactic applies the `kernel_hom` transformation to the goal and then +invokes the `monoidal` tactic to solve or simplify the resulting goal. -/ +syntax (name := kernelMonoidal) "kernel_monoidal" : tactic + +elab_rules : tactic + | `(tactic| kernel_monoidal) => do + evalTactic (← `(tactic| kernel_hom)) + evalTactic (← `(tactic| monoidal)) + +/-- The `kernel_coherence` tactic applies the `kernel_hom` transformation to the goal and then +invokes the `coherence` tactic to solve the resulting goal. -/ +syntax (name := kernelCoherence) "kernel_coherence" : tactic + +elab_rules : tactic + | `(tactic| kernel_coherence) => do + evalTactic (← `(tactic| kernel_hom)) + evalTactic (← `(tactic| coherence)) + +/-- The `kernel_disch` tactic applies the `kernel_hom` transformation to the goal and then +invokes the `cat_disch` tactic to solve the resulting goal. -/ +syntax (name := kernelDisch) "kernel_disch" : tactic + +elab_rules : tactic + | `(tactic| kernel_disch) => do + evalTactic (← `(tactic| kernel_hom)) + evalTactic (← `(tactic| cat_disch)) diff --git a/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelDiagram.lean b/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelDiagram.lean new file mode 100644 index 00000000..6ddf5030 --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelDiagram.lean @@ -0,0 +1,206 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import Lean.Elab.Tactic.Location +public import LeanMachineLearning.Tactic.KernelHom.Tactic.HomKernel +public import Mathlib.Tactic.Widget.StringDiagram + +/-! +# Kernel Diagram Widget + +This file provides meta infrastructure for displaying string diagrams for s-finite kernels in the +infoview. To enable the kernel diagram widget, you need to import this file and inserting +`with_panel_widgets [KernelDiagram]` at the beginning of the proof. Alternatively, you can also +write +```lean +show_panel_widgets [local KernelDiagram] +``` +to enable the string diagram widget in the current section. + +We also have the `#kernel_diagram` command. For example, +```lean +#string_diagram ProbabilityTheory.Kernel.deterministic_comp_copy +``` + +This is an adaptation of the string diagram widget where kernels are transformed into morphisms of +the `SFinKer` monoidal category using the `kernel_hom` tactic. +-/ + +public meta section + +open Lean Meta Elab Command ProofWidgets Mathlib.Tactic.Widget + +open Mathlib.Tactic BicategoryLike Penrose Server + +open MeasureTheory ProbabilityTheory CategoryTheory + +open CategoryTheory + +open scoped MonoidalCategory ComonObj + +namespace Mathlib.Tactic.Widget.StringDiagram + +/-- The kernelized penrose variable associated with a node. -/ +def Node.toPenroseVar_kernel (n : Node) : MetaM PenroseVar := do + let expr ← + try + match n.e.getAppFn with + | Expr.const ``SFinKer.of _ => do + let res ← getTypeFromSFinKer n.e + pure res + | _ => do + let (expr, _) ← transformHomToKernel n.e [] + pure expr + catch _ => + pure n.e + return ⟨"E", [n.vPos, n.hPosSrc, n.hPosTar], expr⟩ + +open scoped Jsx in +/-- Construct a kernelized string diagram from a Penrose `sub`stance program and +expressions `embeds` to display as labels in the diagram. -/ +def mkKernelDiagram (nodes : List (List Node)) (strands : List (List Strand)) : + DiagramBuilderM PUnit := do + /- Add 2-morphisms. -/ + for x in nodes.flatten do + match x with + | .atom _ => do addPenroseVar "Atom" (← x.toPenroseVar_kernel) + | .id _ => do StringDiagram.addPenroseVar "Id" (← x.toPenroseVar_kernel) + /- Add constraints. -/ + for l in nodes do + for (x₁, x₂) in l.consecutivePairs do + DiagramBuilderM.addInstruction + s!"Left({← x₁.toPenroseVar_kernel}, {← x₂.toPenroseVar_kernel})" + /- Add constraints. -/ + for (l₁, l₂) in nodes.consecutivePairs do + if let some x₁ := l₁.head? then + if let some x₂ := l₂.head? then + DiagramBuilderM.addInstruction + s!"Above({← x₁.toPenroseVar_kernel}, {← x₂.toPenroseVar_kernel})" + /- Add 1-morphisms as strings. -/ + for l in strands do + for s in l do + StringDiagram.addConstructor "Mor1" s.toPenroseVar + "MakeString" [← s.startPoint.toPenroseVar_kernel, ← s.endPoint.toPenroseVar_kernel] + +end Mathlib.Tactic.Widget.StringDiagram + +namespace KernelDiagram + +open scoped Jsx in +/-- Given a kernel expression, return a string diagram. Otherwise `none`. -/ +def KernelM? (e : Expr) : MetaM (Option Html) := do + let e ← instantiateMVars e + try + let (e, _) ← transformKernelToHom e [] + let k ← StringDiagram.mkKind e + let x : Option (List (List StringDiagram.Node) × List (List StringDiagram.Strand)) + ← (match k with + | .monoidal => do + let some ctx ← BicategoryLike.mkContext? (ρ := Monoidal.Context) e | return none + CoherenceM.run (ctx := ctx) do + let e' := (← BicategoryLike.eval k.name (← MkMor₂.ofExpr e)).expr + return some (← e'.nodes, ← e'.strands) + | .bicategory => do + let some ctx ← BicategoryLike.mkContext? (ρ := Bicategory.Context) e | return none + CoherenceM.run (ctx := ctx) do + let e' := (← BicategoryLike.eval k.name (← MkMor₂.ofExpr e)).expr + return some (← e'.nodes, ← e'.strands) + | .none => return none) + match x with + | none => return none + | some (nodes, strands) => do + DiagramBuilderM.run do + StringDiagram.mkKernelDiagram nodes strands + trace[string_diagram] "Penrose substance: \n{(← get).sub}" + match ← DiagramBuilderM.buildDiagram StringDiagram.dsl StringDiagram.sty with + | some html => return html + | none => return No non-structural morphisms found. + catch _ => return none + +open scoped Jsx in +/-- Help function for displaying two string diagrams in an equality. -/ +def mkEqHtml (lhs rhs : Html) : Html := +
+
+
+ Kernel diagram for LHS {lhs} +
+
+
+
+ Kernel diagram for RHS {rhs} +
+
+
+ +/-- Given an equality between kernels, return a string diagram of the LHS and RHS. +Otherwise `none`. -/ +def kernelEqM? (e : Expr) : MetaM (Option Html) := do + try + let e ← unfoldKernelOp <| ← instantiateMVars e + let (lifted_e, _) ← liftEquality e + let some (_, lhs, rhs) := lifted_e.eq? | return none + let some lhs ← KernelM? lhs | return none + let some rhs ← KernelM? rhs | return none + return some <| mkEqHtml lhs rhs + catch _ => return none + +/-- Reduce a forall expression and try to display a kernel diagram +if it is an equality of kernels. -/ +def kernelEqMReduce? (e : Expr) : MetaM (Option Html) := do + forallTelescopeReducing (← whnfR <| ← inferType e) fun _ expr => do + kernelEqM? expr + +open scoped Jsx in +/-- The RPC method for displaying kernel diagrams. -/ +@[server_rpc_method] +def rpc (props : PanelWidgetProps) : RequestM (RequestTask Html) := + RequestM.asTask do + let html : Option Html ← (do + if props.goals.isEmpty then + return none + let some g := props.goals[0]? | unreachable! + g.ctx.val.runMetaM {} do + g.mvarId.withContext do + let type ← g.mvarId.getType + kernelEqM? type) + match html with + | none => return No Kernel Diagram. + | some inner => return inner + +end KernelDiagram + +/-- Display the kernel diagrams if the goal is an equality of s-finite kernels. -/ +@[widget_module] +def KernelDiagram : Component PanelWidgetProps := + mk_rpc_widget% KernelDiagram.rpc + +/-- +Display the kernel diagram for a given term. + +Example usage: +``` +/- Kernel diagram for an equality theorem. -/ +#kernel_diagram ProbabilityTheory.Kernel.deterministic_comp_copy +``` +-/ +syntax (name := kernelDiagram) "#kernel_diagram " term : command + +@[command_elab kernelDiagram, inherit_doc kernelDiagram] +def elabKernelDiagramCmd : CommandElab := fun + | stx@`(#kernel_diagram $t:term) => do + let html ← runTermElabM fun _ => do + let e ← try mkConstWithFreshMVarLevels (← realizeGlobalConstNoOverloadWithInfo t) + catch _ => Term.levelMVarToParam (← instantiateMVars (← Term.elabTerm t none)) + match ← KernelDiagram.kernelEqMReduce? e with + | some html => return html + | none => throwError "could not find an equality of kernels: {e}." + liftCoreM <| Widget.savePanelWidgetInfo + (hash HtmlDisplay.javascript) + (return json% { html: $(← Server.RpcEncodable.rpcEncode html) }) + stx + | stx => throwError "Unexpected syntax {stx}." diff --git a/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelHom.lean b/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelHom.lean new file mode 100644 index 00000000..d718e679 --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelHom.lean @@ -0,0 +1,415 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import LeanMachineLearning.Tactic.KernelHom.Kernel.MonoidalComp +public import LeanMachineLearning.Tactic.KernelHom.Tactic.Utils +public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.MeasurableEquiv +public import Lean.Elab.Tactic.Location +public import LeanMachineLearning.Tactic.EqLift.Tactic.Kernel.KernelLift + +/-! +# `kernel_hom` tactic + +This file implements the `kernel_hom` tactic, which transforms equalities of +kernels into equivalent equalities in the monoidal category. + +## Main declarations + +* `transformKernelToHom`: recursive translation from kernel expressions to + categorical morphism expressions. +* `mkKernelHomEqProof`: construction of the equivalence proof used by the + tactic. +* `applyKernelHom`: core implementation of `kernel_hom` on goals and hypotheses. +* `kernel_hom`: user-facing tactic (with location support). +-/ + +public meta section + +open Lean Elab Tactic Meta CategoryTheory Parser.Tactic ProbabilityTheory MonoidalCategory +open ProbabilityTheory.Kernel + +/-- Recursively decompose a product type into `SFinKer` objects with monoidal tensor structure. -/ +partial def decomposeProductToSFinker (X : Expr) (xLvl : Level) : MetaM Expr := do + match X.getAppFn with + | Expr.const ``Prod _ => + let args := X.getAppArgs + let t1 ← decomposeProductToSFinker args[0]! xLvl + let t2 ← decomposeProductToSFinker args[1]! xLvl + mkAppM ``tensorObj #[t1, t2] + | _ => + mkAppOptM ``SFinKer.of #[X, none] + +/-- Compute the `SFinKer` object corresponding to a measurable space. -/ +def computeSFinkerOf (X : Expr) (xLvl : Level) : MetaM Expr := do + match X with + | Expr.const ``PUnit _ | Expr.const ``Unit _ => + let tensorunit := mkConst ``tensorUnit [xLvl, xLvl.succ] + let sfinker := mkConst ``SFinKer [xLvl] + mkAppOptM' tensorunit #[sfinker, none, none] + | _ => + decomposeProductToSFinker X xLvl + +/-- Compute a measurable equivalence between a type and itself by recursively decomposing +products. -/ +partial def idME (X : Expr) : MetaM Expr := do + match X.getAppFn with + | Expr.const ``Prod _ => + let args := X.getAppArgs + let id1 ← idME args[0]! + let id2 ← idME args[1]! + mkAppM ``MeasurableEquiv.prod #[id1, id2] + | Expr.const ``PUnit [xLvl] | Expr.const ``Unit [xLvl] => + let xLvl ← match xLvl with + | Level.succ l => pure l + | _ => throwError "Expected a successor level for PUnit/Unit, got: {xLvl}." + let punitME := mkConst ``MeasurableEquiv.punit [xLvl, xLvl] + mkAppM' punitME #[] + | _ => + mkAppOptM ``MeasurableEquiv.id #[X, none] + +/-- Check if a kernel expression corresponds to a left or right whisker. -/ +def checkWhiskers (κ : Expr) (offset : Nat) : MetaM Bool := do + let κ := κ.consumeMData + let args := κ.getAppArgs + let idKernel := args[args.size - offset]! + if !idKernel.isAppOf ``Kernel.id then + return false + else return true + +/-- Check if a kernel expression corresponds to a left whisker. -/ +def checkWhiskerLeft (κ : Expr) : MetaM Bool := checkWhiskers κ 2 + +/-- Check if a kernel expression corresponds to a right whisker. -/ +def checkWhiskerRight (κ : Expr) : MetaM Bool := checkWhiskers κ 1 + +/-- Construct the relevant data for converting a kernel expression to its whisker morphism +representation. -/ +def constructWhiskersArgs (e X Y : Expr) (left : Bool) : + MetaM (Expr × Expr × Expr × Expr × Expr × Expr × Expr) := do + let (Z, zLvl, X, xLvl) ← match X.getAppFn with + | Expr.const ``Prod univs => + let args := X.getAppArgs + pure (args[left.toNat]!, univs[left.toNat]!, args[1 - left.toNat]!, univs[1 - left.toNat]!) + | _ => + if left then throwError "Expected left whisker with source Z × X, got: {X}." + else throwError "Expected right whisker with source X × Z, got: {X}." + let (Y, yLvl) ← match Y.getAppFn with + | Expr.const ``Prod univs => + let args := Y.getAppArgs + pure (args[1 - left.toNat]!, univs[1 - left.toNat]!) + | _ => + if left then throwError "Expected left whisker with target Z × Y, got: {Y}." + else throwError "Expected right whisker with target Y × Z, got: {Y}." + let κ ← match e.getAppFn with + | Expr.const ``Kernel.parallelComp _ => + let args := e.getAppArgs + pure args[args.size - (left.toNat + 1)]! + | _ => + if left then throwError "Expected left whisker with parallelComp, got: {e}." + else throwError "Expected right whisker with parallelComp, got: {e}." + let SZ ← computeSFinkerOf Z zLvl + let SX ← computeSFinkerOf X xLvl + let SY ← computeSFinkerOf Y yLvl + return (SZ, Z, SX, X, SY, Y, κ) + +/-- Check if a kernel expression corresponds to a left or right unitor. -/ +def checkUnitors (κ : Expr) (offset : Nat) (prod : Name) : MetaM Bool := do + let κ := κ.consumeMData + if !κ.isAppOf ``Kernel.map then + return false + let args := κ.getAppArgs + let fn := args[args.size - 1]! + let idKernel := args[args.size - 2]! + if !fn.isAppOf prod then + return false + if !idKernel.isAppOf ``Kernel.id then + return false + let (src, _, _) ← getTypesFromKernel κ + match src.getAppFn with + | Expr.const ``Prod _ => + let args := src.getAppArgs + if args.size < 2 then + return false + let punit? := args[offset]! + match punit?.getAppFn with + | Expr.const ``PUnit _ | Expr.const ``Unit _ => return true + | _ => return false + | _ => return false + +/-- Check if a kernel expression corresponds to a left unitor. -/ +def checkLeftUnitor (κ : Expr) : MetaM Bool := checkUnitors κ 0 ``Prod.snd + +/-- Check if a kernel expression corresponds to a right unitor. -/ +def checkRightUnitor (κ : Expr) : MetaM Bool := checkUnitors κ 1 ``Prod.fst + +/-- Construct the left or right unitor morphism. -/ +def constructUnitors (X ex₀ : Expr) (xLvl y₀Lvl punitLvl : Level) (offset : Nat) : + MetaM (Expr × Expr) := do + let left ← if offset == 0 then pure true + else if offset == 1 then pure false + else throwError "Invalid offset for unitors." + let SX ← computeSFinkerOf X xLvl + let unitor ← if left then mkAppM ``leftUnitor #[SX] + else mkAppM ``rightUnitor #[SX] + let unitor_hom_const := + if left then mkConst ``leftUnitor_hom [xLvl, y₀Lvl, xLvl, punitLvl] + else mkConst ``rightUnitor_hom [xLvl, y₀Lvl, xLvl, punitLvl] + let unitor_hom_proof ← + if left then mkAppM' unitor_hom_const #[SX, ← idME X, ex₀] + else mkAppM' unitor_hom_const #[SX, ← idME X, ex₀] + return (← mkAppM ``Iso.hom #[unitor], unitor_hom_proof) + +/-- Check if a kernel expression corresponds to an associator morphism or its inverse. -/ +def checkAssociator (κ : Expr) (hom : Bool) : MetaM Bool := do + let κ := κ.consumeMData + if !κ.isAppOf ``Kernel.deterministic then + return false + let args := κ.getAppArgs + let fn := args[args.size - 2]! + if !fn.isAppOf ``DFunLike.coe then + return false + let fn := fn.getAppArgs[fn.getAppApps.size - 1]! + if hom then + if !fn.isAppOf ``MeasurableEquiv.prodAssoc then + return false + else + if !fn.isAppOf ``MeasurableEquiv.symm then + return false + let innerFn := fn.getAppArgs[fn.getAppArgs.size - 1]! + if !innerFn.isAppOf ``MeasurableEquiv.prodAssoc then + return false + return true + +/-- Check if a kernel expression corresponds to an associator morphism. -/ +def checkAssociatorHom (κ : Expr) : MetaM Bool := checkAssociator κ true + +/-- Check if a kernel expression corresponds to an inverse associator morphism. -/ +def checkAssociatorInv (κ : Expr) : MetaM Bool := checkAssociator κ false + +/-- Get the types and universe levels from a expression of the form `X × Y × Z`. -/ +def getTypesFromThreeProds (prod : Expr) : + MetaM (Expr × Expr × Expr × Level × Level × Level) := do + match prod.getAppFn with + | Expr.const ``Prod univs => + let X := prod.getAppArgs[0]! + match prod.getAppArgs[1]!.getAppFn with + | Expr.const ``Prod univs_right => + let Y := prod.getAppArgs[1]!.getAppArgs[0]! + let Z := prod.getAppArgs[1]!.getAppArgs[1]! + return (X, Y, Z, univs[0]!, univs_right[0]!, univs_right[1]!) + | _ => throwError "Expected a product of two types, got: {prod.getAppArgs[1]!}." + | _ => throwError "Expected a product of three types, got: {prod}." + +/-- Get the measurable equivalences from a product of three measurable equivalences. -/ +def getMEFromThreeProds (me_prod : Expr) : + MetaM (Expr × Expr × Expr) := do + match me_prod.getAppFn with + | Expr.const ``MeasurableEquiv.prod _ => + let args := me_prod.getAppArgs + let ex := args[args.size - 2]! + let right := args[args.size - 1]! + match right.getAppFn with + | Expr.const ``MeasurableEquiv.prod _ => + let rightArgs := right.getAppArgs + let ey := rightArgs[rightArgs.size - 2]! + let ez := rightArgs[rightArgs.size - 1]! + return (ex, ey, ez) + | _ => throwError "Expected a product of two measurable equivalences, got: {right}." + | _ => throwError "Expected a product of three measurable equivalences, got: {me_prod}." + +/-- Construct the associator morphism or its inverse. -/ +def constructAssociator (left right ex₀ ey₀ ez₀ : Expr) (hom : Bool) : + MetaM (Expr × Expr) := do + let (X, Y, Z, xLvl, yLvl, zLvl) ← if hom then getTypesFromThreeProds right + else getTypesFromThreeProds left + let SX ← computeSFinkerOf X xLvl + let SY ← computeSFinkerOf Y yLvl + let SZ ← computeSFinkerOf Z zLvl + let associator ← mkAppM ``MonoidalCategory.associator #[SX, SY, SZ] + let associator_hom_proof ← + if hom then mkAppM ``associator_hom #[SX, SY, SZ, ← idME X, ← idME Y, ← idME Z, ex₀, ey₀, ez₀] + else mkAppM ``associator_inv #[SX, SY, SZ, ← idME X, ← idME Y, ← idME Z, ex₀, ey₀, ez₀] + return (← mkAppM (if hom then ``Iso.hom else ``Iso.inv) #[associator], associator_hom_proof) + +/-- Construct the associator morphism. -/ +def constructAssociatorHom (left right ex₀ ey₀ ez₀ : Expr) := + constructAssociator left right ex₀ ey₀ ez₀ true + +/-- Construct the inverse associator morphism. -/ +def constructAssociatorInv (left right ex₀ ey₀ ez₀ : Expr) := + constructAssociator left right ex₀ ey₀ ez₀ false + +/-- Recursive transformation from kernel expressions to morphism expressions in the `SFinKer` +category. -/ +partial def transformKernelToHom (e : Expr) (proofs : List Expr) : + MetaM (Expr × List Expr) := do + match e.getAppFn with + | Expr.const ``Kernel.comp _ => + let args := e.getAppArgs + let η := args[args.size - 2]! + let κ := args[args.size - 1]! + let (X, Y, xLvl, yLvl) ← getTypesFromKernel η + let (Z, _, tLvl, _) ← getTypesFromKernel κ + let SX ← computeSFinkerOf X xLvl + let SY ← computeSFinkerOf Y yLvl + let SZ ← computeSFinkerOf Z tLvl + let comp_hom_proof ← mkAppMInst ``comp_hom #[SX, SY, SZ, ← idME X, ← idME Y, ← idME Z, η, κ] 2 + let (κ', proofs_κ) ← transformKernelToHom κ proofs + let (η', proofs_η) ← transformKernelToHom η proofs_κ + return (← mkAppM ``CategoryStruct.comp #[κ', η'], comp_hom_proof :: proofs_η) + | Expr.const ``Kernel.parallelComp _ => + if ← checkWhiskerLeft e then + let (X, Y, _, _) ← getTypesFromKernel e + let (SZ, Z, SX, X, SY, Y, κ) ← constructWhiskersArgs e X Y false + let (κ', proofs_κ) ← transformKernelToHom κ proofs + let whisker_left_hom_proof ← mkAppMInst ``Kernel.whiskerLeft + #[SX, SY, SZ, ← idME X, ← idME Y, ← idME Z, κ] 1 + let whiskerleft ← mkAppM ``MonoidalCategory.whiskerLeft #[SZ, κ'] + return (whiskerleft, whisker_left_hom_proof :: proofs_κ) + else if ← checkWhiskerRight e then + let (X, Y, _, _) ← getTypesFromKernel e + let (SZ, Z, SX, X, SY, Y, κ) ← constructWhiskersArgs e X Y true + let (κ', proofs_κ) ← transformKernelToHom κ proofs + let whiskerright ← mkAppM ``MonoidalCategory.whiskerRight #[κ', SZ] + let whiskerright_hom_proof ← mkAppMInst ``Kernel.whiskerRight + #[SX, SY, SZ, ← idME X, ← idME Y, ← idME Z, κ] 1 + return (whiskerright, whiskerright_hom_proof :: proofs_κ) + else + let args := e.getAppArgs + let κ := args[args.size - 2]! + let η := args[args.size - 1]! + let (X, Y, xLvl, yLvl) ← getTypesFromKernel κ + let (Z, T, zLvl, tLvl) ← getTypesFromKernel η + let SX ← computeSFinkerOf X xLvl + let SY ← computeSFinkerOf Y yLvl + let SZ ← computeSFinkerOf Z zLvl + let ST ← computeSFinkerOf T tLvl + let parallelComp_hom_proof ← mkAppMInst ``parallelComp_hom + #[SX, SY, SZ, ST, ← idME X, ← idME Y, ← idME Z, ← idME T, κ, η] 2 + let (κ', proofs_κ) ← transformKernelToHom κ proofs + let (η', proofs_η) ← transformKernelToHom η proofs_κ + return (← mkAppM ``tensorHom #[κ', η'], parallelComp_hom_proof :: proofs_η) + | Expr.const ``Kernel.id [xLvl] => + let X := e.getAppArgs[0]! + let SX ← computeSFinkerOf X xLvl + let id_hom_proof ← mkAppM ``id_hom #[SX, ← idME X] + return (← mkAppM ``CategoryStruct.id #[SX], id_hom_proof :: proofs) + | Expr.const ``Kernel.discard [xLvl, punitLvl] => + let X := e.getAppArgs[0]! + let SX ← computeSFinkerOf X xLvl + let discard_const := mkConst ``counit [xLvl, xLvl, punitLvl] + let discard_hom_proof ← mkAppM' discard_const #[SX, ← idME X] + return (← mkAppOptM ``ComonObj.counit #[none, none, none, SX, none], + discard_hom_proof :: proofs) + | Expr.const ``Kernel.copy [xLvl] => + let X := e.getAppArgs[0]! + let SX ← computeSFinkerOf X xLvl + let copy_hom_proof ← mkAppM ``comul #[SX, ← idME X] + return (← mkAppOptM ``ComonObj.comul #[none, none, none, SX, none], copy_hom_proof :: proofs) + | Expr.const ``Kernel.swap [xLvl, yLvl] => + let X := e.getAppArgs[0]! + let Y := e.getAppArgs[1]! + let SX ← computeSFinkerOf X xLvl + let SY ← computeSFinkerOf Y yLvl + let swap_hom_proof ← mkAppM ``braiding_hom #[SX, SY, ← idME X, ← idME Y] + let braiding ← mkAppM ``Iso.hom #[← mkAppM ``BraidedCategory.braiding #[SX, SY]] + return (braiding, swap_hom_proof :: proofs) + | Expr.const ``Kernel.lift [_, y₀Lvl, _] => + let (X, Y, xLvl, yLvl) ← getTypesFromKernel e + let args := e.getAppArgs + let κ := args[args.size - 1]! + if ← checkLeftUnitor κ then + let punitLvl ← match args[0]!.getAppFn with + | Expr.const ``Prod [punitLvl, _] => pure punitLvl + | _ => throwError "Expected a product with PUnit as the first component, got {args[0]!}." + let ey₀ := args[args.size - 2]! + let (leftUnitorExpr, left_unitor_hom_proof) ← constructUnitors Y ey₀ yLvl y₀Lvl punitLvl 0 + return (leftUnitorExpr, left_unitor_hom_proof :: proofs) + else if ← checkRightUnitor κ then + let punitLvl ← match args[0]!.getAppFn with + | Expr.const ``Prod [_, punitLvl] => pure punitLvl + | _ => throwError "Expected a product with PUnit as the first component, got {args[0]!}." + let ey₀ := args[args.size - 2]! + let (rightUnitorExpr, right_unitor_hom_proof) ← constructUnitors Y ey₀ yLvl y₀Lvl punitLvl 1 + return (rightUnitorExpr, right_unitor_hom_proof :: proofs) + else if ← checkAssociatorHom κ then + let (ex₀, ey₀, ez₀) ← getMEFromThreeProds args[args.size - 2]! + let (associatorExpr, associator_hom_proof) ← constructAssociatorHom X Y ex₀ ey₀ ez₀ + return (associatorExpr, associator_hom_proof :: proofs) + else if ← checkAssociatorInv κ then + let (ex₀, ey₀, ez₀) ← getMEFromThreeProds args[args.size - 3]! + let (associatorInvExpr, associator_inv_hom_proof) ← constructAssociatorInv X Y ex₀ ey₀ ez₀ + return (associatorInvExpr, associator_inv_hom_proof :: proofs) + else + let SX ← computeSFinkerOf X xLvl + let SY ← computeSFinkerOf Y yLvl + let homExpr ← mkAppOptM ``ProbabilityTheory.Kernel.hom + #[X, Y, none, none, SX, SY, (← idME X), (← idME Y), e, none] + pure (homExpr, proofs) + | _ => + throwError "Expected a lifted kernel expression, got: {e}." + +/-- Construct the proof of equivalence between the original equality and the transformed one. -/ +def mkKernelHomEqProof (eqProofType lhs rhs : Expr) (proofs : List Expr) : MetaM Expr := do + let mvar ← mkFreshExprSyntheticOpaqueMVar eqProofType + let mvarId := mvar.mvarId! + let propext := mkConst ``propext + match ← mvarId.apply propext with + | [mvarId] => + let proofs := proofs.reverse + let mut mvarId := mvarId + for proof in proofs do + mvarId ← mvarId.nthRewrite 1 proof + let (X, Y, xLvl, yLvl) ← getTypesFromKernel lhs + let SX ← computeSFinkerOf X xLvl + let SY ← computeSFinkerOf Y yLvl + let e ← mkAppMInst ``hom_congr #[SX, SY, ← idME X, ← idME Y, lhs, rhs] 2 + unless ← isDefEq (← mvarId.getType) (← inferType e) do + throwError "Type mismatch: expected {← mvarId.getType}, got {← inferType e}." + mvarId.assign e + instantiateMVars mvar + | _ => + throwError "Failed to apply propext while building kernel_lift equivalence proof for + {eqProofType}." + +/-- Transform a kernel equality into an equivalent equality in `SFinKer`, along with a proof of +equivalence. -/ +def HomEquality (eq : Expr) : MetaM (Expr × Expr) := do + let eq ← unfoldKernelOp eq + let (lifted_expr, lifted_proof) ← liftEquality eq + let some (_, lhs, rhs) := lifted_expr.eq? | throwError "Expected an equality, got: {lifted_expr}." + let (lhs_hom, proofs) ← transformKernelToHom lhs [] + let (rhs_hom, proofs) ← transformKernelToHom rhs proofs + let hom_expr ← mkEq lhs_hom rhs_hom + let hom_eq_proof_type ← mkEq lifted_expr hom_expr + let hom_eq_proof ← mkKernelHomEqProof hom_eq_proof_type lhs rhs proofs + return (hom_expr, ← mkEqTrans lifted_proof hom_eq_proof) + +/-- The `kernel_hom` tactic transforms a kernel equality to an equivalent equality in +the category of measurable spaces and s-finite kernels. + +The tactic supports location specifiers like `rw` or `simp`: +* `kernel_hom` — applies to the goal +* `kernel_hom at h` — applies to hypothesis `h` +* `kernel_hom at h₁ h₂` — applies to multiple hypotheses +* `kernel_hom at h ⊢` — applies to hypothesis `h` and the goal +* `kernel_hom at *` — applies to all hypotheses and the goal + +Example: +```lean +example {W X Y Z : Type*} [MeasurableSpace X] [MeasurableSpace Y] [MeasurableSpace Z] + [MeasurableSpace W] (κ : Kernel X Y) (η : Kernel Y Z) (ξ : Kernel Z W) + [IsFiniteKernel ξ] [IsSFiniteKernel κ] [IsSFiniteKernel η] : + ξ ∘ₖ (η ∘ₖ κ) = ξ ∘ₖ η ∘ₖ κ := by + kernel_hom + exact Category.assoc _ _ _ +``` -/ +syntax (name := kernelHom) "kernel_hom" (ppSpace location)? : tactic + +elab_rules : tactic + | `(tactic| kernel_hom $[$loc]?) => + expandOptLocation (Lean.mkOptionalNode loc) |> applyLocTactic <| HomEquality diff --git a/LeanMachineLearning/Tactic/KernelHom/Tactic/Reassoc.lean b/LeanMachineLearning/Tactic/KernelHom/Tactic/Reassoc.lean new file mode 100644 index 00000000..b0b4c97e --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom/Tactic/Reassoc.lean @@ -0,0 +1,131 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import Mathlib.Tactic.CategoryTheory.Reassoc +public import LeanMachineLearning.Tactic.KernelHom.Tactic.HomKernel + +/-! +# `kernel_reassoc` + +This file extends `Mathlib.Tactic.CategoryTheory.Reassoc` with a kernel-specific variant for +equalities of s-finite kernels. It mirrors the structure of `Mathlib.Tactic.CategoryTheory. +Reassoc`, but targets the kernel language developed in `KernelHom` rather than categorical +morphisms. +-/ + +public meta section + +open Lean Meta Elab Tactic ProbabilityTheory Mathlib.Tactic Reassoc + +/-- Same as `HomEquality`, but allows specifying a universe level that will be taken into account +when computing the maximum universe level. -/ +def HomEqualityToLvl (eq : Expr) (Lvl : Level) : MetaM (Expr × Expr) := do + let eq ← unfoldKernelOp eq + let (lifted_expr, lifted_proof) ← liftEqualityWithLevel Lvl eq + let some (_, lhs, rhs) := lifted_expr.eq? | throwError "Expected an equality, got: {lifted_expr}." + let (lhs_hom, proofs) ← transformKernelToHom lhs [] + let (rhs_hom, proofs) ← transformKernelToHom rhs proofs + let hom_expr ← mkEq lhs_hom rhs_hom + let hom_eq_proof_type ← mkEq lifted_expr hom_expr + let hom_eq_proof ← mkKernelHomEqProof hom_eq_proof_type lhs rhs proofs + return (hom_expr, ← mkEqTrans lifted_proof hom_eq_proof) + +/-- Replace all level metavariables appearing in an expression with named level parameters. -/ +def freshenLevelParam (e : Expr) : MetaM Expr := do + let mvarIds := (Lean.collectLevelMVars {} e).result + for mvarId in mvarIds do + Lean.assignLevelMVar mvarId (Level.param mvarId.name) + instantiateMVars e + +/-- Core handler for `@[kernel_reassoc]`. + +Given an equality between s-finite kernels, this constructs the corresponding reassociated +equality in `SFinKer` category, under the extra `Z` measurable space and instance binders needed to +state the result. The returned array contains the fresh level metavariables that still need to be +added to the declaration's universe levels. +-/ +def kernelReassocHandler (h_eq : Expr) : MetaM (Expr × Array LMVarId) := do + let eq_type ← inferType h_eq + let some (_, lhs, _) := eq_type.eq? | + throwError "Expected an equality, but got {eq_type}" + let (_, Y, _, _) ← getTypesFromKernel lhs + let u ← mkFreshLevelMVar + let proof : Expr ← + withLocalDecl `Z .implicit (mkSort (mkLevelSucc u)) fun Z => do + let mspaceType ← mkAppM ``MeasurableSpace #[Z] + withLocalDecl `inst .instImplicit mspaceType fun _inst => do + let kernelType ← mkAppMInst ``Kernel #[Y, Z] 2 + withLocalDeclD `ξ kernelType fun ξ => do + let sfiniteType ← mkAppM ``IsSFiniteKernel #[ξ] + withLocalDecl `inst_1 BinderInfo.instImplicit sfiniteType fun _inst_1 => do + let (_, hom_proof) ← HomEqualityToLvl eq_type u + let hom_proof ← mkAppM ``Eq.mp #[hom_proof, h_eq] + let (hom_proof_reassoc, _) ← reassocExprHom hom_proof + let univs ← collectExprUniverses eq_type + let maxLvl ← computeMaxLevel <| u :: univs + let (ξ_lift, _) ← liftKernel ξ maxLvl [] + let (ξ_hom, _) ← transformKernelToHom ξ_lift [] + let reassoc_body ← mkAppM' hom_proof_reassoc #[ξ_hom] + let (_, kernel_reassoc_proof) ← KernelEquality <| ← inferType reassoc_body + let kernel_reassoc_proof ← mkAppM ``Eq.mp #[kernel_reassoc_proof, reassoc_body] + mkLambdaFVars #[Z, _inst, ξ, _inst_1] kernel_reassoc_proof + let proof ← freshenLevelParam proof + return (proof, #[u.mvarId!]) + +/-- Same as `@[reassoc]`, but for equalities of s-finite kernels. -/ +syntax (name := kernelReassoc) "kernel_reassoc" optAttrArg : attr + +/-- Registry of kernel reassociation handlers. + +The default handler translates equalities of s-finite kernels, and additional handlers can be +registered to extend the attribute to other kernel-shaped equalities. +-/ +private initialize kernelreassocImplRef : IO.Ref (Array (Expr → MetaM (Expr × Array LMVarId))) ← + IO.mkRef #[kernelReassocHandler] + +/-- IO ref for reassociation handlers `kernel_reassoc` attribute, so that it can be extended +with additional handlers. Handlers take a proof of the equation. -/ +def registerKernelReassocExpr (f : Expr → MetaM (Expr × Array LMVarId)) : IO Unit := do + kernelreassocImplRef.modify (·.push f) + +/-- Reassociates the kernels in the type of `pf` using the registered handlers, +using `kernelReassocHandler` as the default. + +Returns the proof of the lemma along with a list of fresh level metavariables. -/ +def kernelreassocExpr (pf : Expr) : MetaM (Expr × Array LMVarId) := do + forallTelescopeReducing (← inferType pf) fun xs _ => do + let pf := mkAppN pf xs + let handlers ← kernelreassocImplRef.get + let (pf, levels) ← handlers.firstM (fun h => h pf) <|> do + throwError "`kernel_reassoc` can only be used on terms about equality of s-finite kernels." + return (← mkLambdaFVars xs pf, levels) + +private def kernelReassocImpl (src : Name) (ref : Syntax) (kind : AttributeKind) : AttrM Name := + match ref with + | `(attr| kernel_reassoc $optAttr) => MetaM.run' do + unless kind == AttributeKind.global do + throwAttrMustBeGlobal `reassoc kind + let tgt := src.appendAfter "_assoc" + addRelatedDecl src tgt ref optAttr fun value levels => do + Term.TermElabM.run' <| Term.withSynthesize do + let (pf, newLevelMVars) ← kernelreassocExpr value + let newNames := newLevelMVars.map (·.name) + for mvarId in newLevelMVars do + Lean.assignLevelMVar mvarId (Level.param mvarId.name) + let pf ← instantiateMVars pf + pure (pf, levels ++ newNames.toList) + return tgt + | _ => throwUnsupportedSyntax + +initialize + registerGeneratingAttr `kernelReassoc ((#[·]) <$> kernelReassocImpl · · ·) + registerBuiltinAttribute { + name := `kernelReassoc + descr := "" + applicationTime := .afterCompilation + add := (discard <| kernelReassocImpl · · ·) + } diff --git a/LeanMachineLearning/Tactic/KernelHom/Tactic/Utils.lean b/LeanMachineLearning/Tactic/KernelHom/Tactic/Utils.lean new file mode 100644 index 00000000..435015f4 --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom/Tactic/Utils.lean @@ -0,0 +1,37 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import Mathlib.Probability.Kernel.Composition.Prod +public import Mathlib.Probability.Kernel.Composition.CompProd + +/-! +# Kernel transformation utilities +-/ + +public meta section + +open Lean Meta ProbabilityTheory + +/-- Unfold kernel operations in an expression. -/ +def unfoldKernelOp (e : Expr) : MetaM Expr := do + let names := (.empty |> NameSet.insert <| ``Kernel.prod) |> NameSet.insert <| ``Kernel.compProd + transform e (post := fun e => do + let e' ← deltaExpand e names.contains + let e' ← Core.betaReduce e' + return .done e') + +/-- Returns the application `constName` `xs` with `n_impls` last arguments as implicit. -/ +def Lean.Meta.mkAppMInst (constName : Name) (xs : Array Expr) (n_impls : Nat) : MetaM Expr := do + let e ← mkAppM constName xs + let nones : Array (Option Expr) := Array.replicate n_impls none + mkAppOptM' e nones + +/-- Similar to `mkAppMInst`, but takes an `Expr` instead of a constant name. -/ +def Lean.Meta.mkAppMInst' (f : Expr) (xs : Array Expr) (n_insts : Nat) : MetaM Expr := do + let e ← mkAppM' f xs + let nones : Array (Option Expr) := Array.replicate n_insts none + mkAppOptM' e nones diff --git a/scripts/update_tactics.sh b/scripts/update_tactics.sh new file mode 100755 index 00000000..67ce3a78 --- /dev/null +++ b/scripts/update_tactics.sh @@ -0,0 +1,95 @@ +#!/usr/bin/env bash + +# Vendors the standalone tactic projects (EqLift, KernelHom) into +# LeanMachineLearning/Tactic: +# 1. clones each upstream repository into a temporary directory, +# 2. extracts its main folder (e.g. EqLift/EqLift) and its root module file +# (e.g. EqLift/EqLift.lean), +# 3. rewrites every `import EqLift.` / `import KernelHom.` (including the +# `public import` and `meta import` forms) into +# `import LeanMachineLearning.Tactic.EqLift.` / +# `import LeanMachineLearning.Tactic.KernelHom.`, +# 4. prepends the LML copyright header to the root module files, which the +# upstream ones do not carry. +# +# Warning: the destination folders are wiped and replaced, so any local edit +# made to the vendored files is lost. Check `git diff` after running. +# +# Usage: scripts/update_tactics.sh [project ...] (default: all projects) +# DEST= override the destination directory (defaults to +# LeanMachineLearning/Tactic), useful for dry runs. + +set -euo pipefail + +GITHUB_USER="gaetanserre" +PREFIX="LeanMachineLearning.Tactic" +ALL_PROJECTS=(EqLift KernelHom) + +REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +DEST="${DEST:-$REPO_ROOT/LeanMachineLearning/Tactic}" + +if [ "$#" -gt 0 ]; then + PROJECTS=("$@") +else + PROJECTS=("${ALL_PROJECTS[@]}") +fi + +HEADER='/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ + +' + +TMP_DIR="$(mktemp -d)" +trap 'rm -rf "$TMP_DIR"' EXIT + +# Build the sed program rewriting the imports of *every* known project, so that +# cross-project imports (KernelHom depends on EqLift) are rewritten too. +SED_ARGS=() +for project in "${ALL_PROJECTS[@]}"; do + SED_ARGS+=(-e "s/(^|[[:space:]])import ${project}\.(\\w)/\1import ${PREFIX}.${project}.\2/g") +done + +mkdir -p "$DEST" + +for project in "${PROJECTS[@]}"; do + echo "==> $project" + + git clone --quiet --depth 1 \ + "https://github.com/${GITHUB_USER}/${project}.git" "$TMP_DIR/$project" + + src_dir="$TMP_DIR/$project/$project" + src_root="$TMP_DIR/$project/$project.lean" + for path in "$src_dir" "$src_root"; do + if [ ! -e "$path" ]; then + echo " error: $path not found in the cloned repository" >&2 + exit 1 + fi + done + + # Extract the main folder and its root module file. + rm -rf "${DEST:?}/$project" "${DEST:?}/$project.lean" + cp -r "$src_dir" "$DEST/$project" + { printf '%s' "$HEADER"; cat "$src_root"; } > "$DEST/$project.lean" + + # Rewrite the imports. + mapfile -t files < <(find "$DEST/$project" -type f -name '*.lean') + files+=("$DEST/$project.lean") + sed -E -i "${SED_ARGS[@]}" "${files[@]}" + + echo " ${#files[@]} file(s) written to ${DEST#"$REPO_ROOT/"}/$project{,.lean}" + + # Sanity check: no unqualified import of a vendored project may remain. + for other in "${ALL_PROJECTS[@]}"; do + if grep -rEn "(^|[[:space:]])import ${other}\." "$DEST/$project" "$DEST/$project.lean"; then + echo " error: leftover unqualified '${other}.' imports (see above)" >&2 + exit 1 + fi + done +done + +echo +echo "Done. Check \`git diff\` for local edits that were overwritten, and" +echo "regenerate LeanMachineLearning.lean if the module list changed." From 2a6719aa45b38cd19338a44d3b6e5cda568d8c89 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Thu, 20 Aug 2026 15:24:16 +0200 Subject: [PATCH 2/4] `simp` -> `rfl` --- LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean b/LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean index e4ff8685..d305732a 100644 --- a/LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean +++ b/LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean @@ -73,7 +73,7 @@ instance {κ : Kernel X Y} [IsDeterministic κ] [IsMarkovKernel κ] : · rw [comp_apply', comp_apply', copy, deterministic_apply, lintegral_dirac', parallelComp_apply'] at this · convert this - all_goals try simp + all_goals try rfl · ext y simp [MeasurableEquiv.prod] aesop From 1411e72afd1c815db96a316a04141d29b7cde6a4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Thu, 20 Aug 2026 15:28:22 +0200 Subject: [PATCH 3/4] `kernel_lift` -> `kernel_hom` --- LeanMachineLearning/Tactic/KernelHom/Tactic/KernelHom.lean | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelHom.lean b/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelHom.lean index d718e679..5255cf3e 100644 --- a/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelHom.lean +++ b/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelHom.lean @@ -373,7 +373,7 @@ def mkKernelHomEqProof (eqProofType lhs rhs : Expr) (proofs : List Expr) : MetaM mvarId.assign e instantiateMVars mvar | _ => - throwError "Failed to apply propext while building kernel_lift equivalence proof for + throwError "Failed to apply propext while building kernel_hom equivalence proof for {eqProofType}." /-- Transform a kernel equality into an equivalent equality in `SFinKer`, along with a proof of From f2e697867257089055f32cd13404e1a0ca1a7e89 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Thu, 20 Aug 2026 18:52:02 +0200 Subject: [PATCH 4/4] Remove useless measurable equivalences --- LeanMachineLearning.lean | 1 - .../EqLift/ForMathlib/MeasurableEquiv.lean | 23 +----- .../Tactic/EqLift/Kernel/Lift.lean | 52 ++++++------- .../Tactic/EqLift/Tactic/Kernel/Utils.lean | 2 +- LeanMachineLearning/Tactic/KernelHom.lean | 1 - .../KernelHom/ForMathlib/MeasurableEquiv.lean | 26 ------- .../Tactic/KernelHom/Kernel/Hom.lean | 73 ++++++++++--------- .../Tactic/KernelHom/Tactic/KernelHom.lean | 9 +-- 8 files changed, 70 insertions(+), 117 deletions(-) delete mode 100644 LeanMachineLearning/Tactic/KernelHom/ForMathlib/MeasurableEquiv.lean diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index d6e92832..945a88b6 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -57,7 +57,6 @@ public import LeanMachineLearning.Tactic.EqLift.Tactic.Utils public import LeanMachineLearning.Tactic.KernelHom public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.Kernel public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.LIntegral -public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.MeasurableEquiv public import LeanMachineLearning.Tactic.KernelHom.Kernel.Hom public import LeanMachineLearning.Tactic.KernelHom.Kernel.MonoidalComp public import LeanMachineLearning.Tactic.KernelHom.Tactic.Delaborators diff --git a/LeanMachineLearning/Tactic/EqLift/ForMathlib/MeasurableEquiv.lean b/LeanMachineLearning/Tactic/EqLift/ForMathlib/MeasurableEquiv.lean index dea6839f..c131e961 100644 --- a/LeanMachineLearning/Tactic/EqLift/ForMathlib/MeasurableEquiv.lean +++ b/LeanMachineLearning/Tactic/EqLift/ForMathlib/MeasurableEquiv.lean @@ -15,7 +15,6 @@ working with products and unit types. ## Main declarations -* `MeasurableEquiv.prod`: product of measurable equivalences. * `MeasurableEquiv.punit`: measurable equivalence between `PUnit`s. -/ @@ -23,27 +22,9 @@ working with products and unit types. namespace MeasurableEquiv -universe w x y - -variable {X Y X' Y' : Type*} [MeasurableSpace X] [MeasurableSpace Y] [MeasurableSpace X'] - [MeasurableSpace Y'] (ex : X' ≃ᵐ X) (ey : Y' ≃ᵐ Y) - -/-- The product of two measurable equivalences is a measurable equivalence. -/ -def prod : X' × Y' ≃ᵐ X × Y where - toFun := fun (x', y') ↦ (ex x', ey y') - invFun := fun (x, y) ↦ (ex.symm x, ey.symm y) - left_inv := by simp [Function.LeftInverse] - right_inv := by simp [Function.RightInverse, Function.LeftInverse] - measurable_toFun := by simp only [Equiv.coe_fn_mk]; fun_prop - measurable_invFun := by simp only [Equiv.coe_fn_symm_mk]; fun_prop +universe x y /-- The measurable equivalence between two `PUnit`s. -/ -def punit : PUnit.{w + 1} ≃ᵐ PUnit.{x + 1} where - toFun := fun _ ↦ PUnit.unit - invFun := fun _ ↦ PUnit.unit - left_inv := by grind - right_inv := by grind - measurable_toFun := measurable_id - measurable_invFun := measurable_id +abbrev punit : PUnit.{x + 1} ≃ᵐ PUnit.{y + 1} := ofUniqueOfUnique _ _ end MeasurableEquiv diff --git a/LeanMachineLearning/Tactic/EqLift/Kernel/Lift.lean b/LeanMachineLearning/Tactic/EqLift/Kernel/Lift.lean index cae2baba..5a6af29e 100644 --- a/LeanMachineLearning/Tactic/EqLift/Kernel/Lift.lean +++ b/LeanMachineLearning/Tactic/EqLift/Kernel/Lift.lean @@ -113,7 +113,7 @@ lemma comp_lift (η : Kernel X Y) (κ : Kernel Z X) : lemma parallelComp_lift (κ : Kernel X Y) (η : Kernel Z T) : κ.lift (ex := ex) (ey := ey) ∥ₖ η.lift (ex := ez) (ey := et) = - lift (ex := ex.prod ez) (ey := ey.prod et) (κ ∥ₖ η) := by + lift (ex := ex.prodCongr ez) (ey := ey.prodCongr et) (κ ∥ₖ η) := by by_cases hκ : IsSFiniteKernel <| lift (ex := ex) (ey := ey) κ swap · simp only [hκ, not_false_eq_true, parallelComp_of_not_isSFiniteKernel_left, @@ -146,30 +146,30 @@ lemma discard_lift : discard.{_, w} X' = (discard X).lift (ex := ex) (ey := puni exact Set.indicator_eq_indicator (by grind) rfl -lemma copy_lift : copy X' = (copy X).lift (ex := ex) (ey := ex.prod ex) := by +lemma copy_lift : copy X' = (copy X).lift (ex := ex) (ey := ex.prodCongr ex) := by ext _ _ hs rw [lift_apply' _ _ _ _ hs] simp only [copy_apply] rw [Measure.dirac_apply' _ hs, Measure.dirac_apply'] · refine Set.indicator_eq_indicator ?_ rfl - simp [MeasurableEquiv.prod] + simp [MeasurableEquiv.prodCongr] · measurability -lemma swap_lift : swap X' Y' = (swap X Y).lift (ex := ex.prod ey) (ey := ey.prod ex) := by +lemma swap_lift : swap X' Y' = (swap X Y).lift (ex := ex.prodCongr ey) (ey := ey.prodCongr ex) := by ext a s hs rw [lift_apply' _ _ _ _ hs] simp only [swap_apply] rw [Measure.dirac_apply' _ hs, Measure.dirac_apply'] · refine Set.indicator_eq_indicator ?_ rfl - simp only [MeasurableEquiv.prod, MeasurableEquiv.coe_mk, Equiv.coe_fn_mk, Prod.swap_prod_mk, - Set.mem_image, Prod.mk.injEq, EmbeddingLike.apply_eq_iff_eq, Prod.exists, + simp only [prodCongr, MeasurableEquiv.coe_mk, Equiv.prodCongr_apply, Prod.map, coe_toEquiv, + Prod.swap_prod_mk, Set.mem_image, Prod.mk.injEq, EmbeddingLike.apply_eq_iff_eq, Prod.exists, exists_eq_right_right, exists_eq_right] grind · measurability lemma prod_lift (κ : Kernel X Y) (η : Kernel X Z) : κ.lift (ex := ex) (ey := ey) ×ₖ η.lift (ex := ex) (ey := ez) = - lift (ex := ex) (ey := ey.prod ez) (κ ×ₖ η) := by + lift (ex := ex) (ey := ey.prodCongr ez) (κ ×ₖ η) := by by_cases hκ : IsSFiniteKernel <| lift (ex := ex) (ey := ey) κ swap · simp only [hκ, not_false_eq_true, prod_of_not_isSFiniteKernel_left, @@ -181,32 +181,32 @@ lemma prod_lift (κ : Kernel X Y) (η : Kernel X Z) : (isSFinite_lift ex ez η).not.mpr hη] simp [lift] simp only [prod] - rw [← comp_lift (ex := ex.prod ex), ← parallelComp_lift, ← copy_lift] + rw [← comp_lift (ex := ex.prodCongr ex), ← parallelComp_lift, ← copy_lift] lemma compProd_lift (κ : Kernel X Y) (η : Kernel (X × Y) Z) : - κ.lift (ex := ex) (ey := ey) ⊗ₖ η.lift (ex := ex.prod ey) (ey := ez) = - lift (ex := ex) (ey := ey.prod ez) (κ ⊗ₖ η) := by + κ.lift (ex := ex) (ey := ey) ⊗ₖ η.lift (ex := ex.prodCongr ey) (ey := ez) = + lift (ex := ex) (ey := ey.prodCongr ez) (κ ⊗ₖ η) := by by_cases hκ : IsSFiniteKernel <| lift (ex := ex) (ey := ey) κ swap · simp only [hκ, not_false_eq_true, compProd_of_not_isSFiniteKernel_left, (isSFinite_lift ex ey κ).not.mpr hκ] simp [lift] - by_cases hη : IsSFiniteKernel <| lift (ex := ex.prod ey) (ey := ez) η - swap - · simp only [hη, not_false_eq_true, compProd_of_not_isSFiniteKernel_right, - (isSFinite_lift (ex.prod ey) ez η).not.mpr hη] - simp [lift] - simp only [compProd] - rw [← comp_lift (ex := ex.prod ex) (ey := ey.prod ez), ← copy_lift, - ← comp_lift (ex := ex.prod ey), ← parallelComp_lift, ← id_lift, - ← comp_lift (ex := ex.prod (ey.prod ey)), ← parallelComp_lift, ← id_lift, ← copy_lift, - ← comp_lift (ex := (ex.prod ey).prod ey), - ← comp_lift (ex := ez.prod ey), ← parallelComp_lift, ← id_lift, ← swap_lift] - congr - simp only [lift] - rw [deterministic_map (MeasurableEquiv.measurable _) (MeasurableEquiv.measurable _)] - ext _ : 1 - simp [comap_apply, deterministic_apply, MeasurableEquiv.prod, prodAssoc] + · by_cases hη : IsSFiniteKernel <| lift (ex := ex.prodCongr ey) (ey := ez) η + swap + · simp [hη, (isSFinite_lift (ex.prodCongr ey) ez η).not.mpr hη] + simp [lift] + · simp only [compProd] + rw [← comp_lift (ex := ex.prodCongr ex) (ey := ey.prodCongr ez), ← copy_lift, + ← comp_lift (ex := ex.prodCongr ey), ← parallelComp_lift, ← id_lift, + ← comp_lift (ex := ex.prodCongr (ey.prodCongr ey)), ← parallelComp_lift, ← id_lift, + ← copy_lift, + ← comp_lift (ex := (ex.prodCongr ey).prodCongr ey), + ← comp_lift (ex := ez.prodCongr ey), ← parallelComp_lift, ← id_lift, ← swap_lift] + congr + simp only [lift] + rw [deterministic_map (MeasurableEquiv.measurable _) (MeasurableEquiv.measurable _)] + ext _ : 1 + simp [comap_apply, deterministic_apply, MeasurableEquiv.prodCongr, prodAssoc] instance {κ : Kernel X Y} [IsDeterministic κ] : diff --git a/LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/Utils.lean b/LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/Utils.lean index 5f89dd32..49afb3eb 100644 --- a/LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/Utils.lean +++ b/LeanMachineLearning/Tactic/EqLift/Tactic/Kernel/Utils.lean @@ -55,7 +55,7 @@ partial def constructMeasurableEquiv (e : Expr) (eLevel maxLvl : Level) : MetaM let yLevel := univs[1]! let ex ← constructMeasurableEquiv X xLevel maxLvl let ey ← constructMeasurableEquiv Y yLevel maxLvl - let res ← mkAppOptM' (Expr.const ``MeasurableEquiv.prod [xLevel, yLevel, maxLvl, maxLvl]) + let res ← mkAppOptM' (Expr.const ``MeasurableEquiv.prodCongr [maxLvl, xLevel, maxLvl, yLevel]) #[none, none, none, none, none, none, none, none, ex, ey] return res | _ => mkAppOptM' (Expr.const ``MeasurableEquiv.ulift [eLevel, maxLvl]) #[e, none] diff --git a/LeanMachineLearning/Tactic/KernelHom.lean b/LeanMachineLearning/Tactic/KernelHom.lean index efff87e0..3d2e455b 100644 --- a/LeanMachineLearning/Tactic/KernelHom.lean +++ b/LeanMachineLearning/Tactic/KernelHom.lean @@ -8,7 +8,6 @@ module -- shake: keep-all --deprecated_module: ignore public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.Kernel public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.LIntegral -public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.MeasurableEquiv public import LeanMachineLearning.Tactic.KernelHom.Kernel.Hom public import LeanMachineLearning.Tactic.KernelHom.Kernel.MonoidalComp public import LeanMachineLearning.Tactic.KernelHom.Tactic.Delaborators diff --git a/LeanMachineLearning/Tactic/KernelHom/ForMathlib/MeasurableEquiv.lean b/LeanMachineLearning/Tactic/KernelHom/ForMathlib/MeasurableEquiv.lean deleted file mode 100644 index 4a6b95f1..00000000 --- a/LeanMachineLearning/Tactic/KernelHom/ForMathlib/MeasurableEquiv.lean +++ /dev/null @@ -1,26 +0,0 @@ -/- -Copyright (c) 2026 Gaëtan Serré. All rights reserved. -Released under Apache 2.0 license as described in the file LICENSE. -Authors: Gaëtan Serré --/ -module - -public import Mathlib.MeasureTheory.MeasurableSpace.Embedding - -/-! -# Measurable equivalences --/ - -@[expose] public section - -namespace MeasurableEquiv - -variable (α : Type*) [MeasurableSpace α] - -/-- The identity measurable equivalence. -/ -def id : α ≃ᵐ α where - toEquiv := .refl α - measurable_toFun := measurable_id - measurable_invFun := measurable_id - -end MeasurableEquiv diff --git a/LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean b/LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean index d305732a..afbd63f5 100644 --- a/LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean +++ b/LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean @@ -66,7 +66,7 @@ instance {κ : Kernel X Y} [IsDeterministic κ] [IsMarkovKernel κ] : simp only [hom, κ_hom] have := κ.parallelComp_self_comp_copy have := DFunLike.congr_fun (x := ex a) this - have := DFunLike.congr_fun (x := ey.prod ey '' s) this + have := DFunLike.congr_fun (x := ey.prodCongr ey '' s) this rw [comap_parallelComp_comap, map_parallelComp_map, comp_apply', comp_apply', copy, deterministic_apply, lintegral_dirac', comap_apply', map_apply', parallelComp_apply', lintegral_comap, lintegral_map] @@ -75,14 +75,14 @@ instance {κ : Kernel X Y} [IsDeterministic κ] [IsMarkovKernel κ] : · convert this all_goals try rfl · ext y - simp [MeasurableEquiv.prod] + simp [MeasurableEquiv.prodCongr] aesop · simp only [copy, deterministic_apply] rw [Measure.dirac_apply', Measure.dirac_apply'] · refine Set.indicator_eq_indicator ?_ rfl - simp [MeasurableEquiv.prod] + simp [MeasurableEquiv.prodCongr] aesop - · exact (measurableSet_image (ey.prod ey)).mpr hs + · exact (measurableSet_image (ey.prodCongr ey)).mpr hs · exact hs all_goals try measurability · exact Kernel.measurable_coe _ (by measurability) @@ -130,7 +130,7 @@ lemma comp_hom (η : Kernel X Y) (κ : Kernel Z X) [IsSFiniteKernel η] [IsSFini lemma parallelComp_hom (κ : Kernel X Y) (η : Kernel Z T) [IsSFiniteKernel η] [IsSFiniteKernel κ] : κ.hom (ex := ex) (ey := ey) ⊗ₘ η.hom (ex := ez) (ey := et) = - hom (ex := ex.prod ez) (ey := ey.prod et) (κ ∥ₖ η) := by + hom (ex := ex.prodCongr ez) (ey := ey.prodCongr et) (κ ∥ₖ η) := by ext : 1; dsimp simp only [hom] rw [id_parallelComp_comp_parallelComp_id, comap_parallelComp_comap, map_parallelComp_map] @@ -144,28 +144,27 @@ lemma id_hom : 𝟙 SX = Kernel.id.hom (ex := ex) (ey := ex) := by all_goals measurability lemma whiskerLeft (κ : Kernel X Y) [IsSFiniteKernel κ] : SZ ◁ κ.hom (ex := ex) (ey := ey) = - (Kernel.id (α := Z) ∥ₖ κ).hom (ex := ez.prod ex) (ey := ez.prod ey) := by + (Kernel.id (α := Z) ∥ₖ κ).hom (ex := ez.prodCongr ex) (ey := ez.prodCongr ey) := by ext _ _ hs; dsimp simp only [hom] rw [parallelComp_apply, comap_apply, map_apply, id_apply, comap_apply, map_apply, parallelComp_apply, id_apply] - · simp only [Measure.dirac_prod, MeasurableEquiv.prod] + · simp only [Measure.dirac_prod, MeasurableEquiv.prodCongr] rw [Measure.map_map, Measure.map_map, Measure.map_apply, Measure.map_apply] - · congr with y - · simp - · simp + · congr 3 + simp all_goals try fun_prop all_goals exact hs all_goals fun_prop lemma whiskerRight (κ : Kernel X Y) [IsSFiniteKernel κ] : κ.hom (ex := ex) (ey := ey) ▷ SZ = - (κ ∥ₖ Kernel.id (α := Z)).hom (ex := ex.prod ez) (ey := ey.prod ez) := by + (κ ∥ₖ Kernel.id (α := Z)).hom (ex := ex.prodCongr ez) (ey := ey.prodCongr ez) := by ext _ _ hs; dsimp simp only [hom] rw [parallelComp_apply, comap_apply, map_apply, id_apply, comap_apply, map_apply, parallelComp_apply, id_apply] - · simp only [Measure.prod_dirac, MeasurableEquiv.prod] + · simp only [Measure.prod_dirac, MeasurableEquiv.prodCongr] rw [Measure.map_map, Measure.map_map, Measure.map_apply, Measure.map_apply] · congr with y · simp @@ -182,84 +181,86 @@ lemma counit : ε[SX] = (Kernel.discard X).hom (ex := ex) (ey := punit) := by rw [deterministic_map (by fun_prop) (by fun_prop)] rfl -lemma comul : Δ[SX] = (Kernel.copy X).hom (ex := ex) (ey := ex.prod ex) := by +lemma comul : Δ[SX] = (Kernel.copy X).hom (ex := ex) (ey := ex.prodCongr ex) := by ext : 1; dsimp simp only [hom, copy] rw [deterministic_map (by fun_prop) (by fun_prop)] congr with x - all_goals simp [MeasurableEquiv.prod] + all_goals simp [MeasurableEquiv.prodCongr] lemma braiding_hom : (β_ SX SY).hom = - (Kernel.swap X Y).hom (ex := ex.prod ey) (ey := ey.prod ex) := by + (Kernel.swap X Y).hom (ex := ex.prodCongr ey) (ey := ey.prodCongr ex) := by ext : 1; dsimp simp only [hom, swap] rw [deterministic_map (by fun_prop) (by fun_prop)] congr with x - all_goals simp [MeasurableEquiv.prod] + all_goals simp [MeasurableEquiv.prodCongr] variable {X₀ Y₀ Z₀ : Type*} [MeasurableSpace X₀] [MeasurableSpace Y₀] [MeasurableSpace Z₀] (ex₀ : X ≃ᵐ X₀) (ey₀ : Y ≃ᵐ Y₀) (ez₀ : Z ≃ᵐ Z₀) -lemma leftUnitor_hom : (λ_ SX).hom = hom (ex := punit.prod ex) (ey := ex) - (lift (Kernel.id.map (Prod.snd : PUnit × X₀ → X₀)) (ex := punit.prod ex₀) (ey := ex₀)) := by +lemma leftUnitor_hom : (λ_ SX).hom = hom (ex := punit.prodCongr ex) (ey := ex) + (lift (Kernel.id.map (Prod.snd : PUnit × X₀ → X₀)) + (ex := punit.prodCongr ex₀) (ey := ex₀)) := by ext; dsimp rw [hom_apply', lift_apply', id_map (by fun_prop), id_map (by fun_prop), deterministic_apply', deterministic_apply', Set.image] · refine Set.indicator_eq_indicator ?_ rfl - simp [MeasurableEquiv.prod] + simp [MeasurableEquiv.prodCongr] all_goals measurability -lemma leftUnitor_inv : (λ_ SX).inv = hom (ex := ex) (ey := punit.prod ex) - (lift (Kernel.id.map (fun x ↦ (PUnit.unit, x))) (ex := ex₀) (ey := punit.prod ex₀)) := by +lemma leftUnitor_inv : (λ_ SX).inv = hom (ex := ex) (ey := punit.prodCongr ex) + (lift (Kernel.id.map (fun x ↦ (PUnit.unit, x))) (ex := ex₀) (ey := punit.prodCongr ex₀)) := by ext; dsimp rw [hom_apply', lift_apply', id_map (by fun_prop), id_map (by fun_prop), deterministic_apply', deterministic_apply'] · refine Set.indicator_eq_indicator ?_ rfl - simp [Set.image, MeasurableEquiv.prod] + simp [Set.image, MeasurableEquiv.prodCongr] constructor all_goals simp_all all_goals measurability -lemma rightUnitor_hom : (ρ_ SX).hom = hom (ex := ex.prod punit) (ey := ex) - (lift (Kernel.id.map (Prod.fst : X₀ × PUnit → X₀)) (ex := ex₀.prod punit) (ey := ex₀)) := by +lemma rightUnitor_hom : (ρ_ SX).hom = hom (ex := ex.prodCongr punit) (ey := ex) + (lift (Kernel.id.map (Prod.fst : X₀ × PUnit → X₀)) + (ex := ex₀.prodCongr punit) (ey := ex₀)) := by ext; dsimp rw [hom_apply', lift_apply', id_map (by fun_prop), id_map (by fun_prop), deterministic_apply', deterministic_apply'] · refine Set.indicator_eq_indicator ?_ rfl - simp [MeasurableEquiv.prod] + simp [MeasurableEquiv.prodCongr] all_goals measurability -lemma rightUnitor_inv : (ρ_ SX).inv = hom (ex := ex) (ey := ex.prod punit) - (lift (Kernel.id.map (fun x ↦ (x, PUnit.unit))) (ex := ex₀) (ey := ex₀.prod punit)) := by +lemma rightUnitor_inv : (ρ_ SX).inv = hom (ex := ex) (ey := ex.prodCongr punit) + (lift (Kernel.id.map (fun x ↦ (x, PUnit.unit))) (ex := ex₀) (ey := ex₀.prodCongr punit)) := by ext; dsimp rw [hom_apply', lift_apply', id_map (by fun_prop), id_map (by fun_prop), deterministic_apply', deterministic_apply'] · refine Set.indicator_eq_indicator ?_ rfl - simp [Set.image, MeasurableEquiv.prod] + simp [Set.image, MeasurableEquiv.prodCongr] constructor all_goals simp_all all_goals measurability lemma associator_hom : (α_ SX SY SZ).hom = - hom (ex := (ex.prod ey).prod ez) (ey := ex.prod (ey.prod ez)) - (lift (Kernel.deterministic prodAssoc (by fun_prop)) (ex := (ex₀.prod ey₀).prod ez₀) - (ey := ex₀.prod (ey₀.prod ez₀))) := by + hom (ex := (ex.prodCongr ey).prodCongr ez) (ey := ex.prodCongr (ey.prodCongr ez)) + (lift (Kernel.deterministic prodAssoc (by fun_prop)) + (ex := (ex₀.prodCongr ey₀).prodCongr ez₀) (ey := ex₀.prodCongr (ey₀.prodCongr ez₀))) := by ext; dsimp simp only [hom] rw [comap_apply', map_apply', lift_apply', deterministic_apply', deterministic_apply'] · refine Set.indicator_eq_indicator ?_ rfl - simp [MeasurableEquiv.prod, prodAssoc] + simp [MeasurableEquiv.prodCongr, prodAssoc] all_goals measurability lemma associator_inv : (α_ SX SY SZ).inv = - hom (ex := ex.prod (ey.prod ez)) (ey := (ex.prod ey).prod ez) - (lift (Kernel.deterministic prodAssoc.symm (by fun_prop)) (ex := ex₀.prod (ey₀.prod ez₀)) - (ey := (ex₀.prod ey₀).prod ez₀)) := by + hom (ex := ex.prodCongr (ey.prodCongr ez)) (ey := (ex.prodCongr ey).prodCongr ez) + (lift (Kernel.deterministic prodAssoc.symm (by fun_prop)) + (ex := ex₀.prodCongr (ey₀.prodCongr ez₀)) (ey := (ex₀.prodCongr ey₀).prodCongr ez₀)) := by ext; dsimp simp only [hom] rw [comap_apply', map_apply', lift_apply', deterministic_apply', deterministic_apply'] · refine Set.indicator_eq_indicator ?_ rfl - simp [MeasurableEquiv.prod, prodAssoc] + simp [MeasurableEquiv.prodCongr, prodAssoc] all_goals measurability end diff --git a/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelHom.lean b/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelHom.lean index 5255cf3e..e6281507 100644 --- a/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelHom.lean +++ b/LeanMachineLearning/Tactic/KernelHom/Tactic/KernelHom.lean @@ -7,7 +7,6 @@ module public import LeanMachineLearning.Tactic.KernelHom.Kernel.MonoidalComp public import LeanMachineLearning.Tactic.KernelHom.Tactic.Utils -public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.MeasurableEquiv public import Lean.Elab.Tactic.Location public import LeanMachineLearning.Tactic.EqLift.Tactic.Kernel.KernelLift @@ -61,7 +60,7 @@ partial def idME (X : Expr) : MetaM Expr := do let args := X.getAppArgs let id1 ← idME args[0]! let id2 ← idME args[1]! - mkAppM ``MeasurableEquiv.prod #[id1, id2] + mkAppM ``MeasurableEquiv.prodCongr #[id1, id2] | Expr.const ``PUnit [xLvl] | Expr.const ``Unit [xLvl] => let xLvl ← match xLvl with | Level.succ l => pure l @@ -69,7 +68,7 @@ partial def idME (X : Expr) : MetaM Expr := do let punitME := mkConst ``MeasurableEquiv.punit [xLvl, xLvl] mkAppM' punitME #[] | _ => - mkAppOptM ``MeasurableEquiv.id #[X, none] + mkAppOptM ``MeasurableEquiv.refl #[X, none] /-- Check if a kernel expression corresponds to a left or right whisker. -/ def checkWhiskers (κ : Expr) (offset : Nat) : MetaM Bool := do @@ -208,12 +207,12 @@ def getTypesFromThreeProds (prod : Expr) : def getMEFromThreeProds (me_prod : Expr) : MetaM (Expr × Expr × Expr) := do match me_prod.getAppFn with - | Expr.const ``MeasurableEquiv.prod _ => + | Expr.const ``MeasurableEquiv.prodCongr _ => let args := me_prod.getAppArgs let ex := args[args.size - 2]! let right := args[args.size - 1]! match right.getAppFn with - | Expr.const ``MeasurableEquiv.prod _ => + | Expr.const ``MeasurableEquiv.prodCongr _ => let rightArgs := right.getAppArgs let ey := rightArgs[rightArgs.size - 2]! let ez := rightArgs[rightArgs.size - 1]!