tenferro-fft 0.3.0

FFT extension runtime and public concrete/traced FFT APIs for tenferro.
use std::hash::{Hash, Hasher};

use super::super::descriptor::{
    CufftDirection, CufftPlanDescriptor, CufftPlanStructuralKey, CufftTransformKind,
};
use super::super::error::{into_tensor_error, CudaFftError};

fn assert_descriptor(
    kind: CufftTransformKind,
    direction: CufftDirection,
    inembed: [i64; 1],
    onembed: [i64; 1],
) {
    let descriptor = CufftPlanDescriptor::new(kind, direction, 8, 3).unwrap();

    assert_eq!(descriptor.kind, kind);
    assert_eq!(descriptor.direction, direction);
    assert_eq!(descriptor.rank, 1);
    assert_eq!(descriptor.n, [8]);
    assert_eq!(descriptor.inembed, inembed);
    assert_eq!(descriptor.onembed, onembed);
    assert_eq!(descriptor.istride, 3);
    assert_eq!(descriptor.idist, 1);
    assert_eq!(descriptor.ostride, 3);
    assert_eq!(descriptor.odist, 1);
    assert_eq!(descriptor.batch, 3);
}

#[test]
fn descriptor_maps_all_cufft_transform_kinds_to_rank_one_layouts() {
    assert_descriptor(CufftTransformKind::C2c32, CufftDirection::Forward, [8], [8]);
    assert_descriptor(CufftTransformKind::C2c64, CufftDirection::Inverse, [8], [8]);
    assert_descriptor(CufftTransformKind::R2c32, CufftDirection::Forward, [8], [5]);
    assert_descriptor(CufftTransformKind::R2c64, CufftDirection::Forward, [8], [5]);
    assert_descriptor(CufftTransformKind::C2r32, CufftDirection::Inverse, [5], [8]);
    assert_descriptor(CufftTransformKind::C2r64, CufftDirection::Inverse, [5], [8]);
}

#[test]
fn descriptor_uses_ceil_half_spectrum_extent_for_odd_lengths() {
    let r2c =
        CufftPlanDescriptor::new(CufftTransformKind::R2c64, CufftDirection::Forward, 7, 3).unwrap();
    assert_eq!(r2c.kind, CufftTransformKind::R2c64);
    assert_eq!(r2c.direction, CufftDirection::Forward);
    assert_eq!(r2c.rank, 1);
    assert_eq!(r2c.n, [7]);
    assert_eq!(r2c.inembed, [7]);
    assert_eq!(r2c.onembed, [4]);
    assert_eq!(r2c.istride, 3);
    assert_eq!(r2c.idist, 1);
    assert_eq!(r2c.ostride, 3);
    assert_eq!(r2c.odist, 1);
    assert_eq!(r2c.batch, 3);

    let c2r =
        CufftPlanDescriptor::new(CufftTransformKind::C2r64, CufftDirection::Inverse, 7, 3).unwrap();
    assert_eq!(c2r.kind, CufftTransformKind::C2r64);
    assert_eq!(c2r.direction, CufftDirection::Inverse);
    assert_eq!(c2r.rank, 1);
    assert_eq!(c2r.n, [7]);
    assert_eq!(c2r.inembed, [4]);
    assert_eq!(c2r.onembed, [7]);
    assert_eq!(c2r.istride, 3);
    assert_eq!(c2r.idist, 1);
    assert_eq!(c2r.ostride, 3);
    assert_eq!(c2r.odist, 1);
    assert_eq!(c2r.batch, 3);
}

fn assert_invalid(result: Result<CufftPlanDescriptor, CudaFftError>, field: &'static str) {
    assert!(matches!(
        result,
        Err(CudaFftError::InvalidConfiguration { field: actual }) if actual == field
    ));
}

#[test]
fn descriptor_rejects_zero_and_overflowing_configurations() {
    assert_invalid(
        CufftPlanDescriptor::new(CufftTransformKind::C2c32, CufftDirection::Forward, 0, 1),
        "n",
    );
    assert_invalid(
        CufftPlanDescriptor::new(CufftTransformKind::C2c32, CufftDirection::Forward, 1, 0),
        "batch",
    );
    let product_overflow_n = usize::try_from(i64::MAX).unwrap_or(usize::MAX / 2 + 1);
    assert_invalid(
        CufftPlanDescriptor::new(
            CufftTransformKind::C2c32,
            CufftDirection::Forward,
            product_overflow_n,
            2,
        ),
        "element_count",
    );

    if let Ok(outside_signed_width) = usize::try_from(i64::MAX as u128 + 1) {
        assert_invalid(
            CufftPlanDescriptor::new(
                CufftTransformKind::C2c32,
                CufftDirection::Forward,
                outside_signed_width,
                1,
            ),
            "n",
        );
    }
}

#[derive(Default)]
struct ConstantHasher;

impl Hasher for ConstantHasher {
    fn finish(&self) -> u64 {
        0
    }

    fn write(&mut self, _bytes: &[u8]) {}
}

fn constant_hash<T: Hash>(value: &T) -> u64 {
    let mut hasher = ConstantHasher;
    value.hash(&mut hasher);
    hasher.finish()
}

#[test]
fn distinct_structural_keys_require_exact_comparison_even_when_hashes_collide() {
    let first = CufftPlanStructuralKey {
        device_ordinal: 0,
        kind: CufftTransformKind::C2c32,
        direction: CufftDirection::Forward,
        n: 8,
        batch: 3,
        istride: 3,
        idist: 1,
        ostride: 3,
        odist: 1,
    };
    let second = CufftPlanStructuralKey {
        device_ordinal: 0,
        kind: CufftTransformKind::C2c32,
        direction: CufftDirection::Forward,
        n: 7,
        batch: 3,
        istride: 3,
        idist: 1,
        ostride: 3,
        odist: 1,
    };

    assert_eq!(constant_hash(&first), constant_hash(&second));
    assert_ne!(first, second);
}

#[test]
fn invalid_descriptor_configuration_is_structured_tensor_validation() {
    let error = into_tensor_error("fft", CudaFftError::InvalidConfiguration { field: "batch" });

    assert_eq!(
        error.kind(),
        tenferro_tensor::ErrorKind::Validation(tenferro_tensor::ValidationKind::InvalidArgument)
    );
    assert!(matches!(
        error,
        tenferro_tensor::Error::Validation {
            source: tenferro_tensor::ValidationError::InvalidArgument {
                argument: "batch",
                ..
            },
            ..
        }
    ));
}

#[test]
fn internal_invariant_translation_uses_tensor_internal_error() {
    let error = into_tensor_error(
        "fft",
        CudaFftError::InternalInvariant {
            message: "descriptor unexpectedly missing",
        },
    );

    assert!(matches!(
        error,
        tenferro_tensor::Error::Internal(message)
            if message == "descriptor unexpectedly missing"
    ));
}