1use crate::distance::l2_distance_matrix;
21use crate::error::FdarError;
22use crate::helpers::simpsons_weights;
23use crate::matrix::FdMatrix;
24use crate::regression::{fdata_to_pc_1d, FpcaResult};
25use rand::prelude::*;
26
27#[derive(Debug, Clone, PartialEq)]
65#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
66#[non_exhaustive]
67pub struct DbscanConfig {
68 pub eps: f64,
74 pub min_points: usize,
79}
80
81impl Default for DbscanConfig {
82 fn default() -> Self {
83 Self {
84 eps: 0.5,
85 min_points: 3,
86 }
87 }
88}
89
90#[derive(Debug, Clone)]
92#[non_exhaustive]
93pub struct DbscanResult {
94 pub cluster: Vec<Option<usize>>,
99 pub n_clusters: usize,
101 pub n_noise: usize,
103 pub distances: FdMatrix,
105}
106
107#[must_use = "expensive computation whose result should not be discarded"]
157pub fn dbscan_fd(
158 data: &FdMatrix,
159 argvals: &[f64],
160 config: &DbscanConfig,
161) -> Result<DbscanResult, FdarError> {
162 let (n, m) = data.shape();
163
164 if n == 0 || m == 0 {
166 return Err(FdarError::InvalidDimension {
167 parameter: "data",
168 expected: "at least 1 row and 1 column".to_string(),
169 actual: format!("{n} rows, {m} columns"),
170 });
171 }
172 if argvals.len() != m {
173 return Err(FdarError::InvalidDimension {
174 parameter: "argvals",
175 expected: format!("{m}"),
176 actual: format!("{}", argvals.len()),
177 });
178 }
179 if config.eps <= 0.0 {
180 return Err(FdarError::InvalidParameter {
181 parameter: "eps",
182 message: format!("eps must be > 0, got {}", config.eps),
183 });
184 }
185 if config.min_points == 0 {
186 return Err(FdarError::InvalidParameter {
187 parameter: "min_points",
188 message: "min_points must be >= 1".to_string(),
189 });
190 }
191
192 let dist = l2_distance_matrix(data, argvals);
193
194 let mut labels: Vec<Option<usize>> = vec![None; n];
197 let mut visited: Vec<bool> = vec![false; n];
198 let mut cluster_id: usize = 0;
199
200 for i in 0..n {
201 if visited[i] {
202 continue;
203 }
204 visited[i] = true;
205
206 let neighbors: Vec<usize> = (0..n)
208 .filter(|&j| j != i && dist[(i, j)] <= config.eps)
209 .collect();
210
211 if neighbors.len() + 1 < config.min_points {
213 continue;
215 }
216
217 labels[i] = Some(cluster_id);
219
220 let mut queue = neighbors.clone();
222 let mut qi = 0;
223 while qi < queue.len() {
224 let j = queue[qi];
225 qi += 1;
226
227 if !visited[j] {
228 visited[j] = true;
229 let j_neighbors: Vec<usize> = (0..n)
230 .filter(|&k| k != j && dist[(j, k)] <= config.eps)
231 .collect();
232 if j_neighbors.len() + 1 >= config.min_points {
233 for nb in j_neighbors {
235 if !queue.contains(&nb) {
236 queue.push(nb);
237 }
238 }
239 }
240 }
241
242 if labels[j].is_none() {
244 labels[j] = Some(cluster_id);
245 }
246 }
247
248 cluster_id += 1;
249 }
250
251 let n_clusters = cluster_id;
252 let n_noise = labels.iter().filter(|l| l.is_none()).count();
253
254 Ok(DbscanResult {
255 cluster: labels,
256 n_clusters,
257 n_noise,
258 distances: dist,
259 })
260}
261
262#[derive(Debug, Clone, PartialEq)]
275#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
276#[non_exhaustive]
277pub struct KcfcConfig {
278 pub k: usize,
280 pub ncomp: usize,
284 pub max_iter: usize,
286 pub seed: u64,
288}
289
290impl Default for KcfcConfig {
291 fn default() -> Self {
292 Self {
293 k: 2,
294 ncomp: 3,
295 max_iter: 50,
296 seed: 42,
297 }
298 }
299}
300
301#[derive(Debug, Clone)]
303#[non_exhaustive]
304pub struct KcfcResult {
305 pub cluster: Vec<usize>,
307 pub fpca_models: Vec<Option<FpcaResult>>,
311 pub reconstruction_errors: FdMatrix,
316 pub iterations: usize,
318 pub converged: bool,
320}
321
322#[must_use = "expensive computation whose result should not be discarded"]
371pub fn kcfc_cluster(
372 data: &FdMatrix,
373 argvals: &[f64],
374 config: &KcfcConfig,
375) -> Result<KcfcResult, FdarError> {
376 let (n, m) = data.shape();
377
378 if n == 0 || m == 0 {
380 return Err(FdarError::InvalidDimension {
381 parameter: "data",
382 expected: "at least 1 row and 1 column".to_string(),
383 actual: format!("{n} rows, {m} columns"),
384 });
385 }
386 if argvals.len() != m {
387 return Err(FdarError::InvalidDimension {
388 parameter: "argvals",
389 expected: format!("{m}"),
390 actual: format!("{}", argvals.len()),
391 });
392 }
393 if config.k == 0 {
394 return Err(FdarError::InvalidParameter {
395 parameter: "k",
396 message: "k must be >= 1".to_string(),
397 });
398 }
399 if config.k > n {
400 return Err(FdarError::InvalidParameter {
401 parameter: "k",
402 message: format!("k={} exceeds number of curves n={}", config.k, n),
403 });
404 }
405 if config.ncomp == 0 {
406 return Err(FdarError::InvalidParameter {
407 parameter: "ncomp",
408 message: "ncomp must be >= 1".to_string(),
409 });
410 }
411
412 let k = config.k;
413 let weights = simpsons_weights(argvals);
414
415 let row_major = data.to_row_major(); let mut rng = StdRng::seed_from_u64(config.seed);
418
419 let mut center_indices: Vec<usize> = Vec::with_capacity(k);
421 center_indices.push(rng.gen_range(0..n));
422
423 let mut min_dist_sq: Vec<f64> = (0..n)
425 .map(|i| {
426 let c0 = center_indices[0];
427 let d = l2_dist_rowmajor(&row_major, i, c0, m, &weights);
428 d * d
429 })
430 .collect();
431
432 while center_indices.len() < k {
433 let total: f64 = min_dist_sq.iter().sum();
435 let chosen = if total < 1e-15 {
436 rng.gen_range(0..n)
437 } else {
438 let r = rng.gen::<f64>() * total;
439 let mut cumsum = 0.0;
440 let mut sel = n - 1;
441 for (i, &d) in min_dist_sq.iter().enumerate() {
442 cumsum += d;
443 if cumsum >= r {
444 sel = i;
445 break;
446 }
447 }
448 sel
449 };
450 center_indices.push(chosen);
451
452 for i in 0..n {
454 let d = l2_dist_rowmajor(&row_major, i, chosen, m, &weights);
455 let d2 = d * d;
456 if d2 < min_dist_sq[i] {
457 min_dist_sq[i] = d2;
458 }
459 }
460 }
461
462 let mut cluster: Vec<usize> = (0..n)
464 .map(|i| {
465 center_indices
466 .iter()
467 .enumerate()
468 .min_by(|(_, &c1), (_, &c2)| {
469 let d1 = l2_dist_rowmajor(&row_major, i, c1, m, &weights);
470 let d2 = l2_dist_rowmajor(&row_major, i, c2, m, &weights);
471 d1.partial_cmp(&d2).unwrap_or(std::cmp::Ordering::Equal)
472 })
473 .map(|(ki, _)| ki)
474 .unwrap_or(0)
475 })
476 .collect();
477
478 let mut fpca_models: Vec<Option<FpcaResult>> = vec![None; k];
480 let mut reconstruction_errors = FdMatrix::zeros(n, k);
481 let mut converged = false;
482 let mut iterations = 0;
483
484 for _iter in 0..config.max_iter {
485 iterations += 1;
486
487 for ki in 0..k {
489 let member_indices: Vec<usize> = (0..n).filter(|&i| cluster[i] == ki).collect();
490
491 if member_indices.is_empty() {
492 continue;
494 }
495
496 let n_k = member_indices.len();
498 let mut col_major_k = vec![0.0_f64; n_k * m];
499 for (row_in_k, &orig_i) in member_indices.iter().enumerate() {
500 for j in 0..m {
501 col_major_k[row_in_k + j * n_k] = data[(orig_i, j)];
502 }
503 }
504 let data_k = FdMatrix::from_column_major(col_major_k, n_k, m)?;
505
506 match fdata_to_pc_1d(&data_k, config.ncomp, argvals) {
508 Ok(fpca) => {
509 fpca_models[ki] = Some(fpca);
510 }
511 Err(_) => {
512 }
514 }
515 }
516
517 for i in 0..n {
519 let curve_row = data.row(i);
520 let curve_mat = FdMatrix::from_slice(&curve_row, 1, m)?;
521
522 for ki in 0..k {
523 let err = match &fpca_models[ki] {
524 None => f64::INFINITY,
525 Some(fpca) => {
526 let ncomp_eff = fpca.rotation.ncols();
527 match fpca.project(&curve_mat) {
528 Ok(scores) => {
529 match fpca.reconstruct(&scores, ncomp_eff) {
530 Ok(recon) => {
531 let mut err_sq = 0.0;
533 for j in 0..m {
534 let diff = curve_row[j] - recon[(0, j)];
535 err_sq += diff * diff * weights[j];
536 }
537 err_sq
538 }
539 Err(_) => f64::INFINITY,
540 }
541 }
542 Err(_) => f64::INFINITY,
543 }
544 }
545 };
546 reconstruction_errors[(i, ki)] = err;
547 }
548 }
549
550 let mut changed = false;
552 for i in 0..n {
553 let best_k = (0..k)
554 .min_by(|&a, &b| {
555 reconstruction_errors[(i, a)]
556 .partial_cmp(&reconstruction_errors[(i, b)])
557 .unwrap_or(std::cmp::Ordering::Equal)
558 })
559 .unwrap_or(0);
560 if cluster[i] != best_k {
561 cluster[i] = best_k;
562 changed = true;
563 }
564 }
565
566 if !changed {
567 converged = true;
568 break;
569 }
570 }
571
572 Ok(KcfcResult {
573 cluster,
574 fpca_models,
575 reconstruction_errors,
576 iterations,
577 converged,
578 })
579}
580
581#[derive(Debug, Clone, PartialEq)]
633#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
634#[non_exhaustive]
635pub struct FunFemConfig {
636 pub k: usize,
638 pub ncomp: usize,
641 pub p_disc: usize,
644 pub max_iter: usize,
646 pub tol: f64,
648 pub seed: u64,
650}
651
652impl Default for FunFemConfig {
653 fn default() -> Self {
654 Self {
655 k: 2,
656 ncomp: 10,
657 p_disc: 0,
658 max_iter: 50,
659 tol: 1e-6,
660 seed: 42,
661 }
662 }
663}
664
665#[derive(Debug, Clone)]
667#[non_exhaustive]
668pub struct FunFemResult {
669 pub cluster: Vec<usize>,
671 pub membership: FdMatrix,
673 pub disc_subspace: FdMatrix,
675 pub log_likelihood: f64,
677 pub iterations: usize,
679 pub converged: bool,
681}
682
683#[must_use = "expensive computation whose result should not be discarded"]
701pub fn funfem_cluster(
702 data: &FdMatrix,
703 argvals: &[f64],
704 config: &FunFemConfig,
705) -> Result<FunFemResult, FdarError> {
706 let (n, m) = data.shape();
707
708 if n == 0 || m == 0 {
710 return Err(FdarError::InvalidDimension {
711 parameter: "data",
712 expected: "at least 1 row and 1 column".to_string(),
713 actual: format!("{n} rows, {m} columns"),
714 });
715 }
716 if argvals.len() != m {
717 return Err(FdarError::InvalidDimension {
718 parameter: "argvals",
719 expected: format!("{m}"),
720 actual: format!("{}", argvals.len()),
721 });
722 }
723 if config.k == 0 {
724 return Err(FdarError::InvalidParameter {
725 parameter: "k",
726 message: "k must be >= 1".to_string(),
727 });
728 }
729 if config.k > n {
730 return Err(FdarError::InvalidParameter {
731 parameter: "k",
732 message: format!("k={} exceeds number of curves n={}", config.k, n),
733 });
734 }
735 if config.ncomp == 0 {
736 return Err(FdarError::InvalidParameter {
737 parameter: "ncomp",
738 message: "ncomp must be >= 1".to_string(),
739 });
740 }
741
742 let k = config.k;
743
744 let fpca = fdata_to_pc_1d(data, config.ncomp, argvals)?;
746 let scores = &fpca.scores; let ncomp_eff = scores.ncols();
749
750 let p_disc_eff = if config.p_disc == 0 {
752 (k - 1).max(1).min(ncomp_eff)
753 } else {
754 config.p_disc.min(ncomp_eff)
755 };
756
757 let weights_uniform = vec![1.0; ncomp_eff]; let row_major_scores = scores.to_row_major(); let mut rng = StdRng::seed_from_u64(config.seed);
761
762 let mut center_indices: Vec<usize> = Vec::with_capacity(k);
763 center_indices.push(rng.gen_range(0..n));
764 let mut min_dist_sq: Vec<f64> = (0..n)
765 .map(|i| {
766 let c0 = center_indices[0];
767 let d = l2_dist_rowmajor(&row_major_scores, i, c0, ncomp_eff, &weights_uniform);
768 d * d
769 })
770 .collect();
771 while center_indices.len() < k {
772 let total: f64 = min_dist_sq.iter().sum();
773 let chosen = if total < 1e-15 {
774 rng.gen_range(0..n)
775 } else {
776 let r = rng.gen::<f64>() * total;
777 let mut cumsum = 0.0;
778 let mut sel = n - 1;
779 for (i, &d) in min_dist_sq.iter().enumerate() {
780 cumsum += d;
781 if cumsum >= r {
782 sel = i;
783 break;
784 }
785 }
786 sel
787 };
788 center_indices.push(chosen);
789 for i in 0..n {
790 let d = l2_dist_rowmajor(&row_major_scores, i, chosen, ncomp_eff, &weights_uniform);
791 let d2 = d * d;
792 if d2 < min_dist_sq[i] {
793 min_dist_sq[i] = d2;
794 }
795 }
796 }
797
798 let mut cluster: Vec<usize> = (0..n)
800 .map(|i| {
801 center_indices
802 .iter()
803 .enumerate()
804 .min_by(|(_, &c1), (_, &c2)| {
805 let d1 =
806 l2_dist_rowmajor(&row_major_scores, i, c1, ncomp_eff, &weights_uniform);
807 let d2 =
808 l2_dist_rowmajor(&row_major_scores, i, c2, ncomp_eff, &weights_uniform);
809 d1.partial_cmp(&d2).unwrap_or(std::cmp::Ordering::Equal)
810 })
811 .map(|(ki, _)| ki)
812 .unwrap_or(0)
813 })
814 .collect();
815
816 let mut pi: Vec<f64> = vec![1.0 / k as f64; k];
818 let mut mu_k: Vec<Vec<f64>> = vec![vec![0.0; ncomp_eff]; k];
820 let mut sigma_k: Vec<Vec<f64>> = vec![vec![1.0; ncomp_eff]; k];
822
823 update_gmm_params_from_hard(
825 &row_major_scores,
826 &cluster,
827 k,
828 ncomp_eff,
829 &mut pi,
830 &mut mu_k,
831 &mut sigma_k,
832 );
833
834 let mut disc_dirs: Vec<f64> = {
836 let mut v = vec![0.0_f64; ncomp_eff * p_disc_eff];
837 for d in 0..p_disc_eff {
838 if d < ncomp_eff {
839 v[d + d * ncomp_eff] = 1.0; }
841 }
842 v
843 };
844
845 let mut prev_ll = f64::NEG_INFINITY;
846 let mut resp = vec![0.0_f64; n * k];
848 for i in 0..n {
849 let ki = cluster[i].min(k - 1);
850 resp[i * k + ki] = 1.0;
851 }
852 let mut converged = false;
853 let mut iterations = 0;
854
855 for _iter in 0..config.max_iter {
856 iterations += 1;
857
858 let mut proj_scores = vec![0.0_f64; n * p_disc_eff]; for i in 0..n {
863 for d in 0..p_disc_eff {
864 let mut val = 0.0;
865 for j in 0..ncomp_eff {
866 val += scores[(i, j)] * disc_dirs[j + d * ncomp_eff];
867 }
868 proj_scores[i * p_disc_eff + d] = val;
869 }
870 }
871
872 let mut mu_disc: Vec<Vec<f64>> = vec![vec![0.0; p_disc_eff]; k];
875 let mut n_k_soft: Vec<f64> = vec![0.0; k];
876 for i in 0..n {
877 for ki in 0..k {
878 let r = resp[i * k + ki];
879 n_k_soft[ki] += r;
880 for d in 0..p_disc_eff {
881 mu_disc[ki][d] += r * proj_scores[i * p_disc_eff + d];
882 }
883 }
884 }
885 for ki in 0..k {
886 if n_k_soft[ki] > 1e-10 {
887 for d in 0..p_disc_eff {
888 mu_disc[ki][d] /= n_k_soft[ki];
889 }
890 }
891 }
892
893 let mut var_disc: Vec<Vec<f64>> = vec![vec![1.0; p_disc_eff]; k];
895 for ki in 0..k {
896 if n_k_soft[ki] > 1e-10 {
897 for d in 0..p_disc_eff {
898 let mut v = 0.0;
899 for i in 0..n {
900 let diff = proj_scores[i * p_disc_eff + d] - mu_disc[ki][d];
901 v += resp[i * k + ki] * diff * diff;
902 }
903 var_disc[ki][d] = (v / n_k_soft[ki]).max(1e-8);
904 }
905 }
906 }
907
908 let mut log_resp = vec![0.0_f64; n * k];
910 let mut ll = 0.0;
911 for i in 0..n {
912 let mut log_components = vec![0.0_f64; k];
913 for ki in 0..k {
914 let log_pi = if pi[ki] > 1e-300 { pi[ki].ln() } else { -700.0 };
915 let mut log_lik = log_pi;
916 for d in 0..p_disc_eff {
917 let diff = proj_scores[i * p_disc_eff + d] - mu_disc[ki][d];
918 let var = var_disc[ki][d];
919 log_lik -= 0.5 * (var.ln() + diff * diff / var);
920 }
921 log_lik -= 0.5 * (p_disc_eff as f64) * std::f64::consts::TAU.ln();
922 log_components[ki] = log_lik;
923 }
924 let log_sum = log_sum_exp(&log_components);
926 ll += log_sum;
927 for ki in 0..k {
928 log_resp[i * k + ki] = log_components[ki] - log_sum;
929 }
930 }
931
932 resp.fill(0.0);
934 for i in 0..n {
935 for ki in 0..k {
936 resp[i * k + ki] = log_resp[i * k + ki].exp().max(1e-300);
937 }
938 }
939
940 let mut n_k_new: Vec<f64> = vec![0.0; k];
942 for i in 0..n {
943 for ki in 0..k {
944 n_k_new[ki] += resp[i * k + ki];
945 }
946 }
947 let n_total: f64 = n_k_new.iter().sum();
948 for ki in 0..k {
949 pi[ki] = (n_k_new[ki] / n_total).max(1e-300);
950 }
951
952 cluster = (0..n)
954 .map(|i| {
955 (0..k)
956 .max_by(|&a, &b| {
957 resp[i * k + a]
958 .partial_cmp(&resp[i * k + b])
959 .unwrap_or(std::cmp::Ordering::Equal)
960 })
961 .unwrap_or(0)
962 })
963 .collect();
964
965 update_gmm_params_from_soft(
967 &row_major_scores,
968 &resp,
969 k,
970 ncomp_eff,
971 n,
972 &mut pi,
973 &mut mu_k,
974 &mut sigma_k,
975 );
976
977 let global_mean: Vec<f64> = (0..ncomp_eff)
980 .map(|j| {
981 (0..n)
982 .map(|i| row_major_scores[i * ncomp_eff + j])
983 .sum::<f64>()
984 / n as f64
985 })
986 .collect();
987
988 let mut b_soft = vec![0.0_f64; ncomp_eff * ncomp_eff]; let mut w_soft = vec![0.0_f64; ncomp_eff * ncomp_eff];
990
991 for ki in 0..k {
993 let nk = n_k_new[ki].max(1.0);
994 for j in 0..ncomp_eff {
995 for l in 0..ncomp_eff {
996 b_soft[j * ncomp_eff + l] +=
997 nk * (mu_k[ki][j] - global_mean[j]) * (mu_k[ki][l] - global_mean[l]);
998 }
999 }
1000 }
1001
1002 for i in 0..n {
1004 for ki in 0..k {
1005 let r = resp[i * k + ki];
1006 for j in 0..ncomp_eff {
1007 let dj = row_major_scores[i * ncomp_eff + j] - mu_k[ki][j];
1008 for l in 0..ncomp_eff {
1009 let dl = row_major_scores[i * ncomp_eff + l] - mu_k[ki][l];
1010 w_soft[j * ncomp_eff + l] += r * dj * dl;
1011 }
1012 }
1013 }
1014 }
1015
1016 let trace_w: f64 = (0..ncomp_eff)
1018 .map(|j| w_soft[j * ncomp_eff + j])
1019 .sum::<f64>();
1020 let reg_floor = (trace_w / ncomp_eff as f64 * 1e-4).max(1e-8);
1021 for j in 0..ncomp_eff {
1022 w_soft[j * ncomp_eff + j] += reg_floor;
1023 }
1024
1025 match crate::linalg::cholesky_factor(&w_soft, ncomp_eff) {
1028 Ok(l_w) => {
1029 let mut winv_b = vec![0.0_f64; ncomp_eff * ncomp_eff];
1031 for col in 0..ncomp_eff {
1032 let b_col: Vec<f64> = (0..ncomp_eff)
1033 .map(|r| b_soft[r * ncomp_eff + col])
1034 .collect();
1035 let x = crate::linalg::cholesky_forward_back(&l_w, &b_col, ncomp_eff);
1036 for r in 0..ncomp_eff {
1037 winv_b[r * ncomp_eff + col] = x[r];
1038 }
1039 }
1040
1041 use nalgebra::{DMatrix, SVD};
1043 let mat = DMatrix::from_row_slice(ncomp_eff, ncomp_eff, &winv_b);
1044 let svd = SVD::new(mat, true, false);
1045 if let Some(u) = svd.u {
1046 let mut new_dirs = vec![0.0_f64; ncomp_eff * p_disc_eff];
1049 for d in 0..p_disc_eff {
1050 for r in 0..ncomp_eff {
1051 new_dirs[r + d * ncomp_eff] = u[(r, d)];
1052 }
1053 }
1054 disc_dirs = new_dirs;
1055 }
1056 }
1058 Err(_) => {
1059 }
1061 }
1062
1063 let delta = (ll - prev_ll).abs();
1065 prev_ll = ll;
1066 if _iter > 0 && delta < config.tol {
1067 converged = true;
1068 break;
1069 }
1070 }
1071
1072 let mut membership_data = vec![0.0_f64; n * k];
1074 for i in 0..n {
1075 for ki in 0..k {
1076 membership_data[i + ki * n] = resp[i * k + ki];
1078 }
1079 }
1080 let membership = FdMatrix::from_column_major(membership_data, n, k)?;
1081
1082 let disc_subspace = FdMatrix::from_column_major(disc_dirs, ncomp_eff, p_disc_eff)?;
1085
1086 Ok(FunFemResult {
1087 cluster,
1088 membership,
1089 disc_subspace,
1090 log_likelihood: prev_ll,
1091 iterations,
1092 converged,
1093 })
1094}
1095
1096fn update_gmm_params_from_hard(
1098 scores_rm: &[f64],
1099 cluster: &[usize],
1100 k: usize,
1101 d: usize,
1102 pi: &mut [f64],
1103 mu_k: &mut [Vec<f64>],
1104 sigma_k: &mut [Vec<f64>],
1105) {
1106 let n = cluster.len();
1107 let mut counts = vec![0usize; k];
1108 for &c in cluster {
1109 if c < k {
1110 counts[c] += 1;
1111 }
1112 }
1113 for ki in 0..k {
1114 pi[ki] = (counts[ki] as f64 / n as f64).max(1e-300);
1115 mu_k[ki] = vec![0.0; d];
1116 sigma_k[ki] = vec![0.0; d];
1119 for i in 0..n {
1120 if cluster[i] == ki {
1121 for j in 0..d {
1122 mu_k[ki][j] += scores_rm[i * d + j];
1123 }
1124 }
1125 }
1126 if counts[ki] > 0 {
1127 for j in 0..d {
1128 mu_k[ki][j] /= counts[ki] as f64;
1129 }
1130 }
1131 for i in 0..n {
1132 if cluster[i] == ki {
1133 for j in 0..d {
1134 let diff = scores_rm[i * d + j] - mu_k[ki][j];
1135 sigma_k[ki][j] += diff * diff;
1136 }
1137 }
1138 }
1139 if counts[ki] > 1 {
1140 for j in 0..d {
1141 sigma_k[ki][j] = (sigma_k[ki][j] / counts[ki] as f64).max(1e-8);
1142 }
1143 } else {
1144 for j in 0..d {
1145 sigma_k[ki][j] = 1.0;
1146 }
1147 }
1148 }
1149}
1150
1151fn update_gmm_params_from_soft(
1153 scores_rm: &[f64],
1154 resp: &[f64],
1155 k: usize,
1156 d: usize,
1157 n: usize,
1158 pi: &mut [f64],
1159 mu_k: &mut [Vec<f64>],
1160 sigma_k: &mut [Vec<f64>],
1161) {
1162 let mut n_k = vec![0.0_f64; k];
1163 for i in 0..n {
1164 for ki in 0..k {
1165 n_k[ki] += resp[i * k + ki];
1166 }
1167 }
1168 let total: f64 = n_k.iter().sum();
1169 for ki in 0..k {
1170 pi[ki] = (n_k[ki] / total).max(1e-300);
1171 mu_k[ki] = vec![0.0; d];
1172 sigma_k[ki] = vec![1.0; d];
1173 if n_k[ki] > 1e-10 {
1174 for i in 0..n {
1175 let r = resp[i * k + ki];
1176 for j in 0..d {
1177 mu_k[ki][j] += r * scores_rm[i * d + j];
1178 }
1179 }
1180 for j in 0..d {
1181 mu_k[ki][j] /= n_k[ki];
1182 }
1183 let mut var_j = vec![0.0_f64; d];
1184 for i in 0..n {
1185 let r = resp[i * k + ki];
1186 for j in 0..d {
1187 let diff = scores_rm[i * d + j] - mu_k[ki][j];
1188 var_j[j] += r * diff * diff;
1189 }
1190 }
1191 for j in 0..d {
1192 sigma_k[ki][j] = (var_j[j] / n_k[ki]).max(1e-8);
1193 }
1194 }
1195 }
1196}
1197
1198fn log_sum_exp(v: &[f64]) -> f64 {
1200 let max_v = v.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
1201 if max_v == f64::NEG_INFINITY {
1202 return f64::NEG_INFINITY;
1203 }
1204 max_v + v.iter().map(|&x| (x - max_v).exp()).sum::<f64>().ln()
1205}
1206
1207fn l2_dist_rowmajor(buf: &[f64], i: usize, j: usize, m: usize, weights: &[f64]) -> f64 {
1211 let mut sq = 0.0;
1212 for t in 0..m {
1213 let d = buf[i * m + t] - buf[j * m + t];
1214 sq += d * d * weights[t];
1215 }
1216 sq.sqrt()
1217}
1218
1219#[derive(Debug, Clone, PartialEq)]
1265#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
1266#[non_exhaustive]
1267pub struct AlignClusterConfig {
1268 pub k: usize,
1270 pub max_iter: usize,
1272 pub seed: u64,
1274 pub use_amplitude_only: bool,
1277 pub elastic_lambda: f64,
1279 pub karcher_max_iter: usize,
1281 pub karcher_tol: f64,
1283}
1284
1285impl Default for AlignClusterConfig {
1286 fn default() -> Self {
1287 Self {
1288 k: 2,
1289 max_iter: 20,
1290 seed: 42,
1291 use_amplitude_only: true,
1292 elastic_lambda: 0.0,
1293 karcher_max_iter: 15,
1294 karcher_tol: 1e-4,
1295 }
1296 }
1297}
1298
1299#[derive(Debug, Clone)]
1301#[non_exhaustive]
1302pub struct AlignClusterResult {
1303 pub cluster: Vec<usize>,
1305 pub templates: Vec<Vec<f64>>,
1307 pub distances: FdMatrix,
1310 pub iterations: usize,
1312 pub converged: bool,
1314}
1315
1316#[must_use = "expensive computation whose result should not be discarded"]
1335pub fn align_cluster_fd(
1336 data: &FdMatrix,
1337 argvals: &[f64],
1338 config: &AlignClusterConfig,
1339) -> Result<AlignClusterResult, FdarError> {
1340 use crate::alignment::{amplitude_distance, elastic_distance, karcher_mean};
1341
1342 let (n, m) = data.shape();
1343
1344 if n == 0 || m == 0 {
1346 return Err(FdarError::InvalidDimension {
1347 parameter: "data",
1348 expected: "at least 1 row and 1 column".to_string(),
1349 actual: format!("{n} rows, {m} columns"),
1350 });
1351 }
1352 if argvals.len() != m {
1353 return Err(FdarError::InvalidDimension {
1354 parameter: "argvals",
1355 expected: format!("{m}"),
1356 actual: format!("{}", argvals.len()),
1357 });
1358 }
1359 if config.k == 0 {
1360 return Err(FdarError::InvalidParameter {
1361 parameter: "k",
1362 message: "k must be >= 1".to_string(),
1363 });
1364 }
1365 if config.k > n {
1366 return Err(FdarError::InvalidParameter {
1367 parameter: "k",
1368 message: format!("k={} exceeds number of curves n={}", config.k, n),
1369 });
1370 }
1371
1372 let k = config.k;
1373
1374 let mut rng = StdRng::seed_from_u64(config.seed);
1379 let mut shuffled: Vec<usize> = (0..n).collect();
1380 for i in (1..n).rev() {
1382 let j = rng.gen_range(0..=i);
1383 shuffled.swap(i, j);
1384 }
1385 let step = n / k;
1387 let template_indices: Vec<usize> = (0..k).map(|ki| shuffled[(ki * step).min(n - 1)]).collect();
1388
1389 let mut templates: Vec<Vec<f64>> = template_indices.iter().map(|&ci| data.row(ci)).collect();
1391
1392 let mut cluster: Vec<usize> = vec![0; n];
1393 let mut distances = FdMatrix::zeros(n, k);
1394 let mut converged = false;
1395 let mut iterations = 0;
1396
1397 for _iter in 0..config.max_iter {
1398 iterations += 1;
1399
1400 for i in 0..n {
1402 let curve_i = data.row(i);
1403 for ki in 0..k {
1404 let dist = if config.use_amplitude_only {
1405 amplitude_distance(&curve_i, &templates[ki], argvals, config.elastic_lambda)
1406 } else {
1407 elastic_distance(&curve_i, &templates[ki], argvals, config.elastic_lambda)
1408 };
1409 distances[(i, ki)] = dist;
1410 }
1411 }
1412
1413 let mut changed = false;
1415 for i in 0..n {
1416 let best_k = (0..k)
1417 .min_by(|&a, &b| {
1418 distances[(i, a)]
1419 .partial_cmp(&distances[(i, b)])
1420 .unwrap_or(std::cmp::Ordering::Equal)
1421 })
1422 .unwrap_or(0);
1423 if cluster[i] != best_k {
1424 cluster[i] = best_k;
1425 changed = true;
1426 }
1427 }
1428
1429 let mut template_changed = false;
1434 for ki in 0..k {
1435 let member_indices: Vec<usize> = (0..n).filter(|&i| cluster[i] == ki).collect();
1436
1437 if member_indices.is_empty() {
1438 let non_members: Vec<usize> = (0..n).filter(|&i| cluster[i] != ki).collect();
1441 if !non_members.is_empty() {
1442 let rand_idx = rng.gen_range(0..non_members.len());
1443 templates[ki] = data.row(non_members[rand_idx]);
1444 template_changed = true;
1445 }
1446 continue;
1448 }
1449
1450 let n_k = member_indices.len();
1452 let mut col_major_k = vec![0.0_f64; n_k * m];
1453 for (row_in_k, &orig_i) in member_indices.iter().enumerate() {
1454 for j in 0..m {
1455 col_major_k[row_in_k + j * n_k] = data[(orig_i, j)];
1456 }
1457 }
1458 let data_k = FdMatrix::from_column_major(col_major_k, n_k, m)?;
1459
1460 let km = karcher_mean(
1462 &data_k,
1463 argvals,
1464 config.karcher_max_iter,
1465 config.karcher_tol,
1466 config.elastic_lambda,
1467 );
1468 templates[ki] = km.mean;
1469 }
1470
1471 if !changed && !template_changed {
1477 converged = true;
1478 break;
1479 }
1480 }
1481
1482 Ok(AlignClusterResult {
1483 cluster,
1484 templates,
1485 distances,
1486 iterations,
1487 converged,
1488 })
1489}
1490
1491#[cfg(test)]
1496mod tests {
1497 use super::*;
1498 use crate::test_helpers::{adjusted_rand_index, uniform_grid};
1499 use std::f64::consts::PI;
1500
1501 fn two_tight_clusters(n_per: usize, m: usize) -> (FdMatrix, Vec<f64>, Vec<usize>) {
1506 let t = uniform_grid(m);
1507 let n = 2 * n_per;
1508 let mut col_major = vec![0.0_f64; n * m];
1509 for i in 0..n_per {
1510 for (j, &tj) in t.iter().enumerate() {
1511 col_major[i + j * n] = (2.0 * PI * tj).sin();
1512 }
1513 }
1514 for i in 0..n_per {
1515 for (j, &tj) in t.iter().enumerate() {
1516 col_major[(i + n_per) + j * n] = (2.0 * PI * tj).sin() + 5.0;
1517 }
1518 }
1519 let labels: Vec<usize> = (0..n).map(|i| if i < n_per { 0 } else { 1 }).collect();
1520 (
1521 FdMatrix::from_column_major(col_major, n, m).unwrap(),
1522 t,
1523 labels,
1524 )
1525 }
1526
1527 fn clusters_with_noise(n_per: usize, m: usize) -> (FdMatrix, Vec<f64>) {
1529 let t = uniform_grid(m);
1530 let n = 2 * n_per + 2;
1531 let mut col_major = vec![0.0_f64; n * m];
1532 for i in 0..n_per {
1534 for (j, &tj) in t.iter().enumerate() {
1535 col_major[i + j * n] = (2.0 * PI * tj).sin();
1536 }
1537 }
1538 for i in 0..n_per {
1540 for (j, &tj) in t.iter().enumerate() {
1541 col_major[(i + n_per) + j * n] = (2.0 * PI * tj).sin() + 5.0;
1542 }
1543 }
1544 let o0 = 2 * n_per;
1546 for j in 0..m {
1547 col_major[o0 + j * n] = 100.0;
1548 }
1549 let o1 = 2 * n_per + 1;
1551 for j in 0..m {
1552 col_major[o1 + j * n] = -100.0;
1553 }
1554 (FdMatrix::from_column_major(col_major, n, m).unwrap(), t)
1555 }
1556
1557 fn two_separated_clusters(n_per: usize, m: usize) -> (FdMatrix, Vec<f64>, Vec<usize>) {
1559 let t = uniform_grid(m);
1560 let n = 2 * n_per;
1561 let mut col_major = vec![0.0_f64; n * m];
1562 for i in 0..n_per {
1564 for (j, &tj) in t.iter().enumerate() {
1565 col_major[i + j * n] = (2.0 * PI * tj).sin() + 0.05 * (i as f64 / n_per as f64);
1566 }
1567 }
1568 for i in 0..n_per {
1570 for (j, &tj) in t.iter().enumerate() {
1571 col_major[(i + n_per) + j * n] =
1572 (2.0 * PI * tj).cos() + 8.0 + 0.05 * (i as f64 / n_per as f64);
1573 }
1574 }
1575 let labels: Vec<usize> = (0..n).map(|i| if i < n_per { 0 } else { 1 }).collect();
1576 (
1577 FdMatrix::from_column_major(col_major, n, m).unwrap(),
1578 t,
1579 labels,
1580 )
1581 }
1582
1583 #[test]
1586 fn test_dbscan_core_points() {
1587 let m = 30;
1588 let n_per = 5;
1589 let (data, t, _labels) = two_tight_clusters(n_per, m);
1590 let result = dbscan_fd(
1591 &data,
1592 &t,
1593 &DbscanConfig {
1594 eps: 1.0,
1595 min_points: 2,
1596 ..Default::default()
1597 },
1598 )
1599 .unwrap();
1600 assert_eq!(result.n_clusters, 2, "expected 2 clusters");
1601 assert_eq!(result.n_noise, 0, "expected 0 noise points");
1602 assert_eq!(result.cluster.len(), 2 * n_per);
1603 }
1604
1605 #[test]
1606 fn test_dbscan_noise_flagging() {
1607 let m = 30;
1608 let n_per = 5;
1609 let (data, t) = clusters_with_noise(n_per, m);
1610 let result = dbscan_fd(
1611 &data,
1612 &t,
1613 &DbscanConfig {
1614 eps: 1.5,
1615 min_points: 2,
1616 ..Default::default()
1617 },
1618 )
1619 .unwrap();
1620 assert_eq!(
1622 result.n_noise, 2,
1623 "expected exactly 2 noise points, got {}",
1624 result.n_noise
1625 );
1626 assert_eq!(result.n_clusters, 2, "expected 2 clusters");
1627 let n = data.nrows();
1629 assert!(result.cluster[n - 2].is_none(), "outlier 0 should be noise");
1630 assert!(result.cluster[n - 1].is_none(), "outlier 1 should be noise");
1631 }
1632
1633 #[test]
1634 fn test_dbscan_zero_eps_returns_err() {
1635 let m = 20;
1636 let (data, t, _) = two_tight_clusters(5, m);
1637 assert!(
1638 dbscan_fd(
1639 &data,
1640 &t,
1641 &DbscanConfig {
1642 eps: 0.0,
1643 ..Default::default()
1644 }
1645 )
1646 .is_err(),
1647 "eps=0 should return Err"
1648 );
1649 }
1650
1651 #[test]
1652 fn test_dbscan_negative_eps_returns_err() {
1653 let m = 20;
1654 let (data, t, _) = two_tight_clusters(5, m);
1655 assert!(
1656 dbscan_fd(
1657 &data,
1658 &t,
1659 &DbscanConfig {
1660 eps: -1.0,
1661 ..Default::default()
1662 }
1663 )
1664 .is_err(),
1665 "eps=-1 should return Err"
1666 );
1667 }
1668
1669 #[test]
1670 fn test_dbscan_invalid_min_points_zero() {
1671 let m = 20;
1672 let (data, t, _) = two_tight_clusters(5, m);
1673 assert!(
1674 dbscan_fd(
1675 &data,
1676 &t,
1677 &DbscanConfig {
1678 min_points: 0,
1679 ..Default::default()
1680 }
1681 )
1682 .is_err(),
1683 "min_points=0 should return Err"
1684 );
1685 }
1686
1687 #[test]
1688 fn test_dbscan_empty_data() {
1689 let data = FdMatrix::zeros(0, 0);
1690 let t: Vec<f64> = vec![];
1691 assert!(
1692 dbscan_fd(&data, &t, &DbscanConfig::default()).is_err(),
1693 "empty data should return Err"
1694 );
1695 }
1696
1697 #[test]
1698 fn test_dbscan_mismatched_argvals() {
1699 let m = 20;
1700 let (data, _t, _) = two_tight_clusters(5, m);
1701 let wrong_t = uniform_grid(m + 1);
1702 assert!(
1703 dbscan_fd(&data, &wrong_t, &DbscanConfig::default()).is_err(),
1704 "mismatched argvals should return Err"
1705 );
1706 }
1707
1708 #[test]
1709 fn test_dbscan_distances_shape() {
1710 let m = 20;
1711 let n_per = 4;
1712 let (data, t, _) = two_tight_clusters(n_per, m);
1713 let result = dbscan_fd(
1714 &data,
1715 &t,
1716 &DbscanConfig {
1717 eps: 1.0,
1718 min_points: 2,
1719 ..Default::default()
1720 },
1721 )
1722 .unwrap();
1723 let n = 2 * n_per;
1724 assert_eq!(
1725 result.distances.shape(),
1726 (n, n),
1727 "distance matrix must be n x n"
1728 );
1729 }
1730
1731 #[test]
1734 fn test_kcfc_recovery() {
1735 let m = 40;
1736 let n_per = 10;
1737 let (data, t, ground_truth) = two_separated_clusters(n_per, m);
1738 let result = kcfc_cluster(
1739 &data,
1740 &t,
1741 &KcfcConfig {
1742 k: 2,
1743 ncomp: 3,
1744 max_iter: 50,
1745 seed: 42,
1746 ..Default::default()
1747 },
1748 )
1749 .unwrap();
1750 let ari = adjusted_rand_index(&result.cluster, &ground_truth);
1751 assert!(
1752 ari >= 0.90,
1753 "kCFC ARI={ari:.3} should be >= 0.90 on well-separated data"
1754 );
1755 }
1756
1757 #[test]
1758 fn test_kcfc_errors_ordering() {
1759 let m = 40;
1762 let n_per = 10;
1763 let (data, t, ground_truth) = two_separated_clusters(n_per, m);
1764 let result = kcfc_cluster(
1765 &data,
1766 &t,
1767 &KcfcConfig {
1768 k: 2,
1769 ncomp: 3,
1770 max_iter: 50,
1771 seed: 42,
1772 ..Default::default()
1773 },
1774 )
1775 .unwrap();
1776
1777 let gt0_cluster = result.cluster[0]; let gt1_cluster = result.cluster[n_per]; let mut correct_ordering = 0;
1784 let mut total = 0;
1785 for i in 0..data.nrows() {
1786 let expected_cluster = if ground_truth[i] == 0 {
1787 gt0_cluster
1788 } else {
1789 gt1_cluster
1790 };
1791 let other_cluster = (0..2).find(|&c| c != expected_cluster).unwrap_or(0);
1793 let err_own = result.reconstruction_errors[(i, expected_cluster)];
1794 let err_other = result.reconstruction_errors[(i, other_cluster)];
1795 if err_own.is_finite() && err_other.is_finite() {
1796 if err_own < err_other {
1797 correct_ordering += 1;
1798 }
1799 total += 1;
1800 }
1801 }
1802 assert!(
1804 correct_ordering * 10 >= total * 8,
1805 "only {correct_ordering}/{total} curves had smaller error for their true cluster"
1806 );
1807 }
1808
1809 #[test]
1810 fn test_kcfc_deterministic() {
1811 let m = 30;
1812 let n_per = 8;
1813 let (data, t, _) = two_separated_clusters(n_per, m);
1814 let cfg = KcfcConfig {
1815 k: 2,
1816 ncomp: 2,
1817 max_iter: 30,
1818 seed: 7,
1819 ..Default::default()
1820 };
1821 let r1 = kcfc_cluster(&data, &t, &cfg).unwrap();
1822 let r2 = kcfc_cluster(&data, &t, &cfg).unwrap();
1823 assert_eq!(
1824 r1.cluster, r2.cluster,
1825 "identical seed must produce identical assignments"
1826 );
1827 }
1828
1829 #[test]
1830 fn test_kcfc_invalid_k_zero() {
1831 let m = 20;
1832 let (data, t, _) = two_tight_clusters(5, m);
1833 assert!(
1834 kcfc_cluster(
1835 &data,
1836 &t,
1837 &KcfcConfig {
1838 k: 0,
1839 ..Default::default()
1840 }
1841 )
1842 .is_err(),
1843 "k=0 should return Err"
1844 );
1845 }
1846
1847 #[test]
1848 fn test_kcfc_invalid_k_gt_n() {
1849 let m = 20;
1850 let n = 4;
1851 let (data, t, _) = two_tight_clusters(n / 2, m);
1852 assert!(
1853 kcfc_cluster(
1854 &data,
1855 &t,
1856 &KcfcConfig {
1857 k: n + 1,
1858 ..Default::default()
1859 }
1860 )
1861 .is_err(),
1862 "k>n should return Err"
1863 );
1864 }
1865
1866 #[test]
1867 fn test_kcfc_empty_data() {
1868 let data = FdMatrix::zeros(0, 0);
1869 let t: Vec<f64> = vec![];
1870 assert!(
1871 kcfc_cluster(&data, &t, &KcfcConfig::default()).is_err(),
1872 "empty data should return Err"
1873 );
1874 }
1875
1876 #[test]
1877 fn test_kcfc_mismatched_argvals() {
1878 let m = 20;
1879 let (data, _t, _) = two_tight_clusters(5, m);
1880 let wrong_t = uniform_grid(m + 3);
1881 assert!(
1882 kcfc_cluster(&data, &wrong_t, &KcfcConfig::default()).is_err(),
1883 "mismatched argvals should return Err"
1884 );
1885 }
1886
1887 #[test]
1888 fn test_kcfc_result_shapes() {
1889 let m = 20;
1890 let n_per = 5;
1891 let (data, t, _) = two_separated_clusters(n_per, m);
1892 let n = 2 * n_per;
1893 let result = kcfc_cluster(
1894 &data,
1895 &t,
1896 &KcfcConfig {
1897 k: 2,
1898 ncomp: 2,
1899 ..Default::default()
1900 },
1901 )
1902 .unwrap();
1903 assert_eq!(result.cluster.len(), n);
1904 assert_eq!(result.fpca_models.len(), 2);
1905 assert_eq!(result.reconstruction_errors.shape(), (n, 2));
1906 }
1907
1908 #[test]
1909 fn test_kcfc_ncomp_zero_returns_err() {
1910 let m = 20;
1914 let (data, t, _) = two_tight_clusters(5, m);
1915 let result = kcfc_cluster(
1916 &data,
1917 &t,
1918 &KcfcConfig {
1919 k: 2,
1920 ncomp: 0,
1921 ..Default::default()
1922 },
1923 );
1924 assert!(result.is_err(), "ncomp=0 must return Err, got Ok");
1925 if let Err(FdarError::InvalidParameter { parameter, .. }) = result {
1926 assert_eq!(parameter, "ncomp");
1927 } else {
1928 panic!("expected InvalidParameter {{ parameter: \"ncomp\" }}");
1929 }
1930 }
1931
1932 #[test]
1935 fn test_funfem_recovery() {
1936 let m = 40;
1937 let n_per = 10;
1938 let (data, t, ground_truth) = two_separated_clusters(n_per, m);
1939 let result = funfem_cluster(
1940 &data,
1941 &t,
1942 &FunFemConfig {
1943 k: 2,
1944 ncomp: 5,
1945 p_disc: 1,
1946 max_iter: 30,
1947 tol: 1e-5,
1948 seed: 42,
1949 },
1950 )
1951 .unwrap();
1952 let ari = adjusted_rand_index(&result.cluster, &ground_truth);
1953 assert!(
1954 ari >= 0.90,
1955 "funFEM ARI={ari:.3} should be >= 0.90 on well-separated data"
1956 );
1957 }
1958
1959 #[test]
1960 fn test_funfem_deterministic() {
1961 let m = 30;
1962 let n_per = 8;
1963 let (data, t, _) = two_separated_clusters(n_per, m);
1964 let cfg = FunFemConfig {
1965 k: 2,
1966 ncomp: 4,
1967 p_disc: 1,
1968 max_iter: 20,
1969 tol: 1e-5,
1970 seed: 99,
1971 };
1972 let r1 = funfem_cluster(&data, &t, &cfg).unwrap();
1973 let r2 = funfem_cluster(&data, &t, &cfg).unwrap();
1974 assert_eq!(
1975 r1.cluster, r2.cluster,
1976 "same seed must produce identical assignments"
1977 );
1978 }
1979
1980 #[test]
1981 fn test_funfem_invalid_k_zero() {
1982 let m = 20;
1983 let (data, t, _) = two_tight_clusters(5, m);
1984 assert!(
1985 funfem_cluster(
1986 &data,
1987 &t,
1988 &FunFemConfig {
1989 k: 0,
1990 ..FunFemConfig::default()
1991 }
1992 )
1993 .is_err(),
1994 "k=0 must return Err"
1995 );
1996 }
1997
1998 #[test]
1999 fn test_funfem_invalid_k_gt_n() {
2000 let m = 20;
2001 let (data, t, _) = two_tight_clusters(3, m);
2002 assert!(
2003 funfem_cluster(
2004 &data,
2005 &t,
2006 &FunFemConfig {
2007 k: 10,
2008 ..FunFemConfig::default()
2009 }
2010 )
2011 .is_err(),
2012 "k>n must return Err"
2013 );
2014 }
2015
2016 #[test]
2017 fn test_funfem_invalid_ncomp_zero() {
2018 let m = 20;
2019 let (data, t, _) = two_tight_clusters(5, m);
2020 assert!(
2021 funfem_cluster(
2022 &data,
2023 &t,
2024 &FunFemConfig {
2025 ncomp: 0,
2026 ..FunFemConfig::default()
2027 }
2028 )
2029 .is_err(),
2030 "ncomp=0 must return Err"
2031 );
2032 }
2033
2034 #[test]
2035 fn test_funfem_invalid_empty_data() {
2036 let data = FdMatrix::zeros(0, 0);
2037 let t: Vec<f64> = vec![];
2038 assert!(
2039 funfem_cluster(&data, &t, &FunFemConfig::default()).is_err(),
2040 "empty data must return Err"
2041 );
2042 }
2043
2044 #[test]
2045 fn test_funfem_invalid_argvals_mismatch() {
2046 let m = 20;
2047 let (data, _t, _) = two_tight_clusters(5, m);
2048 let wrong_t = uniform_grid(m + 2);
2049 assert!(
2050 funfem_cluster(&data, &wrong_t, &FunFemConfig::default()).is_err(),
2051 "argvals mismatch must return Err"
2052 );
2053 }
2054
2055 #[test]
2056 fn test_funfem_output_shapes() {
2057 let m = 30;
2058 let n_per = 6;
2059 let (data, t, _) = two_separated_clusters(n_per, m);
2060 let n = 2 * n_per;
2061 let result = funfem_cluster(
2062 &data,
2063 &t,
2064 &FunFemConfig {
2065 k: 2,
2066 ncomp: 4,
2067 p_disc: 1,
2068 max_iter: 10,
2069 tol: 1e-4,
2070 seed: 1,
2071 },
2072 )
2073 .unwrap();
2074 assert_eq!(result.cluster.len(), n);
2075 assert_eq!(result.membership.shape(), (n, 2));
2076 }
2077
2078 fn time_warped_clusters(n_per: usize, m: usize) -> (FdMatrix, Vec<f64>, Vec<usize>) {
2089 let t = uniform_grid(m);
2090 let n = 2 * n_per;
2091 let mut col_major = vec![0.0_f64; n * m];
2092 for i in 0..n_per {
2095 let alpha = 1.0 + 0.1 * (i as f64 / n_per as f64); for (j, &tj) in t.iter().enumerate() {
2097 let warped = tj.powf(alpha);
2098 col_major[i + j * n] = (2.0 * PI * warped).sin();
2099 }
2100 }
2101 for i in 0..n_per {
2103 for j in 0..m {
2104 col_major[(i + n_per) + j * n] = 8.0 + 0.05 * i as f64;
2105 }
2106 }
2107 let labels: Vec<usize> = (0..n).map(|i| if i < n_per { 0 } else { 1 }).collect();
2108 (
2109 FdMatrix::from_column_major(col_major, n, m).unwrap(),
2110 t,
2111 labels,
2112 )
2113 }
2114
2115 #[test]
2116 fn test_align_cluster_shape_shift() {
2117 let m = 30;
2119 let n_per = 6;
2120 let (data, t, ground_truth) = time_warped_clusters(n_per, m);
2121 let result = align_cluster_fd(
2122 &data,
2123 &t,
2124 &AlignClusterConfig {
2125 k: 2,
2126 max_iter: 15,
2127 seed: 42,
2128 use_amplitude_only: true,
2129 elastic_lambda: 0.0,
2130 karcher_max_iter: 10,
2131 karcher_tol: 1e-3,
2132 },
2133 )
2134 .unwrap();
2135 let ari = adjusted_rand_index(&result.cluster, &ground_truth);
2136 assert!(
2137 ari >= 0.90,
2138 "align_cluster ARI={ari:.3} on shape-distinct data should be >= 0.90"
2139 );
2140 }
2141
2142 #[test]
2143 fn test_align_cluster_recovery() {
2144 let m = 30;
2145 let n_per = 6;
2146 let (data, t, ground_truth) = two_separated_clusters(n_per, m);
2147 let result = align_cluster_fd(
2148 &data,
2149 &t,
2150 &AlignClusterConfig {
2151 k: 2,
2152 max_iter: 15,
2153 seed: 7,
2154 use_amplitude_only: true,
2155 elastic_lambda: 0.0,
2156 karcher_max_iter: 10,
2157 karcher_tol: 1e-3,
2158 },
2159 )
2160 .unwrap();
2161 let ari = adjusted_rand_index(&result.cluster, &ground_truth);
2162 assert!(
2163 ari >= 0.90,
2164 "align_cluster ARI={ari:.3} on amplitude-separated data should be >= 0.90"
2165 );
2166 }
2167
2168 #[test]
2169 fn test_align_cluster_invalid_k_zero() {
2170 let m = 20;
2171 let (data, t, _) = two_tight_clusters(5, m);
2172 assert!(
2173 align_cluster_fd(
2174 &data,
2175 &t,
2176 &AlignClusterConfig {
2177 k: 0,
2178 ..AlignClusterConfig::default()
2179 }
2180 )
2181 .is_err(),
2182 "k=0 must return Err"
2183 );
2184 }
2185
2186 #[test]
2187 fn test_align_cluster_invalid_k_gt_n() {
2188 let m = 20;
2189 let (data, t, _) = two_tight_clusters(3, m);
2190 assert!(
2191 align_cluster_fd(
2192 &data,
2193 &t,
2194 &AlignClusterConfig {
2195 k: 20,
2196 ..AlignClusterConfig::default()
2197 }
2198 )
2199 .is_err(),
2200 "k>n must return Err"
2201 );
2202 }
2203
2204 #[test]
2205 fn test_align_cluster_invalid_empty_data() {
2206 let data = FdMatrix::zeros(0, 0);
2207 let t: Vec<f64> = vec![];
2208 assert!(
2209 align_cluster_fd(&data, &t, &AlignClusterConfig::default()).is_err(),
2210 "empty data must return Err"
2211 );
2212 }
2213
2214 #[test]
2215 fn test_align_cluster_invalid_argvals_mismatch() {
2216 let m = 20;
2217 let (data, _t, _) = two_tight_clusters(5, m);
2218 let wrong_t = uniform_grid(m + 5);
2219 assert!(
2220 align_cluster_fd(&data, &wrong_t, &AlignClusterConfig::default()).is_err(),
2221 "argvals mismatch must return Err"
2222 );
2223 }
2224
2225 #[test]
2226 fn test_align_cluster_output_shapes() {
2227 let m = 20;
2228 let n_per = 4;
2229 let (data, t, _) = two_separated_clusters(n_per, m);
2230 let n = 2 * n_per;
2231 let result = align_cluster_fd(
2232 &data,
2233 &t,
2234 &AlignClusterConfig {
2235 k: 2,
2236 max_iter: 5,
2237 seed: 1,
2238 karcher_max_iter: 5,
2239 karcher_tol: 1e-2,
2240 ..AlignClusterConfig::default()
2241 },
2242 )
2243 .unwrap();
2244 assert_eq!(result.cluster.len(), n);
2245 assert_eq!(result.templates.len(), 2);
2246 assert!(result.templates.iter().all(|t| t.len() == m));
2247 assert_eq!(result.distances.shape(), (n, 2));
2248 }
2249}