diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index 9bc9ac39..945a88b6 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -42,3 +42,27 @@ 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.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..c131e961 --- /dev/null +++ b/LeanMachineLearning/Tactic/EqLift/ForMathlib/MeasurableEquiv.lean @@ -0,0 +1,30 @@ +/- +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.punit`: measurable equivalence between `PUnit`s. +-/ + +@[expose] public section + +namespace MeasurableEquiv + +universe x y + +/-- The measurable equivalence between two `PUnit`s. -/ +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 new file mode 100644 index 00000000..5a6af29e --- /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.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, + (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.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.prodCongr] + · measurability + +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 [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.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, + (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.prodCongr ex), ← parallelComp_lift, ← copy_lift] + +lemma compProd_lift (κ : Kernel X Y) (η : Kernel (X × Y) Z) : + κ.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.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 κ] : + 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..49afb3eb --- /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.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] + +/-- 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..3d2e455b --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom.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.KernelHom.ForMathlib.Kernel +public import LeanMachineLearning.Tactic.KernelHom.ForMathlib.LIntegral +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/Kernel/Hom.lean b/LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean new file mode 100644 index 00000000..afbd63f5 --- /dev/null +++ b/LeanMachineLearning/Tactic/KernelHom/Kernel/Hom.lean @@ -0,0 +1,268 @@ +/- +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.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] + · rw [comp_apply', comp_apply', copy, deterministic_apply, lintegral_dirac', + parallelComp_apply'] at this + · convert this + all_goals try rfl + · ext y + simp [MeasurableEquiv.prodCongr] + aesop + · simp only [copy, deterministic_apply] + rw [Measure.dirac_apply', Measure.dirac_apply'] + · refine Set.indicator_eq_indicator ?_ rfl + simp [MeasurableEquiv.prodCongr] + aesop + · exact (measurableSet_image (ey.prodCongr 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.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] + · 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.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.prodCongr] + rw [Measure.map_map, Measure.map_map, Measure.map_apply, Measure.map_apply] + · 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.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.prodCongr] + 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.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.prodCongr] + +lemma braiding_hom : (β_ SX SY).hom = + (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.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.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.prodCongr] + all_goals measurability + +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.prodCongr] + constructor + all_goals simp_all + all_goals measurability + +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.prodCongr] + all_goals measurability + +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.prodCongr] + constructor + all_goals simp_all + all_goals measurability + +lemma associator_hom : (α_ SX SY SZ).hom = + 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.prodCongr, prodAssoc] + all_goals measurability + +lemma associator_inv : (α_ SX SY SZ).inv = + 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.prodCongr, 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 := +