use crate::alignment::{kmedoids_from_distances, KMedoidsConfig, KMedoidsResult};
use crate::error::FdarError;
use crate::matrix::FdMatrix;
use crate::metric::sbd::{sbd, sbd_distance_matrix};
use crate::shapelet::z_normalize_window;
use nalgebra::{DMatrix, SymmetricEigen};
use rand::prelude::*;
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct KShapeConfig {
pub n_clusters: usize,
pub n_init: usize,
pub max_iter: usize,
pub tol: f64,
pub seed: u64,
}
impl Default for KShapeConfig {
fn default() -> Self {
Self {
n_clusters: 2,
n_init: 10,
max_iter: 100,
tol: 1e-6,
seed: 0,
}
}
}
impl KShapeConfig {
#[must_use]
pub fn new(n_clusters: usize) -> Self {
Self {
n_clusters,
..Self::default()
}
}
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct KShapeResult {
pub centroids: FdMatrix,
pub cluster: Vec<usize>,
pub inertia: f64,
pub iter: usize,
pub converged: bool,
pub n_init_best: usize,
}
impl KShapeResult {
#[must_use]
pub fn centroids(&self) -> &FdMatrix {
&self.centroids
}
#[must_use]
pub fn cluster(&self) -> &[usize] {
&self.cluster
}
#[must_use]
pub fn inertia(&self) -> f64 {
self.inertia
}
#[must_use]
pub fn n_clusters(&self) -> usize {
self.centroids.nrows()
}
pub fn predict(&self, new_data: &FdMatrix) -> Result<Vec<usize>, FdarError> {
let p = new_data.nrows();
let m = new_data.ncols();
let k = self.centroids.nrows();
let cm = self.centroids.ncols();
if p == 0 || m == 0 {
return Err(FdarError::InvalidDimension {
parameter: "new_data",
expected: "non-empty matrix (nrows > 0, ncols > 0)".to_string(),
actual: format!("{p}x{m}"),
});
}
if m != cm {
return Err(FdarError::InvalidDimension {
parameter: "new_data",
expected: format!("series length m={cm} matching fitted centroids"),
actual: format!("m={m}"),
});
}
let mut centroid_rows: Vec<Vec<f64>> = Vec::with_capacity(k);
for c in 0..k {
centroid_rows.push(self.centroids.row(c));
}
let mut labels = vec![0usize; p];
let mut row = vec![0.0f64; m];
for t in 0..p {
new_data.row_to_buf(t, &mut row);
let z = z_normalize_window(&row);
let mut best = 0usize;
let mut best_d = f64::INFINITY;
for (c, cent) in centroid_rows.iter().enumerate() {
let d = sbd(&z, cent).map(|r| r.distance).unwrap_or(1.0);
if d < best_d {
best_d = d;
best = c;
}
}
labels[t] = best;
}
Ok(labels)
}
}
#[must_use = "expensive computation whose result should not be discarded"]
pub fn kshape_fd(data: &FdMatrix, config: &KShapeConfig) -> Result<KShapeResult, 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 series n={n}"),
});
}
if config.n_init < 1 {
return Err(FdarError::InvalidParameter {
parameter: "n_init",
message: "n_init must be >= 1".to_string(),
});
}
let mut series: Vec<Vec<f64>> = Vec::with_capacity(n);
let mut row = vec![0.0f64; m];
for i in 0..n {
data.row_to_buf(i, &mut row);
series.push(z_normalize_window(&row));
}
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(
&series,
n,
m,
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,
centroids,
inertia,
iter,
converged,
restart_idx,
} = best.expect("n_init >= 1 guarantees a restart outcome");
let mut cmat = FdMatrix::zeros(k, m);
for (c, cent) in centroids.iter().enumerate() {
for (j, &v) in cent.iter().enumerate() {
cmat[(c, j)] = v;
}
}
Ok(KShapeResult {
centroids: cmat,
cluster,
inertia,
iter,
converged,
n_init_best: restart_idx,
})
}
#[must_use = "expensive computation whose result should not be discarded"]
pub fn sbd_kmedoids(data: &FdMatrix, config: &KMedoidsConfig) -> Result<KMedoidsResult, FdarError> {
let dist = sbd_distance_matrix(data)?;
kmedoids_from_distances(&dist, config)
}
struct RestartOutcome {
cluster: Vec<usize>,
centroids: Vec<Vec<f64>>,
inertia: f64,
iter: usize,
converged: bool,
restart_idx: usize,
}
#[allow(clippy::too_many_arguments)]
fn run_restart(
series: &[Vec<f64>],
n: usize,
m: 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 centroids: Vec<Vec<f64>> = vec![vec![0.0f64; m]; k];
refine_centroids(series, &cluster, k, m, &mut centroids);
let mut prev_inertia = f64::INFINITY;
let mut iter = 0usize;
let mut converged = false;
let mut inertia = f64::INFINITY;
while iter < max_iter {
iter += 1;
let mut new_cluster = vec![0usize; n];
let mut dist_to_own = vec![0.0f64; n];
for i in 0..n {
let mut best_c = 0usize;
let mut best_d = f64::INFINITY;
for (c, cent) in centroids.iter().enumerate() {
let d = sbd(&series[i], cent).map(|r| r.distance).unwrap_or(1.0);
if d < best_d {
best_d = d;
best_c = c;
}
}
new_cluster[i] = best_c;
dist_to_own[i] = best_d;
}
recover_empty_clusters(&mut new_cluster, &dist_to_own, n, k);
refine_centroids(series, &new_cluster, k, m, &mut centroids);
inertia = 0.0;
for i in 0..n {
let d = sbd(&series[i], ¢roids[new_cluster[i]])
.map(|r| r.distance)
.unwrap_or(1.0);
inertia += d;
}
let changed = new_cluster != cluster;
cluster = new_cluster;
if !changed || (prev_inertia - inertia).abs() < tol {
converged = true;
break;
}
prev_inertia = inertia;
}
RestartOutcome {
cluster,
centroids,
inertia,
iter,
converged,
restart_idx,
}
}
fn refine_centroids(
series: &[Vec<f64>],
cluster: &[usize],
k: usize,
m: usize,
centroids: &mut [Vec<f64>],
) {
for c in 0..k {
let members: Vec<usize> = (0..series.len()).filter(|&i| cluster[i] == c).collect();
if members.is_empty() {
continue;
}
centroids[c] = shape_extraction(series, &members, ¢roids[c], m);
}
}
fn shape_extraction(
series: &[Vec<f64>],
members: &[usize],
centroid: &[f64],
m: usize,
) -> Vec<f64> {
let n_k = members.len();
let mut x_aligned: Vec<Vec<f64>> = Vec::with_capacity(n_k);
for &i in members {
let shift = sbd(centroid, &series[i]).map(|r| r.shift).unwrap_or(0);
let shifted = circular_shift(&series[i], shift);
x_aligned.push(z_normalize_window(&shifted));
}
let mut s = DMatrix::<f64>::zeros(m, m);
for row_vec in &x_aligned {
for a in 0..m {
let va = row_vec[a];
if va == 0.0 {
continue;
}
for b in 0..m {
s[(a, b)] += va * row_vec[b];
}
}
}
let inv_m = 1.0 / m as f64;
let mut qs = s.clone();
for b in 0..m {
let mut col_mean = 0.0;
for a in 0..m {
col_mean += s[(a, b)];
}
col_mean *= inv_m;
for a in 0..m {
qs[(a, b)] -= col_mean;
}
}
let mut mmat = qs.clone();
for a in 0..m {
let mut row_mean = 0.0;
for b in 0..m {
row_mean += qs[(a, b)];
}
row_mean *= inv_m;
for b in 0..m {
mmat[(a, b)] -= row_mean;
}
}
for a in 0..m {
for b in (a + 1)..m {
let avg = 0.5 * (mmat[(a, b)] + mmat[(b, a)]);
mmat[(a, b)] = avg;
mmat[(b, a)] = avg;
}
}
let eig = SymmetricEigen::new(mmat);
let mut arg = 0usize;
let mut best_eval = f64::NEG_INFINITY;
for (i, &ev) in eig.eigenvalues.iter().enumerate() {
if ev > best_eval {
best_eval = ev;
arg = i;
}
}
let mut v: Vec<f64> = eig.eigenvectors.column(arg).iter().copied().collect();
let neg: Vec<f64> = v.iter().map(|x| -x).collect();
let mut sum_pos = 0.0;
let mut sum_neg = 0.0;
for row_vec in &x_aligned {
sum_pos += sbd(&v, row_vec).map(|r| r.distance).unwrap_or(1.0);
sum_neg += sbd(&neg, row_vec).map(|r| r.distance).unwrap_or(1.0);
}
if sum_neg < sum_pos {
v = neg;
}
z_normalize_window(&v)
}
fn circular_shift(x: &[f64], shift: isize) -> Vec<f64> {
let n = x.len();
if n == 0 {
return Vec::new();
}
let n_i = n as isize;
let s = ((shift % n_i) + n_i) % n_i; let mut out = vec![0.0f64; n];
for (i, &v) in x.iter().enumerate() {
let j = ((i as isize + s) % n_i) as usize;
out[j] = v;
}
out
}
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], dist_to_own: &[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 = dist_to_own[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::*;
use std::f64::consts::PI;
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 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
}
fn shifted_groups(seed: u64) -> (FdMatrix, Vec<usize>) {
let m = 40usize;
let mut rng = StdRng::seed_from_u64(seed);
let mut rows = Vec::new();
let mut truth = Vec::new();
let base_a: Vec<f64> = (0..m)
.map(|j| (2.0 * PI * j as f64 / m as f64).sin())
.collect();
let base_b: Vec<f64> = (0..m)
.map(|j| (4.0 * PI * j as f64 / m as f64).sin())
.collect();
for (label, base) in [(0usize, &base_a), (1usize, &base_b)] {
for _ in 0..8 {
let shift = rng.gen_range(0..m) as isize;
let shifted = circular_shift(base, shift);
let noisy: Vec<f64> = shifted
.iter()
.map(|&v| v + (rng.gen::<f64>() - 0.5) * 0.05)
.collect();
rows.push(noisy);
truth.push(label);
}
}
(matrix_from_rows(&rows), truth)
}
#[test]
fn test_kshape_recovers_shifted_groups() {
let (data, truth) = shifted_groups(7);
let cfg = KShapeConfig {
n_clusters: 2,
n_init: 10,
seed: 3,
..Default::default()
};
let res = kshape_fd(&data, &cfg).unwrap();
assert_eq!(res.cluster.len(), 16);
let p = purity(&res.cluster, &truth, 2);
assert!((p - 1.0).abs() < 1e-12, "purity {p} != 1.0");
assert_eq!(res.n_clusters(), 2);
for c in 0..2 {
let row = res.centroids.row(c);
let mean: f64 = row.iter().sum::<f64>() / row.len() as f64;
assert!(mean.abs() < 1e-8, "centroid {c} not zero-mean: {mean}");
}
}
#[test]
fn test_kshape_centroid_sign() {
let m = 32usize;
let base: Vec<f64> = (0..m)
.map(|j| (2.0 * PI * j as f64 / m as f64).sin())
.collect();
let mut rng = StdRng::seed_from_u64(1);
let rows: Vec<Vec<f64>> = (0..6)
.map(|_| {
base.iter()
.map(|&v| v + (rng.gen::<f64>() - 0.5) * 0.01)
.collect()
})
.collect();
let data = matrix_from_rows(&rows);
let cfg = KShapeConfig::new(1);
let res = kshape_fd(&data, &cfg).unwrap();
let cent = res.centroids.row(0);
let base_z = z_normalize_window(&base);
let cent_z = z_normalize_window(¢);
let corr: f64 = cent_z
.iter()
.zip(base_z.iter())
.map(|(a, b)| a * b)
.sum::<f64>()
/ m as f64;
assert!(
corr > 0.99,
"centroid must correlate positively, corr={corr}"
);
}
#[test]
fn test_kshape_empty_cluster_recovery() {
let (data, _) = shifted_groups(11);
let cfg = KShapeConfig {
n_clusters: 5,
n_init: 3,
seed: 2,
..Default::default()
};
let res = kshape_fd(&data, &cfg).unwrap();
assert_eq!(res.cluster.len(), 16);
assert!(res.cluster.iter().all(|&c| c < 5));
let mut sizes = vec![0usize; 5];
for &c in &res.cluster {
sizes[c] += 1;
}
assert!(
sizes.iter().all(|&s| s >= 1),
"an empty cluster survived: {sizes:?}"
);
}
#[test]
fn test_kshape_deterministic() {
let (data, _) = shifted_groups(5);
let cfg = KShapeConfig {
n_clusters: 2,
n_init: 5,
seed: 42,
..Default::default()
};
let a = kshape_fd(&data, &cfg).unwrap();
let b = kshape_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);
let n = a.centroids.nrows();
let m = a.centroids.ncols();
for i in 0..n {
for j in 0..m {
assert_eq!(a.centroids[(i, j)].to_bits(), b.centroids[(i, j)].to_bits());
}
}
}
#[test]
fn test_kshape_best_of_n_init() {
let (data, _) = shifted_groups(9);
let multi = KShapeConfig {
n_clusters: 2,
n_init: 10,
seed: 4,
..Default::default()
};
let single = KShapeConfig {
n_init: 1,
..multi.clone()
};
let rm = kshape_fd(&data, &multi).unwrap();
let rs = kshape_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_kshape_predict() {
let (data, _) = shifted_groups(13);
let cfg = KShapeConfig {
n_clusters: 2,
n_init: 10,
seed: 6,
..Default::default()
};
let res = kshape_fd(&data, &cfg).unwrap();
let preds = res.predict(&data).unwrap();
assert_eq!(preds, res.cluster, "predict(train) must reproduce cluster");
let m = data.ncols();
let src = data.row(0);
let novel = circular_shift(&src, 7);
let test = matrix_from_rows(&[novel]);
let p = res.predict(&test).unwrap();
assert_eq!(p.len(), 1);
assert_eq!(
p[0], res.cluster[0],
"shifted copy of series 0 should route to its cluster"
);
let _ = m;
}
#[test]
fn test_kshape_validation() {
let (data, _) = shifted_groups(1);
let cfg0 = KShapeConfig::new(0);
assert!(matches!(
kshape_fd(&data, &cfg0),
Err(FdarError::InvalidParameter { .. })
));
let cfg_big = KShapeConfig::new(999);
assert!(matches!(
kshape_fd(&data, &cfg_big),
Err(FdarError::InvalidParameter { .. })
));
let cfg_ni = KShapeConfig {
n_init: 0,
..KShapeConfig::new(2)
};
assert!(matches!(
kshape_fd(&data, &cfg_ni),
Err(FdarError::InvalidParameter { .. })
));
let empty = FdMatrix::zeros(0, 0);
assert!(matches!(
kshape_fd(&empty, &KShapeConfig::new(2)),
Err(FdarError::InvalidDimension { .. })
));
let res = kshape_fd(&data, &KShapeConfig::new(2)).unwrap();
let wrong = matrix_from_rows(&[vec![1.0, 2.0, 3.0]]);
assert!(matches!(
res.predict(&wrong),
Err(FdarError::InvalidDimension { .. })
));
}
#[test]
fn test_sbd_kmedoids_recovers_groups() {
let (data, truth) = shifted_groups(7);
let cfg = KMedoidsConfig {
k: 2,
max_iter: 100,
seed: 3,
};
let res = sbd_kmedoids(&data, &cfg).unwrap();
assert_eq!(res.labels.len(), 16);
assert_eq!(res.medoid_indices.len(), 2);
let p = purity(&res.labels, &truth, 2);
assert!(p >= 0.9, "SBD k-medoids purity {p} too low (< 0.9)");
}
#[test]
fn test_sbd_kmedoids_uses_sbd_matrix() {
let (data, _) = shifted_groups(5);
let cfg = KMedoidsConfig {
k: 2,
max_iter: 100,
seed: 42,
};
let res = sbd_kmedoids(&data, &cfg).unwrap();
let dist = sbd_distance_matrix(&data).unwrap();
let manual = kmedoids_from_distances(&dist, &cfg).unwrap();
assert_eq!(
res.labels, manual.labels,
"labels must match manual composition"
);
assert_eq!(
res.medoid_indices, manual.medoid_indices,
"medoids must match manual composition"
);
assert_eq!(
res.total_within_distance.to_bits(),
manual.total_within_distance.to_bits()
);
}
#[test]
fn test_sbd_kmedoids_validation() {
let (data, _) = shifted_groups(1);
let cfg0 = KMedoidsConfig {
k: 0,
..Default::default()
};
assert!(matches!(
sbd_kmedoids(&data, &cfg0),
Err(FdarError::InvalidParameter { .. })
));
let cfg_big = KMedoidsConfig {
k: 999,
..Default::default()
};
assert!(matches!(
sbd_kmedoids(&data, &cfg_big),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn test_kshape_reexports() {
use crate::{
kshape_fd, sbd, sbd_distance_matrix, sbd_kmedoids, KMedoidsConfig, KMedoidsResult,
KShapeConfig, KShapeResult, SbdResult,
};
let _f: fn(&FdMatrix, &KShapeConfig) -> Result<KShapeResult, FdarError> = kshape_fd;
let _k: fn(&FdMatrix, &KMedoidsConfig) -> Result<KMedoidsResult, FdarError> = sbd_kmedoids;
let _s: fn(&[f64], &[f64]) -> Result<SbdResult, FdarError> = sbd;
let _m: fn(&FdMatrix) -> Result<FdMatrix, FdarError> = sbd_distance_matrix;
}
}