๐Ÿ“ˆ Machine Learning ยท Lecture 25 of 47

k-Means Clustering: Algorithm, Objective and Pitfalls

k-means partitions data into k groups by alternating assignment and update steps. We derive it as coordinate descent, cover k-means++ initialisation, choosing k, and the assumptions that make it fail.

A humanitarian organisation wants to place a limited number of distribution points so that every household is as close as possible to one. A marketing team wants to group customers by behaviour. Both are clustering problems, and the first algorithm everyone reaches for is k-means. It is simple, fast and often effective โ€” but its assumptions are strong, and you must know when they break.

The objective#

Given points $\mathbf{x}_1, \dots, \mathbf{x}_n$ and a number of clusters $k$, find centroids $\boldsymbol{\mu}_1, \dots, \boldsymbol{\mu}_k$ and assignments $c_i \in \{1, \dots, k\}$ minimising the within-cluster sum of squares (inertia):

$$ J = \sum_{i=1}^{n}\|\mathbf{x}_i - \boldsymbol{\mu}_{c_i}\|^2 $$

Minimising $J$ exactly is NP-hard. Lloyd's algorithm finds a local minimum.

Lloyd's algorithm#

  1. Initialise $k$ centroids.
  2. Assignment step: assign each point to its nearest centroid: $c_i = \arg\min_j\|\mathbf{x}_i - \boldsymbol{\mu}_j\|^2$.
  3. Update step: move each centroid to the mean of its assigned points: $\boldsymbol{\mu}_j = \frac{1}{|C_j|}\sum_{i \in C_j}\mathbf{x}_i$.
  4. Repeat until assignments stop changing.

Why it converges: it is coordinate descent on $J$. The assignment step minimises $J$ over assignments with centroids fixed; the update step minimises $J$ over centroids with assignments fixed (the mean minimises squared distances). $J$ never increases, and there are finitely many partitions, so the algorithm terminates โ€” at a local, not necessarily global, minimum. Each iteration costs $O(nkd)$.

python
import numpy as np

def kmeans(X, k, iters=100, seed=0):
    rng = np.random.default_rng(seed)
    centroids = X[rng.choice(len(X), k, replace=False)]
    for _ in range(iters):
        d2 = ((X[:, None, :] - centroids[None]) ** 2).sum(-1)     # (n, k)
        labels = d2.argmin(1)
        new = np.array([X[labels == j].mean(0) if np.any(labels == j) else centroids[j]
                        for j in range(k)])
        if np.allclose(new, centroids):
            break
        centroids = new
    inertia = ((X - centroids[labels]) ** 2).sum()
    return labels, centroids, inertia

Initialisation matters: k-means++#

Random initial centroids can land in the same cluster, producing poor local optima. k-means++ (Arthur & Vassilvitskii, 2007) spreads them out:

  1. Choose the first centroid uniformly at random from the data.
  2. Choose each next centroid with probability proportional to $D(\mathbf{x})^2$, the squared distance to the nearest centroid already chosen.

This gives an expected objective within a factor $O(\log k)$ of optimal, and in practice converges faster and better. Scikit-learn uses k-means++ by default and runs several initialisations (n_init), keeping the best.

Choosing k#

There is no universally correct $k$. Tools:

  • Elbow method โ€” plot inertia against $k$ and look for a bend. Inertia always decreases with $k$, and the elbow is often ambiguous.
  • Silhouette score โ€” for each point, $s = \frac{b - a}{\max(a, b)}$, where $a$ is the mean distance to its own cluster and $b$ the mean distance to the nearest other cluster. Ranges from โˆ’1 to 1; higher is better.
  • Gap statistic โ€” compares inertia with that expected under a null reference distribution.
  • Domain constraints โ€” often decisive: "we can open five centres", "the marketing team can handle four segments".
python
from sklearn.cluster import KMeans
from sklearn.metrics import silhouette_score
from sklearn.datasets import make_blobs

X, _ = make_blobs(n_samples=1500, centers=4, cluster_std=1.0, random_state=7)
for k in range(2, 8):
    km = KMeans(n_clusters=k, n_init=10, random_state=0).fit(X)
    print(f"k={k}: inertia={km.inertia_:9.1f}  silhouette={silhouette_score(X, km.labels_):.3f}")

The hidden assumptions#

k-means implicitly assumes clusters are:

  1. Spherical (isotropic) โ€” it uses Euclidean distance to a centre.
  2. Similar in size and density โ€” boundaries lie halfway between centroids.
  3. Convex and well separated.

Also remember to scale features โ€” k-means is distance-based โ€” and to handle categorical variables appropriately (k-modes / k-prototypes).

Variants#

  • Mini-batch k-means โ€” updates centroids from small random batches; scales to millions of points.
  • k-medoids (PAM) โ€” centres must be actual data points; works with any distance and is more robust to outliers.
  • Spherical k-means โ€” cosine similarity for text embeddings.
  • Soft k-means / Gaussian mixtures โ€” probabilistic assignments (a later lecture).

Applications#

Customer segmentation, image colour quantisation (cluster pixel colours into a palette), document grouping, initialising other algorithms, vector quantisation (compressing embeddings โ€” the idea behind product quantisation in vector search), and facility location as in our opening example.

JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST ยท Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

๐Ÿ“ˆ Machine Learning

Hierarchical Clustering and Dendrograms

Hierarchical clustering builds a whole tree of nested clusters instead of a single partition. We compare linkage criteria, read dendrograms, and learn when a hierarchy is more useful than k flat clusters.

Beginnerโฑ 5 min#075
๐Ÿ“ˆ Machine Learning

DBSCAN and Density-Based Clustering

DBSCAN defines clusters as dense regions separated by sparse ones. It finds arbitrarily shaped clusters, labels outliers as noise and needs no k. We study core points, parameter selection and HDBSCAN.

Intermediateโฑ 5 min#076
๐Ÿ“ˆ Machine Learning

Gaussian Mixture Models and the EM Algorithm

Gaussian mixtures model data as a blend of Gaussian components with soft cluster memberships. We derive the Expectationโ€“Maximisation algorithm, prove it increases likelihood, and choose the number of components with BIC.

Advancedโฑ 5 min#077