diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index d9225776..d8eaac63 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -1,6 +1,11 @@ module -- shake: keep-all --deprecated_module: ignore public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.ChainRule +public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.CompProd +public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.Convex +public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.DataProcessing +public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.MapSequence +public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.Restrict public import LeanMachineLearning.ForMathlib.MeasureTheory.Measurable public import LeanMachineLearning.ForMathlib.MeasureTheory.MeasurableSpace.Embedding public import LeanMachineLearning.ForMathlib.MeasureTheory.Measure.AbsolutelyContinuous @@ -46,6 +51,7 @@ public import LeanMachineLearning.SequentialLearning.Algorithms.RoundRobin public import LeanMachineLearning.SequentialLearning.Algorithms.Uniform public import LeanMachineLearning.SequentialLearning.BayesStationaryEnv public import LeanMachineLearning.SequentialLearning.Deterministic +public import LeanMachineLearning.SequentialLearning.DivergenceDecomposition public import LeanMachineLearning.SequentialLearning.EvaluationEnv public import LeanMachineLearning.SequentialLearning.FeedbackMartingale public import LeanMachineLearning.SequentialLearning.FiniteActions diff --git a/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/ChainRule.lean b/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/ChainRule.lean index a9066f10..7b9f054c 100644 --- a/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/ChainRule.lean +++ b/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/ChainRule.lean @@ -25,8 +25,10 @@ source is countable), the function `a ↦ klDiv (κ a) (η a)` is measurable * `klDiv_compProd_eq_add_lintegral`: `klDiv (μ ⊗ₘ κ) (ν ⊗ₘ η) = klDiv μ ν + ∫⁻ a, klDiv (κ a) (η a) ∂μ`. -We also record the invariance of the divergence under measurable embeddings -(`klDiv_map_measurableEmbedding`) and measurable equivalences (`klDiv_map_measurableEquiv`). +We also record the data processing inequality for the two projections +(`klDiv_le_compProd`, `klDiv_comp_le_compProd`) and the invariance of the divergence under +measurable embeddings (`klDiv_map_measurableEmbedding`) and measurable equivalences +(`klDiv_map_measurableEquiv`). -/ @[expose] public section @@ -56,6 +58,25 @@ lemma klDiv_map_measurableEquiv (μ ν : Measure α) [IsFiniteMeasure μ] [IsFin klDiv (μ.map e) (ν.map e) = klDiv μ ν := klDiv_map_measurableEmbedding μ ν e.measurableEmbedding +/-- **Data processing inequality** for the first projection: for Markov kernels `κ` and `η`, +`μ` and `ν` are the images of `μ ⊗ₘ κ` and `ν ⊗ₘ η` under `Prod.fst`, hence +`klDiv μ ν ≤ klDiv (μ ⊗ₘ κ) (ν ⊗ₘ η)`. -/ +lemma klDiv_le_compProd (μ ν : Measure α) [IsFiniteMeasure μ] [IsFiniteMeasure ν] + (κ η : Kernel α β) [IsMarkovKernel κ] [IsMarkovKernel η] : + klDiv μ ν ≤ klDiv (μ ⊗ₘ κ) (ν ⊗ₘ η) := by + conv_lhs => rw [← Measure.fst_compProd μ κ, ← Measure.fst_compProd ν η] + rw [Measure.fst, Measure.fst] + exact klDiv_map_le _ _ measurable_fst + +/-- **Data processing inequality** for the second projection: `κ ∘ₘ μ` and `η ∘ₘ ν` are the images +of `μ ⊗ₘ κ` and `ν ⊗ₘ η` under `Prod.snd`, hence +`klDiv (κ ∘ₘ μ) (η ∘ₘ ν) ≤ klDiv (μ ⊗ₘ κ) (ν ⊗ₘ η)`. -/ +lemma klDiv_comp_le_compProd (μ ν : Measure α) [IsFiniteMeasure μ] [IsFiniteMeasure ν] + (κ η : Kernel α β) [IsFiniteKernel κ] [IsFiniteKernel η] : + klDiv (κ ∘ₘ μ) (η ∘ₘ ν) ≤ klDiv (μ ⊗ₘ κ) (ν ⊗ₘ η) := by + rw [← Measure.snd_compProd μ κ, ← Measure.snd_compProd ν η, Measure.snd, Measure.snd] + exact klDiv_map_le _ _ measurable_snd + section kernel variable [MeasurableSpace.CountableOrCountablyGenerated α β] {κ η : Kernel α β} [IsFiniteKernel κ] diff --git a/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/CompProd.lean b/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/CompProd.lean new file mode 100644 index 00000000..bb592403 --- /dev/null +++ b/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/CompProd.lean @@ -0,0 +1,113 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import Mathlib.InformationTheory.KullbackLeibler.Basic + +/-! # Lemmas about the Kullback-Leibler divergence of the images of two measures by a measurable map +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Set +open scoped ENNReal NNReal + +namespace InformationTheory + +variable {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {μ ν : Measure α} [IsFiniteMeasure μ] [IsFiniteMeasure ν] + +/-- Transporting `μ ⊗ₘ η.comap f` along `f` in the first coordinate gives `μ.map f ⊗ₘ η`. -/ +lemma _root_.MeasureTheory.Measure.map_compProd_comap (μ : Measure α) [SFinite μ] + (η : Kernel β γ) [IsSFiniteKernel η] {f : α → β} (hf : Measurable f) : + (μ ⊗ₘ η.comap f hf).map (fun p : α × γ ↦ (f p.1, p.2)) = μ.map f ⊗ₘ η := by + ext s hs + rw [Measure.map_apply (by fun_prop) hs, Measure.compProd_apply (hs.preimage (by fun_prop)), + Measure.compProd_apply hs, lintegral_map (Kernel.measurable_kernel_prodMk_left hs) hf] + rfl + +omit [IsFiniteMeasure ν] in +lemma _root_.MeasureTheory.Measure.map_withDensity_comp + {f : β → ℝ≥0∞} (hf : Measurable f) {g : α → β} (hg : Measurable g) : + (ν.withDensity (f ∘ g)).map g = (ν.map g).withDensity f := by + ext s hs + rw [Measure.map_apply hg hs, withDensity_apply _ (hg hs), withDensity_apply _ hs, + ← lintegral_indicator hs, ← lintegral_indicator (hg hs), lintegral_map (hf.indicator hs) hg] + rfl + +lemma klDiv_withDensity_comp_map {f : β → ℝ≥0∞} (hf : Measurable f) {g : α → β} (hg : Measurable g) + [IsFiniteMeasure (ν.withDensity (f ∘ g))] : + klDiv ((ν.withDensity (f ∘ g)).map g) (ν.map g) = klDiv (ν.withDensity (f ∘ g)) ν := by + have hac : ν.withDensity (f ∘ g) ≪ ν := withDensity_absolutelyContinuous ν (f ∘ g) + have h_rnDeriv : ((ν.withDensity (f ∘ g)).map g).rnDeriv (ν.map g) =ᵐ[ν.map g] f := by + rw [Measure.map_withDensity_comp hf hg] + exact Measure.rnDeriv_withDensity _ hf + have hmeas : Measurable fun x : β ↦ + ENNReal.ofReal (klFun (((ν.withDensity (f ∘ g)).map g).rnDeriv (ν.map g) x).toReal) := + (measurable_klFun.comp (Measure.measurable_rnDeriv _ _).ennreal_toReal).ennreal_ofReal + rw [klDiv_eq_lintegral_klFun_of_ac (hac.map hg), klDiv_eq_lintegral_klFun_of_ac hac, + lintegral_map hmeas hg] + refine lintegral_congr_ae ?_ + filter_upwards [Measure.rnDeriv_withDensity ν (hf.comp hg), + ae_of_ae_map hg.aemeasurable h_rnDeriv] with x hx1 hx2 + rw [hx1, hx2] + rfl + +/-- If `μ` has density `f ∘ g` with respect to `ν`, then the divergence of the images by `g` is +the divergence of `μ` and `ν`: `g` is a sufficient statistic. -/ +lemma klDiv_map_of_eq_withDensity_comp {f : β → ℝ≥0∞} (hf : Measurable f) {g : α → β} + (hg : Measurable g) (hμ : μ = ν.withDensity (f ∘ g)) : + klDiv (μ.map g) (ν.map g) = klDiv μ ν := by + rw [hμ] + have : IsFiniteMeasure (ν.withDensity (f ∘ g)) := by rw [← hμ]; infer_instance + exact klDiv_withDensity_comp_map hf hg + +/-- The conditional divergence of two kernels which depend on the conditioning variable only +through a statistic `f` is the conditional divergence given `f`. -/ +lemma klDiv_compProd_comap (μ : Measure α) [IsFiniteMeasure μ] (κ η : Kernel β γ) + [IsFiniteKernel κ] [IsFiniteKernel η] {f : α → β} (hf : Measurable f) : + klDiv (μ ⊗ₘ κ.comap f hf) (μ ⊗ₘ η.comap f hf) = klDiv (μ.map f ⊗ₘ κ) (μ.map f ⊗ₘ η) := by + have hg : Measurable fun p : α × γ ↦ (f p.1, p.2) := by fun_prop + by_cases hac : μ.map f ⊗ₘ κ ≪ μ.map f ⊗ₘ η + swap + · rw [klDiv_of_not_ac hac, klDiv_of_not_ac] + refine fun h ↦ hac ?_ + have := h.map hg + rwa [Measure.map_compProd_comap, Measure.map_compProd_comap] at this + let D := (μ.map f ⊗ₘ κ).rnDeriv (μ.map f ⊗ₘ η) + have hD : Measurable D := Measure.measurable_rnDeriv _ _ + have hDκ : μ.map f ⊗ₘ κ = (μ.map f ⊗ₘ η).withDensity D := + (Measure.withDensity_rnDeriv_eq _ _ hac).symm + -- for every measurable `t`, the sections of the density integrate to `κ (f a) t`, `μ`-a.e. + have h_sect {t : Set γ} (ht : MeasurableSet t) : + ∀ᵐ a ∂μ, ∫⁻ c in t, D (f a, c) ∂(η (f a)) = κ (f a) t := by + refine ae_of_ae_map (p := fun b ↦ ∫⁻ c in t, D (b, c) ∂(η b) = κ b t) hf.aemeasurable ?_ + refine ae_eq_of_forall_setLIntegral_eq_of_sigmaFinite + (Measurable.setLIntegral_kernel_prod_right (f := fun b c ↦ D (b, c)) hD ht) + (Kernel.measurable_coe κ ht) fun u hu _ ↦ ?_ + have h1 := congrArg (fun ρ : Measure (β × γ) ↦ ρ (u ×ˢ t)) hDκ + rw [Measure.compProd_apply_prod hu ht, withDensity_apply _ (hu.prod ht), + Measure.setLIntegral_compProd hD hu ht] at h1 + exact h1.symm + have h_rect s t (hs : MeasurableSet s) (ht : MeasurableSet t) : + (μ ⊗ₘ κ.comap f hf) (s ×ˢ t) = + ((μ ⊗ₘ η.comap f hf).withDensity (D ∘ fun p ↦ (f p.1, p.2))) (s ×ˢ t) := by + rw [Measure.compProd_apply_prod hs ht, withDensity_apply _ (hs.prod ht), + Measure.setLIntegral_compProd (hD.comp hg) hs ht] + refine setLIntegral_congr_fun_ae hs ?_ + filter_upwards [h_sect ht] with a ha _ + simp only [Kernel.comap_apply, Function.comp_apply] + exact ha.symm + have key : μ ⊗ₘ κ.comap f hf = + (μ ⊗ₘ η.comap f hf).withDensity (D ∘ fun p ↦ (f p.1, p.2)) := by + refine ext_of_generate_finite _ generateFrom_prod.symm isPiSystem_prod ?_ ?_ + · rintro _ ⟨s, hs, t, ht, rfl⟩ + exact h_rect s t hs ht + · simpa using h_rect Set.univ Set.univ MeasurableSet.univ MeasurableSet.univ + rw [← Measure.map_compProd_comap μ κ hf, ← Measure.map_compProd_comap μ η hf, + klDiv_map_of_eq_withDensity_comp hD hg key] + +end InformationTheory diff --git a/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/Convex.lean b/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/Convex.lean new file mode 100644 index 00000000..9e4c4e60 --- /dev/null +++ b/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/Convex.lean @@ -0,0 +1,174 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import Mathlib.InformationTheory.KullbackLeibler.Basic + +import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.ChainRule + +/-! +# Convexity of the Kullback–Leibler divergence for mixtures + +For finite measures `μ i`, `ν i` on `Ω` and weights `c i ≥ 0`, the Kullback–Leibler divergence is +convex in the pair of measures: +`klDiv (∑ i, c i • μ i) (∑ i, c i • ν i) ≤ ∑ i, c i * klDiv (μ i) (ν i)`. + +The proof combines the data processing inequality with the integral form of the conditional +divergence. Let `β := ∑ i, c i • δ i` be the measure with weights `c` on the index set and let +`κ`, `η` be the kernels from the index set given by `μ` and `ν`. The two mixtures are the +compositions `κ ∘ₘ β` and `η ∘ₘ β`, that is, the images of `β ⊗ₘ κ` and `β ⊗ₘ η` under the second +projection, so the data processing inequality bounds `klDiv (κ ∘ₘ β) (η ∘ₘ β)` by the conditional +divergence `klDiv (β ⊗ₘ κ) (β ⊗ₘ η)`, which is `∫⁻ i, klDiv (κ i) (η i) ∂β`. + +## Main statements + +* `InformationTheory.klDiv_finsetSum_smul_le`, + `InformationTheory.klDiv_sum_smul_le`: convexity of `klDiv` in the pair of measures, + for mixtures indexed by a `Finset` and by a `Fintype`; +* `InformationTheory.klDiv_smul_add_smul_le`: the same statement for two-point mixtures; +* `InformationTheory.klDiv_finsetSum_smul_left_le`, `InformationTheory.klDiv_sum_smul_left_le`: + convexity of `klDiv` in its first argument; +* `InformationTheory.klDiv_finsetSum_smul_right_le`, `InformationTheory.klDiv_sum_smul_right_le`: + convexity of `klDiv` in its second argument. +-/ + +@[expose] public section + +open Real MeasureTheory ProbabilityTheory Set +open scoped ENNReal NNReal + +namespace InformationTheory + +variable {Ω ι : Type*} {mΩ : MeasurableSpace Ω} + +section pair + +variable {μ ν : ι → Measure Ω} + +/-- **Convexity of the Kullback–Leibler divergence in the pair of measures**, for a mixture +indexed by a `Finset`: for finite measures `μ i`, `ν i` and weights `c i ≥ 0`, +`klDiv (∑ i ∈ s, c i • μ i) (∑ i ∈ s, c i • ν i) ≤ ∑ i ∈ s, c i * klDiv (μ i) (ν i)`. +Weights summing to `1` give the convexity of `klDiv`; no such hypothesis is needed here, since +`klDiv` is positively homogeneous. -/ +lemma klDiv_finsetSum_smul_le [∀ i, IsFiniteMeasure (μ i)] + [∀ i, IsFiniteMeasure (ν i)] (s : Finset ι) (c : ι → ℝ≥0) : + klDiv (∑ i ∈ s, (c i : ℝ≥0∞) • μ i) (∑ i ∈ s, (c i : ℝ≥0∞) • ν i) + ≤ ∑ i ∈ s, (c i : ℝ≥0∞) * klDiv (μ i) (ν i) := by + classical + -- the index set, with the discrete measurable structure + let _ : MeasurableSpace s := ⊤ + have : MeasurableSingletonClass s := ⟨fun _ ↦ trivial⟩ + -- the measure with weights `c` on the index set + set β : Measure s := ∑ i : s, (c i : ℝ≥0∞) • Measure.dirac i with hβ_def + have hβ_lintegral (f : s → ℝ≥0∞) : ∫⁻ i, f i ∂β = ∑ i : s, (c i : ℝ≥0∞) * f i := by + simp only [hβ_def, lintegral_finsetSum_measure, lintegral_smul_measure, lintegral_dirac, + smul_eq_mul] + have : IsFiniteMeasure β := ⟨by + simpa using (hβ_lintegral 1).trans_lt (ENNReal.sum_lt_top.mpr fun i _ ↦ by simp)⟩ + -- the kernels from the index set given by the two families of measures + set κ : Kernel s Ω := Kernel.ofFunOfCountable fun i ↦ μ i + set η : Kernel s Ω := Kernel.ofFunOfCountable fun i ↦ ν i + have hκ_apply (i : s) : κ i = μ i := rfl + have hη_apply (i : s) : η i = ν i := rfl + have h_fin (ρ : Kernel s Ω) (ρ' : ι → Measure Ω) [∀ i, IsFiniteMeasure (ρ' i)] + (hρ : ∀ i : s, ρ i = ρ' i) : IsFiniteKernel ρ := + ⟨∑ i : s, ρ' i univ, ENNReal.sum_lt_top.mpr fun i _ ↦ measure_lt_top _ _, fun i ↦ by + rw [hρ i] + exact Finset.single_le_sum (f := fun j : s ↦ ρ' j univ) (fun _ _ ↦ bot_le) + (Finset.mem_univ i)⟩ + have : IsFiniteKernel κ := h_fin κ μ hκ_apply + have : IsFiniteKernel η := h_fin η ν hη_apply + -- composing a kernel with `β` gives the corresponding mixture + have h_comp (ρ : Kernel s Ω) (ρ' : ι → Measure Ω) (hρ : ∀ i : s, ρ i = ρ' i) : + ρ ∘ₘ β = ∑ i ∈ s, (c i : ℝ≥0∞) • ρ' i := by + ext t ht + rw [Measure.bind_apply ht ρ.aemeasurable, hβ_lintegral, Measure.finsetSum_apply, + ← Finset.sum_coe_sort s] + exact Finset.sum_congr rfl fun i _ ↦ by rw [hρ i, Measure.smul_apply, smul_eq_mul] + calc klDiv (∑ i ∈ s, (c i : ℝ≥0∞) • μ i) (∑ i ∈ s, (c i : ℝ≥0∞) • ν i) + _ = klDiv (κ ∘ₘ β) (η ∘ₘ β) := by rw [h_comp κ μ hκ_apply, h_comp η ν hη_apply] + -- data processing inequality for the second projection + _ ≤ klDiv (β ⊗ₘ κ) (β ⊗ₘ η) := klDiv_comp_le_compProd β β κ η + -- integral form of the conditional divergence + _ = ∫⁻ i, klDiv (κ i) (η i) ∂β := klDiv_compProd_right_eq_lintegral β κ η + _ = ∑ i ∈ s, (c i : ℝ≥0∞) * klDiv (μ i) (ν i) := by + rw [hβ_lintegral] + simp_rw [hκ_apply, hη_apply] + exact Finset.sum_coe_sort s fun i ↦ (c i : ℝ≥0∞) * klDiv (μ i) (ν i) + +/-- **Convexity of the Kullback–Leibler divergence in the pair of measures**: for finite measures +`μ i`, `ν i` and weights `c i ≥ 0`, +`klDiv (∑ i, c i • μ i) (∑ i, c i • ν i) ≤ ∑ i, c i * klDiv (μ i) (ν i)`. -/ +lemma klDiv_sum_smul_le [Fintype ι] [∀ i, IsFiniteMeasure (μ i)] + [∀ i, IsFiniteMeasure (ν i)] (c : ι → ℝ≥0) : + klDiv (∑ i, (c i : ℝ≥0∞) • μ i) (∑ i, (c i : ℝ≥0∞) • ν i) + ≤ ∑ i, (c i : ℝ≥0∞) * klDiv (μ i) (ν i) := + klDiv_finsetSum_smul_le Finset.univ c + +/-- **Convexity of the Kullback–Leibler divergence in the pair of measures**, for a two-point +mixture: for finite measures `μ₀, μ₁, ν₀, ν₁` and weights `a, b ≥ 0`, +`klDiv (a • μ₀ + b • μ₁) (a • ν₀ + b • ν₁) ≤ a * klDiv μ₀ ν₀ + b * klDiv μ₁ ν₁`. -/ +lemma klDiv_smul_add_smul_le (μ₀ μ₁ ν₀ ν₁ : Measure Ω) [IsFiniteMeasure μ₀] [IsFiniteMeasure μ₁] + [IsFiniteMeasure ν₀] [IsFiniteMeasure ν₁] (a b : ℝ≥0) : + klDiv ((a : ℝ≥0∞) • μ₀ + (b : ℝ≥0∞) • μ₁) ((a : ℝ≥0∞) • ν₀ + (b : ℝ≥0∞) • ν₁) + ≤ (a : ℝ≥0∞) * klDiv μ₀ ν₀ + (b : ℝ≥0∞) * klDiv μ₁ ν₁ := by + have : ∀ x : Bool, IsFiniteMeasure (bif x then μ₁ else μ₀) := fun x ↦ by + cases x <;> assumption + have : ∀ x : Bool, IsFiniteMeasure (bif x then ν₁ else ν₀) := fun x ↦ by + cases x <;> assumption + have h := klDiv_sum_smul_le (μ := fun x ↦ bif x then μ₁ else μ₀) + (ν := fun x ↦ bif x then ν₁ else ν₀) (fun x ↦ bif x then b else a) + simpa [Fintype.sum_bool, add_comm] using h + +end pair + +section firstArgument + +variable {μ : ι → Measure Ω} {ν : Measure Ω} {c : ι → ℝ≥0} + +/-- **Convexity of the Kullback–Leibler divergence in its first argument**, for a finite mixture +indexed by a `Finset`: for finite measures `μ i`, `ν` and nonnegative weights `c i` summing to +`1`, `klDiv (∑ i ∈ s, c i • μ i) ν ≤ ∑ i ∈ s, c i * klDiv (μ i) ν`. -/ +lemma klDiv_finsetSum_smul_left_le [∀ i, IsFiniteMeasure (μ i)] [IsFiniteMeasure ν] + {s : Finset ι} (hc : ∑ i ∈ s, c i = 1) : + klDiv (∑ i ∈ s, (c i : ℝ≥0∞) • μ i) ν ≤ ∑ i ∈ s, (c i : ℝ≥0∞) * klDiv (μ i) ν := by + have h := klDiv_finsetSum_smul_le (μ := μ) (ν := fun _ ↦ ν) s c + rwa [← Finset.sum_smul, ← ENNReal.ofNNReal_finsetSum, hc, ENNReal.coe_one, one_smul] at h + +/-- **Convexity of the Kullback–Leibler divergence in its first argument**: for finite measures +`μ i`, `ν` and nonnegative weights `c i` summing to `1`, +`klDiv (∑ i, c i • μ i) ν ≤ ∑ i, c i * klDiv (μ i) ν`. -/ +lemma klDiv_sum_smul_left_le [Fintype ι] [∀ i, IsFiniteMeasure (μ i)] [IsFiniteMeasure ν] + (hc : ∑ i, c i = 1) : + klDiv (∑ i, (c i : ℝ≥0∞) • μ i) ν ≤ ∑ i, (c i : ℝ≥0∞) * klDiv (μ i) ν := + klDiv_finsetSum_smul_left_le hc + +end firstArgument + +section secondArgument + +variable {μ : Measure Ω} {ν : ι → Measure Ω} {c : ι → ℝ≥0} + +/-- **Convexity of the Kullback–Leibler divergence in its second argument**, for a finite mixture +indexed by a `Finset`: for finite measures `μ`, `ν i` and nonnegative weights `c i` summing to +`1`, `klDiv μ (∑ i ∈ s, c i • ν i) ≤ ∑ i ∈ s, c i * klDiv μ (ν i)`. -/ +lemma klDiv_finsetSum_smul_right_le [IsFiniteMeasure μ] [∀ i, IsFiniteMeasure (ν i)] + {s : Finset ι} (hc : ∑ i ∈ s, c i = 1) : + klDiv μ (∑ i ∈ s, (c i : ℝ≥0∞) • ν i) ≤ ∑ i ∈ s, (c i : ℝ≥0∞) * klDiv μ (ν i) := by + have h := klDiv_finsetSum_smul_le (μ := fun _ ↦ μ) (ν := ν) s c + rwa [← Finset.sum_smul, ← ENNReal.ofNNReal_finsetSum, hc, ENNReal.coe_one, one_smul] at h + +/-- **Convexity of the Kullback–Leibler divergence in its second argument**: for finite measures +`μ`, `ν i` and nonnegative weights `c i` summing to `1`, +`klDiv μ (∑ i, c i • ν i) ≤ ∑ i, c i * klDiv μ (ν i)`. -/ +lemma klDiv_sum_smul_right_le [Fintype ι] [IsFiniteMeasure μ] [∀ i, IsFiniteMeasure (ν i)] + (hc : ∑ i, c i = 1) : + klDiv μ (∑ i, (c i : ℝ≥0∞) • ν i) ≤ ∑ i, (c i : ℝ≥0∞) * klDiv μ (ν i) := + klDiv_finsetSum_smul_right_le hc + +end secondArgument + +end InformationTheory diff --git a/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/DataProcessing.lean b/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/DataProcessing.lean new file mode 100644 index 00000000..80d5a9a5 --- /dev/null +++ b/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/DataProcessing.lean @@ -0,0 +1,42 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import Mathlib.InformationTheory.KullbackLeibler.DataProcessing + +import Mathlib.MeasureTheory.Function.ConditionalExpectation.RadonNikodym + +/-! # Lemmas related to the data-processing inequality for the Kullback–Leibler divergence +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Set +open scoped ENNReal NNReal + +namespace InformationTheory + +variable {α β : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {μ ν : Measure α} [IsFiniteMeasure μ] [IsFiniteMeasure ν] + +/-- The divergence of the images of `μ ≪ ν` by a measurable map `g`, as the integral of `klFun` +of the conditional expectation of the density `∂μ/∂ν` given `g`. + +See `klDiv_map_of_ac` for a version of this lemma with a Bochner integral. -/ +lemma klDiv_map_eq_lintegral_klFun_condExp (hμν : μ ≪ ν) {g : α → β} (hg : Measurable g) : + klDiv (μ.map g) (ν.map g) = + ∫⁻ x, ENNReal.ofReal + (klFun ((ν[fun x ↦ (μ.rnDeriv ν x).toReal | mβ.comap g]) x)) ∂ν := by + have hmeas : Measurable fun y : β ↦ + ENNReal.ofReal (klFun ((μ.map g).rnDeriv (ν.map g) y).toReal) := + (measurable_klFun.comp (Measure.measurable_rnDeriv _ _).ennreal_toReal).ennreal_ofReal + rw [klDiv_eq_lintegral_klFun_of_ac (hμν.map hg), lintegral_map hmeas hg] + refine lintegral_congr_ae ?_ + filter_upwards [toReal_rnDeriv_map hμν hg] with x hx + rw [hx] + + +end InformationTheory diff --git a/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/MapSequence.lean b/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/MapSequence.lean new file mode 100644 index 00000000..7b1f1adc --- /dev/null +++ b/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/MapSequence.lean @@ -0,0 +1,195 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.DataProcessing +public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.Restrict +public import Mathlib.MeasureTheory.Integral.Indicator +public import Mathlib.Probability.Martingale.Convergence + +/-! +# The Kullback–Leibler divergence along a generating sequence of maps + +Let `μ, ν` be finite measures on `α` and `g n : α → β n` be measurable maps whose σ-algebras +`comap (g n)` increase to the σ-algebra of `α` (for instance the projections of a sequence space +on its first `n` coordinates). Then the divergence of `μ` and `ν` is the supremum of the +divergences of their images by the `g n`: +`klDiv μ ν = ⨆ n, klDiv (μ.map (g n)) (ν.map (g n))` (`klDiv_eq_iSup_map`). + +The inequality `≥` is the data-processing inequality. For `≤`, if `μ ≪ ν` the divergences of +the images are the integrals of `klFun` of the conditional expectations of the density `∂μ/∂ν` +given `comap (g n)` (`klDiv_map_eq_lintegral_klFun_condExp`), which converge almost everywhere +to the density by Lévy's upward theorem, and Fatou's lemma concludes. If `μ` is not absolutely +continuous with respect to `ν`, a set `A` with `ν A = 0 < μ A` is approximated by +`comap (g n)`-measurable sets `B n` (again by Lévy's upward theorem, applied to the indicator of +`A` under `μ + ν`), and the lower bound `μ B * log (μ B / ν B) + ν B - μ B ≤ klDiv` on `B n` +shows that the divergences of the images tend to infinity. + +The sequence space case is `MeasurableSpace.iSup_comap_restrictFin`: the σ-algebras of the +projections on the first `n` coordinates generate the product σ-algebra of `ℕ → E`. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Filter Topology +open scoped ENNReal + +namespace InformationTheory + +variable {α : Type*} {m0 : MeasurableSpace α} {μ ν : Measure α} [IsFiniteMeasure μ] + [IsFiniteMeasure ν] + +/-- **The divergence along a generating sequence of maps.** If the σ-algebras `comap (g n)` of +measurable maps `g n : α → β n` increase to the σ-algebra of `α`, then the divergence of two +finite measures is the supremum of the divergences of their images by the `g n`. -/ +lemma klDiv_eq_iSup_map {β : ℕ → Type*} [mβ : ∀ n, MeasurableSpace (β n)] {g : ∀ n, α → β n} + (hg : ∀ n, Measurable (g n)) + (hmono : Monotone fun n ↦ (mβ n).comap (g n)) + (hsup : ⨆ n, (mβ n).comap (g n) = m0) : + klDiv μ ν = ⨆ n, klDiv (μ.map (g n)) (ν.map (g n)) := by + let ℱ : Filtration ℕ m0 := ⟨fun n ↦ (mβ n).comap (g n), hmono, fun n ↦ (hg n).comap_le⟩ + have hℱ : ∀ n, ℱ n = (mβ n).comap (g n) := fun _ ↦ rfl + have hℱ_sup : (⨆ n, ℱ n) = m0 := hsup + refine le_antisymm ?_ (iSup_le fun n ↦ klDiv_map_le μ ν (hg n)) + by_cases hμν : μ ≪ ν + · -- Lévy's upward theorem and Fatou's lemma + let f : α → ℝ := fun x ↦ (μ.rnDeriv ν x).toReal + have hf_int : Integrable f ν := Measure.integrable_toReal_rnDeriv + have hf_meas : StronglyMeasurable[⨆ n, ℱ n] f := + (Measure.measurable_rnDeriv μ ν).ennreal_toReal.stronglyMeasurable.mono hℱ_sup.symm.le + have hlim := hf_int.tendsto_ae_condExp (ℱ := ℱ) hf_meas + have hmeas : ∀ n, Measurable fun x ↦ ENNReal.ofReal (klFun ((ν[f | ℱ n]) x)) := fun n ↦ + (measurable_klFun.comp + ((stronglyMeasurable_condExp (m := ℱ n)).measurable.mono (ℱ.le n) le_rfl)).ennreal_ofReal + calc klDiv μ ν = ∫⁻ x, ENNReal.ofReal (klFun (f x)) ∂ν := klDiv_eq_lintegral_klFun_of_ac hμν + _ = ∫⁻ x, liminf (fun n ↦ ENNReal.ofReal (klFun ((ν[f | ℱ n]) x))) atTop ∂ν := by + refine lintegral_congr_ae ?_ + filter_upwards [hlim] with x hx + exact ((ENNReal.continuous_ofReal.tendsto _).comp + ((continuous_klFun.tendsto _).comp hx)).liminf_eq.symm + _ ≤ liminf (fun n ↦ ∫⁻ x, ENNReal.ofReal (klFun ((ν[f | ℱ n]) x)) ∂ν) atTop := + lintegral_liminf_le hmeas + _ ≤ ⨆ n, ∫⁻ x, ENNReal.ofReal (klFun ((ν[f | ℱ n]) x)) ∂ν := + le_trans liminf_le_limsup limsup_le_iSup + _ = ⨆ n, klDiv (μ.map (g n)) (ν.map (g n)) := + iSup_congr fun n ↦ (klDiv_map_eq_lintegral_klFun_condExp hμν (hg n)).symm + · -- the divergences of the images tend to infinity + rw [klDiv_of_not_ac hμν, top_le_iff, iSup_eq_top] + obtain ⟨A, hA, hνA, hμA⟩ : ∃ A, MeasurableSet A ∧ ν A = 0 ∧ μ A ≠ 0 := by + by_contra! h + exact hμν (Measure.AbsolutelyContinuous.mk fun A hA hνA ↦ h A hA hνA) + -- approximation of `A` by `comap (g n)`-measurable sets, by Lévy's upward theorem + let ρ : Measure α := μ + ν + let φ : α → ℝ := A.indicator fun _ ↦ (1 : ℝ) + have hφ_int : Integrable φ ρ := (integrable_const (1 : ℝ)).indicator hA + have hφ_meas : StronglyMeasurable[⨆ n, ℱ n] φ := + (measurable_const.indicator hA).stronglyMeasurable.mono hℱ_sup.symm.le + have hlim := hφ_int.tendsto_ae_condExp (ℱ := ℱ) hφ_meas + set B : ℕ → Set α := fun n ↦ {x | (2⁻¹ : ℝ) < (ρ[φ | ℱ n]) x} with hB + have hB_meas : ∀ n, MeasurableSet[ℱ n] (B n) := fun n ↦ + (stronglyMeasurable_condExp (m := ℱ n)).measurable measurableSet_Ioi + have hB_meas0 : ∀ n, MeasurableSet (B n) := fun n ↦ ℱ.le n _ (hB_meas n) + have hB_lim : ∀ᵐ x ∂ρ, ∀ᶠ n in atTop, x ∈ B n ↔ x ∈ A := by + filter_upwards [hlim] with x hx + by_cases hxA : x ∈ A + · have h1 : Tendsto (fun n ↦ (ρ[φ | ℱ n]) x) atTop (𝓝 1) := by simpa [φ, hxA] using hx + filter_upwards [h1.eventually (lt_mem_nhds (by norm_num : (2⁻¹ : ℝ) < 1))] with n hn + simp [hB, hn, hxA] + · have h0 : Tendsto (fun n ↦ (ρ[φ | ℱ n]) x) atTop (𝓝 0) := by simpa [φ, hxA] using hx + filter_upwards [h0.eventually (gt_mem_nhds (by norm_num : (0 : ℝ) < 2⁻¹))] with n hn + simp [hB, hxA, hn.le] + have hμB : Tendsto (fun n ↦ μ (B n)) atTop (𝓝 (μ A)) := + tendsto_measure_of_ae_tendsto_indicator atTop hA hB_meas0 MeasurableSet.univ + (measure_ne_top μ _) (Eventually.of_forall fun _ ↦ Set.subset_univ _) + (ae_add_measure_iff.1 hB_lim).1 + have hνB : Tendsto (fun n ↦ ν (B n)) atTop (𝓝 0) := by + rw [← hνA] + exact tendsto_measure_of_ae_tendsto_indicator atTop hA hB_meas0 MeasurableSet.univ + (measure_ne_top ν _) (Eventually.of_forall fun _ ↦ Set.subset_univ _) + (ae_add_measure_iff.1 hB_lim).2 + -- lower bounds on the divergences of the images + have hBC : ∀ n, ∃ C : Set (β n), MeasurableSet C ∧ g n ⁻¹' C = B n := fun n ↦ + MeasurableSpace.measurableSet_comap.1 (hB_meas n) + have hkey n : ENNReal.ofReal (μ.real (B n) * Real.log (μ.real (B n) / ν.real (B n)) + + ν.real (B n) - μ.real (B n)) ≤ klDiv (μ.map (g n)) (ν.map (g n)) := by + obtain ⟨C, hC, hCB⟩ := hBC n + refine le_trans ?_ (klDiv_restrict_le hC) + have := mul_log_le_klDiv ((μ.map (g n)).restrict C) ((ν.map (g n)).restrict C) + simpa only [measureReal_def, Measure.restrict_apply_univ, Measure.map_apply (hg n) hC, hCB] + using this + have hkey' n (h0 : ν (B n) = 0) (hpos : μ (B n) ≠ 0) : + klDiv (μ.map (g n)) (ν.map (g n)) = ⊤ := by + obtain ⟨C, hC, hCB⟩ := hBC n + refine klDiv_of_not_ac fun hac ↦ hpos ?_ + have := hac (show (ν.map (g n)) C = 0 by rw [Measure.map_apply (hg n) hC, hCB, h0]) + rwa [Measure.map_apply (hg n) hC, hCB] at this + -- conclusion + intro K hK + set a : ℝ := μ.real A with ha_def + have ha : 0 < a := ENNReal.toReal_pos hμA (measure_ne_top μ A) + set δ : ℝ := Real.exp (-(2 / a) * (K.toReal + 2)) with hδ_def + have hδ : 0 < δ := Real.exp_pos _ + have hδ1 : δ ≤ 1 := by + rw [hδ_def, Real.exp_le_one_iff] + have : 0 ≤ K.toReal + 2 := by positivity + nlinarith [div_pos (by norm_num : (0 : ℝ) < 2) ha] + have hμB' : Tendsto (fun n ↦ μ.real (B n)) atTop (𝓝 a) := + (ENNReal.tendsto_toReal (measure_ne_top μ A)).comp hμB + have hνB' : Tendsto (fun n ↦ ν.real (B n)) atTop (𝓝 0) := by + have := (ENNReal.tendsto_toReal ENNReal.zero_ne_top).comp hνB + simpa [Function.comp_def, measureReal_def] using this + have h_ev : ∀ᶠ n in atTop, a / 2 ≤ μ.real (B n) ∧ ν.real (B n) ≤ δ := + (hμB'.eventually (le_mem_nhds (by linarith))).and (hνB'.eventually (ge_mem_nhds hδ)) + obtain ⟨n, hna, hnδ⟩ := h_ev.exists + refine ⟨n, ?_⟩ + by_cases h0 : ν (B n) = 0 + · rw [hkey' n h0 (fun h ↦ by simp [measureReal_def, h] at hna; linarith)] + exact hK + · have hνpos : 0 < ν.real (B n) := ENNReal.toReal_pos h0 (measure_ne_top ν _) + have hμpos : 0 < μ.real (B n) := by linarith + refine lt_of_lt_of_le ?_ (hkey n) + rw [← ENNReal.ofReal_toReal hK.ne, ENNReal.ofReal_lt_ofReal_iff_of_nonneg + ENNReal.toReal_nonneg] + -- `x log (x / y) + y - x ≥ -1 - x log y ≥ -1 + (a / 2) (-log δ) = K.toReal + 1` + have h1 : μ.real (B n) - 1 ≤ μ.real (B n) * Real.log (μ.real (B n)) := by + have := Real.one_sub_inv_le_log_of_pos hμpos + calc μ.real (B n) - 1 = μ.real (B n) * (1 - (μ.real (B n))⁻¹) := by + field_simp + _ ≤ _ := by gcongr + have h2 : Real.log (ν.real (B n)) ≤ Real.log δ := Real.log_le_log hνpos hnδ + have h3 : Real.log δ = -(2 / a) * (K.toReal + 2) := by rw [hδ_def, Real.log_exp] + have h4 : (a / 2) * (-Real.log δ) = K.toReal + 2 := by + rw [h3] + field_simp + have h5 : (a / 2) * (-Real.log (ν.real (B n))) ≤ + μ.real (B n) * (-Real.log (ν.real (B n))) := by + have hlog : 0 ≤ -Real.log (ν.real (B n)) := by + have := Real.log_nonpos hνpos.le (hnδ.trans hδ1) + linarith + exact mul_le_mul_of_nonneg_right hna hlog + rw [Real.log_div hμpos.ne' hνpos.ne'] + nlinarith [mul_le_mul_of_nonneg_left (neg_le_neg h2) (by positivity : (0 : ℝ) ≤ a / 2), + measureReal_nonneg (μ := ν) (s := B n)] + +end InformationTheory + +namespace MeasurableSpace + +/-- The σ-algebras of the projections of `ℕ → E` on the first `n` coordinates generate the +product σ-algebra. -/ +lemma iSup_comap_restrictFin {E : Type*} [mE : MeasurableSpace E] : + ⨆ n : ℕ, MeasurableSpace.comap (fun f : ℕ → E ↦ fun i : Fin n ↦ f i) + MeasurableSpace.pi = MeasurableSpace.pi := by + refine le_antisymm + (iSup_le fun n ↦ (measurable_pi_lambda _ fun _ ↦ measurable_pi_apply _).comap_le) + (iSup_le fun i ↦ le_iSup_of_le (i + 1) ?_) + have : (fun f : ℕ → E ↦ f i) = + (fun h : Fin (i + 1) → E ↦ h ⟨i, i.lt_succ_self⟩) ∘ + fun f : ℕ → E ↦ fun j : Fin (i + 1) ↦ f j := rfl + rw [this, ← MeasurableSpace.comap_comp] + exact MeasurableSpace.comap_mono (measurable_pi_apply _).comap_le + +end MeasurableSpace diff --git a/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/Restrict.lean b/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/Restrict.lean new file mode 100644 index 00000000..35abb6a2 --- /dev/null +++ b/LeanMachineLearning/ForMathlib/InformationTheory/KullbackLeibler/Restrict.lean @@ -0,0 +1,142 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import Mathlib.InformationTheory.KullbackLeibler.Basic + +/-! +# The Kullback–Leibler divergence of restrictions of measures + +The Kullback–Leibler divergence of two finite measures is the sum of the divergences of their +restrictions to a measurable set and to its complement, so that restricting both measures to +a measurable set does not increase the divergence, and the divergence is the supremum of the +divergences of the restrictions to an increasing sequence of measurable sets covering the space. + +## Main statements + +* `klDiv_restrict_add_restrict_compl`: the divergence is the sum of the divergences of + the restrictions to a measurable set and to its complement. +* `klDiv_restrict_le`: restricting both measures to a measurable set does not increase + the divergence. +* `klDiv_eq_iSup_restrict`: the divergence is the supremum of the divergences of the restrictions + to an increasing sequence of measurable sets covering the space. + +-/ + +@[expose] public section + +open MeasureTheory +open scoped ENNReal + +namespace MeasureTheory.Measure + +variable {α : Type*} {mα : MeasurableSpace α} {μ ν : Measure α} + +/-- The Radon–Nikodym derivative of `μ.restrict s` with respect to `ν.restrict s` is the +Radon–Nikodym derivative of `μ` with respect to `ν`, `ν.restrict s`-almost everywhere. -/ +lemma rnDeriv_restrict_restrict (μ ν : Measure α) [μ.HaveLebesgueDecomposition ν] [SigmaFinite ν] + {s : Set α} (hs : MeasurableSet s) : + (μ.restrict s).rnDeriv (ν.restrict s) =ᵐ[ν.restrict s] μ.rnDeriv ν := by + refine (eq_rnDeriv (s := (μ.singularPart ν).restrict s) (measurable_rnDeriv μ ν) + (((mutuallySingular_singularPart μ ν).restrict s).mono le_rfl restrict_le_self) ?_).symm + rw [← restrict_withDensity hs, ← restrict_add, ← haveLebesgueDecomposition_add μ ν] + +end MeasureTheory.Measure + +namespace InformationTheory + +variable {α : Type*} {mα : MeasurableSpace α} {μ ν : Measure α} [IsFiniteMeasure μ] + [IsFiniteMeasure ν] {s : Set α} + +/-- If the restrictions of `μ` and `ν` to a measurable set `s` satisfy +`μ.restrict s ≪ ν.restrict s`, the divergence of those restrictions is the integral of +`klFun (∂μ/∂ν)` over `s`. -/ +lemma klDiv_restrict_of_ac (hμν : μ.restrict s ≪ ν.restrict s) (hs : MeasurableSet s) : + klDiv (μ.restrict s) (ν.restrict s) = + ∫⁻ x in s, ENNReal.ofReal (klFun (μ.rnDeriv ν x).toReal) ∂ν := by + rw [klDiv_eq_lintegral_klFun_of_ac hμν] + refine lintegral_congr_ae ?_ + filter_upwards [Measure.rnDeriv_restrict_restrict μ ν hs] with x hx + rw [hx] + +/-- The divergence is the sum of the divergences of the restrictions to a measurable set and to +its complement. -/ +lemma klDiv_restrict_add_restrict_compl (hs : MeasurableSet s) : + klDiv (μ.restrict s) (ν.restrict s) + klDiv (μ.restrict sᶜ) (ν.restrict sᶜ) = klDiv μ ν := by + by_cases hμν : μ ≪ ν + · rw [klDiv_restrict_of_ac (hμν.restrict s) hs, klDiv_restrict_of_ac (hμν.restrict sᶜ) hs.compl, + klDiv_eq_lintegral_klFun_of_ac hμν, lintegral_add_compl _ hs] + · rw [klDiv_of_not_ac hμν, ENNReal.add_eq_top] + by_contra! h + refine hμν ?_ + rw [← Measure.restrict_add_restrict_compl (μ := μ) hs, + ← Measure.restrict_add_restrict_compl (μ := ν) hs] + exact (klDiv_ne_top_iff.1 h.1).1.add (klDiv_ne_top_iff.1 h.2).1 + +/-- Restricting both measures to a measurable set does not increase the divergence. -/ +lemma klDiv_restrict_le (hs : MeasurableSet s) : + klDiv (μ.restrict s) (ν.restrict s) ≤ klDiv μ ν := by + rw [← klDiv_restrict_add_restrict_compl (μ := μ) (ν := ν) hs] + exact le_add_right le_rfl + +/-- The divergence of two finite measures supported on a measurable set `s` and its complement +respectively is the sum of the divergences of the two parts. -/ +lemma klDiv_add_add_of_measure_eq_zero {μ' ν' : Measure α} [IsFiniteMeasure μ'] + [IsFiniteMeasure ν'] (hs : MeasurableSet s) (hμ : μ sᶜ = 0) (hν : ν sᶜ = 0) (hμ' : μ' s = 0) + (hν' : ν' s = 0) : + klDiv (μ + μ') (ν + ν') = klDiv μ ν + klDiv μ' ν' := by + have h1 : ∀ (ρ ρ' : Measure α), ρ sᶜ = 0 → ρ' s = 0 → (ρ + ρ').restrict s = ρ := by + intro ρ ρ' hρ hρ' + rw [Measure.restrict_add, Measure.restrict_eq_zero.2 hρ', add_zero, + Measure.restrict_eq_self_of_ae_mem] + exact ae_iff.2 hρ + have h2 : ∀ (ρ ρ' : Measure α), ρ sᶜ = 0 → ρ' s = 0 → (ρ + ρ').restrict sᶜ = ρ' := by + intro ρ ρ' hρ hρ' + rw [Measure.restrict_add, Measure.restrict_eq_zero.2 hρ, zero_add, + Measure.restrict_eq_self_of_ae_mem] + exact ae_iff.2 (by rw [show {a | a ∉ sᶜ} = s from compl_compl s]; exact hρ') + rw [← klDiv_restrict_add_restrict_compl hs, h1 μ μ' hμ hμ', h1 ν ν' hν hν', h2 μ μ' hμ hμ', + h2 ν ν' hν hν'] + +/-- **Monotone convergence** for the Kullback–Leibler divergence: the divergence is the supremum +of the divergences of the restrictions to an increasing sequence of measurable sets covering the +space. -/ +lemma klDiv_eq_iSup_restrict {s : ℕ → Set α} (hs : ∀ n, MeasurableSet (s n)) + (h_mono : Monotone s) (h_univ : ⋃ n, s n = Set.univ) : + klDiv μ ν = ⨆ n, klDiv (μ.restrict (s n)) (ν.restrict (s n)) := by + by_cases hμν : μ ≪ ν + · simp_rw [klDiv_restrict_of_ac (hμν.restrict _) (hs _), klDiv_eq_lintegral_klFun_of_ac hμν, + ← lintegral_indicator (hs _)] + have h_meas : Measurable fun x ↦ ENNReal.ofReal (klFun (μ.rnDeriv ν x).toReal) := + (measurable_klFun.comp (Measure.measurable_rnDeriv μ ν).ennreal_toReal).ennreal_ofReal + rw [← lintegral_iSup (fun n ↦ h_meas.indicator (hs n)) fun n m hnm ↦ + Set.indicator_le_indicator_of_subset (h_mono hnm) fun _ ↦ zero_le] + refine lintegral_congr fun x ↦ ?_ + obtain ⟨n, hn⟩ : ∃ n, x ∈ s n := by + have : x ∈ ⋃ n, s n := h_univ ▸ Set.mem_univ x + simpa using this + refine le_antisymm (le_iSup_of_le n ?_) (iSup_le fun m ↦ Set.indicator_le_self _ _ _) + simp [hn] + · rw [klDiv_of_not_ac hμν, eq_comm, iSup_eq_top] + obtain ⟨t, ht0, htpos⟩ : ∃ t, ν t = 0 ∧ μ t ≠ 0 := by + by_contra! h + exact hμν (Measure.AbsolutelyContinuous.mk fun t _ ht ↦ h t ht) + have h_iUnion : μ t = ⨆ n, μ (t ∩ s n) := by + rw [← Monotone.measure_iUnion (fun n m hnm ↦ Set.inter_subset_inter_right _ (h_mono hnm)), + ← Set.inter_iUnion, h_univ, Set.inter_univ] + obtain ⟨n, hn⟩ : ∃ n, μ (t ∩ s n) ≠ 0 := by + by_contra! h + simp [h_iUnion, h] at htpos + refine fun b hb ↦ ⟨n, lt_of_lt_of_le hb (le_of_eq ?_)⟩ + rw [eq_comm, klDiv_of_not_ac] + intro h + refine hn ?_ + have := h (show ν.restrict (s n) t = 0 by + rw [Measure.restrict_apply' (hs n)] + exact measure_mono_null Set.inter_subset_left ht0) + rwa [Measure.restrict_apply' (hs n)] at this + +end InformationTheory diff --git a/LeanMachineLearning/SequentialLearning/DivergenceDecomposition.lean b/LeanMachineLearning/SequentialLearning/DivergenceDecomposition.lean new file mode 100644 index 00000000..4cb2be59 --- /dev/null +++ b/LeanMachineLearning/SequentialLearning/DivergenceDecomposition.lean @@ -0,0 +1,185 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.ChainRule +public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.CompProd +public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.MapSequence +public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.Restrict +public import LeanMachineLearning.SequentialLearning.StationaryEnv + +/-! +# The divergence decomposition + +Let `alg`, `alg'` be algorithms, `env`, `env'` be environments, and consider two +algorithm-environment sequences `(X, Y)` and `(X', Y')` of `alg` against `env` and of `alg'` +against `env'`, on arbitrary probability spaces `(Ω, P)` and `(Ω', P')`. The Kullback-Leibler +divergence between the laws of the histories of the first `M` rounds is the sum, over the rounds +`t < M`, of the conditional divergences of the step at round `t` given the first `t` rounds. +Note that both arguments of the conditional term use the law of the *first* history, +so that term measures only how the two step kernels differ. +The same identity holds for the whole trajectory `trajectory X Y : Ω → (ℕ → 𝓐 × 𝓨)`, with a series +in place of the finite sum. + +For a single algorithm run against two stationary environments with reward kernels `κ` and +`κ'`, the two step kernels share the policy and differ only in the reward kernel, so the +conditional divergence of a step is the conditional divergence of the reward given the played +action. +This is the *divergence decomposition* of bandit lower bounds. + +## Main statements + +* `IsAlgEnvSeq.klDiv_map_history_stepKernel`, `IsAlgEnvSeq.klDiv_map_trajectory_stepKernel`: + the chain rule for the law of the history of the first `M` rounds and for the law of the + trajectory. +* `IsAlgEnvSeq.klDiv_map_history_compProd`, `IsAlgEnvSeq.klDiv_map_history`: the divergence + decomposition for two stationary environments, in composition-product and in integral form. +* `IsAlgEnvSeq.klDiv_map_trajectory_compProd`, `IsAlgEnvSeq.klDiv_map_trajectory`: the same two + forms for the trajectory. + +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory InformationTheory Finset +open scoped ENNReal RealInnerProductSpace ENat + +namespace Learning + +variable {𝓐 𝓨 : Type*} {m𝓐 : MeasurableSpace 𝓐} {m𝓨 : MeasurableSpace 𝓨} + {Ω Ω' : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} + {P : Measure Ω} {P' : Measure Ω'} [IsProbabilityMeasure P] [IsProbabilityMeasure P'] + {X : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} {X' : ℕ → Ω' → 𝓐} {Y' : ℕ → Ω' → 𝓨} + {alg alg' : Algorithm 𝓐 𝓨} {env env' : Environment 𝓐 𝓨} + {κ κ' : Kernel 𝓐 𝓨} [IsMarkovKernel κ] [IsMarkovKernel κ'] + +section + +variable {α β γ : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} + {mγ : MeasurableSpace γ} {μ : Measure α} [IsFiniteMeasure μ] + +/-- The divergence of one step of a policy/reward decomposition, in composition-product form: +the policy `π` is shared and the reward kernels `κ`, `η` (which ignore the history) differ, so +the divergence is the conditional divergence of the reward kernels given the played action, +whose law is `π ∘ₘ μ`. -/ +lemma klDiv_compProd_compProd_prodMkLeft_eq_klDiv_comp_compProd (μ : Measure α) + [IsFiniteMeasure μ] (π : Kernel α β) [IsMarkovKernel π] (κ η : Kernel β γ) [IsFiniteKernel κ] + [IsFiniteKernel η] : + klDiv (μ ⊗ₘ (π ⊗ₖ κ.prodMkLeft α)) (μ ⊗ₘ (π ⊗ₖ η.prodMkLeft α)) = + klDiv ((π ∘ₘ μ) ⊗ₘ κ) ((π ∘ₘ μ) ⊗ₘ η) := by + rw [← klDiv_map_measurableEquiv _ _ MeasurableEquiv.prodAssoc.symm, Measure.compProd_assoc, + Measure.compProd_assoc, ← Measure.snd_compProd, Measure.snd] + exact klDiv_compProd_comap _ _ _ measurable_snd + +end + +/-- **Chain rule for histories.** For two algorithms `alg`, `alg'` run against two environments +`env`, `env'`, the divergence between the laws of the histories of the first `M` rounds is +the sum over the rounds `t < M` of the conditional divergences of the step at round `t` given +the first `t` rounds. -/ +lemma IsAlgEnvSeq.klDiv_map_history_stepKernel (h : IsAlgEnvSeq X Y alg env P) + (h' : IsAlgEnvSeq X' Y' alg' env' P') (M : ℕ) : + klDiv (P.map (history X Y M)) (P'.map (history X' Y' M)) = + ∑ t ∈ range M, + klDiv (P.map (history X Y t) ⊗ₘ stepKernel alg env t) + (P.map (history X Y t) ⊗ₘ stepKernel alg' env' t) := by + have hX := h.measurable_action + have hY := h.measurable_feedback + have hX' := h'.measurable_action + have hY' := h'.measurable_feedback + induction M with + | zero => simp + | succ M ih => + rw [history_succ, history_succ, ← Measure.map_map (by fun_prop) (by fun_prop), + ← Measure.map_map (by fun_prop) (by fun_prop), klDiv_map_measurableEquiv, + (h.hasCondDistrib_step M).map_eq, (h'.hasCondDistrib_step M).map_eq, + klDiv_compProd_eq_add, ih, sum_range_succ] + +/-- The divergence between the laws of two trajectories is the supremum over `n` of the divergences +between the laws of the histories up to time `n`. -/ +lemma klDiv_map_trajectory_eq_iSup (hX : ∀ n, Measurable (X n)) (hY : ∀ n, Measurable (Y n)) + (hX' : ∀ n, Measurable (X' n)) (hY' : ∀ n, Measurable (Y' n)) : + klDiv (P.map (trajectory X Y)) (P'.map (trajectory X' Y')) = + ⨆ n, klDiv (P.map (history X Y n)) (P'.map (history X' Y' n)) := by + have hg : ∀ n, Measurable fun f : ℕ → 𝓐 × 𝓨 ↦ fun i : Fin n ↦ f i.1 := fun n ↦ + measurable_pi_lambda _ fun i ↦ measurable_pi_apply i.1 + rw [klDiv_eq_iSup_map hg ?_ MeasurableSpace.iSup_comap_restrictFin] + · refine iSup_congr fun n ↦ ?_ + rw [Measure.map_map (hg n) (measurable_trajectory hX hY), + Measure.map_map (hg n) (measurable_trajectory hX' hY')] + rfl + · intro n m hnm + have : (fun f : ℕ → 𝓐 × 𝓨 ↦ fun i : Fin n ↦ f i.1) = + (fun h : Fin m → 𝓐 × 𝓨 ↦ fun i : Fin n ↦ h (Fin.castLE hnm i)) ∘ + fun f : ℕ → 𝓐 × 𝓨 ↦ fun i : Fin m ↦ f i.1 := rfl + beta_reduce + rw [this, ← MeasurableSpace.comap_comp] + exact MeasurableSpace.comap_mono (measurable_pi_lambda _ fun i ↦ + measurable_pi_apply (Fin.castLE hnm i)).comap_le + +/-- **Chain rule for trajectories.** For two algorithms `alg` and `alg'` run against +two environments `env` and `env'`, the divergence between the laws of the trajectories is +the series over the rounds `t` of the conditional divergences of the step at round `t` given +the first `t` rounds. -/ +lemma IsAlgEnvSeq.klDiv_map_trajectory_stepKernel (h : IsAlgEnvSeq X Y alg env P) + (h' : IsAlgEnvSeq X' Y' alg' env' P') : + klDiv (P.map (trajectory X Y)) (P'.map (trajectory X' Y')) = + ∑' t : ℕ, klDiv (P.map (history X Y t) ⊗ₘ stepKernel alg env t) + (P.map (history X Y t) ⊗ₘ stepKernel alg' env' t) := by + have hX := h.measurable_action + have hY := h.measurable_feedback + have hX' := h'.measurable_action + have hY' := h'.measurable_feedback + rw [klDiv_map_trajectory_eq_iSup hX hY hX' hY', ENNReal.tsum_eq_iSup_nat] + exact iSup_congr fun n ↦ h.klDiv_map_history_stepKernel h' n + +section StationaryEnv + +/-- Chain rule for histories of a single algorithm versus two stationary environments. -/ +lemma IsAlgEnvSeq.klDiv_map_history_compProd (h : IsAlgEnvSeq X Y alg (stationaryEnv κ) P) + (h' : IsAlgEnvSeq X' Y' alg (stationaryEnv κ') P') (M : ℕ) : + klDiv (P.map (history X Y M)) (P'.map (history X' Y' M)) = + ∑ t ∈ range M, klDiv (P.map (X t) ⊗ₘ κ) (P.map (X t) ⊗ₘ κ') := by + rw [h.klDiv_map_history_stepKernel h'] + refine sum_congr rfl fun t _ ↦ ?_ + rw [stepKernel_stationaryEnv, stepKernel_stationaryEnv, + klDiv_compProd_compProd_prodMkLeft_eq_klDiv_comp_compProd, + ← (h.hasCondDistrib_action t).hasLaw_comp.map_eq] + +/-- Chain rule for histories of a single algorithm versus two stationary environments. -/ +lemma IsAlgEnvSeq.klDiv_map_history [MeasurableSpace.CountablyGenerated 𝓨] + (h : IsAlgEnvSeq X Y alg (stationaryEnv κ) P) + (h' : IsAlgEnvSeq X' Y' alg (stationaryEnv κ') P') (M : ℕ) : + klDiv (P.map (history X Y M)) (P'.map (history X' Y' M)) = + ∑ t ∈ range M, ∫⁻ ω, klDiv (κ (X t ω)) (κ' (X t ω)) ∂P := by + rw [h.klDiv_map_history_compProd h'] + refine sum_congr rfl fun t _ ↦ ?_ + rw [klDiv_compProd_right_eq_lintegral, + lintegral_map (measurable_klDiv_kernel κ κ') (h.measurable_action t)] + +/-- Chain rule for trajectories of a single algorithm versus two stationary environments. -/ +lemma IsAlgEnvSeq.klDiv_map_trajectory_compProd (h : IsAlgEnvSeq X Y alg (stationaryEnv κ) P) + (h' : IsAlgEnvSeq X' Y' alg (stationaryEnv κ') P') : + klDiv (P.map (trajectory X Y)) (P'.map (trajectory X' Y')) = + ∑' t : ℕ, klDiv (P.map (X t) ⊗ₘ κ) (P.map (X t) ⊗ₘ κ') := by + rw [klDiv_map_trajectory_eq_iSup h.measurable_action h.measurable_feedback + h'.measurable_action h'.measurable_feedback, ENNReal.tsum_eq_iSup_nat] + exact iSup_congr fun n ↦ h.klDiv_map_history_compProd h' n + +/-- Chain rule for trajectories of a single algorithm versus two stationary environments. -/ +lemma IsAlgEnvSeq.klDiv_map_trajectory [MeasurableSpace.CountablyGenerated 𝓨] + (h : IsAlgEnvSeq X Y alg (stationaryEnv κ) P) + (h' : IsAlgEnvSeq X' Y' alg (stationaryEnv κ') P') : + klDiv (P.map (trajectory X Y)) (P'.map (trajectory X' Y')) = + ∑' t : ℕ, ∫⁻ ω, klDiv (κ (X t ω)) (κ' (X t ω)) ∂P := by + rw [h.klDiv_map_trajectory_compProd h'] + refine tsum_congr fun t ↦ ?_ + rw [klDiv_compProd_right_eq_lintegral, + lintegral_map (measurable_klDiv_kernel κ κ') (h.measurable_action t)] + +end StationaryEnv + +end Learning diff --git a/LeanMachineLearning/SequentialLearning/StationaryEnv.lean b/LeanMachineLearning/SequentialLearning/StationaryEnv.lean index 0ac76081..f4108acd 100644 --- a/LeanMachineLearning/SequentialLearning/StationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/StationaryEnv.lean @@ -183,6 +183,10 @@ def stationaryEnv (ν : Kernel 𝓐 𝓨) [IsMarkovKernel ν] : Environment 𝓐 lemma feedback_stationaryEnv (ν : Kernel 𝓐 𝓨) [IsMarkovKernel ν] (n : ℕ) : (stationaryEnv ν).feedback n = ν.prodMkLeft _ := rfl +lemma stepKernel_stationaryEnv (alg : Algorithm 𝓐 𝓨) (η : Kernel 𝓐 𝓨) [IsMarkovKernel η] (n : ℕ) : + stepKernel alg (stationaryEnv η) n = alg.policy n ⊗ₖ η.prodMkLeft _ := by + rw [stepKernel, feedback_stationaryEnv] + @[simp] lemma ν0_stationaryEnv (ν : Kernel 𝓐 𝓨) [IsMarkovKernel ν] : (stationaryEnv ν).ν0 = ν := ν0_obliviousEnv _