diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index 3d683931..7780f5c6 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -1,5 +1,7 @@ module -- shake: keep-all --deprecated_module: ignore +public import LeanMachineLearning.ForMathlib.ConvexAnalysis.Bregman.Basic +public import LeanMachineLearning.ForMathlib.ConvexAnalysis.Subgradient.Basic public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.ChainRule public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.CompProd public import LeanMachineLearning.ForMathlib.InformationTheory.KullbackLeibler.Convex diff --git a/LeanMachineLearning/ForMathlib/ConvexAnalysis/Bregman/Basic.lean b/LeanMachineLearning/ForMathlib/ConvexAnalysis/Bregman/Basic.lean new file mode 100644 index 00000000..d34f4e32 --- /dev/null +++ b/LeanMachineLearning/ForMathlib/ConvexAnalysis/Bregman/Basic.lean @@ -0,0 +1,206 @@ +/- +Copyright (c) 2026 Isidoor Pinillo Esquivel. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Isidoor Pinillo Esquivel +-/ +module + +public import Mathlib.Analysis.Convex.Function +public import Mathlib.Tactic + +/-! +# Bregman Divergences + +Generalized vector-valued Bregman divergences `bregDiv f x y J` (notation: `D_[f](x, y, J)` +in `scoped Bregman`) measure the error of the linear approximation of `f` around `y`. +They satisfy algebraic properties linked to derivatives, including chain rules, +product rules, affine invariance, and convexity preservation. + +## Main definitions + +* `Analysis.Convex.bregDiv f x y J`: The generalized vector-valued Bregman divergence + `f x - f y - J (x - y)` for a function `f : E → F` and + an additive map `J : E →+ F` (e.g. a derivative, gradient, or subgradient). + +## Notation + +* `D_[f](x, y, J)`: Scoped notation in `Bregman` for `bregDiv f x y J`. +-/ + +@[expose] public section + +namespace Analysis.Convex + +variable {E F G : Type*} [AddCommGroup E] [AddCommGroup F] [AddCommGroup G] + +/-- The generalized vector-valued Bregman divergence. +`f` maps `E → F`. `J` is the Jacobian / subgradient additive mapping `E →+ F`. -/ +def bregDiv (f : E → F) (x y : E) (J : E →+ F) : F := + f x - f y - J (x - y) + +/-- Scoped notation for `bregDiv`. -/ +scoped[Bregman] notation "D_[" f "](" x ", " y ", " J ")" => Analysis.Convex.bregDiv f x y J + +open scoped Bregman + +/-- Unexpander for `bregDiv`. -/ +@[app_unexpander bregDiv] +meta def unexpandBregDiv : Lean.PrettyPrinter.Unexpander + | `($_ $f $x $y $J) => `(D_[$f]($x, $y, $J)) + | _ => throw () + +@[simp] +lemma bregDiv_self (f : E → F) (x : E) (J_x : E →+ F) : + D_[f](x, x, J_x) = 0 := by + simp [bregDiv] + +lemma bregDiv_three_point (f : E → F) (x y z : E) (J_x J_y : E →+ F) : + D_[f](z, x, J_x) + D_[f](x, y, J_y) - D_[f](z, y, J_y) = (J_y - J_x) (z - x) := by + simp only [bregDiv, map_sub, AddMonoidHom.sub_apply] + abel + +lemma bregDiv_add_swap (f : E → F) (x y : E) (J_x J_y : E →+ F) : + D_[f](y, x, J_x) + D_[f](x, y, J_y) = (J_x - J_y) (x - y) := by + simp only [bregDiv, map_sub, AddMonoidHom.sub_apply] + abel + +lemma bregDiv_fun_bregDiv (f : E → F) (x y z : E) (J_x J_y : E →+ F) : + D_[fun w ↦ D_[f](w, y, J_y)](z, x, J_x - J_y) = D_[f](z, x, J_x) := by + simp only [bregDiv, map_sub, AddMonoidHom.sub_apply] + abel + +lemma bregDiv_add (f₁ f₂ : E → F) (x y : E) (J₁ J₂ : E →+ F) : + D_[f₁ + f₂](x, y, J₁ + J₂) = D_[f₁](x, y, J₁) + D_[f₂](x, y, J₂) := by + simp only [bregDiv, Pi.add_apply, AddMonoidHom.add_apply, map_sub] + abel + +lemma bregDiv_prod {E₁ E₂ : Type*} [AddCommGroup E₁] [AddCommGroup E₂] + (f₁ : E₁ → F) (f₂ : E₂ → F) (x y : E₁ × E₂) (J₁ : E₁ →+ F) (J₂ : E₂ →+ F) : + D_[fun p : E₁ × E₂ ↦ f₁ p.1 + f₂ p.2](x, y, J₁.coprod J₂) = + D_[f₁](x.1, y.1, J₁) + D_[f₂](x.2, y.2, J₂) := by + simp only [bregDiv, AddMonoidHom.coprod_apply, map_sub] + abel + +@[simp] +lemma bregDiv_const (c : F) (x y : E) : + D_[fun _ ↦ c](x, y, 0) = 0 := by + simp [bregDiv] + +@[simp] +lemma bregDiv_linear (h : E →+ F) (x y : E) : + D_[h](x, y, h) = 0 := by + simp [bregDiv] + +@[simp] +lemma bregDiv_add_const (f : E → F) (c : F) (x y : E) (J : E →+ F) : + D_[fun z ↦ f z + c](x, y, J) = D_[f](x, y, J) := by + simp [bregDiv] + +@[simp] +lemma bregDiv_const_add (c : F) (f : E → F) (x y : E) (J : E →+ F) : + D_[fun z ↦ c + f z](x, y, J) = D_[f](x, y, J) := by + simp [bregDiv] + +@[simp] +lemma bregDiv_add_linear (f : E → F) (h : E →+ F) (x y : E) (J : E →+ F) : + D_[fun z ↦ f z + h z](x, y, J + h) = D_[f](x, y, J) := by + simp only [bregDiv, AddMonoidHom.add_apply, map_sub] + abel + +@[simp] +lemma bregDiv_neg (f : E → F) (x y : E) (J : E →+ F) : + D_[-f](x, y, -J) = - D_[f](x, y, J) := by + simp only [bregDiv, Pi.neg_apply, AddMonoidHom.neg_apply, map_sub] + abel + +lemma bregDiv_comp_neg (f : E → F) (x y : E) (J : E →+ F) : + D_[fun z ↦ f (-z)](x, y, -J) = D_[f](-x, -y, J) := by + simp [bregDiv, map_sub] + +lemma bregDiv_translate (f : E → F) (x₀ : E) (x y : E) (J : E →+ F) : + D_[fun z ↦ f (z + x₀)](x, y, J) = D_[f](x + x₀, y + x₀, J) := by + simp [bregDiv] + +section ChainRules + +lemma bregDiv_comp_affine {E₁ : Type*} [AddCommGroup E₁] + (f : E → F) (A : E₁ →+ E) (b : E) (x y : E₁) (J : E →+ F) : + D_[fun z ↦ f (A z + b)](x, y, J.comp A) = D_[f](A x + b, A y + b, J) := by + simp [bregDiv] + +lemma bregDiv_comp (f₂ : F → G) (f₁ : E → F) (x y : E) (J₂ : F →+ G) (J₁ : E →+ F) : + D_[f₂ ∘ f₁](x, y, J₂.comp J₁) = D_[f₂](f₁ x, f₁ y, J₂) + J₂ (D_[f₁](x, y, J₁)) := by + simp [bregDiv] + +end ChainRules + +section CommRing + +variable {R : Type*} [CommRing R] + +/-- First-order product rule for Bregman divergences. -/ +lemma bregDiv_mul (f₁ f₂ : E → R) (x y : E) (J₁ J₂ : E →+ R) : + D_[f₁ * f₂](x, y, f₂ y • J₁ + f₁ y • J₂) = + D_[f₁](x, y, J₁) * f₂ y + f₁ y * D_[f₂](x, y, J₂) + + (f₁ x - f₁ y) * (f₂ x - f₂ y) := by + dsimp only [bregDiv] + simp only [Pi.mul_apply, AddMonoidHom.add_apply, AddMonoidHom.smul_apply, map_sub, smul_eq_mul] + ring + + +end CommRing + +section Order + +variable {F' : Type*} [AddCommGroup F'] [Preorder F'] [AddRightMono F'] + +lemma bregDiv_le_of_le_of_eq {f₁ f₂ : E → F'} {x y : E} {J_y : E →+ F'} + (h_le : f₁ x ≤ f₂ x) (h_eq : f₁ y = f₂ y) : + D_[f₁](x, y, J_y) ≤ D_[f₂](x, y, J_y) := by + simp [bregDiv, h_eq, h_le] + +end Order + +section Module + +variable {R : Type*} [CommRing R] [Module R F] + +lemma bregDiv_smul (c : R) (f : E → F) (x y : E) (J : E →+ F) : + D_[c • f](x, y, c • J) = c • D_[f](x, y, J) := by + simp [bregDiv, smul_sub] + +lemma bregDiv_convexCombination (f : E → F) (x y : E) (J₁ J₂ : E →+ F) (w : R) : + D_[f](x, y, w • J₁ + (1 - w) • J₂) = w • D_[f](x, y, J₁) + (1 - w) • D_[f](x, y, J₂) := by + calc + D_[f](x, y, w • J₁ + (1 - w) • J₂) + = D_[(w + (1 - w)) • f](x, y, w • J₁ + (1 - w) • J₂) := by + rw [add_sub_cancel, one_smul] + _ = D_[w • f](x, y, w • J₁) + D_[(1 - w) • f](x, y, (1 - w) • J₂) := by + rw [add_smul, bregDiv_add] + _ = w • D_[f](x, y, J₁) + (1 - w) • D_[f](x, y, J₂) := by + simp only [bregDiv_smul] + +end Module + +section Convexity + +variable {𝕜 E' F' : Type*} [Semiring 𝕜] [PartialOrder 𝕜] +variable [AddCommGroup E'] [Module 𝕜 E'] +variable [AddCommGroup F'] [Module 𝕜 F'] [PartialOrder F'] [IsOrderedAddMonoid F'] + +/-- If `f` is convex on `s`, then `x ↦ D_[f](x, y, J)` is convex on `s` for any linear map `J`. -/ +lemma _root_.ConvexOn.bregDiv {s : Set E'} {f : E' → F'} (hf : ConvexOn 𝕜 s f) + (J : E' →ₗ[𝕜] F') (y : E') : + ConvexOn 𝕜 s (fun x ↦ D_[f](x, y, J)) := by + simp only [Analysis.Convex.bregDiv, sub_eq_add_neg, map_add, map_neg, neg_add_rev, neg_neg] + apply ConvexOn.add + · apply ConvexOn.add + · exact hf + · exact convexOn_const (-f y) hf.1 + · apply ConvexOn.add + · exact convexOn_const (J y) hf.1 + · exact (-J).convexOn hf.1 + +end Convexity + +end Analysis.Convex diff --git a/LeanMachineLearning/ForMathlib/ConvexAnalysis/Subgradient/Basic.lean b/LeanMachineLearning/ForMathlib/ConvexAnalysis/Subgradient/Basic.lean new file mode 100644 index 00000000..c8774fca --- /dev/null +++ b/LeanMachineLearning/ForMathlib/ConvexAnalysis/Subgradient/Basic.lean @@ -0,0 +1,218 @@ +/- +Copyright (c) 2026 Isidoor Pinillo Esquivel. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Isidoor Pinillo Esquivel +-/ +module + +public import LeanMachineLearning.ForMathlib.ConvexAnalysis.Bregman.Basic +public import Mathlib.Analysis.Convex.Function +public import Mathlib.Order.ConditionallyCompleteLattice.Basic +public import Mathlib.Tactic + +/-! +# Subgradients and Subdifferentials + +Subgradients of functions `f : E → F` defined +via the non-negativity of the Bregman divergence `0 ≤ D_[f](x, y, g_y)` on an explicit +domain `V` (i.e. the linearization error is non-negative on `V`). + +Carrying the domain `V : Set E` explicitly avoids using indicator functions +while supporting constrained convex analysis. + +## Main definitions + +* `Analysis.Convex.IsSubgradient V f y g_y`: `g_y : E →+ F` is a subgradient of `f` at `y` on `V`. +* `Analysis.Convex.subdifferential V f y`: The subdifferential set `∂[V, y] f`. + +## Main results + +* Chain rules: `Analysis.Convex.IsSubgradient.comp_affine` and `Analysis.Convex.IsSubgradient.comp`. +* Equivalence with classical inequality: `Analysis.Convex.mem_subdifferential_iff_le`. +* Fermat's rule: `Analysis.Convex.zero_mem_subdifferential_iff_isMinOn`. +* Suprema and maxima: `Analysis.Convex.IsSubgradient.max_left`, + `Analysis.Convex.IsSubgradient.max_right`, `Analysis.Convex.IsSubgradient.finset_sup`, + and `Analysis.Convex.IsSubgradient.ciSup`. + +## Notation + +* `∂[V, y] f`: Scoped notation in `Bregman` for `subdifferential V f y`. +-/ + +@[expose] public section + +namespace Analysis.Convex + +open scoped Bregman + +variable {E F : Type*} [AddCommGroup E] [AddCommGroup F] [Preorder F] + +/-- `g_y` is a subgradient of `f` at `y` on domain `V` (non-negativity of Bregman divergence). -/ +def IsSubgradient (V : Set E) (f : E → F) (y : E) (g_y : E →+ F) : Prop := + y ∈ V ∧ ∀ x ∈ V, 0 ≤ D_[f](x, y, g_y) + +/-- The subdifferential `∂[V, y] f` of `f` at `y` on `V`. -/ +def subdifferential (V : Set E) (f : E → F) (y : E) : Set (E →+ F) := + { g | IsSubgradient V f y g } + +/-- Scoped notation for `subdifferential`. -/ +scoped[Bregman] notation:60 "∂[" V ", " y "] " f:50 => Analysis.Convex.subdifferential V f y + +/-- Unexpander for `subdifferential`. -/ +@[app_unexpander subdifferential] +meta def unexpandSubdifferential : Lean.PrettyPrinter.Unexpander + | `($_ $V $f $y) => `(∂[$V, $y] $f) + | _ => throw () + +variable {V : Set E} {f : E → F} {y x : E} {g_y : E →+ F} + +@[local simp] +lemma mem_subdifferential_iff : + g_y ∈ ∂[V, y] f ↔ y ∈ V ∧ ∀ x ∈ V, 0 ≤ D_[f](x, y, g_y) := Iff.rfl + +@[simp] +lemma mem_subdifferential_const_iff {c : F} : + (0 : E →+ F) ∈ ∂[V, y] (fun _ ↦ c) ↔ y ∈ V := by simp + +@[simp] +lemma mem_subdifferential_linear_iff {h : E →+ F} : + h ∈ ∂[V, y] h ↔ y ∈ V := by simp + +@[simp] +lemma mem_subdifferential_add_const_iff {c : F} : + g_y ∈ ∂[V, y] (fun x ↦ f x + c) ↔ g_y ∈ ∂[V, y] f := by simp + +@[simp] +lemma mem_subdifferential_const_add_iff {c : F} : + g_y ∈ ∂[V, y] (fun x ↦ c + f x) ↔ g_y ∈ ∂[V, y] f := by simp + +@[simp] +lemma mem_subdifferential_add_linear_iff {h : E →+ F} : + (g_y + h) ∈ ∂[V, y] (fun x ↦ f x + h x) ↔ g_y ∈ ∂[V, y] f := by simp + +@[simp] +lemma mem_subdifferential_bregDiv_iff {g_x g_y : E →+ F} : + (g_x - g_y) ∈ ∂[V, x] (fun z ↦ D_[f](z, y, g_y)) ↔ g_x ∈ ∂[V, x] f := by + simp [bregDiv_fun_bregDiv] + +lemma zero_mem_subdifferential_iff : + (0 : E →+ F) ∈ ∂[V, y] f ↔ y ∈ V ∧ ∀ x ∈ V, 0 ≤ f x - f y := by simp [bregDiv] + +lemma IsSubgradient.comp_affine {E₁ : Type*} [AddCommGroup E₁] + {V₁ : Set E₁} {y₁ : E₁} {A : E₁ →+ E} {b : E} + (hy₁ : y₁ ∈ V₁) (h_map : ∀ x ∈ V₁, A x + b ∈ V) + (h_sub : g_y ∈ ∂[V, A y₁ + b] f) : + (g_y.comp A) ∈ ∂[V₁, y₁] (fun x ↦ f (A x + b)) := + ⟨hy₁, fun x hx ↦ by simp [bregDiv_comp_affine, h_sub.2 (A x + b) (h_map x hx)]⟩ + +/-- **Chain rule**: `g₂ ∘ g₁` is a subgradient of `f₂ ∘ f₁` when `g₂` is non-negative. -/ +lemma IsSubgradient.comp {G : Type*} [AddCommGroup G] [Preorder G] [IsOrderedAddMonoid G] + {f₂ : F → G} {f₁ : E → F} {g₂ : F →+ G} {g₁ : E →+ F} + (h₂ : g₂ ∈ ∂[f₁ '' V, f₁ y] f₂) (h₁ : g₁ ∈ ∂[V, y] f₁) + (hg₂_nonneg : ∀ z ≥ 0, 0 ≤ g₂ z) : + (g₂.comp g₁) ∈ ∂[V, y] (f₂ ∘ f₁) := + ⟨h₁.1, fun x hx ↦ by + simp only [bregDiv_comp, add_nonneg (h₂.2 (f₁ x) ⟨x, hx, rfl⟩) (hg₂_nonneg _ (h₁.2 x hx))]⟩ + +section OrderedGroup + +variable [IsOrderedAddMonoid F] + +/-- Equivalence with the classical definition. -/ +lemma mem_subdifferential_iff_le : + g_y ∈ ∂[V, y] f ↔ y ∈ V ∧ ∀ x ∈ V, f y + g_y (x - y) ≤ f x := by + simp [bregDiv, sub_sub, sub_nonneg] + +lemma IsSubgradient.monotonicity {g_x : E →+ F} + (hx_sub : g_x ∈ ∂[V, x] f) (hy_sub : g_y ∈ ∂[V, y] f) : + 0 ≤ (g_x - g_y) (x - y) := by + rw [← bregDiv_add_swap f x y g_x g_y] + exact add_nonneg (hx_sub.2 y hy_sub.1) (hy_sub.2 x hx_sub.1) + + +/-- **Fermat's rule**: `0` is a subgradient of `f` at `y` iff `y` is a minimizer of `f` on `V`. -/ +lemma zero_mem_subdifferential_iff_isMinOn : + (0 : E →+ F) ∈ ∂[V, y] f ↔ y ∈ V ∧ IsMinOn f V y := by + simp [bregDiv, sub_nonneg, isMinOn_iff] + +lemma IsSubgradient.add {f₁ f₂ : E → F} {g₁ g₂ : E →+ F} + (h₁ : g₁ ∈ ∂[V, y] f₁) (h₂ : g₂ ∈ ∂[V, y] f₂) : + (g₁ + g₂) ∈ ∂[V, y] (f₁ + f₂) := + ⟨h₁.1, fun x hx ↦ by simp [bregDiv_add, add_nonneg (h₁.2 x hx) (h₂.2 x hx)]⟩ + +lemma IsSubgradient.add_isMinOn {f₁ f₂ : E → F} {x : E} {g : E →+ F} + (h_min : IsMinOn f₁ V x) (hx_mem : x ∈ V) (h_sub : g ∈ ∂[V, x] f₂) : + g ∈ ∂[V, x] (f₁ + f₂) := by + simpa using IsSubgradient.add (zero_mem_subdifferential_iff_isMinOn.mpr ⟨hx_mem, h_min⟩) h_sub + +lemma IsSubgradient.of_le_of_eq {f₁ f₂ : E → F} {g_y : E →+ F} + (h_sub : g_y ∈ ∂[V, y] f₁) (h_le : ∀ x ∈ V, f₁ x ≤ f₂ x) (h_eq : f₁ y = f₂ y) : + g_y ∈ ∂[V, y] f₂ := + ⟨h_sub.1, fun x hx ↦ le_trans (h_sub.2 x hx) + (bregDiv_le_of_le_of_eq (h_le x hx) h_eq)⟩ + +end OrderedGroup + +section ModuleBasic + +variable {R : Type*} [CommRing R] [PartialOrder R] [Module R F] [PosSMulMono R F] + +lemma IsSubgradient.smul {c : R} (hc : 0 ≤ c) {f : E → F} {g_y : E →+ F} + (h_sub : g_y ∈ ∂[V, y] f) : + (c • g_y) ∈ ∂[V, y] (c • f) := + ⟨h_sub.1, fun x hx ↦ by simp [bregDiv_smul, smul_nonneg hc (h_sub.2 x hx)]⟩ + +end ModuleBasic + +section Module + +variable {R : Type*} [CommRing R] [PartialOrder R] [IsOrderedRing R] + [IsOrderedAddMonoid F] [Module R F] [PosSMulMono R F] + +lemma IsSubgradient.convexCombination {f : E → F} {g₁ g₂ : E →+ F} + (h₁ : g₁ ∈ ∂[V, y] f) (h₂ : g₂ ∈ ∂[V, y] f) {w : R} (hw : w ∈ Set.Icc (0 : R) 1) : + (w • g₁ + (1 - w) • g₂) ∈ ∂[V, y] f := + ⟨h₁.1, fun x hx ↦ by + simp [bregDiv_convexCombination, hw.1, sub_nonneg.mpr hw.2, + h₁.2 x hx, h₂.2 x hx, smul_nonneg, add_nonneg]⟩ + +end Module + +section LinearOrder + +variable {F_lin : Type*} [AddCommGroup F_lin] [LinearOrder F_lin] [IsOrderedAddMonoid F_lin] + +lemma IsSubgradient.max_left {f₁ f₂ : E → F_lin} {g_y : E →+ F_lin} + (h_sub : g_y ∈ ∂[V, y] f₁) (h_active : f₁ y = max (f₁ y) (f₂ y)) : + g_y ∈ ∂[V, y] (fun x ↦ max (f₁ x) (f₂ x)) := + IsSubgradient.of_le_of_eq h_sub (fun x _ ↦ le_max_left (f₁ x) (f₂ x)) h_active + +lemma IsSubgradient.max_right {f₁ f₂ : E → F_lin} {g_y : E →+ F_lin} + (h_sub : g_y ∈ ∂[V, y] f₂) (h_active : f₂ y = max (f₁ y) (f₂ y)) : + g_y ∈ ∂[V, y] (fun x ↦ max (f₁ x) (f₂ x)) := + IsSubgradient.of_le_of_eq h_sub (fun x _ ↦ le_max_right (f₁ x) (f₂ x)) h_active + +lemma IsSubgradient.finset_sup {ι : Type*} {s : Finset ι} + {f_i : ι → E → F_lin} {i : ι} {g_y : E →+ F_lin} + (his : i ∈ s) + (h_sub : g_y ∈ ∂[V, y] (f_i i)) (h_active : f_i i y = s.sup' ⟨i, his⟩ (fun j ↦ f_i j y)) : + g_y ∈ ∂[V, y] (fun x ↦ s.sup' ⟨i, his⟩ (fun j ↦ f_i j x)) := + IsSubgradient.of_le_of_eq h_sub (fun _ _ ↦ Finset.le_sup'_of_le _ his (le_refl _)) h_active + +end LinearOrder + +section Lattice + +variable {F_lat : Type*} [AddCommGroup F_lat] + [ConditionallyCompleteLattice F_lat] [IsOrderedAddMonoid F_lat] + +lemma IsSubgradient.ciSup {ι : Type*} {f_i : ι → E → F_lat} {i : ι} {g_y : E →+ F_lat} + (h_sub : g_y ∈ ∂[V, y] (f_i i)) + (h_active : f_i i y = ⨆ j, f_i j y) + (h_bdd : ∀ x ∈ V, BddAbove (Set.range (fun j ↦ f_i j x))) : + g_y ∈ ∂[V, y] (fun x ↦ ⨆ j, f_i j x) := + IsSubgradient.of_le_of_eq h_sub (fun x hx ↦ le_ciSup (h_bdd x hx) i) h_active + +end Lattice + +end Analysis.Convex