Implement K-Means and compare with GMM
Company: Tubi
Role: Machine Learning Engineer
Category: Coding & Algorithms
Difficulty: medium
Interview Round: Technical Screen
Implement K-Means clustering from scratch for a dataset `X` of shape `(n_samples, n_features)` and a target number of clusters `k`.
Your implementation should:
- initialize `k` centroids,
- assign each point to its nearest centroid,
- recompute centroids as the mean of assigned points,
- repeat until convergence or a maximum number of iterations,
- handle edge cases such as empty clusters,
- and include a test or check for convergence, such as unchanged assignments, centroid movement below a tolerance, or non-increasing within-cluster loss.
After coding, answer these follow-up questions:
- How does Gaussian Mixture Modeling differ from K-Means?
- How would you train a Gaussian mixture model using the EM algorithm?
- In what situations would you prefer GMM over K-Means?
Overview: This question evaluates understanding and practical implementation of clustering algorithms, specifically K-Means centroid initialization, assignment and update mechanics, convergence criteria and edge-case handling, along with comparative knowledge of Gaussian Mixture Models and the Expectation-Maximization training framework.
Community answers
Answer by tenbyliu
import numpy as np
class KMeans:
def init(
self,
k,
max_iters=300,
tol=1e-4,
random_state=None,
):
if k <= 0:
raise ValueError("k must be a positive integer.")
self.k = k
self.max_iters = max_iters
self.tol = tol
self.rng = np.random.default_rng(random_state)
self.centroids_ = None
self.labels_ = None
self.inertia_ = None
self.n_iter_ = 0
def _initialize_centroids(self, X):
"""Initialize centroids using k-means++."""
n_samples = X.shape[0]
centroids = np.empty((self.k, X.shape[1]), dtype=float)
# Choose the first centroid randomly.
first_index = self.rng.integers(n_samples)
centroids[0] = X[first_index]
# Choose later centroids with probability proportional to
# squared distance from the nearest existing centroid.
closest_dist_sq = np.sum((X - centroids[0]) ** 2, axis=1)
for i in range(1, self.k):
total = closest_dist_sq.sum()
if total == 0:
# All points are identical to an existing centroid.
index = self.rng.integers(n_samples)
else:
probabilities = closest_dist_sq / total
index = self.rng.choice(n_samples, p=probabilities)
centroids[i] = X[index]
new_dist_sq = np.sum((X - centroids[i]) ** 2, axis=1)
closest_dist_sq = np.minimum(closest_dist_sq, new_dist_sq)
return centroids
@staticmethod
def _assign_clusters(X, centroids):
# Shape: (n_samples, k)
squared_distances = np.sum(
(X[:, np.newaxis, :] - centroids[np.newaxis, :, :]) ** 2,
axis=2,
)
return np.argmin(squared_distances, axis=1)
def _recompute_centroids(self, X, labels, old_centroids):
centroids = np.empty_like(old_cen