pub mod filters;
pub mod regression;
use crate::error::FdarError;
use filters::FilterBank;
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub enum WaveletFamily {
Haar,
Daubechies(usize),
}
impl WaveletFamily {
pub fn from_db_order(order: usize) -> Result<Self, FdarError> {
match order {
1 => Ok(WaveletFamily::Haar),
2..=10 => Ok(WaveletFamily::Daubechies(order)),
_ => Err(FdarError::InvalidParameter {
parameter: "order",
message: format!(
"Daubechies order {order} out of range: supported orders are 1..=10 (1 == Haar)"
),
}),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub enum BoundaryMode {
#[default]
Periodic,
Symmetric,
}
fn extended_len(n: usize, mode: BoundaryMode) -> usize {
match mode {
BoundaryMode::Periodic => {
if n % 2 == 0 {
n
} else {
n + 1
}
}
BoundaryMode::Symmetric => 2 * n,
}
}
fn extend_signal(signal: &[f64], mode: BoundaryMode) -> Vec<f64> {
let n = signal.len();
let m = extended_len(n, mode);
let mut ext = vec![0.0_f64; m];
match mode {
BoundaryMode::Periodic => {
ext[..n].copy_from_slice(signal);
if m > n {
ext[n] = signal[n - 1];
}
}
BoundaryMode::Symmetric => {
for i in 0..n {
ext[i] = signal[i];
ext[m - 1 - i] = signal[i];
}
}
}
ext
}
#[inline]
fn coeff_len(n: usize, mode: BoundaryMode) -> usize {
extended_len(n, mode) / 2
}
fn core_analysis(sig: &[f64], fb: &FilterBank) -> (Vec<f64>, Vec<f64>) {
let m = sig.len();
debug_assert!(m % 2 == 0 && m > 0);
let l = fb.filter_len();
let out = m / 2;
let mut approx = vec![0.0_f64; out];
let mut detail = vec![0.0_f64; out];
for t in 0..out {
let mut a = 0.0_f64;
let mut d = 0.0_f64;
for k in 0..l {
let idx = (2 * t + k) % m;
let s = sig[idx];
a += fb.dec_lo[k] * s;
d += fb.dec_hi[k] * s;
}
approx[t] = a;
detail[t] = d;
}
(approx, detail)
}
fn core_synthesis(approx: &[f64], detail: &[f64], fb: &FilterBank, m: usize) -> Vec<f64> {
let l = fb.filter_len();
let out = approx.len();
let mut signal = vec![0.0_f64; m];
for t in 0..out {
let a = approx[t];
let d = detail[t];
for k in 0..l {
let idx = (2 * t + k) % m;
signal[idx] += fb.dec_lo[k] * a + fb.dec_hi[k] * d;
}
}
signal
}
#[must_use = "the approximation/detail coefficients are the result of the transform"]
pub(crate) fn single_level_analysis(
signal: &[f64],
fb: &FilterBank,
mode: BoundaryMode,
) -> Result<(Vec<f64>, Vec<f64>), FdarError> {
if signal.is_empty() {
return Err(FdarError::InvalidParameter {
parameter: "signal",
message: "signal must be non-empty".to_string(),
});
}
let ext = extend_signal(signal, mode);
Ok(core_analysis(&ext, fb))
}
#[must_use = "the reconstructed signal is the result of the inverse transform"]
pub(crate) fn single_level_synthesis(
approx: &[f64],
detail: &[f64],
fb: &FilterBank,
mode: BoundaryMode,
output_len: usize,
) -> Result<Vec<f64>, FdarError> {
if output_len == 0 {
return Err(FdarError::InvalidParameter {
parameter: "output_len",
message: "output_len must be non-zero".to_string(),
});
}
if approx.len() != detail.len() {
return Err(FdarError::InvalidDimension {
parameter: "detail",
expected: format!("{} (== approx length)", approx.len()),
actual: detail.len().to_string(),
});
}
let expected_coeff_len = coeff_len(output_len, mode);
if approx.len() != expected_coeff_len {
return Err(FdarError::InvalidDimension {
parameter: "approx",
expected: format!("{expected_coeff_len} (mode-dependent coefficient length)"),
actual: approx.len().to_string(),
});
}
let m = extended_len(output_len, mode);
debug_assert!(
output_len <= m,
"output_len {output_len} > extended length {m}; likely off-by-one in odd-signal path"
);
let mut signal = core_synthesis(approx, detail, fb, m);
signal.truncate(output_len);
Ok(signal)
}
pub fn max_level(signal_len: usize, family: &WaveletFamily) -> Result<usize, FdarError> {
if signal_len == 0 {
return Err(FdarError::InvalidParameter {
parameter: "signal_len",
message: "signal length must be non-zero".to_string(),
});
}
let fb = filters::filter_bank(family)?;
let filter_len = fb.filter_len();
if filter_len <= 1 {
return Err(FdarError::InvalidParameter {
parameter: "family",
message: "filter length must exceed 1".to_string(),
});
}
let ratio = signal_len as f64 / (filter_len - 1) as f64;
let level = if ratio < 1.0 {
0
} else {
ratio.log2().floor() as usize
};
if level < 1 {
return Err(FdarError::InvalidParameter {
parameter: "signal_len",
message: format!(
"signal length {signal_len} is too short for even one useful decomposition level \
with filter length {filter_len}"
),
});
}
Ok(level)
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct WaveletCoeffs {
pub approx: Vec<f64>,
pub details: Vec<Vec<f64>>,
pub levels: usize,
pub signal_len: usize,
pub family: WaveletFamily,
pub mode: BoundaryMode,
pub(crate) level_lens: Vec<usize>,
}
impl WaveletCoeffs {
#[must_use]
pub fn levels(&self) -> usize {
self.levels
}
#[must_use]
pub fn signal_len(&self) -> usize {
self.signal_len
}
#[must_use]
pub fn approx(&self) -> &[f64] {
&self.approx
}
#[must_use]
pub fn detail(&self, level: usize) -> Option<&[f64]> {
self.details.get(level).map(Vec::as_slice)
}
#[must_use]
pub fn family(&self) -> &WaveletFamily {
&self.family
}
#[must_use]
pub fn mode(&self) -> BoundaryMode {
self.mode
}
}
#[must_use = "the coefficient pyramid is the result of the transform"]
pub fn decompose(
signal: &[f64],
family: WaveletFamily,
mode: BoundaryMode,
level: Option<usize>,
) -> Result<WaveletCoeffs, FdarError> {
if signal.is_empty() {
return Err(FdarError::InvalidParameter {
parameter: "signal",
message: "signal must be non-empty".to_string(),
});
}
let max_lvl = max_level(signal.len(), &family)?;
let effective_level = match level {
None => max_lvl,
Some(0) => {
return Err(FdarError::InvalidParameter {
parameter: "level",
message: "decomposition level must be at least 1".to_string(),
});
}
Some(l) if l > max_lvl => {
return Err(FdarError::InvalidParameter {
parameter: "level",
message: format!(
"decomposition level {l} exceeds the maximum useful level {max_lvl} \
for signal length {} with this family",
signal.len()
),
});
}
Some(l) => l,
};
let fb = filters::filter_bank(&family)?;
let mut details: Vec<Vec<f64>> = Vec::with_capacity(effective_level);
let mut level_lens: Vec<usize> = Vec::with_capacity(effective_level);
let mut current = signal.to_vec();
for _ in 0..effective_level {
level_lens.push(current.len());
let (approx, detail) = single_level_analysis(¤t, &fb, mode)?;
details.push(detail);
current = approx;
}
Ok(WaveletCoeffs {
approx: current,
details,
levels: effective_level,
signal_len: signal.len(),
family,
mode,
level_lens,
})
}
#[must_use = "the reconstructed signal is the result of the inverse transform"]
pub fn reconstruct(coeffs: &WaveletCoeffs) -> Result<Vec<f64>, FdarError> {
if coeffs.details.len() != coeffs.levels || coeffs.level_lens.len() != coeffs.levels {
return Err(FdarError::InvalidDimension {
parameter: "coeffs",
expected: format!("{} detail bands and level lengths", coeffs.levels),
actual: format!(
"{} detail bands, {} level lengths",
coeffs.details.len(),
coeffs.level_lens.len()
),
});
}
if coeffs.levels == 0 {
return Err(FdarError::InvalidDimension {
parameter: "levels",
expected: "at least 1".to_string(),
actual: "0".to_string(),
});
}
let fb = filters::filter_bank(&coeffs.family)?;
let mut approx = coeffs.approx.clone();
for lvl in (0..coeffs.levels).rev() {
let detail = &coeffs.details[lvl];
let target_len = coeffs.level_lens[lvl];
approx = single_level_synthesis(&approx, detail, &fb, coeffs.mode, target_len)?;
}
Ok(approx)
}
#[must_use = "the per-row coefficient pyramids are the result of the transform"]
pub fn decompose_matrix(
data: &crate::matrix::FdMatrix,
family: WaveletFamily,
mode: BoundaryMode,
level: Option<usize>,
) -> Result<Vec<WaveletCoeffs>, FdarError> {
if data.nrows() == 0 || data.ncols() == 0 {
return Err(FdarError::InvalidDimension {
parameter: "data",
expected: "non-empty matrix (nrows > 0 && ncols > 0)".to_string(),
actual: format!("{}x{}", data.nrows(), data.ncols()),
});
}
let mut out = Vec::with_capacity(data.nrows());
for i in 0..data.nrows() {
let row = data.row(i);
out.push(decompose(&row, family.clone(), mode, level)?);
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::wavelet::filters::filter_bank;
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 rel_err(recon: &[f64], orig: &[f64]) -> f64 {
let num: f64 = recon
.iter()
.zip(orig)
.map(|(r, o)| (r - o) * (r - o))
.sum::<f64>()
.sqrt();
let den: f64 = orig.iter().map(|o| o * o).sum::<f64>().sqrt().max(1e-300);
num / den
}
fn round_trip_families() -> Vec<WaveletFamily> {
vec![
WaveletFamily::Haar,
WaveletFamily::Daubechies(2),
WaveletFamily::Daubechies(4),
WaveletFamily::Daubechies(6),
WaveletFamily::Daubechies(8),
WaveletFamily::Daubechies(10),
]
}
#[test]
fn haar_known_answer_coefficients() {
let a = 1.5_f64;
let b = -0.5;
let c = 3.0;
let d = 2.0;
let signal = [a, b, c, d];
let fb = filter_bank(&WaveletFamily::Haar).unwrap();
let (approx, detail) = single_level_analysis(&signal, &fb, BoundaryMode::Periodic).unwrap();
let s2 = std::f64::consts::SQRT_2;
assert!((approx[0] - (a + b) / s2).abs() < 1e-12);
assert!((approx[1] - (c + d) / s2).abs() < 1e-12);
assert!((detail[0] - (a - b) / s2).abs() < 1e-12);
assert!((detail[1] - (c - d) / s2).abs() < 1e-12);
}
#[test]
fn haar_even_length_round_trip() {
let signal = pseudo_random(8, 42);
let fb = filter_bank(&WaveletFamily::Haar).unwrap();
let (approx, detail) = single_level_analysis(&signal, &fb, BoundaryMode::Periodic).unwrap();
let recon =
single_level_synthesis(&approx, &detail, &fb, BoundaryMode::Periodic, signal.len())
.unwrap();
assert!(rel_err(&recon, &signal) < 1e-10);
}
#[test]
fn empty_signal_is_invalid_parameter() {
let fb = filter_bank(&WaveletFamily::Haar).unwrap();
assert!(matches!(
single_level_analysis(&[], &fb, BoundaryMode::Periodic),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn round_trip_periodic_non_power_of_two() {
let signal = pseudo_random(37, 7);
for fam in round_trip_families() {
let fb = filter_bank(&fam).unwrap();
let (approx, detail) =
single_level_analysis(&signal, &fb, BoundaryMode::Periodic).unwrap();
let recon =
single_level_synthesis(&approx, &detail, &fb, BoundaryMode::Periodic, signal.len())
.unwrap();
let e = rel_err(&recon, &signal);
assert!(e < 1e-10, "{fam:?} periodic n=37 rel err {e}");
}
}
#[test]
fn round_trip_symmetric_non_power_of_two() {
let signal = pseudo_random(37, 11);
for fam in round_trip_families() {
let fb = filter_bank(&fam).unwrap();
let (approx, detail) =
single_level_analysis(&signal, &fb, BoundaryMode::Symmetric).unwrap();
let recon = single_level_synthesis(
&approx,
&detail,
&fb,
BoundaryMode::Symmetric,
signal.len(),
)
.unwrap();
let e = rel_err(&recon, &signal);
assert!(e < 1e-10, "{fam:?} symmetric n=37 rel err {e}");
}
}
#[test]
fn coefficient_lengths_are_ceil_half() {
let fb = filter_bank(&WaveletFamily::Daubechies(4)).unwrap();
for n in [36_usize, 37] {
let signal = pseudo_random(n, 3);
let (approx, detail) =
single_level_analysis(&signal, &fb, BoundaryMode::Periodic).unwrap();
assert_eq!(approx.len(), n.div_ceil(2));
assert_eq!(detail.len(), n.div_ceil(2));
let recon =
single_level_synthesis(&approx, &detail, &fb, BoundaryMode::Periodic, n).unwrap();
assert_eq!(recon.len(), n);
}
}
#[test]
fn reconstruction_has_no_nan_or_inf() {
for &n in &[36_usize, 37] {
let signal = pseudo_random(n, 99);
for fam in round_trip_families() {
let fb = filter_bank(&fam).unwrap();
for mode in [BoundaryMode::Periodic, BoundaryMode::Symmetric] {
let (approx, detail) = single_level_analysis(&signal, &fb, mode).unwrap();
let recon = single_level_synthesis(&approx, &detail, &fb, mode, n).unwrap();
assert!(
recon.iter().all(|x| x.is_finite()),
"{fam:?} {mode:?} n={n} produced non-finite"
);
}
}
}
}
#[test]
fn db4_even_and_odd_round_trip_both_modes() {
let fb = filter_bank(&WaveletFamily::Daubechies(4)).unwrap();
for n in [36_usize, 37] {
let signal = pseudo_random(n, 5);
for mode in [BoundaryMode::Periodic, BoundaryMode::Symmetric] {
let (approx, detail) = single_level_analysis(&signal, &fb, mode).unwrap();
let recon = single_level_synthesis(&approx, &detail, &fb, mode, n).unwrap();
let e = rel_err(&recon, &signal);
assert!(e < 1e-10, "db4 n={n} {mode:?} rel err {e}");
}
}
}
#[test]
fn odd_order_daubechies_round_trip_both_modes() {
let n = 37_usize; let signal = pseudo_random(n, 202);
for order in [3_usize, 5, 7, 9] {
let fam = WaveletFamily::from_db_order(order).unwrap();
let fb = filter_bank(&fam).unwrap();
for mode in [BoundaryMode::Periodic, BoundaryMode::Symmetric] {
let (approx, detail) = single_level_analysis(&signal, &fb, mode).unwrap();
let recon = single_level_synthesis(&approx, &detail, &fb, mode, n).unwrap();
let e = rel_err(&recon, &signal);
assert!(e < 1e-10, "db{order} n={n} {mode:?} rel err {e}");
assert!(recon.iter().all(|x| x.is_finite()));
}
}
}
#[test]
fn odd_order_daubechies_multi_level_round_trip_both_modes() {
let n = 201_usize;
let signal = pseudo_random(n, 303);
for order in [3_usize, 5, 7, 9] {
let fam = WaveletFamily::from_db_order(order).unwrap();
for mode in [BoundaryMode::Periodic, BoundaryMode::Symmetric] {
let coeffs = decompose(&signal, fam.clone(), mode, None).unwrap();
let recon = reconstruct(&coeffs).unwrap();
assert_eq!(recon.len(), n);
let e = rel_err(&recon, &signal);
assert!(
e < 1e-10,
"db{order} n={n} {mode:?} multi-level rel err {e}"
);
assert!(recon.iter().all(|x| x.is_finite()));
}
}
}
#[test]
fn synthesis_rejects_mismatched_coefficient_lengths() {
let fb = filter_bank(&WaveletFamily::Haar).unwrap();
let approx = vec![1.0, 2.0];
let detail = vec![1.0];
assert!(matches!(
single_level_synthesis(&approx, &detail, &fb, BoundaryMode::Periodic, 4),
Err(FdarError::InvalidDimension { .. })
));
}
#[test]
fn synthesis_rejects_zero_output_len() {
let fb = filter_bank(&WaveletFamily::Haar).unwrap();
assert!(matches!(
single_level_synthesis(&[], &[], &fb, BoundaryMode::Periodic, 0),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn from_db_order_maps_correctly() {
assert_eq!(
WaveletFamily::from_db_order(1).unwrap(),
WaveletFamily::Haar
);
assert_eq!(
WaveletFamily::from_db_order(2).unwrap(),
WaveletFamily::Daubechies(2)
);
assert_eq!(
WaveletFamily::from_db_order(10).unwrap(),
WaveletFamily::Daubechies(10)
);
assert!(matches!(
WaveletFamily::from_db_order(0),
Err(FdarError::InvalidParameter { .. })
));
assert!(matches!(
WaveletFamily::from_db_order(11),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn default_boundary_mode_is_periodic() {
assert_eq!(BoundaryMode::default(), BoundaryMode::Periodic);
}
#[test]
fn max_level_haar_is_log2_of_n() {
assert_eq!(max_level(1024, &WaveletFamily::Haar).unwrap(), 10);
}
#[test]
fn max_level_db4_uses_filter_len_minus_one() {
let expected = (1024.0_f64 / 7.0).log2().floor() as usize;
assert_eq!(
max_level(1024, &WaveletFamily::Daubechies(4)).unwrap(),
expected
);
assert_eq!(expected, 7);
}
#[test]
fn max_level_rejects_zero_length() {
assert!(matches!(
max_level(0, &WaveletFamily::Haar),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn max_level_rejects_too_short_signal() {
assert!(matches!(
max_level(1, &WaveletFamily::Haar),
Err(FdarError::InvalidParameter { .. })
));
assert!(matches!(
max_level(6, &WaveletFamily::Daubechies(4)),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn max_level_rejects_unsupported_family() {
assert!(matches!(
max_level(1024, &WaveletFamily::Daubechies(11)),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn wavelet_coeffs_partial_eq_on_identical_inputs() {
let signal = pseudo_random(64, 21);
let a = decompose(
&signal,
WaveletFamily::Daubechies(4),
BoundaryMode::Periodic,
Some(3),
)
.unwrap();
let b = decompose(
&signal,
WaveletFamily::Daubechies(4),
BoundaryMode::Periodic,
Some(3),
)
.unwrap();
assert_eq!(a, b);
assert_eq!(a.levels(), 3);
assert_eq!(a.signal_len(), 64);
assert_eq!(a.detail(0).unwrap().len(), a.details[0].len());
assert!(a.detail(3).is_none());
assert_eq!(a.family(), &WaveletFamily::Daubechies(4));
assert_eq!(a.mode(), BoundaryMode::Periodic);
}
#[test]
fn multi_level_round_trip_periodic_all_families() {
let signal = pseudo_random(256, 123);
for fam in round_trip_families() {
let coeffs = decompose(&signal, fam.clone(), BoundaryMode::Periodic, Some(2)).unwrap();
assert_eq!(coeffs.levels, 2);
let recon = reconstruct(&coeffs).unwrap();
assert_eq!(recon.len(), signal.len());
let e = rel_err(&recon, &signal);
assert!(e < 1e-10, "{fam:?} periodic L=2 rel err {e}");
assert!(recon.iter().all(|x| x.is_finite()));
}
}
#[test]
fn multi_level_round_trip_symmetric_non_power_of_two() {
let signal = pseudo_random(201, 456);
for fam in round_trip_families() {
let coeffs = decompose(&signal, fam.clone(), BoundaryMode::Symmetric, None).unwrap();
assert!(
coeffs.levels >= 2,
"{fam:?} expected >=2 auto levels for n=201"
);
let recon = reconstruct(&coeffs).unwrap();
assert_eq!(recon.len(), 201);
let e = rel_err(&recon, &signal);
assert!(e < 1e-10, "{fam:?} symmetric n=201 auto rel err {e}");
assert!(recon.iter().all(|x| x.is_finite()));
}
}
#[test]
fn multi_level_round_trip_periodic_non_power_of_two() {
let signal = pseudo_random(201, 789);
for fam in round_trip_families() {
let coeffs = decompose(&signal, fam.clone(), BoundaryMode::Periodic, None).unwrap();
assert!(coeffs.levels >= 2, "{fam:?} expected >=2 auto levels");
let recon = reconstruct(&coeffs).unwrap();
assert_eq!(recon.len(), 201);
let e = rel_err(&recon, &signal);
assert!(e < 1e-10, "{fam:?} periodic n=201 rel err {e}");
}
}
#[test]
fn auto_level_equals_max_level() {
let signal = pseudo_random(200, 8);
let coeffs = decompose(
&signal,
WaveletFamily::Daubechies(4),
BoundaryMode::Periodic,
None,
)
.unwrap();
assert_eq!(
coeffs.levels,
max_level(200, &WaveletFamily::Daubechies(4)).unwrap()
);
}
#[test]
fn explicit_level_out_of_range_is_invalid() {
let signal = pseudo_random(64, 8);
let maxl = max_level(64, &WaveletFamily::Daubechies(4)).unwrap();
assert!(matches!(
decompose(
&signal,
WaveletFamily::Daubechies(4),
BoundaryMode::Periodic,
Some(maxl + 1)
),
Err(FdarError::InvalidParameter { .. })
));
assert!(matches!(
decompose(
&signal,
WaveletFamily::Daubechies(4),
BoundaryMode::Periodic,
Some(0)
),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn decompose_rejects_empty_signal() {
assert!(matches!(
decompose(&[], WaveletFamily::Haar, BoundaryMode::Periodic, None),
Err(FdarError::InvalidParameter { .. })
));
}
#[test]
fn signal_len_preserved_odd_and_even() {
for &n in &[37_usize, 64] {
let signal = pseudo_random(n, n as u64);
let coeffs = decompose(
&signal,
WaveletFamily::Daubechies(2),
BoundaryMode::Periodic,
Some(2),
)
.unwrap();
let recon = reconstruct(&coeffs).unwrap();
assert_eq!(recon.len(), n, "n={n}");
}
}
#[test]
fn reconstruct_rejects_inconsistent_coeffs() {
let signal = pseudo_random(64, 3);
let mut coeffs = decompose(
&signal,
WaveletFamily::Haar,
BoundaryMode::Periodic,
Some(3),
)
.unwrap();
coeffs.details.pop();
assert!(matches!(
reconstruct(&coeffs),
Err(FdarError::InvalidDimension { .. })
));
}
#[test]
fn decompose_matrix_matches_per_row_slice_path() {
use crate::matrix::FdMatrix;
let nrows = 5;
let ncols = 48;
let mut flat = vec![0.0_f64; nrows * ncols];
for i in 0..nrows {
let row = pseudo_random(ncols, 1000 + i as u64);
for j in 0..ncols {
flat[i + j * nrows] = row[j];
}
}
let m = FdMatrix::from_column_major(flat, nrows, ncols).unwrap();
let batch = decompose_matrix(
&m,
WaveletFamily::Daubechies(4),
BoundaryMode::Periodic,
None,
)
.unwrap();
assert_eq!(batch.len(), nrows);
for i in 0..nrows {
let per_row = decompose(
&m.row(i),
WaveletFamily::Daubechies(4),
BoundaryMode::Periodic,
None,
)
.unwrap();
assert_eq!(batch[i], per_row, "row {i} batch != per-row");
}
}
#[test]
fn decompose_matrix_round_trip_both_modes() {
use crate::matrix::FdMatrix;
let nrows = 4;
let ncols = 48;
let mut rows = Vec::new();
let mut flat = vec![0.0_f64; nrows * ncols];
for i in 0..nrows {
let row = pseudo_random(ncols, 2000 + i as u64);
for j in 0..ncols {
flat[i + j * nrows] = row[j];
}
rows.push(row);
}
let m = FdMatrix::from_column_major(flat, nrows, ncols).unwrap();
for mode in [BoundaryMode::Periodic, BoundaryMode::Symmetric] {
let batch = decompose_matrix(&m, WaveletFamily::Daubechies(6), mode, None).unwrap();
for (i, coeffs) in batch.iter().enumerate() {
let recon = reconstruct(coeffs).unwrap();
assert_eq!(recon.len(), ncols);
let e = rel_err(&recon, &rows[i]);
assert!(e < 1e-10, "row {i} {mode:?} rel err {e}");
assert!(recon.iter().all(|x| x.is_finite()));
}
}
}
#[test]
fn decompose_matrix_rejects_empty_matrix() {
use crate::matrix::FdMatrix;
let zero_rows = FdMatrix::from_column_major(vec![], 0, 5).unwrap();
assert!(matches!(
decompose_matrix(
&zero_rows,
WaveletFamily::Haar,
BoundaryMode::Periodic,
None
),
Err(FdarError::InvalidDimension { .. })
));
let zero_cols = FdMatrix::from_column_major(vec![], 5, 0).unwrap();
assert!(matches!(
decompose_matrix(
&zero_cols,
WaveletFamily::Haar,
BoundaryMode::Periodic,
None
),
Err(FdarError::InvalidDimension { .. })
));
}
#[test]
fn unsupported_order_surfaces_invalid_parameter() {
assert!(matches!(
WaveletFamily::from_db_order(11),
Err(FdarError::InvalidParameter { .. })
));
let signal = pseudo_random(64, 1);
assert!(matches!(
decompose(
&signal,
WaveletFamily::Daubechies(11),
BoundaryMode::Periodic,
None
),
Err(FdarError::InvalidParameter { .. })
));
}
}