use std::collections::HashMap;
use std::io;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct ExpertKey {
pub layer: u32,
pub expert: u32,
}
pub trait ExpertSource: Send + Sync {
fn expert_len(&self, key: ExpertKey) -> Option<usize>;
fn read_expert(&self, key: ExpertKey) -> io::Result<Vec<u8>>;
}
pub struct ExpertLease {
data: Arc<Vec<u8>>,
}
impl ExpertLease {
pub fn bytes(&self) -> &[u8] {
&self.data
}
pub fn shared_buf(&self) -> Arc<Vec<u8>> {
Arc::clone(&self.data)
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct ExpertStoreStats {
pub hits: u64,
pub misses: u64,
pub evictions: u64,
pub pass_throughs: u64,
pub bytes_read: u64,
pub resident_bytes: u64,
}
struct Entry {
data: Arc<Vec<u8>>,
last_used: u64,
}
struct Inner {
entries: HashMap<ExpertKey, Entry>,
resident_bytes: usize,
clock: u64,
}
pub struct ExpertStore<S: ExpertSource> {
source: S,
budget_bytes: usize,
inner: Mutex<Inner>,
hits: AtomicU64,
misses: AtomicU64,
evictions: AtomicU64,
pass_throughs: AtomicU64,
bytes_read: AtomicU64,
}
impl<S: ExpertSource> ExpertStore<S> {
pub fn new(source: S, budget_bytes: usize) -> Self {
ExpertStore {
source,
budget_bytes,
inner: Mutex::new(Inner {
entries: HashMap::new(),
resident_bytes: 0,
clock: 0,
}),
hits: AtomicU64::new(0),
misses: AtomicU64::new(0),
evictions: AtomicU64::new(0),
pass_throughs: AtomicU64::new(0),
bytes_read: AtomicU64::new(0),
}
}
pub fn budget_bytes(&self) -> usize {
self.budget_bytes
}
pub fn stats(&self) -> ExpertStoreStats {
let resident = self
.inner
.lock()
.unwrap_or_else(|p| p.into_inner())
.resident_bytes as u64;
ExpertStoreStats {
hits: self.hits.load(Ordering::Relaxed),
misses: self.misses.load(Ordering::Relaxed),
evictions: self.evictions.load(Ordering::Relaxed),
pass_throughs: self.pass_throughs.load(Ordering::Relaxed),
bytes_read: self.bytes_read.load(Ordering::Relaxed),
resident_bytes: resident,
}
}
pub fn prefetch(&self, keys: &[ExpertKey]) {
for &key in keys {
let _ = self.acquire(key);
}
}
pub fn acquire(&self, key: ExpertKey) -> io::Result<ExpertLease> {
{
let mut inner = self.inner.lock().unwrap_or_else(|p| p.into_inner());
inner.clock += 1;
let clock = inner.clock;
if let Some(entry) = inner.entries.get_mut(&key) {
entry.last_used = clock;
self.hits.fetch_add(1, Ordering::Relaxed);
return Ok(ExpertLease {
data: Arc::clone(&entry.data),
});
}
}
self.misses.fetch_add(1, Ordering::Relaxed);
let data = self.source.read_expert(key)?;
self.bytes_read
.fetch_add(data.len() as u64, Ordering::Relaxed);
let size = data.len();
let data = Arc::new(data);
let mut inner = self.inner.lock().unwrap_or_else(|p| p.into_inner());
if let Some(entry) = inner.entries.get_mut(&key) {
return Ok(ExpertLease {
data: Arc::clone(&entry.data),
});
}
if size > self.budget_bytes {
self.pass_throughs.fetch_add(1, Ordering::Relaxed);
return Ok(ExpertLease { data });
}
while inner.resident_bytes + size > self.budget_bytes {
let victim = inner
.entries
.iter()
.filter(|(_, e)| Arc::strong_count(&e.data) == 1)
.min_by_key(|(_, e)| e.last_used)
.map(|(k, _)| *k);
match victim {
Some(k) => {
let e = inner.entries.remove(&k).expect("victim key just found");
inner.resident_bytes -= e.data.len();
self.evictions.fetch_add(1, Ordering::Relaxed);
}
None => {
self.pass_throughs.fetch_add(1, Ordering::Relaxed);
return Ok(ExpertLease { data });
}
}
}
inner.clock += 1;
let clock = inner.clock;
inner.resident_bytes += size;
inner.entries.insert(
key,
Entry {
data: Arc::clone(&data),
last_used: clock,
},
);
Ok(ExpertLease { data })
}
}
pub struct FileRangeSource {
file: std::fs::File,
ranges: HashMap<ExpertKey, (u64, usize)>,
#[cfg(not(unix))]
seek_lock: Mutex<()>,
}
impl FileRangeSource {
pub fn new(file: std::fs::File, ranges: HashMap<ExpertKey, (u64, usize)>) -> Self {
FileRangeSource {
file,
ranges,
#[cfg(not(unix))]
seek_lock: Mutex::new(()),
}
}
}
impl ExpertSource for FileRangeSource {
fn expert_len(&self, key: ExpertKey) -> Option<usize> {
self.ranges.get(&key).map(|&(_, len)| len)
}
fn read_expert(&self, key: ExpertKey) -> io::Result<Vec<u8>> {
let &(offset, len) = self
.ranges
.get(&key)
.ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, format!("{key:?}")))?;
let mut buf = vec![0u8; len];
#[cfg(unix)]
{
use std::os::unix::fs::FileExt;
self.file.read_exact_at(&mut buf, offset)?;
}
#[cfg(not(unix))]
{
use std::io::{Read, Seek, SeekFrom};
let _guard = self.seek_lock.lock().unwrap_or_else(|p| p.into_inner());
let mut f = &self.file;
f.seek(SeekFrom::Start(offset))?;
f.read_exact(&mut buf)?;
}
Ok(buf)
}
}
#[cfg(test)]
mod tests {
use super::*;
struct PatternSource {
len: usize,
n_layers: u32,
n_experts: u32,
}
fn expected_bytes(key: ExpertKey, len: usize) -> Vec<u8> {
(0..len)
.map(|i| (key.layer as usize * 31 + key.expert as usize * 7 + i * 13) as u8)
.collect()
}
impl ExpertSource for PatternSource {
fn expert_len(&self, key: ExpertKey) -> Option<usize> {
(key.layer < self.n_layers && key.expert < self.n_experts).then_some(self.len)
}
fn read_expert(&self, key: ExpertKey) -> io::Result<Vec<u8>> {
if self.expert_len(key).is_none() {
return Err(io::Error::new(io::ErrorKind::NotFound, "no such expert"));
}
Ok(expected_bytes(key, self.len))
}
}
fn key(layer: u32, expert: u32) -> ExpertKey {
ExpertKey { layer, expert }
}
fn store(len: usize, budget: usize) -> ExpertStore<PatternSource> {
ExpertStore::new(
PatternSource {
len,
n_layers: 8,
n_experts: 8,
},
budget,
)
}
#[test]
fn hits_misses_and_lru_eviction_order() {
let s = store(100, 250); assert_eq!(
s.acquire(key(0, 0)).unwrap().bytes(),
expected_bytes(key(0, 0), 100)
);
assert_eq!(
s.acquire(key(0, 1)).unwrap().bytes(),
expected_bytes(key(0, 1), 100)
);
s.acquire(key(0, 0)).unwrap();
s.acquire(key(0, 2)).unwrap();
let before = s.stats();
s.acquire(key(0, 0)).unwrap(); let after = s.stats();
assert_eq!(after.hits, before.hits + 1);
assert_eq!(after.misses, before.misses);
assert_eq!(after.evictions, 1);
assert!(after.resident_bytes <= 250);
s.acquire(key(0, 1)).unwrap();
assert_eq!(s.stats().misses, after.misses + 1);
}
#[test]
fn a_live_lease_pins_its_entry_against_eviction() {
let s = store(100, 250); let pinned = s.acquire(key(1, 0)).unwrap();
for e in 1..6 {
s.acquire(key(1, e)).unwrap();
}
assert_eq!(pinned.bytes(), expected_bytes(key(1, 0), 100));
let hits_before = s.stats().hits;
let again = s.acquire(key(1, 0)).unwrap();
assert_eq!(s.stats().hits, hits_before + 1);
assert_eq!(again.bytes(), expected_bytes(key(1, 0), 100));
drop(pinned);
drop(again);
for e in 1..6 {
s.acquire(key(2, e)).unwrap();
}
let misses_before = s.stats().misses;
s.acquire(key(1, 0)).unwrap();
assert_eq!(
s.stats().misses,
misses_before + 1,
"unpinned entry was evictable"
);
}
#[test]
fn budget_smaller_than_one_expert_degrades_to_pass_through_not_failure() {
let s = store(100, 50);
for e in 0..4 {
let lease = s.acquire(key(0, e)).unwrap();
assert_eq!(lease.bytes(), expected_bytes(key(0, e), 100));
}
let st = s.stats();
assert_eq!(st.pass_throughs, 4);
assert_eq!(st.resident_bytes, 0);
assert_eq!(st.evictions, 0);
}
#[test]
fn fully_pinned_cache_serves_new_experts_uncached() {
let s = store(100, 200);
let _a = s.acquire(key(0, 0)).unwrap();
let _b = s.acquire(key(0, 1)).unwrap();
let c = s.acquire(key(0, 2)).unwrap();
assert_eq!(c.bytes(), expected_bytes(key(0, 2), 100));
assert_eq!(s.stats().pass_throughs, 1);
assert_eq!(s.stats().evictions, 0);
let hits_before = s.stats().hits;
s.acquire(key(0, 0)).unwrap();
s.acquire(key(0, 1)).unwrap();
assert_eq!(s.stats().hits, hits_before + 2);
}
#[test]
fn missing_expert_is_a_clean_error() {
let s = store(100, 1000);
assert!(s.acquire(key(99, 0)).is_err());
}
#[test]
fn concurrent_acquires_under_eviction_pressure_never_yield_wrong_bytes() {
let s = Arc::new(store(64, 200)); let mut handles = Vec::new();
for t in 0..8u32 {
let s = Arc::clone(&s);
handles.push(std::thread::spawn(move || {
let mut state = t.wrapping_mul(2654435761).wrapping_add(12345);
for _ in 0..500 {
state = state.wrapping_mul(1664525).wrapping_add(1013904223);
let k = key((state >> 8) % 4, (state >> 16) % 4);
let lease = s.acquire(k).expect("in-range key must read");
assert_eq!(
lease.bytes(),
expected_bytes(k, 64),
"wrong bytes for {k:?} -- slot reuse corruption"
);
}
}));
}
for h in handles {
h.join().unwrap();
}
let st = s.stats();
assert_eq!(
st.hits + st.misses,
8 * 500,
"every acquire is a hit or a miss"
);
assert!(st.resident_bytes <= 200, "budget held under concurrency");
}
#[test]
fn weight_matrix_over_a_lease_computes_identically_and_extends_the_pin() {
use crate::weight_matrix::{QuantKind, WeightBytes, WeightMatrix};
struct QuantSource {
rows: Vec<u8>,
}
impl ExpertSource for QuantSource {
fn expert_len(&self, _k: ExpertKey) -> Option<usize> {
Some(self.rows.len())
}
fn read_expert(&self, _k: ExpertKey) -> io::Result<Vec<u8>> {
Ok(self.rows.clone())
}
}
let cols = 64;
let rows = 2;
let values: Vec<f32> = (0..rows * cols)
.map(|i| ((i as f32) * 0.11).sin())
.collect();
let mut packed = Vec::new();
for r in 0..rows {
packed.extend(ferrox_quant::quantize_q8_0(
&values[r * cols..(r + 1) * cols],
));
}
let store = ExpertStore::new(
QuantSource {
rows: packed.clone(),
},
packed.len(),
);
let lease = store.acquire(key(0, 0)).unwrap();
let matrix = WeightMatrix::Quantized {
data: WeightBytes::Shared {
buf: lease.shared_buf(),
range: 0..packed.len(),
},
rows,
cols,
kind: QuantKind::Q8_0,
};
drop(lease);
let owned = WeightMatrix::Quantized {
data: WeightBytes::Owned(packed),
rows,
cols,
kind: QuantKind::Q8_0,
};
let x: Vec<f32> = (0..cols).map(|i| ((i as f32) * 0.07).cos()).collect();
assert_eq!(
matrix.apply(&x),
owned.apply(&x),
"lease-backed == owned, bit for bit"
);
let hits_before = store.stats().hits;
store.acquire(key(0, 0)).unwrap();
assert_eq!(store.stats().hits, hits_before + 1);
}
#[test]
fn file_range_source_reads_correct_ranges_through_the_store() {
let dir = std::env::temp_dir();
let path = dir.join(format!(
"ferrox_expert_store_test_{}.bin",
std::process::id()
));
let mut file_bytes = Vec::new();
let mut ranges = HashMap::new();
for l in 0..3u32 {
for e in 0..4u32 {
let content = expected_bytes(key(l, e), 96);
ranges.insert(key(l, e), (file_bytes.len() as u64, content.len()));
file_bytes.extend_from_slice(&content);
}
}
std::fs::write(&path, &file_bytes).unwrap();
let source = FileRangeSource::new(std::fs::File::open(&path).unwrap(), ranges);
let store = Arc::new(ExpertStore::new(source, 300));
let mut handles = Vec::new();
for t in 0..4u32 {
let store = Arc::clone(&store);
handles.push(std::thread::spawn(move || {
for i in 0..200u32 {
let k = key((t + i) % 3, i % 4);
let lease = store.acquire(k).unwrap();
assert_eq!(lease.bytes(), expected_bytes(k, 96));
}
}));
}
for h in handles {
h.join().unwrap();
}
std::fs::remove_file(&path).ok();
}
}