use crate::qstar::QStarPolicy;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ExpertId {
pub layer: u32,
pub expert: u32,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct CopyPlan {
pub dst_slots: Vec<u32>,
pub src_rows: Vec<u32>,
}
impl CopyPlan {
pub fn len(&self) -> usize {
self.dst_slots.len()
}
pub fn is_empty(&self) -> bool {
self.dst_slots.is_empty()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct EnsurePlan {
pub slots: Vec<Option<u32>>,
pub copy: CopyPlan,
pub active: usize,
pub missing: usize,
pub fetched: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum FetchOrder {
#[default]
ByRecency,
LowestId,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct ExpertCacheStats {
pub calls: u64,
pub active: u64,
pub missing: u64,
pub fetched: u64,
}
impl ExpertCacheStats {
pub fn miss_rate(&self) -> f64 {
if self.active == 0 {
return 0.0;
}
self.missing as f64 / self.active as f64
}
pub fn fetch_rate(&self) -> f64 {
if self.missing == 0 {
return 0.0;
}
self.fetched as f64 / self.missing as f64
}
pub fn active_per_call(&self) -> f64 {
if self.calls == 0 {
return 0.0;
}
self.active as f64 / self.calls as f64
}
pub fn missing_per_call(&self) -> f64 {
if self.calls == 0 {
return 0.0;
}
self.missing as f64 / self.calls as f64
}
pub fn fetched_per_call(&self) -> f64 {
if self.calls == 0 {
return 0.0;
}
self.fetched as f64 / self.calls as f64
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct LayerRoutingSkew {
pub layer: u32,
pub routed: u64,
pub working_set: usize,
pub experts_for_90pct: usize,
pub norm_entropy: f64,
pub oracle_hit_at_slots: f64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct RoutingSkewReport {
pub slots_per_layer: f64,
pub oracle_slots: usize,
pub working_set_mean: f64,
pub working_set_max: usize,
pub experts_for_90pct: f64,
pub oracle_hit_at_slots: f64,
pub norm_entropy: f64,
pub per_layer: Vec<LayerRoutingSkew>,
}
const COVERAGE_FRACTION: f64 = 0.9;
pub const SMALL_BANK_FEAT_BYTES: u64 = 256 * 1024;
const UNEVICTABLE: i64 = i64::MAX;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct MissRun {
pub start: u32,
pub len: u32,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct GatherPlan {
pub dst_slots: Vec<u32>,
pub src_slots: Vec<u32>,
}
impl GatherPlan {
pub fn len(&self) -> usize {
self.dst_slots.len()
}
pub fn is_empty(&self) -> bool {
self.dst_slots.is_empty()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct BankEntry {
pub bank: usize,
pub dst_slot: u32,
pub src_row: u32,
pub rows: u32,
pub bytes: u64,
pub whole_layer: bool,
}
#[derive(Debug, Clone, PartialEq)]
pub struct PrefillPlan {
pub layer: u32,
pub buffer_id: u32,
pub already_loaded: bool,
pub gather: GatherPlan,
pub miss_runs: Vec<MissRun>,
pub entries: Vec<BankEntry>,
pub gather_banks: Vec<usize>,
}
#[derive(Debug)]
pub struct ExpertCache {
num_layers: usize,
num_experts: usize,
cache_size: usize,
slot_for_id: Vec<i64>,
id_of_slot: Vec<i64>,
usage: Vec<i64>,
step: i64,
recency: Vec<i64>,
fetch_order: FetchOrder,
stats: ExpertCacheStats,
layer_stats: Vec<ExpertCacheStats>,
collect_routing: bool,
routing_freq: Vec<u64>,
prefill_layer: [Option<u32>; 2],
prefill_released: [bool; 2],
prefill_snapshot: Option<Vec<i64>>,
prefill_hit_rows: u64,
prefill_total_rows: u64,
}
impl ExpertCache {
pub fn new(num_layers: usize, num_experts: usize, cache_size: usize) -> Self {
assert!(
num_layers > 0 && num_experts > 0,
"a MoE model has layers and experts"
);
assert!(
cache_size >= num_experts,
"cache_size {cache_size} cannot hold one layer of {num_experts} experts"
);
let total_ids = num_layers * num_experts;
ExpertCache {
num_layers,
num_experts,
cache_size,
slot_for_id: vec![-1; total_ids],
id_of_slot: vec![-1; cache_size],
usage: vec![0; cache_size],
step: 0,
recency: vec![-1; total_ids],
fetch_order: FetchOrder::default(),
stats: ExpertCacheStats::default(),
layer_stats: vec![ExpertCacheStats::default(); num_layers],
collect_routing: false,
routing_freq: vec![0; total_ids],
prefill_layer: [None, None],
prefill_released: [true, true],
prefill_snapshot: None,
prefill_hit_rows: 0,
prefill_total_rows: 0,
}
}
pub fn with_fetch_order(mut self, order: FetchOrder) -> Self {
self.fetch_order = order;
self
}
pub fn num_layers(&self) -> usize {
self.num_layers
}
pub fn num_experts(&self) -> usize {
self.num_experts
}
pub fn cache_size(&self) -> usize {
self.cache_size
}
pub fn stats(&self) -> ExpertCacheStats {
self.stats
}
pub fn layer_stats(&self, layer: u32) -> ExpertCacheStats {
self.layer_stats[layer as usize]
}
pub fn per_layer_stats(&self) -> &[ExpertCacheStats] {
&self.layer_stats
}
pub fn reset_stats(&mut self) {
self.stats = ExpertCacheStats::default();
for layer in &mut self.layer_stats {
*layer = ExpertCacheStats::default();
}
self.prefill_hit_rows = 0;
self.prefill_total_rows = 0;
}
pub fn total_experts(&self) -> usize {
self.num_layers * self.num_experts
}
fn flat(&self, layer: u32, expert: u32) -> usize {
let layer = layer as usize;
let expert = expert as usize;
debug_assert!(layer < self.num_layers && expert < self.num_experts);
layer * self.num_experts + expert
}
pub fn slot_of(&self, layer: u32, expert: u32) -> Option<u32> {
let slot = self.slot_for_id[self.flat(layer, expert)];
(slot >= 0).then_some(slot as u32)
}
pub fn resident_in(&self, slot: u32) -> Option<ExpertId> {
let id = self.id_of_slot[slot as usize];
(id >= 0).then(|| ExpertId {
layer: (id as usize / self.num_experts) as u32,
expert: (id as usize % self.num_experts) as u32,
})
}
pub fn resident_slots(&self) -> usize {
self.id_of_slot.iter().filter(|id| **id >= 0).count()
}
pub fn forget_slot(&mut self, slot: u32) -> Option<ExpertId> {
let id = *self.id_of_slot.get(slot as usize)?;
if id < 0 {
return None;
}
self.id_of_slot[slot as usize] = -1;
self.slot_for_id[id as usize] = -1;
self.usage[slot as usize] = i64::MIN;
Some(ExpertId {
layer: (id as usize / self.num_experts) as u32,
expert: (id as usize % self.num_experts) as u32,
})
}
pub fn reset(&mut self) {
self.slot_for_id.fill(-1);
self.id_of_slot.fill(-1);
self.usage.fill(0);
self.recency.fill(-1);
self.step = 0;
self.reset_stats();
self.reset_routing();
self.forget_prefill_buffers();
}
pub fn rebuild(&mut self, cache_size: usize) -> Result<(), RebuildRejected> {
if cache_size < self.num_experts {
return Err(RebuildRejected {
requested: cache_size,
minimum: self.num_experts,
});
}
self.cache_size = cache_size;
self.id_of_slot = vec![-1; cache_size];
self.usage = vec![0; cache_size];
self.slot_for_id.fill(-1);
self.recency.fill(-1);
self.step = 0;
self.reset_stats();
self.reset_routing();
self.forget_prefill_buffers();
Ok(())
}
pub fn ensure(&mut self, layer: u32, expert_ids: &[u32]) -> EnsurePlan {
self.ensure_split(layer, expert_ids, None)
}
pub fn ensure_hybrid(
&mut self,
layer: u32,
expert_ids: &[u32],
policy: &QStarPolicy,
) -> EnsurePlan {
self.ensure_split(layer, expert_ids, Some(policy))
}
fn ensure_split(
&mut self,
layer: u32,
expert_ids: &[u32],
policy: Option<&QStarPolicy>,
) -> EnsurePlan {
self.step += 1;
let step = self.step;
let base = self.flat(layer, 0);
if self.collect_routing {
for expert in expert_ids {
self.routing_freq[base + *expert as usize] += 1;
}
}
let mut distinct: Vec<u32> = Vec::with_capacity(expert_ids.len());
for id in expert_ids {
if !distinct.contains(id) {
distinct.push(*id);
}
}
let mut missing: Vec<u32> = Vec::new();
for expert in &distinct {
let slot = self.slot_for_id[base + *expert as usize];
if slot >= 0 {
self.usage[slot as usize] = step;
} else {
missing.push(*expert);
}
}
match (policy, self.fetch_order) {
(Some(_), FetchOrder::ByRecency) => {
missing.sort_by_key(|e| (-self.recency[base + *e as usize], *e));
}
_ => missing.sort_unstable(),
}
let num_missing = missing.len();
let num_fetch = match policy {
Some(policy) => policy.split(num_missing).fetch,
None => num_missing,
};
let mut evictable: Vec<i64> = self
.usage
.iter()
.map(|u| if *u == step { UNEVICTABLE } else { *u })
.collect();
if policy.is_some() {
for (evict, id) in evictable.iter_mut().zip(self.id_of_slot.iter()) {
if *id < 0 {
continue;
}
let owner = *id - base as i64;
if (0..self.num_experts as i64).contains(&owner)
&& distinct.contains(&(owner as u32))
{
*evict = UNEVICTABLE;
}
}
}
let mut copy = CopyPlan::default();
for expert in missing.iter().take(num_fetch) {
let victim = argmin_slot(&evictable);
let old = self.id_of_slot[victim];
if old >= 0 {
self.slot_for_id[old as usize] = -1;
}
let id = base + *expert as usize;
self.id_of_slot[victim] = id as i64;
self.slot_for_id[id] = victim as i64;
self.usage[victim] = step;
evictable[victim] = UNEVICTABLE;
copy.dst_slots.push(victim as u32);
copy.src_rows.push(*expert);
}
let slots: Vec<Option<u32>> = expert_ids
.iter()
.map(|expert| {
let slot = self.slot_for_id[base + *expert as usize];
(slot >= 0).then_some(slot as u32)
})
.collect();
if policy.is_some() {
for expert in &distinct {
self.recency[base + *expert as usize] = step;
}
}
self.stats.calls += 1;
self.stats.active += distinct.len() as u64;
self.stats.missing += num_missing as u64;
self.stats.fetched += num_fetch as u64;
let per_layer = &mut self.layer_stats[layer as usize];
per_layer.calls += 1;
per_layer.active += distinct.len() as u64;
per_layer.missing += num_missing as u64;
per_layer.fetched += num_fetch as u64;
EnsurePlan {
slots,
copy,
active: distinct.len(),
missing: num_missing,
fetched: num_fetch,
}
}
pub fn materialize_layer(&mut self, layer: u32) -> CopyPlan {
let base = self.flat(layer, 0);
let owns = |id: i64| id >= base as i64 && id < (base + self.num_experts) as i64;
let previous: Vec<i64> = self.id_of_slot.clone();
for (slot, old) in previous.iter().enumerate() {
if owns(*old) {
self.id_of_slot[slot] = -1;
self.usage[slot] = 0;
}
}
for old in previous.iter().take(self.num_experts) {
if *old >= 0 && !owns(*old) {
self.slot_for_id[*old as usize] = -1;
}
}
self.step += 1;
let mut copy = CopyPlan::default();
for expert in 0..self.num_experts {
self.id_of_slot[expert] = (base + expert) as i64;
self.slot_for_id[base + expert] = expert as i64;
self.usage[expert] = self.step;
copy.dst_slots.push(expert as u32);
copy.src_rows.push(expert as u32);
}
copy
}
pub fn set_collect_routing(&mut self, on: bool) {
self.collect_routing = on;
}
pub fn collects_routing(&self) -> bool {
self.collect_routing
}
pub fn routing_histogram(&self, layer: u32) -> &[u64] {
let base = self.flat(layer, 0);
&self.routing_freq[base..base + self.num_experts]
}
pub fn reset_routing(&mut self) {
self.routing_freq.fill(0);
}
pub fn routing_skew(&self) -> Option<RoutingSkewReport> {
let experts = self.num_experts;
let slots_per_layer = self.cache_size as f64 / self.num_layers as f64;
let oracle_slots = (slots_per_layer.round() as usize).max(1);
let ln_experts = (experts as f64).ln();
let mut per_layer: Vec<LayerRoutingSkew> = Vec::new();
for layer in 0..self.num_layers {
let freq = &self.routing_freq[layer * experts..(layer + 1) * experts];
let routed: u64 = freq.iter().sum();
if routed == 0 {
continue;
}
let total = routed as f64;
let working_set = freq.iter().filter(|f| **f > 0).count();
let mut descending: Vec<u64> = freq.to_vec();
descending.sort_unstable_by(|a, b| b.cmp(a));
let head: u64 = descending.iter().take(oracle_slots).sum();
let oracle_hit_at_slots = head as f64 / total;
let mut cumulative = 0u64;
let mut below = 0usize;
for count in &descending {
cumulative += *count;
if cumulative as f64 / total >= COVERAGE_FRACTION {
break;
}
below += 1;
}
let mut entropy = 0.0f64;
for count in freq {
if *count == 0 {
continue;
}
let p = *count as f64 / total;
entropy -= p * p.max(1e-12).ln();
}
let norm_entropy = if ln_experts > 0.0 {
entropy / ln_experts
} else {
0.0
};
per_layer.push(LayerRoutingSkew {
layer: layer as u32,
routed,
working_set,
experts_for_90pct: below + 1,
norm_entropy,
oracle_hit_at_slots,
});
}
if per_layer.is_empty() {
return None;
}
let layers = per_layer.len() as f64;
Some(RoutingSkewReport {
slots_per_layer,
oracle_slots,
working_set_mean: per_layer.iter().map(|l| l.working_set as f64).sum::<f64>() / layers,
working_set_max: per_layer.iter().map(|l| l.working_set).max().unwrap_or(0),
experts_for_90pct: per_layer
.iter()
.map(|l| l.experts_for_90pct as f64)
.sum::<f64>()
/ layers,
oracle_hit_at_slots: per_layer.iter().map(|l| l.oracle_hit_at_slots).sum::<f64>()
/ layers,
norm_entropy: per_layer.iter().map(|l| l.norm_entropy).sum::<f64>() / layers,
per_layer,
})
}
pub fn prefill_buffer_slots(&self) -> usize {
2 * self.num_experts
}
pub fn prefill_overlap_fits(&self) -> bool {
self.cache_size >= self.prefill_buffer_slots()
}
pub fn prefill_buffer_layer(&self, buffer_id: u32) -> Option<u32> {
self.prefill_layer[buffer_id as usize]
}
pub fn prefill_hit_rows(&self) -> u64 {
self.prefill_hit_rows
}
pub fn prefill_rows(&self) -> u64 {
self.prefill_total_rows
}
pub fn begin_prefill(&mut self) {
assert!(
self.prefill_overlap_fits(),
"cache of {} slots cannot lend {} to the prefill buffers",
self.cache_size,
self.prefill_buffer_slots()
);
self.prefill_layer = [None, None];
self.prefill_released = [true, true];
self.prefill_snapshot = Some(self.slot_for_id.clone());
}
pub fn prefetch_prefill_layer(&mut self, layer: u32, bank_feat_bytes: &[u64]) -> PrefillPlan {
assert!(
(layer as usize) < self.num_layers,
"layer {layer} is outside a model of {} layers",
self.num_layers
);
let experts = self.num_experts;
let buffer_id = (layer as usize) % 2;
let buffer_base = buffer_id * experts;
let gather_banks: Vec<usize> = bank_feat_bytes
.iter()
.enumerate()
.filter(|(_, feat)| **feat >= SMALL_BANK_FEAT_BYTES)
.map(|(bank, _)| bank)
.collect();
if self.prefill_layer[buffer_id] == Some(layer) {
return PrefillPlan {
layer,
buffer_id: buffer_id as u32,
already_loaded: true,
gather: GatherPlan::default(),
miss_runs: Vec::new(),
entries: Vec::new(),
gather_banks,
};
}
if let Some(held) = self.prefill_layer[buffer_id] {
assert!(
self.prefill_released[buffer_id],
"prefill buffer {buffer_id} still holds layer {held}; staging layer \
{layer} into it would overwrite bytes a running GEMM is reading"
);
}
let snapshot = self
.prefill_snapshot
.as_ref()
.expect("begin_prefill must open the chunk before a layer is staged");
let base = layer as usize * experts;
let threshold = self.prefill_buffer_slots() as i64;
let mut gather = GatherPlan::default();
let mut missing: Vec<u32> = Vec::new();
for expert in 0..experts {
let slot = snapshot[base + expert];
if slot >= threshold {
gather.dst_slots.push((buffer_base + expert) as u32);
gather.src_slots.push(slot as u32);
} else {
missing.push(expert as u32);
}
}
self.prefill_hit_rows += gather.len() as u64;
self.prefill_total_rows += experts as u64;
self.invalidate_prefill_buffer(buffer_id);
let miss_runs = coalesce_runs(&missing);
let mut entries: Vec<BankEntry> = Vec::new();
for (bank, feat) in bank_feat_bytes.iter().enumerate() {
if *feat < SMALL_BANK_FEAT_BYTES {
entries.push(BankEntry {
bank,
dst_slot: buffer_base as u32,
src_row: 0,
rows: experts as u32,
bytes: experts as u64 * feat,
whole_layer: true,
});
continue;
}
for run in &miss_runs {
entries.push(BankEntry {
bank,
dst_slot: buffer_base as u32 + run.start,
src_row: run.start,
rows: run.len,
bytes: run.len as u64 * feat,
whole_layer: false,
});
}
}
self.prefill_layer[buffer_id] = Some(layer);
self.prefill_released[buffer_id] = false;
PrefillPlan {
layer,
buffer_id: buffer_id as u32,
already_loaded: false,
gather,
miss_runs,
entries,
gather_banks,
}
}
pub fn release_prefill_layer(&mut self, layer: u32) {
let buffer_id = (layer as usize) % 2;
if self.prefill_layer[buffer_id] == Some(layer) {
self.prefill_released[buffer_id] = true;
}
}
fn invalidate_prefill_buffer(&mut self, buffer_id: usize) {
let start = buffer_id * self.num_experts;
for slot in start..start + self.num_experts {
let old = self.id_of_slot[slot];
if old >= 0 {
self.slot_for_id[old as usize] = -1;
}
self.id_of_slot[slot] = -1;
self.usage[slot] = 0;
}
}
fn forget_prefill_buffers(&mut self) {
self.prefill_layer = [None, None];
self.prefill_released = [true, true];
self.prefill_snapshot = None;
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RebuildRejected {
pub requested: usize,
pub minimum: usize,
}
impl std::fmt::Display for RebuildRejected {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"expert cache of {} slots cannot hold one layer of {} experts",
self.requested, self.minimum
)
}
}
impl std::error::Error for RebuildRejected {}
fn coalesce_runs(experts: &[u32]) -> Vec<MissRun> {
let mut runs: Vec<MissRun> = Vec::new();
for expert in experts {
match runs.last_mut() {
Some(run) if run.start + run.len == *expert => run.len += 1,
_ => runs.push(MissRun {
start: *expert,
len: 1,
}),
}
}
runs
}
fn argmin_slot(usage: &[i64]) -> usize {
let mut best = 0usize;
let mut best_usage = usage[0];
for (slot, value) in usage.iter().enumerate().skip(1) {
if *value < best_usage {
best = slot;
best_usage = *value;
}
}
best
}
#[cfg(test)]
mod tests {
use super::*;
fn cache() -> ExpertCache {
ExpertCache::new(4, 8, 16)
}
#[test]
fn a_cold_miss_is_fetched_and_becomes_resident() {
let mut cache = cache();
let plan = cache.ensure(0, &[3, 5]);
assert_eq!(plan.missing, 2);
assert_eq!(plan.fetched, 2);
assert_eq!(plan.copy.src_rows, vec![3, 5], "layer-local rows");
assert_eq!(plan.slots.len(), 2);
assert!(plan.slots.iter().all(Option::is_some));
let again = cache.ensure(0, &[3, 5]);
assert_eq!(again.missing, 0);
assert!(again.copy.is_empty());
assert_eq!(again.slots, plan.slots);
}
#[test]
fn layers_share_one_pool_through_a_flat_id_space() {
let mut cache = cache();
cache.ensure(0, &[3]);
let plan = cache.ensure(1, &[3]);
assert_eq!(plan.missing, 1, "layer 1's expert 3 is a different expert");
assert_ne!(cache.slot_of(0, 3), cache.slot_of(1, 3));
assert_eq!(cache.resident_slots(), 2);
}
#[test]
fn duplicate_routes_collapse_into_one_copy_and_one_slot() {
let mut cache = cache();
let plan = cache.ensure(0, &[7, 7, 7, 2]);
assert_eq!(plan.active, 2);
assert_eq!(plan.copy.len(), 2, "one copy per distinct expert");
assert_eq!(plan.slots[0], plan.slots[1]);
assert_eq!(plan.slots[1], plan.slots[2]);
assert_ne!(plan.slots[0], plan.slots[3]);
}
#[test]
fn a_step_never_evicts_an_expert_it_is_about_to_read() {
let mut cache = ExpertCache::new(2, 4, 4);
cache.ensure(0, &[0, 1, 2, 3]);
cache.ensure(1, &[0]);
let held = cache.slot_of(1, 0).unwrap();
let plan = cache.ensure(1, &[0, 1]);
assert_eq!(plan.missing, 1);
assert_eq!(cache.slot_of(1, 0), Some(held), "the hit kept its slot");
assert_ne!(plan.copy.dst_slots[0], held);
assert_eq!(plan.slots[0], Some(held));
}
#[test]
fn eviction_takes_the_least_recently_used_slot() {
let mut cache = ExpertCache::new(2, 4, 4);
for expert in 0..4u32 {
cache.ensure(0, &[expert]); }
cache.ensure(0, &[0]); let oldest = cache.slot_of(0, 1).unwrap();
let plan = cache.ensure(1, &[0]);
assert_eq!(plan.copy.dst_slots, vec![oldest]);
assert_eq!(cache.slot_of(0, 1), None, "the oldest went");
assert!(cache.slot_of(0, 0).is_some(), "the refreshed one stayed");
}
#[test]
fn a_large_cold_batch_assigns_unique_slots() {
let mut cache = ExpertCache::new(4, 64, 256);
let ids: Vec<u32> = (0..64).collect();
let plan = cache.ensure(2, &ids);
assert_eq!(plan.missing, 64);
assert_eq!(plan.copy.len(), 64);
let mut slots: Vec<u32> = plan.slots.iter().map(|s| s.unwrap()).collect();
slots.sort_unstable();
slots.dedup();
assert_eq!(slots.len(), 64, "no slot serves two experts");
assert_eq!(
plan.copy.src_rows, ids,
"misses are ranked by ascending expert id"
);
}
#[test]
fn the_hybrid_split_assigns_every_route_to_exactly_one_device() {
let mut cache = ExpertCache::new(1, 32, 40);
let policy = QStarPolicy::from_fraction(0.415);
let routes: Vec<u32> = (0..8).collect();
let plan = cache.ensure_hybrid(0, &routes, &policy);
assert_eq!(plan.missing, 8);
assert_eq!(plan.fetched, 3, "0.415 * 8 = 3.32 -> 3");
assert_eq!(plan.copy.len(), 3);
let on_gpu = plan.slots.iter().filter(|s| s.is_some()).count();
assert_eq!(on_gpu, 3);
assert_eq!(plan.slots.iter().filter(|s| s.is_none()).count(), 5);
}
#[test]
fn a_recurring_miss_climbs_the_fetch_order() {
let mut cache = ExpertCache::new(1, 16, 16);
let policy = QStarPolicy::fixed_cap(1);
for round in 0..6u32 {
cache.ensure_hybrid(0, &[9, round], &policy);
}
assert!(
cache.slot_of(0, 9).is_some(),
"the expert routed every step became resident"
);
}
#[test]
fn the_fixed_cap_is_the_unbenchmarked_default() {
let mut cache = ExpertCache::new(1, 32, 40);
let plan = cache.ensure_hybrid(0, &[0, 1, 2, 3, 4, 5, 6, 7], &QStarPolicy::fixed_cap(1));
assert_eq!(plan.missing, 8);
assert_eq!(plan.fetched, 1);
assert_eq!(plan.copy.len(), 1);
}
#[test]
fn materializing_a_layer_reclaims_other_layers_slots_cleanly() {
let mut cache = ExpertCache::new(2, 4, 4);
cache.ensure(0, &[0, 1, 2, 3]);
assert_eq!(cache.resident_slots(), 4);
let copy = cache.materialize_layer(1);
assert_eq!(copy.dst_slots, vec![0, 1, 2, 3]);
assert_eq!(copy.src_rows, vec![0, 1, 2, 3], "position == expert id");
for expert in 0..4u32 {
assert_eq!(cache.slot_of(1, expert), Some(expert));
assert_eq!(
cache.slot_of(0, expert),
None,
"layer 0 must not hit on layer 1's bytes"
);
}
let plan = cache.ensure(0, &[3]);
assert_eq!(plan.missing, 1);
}
#[test]
fn a_second_materialize_of_the_same_layer_is_idempotent() {
let mut cache = ExpertCache::new(2, 4, 6);
cache.materialize_layer(1);
cache.materialize_layer(1);
for expert in 0..4u32 {
assert_eq!(cache.slot_of(1, expert), Some(expert));
}
assert_eq!(cache.resident_slots(), 4);
}
#[test]
fn stats_answer_whether_the_cache_is_big_enough() {
let mut cache = ExpertCache::new(1, 2, 2);
cache.ensure(0, &[0, 1]); cache.ensure(0, &[0, 1]); let stats = cache.stats();
assert_eq!(stats.calls, 2);
assert_eq!(stats.active, 4);
assert_eq!(stats.missing, 2);
assert_eq!(stats.miss_rate(), 0.5);
assert_eq!(stats.fetch_rate(), 1.0);
assert_eq!(stats.active_per_call(), 2.0);
}
#[test]
fn a_rebuild_that_cannot_hold_a_layer_changes_nothing() {
let mut cache = ExpertCache::new(2, 8, 16);
cache.ensure(0, &[1]);
let err = cache.rebuild(4).unwrap_err();
assert_eq!(err.minimum, 8);
assert_eq!(cache.cache_size(), 16, "still serving from the old cache");
assert!(cache.slot_of(0, 1).is_some());
cache.rebuild(8).expect("one layer fits");
assert_eq!(cache.cache_size(), 8);
assert_eq!(
cache.resident_slots(),
0,
"slot ids no longer mean anything"
);
}
#[test]
#[should_panic(expected = "cannot hold one layer")]
fn a_cache_too_small_for_one_layer_is_rejected_at_construction() {
ExpertCache::new(2, 8, 4);
}
#[test]
fn the_residency_map_stays_bijective_under_pressure() {
let mut cache = ExpertCache::new(4, 16, 20);
let policy = QStarPolicy::from_fraction(0.5);
for step in 0..200u32 {
let layer = step % 4;
let routes: Vec<u32> = (0..6).map(|i| (step * 7 + i * 3) % 16).collect();
let plan = if step % 2 == 0 {
cache.ensure(layer, &routes)
} else {
cache.ensure_hybrid(layer, &routes, &policy)
};
assert_eq!(plan.slots.len(), routes.len());
for slot in 0..cache.cache_size() as u32 {
if let Some(id) = cache.resident_in(slot) {
assert_eq!(
cache.slot_of(id.layer, id.expert),
Some(slot),
"slot {slot} and its occupant disagree"
);
}
}
assert!(cache.resident_slots() <= cache.cache_size());
}
}
#[test]
fn per_layer_stats_tell_one_thrashing_layer_from_a_uniformly_tight_cache() {
let mut thrashing = ExpertCache::new(2, 4, 8);
thrashing.ensure(1, &[0, 1]); thrashing.reset_stats();
thrashing.ensure(0, &[0, 1]); thrashing.ensure(0, &[2, 3]); thrashing.ensure(1, &[0, 1]); thrashing.ensure(1, &[0, 1]);
let mut uniform = ExpertCache::new(2, 4, 8);
uniform.ensure(0, &[0]);
uniform.ensure(1, &[0]);
uniform.reset_stats();
uniform.ensure(0, &[0, 1]); uniform.ensure(0, &[0, 2]); uniform.ensure(1, &[0, 1]);
uniform.ensure(1, &[0, 2]);
assert_eq!(thrashing.stats().active, uniform.stats().active);
assert_eq!(thrashing.stats().missing, uniform.stats().missing);
assert_eq!(thrashing.stats().miss_rate(), 0.5);
assert_eq!(uniform.stats().miss_rate(), 0.5);
assert_eq!(thrashing.layer_stats(0).miss_rate(), 1.0);
assert_eq!(thrashing.layer_stats(1).miss_rate(), 0.0);
assert_eq!(uniform.layer_stats(0).miss_rate(), 0.5);
assert_eq!(uniform.layer_stats(1).miss_rate(), 0.5);
}
#[test]
fn per_layer_stats_sum_to_the_global_stats() {
let mut cache = ExpertCache::new(3, 8, 16);
let policy = QStarPolicy::from_fraction(0.5);
for step in 0..30u32 {
let layer = step % 3;
let routes: Vec<u32> = (0..4).map(|i| (step * 5 + i * 3) % 8).collect();
if step % 2 == 0 {
cache.ensure(layer, &routes);
} else {
cache.ensure_hybrid(layer, &routes, &policy);
}
}
let global = cache.stats();
let summed =
cache
.per_layer_stats()
.iter()
.fold(ExpertCacheStats::default(), |mut acc, layer| {
acc.calls += layer.calls;
acc.active += layer.active;
acc.missing += layer.missing;
acc.fetched += layer.fetched;
acc
});
assert_eq!(summed, global);
assert_eq!(cache.per_layer_stats().len(), 3);
assert_eq!(cache.layer_stats(1).calls, 10);
}
#[test]
fn resetting_stats_opens_a_new_per_layer_window() {
let mut cache = ExpertCache::new(2, 4, 8);
cache.ensure(0, &[0, 1]);
assert_eq!(cache.layer_stats(0).missing, 2);
cache.reset_stats();
assert_eq!(cache.layer_stats(0), ExpertCacheStats::default());
assert_eq!(cache.stats(), ExpertCacheStats::default());
cache.ensure(0, &[0, 1]);
assert_eq!(cache.layer_stats(0).missing, 0);
assert_eq!(cache.layer_stats(0).active, 2);
assert_eq!(cache.layer_stats(0).active_per_call(), 2.0);
assert_eq!(cache.layer_stats(0).missing_per_call(), 0.0);
}
#[test]
fn the_routing_histogram_counts_every_route_not_every_distinct_expert() {
let mut cache = ExpertCache::new(2, 4, 8);
cache.set_collect_routing(true);
cache.ensure(0, &[3, 3, 3, 1]);
assert_eq!(cache.routing_histogram(0), &[0, 1, 0, 3]);
assert_eq!(cache.routing_histogram(1), &[0, 0, 0, 0]);
}
#[test]
fn oracle_hit_at_slots_separates_a_skewed_layer_from_a_flat_one() {
let mut cache = ExpertCache::new(2, 8, 8);
cache.set_collect_routing(true);
let routes = |counts: [usize; 8]| -> Vec<u32> {
let mut out = Vec::new();
for (expert, count) in counts.iter().enumerate() {
for _ in 0..*count {
out.push(expert as u32);
}
}
out
};
cache.ensure(0, &routes([24, 24, 24, 24, 1, 1, 1, 1]));
cache.ensure(1, &routes([13, 13, 13, 13, 12, 12, 12, 12]));
let report = cache.routing_skew().expect("routing was observed");
assert_eq!(report.slots_per_layer, 4.0);
assert_eq!(report.oracle_slots, 4);
assert_eq!(report.per_layer.len(), 2);
let skewed = report.per_layer[0];
let flat = report.per_layer[1];
assert_eq!(skewed.routed, 100);
assert_eq!(flat.routed, 100);
assert_eq!(
skewed.working_set, flat.working_set,
"both layers touch every expert, so the working set cannot tell them apart"
);
assert!((skewed.oracle_hit_at_slots - 0.96).abs() < 1e-12);
assert!((flat.oracle_hit_at_slots - 0.52).abs() < 1e-12);
assert!(skewed.norm_entropy < flat.norm_entropy);
assert!(skewed.experts_for_90pct < flat.experts_for_90pct);
}
#[test]
fn routing_skew_reports_the_working_set_coverage_and_entropy_per_layer() {
let mut cache = ExpertCache::new(2, 8, 8);
cache.set_collect_routing(true);
cache.ensure(0, &[0, 1, 2, 3, 4, 5, 6, 7]); cache.ensure(1, &[2, 2, 2, 2]);
let report = cache.routing_skew().expect("routing was observed");
let flat = report.per_layer[0];
assert_eq!(flat.working_set, 8);
assert_eq!(flat.experts_for_90pct, 8, "7 experts reach 0.875, not 0.9");
assert!((flat.norm_entropy - 1.0).abs() < 1e-12);
assert!((flat.oracle_hit_at_slots - 0.5).abs() < 1e-12);
let single = report.per_layer[1];
assert_eq!(single.working_set, 1);
assert_eq!(single.experts_for_90pct, 1);
assert_eq!(single.norm_entropy, 0.0);
assert_eq!(single.oracle_hit_at_slots, 1.0);
assert_eq!(report.working_set_max, 8);
assert_eq!(report.working_set_mean, 4.5);
assert_eq!(report.experts_for_90pct, 4.5);
assert!((report.oracle_hit_at_slots - 0.75).abs() < 1e-12);
assert!((report.norm_entropy - 0.5).abs() < 1e-12);
}
#[test]
fn routing_skew_is_absent_until_routing_is_observed() {
let mut cache = ExpertCache::new(4, 8, 16);
cache.ensure(0, &[1, 2]);
assert!(
cache.routing_skew().is_none(),
"collection is opt-in, so nothing was recorded"
);
cache.set_collect_routing(true);
cache.ensure(2, &[1, 2]);
let report = cache.routing_skew().expect("layer 2 was observed");
assert_eq!(report.per_layer.len(), 1);
assert_eq!(report.per_layer[0].layer, 2);
cache.reset_routing();
assert!(cache.routing_skew().is_none());
}
const BIG_BANK: u64 = 512 * 1024;
const SMALL_BANK: u64 = 4 * 1024;
fn cache_with_one_layer_per_quarter() -> ExpertCache {
let mut cache = ExpertCache::new(4, 4, 16);
for layer in 0..4u32 {
cache.ensure(layer, &[0, 1, 2, 3]);
}
for layer in 0..4u32 {
for expert in 0..4u32 {
assert_eq!(cache.slot_of(layer, expert), Some(layer * 4 + expert));
}
}
cache
}
#[test]
fn a_slot_inside_the_prefill_buffers_is_a_miss_however_resident_it_looks() {
let mut cache = cache_with_one_layer_per_quarter();
cache.begin_prefill();
assert_eq!(cache.slot_of(1, 0), Some(4));
let plan = cache.prefetch_prefill_layer(1, &[BIG_BANK]);
assert_eq!(plan.buffer_id, 1);
assert!(
plan.gather.is_empty(),
"slots below 2 * num_experts are volatile, not resident"
);
assert_eq!(plan.miss_runs, vec![MissRun { start: 0, len: 4 }]);
assert_eq!(cache.prefill_hit_rows(), 0);
assert_eq!(cache.prefill_rows(), 4);
}
#[test]
fn a_resident_expert_above_the_buffer_slots_is_gathered_device_side() {
let mut cache = cache_with_one_layer_per_quarter();
cache.begin_prefill();
let plan = cache.prefetch_prefill_layer(2, &[BIG_BANK]);
assert_eq!(plan.buffer_id, 0);
assert_eq!(plan.gather.dst_slots, vec![0, 1, 2, 3]);
assert_eq!(plan.gather.src_slots, vec![8, 9, 10, 11]);
assert!(plan.miss_runs.is_empty());
assert!(
plan.entries.is_empty(),
"a big bank with no misses ships nothing"
);
assert_eq!(cache.prefill_hit_rows(), 4);
assert_eq!(cache.prefill_rows(), 4);
}
#[test]
fn invalidating_a_prefill_buffer_makes_its_slots_the_first_victims() {
let mut cache = cache_with_one_layer_per_quarter();
cache.ensure(0, &[0, 1, 2, 3]);
cache.ensure(1, &[0, 1, 2, 3]);
cache.begin_prefill();
cache.prefetch_prefill_layer(0, &[BIG_BANK]); for expert in 0..4u32 {
assert_eq!(
cache.slot_of(0, expert),
None,
"the buffer's old occupants lost their residency"
);
}
let plan = cache.ensure(0, &[0]);
assert_eq!(plan.missing, 1);
assert_eq!(plan.copy.dst_slots, vec![0]);
assert_eq!(
cache.slot_of(2, 0),
Some(8),
"no real resident was evicted while free slots existed"
);
}
#[test]
fn prefill_misses_ship_as_contiguous_expert_runs() {
let mut cache = ExpertCache::new(4, 8, 32);
cache.ensure(0, &[0, 1, 2, 3, 4, 5, 6, 7]); cache.ensure(2, &[0, 1, 2, 3, 4, 5, 6, 7]); cache.ensure(1, &[3, 4, 7]); cache.begin_prefill();
let plan = cache.prefetch_prefill_layer(1, &[BIG_BANK]);
assert_eq!(plan.buffer_id, 1);
assert_eq!(plan.gather.dst_slots, vec![11, 12, 15]);
assert_eq!(plan.gather.src_slots, vec![16, 17, 18]);
assert_eq!(
plan.miss_runs,
vec![MissRun { start: 0, len: 3 }, MissRun { start: 5, len: 2 }]
);
assert_eq!(
plan.entries,
vec![
BankEntry {
bank: 0,
dst_slot: 8,
src_row: 0,
rows: 3,
bytes: 3 * BIG_BANK,
whole_layer: false,
},
BankEntry {
bank: 0,
dst_slot: 13,
src_row: 5,
rows: 2,
bytes: 2 * BIG_BANK,
whole_layer: false,
},
]
);
}
#[test]
fn a_small_bank_ships_its_whole_layer_even_with_no_misses() {
let mut cache = ExpertCache::new(4, 8, 32);
cache.ensure(0, &[0, 1, 2, 3, 4, 5, 6, 7]); cache.ensure(2, &[0, 1, 2, 3, 4, 5, 6, 7]); cache.ensure(1, &[0, 1, 2, 3, 4, 5, 6, 7]); cache.begin_prefill();
let plan = cache.prefetch_prefill_layer(1, &[BIG_BANK, SMALL_BANK]);
assert_eq!(plan.gather.len(), 8, "every row is resident");
assert!(plan.miss_runs.is_empty());
assert_eq!(plan.gather_banks, vec![0], "the small bank is not gathered");
assert_eq!(
plan.entries,
vec![BankEntry {
bank: 1,
dst_slot: 8,
src_row: 0,
rows: 8,
bytes: 8 * SMALL_BANK,
whole_layer: true,
}]
);
}
#[test]
fn the_two_prefill_buffers_rotate_by_layer_parity() {
let mut cache = ExpertCache::new(4, 4, 16);
cache.begin_prefill();
assert_eq!(cache.prefill_buffer_layer(0), None);
assert_eq!(cache.prefetch_prefill_layer(0, &[BIG_BANK]).buffer_id, 0);
assert_eq!(cache.prefetch_prefill_layer(1, &[BIG_BANK]).buffer_id, 1);
assert_eq!(cache.prefill_buffer_layer(0), Some(0));
assert_eq!(cache.prefill_buffer_layer(1), Some(1));
let again = cache.prefetch_prefill_layer(1, &[BIG_BANK]);
assert!(again.already_loaded);
assert!(again.entries.is_empty());
assert_eq!(cache.prefill_rows(), 8, "the no-op staged no rows");
cache.release_prefill_layer(0);
assert_eq!(cache.prefetch_prefill_layer(2, &[BIG_BANK]).buffer_id, 0);
assert_eq!(cache.prefill_buffer_layer(0), Some(2));
}
#[test]
#[should_panic(expected = "still holds layer 0")]
fn reusing_a_prefill_buffer_before_it_is_released_is_refused() {
let mut cache = ExpertCache::new(4, 4, 16);
cache.begin_prefill();
cache.prefetch_prefill_layer(0, &[BIG_BANK]);
cache.prefetch_prefill_layer(2, &[BIG_BANK]);
}
#[test]
#[should_panic(expected = "begin_prefill")]
fn staging_a_layer_without_opening_the_chunk_is_refused() {
let mut cache = ExpertCache::new(4, 4, 16);
cache.prefetch_prefill_layer(0, &[BIG_BANK]);
}
#[test]
fn the_chunk_snapshot_and_the_live_map_agree_on_every_row() {
let mut cache = ExpertCache::new(4, 4, 16);
for layer in 0..4u32 {
cache.ensure(layer, &[0, 1, 2, 3]);
}
cache.begin_prefill();
for layer in [2u32, 3, 2] {
let live: Vec<Option<u32>> = (0..4)
.map(|expert| cache.slot_of(layer, expert))
.map(|slot| slot.filter(|s| *s >= 8))
.collect();
let plan = cache.prefetch_prefill_layer(layer, &[BIG_BANK]);
cache.release_prefill_layer(layer);
let hits: Vec<u32> = plan.gather.src_slots.clone();
let live_hits: Vec<u32> = live.into_iter().flatten().collect();
if !plan.already_loaded {
assert_eq!(hits, live_hits, "layer {layer} classified differently");
}
}
}
#[test]
fn a_pool_that_cannot_lend_the_buffers_their_slots_says_so() {
let cache = ExpertCache::new(2, 8, 12);
assert!(!cache.prefill_overlap_fits());
assert_eq!(cache.prefill_buffer_slots(), 16);
let fits = ExpertCache::new(2, 8, 16);
assert!(fits.prefill_overlap_fits());
}
#[test]
#[should_panic(expected = "cannot lend")]
fn opening_a_chunk_on_a_pool_too_small_for_the_buffers_is_refused() {
let mut cache = ExpertCache::new(2, 8, 12);
cache.begin_prefill();
}
}