HAKO upload: hako/core/kmeanspp.py
Browse files- hako/core/kmeanspp.py +89 -0
hako/core/kmeanspp.py
ADDED
|
@@ -0,0 +1,89 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""K-Means++ seeding + Lloyd refinement.
|
| 2 |
+
|
| 3 |
+
Theorem T-KM (existence, monotonicity, approximation).
|
| 4 |
+
(1) Boundedness: the frozen source features satisfy ||z|| <= C (layer-normed
|
| 5 |
+
activations / bounded quantized weights), hence the k-means cost
|
| 6 |
+
J_KM({mu}) = sum_i min_k ||z_i - mu_k||^2
|
| 7 |
+
is finite and attains its infimum over the compact set
|
| 8 |
+
{mu : ||mu_k|| <= 2C} (Weierstrass).
|
| 9 |
+
(2) Monotone decrease: the assignment step (point -> nearest centroid) never
|
| 10 |
+
increases J_KM for fixed centroids; the update step (centroid -> mean of
|
| 11 |
+
its assigned points) is the exact argmin of J_KM for a fixed assignment
|
| 12 |
+
(d/dmu_k sum ||z - mu_k||^2 = 0 gives mu_k = mean). Therefore every Lloyd
|
| 13 |
+
sweep weakly decreases J_KM. There are finitely many assignments (K^N),
|
| 14 |
+
no assignment repeats under strict decrease, so Lloyd converges to a
|
| 15 |
+
local minimum in finitely many sweeps. QED.
|
| 16 |
+
(3) Approximation: k-means++ seeding (D^2 sampling) yields expected cost
|
| 17 |
+
E[J_seed] <= 8 (ln K + 2) * OPT
|
| 18 |
+
(Arthur & Vassilvitskii, 2007); Lloyd afterwards only improves it by (2).
|
| 19 |
+
|
| 20 |
+
Implementation is numpy-vectorized; K seeds by D^2 sampling with O(N K d)
|
| 21 |
+
total work, then a fixed number of Lloyd sweeps (early-stop on delta).
|
| 22 |
+
"""
|
| 23 |
+
from __future__ import annotations
|
| 24 |
+
|
| 25 |
+
from typing import Tuple
|
| 26 |
+
|
| 27 |
+
import numpy as np
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def _d2_sample(X: np.ndarray, K: int, rng: np.random.Generator) -> np.ndarray:
|
| 31 |
+
"""K-means++ D^2 seeding. Returns (K, d) initial centroids."""
|
| 32 |
+
n = X.shape[0]
|
| 33 |
+
cent = np.empty((K, X.shape[1]), dtype=np.float32)
|
| 34 |
+
cent[0] = X[rng.integers(n)]
|
| 35 |
+
d2 = ((X - cent[0]) ** 2).sum(axis=1)
|
| 36 |
+
for k in range(1, K):
|
| 37 |
+
total = d2.sum()
|
| 38 |
+
if total <= 1e-12:
|
| 39 |
+
cent[k:] = X[rng.integers(n, size=K - k)]
|
| 40 |
+
break
|
| 41 |
+
probs = d2 / total
|
| 42 |
+
idx = rng.choice(n, p=probs)
|
| 43 |
+
cent[k] = X[idx]
|
| 44 |
+
d2 = np.minimum(d2, ((X - cent[k]) ** 2).sum(axis=1))
|
| 45 |
+
return cent
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def kmeanspp_lloyd(X: np.ndarray, K: int, n_init: int = 1, sweeps: int = 24,
|
| 49 |
+
seed: int = 7,
|
| 50 |
+
tol: float = 1e-6) -> Tuple[np.ndarray, np.ndarray, float]:
|
| 51 |
+
"""Returns (centroids (K,d) float32, labels (n,), inertia float).
|
| 52 |
+
|
| 53 |
+
Theorem T-KM(2) is enforced structurally: we only accept a sweep result
|
| 54 |
+
if inertia is non-increasing (guards float noise with a tolerance).
|
| 55 |
+
"""
|
| 56 |
+
assert K >= 1 and X.ndim == 2
|
| 57 |
+
X = np.ascontiguousarray(X, dtype=np.float32)
|
| 58 |
+
n = X.shape[0]
|
| 59 |
+
rng = np.random.default_rng(seed)
|
| 60 |
+
best = None
|
| 61 |
+
for init in range(max(1, n_init)):
|
| 62 |
+
cent = _d2_sample(X, K, rng)
|
| 63 |
+
prev_J = np.inf
|
| 64 |
+
labels = np.zeros(n, dtype=np.int64)
|
| 65 |
+
for _ in range(sweeps):
|
| 66 |
+
# assignment (chunked to bound RAM)
|
| 67 |
+
chunk = max(1, (64 * 1024 * 1024) // max(1, 4 * K * X.shape[1]))
|
| 68 |
+
mins = np.empty(n, dtype=np.float32)
|
| 69 |
+
for s in range(0, n, chunk):
|
| 70 |
+
e = min(n, s + chunk)
|
| 71 |
+
d2 = ((X[s:e, None, :] - cent[None, :, :]) ** 2).sum(axis=2)
|
| 72 |
+
labels[s:e] = d2.argmin(axis=1)
|
| 73 |
+
mins[s:e] = d2.min(axis=1)
|
| 74 |
+
J = float(mins.sum())
|
| 75 |
+
if prev_J - J < tol * max(1.0, prev_J):
|
| 76 |
+
prev_J = min(prev_J, J)
|
| 77 |
+
break
|
| 78 |
+
prev_J = J
|
| 79 |
+
# update
|
| 80 |
+
for k in range(K):
|
| 81 |
+
m = labels == k
|
| 82 |
+
if m.any():
|
| 83 |
+
cent[k] = X[m].mean(axis=0)
|
| 84 |
+
else: # dead centroid -> reseed at farthest point
|
| 85 |
+
far = ((X - cent[labels]) ** 2).sum(axis=1)
|
| 86 |
+
cent[k] = X[int(far.argmax())]
|
| 87 |
+
if best is None or prev_J < best[2]:
|
| 88 |
+
best = (cent.copy(), labels.copy(), prev_J)
|
| 89 |
+
return best
|