Skip to main content

tensor4all_core/cached_function/
mod.rs

1//! Cached function wrapper for expensive function evaluations.
2//!
3//! [`CachedFunction`] wraps a user-supplied function `Fn(&[I]) -> V` with
4//! thread-safe memoization. On first call for a given multi-index, the
5//! function is evaluated and the result is cached; subsequent calls return
6//! the cached value.
7//!
8//! The internal cache key type is automatically selected based on the total
9//! index space size (up to 1024 bits by default). For larger index spaces,
10//! use [`CachedFunction::with_key_type`] with a custom [`CacheKey`]
11//! implementation.
12//!
13//! # Examples
14//!
15//! ```
16//! use tensor4all_core::CachedFunction;
17//!
18//! let cf = CachedFunction::new(
19//!     |idx: &[usize]| idx[0] + idx[1],
20//!     &[3, 4],
21//! ).unwrap();
22//!
23//! assert_eq!(cf.eval(&[1, 2]).unwrap(), 3);
24//! assert_eq!(cf.num_evals(), 1);
25//! assert_eq!(cf.eval(&[1, 2]).unwrap(), 3);
26//! assert_eq!(cf.num_cache_hits(), 1);
27//! ```
28
29pub mod cache_key;
30pub mod error;
31pub mod index_int;
32pub mod multi_index_cache;
33
34use std::collections::HashMap;
35use std::sync::atomic::{AtomicUsize, Ordering};
36use std::sync::{RwLock, RwLockReadGuard, RwLockWriteGuard};
37
38use bnum::types::{U1024, U256, U512};
39
40use cache_key::CacheKey;
41use index_int::IndexInt;
42
43/// Compute total bits needed to represent the index space.
44pub(crate) fn total_bits(local_dims: &[usize]) -> u32 {
45    local_dims
46        .iter()
47        .map(|&d| {
48            if d <= 1 {
49                0
50            } else {
51                ((d - 1) as u64).ilog2() + 1
52            }
53        })
54        .sum()
55}
56
57/// Compute mixed-radix coefficients for flat index computation.
58///
59/// Returns `Err(CacheKeyError::Overflow)` if the index space overflows key type `K`.
60pub(crate) fn compute_coeffs<K: CacheKey>(
61    local_dims: &[usize],
62) -> Result<Vec<K>, error::CacheKeyError> {
63    let bits = total_bits(local_dims);
64    if bits > K::BITS_COUNT {
65        return Err(error::CacheKeyError::Overflow {
66            total_bits: bits,
67            max_bits: K::BITS_COUNT,
68            key_type: std::any::type_name::<K>(),
69        });
70    }
71
72    // INVARIANT: for an index space with cardinality greater than one, each
73    // coefficient consumed by a valid multi-index is at most `cardinality / 2`
74    // and the trailing coordinates only ever contribute zero; the `bits` guard
75    // above caps the cardinality at `2^K`.
76    let last_bond = local_dims.iter().rposition(|&d| d > 1);
77    let mut coeffs = Vec::with_capacity(local_dims.len());
78    let mut prod = K::ONE;
79    for (i, &d) in local_dims.iter().enumerate() {
80        coeffs.push(prod.clone());
81        if Some(i) < last_bond {
82            // Detects coefficient overflow if `total_bits` ever stops bounding
83            // the coefficients this loop produces.
84            let dim = K::from_usize(d);
85            prod = prod
86                .checked_mul(dim)
87                .ok_or_else(|| error::CacheKeyError::Overflow {
88                    total_bits: bits,
89                    max_bits: K::BITS_COUNT,
90                    key_type: std::any::type_name::<K>(),
91                })?;
92        }
93    }
94
95    Ok(coeffs)
96}
97
98fn read_lock<T>(lock: &RwLock<T>) -> RwLockReadGuard<'_, T> {
99    lock.read().unwrap_or_else(|poisoned| poisoned.into_inner())
100}
101
102fn write_lock<T>(lock: &RwLock<T>) -> RwLockWriteGuard<'_, T> {
103    lock.write()
104        .unwrap_or_else(|poisoned| poisoned.into_inner())
105}
106
107/// Compute flat index from multi-index and coefficients.
108fn flat_index<K: CacheKey, I: IndexInt>(idx: &[I], coeffs: &[K]) -> K {
109    idx.iter().zip(coeffs).fold(K::ZERO, |acc, (&i, c)| {
110        acc.wrapping_add(
111            c.clone()
112                .checked_mul(K::from_usize(i.to_usize()))
113                .unwrap_or(K::ZERO),
114        )
115    })
116}
117
118/// Internal cache with automatically selected key type.
119enum InnerCache<V> {
120    U64 {
121        cache: RwLock<HashMap<u64, V>>,
122        coeffs: Vec<u64>,
123    },
124    U128 {
125        cache: RwLock<HashMap<u128, V>>,
126        coeffs: Vec<u128>,
127    },
128    U256 {
129        cache: RwLock<HashMap<U256, V>>,
130        coeffs: Vec<U256>,
131    },
132    U512 {
133        cache: RwLock<HashMap<U512, V>>,
134        coeffs: Vec<U512>,
135    },
136    U1024 {
137        cache: RwLock<HashMap<U1024, V>>,
138        coeffs: Vec<U1024>,
139    },
140}
141
142impl<V: Clone + Send + Sync> InnerCache<V> {
143    /// Size in bytes of one key of the selected key type.
144    fn key_bytes(&self) -> usize {
145        match self {
146            Self::U64 { .. } => std::mem::size_of::<u64>(),
147            Self::U128 { .. } => std::mem::size_of::<u128>(),
148            Self::U256 { .. } => std::mem::size_of::<U256>(),
149            Self::U512 { .. } => std::mem::size_of::<U512>(),
150            Self::U1024 { .. } => std::mem::size_of::<U1024>(),
151        }
152    }
153
154    /// Create a new cache, automatically selecting the key type.
155    fn new(local_dims: &[usize]) -> Result<Self, error::CacheKeyError> {
156        let bits = total_bits(local_dims);
157        if bits <= 64 {
158            Ok(Self::U64 {
159                cache: RwLock::new(HashMap::new()),
160                coeffs: compute_coeffs::<u64>(local_dims)?,
161            })
162        } else if bits <= 128 {
163            Ok(Self::U128 {
164                cache: RwLock::new(HashMap::new()),
165                coeffs: compute_coeffs::<u128>(local_dims)?,
166            })
167        } else if bits <= 256 {
168            Ok(Self::U256 {
169                cache: RwLock::new(HashMap::new()),
170                coeffs: compute_coeffs::<U256>(local_dims)?,
171            })
172        } else if bits <= 512 {
173            Ok(Self::U512 {
174                cache: RwLock::new(HashMap::new()),
175                coeffs: compute_coeffs::<U512>(local_dims)?,
176            })
177        } else if bits <= 1024 {
178            Ok(Self::U1024 {
179                cache: RwLock::new(HashMap::new()),
180                coeffs: compute_coeffs::<U1024>(local_dims)?,
181            })
182        } else {
183            Err(error::CacheKeyError::Overflow {
184                total_bits: bits,
185                max_bits: 1024,
186                key_type: "auto",
187            })
188        }
189    }
190
191    fn get<I: IndexInt>(&self, idx: &[I]) -> Option<V> {
192        match self {
193            Self::U64 { cache, coeffs } => read_lock(cache).get(&flat_index(idx, coeffs)).cloned(),
194            Self::U128 { cache, coeffs } => read_lock(cache).get(&flat_index(idx, coeffs)).cloned(),
195            Self::U256 { cache, coeffs } => read_lock(cache).get(&flat_index(idx, coeffs)).cloned(),
196            Self::U512 { cache, coeffs } => read_lock(cache).get(&flat_index(idx, coeffs)).cloned(),
197            Self::U1024 { cache, coeffs } => {
198                read_lock(cache).get(&flat_index(idx, coeffs)).cloned()
199            }
200        }
201    }
202
203    fn insert<I: IndexInt>(&self, idx: &[I], value: V) {
204        match self {
205            Self::U64 { cache, coeffs } => {
206                write_lock(cache).insert(flat_index(idx, coeffs), value);
207            }
208            Self::U128 { cache, coeffs } => {
209                write_lock(cache).insert(flat_index(idx, coeffs), value);
210            }
211            Self::U256 { cache, coeffs } => {
212                write_lock(cache).insert(flat_index(idx, coeffs), value);
213            }
214            Self::U512 { cache, coeffs } => {
215                write_lock(cache).insert(flat_index(idx, coeffs), value);
216            }
217            Self::U1024 { cache, coeffs } => {
218                write_lock(cache).insert(flat_index(idx, coeffs), value);
219            }
220        }
221    }
222
223    fn contains<I: IndexInt>(&self, idx: &[I]) -> bool {
224        match self {
225            Self::U64 { cache, coeffs } => read_lock(cache).contains_key(&flat_index(idx, coeffs)),
226            Self::U128 { cache, coeffs } => read_lock(cache).contains_key(&flat_index(idx, coeffs)),
227            Self::U256 { cache, coeffs } => read_lock(cache).contains_key(&flat_index(idx, coeffs)),
228            Self::U512 { cache, coeffs } => read_lock(cache).contains_key(&flat_index(idx, coeffs)),
229            Self::U1024 { cache, coeffs } => {
230                read_lock(cache).contains_key(&flat_index(idx, coeffs))
231            }
232        }
233    }
234
235    fn len(&self) -> usize {
236        match self {
237            Self::U64 { cache, .. } => read_lock(cache).len(),
238            Self::U128 { cache, .. } => read_lock(cache).len(),
239            Self::U256 { cache, .. } => read_lock(cache).len(),
240            Self::U512 { cache, .. } => read_lock(cache).len(),
241            Self::U1024 { cache, .. } => read_lock(cache).len(),
242        }
243    }
244
245    fn clear(&self) {
246        match self {
247            Self::U64 { cache, .. } => write_lock(cache).clear(),
248            Self::U128 { cache, .. } => write_lock(cache).clear(),
249            Self::U256 { cache, .. } => write_lock(cache).clear(),
250            Self::U512 { cache, .. } => write_lock(cache).clear(),
251            Self::U1024 { cache, .. } => write_lock(cache).clear(),
252        }
253    }
254
255    fn key_type_name(&self) -> &'static str {
256        match self {
257            Self::U64 { .. } => "u64",
258            Self::U128 { .. } => "u128",
259            Self::U256 { .. } => "U256",
260            Self::U512 { .. } => "U512",
261            Self::U1024 { .. } => "U1024",
262        }
263    }
264}
265
266/// Type-erased cache interface for custom key types.
267trait DynCache<V>: Send + Sync {
268    fn get(&self, idx: &[usize]) -> Option<V>;
269    fn insert(&self, idx: &[usize], value: V);
270    fn contains(&self, idx: &[usize]) -> bool;
271    fn len(&self) -> usize;
272    fn clear(&self);
273}
274
275/// Generic cache for user-specified key types.
276struct GenericCache<K: CacheKey, V> {
277    cache: RwLock<HashMap<K, V>>,
278    coeffs: Vec<K>,
279}
280
281impl<K: CacheKey, V: Clone + Send + Sync> GenericCache<K, V> {
282    fn new(local_dims: &[usize]) -> Result<Self, error::CacheKeyError> {
283        Ok(Self {
284            cache: RwLock::new(HashMap::new()),
285            coeffs: compute_coeffs::<K>(local_dims)?,
286        })
287    }
288}
289
290impl<K: CacheKey, V: Clone + Send + Sync> DynCache<V> for GenericCache<K, V> {
291    fn get(&self, idx: &[usize]) -> Option<V> {
292        let key = flat_index::<K, usize>(idx, &self.coeffs);
293        read_lock(&self.cache).get(&key).cloned()
294    }
295
296    fn insert(&self, idx: &[usize], value: V) {
297        let key = flat_index::<K, usize>(idx, &self.coeffs);
298        write_lock(&self.cache).insert(key, value);
299    }
300
301    fn contains(&self, idx: &[usize]) -> bool {
302        let key = flat_index::<K, usize>(idx, &self.coeffs);
303        read_lock(&self.cache).contains_key(&key)
304    }
305
306    fn len(&self) -> usize {
307        read_lock(&self.cache).len()
308    }
309
310    fn clear(&self) {
311        write_lock(&self.cache).clear();
312    }
313}
314
315/// Internal backend: auto-selected enum or custom type-erased cache.
316enum CacheBackend<V: Clone + Send + Sync + 'static> {
317    Auto(InnerCache<V>),
318    Custom(Box<dyn DynCache<V>>),
319}
320
321impl<V: Clone + Send + Sync + 'static> CacheBackend<V> {
322    fn get<I: IndexInt>(&self, idx: &[I]) -> Option<V> {
323        match self {
324            Self::Auto(inner) => inner.get(idx),
325            Self::Custom(cache) => {
326                let usize_idx: Vec<usize> = idx.iter().map(|&i| i.to_usize()).collect();
327                cache.get(&usize_idx)
328            }
329        }
330    }
331
332    fn insert<I: IndexInt>(&self, idx: &[I], value: V) {
333        match self {
334            Self::Auto(inner) => inner.insert(idx, value),
335            Self::Custom(cache) => {
336                let usize_idx: Vec<usize> = idx.iter().map(|&i| i.to_usize()).collect();
337                cache.insert(&usize_idx, value);
338            }
339        }
340    }
341
342    fn contains<I: IndexInt>(&self, idx: &[I]) -> bool {
343        match self {
344            Self::Auto(inner) => inner.contains(idx),
345            Self::Custom(cache) => {
346                let usize_idx: Vec<usize> = idx.iter().map(|&i| i.to_usize()).collect();
347                cache.contains(&usize_idx)
348            }
349        }
350    }
351
352    fn len(&self) -> usize {
353        match self {
354            Self::Auto(inner) => inner.len(),
355            Self::Custom(cache) => cache.len(),
356        }
357    }
358
359    fn clear(&self) {
360        match self {
361            Self::Auto(inner) => inner.clear(),
362            Self::Custom(cache) => cache.clear(),
363        }
364    }
365
366    /// Size in bytes of one key; the type-erased custom backend does not report one.
367    fn key_bytes(&self) -> usize {
368        match self {
369            Self::Auto(inner) => inner.key_bytes(),
370            Self::Custom(_) => 0,
371        }
372    }
373
374    fn key_type_name(&self) -> &'static str {
375        match self {
376            Self::Auto(inner) => inner.key_type_name(),
377            Self::Custom(_) => "custom",
378        }
379    }
380}
381
382type BatchFunc<I, V> = dyn Fn(&[Vec<I>]) -> Vec<V> + Send + Sync;
383
384/// A wrapper that caches function evaluations for multi-index inputs.
385///
386/// Thread-safe: all methods take `&self`. Multiple threads can call `eval`
387/// concurrently.
388///
389/// # Type parameters
390///
391/// - `V` - cached value type
392/// - `F` - single-evaluation function `Fn(&[I]) -> V`
393/// - `I` - index element type (default `usize`); use `u8` for quantics
394///
395/// # Examples
396///
397/// ```
398/// use tensor4all_core::CachedFunction;
399///
400/// // Cache a 2-site function with local dimensions [3, 4]
401/// let cf = CachedFunction::new(
402///     |idx: &[usize]| (idx[0] * 4 + idx[1]) as f64,
403///     &[3, 4],
404/// ).unwrap();
405///
406/// // First call evaluates and caches
407/// let v00 = cf.eval(&[0, 0]).unwrap();
408/// assert_eq!(v00, 0.0);
409/// assert_eq!(cf.num_evals(), 1);
410/// assert_eq!(cf.num_cache_hits(), 0);
411///
412/// // Second call uses cache
413/// let v00_again = cf.eval(&[0, 0]).unwrap();
414/// assert_eq!(v00_again, 0.0);
415/// assert_eq!(cf.num_cache_hits(), 1);
416///
417/// let v12 = cf.eval(&[1, 2]).unwrap();
418/// assert_eq!(v12, 6.0); // 1*4 + 2
419/// ```
420pub struct CachedFunction<V, F, I = usize>
421where
422    I: IndexInt,
423    V: Clone + Send + Sync + 'static,
424    F: Fn(&[I]) -> V + Send + Sync,
425{
426    func: F,
427    batch_func: Option<Box<BatchFunc<I, V>>>,
428    cache: CacheBackend<V>,
429    local_dims: Vec<usize>,
430    num_evals: AtomicUsize,
431    num_cache_hits: AtomicUsize,
432    _phantom: std::marker::PhantomData<I>,
433}
434
435impl<V, F, I> CachedFunction<V, F, I>
436where
437    I: IndexInt,
438    V: Clone + Send + Sync + 'static,
439    F: Fn(&[I]) -> V + Send + Sync,
440{
441    /// Create a new cached function with automatic key selection (up to 1024 bits).
442    ///
443    /// # Errors
444    ///
445    /// Returns an error when the cache dimensions are invalid (a shape mismatch)
446    /// /// or the construction fails.
447    ///
448    /// # Examples
449    ///
450    /// ```
451    /// use tensor4all_core::CachedFunction;
452    ///
453    /// let cf = CachedFunction::new(|idx: &[usize]| idx[0] + idx[1], &[2, 3]).unwrap();
454    /// assert_eq!(cf.eval(&[1, 2]).unwrap(), 3);
455    /// assert_eq!(cf.num_sites(), 2);
456    /// assert_eq!(cf.local_dims(), &[2, 3]);
457    /// ```
458    pub fn new(func: F, local_dims: &[usize]) -> Result<Self, error::CacheKeyError> {
459        Ok(Self {
460            func,
461            batch_func: None,
462            cache: CacheBackend::Auto(InnerCache::new(local_dims)?),
463            local_dims: local_dims.to_vec(),
464            num_evals: AtomicUsize::new(0),
465            num_cache_hits: AtomicUsize::new(0),
466            _phantom: std::marker::PhantomData,
467        })
468    }
469
470    /// Create with a batch function for efficient multi-point evaluation.
471    ///
472    /// The batch function is used for cache misses during [`eval_batch`](Self::eval_batch)
473    /// calls, enabling amortized cost when evaluating many indices at once
474    /// (e.g., batch FFI calls or vectorized computations).
475    ///
476    /// # Errors
477    ///
478    /// Returns an error when the batch configuration is invalid (an
479    /// /// invalid-configuration failure).
480    ///
481    /// # Examples
482    ///
483    /// ```
484    /// use tensor4all_core::CachedFunction;
485    ///
486    /// let cf = CachedFunction::with_batch(
487    ///     |idx: &[usize]| idx[0] * 10 + idx[1],
488    ///     |indices: &[Vec<usize>]| indices.iter().map(|idx| idx[0] * 10 + idx[1]).collect(),
489    ///     &[3, 4],
490    /// ).unwrap();
491    ///
492    /// let results = cf.eval_batch(&[vec![0, 1], vec![2, 3]]).unwrap();
493    /// assert_eq!(results, vec![1, 23]);
494    /// assert_eq!(cf.num_evals(), 2);
495    /// ```
496    pub fn with_batch<B>(
497        func: F,
498        batch_func: B,
499        local_dims: &[usize],
500    ) -> Result<Self, error::CacheKeyError>
501    where
502        B: Fn(&[Vec<I>]) -> Vec<V> + Send + Sync + 'static,
503    {
504        Ok(Self {
505            func,
506            batch_func: Some(Box::new(batch_func)),
507            cache: CacheBackend::Auto(InnerCache::new(local_dims)?),
508            local_dims: local_dims.to_vec(),
509            num_evals: AtomicUsize::new(0),
510            num_cache_hits: AtomicUsize::new(0),
511            _phantom: std::marker::PhantomData,
512        })
513    }
514
515    /// Create with an explicit key type for index spaces larger than 1024 bits.
516    ///
517    /// # Errors
518    ///
519    /// Returns an error when the key type configuration is invalid (an
520    /// /// invalid-configuration failure).
521    ///
522    /// # Example
523    ///
524    /// ```
525    /// use bnum::types::U2048;
526    /// use tensor4all_core::{CacheKey, CachedFunction};
527    ///
528    /// #[derive(Clone, Hash, PartialEq, Eq)]
529    /// struct U2048Key(U2048);
530    ///
531    /// impl CacheKey for U2048Key {
532    ///     const BITS_COUNT: u32 = 2048;
533    ///     const ZERO: Self = Self(U2048::ZERO);
534    ///     const ONE: Self = Self(U2048::ONE);
535    ///
536    ///     fn from_usize(v: usize) -> Self {
537    ///         Self(U2048::from(v as u64))
538    ///     }
539    ///
540    ///     fn checked_mul(self, rhs: Self) -> Option<Self> {
541    ///         self.0.checked_mul(rhs.0).map(Self)
542    ///     }
543    ///
544    ///     fn wrapping_add(self, rhs: Self) -> Self {
545    ///         Self(self.0.wrapping_add(rhs.0))
546    ///     }
547    /// }
548    ///
549    /// let local_dims = vec![2usize; 1025];
550    /// let cf = CachedFunction::with_key_type::<U2048Key>(
551    ///     |idx: &[usize]| idx.iter().sum::<usize>(),
552    ///     &local_dims,
553    /// ).unwrap();
554    /// let zeros = vec![0usize; 1025];
555    ///
556    /// assert_eq!(cf.eval(&zeros).unwrap(), 0);
557    /// assert_eq!(cf.key_type(), "custom");
558    /// ```
559    pub fn with_key_type<K: CacheKey>(
560        func: F,
561        local_dims: &[usize],
562    ) -> Result<Self, error::CacheKeyError> {
563        Ok(Self {
564            func,
565            batch_func: None,
566            cache: CacheBackend::Custom(Box::new(GenericCache::<K, V>::new(local_dims)?)),
567            local_dims: local_dims.to_vec(),
568            num_evals: AtomicUsize::new(0),
569            num_cache_hits: AtomicUsize::new(0),
570            _phantom: std::marker::PhantomData,
571        })
572    }
573
574    /// Create with explicit key type and batch function.
575    ///
576    /// Combines [`with_key_type`](Self::with_key_type) and
577    /// [`with_batch`](Self::with_batch) for index spaces larger than 1024
578    /// bits that also benefit from batch evaluation.
579    ///
580    /// # Errors
581    ///
582    /// Returns an error when the configuration is invalid (an
583    /// /// invalid-configuration failure).
584    ///
585    /// # Examples
586    ///
587    /// ```
588    /// use tensor4all_core::CachedFunction;
589    ///
590    /// // Use u128 key type with batch support
591    /// let cf = CachedFunction::with_key_type_and_batch::<u128, _>(
592    ///     |idx: &[usize]| idx.iter().sum::<usize>(),
593    ///     |indices: &[Vec<usize>]| indices.iter().map(|idx| idx.iter().sum()).collect(),
594    ///     &[2, 3, 4],
595    /// ).unwrap();
596    ///
597    /// let results = cf.eval_batch(&[vec![0, 0, 0], vec![1, 2, 3]]).unwrap();
598    /// assert_eq!(results, vec![0, 6]);
599    /// ```
600    pub fn with_key_type_and_batch<K: CacheKey, B>(
601        func: F,
602        batch_func: B,
603        local_dims: &[usize],
604    ) -> Result<Self, error::CacheKeyError>
605    where
606        B: Fn(&[Vec<I>]) -> Vec<V> + Send + Sync + 'static,
607    {
608        Ok(Self {
609            func,
610            batch_func: Some(Box::new(batch_func)),
611            cache: CacheBackend::Custom(Box::new(GenericCache::<K, V>::new(local_dims)?)),
612            local_dims: local_dims.to_vec(),
613            num_evals: AtomicUsize::new(0),
614            num_cache_hits: AtomicUsize::new(0),
615            _phantom: std::marker::PhantomData,
616        })
617    }
618
619    /// Evaluate at a given index, using cache if available.
620    ///
621    /// On the first call for a given index, the wrapped function is invoked and
622    /// the result is cached. Subsequent calls with the same index return the
623    /// cached value. This method is thread-safe.
624    ///
625    /// # Errors
626    /// Returns [`error::CacheKeyError::InvalidIndexLength`] for a wrong-rank index or
627    /// [`error::CacheKeyError::IndexOutOfBounds`] for an invalid coordinate.
628    ///
629    /// # Examples
630    ///
631    /// ```
632    /// use tensor4all_core::CachedFunction;
633    ///
634    /// let cf = CachedFunction::new(|idx: &[usize]| idx[0] * idx[1], &[5, 5]).unwrap();
635    /// assert_eq!(cf.eval(&[3, 4]).unwrap(), 12);
636    /// assert_eq!(cf.num_evals(), 1);
637    ///
638    /// // Cache hit
639    /// assert_eq!(cf.eval(&[3, 4]).unwrap(), 12);
640    /// assert_eq!(cf.num_evals(), 1);
641    /// assert_eq!(cf.num_cache_hits(), 1);
642    /// ```
643    pub fn eval(&self, idx: &[I]) -> Result<V, error::CacheKeyError> {
644        self.validate_index(idx)?;
645        if let Some(value) = self.cache.get(idx) {
646            self.num_cache_hits.fetch_add(1, Ordering::Relaxed);
647            return Ok(value);
648        }
649
650        self.num_evals.fetch_add(1, Ordering::Relaxed);
651        let value = (self.func)(idx);
652        self.cache.insert(idx, value.clone());
653        Ok(value)
654    }
655
656    /// Evaluate bypassing the cache.
657    ///
658    /// The result is neither read from nor stored in the cache, and
659    /// evaluation counters are not updated. Useful for verification or
660    /// when the caller intentionally wants a fresh evaluation.
661    ///
662    /// # Errors
663    /// Returns [`error::CacheKeyError::InvalidIndexLength`] for a wrong-rank index or
664    /// [`error::CacheKeyError::IndexOutOfBounds`] for an invalid coordinate.
665    ///
666    /// # Examples
667    ///
668    /// ```
669    /// use tensor4all_core::CachedFunction;
670    ///
671    /// let cf = CachedFunction::new(|idx: &[usize]| idx[0] + 1, &[4]).unwrap();
672    /// assert_eq!(cf.eval_no_cache(&[2]).unwrap(), 3);
673    /// assert_eq!(cf.cache_size(), 0);
674    /// assert_eq!(cf.num_evals(), 0);
675    /// ```
676    pub fn eval_no_cache(&self, idx: &[I]) -> Result<V, error::CacheKeyError> {
677        self.validate_index(idx)?;
678        Ok((self.func)(idx))
679    }
680
681    /// Evaluate at multiple indices. Uses batch function for cache misses if available.
682    ///
683    /// Returns results in the same order as the input indices.
684    ///
685    /// # Errors
686    /// Returns an index validation error for malformed coordinates or
687    /// [`error::CacheKeyError::BatchResultLength`] when a batch callback returns the
688    /// wrong number of values.
689    ///
690    /// # Examples
691    ///
692    /// ```
693    /// use tensor4all_core::CachedFunction;
694    ///
695    /// let cf = CachedFunction::new(|idx: &[usize]| idx[0] * 2 + idx[1], &[2, 2]).unwrap();
696    /// let results = cf.eval_batch(&[vec![0, 0], vec![0, 1], vec![1, 0]]).unwrap();
697    /// assert_eq!(results, vec![0, 1, 2]);
698    /// ```
699    pub fn eval_batch(&self, indices: &[Vec<I>]) -> Result<Vec<V>, error::CacheKeyError> {
700        if indices.is_empty() {
701            return Ok(Vec::new());
702        }
703        for index in indices {
704            self.validate_index(index)?;
705        }
706
707        let mut results: Vec<Option<V>> = Vec::with_capacity(indices.len());
708        let mut miss_positions: Vec<usize> = Vec::new();
709        let mut miss_indices: Vec<Vec<I>> = Vec::new();
710
711        for (pos, idx) in indices.iter().enumerate() {
712            if let Some(value) = self.cache.get(idx) {
713                self.num_cache_hits.fetch_add(1, Ordering::Relaxed);
714                results.push(Some(value));
715            } else {
716                results.push(None);
717                miss_positions.push(pos);
718                miss_indices.push(idx.clone());
719            }
720        }
721
722        if miss_indices.is_empty() {
723            return Ok(results.into_iter().flatten().collect());
724        }
725
726        self.num_evals
727            .fetch_add(miss_indices.len(), Ordering::Relaxed);
728        let miss_values = if let Some(batch_func) = self.batch_func.as_ref() {
729            batch_func(&miss_indices)
730        } else {
731            miss_indices.iter().map(|idx| (self.func)(idx)).collect()
732        };
733        if miss_values.len() != miss_indices.len() {
734            return Err(error::CacheKeyError::BatchResultLength {
735                expected: miss_indices.len(),
736                got: miss_values.len(),
737            });
738        }
739
740        for (i, pos) in miss_positions.iter().enumerate() {
741            self.cache.insert(&miss_indices[i], miss_values[i].clone());
742            results[*pos] = Some(miss_values[i].clone());
743        }
744
745        Ok(results.into_iter().flatten().collect())
746    }
747
748    fn validate_index(&self, idx: &[I]) -> std::result::Result<(), error::CacheKeyError> {
749        if idx.len() != self.local_dims.len() {
750            return Err(error::CacheKeyError::InvalidIndexLength {
751                expected: self.local_dims.len(),
752                got: idx.len(),
753            });
754        }
755        for (axis, (&value, &dim)) in idx.iter().zip(&self.local_dims).enumerate() {
756            let value = value.to_usize();
757            if value >= dim {
758                return Err(error::CacheKeyError::IndexOutOfBounds {
759                    axis,
760                    index: value,
761                    dim,
762                });
763            }
764        }
765        Ok(())
766    }
767
768    /// Get the local dimensions.
769    ///
770    /// # Examples
771    ///
772    /// ```
773    /// use tensor4all_core::CachedFunction;
774    ///
775    /// let cf = CachedFunction::new(|idx: &[usize]| 0, &[3, 4, 5]).unwrap();
776    /// assert_eq!(cf.local_dims(), &[3, 4, 5]);
777    /// ```
778    pub fn local_dims(&self) -> &[usize] {
779        &self.local_dims
780    }
781
782    /// Get the number of sites (length of the multi-index).
783    ///
784    /// # Examples
785    ///
786    /// ```
787    /// use tensor4all_core::CachedFunction;
788    ///
789    /// let cf = CachedFunction::new(|idx: &[usize]| 0, &[2, 3]).unwrap();
790    /// assert_eq!(cf.num_sites(), 2);
791    /// ```
792    pub fn num_sites(&self) -> usize {
793        self.local_dims.len()
794    }
795
796    /// Get the number of function evaluations (cache misses).
797    ///
798    /// # Examples
799    ///
800    /// ```
801    /// use tensor4all_core::CachedFunction;
802    ///
803    /// let cf = CachedFunction::new(|idx: &[usize]| idx[0], &[4]).unwrap();
804    /// cf.eval(&[0]).unwrap();
805    /// cf.eval(&[1]).unwrap();
806    /// cf.eval(&[0]).unwrap(); // cache hit, not a new eval
807    /// assert_eq!(cf.num_evals(), 2);
808    /// ```
809    pub fn num_evals(&self) -> usize {
810        self.num_evals.load(Ordering::Relaxed)
811    }
812
813    /// Get the number of cache hits.
814    ///
815    /// # Examples
816    ///
817    /// ```
818    /// use tensor4all_core::CachedFunction;
819    ///
820    /// let cf = CachedFunction::new(|idx: &[usize]| idx[0], &[4]).unwrap();
821    /// cf.eval(&[0]).unwrap();
822    /// assert_eq!(cf.num_cache_hits(), 0);
823    /// cf.eval(&[0]).unwrap();
824    /// assert_eq!(cf.num_cache_hits(), 1);
825    /// ```
826    pub fn num_cache_hits(&self) -> usize {
827        self.num_cache_hits.load(Ordering::Relaxed)
828    }
829
830    /// Get total calls (evaluations + cache hits).
831    ///
832    /// # Examples
833    ///
834    /// ```
835    /// use tensor4all_core::CachedFunction;
836    ///
837    /// let cf = CachedFunction::new(|idx: &[usize]| idx[0], &[4]).unwrap();
838    /// cf.eval(&[0]).unwrap();
839    /// cf.eval(&[1]).unwrap();
840    /// cf.eval(&[0]).unwrap(); // cache hit
841    /// assert_eq!(cf.total_calls(), 3);
842    /// assert_eq!(cf.total_calls(), cf.num_evals() + cf.num_cache_hits());
843    /// ```
844    pub fn total_calls(&self) -> usize {
845        self.num_evals() + self.num_cache_hits()
846    }
847
848    /// Get cache hit ratio (0.0 when no calls have been made).
849    ///
850    /// Returns `num_cache_hits() / total_calls()` as a value in `[0.0, 1.0]`.
851    ///
852    /// # Examples
853    ///
854    /// ```
855    /// use tensor4all_core::CachedFunction;
856    ///
857    /// let cf = CachedFunction::new(|idx: &[usize]| idx[0], &[4]).unwrap();
858    /// assert_eq!(cf.cache_hit_ratio(), 0.0); // no calls yet
859    ///
860    /// cf.eval(&[0]).unwrap();
861    /// cf.eval(&[0]).unwrap(); // cache hit
862    /// assert!((cf.cache_hit_ratio() - 0.5).abs() < 1e-10);
863    /// ```
864    pub fn cache_hit_ratio(&self) -> f64 {
865        let total = self.total_calls();
866        if total == 0 {
867            0.0
868        } else {
869            self.num_cache_hits() as f64 / total as f64
870        }
871    }
872
873    /// Clear the cache.
874    ///
875    /// # Examples
876    ///
877    /// ```
878    /// use tensor4all_core::CachedFunction;
879    ///
880    /// let cf = CachedFunction::new(|idx: &[usize]| idx[0], &[4]).unwrap();
881    /// cf.eval(&[2]).unwrap();
882    /// assert_eq!(cf.cache_size(), 1);
883    /// cf.clear_cache();
884    /// assert_eq!(cf.cache_size(), 0);
885    /// ```
886    pub fn clear_cache(&self) {
887        self.cache.clear();
888    }
889
890    /// Number of cached entries.
891    ///
892    /// # Examples
893    ///
894    /// ```
895    /// use tensor4all_core::CachedFunction;
896    ///
897    /// let cf = CachedFunction::new(|idx: &[usize]| idx[0], &[4]).unwrap();
898    /// assert_eq!(cf.cache_size(), 0);
899    /// cf.eval(&[0]).unwrap();
900    /// cf.eval(&[1]).unwrap();
901    /// assert_eq!(cf.cache_size(), 2);
902    /// cf.eval(&[0]).unwrap(); // cache hit, no new entry
903    /// assert_eq!(cf.cache_size(), 2);
904    /// ```
905    pub fn cache_size(&self) -> usize {
906        self.cache.len()
907    }
908
909    /// Check if an index is cached.
910    ///
911    /// # Examples
912    ///
913    /// ```
914    /// use tensor4all_core::CachedFunction;
915    ///
916    /// let cf = CachedFunction::new(|idx: &[usize]| idx[0], &[4]).unwrap();
917    /// assert!(!cf.is_cached(&[1]));
918    /// cf.eval(&[1]).unwrap();
919    /// assert!(cf.is_cached(&[1]));
920    /// ```
921    pub fn is_cached(&self, idx: &[I]) -> bool {
922        self.cache.contains(idx)
923    }
924
925    /// Internal key type name (for debugging).
926    ///
927    /// Returns `"u64"`, `"u128"`, `"U256"`, `"U512"`, `"U1024"` for
928    /// automatically selected types, or `"custom"` when constructed with
929    /// [`with_key_type`](Self::with_key_type).
930    ///
931    /// # Examples
932    ///
933    /// ```
934    /// use tensor4all_core::CachedFunction;
935    ///
936    /// // Small index space uses u64
937    /// let cf = CachedFunction::new(|idx: &[usize]| 0, &[2, 3]).unwrap();
938    /// assert_eq!(cf.key_type(), "u64");
939    /// ```
940    pub fn key_type(&self) -> &'static str {
941        self.cache.key_type_name()
942    }
943}
944
945#[cfg(test)]
946mod tests;