1use 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}