1use crate::common::gaussian_kernel;
6use crate::kde::{KdeError, KdeResult};
7use realfft::RealFftPlanner;
8use realfft::num_complex::Complex;
9
10pub fn kde_fft(data: &[f64], grid: &[f64], bandwidth: f64, n: f64) -> KdeResult<Vec<f64>> {
20 let m = grid.len();
21 if m < 2 {
22 return Err(KdeError::StatsError(
23 "Grid must have at least 2 points".to_string(),
24 ));
25 }
26
27 let grid_min = grid[0];
28 let grid_max = grid[m - 1];
29 let grid_spacing = (grid_max - grid_min) / (m - 1) as f64;
30
31 let mut binned = vec![0.0; m];
34 for &x in data {
35 let idx = ((x - grid_min) / grid_spacing).floor() as isize;
36 if idx >= 0 && (idx as usize) < m {
37 binned[idx as usize] += 1.0;
38 }
39 }
40
41 let kernel_center = (m - 1) as f64 / 2.0;
44 let mut kernel = Vec::with_capacity(m);
45 for i in 0..m {
46 let grid_pos = (i as f64 - kernel_center) * grid_spacing;
47 let u = grid_pos / bandwidth;
48 kernel.push(gaussian_kernel(u));
49 }
50
51 let fft_size = (2 * m).next_power_of_two();
54
55 let mut planner = RealFftPlanner::<f64>::new();
57 let r2c = planner.plan_fft_forward(fft_size);
58 let c2r = planner.plan_fft_inverse(fft_size);
59
60 let mut binned_padded = vec![0.0; fft_size];
62 binned_padded[..m].copy_from_slice(&binned);
63
64 let mut kernel_padded = vec![0.0; fft_size];
70 let kernel_start = (fft_size - m) / 2;
71 let first_half = (m + 1) / 2;
73 kernel_padded[kernel_start..kernel_start + first_half].copy_from_slice(&kernel[m / 2..]);
74 let second_half = m / 2;
76 if second_half > 0 {
77 kernel_padded[..second_half].copy_from_slice(&kernel[..second_half]);
78 }
79
80 let mut binned_spectrum = r2c.make_output_vec();
82 r2c.process(&mut binned_padded, &mut binned_spectrum)
83 .map_err(|e| KdeError::FftError(format!("FFT forward failed: {}", e)))?;
84
85 let mut kernel_spectrum = r2c.make_output_vec();
86 r2c.process(&mut kernel_padded, &mut kernel_spectrum)
87 .map_err(|e| KdeError::FftError(format!("FFT forward failed: {}", e)))?;
88
89 let mut conv_spectrum: Vec<Complex<f64>> = binned_spectrum
91 .iter()
92 .zip(kernel_spectrum.iter())
93 .map(|(a, b)| a * b)
94 .collect();
95
96 let mut conv_result = c2r.make_output_vec();
98 c2r.process(&mut conv_spectrum, &mut conv_result)
99 .map_err(|e| KdeError::FftError(format!("FFT inverse failed: {}", e)))?;
100
101 let kernel_start = (fft_size - m) / 2;
105 let mut density = Vec::with_capacity(m);
106 for i in 0..m {
107 let idx = (kernel_start + i) % fft_size;
108 density.push(conv_result[idx]);
109 }
110
111 let density: Vec<f64> = density
114 .iter()
115 .map(|&val| val / (fft_size as f64 * n * bandwidth))
116 .collect();
117
118 Ok(density)
119}