1. Home
  2. AI & Machine Learning
  3. K-Means Clustering

K-Means Clustering

Let the computer discover groups in data by itself. Watch centroids hunt for the centre of each cluster in a 3D feature space.

Interactive 3DBeginner11 min readAI/MLUpdated

Drag to rotate · Right-drag to pan · Click, then scroll to zoom · Space play · ←→ step

What's happening

Pseudocode

    Try this in the 3D model

    • Rotate the cube before running. How many groups can you see? Set k to that number and press Run.
    • Run with a wrong k (too small or too big). What does K-Means do?
    • Press Run several times with Random points start. Do you always get the same result?
    • Switch to k-means++ and compare how many iterations it needs.

    What is clustering?

    So far, many algorithms you’ve seen are given the right answers during training (that’s supervised learning). Clustering is different: we only have raw data and we ask the computer to find natural groups in it by itself. This is unsupervised learning.

    Example: an online shop has data on thousands of customers — age, income and how much they spend. Nobody has labelled them. Clustering might discover groups like “young bargain-hunters”, “high-spending professionals” and “occasional buyers” — and the shop can target each group differently.

    In the 3D model each dot is a customer and each axis is one feature. Points that are close together are similar.

    The K-Means algorithm

    You choose k, the number of groups you want. Then:

    1. Initialise: pick k starting points called centroids (the ◆ shapes).
    2. Assign step: every point joins the cluster of its nearest centroid (straight-line distance).
    3. Update step: move every centroid to the mean (average position) of the points in its cluster.
    4. Repeat steps 2 and 3 until no point changes cluster.

    That’s it! In the model, the thin lines show which centroid each point belongs to. Watch the centroids glide towards the centre of their groups.

    Measuring quality: WCSS

    How good is a clustering? K-Means tries to minimise the Within-Cluster Sum of Squares (WCSS, also called inertia): the total squared distance from every point to its centroid.

    WCSS = Σ over clusters Σ over points in cluster  ‖point − centroid‖²

    Both steps can only lower WCSS (or keep it equal), so the algorithm is guaranteed to stop. The model shows WCSS dropping after every update.

    Choosing k — the elbow method

    K-Means can’t tell you how many groups there should be. A popular trick: run it for k = 1, 2, 3, … and plot WCSS. The curve drops sharply at first and then flattens. The “elbow” where it bends is usually a good k.

    Problems and fixes

    Problem Why Fix
    Different results on each run Random starting centroids Run several times and keep the lowest WCSS; use k-means++
    Bad start → bad result Two centroids start inside the same blob k-means++ spreads the starting centroids apart
    Features on different scales Income (thousands) dominates age (tens) Standardise features first
    Weird-shaped clusters K-Means assumes round, similar-sized blobs Use DBSCAN or Gaussian Mixture Models

    Code

    From scratch with NumPy:

    import numpy as np
    
    def k_means(X, k, iterations=100):
        rng = np.random.default_rng()
        centroids = X[rng.choice(len(X), k, replace=False)]      # random start
        for _ in range(iterations):
            # assign: distance from every point to every centroid
            dists = np.linalg.norm(X[:, None, :] - centroids[None, :, :], axis=2)
            labels = dists.argmin(axis=1)
            # update: mean of each cluster (keep old centroid if a cluster is empty)
            new = np.array([X[labels == j].mean(axis=0) if np.any(labels == j) else centroids[j]
                            for j in range(k)])
            if np.allclose(new, centroids):
                break                                             # converged
            centroids = new
        return labels, centroids
    
    X = np.vstack([np.random.randn(30, 3) + c for c in ([0, 0, 0], [6, 6, 0], [0, 6, 6])])
    labels, centroids = k_means(X, k=3)
    print(centroids.round(2))

    In practice you would use scikit-learn, which uses k-means++ by default:

    from sklearn.cluster import KMeans
    
    model = KMeans(n_clusters=3, n_init=10).fit(X)
    print(model.labels_[:10])        # cluster of the first 10 points
    print(model.inertia_)            # WCSS

    Where is K-Means used?

    • Customer segmentation for marketing.
    • Image compression: cluster the colours of an image into k colours.
    • Document grouping: news articles about similar topics.
    • Anomaly detection: points far from every centroid are suspicious.
    • As a first step for other algorithms (e.g. initialising other models).

    Common mistakes

    • Forgetting to scale the features.
    • Treating cluster numbers as meaningful: “cluster 2” is just a name; it can change between runs.
    • Expecting K-Means to find long, curved or very unequal clusters — it prefers compact round ones.

    Complexity at a glance

    Case / operationTimeWhy
    One iterationO(n · k · d)n points × k centroids × d features distances.
    Whole algorithmO(i · n · k · d)i iterations — usually small (tens).
    Extra spaceO(n · d + k · d)

    Quick check

    Test yourself — pick an answer to see if you got it.

    1. K-Means is an example of which kind of learning?

    2. In the update step, where does each centroid move?

    3. When does K-Means stop?

    4. What is the purpose of k-means++?

    Saved only in this browser — no account needed.
    Spotted a mistake or a bug in the 3D model?

    Report a mistake

    in K-Means Clustering. Thank you — every report makes the lesson better for the next reader.

    We'll also include a link to the step of the 3D model you're on and your browser type, so we can reproduce it.