tenferro-fft 0.3.0

FFT extension runtime and public concrete/traced FFT APIs for tenferro.
use num_complex::Complex64;
use std::num::NonZeroUsize;
use tenferro_cpu::{with_cpu_exec_session, CpuBackend, CpuExecSession};
use tenferro_runtime::ExtensionCacheLimits;
use tenferro_tensor::BackendSessionHost;
use tenferro_tensor::{ErrorKind, Tensor, TensorRead, TensorView, TypedTensorView};

use crate::{FftExecutor, FftNorm, FftPlanCache, TensorFftExt, TensorReadFftExt};

fn with_cpu_session<R>(
    backend: &mut CpuBackend,
    f: impl for<'a> FnOnce(&'a mut CpuExecSession<'a>) -> R + Send,
) -> R
where
    R: Send,
{
    backend.with_backend_session(|session| {
        with_cpu_exec_session(session, f).expect("CpuBackend must expose a CPU execution session")
    })
}

fn assert_complex_close(actual: &[Complex64], expected: &[Complex64]) {
    assert_eq!(actual.len(), expected.len());
    for (actual, expected) in actual.iter().zip(expected) {
        assert!(
            (*actual - *expected).norm() < 1.0e-12,
            "actual {actual:?} expected {expected:?}"
        );
    }
}

fn assert_real_close(actual: &[f64], expected: &[f64]) {
    assert_eq!(actual.len(), expected.len());
    for (actual, expected) in actual.iter().zip(expected) {
        assert!(
            (*actual - *expected).abs() < 1.0e-12,
            "actual {actual:?} expected {expected:?}"
        );
    }
}

#[test]
fn public_tensor_fft_ext_executes_real_and_complex_transforms() {
    let mut backend = CpuBackend::new();
    let real = Tensor::from_vec_col_major(vec![4], vec![1.0_f64, 2.0, 3.0, 4.0]).unwrap();
    let (full, onesided, recovered, recovered_real) = with_cpu_session(&mut backend, |backend| {
        let full = real.fft(None, -1, FftNorm::Backward, backend).unwrap();
        let onesided = real.rfft(None, -1, FftNorm::Backward, backend).unwrap();
        let recovered = full.ifft(None, -1, FftNorm::Backward, backend).unwrap();
        let recovered_real = onesided
            .irfft(Some(4), -1, FftNorm::Backward, backend)
            .unwrap();
        (full, onesided, recovered, recovered_real)
    });

    assert_eq!(full.shape(), &[4]);
    assert_complex_close(
        full.as_slice::<Complex64>().unwrap(),
        &[
            Complex64::new(10.0, 0.0),
            Complex64::new(-2.0, 2.0),
            Complex64::new(-2.0, 0.0),
            Complex64::new(-2.0, -2.0),
        ],
    );
    assert_eq!(onesided.shape(), &[3]);
    assert_complex_close(
        onesided.as_slice::<Complex64>().unwrap(),
        &[
            Complex64::new(10.0, 0.0),
            Complex64::new(-2.0, 2.0),
            Complex64::new(-2.0, 0.0),
        ],
    );

    assert_complex_close(
        recovered.as_slice::<Complex64>().unwrap(),
        &[
            Complex64::new(1.0, 0.0),
            Complex64::new(2.0, 0.0),
            Complex64::new(3.0, 0.0),
            Complex64::new(4.0, 0.0),
        ],
    );

    assert_real_close(
        recovered_real.as_slice::<f64>().unwrap(),
        &[1.0, 2.0, 3.0, 4.0],
    );
}

#[test]
fn public_tensor_read_fft_ext_accepts_strided_host_views() {
    let mut backend = CpuBackend::new();
    let data = [1.0_f64, 99.0, 2.0, 99.0, 3.0, 99.0, 4.0];
    let view = TypedTensorView::from_slice([4], [2], 0, &data).unwrap();
    let input = TensorRead::from_view(TensorView::F64(view));

    let (full, onesided) = with_cpu_session(&mut backend, |backend| {
        let full = input
            .fft_read(None, -1, FftNorm::Backward, backend)
            .unwrap();
        let onesided = input
            .rfft_read(None, -1, FftNorm::Backward, backend)
            .unwrap();
        (full, onesided)
    });

    assert_complex_close(
        full.as_slice::<Complex64>().unwrap(),
        &[
            Complex64::new(10.0, 0.0),
            Complex64::new(-2.0, 2.0),
            Complex64::new(-2.0, 0.0),
            Complex64::new(-2.0, -2.0),
        ],
    );
    assert_complex_close(
        onesided.as_slice::<Complex64>().unwrap(),
        &[
            Complex64::new(10.0, 0.0),
            Complex64::new(-2.0, 2.0),
            Complex64::new(-2.0, 0.0),
        ],
    );
}

#[test]
fn public_tensor_fft_ext_reports_invalid_dtype_and_shape_errors() {
    let mut backend = CpuBackend::new();
    let bools = Tensor::from_vec_col_major(vec![2], vec![true, false]).unwrap();
    let spectrum = Tensor::from_vec_col_major(
        vec![3],
        vec![
            Complex64::new(10.0, 0.0),
            Complex64::new(-2.0, 2.0),
            Complex64::new(-2.0, 0.0),
        ],
    )
    .unwrap();

    let (dtype_err, shape_err) = with_cpu_session(&mut backend, |backend| {
        let dtype_err = bools.fft(None, -1, FftNorm::Backward, backend).unwrap_err();
        let shape_err = spectrum
            .irfft(Some(6), -1, FftNorm::Backward, backend)
            .unwrap_err();
        (dtype_err, shape_err)
    });

    assert!(matches!(
        dtype_err,
        tenferro_tensor::Error::Extension {
            op: "TensorFftExt::fft",
            kind: ErrorKind::Unsupported,
            ..
        }
    ));
    assert!(matches!(
        shape_err,
        tenferro_tensor::Error::Validation { op: "irfft", .. }
    ));
}

