use scirs2_core::ndarray::Array1;
use scirs2_core::numeric::{Float, FromPrimitive};
use std::fmt::Debug;
use super::config::EnhancedPeriodogramConfig;
use super::frequency::{calculate_simple_periodogram, create_window, WindowTypeInfo};
use crate::error::Result;
fn window_reference_characteristics(window_name: &str) -> (f64, f64, f64) {
match window_name {
"Rectangular" => (2.0, -13.3, 3.92),
"Hamming" => (4.0, -42.7, 1.78),
"Blackman" => (6.0, -58.1, 1.10),
_ => (4.0, -31.5, 1.42),
}
}
#[allow(dead_code)]
pub fn calculate_window_analysis<F>(
ts: &Array1<F>,
config: &EnhancedPeriodogramConfig,
) -> Result<WindowTypeInfo<F>>
where
F: Float + FromPrimitive,
{
let window_length = ((ts.len() as f64 * 0.25).round() as usize).max(2);
let window: Vec<F> = create_window(&config.primary_window_type, window_length)?;
let n_f = F::from(window_length).expect("Failed to convert to float");
let sum_w = window.iter().fold(F::zero(), |acc, &w| acc + w);
let sum_w2 = window.iter().fold(F::zero(), |acc, &w| acc + w * w);
let coherent_gain = sum_w / n_f;
let noise_bandwidth = if sum_w > F::zero() {
n_f * sum_w2 / (sum_w * sum_w)
} else {
F::one()
};
let equivalent_noise_bandwidth = noise_bandwidth / n_f;
let processing_gain = if noise_bandwidth > F::zero() {
F::one() / noise_bandwidth
} else {
F::zero()
};
let overlap_shift = (window_length as f64 * 0.5).round() as usize;
let overlap_correlation = if sum_w2 > F::zero() && overlap_shift < window_length {
let mut cross = F::zero();
for i in 0..(window_length - overlap_shift) {
cross = cross + window[i] * window[i + overlap_shift];
}
cross / sum_w2
} else {
F::zero()
};
let (main_lobe_bins, side_lobe_db, scalloping_db) =
window_reference_characteristics(&config.primary_window_type);
Ok(WindowTypeInfo {
window_name: config.primary_window_type.clone(),
main_lobe_width: F::from(main_lobe_bins).expect("Failed to convert to float"),
side_lobe_level: F::from(side_lobe_db).expect("Failed to convert to float"),
scalloping_loss: F::from(scalloping_db).expect("Failed to convert to float"),
processing_gain,
noise_bandwidth,
coherent_gain,
window_length,
equivalent_noise_bandwidth,
overlap_correlation,
})
}
#[allow(dead_code)]
pub fn calculate_window_effectiveness<F>(windowinfo: &WindowTypeInfo<F>) -> F
where
F: Float + FromPrimitive,
{
windowinfo.processing_gain
}
#[allow(dead_code)]
pub fn calculate_spectral_leakage<F>(windowinfo: &WindowTypeInfo<F>) -> F
where
F: Float + FromPrimitive,
{
let ten = F::from(10.0).expect("Failed to convert constant to float");
let twenty = F::from(20.0).expect("Failed to convert constant to float");
ten.powf(windowinfo.side_lobe_level / twenty)
}
#[allow(dead_code)]
pub fn calculate_periodogram_confidence_intervals<F>(
periodogram: &[F],
config: &EnhancedPeriodogramConfig,
) -> Result<Vec<(F, F)>>
where
F: Float + FromPrimitive,
{
if periodogram.is_empty() {
return Ok(Vec::new());
}
let dof = if config.enable_bartlett_method {
(2 * config.bartlett_num_segments).max(2)
} else {
2
};
let alpha = (1.0 - config.confidence_level).clamp(1e-6, 1.0 - 1e-6);
let chi2_hi = F::from(crate::causality::chi_squared_quantile(alpha / 2.0, dof))
.expect("Failed to convert to float");
let chi2_lo = F::from(crate::causality::chi_squared_quantile(
1.0 - alpha / 2.0,
dof,
))
.expect("Failed to convert to float");
let dof_f = F::from(dof).expect("Failed to convert to float");
Ok(periodogram
.iter()
.map(|&value| {
let lower = if chi2_hi > F::zero() {
dof_f * value / chi2_hi
} else {
F::zero()
};
let upper = if chi2_lo > F::zero() {
dof_f * value / chi2_lo
} else {
F::zero()
};
(lower, upper)
})
.collect())
}
#[allow(dead_code)]
pub fn calculate_peak_significance<F>(
periodogram: &[F],
_config: &EnhancedPeriodogramConfig,
) -> Result<Vec<F>>
where
F: Float + FromPrimitive,
{
if periodogram.is_empty() {
return Ok(Vec::new());
}
let n_f = F::from_usize(periodogram.len()).expect("Operation failed");
let mean_power = periodogram.iter().fold(F::zero(), |acc, &x| acc + x) / n_f;
if mean_power <= F::zero() {
return Ok(vec![F::zero(); periodogram.len()]);
}
Ok(periodogram
.iter()
.map(|&value| F::one() - (-(value / mean_power)).exp())
.collect())
}
#[allow(dead_code)]
pub fn calculate_bias_corrected_periodogram<F>(
periodogram: &[F],
config: &EnhancedPeriodogramConfig,
) -> Result<Vec<F>>
where
F: Float + FromPrimitive,
{
if periodogram.is_empty() {
return Ok(Vec::new());
}
let window_length = (periodogram.len() * 2).max(2);
let window: Vec<F> = create_window(&config.primary_window_type, window_length)?;
let n_f = F::from(window_length).expect("Failed to convert to float");
let sum_w2 = window.iter().fold(F::zero(), |acc, &w| acc + w * w);
let power_gain = sum_w2 / n_f;
if power_gain <= F::zero() {
return Ok(periodogram.to_vec());
}
Ok(periodogram
.iter()
.map(|&value| value / power_gain)
.collect())
}
#[allow(dead_code)]
pub fn calculate_variance_reduced_periodogram<F>(
periodogram: &[F],
_config: &EnhancedPeriodogramConfig,
) -> Result<Vec<F>>
where
F: Float + FromPrimitive,
{
let n = periodogram.len();
if n < 3 {
return Ok(periodogram.to_vec());
}
let epsilon = F::from(1e-300).expect("Failed to convert constant to float");
let log_values: Vec<F> = periodogram.iter().map(|&v| v.max(epsilon).ln()).collect();
let mut result = vec![F::zero(); n];
for (i, slot) in result.iter_mut().enumerate() {
let lo = i.saturating_sub(1);
let hi = (i + 1).min(n - 1);
let count = F::from_usize(hi - lo + 1).expect("Operation failed");
let sum = (lo..=hi).fold(F::zero(), |acc, j| acc + log_values[j]);
*slot = (sum / count).exp();
}
Ok(result)
}
#[allow(dead_code)]
pub fn calculate_smoothed_periodogram<F>(
periodogram: &[F],
_config: &EnhancedPeriodogramConfig,
) -> Result<Vec<F>>
where
F: Float + FromPrimitive,
{
let n = periodogram.len();
if n < 3 {
return Ok(periodogram.to_vec());
}
const WEIGHTS: [f64; 5] = [1.0, 2.0, 3.0, 2.0, 1.0];
let half = (WEIGHTS.len() / 2) as isize;
let mut result = vec![F::zero(); n];
for (i, slot) in result.iter_mut().enumerate() {
let mut weighted_sum = F::zero();
let mut weight_total = F::zero();
for (k, &w) in WEIGHTS.iter().enumerate() {
let j = i as isize + (k as isize - half);
if j >= 0 && (j as usize) < n {
let w_f = F::from(w).expect("Failed to convert constant to float");
weighted_sum = weighted_sum + w_f * periodogram[j as usize];
weight_total = weight_total + w_f;
}
}
*slot = if weight_total > F::zero() {
weighted_sum / weight_total
} else {
periodogram[i]
};
}
Ok(result)
}
#[allow(dead_code)]
pub fn calculate_zero_padded_periodogram<F>(
ts: &Array1<F>,
config: &EnhancedPeriodogramConfig,
) -> Result<Vec<F>>
where
F: Float + FromPrimitive + Debug + std::iter::Sum,
for<'a> F: std::iter::Sum<&'a F>,
{
let n = ts.len();
if n == 0 {
return Ok(Vec::new());
}
let padded_len = n * config.zero_padding_factor.max(1);
let mut padded = Array1::<F>::zeros(padded_len);
for (i, &value) in ts.iter().enumerate() {
padded[i] = value;
}
calculate_simple_periodogram(&padded)
}
#[allow(dead_code)]
pub fn calculate_interpolated_periodogram<F>(
periodogram: &[F],
_config: &EnhancedPeriodogramConfig,
) -> Result<Vec<F>>
where
F: Float + FromPrimitive,
{
let n = periodogram.len();
if n < 3 {
return Ok(periodogram.to_vec());
}
let second_derivatives = natural_cubic_spline_second_derivatives(periodogram);
let half = F::from(0.5).expect("Failed to convert constant to float");
let mut interpolated = Vec::with_capacity(2 * n - 1);
for i in 0..(n - 1) {
interpolated.push(periodogram[i]);
interpolated.push(evaluate_natural_cubic_spline(
periodogram,
&second_derivatives,
i,
half,
));
}
interpolated.push(periodogram[n - 1]);
Ok(interpolated)
}
fn natural_cubic_spline_second_derivatives<F>(y: &[F]) -> Vec<F>
where
F: Float + FromPrimitive,
{
let n = y.len();
let mut m = vec![F::zero(); n];
if n < 3 {
return m;
}
let two = F::from(2.0).expect("Failed to convert constant to float");
let four = F::from(4.0).expect("Failed to convert constant to float");
let six = F::from(6.0).expect("Failed to convert constant to float");
let mut c_prime = vec![F::zero(); n];
let mut d_prime = vec![F::zero(); n];
c_prime[1] = F::one() / four;
d_prime[1] = (y[2] - two * y[1] + y[0]) * six / four;
for i in 2..(n - 1) {
let denom = four - c_prime[i - 1];
c_prime[i] = F::one() / denom;
let rhs = (y[i + 1] - two * y[i] + y[i - 1]) * six;
d_prime[i] = (rhs - d_prime[i - 1]) / denom;
}
m[n - 2] = d_prime[n - 2];
for i in (1..(n - 2)).rev() {
m[i] = d_prime[i] - c_prime[i] * m[i + 1];
}
m
}
fn evaluate_natural_cubic_spline<F>(y: &[F], m: &[F], i: usize, t: F) -> F
where
F: Float + FromPrimitive,
{
let six = F::from(6.0).expect("Failed to convert constant to float");
let one_minus_t = F::one() - t;
let term_a = m[i] * one_minus_t * one_minus_t * one_minus_t / six;
let term_b = m[i + 1] * t * t * t / six;
let term_c = (y[i] - m[i] / six) * one_minus_t;
let term_d = (y[i + 1] - m[i + 1] / six) * t;
term_a + term_b + term_c + term_d
}
fn count_local_maxima<F>(spectrum: &[F]) -> usize
where
F: Float,
{
if spectrum.len() < 3 {
return 0;
}
(1..spectrum.len() - 1)
.filter(|&i| spectrum[i] > spectrum[i - 1] && spectrum[i] > spectrum[i + 1])
.count()
}
fn resolution_enhancement_effectiveness<F>(enhanced: &[F], original: &[F]) -> F
where
F: Float + FromPrimitive,
{
let enhanced_peaks = count_local_maxima(enhanced);
let original_peaks = count_local_maxima(original);
if enhanced_peaks == 0 {
return F::zero();
}
let additional = enhanced_peaks.saturating_sub(original_peaks);
let ratio = F::from_usize(additional).expect("Operation failed")
/ F::from_usize(enhanced_peaks).expect("Operation failed");
ratio.min(F::one()).max(F::zero())
}
#[allow(dead_code)]
pub fn calculate_zero_padding_effectiveness<F>(padded: &[F], original: &[F]) -> F
where
F: Float + FromPrimitive,
{
resolution_enhancement_effectiveness(padded, original)
}
#[allow(dead_code)]
pub fn calculate_interpolation_effectiveness<F>(interpolated: &[F], original: &[F]) -> F
where
F: Float + FromPrimitive,
{
resolution_enhancement_effectiveness(interpolated, original)
}
#[cfg(test)]
mod tests {
use super::*;
fn config_with_window(window: &str) -> EnhancedPeriodogramConfig {
EnhancedPeriodogramConfig {
primary_window_type: window.to_string(),
..Default::default()
}
}
#[test]
fn test_window_effectiveness_and_leakage_differ_by_window_type() {
let ts = Array1::<f64>::zeros(100);
let mut effectiveness = std::collections::HashMap::new();
let mut leakage = std::collections::HashMap::new();
for window in ["Rectangular", "Hamming", "Hanning", "Blackman"] {
let config = config_with_window(window);
let info = calculate_window_analysis(&ts, &config)
.unwrap_or_else(|_| panic!("window analysis should succeed for {window}"));
effectiveness.insert(window, calculate_window_effectiveness(&info));
leakage.insert(window, calculate_spectral_leakage(&info));
}
for window in ["Rectangular", "Hamming", "Hanning", "Blackman"] {
assert!(
(effectiveness[window] - 0.8).abs() > 1e-6,
"{window}: effectiveness should not be the old hardcoded 0.8"
);
assert!(
(leakage[window] - 0.1).abs() > 1e-6,
"{window}: leakage should not be the old hardcoded 0.1"
);
}
assert!((effectiveness["Rectangular"] - 1.0).abs() < 1e-9);
assert!(effectiveness["Rectangular"] > effectiveness["Hamming"]);
assert!(effectiveness["Hamming"] > effectiveness["Hanning"]);
assert!(effectiveness["Hanning"] > effectiveness["Blackman"]);
assert!(leakage["Rectangular"] > leakage["Hanning"]);
assert!(leakage["Hanning"] > leakage["Hamming"]);
assert!(leakage["Hamming"] > leakage["Blackman"]);
}
#[test]
fn test_periodogram_confidence_intervals_bracket_the_estimate() {
let periodogram = vec![1.0, 4.0, 9.0, 2.0, 6.0, 3.0];
let config = EnhancedPeriodogramConfig {
confidence_level: 0.95,
enable_bartlett_method: false,
..Default::default()
};
let intervals = calculate_periodogram_confidence_intervals(&periodogram, &config)
.expect("confidence intervals should succeed");
assert_eq!(intervals.len(), periodogram.len());
for (i, &(lower, upper)) in intervals.iter().enumerate() {
assert!(
lower <= periodogram[i] && periodogram[i] <= upper,
"bin {i}: point estimate {} should lie within [{lower}, {upper}]",
periodogram[i]
);
assert!(lower >= 0.0, "bin {i}: lower bound should be non-negative");
}
let averaged_config = EnhancedPeriodogramConfig {
confidence_level: 0.95,
enable_bartlett_method: true,
bartlett_num_segments: 16,
..Default::default()
};
let averaged_intervals =
calculate_periodogram_confidence_intervals(&periodogram, &averaged_config)
.expect("confidence intervals should succeed");
let raw_width: f64 = intervals.iter().map(|&(lo, hi)| hi - lo).sum();
let averaged_width: f64 = averaged_intervals.iter().map(|&(lo, hi)| hi - lo).sum();
assert!(
averaged_width < raw_width,
"averaging over more segments should narrow the confidence interval: \
raw={raw_width}, averaged={averaged_width}"
);
}
#[test]
fn test_peak_significance_detects_injected_peak() {
let mut periodogram = vec![1.0_f64; 20];
periodogram[10] = 200.0;
let config = EnhancedPeriodogramConfig::default();
let significance = calculate_peak_significance(&periodogram, &config)
.expect("peak significance should succeed");
assert_eq!(significance.len(), periodogram.len());
assert!(
significance[10] > 0.999,
"the injected peak should be judged highly significant, got {}",
significance[10]
);
assert!(
significance[0] < significance[10],
"a noise-floor bin should be far less significant than the peak"
);
for &s in &significance {
assert!(
(0.0..=1.0).contains(&s),
"significance must lie in [0, 1], got {s}"
);
}
}
#[test]
fn test_bias_corrected_periodogram_rescales_by_window_power_gain() {
let periodogram = vec![2.0, 5.0, 3.0, 8.0, 1.0];
let config = config_with_window("Hamming");
let corrected = calculate_bias_corrected_periodogram(&periodogram, &config)
.expect("bias correction should succeed");
assert_eq!(corrected.len(), periodogram.len());
assert_ne!(
corrected, periodogram,
"bias correction must actually transform the periodogram"
);
let window_length = periodogram.len() * 2;
let window: Vec<f64> = create_window("Hamming", window_length).expect("window");
let sum_w2: f64 = window.iter().map(|w| w * w).sum();
let power_gain = sum_w2 / window_length as f64;
for (i, &value) in periodogram.iter().enumerate() {
let expected = value / power_gain;
assert!(
(corrected[i] - expected).abs() < 1e-10,
"bin {i}: expected {expected}, got {}",
corrected[i]
);
}
}
#[test]
fn test_variance_reduced_periodogram_reduces_log_domain_variance() {
let periodogram = vec![1.0, 50.0, 2.0, 40.0, 3.0, 60.0, 1.5, 45.0, 2.5, 55.0];
let config = EnhancedPeriodogramConfig::default();
let reduced = calculate_variance_reduced_periodogram(&periodogram, &config)
.expect("variance reduction should succeed");
assert_eq!(reduced.len(), periodogram.len());
assert_ne!(
reduced, periodogram,
"variance reduction must actually transform the periodogram"
);
let log_variance = |values: &[f64]| -> f64 {
let logs: Vec<f64> = values.iter().map(|v| v.ln()).collect();
let mean = logs.iter().sum::<f64>() / logs.len() as f64;
logs.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / logs.len() as f64
};
let original_variance = log_variance(&periodogram);
let reduced_variance = log_variance(&reduced);
assert!(
reduced_variance < original_variance,
"log-domain variance should decrease: original={original_variance}, reduced={reduced_variance}"
);
}
#[test]
fn test_smoothed_periodogram_matches_hand_computed_weights() {
let periodogram = vec![1.0, 2.0, 10.0, 3.0, 1.0, 2.0, 8.0];
let config = EnhancedPeriodogramConfig::default();
let smoothed = calculate_smoothed_periodogram(&periodogram, &config)
.expect("smoothing should succeed");
assert_eq!(smoothed.len(), periodogram.len());
assert_ne!(
smoothed, periodogram,
"smoothing must actually transform the periodogram"
);
let expected_interior = (1.0 * 1.0 + 2.0 * 2.0 + 10.0 * 3.0 + 3.0 * 2.0 + 1.0 * 1.0) / 9.0;
assert!(
(smoothed[2] - expected_interior).abs() < 1e-10,
"expected {expected_interior}, got {}",
smoothed[2]
);
let expected_edge = (1.0 * 3.0 + 2.0 * 2.0 + 10.0 * 1.0) / 6.0;
assert!(
(smoothed[0] - expected_edge).abs() < 1e-10,
"expected {expected_edge}, got {}",
smoothed[0]
);
}
#[test]
fn test_zero_padded_periodogram_increases_resolution() {
let n = 16;
let ts = Array1::from_shape_fn(n, |i| (i as f64 * 0.3).sin());
let unpadded = calculate_simple_periodogram(&ts).expect("periodogram should succeed");
let config = EnhancedPeriodogramConfig {
zero_padding_factor: 4,
..Default::default()
};
let padded =
calculate_zero_padded_periodogram(&ts, &config).expect("zero padding should succeed");
assert_eq!(padded.len(), (n * 4) / 2);
assert!(
padded.len() > unpadded.len(),
"zero padding should genuinely increase spectral resolution: {} vs {}",
padded.len(),
unpadded.len()
);
}
#[test]
fn test_interpolated_periodogram_doubles_resolution_and_preserves_knots() {
let periodogram = vec![1.0, 3.0, 2.0, 5.0, 4.0, 6.0, 2.5, 7.0];
let config = EnhancedPeriodogramConfig::default();
let interpolated = calculate_interpolated_periodogram(&periodogram, &config)
.expect("interpolation should succeed");
assert_eq!(interpolated.len(), 2 * periodogram.len() - 1);
for (i, &value) in periodogram.iter().enumerate() {
assert!((interpolated[2 * i] - value).abs() < 1e-10);
}
let expected_midpoints = [
2.447_655, 2.282_034, 3.549_210, 4.521_127, 5.241_283, 4.201_241, 3.766_253,
];
for (i, &expected) in expected_midpoints.iter().enumerate() {
let got = interpolated[2 * i + 1];
assert!(
(got - expected).abs() < 1e-5,
"segment {i}: expected {expected}, got {got}"
);
}
let linear_midpoint_0 = (periodogram[0] + periodogram[1]) / 2.0;
assert!(
(interpolated[1] - linear_midpoint_0).abs() > 0.05,
"cubic-spline midpoint should differ from the naive linear average"
);
}
#[test]
fn test_resolution_enhancement_effectiveness_reflects_new_peaks() {
let monotonic_original = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let monotonic_enhanced = vec![1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0, 4.5, 5.0];
let zero_effectiveness =
calculate_zero_padding_effectiveness(&monotonic_enhanced, &monotonic_original);
assert_eq!(zero_effectiveness, 0.0);
let coarse_original = vec![1.0, 5.0, 1.0];
let fine_enhanced = vec![1.0, 4.0, 5.0, 4.0, 3.0, 4.0, 5.0, 4.0, 1.0];
let positive_effectiveness =
calculate_interpolation_effectiveness(&fine_enhanced, &coarse_original);
assert!(
positive_effectiveness > 0.0,
"revealing new peaks should score above zero, got {positive_effectiveness}"
);
assert!((zero_effectiveness - 0.9).abs() > 1e-6);
assert!((positive_effectiveness - 0.85).abs() > 1e-6);
}
}