Skip to main content

onnx_runtime_ep_cpu/
weight_offload.rs

1//! Huge-model weight-offload mode and lightweight process-wide observability.
2
3use std::collections::{BTreeMap, BTreeSet};
4use std::ffi::OsStr;
5use std::sync::atomic::{AtomicU64, Ordering};
6use std::sync::{Mutex, OnceLock};
7
8use onnx_runtime_ep_api::ExternalMmapRegion;
9
10pub mod placement;
11pub mod weight_handle;
12
13/// Environment switch for the route-first mmap MoE path.
14pub const WEIGHT_OFFLOAD_ENV: &str = "ONNX_GENAI_WEIGHT_OFFLOAD";
15/// Optional override for the Resource Governor's owned warm-host cache budget.
16pub const WEIGHT_OFFLOAD_HOST_BYTES_ENV: &str = "ONNX_GENAI_WEIGHT_OFFLOAD_HOST_BYTES";
17
18#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
19pub(crate) struct WeightOffloadMode {
20    pub enabled: bool,
21}
22
23impl WeightOffloadMode {
24    pub fn from_env() -> Self {
25        Self::from_value(std::env::var_os(WEIGHT_OFFLOAD_ENV).as_deref())
26    }
27
28    fn from_value(value: Option<&OsStr>) -> Self {
29        Self {
30            enabled: value.is_some_and(|value| value == "1"),
31        }
32    }
33}
34
35/// Best-effort Linux process memory/page-fault counters.
36#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
37pub struct LinuxProcessMemoryStats {
38    pub resident_rss_bytes: u64,
39    pub minor_faults: u64,
40    pub major_faults: u64,
41}
42
43/// Snapshot of route-first weight-offload activity.
44#[derive(Clone, Debug, Default, PartialEq, Eq)]
45pub struct WeightOffloadStats {
46    pub mapped_bytes: u64,
47    pub bytes_read_from_mmap: u64,
48    pub layer_executions: u64,
49    pub active_experts: u64,
50    pub unique_experts_per_batch: u64,
51    pub peak_dequantized_experts: u64,
52    pub host_cache_hits: u64,
53    pub host_cache_misses: u64,
54    pub host_cache_evictions: u64,
55    pub owned_host_cache_bytes: u64,
56    pub peak_owned_host_cache_bytes: u64,
57    pub host_cache_budget_bytes: u64,
58    pub routed_tokens: u64,
59    pub tokens_per_expert: BTreeMap<usize, u64>,
60    pub per_layer: BTreeMap<u32, WeightOffloadLayerStats>,
61    pub linux_process: Option<LinuxProcessMemoryStats>,
62}
63
64#[derive(Clone, Debug, Default, PartialEq, Eq)]
65pub struct WeightOffloadLayerStats {
66    pub executions: u64,
67    pub active_experts: u64,
68    pub unique_experts: u64,
69    pub tokens_per_expert: BTreeMap<usize, u64>,
70}
71
72#[derive(Default)]
73pub(crate) struct WeightOffloadMetrics {
74    mapped_regions: Mutex<MappedRegionState>,
75    bytes_read_from_mmap: AtomicU64,
76    layer_executions: AtomicU64,
77    active_experts: AtomicU64,
78    unique_experts_per_batch: AtomicU64,
79    current_dequantized_experts: AtomicU64,
80    peak_dequantized_experts: AtomicU64,
81    host_cache_hits: AtomicU64,
82    host_cache_misses: AtomicU64,
83    host_cache_evictions: AtomicU64,
84    owned_host_cache_bytes: AtomicU64,
85    peak_owned_host_cache_bytes: AtomicU64,
86    host_cache_budget_bytes: AtomicU64,
87    routed_tokens: AtomicU64,
88    tokens_per_expert: Mutex<BTreeMap<usize, u64>>,
89    per_layer: Mutex<BTreeMap<u32, WeightOffloadLayerStats>>,
90}
91
92#[derive(Default)]
93struct MappedRegionState {
94    regions: BTreeSet<ExternalMmapRegion>,
95    total_bytes: u64,
96}
97
98impl WeightOffloadMetrics {
99    pub fn record_mapped_regions(
100        &self,
101        regions: &[ExternalMmapRegion],
102    ) -> Result<(), &'static str> {
103        let mut state = self
104            .mapped_regions
105            .lock()
106            .expect("weight-offload mapped-region lock poisoned");
107        let mut additions = BTreeSet::new();
108        let mut total = state.total_bytes;
109        for &region in regions {
110            let end = region
111                .offset
112                .checked_add(region.len)
113                .ok_or("mapped region endpoint overflow")?;
114            if end > isize::MAX as usize {
115                return Err("mapped region endpoint exceeds isize::MAX");
116            }
117            if !state.regions.contains(&region) && additions.insert(region) {
118                let len = u64::try_from(region.len).map_err(|_| "mapped region length overflow")?;
119                total = total.checked_add(len).ok_or("mapped byte total overflow")?;
120            }
121        }
122        state.regions.extend(additions);
123        state.total_bytes = total;
124        Ok(())
125    }
126
127    pub fn record_bytes_read(&self, bytes: usize) -> Result<(), &'static str> {
128        let bytes = u64::try_from(bytes).map_err(|_| "mmap read byte count overflow")?;
129        self.bytes_read_from_mmap
130            .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |total| {
131                total.checked_add(bytes)
132            })
133            .map_err(|_| "mmap read byte total overflow")?;
134        Ok(())
135    }
136
137    pub fn record_dequantized_expert_materialized(&self) {
138        let current = self
139            .current_dequantized_experts
140            .fetch_add(1, Ordering::Relaxed)
141            + 1;
142        self.peak_dequantized_experts
143            .fetch_max(current, Ordering::Relaxed);
144    }
145
146    pub fn record_dequantized_expert_released(&self) {
147        let previous = self
148            .current_dequantized_experts
149            .fetch_sub(1, Ordering::Relaxed);
150        debug_assert!(previous > 0, "dequantized expert residency underflow");
151    }
152
153    pub fn record_host_cache_hit(&self) {
154        self.host_cache_hits.fetch_add(1, Ordering::Relaxed);
155    }
156
157    pub fn record_host_cache_miss(&self) {
158        self.host_cache_misses.fetch_add(1, Ordering::Relaxed);
159    }
160
161    pub fn record_host_cache_evictions(&self, count: usize) -> Result<(), &'static str> {
162        let count = u64::try_from(count).map_err(|_| "host-cache eviction count overflow")?;
163        self.host_cache_evictions
164            .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |total| {
165                total.checked_add(count)
166            })
167            .map_err(|_| "host-cache eviction total overflow")?;
168        Ok(())
169    }
170
171    pub fn record_host_cache_residency(
172        &self,
173        previous_owned_bytes: u64,
174        owned_bytes: usize,
175        previous_budget_bytes: u64,
176        budget_bytes: usize,
177    ) -> Result<(u64, u64), &'static str> {
178        let owned =
179            u64::try_from(owned_bytes).map_err(|_| "owned host-cache byte count overflow")?;
180        let budget =
181            u64::try_from(budget_bytes).map_err(|_| "host-cache budget byte count overflow")?;
182
183        adjust_gauge(
184            &self.host_cache_budget_bytes,
185            previous_budget_bytes,
186            budget,
187            "host-cache budget byte total overflow",
188            "host-cache budget byte total underflow",
189        )?;
190        let aggregate_owned = match adjust_gauge(
191            &self.owned_host_cache_bytes,
192            previous_owned_bytes,
193            owned,
194            "owned host-cache byte total overflow",
195            "owned host-cache byte total underflow",
196        ) {
197            Ok(aggregate) => aggregate,
198            Err(failure) => {
199                adjust_gauge(
200                    &self.host_cache_budget_bytes,
201                    budget,
202                    previous_budget_bytes,
203                    "host-cache budget rollback overflow",
204                    "host-cache budget rollback underflow",
205                )
206                .expect("host-cache budget metric rollback must reverse the applied delta");
207                return Err(failure);
208            }
209        };
210        self.peak_owned_host_cache_bytes
211            .fetch_max(aggregate_owned, Ordering::Relaxed);
212        Ok((owned, budget))
213    }
214
215    pub fn release_host_cache_residency(&self, owned_bytes: u64, budget_bytes: u64) {
216        subtract_gauge_saturating(&self.owned_host_cache_bytes, owned_bytes);
217        subtract_gauge_saturating(&self.host_cache_budget_bytes, budget_bytes);
218    }
219
220    pub fn record_routes(&self, layer: u32, token_counts: &BTreeMap<usize, usize>) {
221        let active = token_counts.values().copied().sum::<usize>();
222        self.layer_executions.fetch_add(1, Ordering::Relaxed);
223        self.active_experts
224            .fetch_add(active as u64, Ordering::Relaxed);
225        self.unique_experts_per_batch
226            .fetch_add(token_counts.len() as u64, Ordering::Relaxed);
227        self.routed_tokens
228            .fetch_add(active as u64, Ordering::Relaxed);
229        let mut totals = self
230            .tokens_per_expert
231            .lock()
232            .expect("weight-offload metrics lock poisoned");
233        for (&expert, &tokens) in token_counts {
234            let total = totals.entry(expert).or_default();
235            *total = total.saturating_add(tokens as u64);
236        }
237
238        drop(totals);
239
240        let mut layers = self
241            .per_layer
242            .lock()
243            .expect("weight-offload layer metrics lock poisoned");
244        let layer_stats = layers.entry(layer).or_default();
245        layer_stats.executions = layer_stats.executions.saturating_add(1);
246        layer_stats.active_experts = layer_stats.active_experts.saturating_add(active as u64);
247        layer_stats.unique_experts = layer_stats
248            .unique_experts
249            .saturating_add(token_counts.len() as u64);
250        for (&expert, &tokens) in token_counts {
251            let total = layer_stats.tokens_per_expert.entry(expert).or_default();
252            *total = total.saturating_add(tokens as u64);
253        }
254    }
255
256    fn snapshot(&self) -> WeightOffloadStats {
257        WeightOffloadStats {
258            mapped_bytes: self
259                .mapped_regions
260                .lock()
261                .expect("weight-offload mapped-region lock poisoned")
262                .total_bytes,
263            bytes_read_from_mmap: self.bytes_read_from_mmap.load(Ordering::Relaxed),
264            layer_executions: self.layer_executions.load(Ordering::Relaxed),
265            active_experts: self.active_experts.load(Ordering::Relaxed),
266            unique_experts_per_batch: self.unique_experts_per_batch.load(Ordering::Relaxed),
267            peak_dequantized_experts: self.peak_dequantized_experts.load(Ordering::Relaxed),
268            host_cache_hits: self.host_cache_hits.load(Ordering::Relaxed),
269            host_cache_misses: self.host_cache_misses.load(Ordering::Relaxed),
270            host_cache_evictions: self.host_cache_evictions.load(Ordering::Relaxed),
271            owned_host_cache_bytes: self.owned_host_cache_bytes.load(Ordering::Relaxed),
272            peak_owned_host_cache_bytes: self.peak_owned_host_cache_bytes.load(Ordering::Relaxed),
273            host_cache_budget_bytes: self.host_cache_budget_bytes.load(Ordering::Relaxed),
274            routed_tokens: self.routed_tokens.load(Ordering::Relaxed),
275            tokens_per_expert: self
276                .tokens_per_expert
277                .lock()
278                .expect("weight-offload metrics lock poisoned")
279                .clone(),
280            per_layer: self
281                .per_layer
282                .lock()
283                .expect("weight-offload layer metrics lock poisoned")
284                .clone(),
285            linux_process: linux_process_memory_stats(),
286        }
287    }
288
289    #[cfg(test)]
290    pub fn reset(&self) {
291        *self
292            .mapped_regions
293            .lock()
294            .expect("weight-offload mapped-region lock poisoned") = MappedRegionState::default();
295        self.bytes_read_from_mmap.store(0, Ordering::Relaxed);
296        self.peak_dequantized_experts.store(
297            self.current_dequantized_experts.load(Ordering::Relaxed),
298            Ordering::Relaxed,
299        );
300        self.host_cache_hits.store(0, Ordering::Relaxed);
301        self.host_cache_misses.store(0, Ordering::Relaxed);
302        self.host_cache_evictions.store(0, Ordering::Relaxed);
303        self.owned_host_cache_bytes.store(0, Ordering::Relaxed);
304        self.peak_owned_host_cache_bytes.store(0, Ordering::Relaxed);
305        self.host_cache_budget_bytes.store(0, Ordering::Relaxed);
306        self.layer_executions.store(0, Ordering::Relaxed);
307        self.active_experts.store(0, Ordering::Relaxed);
308        self.unique_experts_per_batch.store(0, Ordering::Relaxed);
309        self.routed_tokens.store(0, Ordering::Relaxed);
310        self.tokens_per_expert
311            .lock()
312            .expect("weight-offload metrics lock poisoned")
313            .clear();
314        self.per_layer
315            .lock()
316            .expect("weight-offload layer metrics lock poisoned")
317            .clear();
318    }
319}
320
321fn adjust_gauge(
322    gauge: &AtomicU64,
323    previous: u64,
324    current: u64,
325    overflow: &'static str,
326    underflow: &'static str,
327) -> Result<u64, &'static str> {
328    if current >= previous {
329        let delta = current - previous;
330        let old = gauge
331            .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |total| {
332                total.checked_add(delta)
333            })
334            .map_err(|_| overflow)?;
335        old.checked_add(delta).ok_or(overflow)
336    } else {
337        let delta = previous - current;
338        let old = gauge
339            .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |total| {
340                total.checked_sub(delta)
341            })
342            .map_err(|_| underflow)?;
343        old.checked_sub(delta).ok_or(underflow)
344    }
345}
346
347fn subtract_gauge_saturating(gauge: &AtomicU64, delta: u64) {
348    gauge
349        .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |total| {
350            Some(total.saturating_sub(delta))
351        })
352        .expect("saturating gauge subtraction always succeeds");
353}
354
355static METRICS: OnceLock<WeightOffloadMetrics> = OnceLock::new();
356
357pub(crate) fn metrics() -> &'static WeightOffloadMetrics {
358    METRICS.get_or_init(WeightOffloadMetrics::default)
359}
360
361/// Read current offload counters. The scheduler governor can poll this without
362/// depending on kernel internals.
363pub fn weight_offload_stats() -> WeightOffloadStats {
364    metrics().snapshot()
365}
366
367/// Set the default CPU provider's owned warm-host cache sub-budget.
368///
369/// `ONNX_GENAI_WEIGHT_OFFLOAD_HOST_BYTES`, when present, overrides this value.
370pub fn set_weight_offload_host_budget(bytes: u64) -> Result<(), &'static str> {
371    crate::kernels::qmoe::default_weight_offload_host_cache()
372        .reconfigure(bytes)
373        .map_err(|_| "cannot lower host-cache budget while entries are leased")
374}
375
376pub(crate) fn weight_offload_host_budget(governor_bytes: u64) -> Result<usize, &'static str> {
377    if let Some(value) = std::env::var_os(WEIGHT_OFFLOAD_HOST_BYTES_ENV) {
378        let value = value
379            .to_str()
380            .ok_or("host-cache byte budget is not valid UTF-8")?;
381        let bytes = value
382            .parse::<u64>()
383            .map_err(|_| "host-cache byte budget must be an unsigned decimal byte count")?;
384        return checked_host_budget(bytes);
385    }
386    checked_host_budget(governor_bytes)
387}
388
389pub(crate) fn checked_host_budget(bytes: u64) -> Result<usize, &'static str> {
390    let bytes = usize::try_from(bytes).map_err(|_| "host-cache byte budget exceeds usize::MAX")?;
391    if bytes > isize::MAX as usize {
392        return Err("host-cache byte budget exceeds isize::MAX");
393    }
394    Ok(bytes)
395}
396
397#[cfg(target_os = "linux")]
398fn linux_process_memory_stats() -> Option<LinuxProcessMemoryStats> {
399    let status = std::fs::read_to_string("/proc/self/status").ok()?;
400    let resident_rss_bytes = status
401        .lines()
402        .find_map(|line| line.strip_prefix("VmRSS:"))
403        .and_then(|value| value.split_whitespace().next())
404        .and_then(|value| value.parse::<u64>().ok())
405        .and_then(|kib| kib.checked_mul(1024))
406        .unwrap_or(0);
407
408    let stat = std::fs::read_to_string("/proc/self/stat").ok()?;
409    let fields = stat.get(stat.rfind(')')?.checked_add(2)?..)?;
410    let fields = fields.split_whitespace().collect::<Vec<_>>();
411    Some(LinuxProcessMemoryStats {
412        resident_rss_bytes,
413        minor_faults: fields.get(7)?.parse().ok()?,
414        major_faults: fields.get(9)?.parse().ok()?,
415    })
416}
417
418#[cfg(not(target_os = "linux"))]
419fn linux_process_memory_stats() -> Option<LinuxProcessMemoryStats> {
420    None
421}
422
423#[cfg(test)]
424mod tests {
425    use super::*;
426
427    #[test]
428    fn weight_offload_flag_is_opt_in() {
429        assert!(!WeightOffloadMode::from_value(None).enabled);
430        assert!(!WeightOffloadMode::from_value(Some(OsStr::new("0"))).enabled);
431        assert!(WeightOffloadMode::from_value(Some(OsStr::new("1"))).enabled);
432    }
433
434    #[test]
435    fn host_cache_budget_rejects_unaddressable_values() {
436        if usize::BITS == 64 {
437            assert_eq!(
438                checked_host_budget(isize::MAX as u64 + 1),
439                Err("host-cache byte budget exceeds isize::MAX")
440            );
441        }
442    }
443
444    #[test]
445    fn mapped_bytes_sum_distinct_ranges_across_layers() {
446        let metrics = WeightOffloadMetrics::default();
447        let first = ExternalMmapRegion {
448            mapping_id: 7,
449            offset: 0,
450            len: 100,
451        };
452        let second = ExternalMmapRegion {
453            mapping_id: 7,
454            offset: 100,
455            len: 200,
456        };
457        metrics.record_mapped_regions(&[first]).unwrap();
458        metrics.record_mapped_regions(&[second, first]).unwrap();
459        assert_eq!(metrics.snapshot().mapped_bytes, 300);
460    }
461
462    #[cfg(target_os = "linux")]
463    #[test]
464    fn linux_process_counters_are_best_effort_readable() {
465        let stats = linux_process_memory_stats().expect("Linux /proc process counters");
466        assert!(stats.resident_rss_bytes > 0);
467    }
468}