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
9fn 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
17pub 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#[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
66pub 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#[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 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 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 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 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 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 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 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 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 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 let x = Array::from_slice::<f32>(&[], &[0]);
225 let shifted = fftshift(&x, None).unwrap();
226 assert!(shifted.eq_exact(&x).unwrap());
227 }
228}