use std::collections::{HashMap, HashSet};
use std::sync::{Arc, Mutex};
use super::runtimes::{LoadError, LoadedModel};
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct CacheKey {
pub digest: String,
pub runtime: &'static str,
pub device: String,
}
struct Entry {
model: Arc<dyn LoadedModel>,
bytes: u64,
last_used: u64,
}
#[derive(Default)]
struct Inner {
entries: HashMap<CacheKey, Entry>,
gone: HashSet<String>,
tick: u64,
}
pub struct LoadedCache {
max_bytes: u64,
inner: Mutex<Inner>,
flights: Mutex<HashMap<CacheKey, Arc<tokio::sync::Mutex<()>>>>,
}
impl LoadedCache {
pub fn new(max_bytes: u64) -> Self {
Self {
max_bytes,
inner: Mutex::new(Inner::default()),
flights: Mutex::new(HashMap::new()),
}
}
pub fn max_bytes(&self) -> u64 {
self.max_bytes
}
pub async fn get_or_load<F, Fut>(
&self,
key: CacheKey,
load: F,
) -> Result<(Arc<dyn LoadedModel>, bool), LoadError>
where
F: FnOnce() -> Fut,
Fut: std::future::Future<Output = Result<Arc<dyn LoadedModel>, LoadError>>,
{
if let Some(model) = self.hit(&key) {
return Ok((model, false));
}
let flight = self
.flights
.lock()
.unwrap_or_else(|e| e.into_inner())
.entry(key.clone())
.or_default()
.clone();
let _in_flight = flight.lock().await;
if let Some(model) = self.hit(&key) {
return Ok((model, true));
}
let model = load().await?;
self.insert(key, model.clone());
Ok((model, true))
}
fn hit(&self, key: &CacheKey) -> Option<Arc<dyn LoadedModel>> {
let mut inner = self.inner.lock().unwrap_or_else(|e| e.into_inner());
inner.tick += 1;
let tick = inner.tick;
let entry = inner.entries.get_mut(key)?;
entry.last_used = tick;
Some(entry.model.clone())
}
fn insert(&self, key: CacheKey, model: Arc<dyn LoadedModel>) {
let bytes = model.resident_bytes() as u64;
let mut inner = self.inner.lock().unwrap_or_else(|e| e.into_inner());
inner.tick += 1;
let tick = inner.tick;
if bytes > self.max_bytes {
tracing::warn!(
digest = %key.digest,
runtime = key.runtime,
resident_bytes = bytes,
max_loaded_bytes = self.max_bytes,
"Model is larger than models.max_loaded_bytes: served for this call and not \
kept resident"
);
inner.gone.insert(key.digest);
crate::metrics::set_model_loaded_bytes(total(&inner));
return;
}
inner.entries.insert(
key,
Entry {
model,
bytes,
last_used: tick,
},
);
while total(&inner) > self.max_bytes {
let Some(victim) = inner
.entries
.iter()
.min_by_key(|(_, e)| e.last_used)
.map(|(k, _)| k.clone())
else {
break;
};
if let Some(evicted) = inner.entries.remove(&victim) {
tracing::info!(
digest = %victim.digest,
runtime = victim.runtime,
resident_bytes = evicted.bytes,
"Model evicted: models.max_loaded_bytes reached"
);
}
inner.gone.insert(victim.digest);
}
crate::metrics::set_model_loaded_bytes(total(&inner));
}
pub fn states(&self) -> Vec<(CacheKey, u64)> {
let inner = self.inner.lock().unwrap_or_else(|e| e.into_inner());
let mut out: Vec<(CacheKey, u64)> = inner
.entries
.iter()
.map(|(k, e)| (k.clone(), e.bytes))
.collect();
out.sort_by(|a, b| {
a.0.digest
.cmp(&b.0.digest)
.then(a.0.runtime.cmp(b.0.runtime))
});
out
}
pub fn loaded_bytes(&self) -> u64 {
total(&self.inner.lock().unwrap_or_else(|e| e.into_inner()))
}
pub fn contains(&self, key: &CacheKey) -> bool {
self.inner
.lock()
.unwrap_or_else(|e| e.into_inner())
.entries
.contains_key(key)
}
pub fn loaded_for(&self, digest: &str) -> Option<(CacheKey, u64)> {
self.states()
.into_iter()
.find(|(key, _)| key.digest == digest)
}
pub fn was_evicted(&self, digest: &str) -> bool {
let inner = self.inner.lock().unwrap_or_else(|e| e.into_inner());
inner.gone.contains(digest) && !inner.entries.keys().any(|k| k.digest == digest)
}
}
fn total(inner: &Inner) -> u64 {
inner.entries.values().map(|e| e.bytes).sum()
}
#[cfg(test)]
mod tests {
use super::*;
use dataflow_rs::datavalue::OwnedDataTensor;
use std::sync::atomic::{AtomicUsize, Ordering};
struct Fake {
digest: String,
bytes: usize,
}
impl LoadedModel for Fake {
fn digest(&self) -> &str {
&self.digest
}
fn resident_bytes(&self) -> usize {
self.bytes
}
fn run(
&self,
_inputs: Vec<OwnedDataTensor>,
) -> Result<Vec<OwnedDataTensor>, super::super::runtimes::RunError> {
Ok(Vec::new())
}
}
fn key(digest: &str) -> CacheKey {
CacheKey {
digest: digest.to_string(),
runtime: "tract",
device: "cpu".to_string(),
}
}
fn fake(digest: &str, bytes: usize) -> Arc<dyn LoadedModel> {
Arc::new(Fake {
digest: digest.to_string(),
bytes,
})
}
#[tokio::test]
async fn a_hit_does_not_reload_and_a_failure_retains_nothing() {
let cache = LoadedCache::new(100);
let loads = AtomicUsize::new(0);
let (_, cold) = cache
.get_or_load(key("a"), || {
loads.fetch_add(1, Ordering::SeqCst);
async { Ok(fake("a", 10)) }
})
.await
.expect("loads");
assert!(cold);
let (model, cold) = cache
.get_or_load(key("a"), || {
loads.fetch_add(1, Ordering::SeqCst);
async { Ok(fake("a", 10)) }
})
.await
.expect("hits");
assert!(!cold);
assert_eq!(model.digest(), "a");
assert_eq!(loads.load(Ordering::SeqCst), 1);
assert!(cache.contains(&key("a")));
assert_eq!(cache.loaded_bytes(), 10);
assert_eq!(cache.loaded_for("a").map(|(_, b)| b), Some(10));
let Err(err) = cache
.get_or_load(key("b"), || async {
Err(LoadError::new("parse", "bad bytes"))
})
.await
else {
unreachable!("a failed load must not produce a model")
};
assert_eq!(err.stage, "parse");
assert!(!cache.contains(&key("b")));
assert!(!cache.was_evicted("b"));
assert_eq!(cache.loaded_bytes(), 10);
}
#[tokio::test]
async fn eviction_is_lru_by_bytes() {
let cache = LoadedCache::new(100);
for (digest, bytes) in [("a", 40), ("b", 40)] {
cache
.get_or_load(key(digest), || async move { Ok(fake(digest, bytes)) })
.await
.expect("loads");
}
cache
.get_or_load(key("a"), || async { unreachable!("a is resident") })
.await
.expect("hit");
cache
.get_or_load(key("c"), || async { Ok(fake("c", 40)) })
.await
.expect("loads");
assert!(cache.contains(&key("a")));
assert!(!cache.contains(&key("b")));
assert!(cache.contains(&key("c")));
assert_eq!(cache.loaded_bytes(), 80);
assert!(cache.was_evicted("b"));
assert!(!cache.was_evicted("a"));
assert_eq!(
cache
.states()
.iter()
.map(|(k, b)| (k.digest.as_str(), *b))
.collect::<Vec<_>>(),
[("a", 40), ("c", 40)]
);
cache
.get_or_load(key("b"), || async { Ok(fake("b", 10)) })
.await
.expect("loads");
assert!(!cache.was_evicted("b"));
}
#[tokio::test]
async fn an_oversized_model_is_served_but_not_retained() {
let cache = LoadedCache::new(100);
cache
.get_or_load(key("small"), || async { Ok(fake("small", 30)) })
.await
.expect("loads");
let (model, cold) = cache
.get_or_load(key("huge"), || async { Ok(fake("huge", 500)) })
.await
.expect("loads");
assert!(cold);
assert_eq!(model.resident_bytes(), 500);
assert!(!cache.contains(&key("huge")));
assert!(cache.was_evicted("huge"));
assert!(cache.contains(&key("small")));
assert_eq!(cache.loaded_bytes(), 30);
}
#[tokio::test]
async fn concurrent_first_callers_share_one_load() {
let cache = Arc::new(LoadedCache::new(100));
let loads = Arc::new(AtomicUsize::new(0));
let gate = Arc::new(tokio::sync::Notify::new());
let mut tasks = Vec::new();
for _ in 0..8 {
let cache = cache.clone();
let loads = loads.clone();
let gate = gate.clone();
tasks.push(tokio::spawn(async move {
cache
.get_or_load(key("shared"), move || async move {
loads.fetch_add(1, Ordering::SeqCst);
gate.notified().await;
Ok(fake("shared", 10))
})
.await
.expect("loads")
}));
}
tokio::task::yield_now().await;
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
gate.notify_one();
for task in tasks {
let (model, cold) = task.await.expect("joins");
assert_eq!(model.digest(), "shared");
assert!(cold, "every caller in the first group waited for the load");
}
assert_eq!(
loads.load(Ordering::SeqCst),
1,
"one load for eight callers"
);
assert_eq!(cache.max_bytes(), 100);
}
}