Skip to main content

flow_density/kde/
fft.rs

1//! FFT-based Kernel Density Estimation
2//!
3//! Uses FFT convolution for O(n log n) performance instead of O(n*m).
4
5use crate::common::gaussian_kernel;
6use crate::kde::{KdeError, KdeResult};
7use realfft::RealFftPlanner;
8use realfft::num_complex::Complex;
9
10/// FFT-based Kernel Density Estimation
11///
12/// Uses FFT convolution for O(n log n) performance instead of O(n*m).
13/// Algorithm:
14/// 1. Bin data onto grid
15/// 2. Create kernel values on grid
16/// 3. Zero-pad both to avoid circular convolution
17/// 4. FFT both, multiply in frequency domain, inverse FFT
18/// 5. Extract and normalize
19pub 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    // Step 1: Bin data onto grid
32    // Count how many data points fall into each grid bin
33    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    // Step 2: Create kernel values on grid
42    // Kernel is centered at grid center (index m/2), symmetric
43    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    // Step 3: Zero-pad to avoid circular convolution
52    // Use next power of 2 >= 2*m for efficient FFT
53    let fft_size = (2 * m).next_power_of_two();
54
55    // Step 4: FFT convolution
56    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    // Prepare padded arrays
61    let mut binned_padded = vec![0.0; fft_size];
62    binned_padded[..m].copy_from_slice(&binned);
63
64    // For linear convolution with a symmetric kernel, we need to place the kernel
65    // such that when convolved, it's centered. Since the kernel is symmetric and
66    // centered at index m/2, we place it starting at position (fft_size - m) / 2
67    // to center it in the padded array, then wrap around for circular convolution
68    // which becomes linear convolution after extraction
69    let mut kernel_padded = vec![0.0; fft_size];
70    let kernel_start = (fft_size - m) / 2;
71    // Place first half of kernel at the end
72    let first_half = (m + 1) / 2;
73    kernel_padded[kernel_start..kernel_start + first_half].copy_from_slice(&kernel[m / 2..]);
74    // Place second half at the beginning (wrapped)
75    let second_half = m / 2;
76    if second_half > 0 {
77        kernel_padded[..second_half].copy_from_slice(&kernel[..second_half]);
78    }
79
80    // Forward FFT
81    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    // Step 5: Multiply in frequency domain
90    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    // Step 6: Inverse FFT
97    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    // Step 7: Extract relevant portion and normalize
102    // With the kernel centered, the valid convolution result starts at kernel_start
103    // Extract m points starting from there (wrapping if needed)
104    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    // Normalize by fft_size (FFT doesn't normalize automatically), n, and bandwidth
112    // This matches the naive implementation: sum(kernel) / (n * bandwidth)
113    let density: Vec<f64> = density
114        .iter()
115        .map(|&val| val / (fft_size as f64 * n * bandwidth))
116        .collect();
117
118    Ok(density)
119}