1use crate::Result;
12use std::collections::HashMap;
13use std::fs::File;
14use std::path::PathBuf;
15use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
16use std::sync::mpsc::SyncSender;
17use std::sync::{Arc, Mutex, OnceLock, Weak};
18
19use super::GgmlDType;
20
21static ENABLED: AtomicBool = AtomicBool::new(false);
23
24pub fn set_enabled(on: bool) {
25 ENABLED.store(on, Ordering::Relaxed);
26}
27
28pub fn enabled() -> bool {
29 ENABLED.load(Ordering::Relaxed)
30}
31
32const PAGE_CACHE_RESERVE: u64 = 2_500_000_000;
34const ACTIVATION_RESERVE: u64 = 1_200_000_000;
35const BUDGET_FRACTION: f64 = 0.88;
37const PIN_FRACTION: f64 = 0.25;
39const REPIN_MAX_SWAPS: usize = 4;
42
43pub struct ExpertStreamBank {
45 name: String,
47 file: File,
48 base_offset: u64,
50 expert_bytes: usize,
52 n_experts: usize,
53 dtype: GgmlDType,
54 n: usize,
55 k: usize,
56 inner: Mutex<BankState>,
57}
58
59struct BankState {
60 cap: usize,
62 clock: u64,
63 lru: HashMap<u32, LruEntry>,
64 pinned: HashMap<u32, Arc<[u8]>>,
65 usage: Vec<u64>,
68 heat: Vec<u32>,
72 hits: u64,
73 misses: u64,
74}
75
76struct LruEntry {
77 bytes: Arc<[u8]>,
78 used: u64,
79}
80
81impl ExpertStreamBank {
82 #[allow(clippy::too_many_arguments)]
83 pub fn open(
84 name: String,
85 path: &std::path::Path,
86 base_offset: u64,
87 expert_bytes: usize,
88 n_experts: usize,
89 dtype: GgmlDType,
90 n: usize,
91 k: usize,
92 ) -> Result<Arc<Self>> {
93 let file = File::open(path)?;
94 let bank = Arc::new(Self {
95 name,
96 file,
97 base_offset,
98 expert_bytes,
99 n_experts,
100 dtype,
101 n,
102 k,
103 inner: Mutex::new(BankState {
105 cap: 1,
106 clock: 0,
107 lru: HashMap::new(),
108 pinned: HashMap::new(),
109 usage: vec![0; n_experts],
110 heat: vec![0; n_experts],
111 hits: 0,
112 misses: 0,
113 }),
114 });
115 registry().lock().expect("registry lock").register(&bank);
116 Ok(bank)
117 }
118
119 pub fn dtype(&self) -> GgmlDType {
120 self.dtype
121 }
122
123 pub fn n_experts(&self) -> usize {
124 self.n_experts
125 }
126
127 pub fn expert_dims(&self) -> (usize, usize) {
128 (self.n, self.k)
129 }
130
131 pub fn expert_bytes(&self) -> usize {
132 self.expert_bytes
133 }
134
135 pub fn logical_bytes(&self) -> usize {
137 self.expert_bytes * self.n_experts
138 }
139
140 pub fn fetch(&self, eid: u32) -> Result<Arc<[u8]>> {
142 let mut guard = self.inner.lock().expect("expert bank lock poisoned");
143 let st = &mut *guard;
145 if (eid as usize) < st.usage.len() {
146 st.usage[eid as usize] = st.usage[eid as usize].saturating_add(1);
147 st.heat[eid as usize] = st.heat[eid as usize].saturating_add(1);
148 }
149 if let Some(bytes) = st.pinned.get(&eid) {
150 st.hits += 1;
151 return Ok(bytes.clone());
152 }
153 if let Some(entry) = st.lru.get_mut(&eid) {
154 st.hits += 1;
155 st.clock += 1;
156 entry.used = st.clock;
157 return Ok(entry.bytes.clone());
158 }
159 st.misses += 1;
160 let bytes = self.read_expert(eid)?;
161 self.admit(st, eid, bytes.clone());
162 Ok(bytes)
163 }
164
165 fn admit(&self, st: &mut BankState, eid: u32, bytes: Arc<[u8]>) {
166 if st.lru.len() >= st.cap {
167 if let Some(&victim) = st.lru.iter().min_by_key(|(_, e)| e.used).map(|(k, _)| k) {
168 st.lru.remove(&victim);
169 }
170 }
171 st.clock += 1;
172 let used = st.clock;
173 st.lru.insert(eid, LruEntry { bytes, used });
174 }
175
176 fn read_expert(&self, eid: u32) -> Result<Arc<[u8]>> {
177 let off = self.base_offset + eid as u64 * self.expert_bytes as u64;
178 let mut buf = vec![0u8; self.expert_bytes];
179 read_exact_at(&self.file, &mut buf, off)?;
180 fadvise_dontneed(&self.file, off, self.expert_bytes as u64);
181 Ok(Arc::from(buf.into_boxed_slice()))
182 }
183
184 pub fn pin(&self, eid: u32) -> Result<()> {
186 let mut st = self.inner.lock().expect("expert bank lock poisoned");
187 if st.pinned.contains_key(&eid) {
188 return Ok(());
189 }
190 let bytes = if let Some(entry) = st.lru.remove(&eid) {
191 entry.bytes
192 } else {
193 drop(st);
194 let b = self.read_expert(eid)?;
195 st = self.inner.lock().expect("expert bank lock poisoned");
196 b
197 };
198 st.pinned.insert(eid, bytes);
199 Ok(())
200 }
201
202 pub fn repin(&self, max_swaps: usize) -> Result<usize> {
210 let mut guard = self.inner.lock().expect("expert bank lock poisoned");
211 let st = &mut *guard;
212 tier_decay(&mut st.heat);
213 let mut swaps = 0usize;
214 while swaps < max_swaps {
215 let pinned: Vec<u32> = st.pinned.keys().copied().collect();
216 let Some((slot, hot, _gain)) = tier_pick_swap(&st.heat, &pinned) else {
217 break;
218 };
219 let cold = pinned[slot];
220 if let Some(bytes) = st.pinned.remove(&cold) {
222 self.admit(st, cold, bytes);
223 }
224 let bytes = match st.lru.remove(&hot) {
226 Some(entry) => entry.bytes,
227 None => self.read_expert(hot)?,
228 };
229 st.pinned.insert(hot, bytes);
230 swaps += 1;
231 }
232 Ok(swaps)
233 }
234
235 fn prefetch_resident(&self, eid: u32) -> Result<()> {
241 {
242 let st = self.inner.lock().expect("expert bank lock poisoned");
243 if st.pinned.contains_key(&eid) || st.lru.contains_key(&eid) {
244 return Ok(());
245 }
246 }
247 let bytes = self.read_expert(eid)?;
249 let mut guard = self.inner.lock().expect("expert bank lock poisoned");
250 let st = &mut *guard;
251 if st.pinned.contains_key(&eid) || st.lru.contains_key(&eid) {
253 return Ok(());
254 }
255 self.admit(st, eid, bytes);
256 Ok(())
257 }
258
259 pub fn prefetch(self: &Arc<Self>, eid: u32) {
263 if !prefetch_enabled() || eid as usize >= self.n_experts {
264 return;
265 }
266 let _ = prefetcher().try_send(PrefetchJob {
267 bank: Arc::downgrade(self),
268 eid,
269 });
270 }
271
272 pub fn set_cap(&self, cap: usize) {
274 let mut st = self.inner.lock().expect("expert bank lock poisoned");
275 st.cap = cap.max(1);
276 }
277
278 pub fn stats(&self) -> (usize, usize, usize, u64, u64) {
280 let st = self.inner.lock().expect("expert bank lock poisoned");
281 (st.pinned.len(), st.lru.len(), st.cap, st.hits, st.misses)
282 }
283
284 pub fn resident_bytes(&self) -> usize {
285 let st = self.inner.lock().expect("expert bank lock poisoned");
286 (st.pinned.len() + st.lru.len()) * self.expert_bytes
287 }
288
289 fn usage_snapshot(&self) -> Vec<(u32, u64)> {
290 let st = self.inner.lock().expect("expert bank lock poisoned");
291 st.usage
292 .iter()
293 .enumerate()
294 .filter(|(_, &c)| c > 0)
295 .map(|(e, &c)| (e as u32, c))
296 .collect()
297 }
298}
299
300struct Registry {
302 banks: Vec<Weak<ExpertStreamBank>>,
303 sidecar: Option<PathBuf>,
304}
305
306fn registry() -> &'static Mutex<Registry> {
307 static REG: OnceLock<Mutex<Registry>> = OnceLock::new();
308 REG.get_or_init(|| {
309 Mutex::new(Registry {
310 banks: Vec::new(),
311 sidecar: None,
312 })
313 })
314}
315
316impl Registry {
317 fn register(&mut self, bank: &Arc<ExpertStreamBank>) {
318 self.banks.push(Arc::downgrade(bank));
319 }
320
321 fn live(&self) -> Vec<Arc<ExpertStreamBank>> {
322 self.banks.iter().filter_map(Weak::upgrade).collect()
323 }
324}
325
326pub fn set_usage_sidecar(path: PathBuf) {
327 registry().lock().expect("registry lock").sidecar = Some(path);
328}
329
330pub fn total_resident_bytes() -> usize {
331 registry()
332 .lock()
333 .expect("registry lock")
334 .live()
335 .iter()
336 .map(|b| b.resident_bytes())
337 .sum()
338}
339
340fn tier_pick_swap(heat: &[u32], pinned: &[u32]) -> Option<(usize, u32, i64)> {
354 if heat.is_empty() || pinned.is_empty() {
355 return None;
356 }
357 let cold = pinned
359 .iter()
360 .enumerate()
361 .min_by_key(|&(_, &p)| heat[p as usize])
362 .map(|(z, _)| z)
363 .expect("pinned is non-empty");
364 let mut hot: Option<usize> = None;
366 let mut hot_heat = 0u32;
367 for (e, &h) in heat.iter().enumerate() {
368 let resident = pinned.iter().any(|&p| p as usize == e);
369 if !resident && h > hot_heat {
370 hot_heat = h;
371 hot = Some(e);
372 }
373 }
374 let hot = hot?;
375 let cold_heat = heat[pinned[cold] as usize];
376 if hot_heat <= cold_heat + (cold_heat >> 2) + 4 {
377 return None;
378 }
379 Some((cold, hot as u32, hot_heat as i64 - cold_heat as i64))
380}
381
382fn tier_decay(heat: &mut [u32]) {
385 for h in heat.iter_mut() {
386 *h >>= 1;
387 }
388}
389
390pub fn repin_interval() -> usize {
393 static N: OnceLock<usize> = OnceLock::new();
394 *N.get_or_init(|| {
395 std::env::var("STREAM_EXPERTS_REPIN")
396 .ok()
397 .and_then(|v| v.parse::<usize>().ok())
398 .unwrap_or(0)
399 })
400}
401
402pub fn repin_all() -> usize {
406 if repin_interval() == 0 {
407 return 0;
408 }
409 let banks = registry().lock().expect("registry lock").live();
414 banks
415 .iter()
416 .map(|b| b.repin(REPIN_MAX_SWAPS).unwrap_or(0))
417 .sum()
418}
419
420pub fn prefetch_enabled() -> bool {
427 static ON: OnceLock<bool> = OnceLock::new();
428 *ON.get_or_init(|| {
429 std::env::var("STREAM_EXPERTS_PREFETCH")
430 .map(|v| !v.is_empty() && v != "0")
431 .unwrap_or(false)
432 })
433}
434
435struct PrefetchJob {
436 bank: Weak<ExpertStreamBank>,
437 eid: u32,
438}
439
440fn prefetcher() -> &'static SyncSender<PrefetchJob> {
445 static P: OnceLock<SyncSender<PrefetchJob>> = OnceLock::new();
446 P.get_or_init(|| {
447 let (tx, rx) = std::sync::mpsc::sync_channel::<PrefetchJob>(256);
448 std::thread::Builder::new()
449 .name("stream-experts-prefetch".into())
450 .spawn(move || {
451 while let Ok(job) = rx.recv() {
452 if let Some(bank) = job.bank.upgrade() {
453 let _ = bank.prefetch_resident(job.eid);
454 }
455 }
456 })
457 .expect("spawn stream-experts prefetch worker");
458 tx
459 })
460}
461
462pub fn finalize() {
465 let banks = registry().lock().expect("registry lock").live();
466 if banks.is_empty() {
467 return;
468 }
469 let expert_bytes = banks
470 .iter()
471 .map(|b| b.expert_bytes)
472 .max()
473 .unwrap_or(1)
474 .max(1);
475 let n_banks = banks.len() as u64;
476
477 let forced_gb = std::env::var("STREAM_EXPERTS_RAM_GB")
479 .ok()
480 .and_then(|v| v.parse::<f64>().ok())
481 .filter(|g| *g > 0.0);
482 let avail = match forced_gb {
483 Some(gb) => (gb * 1e9) as u64,
484 None => mem_available_bytes(),
485 };
486 let budget = (avail as f64 * BUDGET_FRACTION) as u64;
487 let slack = PAGE_CACHE_RESERVE + ACTIVATION_RESERVE;
488 let for_cache = budget.saturating_sub(slack);
489 let mut cap = (for_cache / (n_banks * expert_bytes as u64)) as usize;
490
491 let max_experts = banks.iter().map(|b| b.n_experts).max().unwrap_or(1);
492 cap = cap.clamp(1, max_experts);
493 for b in &banks {
494 b.set_cap(cap);
495 }
496
497 let pinned = load_and_pin(&banks, cap);
498 register_atexit_save();
499 spawn_usage_saver();
500
501 eprintln!(
502 "[stream-experts] {} banks x {:.1} MB/expert; budget {:.1} GB -> cap {}/bank \
503 (cache {:.1} GB), pinned {}",
504 banks.len(),
505 expert_bytes as f64 / 1e6,
506 avail as f64 / 1e9,
507 cap,
508 (cap as u64 * n_banks * expert_bytes as u64) as f64 / 1e9,
509 pinned,
510 );
511}
512
513fn load_and_pin(banks: &[Arc<ExpertStreamBank>], cap: usize) -> usize {
515 let sidecar = match registry().lock().expect("registry lock").sidecar.clone() {
516 Some(p) if p.exists() => p,
517 _ => return 0,
518 };
519 let text = match std::fs::read_to_string(&sidecar) {
520 Ok(t) => t,
521 Err(_) => return 0,
522 };
523 let mut by_name: HashMap<&str, Vec<(u32, u64)>> = HashMap::new();
524 for line in text.lines() {
525 let mut it = line.split_whitespace();
526 let (Some(name), Some(eid), Some(cnt)) = (it.next(), it.next(), it.next()) else {
527 continue;
528 };
529 if let (Ok(eid), Ok(cnt)) = (eid.parse::<u32>(), cnt.parse::<u64>()) {
530 by_name.entry(name).or_default().push((eid, cnt));
531 }
532 }
533 let pin_budget = ((cap as f64 * PIN_FRACTION) as usize).max(1);
534 let mut pinned = 0usize;
535 for b in banks {
536 let Some(rows) = by_name.get(b.name.as_str()) else {
537 continue;
538 };
539 let mut rows = rows.clone();
540 rows.sort_by_key(|r| std::cmp::Reverse(r.1));
543 for (eid, _) in rows.into_iter().take(pin_budget) {
544 if eid as usize >= b.n_experts {
545 continue;
546 }
547 if b.pin(eid).is_ok() {
548 pinned += 1;
549 }
550 }
551 }
552 pinned
553}
554
555pub fn save_usage() -> Result<()> {
557 let (banks, sidecar) = {
558 let reg = registry().lock().expect("registry lock");
559 (reg.live(), reg.sidecar.clone())
560 };
561 let Some(path) = sidecar else {
562 return Ok(());
563 };
564 let mut merged: HashMap<(String, u32), u64> = HashMap::new();
565 if let Ok(text) = std::fs::read_to_string(&path) {
566 for line in text.lines() {
567 let mut it = line.split_whitespace();
568 if let (Some(name), Some(eid), Some(cnt)) = (it.next(), it.next(), it.next()) {
569 if let (Ok(eid), Ok(cnt)) = (eid.parse::<u32>(), cnt.parse::<u64>()) {
570 *merged.entry((name.to_string(), eid)).or_default() += cnt;
571 }
572 }
573 }
574 }
575 for b in &banks {
576 for (eid, cnt) in b.usage_snapshot() {
577 *merged.entry((b.name.clone(), eid)).or_default() += cnt;
578 }
579 }
580 let mut out = String::new();
581 for ((name, eid), cnt) in merged {
582 out.push_str(&format!("{name} {eid} {cnt}\n"));
583 }
584 std::fs::write(&path, out)?;
585 Ok(())
586}
587
588fn register_atexit_save() {
591 #[cfg(unix)]
592 {
593 static ONCE: std::sync::Once = std::sync::Once::new();
594 ONCE.call_once(|| {
595 extern "C" fn on_exit() {
596 let _ = save_usage();
597 }
598 unsafe {
599 libc::atexit(on_exit);
600 }
601 });
602 }
603}
604
605const USAGE_SAVE_INTERVAL: std::time::Duration = std::time::Duration::from_secs(60);
607
608fn spawn_usage_saver() {
622 static ONCE: std::sync::Once = std::sync::Once::new();
623 ONCE.call_once(|| {
624 let _ = std::thread::Builder::new()
625 .name("expert-usage-saver".into())
626 .spawn(|| loop {
627 std::thread::sleep(USAGE_SAVE_INTERVAL);
628 let _ = save_usage();
631 });
632 });
633}
634
635#[cfg(unix)]
636fn read_exact_at(file: &File, buf: &mut [u8], offset: u64) -> Result<()> {
637 use std::os::unix::fs::FileExt;
638 file.read_exact_at(buf, offset)?;
639 Ok(())
640}
641
642#[cfg(not(unix))]
643fn read_exact_at(_file: &File, _buf: &mut [u8], _offset: u64) -> Result<()> {
644 crate::bail!("streaming experts require a unix positional-read (pread) platform")
645}
646
647#[cfg(target_os = "linux")]
651fn fadvise_dontneed(file: &File, offset: u64, len: u64) {
652 use std::os::unix::io::AsRawFd;
653 unsafe {
654 libc::posix_fadvise(
655 file.as_raw_fd(),
656 offset as libc::off_t,
657 len as libc::off_t,
658 libc::POSIX_FADV_DONTNEED,
659 );
660 }
661}
662
663#[cfg(not(target_os = "linux"))]
664fn fadvise_dontneed(_file: &File, _offset: u64, _len: u64) {}
665
666fn mem_available_bytes() -> u64 {
668 #[cfg(target_os = "linux")]
669 {
670 if let Ok(text) = std::fs::read_to_string("/proc/meminfo") {
671 for line in text.lines() {
672 if let Some(rest) = line.strip_prefix("MemAvailable:") {
673 if let Some(kb) = rest
674 .split_whitespace()
675 .next()
676 .and_then(|v| v.parse::<u64>().ok())
677 {
678 return kb.saturating_mul(1024);
679 }
680 }
681 }
682 }
683 }
684 static WARNED: AtomicU64 = AtomicU64::new(0);
685 if WARNED.swap(1, Ordering::Relaxed) == 0 {
686 eprintln!("[stream-experts] MemAvailable unreadable; assuming 8 GB free");
687 }
688 8_000_000_000
689}
690
691#[cfg(test)]
692mod tests {
693 use super::{tier_decay, tier_pick_swap, REPIN_MAX_SWAPS};
694
695 fn simulate(heat: &[u32], pinned: &mut [u32], max_swaps: usize) -> usize {
703 let mut swaps = 0;
704 while swaps < max_swaps {
705 match tier_pick_swap(heat, pinned) {
706 Some((slot, hot, _gain)) => {
707 pinned[slot] = hot;
708 swaps += 1;
709 }
710 None => break,
711 }
712 }
713 swaps
714 }
715
716 #[test]
717 fn swaps_coldest_pinned_for_hottest_streamed() {
718 let heat = [100, 5, 200, 50];
720 let pinned = [0, 1];
721 let (slot, hot, gain) = tier_pick_swap(&heat, &pinned).expect("beneficial swap");
722 assert_eq!(slot, 1, "coldest pinned is at slot 1 (eid 1)");
723 assert_eq!(hot, 2, "hottest streamed is eid 2");
724 assert_eq!(gain, 195, "gain is hot_heat - cold_heat = 200 - 5");
725 }
726
727 #[test]
728 fn hysteresis_blocks_near_equal_experts() {
729 let heat = [100, 100, 110];
732 let pinned = [0, 1];
733 assert!(tier_pick_swap(&heat, &pinned).is_none());
734 }
735
736 #[test]
737 fn hysteresis_margin_is_exclusive_at_the_boundary() {
738 assert!(
740 tier_pick_swap(&[100, 100, 129], &[0, 1]).is_none(),
741 "equal to the threshold must not swap"
742 );
743 let (slot, hot, gain) =
744 tier_pick_swap(&[100, 100, 130], &[0, 1]).expect("one over the threshold swaps");
745 assert_eq!((slot, hot, gain), (0, 2, 30));
746 }
747
748 #[test]
749 fn no_swap_when_every_expert_is_already_pinned() {
750 assert!(tier_pick_swap(&[10, 20], &[0, 1]).is_none());
751 }
752
753 #[test]
754 fn empty_inputs_are_safe() {
755 assert!(tier_pick_swap(&[], &[]).is_none());
756 assert!(tier_pick_swap(&[1, 2, 3], &[]).is_none());
757 assert!(tier_pick_swap(&[], &[0, 1]).is_none());
758 }
759
760 #[test]
761 fn picks_the_coldest_slot_and_hottest_candidate_among_many() {
762 let heat = [50, 10, 30, 200, 80, 5];
764 let pinned = [0, 2, 4];
765 let (slot, hot, gain) = tier_pick_swap(&heat, &pinned).expect("swap");
766 assert_eq!(slot, 1, "eid 2 (heat 30) is the coldest pinned, at slot 1");
767 assert_eq!(hot, 3, "eid 3 (heat 200) is the hottest streamed");
768 assert_eq!(gain, 170);
769 }
770
771 #[test]
772 fn decay_halves_every_expert() {
773 let mut heat = [10, 3, 0, 255, 1];
774 tier_decay(&mut heat);
775 assert_eq!(heat, [5, 1, 0, 127, 0]);
776 }
777
778 #[test]
779 fn repeated_swaps_converge_then_stop_under_hysteresis() {
780 let heat = [100, 90, 5, 4];
783 let mut pinned = vec![2, 3];
784 let swaps = simulate(&heat, &mut pinned, REPIN_MAX_SWAPS);
785 assert_eq!(swaps, 2, "exactly the two hot experts get pinned");
786 pinned.sort_unstable();
787 assert_eq!(
788 pinned,
789 vec![0, 1],
790 "hot-set converged to the two hottest experts"
791 );
792 assert_eq!(simulate(&heat, &mut pinned, REPIN_MAX_SWAPS), 0);
794 }
795
796 #[test]
797 fn one_dominant_expert_settles_after_a_single_swap() {
798 let heat = [50, 10, 30, 200, 80, 5];
801 let mut pinned = vec![0, 2, 4];
802 assert_eq!(simulate(&heat, &mut pinned, REPIN_MAX_SWAPS), 1);
803 pinned.sort_unstable();
804 assert_eq!(pinned, vec![0, 3, 4]);
805 }
806}