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
16pub const FFT_PLAN_CACHE_NAME: &str = "rustfft-plans";
18
19pub const DEFAULT_FFT_PLAN_CACHE_CAPACITY: usize = 64;
21
22pub 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
43pub 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 pub fn with_capacity(capacity: NonZeroUsize) -> Self {
82 Self {
83 store: ExtensionCacheStore::with_limits(ExtensionCacheLimits::new(capacity)),
84 }
85 }
86
87 pub fn capacity(&self) -> NonZeroUsize {
89 self.store.limits().max_entries()
90 }
91
92 pub fn limits(&self) -> ExtensionCacheLimits {
94 self.store.limits()
95 }
96
97 pub fn set_limits(&mut self, limits: ExtensionCacheLimits) {
99 self.store.set_limits(limits);
100 }
101
102 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 pub fn clear(&mut self) {
113 self.store.clear();
114 }
115
116 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}