Skip to main content

memra_engine/
moesd.rs

1//! Diagnostic-only expert-union capture for the MoESD target-efficiency harness.
2//!
3//! The timed forward never enables this collector. The harness rolls its caches back and
4//! replays the same target step with capture enabled, so route D2H cannot contaminate T_T.
5
6use crate::Engine;
7use cudarc::driver::CudaSlice;
8use std::collections::{BTreeMap, HashSet};
9use std::sync::Mutex;
10use std::sync::atomic::{AtomicBool, Ordering};
11
12#[derive(Clone, Debug)]
13pub struct MoesdLayerUnion {
14    pub id: u16,
15    pub union_size: usize,
16    pub n_expert: usize,
17    pub n_used: usize,
18    pub assignments: usize,
19}
20
21#[derive(Default)]
22struct LayerCapture {
23    experts: HashSet<u32>,
24    n_expert: usize,
25    n_used: usize,
26    assignments: usize,
27}
28
29static ACTIVE: AtomicBool = AtomicBool::new(false);
30static CAPTURE: Mutex<Option<BTreeMap<u16, LayerCapture>>> = Mutex::new(None);
31
32pub fn begin_capture() -> Result<(), Box<dyn std::error::Error>> {
33    let mut capture = CAPTURE
34        .lock()
35        .map_err(|_| "MoESD capture lock is poisoned")?;
36    if ACTIVE.load(Ordering::Relaxed) || capture.is_some() {
37        return Err("MoESD capture is already active".into());
38    }
39    *capture = Some(BTreeMap::new());
40    ACTIVE.store(true, Ordering::Release);
41    Ok(())
42}
43
44pub fn finish_capture() -> Result<Vec<MoesdLayerUnion>, Box<dyn std::error::Error>> {
45    if !ACTIVE.swap(false, Ordering::AcqRel) {
46        return Err("MoESD capture is not active".into());
47    }
48    let mut capture = CAPTURE
49        .lock()
50        .map_err(|_| "MoESD capture lock is poisoned")?;
51    let layers = capture.take().ok_or("MoESD capture state is missing")?;
52    Ok(layers
53        .into_iter()
54        .map(|(id, layer)| MoesdLayerUnion {
55            id,
56            union_size: layer.experts.len(),
57            n_expert: layer.n_expert,
58            n_used: layer.n_used,
59            assignments: layer.assignments,
60        })
61        .collect())
62}
63
64pub(crate) fn record_host_routes(
65    il: u16,
66    n_expert: usize,
67    n_used: usize,
68    selected: &[u32],
69) -> Result<(), Box<dyn std::error::Error>> {
70    if !ACTIVE.load(Ordering::Acquire) {
71        return Ok(());
72    }
73    if n_used == 0 {
74        return Err(format!("MoESD layer {il} has router n_used=0").into());
75    }
76    if selected.len() % n_used != 0 {
77        return Err(format!(
78            "MoESD route shape mismatch at layer {il}: {} selections is not divisible by {n_used}",
79            selected.len(),
80        )
81        .into());
82    }
83    if let Some(expert) = selected.iter().find(|&&expert| expert as usize >= n_expert) {
84        return Err(format!(
85            "MoESD layer {il} selected out-of-range expert {expert} from bank {n_expert}",
86        )
87        .into());
88    }
89    let mut capture = CAPTURE
90        .lock()
91        .map_err(|_| "MoESD capture lock is poisoned")?;
92    let layers = capture.as_mut().ok_or("MoESD capture state is missing")?;
93    let layer = layers.entry(il).or_default();
94    if layer.assignments != 0 && (layer.n_expert != n_expert || layer.n_used != n_used) {
95        return Err(format!("MoESD route metadata changed within layer {il}").into());
96    }
97    layer.n_expert = n_expert;
98    layer.n_used = n_used;
99    layer.assignments += selected.len();
100    layer.experts.extend(selected.iter().copied());
101    Ok(())
102}
103
104pub(crate) fn record_device_routes(
105    e: &Engine,
106    il: u16,
107    n_expert: usize,
108    n_used: usize,
109    selected: &CudaSlice<i32>,
110) -> Result<(), Box<dyn std::error::Error>> {
111    if !ACTIVE.load(Ordering::Acquire) {
112        return Ok(());
113    }
114    let selected: Vec<u32> = e
115        .dtoh_i32(selected)?
116        .into_iter()
117        .map(|expert| expert as u32)
118        .collect();
119    record_host_routes(il, n_expert, n_used, &selected)
120}
121
122#[cfg(test)]
123mod tests {
124    use super::*;
125
126    #[test]
127    fn host_capture_counts_distinct_experts_and_assignments() {
128        begin_capture().unwrap();
129        record_host_routes(7, 16, 2, &[1, 2, 2, 3, 3, 4]).unwrap();
130        let layers = finish_capture().unwrap();
131        assert_eq!(layers.len(), 1);
132        assert_eq!(layers[0].id, 7);
133        assert_eq!(layers[0].union_size, 4);
134        assert_eq!(layers[0].n_expert, 16);
135        assert_eq!(layers[0].n_used, 2);
136        assert_eq!(layers[0].assignments, 6);
137    }
138}