use crate::error::FdarError;
use crate::matrix::FdMatrix;
use crate::regression::{fdata_to_pc_1d, fdata_to_pls_1d};
use crate::wavelet::{decompose_matrix, reconstruct, BoundaryMode, WaveletCoeffs, WaveletFamily};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub enum WcrMethod {
#[default]
Pcr,
Pls,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct WcrConfig {
pub family: WaveletFamily,
pub mode: BoundaryMode,
pub level: Option<usize>,
pub ncomp: usize,
pub method: WcrMethod,
}
impl Default for WcrConfig {
fn default() -> Self {
Self {
family: WaveletFamily::Daubechies(4),
mode: BoundaryMode::Periodic,
level: None,
ncomp: 5,
method: WcrMethod::Pcr,
}
}
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct WcrResult {
pub intercept: f64,
pub beta_t: Vec<f64>,
pub fitted_values: Vec<f64>,
pub residuals: Vec<f64>,
pub ncomp: usize,
pub method: WcrMethod,
pub coeff_weights: Vec<f64>,
pub family: WaveletFamily,
pub mode: BoundaryMode,
pub level: usize,
}
#[derive(Debug, Clone, PartialEq)]
pub(crate) struct CoeffLayout {
pub(crate) approx_len: usize,
pub(crate) detail_lens: Vec<usize>,
pub(crate) signal_len: usize,
pub(crate) family: WaveletFamily,
pub(crate) mode: BoundaryMode,
pub(crate) level_lens: Vec<usize>,
}
impl CoeffLayout {
pub(crate) fn total_len(&self) -> usize {
self.approx_len + self.detail_lens.iter().sum::<usize>()
}
pub(crate) fn levels(&self) -> usize {
self.detail_lens.len()
}
}
fn coeffs_to_row(coeffs: &WaveletCoeffs) -> Vec<f64> {
let mut row = Vec::with_capacity(
coeffs.approx.len() + coeffs.details.iter().map(Vec::len).sum::<usize>(),
);
row.extend_from_slice(&coeffs.approx);
for band in &coeffs.details {
row.extend_from_slice(band);
}
row
}
pub(crate) fn curves_to_coeff_design(
data: &FdMatrix,
family: WaveletFamily,
mode: BoundaryMode,
level: Option<usize>,
) -> Result<(FdMatrix, CoeffLayout), FdarError> {
let per_curve = decompose_matrix(data, family.clone(), mode, level)?;
let first = &per_curve[0];
let layout = CoeffLayout {
approx_len: first.approx.len(),
detail_lens: first.details.iter().map(Vec::len).collect(),
signal_len: first.signal_len,
family,
mode,
level_lens: first.level_lens.clone(),
};
let p = layout.total_len();
let n = per_curve.len();
let mut flat = vec![0.0_f64; n * p];
for (i, coeffs) in per_curve.iter().enumerate() {
if coeffs.approx.len() != layout.approx_len
|| coeffs.details.len() != layout.detail_lens.len()
|| coeffs
.details
.iter()
.zip(&layout.detail_lens)
.any(|(band, &len)| band.len() != len)
|| coeffs.signal_len != layout.signal_len
{
return Err(FdarError::InvalidDimension {
parameter: "data",
expected: format!("all curves share curve-0 coefficient layout (P = {p})"),
actual: format!("curve {i} produced a different band structure"),
});
}
let row = coeffs_to_row(coeffs);
for (j, &v) in row.iter().enumerate() {
flat[i + j * n] = v;
}
}
let design = FdMatrix::from_column_major(flat, n, p)?;
Ok((design, layout))
}
pub(crate) fn coeff_weights_to_beta_t(
weights: &[f64],
layout: &CoeffLayout,
) -> Result<Vec<f64>, FdarError> {
let expected = layout.total_len();
if weights.len() != expected {
return Err(FdarError::InvalidDimension {
parameter: "weights",
expected: format!("{expected} coefficients (approx + all detail bands)"),
actual: format!("{} coefficients", weights.len()),
});
}
let approx = weights[..layout.approx_len].to_vec();
let mut details: Vec<Vec<f64>> = Vec::with_capacity(layout.detail_lens.len());
let mut offset = layout.approx_len;
for &len in &layout.detail_lens {
details.push(weights[offset..offset + len].to_vec());
offset += len;
}
let coeffs = WaveletCoeffs {
approx,
details,
levels: layout.levels(),
signal_len: layout.signal_len,
family: layout.family.clone(),
mode: layout.mode,
level_lens: layout.level_lens.clone(),
};
reconstruct(&coeffs)
}
fn design_with_intercept(scores: &FdMatrix, ncomp: usize) -> FdMatrix {
let n = scores.nrows();
let mut design = FdMatrix::zeros(n, 1 + ncomp);
for i in 0..n {
design[(i, 0)] = 1.0;
for k in 0..ncomp {
design[(i, 1 + k)] = scores[(i, k)];
}
}
design
}
fn ols_solve(x: &FdMatrix, y: &[f64]) -> Result<Vec<f64>, FdarError> {
let (n, p) = x.shape();
if n < p || p == 0 {
return Err(FdarError::InvalidDimension {
parameter: "design matrix",
expected: format!("n >= p and p > 0 (p={p})"),
actual: format!("n={n}, p={p}"),
});
}
let mut xtx = vec![0.0_f64; p * p];
let mut xty = vec![0.0_f64; p];
for a in 0..p {
for b in 0..p {
let mut s = 0.0;
for i in 0..n {
s += x[(i, a)] * x[(i, b)];
}
xtx[a + b * p] = s;
}
let mut sy = 0.0;
for i in 0..n {
sy += x[(i, a)] * y[i];
}
xty[a] = sy;
}
let l = cholesky_factor(&xtx, p)?;
Ok(cholesky_solve(&l, &xty, p))
}
fn cholesky_factor(a: &[f64], p: usize) -> Result<Vec<f64>, FdarError> {
let mut l = vec![0.0_f64; p * p];
for j in 0..p {
let mut diag = a[j + j * p];
for k in 0..j {
diag -= l[j + k * p] * l[j + k * p];
}
if diag <= 0.0 {
return Err(FdarError::ComputationFailed {
operation: "Cholesky factorization (wcr OLS)",
detail: "design matrix X'X is not positive definite; try reducing ncomp"
.to_string(),
});
}
let ljj = diag.sqrt();
l[j + j * p] = ljj;
for i in (j + 1)..p {
let mut s = a[i + j * p];
for k in 0..j {
s -= l[i + k * p] * l[j + k * p];
}
l[i + j * p] = s / ljj;
}
}
Ok(l)
}
fn cholesky_solve(l: &[f64], rhs: &[f64], p: usize) -> Vec<f64> {
let mut z = vec![0.0_f64; p];
for i in 0..p {
let mut s = rhs[i];
for k in 0..i {
s -= l[i + k * p] * z[k];
}
z[i] = s / l[i + i * p];
}
let mut b = vec![0.0_f64; p];
for i in (0..p).rev() {
let mut s = z[i];
for k in (i + 1)..p {
s -= l[k + i * p] * b[k];
}
b[i] = s / l[i + i * p];
}
b
}
fn recover_coeff_weights(
design: &FdMatrix,
fitted: &[f64],
intercept: f64,
) -> Result<Vec<f64>, FdarError> {
let (n, p) = design.shape();
let col_means: Vec<f64> = (0..p)
.map(|j| design.column(j).iter().sum::<f64>() / n as f64)
.collect();
let mut xc = FdMatrix::zeros(n, p);
for j in 0..p {
for i in 0..n {
xc[(i, j)] = design[(i, j)] - col_means[j];
}
}
let r: Vec<f64> = fitted.iter().map(|&f| f - intercept).collect();
let mut xtx = vec![0.0_f64; p * p];
let mut xtr = vec![0.0_f64; p];
for a in 0..p {
for b in 0..p {
let mut s = 0.0;
for i in 0..n {
s += xc[(i, a)] * xc[(i, b)];
}
xtx[a + b * p] = s;
}
let mut sr = 0.0;
for i in 0..n {
sr += xc[(i, a)] * r[i];
}
xtr[a] = sr;
}
let trace: f64 = (0..p).map(|j| xtx[j + j * p]).sum();
let eps = 1e-10 * (trace / p as f64).max(1e-12);
for j in 0..p {
xtx[j + j * p] += eps;
}
let l = cholesky_factor(&xtx, p)?;
Ok(cholesky_solve(&l, &xtr, p))
}
fn compute_fitted(design: &FdMatrix, coeffs: &[f64]) -> Vec<f64> {
let (n, p) = design.shape();
(0..n)
.map(|i| {
let mut yhat = 0.0;
for j in 0..p {
yhat += design[(i, j)] * coeffs[j];
}
yhat
})
.collect()
}
#[must_use = "expensive computation whose result should not be discarded"]
pub fn wcr(data: &FdMatrix, y: &[f64], config: &WcrConfig) -> Result<WcrResult, FdarError> {
let (n, m) = data.shape();
if n < 3 {
return Err(FdarError::InvalidDimension {
parameter: "data",
expected: "at least 3 rows (observations)".to_string(),
actual: format!("{n} rows"),
});
}
if m == 0 {
return Err(FdarError::InvalidDimension {
parameter: "data",
expected: "at least 1 column (evaluation point)".to_string(),
actual: format!("{m} columns"),
});
}
if y.len() != n {
return Err(FdarError::InvalidDimension {
parameter: "y",
expected: format!("{n} elements (== data rows)"),
actual: format!("{} elements", y.len()),
});
}
if config.ncomp == 0 {
return Err(FdarError::InvalidParameter {
parameter: "ncomp",
message: "ncomp must be >= 1".to_string(),
});
}
let (design, layout) =
curves_to_coeff_design(data, config.family.clone(), config.mode, config.level)?;
let p = design.ncols();
let argvals: Vec<f64> = (0..p).map(|j| j as f64).collect();
let ncomp = config.ncomp.min(n.saturating_sub(1)).min(p);
let (scores, ncomp) = match config.method {
WcrMethod::Pcr => {
let fpca = fdata_to_pc_1d(&design, ncomp, &argvals)?;
let k = fpca.scores.ncols();
(fpca.scores, k)
}
WcrMethod::Pls => {
let pls = fdata_to_pls_1d(&design, y, ncomp, &argvals)?;
let k = pls.scores.ncols();
(pls.scores, k)
}
};
let ols_design = design_with_intercept(&scores, ncomp);
let coeffs = ols_solve(&ols_design, y)?;
let intercept = coeffs[0];
let fitted_values = compute_fitted(&ols_design, &coeffs);
let coeff_weights = recover_coeff_weights(&design, &fitted_values, intercept)?;
let (n_rows, p_cols) = design.shape();
let intercept = {
let offset: f64 = (0..p_cols)
.map(|j| {
let col_mean = design.column(j).iter().sum::<f64>() / n_rows as f64;
col_mean * coeff_weights[j]
})
.sum();
intercept - offset
};
let beta_t = coeff_weights_to_beta_t(&coeff_weights, &layout)?;
let residuals: Vec<f64> = y
.iter()
.zip(&fitted_values)
.map(|(&yi, &yh)| yi - yh)
.collect();
Ok(WcrResult {
intercept,
beta_t,
fitted_values,
residuals,
ncomp,
method: config.method,
coeff_weights,
family: config.family.clone(),
mode: config.mode,
level: layout.levels(),
})
}
impl WcrResult {
pub fn predict(&self, new: &FdMatrix) -> Result<Vec<f64>, FdarError> {
let train_m = self.beta_t.len();
if new.nrows() == 0 {
return Err(FdarError::InvalidDimension {
parameter: "new",
expected: "at least 1 row (curve)".to_string(),
actual: "0 rows".to_string(),
});
}
if new.ncols() != train_m {
return Err(FdarError::InvalidDimension {
parameter: "new",
expected: format!("{train_m} columns (== training grid length)"),
actual: format!("{} columns", new.ncols()),
});
}
let (design, _layout) =
curves_to_coeff_design(new, self.family.clone(), self.mode, Some(self.level))?;
if design.ncols() != self.coeff_weights.len() {
return Err(FdarError::InvalidDimension {
parameter: "new",
expected: format!(
"coefficient-space width {} (== stored coeff_weights)",
self.coeff_weights.len()
),
actual: format!("{} coefficients", design.ncols()),
});
}
Ok(compute_fitted_affine(
&design,
&self.coeff_weights,
self.intercept,
))
}
#[must_use]
pub fn beta_t(&self) -> &[f64] {
&self.beta_t
}
#[must_use]
pub fn coefficient_function(&self) -> &[f64] {
&self.beta_t
}
#[must_use]
pub fn fitted_values(&self) -> &[f64] {
&self.fitted_values
}
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct WnetConfig {
pub family: WaveletFamily,
pub mode: BoundaryMode,
pub level: Option<usize>,
pub alpha: f64,
pub lambda_grid: Option<Vec<f64>>,
pub n_lambda: usize,
pub n_folds: usize,
pub seed: u64,
pub max_iter: usize,
pub tol: f64,
}
impl Default for WnetConfig {
fn default() -> Self {
Self {
family: WaveletFamily::Daubechies(4),
mode: BoundaryMode::Periodic,
level: None,
alpha: 0.5,
lambda_grid: None,
n_lambda: 50,
n_folds: 5,
seed: 0,
max_iter: 1000,
tol: 1e-6,
}
}
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct WnetResult {
pub intercept: f64,
pub beta_t: Vec<f64>,
pub fitted_values: Vec<f64>,
pub residuals: Vec<f64>,
pub coeff_weights: Vec<f64>,
pub selected: Vec<usize>,
pub lambda: f64,
pub alpha: f64,
pub family: WaveletFamily,
pub mode: BoundaryMode,
pub level: usize,
}
#[inline]
fn soft_threshold(z: f64, gamma: f64) -> f64 {
if z > gamma {
z - gamma
} else if z < -gamma {
z + gamma
} else {
0.0
}
}
pub(crate) fn elastic_net_cd(
design: &FdMatrix,
y: &[f64],
lambda: f64,
alpha: f64,
max_iter: usize,
tol: f64,
) -> Result<(f64, Vec<f64>), FdarError> {
let (n, p) = design.shape();
if y.len() != n {
return Err(FdarError::InvalidDimension {
parameter: "y",
expected: format!("{n} elements (== design rows)"),
actual: format!("{} elements", y.len()),
});
}
if !(0.0..=1.0).contains(&alpha) {
return Err(FdarError::InvalidParameter {
parameter: "alpha",
message: format!("alpha must be in [0, 1], got {alpha}"),
});
}
if lambda < 0.0 || !lambda.is_finite() {
return Err(FdarError::InvalidParameter {
parameter: "lambda",
message: format!("lambda must be finite and >= 0, got {lambda}"),
});
}
if !tol.is_finite() || tol < 0.0 {
return Err(FdarError::InvalidParameter {
parameter: "tol",
message: format!("tol must be finite and >= 0, got {tol}"),
});
}
if max_iter == 0 {
return Err(FdarError::InvalidParameter {
parameter: "max_iter",
message: "max_iter must be >= 1".to_string(),
});
}
let n_f = n as f64;
let mu_y = y.iter().sum::<f64>() / n_f;
let y_centered: Vec<f64> = y.iter().map(|&v| v - mu_y).collect();
let col_means: Vec<f64> = (0..p)
.map(|j| design.column(j).iter().sum::<f64>() / n_f)
.collect();
let mut xc = vec![0.0_f64; n * p]; let mut col_norm_sq_over_n = vec![0.0_f64; p];
for j in 0..p {
let mu = col_means[j];
let mut norm_sq = 0.0;
let col = design.column(j);
for i in 0..n {
let v = col[i] - mu;
xc[i + j * n] = v;
norm_sq += v * v;
}
col_norm_sq_over_n[j] = norm_sq / n_f;
}
let mut beta = vec![0.0_f64; p];
let mut fit = vec![0.0_f64; n]; let l1 = lambda * alpha;
let l2 = lambda * (1.0 - alpha);
for _sweep in 0..max_iter {
let mut max_delta = 0.0_f64;
for j in 0..p {
let denom = col_norm_sq_over_n[j] + l2;
if denom <= 0.0 {
if beta[j] != 0.0 {
let old = beta[j];
for i in 0..n {
fit[i] -= old * xc[i + j * n];
}
max_delta = max_delta.max(old.abs());
beta[j] = 0.0;
}
continue;
}
let old = beta[j];
let mut dot = 0.0;
for i in 0..n {
let r = y_centered[i] - fit[i] + old * xc[i + j * n];
dot += xc[i + j * n] * r;
}
let z = dot / n_f;
let new = soft_threshold(z, l1) / denom;
if new != old {
let diff = new - old;
for i in 0..n {
fit[i] += diff * xc[i + j * n];
}
max_delta = max_delta.max(diff.abs());
beta[j] = new;
}
}
if max_delta < tol {
break;
}
}
let intercept = mu_y - (0..p).map(|j| beta[j] * col_means[j]).sum::<f64>();
Ok((intercept, beta))
}
fn build_lambda_grid(design: &FdMatrix, y: &[f64], config: &WnetConfig) -> Vec<f64> {
if let Some(grid) = &config.lambda_grid {
return grid.clone();
}
let (n, p) = design.shape();
let n_f = n as f64;
let mu_y = y.iter().sum::<f64>() / n_f;
let y_centered: Vec<f64> = y.iter().map(|&v| v - mu_y).collect();
let alpha_eff = config.alpha.max(1e-3);
let mut max_corr = 0.0_f64;
for j in 0..p {
let mu = design.column(j).iter().sum::<f64>() / n_f;
let col = design.column(j);
let dot: f64 = (0..n).map(|i| (col[i] - mu) * y_centered[i]).sum();
max_corr = max_corr.max(dot.abs());
}
let lambda_max = (max_corr / (n_f * alpha_eff)).max(1e-8);
let n_lambda = config.n_lambda.max(1);
if n_lambda == 1 {
return vec![lambda_max];
}
let eps = 1e-3_f64;
let log_max = lambda_max.ln();
let log_min = (lambda_max * eps).ln();
let step = (log_max - log_min) / (n_lambda as f64 - 1.0);
(0..n_lambda)
.map(|k| (log_max - step * k as f64).exp())
.collect()
}
pub(crate) fn wnet_cv_lambda(
design: &FdMatrix,
y: &[f64],
config: &WnetConfig,
) -> Result<f64, FdarError> {
let (n, _p) = design.shape();
if y.len() != n {
return Err(FdarError::InvalidDimension {
parameter: "y",
expected: format!("{n} elements (== design rows)"),
actual: format!("{} elements", y.len()),
});
}
if config.n_folds < 2 {
return Err(FdarError::InvalidParameter {
parameter: "n_folds",
message: format!("n_folds must be >= 2, got {}", config.n_folds),
});
}
if config.n_folds > n {
return Err(FdarError::InvalidParameter {
parameter: "n_folds",
message: format!(
"n_folds ({}) must not exceed the number of observations ({n})",
config.n_folds
),
});
}
if !(0.0..=1.0).contains(&config.alpha) {
return Err(FdarError::InvalidParameter {
parameter: "alpha",
message: format!("alpha must be in [0, 1], got {}", config.alpha),
});
}
if let Some(grid) = &config.lambda_grid {
if grid.is_empty() {
return Err(FdarError::InvalidParameter {
parameter: "lambda_grid",
message: "explicit lambda_grid must be non-empty".to_string(),
});
}
}
let grid = build_lambda_grid(design, y, config);
let folds = crate::cv::create_folds(n, config.n_folds, config.seed);
let fold_sets: Vec<(Vec<usize>, Vec<usize>)> = (0..config.n_folds)
.map(|f| crate::cv::fold_indices(&folds, f))
.collect();
let mut best_lambda = grid[0];
let mut best_mse = f64::INFINITY;
let tie_eps = 1e-12;
for &lam in &grid {
let mut total_sse = 0.0_f64;
let mut scored = 0usize;
for (train_idx, test_idx) in &fold_sets {
if train_idx.is_empty() || test_idx.is_empty() {
continue;
}
let train_data = crate::cv::subset_rows(design, train_idx);
let train_y = crate::cv::subset_vec(y, train_idx);
let (intercept, beta) = elastic_net_cd(
&train_data,
&train_y,
lam,
config.alpha,
config.max_iter,
config.tol,
)?;
for &oi in test_idx {
let mut yhat = intercept;
for j in 0..design.ncols() {
yhat += design[(oi, j)] * beta[j];
}
let e = y[oi] - yhat;
total_sse += e * e;
scored += 1;
}
}
if scored == 0 {
continue;
}
let mse = total_sse / scored as f64;
if mse < best_mse - tie_eps {
best_mse = mse;
best_lambda = lam;
}
}
Ok(best_lambda)
}
#[must_use = "expensive computation whose result should not be discarded"]
pub fn wnet(data: &FdMatrix, y: &[f64], config: &WnetConfig) -> Result<WnetResult, FdarError> {
let (n, m) = data.shape();
if n < 3 {
return Err(FdarError::InvalidDimension {
parameter: "data",
expected: "at least 3 rows (observations)".to_string(),
actual: format!("{n} rows"),
});
}
if m == 0 {
return Err(FdarError::InvalidDimension {
parameter: "data",
expected: "at least 1 column (evaluation point)".to_string(),
actual: format!("{m} columns"),
});
}
if y.len() != n {
return Err(FdarError::InvalidDimension {
parameter: "y",
expected: format!("{n} elements (== data rows)"),
actual: format!("{} elements", y.len()),
});
}
if !(0.0..=1.0).contains(&config.alpha) {
return Err(FdarError::InvalidParameter {
parameter: "alpha",
message: format!("alpha must be in [0, 1], got {}", config.alpha),
});
}
if config.n_folds < 2 {
return Err(FdarError::InvalidParameter {
parameter: "n_folds",
message: format!("n_folds must be >= 2, got {}", config.n_folds),
});
}
if config.n_folds > n {
return Err(FdarError::InvalidParameter {
parameter: "n_folds",
message: format!(
"n_folds ({}) must not exceed the number of observations ({n})",
config.n_folds
),
});
}
if config.max_iter == 0 {
return Err(FdarError::InvalidParameter {
parameter: "max_iter",
message: "max_iter must be >= 1".to_string(),
});
}
if !config.tol.is_finite() || config.tol < 0.0 {
return Err(FdarError::InvalidParameter {
parameter: "tol",
message: format!("tol must be finite and >= 0, got {}", config.tol),
});
}
if let Some(grid) = &config.lambda_grid {
if grid.is_empty() {
return Err(FdarError::InvalidParameter {
parameter: "lambda_grid",
message: "explicit lambda_grid must be non-empty".to_string(),
});
}
}
let (design, layout) =
curves_to_coeff_design(data, config.family.clone(), config.mode, config.level)?;
let lambda = wnet_cv_lambda(&design, y, config)?;
let (intercept, coeff_weights) = elastic_net_cd(
&design,
y,
lambda,
config.alpha,
config.max_iter,
config.tol,
)?;
let selected: Vec<usize> = coeff_weights
.iter()
.enumerate()
.filter(|(_, &b)| b != 0.0)
.map(|(j, _)| j)
.collect();
let fitted_values = compute_fitted_affine(&design, &coeff_weights, intercept);
let residuals: Vec<f64> = y
.iter()
.zip(&fitted_values)
.map(|(&yi, &yh)| yi - yh)
.collect();
let beta_t = coeff_weights_to_beta_t(&coeff_weights, &layout)?;
Ok(WnetResult {
intercept,
beta_t,
fitted_values,
residuals,
coeff_weights,
selected,
lambda,
alpha: config.alpha,
family: config.family.clone(),
mode: config.mode,
level: layout.levels(),
})
}
impl WnetResult {
pub fn predict(&self, new: &FdMatrix) -> Result<Vec<f64>, FdarError> {
let train_m = self.beta_t.len();
if new.nrows() == 0 {
return Err(FdarError::InvalidDimension {
parameter: "new",
expected: "at least 1 row (curve)".to_string(),
actual: "0 rows".to_string(),
});
}
if new.ncols() != train_m {
return Err(FdarError::InvalidDimension {
parameter: "new",
expected: format!("{train_m} columns (== training grid length)"),
actual: format!("{} columns", new.ncols()),
});
}
let (design, _layout) =
curves_to_coeff_design(new, self.family.clone(), self.mode, Some(self.level))?;
if design.ncols() != self.coeff_weights.len() {
return Err(FdarError::InvalidDimension {
parameter: "new",
expected: format!(
"coefficient-space width {} (== stored coeff_weights)",
self.coeff_weights.len()
),
actual: format!("{} coefficients", design.ncols()),
});
}
Ok(compute_fitted_affine(
&design,
&self.coeff_weights,
self.intercept,
))
}
#[must_use]
pub fn beta_t(&self) -> &[f64] {
&self.beta_t
}
#[must_use]
pub fn coefficient_function(&self) -> &[f64] {
&self.beta_t
}
#[must_use]
pub fn fitted_values(&self) -> &[f64] {
&self.fitted_values
}
}
fn compute_fitted_affine(design: &FdMatrix, coeffs: &[f64], intercept: f64) -> Vec<f64> {
let (n, p) = design.shape();
(0..n)
.map(|i| {
let mut yhat = intercept;
for j in 0..p {
yhat += design[(i, j)] * coeffs[j];
}
yhat
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::matrix::FdMatrix;
fn pseudo_random(n: usize, seed: u64) -> Vec<f64> {
let mut state = seed.wrapping_add(0x9E37_79B9_7F4A_7C15);
(0..n)
.map(|_| {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
let u = (state >> 11) as f64 / (1u64 << 53) as f64;
2.0 * u - 1.0
})
.collect()
}
fn spanning_design(n: usize, m: usize, seed0: u64) -> FdMatrix {
let mut flat = vec![0.0_f64; n * m];
for i in 0..n {
let row = pseudo_random(m, seed0 + i as u64);
for j in 0..m {
flat[i + j * n] = row[j];
}
}
FdMatrix::from_column_major(flat, n, m).unwrap()
}
fn rel_l2(recovered: &[f64], truth: &[f64]) -> f64 {
let num: f64 = recovered
.iter()
.zip(truth)
.map(|(a, b)| (a - b) * (a - b))
.sum::<f64>()
.sqrt();
let den: f64 = truth.iter().map(|b| b * b).sum::<f64>().sqrt().max(1e-300);
num / den
}
fn recovery_for_method(method: WcrMethod) {
let (n, m) = (120usize, 32usize);
let data = spanning_design(n, m, 1000);
let family = WaveletFamily::Daubechies(4);
let mode = BoundaryMode::Periodic;
let (design, layout) = curves_to_coeff_design(&data, family.clone(), mode, None).unwrap();
let p = design.ncols();
let beta_coeff = pseudo_random(p, 77);
let true_intercept = 0.37_f64;
let y: Vec<f64> = (0..n)
.map(|i| {
let mut acc = true_intercept;
for j in 0..p {
acc += design[(i, j)] * beta_coeff[j];
}
acc
})
.collect();
let beta_t_true = coeff_weights_to_beta_t(&beta_coeff, &layout).unwrap();
let config = WcrConfig {
family,
mode,
level: None,
ncomp: p.min(n),
method,
..Default::default()
};
let fit = wcr(&data, &y, &config).unwrap();
assert_eq!(fit.method, method);
assert_eq!(fit.beta_t.len(), m);
assert_eq!(fit.coeff_weights.len(), p);
let e = rel_l2(&fit.beta_t, &beta_t_true);
assert!(
e < 1e-6,
"{method:?}: beta_t recovery rel L2 err {e} exceeds tolerance on spanning full-rank design"
);
assert!(fit.beta_t.iter().all(|x| x.is_finite()));
assert!(fit.fitted_values.iter().all(|x| x.is_finite()));
assert!(fit.residuals.iter().all(|x| x.is_finite()));
let max_resid = fit
.residuals
.iter()
.fold(0.0_f64, |acc, &r| acc.max(r.abs()));
assert!(
max_resid < 1e-6,
"{method:?}: residuals not ~0 ({max_resid})"
);
}
#[test]
fn wcr_pcr_recovers_known_beta_t_on_spanning_design() {
recovery_for_method(WcrMethod::Pcr);
}
#[test]
fn wcr_pls_recovers_known_beta_t_on_spanning_design() {
recovery_for_method(WcrMethod::Pls);
}
#[test]
fn wcr_default_config_is_db4_periodic_auto_pcr() {
let c = WcrConfig::default();
assert_eq!(c.family, WaveletFamily::Daubechies(4));
assert_eq!(c.mode, BoundaryMode::Periodic);
assert_eq!(c.level, None);
assert_eq!(c.method, WcrMethod::Pcr);
assert_eq!(WcrMethod::default(), WcrMethod::Pcr);
}
#[test]
fn curves_to_coeff_design_layout_and_shape() {
let (n, m) = (10usize, 48usize);
let data = spanning_design(n, m, 500);
let (design, layout) = curves_to_coeff_design(
&data,
WaveletFamily::Daubechies(4),
BoundaryMode::Periodic,
None,
)
.unwrap();
assert_eq!(design.nrows(), n);
assert_eq!(design.ncols(), layout.total_len());
assert_eq!(layout.signal_len, m);
assert_eq!(layout.levels(), layout.detail_lens.len());
}
#[test]
fn coeff_weights_to_beta_t_inverts_decompose() {
let (n, m) = (4usize, 48usize);
let data = spanning_design(n, m, 900);
let (design, layout) = curves_to_coeff_design(
&data,
WaveletFamily::Daubechies(6),
BoundaryMode::Periodic,
None,
)
.unwrap();
let row0: Vec<f64> = (0..design.ncols()).map(|j| design[(0, j)]).collect();
let recon = coeff_weights_to_beta_t(&row0, &layout).unwrap();
let orig = data.row(0);
assert!(rel_l2(&recon, &orig) < 1e-10);
}
#[test]
fn coeff_weights_to_beta_t_rejects_wrong_length() {
let (n, m) = (4usize, 48usize);
let data = spanning_design(n, m, 901);
let (_design, layout) =
curves_to_coeff_design(&data, WaveletFamily::Haar, BoundaryMode::Periodic, None)
.unwrap();
let wrong = vec![0.0; layout.total_len() + 1];
assert!(matches!(
coeff_weights_to_beta_t(&wrong, &layout),
Err(FdarError::InvalidDimension { .. })
));
}
fn base_config() -> WcrConfig {
WcrConfig {
ncomp: 3,
..Default::default()
}
}
#[test]
fn wcr_rejects_too_few_rows() {
let data = spanning_design(2, 48, 1);
let y = vec![0.0, 1.0];
assert!(matches!(
wcr(&data, &y, &base_config()),
Err(FdarError::InvalidDimension { .. })
));
}
#[test]
fn wcr_rejects_mismatched_y_len() {
let data = spanning_design(10, 48, 2);
let y = vec![0.0; 9];
assert!(matches!(
wcr(&data, &y, &base_config()),
Err(FdarError::InvalidDimension { .. })
));
}
#[test]
fn wcr_rejects_zero_ncomp() {
let data = spanning_design(10, 48, 3);
let y = vec![0.0; 10];
let config = WcrConfig {
ncomp: 0,
..Default::default()
};
assert!(matches!(
wcr(&data, &y, &config),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn wcr_surfaces_unsupported_family() {
let data = spanning_design(10, 48, 4);
let y = vec![0.0; 10];
let config = WcrConfig {
family: WaveletFamily::Daubechies(11),
..base_config()
};
assert!(matches!(
wcr(&data, &y, &config),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn wcr_surfaces_level_out_of_range() {
let data = spanning_design(10, 48, 5);
let y = vec![0.0; 10];
let config = WcrConfig {
level: Some(999),
..base_config()
};
assert!(matches!(
wcr(&data, &y, &config),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn wcr_finite_outputs_both_methods() {
let (n, m) = (100usize, 40usize);
let data = spanning_design(n, m, 4242);
let y = pseudo_random(n, 8080);
for method in [WcrMethod::Pcr, WcrMethod::Pls] {
let config = WcrConfig {
ncomp: 8,
method,
..Default::default()
};
let fit = wcr(&data, &y, &config).unwrap();
assert!(fit.intercept.is_finite());
assert!(fit.beta_t.iter().all(|x| x.is_finite()));
assert!(fit.fitted_values.iter().all(|x| x.is_finite()));
assert!(fit.residuals.iter().all(|x| x.is_finite()));
}
}
#[test]
fn wcr_small_n_default_config_succeeds() {
let (n, m) = (4usize, 32usize);
let data = spanning_design(n, m, 2468);
let y = pseudo_random(n, 1357);
let config = WcrConfig::default(); let fit = wcr(&data, &y, &config).unwrap();
assert!(fit.ncomp < n, "ncomp {} exceeds n - 1", fit.ncomp);
assert_eq!(fit.beta_t.len(), m);
assert!(fit.intercept.is_finite());
assert!(fit.beta_t.iter().all(|x| x.is_finite()));
assert!(fit.fitted_values.iter().all(|x| x.is_finite()));
assert!(fit.residuals.iter().all(|x| x.is_finite()));
let data3 = spanning_design(3, m, 2469);
let y3 = pseudo_random(3, 1358);
let fit3 = wcr(&data3, &y3, &WcrConfig::default()).unwrap();
assert!(fit3.ncomp <= 2);
assert!(fit3.beta_t.iter().all(|x| x.is_finite()));
}
#[test]
fn wcr_predict_reproduces_training_fitted() {
let (n, m) = (120usize, 32usize);
let data = spanning_design(n, m, 6100);
let y = pseudo_random(n, 6101);
let config = WcrConfig {
ncomp: 8,
..Default::default()
};
let fit = wcr(&data, &y, &config).unwrap();
let preds = fit.predict(&data).unwrap();
assert_eq!(preds.len(), fit.fitted_values.len());
for (i, (&p, &f)) in preds.iter().zip(&fit.fitted_values).enumerate() {
assert!(
(p - f).abs() <= 1e-8,
"wcr predict[{i}] {p} != fitted {f} (|Δ| {})",
(p - f).abs()
);
}
}
#[test]
fn wcr_predict_on_new_curves_is_finite_and_rejects_grid_mismatch() {
let (n, m) = (100usize, 32usize);
let data = spanning_design(n, m, 6200);
let y = pseudo_random(n, 6201);
let fit = wcr(
&data,
&y,
&WcrConfig {
ncomp: 6,
..Default::default()
},
)
.unwrap();
let fresh = spanning_design(40, m, 6202);
let preds = fit.predict(&fresh).unwrap();
assert_eq!(preds.len(), 40);
assert!(preds.iter().all(|x| x.is_finite()));
let wrong = spanning_design(10, m + 8, 6203);
assert!(matches!(
fit.predict(&wrong),
Err(FdarError::InvalidDimension { .. })
));
}
#[test]
fn wcr_predict_on_zero_row_input_errors_naming_new_no_panic() {
let (n, m) = (100usize, 32usize);
let data = spanning_design(n, m, 6400);
let y = pseudo_random(n, 6401);
let fit = wcr(
&data,
&y,
&WcrConfig {
ncomp: 6,
..Default::default()
},
)
.unwrap();
let empty = FdMatrix::zeros(0, m);
match fit.predict(&empty) {
Err(FdarError::InvalidDimension { parameter, .. }) => {
assert_eq!(parameter, "new");
}
other => panic!("expected InvalidDimension naming \"new\", got {other:?}"),
}
}
fn sparse_wnet_problem(
n: usize,
m: usize,
seed0: u64,
) -> (FdMatrix, FdMatrix, CoeffLayout, Vec<f64>, Vec<usize>) {
let data = spanning_design(n, m, seed0);
let family = WaveletFamily::Daubechies(4);
let mode = BoundaryMode::Periodic;
let (design, layout) = curves_to_coeff_design(&data, family, mode, None).unwrap();
let p = design.ncols();
let support: Vec<usize> = vec![0, 2, p / 2, p - 3]
.into_iter()
.filter(|&j| j < p)
.collect();
let mut beta_coeff = vec![0.0_f64; p];
let mags = [4.0, -3.5, 5.0, -4.5];
for (k, &j) in support.iter().enumerate() {
beta_coeff[j] = mags[k % mags.len()];
}
(data, design, layout, beta_coeff, support)
}
#[test]
fn wnet_elastic_net_cd_recovers_sparse_support() {
let (n, m) = (256usize, 32usize);
let (_data, design, _layout, beta_coeff, support) = sparse_wnet_problem(n, m, 3000);
let p = design.ncols();
let intercept_true = 0.5_f64;
let y: Vec<f64> = (0..n)
.map(|i| {
let mut acc = intercept_true;
for j in 0..p {
acc += design[(i, j)] * beta_coeff[j];
}
acc
})
.collect();
let (intercept, beta) = elastic_net_cd(&design, &y, 0.05, 0.9, 2000, 1e-8).unwrap();
assert!(intercept.is_finite());
assert!(beta.iter().all(|b| b.is_finite()));
let selected: Vec<usize> = beta
.iter()
.enumerate()
.filter(|(_, &b)| b.abs() > 1e-8)
.map(|(j, _)| j)
.collect();
for &j in &support {
assert!(
selected.contains(&j),
"true-support coeff {j} not selected (selected={selected:?})"
);
}
assert!(
selected.len() < p / 2,
"selection not sparse: |selected|={} of P={p}",
selected.len()
);
}
#[test]
fn wnet_fixed_lambda_end_to_end_finite() {
let (n, m) = (200usize, 32usize);
let (data, design, _layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 3100);
let p = design.ncols();
let y: Vec<f64> = (0..n)
.map(|i| {
let mut acc = 0.25;
for j in 0..p {
acc += design[(i, j)] * beta_coeff[j];
}
acc
})
.collect();
let config = WnetConfig {
n_lambda: 15,
n_folds: 4,
..Default::default()
};
let fit = wnet(&data, &y, &config).unwrap();
assert_eq!(fit.beta_t.len(), m);
assert_eq!(fit.coeff_weights.len(), p);
assert!(fit.intercept.is_finite());
assert!(fit.beta_t.iter().all(|x| x.is_finite()));
assert!(fit.fitted_values.iter().all(|x| x.is_finite()));
assert!(fit.residuals.iter().all(|x| x.is_finite()));
assert!(fit.coeff_weights.iter().all(|x| x.is_finite()));
for &j in &fit.selected {
assert!(fit.coeff_weights[j] != 0.0);
}
}
#[test]
fn wnet_default_config_is_db4_periodic_auto() {
let c = WnetConfig::default();
assert_eq!(c.family, WaveletFamily::Daubechies(4));
assert_eq!(c.mode, BoundaryMode::Periodic);
assert_eq!(c.level, None);
assert!((c.alpha - 0.5).abs() < 1e-15);
assert_eq!(c.lambda_grid, None);
assert_eq!(c.n_lambda, 50);
assert_eq!(c.n_folds, 5);
assert_eq!(c.seed, 0);
}
#[test]
fn wnet_cv_lambda_is_deterministic_across_runs() {
let (n, m) = (200usize, 32usize);
let (data, design, _layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 3200);
let p = design.ncols();
let noise = pseudo_random(n, 9999);
let y: Vec<f64> = (0..n)
.map(|i| {
let mut acc = 0.1;
for j in 0..p {
acc += design[(i, j)] * beta_coeff[j];
}
acc + 0.05 * noise[i]
})
.collect();
let config = WnetConfig {
alpha: 0.8,
n_lambda: 20,
n_folds: 5,
seed: 0,
..Default::default()
};
let fit1 = wnet(&data, &y, &config).unwrap();
let fit2 = wnet(&data, &y, &config).unwrap();
assert_eq!(
fit1.lambda, fit2.lambda,
"CV-selected lambda differs across runs: {} vs {}",
fit1.lambda, fit2.lambda
);
let l1 = wnet_cv_lambda(&design, &y, &config).unwrap();
let l2 = wnet_cv_lambda(&design, &y, &config).unwrap();
assert_eq!(l1, l2);
}
#[test]
fn wnet_recovers_beta_t_on_snr_data() {
let (n, m) = (300usize, 32usize);
let (data, design, layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 3300);
let p = design.ncols();
let beta_t_true = coeff_weights_to_beta_t(&beta_coeff, &layout).unwrap();
let signal: Vec<f64> = (0..n)
.map(|i| {
let mut acc = 0.0;
for j in 0..p {
acc += design[(i, j)] * beta_coeff[j];
}
acc
})
.collect();
let sig_sd = {
let mean = signal.iter().sum::<f64>() / n as f64;
(signal.iter().map(|s| (s - mean).powi(2)).sum::<f64>() / n as f64).sqrt()
};
let noise = pseudo_random(n, 4141);
let noise_scale = 0.05 * sig_sd; let y: Vec<f64> = (0..n)
.map(|i| 0.3 + signal[i] + noise_scale * noise[i])
.collect();
let config = WnetConfig {
alpha: 0.7,
n_lambda: 30,
n_folds: 5,
..Default::default()
};
let fit = wnet(&data, &y, &config).unwrap();
let nonzero = fit.coeff_weights.iter().filter(|&&b| b != 0.0).count();
assert!(nonzero > 0, "degenerate all-zero fit at CV lambda");
let e = rel_l2(&fit.beta_t, &beta_t_true);
assert!(
e < 0.35,
"wnet beta_t recovery rel L2 err {e} exceeds tolerance on SNR data"
);
assert!(fit.beta_t.iter().all(|x| x.is_finite()));
assert!(fit.fitted_values.iter().all(|x| x.is_finite()));
}
fn base_wnet_config() -> WnetConfig {
WnetConfig {
n_lambda: 10,
n_folds: 3,
..Default::default()
}
}
#[test]
fn wnet_rejects_too_few_rows() {
let data = spanning_design(2, 32, 10);
let y = vec![0.0, 1.0];
assert!(matches!(
wnet(&data, &y, &base_wnet_config()),
Err(FdarError::InvalidDimension { .. })
));
}
#[test]
fn wnet_rejects_zero_cols() {
let data = FdMatrix::zeros(5, 0);
let y = vec![0.0; 5];
assert!(matches!(
wnet(&data, &y, &base_wnet_config()),
Err(FdarError::InvalidDimension { .. })
));
}
#[test]
fn wnet_rejects_mismatched_y_len() {
let data = spanning_design(10, 32, 11);
let y = vec![0.0; 9];
assert!(matches!(
wnet(&data, &y, &base_wnet_config()),
Err(FdarError::InvalidDimension { .. })
));
}
#[test]
fn wnet_rejects_alpha_out_of_range() {
let data = spanning_design(10, 32, 12);
let y = vec![0.0; 10];
let config = WnetConfig {
alpha: 1.5,
..base_wnet_config()
};
assert!(matches!(
wnet(&data, &y, &config),
Err(FdarError::InvalidParameter { .. })
));
let config = WnetConfig {
alpha: -0.1,
..base_wnet_config()
};
assert!(matches!(
wnet(&data, &y, &config),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn wnet_rejects_too_few_folds() {
let data = spanning_design(10, 32, 13);
let y = vec![0.0; 10];
let config = WnetConfig {
n_folds: 1,
..base_wnet_config()
};
assert!(matches!(
wnet(&data, &y, &config),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn wnet_rejects_too_many_folds() {
let data = spanning_design(10, 32, 130);
let y = vec![0.0; 10];
let config = WnetConfig {
n_folds: 11,
..base_wnet_config()
};
assert!(matches!(
wnet(&data, &y, &config),
Err(FdarError::InvalidParameter { .. })
));
let (design, _layout) = curves_to_coeff_design(
&data,
WaveletFamily::Daubechies(4),
BoundaryMode::Periodic,
None,
)
.unwrap();
assert!(matches!(
wnet_cv_lambda(&design, &y, &config),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn wnet_rejects_negative_or_nan_tol() {
let data = spanning_design(10, 32, 131);
let y = pseudo_random(10, 5);
for bad in [-1e-6_f64, f64::NAN] {
let config = WnetConfig {
tol: bad,
..base_wnet_config()
};
assert!(matches!(
wnet(&data, &y, &config),
Err(FdarError::InvalidParameter { .. })
));
}
let (design, _layout) = curves_to_coeff_design(
&data,
WaveletFamily::Daubechies(4),
BoundaryMode::Periodic,
None,
)
.unwrap();
assert!(matches!(
elastic_net_cd(&design, &y, 0.1, 0.5, 100, -1.0),
Err(FdarError::InvalidParameter { .. })
));
assert!(matches!(
elastic_net_cd(&design, &y, 0.1, 0.5, 100, f64::NAN),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn wnet_rejects_zero_max_iter() {
let data = spanning_design(10, 32, 132);
let y = pseudo_random(10, 6);
let config = WnetConfig {
max_iter: 0,
..base_wnet_config()
};
assert!(matches!(
wnet(&data, &y, &config),
Err(FdarError::InvalidParameter { .. })
));
let (design, _layout) = curves_to_coeff_design(
&data,
WaveletFamily::Daubechies(4),
BoundaryMode::Periodic,
None,
)
.unwrap();
assert!(matches!(
elastic_net_cd(&design, &y, 0.1, 0.5, 0, 1e-6),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn wnet_rejects_empty_lambda_grid() {
let data = spanning_design(10, 32, 14);
let y = vec![0.0; 10];
let config = WnetConfig {
lambda_grid: Some(vec![]),
..base_wnet_config()
};
assert!(matches!(
wnet(&data, &y, &config),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn wnet_surfaces_unsupported_family() {
let data = spanning_design(10, 32, 15);
let y = vec![0.0; 10];
let config = WnetConfig {
family: WaveletFamily::Daubechies(11),
..base_wnet_config()
};
assert!(matches!(
wnet(&data, &y, &config),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn wnet_surfaces_level_out_of_range() {
let data = spanning_design(10, 32, 16);
let y = vec![0.0; 10];
let config = WnetConfig {
level: Some(999),
..base_wnet_config()
};
assert!(matches!(
wnet(&data, &y, &config),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn wnet_finite_outputs_on_larger_snr_design() {
let (n, m) = (256usize, 48usize);
let (data, design, _layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 3400);
let p = design.ncols();
let noise = pseudo_random(n, 2727);
let y: Vec<f64> = (0..n)
.map(|i| {
let mut acc = 0.2;
for j in 0..p {
acc += design[(i, j)] * beta_coeff[j];
}
acc + 0.1 * noise[i]
})
.collect();
let config = WnetConfig {
alpha: 0.6,
n_lambda: 25,
n_folds: 5,
..Default::default()
};
let fit = wnet(&data, &y, &config).unwrap();
assert!(fit.intercept.is_finite());
assert!(fit.lambda.is_finite());
assert!(fit.beta_t.iter().all(|x| x.is_finite()));
assert!(fit.fitted_values.iter().all(|x| x.is_finite()));
assert!(fit.residuals.iter().all(|x| x.is_finite()));
assert!(fit.coeff_weights.iter().all(|x| x.is_finite()));
}
#[test]
fn wnet_explicit_lambda_grid_is_used() {
let (n, m) = (120usize, 32usize);
let (data, design, _layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 3500);
let p = design.ncols();
let y: Vec<f64> = (0..n)
.map(|i| {
let mut acc = 0.0;
for j in 0..p {
acc += design[(i, j)] * beta_coeff[j];
}
acc
})
.collect();
let config = WnetConfig {
lambda_grid: Some(vec![0.123]),
..base_wnet_config()
};
let fit = wnet(&data, &y, &config).unwrap();
assert!((fit.lambda - 0.123).abs() < 1e-15);
}
#[test]
fn wnet_predict_reproduces_training_fitted() {
let (n, m) = (200usize, 32usize);
let (data, design, _layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 6300);
let p = design.ncols();
let noise = pseudo_random(n, 6301);
let y: Vec<f64> = (0..n)
.map(|i| {
let mut acc = 0.4;
for j in 0..p {
acc += design[(i, j)] * beta_coeff[j];
}
acc + 0.05 * noise[i]
})
.collect();
let config = WnetConfig {
alpha: 0.7,
n_lambda: 20,
n_folds: 5,
..Default::default()
};
let fit = wnet(&data, &y, &config).unwrap();
let preds = fit.predict(&data).unwrap();
assert_eq!(preds.len(), fit.fitted_values.len());
for (i, (&pv, &f)) in preds.iter().zip(&fit.fitted_values).enumerate() {
assert!(
(pv - f).abs() <= 1e-8,
"wnet predict[{i}] {pv} != fitted {f} (|Δ| {})",
(pv - f).abs()
);
}
}
#[test]
fn wnet_predict_on_new_curves_is_finite_and_rejects_grid_mismatch() {
let (n, m) = (150usize, 32usize);
let (data, design, _layout, beta_coeff, _support) = sparse_wnet_problem(n, m, 6400);
let p = design.ncols();
let y: Vec<f64> = (0..n)
.map(|i| {
let mut acc = 0.2;
for j in 0..p {
acc += design[(i, j)] * beta_coeff[j];
}
acc
})
.collect();
let fit = wnet(
&data,
&y,
&WnetConfig {
n_lambda: 12,
n_folds: 4,
..Default::default()
},
)
.unwrap();
let fresh = spanning_design(30, m, 6401);
let preds = fit.predict(&fresh).unwrap();
assert_eq!(preds.len(), 30);
assert!(preds.iter().all(|x| x.is_finite()));
let wrong = spanning_design(10, m + 16, 6402);
assert!(matches!(
fit.predict(&wrong),
Err(FdarError::InvalidDimension { .. })
));
let empty = FdMatrix::zeros(0, m);
match fit.predict(&empty) {
Err(FdarError::InvalidDimension { parameter, .. }) => {
assert_eq!(parameter, "new");
}
other => panic!("expected InvalidDimension naming \"new\", got {other:?}"),
}
}
#[test]
fn accessors_return_stored_slices() {
let (n, m) = (100usize, 32usize);
let data = spanning_design(n, m, 6500);
let y = pseudo_random(n, 6501);
let wcr_fit = wcr(
&data,
&y,
&WcrConfig {
ncomp: 5,
..Default::default()
},
)
.unwrap();
assert_eq!(wcr_fit.beta_t(), wcr_fit.beta_t.as_slice());
assert_eq!(wcr_fit.coefficient_function(), wcr_fit.beta_t.as_slice());
assert_eq!(wcr_fit.beta_t().len(), m);
assert_eq!(wcr_fit.fitted_values(), wcr_fit.fitted_values.as_slice());
assert_eq!(wcr_fit.fitted_values().len(), n);
let wnet_fit = wnet(
&data,
&y,
&WnetConfig {
n_lambda: 10,
n_folds: 4,
..Default::default()
},
)
.unwrap();
assert_eq!(wnet_fit.beta_t(), wnet_fit.beta_t.as_slice());
assert_eq!(wnet_fit.coefficient_function(), wnet_fit.beta_t.as_slice());
assert_eq!(wnet_fit.beta_t().len(), m);
assert_eq!(wnet_fit.fitted_values(), wnet_fit.fitted_values.as_slice());
assert_eq!(wnet_fit.fitted_values().len(), n);
}
}