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.Tactic33.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
simpThe 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.
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.
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) jThe 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.centroidThe 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 CLRSImports
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.Tactic33.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)whereMis the day's expected loss. -
Theorem
totalExpectedLoss_le: the total expected loss of the algorithm is within an additiveln n / ηand a multiplicative(1 + η)factor of the best expert's loss — for every experti,Σₜ 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 expertion dayt(∈ [0, 1]) -
w: a weight vector (Fin n → ℝ) -
Φ: the potential, the sum of the weights -
Mᵗ: the algorithm's expected loss on dayt
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 : ℝ} (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 [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 hgoalend MultiplicativeWeightsend CLRSImports
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.Tactic33.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 fromx₀. -
Definition
avgIterate: the arithmetic mean of the firstKiterates. -
Theorem
gradient_inner_le_sub: the gradient-descent lemma — for a convex differentiablef, the first-order characterization of convexity⟪∇f(x), y - x⟫ ≤ f(y) - f(x)holds for allx, y. -
Theorem
gdStep_potential_le: one gradient-descent step shrinks the squared distance to any pointx*by at least2η·(f(x) - f(x*)), up to the additiveη²·G²term coming from a bound‖∇f‖ ≤ Gon the gradient norm. -
Theorem
gdIterates_potential_le/sum_suboptimality_le: the telescoping potential chain overKsteps and the resulting total-suboptimality bound. -
Theorem
avgIterate_suboptimality_le(Theorem 33.8): the convergence bound — ifx*minimizesf, then the averagex̄of the firstKiterates satisfiesf(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 functionE → ℝto be minimized -
η: the learning rate (0 < η) -
x₀: the starting point -
x*: a point minimizingf -
G: a bound on the norm of the gradient,‖∇f(x)‖ ≤ G -
x̄: the average of the firstKiterates
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⁺.
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] <;> 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).
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
abel
_ = ‖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 CLRSScope 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