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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
32pub struct AdTransformCacheLimits {
33 max_entries: NonZeroUsize,
34 max_retained_bytes: Option<NonZeroUsize>,
35}
36
37impl AdTransformCacheLimits {
38 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 pub fn max_entries(self) -> NonZeroUsize {
69 self.max_entries
70 }
71
72 pub fn max_retained_bytes(self) -> Option<NonZeroUsize> {
82 self.max_retained_bytes
83 }
84
85 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}