Skip to content
Browse chapters

Chapter 33 — Machine-Learning Algorithms

CLRS, fourth edition · Lean 4 formalization

The proofs below use the models and assumptions described in the scope and implementation notes.

Imports
import Mathlib.Analysis.InnerProductSpace.Basic import Mathlib.Analysis.InnerProductSpace.PiL2 import Mathlib.Tactic

33.1. Clustering

This section formalizes the k-means clustering problem and Lloyd's algorithm from CLRS §33.1. A clustering of a finite set of points partitions them into k clusters, each with a centroid, and its cost is the sum over all points of the squared Euclidean distance to the centroid of the point's own cluster.

Main results:

  • Definition sumSqDist: the within-cluster sum of squared distances to a candidate center.

  • Definition mean: the centroid (average) of a finite point family over an index set.

  • Definition Clustering: an assignment of each point to a cluster together with one centroid per cluster.

  • Definition kMeansCost: the cost of a clustering.

  • Theorem mean_minimizes_sumSqDist (Lemma 33.1): the mean of a cluster minimizes the within-cluster sum of squared distances.

  • Theorem assignStep_cost_le: reassigning every point to a nearest centroid never increases the cost.

  • Theorem updateStep_cost_le: replacing every centroid by the mean of its cluster never increases the cost.

  • Theorem lloyd_iteration_cost_le (Theorem 33.2): one Lloyd iteration (assignment, then update) never increases the cost.

The ambient point space is a real inner product space E (CLRS works with ℝᵈ, which is EuclideanSpace ℝ (Fin d) in Mathlib). Points are given by a finite family P : Fin n → E, and a cluster is an index set; the mean of a cluster counts each of its indices, so two distinct indices with coincident coordinates are both included. The development is deliberately abstract: only the monotonicity of the cost under Lloyd's two steps is proved; the termination guarantee and convergence rate are not addressed.

Notation conventions used in this section:

  • E : the ambient real inner product space containing the points

  • n : the number of points

  • k : the number of clusters

  • P : a finite family of points (Fin n → E)

  • C : a clustering of the points

noncomputable sectionopen scoped BigOperatorsopen scoped RealInnerProductSpacenamespace CLRSnamespace KMeansvariable {E : Type*} [SeminormedAddCommGroup E] [InnerProductSpace ℝ E]

The sum of squared distances from the points P i (for i ∈ S) to a candidate center c. This is the within-cluster cost of a cluster whose points are indexed by S and whose center is c (CLRS §33.1).

def sumSqDist {α : Type*} (P : α → E) (S : Finset α) (c : E) : ℝ := ∑ i ∈ S, ‖P i - c‖ ^ 2

The mean (centroid) of the point family P over the index set S: the average of the points P i for i ∈ S. For the empty set it returns the zero vector.

def mean {α : Type*} (P : α → E) (S : Finset α) : E := (S.card : ℝ)⁻¹ • (∑ i ∈ S, P i)

The sum of the displacements of the points of S from their mean vanishes: the mean is the center of mass of the point family P over S.

lemma sum_sub_mean_eq_zero {α : Type*} (P : α → E) (S : Finset α) : ∑ i ∈ S, (P i - mean P S) = 0 := by by_cases h0 : S.card = 0 · rw [Finset.card_eq_zero.mp h0] simp [mean] · have hn : (S.card : ℝ) ≠ 0 := by exact_mod_cast h0 calc (∑ i ∈ S, (P i - mean P S)) = (∑ i ∈ S, P i) - (∑ i ∈ S, mean P S) := by rw [Finset.sum_sub_distrib] _ = (∑ i ∈ S, P i) - S.card • mean P S := by rw [Finset.sum_const] _ = (∑ i ∈ S, P i) - (∑ i ∈ S, P i) := by rw [show S.card • mean P S = ∑ i ∈ S, P i by rw [mean] rw [← Nat.cast_smul_eq_nsmul (R := ℝ)] exact smul_inv_smul₀ hn (∑ i ∈ S, P i)] _ = 0 := by abel

Variance decomposition (parallel-axis theorem). For any index set S, point family P with mean m, and any center c, the sum of squared distances to c exceeds the sum of squared distances to m by |S| · ‖c - m‖²: ∑ i ∈ S ‖P i - c‖² = ∑ i ∈ S ‖P i - m‖² + |S| · ‖c - m‖².

theorem sumSqDist_eq_add_card_mul {α : Type*} (P : α → E) (S : Finset α) (c : E) : sumSqDist P S c = sumSqDist P S (mean P S) + (S.card : ℝ) * ‖c - mean P S‖ ^ 2 := by unfold sumSqDist calc (∑ i ∈ S, ‖P i - c‖ ^ 2) = (∑ i ∈ S, (‖P i - mean P S‖ ^ 2 + 2 * ⟪P i - mean P S, mean P S - c⟫ + ‖mean P S - c‖ ^ 2)) := by apply Finset.sum_congr rfl intro i hi have hpc : P i - c = (P i - mean P S) + (mean P S - c) := by abel rw [hpc, norm_add_sq_real] _ = (∑ i ∈ S, ‖P i - mean P S‖ ^ 2) + 2 * (∑ i ∈ S, ⟪P i - mean P S, mean P S - c⟫) + (S.card : ℝ) * ‖mean P S - c‖ ^ 2 := by rw [Finset.sum_add_distrib] rw [Finset.sum_add_distrib] rw [Finset.mul_sum] rw [Finset.sum_const] simp _ = (∑ i ∈ S, ‖P i - mean P S‖ ^ 2) + 2 * ⟪∑ i ∈ S, (P i - mean P S), mean P S - c⟫ + (S.card : ℝ) * ‖mean P S - c‖ ^ 2 := by rw [← sum_inner S (fun i => P i - mean P S) (mean P S - c)] _ = (∑ i ∈ S, ‖P i - mean P S‖ ^ 2) + (S.card : ℝ) * ‖mean P S - c‖ ^ 2 := by rw [sum_sub_mean_eq_zero P S] simp _ = sumSqDist P S (mean P S) + (S.card : ℝ) * ‖c - mean P S‖ ^ 2 := by rw [← norm_sub_rev (mean P S) c] simp [sumSqDist]

Lemma 33.1. The mean of a cluster minimizes the sum of squared distances to its points: for any index set S, point family P, and candidate center c, ∑ i ∈ S ‖P i - c‖² ≥ ∑ i ∈ S ‖P i - mean P S‖². The variance decomposition sumSqDist_eq_add_card_mul shows the difference is exactly |S| · ‖c - mean P S‖² ≥ 0.

theorem mean_minimizes_sumSqDist {α : Type*} (P : α → E) (S : Finset α) (c : E) : sumSqDist P S (mean P S) ≤ sumSqDist P S c := by rw [sumSqDist_eq_add_card_mul P S c] have h : 0 ≤ (S.card : ℝ) * ‖c - mean P S‖ ^ 2 := by positivity linarith

A clustering of n points into k clusters consists of an assignment of every point to a cluster together with a centroid for every cluster (CLRS §33.1). The ambient point space E is a real inner product space.

The cluster to which each point is assigned.

The centroid of each cluster.

structure Clustering (n k : ℕ) (E : Type*) where assign : Fin n → Fin k centroid : Fin k → E

The points of P assigned to cluster j by C, as a Finset of indices.

def cluster (C : Clustering n k E) (j : Fin k) : Finset (Fin n) := Finset.univ.filter fun i => C.assign i = j

The k-means cost of a clustering C of the points P: the sum over all points of the squared distance to the centroid of the point's assigned cluster. This equals the sum over clusters of the within-cluster squared distances (see kMeansCost_eq_sum_cluster).

