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}