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 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 ®ion 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(®ion) && 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
346pub fn weight_offload_stats() -> WeightOffloadStats {
349 metrics().snapshot()
350}
351
352pub 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}