Skip to main content

mlx_rs/fft/
frequencies.rs

1use crate::error::{Exception, Result};
2use crate::utils::guard::Guarded;
3use crate::{Array, Stream};
4
5fn checked_length(n: usize) -> Result<i32> {
6    i32::try_from(n).map_err(|_| Exception::custom("FFT length exceeds i32::MAX"))
7}
8
9/// Returns the discrete Fourier transform sample frequencies for a transform of length `n`.
10///
11/// The sample spacing is `d`.
12///
13/// # Example
14///
15/// ```rust
16/// use mlx_rs::{fft::fftfreq, Device, Dtype};
17///
18/// Device::set_default(&Device::cpu());
19/// let frequencies = fftfreq(4, 1.0).unwrap();
20/// assert_eq!(frequencies.shape(), &[4]);
21/// assert_eq!(frequencies.dtype(), Dtype::Float32);
22/// ```
23pub fn fftfreq(n: usize, d: f64) -> Result<Array> {
24    let n = checked_length(n)?;
25    let stream = Stream::thread_local_or_default();
26    Array::try_from_op(|res| unsafe { mlx_sys::mlx_fft_fftfreq(res, n, d, stream.as_ptr()) })
27}
28
29/// Returns the nonnegative discrete Fourier transform sample frequencies for a real transform.
30///
31/// The transform length is `n` and the sample spacing is `d`.
32///
33/// # Example
34///
35/// ```rust
36/// use mlx_rs::{fft::rfftfreq, Device, Dtype};
37///
38/// Device::set_default(&Device::cpu());
39/// let frequencies = rfftfreq(4, 0.5).unwrap();
40/// assert_eq!(frequencies.shape(), &[3]);
41/// assert_eq!(frequencies.dtype(), Dtype::Float32);
42/// ```
43pub fn rfftfreq(n: usize, d: f64) -> Result<Array> {
44    let n = checked_length(n)?;
45    let stream = Stream::thread_local_or_default();
46    Array::try_from_op(|res| unsafe { mlx_sys::mlx_fft_rfftfreq(res, n, d, stream.as_ptr()) })
47}
48
49#[cfg(test)]
50mod tests {
51    use super::*;
52
53    #[test]
54    fn rejects_lengths_above_i32_max() {
55        assert!(fftfreq(usize::MAX, 1.0).is_err());
56        assert!(rfftfreq(usize::MAX, 1.0).is_err());
57    }
58}