use crate::matrix::knn::metric::{l2_sq, sqdist_soa_range};
use nalgebra::DMatrix;
use rand::rngs::SmallRng;
use rand::{RngExt, SeedableRng};
use rayon::prelude::*;
const CHUNK_ROWS: usize = 4096;
const MIN_DIRECTION_NORM: f64 = 1e-12;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum KmeansMetric {
Euclidean,
Cosine,
}
#[derive(Debug, Clone)]
pub struct KmeansRowsOpts {
pub k: usize,
pub max_iter: usize,
pub seed: u64,
pub metric: KmeansMetric,
pub min_changed_frac: f64,
pub init_sample: usize,
}
pub struct KmeansRowsFit {
pub centroids: DMatrix<f32>,
pub labels: Vec<usize>,
pub n_iter: usize,
}
struct Partial {
sums: Vec<f64>,
counts: Vec<usize>,
changed: usize,
}
pub fn kmeans_rows_seeded(z: &DMatrix<f32>, opts: &KmeansRowsOpts) -> KmeansRowsFit {
let (n, d, k) = (z.nrows(), z.ncols(), opts.k);
if k <= 1 || n == 0 || d == 0 {
return single_centroid(z, k);
}
let cosine = opts.metric == KmeansMetric::Cosine;
let mut zt = z.transpose();
if cosine {
zt.as_mut_slice().par_chunks_mut(d).for_each(normalise_row);
}
let rows: &[f32] = zt.as_slice();
let mut cents = kmeans_pp_init(rows, n, d, k, opts.seed, opts.init_sample);
let mut labels = vec![usize::MAX; n];
let mut nearest = vec![0f32; n];
let mut n_iter = 0usize;
for _ in 0..opts.max_iter.max(1) {
n_iter += 1;
let cents_soa = to_soa(¢s, k, d);
let Partial {
mut sums,
mut counts,
changed,
} = assign_all(rows, d, ¢s_soa, k, cosine, &mut labels, &mut nearest);
let mut next = vec![0f32; k * d];
for c in 0..k {
if counts[c] == 0 {
continue;
}
let inv = 1.0 / counts[c] as f64;
let mean = &mut sums[c * d..(c + 1) * d];
for m in mean.iter_mut() {
*m *= inv;
}
if cosine {
let norm = mean.iter().map(|v| v * v).sum::<f64>().sqrt();
if norm < MIN_DIRECTION_NORM {
counts[c] = 0;
continue;
}
for m in mean.iter_mut() {
*m /= norm;
}
}
for (dst, &m) in next[c * d..(c + 1) * d].iter_mut().zip(mean.iter()) {
*dst = m as f32;
}
}
reseed_empty_clusters(&mut next, &counts, &mut nearest, rows, d);
cents = next;
if changed == 0 || (changed as f64) < opts.min_changed_frac * n as f64 {
break;
}
}
KmeansRowsFit {
centroids: DMatrix::from_row_slice(k, d, ¢s),
labels,
n_iter,
}
}
pub fn nearest_centroid_rows(
z: &DMatrix<f32>,
centroids: &DMatrix<f32>,
metric: KmeansMetric,
) -> Vec<usize> {
let (n, d, k) = (z.nrows(), z.ncols(), centroids.nrows());
assert_eq!(
centroids.ncols(),
d,
"centroids and rows disagree on the dimension"
);
if n == 0 || k == 0 || d == 0 {
return vec![0; n];
}
let cosine = metric == KmeansMetric::Cosine;
let mut zt = z.transpose();
if cosine {
zt.as_mut_slice().par_chunks_mut(d).for_each(normalise_row);
}
let mut labels = vec![usize::MAX; n];
let mut nearest = vec![0f32; n];
assign_all(
zt.as_slice(),
d,
centroids.as_slice(),
k,
cosine,
&mut labels,
&mut nearest,
);
labels
}
fn single_centroid(z: &DMatrix<f32>, k: usize) -> KmeansRowsFit {
let (n, d) = (z.nrows(), z.ncols());
let mut centroids = DMatrix::<f32>::zeros(k.max(1), d);
if n > 0 {
for j in 0..d {
let s: f64 = z.column(j).iter().map(|&v| f64::from(v)).sum();
centroids[(0, j)] = (s / n as f64) as f32;
}
}
KmeansRowsFit {
centroids,
labels: vec![0; n],
n_iter: 0,
}
}
fn normalise_row(row: &mut [f32]) {
let norm = row.iter().map(|v| v * v).sum::<f32>().sqrt();
if norm > 0.0 {
row.iter_mut().for_each(|v| *v /= norm);
}
}
fn kmeans_pp_init(
rows: &[f32],
n: usize,
d: usize,
k: usize,
seed: u64,
init_sample: usize,
) -> Vec<f32> {
let mut rng = SmallRng::seed_from_u64(seed);
let pool: Option<Vec<usize>> = (init_sample > 0 && init_sample < n).then(|| {
let mut ids = rand::seq::index::sample(&mut rng, n, init_sample).into_vec();
ids.sort_unstable();
ids
});
let m = pool.as_ref().map_or(n, Vec::len);
let row_of = |p: usize| -> &[f32] {
let r = pool.as_ref().map_or(p, |ids| ids[p]);
&rows[r * d..(r + 1) * d]
};
let mut cents = vec![0f32; k * d];
let mut d2 = vec![f32::INFINITY; m];
let first = rng.random_range(0..m);
cents[..d].copy_from_slice(row_of(first));
fold_min_sqdist(&mut d2, ¢s[..d], &row_of);
for c in 1..k {
let sum: f64 = d2.iter().map(|&x| f64::from(x)).sum();
let pick = if sum > 0.0 {
let target = rng.random_range(0.0f64..sum);
let mut acc = 0f64;
let mut idx = m - 1;
for (i, &w) in d2.iter().enumerate() {
acc += f64::from(w);
if acc >= target {
idx = i;
break;
}
}
idx
} else {
rng.random_range(0..m)
};
cents[c * d..(c + 1) * d].copy_from_slice(row_of(pick));
fold_min_sqdist(&mut d2, ¢s[c * d..(c + 1) * d], &row_of);
}
cents
}
fn fold_min_sqdist<'a>(
d2: &mut [f32],
cent: &[f32],
row_of: &(impl Fn(usize) -> &'a [f32] + Sync),
) {
d2.par_iter_mut().enumerate().for_each(|(p, slot)| {
let dd = l2_sq(row_of(p), cent);
if dd < *slot {
*slot = dd;
}
});
}
fn to_soa(cents: &[f32], k: usize, d: usize) -> Vec<f32> {
(0..d)
.flat_map(|dim| (0..k).map(move |c| cents[c * d + dim]))
.collect()
}
fn assign_all(
rows: &[f32],
d: usize,
cents_soa: &[f32],
k: usize,
cosine: bool,
labels: &mut [usize],
nearest: &mut [f32],
) -> Partial {
let partials: Vec<Partial> = rows
.par_chunks(CHUNK_ROWS * d)
.zip(labels.par_chunks_mut(CHUNK_ROWS))
.zip(nearest.par_chunks_mut(CHUNK_ROWS))
.map(|((xs, ls), ns)| assign_chunk(xs, d, cents_soa, k, cosine, ls, ns))
.collect();
let mut total = Partial {
sums: vec![0f64; k * d],
counts: vec![0usize; k],
changed: 0,
};
for p in &partials {
for (s, &v) in total.sums.iter_mut().zip(&p.sums) {
*s += v;
}
for (c, &v) in total.counts.iter_mut().zip(&p.counts) {
*c += v;
}
total.changed += p.changed;
}
total
}
fn assign_chunk(
xs: &[f32],
d: usize,
cents_soa: &[f32],
k: usize,
cosine: bool,
labels: &mut [usize],
nearest: &mut [f32],
) -> Partial {
let mut sums = vec![0f64; k * d];
let mut counts = vec![0usize; k];
let mut changed = 0usize;
let mut dist = vec![0f32; k];
for (r, x) in xs.chunks_exact(d).enumerate() {
sqdist_soa_range(cents_soa, k, x, 0, k, &mut dist);
let mut best = 0usize;
let mut best_d = f32::INFINITY;
for (c, &dd) in dist.iter().enumerate() {
if dd < best_d {
best_d = dd;
best = c;
}
}
nearest[r] = if cosine && x.iter().all(|&v| v == 0.0) {
0.0
} else {
best_d
};
if labels[r] != best {
changed += 1;
labels[r] = best;
}
counts[best] += 1;
for (s, &v) in sums[best * d..(best + 1) * d].iter_mut().zip(x) {
*s += f64::from(v);
}
}
Partial {
sums,
counts,
changed,
}
}
fn reseed_empty_clusters(
next: &mut [f32],
counts: &[usize],
nearest: &mut [f32],
rows: &[f32],
d: usize,
) {
for (c, &count) in counts.iter().enumerate() {
if count > 0 {
continue;
}
let (_, src) = nearest.par_iter().enumerate().map(|(i, &v)| (v, i)).reduce(
|| (f32::NEG_INFINITY, usize::MAX),
|a, b| {
if b.0 > a.0 || (b.0 == a.0 && b.1 < a.1) {
b
} else {
a
}
},
);
if src == usize::MAX {
continue;
}
next[c * d..(c + 1) * d].copy_from_slice(&rows[src * d..(src + 1) * d]);
nearest[src] = 0.0;
}
}
pub fn kmeans_centroids_seeded(
z: &DMatrix<f32>,
k: usize,
max_iter: usize,
seed: u64,
) -> (DMatrix<f32>, Vec<usize>) {
let fit = kmeans_rows_seeded(
z,
&KmeansRowsOpts {
k,
max_iter,
seed,
metric: KmeansMetric::Euclidean,
min_changed_frac: 0.0,
init_sample: 0,
},
);
(fit.centroids, fit.labels)
}
pub fn kmeans_centroids(z: &DMatrix<f32>, k: usize, max_iter: usize) -> (DMatrix<f32>, Vec<usize>) {
kmeans_centroids_seeded(z, k, max_iter, 42)
}
#[cfg(test)]
mod tests;