1#![allow(non_snake_case)]
3use super::covariance::data_scaled_reg;
24use super::em::{compute_bic, compute_icl, hard_assignments, resp_to_membership};
25use super::init::kmeans_init_assignments;
26use crate::error::FdarError;
27use crate::matrix::FdMatrix;
28use crate::regression::fdata_to_pc_1d;
29use nalgebra::{DMatrix, SVD};
30use rand::prelude::*;
31
32#[derive(Debug, Clone, PartialEq)]
45#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
46#[non_exhaustive]
47pub struct FunHddcConfig {
48 pub k: usize,
50 pub d_k: usize,
52 pub max_iter: usize,
54 pub tol: f64,
56 pub n_init: usize,
58 pub seed: u64,
60 pub ncomp_init: usize,
62}
63
64impl Default for FunHddcConfig {
65 fn default() -> Self {
66 FunHddcConfig {
67 k: 2,
68 d_k: 2,
69 max_iter: 100,
70 tol: 1e-6,
71 n_init: 3,
72 seed: 42,
73 ncomp_init: 10,
74 }
75 }
76}
77
78#[derive(Debug, Clone)]
80#[non_exhaustive]
81pub struct FunHddcResult {
82 pub cluster: Vec<usize>,
84 pub membership: FdMatrix,
86 pub subspaces: Vec<FdMatrix>,
88 pub within_vars: Vec<Vec<f64>>,
90 pub noise_vars: Vec<f64>,
92 pub means: Vec<Vec<f64>>,
94 pub weights: Vec<f64>,
96 pub log_likelihood: f64,
98 pub bic: f64,
100 pub icl: f64,
102 pub iterations: usize,
104 pub converged: bool,
106 pub k: usize,
108}
109
110fn log_density_subspace(
123 diff: &[f64],
124 u_k: &[f64],
125 a_k: &[f64],
126 b_k: f64,
127 m: usize,
128 d_k_eff: usize,
129) -> f64 {
130 use std::f64::consts::PI;
131 let mut z = vec![0.0_f64; d_k_eff];
133 for j in 0..d_k_eff {
134 for r in 0..m {
135 z[j] += u_k[r + j * m] * diff[r];
136 }
137 }
138
139 let mut ll = 0.0_f64;
141 for j in 0..d_k_eff {
142 if a_k[j] <= 0.0 {
143 return f64::NEG_INFINITY;
144 }
145 ll -= 0.5 * (a_k[j].ln() + z[j].powi(2) / a_k[j]);
146 }
147
148 let diff_sq: f64 = diff.iter().map(|v| v * v).sum();
153 let z_sq: f64 = z.iter().map(|v| v * v).sum();
154 let complement_sq = (diff_sq - z_sq).max(0.0);
155
156 if b_k <= 0.0 {
157 return f64::NEG_INFINITY;
158 }
159 let m_minus_dk = (m - d_k_eff) as f64;
160 ll -= 0.5 * (m_minus_dk * b_k.ln() + complement_sq / b_k);
161
162 ll -= 0.5 * (m as f64) * (2.0 * PI).ln();
164 ll
165}
166
167fn normalize_log_probs(log_probs: &[f64], resp: &mut [f64]) -> f64 {
170 let k = log_probs.len();
171 let max_lp = log_probs.iter().copied().fold(f64::NEG_INFINITY, f64::max);
172 if max_lp == f64::NEG_INFINITY {
173 let uniform = 1.0 / k as f64;
174 for r in resp.iter_mut() {
175 *r = uniform;
176 }
177 return 0.0;
178 }
179 let lse = max_lp
180 + log_probs
181 .iter()
182 .map(|&lp| (lp - max_lp).exp())
183 .sum::<f64>()
184 .ln();
185 for c in 0..k {
186 resp[c] = (log_probs[c] - lse).exp();
187 }
188 lse
189}
190
191fn e_step_subspace(
195 data_rows: &[Vec<f64>], means: &[Vec<f64>], subspaces: &[Vec<f64>], within_vars: &[Vec<f64>], noise_vars: &[f64], weights: &[f64], k: usize,
202 m: usize,
203) -> (Vec<f64>, f64) {
204 let n = data_rows.len();
205 let mut resp = vec![0.0_f64; n * k];
206 let mut total_ll = 0.0_f64;
207
208 for i in 0..n {
209 let x = &data_rows[i];
210 let mut log_probs = vec![f64::NEG_INFINITY; k];
211 for c in 0..k {
212 if weights[c] > 1e-15 {
213 let d_k_eff = within_vars[c].len();
214 let diff: Vec<f64> = x
215 .iter()
216 .zip(means[c].iter())
217 .map(|(&xi, &mi)| xi - mi)
218 .collect();
219 let ld = log_density_subspace(
220 &diff,
221 &subspaces[c],
222 &within_vars[c],
223 noise_vars[c],
224 m,
225 d_k_eff,
226 );
227 log_probs[c] = weights[c].ln() + ld;
228 }
229 }
230 let mut r = vec![0.0_f64; k];
231 let ll_i = normalize_log_probs(&log_probs, &mut r);
232 resp[i * k..(i + 1) * k].copy_from_slice(&r);
233 total_ll += ll_i;
234 }
235 (resp, total_ll)
236}
237
238fn per_group_svd(
244 centered_rows: &[Vec<f64>],
245 d_k_req: usize,
246 m: usize,
247 reg: f64,
248) -> Option<(Vec<f64>, Vec<f64>)> {
249 let n_k = centered_rows.len();
250 if n_k == 0 {
251 return None;
252 }
253 let d_k_eff = d_k_req.min(n_k).min(m);
254 if d_k_eff == 0 {
255 return None;
256 }
257
258 let mut mat = DMatrix::<f64>::zeros(n_k, m);
260 for (i, row) in centered_rows.iter().enumerate() {
261 for j in 0..m {
262 mat[(i, j)] = row[j];
263 }
264 }
265
266 let svd = SVD::new(mat, true, true);
267 let v_t = svd.v_t?;
268 let singular_values = &svd.singular_values;
269
270 let mut u_k_flat = vec![0.0_f64; m * d_k_eff];
272 for j in 0..d_k_eff {
273 for r in 0..m {
274 u_k_flat[r + j * m] = v_t[(j, r)];
276 }
277 }
278
279 let n_k_f = n_k as f64;
281 let a_k: Vec<f64> = (0..d_k_eff)
282 .map(|j| {
283 let sv = singular_values[j];
284 (sv * sv / n_k_f).max(reg)
285 })
286 .collect();
287
288 Some((u_k_flat, a_k))
289}
290
291#[allow(clippy::too_many_arguments)]
293fn run_one_em(
294 data_rows: &[Vec<f64>],
295 k: usize,
296 m: usize,
297 d_k_req: usize,
298 max_iter: usize,
299 tol: f64,
300 init_assignments: &[usize],
301 reg: f64,
302) -> Option<(
303 Vec<f64>, Vec<Vec<f64>>, Vec<Vec<f64>>, Vec<Vec<f64>>, Vec<f64>, Vec<f64>, f64, usize, bool, )> {
313 let n = data_rows.len();
314
315 let mut means: Vec<Vec<f64>> = vec![vec![0.0_f64; m]; k];
317 let mut counts = vec![0usize; k];
318 for (i, &c) in init_assignments.iter().enumerate() {
319 counts[c] += 1;
320 for j in 0..m {
321 means[c][j] += data_rows[i][j];
322 }
323 }
324 for c in 0..k {
325 let nc = counts[c].max(1);
326 for j in 0..m {
327 means[c][j] /= nc as f64;
328 }
329 }
330 let mut weights: Vec<f64> = counts.iter().map(|&c| c.max(1) as f64 / n as f64).collect();
331
332 let mut subspaces: Vec<Vec<f64>> = vec![vec![0.0_f64; m * d_k_req.min(m)]; k];
334 let mut within_vars: Vec<Vec<f64>> = vec![vec![reg; d_k_req.min(m)]; k];
335 let mut noise_vars: Vec<f64> = vec![reg; k];
336
337 for c in 0..k {
339 let member_rows: Vec<Vec<f64>> = (0..n)
340 .filter(|&i| init_assignments[i] == c)
341 .map(|i| {
342 data_rows[i]
343 .iter()
344 .zip(means[c].iter())
345 .map(|(&x, &mu)| x - mu)
346 .collect()
347 })
348 .collect();
349
350 if let Some((u_k, a_k)) = per_group_svd(&member_rows, d_k_req, m, reg) {
351 let d_k_eff = a_k.len();
352 subspaces[c] = u_k;
353 within_vars[c] = a_k.clone();
354 let total_var: f64 = member_rows
356 .iter()
357 .flat_map(|r| r.iter())
358 .map(|v| v * v)
359 .sum::<f64>()
360 / member_rows.len().max(1) as f64;
361 let subspace_var: f64 = a_k.iter().sum();
362 let complement_var = (total_var - subspace_var).max(0.0);
363 let m_minus_dk = (m - d_k_eff) as f64;
364 noise_vars[c] = if m_minus_dk > 0.0 {
365 (complement_var / m_minus_dk).max(reg)
366 } else {
367 reg
368 };
369 }
370 }
371
372 let mut resp = vec![0.0_f64; n * k];
373 let mut prev_ll = f64::NEG_INFINITY;
374 let mut converged = false;
375 let mut iterations = 0usize;
376
377 for iter in 0..max_iter {
378 iterations = iter + 1;
379
380 let (new_resp, ll) = e_step_subspace(
382 data_rows,
383 &means,
384 &subspaces,
385 &within_vars,
386 &noise_vars,
387 &weights,
388 k,
389 m,
390 );
391 resp = new_resp;
392
393 if (ll - prev_ll).abs() < tol && iter > 0 {
394 converged = true;
395 break;
396 }
397 prev_ll = ll;
398
399 let mut new_means = vec![vec![0.0_f64; m]; k];
402 let mut nk_vec = vec![0.0_f64; k];
403 for i in 0..n {
404 for c in 0..k {
405 let r = resp[i * k + c];
406 nk_vec[c] += r;
407 for j in 0..m {
408 new_means[c][j] += r * data_rows[i][j];
409 }
410 }
411 }
412 for c in 0..k {
413 let nk = nk_vec[c];
414 if nk > 1e-15 {
415 for j in 0..m {
416 new_means[c][j] /= nk;
417 }
418 }
419 }
420 let n_f = n as f64;
421 weights = nk_vec.iter().map(|&nk| nk / n_f).collect();
422 means = new_means;
423
424 for c in 0..k {
426 let nk = nk_vec[c];
427 if nk < 1e-15 {
428 let d_k_eff = d_k_req.min(m);
430 subspaces[c] = vec![0.0_f64; m * d_k_eff];
431 within_vars[c] = vec![reg; d_k_eff];
432 noise_vars[c] = reg;
433 continue;
434 }
435
436 let mut w_rows: Vec<Vec<f64>> = Vec::with_capacity(n);
438 for i in 0..n {
439 let sqrt_r = resp[i * k + c].sqrt();
440 if sqrt_r > 1e-15 {
441 let row: Vec<f64> = data_rows[i]
442 .iter()
443 .zip(means[c].iter())
444 .map(|(&x, &mu)| sqrt_r * (x - mu))
445 .collect();
446 w_rows.push(row);
447 }
448 }
449
450 if w_rows.is_empty() {
451 let d_k_eff = d_k_req.min(m);
452 subspaces[c] = vec![0.0_f64; m * d_k_eff];
453 within_vars[c] = vec![reg; d_k_eff];
454 noise_vars[c] = reg;
455 continue;
456 }
457
458 if let Some((u_k, a_k)) = per_group_svd(&w_rows, d_k_req, m, reg) {
459 let d_k_eff = a_k.len();
460 let a_k_rescaled: Vec<f64> = a_k.iter().map(|&a| a.max(reg)).collect();
462 subspaces[c] = u_k;
463 within_vars[c] = a_k_rescaled.clone();
464
465 let total_wvar: f64 = w_rows
467 .iter()
468 .flat_map(|r| r.iter())
469 .map(|v| v * v)
470 .sum::<f64>()
471 / w_rows.len() as f64;
472 let subspace_var: f64 = a_k_rescaled.iter().sum();
473 let complement_var = (total_wvar - subspace_var).max(0.0);
474 let m_minus_dk = (m - d_k_eff) as f64;
475 noise_vars[c] = if m_minus_dk > 0.0 {
476 (complement_var / m_minus_dk).max(reg)
477 } else {
478 reg
479 };
480 }
481 }
482 }
483
484 let (final_resp, final_ll) = e_step_subspace(
486 data_rows,
487 &means,
488 &subspaces,
489 &within_vars,
490 &noise_vars,
491 &weights,
492 k,
493 m,
494 );
495
496 Some((
497 final_resp,
498 means,
499 subspaces,
500 within_vars,
501 noise_vars,
502 weights,
503 final_ll,
504 iterations,
505 converged,
506 ))
507}
508
509#[must_use = "expensive computation whose result should not be discarded"]
554pub fn funhddC_cluster(
555 data: &FdMatrix,
556 argvals: &[f64],
557 config: &FunHddcConfig,
558) -> Result<FunHddcResult, FdarError> {
559 let (n, m) = data.shape();
560
561 if n == 0 || m == 0 {
563 return Err(FdarError::InvalidDimension {
564 parameter: "data",
565 expected: "non-empty matrix".to_string(),
566 actual: format!("{n}x{m}"),
567 });
568 }
569 if argvals.len() != m {
570 return Err(FdarError::InvalidDimension {
571 parameter: "argvals",
572 expected: format!("{m} elements"),
573 actual: format!("{} elements", argvals.len()),
574 });
575 }
576 if config.k == 0 {
577 return Err(FdarError::InvalidParameter {
578 parameter: "k",
579 message: "must be >= 1".to_string(),
580 });
581 }
582 if config.k > n {
583 return Err(FdarError::InvalidParameter {
584 parameter: "k",
585 message: format!("must be <= n ({n}), got {}", config.k),
586 });
587 }
588 if config.d_k == 0 {
589 return Err(FdarError::InvalidParameter {
590 parameter: "d_k",
591 message: "must be >= 1".to_string(),
592 });
593 }
594 if config.d_k >= m {
595 return Err(FdarError::InvalidParameter {
596 parameter: "d_k",
597 message: format!("must be < m ({m}), got {}", config.d_k),
598 });
599 }
600
601 let data_rm = data.to_row_major();
603 let data_rows: Vec<Vec<f64>> = (0..n)
604 .map(|i| data_rm[i * m..(i + 1) * m].to_vec())
605 .collect();
606
607 let reg = data_scaled_reg(&data_rows, m);
609
610 let ncomp_init = config.ncomp_init.min(n).min(m).max(1);
612 let fpca = fdata_to_pc_1d(data, ncomp_init, argvals)?;
613 let score_mat = &fpca.scores;
614 let d_feat = score_mat.ncols();
615 let features: Vec<Vec<f64>> = (0..n)
616 .map(|i| (0..d_feat).map(|j| score_mat[(i, j)]).collect())
617 .collect();
618
619 let k = config.k;
620
621 let mut best: Option<(
623 Vec<f64>,
624 Vec<Vec<f64>>,
625 Vec<Vec<f64>>,
626 Vec<Vec<f64>>,
627 Vec<f64>,
628 Vec<f64>,
629 f64,
630 usize,
631 bool,
632 )> = None;
633
634 for init_idx in 0..config.n_init {
635 let seed = config.seed.wrapping_add(init_idx as u64 * 1000);
636 let mut rng = StdRng::seed_from_u64(seed);
637 let init_assignments = kmeans_init_assignments(&features, k, &mut rng);
638
639 if let Some(result) = run_one_em(
640 &data_rows,
641 k,
642 m,
643 config.d_k,
644 config.max_iter,
645 config.tol,
646 &init_assignments,
647 reg,
648 ) {
649 let ll = result.6;
650 let is_better = best.as_ref().map_or(true, |b| ll > b.6);
651 if is_better {
652 best = Some(result);
653 }
654 }
655 }
656
657 let (
658 resp,
659 means,
660 subspaces_flat,
661 within_vars,
662 noise_vars,
663 weights,
664 log_likelihood,
665 iterations,
666 converged,
667 ) = best.ok_or_else(|| FdarError::ComputationFailed {
668 operation: "funhddC_cluster",
669 detail: "all EM restarts failed".to_string(),
670 })?;
671
672 let d_k_eff = within_vars.first().map_or(1, |v| v.len());
676 let subspace_params = k * (m * d_k_eff - d_k_eff * (d_k_eff.saturating_sub(1)) / 2);
677 let var_params = k * d_k_eff + k; let n_params = subspace_params + var_params + (k - 1);
679
680 let bic = compute_bic(log_likelihood, n, n_params);
681 let icl = compute_icl(bic, &resp, n, k);
682
683 let cluster = hard_assignments(&resp, n, k);
684 let membership = resp_to_membership(&resp, n, k);
685
686 let subspaces: Vec<FdMatrix> = subspaces_flat
688 .into_iter()
689 .zip(within_vars.iter())
690 .map(|(flat, av)| {
691 let d = av.len();
692 if d == 0 || flat.is_empty() {
693 FdMatrix::zeros(m, d.max(1))
696 } else {
697 FdMatrix::from_column_major(flat, m, d).unwrap_or_else(|_| FdMatrix::zeros(m, d))
698 }
699 })
700 .collect();
701
702 Ok(FunHddcResult {
703 cluster,
704 membership,
705 subspaces,
706 within_vars,
707 noise_vars,
708 means,
709 weights,
710 log_likelihood,
711 bic,
712 icl,
713 iterations,
714 converged,
715 k,
716 })
717}
718
719#[cfg(test)]
724mod tests {
725 use super::*;
726 use crate::test_helpers::{adjusted_rand_index, uniform_grid};
727
728 fn two_separated_clusters(n_per: usize, m: usize) -> (FdMatrix, Vec<f64>, Vec<usize>) {
731 let argvals = uniform_grid(m);
732 let n = 2 * n_per;
733 let mut data_rm = vec![0.0_f64; n * m];
734 let mut labels = vec![0usize; n];
735 for i in 0..n_per {
736 for j in 0..m {
737 data_rm[i * m + j] = argvals[j].sin();
738 }
739 labels[i] = 0;
740 }
741 for i in 0..n_per {
742 for j in 0..m {
743 data_rm[(n_per + i) * m + j] = argvals[j].sin() + 5.0;
744 }
745 labels[n_per + i] = 1;
746 }
747 let mut col_major = vec![0.0_f64; n * m];
749 for ii in 0..n {
750 for jj in 0..m {
751 col_major[ii + jj * n] = data_rm[ii * m + jj];
752 }
753 }
754 let data = FdMatrix::from_column_major(col_major, n, m).unwrap();
755 (data, argvals, labels)
756 }
757
758 #[test]
759 fn test_funhddC_recovery() {
760 let (data, argvals, labels) = two_separated_clusters(15, 20);
761 let config = FunHddcConfig {
762 k: 2,
763 d_k: 2,
764 max_iter: 100,
765 tol: 1e-6,
766 n_init: 3,
767 seed: 42,
768 ncomp_init: 8,
769 };
770 let result = funhddC_cluster(&data, &argvals, &config).unwrap();
771 let ari = adjusted_rand_index(&labels, &result.cluster);
772 assert!(ari >= 0.90, "Recovery ARI should be >= 0.90, got {ari:.4}");
773 }
774
775 #[test]
776 fn test_funhddC_bic_finite() {
777 let (data, argvals, _) = two_separated_clusters(15, 20);
778 let config = FunHddcConfig {
779 k: 2,
780 d_k: 2,
781 max_iter: 100,
782 tol: 1e-6,
783 n_init: 3,
784 seed: 42,
785 ncomp_init: 8,
786 };
787 let result = funhddC_cluster(&data, &argvals, &config).unwrap();
788 assert!(
789 result.bic.is_finite(),
790 "BIC should be finite, got {}",
791 result.bic
792 );
793 assert!(
794 result.icl.is_finite(),
795 "ICL should be finite, got {}",
796 result.icl
797 );
798 assert!(
799 result.log_likelihood.is_finite(),
800 "log-likelihood should be finite, got {}",
801 result.log_likelihood
802 );
803 }
804
805 #[test]
806 fn test_funhddC_deterministic() {
807 let (data, argvals, _) = two_separated_clusters(15, 20);
808 let config = FunHddcConfig {
809 k: 2,
810 d_k: 2,
811 max_iter: 100,
812 tol: 1e-6,
813 n_init: 3,
814 seed: 99,
815 ncomp_init: 8,
816 };
817 let r1 = funhddC_cluster(&data, &argvals, &config).unwrap();
818 let r2 = funhddC_cluster(&data, &argvals, &config).unwrap();
819 assert_eq!(
820 r1.cluster, r2.cluster,
821 "Same seed must give identical cluster assignments"
822 );
823 }
824
825 #[test]
826 fn test_funhddC_invalid_empty() {
827 let data = FdMatrix::zeros(0, 10);
828 let argvals = uniform_grid(10);
829 let config = FunHddcConfig {
830 k: 2,
831 ..Default::default()
832 };
833 assert!(funhddC_cluster(&data, &argvals, &config).is_err());
834 }
835
836 #[test]
837 fn test_funhddC_invalid_k_zero() {
838 let data = FdMatrix::zeros(5, 10);
839 let argvals = uniform_grid(10);
840 let config = FunHddcConfig {
841 k: 0,
842 ..Default::default()
843 };
844 assert!(funhddC_cluster(&data, &argvals, &config).is_err());
845 }
846
847 #[test]
848 fn test_funhddC_invalid_k_exceeds_n() {
849 let data = FdMatrix::zeros(3, 10);
850 let argvals = uniform_grid(10);
851 let config = FunHddcConfig {
852 k: 5,
853 ..Default::default()
854 };
855 assert!(funhddC_cluster(&data, &argvals, &config).is_err());
856 }
857
858 #[test]
859 fn test_funhddC_invalid_dk_ge_m() {
860 let data = FdMatrix::zeros(5, 10);
861 let argvals = uniform_grid(10);
862 let config = FunHddcConfig {
863 k: 2,
864 d_k: 10,
865 ..Default::default()
866 };
867 assert!(funhddC_cluster(&data, &argvals, &config).is_err());
868 }
869
870 #[test]
871 fn test_funhddC_invalid_argvals_mismatch() {
872 let data = FdMatrix::zeros(5, 10);
873 let argvals = uniform_grid(8); let config = FunHddcConfig {
875 k: 2,
876 ..Default::default()
877 };
878 assert!(funhddC_cluster(&data, &argvals, &config).is_err());
879 }
880}