use crate::error::FdarError;
use crate::matrix::FdMatrix;
use crate::metric::gak::{gak_gram_predict, gak_gram_train, GakConfig, GakGramTrain};
use rand::prelude::*;
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct KernelKmeansConfig {
pub n_clusters: usize,
pub n_init: usize,
pub max_iter: usize,
pub tol: f64,
pub seed: u64,
pub gak: GakConfig,
}
impl Default for KernelKmeansConfig {
fn default() -> Self {
Self {
n_clusters: 2,
n_init: 10,
max_iter: 300,
tol: 1e-4,
seed: 0,
gak: GakConfig::default(),
}
}
}
impl KernelKmeansConfig {
#[must_use]
pub fn new(n_clusters: usize, sigma: f64) -> Self {
Self {
n_clusters,
gak: GakConfig::with_sigma(sigma),
..Self::default()
}
}
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct KernelKmeansResult {
pub cluster: Vec<usize>,
pub inertia: f64,
pub iter: usize,
pub converged: bool,
pub n_init_best: usize,
train: GakGramTrain,
within: Vec<f64>,
sizes: Vec<usize>,
}
impl KernelKmeansResult {
#[must_use]
pub fn n_clusters(&self) -> usize {
self.within.len()
}
pub fn predict(&self, new_data: &FdMatrix) -> Result<Vec<usize>, FdarError> {
let kcross = gak_gram_predict(&self.train, new_data)?;
let n_test = kcross.nrows();
let n_train = kcross.ncols();
let k = self.within.len();
let mut labels = vec![0usize; n_test];
for t in 0..n_test {
let mut cross_sum = vec![0.0f64; k];
for j in 0..n_train {
cross_sum[self.cluster[j]] += kcross[(t, j)];
}
let mut best = 0usize;
let mut best_d2 = f64::INFINITY;
for (c, &sz) in self.sizes.iter().enumerate() {
let d2 = if sz == 0 {
f64::INFINITY
} else {
1.0 - (2.0 / sz as f64) * cross_sum[c] + self.within[c]
};
if d2 < best_d2 {
best_d2 = d2;
best = c;
}
}
labels[t] = best;
}
Ok(labels)
}
}
#[must_use = "expensive computation whose result should not be discarded"]
pub fn kernel_kmeans_fd(
data: &FdMatrix,
config: &KernelKmeansConfig,
) -> Result<KernelKmeansResult, FdarError> {
let n = data.nrows();
let m = data.ncols();
if n == 0 || m == 0 {
return Err(FdarError::InvalidDimension {
parameter: "data",
expected: "non-empty matrix (nrows > 0, ncols > 0)".to_string(),
actual: format!("{n}x{m}"),
});
}
let k = config.n_clusters;
if k < 1 {
return Err(FdarError::InvalidParameter {
parameter: "n_clusters",
message: "number of clusters must be >= 1".to_string(),
});
}
if k > n {
return Err(FdarError::InvalidParameter {
parameter: "n_clusters",
message: format!("n_clusters={k} exceeds number of curves n={n}"),
});
}
if config.n_init < 1 {
return Err(FdarError::InvalidParameter {
parameter: "n_init",
message: "n_init must be >= 1".to_string(),
});
}
let train = gak_gram_train(data, &config.gak)?;
let gram = &train.gram;
let mut best: Option<RestartOutcome> = None;
for restart in 0..config.n_init {
let mut rng = StdRng::seed_from_u64(config.seed.wrapping_add(restart as u64));
let outcome = run_restart(gram, n, k, config.max_iter, config.tol, &mut rng, restart);
let take = match &best {
None => true,
Some(b) => outcome.inertia < b.inertia,
};
if take {
best = Some(outcome);
}
}
let RestartOutcome {
cluster,
inertia,
iter,
converged,
within,
sizes,
restart_idx,
} = best.expect("n_init >= 1 guarantees a restart outcome");
Ok(KernelKmeansResult {
cluster,
inertia,
iter,
converged,
n_init_best: restart_idx,
train,
within,
sizes,
})
}
struct RestartOutcome {
cluster: Vec<usize>,
inertia: f64,
iter: usize,
converged: bool,
within: Vec<f64>,
sizes: Vec<usize>,
restart_idx: usize,
}
fn run_restart(
gram: &FdMatrix,
n: usize,
k: usize,
max_iter: usize,
tol: f64,
rng: &mut StdRng,
restart_idx: usize,
) -> RestartOutcome {
let mut cluster: Vec<usize> = (0..n).map(|_| rng.gen_range(0..k)).collect();
ensure_no_empty_random(&mut cluster, n, k, rng);
let mut sizes = vec![0usize; k];
let mut within = vec![0.0f64; k];
let mut d2 = vec![0.0f64; n * k]; let mut prev_inertia = f64::INFINITY;
let mut iter = 0usize;
let mut converged = false;
while iter < max_iter {
iter += 1;
compute_cluster_stats(gram, &cluster, n, k, &mut sizes, &mut within);
for i in 0..n {
let mut cross = vec![0.0f64; k];
for j in 0..n {
cross[cluster[j]] += gram[(i, j)];
}
let kii = gram[(i, i)];
for c in 0..k {
d2[i * k + c] = if sizes[c] == 0 {
f64::INFINITY
} else {
kii - (2.0 / sizes[c] as f64) * cross[c] + within[c]
};
}
}
let mut new_cluster = vec![0usize; n];
for i in 0..n {
let mut best_c = 0usize;
let mut best_d = f64::INFINITY;
for c in 0..k {
let v = d2[i * k + c];
if v < best_d {
best_d = v;
best_c = c;
}
}
new_cluster[i] = best_c;
}
recover_empty_clusters(&mut new_cluster, &d2, n, k);
let inertia: f64 = (0..n).map(|i| d2[i * k + new_cluster[i]]).sum();
let changed = new_cluster != cluster;
cluster = new_cluster;
let rel = if prev_inertia.is_finite() && prev_inertia.abs() > 0.0 {
(prev_inertia - inertia).abs() / prev_inertia.abs()
} else {
f64::INFINITY
};
if !changed || rel < tol {
converged = true;
break;
}
prev_inertia = inertia;
}
compute_cluster_stats(gram, &cluster, n, k, &mut sizes, &mut within);
let inertia = final_inertia(gram, &cluster, &sizes, &within, n, k);
RestartOutcome {
cluster,
inertia,
iter,
converged,
within,
sizes,
restart_idx,
}
}
fn compute_cluster_stats(
gram: &FdMatrix,
cluster: &[usize],
n: usize,
k: usize,
sizes: &mut [usize],
within: &mut [f64],
) {
sizes.iter_mut().for_each(|s| *s = 0);
within.iter_mut().for_each(|w| *w = 0.0);
for &c in cluster.iter() {
sizes[c] += 1;
}
let mut sums = vec![0.0f64; k];
for j in 0..n {
let cj = cluster[j];
for l in 0..n {
if cluster[l] == cj {
sums[cj] += gram[(j, l)];
}
}
}
for c in 0..k {
if sizes[c] > 0 {
let sz = sizes[c] as f64;
within[c] = sums[c] / (sz * sz);
} else {
within[c] = 0.0;
}
}
}
fn final_inertia(
gram: &FdMatrix,
cluster: &[usize],
sizes: &[usize],
within: &[f64],
n: usize,
k: usize,
) -> f64 {
let mut total = 0.0;
for i in 0..n {
let mut cross = vec![0.0f64; k];
for j in 0..n {
cross[cluster[j]] += gram[(i, j)];
}
let c = cluster[i];
if sizes[c] > 0 {
let d2 = gram[(i, i)] - (2.0 / sizes[c] as f64) * cross[c] + within[c];
total += d2;
}
}
total
}
fn ensure_no_empty_random(cluster: &mut [usize], n: usize, k: usize, rng: &mut StdRng) {
loop {
let mut sizes = vec![0usize; k];
for &c in cluster.iter() {
sizes[c] += 1;
}
let empty: Vec<usize> = (0..k).filter(|&c| sizes[c] == 0).collect();
if empty.is_empty() {
return;
}
for c in empty {
let donors: Vec<usize> = (0..n).filter(|&i| sizes[cluster[i]] > 1).collect();
if donors.is_empty() {
return;
}
let pick = donors[rng.gen_range(0..donors.len())];
sizes[cluster[pick]] -= 1;
cluster[pick] = c;
sizes[c] += 1;
}
}
}
fn recover_empty_clusters(cluster: &mut [usize], d2: &[f64], n: usize, k: usize) {
loop {
let mut sizes = vec![0usize; k];
for &c in cluster.iter() {
sizes[c] += 1;
}
let Some(empty) = (0..k).find(|&c| sizes[c] == 0) else {
return;
};
let mut best_i = None;
let mut best_d = f64::NEG_INFINITY;
for i in 0..n {
if sizes[cluster[i]] <= 1 {
continue;
}
let d = d2[i * k + cluster[i]];
if d > best_d {
best_d = d;
best_i = Some(i);
}
}
match best_i {
Some(i) => {
cluster[i] = empty;
}
None => return, }
}
}
#[cfg(test)]
mod tests {
use super::*;
fn matrix_from_rows(rows: &[Vec<f64>]) -> FdMatrix {
let n = rows.len();
let m = rows[0].len();
let mut data = vec![0.0; n * m];
for (i, r) in rows.iter().enumerate() {
for (j, &v) in r.iter().enumerate() {
data[i + j * n] = v; }
}
FdMatrix::from_slice(&data, n, m).unwrap()
}
fn two_groups() -> (FdMatrix, Vec<usize>) {
let m = 20;
let mut rows = Vec::new();
let mut truth = Vec::new();
for i in 0..5 {
let off = i as f64 * 0.01;
rows.push(
(0..m)
.map(|k| (k as f64 * 0.05).sin() * 0.2 + off)
.collect(),
);
truth.push(0);
}
for i in 0..5 {
let off = i as f64 * 0.01;
rows.push(
(0..m)
.map(|k| (k as f64 * 0.05).sin() * 0.2 + 10.0 + off)
.collect(),
);
truth.push(1);
}
(matrix_from_rows(&rows), truth)
}
fn purity(labels: &[usize], truth: &[usize], k: usize) -> f64 {
let n = labels.len();
let n_truth = truth.iter().copied().max().unwrap_or(0) + 1;
let mut correct = 0usize;
for c in 0..k {
let mut counts = vec![0usize; n_truth];
for i in 0..n {
if labels[i] == c {
counts[truth[i]] += 1;
}
}
correct += counts.iter().copied().max().unwrap_or(0);
}
correct as f64 / n as f64
}
#[test]
fn test_kernel_kmeans_recovers_groups() {
let (data, truth) = two_groups();
let cfg = KernelKmeansConfig::new(2, 1.0);
let res = kernel_kmeans_fd(&data, &cfg).unwrap();
assert_eq!(res.cluster.len(), 10);
let p = purity(&res.cluster, &truth, 2);
assert!((p - 1.0).abs() < 1e-12, "purity {p} != 1.0");
assert_eq!(res.n_clusters(), 2);
}
#[test]
fn test_kernel_kmeans_deterministic() {
let (data, _) = two_groups();
let cfg = KernelKmeansConfig::new(2, 1.0);
let a = kernel_kmeans_fd(&data, &cfg).unwrap();
let b = kernel_kmeans_fd(&data, &cfg).unwrap();
assert_eq!(a.cluster, b.cluster, "same seed must give identical labels");
assert_eq!(a.inertia.to_bits(), b.inertia.to_bits());
assert_eq!(a.n_init_best, b.n_init_best);
}
#[test]
fn test_kernel_kmeans_empty_cluster_recovery() {
let (data, _) = two_groups();
let cfg = KernelKmeansConfig {
n_clusters: 4,
..KernelKmeansConfig::new(4, 1.0)
};
let res = kernel_kmeans_fd(&data, &cfg).unwrap();
assert_eq!(res.cluster.len(), 10);
assert!(res.cluster.iter().all(|&c| c < 4));
let mut sizes = vec![0usize; 4];
for &c in &res.cluster {
sizes[c] += 1;
}
assert!(
sizes.iter().all(|&s| s >= 1),
"an empty cluster survived: {sizes:?}"
);
}
#[test]
fn test_kernel_kmeans_empty_cluster_k_equals_n() {
let rows: Vec<Vec<f64>> = (0..4)
.map(|i| (0..12).map(|k| (k as f64 * 0.1 + i as f64).sin()).collect())
.collect();
let data = matrix_from_rows(&rows);
let cfg = KernelKmeansConfig::new(4, 1.0);
let res = kernel_kmeans_fd(&data, &cfg).unwrap();
assert_eq!(res.cluster.len(), 4);
assert!(res.cluster.iter().all(|&c| c < 4));
}
#[test]
fn test_kernel_kmeans_n_init() {
let (data, _) = two_groups();
let multi = KernelKmeansConfig {
n_init: 10,
..KernelKmeansConfig::new(2, 1.0)
};
let single = KernelKmeansConfig {
n_init: 1,
..KernelKmeansConfig::new(2, 1.0)
};
let rm = kernel_kmeans_fd(&data, &multi).unwrap();
let rs = kernel_kmeans_fd(&data, &single).unwrap();
assert!(
rm.inertia <= rs.inertia + 1e-12,
"multi-init inertia {} worse than single-init {}",
rm.inertia,
rs.inertia
);
}
#[test]
fn test_kernel_kmeans_predict() {
let (data, _) = two_groups();
let cfg = KernelKmeansConfig::new(2, 1.0);
let res = kernel_kmeans_fd(&data, &cfg).unwrap();
let low_label = res.cluster[0];
let high_label = res.cluster[5];
assert_ne!(low_label, high_label);
let m = 20;
let low_curve: Vec<f64> = (0..m)
.map(|k| (k as f64 * 0.05).sin() * 0.2 + 0.03)
.collect();
let high_curve: Vec<f64> = (0..m)
.map(|k| (k as f64 * 0.05).sin() * 0.2 + 10.03)
.collect();
let copy0 = data.row(0);
let test = matrix_from_rows(&[low_curve, high_curve, copy0]);
let preds = res.predict(&test).unwrap();
assert_eq!(preds.len(), 3);
assert_eq!(
preds[0], low_label,
"novel low curve should route to low cluster"
);
assert_eq!(
preds[1], high_label,
"novel high curve should route to high cluster"
);
assert_eq!(
preds[2], res.cluster[0],
"exact copy should match its training label"
);
}
#[test]
fn test_kernel_kmeans_validation() {
let (data, _) = two_groups();
let cfg0 = KernelKmeansConfig {
n_clusters: 0,
..KernelKmeansConfig::new(0, 1.0)
};
assert!(matches!(
kernel_kmeans_fd(&data, &cfg0),
Err(FdarError::InvalidParameter { .. })
));
let cfg_big = KernelKmeansConfig {
n_clusters: 999,
..KernelKmeansConfig::new(999, 1.0)
};
assert!(matches!(
kernel_kmeans_fd(&data, &cfg_big),
Err(FdarError::InvalidParameter { .. })
));
let cfg_ni = KernelKmeansConfig {
n_init: 0,
..KernelKmeansConfig::new(2, 1.0)
};
assert!(matches!(
kernel_kmeans_fd(&data, &cfg_ni),
Err(FdarError::InvalidParameter { .. })
));
let empty = FdMatrix::zeros(0, 0);
assert!(matches!(
kernel_kmeans_fd(&empty, &KernelKmeansConfig::new(2, 1.0)),
Err(FdarError::InvalidDimension { .. })
));
}
#[test]
fn test_kernel_kmeans_no_centroid() {
let (data, _) = two_groups();
let res = kernel_kmeans_fd(&data, &KernelKmeansConfig::new(2, 1.0)).unwrap();
let KernelKmeansResult {
cluster,
inertia,
iter,
converged,
n_init_best,
..
} = &res;
assert_eq!(cluster.len(), 10);
assert!(inertia.is_finite());
assert!(*iter >= 1);
let _ = converged;
let _ = n_init_best;
let preds = res.predict(&data).unwrap();
assert_eq!(preds, res.cluster);
}
}