Skip to article frontmatterSkip to article content
Site not loading correctly?

This may be due to an incorrect BASE_URL configuration. See the MyST Documentation for reference.

k-means Clustering

k-means Clustering

Mahmood Amintoosi, Fall 2026

Computer Science Dept, Ferdowsi University of Mashhad

Definition of Clustering

A clustering of a set of datapoints {x1,x2,…,xn}\{x_1, x_2, \ldots, x_n\} is a partition of the datapoints into k disjoint subsets (or clusters) C={C1,C2,…,Ck}\mathcal{C} = \{C_1, C_2, \ldots, C_k\} in such a way that datapoints in the same group are more similar (in some specific sense defined by the analyst) to each other than to those in other groups (clusters).

Representative-based Clustering

Given a dataset with nn points in a dd-dimensional space, D={xi}i=1n\textbf{D} = \{\textbf{x}_i\}_{i=1}^n, and given the number of desired clusters kk, the goal of representative-based clustering is to partition the dataset into kk groups or clusters, which is called a clustering and is denoted as C={C1,C2,…,Ck}\mathcal{C} = \{C_1, C_2, \ldots, C_k\}.

For each cluster CiC_i there exists a representative point that summarizes the cluster, a common choice being the mean (also called the centroid) μi{\mu}_i of all points in the cluster, that is,

μi=1ni∑xj∈Cixj\begin{align*} {\mu_i} = {1\over n_i} \sum_{x_{j} \in C_i} \textbf{x}_{j} \end{align*}

where ni=∣Ci∣n_i = |C_i| is the number of points in cluster CiC_i.

k-means Clustering

k-means clustering is Representative-based Clustering that partitions the datapoints into k clusters based on their similarity to the centroid of each cluster. The goal of k-means clustering is to find the partition C\mathcal{C} that minimizes the sum of squared distances between each datapoint and the centroid of its assigned cluster.

Lemma 1: Sum of Squared Distances to Any Point

Let {x1,x2,…,xn}\{x_1, x_2, \ldots, x_n\} be a set of datapoints and let μ\mu be the centroid of the datapoints. The sum of squared distances of the datapoints to any arbitrary point zz equals the sum of squared distances to the centroid plus the number of datapoints times the squared distance from the point zz to the centroid. That is,

∑i=1n∣xi−z∣2=∑i=1n∣xi−μ∣2+n∣μ−z∣2\sum_{i=1}^n |x_i - z|^2 = \sum_{i=1}^n |x_i - \mu|^2 + n |\mu - z|^2

Proof:

∑i=1n∣xi−z∣2=∑i=1n∣xi−μ+μ−z∣2=∑i=1n∣xi−μ∣2+2(μ−z)⋅∑i=1n(xi−μ)+n∣μ−z∣2=∑i=1n∣xi−μ∣2+n∣μ−z∣2\begin{align*} \sum_{i=1}^n |x_i - z|^2 &= \sum_{i=1}^n |x_i - \mu + \mu - z|^2 \\ &= \sum_{i=1}^n |x_i - \mu|^2 + 2(\mu - z) \cdot \sum_{i=1}^n (x_i - \mu) + n |\mu - z|^2 \\ &= \sum_{i=1}^n |x_i - \mu|^2 + n |\mu - z|^2 \end{align*}

since ∑i=1n(xi−μ)=0\sum_{i=1}^n (x_i - \mu) = 0.

Corollary 1: Centroid Minimizes Sum of Squared Distances

The centroid minimizes the sum of squared distances since the second term, n∣μ−z∣2n |\mu - z|^2, is always positive.

Proof:

This follows directly from Lemma 1, since the second term, n∣μ−z∣2n |\mu - z|^2, is always positive.

Lemma 2: Sum of Squared Distances Between All Pairs of Points

Let {x1,x2,…,xn}\{x_1, x_2, \ldots, x_n\} be a set of datapoints and let μ\mu be the centroid of the datapoints. The sum of squared distances between all pairs of points equals the number of points times the sum of squared distances of the points to the centroid of the points. That is,

∑i=1n∑j>i∣xi−xj∣2=n∑i=1n∣xi−μ∣2\sum_{i=1}^n \sum_{j>i} |x_i - x_j|^2 = n \sum_{i=1}^n |x_i - \mu|^2

Proof:

∑i=1n∑j>i∣xi−xj∣2=12∑i=1n∑j=1n∣xi−xj∣2=12∑j=1n(∑i=1n∣xi−xj∣2)=12∑j=1n(∑i=1n∣xi−μ∣2+n∣μ−xj∣2)=12(∑j=1n(∑i=1n∣xi−μ∣2)+n∑j=1n(∣μ−xj∣2))=12(n∑i=1n∣xi−μ∣2)+n∑i=1n∣xi−μ∣2))=n∑i=1n∣xi−μ∣2\begin{align*} \sum_{i=1}^n \sum_{j>i} |x_i - x_j|^2 &= \frac{1}{2} \sum_{i=1}^n \sum_{j=1}^n |x_i - x_j|^2 \\ &= \frac{1}{2} \sum_{j=1}^n (\sum_{i=1}^n |x_i - x_j|^2) \\ &= \frac{1}{2} \sum_{j=1}^n (\sum_{i=1}^n |x_i - \mu|^2 + n |\mu - x_j|^2) \tag{Lemma 1}\\ &= \frac{1}{2} \left(\sum_{j=1}^n (\sum_{i=1}^n |x_i - \mu|^2) + n \sum_{j=1}^n(|\mu - x_j|^2)\right)\\ &= \frac{1}{2} \left(n\sum_{i=1}^n |x_i - \mu|^2) + n\sum_{i=1}^n |x_i - \mu|^2)\right)\\ &= n \sum_{i=1}^n |x_i - \mu|^2 \end{align*}

k-means Clustering Algorithm

The k-means clustering algorithm starts with k initial centroids and iteratively updates the centroids and assigns each datapoint to its nearest centroid until convergence.

  1. Initialize k initial centroids μ1,…,μk\mu_1, \ldots, \mu_k randomly.

  2. For each iteration, perform the following steps:

    1. Assign each datapoint xix_i to the cluster CjC_j with the nearest centroid μj\mu_j.

    2. Update the centroids μ1,…,μk\mu_1, \ldots, \mu_k as the mean of all datapoints assigned to each cluster.

  3. Repeat step 2 until convergence.

Convergence of k-means Algorithm

The k-means algorithm always converges, but possibly to a local minimum. To show convergence, we argue that the cost of the clustering, the sum of the squares of the distances of each datapoint to its cluster centroid, always improves.

Here is a Python implementation of the k-means clustering algorithm:

Using OOP for implememntation of k-means

This notebook first generates some sample data and defines a KMeansClustering class that implements the K-Means algorithm. The fit method initializes the centroids randomly and then iteratively updates them until convergence.

The following script creates an instance of the KMeansClustering class, fits the model to the data, and plots the data with centroids and labels.

<Figure size 640x480 with 1 Axes>

Now, let’s use the Lemma 2 to prove that the centroid minimizes the sum of squared distances

∑i=1n∑j>i∣xi−xj∣2=n∑i=1n∣xi−μ∣2\sum_{i=1}^n \sum_{j>i} |x_i - x_j|^2 = n \sum_{i=1}^n |x_i - \mu|^2

function lemma_2 implements the Lemma 2. The lemma states that the sum of squared distances of all data points is equal to the number of data points times the squared distance from the centroid to the mean of the data points. The notebook tests the lemma on a subset of the data (cluster 0).

1812.8736866847603 1812.8736866847605