1use 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
13pub const WEIGHT_OFFLOAD_ENV: &str = "ONNX_GENAI_WEIGHT_OFFLOAD";
15pub 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#[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#[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 ®ion 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(®ion) && 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
361pub fn weight_offload_stats() -> WeightOffloadStats {
364 metrics().snapshot()
365}
366
367pub 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}