Skip to main content

mlx_rs/fft/
shift.rs

1use mlx_internal_macros::generate_macro;
2use smallvec::SmallVec;
3
4use crate::{
5    array::Array, constants::DEFAULT_STACK_VEC_LEN, error::Result, utils::guard::Guarded,
6    utils::IntoOption, with_stream, Stream,
7};
8
9/// Resolve axes for shift operations - when None, returns all axes
10fn resolve_axes(a: &Array, axes: Option<&[i32]>) -> SmallVec<[i32; DEFAULT_STACK_VEC_LEN]> {
11    match axes {
12        Some(axes) => SmallVec::from_slice(axes),
13        None => (0..a.ndim() as i32).collect(),
14    }
15}
16
17/// Shift the zero-frequency component to the center of the spectrum.
18///
19/// This function swaps half-spaces for all axes listed (defaults to all).
20/// Note that `y[0]` is the Nyquist component only if `len(x)` is even.
21///
22/// # Params
23///
24/// - `a`: The input array.
25/// - `axes`: Axes over which to shift. The default is `None` which shifts all axes.
26///
27/// # Example
28///
29/// ```rust
30/// use mlx_rs::{Array, fft::*};
31///
32/// let a = Array::from_slice(&[0.0f32, 1.0, 2.0, 3.0, 4.0, -4.0, -3.0, -2.0, -1.0], &[9]);
33/// let shifted = fftshift(&a, None).unwrap();
34/// // shifted contains: [-4.0, -3.0, -2.0, -1.0, 0.0, 1.0, 2.0, 3.0, 4.0]
35/// ```
36pub fn fftshift<'a>(a: impl AsRef<Array>, axes: impl IntoOption<&'a [i32]>) -> Result<Array> {
37    let a = a.as_ref();
38    let axes = resolve_axes(a, axes.into_option());
39    let stream = Stream::thread_local_or_default();
40
41    Array::try_from_op(|res| unsafe {
42        mlx_sys::mlx_fft_fftshift(
43            res,
44            a.as_ptr(),
45            axes.as_ptr(),
46            axes.len(),
47            stream.as_ref().as_ptr(),
48        )
49    })
50}
51
52/// Compatibility shim for [`fftshift`].
53#[generate_macro(customize(forwarding_shim = true, root = "$crate::fft"))]
54#[deprecated(
55    since = "0.26.0",
56    note = "use `with_stream` or `with_device` around `fftshift`"
57)]
58pub fn fftshift_device<'a>(
59    a: impl AsRef<Array>,
60    #[optional] axes: impl IntoOption<&'a [i32]>,
61    #[optional] stream: impl AsRef<Stream>,
62) -> Result<Array> {
63    with_stream(stream.as_ref(), || fftshift(a, axes))
64}
65
66/// The inverse of `fftshift`.
67///
68/// Although identical for even-length `x`, the functions differ by one sample for odd-length `x`.
69///
70/// # Params
71///
72/// - `a`: The input array.
73/// - `axes`: Axes over which to calculate. The default is `None` which shifts all axes.
74///
75/// # Example
76///
77/// ```rust
78/// use mlx_rs::{Array, fft::*};
79///
80/// let a = Array::from_slice(&[-4.0f32, -3.0, -2.0, -1.0, 0.0, 1.0, 2.0, 3.0, 4.0], &[9]);
81/// let unshifted = ifftshift(&a, None).unwrap();
82/// // unshifted contains: [0.0, 1.0, 2.0, 3.0, 4.0, -4.0, -3.0, -2.0, -1.0]
83/// ```
84pub fn ifftshift<'a>(a: impl AsRef<Array>, axes: impl IntoOption<&'a [i32]>) -> Result<Array> {
85    let a = a.as_ref();
86    let axes = resolve_axes(a, axes.into_option());
87    let stream = Stream::thread_local_or_default();
88
89    Array::try_from_op(|res| unsafe {
90        mlx_sys::mlx_fft_ifftshift(
91            res,
92            a.as_ptr(),
93            axes.as_ptr(),
94            axes.len(),
95            stream.as_ref().as_ptr(),
96        )
97    })
98}
99
100/// Compatibility shim for [`ifftshift`].
101#[generate_macro(customize(forwarding_shim = true, root = "$crate::fft"))]
102#[deprecated(
103    since = "0.26.0",
104    note = "use `with_stream` or `with_device` around `ifftshift`"
105)]
106pub fn ifftshift_device<'a>(
107    a: impl AsRef<Array>,
108    #[optional] axes: impl IntoOption<&'a [i32]>,
109    #[optional] stream: impl AsRef<Stream>,
110) -> Result<Array> {
111    with_stream(stream.as_ref(), || ifftshift(a, axes))
112}
113
114#[cfg(test)]
115mod tests {
116    use super::*;
117    use crate::random;
118
119    // Helper to check fftshift matches expected behavior
120    fn check_fftshift(a: &Array, axes: Option<&[i32]>) {
121        let shifted = fftshift(a, axes).unwrap();
122        let unshifted = ifftshift(&shifted, axes).unwrap();
123        assert!(
124            unshifted.all_close(a, 1e-5, 1e-6, None).unwrap(),
125            "ifftshift(fftshift(x)) should equal x"
126        );
127    }
128
129    #[test]
130    fn test_fftshift_1d() {
131        // Test 1D arrays (matches Python test)
132        random::seed(42).unwrap();
133        let r = random::uniform::<_, f32>(0.0, 1.0, &[100], None).unwrap();
134        check_fftshift(&r, None);
135    }
136
137    #[test]
138    fn test_fftshift_with_axes() {
139        // Test with specific axis (matches Python test)
140        random::seed(42).unwrap();
141        let r = random::uniform::<_, f32>(0.0, 1.0, &[4, 6], None).unwrap();
142        check_fftshift(&r, Some(&[0]));
143        check_fftshift(&r, Some(&[1]));
144        check_fftshift(&r, Some(&[0, 1]));
145    }
146
147    #[test]
148    fn test_fftshift_negative_axes() {
149        // Test with negative axes (matches Python test)
150        random::seed(42).unwrap();
151        let r = random::uniform::<_, f32>(0.0, 1.0, &[4, 6], None).unwrap();
152        check_fftshift(&r, Some(&[-1]));
153    }
154
155    #[test]
156    fn test_fftshift_odd_lengths() {
157        // Test with odd lengths (matches Python test)
158        random::seed(42).unwrap();
159        let r = random::uniform::<_, f32>(0.0, 1.0, &[5, 7], None).unwrap();
160        check_fftshift(&r, None);
161        check_fftshift(&r, Some(&[0]));
162    }
163
164    #[test]
165    fn test_ifftshift_1d() {
166        // Test 1D arrays (matches Python test)
167        random::seed(42).unwrap();
168        let r = random::uniform::<_, f32>(0.0, 1.0, &[100], None).unwrap();
169
170        let shifted = ifftshift(&r, None).unwrap();
171        let unshifted = fftshift(&shifted, None).unwrap();
172        assert!(
173            unshifted.all_close(&r, 1e-5, 1e-6, None).unwrap(),
174            "fftshift(ifftshift(x)) should equal x"
175        );
176    }
177
178    #[test]
179    fn test_ifftshift_with_axes() {
180        // Test with specific axis (matches Python test)
181        random::seed(42).unwrap();
182        let r = random::uniform::<_, f32>(0.0, 1.0, &[4, 6], None).unwrap();
183
184        for axes in [&[0][..], &[1][..], &[0, 1][..]] {
185            let shifted = ifftshift(&r, axes).unwrap();
186            let unshifted = fftshift(&shifted, axes).unwrap();
187            assert!(
188                unshifted.all_close(&r, 1e-5, 1e-6, None).unwrap(),
189                "fftshift(ifftshift(x)) should equal x for axes {:?}",
190                axes
191            );
192        }
193    }
194
195    #[test]
196    fn test_ifftshift_negative_axes() {
197        // Test with negative axes (matches Python test)
198        random::seed(42).unwrap();
199        let r = random::uniform::<_, f32>(0.0, 1.0, &[4, 6], None).unwrap();
200
201        let shifted = ifftshift(&r, &[-1]).unwrap();
202        let unshifted = fftshift(&shifted, &[-1]).unwrap();
203        assert!(unshifted.all_close(&r, 1e-5, 1e-6, None).unwrap(),);
204    }
205
206    #[test]
207    fn test_ifftshift_odd_lengths() {
208        // Test with odd lengths (matches Python test)
209        random::seed(42).unwrap();
210        let r = random::uniform::<_, f32>(0.0, 1.0, &[5, 7], None).unwrap();
211
212        let shifted = ifftshift(&r, None).unwrap();
213        let unshifted = fftshift(&shifted, None).unwrap();
214        assert!(unshifted.all_close(&r, 1e-5, 1e-6, None).unwrap(),);
215
216        let shifted = ifftshift(&r, &[0]).unwrap();
217        let unshifted = fftshift(&shifted, &[0]).unwrap();
218        assert!(unshifted.all_close(&r, 1e-5, 1e-6, None).unwrap(),);
219    }
220
221    #[test]
222    fn test_fftshift_empty_array() {
223        // Test empty array (matches Python test)
224        let x = Array::from_slice::<f32>(&[], &[0]);
225        let shifted = fftshift(&x, None).unwrap();
226        assert!(shifted.eq_exact(&x).unwrap());
227    }
228}