use crate::api::{Direction, Flags, Plan};
use crate::kernel::{Complex, Float};
use crate::prelude::*;
pub mod nufft2d;
pub mod nufft3d;
pub use nufft2d::{nufft2d_type1, nufft2d_type2};
pub use nufft3d::nufft3d_type1;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum NufftType {
Type1,
Type2,
Type3,
}
#[derive(Debug, Clone, Copy)]
pub struct NufftOptions {
pub oversampling: f64,
pub kernel_width: usize,
pub tolerance: f64,
pub threaded: bool,
}
impl Default for NufftOptions {
fn default() -> Self {
Self {
oversampling: 2.0,
kernel_width: 6,
tolerance: 1e-6,
threaded: true,
}
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum NufftError {
InvalidSize(usize),
PointsOutOfRange,
PlanFailed,
ExecutionFailed(String),
InvalidTolerance,
}
impl core::fmt::Display for NufftError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::InvalidSize(n) => write!(f, "Invalid NUFFT size: {n}"),
Self::PointsOutOfRange => write!(f, "Non-uniform points must be in [-π, π]"),
Self::PlanFailed => write!(f, "Failed to create FFT plan"),
Self::ExecutionFailed(msg) => write!(f, "NUFFT execution failed: {msg}"),
Self::InvalidTolerance => write!(f, "Tolerance must be positive"),
}
}
}
pub type NufftResult<T> = Result<T, NufftError>;
#[allow(clippy::struct_field_names)] pub struct Nufft<T: Float> {
nufft_type: NufftType,
n_uniform: usize,
n_nonuniform: usize,
n_oversampled: usize,
points: Vec<f64>,
spread_coeffs: Vec<Vec<(usize, T)>>,
deconv_factors: Vec<Complex<T>>,
fft_plan: Option<Plan<T>>,
inv_fft_plan: Option<Plan<T>>,
flags: Flags,
options: NufftOptions,
}
impl<T: Float> Nufft<T> {
pub fn new(
nufft_type: NufftType,
n_uniform: usize,
points: &[f64],
tolerance: f64,
) -> NufftResult<Self> {
let options = NufftOptions {
tolerance,
..Default::default()
};
Self::with_options(nufft_type, n_uniform, points, &options)
}
pub fn with_options(
nufft_type: NufftType,
n_uniform: usize,
points: &[f64],
options: &NufftOptions,
) -> NufftResult<Self> {
Self::with_options_and_flags(nufft_type, n_uniform, points, options, Flags::ESTIMATE)
}
pub fn with_flags(
nufft_type: NufftType,
n_uniform: usize,
points: &[f64],
tolerance: f64,
flags: Flags,
) -> NufftResult<Self> {
let options = NufftOptions {
tolerance,
..Default::default()
};
Self::with_options_and_flags(nufft_type, n_uniform, points, &options, flags)
}
pub fn with_options_and_flags(
nufft_type: NufftType,
n_uniform: usize,
points: &[f64],
options: &NufftOptions,
flags: Flags,
) -> NufftResult<Self> {
if n_uniform == 0 {
return Err(NufftError::InvalidSize(0));
}
if options.tolerance <= 0.0 {
return Err(NufftError::InvalidTolerance);
}
let kernel_width = compute_kernel_width(
options.tolerance,
options.oversampling,
options.kernel_width,
);
let n_oversampled = ((n_uniform as f64) * options.oversampling).ceil() as usize;
let n_oversampled = next_smooth_number(n_oversampled);
let mut normalized_points = Vec::with_capacity(points.len());
for &p in points {
if !(-core::f64::consts::PI..=core::f64::consts::PI).contains(&p) {
return Err(NufftError::PointsOutOfRange);
}
normalized_points.push(p + core::f64::consts::PI);
}
let spread_coeffs =
precompute_spreading_coeffs(&normalized_points, n_oversampled, kernel_width);
let deconv_factors = precompute_deconv_factors(n_uniform, n_oversampled, kernel_width);
let fft_plan = Plan::dft_1d(n_oversampled, Direction::Forward, flags);
let inv_fft_plan = Plan::dft_1d(n_oversampled, Direction::Backward, flags);
Ok(Self {
nufft_type,
n_uniform,
n_nonuniform: points.len(),
n_oversampled,
points: normalized_points,
spread_coeffs,
deconv_factors,
fft_plan,
inv_fft_plan,
flags,
options: NufftOptions {
kernel_width,
..*options
},
})
}
pub fn type1(&self, values: &[Complex<T>]) -> NufftResult<Vec<Complex<T>>> {
if values.len() != self.n_nonuniform {
return Err(NufftError::ExecutionFailed(format!(
"Expected {} values, got {}",
self.n_nonuniform,
values.len()
)));
}
let mut grid = vec![Complex::<T>::zero(); self.n_oversampled];
self.spread_to_grid(values, &mut grid);
let mut fft_result = vec![Complex::<T>::zero(); self.n_oversampled];
if let Some(ref plan) = self.fft_plan {
plan.execute(&grid, &mut fft_result);
} else {
return Err(NufftError::PlanFailed);
}
let mut result = Vec::with_capacity(self.n_uniform);
for k in 0..self.n_uniform {
let (grid_idx, deconv_idx) =
centered_freq_indices(k, self.n_uniform, self.n_oversampled);
result.push(fft_result[grid_idx] * self.deconv_factors[deconv_idx]);
}
Ok(result)
}
pub fn type2(&self, coeffs: &[Complex<T>]) -> NufftResult<Vec<Complex<T>>> {
if coeffs.len() != self.n_uniform {
return Err(NufftError::ExecutionFailed(format!(
"Expected {} coefficients, got {}",
self.n_uniform,
coeffs.len()
)));
}
let mut grid = vec![Complex::<T>::zero(); self.n_oversampled];
let n_os_scale = T::from_usize(self.n_oversampled);
for (k, &coeff) in coeffs.iter().enumerate() {
let (grid_idx, deconv_idx) =
centered_freq_indices(k, self.n_uniform, self.n_oversampled);
let scaled_deconv = Complex::new(
self.deconv_factors[deconv_idx].re * n_os_scale,
self.deconv_factors[deconv_idx].im * n_os_scale,
);
grid[grid_idx] = coeff * scaled_deconv;
}
let mut ifft_result = vec![Complex::<T>::zero(); self.n_oversampled];
if let Some(ref inv_plan) = self.inv_fft_plan {
inv_plan.execute(&grid, &mut ifft_result);
} else {
return Err(NufftError::PlanFailed);
}
let scale = T::ONE / T::from_usize(self.n_oversampled);
for c in &mut ifft_result {
*c = Complex::new(c.re * scale, c.im * scale);
}
let result = self.interpolate_from_grid(&ifft_result);
Ok(result)
}
pub fn execute(&self, input: &[Complex<T>]) -> NufftResult<Vec<Complex<T>>> {
match self.nufft_type {
NufftType::Type1 => self.type1(input),
NufftType::Type2 => self.type2(input),
NufftType::Type3 => {
Err(NufftError::ExecutionFailed(
"Type 3 requires separate execute_type3 call".into(),
))
}
}
}
pub fn execute_type3(
&self,
values: &[Complex<T>],
target_points: &[f64],
) -> NufftResult<Vec<Complex<T>>> {
let uniform_coeffs = self.type1(values)?;
let type2_plan = Self::with_options_and_flags(
NufftType::Type2,
self.n_uniform,
target_points,
&self.options,
self.flags,
)?;
type2_plan.type2(&uniform_coeffs)
}
fn spread_to_grid(&self, values: &[Complex<T>], grid: &mut [Complex<T>]) {
for (j, &val) in values.iter().enumerate() {
for &(idx, weight) in &self.spread_coeffs[j] {
grid[idx] = grid[idx] + Complex::new(val.re * weight, val.im * weight);
}
}
}
fn interpolate_from_grid(&self, grid: &[Complex<T>]) -> Vec<Complex<T>> {
let mut result = Vec::with_capacity(self.n_nonuniform);
for j in 0..self.n_nonuniform {
let mut sum = Complex::<T>::zero();
for &(idx, weight) in &self.spread_coeffs[j] {
sum = sum + Complex::new(grid[idx].re * weight, grid[idx].im * weight);
}
result.push(sum);
}
result
}
pub fn n_uniform(&self) -> usize {
self.n_uniform
}
pub fn n_nonuniform(&self) -> usize {
self.n_nonuniform
}
pub fn nufft_type(&self) -> NufftType {
self.nufft_type
}
pub fn flags(&self) -> Flags {
self.flags
}
pub fn points(&self) -> &[f64] {
&self.points
}
}
pub(crate) fn compute_kernel_width(tolerance: f64, oversampling: f64, default: usize) -> usize {
let sigma = oversampling.clamp(1.05, 4.0); let f_sigma = 2.0 - sigma / 2.0;
let w = ((-tolerance.log10()) * f_sigma).ceil() as usize;
(2 * w).max(4).max(default)
}
pub(crate) fn centered_freq_indices(k: usize, n: usize, n_oversampled: usize) -> (usize, usize) {
let half_n = n / 2;
let freq = (k as isize) - (half_n as isize);
let grid_idx = if freq >= 0 {
freq as usize
} else {
(n_oversampled as isize + freq) as usize
};
let deconv_idx = if freq >= 0 {
freq as usize
} else {
(n as isize + freq) as usize
};
(grid_idx, deconv_idx)
}
pub(crate) fn next_smooth_number(n: usize) -> usize {
let mut candidate = n;
loop {
let mut temp = candidate;
while temp.is_multiple_of(2) {
temp /= 2;
}
while temp.is_multiple_of(3) {
temp /= 3;
}
while temp.is_multiple_of(5) {
temp /= 5;
}
if temp == 1 {
return candidate;
}
candidate += 1;
}
}
pub(crate) fn precompute_spreading_coeffs<T: Float>(
points: &[f64],
n_grid: usize,
kernel_width: usize,
) -> Vec<Vec<(usize, T)>> {
let grid_spacing = 2.0 * core::f64::consts::PI / (n_grid as f64);
let half_width = kernel_width / 2;
let beta = 2.3 * (half_width as f64);
points
.iter()
.map(|&x| {
let grid_pos = x / grid_spacing;
let center = grid_pos.round() as isize;
let mut coeffs = Vec::with_capacity(kernel_width);
for offset in -(half_width as isize)..=(half_width as isize) {
let grid_idx = (center + offset).rem_euclid(n_grid as isize) as usize;
let grid_x = (grid_idx as f64) * grid_spacing;
let mut dx = x - grid_x;
if dx > core::f64::consts::PI {
dx -= 2.0 * core::f64::consts::PI;
} else if dx < -core::f64::consts::PI {
dx += 2.0 * core::f64::consts::PI;
}
let normalized_dx = dx / (grid_spacing * (half_width as f64));
let weight = (-beta * normalized_dx * normalized_dx).exp();
if weight > 1e-15 {
coeffs.push((grid_idx, T::from_f64(weight)));
}
}
coeffs
})
.collect()
}
pub(crate) fn precompute_deconv_factors<T: Float>(
n_uniform: usize,
n_oversampled: usize,
kernel_width: usize,
) -> Vec<Complex<T>> {
let w_int = (kernel_width / 2) as isize; let beta = 2.3 * (w_int as f64);
let two_pi_over_nos = 2.0 * core::f64::consts::PI / (n_oversampled as f64);
(0..n_uniform)
.map(|d| {
let (freq, grid_bin) = if d < n_uniform / 2 {
(d as isize, d)
} else {
let f = (d as isize) - (n_uniform as isize);
(f, n_oversampled + d - n_uniform)
};
let kernel_dft: f64 = (-w_int..=w_int)
.map(|j| {
let w = (-beta * ((j * j) as f64 / (w_int * w_int) as f64)).exp();
let angle = two_pi_over_nos * (freq * j) as f64;
w * angle.cos()
})
.sum();
let phase_sign = if grid_bin % 2 == 0 { 1.0_f64 } else { -1.0_f64 };
let deconv = if kernel_dft.abs() > f64::EPSILON {
phase_sign / kernel_dft
} else {
0.0_f64
};
Complex::new(T::from_f64(deconv), T::ZERO)
})
.collect()
}
pub fn nufft_type1<T: Float>(
points: &[f64],
values: &[Complex<T>],
n_output: usize,
tolerance: f64,
) -> NufftResult<Vec<Complex<T>>> {
let plan = Nufft::new(NufftType::Type1, n_output, points, tolerance)?;
plan.type1(values)
}
pub fn nufft_type2<T: Float>(
coeffs: &[Complex<T>],
points: &[f64],
tolerance: f64,
) -> NufftResult<Vec<Complex<T>>> {
let plan = Nufft::new(NufftType::Type2, coeffs.len(), points, tolerance)?;
plan.type2(coeffs)
}
pub fn nufft_type3<T: Float>(
source_points: &[f64],
values: &[Complex<T>],
target_points: &[f64],
tolerance: f64,
) -> NufftResult<Vec<Complex<T>>> {
let n_uniform = (source_points.len() + target_points.len()).next_power_of_two();
let plan = Nufft::new(NufftType::Type1, n_uniform, source_points, tolerance)?;
plan.execute_type3(values, target_points)
}
#[cfg(test)]
mod tests {
use super::*;
fn dense_ndft_type1(points: &[f64], values: &[Complex<f64>], n: usize) -> Vec<Complex<f64>> {
let half = (n / 2) as isize;
(0..n)
.map(|k| {
let freq = (k as isize - half) as f64;
values
.iter()
.zip(points.iter())
.fold(Complex::new(0.0, 0.0), |acc, (&val, &xj)| {
let angle = -freq * xj;
acc + val * Complex::new(angle.cos(), angle.sin())
})
})
.collect()
}
fn dense_ndft_type2(fhat: &[Complex<f64>], points: &[f64]) -> Vec<Complex<f64>> {
let n = fhat.len();
let half = (n / 2) as isize;
points
.iter()
.map(|&xj| {
(0..n).fold(Complex::new(0.0, 0.0), |acc, k| {
let freq = (k as isize - half) as f64;
let angle = freq * xj;
acc + fhat[k] * Complex::new(angle.cos(), angle.sin())
})
})
.collect()
}
fn max_relative_error(nufft_out: &[Complex<f64>], reference: &[Complex<f64>]) -> f64 {
let ref_max = reference.iter().map(|c| c.norm()).fold(0.0_f64, f64::max);
if ref_max < 1e-30 {
return 0.0;
}
nufft_out
.iter()
.zip(reference.iter())
.map(|(n, r)| (*n - *r).norm() / ref_max)
.fold(0.0_f64, f64::max)
}
const HEADROOM: f64 = 1e-5;
#[test]
fn test_nufft_type1_uniform_points_matches_dense_ndft() {
let n = 8;
let points: Vec<f64> = (0..n)
.map(|k| -core::f64::consts::PI + (k as f64) * 2.0 * core::f64::consts::PI / (n as f64))
.collect();
let values: Vec<Complex<f64>> = (0..n)
.map(|k| Complex::new((k as f64).cos(), (k as f64).sin()))
.collect();
let result = nufft_type1(&points, &values, n, 1e-6).expect("NUFFT Type1 failed");
assert_eq!(result.len(), n);
let reference = dense_ndft_type1(&points, &values, n);
let rel_err = max_relative_error(&result, &reference);
assert!(
rel_err <= HEADROOM,
"uniform-points Type1 rel_err {rel_err:.2e} exceeds headroom {HEADROOM:.2e}"
);
}
#[test]
fn test_nufft_type2_single_frequency_matches_dense_ndft() {
let n = 16;
let mut coeffs = vec![Complex::<f64>::zero(); n];
coeffs[1] = Complex::new(1.0, 0.0);
let points: Vec<f64> = (0..5)
.map(|k| -core::f64::consts::PI + f64::from(k) * 0.5)
.collect();
let result = nufft_type2(&coeffs, &points, 1e-6).expect("NUFFT Type2 failed");
assert_eq!(result.len(), 5);
let reference = dense_ndft_type2(&coeffs, &points);
let rel_err = max_relative_error(&result, &reference);
assert!(
rel_err <= HEADROOM,
"single-frequency Type2 rel_err {rel_err:.2e} exceeds headroom {HEADROOM:.2e}"
);
}
#[test]
fn test_nufft_type1_and_type2_independently_match_dense_ndft() {
let n = 32;
let points: Vec<f64> = (0..10).map(|k| -2.5 + f64::from(k) * 0.5).collect();
let values: Vec<Complex<f64>> = points
.iter()
.map(|&x| Complex::new(x.cos(), x.sin()))
.collect();
let type1_out = nufft_type1(&points, &values, n, 1e-6).expect("Type1 failed");
let type1_ref = dense_ndft_type1(&points, &values, n);
let type1_rel_err = max_relative_error(&type1_out, &type1_ref);
assert!(
type1_rel_err <= HEADROOM,
"Type1 rel_err {type1_rel_err:.2e} exceeds headroom {HEADROOM:.2e}"
);
let type2_out = nufft_type2(&type1_ref, &points, 1e-6).expect("Type2 failed");
let type2_ref = dense_ndft_type2(&type1_ref, &points);
let type2_rel_err = max_relative_error(&type2_out, &type2_ref);
assert!(
type2_rel_err <= HEADROOM,
"Type2 rel_err {type2_rel_err:.2e} exceeds headroom {HEADROOM:.2e}"
);
assert_eq!(type2_out.len(), values.len());
}
#[test]
fn test_nufft_type3_matches_composed_dense_reference() {
let n_uniform = 32;
let source_points: Vec<f64> = (0..9).map(|k| -2.8 + f64::from(k) * 0.6).collect();
let target_points: Vec<f64> = (0..6).map(|k| -1.9 + f64::from(k) * 0.55).collect();
let values: Vec<Complex<f64>> = source_points
.iter()
.map(|&x| Complex::new((x * 0.5).cos(), (x * 0.5).sin()))
.collect();
let plan = Nufft::new(NufftType::Type1, n_uniform, &source_points, 1e-6)
.expect("Type3 plan failed");
let result = plan
.execute_type3(&values, &target_points)
.expect("Type3 execute failed");
assert_eq!(result.len(), target_points.len());
let intermediate_ref = dense_ndft_type1(&source_points, &values, n_uniform);
let composed_ref = dense_ndft_type2(&intermediate_ref, &target_points);
let rel_err = max_relative_error(&result, &composed_ref);
assert!(
rel_err <= HEADROOM,
"Type3 rel_err {rel_err:.2e} exceeds headroom {HEADROOM:.2e}"
);
}
#[test]
fn test_nufft_type1_coincident_points_linearity() {
let n = 16;
let x_single = vec![0.4f64];
let c_combined = vec![Complex::new(1.7, -0.5)];
let x_dup = vec![0.4f64, 0.4f64];
let c_dup = vec![Complex::new(1.0, -0.2), Complex::new(0.7, -0.3)];
let result_single = nufft_type1(&x_single, &c_combined, n, 1e-6).expect("single failed");
let result_dup = nufft_type1(&x_dup, &c_dup, n, 1e-6).expect("dup failed");
for (a, b) in result_single.iter().zip(result_dup.iter()) {
assert!(
(*a - *b).norm() < 1e-9,
"coincident-point linearity violated: {a:?} vs {b:?}"
);
}
}
#[test]
fn test_nufft_type1_edge_points_matches_dense_ndft() {
let pi = core::f64::consts::PI;
let points = vec![-pi, pi, 0.0];
let values = vec![
Complex::new(1.0, 0.0),
Complex::new(0.5, -0.5),
Complex::new(-0.3, 0.2),
];
let n = 16;
let result = nufft_type1(&points, &values, n, 1e-6).expect("edge-point Type1 failed");
let reference = dense_ndft_type1(&points, &values, n);
let rel_err = max_relative_error(&result, &reference);
assert!(
rel_err <= HEADROOM,
"edge-points rel_err {rel_err:.2e} exceeds headroom {HEADROOM:.2e}"
);
}
#[test]
fn test_nufft_type1_empty_input_returns_zero_grid() {
let points: Vec<f64> = vec![];
let values: Vec<Complex<f64>> = vec![];
let n = 8;
let result = nufft_type1(&points, &values, n, 1e-6).expect("empty Type1 failed");
assert_eq!(result.len(), n);
for v in &result {
assert_eq!(v.re, 0.0);
assert_eq!(v.im, 0.0);
}
}
#[test]
fn test_nufft_type2_empty_points_returns_empty() {
let n = 8;
let coeffs = vec![Complex::<f64>::zero(); n];
let points: Vec<f64> = vec![];
let result = nufft_type2(&coeffs, &points, 1e-6).expect("empty Type2 failed");
assert!(result.is_empty());
}
#[test]
fn test_nufft_error_handling() {
let points = vec![0.0, 0.5, 1.0];
let result = Nufft::<f64>::new(NufftType::Type1, 0, &points, 1e-6);
assert!(result.is_err());
let bad_points = vec![0.0, 5.0]; let result = Nufft::<f64>::new(NufftType::Type1, 16, &bad_points, 1e-6);
assert!(result.is_err());
let result = Nufft::<f64>::new(NufftType::Type1, 16, &points, -1e-6);
assert!(result.is_err());
}
#[test]
fn test_smooth_number() {
assert_eq!(next_smooth_number(100), 100); assert_eq!(next_smooth_number(101), 108); assert_eq!(next_smooth_number(7), 8); }
#[test]
fn test_nufft_new_defaults_to_estimate_flag() {
let points = vec![0.0, 0.3, -0.5];
let plan = Nufft::<f64>::new(NufftType::Type1, 16, &points, 1e-6).expect("plan failed");
assert_eq!(plan.flags(), Flags::ESTIMATE);
}
#[test]
fn test_nufft_with_flags_estimate_and_measure_match() {
let n = 16;
let points: Vec<f64> = (0..6).map(|k| -2.0 + f64::from(k) * 0.6).collect();
let values: Vec<Complex<f64>> = points
.iter()
.map(|&x| Complex::new(x.cos(), x.sin()))
.collect();
let plan_estimate =
Nufft::<f64>::with_flags(NufftType::Type1, n, &points, 1e-6, Flags::ESTIMATE)
.expect("ESTIMATE plan failed");
let plan_measure =
Nufft::<f64>::with_flags(NufftType::Type1, n, &points, 1e-6, Flags::MEASURE)
.expect("MEASURE plan failed");
assert_eq!(plan_estimate.flags(), Flags::ESTIMATE);
assert_eq!(plan_measure.flags(), Flags::MEASURE);
let out_estimate = plan_estimate
.type1(&values)
.expect("ESTIMATE execute failed");
let out_measure = plan_measure.type1(&values).expect("MEASURE execute failed");
for (a, b) in out_estimate.iter().zip(out_measure.iter()) {
assert!((*a - *b).norm() < 1e-9);
}
}
#[test]
fn test_nufft_type2_forward_and_inverse_plans_share_flags() {
let n = 16;
let coeffs: Vec<Complex<f64>> = (0..n)
.map(|k| Complex::new(((k as f64) * 0.1).cos(), 0.0))
.collect();
let points: Vec<f64> = (0..5).map(|k| -2.0 + f64::from(k) * 0.7).collect();
let plan = Nufft::<f64>::with_flags(NufftType::Type2, n, &points, 1e-6, Flags::MEASURE)
.expect("plan failed");
assert_eq!(plan.flags(), Flags::MEASURE);
let out1 = plan.type2(&coeffs).expect("first type2 call failed");
let out2 = plan.type2(&coeffs).expect("second type2 call failed");
assert_eq!(out1.len(), out2.len());
for (a, b) in out1.iter().zip(out2.iter()) {
assert_eq!(a.re, b.re);
assert_eq!(a.im, b.im);
}
}
}