#[test]
fn fft_plan_cache_is_bounded_lru_and_reports_known_retention() {
    let mut cache = FftPlanCache::with_capacity(NonZeroUsize::new(2).unwrap());
    assert_eq!(cache.capacity().get(), 2);
    let first = cache.plan_f64(4, true);
    cache.plan_f64(8, true);
    assert!(std::sync::Arc::ptr_eq(&first, &cache.plan_f64(4, true)));
    cache.plan_f64(16, true);

    assert!(cache.contains_f64(4, true));
    assert!(!cache.contains_f64(8, true));
    assert!(cache.contains_f64(16, true));
    assert_eq!(cache.stats().entries, 2);
    assert!(cache.stats().retained_bytes > 0);
    assert!(std::sync::Arc::ptr_eq(&first, &cache.plan_f64(4, true)));

    cache.clear();
    let stats = cache.stats();
    assert_eq!(stats.entries, 0);
    assert_eq!(stats.retained_bytes, 0);
    assert_eq!(stats.hits, 4);
    assert_eq!(stats.misses, 4);
    assert_eq!(stats.evictions, 1);
    assert_eq!(stats.clears, 1);
    cache.set_capacity(NonZeroUsize::MIN);
    assert_eq!(cache.capacity(), NonZeroUsize::MIN);
}

#[test]
fn fft_plan_cache_exposes_full_extension_cache_limits() {
    let mut cache = FftPlanCache::with_capacity(NonZeroUsize::new(8).unwrap());
    let limits = ExtensionCacheLimits::new(NonZeroUsize::new(8).unwrap())
        .with_max_retained_bytes(NonZeroUsize::new(1).unwrap());

    cache.set_limits(limits);
    assert_eq!(cache.limits(), limits);
    cache.plan_f64(4, true);

    assert_eq!(cache.stats().entries, 0);
    assert_eq!(cache.stats().evictions, 1);
}

#[test]
fn fft_plan_cache_key_distinguishes_scalar_dtype_length_and_direction() {
    let mut cache = FftPlanCache::with_capacity(NonZeroUsize::new(8).unwrap());

    let f32_forward_4 = cache.plan_f32(4, true);
    let f32_forward_8 = cache.plan_f32(8, true);
    let f64_forward_4 = cache.plan_f64(4, true);
    let f64_inverse_4 = cache.plan_f64(4, false);

    assert_eq!(cache.stats().entries, 4);
    assert!(std::sync::Arc::ptr_eq(
        &f32_forward_4,
        &cache.plan_f32(4, true)
    ));
    assert!(std::sync::Arc::ptr_eq(
        &f32_forward_8,
        &cache.plan_f32(8, true)
    ));
    assert!(std::sync::Arc::ptr_eq(
        &f64_forward_4,
        &cache.plan_f64(4, true)
    ));
    assert!(std::sync::Arc::ptr_eq(
        &f64_inverse_4,
        &cache.plan_f64(4, false)
    ));
    assert!(!std::sync::Arc::ptr_eq(&f64_forward_4, &f64_inverse_4));
}

#[test]
fn caller_owned_fft_executor_reuses_plans() {
    let mut backend = CpuBackend::new();
    let input = Tensor::from_vec_col_major(vec![4], vec![1.0_f64, 2.0, 3.0, 4.0]).unwrap();
    let mut executor = FftExecutor::default();

    with_cpu_session(&mut backend, |backend| {
        let full = executor
            .fft(&input, None, -1, FftNorm::Backward, backend)
            .unwrap();
        executor
            .ifft(&full, None, -1, FftNorm::Backward, backend)
            .unwrap();
        let onesided = executor
            .rfft(&input, None, -1, FftNorm::Backward, backend)
            .unwrap();
        executor
            .irfft(&onesided, Some(4), -1, FftNorm::Backward, backend)
            .unwrap();
    });

    assert_eq!(executor.cache_stats().entries, 2);
    assert_eq!(executor.plan_cache().stats().entries, 2);
    executor.plan_cache_mut().set_capacity(NonZeroUsize::MIN);
    assert_eq!(executor.cache_stats().entries, 1);
    executor.clear_cache();
    assert_eq!(executor.cache_stats().entries, 0);
}

#[test]
fn configured_fft_executor_validates_each_public_operation_before_dispatch() {
    let mut executor = FftExecutor::new(FftPlanCache::default());
    let mut backend = CpuBackend::new();
    let real = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
    let complex = Tensor::from_vec_col_major(
        vec![2],
        vec![Complex64::new(1.0, 0.0), Complex64::new(2.0, 0.0)],
    )
    .unwrap();

    with_cpu_session(&mut backend, |backend| {
        assert!(executor
            .fft(&real, Some(0), -1, FftNorm::Backward, backend)
            .is_err());
        assert!(executor
            .ifft(&real, None, -1, FftNorm::Backward, backend)
            .is_err());
        assert!(executor
            .rfft(&complex, None, -1, FftNorm::Backward, backend)
            .is_err());
        assert!(executor
            .irfft(&complex, Some(0), -1, FftNorm::Backward, backend)
            .is_err());
    });
}