बड़े पैमाने पर K-Means: यह क्यों टूटता है और आगे क्या?
एक अरब डेटा पॉइंट्स पर K-means क्यों फेल होता है? जानिए आधुनिक एल्गोरिदम मेमोरी, स्पीड और इनिशियलाइज़ेशन की समस्याएं कैसे सुलझाते हैं।

K-means वह एल्गोरिदम है जिसे हर कोई इंट्रो ML कोर्स में सीखता है और फिर करियर भर गलत तरीके से इस्तेमाल करता है। यह सुनने में बेहद सरल है: K क्लस्टर सेंटर चुनो, हर पॉइंट को उसके सबसे नज़दीकी सेंटर को असाइन करो, सेंटर्स को उनके पॉइंट्स के माध्य पर खिसकाओ, और दोहराओ। बस तीन लाइन का pseudocode। दिक्कत यह है कि यह सादगी कई बारूदी सुरंगें छुपाए रखती है, जो असली डेटा पर बड़े पैमाने पर लगाने पर फट पड़ती हैं।
10,000 पॉइंट्स पर k-means मिलीसेकंड में खत्म हो जाता है और कोई इम्प्लीमेंटेशन की डिटेल्स की फिक्र नहीं करता। 1 करोड़ पॉइंट्स पर इनिशियलाइज़ेशन का तरीका रनटाइम को 100 गुना तक बदल सकता है। 1 अरब पॉइंट्स पर स्टैंडर्ड एल्गोरिदम मेमोरी में फिट ही नहीं होता, और आपको बिल्कुल अलग तरीकों की ज़रूरत पड़ती है। Flash-KMeans पर एक हालिया पेपर ठीक इसी समस्या पर काम करता है, जो कम मेमोरी और तेज़ कन्वर्जेंस के साथ सटीक k-means नतीजे देता है। चलिए देखते हैं कि बड़े पैमाने पर k-means को कठिन क्या बनाता है और आधुनिक वेरिएंट इसे कैसे हल करते हैं।
स्टैंडर्ड K-Means असल में क्या करता है
Lloyd's algorithm, जिसे लोग 'k-means' कहते हैं, हर इटरेशन में दो चरण होते हैं। असाइनमेंट चरण: हर डेटा पॉइंट के लिए सभी K सेंट्रॉइड्स की दूरी निकालो और उसे सबसे नज़दीकी सेंट्रॉइड को सौंप दो। अपडेट चरण: हर सेंट्रॉइड को उसे सौंपे गए सभी पॉइंट्स के माध्य के रूप में दोबारा गणना करो।
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
असाइनमेंट चरण की लागत हर इटरेशन में O(n × K × d) है, जहां n पॉइंट्स की संख्या, K क्लस्टर की संख्या और d डायमेंशन है। n = 1 अरब, K = 1000, d = 128 (एम्बेडिंग क्लस्टरिंग का एक असली परिदृश्य) के लिए यह प्रति इटरेशन 128 ट्रिलियन फ्लोटिंग-पॉइंट ऑपरेशन है। एक टेराफ्लॉप कंप्यूट पर भी यह प्रति इटरेशन 128 सेकंड है, और k-means आमतौर पर कन्वर्ज होने में 10-50 इटरेशन लेता है।
इनिशियलाइज़ेशन की समस्या
k-means एक भी इटरेशन चलाने से पहले सेंट्रॉइड्स की शुरुआती पोज़ीशन चुनता है। यह चुनाव ज़्यादातर लोगों की सोच से कहीं ज़्यादा मायने रखता है। खराब इनिशियलाइज़ेशन ऐसे समाधान पर कन्वर्ज करवा सकता है जो optimal से मनमाने ढंग से खराब हो।
रैंडम इनिशियलाइज़ेशन (शुरुआती सेंट्रॉइड्स के रूप में K रैंडम डेटा पॉइंट्स चुनना) सरल है, पर भरोसेमंद नहीं। अगर दो शुरुआती सेंट्रॉइड्स एक ही क्लस्टर में गिर जाएं, तो एक क्लस्टर का प्रतिनिधित्व ही नहीं होगा। एल्गोरिदम कन्वर्ज तो हो जाएगा, पर सबऑप्टिमल समाधान पर। K = 100 पर कम से कम एक टकराव की संभावना काफी ज़्यादा होती है।
K-means++ (Arthur और Vassilvitskii, 2007) इसका स्टैंडर्ड हल है। यह पहला सेंट्रॉइड रैंडम चुनता है, और फिर हर अगला सेंट्रॉइड उसकी मौजूदा सबसे नज़दीकी सेंट्रॉइड से वर्ग-दूरी के अनुपात में चुनता है। जो पॉइंट्स किसी मौजूदा सेंट्रॉइड से दूर हैं, उनके चुने जाने की संभावना ज़्यादा होती है, जिससे शुरुआती सेंट्रॉइड्स पूरे डेटा में फैल जाते हैं। इससे O(log K) अनुमान की गारंटी मिलती है, यानी समाधान सिद्ध रूप से optimal के O(log K) के भीतर रहता है।
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)
पकड़ यह है: k-means++ इनिशियलाइज़ेशन खुद O(n × K × d) है, यानी इसे पूरे डेटासेट से K बार गुज़रना पड़ता है। बड़े K और बड़े n पर इनिशियलाइज़ेशन में Lloyd's एल्गोरिदम के कई इटरेशन से भी ज़्यादा समय लग सकता है। k-means|| (Bahmani et al., 2012) जैसे स्केलेबल वेरिएंट उम्मीदवारों को समानांतर में ओवर-सैंपल करके पास की संख्या घटाते हैं, और फिर उन्हें K सेंट्रॉइड्स में समेट देते हैं।
मेमोरी की दीवार
स्टैंडर्ड k-means को पूरा डेटासेट एक साथ मेमोरी में चाहिए। असाइनमेंट चरण के लिए आपको हर डेटा पॉइंट तक पहुंचकर उसे हर सेंट्रॉइड से तुलना करनी होती है। 1 अरब 128-डायमेंशनल float32 वेक्टर्स का डेटा अकेले 512 GB लेता है, जो ज़्यादातर सिंगल मशीनों की क्षमता से ज़्यादा है।
सबसे आसान हल है mini-batch k-means: हर इटरेशन में पॉइंट्स का रैंडम सबसेट सैंपल करो और उसी के आधार पर सेंट्रॉइड्स अपडेट करो। यह काम करता है, पर धीरे कन्वर्ज होता है और सटीक k-means से थोड़े अलग (आमतौर पर थोड़े खराब) समाधान पर पहुंचता है। कई एप्लिकेशन्स के लिए क्वालिटी का यह नुकसान स्वीकार्य है। पर कुछ के लिए, जैसे वेक्टर सर्च में quantization codebooks, यह फर्क मायने रखता है।
Flash-KMeans अलग तरीका अपनाता है। पूरे डेटासेट को मेमोरी में प्रोसेस करने के बजाय यह समझदार कैशिंग के साथ स्ट्रीमिंग पास का उपयोग करता है। मुख्य समझ यह है कि ज़्यादातर पॉइंट्स इटरेशन के बीच अपना क्लस्टर नहीं बदलते। अगर कोई पॉइंट पक्के तौर पर क्लस्टर 7 में है (सेंट्रॉइड 7 के बहुत नज़दीक, बाकी किसी से कहीं ज़्यादा), तो सभी K दूरियां निकालना बेकार काम है। दूरी की सीमाएं (distance bounds) बनाए रखकर Flash-KMeans ज़्यादातर इटरेशन में ज़्यादातर पॉइंट्स के लिए पूरी दूरी गणना छोड़ सकता है।
ट्रायंगल इनइक्वालिटी ऑप्टिमाइज़ेशन
Elkan का एल्गोरिदम (2003) ट्रायंगल इनइक्वालिटी का उपयोग करके अनावश्यक दूरी गणनाएं छोड़ देता है। ट्रायंगल इनइक्वालिटी कहती है: पॉइंट P से सेंट्रॉइड A की दूरी, P से सेंट्रॉइड B की दूरी और B से A की दूरी के योग से ज़्यादा नहीं हो सकती। अगर आपको पता है कि P अभी सेंट्रॉइड B को असाइन है, और सेंट्रॉइड A व B के बीच की दूरी भी पता है, तो आप कभी-कभी बिना वास्तविक दूरी निकाले सिद्ध कर सकते हैं कि A, P से बहुत दूर है।
व्यवहार में, शुरुआती कुछ इटरेशन के बाद, जब ज़्यादातर पॉइंट्स पहले से अपने सही सेंट्रॉइड के पास होते हैं, यह 80-95% दूरी गणनाओं को खत्म कर देता है। बची हुई गणनाएं उन पॉइंट्स के लिए हैं जो क्लस्टर की सीमाओं के पास हैं, यानी वे जो असल में असाइनमेंट बदल सकते हैं।
ट्रेड-ऑफ यह है: Elkan के एल्गोरिदम को दूरी सीमाएं स्टोर करने के लिए अतिरिक्त O(n × K) मेमोरी चाहिए (हर पॉइंट से हर सेंट्रॉइड तक निचली सीमाएं, और असाइन किए गए सेंट्रॉइड तक ऊपरी सीमा)। बड़े K के लिए यह मेमोरी खर्च काफी हो सकता है। Hamerly का एल्गोरिदम हर पॉइंट के लिए सिर्फ एक निचली सीमा रखकर इसे O(n) मेमोरी तक ले आता है, पर इसमें प्रून होने वाली गणनाएं कम होती हैं।
GPU एक्सेलरेशन
असाइनमेंट चरण में k-means स्वाभाविक रूप से एम्बैरेसिंगली पैरेलल है, क्योंकि हर पॉइंट की दूरी गणना दूसरे से स्वतंत्र है। इसलिए यह GPU एक्सेलरेशन के लिए बिल्कुल फिट बैठता है। NVIDIA की cuML लाइब्रेरी और Facebook की FAISS, दोनों में GPU k-means इम्प्लीमेंटेशन हैं, जो बड़े डेटासेट पर CPU इम्प्लीमेंटेशन से 10-50 गुना स्पीडअप देते हैं।
पेच यह है कि GPU मेमोरी सीमित है। A100 में 80 GB मेमोरी है, जो 128-डायमेंशनल लगभग 15 करोड़ वेक्टर्स के लिए काफी है। इससे बड़े डेटासेट के लिए या तो मल्टी-GPU पार्टीशनिंग चाहिए, या स्ट्रीमिंग तरीका जिसमें डेटा चंक्स में लोड होकर GPU पर प्रोसेस होता है और नतीजे 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
कब K-Means गलत चुनाव है
अपने k-means इम्प्लीमेंटेशन को ऑप्टिमाइज़ करने से पहले सोचिए कि क्या k-means सही एल्गोरिदम है। यह कुछ मज़बूत मान्यताएं लेकर चलता है, जो कई असली डेटासेट पर सही नहीं बैठतीं।
- गोलाकार क्लस्टर। K-means मानता है कि क्लस्टर लगभग गोलाकार और बराबर आकार के हैं। अगर आपके क्लस्टर लंबे, अनियमित आकार के हैं या उनके आकार में बहुत फर्क है, तो k-means बड़े क्लस्टर को तोड़ेगा और छोटे को मिला देगा। Gaussian Mixture Models (GMMs) दीर्घवृत्ताकार क्लस्टर संभालते हैं। DBSCAN किसी भी आकार को संभालता है।
- K पहले से ज्ञात होना। क्लस्टरों की संख्या पहले से बतानी पड़ती है। अगर आपको K नहीं पता, तो अलग-अलग मानों के साथ k-means कई बार चलाना होगा और सिल्हूट स्कोर या एल्बो मेथड जैसे मेट्रिक से सबसे अच्छा चुनना होगा। इससे कुल कंप्यूट उतने K मानों की संख्या से गुना हो जाता है जितने आप आज़माते हैं।
- यूक्लिडियन दूरी। K-means डिफ़ॉल्ट रूप से यूक्लिडियन दूरी इस्तेमाल करता है। टेक्स्ट एम्बेडिंग के लिए आमतौर पर cosine similarity बेहतर होती है। आप अपने वेक्टर्स को L2-normalize करके इसका हल निकाल सकते हैं (जिससे यूक्लिडियन दूरी cosine distance के बराबर हो जाती है), पर इसे भूलना आसान है।
- आउटलायर के प्रति संवेदनशीलता। किसी क्लस्टर से दूर एक अकेला आउटलायर उसके सेंट्रॉइड को अपनी ओर खींच लेता है। k-medoids जैसे मज़बूत वेरिएंट (जो mean की जगह median इस्तेमाल करते हैं) आउटलायर्स को बेहतर संभालते हैं, पर वे महंगे हैं।
व्यावहारिक सलाह
हज़ारों से लेकर अरबों पॉइंट्स के डेटासेट पर k-means इस्तेमाल करने के बाद, मेरे अनुभव से ये बातें सबसे ज़्यादा मायने रखती हैं:
- हमेशा k-means++ इनिशियलाइज़ेशन इस्तेमाल करें। रैंडम इनिशियलाइज़ेशन का जोखिम कभी लायक नहीं। अंतिम क्लस्टर क्वालिटी में अक्सर 10-30% का फर्क पड़ता है, और छोटे से मध्यम K के लिए k-means++ का अतिरिक्त खर्च नगण्य है।
- कई बार चलाएं। K-means एक लोकल ऑप्टिमम खोजता है, ग्लोबल नहीं। अलग-अलग रैंडम सीड्स के साथ 5-10 बार चलाएं और सबसे अच्छा नतीजा चुनें (कुल वितरण दूरी सबसे कम)। खराब रन के खिलाफ यह सस्ता बीमा है।
- फीचर्स को नॉर्मलाइज़ करें। अगर एक फीचर की रेंज [0, 1000000] है और दूसरे की [0, 1], तो बड़ी रेंज वाला फीचर दूरी गणना पर हावी हो जाएगा। क्लस्टरिंग से पहले स्टैंडर्डाइज़ करें (माध्य घटाएं, मानक विचलन से भाग दें) या min-max नॉर्मलाइज़ करें।
- बड़े पैमाने की क्लस्टरिंग के लिए FAISS इस्तेमाल करें। अगर आपके पास दस लाख से ज़्यादा पॉइंट्स हैं, तो scikit-learn का k-means धीमा होगा। FAISS का इम्प्लीमेंटेशन बहुत ऑप्टिमाइज़्ड है और बिना अतिरिक्त सेटअप के GPU एक्सेलरेशन सपोर्ट करता है।
- एक्सप्लोरेटरी काम के लिए अनुमानित तरीकों पर विचार करें। Mini-batch k-means सटीक k-means से लगभग 10 गुना तेज़ है और एक्सप्लोरेशन के लिए अक्सर काफी अच्छे नतीजे देता है। प्रोडक्शन क्लस्टरिंग में, जहां क्वालिटी मायने रखती है, सटीक k-means इस्तेमाल करें।
K-means उन एल्गोरिदम में से है जिन्हें इस्तेमाल करना आसान है, पर सही तरीके से इस्तेमाल करना कठिन, और जिन्हें गहराई से समझना सार्थक है। नेव इम्प्लीमेंटेशन और ऑप्टिमाइज़्ड इम्प्लीमेंटेशन के बीच का फासला, स्पीड और नतीजों की क्वालिटी दोनों में, बहुत बड़ा है। छोटे पैमाने पर इनमें से कुछ मायने नहीं रखता। जिस पैमाने पर ये मायने रखते हैं, वहां इनिशियलाइज़ेशन, मेमोरी मैनेजमेंट, दूरी प्रूनिंग और GPU एक्सेलरेशन को समझना ही उस क्लस्टरिंग के बीच फर्क है जो घंटों लेती है और उसके बीच जो मिनटों में हो जाती है।