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