Skip to main content

tenferro_ad/
transform_cache.rs

1use std::fmt;
2use std::mem::{size_of, size_of_val};
3use std::num::NonZeroUsize;
4use std::sync::{Arc, Mutex, MutexGuard};
5
6use lru::LruCache;
7use tenferro_runtime::program::{
8    FrozenProgram, ProgramValueMetadata, SemanticFingerprint, SemanticProgram,
9};
10use tenferro_runtime::{CacheStats, Error, ErrorPhase, Result};
11
12use crate::semantic_transform::SemanticAdProgram;
13
14const DEFAULT_AD_TRANSFORM_CACHE_ENTRIES: usize = 128;
15const DEFAULT_AD_TRANSFORM_CACHE_RETAINED_BYTES: usize = 64 * 1024 * 1024;
16
17/// Retention limits for AD transform graph caches.
18///
19/// The retained-byte limit is a logical payload estimate, not process RSS.
20///
21/// # Examples
22///
23/// ```rust
24/// use std::num::NonZeroUsize;
25/// use tenferro_ad::AdTransformCacheLimits;
26///
27/// let limits = AdTransformCacheLimits::new(NonZeroUsize::new(4).unwrap());
28/// assert_eq!(limits.max_entries().get(), 4);
29/// assert!(limits.max_retained_bytes().is_some());
30/// ```
31#[derive(Clone, Copy, Debug, PartialEq, Eq)]
32pub struct AdTransformCacheLimits {
33    max_entries: NonZeroUsize,
34    max_retained_bytes: Option<NonZeroUsize>,
35}
36
37impl AdTransformCacheLimits {
38    /// Create AD transform cache limits with the default retained-byte bound.
39    ///
40    /// # Examples
41    ///
42    /// ```rust
43    /// use std::num::NonZeroUsize;
44    /// use tenferro_ad::AdTransformCacheLimits;
45    ///
46    /// let limits = AdTransformCacheLimits::new(NonZeroUsize::new(2).unwrap());
47    /// assert_eq!(limits.max_entries().get(), 2);
48    /// ```
49    pub fn new(max_entries: NonZeroUsize) -> Self {
50        Self {
51            max_entries,
52            max_retained_bytes: Some(
53                NonZeroUsize::new(DEFAULT_AD_TRANSFORM_CACHE_RETAINED_BYTES)
54                    .unwrap_or(NonZeroUsize::MIN),
55            ),
56        }
57    }
58
59    /// Return the maximum number of retained AD transform entries.
60    ///
61    /// # Examples
62    ///
63    /// ```rust
64    /// use tenferro_ad::AdTransformCacheLimits;
65    ///
66    /// assert!(AdTransformCacheLimits::default().max_entries().get() > 0);
67    /// ```
68    pub fn max_entries(self) -> NonZeroUsize {
69        self.max_entries
70    }
71
72    /// Return the logical retained-byte bound, when one is configured.
73    ///
74    /// # Examples
75    ///
76    /// ```rust
77    /// use tenferro_ad::AdTransformCacheLimits;
78    ///
79    /// assert!(AdTransformCacheLimits::default().max_retained_bytes().is_some());
80    /// ```
81    pub fn max_retained_bytes(self) -> Option<NonZeroUsize> {
82        self.max_retained_bytes
83    }
84
85    /// Return limits with a new logical retained-byte bound.
86    ///
87    /// # Examples
88    ///
89    /// ```rust
90    /// use std::num::NonZeroUsize;
91    /// use tenferro_ad::AdTransformCacheLimits;
92    ///
93    /// let limits = AdTransformCacheLimits::default()
94    ///     .with_max_retained_bytes(NonZeroUsize::new(1024).unwrap());
95    /// assert_eq!(limits.max_retained_bytes().unwrap().get(), 1024);
96    /// ```
97    pub fn with_max_retained_bytes(mut self, max_retained_bytes: NonZeroUsize) -> Self {
98        self.max_retained_bytes = Some(max_retained_bytes);
99        self
100    }
101}
102
103impl Default for AdTransformCacheLimits {
104    fn default() -> Self {
105        Self::new(
106            NonZeroUsize::new(DEFAULT_AD_TRANSFORM_CACHE_ENTRIES).unwrap_or(NonZeroUsize::MIN),
107        )
108    }
109}
110
111#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
112pub(crate) enum SemanticAdTransformKind {
113    Jvp,
114    Vjp,
115}
116
117#[derive(Clone, Debug, PartialEq, Eq, Hash)]
118pub(crate) struct SemanticAdTransformCacheKey {
119    kind: SemanticAdTransformKind,
120    input_fingerprint: SemanticFingerprint,
121    input_metadata: Box<[ProgramValueMetadata]>,
122    active_inputs: Box<[bool]>,
123    active_outputs: Box<[bool]>,
124}
125
126impl SemanticAdTransformCacheKey {
127    pub(crate) fn jvp(input: &FrozenProgram, active_inputs: &[bool]) -> Self {
128        Self {
129            kind: SemanticAdTransformKind::Jvp,
130            input_fingerprint: input.program.semantic_fingerprint(),
131            input_metadata: semantic_input_metadata(input),
132            active_inputs: active_inputs.into(),
133            active_outputs: Box::new([]),
134        }
135    }
136
137    pub(crate) fn vjp(
138        input: &FrozenProgram,
139        active_inputs: &[bool],
140        active_outputs: &[bool],
141    ) -> Self {
142        Self {
143            kind: SemanticAdTransformKind::Vjp,
144            input_fingerprint: input.program.semantic_fingerprint(),
145            input_metadata: semantic_input_metadata(input),
146            active_inputs: active_inputs.into(),
147            active_outputs: active_outputs.into(),
148        }
149    }
150}
151
152fn semantic_input_metadata(input: &FrozenProgram) -> Box<[ProgramValueMetadata]> {
153    input.input_metadata_with_bound_shapes()
154}
155
156#[derive(Clone)]
157struct CachedSemanticAdTransform {
158    input: Arc<SemanticProgram>,
159    output: Arc<SemanticAdProgram>,
160}
161
162#[derive(Debug)]
163pub(crate) struct AdTransformCache {
164    store: Mutex<AdTransformCacheStore>,
165}
166
167impl AdTransformCache {
168    pub(crate) fn new() -> Self {
169        Self {
170            store: Mutex::new(AdTransformCacheStore::default()),
171        }
172    }
173
174    pub(crate) fn limits(&self) -> Result<AdTransformCacheLimits> {
175        Ok(self.lock_store()?.limits)
176    }
177
178    pub(crate) fn set_limits(&self, limits: AdTransformCacheLimits) -> Result<()> {
179        self.lock_store()?.set_limits(limits);
180        Ok(())
181    }
182
183    pub(crate) fn clear(&self) -> Result<()> {
184        self.lock_store()?.clear();
185        Ok(())
186    }
187
188    pub(crate) fn stats(&self) -> Result<CacheStats> {
189        Ok(self.lock_store()?.stats())
190    }
191
192    pub(crate) fn get_semantic(
193        &self,
194        key: &SemanticAdTransformCacheKey,
195        input: &FrozenProgram,
196    ) -> Result<Option<Arc<SemanticAdProgram>>> {
197        Ok(self.lock_store()?.get_semantic(key, input))
198    }
199
200    pub(crate) fn put_semantic(
201        &self,
202        key: SemanticAdTransformCacheKey,
203        input: &FrozenProgram,
204        output: Arc<SemanticAdProgram>,
205    ) -> Result<()> {
206        self.lock_store()?.put_semantic(key, input, output);
207        Ok(())
208    }
209
210    fn lock_store(&self) -> Result<MutexGuard<'_, AdTransformCacheStore>> {
211        self.store.lock().map_err(|_| {
212            Error::runtime_state("ad_transform_cache", ErrorPhase::Compile, "lock poisoned")
213        })
214    }
215}
216
217#[derive(Debug)]
218struct AdTransformCacheStore {
219    limits: AdTransformCacheLimits,
220    entries: LruCache<AdTransformCacheKey, AdTransformCacheEntryWithStats>,
221    stats: CacheStats,
222}
223
224impl AdTransformCacheStore {
225    fn set_limits(&mut self, limits: AdTransformCacheLimits) {
226        self.limits = limits;
227        self.evict_to_limits();
228    }
229
230    fn clear(&mut self) {
231        let clears = self.stats.clears.saturating_add(1);
232        self.entries.clear();
233        self.stats = CacheStats {
234            clears,
235            ..CacheStats::empty()
236        };
237    }
238
239    fn stats(&self) -> CacheStats {
240        self.stats
241    }
242
243    fn get_semantic(
244        &mut self,
245        key: &SemanticAdTransformCacheKey,
246        input: &FrozenProgram,
247    ) -> Option<Arc<SemanticAdProgram>> {
248        let cache_key = AdTransformCacheKey::Semantic(key.clone());
249        let result = self
250            .entries
251            .get(&cache_key)
252            .and_then(|entry| match &entry.entry {
253                AdTransformCacheEntry::Semantic(bucket) => bucket
254                    .iter()
255                    .find(|cached| cached.input.semantic_eq(input.program.as_ref()))
256                    .map(|cached| Arc::clone(&cached.output)),
257            });
258        if result.is_some() {
259            self.stats.hits = self.stats.hits.saturating_add(1);
260        } else {
261            self.stats.misses = self.stats.misses.saturating_add(1);
262        }
263        result
264    }
265
266    fn put_semantic(
267        &mut self,
268        key: SemanticAdTransformCacheKey,
269        input: &FrozenProgram,
270        output: Arc<SemanticAdProgram>,
271    ) {
272        let cache_key = AdTransformCacheKey::Semantic(key);
273        let mut bucket = match self.entries.pop(&cache_key) {
274            Some(entry) => {
275                self.stats.retained_bytes = self
276                    .stats
277                    .retained_bytes
278                    .saturating_sub(entry.retained_bytes);
279                match entry.entry {
280                    AdTransformCacheEntry::Semantic(bucket) => bucket,
281                }
282            }
283            None => Vec::new(),
284        };
285        if let Some(cached) = bucket
286            .iter_mut()
287            .find(|cached| cached.input.semantic_eq(input.program.as_ref()))
288        {
289            cached.output = output;
290        } else {
291            bucket.push(CachedSemanticAdTransform {
292                input: Arc::clone(&input.program),
293                output,
294            });
295        }
296        self.put_entry(cache_key, AdTransformCacheEntry::Semantic(bucket));
297    }
298
299    fn put_entry(&mut self, key: AdTransformCacheKey, entry: AdTransformCacheEntry) {
300        let retained_bytes = ad_transform_cache_entry_retained_bytes(&key, &entry);
301        let entry = AdTransformCacheEntryWithStats {
302            entry,
303            retained_bytes,
304        };
305        self.stats.entries = self.entries.len();
306        self.stats.retained_bytes = self.stats.retained_bytes.saturating_add(retained_bytes);
307        if let Some((_old_key, old_entry)) = self.entries.push(key, entry) {
308            self.stats.retained_bytes = self
309                .stats
310                .retained_bytes
311                .saturating_sub(old_entry.retained_bytes);
312        }
313        self.stats.entries = self.entries.len();
314        self.evict_to_limits();
315    }
316
317    fn evict_to_limits(&mut self) {
318        while self.entries.len() > self.limits.max_entries.get()
319            || self
320                .limits
321                .max_retained_bytes
322                .is_some_and(|limit| self.stats.retained_bytes > limit.get())
323        {
324            let Some((_key, entry)) = self.entries.pop_lru() else {
325                break;
326            };
327            self.stats.retained_bytes = self
328                .stats
329                .retained_bytes
330                .saturating_sub(entry.retained_bytes);
331            self.stats.evictions = self.stats.evictions.saturating_add(1);
332        }
333        self.stats.entries = self.entries.len();
334    }
335}
336
337impl Default for AdTransformCacheStore {
338    fn default() -> Self {
339        Self {
340            limits: AdTransformCacheLimits::default(),
341            entries: LruCache::unbounded(),
342            stats: CacheStats::empty(),
343        }
344    }
345}
346
347#[derive(Clone, Debug, PartialEq, Eq, Hash)]
348enum AdTransformCacheKey {
349    Semantic(SemanticAdTransformCacheKey),
350}
351
352enum AdTransformCacheEntry {
353    Semantic(Vec<CachedSemanticAdTransform>),
354}
355
356impl fmt::Debug for AdTransformCacheEntry {
357    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
358        match self {
359            Self::Semantic(bucket) => f
360                .debug_tuple("Semantic")
361                .field(&format_args!("{} entries", bucket.len()))
362                .finish(),
363        }
364    }
365}
366
367#[derive(Debug)]
368struct AdTransformCacheEntryWithStats {
369    entry: AdTransformCacheEntry,
370    retained_bytes: usize,
371}
372
373fn ad_transform_cache_entry_retained_bytes(
374    key: &AdTransformCacheKey,
375    entry: &AdTransformCacheEntry,
376) -> usize {
377    size_of::<AdTransformCacheKey>()
378        + ad_transform_cache_key_retained_bytes(key)
379        + size_of::<AdTransformCacheEntry>()
380        + ad_transform_cache_value_retained_bytes(entry)
381}
382
383fn ad_transform_cache_key_retained_bytes(key: &AdTransformCacheKey) -> usize {
384    match key {
385        AdTransformCacheKey::Semantic(key) => {
386            key.input_metadata.len() * size_of::<ProgramValueMetadata>()
387                + key.active_inputs.len() * size_of::<bool>()
388                + key.active_outputs.len() * size_of::<bool>()
389        }
390    }
391}
392
393fn ad_transform_cache_value_retained_bytes(entry: &AdTransformCacheEntry) -> usize {
394    match entry {
395        AdTransformCacheEntry::Semantic(bucket) => {
396            size_of_val(bucket.as_slice())
397                + bucket
398                    .iter()
399                    .map(|cached| {
400                        size_of::<CachedSemanticAdTransform>()
401                            + semantic_program_retained_bytes(cached.input.as_ref())
402                            + semantic_program_retained_bytes(
403                                cached.output.frozen().program.as_ref(),
404                            )
405                            + size_of_val(cached.output.derivative_input_indices())
406                            + size_of_val(cached.output.derivative_output_indices())
407                    })
408                    .fold(0usize, usize::saturating_add)
409        }
410    }
411}
412
413fn semantic_program_retained_bytes(program: &SemanticProgram) -> usize {
414    size_of::<SemanticProgram>()
415        + size_of_val(program.inputs())
416        + size_of_val(program.outputs())
417        + program.operations().len() * size_of::<usize>()
418        + program.shape_guards().len() * size_of::<usize>()
419}