1use std::sync::Arc;
2
3use num_complex::Complex64;
4use rustfft::{Fft, FftPlanner};
5
6use crate::{Result, array_layout::checked_len_2d, error::Error};
7
8use super::{
9 Backend, BackendCapabilities, ComplexBuffer, FftDirection, MemoryLocation, RealBuffer,
10 ResidentBackend,
11};
12
13#[derive(Debug)]
14struct CpuComplexBuffer {
15 values: Vec<Complex64>,
16}
17
18impl ComplexBuffer for CpuComplexBuffer {
19 fn len(&self) -> usize {
20 self.values.len()
21 }
22
23 fn location(&self) -> MemoryLocation {
24 MemoryLocation::Host
25 }
26
27 fn as_any(&self) -> &dyn std::any::Any {
28 self
29 }
30
31 fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
32 self
33 }
34}
35
36#[derive(Debug)]
37struct CpuRealBuffer {
38 values: Vec<f64>,
39}
40
41impl RealBuffer for CpuRealBuffer {
42 fn len(&self) -> usize {
43 self.values.len()
44 }
45
46 fn location(&self) -> MemoryLocation {
47 MemoryLocation::Host
48 }
49
50 fn as_any(&self) -> &dyn std::any::Any {
51 self
52 }
53
54 fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
55 self
56 }
57}
58
59#[derive(Clone)]
60struct CpuFftPlan {
61 shape: (usize, usize),
62 row_forward: Arc<dyn Fft<f64>>,
63 row_inverse: Arc<dyn Fft<f64>>,
64 column_forward: Arc<dyn Fft<f64>>,
65 column_inverse: Arc<dyn Fft<f64>>,
66}
67
68impl CpuFftPlan {
69 fn new(shape: (usize, usize)) -> Self {
70 let mut planner = FftPlanner::new();
71 Self {
72 shape,
73 row_forward: planner.plan_fft_forward(shape.1),
74 row_inverse: planner.plan_fft_inverse(shape.1),
75 column_forward: planner.plan_fft_forward(shape.0),
76 column_inverse: planner.plan_fft_inverse(shape.0),
77 }
78 }
79}
80
81#[derive(Clone)]
83pub struct CpuBackend {
84 low: CpuFftPlan,
85 high: CpuFftPlan,
86}
87
88impl CpuBackend {
89 pub fn new(low_shape: (usize, usize), high_shape: (usize, usize)) -> Result<Self> {
91 if low_shape.0 == 0 || low_shape.1 == 0 || high_shape.0 == 0 || high_shape.1 == 0 {
92 return Err(Error::InvalidShape(
93 "FFT dimensions must be non-zero".into(),
94 ));
95 }
96 checked_len_2d(low_shape)?;
97 checked_len_2d(high_shape)?;
98 Ok(Self {
99 low: CpuFftPlan::new(low_shape),
100 high: CpuFftPlan::new(high_shape),
101 })
102 }
103
104 fn plan(&self, shape: (usize, usize)) -> Result<&CpuFftPlan> {
105 if shape == self.low.shape {
106 Ok(&self.low)
107 } else if shape == self.high.shape {
108 Ok(&self.high)
109 } else {
110 Err(Error::InvalidShape(format!(
111 "FFT shape {shape:?} is neither configured shape {:?} nor {:?}",
112 self.low.shape, self.high.shape
113 )))
114 }
115 }
116}
117
118impl Backend for CpuBackend {
119 fn capabilities(&self) -> BackendCapabilities {
120 BackendCapabilities {
121 preferred_memory: MemoryLocation::Host,
122 resident_buffers: true,
123 }
124 }
125
126 fn resident_backend(&self) -> Option<&dyn ResidentBackend> {
127 Some(self)
128 }
129
130 fn fft2(
131 &self,
132 values: &mut [Complex64],
133 shape: (usize, usize),
134 direction: FftDirection,
135 column_scratch: &mut [Complex64],
136 ) -> Result<()> {
137 let expected = checked_len_2d(shape)?;
138 if values.len() != expected || column_scratch.len() < shape.0 {
139 return Err(Error::InvalidShape(format!(
140 "FFT {:?} requires {expected} values and {} column scratch values",
141 shape, shape.0
142 )));
143 }
144 let plan = self.plan(shape)?;
145 let (row_fft, column_fft) = match direction {
146 FftDirection::Forward => (&plan.row_forward, &plan.column_forward),
147 FftDirection::Inverse => (&plan.row_inverse, &plan.column_inverse),
148 };
149 for row in values.chunks_exact_mut(shape.1) {
150 row_fft.process(row);
151 }
152 let column = &mut column_scratch[..shape.0];
153 for column_index in 0..shape.1 {
154 for row_index in 0..shape.0 {
155 column[row_index] = values[row_index * shape.1 + column_index];
156 }
157 column_fft.process(column);
158 for row_index in 0..shape.0 {
159 values[row_index * shape.1 + column_index] = column[row_index];
160 }
161 }
162 if direction == FftDirection::Forward {
168 let normalization = expected as f64;
169 for value in values {
170 *value /= normalization;
171 }
172 }
173 Ok(())
174 }
175}
176
177impl ResidentBackend for CpuBackend {
178 fn allocate_complex(&self, len: usize) -> Result<Box<dyn ComplexBuffer>> {
179 Ok(Box::new(CpuComplexBuffer {
180 values: vec![Complex64::default(); len],
181 }))
182 }
183
184 fn allocate_real(&self, len: usize) -> Result<Box<dyn RealBuffer>> {
185 Ok(Box::new(CpuRealBuffer {
186 values: vec![0.0; len],
187 }))
188 }
189
190 fn upload_complex(
191 &self,
192 destination: &mut dyn ComplexBuffer,
193 source: &[Complex64],
194 ) -> Result<()> {
195 let destination = cpu_complex_mut(destination)?;
196 validate_transfer_length(destination.values.len(), source.len())?;
197 destination.values.copy_from_slice(source);
198 Ok(())
199 }
200
201 fn download_complex(
202 &self,
203 source: &dyn ComplexBuffer,
204 destination: &mut [Complex64],
205 ) -> Result<()> {
206 let source = cpu_complex(source)?;
207 validate_transfer_length(destination.len(), source.values.len())?;
208 destination.copy_from_slice(&source.values);
209 Ok(())
210 }
211
212 fn upload_real(&self, destination: &mut dyn RealBuffer, source: &[f64]) -> Result<()> {
213 let destination = cpu_real_mut(destination)?;
214 validate_transfer_length(destination.values.len(), source.len())?;
215 destination.values.copy_from_slice(source);
216 Ok(())
217 }
218
219 fn download_real(&self, source: &dyn RealBuffer, destination: &mut [f64]) -> Result<()> {
220 let source = cpu_real(source)?;
221 validate_transfer_length(destination.len(), source.values.len())?;
222 destination.copy_from_slice(&source.values);
223 Ok(())
224 }
225
226 fn fft2_resident(
227 &self,
228 values: &mut dyn ComplexBuffer,
229 shape: (usize, usize),
230 direction: FftDirection,
231 ) -> Result<()> {
232 let values = cpu_complex_mut(values)?;
233 let mut column = vec![Complex64::default(); shape.0];
234 self.fft2(&mut values.values, shape, direction, &mut column)
235 }
236}
237
238fn cpu_complex(buffer: &dyn ComplexBuffer) -> Result<&CpuComplexBuffer> {
239 buffer
240 .as_any()
241 .downcast_ref()
242 .ok_or_else(|| Error::InvalidParameter {
243 name: "complex buffer",
244 reason: "buffer was not allocated by CpuBackend".into(),
245 })
246}
247
248fn cpu_complex_mut(buffer: &mut dyn ComplexBuffer) -> Result<&mut CpuComplexBuffer> {
249 buffer
250 .as_any_mut()
251 .downcast_mut()
252 .ok_or_else(|| Error::InvalidParameter {
253 name: "complex buffer",
254 reason: "buffer was not allocated by CpuBackend".into(),
255 })
256}
257
258fn cpu_real(buffer: &dyn RealBuffer) -> Result<&CpuRealBuffer> {
259 buffer
260 .as_any()
261 .downcast_ref()
262 .ok_or_else(|| Error::InvalidParameter {
263 name: "real buffer",
264 reason: "buffer was not allocated by CpuBackend".into(),
265 })
266}
267
268fn cpu_real_mut(buffer: &mut dyn RealBuffer) -> Result<&mut CpuRealBuffer> {
269 buffer
270 .as_any_mut()
271 .downcast_mut()
272 .ok_or_else(|| Error::InvalidParameter {
273 name: "real buffer",
274 reason: "buffer was not allocated by CpuBackend".into(),
275 })
276}
277
278fn validate_transfer_length(destination: usize, source: usize) -> Result<()> {
279 if destination != source {
280 return Err(Error::InvalidShape(format!(
281 "buffer transfer length {source} does not match destination length {destination}"
282 )));
283 }
284 Ok(())
285}