Skip to main content

fpm_rs/backend/
cpu.rs

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/// rustfft-based CPU backend with plans cached for low- and high-resolution grids.
82#[derive(Clone)]
83pub struct CpuBackend {
84    low: CpuFftPlan,
85    high: CpuFftPlan,
86}
87
88impl CpuBackend {
89    /// Creates cached FFT plans for non-zero low- and high-resolution `(height, width)` grids.
90    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        // The forward transform is normalized by 1/N and the inverse is
163        // unnormalized. This pair is exactly invertible and, unlike the common
164        // inverse-normalized convention, preserves the amplitude of a constant
165        // object when a high-resolution spectrum is cropped and inverse-
166        // transformed on a smaller grid.
167        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}