use crate::Engine;
use crate::model::{ExpertKeepalive, ExpertSource};
use crate::spill_pread::{PreadPool, PreadStats, ReadTicket, SpillIoMode};
use cudarc::driver::{CudaEvent, CudaSlice, CudaStream, HostSlice, SyncOnDrop};
use std::collections::{BTreeMap, HashMap, HashSet};
use std::sync::Arc;
pub const PROJ_GATE: u8 = 0;
pub const PROJ_UP: u8 = 1;
pub const PROJ_DOWN: u8 = 2;
#[derive(Clone, Copy, PartialEq, Eq, Hash, Debug)]
pub struct BlockId {
pub layer: u16,
pub proj: u8,
pub ex: u16,
}
impl BlockId {
#[inline]
pub fn new(layer: u16, proj: u8, ex: u16) -> Self {
BlockId { layer, proj, ex }
}
}
#[derive(Clone, Copy, Debug)]
pub enum DispatchSlot {
Resident(usize),
}
const NIL: u32 = u32::MAX;
const SEG_NONE: u8 = 0;
const SEG_PROBATION: u8 = 1;
const SEG_PROTECTED: u8 = 2;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct SlotLink {
prev: u32,
next: u32,
seg: u8,
}
impl SlotLink {
const fn none() -> Self {
SlotLink {
prev: NIL,
next: NIL,
seg: SEG_NONE,
}
}
}
#[derive(Debug)]
struct SlruList {
head: u32,
tail: u32,
len: usize,
}
impl SlruList {
const fn new() -> Self {
SlruList {
head: NIL,
tail: NIL,
len: 0,
}
}
fn push_back(&mut self, slot: usize, seg: u8, links: &mut [SlotLink]) {
debug_assert_eq!(
links[slot].seg, SEG_NONE,
"slot {slot} already in a segment"
);
let s = slot as u32;
links[slot] = SlotLink {
prev: self.tail,
next: NIL,
seg,
};
if self.tail != NIL {
links[self.tail as usize].next = s;
} else {
self.head = s;
}
self.tail = s;
self.len += 1;
}
fn pop_front(&mut self, links: &mut [SlotLink]) -> Option<usize> {
if self.head == NIL {
return None;
}
let s = self.head as usize;
self.unlink(s, links);
Some(s)
}
fn unlink(&mut self, slot: usize, links: &mut [SlotLink]) {
let l = links[slot];
debug_assert_ne!(l.seg, SEG_NONE, "unlink of slot {slot} not in a segment");
if l.prev != NIL {
links[l.prev as usize].next = l.next;
} else {
debug_assert_eq!(self.head, slot as u32);
self.head = l.next;
}
if l.next != NIL {
links[l.next as usize].prev = l.prev;
} else {
debug_assert_eq!(self.tail, slot as u32);
self.tail = l.prev;
}
links[slot] = SlotLink::none();
self.len -= 1;
}
fn iter<'a>(&self, links: &'a [SlotLink]) -> SlruIter<'a> {
SlruIter {
links,
cur: self.head,
}
}
}
struct SlruIter<'a> {
links: &'a [SlotLink],
cur: u32,
}
impl Iterator for SlruIter<'_> {
type Item = usize;
fn next(&mut self) -> Option<usize> {
if self.cur == NIL {
return None;
}
let s = self.cur as usize;
self.cur = self.links[s].next;
Some(s)
}
}
struct SlotClass {
capacity: usize,
probation: SlruList,
protected: SlruList,
free: Vec<usize>,
protected_cap: usize,
}
impl SlotClass {
fn on_hit_full(&mut self, slot: usize, links: &mut [SlotLink]) {
match links[slot].seg {
SEG_PROBATION => {
self.probation.unlink(slot, links);
self.push_protected(slot, links);
}
SEG_PROTECTED => {
self.protected.unlink(slot, links);
self.protected.push_back(slot, SEG_PROTECTED, links); }
_ => self.push_protected(slot, links),
}
}
fn push_protected(&mut self, slot: usize, links: &mut [SlotLink]) {
self.protected.push_back(slot, SEG_PROTECTED, links);
while self.protected.len > self.protected_cap {
if let Some(demoted) = self.protected.pop_front(links) {
self.probation.push_back(demoted, SEG_PROBATION, links);
} else {
break;
}
}
}
fn pop_lru(&mut self, links: &mut [SlotLink]) -> Option<usize> {
self.probation
.pop_front(links)
.or_else(|| self.protected.pop_front(links))
}
fn unlink_from_segment(&mut self, slot: usize, links: &mut [SlotLink]) {
match links[slot].seg {
SEG_PROBATION => self.probation.unlink(slot, links),
SEG_PROTECTED => self.protected.unlink(slot, links),
_ => {}
}
}
}
pub struct MoeSlotCache {
slots: Vec<CudaSlice<u8>>, slot_class: Vec<usize>, classes: Vec<SlotClass>,
links: Vec<SlotLink>,
occupant: Vec<Option<BlockId>>, table: HashMap<BlockId, usize>, frequencies: HashMap<BlockId, f32>,
pending: HashMap<BlockId, PendingBlock>,
inflight_sources: Vec<(Arc<CudaEvent>, ExpertKeepalive)>,
quarantined_sources: Vec<ExpertKeepalive>,
compute_sources: HashMap<KeepaliveKey, ExpertKeepalive>,
pread: Option<PreadPool>,
worker_reads: HashMap<BlockId, WorkerRead>,
pread_requested: bool,
pread_fallbacks: u64,
copy_stream: Arc<CudaStream>,
copy_stream_unknown: bool,
compute_stream: Arc<CudaStream>,
compute_stream_unknown: bool,
n: usize,
max_block_bytes: usize,
size_aware: bool,
frequency_evict: bool,
frequency_decay: Option<f32>,
mtp_frequency_weight: f32,
last_forward_layer: Option<u16>,
last_forward_t: usize,
frozen: bool,
per_layer: HashMap<u16, u32>,
dev_rows: HashMap<u16, CudaSlice<u64>>,
prewarm_tried: HashSet<u16>,
pub hits: u64,
pub misses: u64,
pub staged_bytes: u64, }
struct PendingBlock {
slot: usize,
ready: Arc<CudaEvent>,
keepalive: Option<ExpertKeepalive>,
}
#[derive(Clone, Copy)]
struct WorkerRead {
ticket: ReadTicket,
len: usize,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
enum KeepaliveKey {
Pinned(usize),
Buffer(usize),
Mmap(usize),
}
impl KeepaliveKey {
fn from_owner(owner: &ExpertKeepalive) -> Self {
match owner {
ExpertKeepalive::Pinned(value) => Self::Pinned(Arc::as_ptr(value) as usize),
ExpertKeepalive::Buffer(value) => Self::Buffer(Arc::as_ptr(value) as usize),
ExpertKeepalive::Mmap(value) => Self::Mmap(Arc::as_ptr(value) as usize),
}
}
}
struct ExactPinnedPrefix<'a>(&'a [u8]);
impl HostSlice<u8> for ExactPinnedPrefix<'_> {
fn len(&self) -> usize {
self.0.len()
}
unsafe fn stream_synced_slice<'a>(
&'a self,
_stream: &'a CudaStream,
) -> (&'a [u8], SyncOnDrop<'a>) {
(self.0, SyncOnDrop::Record(None))
}
unsafe fn stream_synced_mut_slice<'a>(
&'a mut self,
_stream: &'a CudaStream,
) -> (&'a mut [u8], SyncOnDrop<'a>) {
panic!("ExactPinnedPrefix is a source-only HostSlice")
}
}
fn stage_on_copy_stream(
e: &Engine,
host_bytes: &[u8],
slot: &mut CudaSlice<u8>,
) -> Result<Arc<CudaEvent>, (Box<dyn std::error::Error>, bool)> {
let prior = match e.stream().record_event(None) {
Ok(prior) => prior,
Err(err) => return Err((err.into(), true)),
};
if let Err(err) = e.copy_stream.wait(&prior) {
return Err((err.into(), true));
}
match e.stage_expert_async(host_bytes, slot, 0) {
Ok(ready) => Ok(Arc::new(ready)),
Err(err) => {
match e.copy_stream.synchronize() {
Ok(()) => Err((err, true)),
Err(sync_err) => Err((std::io::Error::other(format!(
"copy-stream H2D setup failed ({err}); stream drain also failed ({sync_err})"
)).into(), false)),
}
}
}
}
fn stage_pread_on_compute_stream(
e: &Engine,
host_bytes: &[u8],
slot: &mut CudaSlice<u8>,
) -> Result<Arc<CudaEvent>, Box<dyn std::error::Error>> {
let ready = Arc::new(e.ctx().new_event(None)?);
let source = ExactPinnedPrefix(host_bytes);
let mut dst = slot.slice_mut(0..host_bytes.len());
e.stream().memcpy_htod(&source, &mut dst)?;
ready.record(&e.stream())?;
Ok(ready)
}
fn stage_pread_prefetch_on_copy_stream(
e: &Engine,
host_bytes: &[u8],
slot: &mut CudaSlice<u8>,
) -> Result<Arc<CudaEvent>, (Box<dyn std::error::Error>, bool)> {
let ready = match e.ctx().new_event(None) {
Ok(ready) => Arc::new(ready),
Err(err) => return Err((err.into(), true)),
};
let source = ExactPinnedPrefix(host_bytes);
let mut dst = slot.slice_mut(0..host_bytes.len());
let submitted = e
.copy_stream
.memcpy_htod(&source, &mut dst)
.and_then(|()| ready.record(&e.copy_stream));
match submitted {
Ok(()) => Ok(ready),
Err(err) => match e.copy_stream.synchronize() {
Ok(()) => Err((err.into(), true)),
Err(sync_err) => Err((
std::io::Error::other(format!(
"pread copy-stream H2D setup failed ({err}); stream drain also failed ({sync_err})"
))
.into(),
false,
)),
},
}
}
fn size_class_plan(block_bytes: &[usize], budget_bytes: usize) -> Vec<(usize, usize)> {
let mut counts: BTreeMap<usize, usize> = BTreeMap::new();
for &bytes in block_bytes.iter().filter(|&&bytes| bytes > 0) {
*counts.entry(bytes).or_insert(0) += 1;
}
if counts.is_empty() || budget_bytes == 0 {
return Vec::new();
}
let total_bytes: u128 = counts
.iter()
.map(|(&bytes, &count)| (bytes as u128 + 8) * count as u128)
.sum();
let budget = budget_bytes as u128;
let mut plan: Vec<(usize, usize, u128)> = counts
.iter()
.map(|(&bytes, &count)| {
let scaled = count as u128 * budget;
(
bytes,
(scaled / total_bytes).min(count as u128) as usize,
scaled % total_bytes,
)
})
.collect();
let mut used: u128 = plan
.iter()
.map(|(bytes, count, _)| (*bytes as u128 + 8) * *count as u128)
.sum();
let mut order: Vec<usize> = (0..plan.len()).collect();
order.sort_by(|&a, &b| plan[b].2.cmp(&plan[a].2).then(a.cmp(&b)));
for index in order {
let (bytes, count, _) = plan[index];
let available = counts[&bytes];
let required = bytes as u128 + 8;
if count < available && used + required <= budget {
plan[index].1 += 1;
used += required;
}
}
plan.into_iter()
.filter_map(|(bytes, count, _)| (count > 0).then_some((bytes, count)))
.collect()
}
impl MoeSlotCache {
pub fn new(e: &Engine, max_block_bytes: usize) -> Result<Self, Box<dyn std::error::Error>> {
let (free, _total) = e.ctx().mem_get_info()?;
let hard_frac = cache_hard_vram_frac();
let hard_bytes =
((free as f64 * hard_frac) as usize).saturating_sub(2 * (max_block_bytes + 8));
let forced_slots = std::env::var("MEMRA_MOE_SLOTS")
.ok()
.and_then(|s| s.parse::<usize>().ok());
let requested_bytes = if let Some(n) = forced_slots {
n.saturating_mul(max_block_bytes + 8)
} else {
let frac = std::env::var("MEMRA_MOE_VRAM_FRAC")
.ok()
.and_then(|s| s.parse::<f64>().ok())
.unwrap_or(0.85);
(free as f64 * frac) as usize
};
let budget_bytes = requested_bytes.min(hard_bytes);
let layout = e.moe_cache_layout().unwrap_or_default();
let size_aware = forced_slots.is_none()
&& std::env::var("MEMRA_MOE_SIZE_AWARE").as_deref() == Ok("1")
&& !layout.is_empty();
let frequency_evict = std::env::var("MEMRA_MOE_LFU").as_deref() == Ok("1");
let frequency_decay = if frequency_evict {
cache_lfu_decay()
} else {
None
};
let mtp_frequency_weight = cache_lfu_mtp_weight();
let mut class_plan = if size_aware {
size_class_plan(&layout, budget_bytes)
} else {
Vec::new()
};
if class_plan.iter().map(|(_, count)| count).sum::<usize>() < 8 {
let n = (budget_bytes / (max_block_bytes + 8)).max(8);
class_plan = vec![(max_block_bytes, n)];
}
let n: usize = class_plan.iter().map(|(_, count)| count).sum();
let mut slots = Vec::with_capacity(n);
let mut slot_class = Vec::with_capacity(n);
let mut classes = Vec::with_capacity(class_plan.len());
let mut occupant = Vec::with_capacity(n);
for (class_index, &(capacity, count)) in class_plan.iter().enumerate() {
let start = slots.len();
for _ in 0..count {
slots.push(e.alloc_u8(capacity + 8)?);
slot_class.push(class_index);
occupant.push(None);
}
let free_slots = (start..start + count).rev().collect();
classes.push(SlotClass {
capacity,
probation: SlruList::new(),
protected: SlruList::new(),
free: free_slots,
protected_cap: ((count as f64 * 0.8) as usize).max(1),
});
}
let links = vec![SlotLink::none(); n];
if size_aware {
let allocated: usize = class_plan
.iter()
.map(|(bytes, count)| (bytes + 8) * count)
.sum();
eprintln!(
"[moe-cache] size-aware fixed slots: {n} slots in {} classes, {:.2} GB / {:.2} GB budget",
class_plan.len(),
allocated as f64 / 1e9,
budget_bytes as f64 / 1e9
);
}
let pread_mode = crate::spill_pread::configured_mode();
let pread_requested = pread_mode != SpillIoMode::Mmap;
let pread = if pread_requested {
match PreadPool::try_new(e, max_block_bytes, pread_mode) {
Ok(pool) => Some(pool),
Err(err) => {
eprintln!(
"[spill-pread] pinned-buffer initialization failed ({err}); using mmap"
);
None
}
}
} else {
None
};
Ok(MoeSlotCache {
slots,
slot_class,
classes,
links,
occupant,
table: HashMap::with_capacity(n * 2),
frequencies: HashMap::with_capacity(layout.len().max(n * 2)),
pending: HashMap::new(),
inflight_sources: Vec::new(),
quarantined_sources: Vec::new(),
compute_sources: HashMap::new(),
pread,
worker_reads: HashMap::new(),
pread_requested,
pread_fallbacks: 0,
copy_stream: e.copy_stream.clone(),
copy_stream_unknown: false,
compute_stream: e.stream().clone(),
compute_stream_unknown: false,
n,
max_block_bytes,
size_aware,
frequency_evict,
frequency_decay,
mtp_frequency_weight,
last_forward_layer: None,
last_forward_t: 0,
frozen: false,
per_layer: HashMap::new(),
dev_rows: HashMap::new(),
prewarm_tried: HashSet::new(),
hits: 0,
misses: 0,
staged_bytes: 0,
})
}
#[inline]
pub fn n_slots(&self) -> usize {
self.n
}
#[inline]
pub fn is_frozen(&self) -> bool {
self.frozen
}
pub fn freeze(&mut self) {
if !self.frozen {
self.frozen = true;
let (_, complete, one_projection, two_projections, stranded_blocks) =
self.expert_residency_shape();
eprintln!(
"[moe-cache] residency frozen: {} slots, {} resident blocks; \
{complete} complete experts, {one_projection} one-projection fragments, \
{two_projections} two-projection fragments ({stranded_blocks} stranded blocks)",
self.n,
self.table.len()
);
let mut mtp_masks = HashMap::<u16, u8>::new();
for id in self.table.keys().filter(|id| id.layer == u16::MAX) {
*mtp_masks.entry(id.ex).or_insert(0) |= 1u8 << id.proj;
}
if !mtp_masks.is_empty() {
let complete = mtp_masks.values().filter(|&&mask| mask == 0b111).count();
eprintln!(
"[moe-cache] frozen MTP residency: {} blocks, {complete} complete experts",
mtp_masks
.values()
.map(|mask| mask.count_ones() as usize)
.sum::<usize>()
);
}
}
}
pub(crate) fn expert_residency_shape(&self) -> (usize, usize, usize, usize, usize) {
let mut masks = HashMap::<(u16, u16), u8>::new();
for id in self.table.keys() {
*masks.entry((id.layer, id.ex)).or_insert(0) |= 1u8 << id.proj;
}
let complete = masks.values().filter(|&&mask| mask == 0b111).count();
let one_projection = masks
.values()
.filter(|&&mask| mask.count_ones() == 1)
.count();
let two_projections = masks
.values()
.filter(|&&mask| mask.count_ones() == 2)
.count();
let stranded_blocks = one_projection + 2 * two_projections;
(
masks.len(),
complete,
one_projection,
two_projections,
stranded_blocks,
)
}
#[inline]
pub fn max_block_bytes(&self) -> usize {
self.max_block_bytes
}
#[inline]
pub fn resident(&self, id: BlockId) -> Option<usize> {
self.table.get(&id).copied()
}
#[inline]
fn frequency_increment(&self, id: BlockId) -> f32 {
if id.layer == u16::MAX {
self.mtp_frequency_weight
} else {
1.0
}
}
pub(crate) fn note_profile_hit(&mut self, id: BlockId) {
if self.frozen || !self.table.contains_key(&id) {
return;
}
let increment = self.frequency_increment(id);
*self.frequencies.entry(id).or_insert(0.0) += increment;
}
fn on_hit(&mut self, slot: usize) {
let class_index = self.slot_class[slot];
let class = &mut self.classes[class_index];
if !class.free.is_empty() {
return;
}
class.on_hit_full(slot, &mut self.links);
}
fn remove_occupant(&mut self, slot: usize) {
if let Some(old) = self.occupant[slot].take() {
self.table.remove(&old);
self.on_block_evicted(old.layer);
}
}
fn frequency_victim_in_class(&mut self, class_index: usize, keep: &[BlockId]) -> Option<usize> {
let class = &self.classes[class_index];
let candidate = class
.probation
.iter(&self.links)
.chain(class.protected.iter(&self.links))
.enumerate()
.filter_map(|(position, slot)| {
let id = self.occupant[slot]?;
(!keep.contains(&id)).then_some((
self.frequencies.get(&id).copied().unwrap_or(0.0),
position,
slot,
))
})
.min_by(|a, b| a.0.total_cmp(&b.0).then(a.1.cmp(&b.1)));
let (_, _, slot) = candidate?;
self.classes[class_index].unlink_from_segment(slot, &mut self.links);
Some(slot)
}
fn evict_one(&mut self, required: usize) -> Option<usize> {
for class_index in 0..self.classes.len() {
if self.classes[class_index].capacity < required {
continue;
}
let slot = if self.frequency_evict {
self.frequency_victim_in_class(class_index, &[])
} else {
self.classes[class_index].pop_lru(&mut self.links)
};
if let Some(slot) = slot {
self.remove_occupant(slot);
return Some(slot);
}
}
None
}
fn evict_one_excluding(&mut self, required: usize, keep: &[BlockId]) -> Option<usize> {
fn take(
q: &mut SlruList,
links: &mut [SlotLink],
occupant: &[Option<BlockId>],
keep: &[BlockId],
) -> Option<usize> {
let slot = q
.iter(links)
.find(|&s| occupant[s].is_some_and(|id| !keep.contains(&id)))?;
q.unlink(slot, links);
Some(slot)
}
for class_index in 0..self.classes.len() {
if self.classes[class_index].capacity < required {
continue;
}
let slot = if self.frequency_evict {
self.frequency_victim_in_class(class_index, keep)
} else {
let class = &mut self.classes[class_index];
take(&mut class.probation, &mut self.links, &self.occupant, keep)
.or_else(|| take(&mut class.protected, &mut self.links, &self.occupant, keep))
};
if let Some(slot) = slot {
self.remove_occupant(slot);
return Some(slot);
}
}
None
}
fn on_block_evicted(&mut self, layer: u16) {
if let Some(c) = self.per_layer.get_mut(&layer) {
*c -= 1;
}
self.dev_rows.remove(&layer);
}
fn reserve_slot(&mut self, required: usize) -> Option<usize> {
for class in &mut self.classes {
if class.capacity >= required
&& let Some(slot) = class.free.pop()
{
return Some(slot);
}
}
self.evict_one(required)
}
fn release_reserved_slot(&mut self, slot: usize) {
debug_assert!(self.occupant[slot].is_none());
self.classes[self.slot_class[slot]].free.push(slot);
}
fn publish(&mut self, id: BlockId, slot: usize) {
self.occupant[slot] = Some(id);
self.table.insert(id, slot);
self.classes[self.slot_class[slot]].probation.push_back(
slot,
SEG_PROBATION,
&mut self.links,
);
*self.per_layer.entry(id.layer).or_insert(0) += 1;
}
fn reap_copy_sources(&mut self) {
self.inflight_sources
.retain(|(ready, _)| !ready.is_complete());
}
fn retain_compute_source(&mut self, owner: Option<ExpertKeepalive>) {
if let Some(owner) = owner {
let key = KeepaliveKey::from_owner(&owner);
self.compute_sources.entry(key).or_insert(owner);
}
}
fn admit(
&mut self,
id: BlockId,
host_bytes: &[u8],
e: &Engine,
) -> Result<usize, Box<dyn std::error::Error>> {
let slot = self.reserve_slot(host_bytes.len()).ok_or_else(|| {
std::io::Error::other(format!(
"no MoE cache slot can hold {} bytes (max class {})",
host_bytes.len(),
self.classes.last().map(|class| class.capacity).unwrap_or(0)
))
})?;
if let Err(err) = e.stage_expert(host_bytes, &mut self.slots[slot], 0) {
return match e.stream().synchronize() {
Ok(()) => {
self.release_reserved_slot(slot);
Err(err)
}
Err(sync_err) => {
self.compute_stream_unknown = true;
Err(std::io::Error::other(format!(
"compute-stream H2D setup failed ({err}); stream drain also failed ({sync_err})"
)).into())
}
};
}
self.staged_bytes += host_bytes.len() as u64;
self.publish(id, slot);
Ok(slot)
}
fn note_pread_fallback(&mut self, reason: &dyn std::fmt::Display) {
self.pread_fallbacks += 1;
if let Some(pool) = self.pread.as_mut() {
pool.note_fallback();
}
if self.pread_fallbacks <= 3 {
eprintln!("[spill-pread] falling back to mmap: {reason}");
}
}
pub(crate) fn begin_worker_scope(&mut self) {
if self.worker_reads.is_empty() {
return;
}
let tickets: Vec<_> = self
.worker_reads
.drain()
.map(|(_, read)| read.ticket)
.collect();
if let Some(pool) = self.pread.as_mut().filter(|pool| pool.is_worker()) {
for ticket in tickets {
let _ = pool.cancel_worker(ticket);
}
}
}
pub(crate) fn begin_forward_epoch(&mut self, layer: u16, t: usize) {
let Some(decay) = self.frequency_decay else {
self.last_forward_layer = Some(layer);
self.last_forward_t = t;
return;
};
let new_sweep = self
.last_forward_layer
.is_some_and(|previous| layer <= previous);
if new_sweep && t == 1 {
if self.last_forward_t != 1 {
self.frequencies.clear();
} else {
self.frequencies.retain(|_, score| {
*score *= decay;
*score >= 1.0e-3
});
}
}
self.last_forward_layer = Some(layer);
self.last_forward_t = t;
}
pub(crate) fn promote_worker_reads_at_safe_boundary(
&mut self,
order: &[BlockId],
keep: &[BlockId],
e: &Engine,
) -> Result<usize, Box<dyn std::error::Error>> {
if !crate::spill_pread::copy_h2d_enabled() {
return Ok(0);
}
let mut promoted = 0usize;
for &id in order {
if self.table.contains_key(&id) || self.pending.contains_key(&id) {
if let Some(read) = self.worker_reads.remove(&id)
&& let Some(pool) = self.pread.as_mut()
{
let _ = pool.cancel_worker(read.ticket);
}
continue;
}
let Some(read) = self.worker_reads.get(&id).copied() else {
continue;
};
let Some(slot) = self.reserve_prefetch_slot(read.len, keep) else {
continue;
};
self.worker_reads.remove(&id);
let index = match self.pread.as_mut().unwrap().wait_worker(read.ticket) {
Ok(index) => index,
Err(err) => {
let _ = self.pread.as_mut().unwrap().cancel_worker(read.ticket);
self.release_reserved_slot(slot);
self.note_pread_fallback(err.as_ref());
continue;
}
};
let ready = {
let bytes = match self.pread.as_ref().unwrap().bytes(index, read.len) {
Ok(bytes) => bytes,
Err(err) => {
self.pread.as_mut().unwrap().abort_read(index);
self.release_reserved_slot(slot);
self.note_pread_fallback(err.as_ref());
continue;
}
};
stage_pread_prefetch_on_copy_stream(e, bytes, &mut self.slots[slot])
};
let ready = match ready {
Ok(ready) => ready,
Err((err, reusable)) => {
if reusable {
self.pread.as_mut().unwrap().abort_read(index);
self.release_reserved_slot(slot);
self.note_pread_fallback(err.as_ref());
continue;
}
self.pread.as_mut().unwrap().mark_unknown_h2d(index);
self.copy_stream_unknown = true;
return Err(err);
}
};
self.pread.as_mut().unwrap().mark_h2d(index, ready.clone());
self.occupant[slot] = Some(id);
self.pending.insert(
id,
PendingBlock {
slot,
ready,
keepalive: None,
},
);
self.staged_bytes += read.len as u64;
promoted += 1;
}
Ok(promoted)
}
fn dispatch_disk(
&mut self,
id: BlockId,
file: &Arc<std::fs::File>,
offset: u64,
len: usize,
fallback: &[u8],
e: &Engine,
) -> Result<DispatchSlot, Box<dyn std::error::Error>> {
if self.pread.is_none() {
if self.pread_requested {
self.note_pread_fallback(&"pinned-buffer backend unavailable");
}
return Ok(DispatchSlot::Resident(self.admit(id, fallback, e)?));
}
let pending = self.worker_reads.remove(&id);
let pool = self.pread.as_mut().unwrap();
let read = if pool.is_worker() {
let ticket = match pending {
Some(read) => Ok(Some(read.ticket)),
None => pool.submit_worker(file.clone(), offset, len),
};
match ticket {
Ok(Some(ticket)) => match pool.wait_worker(ticket) {
Ok(index) => Ok(index),
Err(err) => {
let _ = pool.cancel_worker(ticket);
Err(err)
}
},
Ok(None) => Err(std::io::Error::other("worker read ring is busy").into()),
Err(err) => Err(err),
}
} else {
debug_assert!(pending.is_none());
pool.read(file.as_ref(), offset, len)
};
let index = match read {
Ok(index) => index,
Err(err) => {
self.note_pread_fallback(err.as_ref());
return Ok(DispatchSlot::Resident(self.admit(id, fallback, e)?));
}
};
let slot = self.reserve_slot(len).ok_or_else(|| {
std::io::Error::other(format!(
"no MoE cache slot can hold {len} bytes (max class {})",
self.classes.last().map(|class| class.capacity).unwrap_or(0)
))
})?;
let ready = {
let bytes = match self.pread.as_ref().unwrap().bytes(index, len) {
Ok(bytes) => bytes,
Err(err) => {
self.pread.as_mut().unwrap().abort_read(index);
self.release_reserved_slot(slot);
self.note_pread_fallback(err.as_ref());
return Ok(DispatchSlot::Resident(self.admit(id, fallback, e)?));
}
};
stage_pread_on_compute_stream(e, bytes, &mut self.slots[slot])
};
let ready = match ready {
Ok(ready) => ready,
Err(err) => {
match e.stream().synchronize() {
Ok(()) => {
self.pread.as_mut().unwrap().abort_read(index);
self.release_reserved_slot(slot);
self.note_pread_fallback(err.as_ref());
return Ok(DispatchSlot::Resident(self.admit(id, fallback, e)?));
}
Err(sync_err) => {
self.pread.as_mut().unwrap().mark_unknown_h2d(index);
return Err(std::io::Error::other(format!(
"pread H2D setup failed ({err}); CUDA stream drain also failed ({sync_err})"
)).into());
}
}
}
};
self.pread.as_mut().unwrap().mark_h2d(index, ready);
self.staged_bytes += len as u64;
self.publish(id, slot);
Ok(DispatchSlot::Resident(slot))
}
pub fn dispatch(
&mut self,
id: BlockId,
host_bytes: &[u8],
e: &Engine,
) -> Result<DispatchSlot, Box<dyn std::error::Error>> {
self.dispatch_source(
id,
ExpertSource::Memory {
bytes: host_bytes,
keepalive: None,
},
e,
)
}
pub(crate) fn dispatch_source(
&mut self,
id: BlockId,
source: ExpertSource<'_>,
e: &Engine,
) -> Result<DispatchSlot, Box<dyn std::error::Error>> {
self.reap_copy_sources();
let increment = self.frequency_increment(id);
*self.frequencies.entry(id).or_insert(0.0) += increment;
if let Some(s) = self.table.get(&id).copied() {
self.hits += 1;
self.on_hit(s);
return Ok(DispatchSlot::Resident(s));
}
if let Some(pending) = self.pending.remove(&id) {
if let Err(err) = e.compute_wait(pending.ready.as_ref()) {
self.pending.insert(id, pending);
return Err(err);
}
self.misses += 1;
let slot = pending.slot;
if let Some(keepalive) = pending.keepalive {
self.inflight_sources.push((pending.ready, keepalive));
}
self.publish(id, slot);
return Ok(DispatchSlot::Resident(slot));
}
self.misses += 1;
match source {
ExpertSource::Memory { bytes, keepalive } => {
self.retain_compute_source(keepalive);
let slot = self.admit(id, bytes, e)?;
Ok(DispatchSlot::Resident(slot))
}
ExpertSource::Disk {
file,
offset,
len,
fallback,
keepalive,
} => {
self.retain_compute_source(Some(keepalive));
self.dispatch_disk(id, file, offset, len, fallback, e)
}
}
}
pub fn prefetch(
&mut self,
id: BlockId,
host_bytes: &[u8],
keep: &[BlockId],
e: &Engine,
) -> Result<bool, Box<dyn std::error::Error>> {
self.prefetch_source(
id,
ExpertSource::Memory {
bytes: host_bytes,
keepalive: None,
},
keep,
e,
)
}
fn reserve_prefetch_slot(&mut self, required: usize, keep: &[BlockId]) -> Option<usize> {
for class in &mut self.classes {
if class.capacity >= required
&& let Some(slot) = class.free.pop()
{
return Some(slot);
}
}
self.evict_one_excluding(required, keep)
}
fn prefetch_bytes(
&mut self,
id: BlockId,
host_bytes: &[u8],
keepalive: Option<ExpertKeepalive>,
keep: &[BlockId],
e: &Engine,
) -> Result<bool, Box<dyn std::error::Error>> {
let Some(slot) = self.reserve_prefetch_slot(host_bytes.len(), keep) else {
return Ok(false);
};
let ready = match stage_on_copy_stream(e, host_bytes, &mut self.slots[slot]) {
Ok(ready) => ready,
Err((err, reusable)) => {
if reusable {
self.release_reserved_slot(slot);
} else {
self.copy_stream_unknown = true;
if let Some(keepalive) = keepalive {
self.quarantined_sources.push(keepalive);
}
eprintln!(
"[moe-cache] quarantining slot {slot} after unprovable copy completion"
);
}
return Err(err);
}
};
self.occupant[slot] = Some(id);
self.pending.insert(
id,
PendingBlock {
slot,
ready,
keepalive,
},
);
self.staged_bytes += host_bytes.len() as u64;
Ok(true)
}
pub(crate) fn prefetch_source(
&mut self,
id: BlockId,
source: ExpertSource<'_>,
keep: &[BlockId],
e: &Engine,
) -> Result<bool, Box<dyn std::error::Error>> {
self.reap_copy_sources();
if self.table.contains_key(&id)
|| self.pending.contains_key(&id)
|| self.worker_reads.contains_key(&id)
{
return Ok(false);
}
match source {
ExpertSource::Memory { bytes, keepalive } => {
self.prefetch_bytes(id, bytes, keepalive, keep, e)
}
ExpertSource::Disk {
file,
offset,
len,
fallback,
keepalive,
} => {
if self.pread.as_ref().is_some_and(PreadPool::is_worker) {
match self.pread.as_mut().unwrap().submit_worker_speculative(
file.clone(),
offset,
len,
) {
Ok(Some(ticket)) => {
self.worker_reads.insert(id, WorkerRead { ticket, len });
Ok(true)
}
Ok(None) => Ok(false),
Err(err) => {
self.note_pread_fallback(err.as_ref());
Ok(false)
}
}
} else if self.pread.is_some() {
Ok(false)
} else {
self.prefetch_bytes(id, fallback, Some(keepalive), keep, e)
}
}
}
}
pub fn force_admit(
&mut self,
id: BlockId,
host_bytes: &[u8],
e: &Engine,
) -> Result<usize, Box<dyn std::error::Error>> {
if let Some(s) = self.table.get(&id).copied() {
return Ok(s);
}
self.admit(id, host_bytes, e)
}
pub fn export_residency(&self) -> Vec<(u16, u8, u16)> {
self.occupant
.iter()
.flatten()
.map(|id| (id.layer, id.proj, id.ex))
.collect()
}
pub fn restage_block(
&mut self,
id: BlockId,
m: &crate::hybrid::MoeWeights,
e: &Engine,
) -> Result<bool, Box<dyn std::error::Error>> {
if self.table.contains_key(&id) {
return Ok(true);
}
let exps = match id.proj {
PROJ_GATE => &m.gate_exps,
PROJ_UP => &m.up_exps,
PROJ_DOWN => &m.down_exps,
_ => return Ok(false),
};
if id.ex as usize >= exps.n_expert {
return Ok(false);
}
if m.active_experts
.as_ref()
.is_some_and(|active| !active[id.ex as usize])
{
return Ok(false);
}
if exps.expert_layout(id.ex as usize).len == 0 {
return Ok(false);
}
match exps.expert_source(id.ex as usize) {
ExpertSource::Memory { bytes, keepalive } => {
self.retain_compute_source(keepalive);
self.admit(id, bytes, e)?;
}
ExpertSource::Disk {
fallback,
keepalive,
..
} => {
self.retain_compute_source(Some(keepalive));
self.admit(id, fallback, e)?;
}
}
Ok(true)
}
pub fn prewarm_layer(
&mut self,
layer: u16,
m: &crate::hybrid::MoeWeights,
e: &Engine,
) -> Result<(), Box<dyn std::error::Error>> {
if !self.prewarm_tried.insert(layer) {
return Ok(());
}
let n_expert = m.gate_exps.n_expert;
if self.pread.is_some()
&& (0..n_expert).any(|ex| {
matches!(m.gate_exps.expert_source(ex), ExpertSource::Disk { .. })
|| matches!(m.up_exps.expert_source(ex), ExpertSource::Disk { .. })
|| matches!(m.down_exps.expert_source(ex), ExpertSource::Disk { .. })
})
{
return Ok(());
}
let resident = self.per_layer.get(&layer).copied().unwrap_or(0) as usize;
let missing = 3 * n_expert - resident;
if self.size_aware {
return Ok(());
} if self
.classes
.iter()
.map(|class| class.free.len())
.sum::<usize>()
< missing
{
return Ok(()); }
for ex in 0..n_expert {
for (proj, exps) in [
(PROJ_GATE, &m.gate_exps),
(PROJ_UP, &m.up_exps),
(PROJ_DOWN, &m.down_exps),
] {
let id = BlockId::new(layer, proj, ex as u16);
if self.table.contains_key(&id) {
continue;
}
match exps.expert_source(ex) {
ExpertSource::Memory { bytes, keepalive } => {
self.retain_compute_source(keepalive);
self.admit(id, bytes, e)?;
}
ExpertSource::Disk {
fallback,
keepalive,
..
} => {
self.retain_compute_source(Some(keepalive));
self.admit(id, fallback, e)?;
}
}
}
}
Ok(())
}
pub fn layer_dev_row(
&mut self,
layer: u16,
n_expert: usize,
e: &Engine,
) -> Result<Option<&CudaSlice<u64>>, Box<dyn std::error::Error>> {
if self.per_layer.get(&layer).copied().unwrap_or(0) as usize != 3 * n_expert {
return Ok(None);
}
if !self.dev_rows.contains_key(&layer) {
use cudarc::driver::DevicePtr;
let mut host = vec![0u64; 3 * n_expert];
for proj in 0..3u8 {
for ex in 0..n_expert {
let Some(&s) = self.table.get(&BlockId::new(layer, proj, ex as u16)) else {
return Ok(None);
};
let __s_ev = e.stream();
let (p, _ev) = self.slots[s].device_ptr(&__s_ev);
host[proj as usize * n_expert + ex] = p;
}
}
let row = e.stream().clone_htod(&host)?;
self.dev_rows.insert(layer, row);
}
Ok(self.dev_rows.get(&layer))
}
#[inline]
pub fn buf(&self, d: DispatchSlot) -> &CudaSlice<u8> {
match d {
DispatchSlot::Resident(s) => &self.slots[s],
}
}
#[inline]
pub fn slot(&self, s: usize) -> &CudaSlice<u8> {
&self.slots[s]
}
pub fn hit_rate(&self) -> f64 {
let tot = self.hits + self.misses;
if tot == 0 {
0.0
} else {
self.hits as f64 / tot as f64
}
}
pub fn reset_counters(&mut self) {
self.hits = 0;
self.misses = 0;
self.staged_bytes = 0;
}
pub(crate) fn pread_stats(&self) -> Option<PreadStats> {
if !self.pread_requested {
return None;
}
let mut stats = self
.pread
.as_ref()
.map(PreadPool::stats)
.unwrap_or_default();
stats.fallbacks = self.pread_fallbacks;
Some(stats)
}
}
fn cache_lfu_decay() -> Option<f32> {
let raw = std::env::var("MEMRA_MOE_LFU_DECAY").ok()?;
match parse_cache_lfu_decay(Some(&raw)) {
Ok(value) => value,
Err(reason) => {
eprintln!(
"[moe-cache] invalid MEMRA_MOE_LFU_DECAY={raw:?} ({reason}); disabling LFU decay"
);
None
}
}
}
fn cache_lfu_mtp_weight() -> f32 {
const DEFAULT: f32 = 1.0;
let raw = std::env::var("MEMRA_MOE_LFU_MTP_WEIGHT").ok();
match parse_cache_lfu_mtp_weight(raw.as_deref()) {
Ok(value) => value,
Err(reason) => {
eprintln!(
"[moe-cache] invalid MEMRA_MOE_LFU_MTP_WEIGHT={:?} ({reason}); using {DEFAULT}",
raw.as_deref().unwrap_or("")
);
DEFAULT
}
}
}
fn parse_cache_lfu_mtp_weight(raw: Option<&str>) -> Result<f32, &'static str> {
let value = raw
.unwrap_or("1")
.parse::<f32>()
.map_err(|_| "expected a number")?;
if value.is_finite() && (0.25..=64.0).contains(&value) {
Ok(value)
} else {
Err("expected a finite multiplier from 0.25 through 64")
}
}
fn parse_cache_lfu_decay(raw: Option<&str>) -> Result<Option<f32>, &'static str> {
let Some(raw) = raw else { return Ok(None) };
let value = raw.parse::<f32>().map_err(|_| "expected a number")?;
if value.is_finite() && value > 0.0 && value <= 1.0 {
Ok(Some(value))
} else {
Err("expected a finite fraction greater than 0 and at most 1")
}
}
fn cache_hard_vram_frac() -> f64 {
const DEFAULT: f64 = 0.80;
let raw = std::env::var("MEMRA_MOE_HARD_VRAM_FRAC").ok();
match parse_cache_hard_vram_frac(raw.as_deref()) {
Ok(value) => value,
Err(reason) => {
eprintln!(
"[moe-cache] invalid MEMRA_MOE_HARD_VRAM_FRAC={:?} ({reason}); using {DEFAULT}",
raw.as_deref().unwrap_or("")
);
DEFAULT
}
}
}
fn parse_cache_hard_vram_frac(raw: Option<&str>) -> Result<f64, &'static str> {
let value = raw
.unwrap_or("0.80")
.parse::<f64>()
.map_err(|_| "expected a number")?;
if value.is_finite() && (0.10..=0.95).contains(&value) {
Ok(value)
} else {
Err("expected a finite fraction from 0.10 through 0.95")
}
}
#[cfg(test)]
mod slru_intrusive_tests {
use super::{SEG_PROBATION, SEG_PROTECTED, SlotClass, SlotLink, SlruList};
use std::collections::VecDeque;
struct OldSlru {
probation: VecDeque<usize>,
protected: VecDeque<usize>,
protected_cap: usize,
}
impl OldSlru {
fn on_hit_full(&mut self, slot: usize) {
if let Some(pos) = self.probation.iter().position(|&x| x == slot) {
self.probation.remove(pos);
self.push_protected(slot);
} else if let Some(pos) = self.protected.iter().position(|&x| x == slot) {
self.protected.remove(pos);
self.protected.push_back(slot); } else {
self.push_protected(slot);
}
}
fn push_protected(&mut self, slot: usize) {
self.protected.push_back(slot);
while self.protected.len() > self.protected_cap {
if let Some(demoted) = self.protected.pop_front() {
self.probation.push_back(demoted);
} else {
break;
}
}
}
fn pop_lru(&mut self) -> Option<usize> {
self.probation
.pop_front()
.or_else(|| self.protected.pop_front())
}
fn take_excluding(&mut self, banned: &[usize]) -> Option<usize> {
let take = |q: &mut VecDeque<usize>| {
q.iter()
.position(|&s| !banned.contains(&s))
.and_then(|pos| q.remove(pos))
};
take(&mut self.probation).or_else(|| take(&mut self.protected))
}
}
fn new_pair(n: usize, protected_cap: usize) -> (SlotClass, Vec<SlotLink>, OldSlru) {
let class = SlotClass {
capacity: 1,
probation: SlruList::new(),
protected: SlruList::new(),
free: Vec::new(),
protected_cap,
};
let links = vec![SlotLink::none(); n];
let old = OldSlru {
probation: VecDeque::new(),
protected: VecDeque::new(),
protected_cap,
};
(class, links, old)
}
fn orders_match(class: &SlotClass, links: &[SlotLink], old: &OldSlru) -> bool {
let np: Vec<usize> = class.probation.iter(links).collect();
let nt: Vec<usize> = class.protected.iter(links).collect();
let op: Vec<usize> = old.probation.iter().copied().collect();
let ot: Vec<usize> = old.protected.iter().copied().collect();
np == op && nt == ot
}
struct Rng(u64);
impl Rng {
fn next(&mut self) -> u64 {
self.0 = self.0.wrapping_add(0x9E3779B97F4A7C15);
let mut z = self.0;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
z ^ (z >> 31)
}
fn below(&mut self, n: usize) -> usize {
(self.next() % n as u64) as usize
}
}
#[test]
fn same_eviction_decisions_randomized_soak() {
for &(n, cap) in &[(8usize, 1usize), (16, 12), (64, 51), (128, 102)] {
let (mut class, mut links, mut old) = new_pair(n, cap);
let mut rng = Rng(0xC0FFEE ^ (n as u64) << 8 ^ cap as u64);
let mut resident: Vec<usize> = Vec::new();
let mut free: Vec<usize> = (0..n).rev().collect();
for step in 0..200_000 {
let op = rng.below(100);
if op < 55 && !resident.is_empty() {
let slot = resident[rng.below(resident.len())];
class.on_hit_full(slot, &mut links);
old.on_hit_full(slot);
} else if op < 80 {
let slot = if let Some(s) = free.pop() {
s
} else {
let v_new = class.pop_lru(&mut links);
let v_old = old.pop_lru();
assert_eq!(
v_new, v_old,
"victim diverged at step {step} (n={n} cap={cap})"
);
let v = v_new.unwrap();
resident.retain(|&s| s != v);
v
};
class.probation.push_back(slot, SEG_PROBATION, &mut links);
old.probation.push_back(slot);
resident.push(slot);
} else if op < 92 && resident.len() > 2 {
let banned: Vec<usize> = (0..3.min(resident.len()))
.map(|_| resident[rng.below(resident.len())])
.collect();
let take_new = {
let q = &mut class.probation;
let found = q.iter(&links).find(|s| !banned.contains(s));
match found {
Some(s) => {
q.unlink(s, &mut links);
Some(s)
}
None => {
let q = &mut class.protected;
q.iter(&links).find(|s| !banned.contains(s)).inspect(|&s| {
q.unlink(s, &mut links);
})
}
}
};
let take_old = old.take_excluding(&banned);
assert_eq!(
take_new, take_old,
"excluding-victim diverged at step {step}"
);
if let Some(v) = take_new {
resident.retain(|&s| s != v);
free.push(v);
}
} else if !resident.is_empty() {
let slot = resident[rng.below(resident.len())];
match links[slot].seg {
SEG_PROBATION => class.probation.unlink(slot, &mut links),
SEG_PROTECTED => class.protected.unlink(slot, &mut links),
_ => {}
}
if let Some(pos) = old.probation.iter().position(|&x| x == slot) {
old.probation.remove(pos);
} else if let Some(pos) = old.protected.iter().position(|&x| x == slot) {
old.protected.remove(pos);
}
class.on_hit_full(slot, &mut links);
old.on_hit_full(slot);
}
assert!(
orders_match(&class, &links, &old),
"segment order diverged at step {step} (n={n} cap={cap})"
);
}
}
}
#[test]
fn hit_promotion_is_o1_not_on() {
fn bench(n: usize, hits: usize) -> std::time::Duration {
let (mut class, mut links, _) = new_pair(n, (n as f64 * 0.8) as usize);
for s in 0..n {
class.probation.push_back(s, SEG_PROBATION, &mut links);
}
let mut rng = Rng(0xBEEF);
let t0 = std::time::Instant::now();
for _ in 0..hits {
class.on_hit_full(rng.below(n), &mut links);
}
t0.elapsed()
}
bench(1_000, 10_000);
bench(46_000, 10_000);
let small = bench(1_000, 850_000).as_secs_f64() / 850_000.0;
let large = bench(46_000, 850_000).as_secs_f64() / 850_000.0;
assert!(
large < small * 8.0,
"per-hit cost scaled with n_slots: {:.1}ns @1k vs {:.1}ns @46k",
small * 1e9,
large * 1e9
);
}
#[test]
fn slru_list_basic_invariants() {
let mut links = vec![SlotLink::none(); 4];
let mut l = SlruList::new();
assert_eq!(l.pop_front(&mut links), None);
l.push_back(2, SEG_PROBATION, &mut links);
l.push_back(0, SEG_PROBATION, &mut links);
l.push_back(3, SEG_PROBATION, &mut links);
assert_eq!(l.iter(&links).collect::<Vec<_>>(), vec![2, 0, 3]);
assert_eq!(l.len, 3);
l.unlink(0, &mut links); assert_eq!(l.iter(&links).collect::<Vec<_>>(), vec![2, 3]);
l.unlink(3, &mut links); assert_eq!(l.iter(&links).collect::<Vec<_>>(), vec![2]);
assert_eq!(l.pop_front(&mut links), Some(2)); assert_eq!(l.len, 0);
assert_eq!(l.head, super::NIL);
assert_eq!(l.tail, super::NIL);
assert!(links.iter().all(|k| k.seg == super::SEG_NONE));
}
}
#[cfg(test)]
mod vram_fraction_tests {
use super::{
parse_cache_hard_vram_frac, parse_cache_lfu_decay, parse_cache_lfu_mtp_weight,
size_class_plan,
};
#[test]
fn hard_vram_fraction_defaults_and_rejects_unsafe_values() {
assert_eq!(parse_cache_hard_vram_frac(None), Ok(0.80));
assert_eq!(parse_cache_hard_vram_frac(Some("0.82")), Ok(0.82));
assert!(parse_cache_hard_vram_frac(Some("NaN")).is_err());
assert_eq!(parse_cache_hard_vram_frac(Some("0.95")), Ok(0.95));
assert!(parse_cache_hard_vram_frac(Some("0.96")).is_err());
assert!(parse_cache_hard_vram_frac(Some("1.0")).is_err());
assert!(parse_cache_hard_vram_frac(Some("bad")).is_err());
}
#[test]
fn lfu_decay_is_opt_in_and_bounded() {
assert_eq!(parse_cache_lfu_decay(None), Ok(None));
assert_eq!(parse_cache_lfu_decay(Some("0.8")), Ok(Some(0.8)));
assert_eq!(parse_cache_lfu_decay(Some("1")), Ok(Some(1.0)));
for value in ["0", "-0.1", "1.1", "NaN", "bad"] {
assert!(
parse_cache_lfu_decay(Some(value)).is_err(),
"accepted {value}"
);
}
}
#[test]
fn lfu_mtp_weight_defaults_and_is_bounded() {
assert_eq!(parse_cache_lfu_mtp_weight(None), Ok(1.0));
assert_eq!(parse_cache_lfu_mtp_weight(Some("4")), Ok(4.0));
for value in ["0", "0.1", "65", "NaN", "bad"] {
assert!(
parse_cache_lfu_mtp_weight(Some(value)).is_err(),
"accepted {value}"
);
}
}
#[test]
fn size_class_plan_preserves_classes_and_never_exceeds_budget() {
let blocks = [100usize, 100, 100, 200, 200, 400];
let budget = (108 * 2) + 208 + 408;
let plan = size_class_plan(&blocks, budget);
assert!(plan.iter().all(|(_, count)| *count > 0));
assert!(
plan.iter()
.map(|(bytes, count)| (bytes + 8) * count)
.sum::<usize>()
<= budget
);
assert!(plan.iter().all(|(bytes, count)| {
*count <= blocks.iter().filter(|block| **block == *bytes).count()
}));
}
#[test]
fn size_class_plan_returns_full_inventory_when_it_fits() {
let blocks = [100usize, 100, 200, 400];
let budget: usize = blocks.iter().map(|bytes| bytes + 8).sum();
assert_eq!(
size_class_plan(&blocks, budget),
vec![(100, 2), (200, 1), (400, 1)]
);
}
#[test]
fn size_class_plan_does_not_overflow_on_pathological_sizes() {
let plan = size_class_plan(&[usize::MAX, usize::MAX], usize::MAX);
assert!(plan.is_empty());
}
}
impl Drop for MoeSlotCache {
fn drop(&mut self) {
let mut safe_to_drop_slots = true;
if self.compute_stream_unknown || !self.compute_sources.is_empty() {
if let Err(err) = self.compute_stream.synchronize() {
safe_to_drop_slots = false;
eprintln!(
"[moe-cache] unknown compute-stream drain failed ({err}); leaking GPU slots for safety"
);
for (_, keepalive) in self.compute_sources.drain() {
std::mem::forget(keepalive);
}
} else {
self.compute_stream_unknown = false;
self.compute_sources.clear();
}
}
let need_copy_drain = self.copy_stream_unknown
|| !self.pending.is_empty()
|| !self.inflight_sources.is_empty()
|| !self.quarantined_sources.is_empty();
if need_copy_drain {
if let Err(err) = self.copy_stream.synchronize() {
safe_to_drop_slots = false;
eprintln!(
"[moe-cache] unknown copy-stream drain failed ({err}); leaking GPU slots for safety"
);
for (_, keepalive) in self.inflight_sources.drain(..) {
std::mem::forget(keepalive);
}
for keepalive in self.quarantined_sources.drain(..) {
std::mem::forget(keepalive);
}
for (_, pending) in self.pending.drain() {
if let Some(keepalive) = pending.keepalive {
std::mem::forget(keepalive);
}
}
} else {
self.copy_stream_unknown = false;
self.inflight_sources.clear();
self.quarantined_sources.clear();
self.pending.clear();
}
}
if let Some(pool) = self.pread.as_mut() {
safe_to_drop_slots &= pool.drain();
} else if self.pread_requested && self.pread_fallbacks != 0 {
eprintln!(
"[spill-pread] backend unavailable; mmap_fallbacks={}",
self.pread_fallbacks
);
}
if !safe_to_drop_slots {
for slot in self.slots.drain(..) {
std::mem::forget(slot);
}
}
}
}