Skip to main content

tenferro_fft/
cache.rs

1use std::collections::hash_map::DefaultHasher;
2use std::fmt;
3use std::hash::{Hash, Hasher};
4use std::num::NonZeroUsize;
5use std::sync::Arc;
6
7use num_traits::{Float, FromPrimitive};
8use rustfft::{Fft, FftNum, FftPlanner};
9use tenferro_runtime::{
10    ExtensionCacheKey, ExtensionCacheLimits, ExtensionCacheSelector, ExtensionCacheStore,
11};
12use tenferro_tensor::{CacheStats, RuntimeCacheControl};
13
14use crate::FFT_EXTENSION_FAMILY_ID;
15
16/// Runtime cache namespace used for private RustFFT plans.
17pub const FFT_PLAN_CACHE_NAME: &str = "rustfft-plans";
18
19/// Default number of typed entries retained by a caller-owned [`FftPlanCache`].
20pub const DEFAULT_FFT_PLAN_CACHE_CAPACITY: usize = 64;
21
22/// Select the private CPU RustFFT plan entries in an extension runtime cache.
23///
24/// Other [`crate::FftBackend`] implementations use distinct cache names in the
25/// same FFT extension family.
26///
27/// # Examples
28///
29/// ```
30/// let selector = tenferro_fft::fft_plan_cache_selector();
31/// assert!(matches!(
32///     selector,
33///     tenferro_runtime::ExtensionCacheSelector::Cache { .. }
34/// ));
35/// ```
36pub const fn fft_plan_cache_selector() -> ExtensionCacheSelector {
37    ExtensionCacheSelector::Cache {
38        family_id: FFT_EXTENSION_FAMILY_ID,
39        cache_name: FFT_PLAN_CACHE_NAME,
40    }
41}
42
43/// Bounded, caller-owned typed cache for backend FFT plans and workspaces.
44///
45/// [`crate::FftExecutor`] owns one cache and passes its store to every
46/// [`crate::FftBackend`] through [`crate::FftExecutionCache`]. Backends retain
47/// private `Send + Sync + 'static` values under their own
48/// [`ExtensionCacheKey`] namespace. Entry limits, LRU eviction, clearing,
49/// aggregate entry counts, and retained-byte accounting apply uniformly to CPU
50/// and non-CPU entries.
51///
52/// The CPU backend stores RustFFT plans in the private `rustfft-plans`
53/// namespace. Its retained-byte estimate includes the exact plan key and the
54/// cache-owned `Arc` handle; allocations opaque to RustFFT are excluded.
55///
56/// # Examples
57///
58/// ```
59/// use std::num::NonZeroUsize;
60/// use tenferro_fft::FftPlanCache;
61///
62/// let cache = FftPlanCache::with_capacity(NonZeroUsize::new(2).unwrap());
63/// assert_eq!(cache.capacity().get(), 2);
64/// assert_eq!(cache.stats().entries, 0);
65/// ```
66pub struct FftPlanCache {
67    store: ExtensionCacheStore,
68}
69
70impl fmt::Debug for FftPlanCache {
71    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
72        f.debug_struct("FftPlanCache")
73            .field("capacity", &self.capacity())
74            .field("stats", &self.stats())
75            .finish_non_exhaustive()
76    }
77}
78
79impl FftPlanCache {
80    /// Create an empty typed cache with an explicit maximum entry count.
81    pub fn with_capacity(capacity: NonZeroUsize) -> Self {
82        Self {
83            store: ExtensionCacheStore::with_limits(ExtensionCacheLimits::new(capacity)),
84        }
85    }
86
87    /// Maximum number of retained backend entries across all namespaces.
88    pub fn capacity(&self) -> NonZeroUsize {
89        self.store.limits().max_entries()
90    }
91
92    /// Return complete cache retention limits.
93    pub fn limits(&self) -> ExtensionCacheLimits {
94        self.store.limits()
95    }
96
97    /// Replace complete cache retention limits.
98    pub fn set_limits(&mut self, limits: ExtensionCacheLimits) {
99        self.store.set_limits(limits);
100    }
101
102    /// Resize the cache, evicting least-recently-used entries when necessary.
103    pub fn set_capacity(&mut self, capacity: NonZeroUsize) {
104        let mut limits = ExtensionCacheLimits::new(capacity);
105        if let Some(max_retained_bytes) = self.store.limits().max_retained_bytes() {
106            limits = limits.with_max_retained_bytes(max_retained_bytes);
107        }
108        self.store.set_limits(limits);
109    }
110
111    /// Remove every retained backend plan or workspace.
112    pub fn clear(&mut self) {
113        self.store.clear();
114    }
115
116    /// Snapshot aggregate entries and backend-reported retained bytes.
117    pub fn stats(&self) -> CacheStats {
118        self.store.stats(ExtensionCacheSelector::All)
119    }
120
121    pub(crate) fn store_mut(&mut self) -> &mut ExtensionCacheStore {
122        &mut self.store
123    }
124
125    pub(crate) fn plan_f32(&mut self, len: usize, forward: bool) -> Arc<dyn Fft<f32>> {
126        ExtensionFftPlanCache::new(&mut self.store).plan_f32(len, forward)
127    }
128
129    pub(crate) fn plan_f64(&mut self, len: usize, forward: bool) -> Arc<dyn Fft<f64>> {
130        ExtensionFftPlanCache::new(&mut self.store).plan_f64(len, forward)
131    }
132
133    #[cfg(test)]
134    pub(crate) fn contains_f64(&mut self, len: usize, forward: bool) -> bool {
135        let key = FftPlanKey {
136            len,
137            forward,
138            dtype: FftPlanDType::F64,
139        };
140        self.store
141            .get::<ExtensionFftPlanEntry>(&extension_plan_key(key))
142            .is_some_and(|entry| entry.matches_f64(key))
143    }
144}
145
146impl Default for FftPlanCache {
147    fn default() -> Self {
148        Self::with_capacity(
149            NonZeroUsize::new(DEFAULT_FFT_PLAN_CACHE_CAPACITY).unwrap_or(NonZeroUsize::MIN),
150        )
151    }
152}
153
154impl RuntimeCacheControl for FftPlanCache {
155    fn clear(&mut self) {
156        Self::clear(self);
157    }
158
159    fn stats(&self) -> CacheStats {
160        Self::stats(self)
161    }
162}
163
164#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
165enum FftPlanDType {
166    F32,
167    F64,
168}
169
170#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
171struct FftPlanKey {
172    len: usize,
173    forward: bool,
174    dtype: FftPlanDType,
175}
176
177enum CachedFftPlan {
178    F32(Arc<dyn Fft<f32>>),
179    F64(Arc<dyn Fft<f64>>),
180}
181
182struct ExtensionFftPlanEntry {
183    key: FftPlanKey,
184    plan: CachedFftPlan,
185}
186
187impl ExtensionFftPlanEntry {
188    #[cfg(test)]
189    fn matches_f64(&self, key: FftPlanKey) -> bool {
190        self.key == key && matches!(self.plan, CachedFftPlan::F64(_))
191    }
192}
193
194pub(crate) trait FftPlanProvider: Send {
195    fn plan_f32(&mut self, len: usize, forward: bool) -> Arc<dyn Fft<f32>>;
196    fn plan_f64(&mut self, len: usize, forward: bool) -> Arc<dyn Fft<f64>>;
197}
198
199impl FftPlanProvider for FftPlanCache {
200    fn plan_f32(&mut self, len: usize, forward: bool) -> Arc<dyn Fft<f32>> {
201        Self::plan_f32(self, len, forward)
202    }
203
204    fn plan_f64(&mut self, len: usize, forward: bool) -> Arc<dyn Fft<f64>> {
205        Self::plan_f64(self, len, forward)
206    }
207}
208
209pub(crate) trait CachedFftPlanScalar: FftNum + Float + FromPrimitive + 'static {
210    fn plan<P: FftPlanProvider + ?Sized>(
211        plans: &mut P,
212        len: usize,
213        forward: bool,
214    ) -> Arc<dyn Fft<Self>>;
215}
216
217impl CachedFftPlanScalar for f32 {
218    fn plan<P: FftPlanProvider + ?Sized>(
219        plans: &mut P,
220        len: usize,
221        forward: bool,
222    ) -> Arc<dyn Fft<Self>> {
223        plans.plan_f32(len, forward)
224    }
225}
226
227impl CachedFftPlanScalar for f64 {
228    fn plan<P: FftPlanProvider + ?Sized>(
229        plans: &mut P,
230        len: usize,
231        forward: bool,
232    ) -> Arc<dyn Fft<Self>> {
233        plans.plan_f64(len, forward)
234    }
235}
236
237pub(crate) fn cached_fft_plan<T: CachedFftPlanScalar, P: FftPlanProvider + ?Sized>(
238    plans: &mut P,
239    len: usize,
240    forward: bool,
241) -> Arc<dyn Fft<T>> {
242    T::plan(plans, len, forward)
243}
244
245pub(crate) struct ExtensionFftPlanCache<'a> {
246    entries: &'a mut ExtensionCacheStore,
247}
248
249impl<'a> ExtensionFftPlanCache<'a> {
250    pub(crate) fn new(entries: &'a mut ExtensionCacheStore) -> Self {
251        Self { entries }
252    }
253}
254
255fn extension_plan_key(key: FftPlanKey) -> ExtensionCacheKey {
256    let mut hasher = DefaultHasher::new();
257    key.hash(&mut hasher);
258    ExtensionCacheKey::new(
259        FFT_EXTENSION_FAMILY_ID,
260        FFT_PLAN_CACHE_NAME,
261        hasher.finish(),
262    )
263}
264
265impl FftPlanProvider for ExtensionFftPlanCache<'_> {
266    fn plan_f32(&mut self, len: usize, forward: bool) -> Arc<dyn Fft<f32>> {
267        let key = FftPlanKey {
268            len,
269            forward,
270            dtype: FftPlanDType::F32,
271        };
272        let cache_key = extension_plan_key(key);
273        if let Some(cached) = self.entries.get::<ExtensionFftPlanEntry>(&cache_key) {
274            if cached.key == key {
275                if let CachedFftPlan::F32(plan) = &cached.plan {
276                    return Arc::clone(plan);
277                }
278            }
279        }
280        let plan = build_fft_plan::<f32>(len, forward);
281        self.entries.put(
282            cache_key,
283            ExtensionFftPlanEntry {
284                key,
285                plan: CachedFftPlan::F32(Arc::clone(&plan)),
286            },
287            fft_plan_retained_bytes(),
288        );
289        plan
290    }
291
292    fn plan_f64(&mut self, len: usize, forward: bool) -> Arc<dyn Fft<f64>> {
293        let key = FftPlanKey {
294            len,
295            forward,
296            dtype: FftPlanDType::F64,
297        };
298        let cache_key = extension_plan_key(key);
299        if let Some(cached) = self.entries.get::<ExtensionFftPlanEntry>(&cache_key) {
300            if cached.key == key {
301                if let CachedFftPlan::F64(plan) = &cached.plan {
302                    return Arc::clone(plan);
303                }
304            }
305        }
306        let plan = build_fft_plan::<f64>(len, forward);
307        self.entries.put(
308            cache_key,
309            ExtensionFftPlanEntry {
310                key,
311                plan: CachedFftPlan::F64(Arc::clone(&plan)),
312            },
313            fft_plan_retained_bytes(),
314        );
315        plan
316    }
317}
318
319fn build_fft_plan<T: FftNum + 'static>(len: usize, forward: bool) -> Arc<dyn Fft<T>> {
320    let mut planner = FftPlanner::<T>::new();
321    if forward {
322        planner.plan_fft_forward(len)
323    } else {
324        planner.plan_fft_inverse(len)
325    }
326}
327
328const fn fft_plan_retained_bytes() -> usize {
329    std::mem::size_of::<FftPlanKey>() + std::mem::size_of::<CachedFftPlan>()
330}