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;