1use crate::error::FdarError;
66use crate::matrix::FdMatrix;
67use crate::regression::{fdata_to_pc_1d, fdata_to_pls_1d};
68use crate::wavelet::{decompose_matrix, reconstruct, BoundaryMode, WaveletCoeffs, WaveletFamily};
69
70#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
72#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
73#[non_exhaustive]
74pub enum WcrMethod {
75 #[default]
78 Pcr,
79 Pls,
82}
83
84#[derive(Debug, Clone, PartialEq)]
90#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
91#[non_exhaustive]
92pub struct WcrConfig {
93 pub family: WaveletFamily,
95 pub mode: BoundaryMode,
97 pub level: Option<usize>,
99 pub ncomp: usize,
101 pub method: WcrMethod,
103}
104
105impl Default for WcrConfig {
106 fn default() -> Self {
107 Self {
108 family: WaveletFamily::Daubechies(4),
109 mode: BoundaryMode::Periodic,
110 level: None,
111 ncomp: 5,
112 method: WcrMethod::Pcr,
113 }
114 }
115}
116
117#[derive(Debug, Clone, PartialEq)]
123#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
124#[non_exhaustive]
125pub struct WcrResult {
126 pub intercept: f64,
136 pub beta_t: Vec<f64>,
138 pub fitted_values: Vec<f64>,
140 pub residuals: Vec<f64>,
142 pub ncomp: usize,
144 pub method: WcrMethod,
146 pub coeff_weights: Vec<f64>,
148 pub family: WaveletFamily,
150 pub mode: BoundaryMode,
152 pub level: usize,
154}
155
156#[derive(Debug, Clone, PartialEq)]
162pub(crate) struct CoeffLayout {
163 pub(crate) approx_len: usize,
165 pub(crate) detail_lens: Vec<usize>,
167 pub(crate) signal_len: usize,
169 pub(crate) family: WaveletFamily,
171 pub(crate) mode: BoundaryMode,
173 pub(crate) level_lens: Vec<usize>,
175}
176
177impl CoeffLayout {
178 pub(crate) fn total_len(&self) -> usize {
180 self.approx_len + self.detail_lens.iter().sum::<usize>()
181 }
182
183 pub(crate) fn levels(&self) -> usize {
185 self.detail_lens.len()
186 }
187}
188
189fn coeffs_to_row(coeffs: &WaveletCoeffs) -> Vec<f64> {
191 let mut row = Vec::with_capacity(
192 coeffs.approx.len() + coeffs.details.iter().map(Vec::len).sum::<usize>(),
193 );
194 row.extend_from_slice(&coeffs.approx);
195 for band in &coeffs.details {
196 row.extend_from_slice(band);
197 }
198 row
199}
200
201pub(crate) fn curves_to_coeff_design(
217 data: &FdMatrix,
218 family: WaveletFamily,
219 mode: BoundaryMode,
220 level: Option<usize>,
221) -> Result<(FdMatrix, CoeffLayout), FdarError> {
222 let per_curve = decompose_matrix(data, family.clone(), mode, level)?;
223 let first = &per_curve[0];
225 let layout = CoeffLayout {
226 approx_len: first.approx.len(),
227 detail_lens: first.details.iter().map(Vec::len).collect(),
228 signal_len: first.signal_len,
229 family,
230 mode,
231 level_lens: first.level_lens.clone(),
232 };
233 let p = layout.total_len();
234 let n = per_curve.len();
235
236 let mut flat = vec![0.0_f64; n * p];
239 for (i, coeffs) in per_curve.iter().enumerate() {
240 if coeffs.approx.len() != layout.approx_len
241 || coeffs.details.len() != layout.detail_lens.len()
242 || coeffs
243 .details
244 .iter()
245 .zip(&layout.detail_lens)
246 .any(|(band, &len)| band.len() != len)
247 || coeffs.signal_len != layout.signal_len
248 {
249 return Err(FdarError::InvalidDimension {
250 parameter: "data",
251 expected: format!("all curves share curve-0 coefficient layout (P = {p})"),
252 actual: format!("curve {i} produced a different band structure"),
253 });
254 }
255 let row = coeffs_to_row(coeffs);
256 for (j, &v) in row.iter().enumerate() {
257 flat[i + j * n] = v;
258 }
259 }
260
261 let design = FdMatrix::from_column_major(flat, n, p)?;
262 Ok((design, layout))
263}
264
265pub(crate) fn coeff_weights_to_beta_t(
278 weights: &[f64],
279 layout: &CoeffLayout,
280) -> Result<Vec<f64>, FdarError> {
281 let expected = layout.total_len();
282 if weights.len() != expected {
283 return Err(FdarError::InvalidDimension {
284 parameter: "weights",
285 expected: format!("{expected} coefficients (approx + all detail bands)"),
286 actual: format!("{} coefficients", weights.len()),
287 });
288 }
289
290 let approx = weights[..layout.approx_len].to_vec();
291 let mut details: Vec<Vec<f64>> = Vec::with_capacity(layout.detail_lens.len());
292 let mut offset = layout.approx_len;
293 for &len in &layout.detail_lens {
294 details.push(weights[offset..offset + len].to_vec());
295 offset += len;
296 }
297
298 let coeffs = WaveletCoeffs {
299 approx,
300 details,
301 levels: layout.levels(),
302 signal_len: layout.signal_len,
303 family: layout.family.clone(),
304 mode: layout.mode,
305 level_lens: layout.level_lens.clone(),
306 };
307 reconstruct(&coeffs)
308}
309
310fn design_with_intercept(scores: &FdMatrix, ncomp: usize) -> FdMatrix {
318 let n = scores.nrows();
319 let mut design = FdMatrix::zeros(n, 1 + ncomp);
320 for i in 0..n {
321 design[(i, 0)] = 1.0;
322 for k in 0..ncomp {
323 design[(i, 1 + k)] = scores[(i, k)];
324 }
325 }
326 design
327}
328
329fn ols_solve(x: &FdMatrix, y: &[f64]) -> Result<Vec<f64>, FdarError> {
331 let (n, p) = x.shape();
332 if n < p || p == 0 {
333 return Err(FdarError::InvalidDimension {
334 parameter: "design matrix",
335 expected: format!("n >= p and p > 0 (p={p})"),
336 actual: format!("n={n}, p={p}"),
337 });
338 }
339 let mut xtx = vec![0.0_f64; p * p];
341 let mut xty = vec![0.0_f64; p];
342 for a in 0..p {
343 for b in 0..p {
344 let mut s = 0.0;
345 for i in 0..n {
346 s += x[(i, a)] * x[(i, b)];
347 }
348 xtx[a + b * p] = s;
349 }
350 let mut sy = 0.0;
351 for i in 0..n {
352 sy += x[(i, a)] * y[i];
353 }
354 xty[a] = sy;
355 }
356 let l = cholesky_factor(&xtx, p)?;
357 Ok(cholesky_solve(&l, &xty, p))
358}
359
360fn cholesky_factor(a: &[f64], p: usize) -> Result<Vec<f64>, FdarError> {
362 let mut l = vec![0.0_f64; p * p];
363 for j in 0..p {
364 let mut diag = a[j + j * p];
365 for k in 0..j {
366 diag -= l[j + k * p] * l[j + k * p];
367 }
368 if diag <= 0.0 {
369 return Err(FdarError::ComputationFailed {
370 operation: "Cholesky factorization (wcr OLS)",
371 detail: "design matrix X'X is not positive definite; try reducing ncomp"
372 .to_string(),
373 });
374 }
375 let ljj = diag.sqrt();
376 l[j + j * p] = ljj;
377 for i in (j + 1)..p {
378 let mut s = a[i + j * p];
379 for k in 0..j {
380 s -= l[i + k * p] * l[j + k * p];
381 }
382 l[i + j * p] = s / ljj;
383 }
384 }
385 Ok(l)
386}
387
388fn cholesky_solve(l: &[f64], rhs: &[f64], p: usize) -> Vec<f64> {
390 let mut z = vec![0.0_f64; p];
392 for i in 0..p {
393 let mut s = rhs[i];
394 for k in 0..i {
395 s -= l[i + k * p] * z[k];
396 }
397 z[i] = s / l[i + i * p];
398 }
399 let mut b = vec![0.0_f64; p];
401 for i in (0..p).rev() {
402 let mut s = z[i];
403 for k in (i + 1)..p {
404 s -= l[k + i * p] * b[k];
405 }
406 b[i] = s / l[i + i * p];
407 }
408 b
409}
410
411fn recover_coeff_weights(
420 design: &FdMatrix,
421 fitted: &[f64],
422 intercept: f64,
423) -> Result<Vec<f64>, FdarError> {
424 let (n, p) = design.shape();
425 let col_means: Vec<f64> = (0..p)
427 .map(|j| design.column(j).iter().sum::<f64>() / n as f64)
428 .collect();
429 let mut xc = FdMatrix::zeros(n, p);
431 for j in 0..p {
432 for i in 0..n {
433 xc[(i, j)] = design[(i, j)] - col_means[j];
434 }
435 }
436 let r: Vec<f64> = fitted.iter().map(|&f| f - intercept).collect();
437 let mut xtx = vec![0.0_f64; p * p];
439 let mut xtr = vec![0.0_f64; p];
440 for a in 0..p {
441 for b in 0..p {
442 let mut s = 0.0;
443 for i in 0..n {
444 s += xc[(i, a)] * xc[(i, b)];
445 }
446 xtx[a + b * p] = s;
447 }
448 let mut sr = 0.0;
449 for i in 0..n {
450 sr += xc[(i, a)] * r[i];
451 }
452 xtr[a] = sr;
453 }
454 let trace: f64 = (0..p).map(|j| xtx[j + j * p]).sum();
458 let eps = 1e-10 * (trace / p as f64).max(1e-12);
459 for j in 0..p {
460 xtx[j + j * p] += eps;
461 }
462 let l = cholesky_factor(&xtx, p)?;
463 Ok(cholesky_solve(&l, &xtr, p))
464}
465
466fn compute_fitted(design: &FdMatrix, coeffs: &[f64]) -> Vec<f64> {
468 let (n, p) = design.shape();
469 (0..n)
470 .map(|i| {
471 let mut yhat = 0.0;
472 for j in 0..p {
473 yhat += design[(i, j)] * coeffs[j];
474 }
475 yhat
476 })
477 .collect()
478}
479
480#[must_use = "expensive computation whose result should not be discarded"]
506pub fn wcr(data: &FdMatrix, y: &[f64], config: &WcrConfig) -> Result<WcrResult, FdarError> {
507 let (n, m) = data.shape();
508 if n < 3 {
509 return Err(FdarError::InvalidDimension {
510 parameter: "data",
511 expected: "at least 3 rows (observations)".to_string(),
512 actual: format!("{n} rows"),
513 });
514 }
515 if m == 0 {
516 return Err(FdarError::InvalidDimension {
517 parameter: "data",
518 expected: "at least 1 column (evaluation point)".to_string(),
519 actual: format!("{m} columns"),
520 });
521 }
522 if y.len() != n {
523 return Err(FdarError::InvalidDimension {
524 parameter: "y",
525 expected: format!("{n} elements (== data rows)"),
526 actual: format!("{} elements", y.len()),
527 });
528 }
529 if config.ncomp == 0 {
530 return Err(FdarError::InvalidParameter {
531 parameter: "ncomp",
532 message: "ncomp must be >= 1".to_string(),
533 });
534 }
535
536 let (design, layout) =
538 curves_to_coeff_design(data, config.family.clone(), config.mode, config.level)?;
539 let p = design.ncols();
540
541 let argvals: Vec<f64> = (0..p).map(|j| j as f64).collect();
543
544 let ncomp = config.ncomp.min(n.saturating_sub(1)).min(p);
550
551 let (scores, ncomp) = match config.method {
554 WcrMethod::Pcr => {
555 let fpca = fdata_to_pc_1d(&design, ncomp, &argvals)?;
556 let k = fpca.scores.ncols();
557 (fpca.scores, k)
558 }
559 WcrMethod::Pls => {
560 let pls = fdata_to_pls_1d(&design, y, ncomp, &argvals)?;
561 let k = pls.scores.ncols();
562 (pls.scores, k)
563 }
564 };
565 let ols_design = design_with_intercept(&scores, ncomp);
566 let coeffs = ols_solve(&ols_design, y)?;
567 let intercept = coeffs[0];
568 let fitted_values = compute_fitted(&ols_design, &coeffs);
569
570 let coeff_weights = recover_coeff_weights(&design, &fitted_values, intercept)?;
581
582 let (n_rows, p_cols) = design.shape();
590 let intercept = {
591 let offset: f64 = (0..p_cols)
592 .map(|j| {
593 let col_mean = design.column(j).iter().sum::<f64>() / n_rows as f64;
594 col_mean * coeff_weights[j]
595 })
596 .sum();
597 intercept - offset
598 };
599
600 let beta_t = coeff_weights_to_beta_t(&coeff_weights, &layout)?;
602
603 let residuals: Vec<f64> = y
604 .iter()
605 .zip(&fitted_values)
606 .map(|(&yi, &yh)| yi - yh)
607 .collect();
608
609 Ok(WcrResult {
610 intercept,
611 beta_t,
612 fitted_values,
613 residuals,
614 ncomp,
615 method: config.method,
616 coeff_weights,
617 family: config.family.clone(),
618 mode: config.mode,
619 level: layout.levels(),
620 })
621}
622
623impl WcrResult {
624 pub fn predict(&self, new: &FdMatrix) -> Result<Vec<f64>, FdarError> {
644 let train_m = self.beta_t.len();
645 if new.nrows() == 0 {
646 return Err(FdarError::InvalidDimension {
647 parameter: "new",
648 expected: "at least 1 row (curve)".to_string(),
649 actual: "0 rows".to_string(),
650 });
651 }
652 if new.ncols() != train_m {
653 return Err(FdarError::InvalidDimension {
654 parameter: "new",
655 expected: format!("{train_m} columns (== training grid length)"),
656 actual: format!("{} columns", new.ncols()),
657 });
658 }
659 let (design, _layout) =
662 curves_to_coeff_design(new, self.family.clone(), self.mode, Some(self.level))?;
663 if design.ncols() != self.coeff_weights.len() {
664 return Err(FdarError::InvalidDimension {
665 parameter: "new",
666 expected: format!(
667 "coefficient-space width {} (== stored coeff_weights)",
668 self.coeff_weights.len()
669 ),
670 actual: format!("{} coefficients", design.ncols()),
671 });
672 }
673 Ok(compute_fitted_affine(
674 &design,
675 &self.coeff_weights,
676 self.intercept,
677 ))
678 }
679
680 #[must_use]
682 pub fn beta_t(&self) -> &[f64] {
683 &self.beta_t
684 }
685
686 #[must_use]
688 pub fn coefficient_function(&self) -> &[f64] {
689 &self.beta_t
690 }
691
692 #[must_use]
694 pub fn fitted_values(&self) -> &[f64] {
695 &self.fitted_values
696 }
697}
698
699#[derive(Debug, Clone, PartialEq)]
727#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
728#[non_exhaustive]
729pub struct WnetConfig {
730 pub family: WaveletFamily,
732 pub mode: BoundaryMode,
734 pub level: Option<usize>,
736 pub alpha: f64,
739 pub lambda_grid: Option<Vec<f64>>,
741 pub n_lambda: usize,
744 pub n_folds: usize,
746 pub seed: u64,
748 pub max_iter: usize,
750 pub tol: f64,
752}
753
754impl Default for WnetConfig {
755 fn default() -> Self {
756 Self {
757 family: WaveletFamily::Daubechies(4),
758 mode: BoundaryMode::Periodic,
759 level: None,
760 alpha: 0.5,
761 lambda_grid: None,
762 n_lambda: 50,
763 n_folds: 5,
764 seed: 0,
765 max_iter: 1000,
766 tol: 1e-6,
767 }
768 }
769}
770
771#[derive(Debug, Clone, PartialEq)]
779#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
780#[non_exhaustive]
781pub struct WnetResult {
782 pub intercept: f64,
787 pub beta_t: Vec<f64>,
789 pub fitted_values: Vec<f64>,
791 pub residuals: Vec<f64>,
793 pub coeff_weights: Vec<f64>,
796 pub selected: Vec<usize>,
798 pub lambda: f64,
800 pub alpha: f64,
802 pub family: WaveletFamily,
804 pub mode: BoundaryMode,
806 pub level: usize,
808}
809
810#[inline]
812fn soft_threshold(z: f64, gamma: f64) -> f64 {
813 if z > gamma {
814 z - gamma
815 } else if z < -gamma {
816 z + gamma
817 } else {
818 0.0
819 }
820}
821
822pub(crate) fn elastic_net_cd(
844 design: &FdMatrix,
845 y: &[f64],
846 lambda: f64,
847 alpha: f64,
848 max_iter: usize,
849 tol: f64,
850) -> Result<(f64, Vec<f64>), FdarError> {
851 let (n, p) = design.shape();
852 if y.len() != n {
853 return Err(FdarError::InvalidDimension {
854 parameter: "y",
855 expected: format!("{n} elements (== design rows)"),
856 actual: format!("{} elements", y.len()),
857 });
858 }
859 if !(0.0..=1.0).contains(&alpha) {
860 return Err(FdarError::InvalidParameter {
861 parameter: "alpha",
862 message: format!("alpha must be in [0, 1], got {alpha}"),
863 });
864 }
865 if lambda < 0.0 || !lambda.is_finite() {
866 return Err(FdarError::InvalidParameter {
867 parameter: "lambda",
868 message: format!("lambda must be finite and >= 0, got {lambda}"),
869 });
870 }
871 if !tol.is_finite() || tol < 0.0 {
872 return Err(FdarError::InvalidParameter {
873 parameter: "tol",
874 message: format!("tol must be finite and >= 0, got {tol}"),
875 });
876 }
877 if max_iter == 0 {
878 return Err(FdarError::InvalidParameter {
879 parameter: "max_iter",
880 message: "max_iter must be >= 1".to_string(),
881 });
882 }
883
884 let n_f = n as f64;
885 let mu_y = y.iter().sum::<f64>() / n_f;
886 let y_centered: Vec<f64> = y.iter().map(|&v| v - mu_y).collect();
887
888 let col_means: Vec<f64> = (0..p)
890 .map(|j| design.column(j).iter().sum::<f64>() / n_f)
891 .collect();
892 let mut xc = vec![0.0_f64; n * p]; let mut col_norm_sq_over_n = vec![0.0_f64; p];
894 for j in 0..p {
895 let mu = col_means[j];
896 let mut norm_sq = 0.0;
897 let col = design.column(j);
898 for i in 0..n {
899 let v = col[i] - mu;
900 xc[i + j * n] = v;
901 norm_sq += v * v;
902 }
903 col_norm_sq_over_n[j] = norm_sq / n_f;
904 }
905
906 let mut beta = vec![0.0_f64; p];
909 let mut fit = vec![0.0_f64; n]; let l1 = lambda * alpha;
911 let l2 = lambda * (1.0 - alpha);
912
913 for _sweep in 0..max_iter {
914 let mut max_delta = 0.0_f64;
915 for j in 0..p {
916 let denom = col_norm_sq_over_n[j] + l2;
917 if denom <= 0.0 {
918 if beta[j] != 0.0 {
920 let old = beta[j];
921 for i in 0..n {
922 fit[i] -= old * xc[i + j * n];
923 }
924 max_delta = max_delta.max(old.abs());
925 beta[j] = 0.0;
926 }
927 continue;
928 }
929 let old = beta[j];
932 let mut dot = 0.0;
933 for i in 0..n {
934 let r = y_centered[i] - fit[i] + old * xc[i + j * n];
935 dot += xc[i + j * n] * r;
936 }
937 let z = dot / n_f;
938 let new = soft_threshold(z, l1) / denom;
939 if new != old {
940 let diff = new - old;
941 for i in 0..n {
942 fit[i] += diff * xc[i + j * n];
943 }
944 max_delta = max_delta.max(diff.abs());
945 beta[j] = new;
946 }
947 }
948 if max_delta < tol {
949 break;
950 }
951 }
952
953 let intercept = mu_y - (0..p).map(|j| beta[j] * col_means[j]).sum::<f64>();
955 Ok((intercept, beta))
956}
957
958fn build_lambda_grid(design: &FdMatrix, y: &[f64], config: &WnetConfig) -> Vec<f64> {
969 if let Some(grid) = &config.lambda_grid {
970 return grid.clone();
971 }
972 let (n, p) = design.shape();
973 let n_f = n as f64;
974 let mu_y = y.iter().sum::<f64>() / n_f;
975 let y_centered: Vec<f64> = y.iter().map(|&v| v - mu_y).collect();
976
977 let alpha_eff = config.alpha.max(1e-3);
979 let mut max_corr = 0.0_f64;
980 for j in 0..p {
981 let mu = design.column(j).iter().sum::<f64>() / n_f;
982 let col = design.column(j);
983 let dot: f64 = (0..n).map(|i| (col[i] - mu) * y_centered[i]).sum();
984 max_corr = max_corr.max(dot.abs());
985 }
986 let lambda_max = (max_corr / (n_f * alpha_eff)).max(1e-8);
987
988 let n_lambda = config.n_lambda.max(1);
989 if n_lambda == 1 {
990 return vec![lambda_max];
991 }
992 let eps = 1e-3_f64;
993 let log_max = lambda_max.ln();
994 let log_min = (lambda_max * eps).ln();
995 let step = (log_max - log_min) / (n_lambda as f64 - 1.0);
996 (0..n_lambda)
997 .map(|k| (log_max - step * k as f64).exp())
998 .collect()
999}
1000
1001pub(crate) fn wnet_cv_lambda(
1016 design: &FdMatrix,
1017 y: &[f64],
1018 config: &WnetConfig,
1019) -> Result<f64, FdarError> {
1020 let (n, _p) = design.shape();
1021 if y.len() != n {
1022 return Err(FdarError::InvalidDimension {
1023 parameter: "y",
1024 expected: format!("{n} elements (== design rows)"),
1025 actual: format!("{} elements", y.len()),
1026 });
1027 }
1028 if config.n_folds < 2 {
1029 return Err(FdarError::InvalidParameter {
1030 parameter: "n_folds",
1031 message: format!("n_folds must be >= 2, got {}", config.n_folds),
1032 });
1033 }
1034 if config.n_folds > n {
1035 return Err(FdarError::InvalidParameter {
1036 parameter: "n_folds",
1037 message: format!(
1038 "n_folds ({}) must not exceed the number of observations ({n})",
1039 config.n_folds
1040 ),
1041 });
1042 }
1043 if !(0.0..=1.0).contains(&config.alpha) {
1044 return Err(FdarError::InvalidParameter {
1045 parameter: "alpha",
1046 message: format!("alpha must be in [0, 1], got {}", config.alpha),
1047 });
1048 }
1049 if let Some(grid) = &config.lambda_grid {
1050 if grid.is_empty() {
1051 return Err(FdarError::InvalidParameter {
1052 parameter: "lambda_grid",
1053 message: "explicit lambda_grid must be non-empty".to_string(),
1054 });
1055 }
1056 }
1057
1058 let grid = build_lambda_grid(design, y, config);
1059 let folds = crate::cv::create_folds(n, config.n_folds, config.seed);
1060
1061 let fold_sets: Vec<(Vec<usize>, Vec<usize>)> = (0..config.n_folds)
1063 .map(|f| crate::cv::fold_indices(&folds, f))
1064 .collect();
1065
1066 let mut best_lambda = grid[0];
1067 let mut best_mse = f64::INFINITY;
1068 let tie_eps = 1e-12;
1069
1070 for &lam in &grid {
1071 let mut total_sse = 0.0_f64;
1072 let mut scored = 0usize;
1073 for (train_idx, test_idx) in &fold_sets {
1074 if train_idx.is_empty() || test_idx.is_empty() {
1075 continue;
1076 }
1077 let train_data = crate::cv::subset_rows(design, train_idx);
1078 let train_y = crate::cv::subset_vec(y, train_idx);
1079 let (intercept, beta) = elastic_net_cd(
1080 &train_data,
1081 &train_y,
1082 lam,
1083 config.alpha,
1084 config.max_iter,
1085 config.tol,
1086 )?;
1087 for &oi in test_idx {
1088 let mut yhat = intercept;
1089 for j in 0..design.ncols() {
1090 yhat += design[(oi, j)] * beta[j];
1091 }
1092 let e = y[oi] - yhat;
1093 total_sse += e * e;
1094 scored += 1;
1095 }
1096 }
1097 if scored == 0 {
1098 continue;
1099 }
1100 let mse = total_sse / scored as f64;
1101 if mse < best_mse - tie_eps {
1105 best_mse = mse;
1106 best_lambda = lam;
1107 }
1108 }
1109
1110 Ok(best_lambda)
1111}
1112
1113#[must_use = "expensive computation whose result should not be discarded"]
1145pub fn wnet(data: &FdMatrix, y: &[f64], config: &WnetConfig) -> Result<WnetResult, FdarError> {
1146 let (n, m) = data.shape();
1147 if n < 3 {
1148 return Err(FdarError::InvalidDimension {
1149 parameter: "data",
1150 expected: "at least 3 rows (observations)".to_string(),
1151 actual: format!("{n} rows"),
1152 });
1153 }
1154 if m == 0 {
1155 return Err(FdarError::InvalidDimension {
1156 parameter: "data",
1157 expected: "at least 1 column (evaluation point)".to_string(),
1158 actual: format!("{m} columns"),
1159 });
1160 }
1161 if y.len() != n {
1162 return Err(FdarError::InvalidDimension {
1163 parameter: "y",
1164 expected: format!("{n} elements (== data rows)"),
1165 actual: format!("{} elements", y.len()),
1166 });
1167 }
1168 if !(0.0..=1.0).contains(&config.alpha) {
1169 return Err(FdarError::InvalidParameter {
1170 parameter: "alpha",
1171 message: format!("alpha must be in [0, 1], got {}", config.alpha),
1172 });
1173 }
1174 if config.n_folds < 2 {
1175 return Err(FdarError::InvalidParameter {
1176 parameter: "n_folds",
1177 message: format!("n_folds must be >= 2, got {}", config.n_folds),
1178 });
1179 }
1180 if config.n_folds > n {
1181 return Err(FdarError::InvalidParameter {
1182 parameter: "n_folds",
1183 message: format!(
1184 "n_folds ({}) must not exceed the number of observations ({n})",
1185 config.n_folds
1186 ),
1187 });
1188 }
1189 if config.max_iter == 0 {
1190 return Err(FdarError::InvalidParameter {
1191 parameter: "max_iter",
1192 message: "max_iter must be >= 1".to_string(),
1193 });
1194 }
1195 if !config.tol.is_finite() || config.tol < 0.0 {
1196 return Err(FdarError::InvalidParameter {
1197 parameter: "tol",
1198 message: format!("tol must be finite and >= 0, got {}", config.tol),
1199 });
1200 }
1201 if let Some(grid) = &config.lambda_grid {
1202 if grid.is_empty() {
1203 return Err(FdarError::InvalidParameter {
1204 parameter: "lambda_grid",
1205 message: "explicit lambda_grid must be non-empty".to_string(),
1206 });
1207 }
1208 }
1209
1210 let (design, layout) =
1212 curves_to_coeff_design(data, config.family.clone(), config.mode, config.level)?;
1213
1214 let lambda = wnet_cv_lambda(&design, y, config)?;
1216 let (intercept, coeff_weights) = elastic_net_cd(
1217 &design,
1218 y,
1219 lambda,
1220 config.alpha,
1221 config.max_iter,
1222 config.tol,
1223 )?;
1224
1225 let selected: Vec<usize> = coeff_weights
1227 .iter()
1228 .enumerate()
1229 .filter(|(_, &b)| b != 0.0)
1230 .map(|(j, _)| j)
1231 .collect();
1232
1233 let fitted_values = compute_fitted_affine(&design, &coeff_weights, intercept);
1235 let residuals: Vec<f64> = y
1236 .iter()
1237 .zip(&fitted_values)
1238 .map(|(&yi, &yh)| yi - yh)
1239 .collect();
1240
1241 let beta_t = coeff_weights_to_beta_t(&coeff_weights, &layout)?;
1243
1244 Ok(WnetResult {
1245 intercept,
1246 beta_t,
1247 fitted_values,
1248 residuals,
1249 coeff_weights,
1250 selected,
1251 lambda,
1252 alpha: config.alpha,
1253 family: config.family.clone(),
1254 mode: config.mode,
1255 level: layout.levels(),
1256 })
1257}
1258
1259impl WnetResult {
1260 pub fn predict(&self, new: &FdMatrix) -> Result<Vec<f64>, FdarError> {
1280 let train_m = self.beta_t.len();
1281 if new.nrows() == 0 {
1282 return Err(FdarError::InvalidDimension {
1283 parameter: "new",
1284 expected: "at least 1 row (curve)".to_string(),
1285 actual: "0 rows".to_string(),
1286 });
1287 }
1288 if new.ncols() != train_m {
1289 return Err(FdarError::InvalidDimension {
1290 parameter: "new",
1291 expected: format!("{train_m} columns (== training grid length)"),
1292 actual: format!("{} columns", new.ncols()),
1293 });
1294 }
1295 let (design, _layout) =
1298 curves_to_coeff_design(new, self.family.clone(), self.mode, Some(self.level))?;
1299 if design.ncols() != self.coeff_weights.len() {
1300 return Err(FdarError::InvalidDimension {
1301 parameter: "new",
1302 expected: format!(
1303 "coefficient-space width {} (== stored coeff_weights)",
1304 self.coeff_weights.len()
1305 ),
1306 actual: format!("{} coefficients", design.ncols()),
1307 });
1308 }
1309 Ok(compute_fitted_affine(
1310 &design,
1311 &self.coeff_weights,
1312 self.intercept,
1313 ))
1314 }
1315
1316 #[must_use]
1318 pub fn beta_t(&self) -> &[f64] {
1319 &self.beta_t
1320 }
1321
1322 #[must_use]
1324 pub fn coefficient_function(&self) -> &[f64] {
1325 &self.beta_t
1326 }
1327
1328 #[must_use]
1330 pub fn fitted_values(&self) -> &[f64] {
1331 &self.fitted_values
1332 }
1333}
1334
1335fn compute_fitted_affine(design: &FdMatrix, coeffs: &[f64], intercept: f64) -> Vec<f64> {
1337 let (n, p) = design.shape();
1338 (0..n)
1339 .map(|i| {
1340 let mut yhat = intercept;
1341 for j in 0..p {
1342 yhat += design[(i, j)] * coeffs[j];
1343 }
1344 yhat
1345 })
1346 .collect()
1347}
1348
1349#[cfg(test)]
1350mod tests {
1351 use super::*;
1352 use crate::matrix::FdMatrix;
1353
1354 fn pseudo_random(n: usize, seed: u64) -> Vec<f64> {
1357 let mut state = seed.wrapping_add(0x9E37_79B9_7F4A_7C15);
1358 (0..n)
1359 .map(|_| {
1360 state = state
1361 .wrapping_mul(6_364_136_223_846_793_005)
1362 .wrapping_add(1_442_695_040_888_963_407);
1363 let u = (state >> 11) as f64 / (1u64 << 53) as f64;
1364 2.0 * u - 1.0
1365 })
1366 .collect()
1367 }
1368
1369 fn spanning_design(n: usize, m: usize, seed0: u64) -> FdMatrix {
1372 let mut flat = vec![0.0_f64; n * m];
1373 for i in 0..n {
1374 let row = pseudo_random(m, seed0 + i as u64);
1375 for j in 0..m {
1376 flat[i + j * n] = row[j];
1377 }
1378 }
1379 FdMatrix::from_column_major(flat, n, m).unwrap()
1380 }
1381
1382 fn rel_l2(recovered: &[f64], truth: &[f64]) -> f64 {
1383 let num: f64 = recovered
1384 .iter()
1385 .zip(truth)
1386 .map(|(a, b)| (a - b) * (a - b))
1387 .sum::<f64>()
1388 .sqrt();
1389 let den: f64 = truth.iter().map(|b| b * b).sum::<f64>().sqrt().max(1e-300);
1390 num / den
1391 }
1392
1393 fn recovery_for_method(method: WcrMethod) {
1397 let (n, m) = (120usize, 32usize);
1398 let data = spanning_design(n, m, 1000);
1399 let family = WaveletFamily::Daubechies(4);
1400 let mode = BoundaryMode::Periodic;
1401
1402 let (design, layout) = curves_to_coeff_design(&data, family.clone(), mode, None).unwrap();
1404 let p = design.ncols();
1405
1406 let beta_coeff = pseudo_random(p, 77);
1409 let true_intercept = 0.37_f64;
1410 let y: Vec<f64> = (0..n)
1411 .map(|i| {
1412 let mut acc = true_intercept;
1413 for j in 0..p {
1414 acc += design[(i, j)] * beta_coeff[j];
1415 }
1416 acc
1417 })
1418 .collect();
1419
1420 let beta_t_true = coeff_weights_to_beta_t(&beta_coeff, &layout).unwrap();
1422
1423 let config = WcrConfig {
1425 family,
1426 mode,
1427 level: None,
1428 ncomp: p.min(n),
1429 method,
1430 ..Default::default()
1431 };
1432 let fit = wcr(&data, &y, &config).unwrap();
1433
1434 assert_eq!(fit.method, method);
1435 assert_eq!(fit.beta_t.len(), m);
1436 assert_eq!(fit.coeff_weights.len(), p);
1437
1438 let e = rel_l2(&fit.beta_t, &beta_t_true);
1439 assert!(
1440 e < 1e-6,
1441 "{method:?}: beta_t recovery rel L2 err {e} exceeds tolerance on spanning full-rank design"
1442 );
1443
1444 assert!(fit.beta_t.iter().all(|x| x.is_finite()));
1446 assert!(fit.fitted_values.iter().all(|x| x.is_finite()));
1447 assert!(fit.residuals.iter().all(|x| x.is_finite()));
1448 let max_resid = fit
1450 .residuals
1451 .iter()
1452 .fold(0.0_f64, |acc, &r| acc.max(r.abs()));
1453 assert!(
1454 max_resid < 1e-6,
1455 "{method:?}: residuals not ~0 ({max_resid})"
1456 );
1457 }
1458
1459 #[test]
1460 fn wcr_pcr_recovers_known_beta_t_on_spanning_design() {
1461 recovery_for_method(WcrMethod::Pcr);
1462 }
1463
1464 #[test]
1465 fn wcr_pls_recovers_known_beta_t_on_spanning_design() {
1466 recovery_for_method(WcrMethod::Pls);
1467 }
1468
1469 #[test]
1470 fn wcr_default_config_is_db4_periodic_auto_pcr() {
1471 let c = WcrConfig::default();
1472 assert_eq!(c.family, WaveletFamily::Daubechies(4));
1473 assert_eq!(c.mode, BoundaryMode::Periodic);
1474 assert_eq!(c.level, None);
1475 assert_eq!(c.method, WcrMethod::Pcr);
1476 assert_eq!(WcrMethod::default(), WcrMethod::Pcr);
1477 }
1478
1479 #[test]
1480 fn curves_to_coeff_design_layout_and_shape() {
1481 let (n, m) = (10usize, 48usize);
1482 let data = spanning_design(n, m, 500);
1483 let (design, layout) = curves_to_coeff_design(
1484 &data,
1485 WaveletFamily::Daubechies(4),
1486 BoundaryMode::Periodic,
1487 None,
1488 )
1489 .unwrap();
1490 assert_eq!(design.nrows(), n);
1491 assert_eq!(design.ncols(), layout.total_len());
1492 assert_eq!(layout.signal_len, m);
1493 assert_eq!(layout.levels(), layout.detail_lens.len());
1494 }
1495
1496 #[test]
1497 fn coeff_weights_to_beta_t_inverts_decompose() {
1498 let (n, m) = (4usize, 48usize);
1500 let data = spanning_design(n, m, 900);
1501 let (design, layout) = curves_to_coeff_design(
1502 &data,
1503 WaveletFamily::Daubechies(6),
1504 BoundaryMode::Periodic,
1505 None,
1506 )
1507 .unwrap();
1508 let row0: Vec<f64> = (0..design.ncols()).map(|j| design[(0, j)]).collect();
1509 let recon = coeff_weights_to_beta_t(&row0, &layout).unwrap();
1510 let orig = data.row(0);
1511 assert!(rel_l2(&recon, &orig) < 1e-10);
1512 }
1513
1514 #[test]
1515 fn coeff_weights_to_beta_t_rejects_wrong_length() {
1516 let (n, m) = (4usize, 48usize);
1517 let data = spanning_design(n, m, 901);
1518 let (_design, layout) =
1519 curves_to_coeff_design(&data, WaveletFamily::Haar, BoundaryMode::Periodic, None)
1520 .unwrap();
1521 let wrong = vec![0.0; layout.total_len() + 1];
1522 assert!(matches!(
1523 coeff_weights_to_beta_t(&wrong, &layout),
1524 Err(FdarError::InvalidDimension { .. })
1525 ));
1526 }
1527
1528 fn base_config() -> WcrConfig {
1531 WcrConfig {
1532 ncomp: 3,
1533 ..Default::default()
1534 }
1535 }
1536
1537 #[test]
1538 fn wcr_rejects_too_few_rows() {
1539 let data = spanning_design(2, 48, 1);
1540 let y = vec![0.0, 1.0];
1541 assert!(matches!(
1542 wcr(&data, &y, &base_config()),
1543 Err(FdarError::InvalidDimension { .. })
1544 ));
1545 }
1546
1547 #[test]
1548 fn wcr_rejects_mismatched_y_len() {
1549 let data = spanning_design(10, 48, 2);
1550 let y = vec![0.0; 9];
1551 assert!(matches!(
1552 wcr(&data, &y, &base_config()),
1553 Err(FdarError::InvalidDimension { .. })
1554 ));
1555 }
1556
1557 #[test]
1558 fn wcr_rejects_zero_ncomp() {
1559 let data = spanning_design(10, 48, 3);
1560 let y = vec![0.0; 10];
1561 let config = WcrConfig {
1562 ncomp: 0,
1563 ..Default::default()
1564 };
1565 assert!(matches!(
1566 wcr(&data, &y, &config),
1567 Err(FdarError::InvalidParameter { .. })
1568 ));
1569 }
1570
1571 #[test]
1572 fn wcr_surfaces_unsupported_family() {
1573 let data = spanning_design(10, 48, 4);
1574 let y = vec![0.0; 10];
1575 let config = WcrConfig {
1576 family: WaveletFamily::Daubechies(11),
1577 ..base_config()
1578 };
1579 assert!(matches!(
1580 wcr(&data, &y, &config),
1581 Err(FdarError::InvalidParameter { .. })
1582 ));
1583 }
1584
1585 #[test]
1586 fn wcr_surfaces_level_out_of_range() {
1587 let data = spanning_design(10, 48, 5);
1588 let y = vec![0.0; 10];
1589 let config = WcrConfig {
1590 level: Some(999),
1591 ..base_config()
1592 };
1593 assert!(matches!(
1594 wcr(&data, &y, &config),
1595 Err(FdarError::InvalidParameter { .. })
1596 ));
1597 }
1598
1599 #[test]
1600 fn wcr_finite_outputs_both_methods() {
1601 let (n, m) = (100usize, 40usize);
1602 let data = spanning_design(n, m, 4242);
1603 let y = pseudo_random(n, 8080);
1604 for method in [WcrMethod::Pcr, WcrMethod::Pls] {
1605 let config = WcrConfig {
1606 ncomp: 8,
1607 method,
1608 ..Default::default()
1609 };
1610 let fit = wcr(&data, &y, &config).unwrap();
1611 assert!(fit.intercept.is_finite());
1612 assert!(fit.beta_t.iter().all(|x| x.is_finite()));
1613 assert!(fit.fitted_values.iter().all(|x| x.is_finite()));
1614 assert!(fit.residuals.iter().all(|x| x.is_finite()));
1615 }
1616 }
1617
1618 #[test]
1619 fn wcr_small_n_default_config_succeeds() {
1620 let (n, m) = (4usize, 32usize);
1625 let data = spanning_design(n, m, 2468);
1626 let y = pseudo_random(n, 1357);
1627 let config = WcrConfig::default(); let fit = wcr(&data, &y, &config).unwrap();
1629 assert!(fit.ncomp < n, "ncomp {} exceeds n - 1", fit.ncomp);
1631 assert_eq!(fit.beta_t.len(), m);
1632 assert!(fit.intercept.is_finite());
1633 assert!(fit.beta_t.iter().all(|x| x.is_finite()));
1634 assert!(fit.fitted_values.iter().all(|x| x.is_finite()));
1635 assert!(fit.residuals.iter().all(|x| x.is_finite()));
1636
1637 let data3 = spanning_design(3, m, 2469);
1639 let y3 = pseudo_random(3, 1358);
1640 let fit3 = wcr(&data3, &y3, &WcrConfig::default()).unwrap();
1641 assert!(fit3.ncomp <= 2);
1642 assert!(fit3.beta_t.iter().all(|x| x.is_finite()));
1643 }
1644
1645 #[test]
1648 fn wcr_predict_reproduces_training_fitted() {
1649 let (n, m) = (120usize, 32usize);
1650 let data = spanning_design(n, m, 6100);
1651 let y = pseudo_random(n, 6101);
1652 let config = WcrConfig {
1653 ncomp: 8,
1654 ..Default::default()
1655 };
1656 let fit = wcr(&data, &y, &config).unwrap();
1657 let preds = fit.predict(&data).unwrap();
1658 assert_eq!(preds.len(), fit.fitted_values.len());
1659 for (i, (&p, &f)) in preds.iter().zip(&fit.fitted_values).enumerate() {
1660 assert!(
1661 (p - f).abs() <= 1e-8,
1662 "wcr predict[{i}] {p} != fitted {f} (|Δ| {})",
1663 (p - f).abs()
1664 );
1665 }
1666 }
1667
1668 #[test]
1669 fn wcr_predict_on_new_curves_is_finite_and_rejects_grid_mismatch() {
1670 let (n, m) = (100usize, 32usize);
1671 let data = spanning_design(n, m, 6200);
1672 let y = pseudo_random(n, 6201);
1673 let fit = wcr(
1674 &data,
1675 &y,
1676 &WcrConfig {
1677 ncomp: 6,
1678 ..Default::default()
1679 },
1680 )
1681 .unwrap();
1682
1683 let fresh = spanning_design(40, m, 6202);
1685 let preds = fit.predict(&fresh).unwrap();
1686 assert_eq!(preds.len(), 40);
1687 assert!(preds.iter().all(|x| x.is_finite()));
1688
1689 let wrong = spanning_design(10, m + 8, 6203);
1691 assert!(matches!(
1692 fit.predict(&wrong),
1693 Err(FdarError::InvalidDimension { .. })
1694 ));
1695 }
1696
1697 #[test]
1698 fn wcr_predict_on_zero_row_input_errors_naming_new_no_panic() {
1699 let (n, m) = (100usize, 32usize);
1700 let data = spanning_design(n, m, 6400);
1701 let y = pseudo_random(n, 6401);
1702 let fit = wcr(
1703 &data,
1704 &y,
1705 &WcrConfig {
1706 ncomp: 6,
1707 ..Default::default()
1708 },
1709 )
1710 .unwrap();
1711
1712 let empty = FdMatrix::zeros(0, m);
1715 match fit.predict(&empty) {
1716 Err(FdarError::InvalidDimension { parameter, .. }) => {
1717 assert_eq!(parameter, "new");
1718 }
1719 other => panic!("expected InvalidDimension naming \"new\", got {other:?}"),
1720 }
1721 }
1722
1723 fn sparse_wnet_problem(
1731 n: usize,
1732 m: usize,
1733 seed0: u64,
1734 ) -> (FdMatrix, FdMatrix, CoeffLayout, Vec<f64>, Vec<usize>) {
1735 let data = spanning_design(n, m, seed0);
1736 let family = WaveletFamily::Daubechies(4);
1737 let mode = BoundaryMode::Periodic;
1738 let (design, layout) = curves_to_coeff_design(&data, family, mode, None).unwrap();
1739 let p = design.ncols();
1740
1741 let support: Vec<usize> = vec![0, 2, p / 2, p - 3]
1743 .into_iter()
1744 .filter(|&j| j < p)
1745 .collect();
1746 let mut beta_coeff = vec![0.0_f64; p];
1747 let mags = [4.0, -3.5, 5.0, -4.5];
1750 for (k, &j) in support.iter().enumerate() {
1751 beta_coeff[j] = mags[k % mags.len()];
1752 }
1753 (data, design, layout, beta_coeff, support)
1754 }
1755
1756 #[test]
1757 fn wnet_elastic_net_cd_recovers_sparse_support() {
1758 let (n, m) = (256usize, 32usize);
1761 let (_data, design, _layout, beta_coeff, support) = sparse_wnet_problem(n, m, 3000);
1762 let p = design.ncols();
1763
1764 let intercept_true = 0.5_f64;
1766 let y: Vec<f64> = (0..n)
1767 .map(|i| {
1768 let mut acc = intercept_true;
1769 for j in 0..p {
1770 acc += design[(i, j)] * beta_coeff[j];
1771 }
1772 acc
1773 })
1774 .collect();
1775
1776 let (intercept, beta) = elastic_net_cd(&design, &y, 0.05, 0.9, 2000, 1e-8).unwrap();
1778
1779 assert!(intercept.is_finite());
1780 assert!(beta.iter().all(|b| b.is_finite()));
1781
1782 let selected: Vec<usize> = beta
1783 .iter()
1784 .enumerate()
1785 .filter(|(_, &b)| b.abs() > 1e-8)
1786 .map(|(j, _)| j)
1787 .collect();
1788
1789 for &j in &support {
1791 assert!(
1792 selected.contains(&j),
1793 "true-support coeff {j} not selected (selected={selected:?})"
1794 );
1795 }
1796 assert!(
1798 selected.len() < p / 2,
1799 "selection not sparse: |selected|={} of P={p}",
1800 selected.len()
1801 );
1802 }
1803
1804 #[test]
1805 fn wnet_fixed_lambda_end_to_end_finite() {
1806 let (n, m) = (200usize, 32usize);
1809 let (data, design, _layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 3100);
1810 let p = design.ncols();
1811 let y: Vec<f64> = (0..n)
1812 .map(|i| {
1813 let mut acc = 0.25;
1814 for j in 0..p {
1815 acc += design[(i, j)] * beta_coeff[j];
1816 }
1817 acc
1818 })
1819 .collect();
1820
1821 let config = WnetConfig {
1822 n_lambda: 15,
1823 n_folds: 4,
1824 ..Default::default()
1825 };
1826 let fit = wnet(&data, &y, &config).unwrap();
1827 assert_eq!(fit.beta_t.len(), m);
1828 assert_eq!(fit.coeff_weights.len(), p);
1829 assert!(fit.intercept.is_finite());
1830 assert!(fit.beta_t.iter().all(|x| x.is_finite()));
1831 assert!(fit.fitted_values.iter().all(|x| x.is_finite()));
1832 assert!(fit.residuals.iter().all(|x| x.is_finite()));
1833 assert!(fit.coeff_weights.iter().all(|x| x.is_finite()));
1834 for &j in &fit.selected {
1836 assert!(fit.coeff_weights[j] != 0.0);
1837 }
1838 }
1839
1840 #[test]
1841 fn wnet_default_config_is_db4_periodic_auto() {
1842 let c = WnetConfig::default();
1843 assert_eq!(c.family, WaveletFamily::Daubechies(4));
1844 assert_eq!(c.mode, BoundaryMode::Periodic);
1845 assert_eq!(c.level, None);
1846 assert!((c.alpha - 0.5).abs() < 1e-15);
1847 assert_eq!(c.lambda_grid, None);
1848 assert_eq!(c.n_lambda, 50);
1849 assert_eq!(c.n_folds, 5);
1850 assert_eq!(c.seed, 0);
1851 }
1852
1853 #[test]
1856 fn wnet_cv_lambda_is_deterministic_across_runs() {
1857 let (n, m) = (200usize, 32usize);
1858 let (data, design, _layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 3200);
1859 let p = design.ncols();
1860 let noise = pseudo_random(n, 9999);
1862 let y: Vec<f64> = (0..n)
1863 .map(|i| {
1864 let mut acc = 0.1;
1865 for j in 0..p {
1866 acc += design[(i, j)] * beta_coeff[j];
1867 }
1868 acc + 0.05 * noise[i]
1869 })
1870 .collect();
1871
1872 let config = WnetConfig {
1873 alpha: 0.8,
1874 n_lambda: 20,
1875 n_folds: 5,
1876 seed: 0,
1877 ..Default::default()
1878 };
1879 let fit1 = wnet(&data, &y, &config).unwrap();
1880 let fit2 = wnet(&data, &y, &config).unwrap();
1881 assert_eq!(
1882 fit1.lambda, fit2.lambda,
1883 "CV-selected lambda differs across runs: {} vs {}",
1884 fit1.lambda, fit2.lambda
1885 );
1886 let l1 = wnet_cv_lambda(&design, &y, &config).unwrap();
1888 let l2 = wnet_cv_lambda(&design, &y, &config).unwrap();
1889 assert_eq!(l1, l2);
1890 }
1891
1892 #[test]
1895 fn wnet_recovers_beta_t_on_snr_data() {
1896 let (n, m) = (300usize, 32usize);
1899 let (data, design, layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 3300);
1900 let p = design.ncols();
1901 let beta_t_true = coeff_weights_to_beta_t(&beta_coeff, &layout).unwrap();
1902
1903 let signal: Vec<f64> = (0..n)
1905 .map(|i| {
1906 let mut acc = 0.0;
1907 for j in 0..p {
1908 acc += design[(i, j)] * beta_coeff[j];
1909 }
1910 acc
1911 })
1912 .collect();
1913 let sig_sd = {
1914 let mean = signal.iter().sum::<f64>() / n as f64;
1915 (signal.iter().map(|s| (s - mean).powi(2)).sum::<f64>() / n as f64).sqrt()
1916 };
1917 let noise = pseudo_random(n, 4141);
1918 let noise_scale = 0.05 * sig_sd; let y: Vec<f64> = (0..n)
1920 .map(|i| 0.3 + signal[i] + noise_scale * noise[i])
1921 .collect();
1922
1923 let config = WnetConfig {
1924 alpha: 0.7,
1925 n_lambda: 30,
1926 n_folds: 5,
1927 ..Default::default()
1928 };
1929 let fit = wnet(&data, &y, &config).unwrap();
1930
1931 let nonzero = fit.coeff_weights.iter().filter(|&&b| b != 0.0).count();
1933 assert!(nonzero > 0, "degenerate all-zero fit at CV lambda");
1934
1935 let e = rel_l2(&fit.beta_t, &beta_t_true);
1937 assert!(
1938 e < 0.35,
1939 "wnet beta_t recovery rel L2 err {e} exceeds tolerance on SNR data"
1940 );
1941 assert!(fit.beta_t.iter().all(|x| x.is_finite()));
1942 assert!(fit.fitted_values.iter().all(|x| x.is_finite()));
1943 }
1944
1945 fn base_wnet_config() -> WnetConfig {
1948 WnetConfig {
1949 n_lambda: 10,
1950 n_folds: 3,
1951 ..Default::default()
1952 }
1953 }
1954
1955 #[test]
1956 fn wnet_rejects_too_few_rows() {
1957 let data = spanning_design(2, 32, 10);
1958 let y = vec![0.0, 1.0];
1959 assert!(matches!(
1960 wnet(&data, &y, &base_wnet_config()),
1961 Err(FdarError::InvalidDimension { .. })
1962 ));
1963 }
1964
1965 #[test]
1966 fn wnet_rejects_zero_cols() {
1967 let data = FdMatrix::zeros(5, 0);
1969 let y = vec![0.0; 5];
1970 assert!(matches!(
1971 wnet(&data, &y, &base_wnet_config()),
1972 Err(FdarError::InvalidDimension { .. })
1973 ));
1974 }
1975
1976 #[test]
1977 fn wnet_rejects_mismatched_y_len() {
1978 let data = spanning_design(10, 32, 11);
1979 let y = vec![0.0; 9];
1980 assert!(matches!(
1981 wnet(&data, &y, &base_wnet_config()),
1982 Err(FdarError::InvalidDimension { .. })
1983 ));
1984 }
1985
1986 #[test]
1987 fn wnet_rejects_alpha_out_of_range() {
1988 let data = spanning_design(10, 32, 12);
1989 let y = vec![0.0; 10];
1990 let config = WnetConfig {
1991 alpha: 1.5,
1992 ..base_wnet_config()
1993 };
1994 assert!(matches!(
1995 wnet(&data, &y, &config),
1996 Err(FdarError::InvalidParameter { .. })
1997 ));
1998 let config = WnetConfig {
1999 alpha: -0.1,
2000 ..base_wnet_config()
2001 };
2002 assert!(matches!(
2003 wnet(&data, &y, &config),
2004 Err(FdarError::InvalidParameter { .. })
2005 ));
2006 }
2007
2008 #[test]
2009 fn wnet_rejects_too_few_folds() {
2010 let data = spanning_design(10, 32, 13);
2011 let y = vec![0.0; 10];
2012 let config = WnetConfig {
2013 n_folds: 1,
2014 ..base_wnet_config()
2015 };
2016 assert!(matches!(
2017 wnet(&data, &y, &config),
2018 Err(FdarError::InvalidParameter { .. })
2019 ));
2020 }
2021
2022 #[test]
2023 fn wnet_rejects_too_many_folds() {
2024 let data = spanning_design(10, 32, 130);
2026 let y = vec![0.0; 10];
2027 let config = WnetConfig {
2028 n_folds: 11,
2029 ..base_wnet_config()
2030 };
2031 assert!(matches!(
2032 wnet(&data, &y, &config),
2033 Err(FdarError::InvalidParameter { .. })
2034 ));
2035 let (design, _layout) = curves_to_coeff_design(
2037 &data,
2038 WaveletFamily::Daubechies(4),
2039 BoundaryMode::Periodic,
2040 None,
2041 )
2042 .unwrap();
2043 assert!(matches!(
2044 wnet_cv_lambda(&design, &y, &config),
2045 Err(FdarError::InvalidParameter { .. })
2046 ));
2047 }
2048
2049 #[test]
2050 fn wnet_rejects_negative_or_nan_tol() {
2051 let data = spanning_design(10, 32, 131);
2053 let y = pseudo_random(10, 5);
2054 for bad in [-1e-6_f64, f64::NAN] {
2055 let config = WnetConfig {
2056 tol: bad,
2057 ..base_wnet_config()
2058 };
2059 assert!(matches!(
2060 wnet(&data, &y, &config),
2061 Err(FdarError::InvalidParameter { .. })
2062 ));
2063 }
2064 let (design, _layout) = curves_to_coeff_design(
2066 &data,
2067 WaveletFamily::Daubechies(4),
2068 BoundaryMode::Periodic,
2069 None,
2070 )
2071 .unwrap();
2072 assert!(matches!(
2073 elastic_net_cd(&design, &y, 0.1, 0.5, 100, -1.0),
2074 Err(FdarError::InvalidParameter { .. })
2075 ));
2076 assert!(matches!(
2077 elastic_net_cd(&design, &y, 0.1, 0.5, 100, f64::NAN),
2078 Err(FdarError::InvalidParameter { .. })
2079 ));
2080 }
2081
2082 #[test]
2083 fn wnet_rejects_zero_max_iter() {
2084 let data = spanning_design(10, 32, 132);
2086 let y = pseudo_random(10, 6);
2087 let config = WnetConfig {
2088 max_iter: 0,
2089 ..base_wnet_config()
2090 };
2091 assert!(matches!(
2092 wnet(&data, &y, &config),
2093 Err(FdarError::InvalidParameter { .. })
2094 ));
2095 let (design, _layout) = curves_to_coeff_design(
2097 &data,
2098 WaveletFamily::Daubechies(4),
2099 BoundaryMode::Periodic,
2100 None,
2101 )
2102 .unwrap();
2103 assert!(matches!(
2104 elastic_net_cd(&design, &y, 0.1, 0.5, 0, 1e-6),
2105 Err(FdarError::InvalidParameter { .. })
2106 ));
2107 }
2108
2109 #[test]
2110 fn wnet_rejects_empty_lambda_grid() {
2111 let data = spanning_design(10, 32, 14);
2112 let y = vec![0.0; 10];
2113 let config = WnetConfig {
2114 lambda_grid: Some(vec![]),
2115 ..base_wnet_config()
2116 };
2117 assert!(matches!(
2118 wnet(&data, &y, &config),
2119 Err(FdarError::InvalidParameter { .. })
2120 ));
2121 }
2122
2123 #[test]
2124 fn wnet_surfaces_unsupported_family() {
2125 let data = spanning_design(10, 32, 15);
2126 let y = vec![0.0; 10];
2127 let config = WnetConfig {
2128 family: WaveletFamily::Daubechies(11),
2129 ..base_wnet_config()
2130 };
2131 assert!(matches!(
2132 wnet(&data, &y, &config),
2133 Err(FdarError::InvalidParameter { .. })
2134 ));
2135 }
2136
2137 #[test]
2138 fn wnet_surfaces_level_out_of_range() {
2139 let data = spanning_design(10, 32, 16);
2140 let y = vec![0.0; 10];
2141 let config = WnetConfig {
2142 level: Some(999),
2143 ..base_wnet_config()
2144 };
2145 assert!(matches!(
2146 wnet(&data, &y, &config),
2147 Err(FdarError::InvalidParameter { .. })
2148 ));
2149 }
2150
2151 #[test]
2152 fn wnet_finite_outputs_on_larger_snr_design() {
2153 let (n, m) = (256usize, 48usize);
2154 let (data, design, _layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 3400);
2155 let p = design.ncols();
2156 let noise = pseudo_random(n, 2727);
2157 let y: Vec<f64> = (0..n)
2158 .map(|i| {
2159 let mut acc = 0.2;
2160 for j in 0..p {
2161 acc += design[(i, j)] * beta_coeff[j];
2162 }
2163 acc + 0.1 * noise[i]
2164 })
2165 .collect();
2166
2167 let config = WnetConfig {
2168 alpha: 0.6,
2169 n_lambda: 25,
2170 n_folds: 5,
2171 ..Default::default()
2172 };
2173 let fit = wnet(&data, &y, &config).unwrap();
2174 assert!(fit.intercept.is_finite());
2175 assert!(fit.lambda.is_finite());
2176 assert!(fit.beta_t.iter().all(|x| x.is_finite()));
2177 assert!(fit.fitted_values.iter().all(|x| x.is_finite()));
2178 assert!(fit.residuals.iter().all(|x| x.is_finite()));
2179 assert!(fit.coeff_weights.iter().all(|x| x.is_finite()));
2180 }
2181
2182 #[test]
2183 fn wnet_explicit_lambda_grid_is_used() {
2184 let (n, m) = (120usize, 32usize);
2186 let (data, design, _layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 3500);
2187 let p = design.ncols();
2188 let y: Vec<f64> = (0..n)
2189 .map(|i| {
2190 let mut acc = 0.0;
2191 for j in 0..p {
2192 acc += design[(i, j)] * beta_coeff[j];
2193 }
2194 acc
2195 })
2196 .collect();
2197 let config = WnetConfig {
2198 lambda_grid: Some(vec![0.123]),
2199 ..base_wnet_config()
2200 };
2201 let fit = wnet(&data, &y, &config).unwrap();
2202 assert!((fit.lambda - 0.123).abs() < 1e-15);
2203 }
2204
2205 #[test]
2208 fn wnet_predict_reproduces_training_fitted() {
2209 let (n, m) = (200usize, 32usize);
2210 let (data, design, _layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 6300);
2211 let p = design.ncols();
2212 let noise = pseudo_random(n, 6301);
2213 let y: Vec<f64> = (0..n)
2214 .map(|i| {
2215 let mut acc = 0.4;
2216 for j in 0..p {
2217 acc += design[(i, j)] * beta_coeff[j];
2218 }
2219 acc + 0.05 * noise[i]
2220 })
2221 .collect();
2222 let config = WnetConfig {
2223 alpha: 0.7,
2224 n_lambda: 20,
2225 n_folds: 5,
2226 ..Default::default()
2227 };
2228 let fit = wnet(&data, &y, &config).unwrap();
2229 let preds = fit.predict(&data).unwrap();
2230 assert_eq!(preds.len(), fit.fitted_values.len());
2231 for (i, (&pv, &f)) in preds.iter().zip(&fit.fitted_values).enumerate() {
2232 assert!(
2233 (pv - f).abs() <= 1e-8,
2234 "wnet predict[{i}] {pv} != fitted {f} (|Δ| {})",
2235 (pv - f).abs()
2236 );
2237 }
2238 }
2239
2240 #[test]
2241 fn wnet_predict_on_new_curves_is_finite_and_rejects_grid_mismatch() {
2242 let (n, m) = (150usize, 32usize);
2243 let (data, design, _layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 6400);
2244 let p = design.ncols();
2245 let y: Vec<f64> = (0..n)
2246 .map(|i| {
2247 let mut acc = 0.2;
2248 for j in 0..p {
2249 acc += design[(i, j)] * beta_coeff[j];
2250 }
2251 acc
2252 })
2253 .collect();
2254 let fit = wnet(
2255 &data,
2256 &y,
2257 &WnetConfig {
2258 n_lambda: 12,
2259 n_folds: 4,
2260 ..Default::default()
2261 },
2262 )
2263 .unwrap();
2264
2265 let fresh = spanning_design(30, m, 6401);
2267 let preds = fit.predict(&fresh).unwrap();
2268 assert_eq!(preds.len(), 30);
2269 assert!(preds.iter().all(|x| x.is_finite()));
2270
2271 let wrong = spanning_design(10, m + 16, 6402);
2273 assert!(matches!(
2274 fit.predict(&wrong),
2275 Err(FdarError::InvalidDimension { .. })
2276 ));
2277
2278 let empty = FdMatrix::zeros(0, m);
2280 match fit.predict(&empty) {
2281 Err(FdarError::InvalidDimension { parameter, .. }) => {
2282 assert_eq!(parameter, "new");
2283 }
2284 other => panic!("expected InvalidDimension naming \"new\", got {other:?}"),
2285 }
2286 }
2287
2288 #[test]
2289 fn accessors_return_stored_slices() {
2290 let (n, m) = (100usize, 32usize);
2291 let data = spanning_design(n, m, 6500);
2292 let y = pseudo_random(n, 6501);
2293
2294 let wcr_fit = wcr(
2295 &data,
2296 &y,
2297 &WcrConfig {
2298 ncomp: 5,
2299 ..Default::default()
2300 },
2301 )
2302 .unwrap();
2303 assert_eq!(wcr_fit.beta_t(), wcr_fit.beta_t.as_slice());
2304 assert_eq!(wcr_fit.coefficient_function(), wcr_fit.beta_t.as_slice());
2305 assert_eq!(wcr_fit.beta_t().len(), m);
2306 assert_eq!(wcr_fit.fitted_values(), wcr_fit.fitted_values.as_slice());
2307 assert_eq!(wcr_fit.fitted_values().len(), n);
2308
2309 let wnet_fit = wnet(
2310 &data,
2311 &y,
2312 &WnetConfig {
2313 n_lambda: 10,
2314 n_folds: 4,
2315 ..Default::default()
2316 },
2317 )
2318 .unwrap();
2319 assert_eq!(wnet_fit.beta_t(), wnet_fit.beta_t.as_slice());
2320 assert_eq!(wnet_fit.coefficient_function(), wnet_fit.beta_t.as_slice());
2321 assert_eq!(wnet_fit.beta_t().len(), m);
2322 assert_eq!(wnet_fit.fitted_values(), wnet_fit.fitted_values.as_slice());
2323 assert_eq!(wnet_fit.fitted_values().len(), n);
2324 }
2325}