K-means à grande échelle : où il craque et que faire
K-means semble simple jusqu'à un milliard de points. Découvrez comment les algorithmes modernes résolvent mémoire, vitesse et initialisation.

K-means est l'algorithme que tout le monde découvre dans son cours d'introduction au ML, puis utilise mal pendant le reste de sa carrière. Il est trompeusement simple : choisir K centres de clusters, affecter chaque point au centre le plus proche, déplacer les centres vers la moyenne des points qui leur sont affectés, et recommencer. Trois lignes de pseudo-code. Le problème, c'est que cette simplicité cache une série de pièges qui explosent dès qu'on l'applique à des données réelles à grande échelle.
À 10 000 points, k-means termine en quelques millisecondes et personne ne se soucie des détails d'implémentation. À 10 millions de points, le choix de la méthode d'initialisation peut changer le temps d'exécution d'un facteur 100. À 1 milliard de points, l'algorithme standard ne tient plus en mémoire et il faut adopter des approches fondamentalement différentes. Un article récent sur Flash-KMeans s'attaque exactement à ce problème : il obtient des résultats exacts de k-means avec une mémoire fortement réduite et une convergence plus rapide. Voyons ce qui rend k-means difficile à grande échelle et comment les variantes modernes le résolvent.
Ce que fait vraiment le K-means standard
L'algorithme de Lloyd, ce qu'on entend généralement par « k-means », comporte deux étapes par itération. L'étape d'affectation : pour chaque point, on calcule sa distance à tous les K centroïdes et on l'affecte au plus proche. L'étape de mise à jour : on recalcule chaque centroïde comme la moyenne de tous les points qui lui sont affectés.
import numpy as np
def kmeans_lloyd(X, K, max_iter=100):
n, d = X.shape
# Random initialization (bad, but we'll fix this later)
centroids = X[np.random.choice(n, K, replace=False)]
for _ in range(max_iter):
# Assignment: O(n * K * d) — this is the bottleneck
distances = np.linalg.norm(X[:, None] - centroids[None, :], axis=2)
labels = np.argmin(distances, axis=1)
# Update: O(n * d)
new_centroids = np.array([
X[labels == k].mean(axis=0) for k in range(K)
])
if np.allclose(centroids, new_centroids):
break
centroids = new_centroids
return labels, centroids
L'étape d'affectation coûte O(n × K × d) par itération, où n est le nombre de points, K le nombre de clusters et d la dimensionnalité. Pour n = 1 milliard, K = 1000, d = 128 (un scénario réaliste de clustering d'embeddings), cela représente 128 000 milliards d'opérations en virgule flottante par itération. Même avec une capacité de calcul d'un téraflop, cela fait 128 secondes par itération, et k-means a généralement besoin de 10 à 50 itérations pour converger.
Le problème de l'initialisation
Avant que k-means n'exécute la moindre itération, il doit choisir la position initiale des centroïdes. Ce choix compte beaucoup plus que la plupart des gens ne le pensent : une mauvaise initialisation peut mener à une solution arbitrairement moins bonne que l'optimum.
L'initialisation aléatoire (choisir K points de données au hasard comme centroïdes initiaux) est simple mais peu fiable. Si deux centroïdes initiaux tombent par hasard dans le même cluster, un autre cluster restera sans représentant. L'algorithme convergera, mais vers une solution sous-optimale. Avec K = 100, la probabilité d'au moins une collision est élevée.
K-means++ (Arthur et Vassilvitskii, 2007) est la solution standard. Il choisit le premier centroïde au hasard, puis chaque centroïde suivant avec une probabilité proportionnelle à sa distance au carré au centroïde existant le plus proche. Les points éloignés de tout centroïde existant ont plus de chances d'être choisis, ce qui répartit les centroïdes initiaux dans les données. On obtient ainsi une garantie d'approximation en O(log K) : la solution est prouvablement à un facteur O(log K) de l'optimum.
def kmeans_plus_plus_init(X, K):
n, d = X.shape
centroids = [X[np.random.randint(n)]]
for _ in range(1, K):
# Distance from each point to nearest existing centroid
dists = np.min([
np.sum((X - c) ** 2, axis=1) for c in centroids
], axis=0)
# Sample proportional to distance squared
probs = dists / dists.sum()
next_idx = np.random.choice(n, p=probs)
centroids.append(X[next_idx])
return np.array(centroids)
Le piège : l'initialisation k-means++ elle-même coûte O(n × K × d), car elle nécessite K passes sur l'ensemble des données. Pour un K et un n importants, l'initialisation prend plus de temps que plusieurs itérations de l'algorithme de Lloyd. Des variantes passant à l'échelle, comme k-means|| (Bahmani et al., 2012), réduisent le nombre de passes en sur-échantillonnant des candidats en parallèle, puis en les consolidant en K centroïdes.
Le mur de la mémoire
Le k-means standard a besoin de l'ensemble des données en mémoire simultanément. Pour l'étape d'affectation, il faut accéder à chaque point et le comparer à chaque centroïde. Avec 1 milliard de vecteurs float32 de 128 dimensions, les seules données occupent 512 Go, soit plus que ce que possèdent la plupart des machines.
La solution naïve est le mini-batch k-means : à chaque itération, on tire un sous-ensemble aléatoire de points et on met à jour les centroïdes à partir de cet échantillon. Cela fonctionne, mais converge plus lentement et vers une solution légèrement différente (généralement un peu moins bonne) que le k-means exact. Pour beaucoup d'applications, la perte de qualité est acceptable. Pour d'autres, comme les codebooks de quantification dans la recherche vectorielle, la différence compte.
Flash-KMeans adopte une approche différente. Plutôt que de traiter l'ensemble des données en mémoire, il utilise une passe en flux avec un cache intelligent. L'idée clé : la plupart des points ne changent pas de cluster d'une itération à l'autre. Si un point appartient fermement au cluster 7 (bien plus proche du centroïde 7 que de tout autre centroïde), calculer les K distances est du travail inutile. En maintenant des bornes sur les distances, Flash-KMeans peut éviter le calcul complet pour la plupart des points, à la plupart des itérations.
Optimisations par l'inégalité triangulaire
L'algorithme d'Elkan (2003) utilise l'inégalité triangulaire pour éviter des calculs de distance inutiles. Elle dit que la distance entre le point P et le centroïde A est au plus égale à la distance de P au centroïde B plus la distance entre B et A. Si l'on sait que P est actuellement affecté au centroïde B, et que l'on connaît la distance entre A et B, on peut parfois démontrer que A est trop loin de P sans calculer la distance réelle.
En pratique, cela élimine 80 à 95 % des calculs de distance après les premières itérations, lorsque la plupart des points sont déjà proches de leur bon centroïde. Les calculs restants concernent les points proches des frontières entre clusters, les seuls susceptibles de changer réellement d'affectation.
Le compromis : l'algorithme d'Elkan nécessite O(n × K) de mémoire supplémentaire pour stocker les bornes de distance (bornes inférieures de chaque point vers chaque centroïde, plus bornes supérieures vers le centroïde affecté). Pour un K élevé, ce coût mémoire peut devenir considérable. L'algorithme de Hamerly le réduit à O(n) en ne conservant qu'une seule borne inférieure par point, au prix de moins de calculs évités.
Accélération GPU
Le k-means est embarrassingly parallel dans l'étape d'affectation : le calcul de distance de chaque point est indépendant. C'est donc un candidat naturel pour l'accélération GPU. La bibliothèque cuML de NVIDIA et FAISS de Facebook proposent toutes deux des implémentations GPU de k-means, qui atteignent un gain de 10 à 50 fois par rapport aux implémentations CPU sur les grands jeux de données.
Le problème, c'est que la mémoire GPU est limitée. Un A100 dispose de 80 Go de mémoire, de quoi stocker environ 150 millions de vecteurs de 128 dimensions. Les jeux de données plus volumineux nécessitent soit un partitionnement sur plusieurs GPU, soit une approche en flux où les données sont chargées par blocs, traitées sur le GPU, puis les résultats agrégés sur le CPU.
# Using FAISS for GPU-accelerated k-means
import faiss
import numpy as np
# 10 million 128-dimensional vectors
n, d, K = 10_000_000, 128, 1000
X = np.random.randn(n, d).astype('float32')
# CPU k-means (for comparison)
kmeans_cpu = faiss.Kmeans(d, K, niter=20, verbose=True)
kmeans_cpu.train(X) # ~120 seconds
# GPU k-means (single GPU)
kmeans_gpu = faiss.Kmeans(d, K, niter=20, verbose=True, gpu=True)
kmeans_gpu.train(X) # ~8 seconds — 15x faster
Quand K-Means est le mauvais choix
Avant d'optimiser votre implémentation de k-means, demandez-vous si c'est vraiment le bon algorithme. Il repose sur des hypothèses fortes qui ne tiennent pas pour de nombreux jeux de données réels.
- Clusters sphériques. K-means suppose que les clusters sont à peu près sphériques et de taille similaire. Si vos clusters sont allongés, de forme irrégulière ou de tailles très différentes, k-means découpera les gros clusters et fusionnera les petits. Les modèles de mélange gaussien (GMM) gèrent les clusters elliptiques, et DBSCAN gère les formes arbitraires.
- K connu. Il faut spécifier le nombre de clusters à l'avance. Si vous ne connaissez pas K, il faut lancer k-means plusieurs fois avec différentes valeurs et utiliser une métrique (score de silhouette, méthode du coude) pour choisir la meilleure. Cela multiplie le coût de calcul total par le nombre de valeurs de K testées.
- Distance euclidienne. K-means utilise par défaut la distance euclidienne. Pour les embeddings de texte, la similarité cosinus est généralement plus appropriée. On peut contourner le problème en normalisant les vecteurs en L2 (ce qui rend la distance euclidienne équivalente à la distance cosinus), mais c'est facile à oublier.
- Sensibilité aux valeurs aberrantes. Un seul point aberrant, éloigné de tout cluster, attirera vers lui le centroïde auquel il est affecté. Des variantes robustes comme k-medoids (qui utilise la médiane au lieu de la moyenne) gèrent mieux les aberrations, mais coûtent plus cher.
Conseils pratiques
Ayant utilisé k-means sur des jeux de données allant de quelques milliers à des milliards de points, voici ce que j'ai appris de plus important :
- Utilisez toujours l'initialisation k-means++. L'initialisation aléatoire ne vaut jamais le risque. L'écart de qualité finale des clusters est souvent de 10 à 30 %, et k-means++ ajoute un surcoût négligeable pour des K petits à moyens.
- Lancez plusieurs exécutions. K-means trouve un optimum local, pas global. Lancez-le 5 à 10 fois avec différentes graines aléatoires et gardez le meilleur résultat (somme des distances intra-cluster la plus faible). C'est une assurance peu coûteuse contre les mauvaises exécutions.
- Normalisez vos features. Si une feature a une plage [0, 1000000] et une autre [0, 1], la feature à grande plage dominera le calcul des distances. Standardisez (soustraire la moyenne, diviser par l'écart type) ou normalisez en min-max avant le clustering.
- Utilisez FAISS pour le clustering à grande échelle. Si vous avez plus d'un million de points, le k-means de scikit-learn sera lent. L'implémentation de FAISS est très optimisée et prend en charge l'accélération GPU dès le départ.
- Envisagez des méthodes approximatives pour l'exploration. Le mini-batch k-means est 10 fois plus rapide que le k-means exact et donne des résultats généralement suffisants pour explorer. Utilisez le k-means exact pour le clustering en production, là où la qualité compte.
K-means fait partie de ces algorithmes faciles à utiliser, difficiles à bien utiliser, et qui valent la peine d'être compris en profondeur. L'écart entre une implémentation naïve et une implémentation optimisée, en vitesse comme en qualité de résultat, est énorme. À petite échelle, rien de tout cela n'a d'importance. À l'échelle où cela compte, comprendre l'initialisation, la gestion de la mémoire, l'élagage des distances et l'accélération GPU fait la différence entre un clustering qui prend des heures et un clustering qui prend des minutes.