def kMeansCost (P : Fin n → E) (C : Clustering n k E) : ℝ := ∑ i : Fin n, ‖P i - C.centroid (C.assign i)‖ ^ 2

Sum over the fibers of a function on a finite type: the sum over β of the sum of f on the preimage of each element equals the total sum of f.

lemma sum_fiberwise {α β : Type*} [Fintype β] [DecidableEq β] {γ : Type*} [AddCommMonoid γ] (s : Finset α) (g : α → β) (f : α → γ) : (∑ b : β, ∑ a ∈ s.filter (fun a => g a = b), f a) = ∑ a ∈ s, f a := by calc (∑ b : β, ∑ a ∈ s.filter (fun a => g a = b), f a) = (∑ b : β, ∑ a ∈ s, if g a = b then f a else 0) := by simp_rw [Finset.sum_filter] _ = (∑ a ∈ s, ∑ b : β, if g a = b then f a else 0) := by rw [Finset.sum_comm] _ = ∑ a ∈ s, f a := by apply Finset.sum_congr rfl intro a ha simp

The k-means cost of a clustering is the sum, over clusters, of the sum of squared distances from each cluster's points to its own centroid.

automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false` automatically included section variable(s) unused in theorem `CLRS.KMeans.kMeansCost_eq_sum_cluster`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`theorem kMeansCost_eq_sum_cluster (P : Fin n → E) (C : Clustering n k E) : kMeansCost P C = ∑ j : Fin k, ∑ i ∈ cluster C j, ‖P i - C.centroid j‖ ^ 2 := by unfold kMeansCost have hfw : (∑ j : Fin k, ∑ i ∈ (Finset.univ : Finset (Fin n)).filter (fun i => C.assign i = j), ‖P i - C.centroid (C.assign i)‖ ^ 2) = ∑ i : Fin n, ‖P i - C.centroid (C.assign i)‖ ^ 2 := by exact sum_fiberwise (s := (Finset.univ : Finset (Fin n))) (g := C.assign) (f := fun i => ‖P i - C.centroid (C.assign i)‖ ^ 2) calc (∑ i : Fin n, ‖P i - C.centroid (C.assign i)‖ ^ 2) = (∑ j : Fin k, ∑ i ∈ cluster C j, ‖P i - C.centroid (C.assign i)‖ ^ 2) := by unfold cluster rw [hfw] _ = (∑ j : Fin k, ∑ i ∈ cluster C j, ‖P i - C.centroid j‖ ^ 2) := by apply Finset.sum_congr rfl intro j hj apply Finset.sum_congr rfl intro i hi rw [(Finset.mem_filter.mp hi).2]

There is a cluster whose centroid is a nearest centroid to p: among the k candidates some one minimizes the squared distance to p.

automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false` automatically included section variable(s) unused in theorem `CLRS.KMeans.exists_nearest`: [InnerProductSpace ℝ E] consider restructuring your `variable` declarations so that the variables are not in scope or explicitly omit them: omit [InnerProductSpace ℝ E] in theorem ... Note: This linter can be disabled with `set_option linter.unusedSectionVars false`lemma exists_nearest (centroids : Fin k → E) (p : E) (hk : 0 < k) : ∃ j : Fin k, ∀ j' : Fin k, ‖p - centroids j‖ ^ 2 ≤ ‖p - centroids j'‖ ^ 2 := by let vals : Finset ℝ := Finset.univ.image fun j : Fin k => ‖p - centroids j‖ ^ 2 have hne : vals.Nonempty := by haveI : Nonempty (Fin k) := ⟨⟨0, hk⟩⟩ obtain ⟨j⟩ := (Finset.univ_nonempty : (Finset.univ : Finset (Fin k)).Nonempty) exact ⟨‖p - centroids j‖ ^ 2, Finset.mem_image.mpr ⟨j, Finset.mem_univ j, rfl⟩⟩ let m : ℝ := vals.min' hne have hm_le : ∀ y ∈ vals, m ≤ y := by intro y hy simpa [m] using vals.min'_le y hy have hm_mem : m ∈ vals := Finset.min'_mem vals hne rcases Finset.mem_image.mp hm_mem with ⟨j, hj, hjm⟩ refine ⟨j, ?_⟩ intro j' exact (hjm.symm ▸ hm_le (‖p - centroids j'‖ ^ 2) (Finset.mem_image.mpr ⟨j', Finset.mem_univ j', rfl⟩))

The index of a nearest centroid to p (CLRS §33.1, assignment step). Ties are broken by the choice made in exists_nearest.

noncomputable def nearestIndex (centroids : Fin k → E) (p : E) (hk : 0 < k) : Fin k := Classical.choose (exists_nearest centroids p hk)

The nearest centroid is no farther from p than any other centroid.

lemma nearestIndex_sq_le (centroids : Fin k → E) (p : E) (hk : 0 < k) (j : Fin k) : ‖p - centroids (nearestIndex centroids p hk)‖ ^ 2 ≤ ‖p - centroids j‖ ^ 2 := by exact Classical.choose_spec (exists_nearest centroids p hk) j

The assignment step of Lloyd's algorithm: keep the centroids fixed and move every point to a nearest centroid.

def assignStep (P : Fin n → E) (C : Clustering n k E) (hk : 0 < k) : Clustering n k E where assign := fun i => nearestIndex C.centroid (P i) hk centroid := C.centroid

The assignment step never increases the cost: moving every point to a nearest centroid can only bring each point closer to its centroid.

theorem assignStep_cost_le (P : Fin n → E) (C : Clustering n k E) (hk : 0 < k) : kMeansCost P (assignStep P C hk) ≤ kMeansCost P C := by unfold kMeansCost apply Finset.sum_le_sum intro i hi simpa [assignStep] using (nearestIndex_sq_le C.centroid (P i) hk (C.assign i))

The mean of cluster j in the clustering C: the average of the points assigned to j.

def clusterMean (P : Fin n → E) (C : Clustering n k E) (j : Fin k) : E := mean P (cluster C j)

The update step of Lloyd's algorithm: keep the assignment fixed and replace every centroid by the mean of its cluster's points.

def updateStep (P : Fin n → E) (C : Clustering n k E) : Clustering n k E where assign := C.assign centroid := fun j => clusterMean P C j

The update step never increases the cost: by mean_minimizes_sumSqDist (Lemma 33.1) each cluster's within-cluster cost is minimized when its centroid is the mean of its points.

theorem updateStep_cost_le (P : Fin n → E) (C : Clustering n k E) : kMeansCost P (updateStep P C) ≤ kMeansCost P C := by rw [kMeansCost_eq_sum_cluster P (updateStep P C)] rw [kMeansCost_eq_sum_cluster P C] calc (∑ j : Fin k, ∑ i ∈ cluster (updateStep P C) j, ‖P i - (updateStep P C).centroid j‖ ^ 2) = (∑ j : Fin k, ∑ i ∈ cluster C j, ‖P i - mean P (cluster C j)‖ ^ 2) := by apply Finset.sum_congr rfl intro j hj have hcl : cluster (updateStep P C) j = cluster C j := by simp [cluster, updateStep] rw [hcl] apply Finset.sum_congr rfl intro i hi simp [updateStep, clusterMean] _ ≤ (∑ j : Fin k, ∑ i ∈ cluster C j, ‖P i - C.centroid j‖ ^ 2) := by apply Finset.sum_le_sum intro j hj simpa [clusterMean, sumSqDist] using mean_minimizes_sumSqDist P (cluster C j) (C.centroid j)

Theorem 33.2. Lloyd's algorithm never increases the cost: one full iteration (assign every point to a nearest centroid, then replace each centroid by the mean of its cluster) does not increase the k-means cost.

theorem lloyd_iteration_cost_le (P : Fin n → E) (C : Clustering n k E) (hk : 0 < k) : kMeansCost P (updateStep P (assignStep P C hk)) ≤ kMeansCost P C := by exact le_trans (updateStep_cost_le P (assignStep P C hk)) (assignStep_cost_le P C hk)
end KMeansend CLRS
Imports
import Mathlib.Analysis.SpecialFunctions.Log.Basic import Mathlib.Analysis.SpecialFunctions.Log.Deriv import Mathlib.Analysis.SpecialFunctions.Pow.Real import Mathlib.Analysis.SpecialFunctions.Exp import Mathlib.Analysis.Convex.SpecificFunctions.Basic import Mathlib.Analysis.Calculus.Deriv.MeanValue import Mathlib.Tactic

33.2. Multiplicative-Weights Algorithms

This section formalizes the multiplicative weights (MW) update method from CLRS §33.2. The setting has n experts. Over T days, each expert i incurs a loss m t i ∈ [0, 1] on day t, and the algorithm maintains a weight w i for each expert, initially 1. On each day it incurs the expected loss Σᵢ (w i / Φ) · m t i under the normalized distribution w i / Φ, where Φ = Σᵢ w i is the potential; it then multiplies expert i's weight by (1 - η)^(m t i) for a learning rate η ∈ (0, 1/2].

Main results:

  • Definition weights: the weight of each expert after any number of days.

  • Definition potential: the sum of the weights (Φ).

  • Definition expectedLoss / totalExpectedLoss: the algorithm's daily and total expected loss.

  • Definition expertLoss: the total loss of a single expert.

  • Theorem potential_update_le_exp: one MW update shrinks the potential: Φ' ≤ Φ · exp(-η · M) where M is the day's expected loss.

  • Theorem totalExpectedLoss_le: the total expected loss of the algorithm is within an additive ln n / η and a multiplicative (1 + η) factor of the best expert's loss — for every expert i, Σₜ Mᵗ ≤ (1 + η) · Σₜ m t i + ln n / η.

The proof follows CLRS §33.2: three analytic inequalities — (1-x)^y ≤ 1 - x·y (convexity of the exponential), 1 - x ≤ e^(-x), and -ln (1 - x) ≤ x + x² for 0 ≤ x ≤ 1/2 — feed the potential chain Φ^{t+1} ≤ Φᵗ·e^(-η·Mᵗ), which is sandwiched between the weight of the best expert, (1-η)^L, and the initial potential n, and the resulting logarithmic inequality is rearranged.

Notation conventions used in this section:

  • n : the number of experts

  • T : the number of days

  • η : the learning rate (0 < η ≤ 1/2)

  • m t i : the loss of expert i on day t (∈ [0, 1])

  • w : a weight vector (Fin n → ℝ)

  • Φ : the potential, the sum of the weights

  • Mᵗ : the algorithm's expected loss on day t

noncomputable sectionopen scoped BigOperatorsnamespace CLRSnamespace MultiplicativeWeightsvariable {n : ℕ}

The initial weight vector: every expert starts with weight 1 (CLRS §33.2).

def initialWeights (n : ℕ) : Fin n → ℝ := fun _ => 1

The potential Φ(w) of a weight vector w: the sum of all expert weights (CLRS §33.2).

def potential {n : ℕ} (w : Fin n → ℝ) : ℝ := ∑ i : Fin n, w i

The potential of the initial weight vector is the number of experts n.

lemma potential_initialWeights (n : ℕ) : potential (initialWeights n) = n := by simp [potential, initialWeights, Finset.sum_const, Finset.card_univ]

One multiplicative-weights update: expert i's weight is multiplied by (1 - η) ^ (l i), where l i ∈ [0, 1] is expert i's loss on the current day (CLRS §33.2).

def updateWeight (η : ℝ) {n : ℕ} (w : Fin n → ℝ) (l : Fin n → ℝ) : Fin n → ℝ := fun i => w i * (1 - η) ^ l i

The losses of the n experts on day t, totalized to every natural index: for days t ≥ T the zero loss vector is returned (a junk value that the lemmas only ever use inside Finset.range T).

def dayLoss {T n : ℕ} (m : Fin T → Fin n → ℝ) (t : ℕ) : Fin n → ℝ := if h : t < T then m ⟨t, h⟩ else fun _ => 0

The weight vector after t days of updates (meaningful for t ≤ T): the result of applying the update rule for days 0, ..., t - 1 to initialWeights n.

def weights {T n : ℕ} (η : ℝ) (m : Fin T → Fin n → ℝ) : ℕ → Fin n → ℝ | 0 => initialWeights n | t + 1 => updateWeight η (weights η m t) (dayLoss m t)

The expected loss of the algorithm on a day with current weight vector w and losses l: Σᵢ (w i / Φ(w)) · l i, the average loss under the weight-normalized distribution (CLRS §33.2).

def expectedLossAt {n : ℕ} (w : Fin n → ℝ) (l : Fin n → ℝ) : ℝ := ∑ i : Fin n, (w i / potential w) * l i

The algorithm's expected loss on day t (with t < T the loss is m t i; for t ≥ T the totalized loss vector is zero).

def expectedLoss {T n : ℕ} (η : ℝ) (m : Fin T → Fin n → ℝ) (t : ℕ) : ℝ := expectedLossAt (weights η m t) (dayLoss m t)

The algorithm's total expected loss over days 0, ..., T - 1.

def totalExpectedLoss {T n : ℕ} (η : ℝ) (m : Fin T → Fin n → ℝ) : ℝ := ∑ t ∈ Finset.range T, expectedLoss η m t

The total loss of expert i over days 0, ..., T - 1.

def expertLoss {T n : ℕ} (m : Fin T → Fin n → ℝ) (i : Fin n) : ℝ := ∑ t ∈ Finset.range T, dayLoss m t i

The day-t loss of every expert lies in [0, 1] for every totalized index.

lemma dayLoss_mem_Icc {T n : ℕ} {m : Fin T → Fin n → ℝ} (hm : ∀ t j, m t j ∈ Set.Icc (0 : ℝ) 1) (t : ℕ) : ∀ i, dayLoss m t i ∈ Set.Icc (0 : ℝ) 1 := by intro i by_cases ht : t < T · simpa [dayLoss, ht] using hm ⟨t, ht⟩ i · simp [dayLoss, ht]

For 0 ≤ x < 1 and 0 ≤ y ≤ 1, (1 - x) ^ y ≤ 1 - x·y. This is the fact that the exponential function is convex, applied between the points 0 and log (1 - x) with weights 1 - y and y.

lemma one_sub_rpow_le_one_sub_mul {x y : ℝ} (Variable name `hx0` is not explicitly referenced. The binding can be removed (if unused) or named `_` (if used implicitly). Note: This linter can be disabled with `set_option linter.unusedVariables false`hx0 : 0 ≤ x) (hx1 : x < 1) (hy0 : 0 ≤ y) (hy1 : y ≤ 1) : (1 - x) ^ y ≤ 1 - x * y := by calc (1 - x) ^ y = Real.exp (Real.log (1 - x) * y) := by exact Real.rpow_def_of_pos (by linarith : 0 < 1 - x) y _ = Real.exp (y * Real.log (1 - x)) := by rw [mul_comm] _ ≤ (1 - y) * Real.exp 0 + y * Real.exp (Real.log (1 - x)) := by have hc0 : Real.exp (y * Real.log (1 - x)) ≤ (1 - y) * Real.exp 0 + y * Real.exp (Real.log (1 - x)) := by simpa [smul_eq_mul, zero_mul, add_zero] using (convexOn_exp.2 (by simp : (0 : ℝ) ∈ Set.univ) (by simp : Real.log (1 - x) ∈ Set.univ) (show 0 ≤ 1 - y by linarith) (by exact hy0) (by ring)) exact hc0 _ = (1 - y) * 1 + y * (1 - x) := by rw [Real.exp_zero, Real.exp_log (by linarith : 0 < 1 - x)] _ = 1 - x * y := by ring

For 0 ≤ x ≤ 1/2, -ln (1 - x) ≤ x + x². This is the logarithmic approximation used to pass from the factor -ln (1 - η)/η to the clean 1 + η in the regret bound.

lemma neg_log_one_sub_le_add_sq {x : ℝ} (hx0 : 0 ≤ x) (hx : x ≤ 1 / 2) : -Real.log (1 - x) ≤ x + x ^ 2 := by let f : ℝ → ℝ := fun t => t + t ^ 2 + Real.log (1 - t) have hd : ∀ t : ℝ, t < 1 → HasDerivAt f (1 + 2 * t - (1 - t)⁻¹) t := by intro t ht have hid : HasDerivAt (fun u : ℝ => u) 1 t := hasDerivAt_id t have hsq : HasDerivAt (fun u : ℝ => u ^ 2) (2 * t) t := by simpa using (hasDerivAt_pow 2 t) have hlog : HasDerivAt (fun u : ℝ => Real.log (1 - u)) (-(1 - t)⁻¹) t := by have hsub : HasDerivAt (fun u : ℝ => (1 : ℝ) - u) (-1) t := by convert ((hasDerivAt_const t (1 : ℝ)).sub hid) using 1 · rfl · rfl · rfl · norm_num have hlog1 : HasDerivAt Real.log ((1 - t)⁻¹) (1 - t) := by exact Real.hasDerivAt_log (by linarith : (1 - t) ≠ 0) convert (hlog1.comp t hsub) using 1 · rfl · rfl · rfl · ring have hsum : HasDerivAt (fun u : ℝ => u + u ^ 2 + Real.log (1 - u)) (1 + 2 * t - (1 - t)⁻¹) t := by convert ((hid.add hsq).add hlog) using 1 · rfl · rfl · rfl · ring simpa [f] using hsum have hmono : MonotoneOn f (Set.Icc 0 (1 / 2)) := by refine monotoneOn_of_deriv_nonneg (convex_Icc 0 (1 / 2)) ?_ ?_ ?_ · have hc_poly : ContinuousOn (fun t : ℝ => t + t ^ 2) (Set.Icc 0 (1 / 2)) := by exact continuousOn_id.add (continuousOn_id.pow 2) have hc_log : ContinuousOn (fun t : ℝ => Real.log (1 - t)) (Set.Icc 0 (1 / 2)) := by refine ContinuousOn.log ?_ ?_ · exact continuousOn_const.sub continuousOn_id · intro t ht exact ne_of_gt (by linarith [ht.2]) have hc_sum : ContinuousOn (fun t : ℝ => (t + t ^ 2) + Real.log (1 - t)) (Set.Icc 0 (1 / 2)) := by exact (hc_poly.add hc_log) simpa [f] using hc_sum · rw [interior_Icc] intro t ht exact (hd t (by linarith [ht.2])).differentiableAt.differentiableWithinAt · rw [interior_Icc] intro t ht have hd_t : HasDerivAt f (1 + 2 * t - (1 - t)⁻¹) t := hd t (by linarith [ht.2]) rw [hd_t.deriv] have hden : 0 < 1 - t := by linarith [ht.2] have hnum : 0 ≤ t * (1 - 2 * t) := by have ht0 : 0 < t := ht.1 have ht1 : 0 < 1 - 2 * t := by linarith [ht.2] positivity field_simp [hden.ne'] nlinarith [hnum] have hmem0 : (0 : ℝ) ∈ Set.Icc 0 (1 / 2) := by norm_num have hmemx : x ∈ Set.Icc 0 (1 / 2) := by exact ⟨hx0, hx⟩ have hle : f 0 ≤ f x := hmono hmem0 hmemx hx0 have hf0 : f 0 = 0 := by simp [f] have hfx : f x = x + x ^ 2 + Real.log (1 - x) := by simp [f] nlinarith [hle, hf0, hfx]

Every weight is strictly positive: weights only ever get multiplied by powers of 1 - η > 0.

lemma weights_pos {T n : ℕ} {η : ℝ} {m : Fin T → Fin n → ℝ} (hη1 : η < 1) : ∀ t i, 0 < weights η m t i := by intro t i induction t with | zero => simp [weights, initialWeights] | succ t ih => have hw : 0 < weights η m t i := ih have hp : 0 < (1 - η) ^ dayLoss m t i := Real.rpow_pos_of_pos (by linarith : 0 < 1 - η) (dayLoss m t i) simpa [weights, updateWeight] using mul_pos hw hp

The potential at any time is positive: it is a sum of strictly positive weights, of which there is at least one when n > 0.

lemma potential_weights_pos {T n : ℕ} {η : ℝ} {m : Fin T → Fin n → ℝ} (hn : 0 < n) (hη1 : η < 1) (t : ℕ) : 0 < potential (weights η m t) := by have h0 : 0 < weights η m t ⟨0, hn⟩ := weights_pos hη1 t ⟨0, hn⟩ have hs : weights η m t ⟨0, hn⟩ ≤ potential (weights η m t) := Finset.single_le_sum (fun i hi => le_of_lt (weights_pos hη1 t i)) (by simp) exact lt_of_lt_of_le h0 hs

Σᵢ w i · l i = Φ · M, where M is the expected loss at w, l: the unweighted sum of weighted losses equals the potential times the expected loss.

lemma sum_weight_mul_loss_eq {n : ℕ} {w : Fin n → ℝ} {l : Fin n → ℝ} (hpw : 0 < potential w) : (∑ i : Fin n, w i * l i) = potential w * expectedLossAt w l := by calc (∑ i : Fin n, w i * l i) = ∑ i : Fin n, potential w * ((w i / potential w) * l i) := by apply Finset.sum_congr rfl intro i hi field_simp [ne_of_gt hpw] _ = potential w * (∑ i : Fin n, (w i / potential w) * l i) := by rw [← Finset.mul_sum] _ = potential w * expectedLossAt w l := by rfl

Per-round potential bound. One MW update shrinks the potential by at least the weighted loss: Φ(w') ≤ Φ(w) - η · Σᵢ w i · l i, using the pointwise inequality (1 - η)^(l i) ≤ 1 - η·(l i).

lemma potential_update_le {n : ℕ} {η : ℝ} {w : Fin n → ℝ} {l : Fin n → ℝ} (hw : ∀ i, 0 ≤ w i) (hη0 : 0 ≤ η) (hη1 : η < 1) (hl : ∀ i, l i ∈ Set.Icc (0 : ℝ) 1) : potential (updateWeight η w l) ≤ potential w - η * (∑ i : Fin n, w i * l i) := by unfold potential updateWeight calc (∑ i : Fin n, w i * (1 - η) ^ l i) ≤ ∑ i : Fin n, (w i - η * (w i * l i)) := by apply Finset.sum_le_sum intro i hi have hterm : (1 - η) ^ l i ≤ 1 - η * l i := by exact one_sub_rpow_le_one_sub_mul hη0 hη1 (hl i).1 (hl i).2 calc w i * (1 - η) ^ l i ≤ w i * (1 - η * l i) := by exact mul_le_mul_of_nonneg_left hterm (hw i) _ = w i - η * (w i * l i) := by ring _ = (∑ i : Fin n, w i) - η * (∑ i : Fin n, w i * l i) := by rw [Finset.sum_sub_distrib] rw [← Finset.mul_sum] _ = potential w - η * (∑ i : Fin n, w i * l i) := by rfl

Exponential potential bound. One MW update shrinks the potential multiplicatively: Φ(w') ≤ Φ(w) · exp (-η · M), where M is the day's expected loss. This chains the algebraic per-round bound with 1 - x ≤ e^(-x).

lemma potential_update_le_exp {n : ℕ} {η : ℝ} {w : Fin n → ℝ} {l : Fin n → ℝ} (hw : ∀ i, 0 ≤ w i) (hpw : 0 < potential w) (hη0 : 0 ≤ η) (hη1 : η < 1) (hl : ∀ i, l i ∈ Set.Icc (0 : ℝ) 1) : potential (updateWeight η w l) ≤ potential w * Real.exp (-η * expectedLossAt w l) := by calc potential (updateWeight η w l) ≤ potential w - η * (∑ i : Fin n, w i * l i) := potential_update_le hw hη0 hη1 hl _ = potential w * (1 - η * expectedLossAt w l) := by rw [sum_weight_mul_loss_eq hpw] ring _ ≤ potential w * Real.exp (-η * expectedLossAt w l) := by simpa using (mul_le_mul_of_nonneg_left (Real.one_sub_le_exp_neg (η * expectedLossAt w l)) (le_of_lt hpw))

Potential chain. After t ≤ T days, the potential is at most n · exp (-η · Σ_{s < t} Mˢ): the initial potential n decays by a factor e^(-η·Mˢ) on every day. This is the induction that iterates potential_update_le_exp over the days.

lemma potential_weights_le {T n : ℕ} {η : ℝ} {m : Fin T → Fin n → ℝ} (hn : 0 < n) (hη0 : 0 ≤ η) (hη1 : η < 1) (hm : ∀ t j, m t j ∈ Set.Icc (0 : ℝ) 1) : ∀ t : ℕ, t ≤ T → potential (weights η m t) ≤ n * Real.exp (-η * ∑ s ∈ Finset.range t, expectedLoss η m s) := by intro t induction t with | zero => intro _ht0 calc potential (weights η m 0) = n := by simp [weights, potential_initialWeights] _ ≤ n * Real.exp (-η * (∑ s ∈ Finset.range 0, expectedLoss η m s)) := by simp [This simp argument is unused: Finset.sum_range_zero Hint: Omit it from the simp argument list. simp ̵[̵F̵i̵n̵s̵e̵t̵.̵s̵u̵m̵_̵r̵a̵n̵g̵e̵_̵z̵e̵r̵o̵]̵ Note: This linter can be disabled with `set_option linter.unusedSimpArgs false`Finset.sum_range_zero] | succ t ih => intro hts have htT : t < T := Nat.lt_of_succ_le hts have ht_le : t ≤ T := le_of_lt htT have ih_t := ih ht_le have hup : potential (updateWeight η (weights η m t) (dayLoss m t)) ≤ potential (weights η m t) * Real.exp (-η * expectedLossAt (weights η m t) (dayLoss m t)) := by exact potential_update_le_exp (fun i => le_of_lt (weights_pos hη1 t i)) (potential_weights_pos hn hη1 t) hη0 hη1 (dayLoss_mem_Icc hm t) have h1 : potential (weights η m (t + 1)) ≤ potential (weights η m t) * Real.exp (-η * expectedLoss η m t) := by simpa [weights, expectedLoss] using hup have h3 : potential (weights η m t) * Real.exp (-η * expectedLoss η m t) ≤ n * Real.exp (-η * (∑ s ∈ Finset.range (t + 1), expectedLoss η m s)) := by calc potential (weights η m t) * Real.exp (-η * expectedLoss η m t) ≤ (n * Real.exp (-η * (∑ s ∈ Finset.range t, expectedLoss η m s))) * Real.exp (-η * expectedLoss η m t) := by exact mul_le_mul_of_nonneg_right ih_t (Real.exp_pos _).le _ = n * (Real.exp (-η * (∑ s ∈ Finset.range t, expectedLoss η m s)) * Real.exp (-η * expectedLoss η m t)) := by ring _ = n * Real.exp ((-η * (∑ s ∈ Finset.range t, expectedLoss η m s)) + (-η) * expectedLoss η m t) := by rw [← Real.exp_add] _ = n * Real.exp (-η * (∑ s ∈ Finset.range (t + 1), expectedLoss η m s)) := by congr 1 congr 1 rw [Finset.sum_range_succ] ring exact le_trans h1 h3

The weight of expert i after T days is (1 - η) ^ (expertLoss m i): each day multiplies the weight by (1 - η)^(m t i), so the accumulated weight is (1 - η) raised to the expert's total loss.

lemma weights_eq_rpow_expertLoss {T n : ℕ} {η : ℝ} {m : Fin T → Fin n → ℝ} (hη1 : η < 1) (i : Fin n) : weights η m T i = (1 - η) ^ expertLoss m i := by have hbase : 0 < 1 - η := by linarith have aux : ∀ t : ℕ, weights η m t i = ∏ s ∈ Finset.range t, (1 - η) ^ dayLoss m s i := by intro t induction t with | zero => simp [weights, initialWeights] | succ t ih => rw [weights] rw [updateWeight] rw [ih] rw [Finset.prod_range_succ] calc weights η m T i = ∏ s ∈ Finset.range T, (1 - η) ^ dayLoss m s i := aux T _ = (1 - η) ^ (∑ s ∈ Finset.range T, dayLoss m s i) := by rw [← Real.rpow_sum_of_pos hbase] _ = (1 - η) ^ expertLoss m i := by simp [expertLoss]

The total loss of every expert is nonnegative, since every daily loss lies in [0, 1].

lemma expertLoss_nonneg {T n : ℕ} {m : Fin T → Fin n → ℝ} (hm : ∀ t j, m t j ∈ Set.Icc (0 : ℝ) 1) (i : Fin n) : 0 ≤ expertLoss m i := by unfold expertLoss exact Finset.sum_nonneg (fun t ht => (dayLoss_mem_Icc hm t i).1)

Regret bound (Theorem 33.3). For every expert i, the total expected loss of the multiplicative-weights algorithm is within an additive ln n / η and a multiplicative (1 + η) factor of expert i's total loss:

Σₜ Mᵗ ≤ (1 + η) · Σₜ m t i + ln n / η.

The proof sandwiches the final potential between the best expert's remaining weight, (1 - η)^L, and the decayed initial potential, n·e^(-η·E), takes logarithms, and uses -ln (1 - η) ≤ η + η².

theorem totalExpectedLoss_le {T n : ℕ} {η : ℝ} {m : Fin T → Fin n → ℝ} {i : Fin n} (hn : 0 < n) (hη0 : 0 < η) (hη1 : η ≤ 1 / 2) (hm : ∀ t j, m t j ∈ Set.Icc (0 : ℝ) 1) : totalExpectedLoss η m ≤ (1 + η) * expertLoss m i + Real.log n / η := by let E := totalExpectedLoss η m let L := expertLoss m i have hη0' : 0 ≤ η := le_of_lt hη0 have hη1' : η < 1 := by linarith have hbase : 0 < 1 - η := by linarith have hub : potential (weights η m T) ≤ n * Real.exp (-η * E) := by have hpb := potential_weights_le hn hη0' hη1' hm T (le_rfl) simpa [E, totalExpectedLoss] using hpb have hlb : (1 - η) ^ L ≤ potential (weights η m T) := by have hprod : (1 - η) ^ L = weights η m T i := (weights_eq_rpow_expertLoss hη1' i).symm have hpot_ge : weights η m T i ≤ potential (weights η m T) := Finset.single_le_sum (fun j hj => le_of_lt (weights_pos hη1' T j)) (by simp) rw [hprod] exact hpot_ge have htot : (1 - η) ^ L ≤ n * Real.exp (-η * E) := le_trans hlb hub have hE : L * Real.log (1 - η) ≤ Real.log n - η * E := by have hle := (Real.rpow_le_iff_le_log hbase (by positivity : 0 < (n : ℝ) * Real.exp (-η * E))).mp htot rw [Real.log_mul (by exact_mod_cast hn.ne') (Real.exp_pos (-η * E)).ne', Real.log_exp] at hle simpa [sub_eq_add_neg] using hle have hlog2 : η * E ≤ Real.log n + L * (-Real.log (1 - η)) := by nlinarith [hE] have hdiv : E ≤ (Real.log n + L * (-Real.log (1 - η))) / η := by rw [le_div_iff₀ hη0] nlinarith [hlog2] have hlg : -Real.log (1 - η) ≤ η + η ^ 2 := neg_log_one_sub_le_add_sq hη0' hη1 have hLge : 0 ≤ L := expertLoss_nonneg hm i have hbound : (Real.log n + L * (-Real.log (1 - η))) / η ≤ (Real.log n + L * (η + η ^ 2)) / η := by gcongr have hsplit : (Real.log n + L * (η + η ^ 2)) / η = Real.log n / η + L * (1 + η) := by field_simp [ne_of_gt hη0] have hfinal : E ≤ Real.log n / η + L * (1 + η) := by calc E ≤ (Real.log n + L * (-Real.log (1 - η))) / η := hdiv _ ≤ (Real.log n + L * (η + η ^ 2)) / η := hbound _ = Real.log n / η + L * (1 + η) := hsplit have hgoal : E ≤ (1 + η) * L + Real.log n / η := by nlinarith [hfinal] simpa [E, L] using hgoal
end MultiplicativeWeightsend CLRS
Imports
import Mathlib.Analysis.Calculus.Gradient.Basic import Mathlib.Analysis.Calculus.LineDeriv.Basic import Mathlib.Analysis.Convex.Function import Mathlib.Analysis.Convex.Deriv import Mathlib.Analysis.Convex.Slope import Mathlib.Analysis.Convex.Jensen import Mathlib.Analysis.InnerProductSpace.Basic import Mathlib.Tactic

33.3. Gradient Descent

This section formalizes the gradient-descent algorithm for minimizing a convex function from CLRS §33.3. Given a convex differentiable function f : E → ℝ on a real inner product space E and a starting point x₀, the algorithm repeatedly moves in the direction of steepest descent, the negative gradient: x_{k+1} = x_k - η · ∇f(x_k) for a fixed learning rate η > 0.

Main results:

  • Definition gradientStep: one gradient-descent update.

  • Definition gdIterates: the sequence of iterates generated from x₀.

  • Definition avgIterate: the arithmetic mean of the first K iterates.

  • Theorem gradient_inner_le_sub: the gradient-descent lemma — for a convex differentiable f, the first-order characterization of convexity ⟪∇f(x), y - x⟫ ≤ f(y) - f(x) holds for all x, y.

  • Theorem gdStep_potential_le: one gradient-descent step shrinks the squared distance to any point x* by at least 2η·(f(x) - f(x*)), up to the additive η²·G² term coming from a bound ‖∇f‖ ≤ G on the gradient norm.

  • Theorem gdIterates_potential_le / sum_suboptimality_le: the telescoping potential chain over K steps and the resulting total-suboptimality bound.

  • Theorem avgIterate_suboptimality_le (Theorem 33.8): the convergence bound — if x* minimizes f, then the average x̄ of the first K iterates satisfies f(x̄) - f(x*) ≤ ‖x₀ - x*‖² / (2·η·K) + η·G² / 2.

The proof follows CLRS §33.3. The gradient-descent lemma is the analytic engine: it is proved from ConvexOn and differentiability by restricting f to the segment [x, y] and taking the limit, as t → 0⁺, of the convexity chord inequality (f(x + t(y-x)) - f(x)) / t ≤ f(y) - f(x). This feeds the per-step potential inequality ‖x_{k+1} - x*‖² ≤ ‖x_k - x*‖² - 2η(f(x_k) - f(x*)) + η²G², which telescopes over the K steps and combines with Jensen's inequality for the average iterate to give the convergence bound. The bound is stated for the average iterate; individual iterates can overshoot, and only the amortized progress is guaranteed to converge.

Notation conventions used in this section:

  • E : the ambient real inner product (Hilbert) space

  • f : a convex differentiable function E → ℝ to be minimized

  • η : the learning rate (0 < η)

  • x₀ : the starting point

  • x* : a point minimizing f

  • G : a bound on the norm of the gradient, ‖∇f(x)‖ ≤ G

  • x̄ : the average of the first K iterates

noncomputable sectionopen scoped BigOperatorsopen scoped RealInnerProductSpacenamespace CLRSnamespace GradientDescentvariable {E : Type*} [NormedAddCommGroup E] [InnerProductSpace ℝ E] [CompleteSpace E]

The gradient-descent lemma. For a convex differentiable function f, the gradient supports the graph from below at every point: for all x, y, ⟪∇f(x), y - x⟫ ≤ f(y) - f(x) (CLRS §33.3). This is the first-order characterization of convexity. The proof restricts f to the segment from x to y; convexity gives the chord bound (f(x + t·(y - x)) - f(x)) / t ≤ f(y) - f(x) for every t ∈ (0, 1], and differentiability of f at x lets t → 0⁺.

Used `tac1 <;> tac2` where `(tac1; tac2)` would suffice Note: This linter can be disabled with `set_option linter.unnecessarySeqFocus false` theorem gradient_inner_le_sub {x y : E} {f : E → ℝ} (hf : ConvexOn ℝ Set.univ f) (hd : DifferentiableAt ℝ f x) : ⟪gradient f x, y - x⟫ ≤ f y - f x := by rw [inner_gradient_left] let g : ℝ → ℝ := fun t => f (x + t • (y - x)) have hg_convex : ConvexOn ℝ (Set.univ : Set ℝ) g := by refine ⟨convex_univ, ?_⟩ intro t₁ ht₁ t₂ ht₂ a b ha hb hab have hseg : x + (a • t₁ + b • t₂) • (y - x) = a • (x + t₁ • (y - x)) + b • (x + t₂ • (y - x)) := by match_scalars <;> simp [smul_eq_mul] Used `tac1 <;> tac2` where `(tac1; tac2)` would suffice Note: This linter can be disabled with `set_option linter.unnecessarySeqFocus false`<;> nlinarith [hab] calc g (a • t₁ + b • t₂) = f (x + (a • t₁ + b • t₂) • (y - x)) := rfl _ = f (a • (x + t₁ • (y - x)) + b • (x + t₂ • (y - x))) := by rw [hseg] _ ≤ a • f (x + t₁ • (y - x)) + b • f (x + t₂ • (y - x)) := by have hx : x + t₁ • (y - x) ∈ (Set.univ : Set E) := by simp have hy : x + t₂ • (y - x) ∈ (Set.univ : Set E) := by simp exact hf.2 hx hy ha hb hab _ = a • g t₁ + b • g t₂ := rfl have hg_diff : HasDerivAt g (fderiv ℝ f x (y - x)) 0 := by have hf' : HasFDerivAt f (fderiv ℝ f x) x := hd.hasFDerivAt exact (hf'.hasLineDerivAt (y - x)) have hslope_le : fderiv ℝ f x (y - x) ≤ slope g 0 1 := by exact ConvexOn.le_slope_of_hasDerivWithinAt hg_convex (Set.mem_univ 0) (Set.mem_univ 1) (by norm_num) hg_diff.hasDerivWithinAt have hslope : slope g 0 1 = f y - f x := by rw [slope_def_field] have hg1 : g 1 = f y := by simp [g] have hg0 : g 0 = f x := by simp [g, zero_smul] rw [hg1, hg0] norm_num rw [← hslope] exact hslope_le

One gradient-descent step: x ↦ x - η·∇f(x), moving from x in the direction of steepest descent (the negative gradient) with learning rate η (CLRS §33.3).

def gradientStep (η : ℝ) (x : E) (f : E → ℝ) : E := x - η • gradient f x

The sequence of gradient-descent iterates generated from x₀: x_{k+1} = x_k - η·∇f(x_k) (CLRS §33.3).

def gdIterates (η : ℝ) (x₀ : E) (f : E → ℝ) : ℕ → E | 0 => x₀ | k + 1 => gradientStep η (gdIterates η x₀ f k) f

The average iterate x̄ after K steps: the arithmetic mean of the first K iterates x₀, ..., x_{K-1} (CLRS §33.3). For K = 0 the junk value 0 is returned, since (0 : ℝ)⁻¹ = 0.

def avgIterate (η : ℝ) (x₀ : E) (f : E → ℝ) (K : ℕ) : E := (K : ℝ)⁻¹ • ∑ k ∈ Finset.range K, gdIterates η x₀ f k

Per-step potential bound. One gradient-descent step from x shrinks the squared distance to any point x* by at least 2·η·(f x - f x*), up to the additive η²·G² term:

‖gradientStep η x f - x*‖² ≤ ‖x - x*‖² - 2·η·(f x - f x*) + η²·G²,

where G bounds the gradient norm at x. This is the engine of the convergence analysis: each step makes progress proportional to the current suboptimality, minus a small error term (CLRS §33.3).

Try this: [apply] abel_nf lemma gdStep_potential_le {x xstar : E} {η G : ℝ} {f : E → ℝ} (hf : ConvexOn ℝ Set.univ f) (hd : DifferentiableAt ℝ f x) (hη : 0 ≤ η) (hG : ‖gradient f x‖ ≤ G) : ‖gradientStep η x f - xstar‖ ^ 2 ≤ ‖x - xstar‖ ^ 2 - 2 * η * (f x - f xstar) + η ^ 2 * G ^ 2 := by have hg : f x - f xstar ≤ ⟪gradient f x, x - xstar⟫ := by have hle := gradient_inner_le_sub (x := x) (y := xstar) hf hd have hrewrite : ⟪gradient f x, x - xstar⟫ = -⟪gradient f x, xstar - x⟫ := by rw [inner_sub_right] rw [inner_sub_right] ring rw [hrewrite] linarith have hgnorm : ‖gradient f x‖ ^ 2 ≤ G ^ 2 := by nlinarith [hG, norm_nonneg (gradient f x)] calc ‖gradientStep η x f - xstar‖ ^ 2 = ‖(x - xstar) - η • gradient f x‖ ^ 2 := by rw [gradientStep] congr 1 Try this: [apply] abel_nfabel _ = ‖x - xstar‖ ^ 2 - 2 * ⟪η • gradient f x, x - xstar⟫ + ‖η • gradient f x‖ ^ 2 := by rw [norm_sub_sq (𝕜 := ℝ)] simp [real_inner_smul_left, real_inner_comm] _ = ‖x - xstar‖ ^ 2 - 2 * η * ⟪gradient f x, x - xstar⟫ + ‖η • gradient f x‖ ^ 2 := by rw [real_inner_smul_left] ring _ ≤ ‖x - xstar‖ ^ 2 - 2 * η * (f x - f xstar) + ‖η • gradient f x‖ ^ 2 := by have hinner : 2 * η * (f x - f xstar) ≤ 2 * η * ⟪gradient f x, x - xstar⟫ := by exact mul_le_mul_of_nonneg_left hg (by positivity) nlinarith _ ≤ ‖x - xstar‖ ^ 2 - 2 * η * (f x - f xstar) + η ^ 2 * G ^ 2 := by have hsmul : ‖η • gradient f x‖ ^ 2 = η ^ 2 * ‖gradient f x‖ ^ 2 := by rw [norm_smul, Real.norm_eq_abs, abs_of_nonneg hη] ring have hsqG : η ^ 2 * ‖gradient f x‖ ^ 2 ≤ η ^ 2 * G ^ 2 := by exact mul_le_mul_of_nonneg_left hgnorm (sq_nonneg η) rw [hsmul] nlinarith

Telescoping potential chain. After K gradient-descent steps the squared distance to x* satisfies

‖x_K - x*‖² ≤ ‖x₀ - x*‖² - 2·η·Σ_{k<K}(f(x_k) - f(x*)) + K·η²·G².

Each step's potential loss compounds into a cumulative suboptimality sum (CLRS §33.3).

lemma gdIterates_potential_le {x₀ xstar : E} {η G : ℝ} {f : E → ℝ} (hf : ConvexOn ℝ Set.univ f) (hd : ∀ y, DifferentiableAt ℝ f y) (hη : 0 ≤ η) (hG : ∀ y, ‖gradient f y‖ ≤ G) (K : ℕ) : ‖gdIterates η x₀ f K - xstar‖ ^ 2 ≤ ‖x₀ - xstar‖ ^ 2 - 2 * η * (∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar)) + (K : ℝ) * η ^ 2 * G ^ 2 := by induction K with | zero => simp [gdIterates] | succ K ih => let xK : E := gdIterates η x₀ f K have hrec : ‖gdIterates η x₀ f (K + 1) - xstar‖ ^ 2 ≤ ‖xK - xstar‖ ^ 2 - 2 * η * (f xK - f xstar) + η ^ 2 * G ^ 2 := by simpa [xK, gdIterates] using (gdStep_potential_le hf (hd xK) hη (hG xK)) calc ‖gdIterates η x₀ f (K + 1) - xstar‖ ^ 2 ≤ ‖xK - xstar‖ ^ 2 - 2 * η * (f xK - f xstar) + η ^ 2 * G ^ 2 := hrec _ ≤ ‖x₀ - xstar‖ ^ 2 - 2 * η * (∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar)) + (K : ℝ) * η ^ 2 * G ^ 2 - 2 * η * (f xK - f xstar) + η ^ 2 * G ^ 2 := by dsimp [xK] nlinarith [ih] _ = ‖x₀ - xstar‖ ^ 2 - 2 * η * (∑ k ∈ Finset.range (K + 1), (f (gdIterates η x₀ f k) - f xstar)) + ((K + 1 : ℕ) : ℝ) * η ^ 2 * G ^ 2 := by rw [Finset.sum_range_succ] simp [xK] ring

Total-suboptimality bound. The sum of suboptimalities over the first K iterates is bounded by ‖x₀ - x*‖²/(2·η) + K·η·G²/2:

Σ_{k<K} (f(x_k) - f(x*)) ≤ ‖x₀ - x*‖²/(2·η) + (K:ℝ)·η·G²/2.

This is the telescoped potential chain with the (nonnegative) final-distance term dropped (CLRS §33.3).

lemma sum_suboptimality_le {x₀ xstar : E} {η G : ℝ} {f : E → ℝ} (hf : ConvexOn ℝ Set.univ f) (hd : ∀ y, DifferentiableAt ℝ f y) (hη : 0 < η) (hG : ∀ y, ‖gradient f y‖ ≤ G) (K : ℕ) : (∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar)) ≤ ‖x₀ - xstar‖ ^ 2 / (2 * η) + (K : ℝ) * η * G ^ 2 / 2 := by have hpot : ‖gdIterates η x₀ f K - xstar‖ ^ 2 ≤ ‖x₀ - xstar‖ ^ 2 - 2 * η * (∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar)) + (K : ℝ) * η ^ 2 * G ^ 2 := gdIterates_potential_le hf hd (le_of_lt hη) hG K have h0 : 0 ≤ ‖gdIterates η x₀ f K - xstar‖ ^ 2 := sq_nonneg _ have hlin : 2 * η * (∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar)) ≤ ‖x₀ - xstar‖ ^ 2 + (K : ℝ) * η ^ 2 * G ^ 2 := by nlinarith [hpot, h0] have hdiv : (∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar)) ≤ (‖x₀ - xstar‖ ^ 2 + (K : ℝ) * η ^ 2 * G ^ 2) / (2 * η) := by rw [le_div_iff₀ (by positivity : 0 < 2 * η)] simpa [mul_comm] using hlin exact le_trans hdiv (by field_simp [hη.ne'] exact le_rfl)

Convergence bound (Theorem 33.8). If x* minimizes f, then the average iterate x̄ of the first K iterates satisfies

f(x̄) - f(x*) ≤ ‖x₀ - x*‖²/(2·η·K) + η·G²/2.

The proof combines Jensen's inequality for the convex f — the average suboptimality is at most the average of the suboptimalities — with the total-suboptimality bound divided by K (CLRS §33.3).

theorem avgIterate_suboptimality_le {x₀ xstar : E} {η G : ℝ} {f : E → ℝ} (hf : ConvexOn ℝ Set.univ f) (hd : ∀ y, DifferentiableAt ℝ f y) (hη : 0 < η) (hG : ∀ y, ‖gradient f y‖ ≤ G) {K : ℕ} (hK : 0 < K) : f (avgIterate η x₀ f K) - f xstar ≤ ‖x₀ - xstar‖ ^ 2 / (2 * η * K) + η * G ^ 2 / 2 := by have hKc : (K : ℝ) ≠ 0 := by exact_mod_cast (ne_of_gt hK) have hJ : f (∑ k ∈ Finset.range K, (K : ℝ)⁻¹ • gdIterates η x₀ f k) ≤ ∑ k ∈ Finset.range K, (K : ℝ)⁻¹ • f (gdIterates η x₀ f k) := by refine hf.map_sum_le ?_ ?_ ?_ · intro i hi exact inv_nonneg.mpr (Nat.cast_nonneg K) · rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul, mul_inv_cancel₀ hKc] · intro i hi exact Set.mem_univ _ have hJensen : f (avgIterate η x₀ f K) ≤ (K : ℝ)⁻¹ • (∑ k ∈ Finset.range K, f (gdIterates η x₀ f k)) := by calc f (avgIterate η x₀ f K) = f (∑ k ∈ Finset.range K, (K : ℝ)⁻¹ • gdIterates η x₀ f k) := by rw [avgIterate, Finset.smul_sum] _ ≤ ∑ k ∈ Finset.range K, (K : ℝ)⁻¹ • f (gdIterates η x₀ f k) := hJ _ = (K : ℝ)⁻¹ • (∑ k ∈ Finset.range K, f (gdIterates η x₀ f k)) := by rw [Finset.smul_sum] have hCancel : (K : ℝ)⁻¹ * ((K : ℝ) * f xstar) = f xstar := by rw [← mul_assoc, inv_mul_cancel₀ hKc] simp have hshift_in : (∑ k ∈ Finset.range K, f (gdIterates η x₀ f k)) - (K : ℝ) * f xstar = ∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar) := by rw [show (K : ℝ) * f xstar = ∑ k ∈ Finset.range K, f xstar by rw [Finset.sum_const, Finset.card_range, nsmul_eq_mul]] rw [← Finset.sum_sub_distrib] have hshift : (K : ℝ)⁻¹ • (∑ k ∈ Finset.range K, f (gdIterates η x₀ f k)) - f xstar = (K : ℝ)⁻¹ • (∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar)) := by rw [smul_eq_mul, smul_eq_mul, ← hshift_in, mul_sub, hCancel] have hmain : f (avgIterate η x₀ f K) - f xstar ≤ (K : ℝ)⁻¹ • (∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar)) := by calc f (avgIterate η x₀ f K) - f xstar ≤ (K : ℝ)⁻¹ • (∑ k ∈ Finset.range K, f (gdIterates η x₀ f k)) - f xstar := by linarith [hJensen] _ = (K : ℝ)⁻¹ • (∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar)) := hshift have hsum : (∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar)) ≤ ‖x₀ - xstar‖ ^ 2 / (2 * η) + (K : ℝ) * η * G ^ 2 / 2 := sum_suboptimality_le hf hd hη hG K have hscale : (K : ℝ)⁻¹ * (∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar)) ≤ (K : ℝ)⁻¹ * (‖x₀ - xstar‖ ^ 2 / (2 * η) + (K : ℝ) * η * G ^ 2 / 2) := by exact mul_le_mul_of_nonneg_left hsum (inv_nonneg.mpr (Nat.cast_nonneg K)) have hfinal : f (avgIterate η x₀ f K) - f xstar ≤ (K : ℝ)⁻¹ * (‖x₀ - xstar‖ ^ 2 / (2 * η) + (K : ℝ) * η * G ^ 2 / 2) := by calc f (avgIterate η x₀ f K) - f xstar ≤ (K : ℝ)⁻¹ • (∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar)) := hmain _ = (K : ℝ)⁻¹ * (∑ k ∈ Finset.range K, (f (gdIterates η x₀ f k) - f xstar)) := by rfl _ ≤ (K : ℝ)⁻¹ * (‖x₀ - xstar‖ ^ 2 / (2 * η) + (K : ℝ) * η * G ^ 2 / 2) := hscale have hsimpl : (K : ℝ)⁻¹ * (‖x₀ - xstar‖ ^ 2 / (2 * η) + (K : ℝ) * η * G ^ 2 / 2) = ‖x₀ - xstar‖ ^ 2 / (2 * η * K) + η * G ^ 2 / 2 := by field_simp [hKc, hη.ne'] exact le_trans hfinal (le_of_eq hsimpl)
end GradientDescentend CLRS

Scope and implementation notes

Imports

Current source

Section 33.1 (Clustering) is formalized natively in CLRSLean.FourthEdition.Chapter_33.Section_33_1_Clustering: the k-means clustering problem, Lloyd's algorithm, and the two monotonicity theorems of the cost under the assignment and update steps. Section 33.2 (Multiplicative-weights algorithms) is formalized natively in CLRSLean.FourthEdition.Chapter_33.Section_33_2_Multiplicative_Weights: the potential-based analysis of the multiplicative-weights update method and its regret bound against the best expert. Section 33.3 (Gradient descent) is formalized natively in CLRSLean.FourthEdition.Chapter_33.Section_33_3_Gradient_Descent: the gradient-descent method for minimizing a convex differentiable function and the convergence bound on the average iterate (Theorem 33.8). No legacy source is promoted into this chapter.

Coverage boundary

Status: complete. Represented sections: 33.1 (Clustering) — the k-means cost, its variance decomposition, Lemma 33.1 (the mean minimizes the within-cluster sum of squared distances), and Theorem 33.2 (a Lloyd iteration never increases the cost). 33.2 (Multiplicative-weights algorithms) — the multiplicative-weights update rule, the potential and expected-loss accounting, the exponential potential chain, and Theorem 33.3 (the total expected loss is within an additive ln n / η and a multiplicative (1 + η) factor of the best expert's loss). 33.3 (Gradient descent) — the gradient-descent lemma, the per-step and telescoping potential inequalities, and Theorem 33.8 (the average iterate converges with bound ‖x₀ - x*‖²/(2·η·K) + η·G²/2).

See docs/clrs-fourth-edition-map.csv for the section-level mapping and docs/migrations/clrs4.md for compatibility and deprecation policy.

CLRS, fourth edition · Chapter 33 of 35