use std::any::Any;
use std::collections::{HashMap, HashSet};
use std::path::PathBuf;
use std::ptr::NonNull;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use std::time::Duration;
use cudarc::driver::sys;
use cudarc::driver::sys::CUdeviceptr;
use onnx_runtime_ep_api::{
ExecutorArtifactGeneration, ExecutorArtifactProviderId, ExecutorInstanceId, ExternalMmapRegion,
FinalizedExpertBank, FinalizedExpertWeight, LazyDeviceWeightBinder, LazyWeight,
LazyWeightBoundary, MmapRegionSource, PagedWeight, WeightHandleError,
};
use onnx_runtime_ir::{DataType, DeviceId, DeviceType, NodeId, ValueId};
use onnx_runtime_memory_governor::{AllocationReleaseState, Tier, VirtualBacking};
use crate::deferred_release::{
CudaDeferredReleaseQueue, DeferredActionOutcome, DeferredReleaseAction, RetainedOwnership,
};
use crate::pinned_pool::{PinnedStagingPool, PooledStaging};
use crate::prefill_double_buffer::{
CudaPrefillError, CudaPrefillTransfer, LayerTicket, PrefillDoubleBuffer, PrefillLayerRequest,
PrefillReject,
};
use crate::runtime::{CopyCompleted, CudaRuntime, FailedHtodCompletion, PinnedStaging, raw_ptr};
const WEIGHT_SLOT_ALIGN: usize = 256;
const DEFERRED_RELEASE_WAIT_TIMEOUT: Duration = Duration::from_secs(30);
static GLOBAL_PAGE_INS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_HITS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_HIT_BYTES: AtomicU64 = AtomicU64::new(0);
static GLOBAL_EVICTIONS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_BYPASSED_PAGE_INS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_BYPASSED_PAGE_IN_BYTES: AtomicU64 = AtomicU64::new(0);
static GLOBAL_MATERIALIZE_NS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_HTOD_NS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_ADMIT_SYNC_NS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_STAGING_FILL_BYTES: AtomicU64 = AtomicU64::new(0);
static GLOBAL_STAGING_FILL_REGIONS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_STAGING_FILL_CALLS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_MATERIALIZE_FALLBACK_CALLS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_HTOD_BYTES: AtomicU64 = AtomicU64::new(0);
static GLOBAL_VRAM_ALLOC_NS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_VRAM_FREE_NS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_VRAM_FREE_SYNC_NS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_PEAK_RESIDENT_BYTES: AtomicU64 = AtomicU64::new(0);
static GLOBAL_BUDGET_BYTES: AtomicU64 = AtomicU64::new(0);
static GLOBAL_CONTENT_RESIDENT_BYTES: AtomicU64 = AtomicU64::new(0);
static GLOBAL_WEIGHT_MAPPED_BYTES: AtomicU64 = AtomicU64::new(0);
static GLOBAL_ZERO_COPY_BINDS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_ZERO_COPY_READS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_ZERO_COPY_BYTES: AtomicU64 = AtomicU64::new(0);
static GLOBAL_HOST_REGISTERED_BYTES: AtomicU64 = AtomicU64::new(0);
static GLOBAL_ZERO_COPY_BOUND_BYTES: AtomicU64 = AtomicU64::new(0);
static GLOBAL_PREFETCH_ISSUED: AtomicU64 = AtomicU64::new(0);
static GLOBAL_PREFETCH_ISSUED_BYTES: AtomicU64 = AtomicU64::new(0);
static GLOBAL_PREFETCH_PROMOTED: AtomicU64 = AtomicU64::new(0);
static GLOBAL_PREFETCH_PROMOTE_WAIT_NS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_PREFETCH_DECLINED_BUDGET: AtomicU64 = AtomicU64::new(0);
static GLOBAL_PREFETCH_DECLINED_BUSY: AtomicU64 = AtomicU64::new(0);
static GLOBAL_PREFETCH_DECLINED_UNSUPPORTED: AtomicU64 = AtomicU64::new(0);
static GLOBAL_PREFETCH_DECLINED_RESIDENT: AtomicU64 = AtomicU64::new(0);
static GLOBAL_PREFETCH_DECLINED_POOL_CAPACITY: AtomicU64 = AtomicU64::new(0);
fn add_duration(counter: &AtomicU64, elapsed: Duration) {
let nanos = elapsed.as_nanos().min(u128::from(u64::MAX)) as u64;
counter.fetch_add(nanos, Ordering::Relaxed);
}
fn replace_global_budget(old: u64, new: u64) {
if new >= old {
GLOBAL_BUDGET_BYTES.fetch_add(new - old, Ordering::Relaxed);
} else {
let decrease = old - new;
let _ = GLOBAL_BUDGET_BYTES.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| {
Some(current.saturating_sub(decrease))
});
}
}
fn committed_admission_fits(
required_mapped: u64,
zone_available: u64,
required_owned: u64,
global_available: u64,
) -> bool {
required_mapped <= zone_available && required_owned <= global_available
}
fn eviction_made_committed_progress(
before_owned: u64,
after_owned: u64,
before_required_owned: u64,
after_required_owned: u64,
before_required_mapped: u64,
after_required_mapped: u64,
) -> bool {
after_owned < before_owned
|| after_required_owned < before_required_owned
|| after_required_mapped < before_required_mapped
}
fn vmm_committed_authority_matches(
allocator: &crate::vmm_allocator::CudaVmmAllocator,
governor: &dyn onnx_runtime_memory_governor::MemoryGovernor,
) -> bool {
allocator.committed_byte_authority() == Some(governor.authority_id())
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct GlobalOffloadStats {
pub page_ins: u64,
pub hits: u64,
pub hit_bytes: u64,
pub evictions: u64,
pub bypassed_page_ins: u64,
pub bypassed_page_in_bytes: u64,
pub materialize_ns: u64,
pub htod_ns: u64,
pub admit_sync_ns: u64,
pub staging_fill_bytes: u64,
pub staging_fill_regions: u64,
pub staging_fill_calls: u64,
pub materialize_fallback_calls: u64,
pub htod_bytes: u64,
pub vram_alloc_ns: u64,
pub vram_free_ns: u64,
pub vram_free_sync_ns: u64,
pub budget_bytes: u64,
pub peak_resident_bytes: u64,
pub content_resident_bytes: u64,
pub physical_owned_bytes: u64,
pub mapped_physical_bytes: u64,
pub pinned_alloc_calls: u64,
pub pinned_reuses: u64,
pub zero_copy_binds: u64,
pub zero_copy_reads: u64,
pub zero_copy_bytes: u64,
pub host_registered_bytes: u64,
pub prefetch_issued: u64,
pub prefetch_issued_bytes: u64,
pub prefetch_promoted: u64,
pub prefetch_promote_wait_ns: u64,
pub prefetch_declined_budget: u64,
pub prefetch_declined_busy: u64,
pub prefetch_declined_unsupported: u64,
pub prefetch_declined_resident: u64,
pub prefetch_declined_pool_capacity: u64,
}
impl GlobalOffloadStats {
#[must_use]
pub fn byte_hit_rate(&self) -> Option<f64> {
let requested = self.hit_bytes.checked_add(self.htod_bytes)?;
(requested > 0).then(|| self.hit_bytes as f64 / requested as f64)
}
#[must_use]
pub fn bypassed_byte_share(&self) -> Option<f64> {
(self.htod_bytes > 0).then(|| self.bypassed_page_in_bytes as f64 / self.htod_bytes as f64)
}
#[must_use]
pub fn zero_copy_byte_hit_rate(&self) -> Option<f64> {
let requested = self
.hit_bytes
.checked_add(self.htod_bytes)?
.checked_add(self.zero_copy_bytes)?;
(requested > 0).then(|| self.hit_bytes as f64 / requested as f64)
}
}
pub fn global_offload_stats() -> GlobalOffloadStats {
GlobalOffloadStats {
page_ins: GLOBAL_PAGE_INS.load(Ordering::Relaxed),
hits: GLOBAL_HITS.load(Ordering::Relaxed),
hit_bytes: GLOBAL_HIT_BYTES.load(Ordering::Relaxed),
evictions: GLOBAL_EVICTIONS.load(Ordering::Relaxed),
bypassed_page_ins: GLOBAL_BYPASSED_PAGE_INS.load(Ordering::Relaxed),
bypassed_page_in_bytes: GLOBAL_BYPASSED_PAGE_IN_BYTES.load(Ordering::Relaxed),
materialize_ns: GLOBAL_MATERIALIZE_NS.load(Ordering::Relaxed),
htod_ns: GLOBAL_HTOD_NS.load(Ordering::Relaxed),
admit_sync_ns: GLOBAL_ADMIT_SYNC_NS.load(Ordering::Relaxed),
staging_fill_bytes: GLOBAL_STAGING_FILL_BYTES.load(Ordering::Relaxed),
staging_fill_regions: GLOBAL_STAGING_FILL_REGIONS.load(Ordering::Relaxed),
staging_fill_calls: GLOBAL_STAGING_FILL_CALLS.load(Ordering::Relaxed),
materialize_fallback_calls: GLOBAL_MATERIALIZE_FALLBACK_CALLS.load(Ordering::Relaxed),
htod_bytes: GLOBAL_HTOD_BYTES.load(Ordering::Relaxed),
vram_alloc_ns: GLOBAL_VRAM_ALLOC_NS.load(Ordering::Relaxed),
vram_free_ns: GLOBAL_VRAM_FREE_NS.load(Ordering::Relaxed),
vram_free_sync_ns: GLOBAL_VRAM_FREE_SYNC_NS.load(Ordering::Relaxed),
budget_bytes: GLOBAL_BUDGET_BYTES.load(Ordering::Relaxed),
peak_resident_bytes: GLOBAL_PEAK_RESIDENT_BYTES.load(Ordering::Relaxed),
content_resident_bytes: GLOBAL_CONTENT_RESIDENT_BYTES.load(Ordering::Relaxed),
physical_owned_bytes: crate::virtual_memory::total_physical_pool_owned_bytes(),
mapped_physical_bytes: GLOBAL_WEIGHT_MAPPED_BYTES.load(Ordering::Relaxed),
pinned_alloc_calls: crate::pinned_pool::global_pinned_alloc_calls(),
pinned_reuses: crate::pinned_pool::global_pinned_reuses(),
zero_copy_binds: GLOBAL_ZERO_COPY_BINDS.load(Ordering::Relaxed),
zero_copy_reads: GLOBAL_ZERO_COPY_READS.load(Ordering::Relaxed),
zero_copy_bytes: GLOBAL_ZERO_COPY_BYTES.load(Ordering::Relaxed),
host_registered_bytes: GLOBAL_HOST_REGISTERED_BYTES.load(Ordering::Relaxed),
prefetch_issued: GLOBAL_PREFETCH_ISSUED.load(Ordering::Relaxed),
prefetch_issued_bytes: GLOBAL_PREFETCH_ISSUED_BYTES.load(Ordering::Relaxed),
prefetch_promoted: GLOBAL_PREFETCH_PROMOTED.load(Ordering::Relaxed),
prefetch_promote_wait_ns: GLOBAL_PREFETCH_PROMOTE_WAIT_NS.load(Ordering::Relaxed),
prefetch_declined_budget: GLOBAL_PREFETCH_DECLINED_BUDGET.load(Ordering::Relaxed),
prefetch_declined_busy: GLOBAL_PREFETCH_DECLINED_BUSY.load(Ordering::Relaxed),
prefetch_declined_unsupported: GLOBAL_PREFETCH_DECLINED_UNSUPPORTED.load(Ordering::Relaxed),
prefetch_declined_resident: GLOBAL_PREFETCH_DECLINED_RESIDENT.load(Ordering::Relaxed),
prefetch_declined_pool_capacity: GLOBAL_PREFETCH_DECLINED_POOL_CAPACITY
.load(Ordering::Relaxed),
}
}
pub fn reset_global_offload_stats() {
GLOBAL_PAGE_INS.store(0, Ordering::Relaxed);
GLOBAL_HITS.store(0, Ordering::Relaxed);
GLOBAL_HIT_BYTES.store(0, Ordering::Relaxed);
GLOBAL_EVICTIONS.store(0, Ordering::Relaxed);
GLOBAL_BYPASSED_PAGE_INS.store(0, Ordering::Relaxed);
GLOBAL_BYPASSED_PAGE_IN_BYTES.store(0, Ordering::Relaxed);
GLOBAL_MATERIALIZE_NS.store(0, Ordering::Relaxed);
GLOBAL_HTOD_NS.store(0, Ordering::Relaxed);
GLOBAL_ADMIT_SYNC_NS.store(0, Ordering::Relaxed);
GLOBAL_STAGING_FILL_BYTES.store(0, Ordering::Relaxed);
GLOBAL_STAGING_FILL_REGIONS.store(0, Ordering::Relaxed);
GLOBAL_STAGING_FILL_CALLS.store(0, Ordering::Relaxed);
GLOBAL_MATERIALIZE_FALLBACK_CALLS.store(0, Ordering::Relaxed);
GLOBAL_HTOD_BYTES.store(0, Ordering::Relaxed);
GLOBAL_VRAM_ALLOC_NS.store(0, Ordering::Relaxed);
GLOBAL_VRAM_FREE_NS.store(0, Ordering::Relaxed);
GLOBAL_VRAM_FREE_SYNC_NS.store(0, Ordering::Relaxed);
GLOBAL_ZERO_COPY_READS.store(0, Ordering::Relaxed);
GLOBAL_ZERO_COPY_BYTES.store(0, Ordering::Relaxed);
GLOBAL_PREFETCH_ISSUED.store(0, Ordering::Relaxed);
GLOBAL_PREFETCH_ISSUED_BYTES.store(0, Ordering::Relaxed);
GLOBAL_PREFETCH_PROMOTED.store(0, Ordering::Relaxed);
GLOBAL_PREFETCH_PROMOTE_WAIT_NS.store(0, Ordering::Relaxed);
GLOBAL_PREFETCH_DECLINED_BUDGET.store(0, Ordering::Relaxed);
GLOBAL_PREFETCH_DECLINED_BUSY.store(0, Ordering::Relaxed);
GLOBAL_PREFETCH_DECLINED_UNSUPPORTED.store(0, Ordering::Relaxed);
GLOBAL_PREFETCH_DECLINED_RESIDENT.store(0, Ordering::Relaxed);
GLOBAL_PREFETCH_DECLINED_POOL_CAPACITY.store(0, Ordering::Relaxed);
reset_key_trace();
crate::pinned_pool::reset_pinned_pool_counters();
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct WeightKeyTrace {
pub len: u64,
pub hits: u64,
pub retained_page_ins: u64,
pub bypass_page_ins: u64,
}
impl WeightKeyTrace {
#[must_use]
pub fn reads(&self) -> u64 {
self.hits
.saturating_add(self.retained_page_ins)
.saturating_add(self.bypass_page_ins)
}
}
fn key_trace_enabled() -> bool {
static CACHE: OnceLock<bool> = OnceLock::new();
*CACHE.get_or_init(|| {
matches!(
std::env::var("ONNX_GENAI_WEIGHT_PAGING_KEY_TRACE")
.ok()
.as_deref()
.map(|value| value.trim().to_ascii_lowercase()),
Some(ref value) if matches!(value.as_str(), "1" | "true" | "yes" | "on")
)
})
}
fn key_trace_map() -> &'static Mutex<HashMap<u64, WeightKeyTrace>> {
static MAP: OnceLock<Mutex<HashMap<u64, WeightKeyTrace>>> = OnceLock::new();
MAP.get_or_init(|| Mutex::new(HashMap::new()))
}
#[derive(Clone, Copy)]
enum KeyTraceEvent {
Hit,
Retained,
Bypass,
}
fn record_key_trace(key: u64, len: u64, event: KeyTraceEvent) {
if !key_trace_enabled() {
return;
}
let mut map = key_trace_map().lock().unwrap_or_else(|e| e.into_inner());
let row = map.entry(key).or_default();
row.len = len;
match event {
KeyTraceEvent::Hit => row.hits = row.hits.saturating_add(1),
KeyTraceEvent::Retained => row.retained_page_ins = row.retained_page_ins.saturating_add(1),
KeyTraceEvent::Bypass => row.bypass_page_ins = row.bypass_page_ins.saturating_add(1),
}
}
fn reset_key_trace() {
if key_trace_enabled() {
key_trace_map()
.lock()
.unwrap_or_else(|e| e.into_inner())
.clear();
}
}
#[must_use]
pub fn weight_paging_key_trace() -> Vec<(u64, WeightKeyTrace)> {
let mut rows: Vec<(u64, WeightKeyTrace)> = key_trace_map()
.lock()
.unwrap_or_else(|e| e.into_inner())
.iter()
.map(|(&key, &row)| (key, row))
.collect();
rows.sort_by_key(|(_, row)| std::cmp::Reverse(row.bypass_page_ins.saturating_mul(row.len)));
rows
}
pub const WEIGHT_OFFLOAD_ENV: &str = onnx_runtime_ep_cpu::WEIGHT_OFFLOAD_ENV;
pub const WEIGHT_OFFLOAD_DEVICE_BYTES_ENV: &str = "ONNX_GENAI_WEIGHT_OFFLOAD_DEVICE_BYTES";
pub const WEIGHT_OFFLOAD_ASYNC_PAGEIN_ENV: &str = "ONNX_GENAI_WEIGHT_OFFLOAD_ASYNC_PAGEIN";
pub const WEIGHT_OFFLOAD_SCAN_RESISTANT_ENV: &str = "ONNX_GENAI_WEIGHT_OFFLOAD_SCAN_RESISTANT";
pub const WEIGHT_OFFLOAD_BYTE_AWARE_ENV: &str = "ONNX_GENAI_WEIGHT_OFFLOAD_BYTE_AWARE";
pub const PREFILL_DOUBLE_BUFFER_ENV: &str = "ONNX_GENAI_PREFILL_DOUBLE_BUFFER";
pub fn prefill_double_buffer_enabled() -> bool {
std::env::var(PREFILL_DOUBLE_BUFFER_ENV).is_ok_and(|value| {
matches!(
value.trim().to_ascii_lowercase().as_str(),
"1" | "true" | "yes" | "on"
)
})
}
pub(crate) fn async_pagein_from_env_value(value: Option<&str>) -> bool {
match value {
Some(value) => matches!(
value.trim().to_ascii_lowercase().as_str(),
"1" | "true" | "yes" | "on"
),
None => true,
}
}
pub(crate) fn scan_resistant_from_env_value(value: Option<&str>) -> bool {
match value {
Some(value) => !matches!(
value.trim().to_ascii_lowercase().as_str(),
"0" | "false" | "no" | "off"
),
None => true,
}
}
pub(crate) fn byte_aware_from_env_value(value: Option<&str>) -> bool {
match value {
Some(value) => matches!(
value.trim().to_ascii_lowercase().as_str(),
"1" | "true" | "yes" | "on"
),
None => false,
}
}
#[must_use]
pub fn byte_aware_residency_from_env() -> bool {
byte_aware_from_env_value(std::env::var(WEIGHT_OFFLOAD_BYTE_AWARE_ENV).ok().as_deref())
}
pub const WEIGHT_OFFLOAD_ZERO_COPY_HYBRID_ENV: &str = "ONNX_GENAI_ZERO_COPY_HYBRID";
pub(crate) fn zero_copy_hybrid_from_env_value(value: Option<&str>) -> bool {
match value {
Some(value) => matches!(
value.trim().to_ascii_lowercase().as_str(),
"1" | "true" | "yes" | "on"
),
None => false,
}
}
#[must_use]
pub fn zero_copy_hybrid_from_env() -> bool {
zero_copy_hybrid_from_env_value(
std::env::var(WEIGHT_OFFLOAD_ZERO_COPY_HYBRID_ENV)
.ok()
.as_deref(),
)
}
fn zero_copy_copy_instead() -> bool {
static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*V.get_or_init(|| {
zero_copy_hybrid_from_env_value(
std::env::var("ONNX_GENAI_ZERO_COPY_HYBRID_COPY_INSTEAD")
.ok()
.as_deref(),
)
})
}
fn zero_copy_debug() -> bool {
static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*V.get_or_init(|| {
zero_copy_hybrid_from_env_value(
std::env::var("ONNX_GENAI_ZERO_COPY_HYBRID_DEBUG")
.ok()
.as_deref(),
)
})
}
fn zero_copy_prefault() -> bool {
static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*V.get_or_init(|| {
zero_copy_hybrid_from_env_value(
std::env::var("ONNX_GENAI_ZERO_COPY_HYBRID_PREFAULT")
.ok()
.as_deref(),
)
})
}
fn zero_copy_no_readonly() -> bool {
static V: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*V.get_or_init(|| {
zero_copy_hybrid_from_env_value(
std::env::var("ONNX_GENAI_ZERO_COPY_HYBRID_NO_READONLY")
.ok()
.as_deref(),
)
})
}
#[derive(Debug, Clone, PartialEq, Eq)]
enum NumericEnv {
Unset,
Invalid(String),
Value(u64),
}
impl NumericEnv {
fn or_default(self, name: &str, default: u64) -> u64 {
match self {
NumericEnv::Value(value) => value,
NumericEnv::Unset => default,
NumericEnv::Invalid(raw) => {
eprintln!(
"cuda_ep: {name}={raw:?} is not a base-10 byte count; using {default} instead. \
Set a plain integer (e.g. 1073741824 for 1 GiB) — this value did NOT take \
effect, so any measurement taken under it describes the default."
);
default
}
}
}
fn into_option(self, name: &str) -> Option<u64> {
match self {
NumericEnv::Value(value) => Some(value),
NumericEnv::Unset => None,
NumericEnv::Invalid(raw) => {
eprintln!(
"cuda_ep: {name}={raw:?} is not a base-10 integer; ignoring it. This value did \
NOT take effect."
);
None
}
}
}
}
fn parse_numeric_env(name: &str) -> NumericEnv {
let Ok(raw) = std::env::var(name) else {
return NumericEnv::Unset;
};
match raw.trim().parse::<u64>() {
Ok(value) => NumericEnv::Value(value),
Err(_) => NumericEnv::Invalid(raw),
}
}
fn zero_copy_max_binds() -> Option<u64> {
const NAME: &str = "ONNX_GENAI_ZERO_COPY_HYBRID_MAX_BINDS";
static V: std::sync::OnceLock<Option<u64>> = std::sync::OnceLock::new();
*V.get_or_init(|| parse_numeric_env(NAME).into_option(NAME))
}
const ZERO_COPY_SAFE_BUDGET_BYTES_WDDM: u64 = 256 * 1024 * 1024;
const ZERO_COPY_SAFE_BUDGET_BYTES_NON_WINDOWS: u64 = 2 * 1024 * 1024 * 1024;
const ZERO_COPY_SAFE_BUDGET_BYTES: u64 = if cfg!(target_os = "windows") {
ZERO_COPY_SAFE_BUDGET_BYTES_WDDM
} else {
ZERO_COPY_SAFE_BUDGET_BYTES_NON_WINDOWS
};
fn zero_copy_budget_bytes() -> u64 {
const NAME: &str = "ONNX_GENAI_ZERO_COPY_HYBRID_BUDGET_BYTES";
static V: std::sync::OnceLock<u64> = std::sync::OnceLock::new();
*V.get_or_init(|| parse_numeric_env(NAME).or_default(NAME, ZERO_COPY_SAFE_BUDGET_BYTES))
}
pub const WEIGHT_OFFLOAD_EVICT_ORDER_ENV: &str = "ONNX_GENAI_WEIGHT_OFFLOAD_EVICT_ORDER";
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
pub enum EvictOrderProbe {
#[default]
Lru,
Mru,
Smallest,
Largest,
}
pub(crate) fn evict_order_from_env_value(value: Option<&str>) -> EvictOrderProbe {
match value.map(|value| value.trim().to_ascii_lowercase()) {
Some(value) => match value.as_str() {
"mru" | "reverse" => EvictOrderProbe::Mru,
"smallest" | "small" => EvictOrderProbe::Smallest,
"largest" | "large" => EvictOrderProbe::Largest,
_ => EvictOrderProbe::Lru,
},
None => EvictOrderProbe::Lru,
}
}
#[must_use]
pub fn evict_order_probe_from_env() -> EvictOrderProbe {
evict_order_from_env_value(
std::env::var(WEIGHT_OFFLOAD_EVICT_ORDER_ENV)
.ok()
.as_deref(),
)
}
const WEIGHT_OFFLOAD_SYNC_BEFORE_FILL_ENV: &str = "ONNX_GENAI_WEIGHT_OFFLOAD_SYNC_BEFORE_FILL";
fn sync_before_fill_enabled() -> bool {
static CACHE: OnceLock<bool> = OnceLock::new();
*CACHE.get_or_init(|| {
matches!(
std::env::var(WEIGHT_OFFLOAD_SYNC_BEFORE_FILL_ENV)
.ok()
.as_deref()
.map(|value| value.trim().to_ascii_lowercase()),
Some(ref value) if matches!(value.as_str(), "1" | "true" | "yes" | "on")
)
})
}
const WEIGHT_OFFLOAD_RETAIN_SLOTTED_ENV: &str = "ONNX_GENAI_WEIGHT_OFFLOAD_RETAIN_SLOTTED";
fn retain_slotted_enabled() -> bool {
static CACHE: OnceLock<bool> = OnceLock::new();
*CACHE.get_or_init(|| {
matches!(
std::env::var(WEIGHT_OFFLOAD_RETAIN_SLOTTED_ENV)
.ok()
.as_deref()
.map(|value| value.trim().to_ascii_lowercase()),
Some(ref value) if matches!(value.as_str(), "1" | "true" | "yes" | "on")
)
})
}
const WEIGHT_OFFLOAD_PIN_THRESHOLD_ENV: &str = "ONNX_GENAI_WEIGHT_OFFLOAD_PIN_THRESHOLD_BYTES";
const WEIGHT_OFFLOAD_PIN_BUDGET_ENV: &str = "ONNX_GENAI_WEIGHT_OFFLOAD_PIN_BUDGET_BYTES";
fn static_pin_config() -> Option<(u64, u64)> {
static CACHE: OnceLock<Option<(u64, u64)>> = OnceLock::new();
*CACHE.get_or_init(|| {
let threshold = parse_numeric_env(WEIGHT_OFFLOAD_PIN_THRESHOLD_ENV)
.into_option(WEIGHT_OFFLOAD_PIN_THRESHOLD_ENV)?;
let budget = parse_numeric_env(WEIGHT_OFFLOAD_PIN_BUDGET_ENV)
.into_option(WEIGHT_OFFLOAD_PIN_BUDGET_ENV)?;
(threshold > 0 && budget > 0).then_some((threshold, budget))
})
}
const WEIGHT_OFFLOAD_PIN_KEYS_ENV: &str = "ONNX_GENAI_WEIGHT_OFFLOAD_PIN_KEYS";
fn static_pin_keys() -> Option<&'static HashSet<u64>> {
static CACHE: OnceLock<Option<HashSet<u64>>> = OnceLock::new();
CACHE
.get_or_init(|| {
let raw = std::env::var(WEIGHT_OFFLOAD_PIN_KEYS_ENV).ok()?;
let keys: HashSet<u64> = raw
.split(',')
.filter_map(|token| token.trim().parse::<u64>().ok())
.collect();
(!keys.is_empty()).then_some(keys)
})
.as_ref()
}
const PIN_PROBE_GRANULE_BYTES: usize = 2 * 1024 * 1024;
const WEIGHT_PIN_CHECKSUM_ENV: &str = "ONNX_GENAI_WEIGHT_PIN_CHECKSUM";
const WEIGHT_PIN_REFILL_EVERY_ENV: &str = "ONNX_GENAI_WEIGHT_PIN_REFILL_EVERY";
fn pin_checksum_keys() -> Option<&'static HashSet<u64>> {
static CACHE: OnceLock<Option<HashSet<u64>>> = OnceLock::new();
CACHE
.get_or_init(|| {
let raw = std::env::var(WEIGHT_PIN_CHECKSUM_ENV).ok()?;
let keys: HashSet<u64> = raw
.split(',')
.filter_map(|token| token.trim().parse::<u64>().ok())
.collect();
(!keys.is_empty()).then_some(keys)
})
.as_ref()
}
fn pin_refill_every() -> Option<u64> {
std::env::var(WEIGHT_PIN_REFILL_EVERY_ENV)
.ok()
.and_then(|raw| raw.trim().parse::<u64>().ok())
.filter(|n| *n > 0)
}
const WEIGHT_SLOT_BYPASS_RETAIN_ENV: &str = "ONNX_GENAI_WEIGHT_SLOT_BYPASS_RETAIN";
fn slot_bypass_retain_enabled() -> bool {
static CACHE: OnceLock<bool> = OnceLock::new();
*CACHE.get_or_init(|| {
matches!(
std::env::var(WEIGHT_SLOT_BYPASS_RETAIN_ENV)
.ok()
.as_deref()
.map(|value| value.trim().to_ascii_lowercase()),
Some(ref v) if v == "1" || v == "true" || v == "yes" || v == "on"
)
})
}
fn fnv1a_64(bytes: &[u8]) -> u64 {
let mut hash = 0xcbf2_9ce4_8422_2325u64;
for &byte in bytes {
hash ^= u64::from(byte);
hash = hash.wrapping_mul(0x0000_0100_0000_01b3);
}
hash
}
fn report_pin_granule_diff(key: u64, step: u64, baseline: &[u64], current: &[u64]) {
if baseline == current {
eprintln!(
"weight_pin_checksum[#945]: key={key} step={step} MATCH granules={}",
baseline.len()
);
return;
}
let mut first_changed: Option<usize> = None;
let mut changed = 0usize;
for (index, (want, got)) in baseline.iter().zip(current.iter()).enumerate() {
if want != got {
first_changed.get_or_insert(index);
changed += 1;
}
}
let first = first_changed.unwrap_or(0);
eprintln!(
"weight_pin_checksum[#945]: key={key} step={step} CHANGED first_granule={first} \
byte_offset={} changed_granules={} of {} (len_delta={})",
first * PIN_PROBE_GRANULE_BYTES,
changed,
baseline.len(),
current.len() as isize - baseline.len() as isize,
);
}
static GLOBAL_PINNED_KEYS: AtomicU64 = AtomicU64::new(0);
static GLOBAL_PINNED_BYTES: AtomicU64 = AtomicU64::new(0);
#[must_use]
pub fn pinned_hot_set() -> (u64, u64) {
(
GLOBAL_PINNED_KEYS.load(Ordering::Relaxed),
GLOBAL_PINNED_BYTES.load(Ordering::Relaxed),
)
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct DeviceOffloadPolicy {
pub enabled: bool,
pub managed_no_spill: bool,
pub managed_limit_bytes: Option<u64>,
pub device_budget_bytes: Option<u64>,
pub async_pagein: bool,
pub scan_resistant_dense: bool,
pub byte_aware_residency: bool,
pub evict_order_probe: EvictOrderProbe,
pub zero_copy_hybrid: bool,
}
impl Default for DeviceOffloadPolicy {
fn default() -> Self {
Self {
enabled: false,
managed_no_spill: false,
managed_limit_bytes: None,
device_budget_bytes: None,
async_pagein: false,
scan_resistant_dense: true,
byte_aware_residency: false,
evict_order_probe: EvictOrderProbe::Lru,
zero_copy_hybrid: false,
}
}
}
impl DeviceOffloadPolicy {
pub fn from_env() -> Self {
let enabled = std::env::var_os(WEIGHT_OFFLOAD_ENV).is_some_and(|value| value == "1");
let device_budget_bytes = std::env::var(WEIGHT_OFFLOAD_DEVICE_BYTES_ENV)
.ok()
.and_then(|value| parse_budget_bytes(&value));
let async_pagein = async_pagein_from_env_value(
std::env::var(WEIGHT_OFFLOAD_ASYNC_PAGEIN_ENV)
.ok()
.as_deref(),
);
let scan_resistant_dense = scan_resistant_from_env_value(
std::env::var(WEIGHT_OFFLOAD_SCAN_RESISTANT_ENV)
.ok()
.as_deref(),
);
let byte_aware_residency =
byte_aware_from_env_value(std::env::var(WEIGHT_OFFLOAD_BYTE_AWARE_ENV).ok().as_deref());
let evict_order_probe = evict_order_from_env_value(
std::env::var(WEIGHT_OFFLOAD_EVICT_ORDER_ENV)
.ok()
.as_deref(),
);
let zero_copy_hybrid = zero_copy_hybrid_from_env_value(
std::env::var(WEIGHT_OFFLOAD_ZERO_COPY_HYBRID_ENV)
.ok()
.as_deref(),
);
Self {
enabled,
managed_no_spill: false,
managed_limit_bytes: None,
device_budget_bytes,
async_pagein,
scan_resistant_dense,
byte_aware_residency,
evict_order_probe,
zero_copy_hybrid,
}
}
}
fn parse_budget_bytes(value: &str) -> Option<u64> {
match value.trim().parse::<u64>() {
Ok(bytes) if bytes > 0 => Some(bytes),
_ => None,
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct CudaResidencyStats {
pub budget_bytes: u64,
pub resident_bytes: u64,
pub peak_resident_bytes: u64,
pub pages_resident: u64,
pub page_ins: u64,
pub hits: u64,
pub evictions: u64,
pub physical_owned_bytes: u64,
pub mapped_physical_bytes: u64,
pub admission_no_progress: u64,
pub prefetch_issued: u64,
pub prefetch_issued_bytes: u64,
pub prefetch_promoted: u64,
pub prefetch_promote_wait_ns: u64,
pub prefetch_declined_budget: u64,
pub prefetch_declined_busy: u64,
pub prefetch_declined_unsupported: u64,
pub prefetch_declined_resident: u64,
pub prefetch_declined_pool_capacity: u64,
pub pinned_pool_alloc_calls: u64,
pub pinned_pool_reuses: u64,
}
pub struct CudaWeightPage {
runtime: Arc<CudaRuntime>,
queue: Option<Arc<CudaDeferredReleaseQueue>>,
allocation: WeightAllocation,
ptr: CUdeviceptr,
len: usize,
dtype: DataType,
shape: Vec<usize>,
}
enum WeightAllocation {
Runtime,
Retired,
HostMapped,
Vmm {
allocator: Arc<crate::vmm_allocator::CudaVmmAllocator>,
allowance: onnx_runtime_memory_governor::MappedAllowance,
stable_slot: bool,
slot_state: Option<Arc<SlotOperationState>>,
},
}
#[derive(Debug, Default)]
pub(crate) struct SlotOperationState {
state: std::sync::atomic::AtomicU8,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum SlotStatus {
Idle,
Pending,
Poisoned,
}
impl SlotOperationState {
const IDLE: u8 = 0;
const PENDING: u8 = 1;
const POISONED: u8 = 2;
pub(crate) fn status(&self) -> SlotStatus {
match self.state.load(Ordering::Acquire) {
Self::PENDING => SlotStatus::Pending,
Self::POISONED => SlotStatus::Poisoned,
_ => SlotStatus::Idle,
}
}
fn begin_pending(&self) -> bool {
self.state
.compare_exchange(
Self::IDLE,
Self::PENDING,
Ordering::AcqRel,
Ordering::Acquire,
)
.is_ok()
}
fn begin_release(&self) -> bool {
self.begin_pending()
}
fn begin_refill(&self) -> bool {
self.begin_pending()
}
fn finish_pending(&self) {
let _ = self.state.compare_exchange(
Self::PENDING,
Self::IDLE,
Ordering::AcqRel,
Ordering::Acquire,
);
}
fn finish_release(&self) {
self.finish_pending();
}
fn finish_refill(&self) {
self.finish_pending();
}
fn poison(&self) {
self.state.store(Self::POISONED, Ordering::Release);
}
}
#[derive(Clone, Debug)]
struct StableWeightSlot {
va: CUdeviceptr,
len: usize,
state: Arc<SlotOperationState>,
}
#[derive(Debug)]
enum WeightReleaseAction {
Runtime {
runtime: Arc<CudaRuntime>,
ptr: CUdeviceptr,
len: usize,
},
VmmSpan {
allocator: Arc<crate::vmm_allocator::CudaVmmAllocator>,
allowance: onnx_runtime_memory_governor::MappedAllowance,
ptr: CUdeviceptr,
len: usize,
},
VmmSlot {
allocator: Arc<crate::vmm_allocator::CudaVmmAllocator>,
allowance: onnx_runtime_memory_governor::MappedAllowance,
ptr: CUdeviceptr,
len: usize,
slot_state: Option<Arc<SlotOperationState>>,
},
RetainedVmmSlot {
allocator: Arc<crate::vmm_allocator::CudaVmmAllocator>,
allowance: onnx_runtime_memory_governor::MappedAllowance,
ptr: CUdeviceptr,
len: usize,
slot_state: Option<Arc<SlotOperationState>>,
},
}
impl WeightReleaseAction {
fn bytes(&self) -> u64 {
match self {
Self::Runtime { len, .. }
| Self::VmmSpan { len, .. }
| Self::VmmSlot { len, .. }
| Self::RetainedVmmSlot { len, .. } => *len as u64,
}
}
fn poison_slot(&self) {
match self {
Self::VmmSlot {
slot_state: Some(state),
..
}
| Self::RetainedVmmSlot {
slot_state: Some(state),
..
} => state.poison(),
_ => {}
}
}
fn refund(allowance: &onnx_runtime_memory_governor::MappedAllowance, unmapped: u64) {
if unmapped == 0 {
return;
}
allowance.unmap(unmapped);
let _ = GLOBAL_WEIGHT_MAPPED_BYTES.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|current| Some(current.saturating_sub(unmapped)),
);
}
fn run(self) -> DeferredActionOutcome {
let free_start = std::time::Instant::now();
let outcome = self.run_inner();
add_duration(&GLOBAL_VRAM_FREE_NS, free_start.elapsed());
outcome
}
fn run_inner(self) -> DeferredActionOutcome {
match self {
Self::Runtime { runtime, ptr, len } => {
match unsafe { runtime.free_raw(ptr) } {
Ok(()) => DeferredActionOutcome::released(0),
Err(error) => {
let detail =
format!("cuMemFree of a {len} byte weight page failed: {error}");
DeferredActionOutcome::quarantined(
AllocationReleaseState::Quarantined,
0,
detail.clone(),
Some(RetainedOwnership {
bytes: len as u64,
detail,
keep_alive: Box::new((runtime, ptr, len)),
}),
)
}
}
}
Self::VmmSpan {
allocator,
allowance,
ptr,
len,
} => {
let Some(ptr) = NonNull::new(ptr as *mut u8) else {
return DeferredActionOutcome::released(0);
};
let outcome = unsafe {
onnx_runtime_memory_governor::DeviceAllocator::release(
allocator.as_ref(),
ptr,
len,
WEIGHT_SLOT_ALIGN,
)
};
match outcome {
onnx_runtime_memory_governor::AllocationReleaseOutcome::Complete {
accounting,
} => {
Self::refund(&allowance, accounting.unmapped_bytes);
DeferredActionOutcome::released(accounting.unmapped_bytes)
}
onnx_runtime_memory_governor::AllocationReleaseOutcome::Quarantined {
accounting,
residual,
} => {
Self::refund(&allowance, accounting.unmapped_bytes);
DeferredActionOutcome::quarantined(
residual.state,
accounting.unmapped_bytes,
format!(
"{} ({} byte(s) retained at {:#x})",
residual.reason, residual.retained_bytes, residual.address
),
Some(RetainedOwnership {
bytes: residual.retained_bytes,
detail: String::from("transient VMM weight page"),
keep_alive: Box::new((allocator, allowance)),
}),
)
}
onnx_runtime_memory_governor::AllocationReleaseOutcome::Failed { failure } => {
DeferredActionOutcome::quarantined(
AllocationReleaseState::Quarantined,
0,
failure.to_string(),
Some(RetainedOwnership {
bytes: len as u64,
detail: String::from("transient VMM weight page"),
keep_alive: Box::new((allocator, allowance)),
}),
)
}
}
}
Self::VmmSlot {
allocator,
allowance,
ptr,
len,
slot_state,
} => {
let Some(ptr) = NonNull::new(ptr as *mut u8) else {
if let Some(state) = slot_state.as_ref() {
state.finish_release();
}
return DeferredActionOutcome::released(0);
};
let outcome = allocator.decommit_allocation_range_outcome(ptr, len, 0, len);
match outcome {
Ok(crate::vmm_allocator::DecommitOutcome::Complete { accounting }) => {
Self::refund(&allowance, accounting.unmapped_bytes);
if let Some(state) = slot_state.as_ref() {
state.finish_release();
}
if accounting.quarantined_owned_bytes == 0 {
DeferredActionOutcome::released(accounting.unmapped_bytes)
} else {
DeferredActionOutcome::quarantined(
AllocationReleaseState::Quarantined,
accounting.unmapped_bytes,
format!(
"stable-slot decommit unmapped successfully but retained {} \
byte(s) of quarantined physical handles",
accounting.quarantined_owned_bytes
),
Some(RetainedOwnership {
bytes: accounting.quarantined_owned_bytes,
detail: String::from(
"stable-VA weight-slot physical-handle quarantine",
),
keep_alive: Box::new((allocator, allowance)),
}),
)
}
}
Ok(crate::vmm_allocator::DecommitOutcome::RolledBack { reason }) => {
if let Some(state) = slot_state.as_ref() {
state.poison();
}
DeferredActionOutcome::quarantined(
AllocationReleaseState::Quarantined,
0,
format!("stable-slot decommit rolled back: {reason}"),
Some(RetainedOwnership {
bytes: len as u64,
detail: String::from("stable-VA weight slot"),
keep_alive: Box::new((allocator, allowance)),
}),
)
}
Ok(crate::vmm_allocator::DecommitOutcome::Quarantined {
accounting,
residual,
reason,
}) => {
Self::refund(&allowance, accounting.unmapped_bytes);
if let Some(state) = slot_state.as_ref() {
state.poison();
}
DeferredActionOutcome::quarantined(
residual.state,
accounting.unmapped_bytes,
format!(
"stable-slot decommit quarantined: {reason} ({} byte(s) retained \
at {:#x})",
residual.retained_bytes, residual.address
),
Some(RetainedOwnership {
bytes: residual.retained_bytes,
detail: String::from("stable-VA weight slot"),
keep_alive: Box::new((allocator, allowance)),
}),
)
}
Err(error) => {
if let Some(state) = slot_state.as_ref() {
state.poison();
}
DeferredActionOutcome::quarantined(
AllocationReleaseState::Quarantined,
0,
format!("stable-slot decommit refused: {error}"),
Some(RetainedOwnership {
bytes: len as u64,
detail: String::from("stable-VA weight slot"),
keep_alive: Box::new((allocator, allowance)),
}),
)
}
}
}
Self::RetainedVmmSlot {
allocator,
allowance,
ptr,
len,
slot_state,
} => DeferredActionOutcome::quarantined(
AllocationReleaseState::Quarantined,
0,
format!("stable weight slot at {ptr:#x} already had a pending or poisoned release"),
Some(RetainedOwnership {
bytes: len as u64,
detail: String::from("duplicate stable-VA weight-slot release"),
keep_alive: Box::new((allocator, allowance, slot_state)),
}),
),
}
}
}
#[derive(Debug)]
struct WeightPageRelease {
action: Option<WeightReleaseAction>,
}
impl DeferredReleaseAction for WeightPageRelease {
fn execute(mut self: Box<Self>) -> DeferredActionOutcome {
let Some(action) = self.action.take() else {
return DeferredActionOutcome::quarantined(
AllocationReleaseState::Quarantined,
0,
"a weight page release ran without its action",
None,
);
};
action.run()
}
fn label(&self) -> &'static str {
"weight page"
}
fn bytes(&self) -> u64 {
self.action.as_ref().map_or(0, WeightReleaseAction::bytes)
}
}
impl Drop for WeightPageRelease {
fn drop(&mut self) {
if let Some(action) = self.action.take() {
eprintln!(
"cuda_ep: WARNING: a deferred weight page release of {} byte(s) was abandoned; \
its device memory is retained rather than freed",
action.bytes()
);
action.poison_slot();
std::mem::forget(action);
}
}
}
impl CudaWeightPage {
pub fn with_deferred_release_queue(mut self, queue: Arc<CudaDeferredReleaseQueue>) -> Self {
self.queue = Some(queue);
self
}
fn take_release_action(&mut self) -> Option<WeightReleaseAction> {
let allocation = std::mem::replace(&mut self.allocation, WeightAllocation::Retired);
match allocation {
WeightAllocation::Retired => None,
WeightAllocation::HostMapped => {
None
}
WeightAllocation::Runtime => Some(WeightReleaseAction::Runtime {
runtime: Arc::clone(&self.runtime),
ptr: self.ptr,
len: self.len,
}),
WeightAllocation::Vmm {
allocator,
allowance,
stable_slot,
slot_state,
} => {
if stable_slot {
if let Some(state) = slot_state.as_ref()
&& !state.begin_release()
{
eprintln!(
"cuda_ep: WARNING: a stable weight slot at {:#x} already has an \
unfinished release; retaining this page's physical granules rather \
than decommitting them twice",
self.ptr
);
return Some(WeightReleaseAction::RetainedVmmSlot {
allocator,
allowance,
ptr: self.ptr,
len: self.len,
slot_state,
});
}
Some(WeightReleaseAction::VmmSlot {
allocator,
allowance,
ptr: self.ptr,
len: self.len,
slot_state,
})
} else {
Some(WeightReleaseAction::VmmSpan {
allocator,
allowance,
ptr: self.ptr,
len: self.len,
})
}
}
}
}
fn release_allocation(&mut self) {
let Some(action) = self.take_release_action() else {
return;
};
let Some(queue) = self.queue.clone() else {
eprintln!(
"cuda_ep: WARNING: retaining a {} byte weight page because no deferred-release \
queue is installed",
action.bytes()
);
action.poison_slot();
std::mem::forget(action);
return;
};
let bytes = action.bytes();
if let Err(refused) = queue.enqueue(WeightPageRelease {
action: Some(action),
}) {
eprintln!(
"cuda_ep: WARNING: the deferred release queue refused a {bytes} byte weight page \
release ({}); its device memory is retained rather than freed before in-flight \
work has finished",
refused.rejection.name()
);
}
}
fn retire_after_stream_sync(&mut self) {
self.release_allocation();
}
pub fn upload(
runtime: &Arc<CudaRuntime>,
dtype: DataType,
shape: Vec<usize>,
bytes: &[u8],
) -> Result<Self, WeightHandleError> {
Self::upload_inner(runtime, dtype, shape, bytes, None)
}
pub(crate) fn upload_queued(
runtime: &Arc<CudaRuntime>,
dtype: DataType,
shape: Vec<usize>,
bytes: &[u8],
queue: Arc<CudaDeferredReleaseQueue>,
) -> Result<Self, WeightHandleError> {
Self::upload_inner(runtime, dtype, shape, bytes, Some(queue))
}
fn upload_inner(
runtime: &Arc<CudaRuntime>,
dtype: DataType,
shape: Vec<usize>,
bytes: &[u8],
queue: Option<Arc<CudaDeferredReleaseQueue>>,
) -> Result<Self, WeightHandleError> {
if bytes.is_empty() {
return Err(WeightHandleError::MissingRegions);
}
let alloc_start = std::time::Instant::now();
let ptr = runtime
.alloc_raw(bytes.len())
.map_err(|error| WeightHandleError::DeviceBinding(format!("VRAM alloc: {error}")))?;
add_duration(&GLOBAL_VRAM_ALLOC_NS, alloc_start.elapsed());
let page = Self {
runtime: Arc::clone(runtime),
queue,
allocation: WeightAllocation::Runtime,
ptr,
len: bytes.len(),
dtype,
shape,
};
let copy_start = std::time::Instant::now();
unsafe { runtime.htod(bytes, ptr) }
.map_err(|error| WeightHandleError::DeviceBinding(format!("H2D copy: {error}")))?;
add_duration(&GLOBAL_HTOD_NS, copy_start.elapsed());
GLOBAL_HTOD_BYTES.fetch_add(bytes.len() as u64, Ordering::Relaxed);
Ok(page)
}
pub fn upload_async(
runtime: &Arc<CudaRuntime>,
dtype: DataType,
shape: Vec<usize>,
bytes: &[u8],
staging: PinnedStaging,
) -> Result<(Self, u64, PinnedStaging), WeightHandleError> {
Self::upload_async_inner(runtime, dtype, shape, bytes, staging, None)
}
fn upload_async_inner(
runtime: &Arc<CudaRuntime>,
dtype: DataType,
shape: Vec<usize>,
bytes: &[u8],
mut staging: PinnedStaging,
queue: Option<Arc<CudaDeferredReleaseQueue>>,
) -> Result<(Self, u64, PinnedStaging), WeightHandleError> {
if bytes.is_empty() {
return Err(WeightHandleError::MissingRegions);
}
if staging.len() < bytes.len() {
return Err(WeightHandleError::InvalidResident(format!(
"pinned staging buffer is too small: {} < {}",
staging.len(),
bytes.len()
)));
}
let alloc_start = std::time::Instant::now();
let ptr = runtime
.alloc_raw(bytes.len())
.map_err(|error| WeightHandleError::DeviceBinding(format!("VRAM alloc: {error}")))?;
add_duration(&GLOBAL_VRAM_ALLOC_NS, alloc_start.elapsed());
staging.as_mut_slice()[..bytes.len()].copy_from_slice(bytes);
let page = Self {
runtime: Arc::clone(runtime),
queue,
allocation: WeightAllocation::Runtime,
ptr,
len: bytes.len(),
dtype,
shape,
};
let staged = &staging.as_slice()[..bytes.len()];
let copy_start = std::time::Instant::now();
unsafe { runtime.htod_async(staged, ptr) }.map_err(|error| {
WeightHandleError::DeviceBinding(format!("async H2D copy: {error}"))
})?;
let fence = runtime
.record_copy_fence()
.map_err(|error| WeightHandleError::DeviceBinding(format!("copy fence: {error}")))?;
add_duration(&GLOBAL_HTOD_NS, copy_start.elapsed());
GLOBAL_HTOD_BYTES.fetch_add(bytes.len() as u64, Ordering::Relaxed);
Ok((page, fence, staging))
}
pub fn upload_staged_async(
runtime: &Arc<CudaRuntime>,
dtype: DataType,
shape: Vec<usize>,
len: usize,
staging: PinnedStaging,
) -> Result<(Self, u64, PinnedStaging, CopyCompleted), WeightHandleError> {
Self::upload_staged_async_inner(runtime, dtype, shape, len, staging, None)
}
pub(crate) fn upload_staged_async_queued(
runtime: &Arc<CudaRuntime>,
dtype: DataType,
shape: Vec<usize>,
len: usize,
staging: PinnedStaging,
queue: Arc<CudaDeferredReleaseQueue>,
) -> Result<(Self, u64, PinnedStaging, CopyCompleted), WeightHandleError> {
Self::upload_staged_async_inner(runtime, dtype, shape, len, staging, Some(queue))
}
fn upload_staged_async_inner(
runtime: &Arc<CudaRuntime>,
dtype: DataType,
shape: Vec<usize>,
len: usize,
staging: PinnedStaging,
queue: Option<Arc<CudaDeferredReleaseQueue>>,
) -> Result<(Self, u64, PinnedStaging, CopyCompleted), WeightHandleError> {
if len == 0 {
return Err(WeightHandleError::MissingRegions);
}
if staging.len() < len {
return Err(WeightHandleError::InvalidResident(format!(
"pinned staging buffer is too small: {} < {}",
staging.len(),
len
)));
}
let alloc_start = std::time::Instant::now();
let ptr = runtime
.alloc_raw(len)
.map_err(|error| WeightHandleError::DeviceBinding(format!("VRAM alloc: {error}")))?;
add_duration(&GLOBAL_VRAM_ALLOC_NS, alloc_start.elapsed());
let page = Self {
runtime: Arc::clone(runtime),
queue,
allocation: WeightAllocation::Runtime,
ptr,
len,
dtype,
shape,
};
let staged = &staging.as_slice()[..len];
let (copy_ms, completed) = match unsafe { runtime.htod_async_elapsed_ms(staged, ptr) } {
Ok(result) => result,
Err(error) => {
let (detail, completion) = error.into_parts();
let error =
WeightHandleError::DeviceBinding(format!("measured H2D copy: {detail}"));
match completion {
FailedHtodCompletion::NotSubmitted => return Err(error),
FailedHtodCompletion::Completed(_) => return Err(error),
FailedHtodCompletion::MayBeInFlight => {
quarantine_in_flight_fill(Box::new((page, staging)));
return Err(WeightHandleError::DeviceBinding(format!(
"{error}; destination and staging source were quarantined because \
copy-stream completion could not be established"
)));
}
}
}
};
GLOBAL_HTOD_NS.fetch_add((copy_ms * 1_000_000.0) as u64, Ordering::Relaxed);
GLOBAL_HTOD_BYTES.fetch_add(len as u64, Ordering::Relaxed);
Ok((page, 0, staging, completed))
}
pub fn device_ptr(&self) -> *const std::ffi::c_void {
raw_ptr(self.ptr)
}
pub fn len(&self) -> usize {
self.len
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
pub fn dtype(&self) -> DataType {
self.dtype
}
pub fn shape(&self) -> &[usize] {
&self.shape
}
}
pub(crate) fn fill_staging_from_regions(
weight: &LazyWeight,
source: &dyn MmapRegionSource,
staging: &mut PinnedStaging,
) -> Result<(), WeightHandleError> {
let total = weight.region_bytes_len();
if total == 0 {
return Err(WeightHandleError::MissingRegions);
}
if staging.len() < total {
return Err(WeightHandleError::InvalidResident(format!(
"pinned staging buffer is too small: {} < {}",
staging.len(),
total
)));
}
let mut offset = 0usize;
for region in &weight.regions {
let bytes = source.region_bytes(region)?;
if bytes.len() != region.len {
return Err(WeightHandleError::DeviceBinding(format!(
"region source returned {} bytes for a {}-byte region",
bytes.len(),
region.len
)));
}
let end = offset.checked_add(bytes.len()).ok_or_else(|| {
WeightHandleError::InvalidResident("staging byte count overflow".into())
})?;
staging.as_mut_slice()[offset..end].copy_from_slice(bytes);
offset = end;
}
GLOBAL_STAGING_FILL_CALLS.fetch_add(1, Ordering::Relaxed);
GLOBAL_STAGING_FILL_REGIONS.fetch_add(weight.regions.len() as u64, Ordering::Relaxed);
GLOBAL_STAGING_FILL_BYTES.fetch_add(total as u64, Ordering::Relaxed);
Ok(())
}
impl std::fmt::Debug for CudaWeightPage {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("CudaWeightPage")
.field("len", &self.len)
.field("dtype", &self.dtype)
.field("shape", &self.shape)
.finish_non_exhaustive()
}
}
impl Drop for CudaWeightPage {
fn drop(&mut self) {
self.release_allocation();
}
}
pub struct CudaWeightPager<'a, S: MmapRegionSource + ?Sized> {
runtime: Arc<CudaRuntime>,
queue: Option<Arc<CudaDeferredReleaseQueue>>,
context_scope: Option<onnx_runtime_memory_governor::MemoryContextScope>,
source: &'a S,
}
impl<'a, S: MmapRegionSource + ?Sized> CudaWeightPager<'a, S> {
pub fn new(runtime: Arc<CudaRuntime>, source: &'a S) -> Self {
runtime.set_weights_may_be_paged();
Self {
runtime,
queue: None,
context_scope: None,
source,
}
}
pub fn with_deferred_release_queue(mut self, queue: Arc<CudaDeferredReleaseQueue>) -> Self {
self.queue = Some(queue);
self
}
pub fn with_context_scope(
mut self,
scope: onnx_runtime_memory_governor::MemoryContextScope,
) -> Self {
self.context_scope = Some(scope);
self
}
}
impl<S: MmapRegionSource + ?Sized> LazyDeviceWeightBinder for CudaWeightPager<'_, S> {
type Binding = CudaWeightPage;
fn bind_block_quantized_moe(
&self,
weight: &LazyWeight,
) -> Result<Self::Binding, WeightHandleError> {
let _context_operation = self
.context_scope
.as_ref()
.map(onnx_runtime_memory_governor::MemoryContextScope::enter)
.transpose()
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))?;
let total = weight.region_bytes_len();
if total == 0 {
return Err(WeightHandleError::MissingRegions);
}
let ptr = self
.runtime
.alloc_raw(total)
.map_err(|error| WeightHandleError::DeviceBinding(format!("VRAM alloc: {error}")))?;
let page = CudaWeightPage {
runtime: Arc::clone(&self.runtime),
queue: self.queue.clone(),
allocation: WeightAllocation::Runtime,
ptr,
len: total,
dtype: weight.dtype,
shape: weight.shape.clone(),
};
let mut offset: usize = 0;
for region in &weight.regions {
let bytes = self.source.region_bytes(region)?;
if bytes.len() != region.len {
return Err(WeightHandleError::DeviceBinding(format!(
"region source returned {} bytes for a {}-byte region",
bytes.len(),
region.len
)));
}
let dst = ptr + offset as CUdeviceptr;
let copy_start = std::time::Instant::now();
unsafe { self.runtime.htod(bytes, dst) }
.map_err(|error| WeightHandleError::DeviceBinding(format!("H2D copy: {error}")))?;
add_duration(&GLOBAL_HTOD_NS, copy_start.elapsed());
GLOBAL_HTOD_BYTES.fetch_add(bytes.len() as u64, Ordering::Relaxed);
offset += region.len;
}
Ok(page)
}
}
const HOST_REGISTER_PAGE: usize = 4096;
const CU_MEMHOSTREGISTER_DEVICEMAP: u32 = 0x02;
const CU_MEMHOSTREGISTER_READ_ONLY: u32 = 0x08;
struct HostMapRegistry {
registered: HashMap<usize, (usize, usize)>,
}
impl HostMapRegistry {
fn new() -> Self {
Self {
registered: HashMap::new(),
}
}
fn device_ptr_for(
&mut self,
mapping_id: usize,
mapping: &[u8],
host_ptr: *const u8,
) -> Result<CUdeviceptr, WeightHandleError> {
if let std::collections::hash_map::Entry::Vacant(entry) = self.registered.entry(mapping_id)
{
let base = mapping.as_ptr();
if (base as usize) & (HOST_REGISTER_PAGE - 1) != 0 {
return Err(WeightHandleError::DeviceBinding(format!(
"mmap base {base:p} is not {HOST_REGISTER_PAGE}-byte page aligned; \
cannot host-register for zero-copy"
)));
}
let len = mapping.len();
if zero_copy_prefault() {
let mut acc: u64 = 0;
let mut off = 0usize;
while off < len {
acc = acc.wrapping_add(unsafe { *mapping.get_unchecked(off) } as u64);
off += HOST_REGISTER_PAGE;
}
std::hint::black_box(acc);
}
let flags = if zero_copy_no_readonly() {
CU_MEMHOSTREGISTER_DEVICEMAP
} else {
CU_MEMHOSTREGISTER_DEVICEMAP | CU_MEMHOSTREGISTER_READ_ONLY
};
unsafe { sys::cuMemHostRegister_v2(base as *mut std::ffi::c_void, len, flags) }
.result()
.map_err(|error| {
WeightHandleError::DeviceBinding(format!(
"cuMemHostRegister(DEVICEMAP) of {len} bytes failed: {error}"
))
})?;
entry.insert((base as usize, len));
GLOBAL_HOST_REGISTERED_BYTES.fetch_add(len as u64, Ordering::Relaxed);
}
let mut device_ptr: CUdeviceptr = 0;
unsafe {
sys::cuMemHostGetDevicePointer_v2(&mut device_ptr, host_ptr as *mut std::ffi::c_void, 0)
}
.result()
.map_err(|error| {
WeightHandleError::DeviceBinding(format!("cuMemHostGetDevicePointer failed: {error}"))
})?;
Ok(device_ptr)
}
}
impl Drop for HostMapRegistry {
fn drop(&mut self) {
for &(base, _len) in self.registered.values() {
let _ = unsafe { sys::cuMemHostUnregister(base as *mut std::ffi::c_void) };
}
}
}
enum SpanAdmit {
Filled { bypass: bool },
DeferToZeroCopy,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum FreshSpanCleanup {
CallerOwns,
Consumed,
}
#[derive(Debug)]
struct SpanAdmitError {
error: WeightHandleError,
fresh_span: FreshSpanCleanup,
}
impl From<WeightHandleError> for SpanAdmitError {
fn from(error: WeightHandleError) -> Self {
Self {
error,
fresh_span: FreshSpanCleanup::CallerOwns,
}
}
}
struct VmmFillFailure {
error: WeightHandleError,
in_flight_source: Option<Box<dyn Any + Send>>,
}
impl VmmFillFailure {
fn completed(error: WeightHandleError) -> Self {
Self {
error,
in_flight_source: None,
}
}
fn may_be_in_flight(error: WeightHandleError, source: Box<dyn Any + Send>) -> Self {
Self {
error,
in_flight_source: Some(source),
}
}
}
trait IntoVmmFillResult {
fn into_vmm_fill_result(self) -> Result<(), VmmFillFailure>;
}
impl IntoVmmFillResult for Result<(), WeightHandleError> {
fn into_vmm_fill_result(self) -> Result<(), VmmFillFailure> {
self.map_err(VmmFillFailure::completed)
}
}
impl IntoVmmFillResult for Result<(), VmmFillFailure> {
fn into_vmm_fill_result(self) -> Result<(), VmmFillFailure> {
self
}
}
fn in_flight_fill_quarantine() -> &'static Mutex<Vec<Box<dyn Any + Send>>> {
static QUARANTINE: OnceLock<Mutex<Vec<Box<dyn Any + Send>>>> = OnceLock::new();
QUARANTINE.get_or_init(|| Mutex::new(Vec::new()))
}
fn quarantine_in_flight_fill(ownership: Box<dyn Any + Send>) {
in_flight_fill_quarantine()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.push(ownership);
}
#[cfg(test)]
fn in_flight_fill_quarantine_count() -> usize {
in_flight_fill_quarantine()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.len()
}
#[derive(Debug)]
enum VmmAdmit {
Page(Arc<CudaWeightPage>),
DeferToZeroCopy,
}
impl VmmAdmit {
fn expect_page(self) -> Arc<CudaWeightPage> {
match self {
VmmAdmit::Page(page) => page,
VmmAdmit::DeferToZeroCopy => {
unreachable!("non-hybrid VMM admission never defers to zero-copy")
}
}
}
}
pub struct CudaWeightResidency {
runtime: Arc<CudaRuntime>,
queue: Option<Arc<CudaDeferredReleaseQueue>>,
scan_resistant_dense: bool,
byte_aware: bool,
evict_order_probe: EvictOrderProbe,
zero_copy_hybrid: bool,
host_registry: Mutex<HostMapRegistry>,
physical: OnceLock<PhysicalAdmission>,
context_scope: OnceLock<onnx_runtime_memory_governor::MemoryContextScope>,
context_terminated: AtomicBool,
staging_pool: Arc<PinnedStagingPool>,
route_reservations: Mutex<HashMap<ExecutorInstanceId, Arc<RouteReservationSet>>>,
route_retirement_counters: Arc<RouteReservationRetirementCounters>,
inner: Mutex<ResidencyInner>,
routed_guards_active: AtomicU64,
prefill_double_buffer_enabled: bool,
prefill_pipeline: Mutex<Option<PrefillPipeline>>,
}
pub struct RoutedResidencyGuard {
proof: onnx_runtime_ep_api::RoutedResidencyProof,
residency: Arc<CudaWeightResidency>,
_reservation_use: Option<RouteReservationUseGuard>,
}
impl RoutedResidencyGuard {
pub fn proof(&self) -> &onnx_runtime_ep_api::RoutedResidencyProof {
&self.proof
}
}
impl onnx_runtime_ep_api::RoutedResidencyGuardHandle for RoutedResidencyGuard {
fn proof(&self) -> &onnx_runtime_ep_api::RoutedResidencyProof {
&self.proof
}
}
impl Drop for RoutedResidencyGuard {
fn drop(&mut self) {
self.residency
.routed_guards_active
.fetch_sub(1, Ordering::SeqCst);
}
}
struct PhysicalAdmission {
allocator: Arc<crate::vmm_allocator::CudaVmmAllocator>,
governor: Arc<dyn onnx_runtime_memory_governor::MemoryGovernor + Send + Sync>,
}
struct RouteWeightReservation {
identity: FinalizedExpertWeight,
allocator: Arc<crate::vmm_allocator::CudaVmmAllocator>,
health: Arc<RouteReservationHealth>,
ptr: CUdeviceptr,
len: usize,
}
struct RouteReservationSet {
by_key: HashMap<u64, Arc<RouteWeightReservation>>,
catalogs: HashMap<ValueId, onnx_runtime_loader::WeightRegionCatalog>,
allocators: HashMap<ValueId, Arc<crate::vmm_allocator::CudaVmmAllocator>>,
device_pool: Arc<onnx_runtime_cuda_memory::virtual_memory::PhysicalHandlePool>,
host_pool: Arc<onnx_runtime_cuda_memory::virtual_memory::PhysicalHandlePool>,
groups: Vec<onnx_runtime_ep_api::ExpertWeightGroup>,
health: Arc<RouteReservationHealth>,
}
#[derive(Debug)]
pub(crate) struct RouteReservationRetirementResources {
_allocators: Vec<Arc<crate::vmm_allocator::CudaVmmAllocator>>,
_device_pool: Arc<onnx_runtime_cuda_memory::virtual_memory::PhysicalHandlePool>,
_host_pool: Arc<onnx_runtime_cuda_memory::virtual_memory::PhysicalHandlePool>,
}
#[doc(hidden)]
pub struct RouteReservationHealth {
identity: Option<RouteReservationIdentity>,
lifecycle: AtomicU64,
reason: Mutex<Option<String>>,
retirement_cleanup: Mutex<Option<Box<dyn RouteReservationRetirementCleanup>>>,
retirement_counters: Arc<RouteReservationRetirementCounters>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct RouteReservationIdentity {
provider: ExecutorArtifactProviderId,
executor: ExecutorInstanceId,
artifact_generation: ExecutorArtifactGeneration,
device_ordinal: u32,
reservation_generation: u64,
}
const ROUTE_HEALTH_POISONED: u64 = 1;
const ROUTE_HEALTH_TRANSITIONING: u64 = 1 << 1;
const ROUTE_HEALTH_RETIRING: u64 = 1 << 2;
const ROUTE_HEALTH_RETIRED: u64 = 1 << 3;
const ROUTE_HEALTH_USE: u64 = 1 << 4;
static NEXT_ROUTE_RESERVATION_GENERATION: AtomicU64 = AtomicU64::new(1);
fn next_route_reservation_generation(counter: &AtomicU64) -> Option<u64> {
counter
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |next| {
next.checked_add(1)
})
.ok()
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct RouteReservationRetirementStats {
pub retirements_started: u64,
pub deferred_cleanups: u64,
pub cleanups_scheduled: u64,
pub cleanups_executed: u64,
}
#[derive(Debug, Default)]
pub(crate) struct RouteReservationRetirementCounters {
retirements_started: AtomicU64,
deferred_cleanups: AtomicU64,
cleanups_scheduled: AtomicU64,
cleanups_executed: AtomicU64,
}
impl RouteReservationRetirementCounters {
fn snapshot(&self) -> RouteReservationRetirementStats {
RouteReservationRetirementStats {
retirements_started: self.retirements_started.load(Ordering::Acquire),
deferred_cleanups: self.deferred_cleanups.load(Ordering::Acquire),
cleanups_scheduled: self.cleanups_scheduled.load(Ordering::Acquire),
cleanups_executed: self.cleanups_executed.load(Ordering::Acquire),
}
}
pub(crate) fn record_cleanup_executed(&self) {
self.cleanups_executed.fetch_add(1, Ordering::AcqRel);
}
}
pub(crate) trait RouteReservationRetirementCleanup: Send + std::fmt::Debug {
fn schedule(self: Box<Self>);
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum RouteReservationRetirementStart {
Started,
AlreadyRetiring,
Retired,
}
pub(crate) struct RouteReservationUseGuard {
health: Arc<RouteReservationHealth>,
}
impl onnx_runtime_ep_api::ExecutorArtifactUseGuard for RouteReservationUseGuard {}
impl Drop for RouteReservationUseGuard {
fn drop(&mut self) {
self.health
.lifecycle
.fetch_sub(ROUTE_HEALTH_USE, Ordering::AcqRel);
self.health.try_schedule_retirement_cleanup();
}
}
pub(crate) struct RouteReservationTransitionGuard {
health: Arc<RouteReservationHealth>,
completed: bool,
}
impl RouteReservationTransitionGuard {
pub(crate) fn complete(mut self) {
self.completed = true;
self.health
.lifecycle
.fetch_and(!ROUTE_HEALTH_TRANSITIONING, Ordering::AcqRel);
self.health.try_schedule_retirement_cleanup();
}
pub(crate) fn poison(mut self, reason: String) {
self.health.mark_unusable(reason);
self.completed = true;
self.health
.lifecycle
.fetch_and(!ROUTE_HEALTH_TRANSITIONING, Ordering::Release);
self.health.try_schedule_retirement_cleanup();
}
}
impl Drop for RouteReservationTransitionGuard {
fn drop(&mut self) {
if !self.completed {
self.health.mark_unusable(
"reservation transition authority exited without publishing a terminal outcome"
.to_string(),
);
self.health
.lifecycle
.fetch_and(!ROUTE_HEALTH_TRANSITIONING, Ordering::Release);
self.health.try_schedule_retirement_cleanup();
}
}
}
pub(crate) struct RouteReservationRequirement {
health: Arc<RouteReservationHealth>,
identity: RouteReservationIdentity,
}
impl onnx_runtime_ep_api::ExecutorArtifactRequirementState for RouteReservationRequirement {
fn acquire_use(
&self,
) -> onnx_runtime_ep_api::Result<Box<dyn onnx_runtime_ep_api::ExecutorArtifactUseGuard>> {
self.health
.acquire_use_identity(self.identity)
.map(|guard| Box::new(guard) as Box<dyn onnx_runtime_ep_api::ExecutorArtifactUseGuard>)
.map_err(|reason| {
onnx_runtime_ep_api::EpError::KernelFailed(format!(
"cuda_ep: executor {} CUDA:{} route-bank reservation is unusable \
(generation {}): {reason}; tear down and rebuild the executor before \
dispatch, capture, or replay",
self.identity.executor.get(),
self.identity.device_ordinal,
self.identity.reservation_generation
))
})
}
}
impl RouteReservationHealth {
pub fn new() -> Arc<Self> {
Arc::new(Self {
identity: None,
lifecycle: AtomicU64::new(0),
reason: Mutex::new(None),
retirement_cleanup: Mutex::new(None),
retirement_counters: Arc::new(RouteReservationRetirementCounters::default()),
})
}
fn new_scoped(
provider: ExecutorArtifactProviderId,
executor: ExecutorInstanceId,
artifact_generation: ExecutorArtifactGeneration,
device_ordinal: u32,
retirement_counters: Arc<RouteReservationRetirementCounters>,
) -> Option<Arc<Self>> {
let reservation_generation =
next_route_reservation_generation(&NEXT_ROUTE_RESERVATION_GENERATION)?;
Some(Arc::new(Self {
identity: Some(RouteReservationIdentity {
provider,
executor,
artifact_generation,
device_ordinal,
reservation_generation,
}),
lifecycle: AtomicU64::new(0),
reason: Mutex::new(None),
retirement_cleanup: Mutex::new(None),
retirement_counters,
}))
}
pub(crate) fn generation(&self) -> Option<u64> {
self.identity
.map(|identity| identity.reservation_generation)
}
pub(crate) fn requirement_state(
self: &Arc<Self>,
) -> Result<Arc<dyn onnx_runtime_ep_api::ExecutorArtifactRequirementState>, String> {
let identity = self.identity.ok_or_else(|| {
"reservation health has no provider/executor/artifact/device/generation identity; \
rebuild the executor"
.to_string()
})?;
Ok(Arc::new(RouteReservationRequirement {
health: Arc::clone(self),
identity,
})
as Arc<
dyn onnx_runtime_ep_api::ExecutorArtifactRequirementState,
>)
}
pub(crate) fn mark_unusable(&self, reason: String) {
let mut stored = self
.reason
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if self.lifecycle.load(Ordering::Acquire) & ROUTE_HEALTH_POISONED == 0 {
*stored = Some(reason);
self.lifecycle
.fetch_or(ROUTE_HEALTH_POISONED, Ordering::Release);
}
}
pub(crate) fn ensure_usable(&self) -> Result<(), String> {
let state = self.lifecycle.load(Ordering::Acquire);
if state & ROUTE_HEALTH_RETIRED != 0 {
return Err("reservation is retired and its mappings are no longer launchable".into());
}
if state & ROUTE_HEALTH_RETIRING != 0 {
return Err("reservation teardown is retiring its mappings".into());
}
if state & ROUTE_HEALTH_POISONED == 0 {
return Ok(());
}
let reason = self
.reason
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.clone()
.unwrap_or_else(|| "reservation health was invalidated".to_string());
Err(reason)
}
fn validate_identity(
&self,
executor: ExecutorInstanceId,
device_ordinal: u32,
reservation_generation: u64,
) -> Result<RouteReservationIdentity, String> {
let Some(identity) = self.identity else {
return Err(
"reservation health has no provider/executor/artifact/device/generation identity; \
rebuild the executor"
.to_string(),
);
};
if identity.executor != executor
|| identity.device_ordinal != device_ordinal
|| identity.reservation_generation != reservation_generation
{
return Err(format!(
"reservation generation {} belongs to executor {} on CUDA:{}, not executor {} on \
CUDA:{} generation {}",
identity.reservation_generation,
identity.executor.get(),
identity.device_ordinal,
executor.get(),
device_ordinal,
reservation_generation
));
}
Ok(identity)
}
pub(crate) fn validate_artifact_scope(
&self,
provider: ExecutorArtifactProviderId,
executor: ExecutorInstanceId,
artifact_generation: ExecutorArtifactGeneration,
device_ordinal: u32,
) -> Result<(), String> {
let identity = self.identity.ok_or_else(|| {
"reservation health has no provider/executor/artifact/device/generation identity; \
rebuild the executor"
.to_string()
})?;
if identity.provider != provider
|| identity.executor != executor
|| identity.artifact_generation != artifact_generation
|| identity.device_ordinal != device_ordinal
{
return Err(format!(
"reservation belongs to provider {} executor {} artifact generation {} on CUDA:{}, \
not provider {} executor {} artifact generation {} on CUDA:{}",
identity.provider.get(),
identity.executor.get(),
identity.artifact_generation.get(),
identity.device_ordinal,
provider.get(),
executor.get(),
artifact_generation.get(),
device_ordinal,
));
}
Ok(())
}
pub(crate) fn acquire_use(
self: &Arc<Self>,
executor: ExecutorInstanceId,
device_ordinal: u32,
reservation_generation: u64,
) -> Result<RouteReservationUseGuard, String> {
let identity = self.validate_identity(executor, device_ordinal, reservation_generation)?;
self.acquire_use_identity(identity)
}
fn acquire_use_identity(
self: &Arc<Self>,
identity: RouteReservationIdentity,
) -> Result<RouteReservationUseGuard, String> {
if self.identity != Some(identity) {
return Err(
"baked route-reservation identity no longer matches its health owner".into(),
);
}
loop {
let state = self.lifecycle.load(Ordering::Acquire);
if state & ROUTE_HEALTH_RETIRED != 0 {
return Err(format!(
"reservation generation {} is retired; its mappings cannot be used",
identity.reservation_generation
));
}
if state & ROUTE_HEALTH_RETIRING != 0 {
return Err(format!(
"reservation generation {} is retiring; no new launch may begin",
identity.reservation_generation
));
}
if state & ROUTE_HEALTH_POISONED != 0 {
return Err(format!(
"reservation generation {} is poisoned: {}",
identity.reservation_generation,
self.ensure_usable().unwrap_err()
));
}
if state & ROUTE_HEALTH_TRANSITIONING != 0 {
return Err(format!(
"reservation generation {} is in an atomic group transition; retry after the \
request boundary completes",
identity.reservation_generation
));
}
let next = state.checked_add(ROUTE_HEALTH_USE).ok_or_else(|| {
format!(
"reservation generation {} use counter overflowed",
identity.reservation_generation
)
})?;
if self
.lifecycle
.compare_exchange_weak(state, next, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
return Ok(RouteReservationUseGuard {
health: Arc::clone(self),
});
}
}
}
pub(crate) fn begin_transition(
self: &Arc<Self>,
) -> Result<RouteReservationTransitionGuard, String> {
match self.lifecycle.compare_exchange(
0,
ROUTE_HEALTH_TRANSITIONING,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => Ok(RouteReservationTransitionGuard {
health: Arc::clone(self),
completed: false,
}),
Err(state) if state & ROUTE_HEALTH_POISONED != 0 => self
.ensure_usable()
.map(|_| unreachable!("poisoned reservation unexpectedly reported usable")),
Err(state) if state & ROUTE_HEALTH_TRANSITIONING != 0 => {
Err("another route-bank group transition is already active".to_string())
}
Err(state) if state & ROUTE_HEALTH_RETIRED != 0 => {
Err("route-bank reservation is retired".to_string())
}
Err(state) if state & ROUTE_HEALTH_RETIRING != 0 => {
Err("route-bank reservation teardown is in progress".to_string())
}
Err(state) => Err(format!(
"{} reservation-backed dispatch/replay lease(s) are still active",
state / ROUTE_HEALTH_USE
)),
}
}
pub(crate) fn begin_retirement(&self) -> RouteReservationRetirementStart {
loop {
let state = self.lifecycle.load(Ordering::Acquire);
if state & ROUTE_HEALTH_RETIRED != 0 {
return RouteReservationRetirementStart::Retired;
}
if state & ROUTE_HEALTH_RETIRING != 0 {
return RouteReservationRetirementStart::AlreadyRetiring;
}
if self
.lifecycle
.compare_exchange_weak(
state,
state | ROUTE_HEALTH_RETIRING,
Ordering::AcqRel,
Ordering::Acquire,
)
.is_err()
{
continue;
}
self.retirement_counters
.retirements_started
.fetch_add(1, Ordering::AcqRel);
return RouteReservationRetirementStart::Started;
}
}
pub(crate) fn install_retirement_cleanup(
&self,
cleanup: Box<dyn RouteReservationRetirementCleanup>,
) {
let state = self.lifecycle.load(Ordering::Acquire);
if state / ROUTE_HEALTH_USE != 0 || state & ROUTE_HEALTH_TRANSITIONING != 0 {
self.retirement_counters
.deferred_cleanups
.fetch_add(1, Ordering::AcqRel);
}
let mut slot = self
.retirement_cleanup
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if slot.is_some() {
self.mark_unusable(
"reservation retirement attempted to install cleanup ownership twice".to_string(),
);
return;
}
*slot = Some(cleanup);
drop(slot);
self.try_schedule_retirement_cleanup();
}
fn try_schedule_retirement_cleanup(&self) {
let state = self.lifecycle.load(Ordering::Acquire);
if state & ROUTE_HEALTH_RETIRING == 0
|| state & ROUTE_HEALTH_RETIRED != 0
|| state / ROUTE_HEALTH_USE != 0
|| state & ROUTE_HEALTH_TRANSITIONING != 0
{
return;
}
let cleanup = {
let mut slot = self
.retirement_cleanup
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let state = self.lifecycle.load(Ordering::Acquire);
if state & ROUTE_HEALTH_RETIRING == 0
|| state & ROUTE_HEALTH_RETIRED != 0
|| state / ROUTE_HEALTH_USE != 0
|| state & ROUTE_HEALTH_TRANSITIONING != 0
{
return;
}
slot.take()
};
let Some(cleanup) = cleanup else {
return;
};
self.retirement_counters
.cleanups_scheduled
.fetch_add(1, Ordering::AcqRel);
cleanup.schedule();
}
pub(crate) fn complete_retirement(&self) {
let poisoned = self.lifecycle.load(Ordering::Acquire) & ROUTE_HEALTH_POISONED;
self.lifecycle
.store(poisoned | ROUTE_HEALTH_RETIRED, Ordering::Release);
self.retirement_counters.record_cleanup_executed();
}
pub(crate) fn retirement_status(&self) -> (bool, bool) {
let state = self.lifecycle.load(Ordering::Acquire);
(
state & ROUTE_HEALTH_RETIRING != 0,
state & ROUTE_HEALTH_RETIRED != 0,
)
}
}
#[derive(Clone)]
pub struct RouteReservationAuthorities {
pub catalogs: HashMap<ValueId, onnx_runtime_loader::WeightRegionCatalog>,
pub allocators: HashMap<ValueId, Arc<crate::vmm_allocator::CudaVmmAllocator>>,
pub device_pool: Arc<onnx_runtime_cuda_memory::virtual_memory::PhysicalHandlePool>,
pub host_pool: Arc<onnx_runtime_cuda_memory::virtual_memory::PhysicalHandlePool>,
pub groups: Vec<onnx_runtime_ep_api::ExpertWeightGroup>,
pub(crate) health: Arc<RouteReservationHealth>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RouteBankReservationReject {
NoPhysicalAdmission,
NoDeferredReleaseQueue,
NoBanks,
BlockQuantizedMoeUnavailable {
node: NodeId,
},
MissingFinalizedMember {
node: NodeId,
value: ValueId,
},
UnexpectedFinalizedMember {
node: NodeId,
value: ValueId,
},
DuplicateValue {
value: ValueId,
},
OverlappingExternalRange {
first: ValueId,
second: ValueId,
start: usize,
end: usize,
},
NonPageable {
value: ValueId,
},
IdentityMismatch {
value: ValueId,
reason: String,
},
InconsistentExpertCount {
node: NodeId,
},
InvalidDeviceOrdinal {
ordinal: i32,
},
DeviceMismatch {
expected: onnx_runtime_memory_governor::DeviceKey,
actual: onnx_runtime_memory_governor::DeviceKey,
},
GranularityMismatch {
device: usize,
host: usize,
},
UnalignedExpertRange {
value: ValueId,
expert: usize,
offset: usize,
len: usize,
granularity: usize,
},
ExistingReservationMismatch {
reason: String,
},
Materialization {
value: ValueId,
reason: String,
},
Reservation(String),
}
impl std::fmt::Display for RouteBankReservationReject {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::NoPhysicalAdmission => {
write!(formatter, "weight residency has no VMM physical admission")
}
Self::NoDeferredReleaseQueue => {
write!(
formatter,
"weight residency has no deferred reservation queue"
)
}
Self::NoBanks => write!(formatter, "no finalized routed expert banks"),
Self::BlockQuantizedMoeUnavailable { node } => write!(
formatter,
"BlockQuantizedMoE bank {node:?} has no finalized production telemetry/catalog contract"
),
Self::MissingFinalizedMember { node, value } => {
write!(
formatter,
"bank {node:?} is missing finalized member {value:?}"
)
}
Self::UnexpectedFinalizedMember { node, value } => write!(
formatter,
"bank {node:?} contains unexpected finalized member {value:?}"
),
Self::DuplicateValue { value } => {
write!(
formatter,
"bank value {value:?} appears in more than one group"
)
}
Self::OverlappingExternalRange {
first,
second,
start,
end,
} => write!(
formatter,
"bank values {first:?} and {second:?} overlap external bytes {start}..{end}"
),
Self::NonPageable { value } => {
write!(formatter, "bank value {value:?} is not expert-pageable")
}
Self::IdentityMismatch { value, reason } => {
write!(
formatter,
"bank value {value:?} identity mismatch: {reason}"
)
}
Self::InconsistentExpertCount { node } => {
write!(formatter, "bank {node:?} members disagree on expert count")
}
Self::InvalidDeviceOrdinal { ordinal } => {
write!(formatter, "invalid CUDA device ordinal {ordinal}")
}
Self::DeviceMismatch { expected, actual } => write!(
formatter,
"route-bank reservation device mismatch: expected {expected:?}, got {actual:?}"
),
Self::GranularityMismatch { device, host } => write!(
formatter,
"device granularity {device} differs from host-NUMA granularity {host}"
),
Self::UnalignedExpertRange {
value,
expert,
offset,
len,
granularity,
} => write!(
formatter,
"bank value {value:?} expert {expert} range {offset}..{} is not aligned to VMM granularity {granularity}",
offset.saturating_add(*len)
),
Self::ExistingReservationMismatch { reason } => {
write!(formatter, "existing route reservation mismatch: {reason}")
}
Self::Materialization { value, reason } => {
write!(
formatter,
"bank value {value:?} materialization failed: {reason}"
)
}
Self::Reservation(reason) => write!(formatter, "route reservation failed: {reason}"),
}
}
}
fn validate_existing_route_reservations(
set: &RouteReservationSet,
banks: &[FinalizedExpertBank],
device_ordinal: i32,
) -> Result<(), String> {
let groups: Vec<_> = banks.iter().map(|bank| bank.group.clone()).collect();
if set.groups != groups {
return Err("expert-bank topology differs from the installed reservation set".into());
}
let member_count = banks.iter().map(|bank| bank.members.len()).sum::<usize>();
if set.by_key.len() != member_count {
return Err(format!(
"installed reservation has {} values, finalized banks have {member_count}",
set.by_key.len()
));
}
let expected_device = onnx_runtime_memory_governor::DeviceKey::device(device_ordinal as u32);
for member in banks.iter().flat_map(|bank| &bank.members) {
let Some(existing) = set.by_key.get(&(member.value.0 as u64)) else {
return Err(format!(
"finalized value {:?} has no installed reservation",
member.value
));
};
if existing.identity.value != member.value
|| existing.identity.external_path != member.external_path
|| existing.identity.weight.boundary != member.weight.boundary
|| existing.identity.weight.dtype != member.weight.dtype
|| existing.identity.weight.shape != member.weight.shape
|| existing.identity.weight.regions != member.weight.regions
|| existing.identity.catalog != member.catalog
|| existing.len != member.catalog.tensor_len()
{
return Err(format!(
"finalized value {:?} differs from its installed property identity",
member.value
));
}
if existing.allocator.device_key() != expected_device {
return Err(format!(
"finalized value {:?} reservation is on {:?}, expected {expected_device:?}",
member.value,
existing.allocator.device_key()
));
}
}
Ok(())
}
struct ResidencyInner {
policy: WeightResidencyPolicy,
pages: HashMap<u64, Arc<CudaWeightPage>>,
mapped_allowance: Option<onnx_runtime_memory_governor::MappedAllowance>,
lease: Option<onnx_runtime_memory_governor::MemoryLease>,
admission_no_progress: u64,
slots: HashMap<u64, StableWeightSlot>,
cold_pages: HashMap<u64, Arc<CudaWeightPage>>,
pinned: HashSet<u64>,
pinned_bytes: u64,
pin_hit_step: HashMap<u64, u64>,
pin_granule_hashes: HashMap<u64, Vec<u64>>,
pending_prefetch: Option<PendingPrefetch>,
prefetch_stats: PrefetchStats,
}
#[derive(Debug, Default, Clone, Copy)]
struct PrefetchStats {
issued: u64,
issued_bytes: u64,
promoted: u64,
promote_wait_ns: u64,
declined_budget: u64,
declined_busy: u64,
declined_unsupported: u64,
declined_resident: u64,
declined_pool_capacity: u64,
}
struct PendingPrefetch {
key: u64,
page: Arc<CudaWeightPage>,
fence_id: u64,
staging: PooledStaging,
bytes: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum WeightEvictionPolicy {
Lru,
StableResident,
}
#[cfg(test)]
#[derive(Debug, PartialEq, Eq)]
struct WeightPolicyAccess {
hit: bool,
admitted: bool,
evicted: Vec<u64>,
}
#[derive(Debug)]
struct WeightResidencyPolicy {
budget: u64,
resident_bytes: u64,
peak_resident_bytes: u64,
page_ins: u64,
hits: u64,
evictions: u64,
order: Vec<u64>,
bytes_by_key: HashMap<u64, u64>,
}
impl WeightResidencyPolicy {
fn new(budget: u64) -> Self {
Self {
budget,
resident_bytes: 0,
peak_resident_bytes: 0,
page_ins: 0,
hits: 0,
evictions: 0,
order: Vec::new(),
bytes_by_key: HashMap::new(),
}
}
#[cfg(test)]
fn access(
&mut self,
key: u64,
bytes: u64,
eviction: WeightEvictionPolicy,
) -> WeightPolicyAccess {
if self.bytes_by_key.contains_key(&key) {
self.record_hit(key);
return WeightPolicyAccess {
hit: true,
admitted: false,
evicted: Vec::new(),
};
}
if eviction == WeightEvictionPolicy::StableResident && !self.can_fit(bytes) {
self.record_page_in();
return WeightPolicyAccess {
hit: false,
admitted: false,
evicted: Vec::new(),
};
}
let evicted = self.evict_to_fit(bytes, eviction, |_| true);
self.insert_page(key, bytes);
WeightPolicyAccess {
hit: false,
admitted: true,
evicted,
}
}
#[cfg(test)]
fn access_byte_aware(&mut self, key: u64, bytes: u64) -> WeightPolicyAccess {
if self.bytes_by_key.contains_key(&key) {
self.record_hit(key);
return WeightPolicyAccess {
hit: true,
admitted: false,
evicted: Vec::new(),
};
}
if !self.can_fit(bytes) {
let reclaimable: u64 = self
.bytes_by_key
.values()
.filter(|&&resident| resident < bytes)
.sum();
let free = self.budget.saturating_sub(self.resident_bytes);
if free.saturating_add(reclaimable) < bytes {
self.record_page_in();
return WeightPolicyAccess {
hit: false,
admitted: false,
evicted: Vec::new(),
};
}
}
let mut evicted = Vec::new();
while !self.can_fit(bytes) {
let Some((&smallest_key, _)) = self
.bytes_by_key
.iter()
.filter(|&(_, &resident)| resident < bytes)
.min_by_key(|&(_, &resident)| resident)
else {
break;
};
self.remove_page(smallest_key);
evicted.push(smallest_key);
}
self.insert_page(key, bytes);
WeightPolicyAccess {
hit: false,
admitted: true,
evicted,
}
}
fn can_fit(&self, incoming: u64) -> bool {
self.resident_bytes.saturating_add(incoming) <= self.budget
}
fn touch(&mut self, key: u64) {
if let Some(position) = self.order.iter().position(|&k| k == key) {
let key = self.order.remove(position);
self.order.push(key);
}
}
fn record_hit(&mut self, key: u64) {
self.touch(key);
self.hits += 1;
}
fn insert_page(&mut self, key: u64, bytes: u64) {
self.bytes_by_key.insert(key, bytes);
self.order.push(key);
self.resident_bytes += bytes;
self.peak_resident_bytes = self.peak_resident_bytes.max(self.resident_bytes);
self.record_page_in();
}
fn remove_page(&mut self, key: u64) -> Option<u64> {
if let Some(position) = self.order.iter().position(|&candidate| candidate == key) {
self.order.remove(position);
}
let bytes = self.bytes_by_key.remove(&key)?;
self.resident_bytes = self.resident_bytes.saturating_sub(bytes);
self.evictions = self.evictions.saturating_add(1);
Some(bytes)
}
fn record_page_in(&mut self) {
self.page_ins += 1;
}
fn evict_to_fit<F>(
&mut self,
incoming: u64,
eviction: WeightEvictionPolicy,
mut evictable: F,
) -> Vec<u64>
where
F: FnMut(u64) -> bool,
{
let mut evicted = Vec::new();
while self.resident_bytes.saturating_add(incoming) > self.budget && !self.order.is_empty() {
let Some(index) = self.next_evictable_index(eviction, &mut evictable) else {
break;
};
let key = self.order.remove(index);
if let Some(bytes) = self.bytes_by_key.remove(&key) {
self.resident_bytes = self.resident_bytes.saturating_sub(bytes);
self.evictions += 1;
evicted.push(key);
}
}
evicted
}
fn next_evictable_index<F>(
&self,
eviction: WeightEvictionPolicy,
evictable: &mut F,
) -> Option<usize>
where
F: FnMut(u64) -> bool,
{
match eviction {
WeightEvictionPolicy::Lru | WeightEvictionPolicy::StableResident => {
self.order.iter().position(|&key| evictable(key))
}
}
}
}
impl From<onnx_runtime_ep_api::EvictionClass> for WeightEvictionPolicy {
fn from(class: onnx_runtime_ep_api::EvictionClass) -> Self {
match class {
onnx_runtime_ep_api::EvictionClass::Lru => WeightEvictionPolicy::Lru,
onnx_runtime_ep_api::EvictionClass::StableResident => {
WeightEvictionPolicy::StableResident
}
}
}
}
#[derive(Clone, Copy, Debug)]
struct CudaHotSetResidencyPolicy {
scan_resistant_dense: bool,
static_pin_keys: Option<&'static HashSet<u64>>,
static_pin_config: Option<(u64, u64)>,
}
impl CudaHotSetResidencyPolicy {
fn from_env(scan_resistant_dense: bool) -> Self {
Self {
scan_resistant_dense,
static_pin_keys: static_pin_keys(),
static_pin_config: static_pin_config(),
}
}
}
impl onnx_runtime_ep_api::ResidencyPolicy for CudaHotSetResidencyPolicy {
fn name(&self) -> &'static str {
"cuda_hot_set"
}
fn decide(
&self,
input: &onnx_runtime_ep_api::ResidencyPolicyInput<'_>,
) -> onnx_runtime_ep_api::ResidencyDecision {
onnx_runtime_ep_api::WholeBankResidentPolicy.decide(input)
}
fn eviction_class(&self, boundary: LazyWeightBoundary) -> onnx_runtime_ep_api::EvictionClass {
if self.scan_resistant_dense
&& matches!(
boundary,
LazyWeightBoundary::MatMul | LazyWeightBoundary::MatMulNBits
)
{
onnx_runtime_ep_api::EvictionClass::StableResident
} else {
onnx_runtime_ep_api::EvictionClass::Lru
}
}
fn should_pin(&self, input: &onnx_runtime_ep_api::AdmissionPolicyInput) -> bool {
if let Some(keys) = self.static_pin_keys {
keys.contains(&input.key) && !input.already_pinned
} else if let Some((threshold, budget)) = self.static_pin_config {
input.len_bytes >= threshold
&& !input.already_pinned
&& input.pinned_bytes_used.saturating_add(input.len_bytes) <= budget
} else {
false
}
}
}
fn eviction_for_boundary(
scan_resistant_dense: bool,
boundary: LazyWeightBoundary,
) -> WeightEvictionPolicy {
use onnx_runtime_ep_api::ResidencyPolicy as _;
CudaHotSetResidencyPolicy::from_env(scan_resistant_dense)
.eviction_class(boundary)
.into()
}
impl CudaWeightResidency {
pub fn new(runtime: Arc<CudaRuntime>, budget_bytes: u64) -> Self {
replace_global_budget(0, budget_bytes);
runtime.set_weights_may_be_paged();
Self {
runtime: Arc::clone(&runtime),
queue: None,
scan_resistant_dense: false,
byte_aware: false,
evict_order_probe: EvictOrderProbe::Lru,
zero_copy_hybrid: false,
host_registry: Mutex::new(HostMapRegistry::new()),
physical: OnceLock::new(),
context_scope: OnceLock::new(),
context_terminated: AtomicBool::new(false),
staging_pool: PinnedStagingPool::new(Arc::clone(&runtime)),
route_reservations: Mutex::new(HashMap::new()),
route_retirement_counters: Arc::new(RouteReservationRetirementCounters::default()),
inner: Mutex::new(ResidencyInner {
policy: WeightResidencyPolicy::new(budget_bytes),
lease: None,
pages: HashMap::new(),
mapped_allowance: None,
admission_no_progress: 0,
slots: HashMap::new(),
cold_pages: HashMap::new(),
pinned: HashSet::new(),
pinned_bytes: 0,
pin_hit_step: HashMap::new(),
pin_granule_hashes: HashMap::new(),
pending_prefetch: None,
prefetch_stats: PrefetchStats::default(),
}),
routed_guards_active: AtomicU64::new(0),
prefill_double_buffer_enabled: prefill_double_buffer_enabled(),
prefill_pipeline: Mutex::new(None),
}
}
pub fn new_leased(
runtime: Arc<CudaRuntime>,
budget_bytes: u64,
governor: &dyn onnx_runtime_memory_governor::MemoryGovernor,
tier: onnx_runtime_memory_governor::Tier,
holder: onnx_runtime_memory_governor::HolderId,
) -> Result<Self, onnx_runtime_memory_governor::MemoryError> {
let lease = governor.reserve(
tier,
budget_bytes,
onnx_runtime_memory_governor::MemoryRole::Weights,
holder,
)?;
replace_global_budget(0, lease.bytes());
Ok(Self {
runtime: Arc::clone(&runtime),
queue: None,
scan_resistant_dense: false,
byte_aware: false,
evict_order_probe: EvictOrderProbe::Lru,
zero_copy_hybrid: false,
host_registry: Mutex::new(HostMapRegistry::new()),
physical: OnceLock::new(),
context_scope: OnceLock::new(),
context_terminated: AtomicBool::new(false),
staging_pool: PinnedStagingPool::new(Arc::clone(&runtime)),
route_reservations: Mutex::new(HashMap::new()),
route_retirement_counters: Arc::new(RouteReservationRetirementCounters::default()),
inner: Mutex::new(ResidencyInner {
policy: WeightResidencyPolicy::new(lease.bytes()),
lease: Some(lease),
pages: HashMap::new(),
mapped_allowance: None,
admission_no_progress: 0,
slots: HashMap::new(),
cold_pages: HashMap::new(),
pinned: HashSet::new(),
pinned_bytes: 0,
pin_hit_step: HashMap::new(),
pin_granule_hashes: HashMap::new(),
pending_prefetch: None,
prefetch_stats: PrefetchStats::default(),
}),
routed_guards_active: AtomicU64::new(0),
prefill_double_buffer_enabled: prefill_double_buffer_enabled(),
prefill_pipeline: Mutex::new(None),
})
}
pub fn budget(&self) -> (u64, bool) {
let inner = self.inner.lock().expect("residency lock poisoned");
(
inner.policy.budget,
inner.lease.is_some() || inner.mapped_allowance.is_some(),
)
}
pub fn set_ungoverned_budget(
&self,
budget_bytes: u64,
) -> Result<u64, onnx_runtime_memory_governor::MemoryError> {
let mut inner = self.inner.lock().expect("residency lock poisoned");
if inner.lease.is_some() || inner.mapped_allowance.is_some() {
return Ok(inner.policy.budget);
}
let old = inner.policy.budget;
inner.policy.budget = budget_bytes;
replace_global_budget(old, budget_bytes);
Ok(inner.policy.budget)
}
pub fn adopt_governed_budget(
&self,
governor: &dyn onnx_runtime_memory_governor::MemoryGovernor,
tier: onnx_runtime_memory_governor::Tier,
holder: onnx_runtime_memory_governor::HolderId,
) -> Result<u64, onnx_runtime_memory_governor::MemoryError> {
if let Some(physical) = self.physical.get() {
if physical.governor.authority_id() != governor.authority_id() {
return Err(onnx_runtime_memory_governor::MemoryError::InvalidRequest {
tier: tier.name(),
requested: 0,
reason: "VMM weight residency was built with a different physical-memory authority",
});
}
let mut inner = self.lock();
if inner.mapped_allowance.is_none() {
inner.mapped_allowance = Some(governor.reserve_mapped_allowance(
tier,
inner.policy.budget,
onnx_runtime_memory_governor::MemoryRole::Weights,
holder,
)?);
}
return Ok(inner.policy.budget);
}
let requested = {
let inner = self.inner.lock().expect("residency lock poisoned");
if inner.lease.is_some() {
return Ok(inner.policy.budget);
}
inner.policy.budget
};
let lease = governor.reserve(
tier,
requested,
onnx_runtime_memory_governor::MemoryRole::Weights,
holder,
)?;
let granted = lease.bytes();
let mut inner = self.inner.lock().expect("residency lock poisoned");
let old = inner.policy.budget;
inner.policy.budget = granted;
inner.lease = Some(lease);
replace_global_budget(old, granted);
Ok(granted)
}
pub fn resize_safe_point(&self, device_count: usize) -> onnx_runtime_ep_api::ResizeSafePoint {
let capturing = self.runtime.is_capturing().unwrap_or(true);
let pending_deferred_releases = self
.queue
.as_ref()
.map_or(0, |queue| queue.pending() as u64);
let admission_in_flight = !in_flight_fill_quarantine()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.is_empty();
onnx_runtime_ep_api::ResizeSafePoint {
capturing,
pending_deferred_releases,
admission_in_flight,
multi_device: device_count > 1,
routed_guards_active: self.routed_guards_active.load(Ordering::SeqCst),
}
}
pub fn install_route_bank_reservations(
&self,
provider: ExecutorArtifactProviderId,
executor: ExecutorInstanceId,
artifact_generation: ExecutorArtifactGeneration,
banks: &[FinalizedExpertBank],
device_ordinal: i32,
) -> Result<RouteReservationAuthorities, RouteBankReservationReject> {
let physical = self
.physical
.get()
.ok_or(RouteBankReservationReject::NoPhysicalAdmission)?;
let queue = self
.queue
.as_ref()
.ok_or(RouteBankReservationReject::NoDeferredReleaseQueue)?;
if banks.is_empty() {
return Err(RouteBankReservationReject::NoBanks);
}
if device_ordinal < 0 {
return Err(RouteBankReservationReject::InvalidDeviceOrdinal {
ordinal: device_ordinal,
});
}
let expected_device =
onnx_runtime_memory_governor::DeviceKey::device(device_ordinal as u32);
let actual_device = physical.allocator.device_key();
if actual_device != expected_device {
return Err(RouteBankReservationReject::DeviceMismatch {
expected: expected_device,
actual: actual_device,
});
}
let mut installed = self
.route_reservations
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if let Some(existing) = installed.get(&executor) {
existing
.health
.validate_artifact_scope(
provider,
executor,
artifact_generation,
device_ordinal as u32,
)
.map_err(
|reason| RouteBankReservationReject::ExistingReservationMismatch { reason },
)?;
validate_existing_route_reservations(existing, banks, device_ordinal).map_err(
|reason| RouteBankReservationReject::ExistingReservationMismatch { reason },
)?;
return Ok(Self::authorities_from_set(existing));
}
let host_numa = onnx_runtime_cuda_memory::capability::host_numa_capability(device_ordinal)
.map_err(|error| RouteBankReservationReject::Reservation(error.to_string()))?;
let device_location = onnx_runtime_cuda_memory::virtual_memory::PhysicalLocation::Device {
ordinal: device_ordinal,
};
let host_location = onnx_runtime_cuda_memory::virtual_memory::PhysicalLocation::HostNuma {
node: host_numa.host_numa_id,
};
let device_granularity =
onnx_runtime_cuda_memory::virtual_memory::allocation_granularity_for_location(
device_location,
);
let host_granularity =
onnx_runtime_cuda_memory::virtual_memory::allocation_granularity_for_location(
host_location,
);
if device_granularity != host_granularity {
return Err(RouteBankReservationReject::GranularityMismatch {
device: device_granularity,
host: host_granularity,
});
}
let granularity = device_granularity;
let mut seen = HashSet::new();
let mut external_ranges = Vec::<(PathBuf, usize, usize, ValueId)>::new();
for bank in banks {
if bank.group.boundary == LazyWeightBoundary::BlockQuantizedMoe {
return Err(RouteBankReservationReject::BlockQuantizedMoeUnavailable {
node: bank.group.node,
});
}
if bank.group.boundary != LazyWeightBoundary::QMoe {
return Err(RouteBankReservationReject::IdentityMismatch {
value: bank
.group
.members
.first()
.copied()
.unwrap_or(ValueId(u32::MAX)),
reason: "only the production QMoE boundary has routed-bank telemetry".into(),
});
}
let members: HashMap<_, _> = bank
.members
.iter()
.map(|member| (member.value, member))
.collect();
if members.len() != bank.members.len() {
let value = bank
.members
.iter()
.find(|candidate| {
bank.members
.iter()
.filter(|other| other.value == candidate.value)
.count()
> 1
})
.map(|member| member.value)
.unwrap_or(ValueId(u32::MAX));
return Err(RouteBankReservationReject::DuplicateValue { value });
}
for value in &bank.group.members {
if !members.contains_key(value) {
return Err(RouteBankReservationReject::MissingFinalizedMember {
node: bank.group.node,
value: *value,
});
}
}
for member in &bank.members {
if !bank.group.contains(member.value) {
return Err(RouteBankReservationReject::UnexpectedFinalizedMember {
node: bank.group.node,
value: member.value,
});
}
if !seen.insert(member.value) {
return Err(RouteBankReservationReject::DuplicateValue {
value: member.value,
});
}
if !member.catalog.is_pageable() {
return Err(RouteBankReservationReject::NonPageable {
value: member.value,
});
}
if member.catalog.path() != Some(member.external_path.as_path()) {
return Err(RouteBankReservationReject::IdentityMismatch {
value: member.value,
reason:
"pageable bank catalog path differs from the finalized external property"
.into(),
});
}
let region = match member.weight.regions.as_slice() {
[region] => region,
_ => {
return Err(RouteBankReservationReject::IdentityMismatch {
value: member.value,
reason: "coarse bank members require one contiguous mmap identity"
.into(),
});
}
};
let expected_shape = vec![
member.catalog.layout().experts,
member.catalog.layout().rows_per_expert,
member.catalog.layout().storage_elements_per_row,
];
let endpoint = region.offset.checked_add(region.len).ok_or_else(|| {
RouteBankReservationReject::IdentityMismatch {
value: member.value,
reason: "mmap range endpoint overflow".into(),
}
})?;
for (path, start, end, value) in &external_ranges {
if path == &member.external_path && region.offset < *end && *start < endpoint {
return Err(RouteBankReservationReject::OverlappingExternalRange {
first: *value,
second: member.value,
start: region.offset.max(*start),
end: endpoint.min(*end),
});
}
}
if endpoint > isize::MAX as usize
|| region.mapping_id == 0
|| member.weight.boundary != bank.group.boundary
|| member.weight.dtype != member.catalog.dtype()
|| member.weight.shape != expected_shape
|| region.offset != member.catalog.tensor_offset()
|| region.len != member.catalog.tensor_len()
|| member.weight.region_bytes_len() != member.catalog.tensor_len()
|| region.len % granularity != 0
{
return Err(RouteBankReservationReject::IdentityMismatch {
value: member.value,
reason:
"boundary/format/shape/mmap mapping and catalog byte properties disagree"
.into(),
});
}
external_ranges.push((
member.external_path.clone(),
region.offset,
endpoint,
member.value,
));
let mut prior_end = 0usize;
for expert in 0..member.catalog.layout().experts {
let range = member.catalog.relative_range(expert).ok_or_else(|| {
RouteBankReservationReject::IdentityMismatch {
value: member.value,
reason: format!("expert {expert} has no exact relative byte range"),
}
})?;
let len = range.end.saturating_sub(range.start);
if range.start != prior_end
|| range.end > member.catalog.tensor_len()
|| range.start % granularity != 0
|| len == 0
|| len % granularity != 0
{
return Err(RouteBankReservationReject::UnalignedExpertRange {
value: member.value,
expert,
offset: range.start,
len,
granularity,
});
}
prior_end = range.end;
}
if prior_end != member.catalog.tensor_len() {
return Err(RouteBankReservationReject::IdentityMismatch {
value: member.value,
reason: "expert ranges do not exactly partition the tensor byte range"
.into(),
});
}
}
let mut expert_counts = bank
.members
.iter()
.map(|member| member.catalog.layout().experts);
if let Some(first) = expert_counts.next()
&& expert_counts.any(|experts| experts != first)
{
return Err(RouteBankReservationReject::InconsistentExpertCount {
node: bank.group.node,
});
}
}
let materialized = banks
.iter()
.flat_map(|bank| &bank.members)
.map(|member| {
let resident = member.weight.materialize().map_err(|error| {
RouteBankReservationReject::Materialization {
value: member.value,
reason: error.to_string(),
}
})?;
if resident.dtype != member.weight.dtype
|| resident.shape != member.weight.shape
|| resident.bytes().len() != member.catalog.tensor_len()
{
return Err(RouteBankReservationReject::Materialization {
value: member.value,
reason: "resolved bytes differ from finalized dtype/shape/catalog".into(),
});
}
Ok((member.clone(), resident))
})
.collect::<Result<Vec<_>, RouteBankReservationReject>>()?;
let context = self.runtime.cuda_context();
let role = onnx_runtime_memory_governor::MemoryRole::Weights;
let holder_seed = executor.get().wrapping_mul(2);
let device_pool =
onnx_runtime_cuda_memory::virtual_memory::PhysicalHandlePool::new_isolated_at_location(
Arc::clone(&context),
device_ordinal,
device_location,
0,
physical.governor.as_ref(),
onnx_runtime_memory_governor::HolderId::new(holder_seed.wrapping_add(0x1810)),
role,
)
.map_err(|error| RouteBankReservationReject::Reservation(error.to_string()))?;
let host_pool =
onnx_runtime_cuda_memory::virtual_memory::PhysicalHandlePool::new_isolated_at_location(
Arc::clone(&context),
device_ordinal,
host_location,
0,
physical.governor.as_ref(),
onnx_runtime_memory_governor::HolderId::new(holder_seed.wrapping_add(0x1811)),
role,
)
.map_err(|error| RouteBankReservationReject::Reservation(error.to_string()))?;
let reservation_queue: Arc<
dyn onnx_runtime_cuda_memory::virtual_memory::DeferredReservationQueue,
> = Arc::clone(queue)
as Arc<dyn onnx_runtime_cuda_memory::virtual_memory::DeferredReservationQueue>;
let health = RouteReservationHealth::new_scoped(
provider,
executor,
artifact_generation,
device_ordinal as u32,
Arc::clone(&self.route_retirement_counters),
)
.ok_or_else(|| {
RouteBankReservationReject::Reservation(
"route-reservation generation identity exhausted".to_string(),
)
})?;
let mut by_key = HashMap::new();
let mut catalogs = HashMap::new();
let mut allocators = HashMap::new();
for (member, resident) in materialized {
let allocator = Arc::new(
onnx_runtime_cuda_memory::vmm_allocator::CudaVmmAllocator::new_with_physical_pool_and_reservation_queue(
Arc::clone(&context),
expected_device,
device_ordinal,
member.catalog.tensor_len(),
physical.governor.as_ref(),
onnx_runtime_memory_governor::HolderId::new(holder_seed.wrapping_add(0x1810)),
role,
Arc::clone(&device_pool),
Arc::clone(&reservation_queue),
)
.map_err(|error| RouteBankReservationReject::Reservation(error.to_string()))?,
);
let full_range = 0..member.catalog.tensor_len();
let ptr = allocator
.allocate_committed(
member.catalog.tensor_len(),
granularity,
std::slice::from_ref(&full_range),
)
.map_err(|error| RouteBankReservationReject::Reservation(error.to_string()))?;
let mut staging = self
.runtime
.alloc_pinned(member.catalog.tensor_len())
.map_err(|error| RouteBankReservationReject::Reservation(error.to_string()))?;
staging.as_mut_slice().copy_from_slice(resident.bytes());
unsafe {
self.runtime
.htod_async(staging.as_slice(), ptr.as_ptr() as CUdeviceptr)
}
.map_err(|error| RouteBankReservationReject::Reservation(error.to_string()))?;
self.runtime
.sync_copy_stream()
.map_err(|error| RouteBankReservationReject::Reservation(error.to_string()))?;
let reservation = Arc::new(RouteWeightReservation {
identity: member.clone(),
allocator: Arc::clone(&allocator),
health: Arc::clone(&health),
ptr: ptr.as_ptr() as CUdeviceptr,
len: member.catalog.tensor_len(),
});
catalogs.insert(member.value, member.catalog.clone());
allocators.insert(member.value, allocator);
by_key.insert(member.value.0 as u64, reservation);
}
let set = Arc::new(RouteReservationSet {
by_key,
catalogs,
allocators,
device_pool,
host_pool,
groups: banks.iter().map(|bank| bank.group.clone()).collect(),
health,
});
installed.insert(executor, Arc::clone(&set));
Ok(Self::authorities_from_set(&set))
}
fn authorities_from_set(set: &RouteReservationSet) -> RouteReservationAuthorities {
RouteReservationAuthorities {
catalogs: set.catalogs.clone(),
allocators: set.allocators.clone(),
device_pool: Arc::clone(&set.device_pool),
host_pool: Arc::clone(&set.host_pool),
groups: set.groups.clone(),
health: Arc::clone(&set.health),
}
}
pub fn route_reservation_authorities(
&self,
executor: ExecutorInstanceId,
) -> Option<RouteReservationAuthorities> {
let installed = self
.route_reservations
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
installed
.get(&executor)
.map(|set| Self::authorities_from_set(set))
}
pub(crate) fn take_route_bank_reservations(
&self,
executor: ExecutorInstanceId,
) -> Option<RouteReservationRetirementResources> {
let removed = self
.route_reservations
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.remove(&executor);
removed.map(|set| RouteReservationRetirementResources {
_allocators: set.allocators.values().cloned().collect(),
_device_pool: Arc::clone(&set.device_pool),
_host_pool: Arc::clone(&set.host_pool),
})
}
pub fn remove_route_bank_reservations(&self, executor: ExecutorInstanceId) -> bool {
self.take_route_bank_reservations(executor).is_some()
}
pub fn route_reservation_count(&self) -> usize {
self.route_reservations
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.len()
}
pub(crate) fn route_reservation_resource_stats(
&self,
executor: ExecutorInstanceId,
) -> Option<(usize, u64)> {
self.route_reservations
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.get(&executor)
.map(|set| {
(
set.allocators.len(),
set.by_key
.values()
.map(|reservation| reservation.len as u64)
.sum(),
)
})
}
pub fn route_reservation_retirement_stats(&self) -> RouteReservationRetirementStats {
self.route_retirement_counters.snapshot()
}
pub(crate) fn route_weight_page(
&self,
executor: ExecutorInstanceId,
key: u64,
weight: &LazyWeight,
device: DeviceId,
) -> Result<Option<PagedWeight>, WeightHandleError> {
let reservation = {
let installed = self
.route_reservations
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
installed
.get(&executor)
.and_then(|set| set.by_key.get(&key))
.cloned()
};
let Some(reservation) = reservation else {
return Ok(None);
};
reservation.health.ensure_usable().map_err(|reason| {
WeightHandleError::DeviceBinding(format!(
"route-bank value key {key} is unavailable because its executor-scoped \
reservation was invalidated after a failed atomic residency transition: \
{reason}; tear down and rebuild the executor"
))
})?;
if weight.boundary != reservation.identity.weight.boundary
|| weight.dtype != reservation.identity.weight.dtype
|| weight.shape != reservation.identity.weight.shape
|| weight.regions != reservation.identity.weight.regions
|| weight.region_bytes_len() != reservation.len
{
return Err(WeightHandleError::DeviceBinding(format!(
"route-bank value key {key} no longer matches its finalized property identity"
)));
}
let reservation_device = reservation.allocator.device_key();
let requested_device = onnx_runtime_memory_governor::DeviceKey::device(device.index);
if device.device_type != DeviceType::Cuda || requested_device != reservation_device {
return Err(WeightHandleError::DeviceBinding(format!(
"route-bank value key {key} requested on {requested_device:?}, reservation is on \
{reservation_device:?}"
)));
}
let ptr = reservation.ptr;
let keep_alive: Arc<dyn std::any::Any + Send + Sync> = reservation;
Ok(Some(PagedWeight::new(
ptr as *const std::ffi::c_void,
device,
weight.region_bytes_len(),
keep_alive,
)))
}
pub fn coarse_route_bank_reservation(
&self,
executor: ExecutorInstanceId,
value: onnx_runtime_ir::ValueId,
) -> Option<Arc<crate::vmm_allocator::CudaVmmAllocator>> {
self.route_reservations
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.get(&executor)?
.allocators
.get(&value)
.cloned()
}
pub(crate) fn route_reservation_health(
&self,
executor: ExecutorInstanceId,
value: ValueId,
) -> Option<Arc<RouteReservationHealth>> {
let installed = self
.route_reservations
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let set = installed.get(&executor)?;
set.allocators
.contains_key(&value)
.then(|| Arc::clone(&set.health))
}
#[allow(clippy::too_many_arguments)]
pub fn apply_coarse_residency_plan(
&self,
plan: &onnx_runtime_ep_api::ResidencyPlan,
catalogs: &std::collections::HashMap<
onnx_runtime_ir::ValueId,
onnx_runtime_loader::WeightRegionCatalog,
>,
allocators: &std::collections::HashMap<
onnx_runtime_ir::ValueId,
Arc<onnx_runtime_cuda_memory::vmm_allocator::CudaVmmAllocator>,
>,
device_pool: &Arc<onnx_runtime_cuda_memory::virtual_memory::PhysicalHandlePool>,
host_pool: &Arc<onnx_runtime_cuda_memory::virtual_memory::PhysicalHandlePool>,
device_count: usize,
device_ordinal: i32,
expert_groups: &[onnx_runtime_ep_api::ExpertWeightGroup],
) -> crate::coarse_residency::BoundaryApplicationOutcome {
crate::coarse_residency::apply_residency_plan_at_boundary(
&self.runtime,
self,
plan,
catalogs,
allocators,
device_pool,
host_pool,
device_count,
device_ordinal,
expert_groups,
)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn apply_resolved_coarse_residency_plan(
&self,
plan: &onnx_runtime_ep_api::ResidencyPlan,
catalogs: &std::collections::HashMap<
onnx_runtime_ir::ValueId,
onnx_runtime_loader::WeightRegionCatalog,
>,
allocators: &std::collections::HashMap<
onnx_runtime_ir::ValueId,
Arc<onnx_runtime_cuda_memory::vmm_allocator::CudaVmmAllocator>,
>,
device_pool: &Arc<onnx_runtime_cuda_memory::virtual_memory::PhysicalHandlePool>,
host_pool: &Arc<onnx_runtime_cuda_memory::virtual_memory::PhysicalHandlePool>,
device_count: usize,
device_ordinal: i32,
expert_groups: &[onnx_runtime_ep_api::ExpertWeightGroup],
) -> crate::coarse_residency::BoundaryApplicationOutcome {
crate::coarse_residency::apply_resolved_residency_plan_at_boundary(
&self.runtime,
self,
plan,
catalogs,
allocators,
device_pool,
host_pool,
device_count,
device_ordinal,
expert_groups,
)
}
pub fn acquire_routed_residency(
self: &Arc<Self>,
requirement: onnx_runtime_ep_api::RoutedResidencyRequirement,
catalog: &onnx_runtime_loader::WeightRegionCatalog,
) -> RoutedResidencyGuard {
let proof = onnx_runtime_ep_api::prove_routed_residency(requirement, catalog);
self.routed_guards_active.fetch_add(1, Ordering::SeqCst);
RoutedResidencyGuard {
proof,
residency: Arc::clone(self),
_reservation_use: None,
}
}
pub(crate) fn acquire_route_reservation_use(
self: &Arc<Self>,
executor: ExecutorInstanceId,
device_ordinal: u32,
requirement: onnx_runtime_ep_api::RoutedResidencyRequirement,
catalog: &onnx_runtime_loader::WeightRegionCatalog,
health: Arc<RouteReservationHealth>,
) -> Result<RoutedResidencyGuard, String> {
let generation = health.generation().ok_or_else(|| {
"route reservation use requires an executor/device/generation identity".to_string()
})?;
let reservation_use = health.acquire_use(executor, device_ordinal, generation)?;
let proof = onnx_runtime_ep_api::prove_routed_residency(requirement, catalog);
self.routed_guards_active.fetch_add(1, Ordering::SeqCst);
Ok(RoutedResidencyGuard {
proof,
residency: Arc::clone(self),
_reservation_use: Some(reservation_use),
})
}
pub fn execute_resize(
&self,
plan: onnx_runtime_ep_api::ResidencyResizePlan,
device_count: usize,
) -> onnx_runtime_ep_api::ResidencyResizeOutcome {
use onnx_runtime_ep_api::{ResidencyResizeOutcome, ResidencyResizePlan, ResizeRejection};
let safe_point = self.resize_safe_point(device_count);
let request = match plan {
ResidencyResizePlan::Rejected { request, reason } => {
let (before, _) = self.budget();
return ResidencyResizeOutcome {
direction: request.direction,
requested_bytes: request.target_bytes,
accepted_bytes: 0,
before_bytes: before,
after_bytes: before,
rejection: Some(reason),
rollback_count: 0,
safe_point,
};
}
ResidencyResizePlan::Accepted(request) => request,
};
let (before, governed) = self.budget();
if let Some(reason) = safe_point.blocking_reason() {
return ResidencyResizeOutcome {
direction: request.direction,
requested_bytes: request.target_bytes,
accepted_bytes: 0,
before_bytes: before,
after_bytes: before,
rejection: Some(ResizeRejection::NotSafePoint(reason)),
rollback_count: 0,
safe_point,
};
}
match request.direction {
onnx_runtime_ep_api::ResizeDirection::Grow => {
let mut inner = self.inner.lock().expect("residency lock poisoned");
match inner.lease.as_mut() {
Some(lease) => match lease.grow(request.target_bytes) {
Ok(()) => {
inner.policy.budget =
inner.policy.budget.saturating_add(request.target_bytes);
let after = inner.policy.budget;
replace_global_budget(before, after);
ResidencyResizeOutcome {
direction: request.direction,
requested_bytes: request.target_bytes,
accepted_bytes: request.target_bytes,
before_bytes: before,
after_bytes: after,
rejection: None,
rollback_count: 0,
safe_point,
}
}
Err(error) => ResidencyResizeOutcome {
direction: request.direction,
requested_bytes: request.target_bytes,
accepted_bytes: 0,
before_bytes: before,
after_bytes: before,
rejection: Some(ResizeRejection::ExecutionFailed(error.to_string())),
rollback_count: 1,
safe_point,
},
},
None => ResidencyResizeOutcome {
direction: request.direction,
requested_bytes: request.target_bytes,
accepted_bytes: 0,
before_bytes: before,
after_bytes: before,
rejection: Some(ResizeRejection::ExecutionFailed(
"growing an ungoverned budget has no lease to grow; adopt a governed \
budget first"
.into(),
)),
rollback_count: 0,
safe_point,
},
}
}
onnx_runtime_ep_api::ResizeDirection::Shrink => {
if !governed {
return ResidencyResizeOutcome {
direction: request.direction,
requested_bytes: request.target_bytes,
accepted_bytes: 0,
before_bytes: before,
after_bytes: before,
rejection: Some(ResizeRejection::ExecutionFailed(
"shrinking an ungoverned budget has no lease/allowance to release \
into; refusing rather than silently shrinking a resident cache"
.into(),
)),
rollback_count: 0,
safe_point,
};
}
let has_mapped_allowance = self
.inner
.lock()
.expect("residency lock poisoned")
.mapped_allowance
.is_some();
if has_mapped_allowance {
use onnx_runtime_memory_governor::ReclaimableMappedHolder as _;
match self.reclaim_mapped(request.target_bytes) {
Ok(report) => {
let after = before.saturating_sub(report.reclaimed_bytes);
ResidencyResizeOutcome {
direction: request.direction,
requested_bytes: request.target_bytes,
accepted_bytes: report.reclaimed_bytes,
before_bytes: before,
after_bytes: after,
rejection: None,
rollback_count: 0,
safe_point,
}
}
Err(error) => ResidencyResizeOutcome {
direction: request.direction,
requested_bytes: request.target_bytes,
accepted_bytes: 0,
before_bytes: before,
after_bytes: before,
rejection: Some(ResizeRejection::ExecutionFailed(error.to_string())),
rollback_count: 1,
safe_point,
},
}
} else {
let mut inner = self.inner.lock().expect("residency lock poisoned");
let mut reclaimed = 0u64;
let max_attempts = inner.pages.len();
let mut attempts = 0usize;
while reclaimed < request.target_bytes && attempts < max_attempts {
let Some(key) = inner.next_evictable_key(WeightEvictionPolicy::Lru) else {
break;
};
let bytes = inner.policy.bytes_by_key.get(&key).copied().unwrap_or(0);
inner.remove_page(key);
reclaimed = reclaimed.saturating_add(bytes);
attempts += 1;
}
let returned = match inner.lease.as_mut() {
Some(lease) => lease.shrink(reclaimed),
None => 0,
};
inner.policy.budget = inner.policy.budget.saturating_sub(returned);
let after = inner.policy.budget;
drop(inner);
replace_global_budget(before, after);
ResidencyResizeOutcome {
direction: request.direction,
requested_bytes: request.target_bytes,
accepted_bytes: returned,
before_bytes: before,
after_bytes: after,
rejection: None,
rollback_count: 0,
safe_point,
}
}
}
}
}
pub fn with_deferred_release_queue(mut self, queue: Arc<CudaDeferredReleaseQueue>) -> Self {
self.queue = Some(queue);
self
}
pub fn install_context_scope(
&self,
scope: onnx_runtime_memory_governor::MemoryContextScope,
) -> Result<(), &'static str> {
self.context_scope
.set(scope)
.map_err(|_| "weight residency context scope was already installed")
}
fn enter_context_operation(
&self,
) -> Result<Option<onnx_runtime_memory_governor::MemoryContextOperation>, WeightHandleError>
{
if self.context_terminated.load(Ordering::Acquire) {
return Err(WeightHandleError::DeviceBinding(
"weight residency belongs to a terminated provider context".into(),
));
}
self.context_scope
.get()
.map(onnx_runtime_memory_governor::MemoryContextScope::enter)
.transpose()
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))
}
pub fn confirm_context_terminated(&self) {
if self.context_terminated.swap(true, Ordering::AcqRel) {
return;
}
let (allowance, lease) = {
let mut inner = self.lock();
for slot in inner.slots.values() {
slot.state.poison();
}
let allowance = inner.mapped_allowance.take();
let lease = inner.lease.take();
inner.policy.budget = 0;
(allowance, lease)
};
if let Some(allowance) = allowance {
allowance.confirm_context_terminated();
}
drop(lease);
}
pub fn deferred_release_queue(&self) -> Option<&Arc<CudaDeferredReleaseQueue>> {
self.queue.as_ref()
}
pub fn with_async_pagein(self, async_pagein: bool) -> Self {
let _ = async_pagein;
self
}
pub fn with_scan_resistant_dense(mut self, scan_resistant_dense: bool) -> Self {
self.scan_resistant_dense = scan_resistant_dense;
self
}
pub fn with_byte_aware_residency(mut self, byte_aware: bool) -> Self {
self.byte_aware = byte_aware;
self
}
pub fn with_evict_order_probe(mut self, evict_order_probe: EvictOrderProbe) -> Self {
self.evict_order_probe = evict_order_probe;
self
}
pub fn with_zero_copy_hybrid(mut self, zero_copy_hybrid: bool) -> Self {
self.zero_copy_hybrid = zero_copy_hybrid;
self
}
pub fn with_vmm_admission(
self,
allocator: Arc<crate::vmm_allocator::CudaVmmAllocator>,
governor: Arc<dyn onnx_runtime_memory_governor::MemoryGovernor + Send + Sync>,
) -> Result<Self, WeightHandleError> {
if self.lock().lease.is_some() {
return Err(WeightHandleError::DeviceBinding(
"cannot install VMM physical admission on a cache that already holds a \
content-byte governor lease"
.into(),
));
}
if !vmm_committed_authority_matches(allocator.as_ref(), governor.as_ref()) {
return Err(WeightHandleError::DeviceBinding(format!(
"VMM weight residency and its governor must share one committed-byte \
authority (allocator: {:?}, governor: {:?})",
allocator.committed_byte_authority(),
governor.authority_id()
)));
}
let _ = self.physical.set(PhysicalAdmission {
allocator,
governor,
});
Ok(self)
}
pub(crate) fn install_vmm_admission(
&self,
allocator: Arc<crate::vmm_allocator::CudaVmmAllocator>,
governor: Arc<dyn onnx_runtime_memory_governor::MemoryGovernor + Send + Sync>,
) -> Result<(), WeightHandleError> {
if self.lock().lease.is_some() {
return Err(WeightHandleError::DeviceBinding(
"cannot install VMM physical admission on a cache that already holds a \
content-byte governor lease"
.into(),
));
}
if !vmm_committed_authority_matches(allocator.as_ref(), governor.as_ref()) {
return Err(WeightHandleError::DeviceBinding(format!(
"VMM weight residency and its governor must share one committed-byte \
authority (allocator: {:?}, governor: {:?})",
allocator.committed_byte_authority(),
governor.authority_id()
)));
}
self.physical
.set(PhysicalAdmission {
allocator,
governor,
})
.map_err(|_| {
WeightHandleError::DeviceBinding(
"VMM weight residency admission was installed more than once".into(),
)
})
}
pub fn prefetch_block_quantized_moe(
&self,
key: u64,
weight: &LazyWeight,
source: &dyn MmapRegionSource,
) -> Result<bool, WeightHandleError> {
if weight.boundary != LazyWeightBoundary::BlockQuantizedMoe {
return Ok(false);
}
if self.physical.get().is_some() || self.zero_copy_hybrid {
GLOBAL_PREFETCH_DECLINED_UNSUPPORTED.fetch_add(1, Ordering::Relaxed);
self.lock().prefetch_stats.declined_unsupported += 1;
return Ok(false);
}
let bytes = weight.region_bytes_len() as u64;
if bytes == 0 {
return Ok(false);
}
{
let mut inner = self.lock();
if inner.pages.contains_key(&key) {
GLOBAL_PREFETCH_DECLINED_RESIDENT.fetch_add(1, Ordering::Relaxed);
inner.prefetch_stats.declined_resident += 1;
return Ok(false);
}
if inner.pending_prefetch.is_some() {
GLOBAL_PREFETCH_DECLINED_BUSY.fetch_add(1, Ordering::Relaxed);
inner.prefetch_stats.declined_busy += 1;
return Ok(false);
}
if !inner.policy.can_fit(bytes) {
GLOBAL_PREFETCH_DECLINED_BUDGET.fetch_add(1, Ordering::Relaxed);
inner.prefetch_stats.declined_budget += 1;
return Ok(false);
}
}
if !self.staging_pool.can_retain_concurrent(bytes as usize, 2) {
let mut inner = self.lock();
GLOBAL_PREFETCH_DECLINED_POOL_CAPACITY.fetch_add(1, Ordering::Relaxed);
inner.prefetch_stats.declined_pool_capacity += 1;
return Ok(false);
}
let mut staging = self
.staging_pool
.acquire(bytes as usize)
.map_err(|error| WeightHandleError::DeviceBinding(format!("pinned alloc: {error}")))?;
fill_staging_from_regions(weight, source, staging.staging_mut())?;
let ptr = self
.runtime
.alloc_raw(bytes as usize)
.map_err(|error| WeightHandleError::DeviceBinding(format!("VRAM alloc: {error}")))?;
let page = Arc::new(CudaWeightPage {
runtime: Arc::clone(&self.runtime),
queue: self.queue.clone(),
allocation: WeightAllocation::Runtime,
ptr,
len: bytes as usize,
dtype: weight.dtype,
shape: weight.shape.clone(),
});
if let Err(error) = unsafe { self.runtime.htod_async(staging.as_slice(), ptr) } {
return Err(WeightHandleError::DeviceBinding(format!(
"H2D prefetch enqueue: {error}"
)));
}
GLOBAL_HTOD_BYTES.fetch_add(bytes, Ordering::Relaxed);
let fence_id = self
.runtime
.record_copy_fence()
.map_err(|error| WeightHandleError::DeviceBinding(format!("record fence: {error}")))?;
let mut inner = self.lock();
if inner.pages.contains_key(&key) || inner.pending_prefetch.is_some() {
drop(inner);
match self.runtime.resolve_prefetch_fence(fence_id) {
Ok(completed) => staging.retire(completed),
Err(error) => {
let (detail, completion) = error.into_parts();
if matches!(completion, FailedHtodCompletion::MayBeInFlight) {
quarantine_in_flight_fill(Box::new((page, staging)));
} else {
eprintln!(
"cuda_ep: note: prefetch fence resolution recovered via a fallback \
copy-stream sync while discarding a lost-race prefetch: {detail}"
);
}
}
}
return Ok(false);
}
inner.pending_prefetch = Some(PendingPrefetch {
key,
page,
fence_id,
staging,
bytes,
});
inner.prefetch_stats.issued += 1;
inner.prefetch_stats.issued_bytes += bytes;
drop(inner);
GLOBAL_PREFETCH_ISSUED.fetch_add(1, Ordering::Relaxed);
GLOBAL_PREFETCH_ISSUED_BYTES.fetch_add(bytes, Ordering::Relaxed);
Ok(true)
}
fn promote_pending_prefetch(
&self,
key: u64,
) -> Option<Result<Arc<CudaWeightPage>, WeightHandleError>> {
let pending = {
let mut inner = self.lock();
match inner.pending_prefetch.as_ref() {
Some(pending) if pending.key == key => inner.pending_prefetch.take(),
_ => None,
}
}?;
let wait_start = std::time::Instant::now();
let resolved = self.runtime.resolve_prefetch_fence(pending.fence_id);
let wait_elapsed = wait_start.elapsed();
add_duration(&GLOBAL_PREFETCH_PROMOTE_WAIT_NS, wait_elapsed);
let completed = match resolved {
Ok(completed) => completed,
Err(error) => {
let (detail, completion) = error.into_parts();
let error =
WeightHandleError::DeviceBinding(format!("resolving prefetch fence: {detail}"));
if matches!(completion, FailedHtodCompletion::MayBeInFlight) {
quarantine_in_flight_fill(Box::new((pending.page, pending.staging)));
return Some(Err(WeightHandleError::DeviceBinding(format!(
"{error}; the pending prefetch's destination page and staging source \
were quarantined because copy-stream completion could not be \
established"
))));
}
return Some(Err(error));
}
};
debug_assert_eq!(
pending.bytes,
pending.page.len() as u64,
"pending prefetch byte count does not match its device allocation"
);
debug_assert_eq!(
pending.bytes,
pending.staging.as_slice().len() as u64,
"pending prefetch byte count does not match its staging buffer"
);
pending.staging.retire(completed);
GLOBAL_PREFETCH_PROMOTED.fetch_add(1, Ordering::Relaxed);
{
let mut inner = self.lock();
inner.prefetch_stats.promoted += 1;
inner.prefetch_stats.promote_wait_ns +=
wait_elapsed.as_nanos().min(u128::from(u64::MAX)) as u64;
}
Some(self.admit(
pending.key,
pending.page,
self.eviction_for(LazyWeightBoundary::BlockQuantizedMoe),
))
}
pub fn resident<S: MmapRegionSource>(
&self,
key: u64,
weight: &LazyWeight,
source: &S,
) -> Result<Arc<CudaWeightPage>, WeightHandleError> {
let _context_operation = self.enter_context_operation()?;
if let Some(hit) = self.get_hit(key) {
return Ok(hit);
}
if let Some(promoted) = self.promote_pending_prefetch(key) {
return promoted;
}
if self.physical.get().is_some() {
let resident = weight.materialize()?;
GLOBAL_MATERIALIZE_FALLBACK_CALLS.fetch_add(1, Ordering::Relaxed);
let bytes = resident.bytes().to_vec();
return self
.resident_vmm_with(
key,
resident.dtype,
resident.shape.clone(),
bytes.len(),
self.eviction_for(weight.boundary),
false,
move |runtime, ptr| {
unsafe { runtime.htod(&bytes, ptr) }.map_err(|error| {
WeightHandleError::DeviceBinding(format!("H2D copy: {error}"))
})
},
)
.map(VmmAdmit::expect_page);
}
let mut pager = CudaWeightPager::new(Arc::clone(&self.runtime), source);
if let Some(queue) = self.queue.as_ref() {
pager = pager.with_deferred_release_queue(Arc::clone(queue));
}
let page = Arc::new(pager.bind_block_quantized_moe(weight)?);
self.admit(key, page, self.eviction_for(weight.boundary))
}
pub fn resident_mapped(
&self,
key: u64,
weight: &LazyWeight,
source: &dyn MmapRegionSource,
) -> Result<Arc<CudaWeightPage>, WeightHandleError> {
let _context_operation = self.enter_context_operation()?;
if self.zero_copy_hybrid && self.physical.get().is_some() {
return self.resident_mapped_hybrid(key, weight, source);
}
self.resident_mapped_inner(key, weight, source, false)
.map(VmmAdmit::expect_page)
}
fn resident_mapped_inner(
&self,
key: u64,
weight: &LazyWeight,
source: &dyn MmapRegionSource,
hybrid_zero_copy: bool,
) -> Result<VmmAdmit, WeightHandleError> {
if let Some(hit) = self.get_hit(key) {
return Ok(VmmAdmit::Page(hit));
}
if let Some(promoted) = self.promote_pending_prefetch(key) {
return promoted.map(VmmAdmit::Page);
}
let len = weight.region_bytes_len();
let mut staging = self
.staging_pool
.acquire(len)
.map_err(|error| WeightHandleError::DeviceBinding(format!("pinned alloc: {error}")))?;
let materialize_start = std::time::Instant::now();
fill_staging_from_regions(weight, source, staging.staging_mut())?;
add_duration(&GLOBAL_MATERIALIZE_NS, materialize_start.elapsed());
if self.physical.get().is_some() {
return self.resident_vmm_with(
key,
weight.dtype,
weight.shape.clone(),
len,
self.eviction_for(weight.boundary),
hybrid_zero_copy,
move |runtime, ptr| {
let staged = &staging.as_slice()[..len];
let (copy_ms, completed) = match unsafe {
runtime.htod_async_elapsed_ms(staged, ptr)
} {
Ok(result) => result,
Err(error) => {
let (detail, completion) = error.into_parts();
let error = WeightHandleError::DeviceBinding(format!(
"measured H2D copy: {detail}"
));
return match completion {
FailedHtodCompletion::NotSubmitted => {
Err(VmmFillFailure::completed(error))
}
FailedHtodCompletion::Completed(completed) => {
staging.retire(completed);
Err(VmmFillFailure::completed(error))
}
FailedHtodCompletion::MayBeInFlight => {
let source = staging.into_inner();
Err(VmmFillFailure::may_be_in_flight(error, Box::new(source)))
}
};
}
};
GLOBAL_HTOD_NS.fetch_add((copy_ms * 1_000_000.0) as u64, Ordering::Relaxed);
GLOBAL_HTOD_BYTES.fetch_add(len as u64, Ordering::Relaxed);
staging.retire(completed);
Ok(())
},
);
}
let raw_staging = staging.into_inner();
let (page, _, raw_staging, completed) = match self.queue.as_ref() {
Some(queue) => CudaWeightPage::upload_staged_async_queued(
&self.runtime,
weight.dtype,
weight.shape.clone(),
len,
raw_staging,
Arc::clone(queue),
),
None => CudaWeightPage::upload_staged_async(
&self.runtime,
weight.dtype,
weight.shape.clone(),
len,
raw_staging,
),
}?;
self.staging_pool.release(raw_staging, completed);
self.admit(key, Arc::new(page), self.eviction_for(weight.boundary))
.map(VmmAdmit::Page)
}
fn resident_mapped_hybrid(
&self,
key: u64,
weight: &LazyWeight,
source: &dyn MmapRegionSource,
) -> Result<Arc<CudaWeightPage>, WeightHandleError> {
if let Some(hit) = self.get_hit(key) {
return Ok(hit);
}
{
let inner = self.lock();
if let Some(cold) = inner.cold_pages.get(&key).cloned() {
let len = cold.len();
drop(inner);
self.record_zero_copy_read(len);
return Ok(cold);
}
}
let len = weight.region_bytes_len();
match self.resident_mapped_inner(key, weight, source, true)? {
VmmAdmit::Page(page) => Ok(page),
VmmAdmit::DeferToZeroCopy => self.bind_zero_copy(key, weight, source, len),
}
}
fn bind_zero_copy(
&self,
key: u64,
weight: &LazyWeight,
source: &dyn MmapRegionSource,
len: usize,
) -> Result<Arc<CudaWeightPage>, WeightHandleError> {
if GLOBAL_ZERO_COPY_BOUND_BYTES
.load(Ordering::Relaxed)
.saturating_add(len as u64)
> zero_copy_budget_bytes()
{
return self
.resident_mapped_inner(key, weight, source, false)
.map(VmmAdmit::expect_page);
}
let device_ptr = match self.zero_copy_device_ptr(weight, source)? {
Some(ptr) => ptr,
None => {
return self
.resident_mapped_inner(key, weight, source, false)
.map(VmmAdmit::expect_page);
}
};
if zero_copy_max_binds()
.is_some_and(|max| GLOBAL_ZERO_COPY_BINDS.load(Ordering::Relaxed) >= max)
{
return self
.resident_mapped_inner(key, weight, source, false)
.map(VmmAdmit::expect_page);
} if zero_copy_copy_instead() {
return self
.resident_mapped_inner(key, weight, source, false)
.map(VmmAdmit::expect_page);
}
if zero_copy_debug() {
let host_align = (device_ptr as usize) & 0xff;
eprintln!(
"zero_copy_bind: key={key} len={len} dptr=0x{device_ptr:x} dptr_align256={host_align} \
regions={} first_off={}",
weight.regions.len(),
weight.regions.first().map(|r| r.offset).unwrap_or(0)
);
}
let page = Arc::new(CudaWeightPage {
runtime: Arc::clone(&self.runtime),
queue: self.queue.clone(),
allocation: WeightAllocation::HostMapped,
ptr: device_ptr,
len,
dtype: weight.dtype,
shape: weight.shape.clone(),
});
{
let mut inner = self.lock();
if let Some(existing) = inner.cold_pages.get(&key).cloned() {
let existing_len = existing.len();
drop(inner);
self.record_zero_copy_read(existing_len);
return Ok(existing);
}
inner.cold_pages.insert(key, Arc::clone(&page));
}
GLOBAL_ZERO_COPY_BINDS.fetch_add(1, Ordering::Relaxed);
GLOBAL_ZERO_COPY_BOUND_BYTES.fetch_add(len as u64, Ordering::Relaxed);
self.record_zero_copy_read(len);
Ok(page)
}
fn zero_copy_device_ptr(
&self,
weight: &LazyWeight,
source: &dyn MmapRegionSource,
) -> Result<Option<CUdeviceptr>, WeightHandleError> {
let Some(first) = weight.regions.first() else {
return Ok(None);
};
let mapping_id = first.mapping_id;
let mut expected = first.offset;
for region in &weight.regions {
if region.mapping_id != mapping_id || region.offset != expected {
return Ok(None);
}
expected = match expected.checked_add(region.len) {
Some(next) => next,
None => return Ok(None),
};
}
let span = ExternalMmapRegion {
mapping_id,
offset: first.offset,
len: weight.region_bytes_len(),
};
let bytes = source.region_bytes(&span)?;
let host_ptr = bytes.as_ptr();
let Some(mapping) = source.full_mapping_bytes(mapping_id) else {
return Ok(None);
};
let device_ptr = self
.host_registry
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.device_ptr_for(mapping_id, mapping, host_ptr)?;
Ok(Some(device_ptr))
}
fn record_zero_copy_read(&self, len: usize) {
GLOBAL_ZERO_COPY_READS.fetch_add(1, Ordering::Relaxed);
GLOBAL_ZERO_COPY_BYTES.fetch_add(len as u64, Ordering::Relaxed);
}
pub fn resident_materialized(
&self,
key: u64,
weight: &LazyWeight,
) -> Result<Arc<CudaWeightPage>, WeightHandleError> {
let _context_operation = self.enter_context_operation()?;
if let Some(hit) = self.get_hit(key) {
return Ok(hit);
}
let materialize_start = std::time::Instant::now();
let resident = weight.materialize()?;
GLOBAL_MATERIALIZE_FALLBACK_CALLS.fetch_add(1, Ordering::Relaxed);
add_duration(&GLOBAL_MATERIALIZE_NS, materialize_start.elapsed());
if self.physical.get().is_some() {
let bytes = resident.bytes().to_vec();
return self
.resident_vmm_with(
key,
resident.dtype,
resident.shape.clone(),
bytes.len(),
self.eviction_for(weight.boundary),
false,
move |runtime, ptr| {
unsafe { runtime.htod(&bytes, ptr) }.map_err(|error| {
WeightHandleError::DeviceBinding(format!("H2D copy: {error}"))
})
},
)
.map(VmmAdmit::expect_page);
}
let page = match self.queue.as_ref() {
Some(queue) => CudaWeightPage::upload_queued(
&self.runtime,
resident.dtype,
resident.shape.clone(),
resident.bytes(),
Arc::clone(queue),
),
None => CudaWeightPage::upload(
&self.runtime,
resident.dtype,
resident.shape.clone(),
resident.bytes(),
),
}?;
let page = Arc::new(page);
self.admit(key, page, self.eviction_for(weight.boundary))
}
#[allow(clippy::too_many_arguments)]
fn resident_vmm_with<F, R>(
&self,
key: u64,
dtype: DataType,
shape: Vec<usize>,
len: usize,
eviction: WeightEvictionPolicy,
hybrid_zero_copy: bool,
fill: F,
) -> Result<VmmAdmit, WeightHandleError>
where
F: FnOnce(&CudaRuntime, CUdeviceptr) -> R,
R: IntoVmmFillResult,
{
let physical = self
.physical
.get()
.expect("VMM residency helper requires physical admission");
let mut inner = self.lock();
if let Some(existing) = inner.pages.get(&key).cloned() {
let slot_state = inner.slots.get(&key).map(|slot| Arc::clone(&slot.state));
if let Some(state) = slot_state.as_ref() {
match state.status() {
SlotStatus::Idle => {}
SlotStatus::Pending => {
return Err(WeightHandleError::DeviceBinding(format!(
"stable weight slot for key {key} has a refill or release pending; \
retry after that operation terminally completes"
)));
}
SlotStatus::Poisoned => {
return Err(WeightHandleError::DeviceBinding(format!(
"stable weight slot for key {key} is poisoned; its device pointer \
must never be returned or written again"
)));
}
}
}
inner.record_hit(key);
let is_pinned = inner.pinned.contains(&key);
let probe_checksum =
is_pinned && pin_checksum_keys().is_some_and(|keys| keys.contains(&key));
let refill_every = if is_pinned { pin_refill_every() } else { None };
if is_pinned && (probe_checksum || refill_every.is_some()) {
let step = {
let counter = inner.pin_hit_step.entry(key).or_insert(0);
*counter += 1;
*counter
};
let baseline = if probe_checksum {
inner.pin_granule_hashes.get(&key).cloned()
} else {
None
};
let refill_claim = if refill_every.is_some_and(|every| step % every == 0) {
let state = slot_state.ok_or_else(|| {
WeightHandleError::DeviceBinding(format!(
"pinned resident key {key} has no stable-slot refill state"
))
})?;
if Arc::strong_count(&existing) != 2 {
None
} else if !state.begin_refill() {
return Err(WeightHandleError::DeviceBinding(format!(
"stable weight slot for key {key} could not claim its refill; \
another operation is pending or the slot is poisoned"
)));
} else {
Some(state)
}
} else {
None
};
drop(inner);
if let Some(baseline) = baseline {
let current = match self.hash_page_granules(existing.ptr, existing.len) {
Ok(current) => current,
Err(error) => {
if let Some(state) = refill_claim.as_ref() {
state.finish_refill();
}
return Err(error);
}
};
report_pin_granule_diff(key, step, &baseline, ¤t);
}
if let Some(state) = refill_claim {
if let Err(error) = self.runtime.drain_for_unmap() {
state.finish_refill();
return Err(WeightHandleError::DeviceBinding(format!(
"pin-refill compute drain: {error}"
)));
}
match fill(&self.runtime, existing.ptr).into_vmm_fill_result() {
Ok(()) => state.finish_refill(),
Err(failure) => {
return Err(
self.handle_pinned_refill_failure(key, existing, state, failure)
);
}
}
if state.status() != SlotStatus::Idle {
return Err(WeightHandleError::DeviceBinding(format!(
"stable weight slot for key {key} became poisoned during refill; \
refusing to return its device pointer"
)));
}
eprintln!(
"weight_pin_refill[#945]: key={key} step={step} re-filled \
pinned page from host source"
);
}
return Ok(VmmAdmit::Page(existing));
}
return Ok(VmmAdmit::Page(existing));
}
let allowance = inner.mapped_allowance.clone().ok_or_else(|| {
WeightHandleError::DeviceBinding(
"VMM weight residency has no authority-scoped mapped-byte allowance; \
adopt the memory governor before page-in"
.into(),
)
})?;
let reused_slot = inner.slots.get(&key).cloned();
let ptr = match reused_slot.as_ref() {
Some(slot) => {
if slot.len != len {
return Err(WeightHandleError::DeviceBinding(format!(
"stable weight slot for key {key} was reserved for {} bytes but a \
{len}-byte page-in requested it",
slot.len
)));
}
match slot.state.status() {
SlotStatus::Idle => {}
SlotStatus::Pending => {
let Some(queue) = self.queue.as_ref().cloned() else {
return Err(WeightHandleError::DeviceBinding(format!(
"stable weight slot for key {key} has a deferred decommit in \
flight but this residency has no release queue to settle it"
)));
};
drop(inner);
if !queue.wait_until_idle(DEFERRED_RELEASE_WAIT_TIMEOUT) {
return Err(WeightHandleError::DeviceBinding(format!(
"stable weight slot for key {key} did not finish its deferred \
decommit within {DEFERRED_RELEASE_WAIT_TIMEOUT:?}; the address \
remains unavailable rather than being remapped underneath it"
)));
}
inner = self.lock();
if let Some(existing) = inner.pages.get(&key).cloned() {
inner.record_hit(key);
return Ok(VmmAdmit::Page(existing));
}
match slot.state.status() {
SlotStatus::Idle => {}
SlotStatus::Pending => {
return Err(WeightHandleError::DeviceBinding(format!(
"stable weight slot for key {key} remains pending after its \
deferred-release queue drained"
)));
}
SlotStatus::Poisoned => {
return Err(WeightHandleError::DeviceBinding(format!(
"stable weight slot for key {key} is poisoned: its previous \
decommit retained ownership"
)));
}
}
}
SlotStatus::Poisoned => {
return Err(WeightHandleError::DeviceBinding(format!(
"stable weight slot for key {key} is poisoned: a previous decommit \
did not complete and its physical ownership is retained, so this \
address can never be mapped again"
)));
}
}
NonNull::new(slot.va as *mut u8).ok_or_else(|| {
WeightHandleError::DeviceBinding("stable weight slot has a null VA".into())
})?
}
None => physical
.allocator
.allocate_committed(len, WEIGHT_SLOT_ALIGN, &[])
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))?,
};
let mut bypass = match self.admit_committed_span(
&mut inner,
physical,
&allowance,
key,
ptr,
len,
eviction,
reused_slot.is_some(),
hybrid_zero_copy,
fill,
) {
Ok(SpanAdmit::Filled { bypass }) => bypass,
Ok(SpanAdmit::DeferToZeroCopy) => {
if reused_slot.is_none() {
let _ = physical.allocator.deallocate_span(ptr);
}
return Ok(VmmAdmit::DeferToZeroCopy);
}
Err(failure) => {
if reused_slot.is_none() && failure.fresh_span == FreshSpanCleanup::CallerOwns {
let _ = physical.allocator.deallocate_span(ptr);
}
return Err(failure.error);
}
};
if reused_slot.is_some() && bypass {
static SLOTTED_BYPASS_SEEN: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(false);
if !SLOTTED_BYPASS_SEEN.swap(true, Ordering::Relaxed) {
eprintln!(
"weight_paging_diag[#888]: slotted-key bypass occurred (key={key}, \
len={len}) — stable_slot/residency disagreement reachable under byte-aware"
);
}
if slot_bypass_retain_enabled() {
bypass = false;
}
}
let stable_slot = reused_slot.is_some() || !bypass;
let slot_state = match reused_slot.as_ref() {
Some(slot) => Some(Arc::clone(&slot.state)),
None if stable_slot => {
let state = Arc::new(SlotOperationState::default());
inner.slots.insert(
key,
StableWeightSlot {
va: ptr.as_ptr() as CUdeviceptr,
len,
state: Arc::clone(&state),
},
);
Some(state)
}
None => None,
};
let page = Arc::new(CudaWeightPage {
runtime: Arc::clone(&self.runtime),
queue: self.queue.clone(),
allocation: WeightAllocation::Vmm {
allocator: Arc::clone(&physical.allocator),
allowance: allowance.clone(),
stable_slot,
slot_state,
},
ptr: ptr.as_ptr() as CUdeviceptr,
len,
dtype,
shape,
});
if bypass {
inner.record_bypassed_page_in(key, len as u64);
} else {
inner.insert_page(key, Arc::clone(&page), len as u64);
}
let want_baseline = !bypass
&& inner.pinned.contains(&key)
&& pin_checksum_keys().is_some_and(|keys| keys.contains(&key));
drop(inner);
if want_baseline {
let hashes = self.hash_page_granules(page.ptr, page.len)?;
eprintln!(
"weight_pin_checksum[#945]: key={key} admitted len={} granules={} \
(baseline snapshot)",
page.len,
hashes.len()
);
self.lock().pin_granule_hashes.insert(key, hashes);
}
Ok(VmmAdmit::Page(page))
}
fn handle_pinned_refill_failure(
&self,
key: u64,
existing: Arc<CudaWeightPage>,
state: Arc<SlotOperationState>,
failure: VmmFillFailure,
) -> WeightHandleError {
let VmmFillFailure {
error,
in_flight_source,
} = failure;
let Some(source) = in_flight_source else {
state.finish_refill();
return error;
};
let mut inner = self.lock();
state.poison();
inner.remove_page(key);
drop(inner);
quarantine_in_flight_fill(Box::new((existing, source)));
WeightHandleError::DeviceBinding(format!(
"{error}; pinned-page refill completion could not be established; the destination \
mapping and staging source remain charged and quarantined"
))
}
#[allow(clippy::too_many_arguments)]
fn admit_committed_span<F, R>(
&self,
inner: &mut ResidencyInner,
physical: &PhysicalAdmission,
allowance: &onnx_runtime_memory_governor::MappedAllowance,
key: u64,
ptr: NonNull<u8>,
len: usize,
eviction: WeightEvictionPolicy,
has_stable_slot: bool,
hybrid_zero_copy: bool,
fill: F,
) -> Result<SpanAdmit, SpanAdmitError>
where
F: FnOnce(&CudaRuntime, CUdeviceptr) -> R,
R: IntoVmmFillResult,
{
let mut fill = Some(fill);
let max_evictions = inner.pages.len();
let mut evictions = 0usize;
let mut bypass = false;
let pin_this = {
use onnx_runtime_ep_api::ResidencyPolicy as _;
CudaHotSetResidencyPolicy::from_env(self.scan_resistant_dense).should_pin(
&onnx_runtime_ep_api::AdmissionPolicyInput {
key,
len_bytes: len as u64,
already_pinned: inner.pinned.contains(&key),
pinned_bytes_used: inner.pinned_bytes,
},
)
};
loop {
let required_owned = physical
.allocator
.incremental_owned_bytes_for_span(ptr, len, 0, len)
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))?;
let required_mapped = physical
.allocator
.incremental_mapped_bytes_for_span(ptr, len, 0, len)
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))?;
let global_available = physical.governor.available(Tier::Device);
let zone_available = allowance.available();
if committed_admission_fits(
required_mapped,
zone_available,
required_owned,
global_available,
) {
allowance
.try_map(required_mapped)
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))?;
GLOBAL_WEIGHT_MAPPED_BYTES.fetch_add(required_mapped, Ordering::Relaxed);
match physical.allocator.try_commit_span(
ptr,
len,
0,
len,
required_mapped,
global_available,
) {
Ok(commit) => {
let excess = required_mapped.saturating_sub(commit.newly_mapped_bytes);
if excess > 0 {
allowance.unmap(excess);
let _ = GLOBAL_WEIGHT_MAPPED_BYTES.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|current| Some(current.saturating_sub(excess)),
);
}
if sync_before_fill_enabled() {
if let Err(error) = self.runtime.drain_for_unmap() {
return Err(self.rollback_failed_vmm_fill(
physical,
allowance,
ptr,
len,
has_stable_slot,
inner.slots.get(&key).map(|slot| Arc::clone(&slot.state)),
VmmFillFailure::completed(WeightHandleError::DeviceBinding(
format!("sync-before-fill compute drain: {error}"),
)),
));
}
if let Err(error) = self.runtime.copy_stream().synchronize() {
return Err(self.rollback_failed_vmm_fill(
physical,
allowance,
ptr,
len,
has_stable_slot,
inner.slots.get(&key).map(|slot| Arc::clone(&slot.state)),
VmmFillFailure::completed(WeightHandleError::DeviceBinding(
format!("sync-before-fill copy drain: {error}"),
)),
));
}
}
if let Err(failure) = fill.take().expect("VMM page fill runs once")(
&self.runtime,
ptr.as_ptr() as CUdeviceptr,
)
.into_vmm_fill_result()
{
return Err(self.rollback_failed_vmm_fill(
physical,
allowance,
ptr,
len,
has_stable_slot,
inner.slots.get(&key).map(|slot| Arc::clone(&slot.state)),
failure,
));
}
if pin_this && !bypass {
inner.mark_pinned(key, len as u64);
}
return Ok(SpanAdmit::Filled { bypass });
}
Err(error) => {
allowance.unmap(required_mapped);
let _ = GLOBAL_WEIGHT_MAPPED_BYTES.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|current| Some(current.saturating_sub(required_mapped)),
);
if evictions >= max_evictions {
return Err(WeightHandleError::DeviceBinding(error.to_string()).into());
}
}
}
} else if eviction == WeightEvictionPolicy::StableResident {
bypass = !(pin_this
|| self.byte_aware
&& (has_stable_slot && retain_slotted_enabled()
|| inner
.smallest_evictable()
.is_some_and(|(_, smallest)| (len as u64) > smallest)));
}
if hybrid_zero_copy {
return Ok(SpanAdmit::DeferToZeroCopy);
}
if evictions >= max_evictions {
return Err(WeightHandleError::DeviceBinding(format!(
"weight residency requires {required_mapped} incremental mapped bytes with \
{zone_available} bytes of weight-zone headroom and {required_owned} \
incremental committed bytes with {global_available} bytes of physical \
headroom after {evictions} eviction(s)"
))
.into());
}
let before_owned = physical
.allocator
.physical_pool_stats()
.map_or(0, |stats| stats.snapshot().total_owned_bytes);
let before_required_owned = required_owned;
let before_required_mapped = required_mapped;
let evicted_key = if self.byte_aware || pin_this {
inner.smallest_evictable().map(|(key, _)| key)
} else {
inner.evictable_key_by_probe(self.evict_order_probe, eviction)
};
let Some(evicted_key) = evicted_key else {
return Err(WeightHandleError::DeviceBinding(format!(
"weight residency requires {required_mapped} incremental mapped bytes with \
{zone_available} bytes of weight-zone headroom and {required_owned} \
incremental committed bytes with {global_available} bytes of physical \
headroom, and no page is evictable"
))
.into());
};
let Some(queue) = self.queue.as_ref().cloned() else {
return Err(WeightHandleError::DeviceBinding(
"weight residency must evict a page, but no deferred-release queue is \
installed; refusing rather than leaking the page or freeing it before \
in-flight CUDA work completes"
.into(),
)
.into());
};
let evict_start = std::time::Instant::now();
inner.remove_page_after_stream_sync(evicted_key);
if !queue.wait_until_idle(DEFERRED_RELEASE_WAIT_TIMEOUT) {
return Err(WeightHandleError::DeviceBinding(format!(
"weight eviction did not settle its deferred release within \
{DEFERRED_RELEASE_WAIT_TIMEOUT:?}; admission cannot claim the mapped or \
physical headroom yet"
))
.into());
}
add_duration(&GLOBAL_ADMIT_SYNC_NS, evict_start.elapsed());
evictions += 1;
let after_owned = physical
.allocator
.physical_pool_stats()
.map_or(0, |stats| stats.snapshot().total_owned_bytes);
let after_required_owned = physical
.allocator
.incremental_owned_bytes_for_span(ptr, len, 0, len)
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))?;
let after_required_mapped = physical
.allocator
.incremental_mapped_bytes_for_span(ptr, len, 0, len)
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))?;
if !eviction_made_committed_progress(
before_owned,
after_owned,
before_required_owned,
after_required_owned,
before_required_mapped,
after_required_mapped,
) {
inner.admission_no_progress = inner.admission_no_progress.saturating_add(1);
}
}
}
#[allow(clippy::too_many_arguments)]
fn rollback_failed_vmm_fill(
&self,
physical: &PhysicalAdmission,
allowance: &onnx_runtime_memory_governor::MappedAllowance,
ptr: NonNull<u8>,
len: usize,
has_stable_slot: bool,
slot_state: Option<Arc<SlotOperationState>>,
failure: VmmFillFailure,
) -> SpanAdmitError {
if let Some(source) = failure.in_flight_source {
if let Some(state) = slot_state.as_ref() {
state.poison();
}
quarantine_in_flight_fill(Box::new((
Arc::clone(&physical.allocator),
allowance.clone(),
ptr.as_ptr() as usize,
len,
slot_state,
source,
)));
return SpanAdmitError {
error: WeightHandleError::DeviceBinding(format!(
"{}; VMM fill rollback: copy-stream completion could not be established; \
the destination mapping and staging source remain charged and quarantined",
failure.error
)),
fresh_span: FreshSpanCleanup::Consumed,
};
}
let cleanup = if has_stable_slot {
match physical
.allocator
.decommit_allocation_range_outcome(ptr, len, 0, len)
{
Ok(crate::vmm_allocator::DecommitOutcome::Complete { accounting }) => {
WeightReleaseAction::refund(allowance, accounting.unmapped_bytes);
if accounting.quarantined_owned_bytes == 0 {
format!("rolled back {} mapped byte(s)", accounting.unmapped_bytes)
} else {
format!(
"unmapped {} byte(s), but quarantined {} byte(s) of residual physical \
ownership",
accounting.unmapped_bytes, accounting.quarantined_owned_bytes
)
}
}
Ok(crate::vmm_allocator::DecommitOutcome::RolledBack { reason }) => {
if let Some(state) = slot_state.as_ref() {
state.poison();
}
format!(
"cleanup rolled back to the partially filled mapping ({reason}); the \
stable slot was poisoned and remains charged"
)
}
Ok(crate::vmm_allocator::DecommitOutcome::Quarantined {
accounting,
residual,
reason,
}) => {
WeightReleaseAction::refund(allowance, accounting.unmapped_bytes);
if let Some(state) = slot_state.as_ref() {
state.poison();
}
format!(
"cleanup quarantined the stable slot ({reason}); refunded {} unmapped \
byte(s) and retained {} byte(s) at {:#x}",
accounting.unmapped_bytes, residual.retained_bytes, residual.address
)
}
Err(cleanup_error) => {
if let Some(state) = slot_state.as_ref() {
state.poison();
}
format!(
"cleanup was refused ({cleanup_error}); the stable slot was poisoned and \
remains charged"
)
}
}
} else {
match physical.allocator.deallocate_span_outcome(ptr) {
onnx_runtime_memory_governor::AllocationReleaseOutcome::Complete { accounting } => {
WeightReleaseAction::refund(allowance, accounting.unmapped_bytes);
format!(
"released the fresh span and refunded {} mapped byte(s)",
accounting.unmapped_bytes
)
}
onnx_runtime_memory_governor::AllocationReleaseOutcome::Quarantined {
accounting,
residual,
} => {
WeightReleaseAction::refund(allowance, accounting.unmapped_bytes);
format!(
"quarantined the fresh span; refunded {} unmapped byte(s) and retained {} \
byte(s) at {:#x}: {}",
accounting.unmapped_bytes,
residual.retained_bytes,
residual.address,
residual.reason
)
}
onnx_runtime_memory_governor::AllocationReleaseOutcome::Failed { failure } => {
format!(
"cleanup failed before a release outcome was established ({failure}); the \
allocator retains the span and its charge"
)
}
}
};
SpanAdmitError {
error: WeightHandleError::DeviceBinding(format!(
"{}; VMM fill rollback: {cleanup}",
failure.error
)),
fresh_span: FreshSpanCleanup::Consumed,
}
}
fn eviction_for(&self, boundary: LazyWeightBoundary) -> WeightEvictionPolicy {
eviction_for_boundary(self.scan_resistant_dense, boundary)
}
fn get_hit(&self, key: u64) -> Option<Arc<CudaWeightPage>> {
let mut inner = self.lock();
if let Some(page) = inner.pages.get(&key).cloned() {
if inner
.slots
.get(&key)
.is_some_and(|slot| slot.state.status() != SlotStatus::Idle)
{
return None;
}
let is_pinned = inner.pinned.contains(&key);
if is_pinned && pin_refill_every().is_some() {
return None;
}
inner.record_hit(key);
if is_pinned && pin_checksum_keys().is_some_and(|keys| keys.contains(&key)) {
let step = {
let counter = inner.pin_hit_step.entry(key).or_insert(0);
*counter += 1;
*counter
};
let baseline = inner.pin_granule_hashes.get(&key).cloned();
let (ptr, len) = (page.ptr, page.len);
drop(inner);
if let Some(baseline) = baseline {
match self.hash_page_granules(ptr, len) {
Ok(current) => report_pin_granule_diff(key, step, &baseline, ¤t),
Err(error) => eprintln!(
"weight_pin_checksum[#945]: key={key} step={step} readback failed: {error}"
),
}
}
return Some(page);
}
Some(page)
} else {
None
}
}
fn hash_page_granules(
&self,
ptr: CUdeviceptr,
len: usize,
) -> Result<Vec<u64>, WeightHandleError> {
self.runtime.drain_for_unmap().map_err(|error| {
WeightHandleError::DeviceBinding(format!("pin-checksum compute drain: {error}"))
})?;
self.runtime.copy_stream().synchronize().map_err(|error| {
WeightHandleError::DeviceBinding(format!("pin-checksum copy drain: {error}"))
})?;
let mut host = vec![0u8; len];
unsafe { self.runtime.dtoh(&mut host, ptr) }.map_err(|error| {
WeightHandleError::DeviceBinding(format!("pin-checksum dtoh: {error}"))
})?;
Ok(host.chunks(PIN_PROBE_GRANULE_BYTES).map(fnv1a_64).collect())
}
fn admit(
&self,
key: u64,
page: Arc<CudaWeightPage>,
eviction: WeightEvictionPolicy,
) -> Result<Arc<CudaWeightPage>, WeightHandleError> {
let bytes = page.len() as u64;
{
let mut inner = self.lock();
if let Some(existing) = inner.pages.get(&key).cloned() {
inner.record_hit(key);
drop(inner);
return Ok(existing);
}
if inner.policy.can_fit(bytes) {
inner.insert_page(key, Arc::clone(&page), bytes);
return Ok(page);
}
}
let mut inner = self.lock();
if let Some(existing) = inner.pages.get(&key).cloned() {
inner.record_hit(key);
drop(inner);
return Ok(existing);
}
if eviction == WeightEvictionPolicy::StableResident && !inner.policy.can_fit(bytes) {
inner.record_bypassed_page_in(key, bytes);
return Ok(page);
}
inner.evict_to_fit(bytes, eviction);
if let Some(over) = inner
.policy
.resident_bytes
.saturating_add(bytes)
.checked_sub(inner.policy.budget)
.filter(|over| *over > 0)
{
match inner.lease.as_mut() {
Some(lease) => {
lease.grow(over).map_err(|error| {
WeightHandleError::DeviceBinding(format!(
"the weight-residency cache needs {over} bytes beyond its \
{} byte budget for a {bytes} byte page, and eviction could not \
free them: {error}",
inner.policy.budget
))
})?;
inner.policy.budget = inner.policy.budget.saturating_add(over);
replace_global_budget(inner.policy.budget - over, inner.policy.budget);
}
None => {
inner.policy.budget = inner.policy.budget.saturating_add(over);
replace_global_budget(inner.policy.budget - over, inner.policy.budget);
}
}
}
inner.insert_page(key, Arc::clone(&page), bytes);
Ok(page)
}
pub fn stats(&self) -> CudaResidencyStats {
let inner = self.lock();
let physical_owned_bytes = self
.physical
.get()
.and_then(|physical| physical.allocator.physical_pool_stats())
.map_or(inner.policy.resident_bytes, |stats| {
stats.snapshot().total_owned_bytes
});
let mapped_physical_bytes = inner
.mapped_allowance
.as_ref()
.map_or(inner.policy.resident_bytes, |allowance| {
allowance.mapped_bytes()
});
CudaResidencyStats {
budget_bytes: inner.policy.budget,
resident_bytes: inner.policy.resident_bytes,
peak_resident_bytes: inner.policy.peak_resident_bytes,
pages_resident: inner.pages.len() as u64,
page_ins: inner.policy.page_ins,
hits: inner.policy.hits,
evictions: inner.policy.evictions,
physical_owned_bytes,
mapped_physical_bytes,
admission_no_progress: inner.admission_no_progress,
prefetch_issued: inner.prefetch_stats.issued,
prefetch_issued_bytes: inner.prefetch_stats.issued_bytes,
prefetch_promoted: inner.prefetch_stats.promoted,
prefetch_promote_wait_ns: inner.prefetch_stats.promote_wait_ns,
prefetch_declined_budget: inner.prefetch_stats.declined_budget,
prefetch_declined_busy: inner.prefetch_stats.declined_busy,
prefetch_declined_unsupported: inner.prefetch_stats.declined_unsupported,
prefetch_declined_resident: inner.prefetch_stats.declined_resident,
prefetch_declined_pool_capacity: inner.prefetch_stats.declined_pool_capacity,
pinned_pool_alloc_calls: self.staging_pool.alloc_calls(),
pinned_pool_reuses: self.staging_pool.reuses(),
}
}
fn lock(&self) -> std::sync::MutexGuard<'_, ResidencyInner> {
self.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
pub fn stable_va_paging_active(&self) -> bool {
self.physical.get().is_some()
}
pub fn prefill_double_buffer(
&self,
layer_bytes: u64,
) -> Result<PrefillDoubleBuffer<CudaPrefillTransfer>, PrefillReject<CudaPrefillError>> {
if !prefill_double_buffer_enabled() {
return Err(PrefillReject::Disabled);
}
self.build_prefill_double_buffer(layer_bytes)
}
#[doc(hidden)]
pub fn build_prefill_double_buffer(
&self,
layer_bytes: u64,
) -> Result<PrefillDoubleBuffer<CudaPrefillTransfer>, PrefillReject<CudaPrefillError>> {
let transfer =
CudaPrefillTransfer::new(Arc::clone(&self.runtime), Arc::clone(&self.staging_pool));
PrefillDoubleBuffer::new(transfer, layer_bytes)
}
#[doc(hidden)]
#[must_use]
pub fn with_prefill_double_buffer_enabled(mut self, enabled: bool) -> Self {
self.prefill_double_buffer_enabled = enabled;
self
}
pub fn prefill_pipeline_prefetch(
&self,
key: u64,
weight: &LazyWeight,
source: &dyn MmapRegionSource,
) -> PrefillRoute {
if !self.prefill_double_buffer_enabled {
return PrefillRoute::Declined(PrefillWireDecline::Disabled);
}
if weight.boundary != LazyWeightBoundary::BlockQuantizedMoe {
return PrefillRoute::Declined(PrefillWireDecline::Boundary);
}
if self.physical.get().is_some() || self.zero_copy_hybrid {
return PrefillRoute::Declined(PrefillWireDecline::Unsupported);
}
let bytes = weight.region_bytes_len() as u64;
if bytes == 0 {
return PrefillRoute::Declined(PrefillWireDecline::Empty);
}
let mut guard = self
.prefill_pipeline
.lock()
.expect("prefill pipeline lock poisoned");
if guard.is_none() {
match self.build_prefill_double_buffer(bytes) {
Ok(inner) => {
*guard = Some(PrefillPipeline {
inner,
tickets: HashMap::new(),
stats: PrefillWireStats::default(),
});
}
Err(reject) => {
return PrefillRoute::Declined(PrefillWireDecline::Build(reject));
}
}
}
let pipeline = guard.as_mut().expect("pipeline built above");
let capacity = pipeline.inner.layer_bytes();
if bytes > capacity {
pipeline.stats.declined_oversize += 1;
return PrefillRoute::Declined(PrefillWireDecline::Oversize {
layer_bytes: bytes,
capacity,
});
}
let req = PrefillLayerRequest::new(weight, source);
match pipeline.inner.prefetch(key, &req) {
Ok(ticket) => {
pipeline.tickets.insert(key, ticket);
pipeline.stats.routed += 1;
PrefillRoute::Prefetched
}
Err(PrefillReject::SlotsBusy) => {
pipeline.stats.declined_slots_busy += 1;
PrefillRoute::Declined(PrefillWireDecline::Prefetch(PrefillReject::SlotsBusy))
}
Err(reject) => {
pipeline.stats.declined_prefetch += 1;
PrefillRoute::Declined(PrefillWireDecline::Prefetch(reject))
}
}
}
pub fn prefill_pipeline_page(
self: &Arc<Self>,
key: u64,
device: DeviceId,
) -> Result<Option<PagedWeight>, WeightHandleError> {
if !self.prefill_double_buffer_enabled {
return Ok(None);
}
let mut guard = self
.prefill_pipeline
.lock()
.expect("prefill pipeline lock poisoned");
let Some(pipeline) = guard.as_mut() else {
return Ok(None);
};
let Some(ticket) = pipeline.tickets.remove(&key) else {
return Ok(None);
};
let view = match pipeline.inner.wait(&ticket) {
Ok(view) => view,
Err(reject) => {
pipeline.stats.wait_failed += 1;
return Err(WeightHandleError::DeviceBinding(format!(
"prefill double-buffer wait for key {key}: {reject:?}"
)));
}
};
pipeline.stats.consumed += 1;
let device_ptr = raw_ptr(view.device_ptr) as *const std::ffi::c_void;
let len = view.len;
drop(guard);
let keep_alive: Arc<dyn Any + Send + Sync> = Arc::new(PrefillSlotGuard {
residency: Arc::clone(self),
ticket: Some(ticket),
});
Ok(Some(PagedWeight::new(device_ptr, device, len, keep_alive)))
}
fn prefill_pipeline_release(&self, ticket: LayerTicket) {
if let Ok(mut guard) = self.prefill_pipeline.lock()
&& let Some(pipeline) = guard.as_mut()
&& pipeline.inner.release(ticket).is_ok()
{
pipeline.stats.released += 1;
}
}
pub fn prefill_pipeline_stats(&self) -> Option<PrefillWireStats> {
self.prefill_pipeline
.lock()
.ok()?
.as_ref()
.map(|pipeline| pipeline.stats)
}
pub fn prefill_pipeline_metrics(&self) -> Option<crate::prefill_double_buffer::PrefillMetrics> {
self.prefill_pipeline
.lock()
.ok()?
.as_ref()
.map(|pipeline| pipeline.inner.metrics())
}
pub fn prefill_pipeline_quarantined(&self) -> Option<usize> {
self.prefill_pipeline
.lock()
.ok()?
.as_ref()
.map(|pipeline| pipeline.inner.transfer().quarantined_len())
}
pub fn prefill_pipeline_active(&self) -> bool {
self.prefill_pipeline
.lock()
.map(|guard| guard.is_some())
.unwrap_or(false)
}
}
struct PrefillPipeline {
inner: PrefillDoubleBuffer<CudaPrefillTransfer>,
tickets: HashMap<u64, LayerTicket>,
stats: PrefillWireStats,
}
struct PrefillSlotGuard {
residency: Arc<CudaWeightResidency>,
ticket: Option<LayerTicket>,
}
impl Drop for PrefillSlotGuard {
fn drop(&mut self) {
if let Some(ticket) = self.ticket.take() {
self.residency.prefill_pipeline_release(ticket);
}
}
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct PrefillWireStats {
pub routed: u64,
pub consumed: u64,
pub released: u64,
pub declined_oversize: u64,
pub declined_slots_busy: u64,
pub declined_prefetch: u64,
pub wait_failed: u64,
}
#[derive(Debug)]
pub enum PrefillWireDecline {
Disabled,
Boundary,
Unsupported,
Empty,
Oversize {
layer_bytes: u64,
capacity: u64,
},
Build(PrefillReject<CudaPrefillError>),
Prefetch(PrefillReject<CudaPrefillError>),
}
#[derive(Debug)]
pub enum PrefillRoute {
Prefetched,
Declined(PrefillWireDecline),
}
impl std::fmt::Debug for CudaWeightResidency {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("CudaWeightResidency")
.field("stats", &self.stats())
.finish_non_exhaustive()
}
}
impl onnx_runtime_memory_governor::ReclaimableMappedHolder for CudaWeightResidency {
fn allowance(&self) -> onnx_runtime_memory_governor::MappedAllowance {
self.lock()
.mapped_allowance
.clone()
.expect("registered VMM weight residency must have a mapped allowance")
}
fn reclaim_priority(&self) -> u32 {
0
}
fn mapped_bytes(&self) -> u64 {
self.lock()
.mapped_allowance
.as_ref()
.map_or(0, |allowance| allowance.mapped_bytes())
}
fn reclaim_mapped(
&self,
target_bytes: u64,
) -> Result<
onnx_runtime_memory_governor::MappedReclaimReport,
onnx_runtime_memory_governor::MemoryError,
> {
let allowance = self.allowance();
let before = allowance.mapped_bytes();
let mut inner = self.lock();
let max_attempts = inner.pages.len();
let mut attempts = 0usize;
while before.saturating_sub(allowance.mapped_bytes()) < target_bytes
&& attempts < max_attempts
{
let Some(queue) = self.queue.as_ref().cloned() else {
return Err(onnx_runtime_memory_governor::MemoryError::InvalidRequest {
tier: Tier::Device.name(),
requested: target_bytes,
reason: "mapped weight reclaim requires the provider deferred-release queue",
});
};
let Some(key) = inner.next_evictable_key(WeightEvictionPolicy::Lru) else {
break;
};
inner.remove_page(key);
attempts += 1;
if !queue.wait_until_idle(DEFERRED_RELEASE_WAIT_TIMEOUT) {
let reclaimed = before.saturating_sub(allowance.mapped_bytes());
return Err(
onnx_runtime_memory_governor::MemoryError::CapacityUnavailable {
tier: Tier::Device.name(),
requested: target_bytes,
available: reclaimed,
role: allowance.role(),
detail: format!(
"deferred weight reclaim did not settle within \
{DEFERRED_RELEASE_WAIT_TIMEOUT:?}; mapped capacity remains charged"
),
source: None,
},
);
}
}
let reclaimed = before.saturating_sub(allowance.mapped_bytes());
Ok(onnx_runtime_memory_governor::MappedReclaimReport {
target_bytes,
reclaimed_bytes: reclaimed,
})
}
}
impl Drop for CudaWeightResidency {
fn drop(&mut self) {
let budget = self
.inner
.get_mut()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.policy
.budget;
let resident = self
.inner
.get_mut()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.policy
.resident_bytes;
replace_global_budget(budget, 0);
let _ = GLOBAL_CONTENT_RESIDENT_BYTES.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|current| Some(current.saturating_sub(resident)),
);
}
}
impl ResidencyInner {
fn record_hit(&mut self, key: u64) {
self.policy.record_hit(key);
GLOBAL_HITS.fetch_add(1, Ordering::Relaxed);
if let Some(page) = self.pages.get(&key) {
GLOBAL_HIT_BYTES.fetch_add(page.len as u64, Ordering::Relaxed);
record_key_trace(key, page.len as u64, KeyTraceEvent::Hit);
}
}
fn insert_page(&mut self, key: u64, page: Arc<CudaWeightPage>, bytes: u64) {
record_key_trace(key, bytes, KeyTraceEvent::Retained);
self.pages.insert(key, page);
self.policy.insert_page(key, bytes);
let global_resident = GLOBAL_CONTENT_RESIDENT_BYTES
.fetch_add(bytes, Ordering::Relaxed)
.saturating_add(bytes);
GLOBAL_PEAK_RESIDENT_BYTES.fetch_max(global_resident, Ordering::Relaxed);
GLOBAL_PAGE_INS.fetch_add(1, Ordering::Relaxed);
}
fn record_bypassed_page_in(&mut self, key: u64, bytes: u64) {
record_key_trace(key, bytes, KeyTraceEvent::Bypass);
self.policy.record_page_in();
GLOBAL_PAGE_INS.fetch_add(1, Ordering::Relaxed);
GLOBAL_BYPASSED_PAGE_INS.fetch_add(1, Ordering::Relaxed);
GLOBAL_BYPASSED_PAGE_IN_BYTES.fetch_add(bytes, Ordering::Relaxed);
}
fn mark_pinned(&mut self, key: u64, bytes: u64) {
if self.pinned.insert(key) {
self.pinned_bytes = self.pinned_bytes.saturating_add(bytes);
GLOBAL_PINNED_KEYS.fetch_add(1, Ordering::Relaxed);
GLOBAL_PINNED_BYTES.fetch_add(bytes, Ordering::Relaxed);
}
}
fn next_evictable_key(&self, eviction: WeightEvictionPolicy) -> Option<u64> {
self.policy
.next_evictable_index(eviction, &mut |key| {
!self.pinned.contains(&key)
&& self.slot_is_idle(key)
&& self
.pages
.get(&key)
.is_some_and(|page| Arc::strong_count(page) == 1)
})
.map(|index| self.policy.order[index])
}
fn smallest_evictable(&self) -> Option<(u64, u64)> {
let mut best: Option<(u64, u64)> = None;
for (&key, &bytes) in &self.policy.bytes_by_key {
let evictable = !self.pinned.contains(&key)
&& self.slot_is_idle(key)
&& self
.pages
.get(&key)
.is_some_and(|page| Arc::strong_count(page) == 1);
if evictable && best.is_none_or(|(_, best_bytes)| bytes < best_bytes) {
best = Some((key, bytes));
}
}
best
}
fn evictable_key_by_probe(
&self,
probe: EvictOrderProbe,
eviction: WeightEvictionPolicy,
) -> Option<u64> {
let is_evictable = |key: u64| -> bool {
!self.pinned.contains(&key)
&& self.slot_is_idle(key)
&& self
.pages
.get(&key)
.is_some_and(|page| Arc::strong_count(page) == 1)
};
match probe {
EvictOrderProbe::Lru => self.next_evictable_key(eviction),
EvictOrderProbe::Mru => self
.policy
.order
.iter()
.rev()
.copied()
.find(|&k| is_evictable(k)),
EvictOrderProbe::Smallest => self.smallest_evictable().map(|(key, _)| key),
EvictOrderProbe::Largest => {
let mut best: Option<(u64, u64)> = None;
for (&key, &bytes) in &self.policy.bytes_by_key {
if is_evictable(key) && best.is_none_or(|(_, best_bytes)| bytes > best_bytes) {
best = Some((key, bytes));
}
}
best.map(|(key, _)| key)
}
}
}
fn remove_page(&mut self, key: u64) {
if self.pages.remove(&key).is_some()
&& let Some(bytes) = self.policy.remove_page(key)
{
let _ = GLOBAL_CONTENT_RESIDENT_BYTES.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|current| Some(current.saturating_sub(bytes)),
);
GLOBAL_EVICTIONS.fetch_add(1, Ordering::Relaxed);
}
}
fn remove_page_after_stream_sync(&mut self, key: u64) {
let Some(page) = self.pages.remove(&key) else {
return;
};
let bytes = self.policy.remove_page(key);
if let Ok(mut page) = Arc::try_unwrap(page) {
page.retire_after_stream_sync();
}
if let Some(bytes) = bytes {
let _ = GLOBAL_CONTENT_RESIDENT_BYTES.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|current| Some(current.saturating_sub(bytes)),
);
GLOBAL_EVICTIONS.fetch_add(1, Ordering::Relaxed);
}
}
fn evict_to_fit(&mut self, incoming: u64, eviction: WeightEvictionPolicy) {
let evicted = {
let pages = &self.pages;
let slots = &self.slots;
self.policy.evict_to_fit(incoming, eviction, |key| {
slots
.get(&key)
.is_none_or(|slot| slot.state.status() == SlotStatus::Idle)
&& pages
.get(&key)
.is_some_and(|page| Arc::strong_count(page) == 1)
})
};
for key in evicted {
if let Some(page) = self.pages.remove(&key) {
let bytes = page.len() as u64;
let _ = GLOBAL_CONTENT_RESIDENT_BYTES.fetch_update(
Ordering::Relaxed,
Ordering::Relaxed,
|current| Some(current.saturating_sub(bytes)),
);
GLOBAL_EVICTIONS.fetch_add(1, Ordering::Relaxed);
}
}
}
fn slot_is_idle(&self, key: u64) -> bool {
self.slots
.get(&key)
.is_none_or(|slot| slot.state.status() == SlotStatus::Idle)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::EnvVarGuard;
#[derive(Debug)]
struct CountingRetirementCleanup {
runs: Arc<AtomicU64>,
health: std::sync::Weak<RouteReservationHealth>,
}
impl RouteReservationRetirementCleanup for CountingRetirementCleanup {
fn schedule(self: Box<Self>) {
self.runs.fetch_add(1, Ordering::SeqCst);
if let Some(health) = self.health.upgrade() {
health.complete_retirement();
}
}
}
fn scoped_health(
executor: ExecutorInstanceId,
device: u32,
) -> (
Arc<RouteReservationHealth>,
Arc<RouteReservationRetirementCounters>,
) {
let counters = Arc::new(RouteReservationRetirementCounters::default());
(
RouteReservationHealth::new_scoped(
ExecutorArtifactProviderId::from_raw(1),
executor,
ExecutorArtifactGeneration::from_raw(1),
device,
Arc::clone(&counters),
)
.expect("test route generation"),
counters,
)
}
fn fresh_executor() -> ExecutorInstanceId {
static NEXT: AtomicU64 = AtomicU64::new(1);
ExecutorInstanceId::from_raw(NEXT.fetch_add(1, Ordering::Relaxed))
}
#[test]
fn route_reservation_generation_exhaustion_never_wraps_or_reuses_identity() {
let counter = AtomicU64::new(u64::MAX - 1);
assert_eq!(
next_route_reservation_generation(&counter),
Some(u64::MAX - 1)
);
assert_eq!(counter.load(Ordering::Acquire), u64::MAX);
assert_eq!(next_route_reservation_generation(&counter), None);
assert_eq!(next_route_reservation_generation(&counter), None);
assert_eq!(counter.load(Ordering::Acquire), u64::MAX);
}
#[test]
fn route_reservation_lifecycle_linearizes_use_transition_and_poison() {
let executor = fresh_executor();
let (health, _) = scoped_health(executor, 3);
let generation = health.generation().expect("scoped generation");
let use_guard = health
.acquire_use(executor, 3, generation)
.expect("exact owner acquires use");
let blocked = health
.begin_transition()
.err()
.expect("active use must block transition");
assert!(blocked.contains("dispatch/replay lease"));
drop(use_guard);
let transition = health
.begin_transition()
.expect("transition starts after use completes");
let blocked = health
.acquire_use(executor, 3, generation)
.err()
.expect("transition must block a new use");
assert!(blocked.contains("atomic group transition"));
transition.poison("incomplete logical expert group".to_string());
let poisoned = health
.acquire_use(executor, 3, generation)
.err()
.expect("poison is irreversible");
assert!(poisoned.contains(&format!("generation {generation}")));
assert!(poisoned.contains("incomplete logical expert group"));
}
#[test]
fn route_reservation_lifecycle_rejects_sibling_and_device_aliases() {
let owner = fresh_executor();
let sibling = fresh_executor();
let (health, _) = scoped_health(owner, 1);
let generation = health.generation().expect("scoped generation");
let sibling_error = health
.acquire_use(sibling, 1, generation)
.err()
.expect("sibling executor must not borrow owner's generation");
assert!(sibling_error.contains(&format!("executor {}", owner.get())));
assert!(sibling_error.contains(&format!("not executor {}", sibling.get())));
let device_error = health
.acquire_use(owner, 2, generation)
.err()
.expect("wrong device must not borrow owner's generation");
assert!(device_error.contains("CUDA:1"));
assert!(device_error.contains("CUDA:2"));
let generation_error = health
.acquire_use(owner, 1, generation + 1)
.err()
.expect("stale generation must not borrow a replacement reservation");
assert!(generation_error.contains(&format!("generation {generation}")));
assert!(generation_error.contains(&format!("generation {}", generation + 1)));
}
#[test]
fn route_reservation_requirement_rejects_foreign_provider_and_artifact_generation() {
let executor = fresh_executor();
let (health, _) = scoped_health(executor, 2);
let provider = ExecutorArtifactProviderId::from_raw(1);
let generation = ExecutorArtifactGeneration::from_raw(1);
health
.validate_artifact_scope(provider, executor, generation, 2)
.expect("exact private session scope");
let foreign_provider = health
.validate_artifact_scope(
ExecutorArtifactProviderId::from_raw(2),
executor,
generation,
2,
)
.expect_err("foreign provider cannot retain another owner's requirement");
assert!(foreign_provider.contains("provider 1"));
assert!(foreign_provider.contains("not provider 2"));
let stale_generation = health
.validate_artifact_scope(
provider,
executor,
ExecutorArtifactGeneration::from_raw(2),
2,
)
.expect_err("stale artifact generation cannot retain a replacement reservation");
assert!(stale_generation.contains("artifact generation 1"));
assert!(stale_generation.contains("artifact generation 2"));
}
#[test]
fn route_reservation_retirement_returns_with_public_holder_and_cleans_on_last_release() {
let executor = fresh_executor();
let (health, counters) = scoped_health(executor, 0);
let requirement = health
.requirement_state()
.expect("retained requirement state");
let holder = requirement.acquire_use().expect("pre-retirement holder");
let cleanups = Arc::new(AtomicU64::new(0));
let teardown_health = Arc::clone(&health);
let teardown_cleanups = Arc::clone(&cleanups);
let (returned_tx, returned_rx) = std::sync::mpsc::channel();
let teardown = std::thread::spawn(move || {
assert_eq!(
teardown_health.begin_retirement(),
RouteReservationRetirementStart::Started
);
teardown_health.install_retirement_cleanup(Box::new(CountingRetirementCleanup {
runs: teardown_cleanups,
health: Arc::downgrade(&teardown_health),
}));
returned_tx.send(()).unwrap();
});
returned_rx
.recv_timeout(Duration::from_secs(1))
.expect("retirement must not wait for a public use guard");
let rejection = requirement
.acquire_use()
.err()
.expect("retirement rejects later public acquisition");
assert!(rejection.to_string().contains("retiring"));
assert_eq!(
cleanups.load(Ordering::SeqCst),
0,
"cleanup must remain quarantined while the original lease is active"
);
let replay_launches = AtomicU64::new(0);
if requirement.acquire_use().is_ok() {
replay_launches.fetch_add(1, Ordering::SeqCst);
}
assert_eq!(
replay_launches.load(Ordering::SeqCst),
0,
"a replay racing retirement must not reach its launch"
);
drop(holder);
teardown.join().unwrap();
assert_eq!(cleanups.load(Ordering::SeqCst), 1);
let retired = requirement
.acquire_use()
.err()
.expect("completed teardown remains a requirement, but is retired");
assert!(retired.to_string().contains("retired"));
assert_eq!(
counters.snapshot(),
RouteReservationRetirementStats {
retirements_started: 1,
deferred_cleanups: 1,
cleanups_scheduled: 1,
cleanups_executed: 1,
}
);
}
#[test]
fn route_reservation_stalled_transition_defers_and_releases_cleanup_once() {
let executor = fresh_executor();
let (health, counters) = scoped_health(executor, 0);
let transition = health.begin_transition().expect("transition authority");
let cleanups = Arc::new(AtomicU64::new(0));
assert_eq!(
health.begin_retirement(),
RouteReservationRetirementStart::Started
);
health.install_retirement_cleanup(Box::new(CountingRetirementCleanup {
runs: Arc::clone(&cleanups),
health: Arc::downgrade(&health),
}));
assert_eq!(cleanups.load(Ordering::SeqCst), 0);
transition.complete();
assert_eq!(cleanups.load(Ordering::SeqCst), 1);
assert_eq!(counters.snapshot().deferred_cleanups, 1);
assert_eq!(counters.snapshot().cleanups_scheduled, 1);
}
#[test]
fn route_reservation_acquire_racing_retirement_never_launches_after_unmap() {
for _ in 0..64 {
let executor = fresh_executor();
let (health, _) = scoped_health(executor, 0);
let generation = health.generation().expect("scoped generation");
let gate = Arc::new(std::sync::Barrier::new(3));
let launches = Arc::new(AtomicU64::new(0));
let unmaps = Arc::new(AtomicU64::new(0));
let use_health = Arc::clone(&health);
let use_gate = Arc::clone(&gate);
let use_launches = Arc::clone(&launches);
let use_unmaps = Arc::clone(&unmaps);
let acquire = std::thread::spawn(move || {
use_gate.wait();
match use_health.acquire_use(executor, 0, generation) {
Ok(guard) => {
assert_eq!(
use_unmaps.load(Ordering::SeqCst),
0,
"a valid pre-retirement lease must launch before unmap"
);
use_launches.fetch_add(1, Ordering::SeqCst);
drop(guard);
true
}
Err(error) => {
assert!(error.contains("retiring") || error.contains("retired"));
false
}
}
});
let retire_health = Arc::clone(&health);
let retire_gate = Arc::clone(&gate);
let retire_unmaps = Arc::clone(&unmaps);
let retire = std::thread::spawn(move || {
retire_gate.wait();
assert_eq!(
retire_health.begin_retirement(),
RouteReservationRetirementStart::Started
);
retire_health.install_retirement_cleanup(Box::new(CountingRetirementCleanup {
runs: retire_unmaps,
health: Arc::downgrade(&retire_health),
}));
});
gate.wait();
let acquired = acquire.join().unwrap();
retire.join().unwrap();
assert_eq!(unmaps.load(Ordering::SeqCst), 1);
assert_eq!(launches.load(Ordering::SeqCst), u64::from(acquired));
assert!(
health
.acquire_use(executor, 0, generation)
.err()
.expect("completed retirement rejects acquisition")
.contains("retired")
);
}
}
#[test]
fn slot_operation_state_never_reopens_after_poison() {
let state = SlotOperationState::default();
assert!(state.begin_refill());
assert_eq!(state.status(), SlotStatus::Pending);
assert!(!state.begin_release());
state.poison();
state.finish_refill();
assert_eq!(state.status(), SlotStatus::Poisoned);
let completed = SlotOperationState::default();
assert!(completed.begin_refill());
completed.finish_refill();
assert_eq!(completed.status(), SlotStatus::Idle);
}
#[test]
fn reset_clears_window_counters_and_preserves_live_gauges() {
const SENTINEL: u64 = 1 << 40;
GLOBAL_PAGE_INS.fetch_add(SENTINEL, Ordering::Relaxed);
GLOBAL_HITS.fetch_add(SENTINEL, Ordering::Relaxed);
GLOBAL_EVICTIONS.fetch_add(SENTINEL, Ordering::Relaxed);
GLOBAL_BUDGET_BYTES.fetch_add(SENTINEL, Ordering::Relaxed);
GLOBAL_CONTENT_RESIDENT_BYTES.fetch_add(SENTINEL, Ordering::Relaxed);
GLOBAL_WEIGHT_MAPPED_BYTES.fetch_add(SENTINEL, Ordering::Relaxed);
let planted_peak = GLOBAL_CONTENT_RESIDENT_BYTES
.load(Ordering::Relaxed)
.saturating_add(SENTINEL);
GLOBAL_PEAK_RESIDENT_BYTES.fetch_max(planted_peak, Ordering::Relaxed);
reset_global_offload_stats();
let after = global_offload_stats();
assert!(
after.page_ins < SENTINEL,
"reset left page_ins at {}, above the planted sentinel",
after.page_ins
);
assert!(
after.hits < SENTINEL,
"reset left hits at {}, above the planted sentinel",
after.hits
);
assert!(
after.evictions < SENTINEL,
"reset left evictions at {}, above the planted sentinel",
after.evictions
);
assert!(
after.budget_bytes >= SENTINEL,
"reset dropped budget_bytes to {}, but it is a live gauge",
after.budget_bytes
);
assert!(
after.content_resident_bytes >= SENTINEL,
"reset dropped content_resident_bytes to {}, but it is a live gauge",
after.content_resident_bytes
);
assert!(
after.mapped_physical_bytes >= SENTINEL,
"reset dropped mapped_physical_bytes to {}, but it is a live gauge",
after.mapped_physical_bytes
);
assert!(
after.peak_resident_bytes >= planted_peak,
"reset wrote {} over a lifetime peak of {}; it must not write the \
peak at all",
after.peak_resident_bytes,
planted_peak
);
GLOBAL_BUDGET_BYTES.fetch_sub(SENTINEL, Ordering::Relaxed);
GLOBAL_CONTENT_RESIDENT_BYTES.fetch_sub(SENTINEL, Ordering::Relaxed);
GLOBAL_WEIGHT_MAPPED_BYTES.fetch_sub(SENTINEL, Ordering::Relaxed);
}
#[test]
fn device_policy_defaults_to_disabled() {
let policy = DeviceOffloadPolicy::default();
assert!(!policy.enabled);
assert_eq!(policy.device_budget_bytes, None);
assert!(!policy.async_pagein);
assert!(policy.scan_resistant_dense);
}
#[test]
fn async_pagein_env_is_default_on_with_explicit_opt_out() {
assert!(async_pagein_from_env_value(None));
assert!(async_pagein_from_env_value(Some("1")));
assert!(async_pagein_from_env_value(Some("true")));
assert!(async_pagein_from_env_value(Some("YES")));
assert!(async_pagein_from_env_value(Some(" On ")));
assert!(!async_pagein_from_env_value(Some("0")));
assert!(!async_pagein_from_env_value(Some("false")));
assert!(!async_pagein_from_env_value(Some("")));
assert!(!async_pagein_from_env_value(Some("maybe")));
}
#[test]
fn byte_hit_rate_diverges_from_the_count_based_rate() {
let stats = GlobalOffloadStats {
hits: 9,
page_ins: 1,
hit_bytes: 9 * 10 * 1024, htod_bytes: 12 * 1024 * 1024, ..GlobalOffloadStats::default()
};
let count_rate = stats.hits as f64 / (stats.hits + stats.page_ins) as f64;
let byte_rate = stats.byte_hit_rate().expect("bytes were requested");
assert!(
(count_rate - 0.90).abs() < 1e-9,
"count-based rate looks excellent: {count_rate}"
);
assert!(
byte_rate < 0.01,
"byte-weighted rate tells the truth about streaming cost: {byte_rate}"
);
}
#[test]
fn byte_hit_rate_is_none_when_no_bytes_were_requested() {
assert_eq!(GlobalOffloadStats::default().byte_hit_rate(), None);
}
#[test]
fn bypassed_byte_share_attributes_streamed_bytes() {
assert_eq!(GlobalOffloadStats::default().bypassed_byte_share(), None);
let stats = GlobalOffloadStats {
htod_bytes: 12 * 1024 * 1024,
bypassed_page_in_bytes: 3 * 1024 * 1024,
..GlobalOffloadStats::default()
};
let share = stats.bypassed_byte_share().expect("bytes were streamed");
assert!(
(share - 0.25).abs() < 1e-9,
"one quarter of stream: {share}"
);
}
#[test]
fn scan_resistant_env_defaults_on_with_lru_opt_out() {
assert!(scan_resistant_from_env_value(None));
assert!(scan_resistant_from_env_value(Some("1")));
assert!(scan_resistant_from_env_value(Some("true")));
assert!(scan_resistant_from_env_value(Some("YES")));
assert!(scan_resistant_from_env_value(Some(" On ")));
assert!(!scan_resistant_from_env_value(Some("0")));
assert!(!scan_resistant_from_env_value(Some("false")));
assert!(!scan_resistant_from_env_value(Some("NO")));
assert!(!scan_resistant_from_env_value(Some(" off ")));
assert!(scan_resistant_from_env_value(Some("")));
assert!(scan_resistant_from_env_value(Some("maybe")));
}
#[test]
fn byte_aware_env_defaults_off_and_opts_in() {
assert!(!byte_aware_from_env_value(None));
assert!(!byte_aware_from_env_value(Some("0")));
assert!(!byte_aware_from_env_value(Some("false")));
assert!(!byte_aware_from_env_value(Some("")));
assert!(!byte_aware_from_env_value(Some("maybe")));
assert!(byte_aware_from_env_value(Some("1")));
assert!(byte_aware_from_env_value(Some("true")));
assert!(byte_aware_from_env_value(Some("YES")));
assert!(byte_aware_from_env_value(Some(" On ")));
}
#[test]
fn zero_copy_hybrid_env_defaults_off_and_opts_in() {
assert!(!zero_copy_hybrid_from_env_value(None));
assert!(!zero_copy_hybrid_from_env_value(Some("0")));
assert!(!zero_copy_hybrid_from_env_value(Some("false")));
assert!(!zero_copy_hybrid_from_env_value(Some("")));
assert!(!zero_copy_hybrid_from_env_value(Some("maybe")));
assert!(zero_copy_hybrid_from_env_value(Some("1")));
assert!(zero_copy_hybrid_from_env_value(Some("true")));
assert!(zero_copy_hybrid_from_env_value(Some("YES")));
assert!(zero_copy_hybrid_from_env_value(Some(" On ")));
}
#[test]
fn zero_copy_safe_budget_defaults_are_platform_aware() {
let wddm = ZERO_COPY_SAFE_BUDGET_BYTES_WDDM;
let non_windows = ZERO_COPY_SAFE_BUDGET_BYTES_NON_WINDOWS;
let selected = ZERO_COPY_SAFE_BUDGET_BYTES;
let wddm_observed_safe_bytes: u64 = 436_633_600; assert_eq!(wddm, 256 * 1024 * 1024);
assert!(
wddm < wddm_observed_safe_bytes,
"WDDM default {wddm} must stay under the WDDM ceiling"
);
let linux_measured_safe_bytes: u64 = 6_795_458_560; assert_eq!(non_windows, 2 * 1024 * 1024 * 1024);
assert!(
non_windows < linux_measured_safe_bytes,
"non-Windows default {non_windows} must stay under the #925 measured-safe ceiling"
);
assert!(
non_windows > wddm_observed_safe_bytes,
"non-Windows default must clear the WDDM corruption band to unlock the lever"
);
#[cfg(target_os = "windows")]
assert_eq!(selected, wddm);
#[cfg(not(target_os = "windows"))]
assert_eq!(selected, non_windows);
}
#[test]
fn numeric_env_keeps_unset_and_unparseable_distinct() {
const NAME: &str = "ONNX_GENAI_TEST_NUMERIC_ENV_PROBE";
unsafe {
std::env::remove_var(NAME);
}
assert_eq!(parse_numeric_env(NAME), NumericEnv::Unset);
unsafe {
std::env::set_var(NAME, " 1073741824 ");
}
assert_eq!(
parse_numeric_env(NAME),
NumericEnv::Value(1_073_741_824),
"a plain integer, surrounding whitespace included, must be honoured"
);
for bad in ["2GB", "1_073_741_824", "0x10", "", "1.5", "-1"] {
unsafe {
std::env::set_var(NAME, bad);
}
assert_eq!(
parse_numeric_env(NAME),
NumericEnv::Invalid(bad.to_string()),
"{bad:?} must be reported as supplied-but-unusable, not as absent"
);
}
assert_eq!(parse_numeric_env(NAME).or_default(NAME, 4096), 4096);
assert_eq!(parse_numeric_env(NAME).into_option(NAME), None);
unsafe {
std::env::remove_var(NAME);
}
assert_eq!(parse_numeric_env(NAME).or_default(NAME, 4096), 4096);
}
#[test]
fn evict_order_env_defaults_lru_and_parses_variants() {
assert_eq!(evict_order_from_env_value(None), EvictOrderProbe::Lru);
assert_eq!(evict_order_from_env_value(Some("")), EvictOrderProbe::Lru);
assert_eq!(
evict_order_from_env_value(Some("lru")),
EvictOrderProbe::Lru
);
assert_eq!(
evict_order_from_env_value(Some("nonsense")),
EvictOrderProbe::Lru
);
assert_eq!(
evict_order_from_env_value(Some("mru")),
EvictOrderProbe::Mru
);
assert_eq!(
evict_order_from_env_value(Some(" Reverse ")),
EvictOrderProbe::Mru
);
assert_eq!(
evict_order_from_env_value(Some("SMALLEST")),
EvictOrderProbe::Smallest
);
assert_eq!(
evict_order_from_env_value(Some("large")),
EvictOrderProbe::Largest
);
}
#[test]
fn budget_parsing_rejects_zero_and_garbage() {
assert_eq!(parse_budget_bytes("1048576"), Some(1_048_576));
assert_eq!(parse_budget_bytes(" 4096 "), Some(4096));
assert_eq!(parse_budget_bytes("0"), None);
assert_eq!(parse_budget_bytes(""), None);
assert_eq!(parse_budget_bytes("lots"), None);
assert_eq!(parse_budget_bytes("-5"), None);
}
#[test]
fn vmm_admission_requires_both_zone_and_global_headroom() {
assert!(committed_admission_fits(0, 0, 0, 0));
assert!(
!committed_admission_fits(2 << 20, 0, 0, 8 << 30),
"pooled ownership cannot bypass a full weight zone"
);
assert!(
!committed_admission_fits(0, 8 << 30, 2 << 20, 742 << 10),
"zone room cannot bypass missing global creation headroom"
);
}
#[test]
fn zero_committed_byte_eviction_is_observable_no_progress() {
assert!(!eviction_made_committed_progress(8, 8, 4, 4, 4, 4));
assert!(
eviction_made_committed_progress(8, 8, 4, 0, 4, 4),
"returning an owned handle to the pool is useful even though owned bytes stay flat"
);
assert!(eviction_made_committed_progress(8, 4, 4, 4, 4, 4));
assert!(eviction_made_committed_progress(8, 8, 4, 4, 4, 0));
}
#[cfg_attr(
not(feature = "gpu-tests"),
ignore = "requires CUDA device; enable the gpu-tests feature on a CUDA runner"
)]
#[test]
fn failed_vmm_weight_fills_refund_fresh_and_reused_slot_charges() {
use onnx_runtime_memory_governor::{
DeviceKey, HolderId, LeaseLedger, LedgerGovernor, MemoryGovernor, MemoryRole,
};
let mut env = EnvVarGuard::acquire();
env.unset(crate::vmm_allocator::CUDA_PHYSICAL_HANDLE_POOL_BYTES_ENV);
let Ok(runtime) = CudaRuntime::new(0).map(Arc::new) else {
eprintln!("SKIPPED (CUDA runtime dependencies unavailable): VMM fill rollback test");
return;
};
let granule = 2usize << 20;
let governor = Arc::new(LedgerGovernor::new(LeaseLedger::new(
(granule * 2) as u64,
0,
0,
)));
let allocator = Arc::new(
crate::vmm_allocator::CudaVmmAllocator::new(
runtime.cuda_context(),
DeviceKey::device(0),
0,
64 << 20,
governor.as_ref(),
HolderId::new(738),
MemoryRole::Weights,
)
.expect("no-pool VMM allocator"),
);
let authority: Arc<dyn MemoryGovernor + Send + Sync> = governor.clone();
let release_queue = CudaDeferredReleaseQueue::new(
Box::new(crate::deferred_release::CudaStreamFences::new(Arc::clone(
&runtime,
))),
crate::deferred_release::DEFAULT_DEFERRED_RELEASE_CAPACITY,
);
let residency = CudaWeightResidency::new(Arc::clone(&runtime), granule as u64)
.with_deferred_release_queue(Arc::clone(&release_queue))
.with_vmm_admission(Arc::clone(&allocator), authority)
.expect("install VMM admission");
residency
.adopt_governed_budget(governor.as_ref(), Tier::Device, HolderId::new(738))
.expect("reserve mapped allowance");
let global_before = GLOBAL_WEIGHT_MAPPED_BYTES.load(Ordering::Relaxed);
let fresh_error = residency
.resident_vmm_with(
1,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
|_, _| {
Err(WeightHandleError::DeviceBinding(
"injected H2D fill failure".into(),
))
},
)
.expect_err("fresh fill failure");
assert!(fresh_error.to_string().contains("released the fresh span"));
assert_eq!(residency.stats().mapped_physical_bytes, 0);
assert_eq!(governor.used(Tier::Device), 0);
assert_eq!(
GLOBAL_WEIGHT_MAPPED_BYTES.load(Ordering::Relaxed),
global_before
);
let allocations: Vec<usize> = (0..2)
.map(|_| {
let allocator = Arc::clone(&allocator);
std::thread::spawn(move || {
allocator
.allocate_committed(granule, WEIGHT_SLOT_ALIGN, &[])
.expect("post-rollback allocation")
.as_ptr() as usize
})
})
.map(|thread| thread.join().expect("allocation thread panicked"))
.collect();
assert_ne!(
allocations[0], allocations[1],
"fresh rollback double-released one VA into the arena free list"
);
for address in allocations {
let ptr = NonNull::new(address as *mut u8).expect("non-null allocation");
let outcome = allocator.deallocate_span_outcome(ptr);
assert!(
matches!(
outcome,
onnx_runtime_memory_governor::AllocationReleaseOutcome::Complete { .. }
),
"post-rollback probe span must release exactly once: {outcome:?}"
);
}
let first_bytes = vec![0x31u8; granule];
let first = residency
.resident_vmm_with(
1,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |runtime, ptr| {
unsafe { runtime.htod(&first_bytes, ptr) }
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))
},
)
.map(VmmAdmit::expect_page)
.expect("first stable slot");
drop(first);
let second_bytes = vec![0x42u8; granule];
let second = residency
.resident_vmm_with(
2,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |runtime, ptr| {
unsafe { runtime.htod(&second_bytes, ptr) }
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))
},
)
.map(VmmAdmit::expect_page)
.expect("evict first slot");
assert!(
release_queue.wait_until_idle(DEFERRED_RELEASE_WAIT_TIMEOUT),
"first stable-slot decommit"
);
drop(second);
residency.lock().remove_page_after_stream_sync(2);
assert!(
release_queue.wait_until_idle(DEFERRED_RELEASE_WAIT_TIMEOUT),
"second stable-slot decommit"
);
let mapped_baseline = residency.stats().mapped_physical_bytes;
let charged_baseline = governor.used(Tier::Device);
let global_baseline = GLOBAL_WEIGHT_MAPPED_BYTES.load(Ordering::Relaxed);
let reused_error = residency
.resident_vmm_with(
1,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
|_, _| {
Err(WeightHandleError::DeviceBinding(
"injected synchronization failure".into(),
))
},
)
.expect_err("reused-slot fill failure");
assert!(reused_error.to_string().contains("rolled back"));
assert_eq!(residency.stats().mapped_physical_bytes, mapped_baseline);
assert_eq!(governor.used(Tier::Device), charged_baseline);
assert_eq!(
GLOBAL_WEIGHT_MAPPED_BYTES.load(Ordering::Relaxed),
global_baseline
);
}
#[cfg_attr(
not(feature = "gpu-tests"),
ignore = "requires CUDA device; enable the gpu-tests feature on a CUDA runner"
)]
#[test]
fn possible_in_flight_vmm_fills_quarantine_fresh_and_reused_destinations() {
use onnx_runtime_memory_governor::{
DeviceKey, HolderId, LeaseLedger, LedgerGovernor, MemoryGovernor, MemoryRole,
};
#[derive(Debug)]
struct DropProbe(Arc<AtomicBool>);
impl Drop for DropProbe {
fn drop(&mut self) {
self.0.store(true, Ordering::Release);
}
}
let mut env = EnvVarGuard::acquire();
env.unset(crate::vmm_allocator::CUDA_PHYSICAL_HANDLE_POOL_BYTES_ENV);
let Ok(runtime) = CudaRuntime::new(0).map(Arc::new) else {
eprintln!("SKIPPED (CUDA runtime dependencies unavailable): in-flight fill test");
return;
};
let granule = 2usize << 20;
let governor = Arc::new(LedgerGovernor::new(LeaseLedger::new(
(granule * 2) as u64,
0,
0,
)));
let allocator = Arc::new(
crate::vmm_allocator::CudaVmmAllocator::new(
runtime.cuda_context(),
DeviceKey::device(0),
0,
64 << 20,
governor.as_ref(),
HolderId::new(740),
MemoryRole::Weights,
)
.expect("no-pool VMM allocator"),
);
let authority: Arc<dyn MemoryGovernor + Send + Sync> = governor.clone();
let release_queue = CudaDeferredReleaseQueue::new(
Box::new(crate::deferred_release::CudaStreamFences::new(Arc::clone(
&runtime,
))),
crate::deferred_release::DEFAULT_DEFERRED_RELEASE_CAPACITY,
);
let residency = CudaWeightResidency::new(Arc::clone(&runtime), (granule * 2) as u64)
.with_deferred_release_queue(Arc::clone(&release_queue))
.with_vmm_admission(Arc::clone(&allocator), authority)
.expect("install VMM admission");
residency
.adopt_governed_budget(governor.as_ref(), Tier::Device, HolderId::new(740))
.expect("reserve mapped allowance");
let quarantine_before = in_flight_fill_quarantine_count();
let global_before = GLOBAL_WEIGHT_MAPPED_BYTES.load(Ordering::Relaxed);
let fresh_dropped = Arc::new(AtomicBool::new(false));
let fresh_probe = Arc::clone(&fresh_dropped);
let fresh_error = residency
.resident_vmm_with(
90,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |_, _| -> Result<(), VmmFillFailure> {
Err(VmmFillFailure::may_be_in_flight(
WeightHandleError::DeviceBinding(
"injected end-event record and copy-stream sync failure".into(),
),
Box::new(DropProbe(fresh_probe)),
))
},
)
.expect_err("fresh destination must be quarantined");
assert!(
fresh_error
.to_string()
.contains("remain charged and quarantined")
);
assert!(!fresh_dropped.load(Ordering::Acquire));
assert_eq!(governor.used(Tier::Device), granule as u64);
assert_eq!(residency.stats().mapped_physical_bytes, granule as u64);
let bytes = vec![0x51u8; granule];
let page = residency
.resident_vmm_with(
91,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |runtime, ptr| {
unsafe { runtime.htod(&bytes, ptr) }
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))
},
)
.map(VmmAdmit::expect_page)
.expect("create stable slot for reuse");
drop(page);
residency.lock().remove_page_after_stream_sync(91);
assert!(
release_queue.wait_until_idle(DEFERRED_RELEASE_WAIT_TIMEOUT),
"stable slot must become reusable"
);
let reused_dropped = Arc::new(AtomicBool::new(false));
let reused_probe = Arc::clone(&reused_dropped);
let reused_error = residency
.resident_vmm_with(
91,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |_, _| -> Result<(), VmmFillFailure> {
Err(VmmFillFailure::may_be_in_flight(
WeightHandleError::DeviceBinding(
"injected reused-slot event and synchronization failure".into(),
),
Box::new(DropProbe(reused_probe)),
))
},
)
.expect_err("reused destination must be quarantined");
assert!(
reused_error
.to_string()
.contains("remain charged and quarantined")
);
assert!(!reused_dropped.load(Ordering::Acquire));
assert_eq!(governor.used(Tier::Device), (granule * 2) as u64);
assert_eq!(
residency.stats().mapped_physical_bytes,
(granule * 2) as u64
);
assert_eq!(
residency
.lock()
.slots
.get(&91)
.expect("reused slot")
.state
.status(),
SlotStatus::Poisoned
);
assert_eq!(in_flight_fill_quarantine_count(), quarantine_before + 2);
assert_eq!(
GLOBAL_WEIGHT_MAPPED_BYTES.load(Ordering::Relaxed),
global_before + (granule * 2) as u64
);
}
#[cfg_attr(
not(feature = "gpu-tests"),
ignore = "requires CUDA device; enable the gpu-tests feature on a CUDA runner"
)]
#[test]
fn concurrent_pinned_refill_state_machine_quarantines_only_unresolved_copies() {
use onnx_runtime_memory_governor::{
DeviceKey, HolderId, LeaseLedger, LedgerGovernor, MemoryGovernor, MemoryRole,
};
#[derive(Debug)]
struct DropProbe(Arc<AtomicBool>);
impl Drop for DropProbe {
fn drop(&mut self) {
self.0.store(true, Ordering::Release);
}
}
let mut env = EnvVarGuard::acquire();
env.unset(crate::vmm_allocator::CUDA_PHYSICAL_HANDLE_POOL_BYTES_ENV)
.set(WEIGHT_PIN_REFILL_EVERY_ENV, "1");
let Ok(runtime) = CudaRuntime::new(0).map(Arc::new) else {
eprintln!("SKIPPED (CUDA runtime dependencies unavailable): pinned refill fault test");
return;
};
let granule = 2usize << 20;
let governor = Arc::new(LedgerGovernor::new(LeaseLedger::new(
(granule * 2) as u64,
0,
0,
)));
let allocator = Arc::new(
crate::vmm_allocator::CudaVmmAllocator::new(
runtime.cuda_context(),
DeviceKey::device(0),
0,
64 << 20,
governor.as_ref(),
HolderId::new(741),
MemoryRole::Weights,
)
.expect("no-pool VMM allocator"),
);
let authority: Arc<dyn MemoryGovernor + Send + Sync> = governor.clone();
let release_queue = CudaDeferredReleaseQueue::new(
Box::new(crate::deferred_release::CudaStreamFences::new(Arc::clone(
&runtime,
))),
crate::deferred_release::DEFAULT_DEFERRED_RELEASE_CAPACITY,
);
let residency = CudaWeightResidency::new(Arc::clone(&runtime), (granule * 2) as u64)
.with_deferred_release_queue(Arc::clone(&release_queue))
.with_vmm_admission(Arc::clone(&allocator), authority)
.expect("install VMM admission");
residency
.adopt_governed_budget(governor.as_ref(), Tier::Device, HolderId::new(741))
.expect("reserve mapped allowance");
let completed_bytes = vec![0x31u8; granule];
let completed_page = residency
.resident_vmm_with(
100,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
|runtime, ptr| {
unsafe { runtime.htod(&completed_bytes, ptr) }
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))
},
)
.map(VmmAdmit::expect_page)
.expect("admit completed-control page");
residency.lock().mark_pinned(100, granule as u64);
let completed_ptr = completed_page.ptr;
let held_fill_called = Arc::new(AtomicBool::new(false));
let held_fill_probe = Arc::clone(&held_fill_called);
std::thread::scope(|scope| {
scope
.spawn(|| {
let hit = residency
.resident_vmm_with(
100,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |_, _| -> Result<(), WeightHandleError> {
held_fill_probe.store(true, Ordering::Release);
Ok(())
},
)
.map(VmmAdmit::expect_page)
.expect("consumer-owned page must remain a valid hit");
assert_eq!(hit.ptr, completed_ptr);
})
.join()
.expect("consumer-owned refill thread");
});
assert!(!held_fill_called.load(Ordering::Acquire));
assert_eq!(
residency
.lock()
.slots
.get(&100)
.expect("consumer-owned slot")
.state
.status(),
SlotStatus::Idle
);
let mut held_bytes = vec![0u8; granule];
unsafe { runtime.dtoh(&mut held_bytes, completed_page.ptr) }
.expect("read consumer-owned page");
assert_eq!(held_bytes, completed_bytes);
drop(completed_page);
let refill_bytes = vec![0x32u8; granule];
let exclusive_fill_called = Arc::new(AtomicBool::new(false));
let exclusive_fill_probe = Arc::clone(&exclusive_fill_called);
let refilled_page = residency
.resident_vmm_with(
100,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
|runtime, ptr| {
exclusive_fill_probe.store(true, Ordering::Release);
unsafe { runtime.htod(&refill_bytes, ptr) }
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))
},
)
.map(VmmAdmit::expect_page)
.expect("exclusive refill must execute");
assert!(exclusive_fill_called.load(Ordering::Acquire));
assert_eq!(refilled_page.ptr, completed_ptr);
let mut actual_refill_bytes = vec![0u8; granule];
unsafe { runtime.dtoh(&mut actual_refill_bytes, refilled_page.ptr) }
.expect("read exclusively refilled page");
assert_eq!(actual_refill_bytes, refill_bytes);
assert_eq!(
residency
.lock()
.slots
.get(&100)
.expect("successfully refilled slot")
.state
.status(),
SlotStatus::Idle
);
drop(refilled_page);
let completed_error = residency
.resident_vmm_with(
100,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
|_, _| {
Err(VmmFillFailure::completed(WeightHandleError::DeviceBinding(
"injected completed refill measurement failure".into(),
)))
},
)
.expect_err("completed refill fault must be reported");
assert!(completed_error.to_string().contains("injected completed"));
assert_eq!(
residency
.lock()
.slots
.get(&100)
.expect("completed-control slot")
.state
.status(),
SlotStatus::Idle
);
let unresolved_bytes = vec![0x41u8; granule];
let unresolved_page = residency
.resident_vmm_with(
101,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
|runtime, ptr| {
unsafe { runtime.htod(&unresolved_bytes, ptr) }
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))
},
)
.map(VmmAdmit::expect_page)
.expect("admit unresolved-control page");
residency.lock().mark_pinned(101, granule as u64);
drop(unresolved_page);
let source_dropped = Arc::new(AtomicBool::new(false));
let source_probe = Arc::clone(&source_dropped);
let charged_before = governor.used(Tier::Device);
let mapped_before = residency.stats().mapped_physical_bytes;
let active_refills = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let max_active_refills = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let second_fill_called = Arc::new(AtomicBool::new(false));
let (first_entered_tx, first_entered_rx) = std::sync::mpsc::channel();
let (release_first_tx, release_first_rx) = std::sync::mpsc::channel();
let (unresolved_error, concurrent_error) = std::thread::scope(|scope| {
let active_refills = Arc::clone(&active_refills);
let max_active_refills = Arc::clone(&max_active_refills);
let first = scope.spawn(|| {
residency
.resident_vmm_with(
101,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |runtime, ptr| -> Result<(), VmmFillFailure> {
let active = active_refills.fetch_add(1, Ordering::AcqRel) + 1;
max_active_refills.fetch_max(active, Ordering::AcqRel);
first_entered_tx.send(()).expect("signal first refill");
release_first_rx.recv().expect("release first refill");
unsafe { runtime.htod(&unresolved_bytes, ptr) }.map_err(|error| {
VmmFillFailure::completed(WeightHandleError::DeviceBinding(
error.to_string(),
))
})?;
active_refills.fetch_sub(1, Ordering::AcqRel);
Err(VmmFillFailure::may_be_in_flight(
WeightHandleError::DeviceBinding(
"injected pinned refill synchronization failure".into(),
),
Box::new(DropProbe(source_probe)),
))
},
)
.expect_err("unresolved refill must quarantine the resident page")
});
first_entered_rx.recv().expect("first refill entered");
assert_eq!(
residency
.lock()
.slots
.get(&101)
.expect("pending refill slot")
.state
.status(),
SlotStatus::Pending
);
let second_fill_probe = Arc::clone(&second_fill_called);
let concurrent_error = residency
.resident_vmm_with(
101,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |_, _| -> Result<(), WeightHandleError> {
second_fill_probe.store(true, Ordering::Release);
Ok(())
},
)
.expect_err("concurrent lookup must not return or rewrite a pending slot");
release_first_tx.send(()).expect("release first refill");
(first.join().expect("first refill thread"), concurrent_error)
});
assert!(concurrent_error.to_string().contains("pending"));
assert!(!second_fill_called.load(Ordering::Acquire));
assert_eq!(max_active_refills.load(Ordering::Acquire), 1);
assert!(
unresolved_error
.to_string()
.contains("destination mapping and staging source remain charged and quarantined")
);
assert!(!source_dropped.load(Ordering::Acquire));
assert!(!residency.lock().pages.contains_key(&101));
assert_eq!(
residency
.lock()
.slots
.get(&101)
.expect("unresolved slot")
.state
.status(),
SlotStatus::Poisoned
);
assert_eq!(governor.used(Tier::Device), charged_before);
assert_eq!(residency.stats().mapped_physical_bytes, mapped_before);
assert!(!source_dropped.load(Ordering::Acquire));
let fill_called = Arc::new(AtomicBool::new(false));
let fill_probe = Arc::clone(&fill_called);
let poisoned_error = residency
.resident_vmm_with(
101,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |_, _| -> Result<(), WeightHandleError> {
fill_probe.store(true, Ordering::Release);
Ok(())
},
)
.expect_err("quarantined destination must never be refilled");
assert!(poisoned_error.to_string().contains("poisoned"));
assert!(!fill_called.load(Ordering::Acquire));
let (teardown_entered_tx, teardown_entered_rx) = std::sync::mpsc::channel();
let (finish_teardown_tx, finish_teardown_rx) = std::sync::mpsc::channel();
std::thread::scope(|scope| {
let refill = scope.spawn(|| {
residency
.resident_vmm_with(
100,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |_, _| -> Result<(), WeightHandleError> {
teardown_entered_tx
.send(())
.expect("signal teardown refill");
finish_teardown_rx.recv().expect("finish teardown refill");
Ok(())
},
)
.expect_err("teardown poison must prevent the pending refill pointer return")
});
teardown_entered_rx.recv().expect("teardown refill entered");
assert_eq!(
residency
.lock()
.slots
.get(&100)
.expect("teardown pending slot")
.state
.status(),
SlotStatus::Pending
);
residency.confirm_context_terminated();
finish_teardown_tx.send(()).expect("finish teardown refill");
let error = refill.join().expect("teardown refill thread");
assert!(error.to_string().contains("poisoned"));
});
assert_eq!(
residency
.lock()
.slots
.get(&100)
.expect("teardown poisoned slot")
.state
.status(),
SlotStatus::Poisoned
);
assert_eq!(
residency
.lock()
.slots
.get(&101)
.expect("unresolved poisoned slot")
.state
.status(),
SlotStatus::Poisoned
);
}
#[cfg(feature = "gpu-tests")]
#[test]
fn failed_vmm_fill_keeps_cleanup_residuals_charged_and_quarantined() {
use onnx_runtime_cuda_memory::release::{DriverFaultPlan, DriverOperation};
use onnx_runtime_memory_governor::{
DeviceKey, HolderId, LeaseLedger, LedgerGovernor, MemoryGovernor, MemoryRole,
};
let mut env = EnvVarGuard::acquire();
env.unset(crate::vmm_allocator::CUDA_PHYSICAL_HANDLE_POOL_BYTES_ENV);
let Ok(runtime) = CudaRuntime::new(0).map(Arc::new) else {
eprintln!("SKIPPED (CUDA runtime dependencies unavailable): VMM cleanup fault test");
return;
};
let granule = 2usize << 20;
let governor = Arc::new(LedgerGovernor::new(LeaseLedger::new(granule as u64, 0, 0)));
let mut allocator = crate::vmm_allocator::CudaVmmAllocator::new(
runtime.cuda_context(),
DeviceKey::device(0),
0,
64 << 20,
governor.as_ref(),
HolderId::new(739),
MemoryRole::Weights,
)
.expect("no-pool VMM allocator");
allocator.install_driver_faults(Arc::new(
DriverFaultPlan::new().fail_nth(DriverOperation::Unmap, 1),
));
let allocator = Arc::new(allocator);
let authority: Arc<dyn MemoryGovernor + Send + Sync> = governor.clone();
let residency = CudaWeightResidency::new(Arc::clone(&runtime), granule as u64)
.with_vmm_admission(Arc::clone(&allocator), authority)
.expect("install VMM admission");
residency
.adopt_governed_budget(governor.as_ref(), Tier::Device, HolderId::new(739))
.expect("reserve mapped allowance");
let global_before = GLOBAL_WEIGHT_MAPPED_BYTES.load(Ordering::Relaxed);
let error = residency
.resident_vmm_with(
1,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
|_, _| {
Err(WeightHandleError::DeviceBinding(
"injected fill failure before cleanup fault".into(),
))
},
)
.expect_err("fill and cleanup failure");
assert!(error.to_string().contains("quarantined"));
assert_eq!(residency.stats().mapped_physical_bytes, granule as u64);
assert_eq!(governor.used(Tier::Device), granule as u64);
assert_eq!(allocator.quarantined_owned_bytes(), granule as u64);
assert_eq!(
GLOBAL_WEIGHT_MAPPED_BYTES.load(Ordering::Relaxed),
global_before + granule as u64,
"a still-mapped quarantined residual must remain in the global mapped gauge"
);
}
#[cfg_attr(
not(feature = "gpu-tests"),
ignore = "requires CUDA device; enable the gpu-tests feature on a CUDA runner"
)]
#[test]
fn vmm_weight_admission_reuses_owned_granules_without_runtime_alloc_free() {
use onnx_runtime_memory_governor::{
DeviceKey, HolderId, LeaseLedger, LedgerGovernor, MemoryGovernor, MemoryRole,
};
let mut env = EnvVarGuard::acquire();
env.set(
crate::vmm_allocator::CUDA_PHYSICAL_HANDLE_POOL_BYTES_ENV,
&(64usize << 20).to_string(),
);
let Ok(runtime) = CudaRuntime::new(0).map(Arc::new) else {
eprintln!(
"SKIPPED (CUDA runtime dependencies unavailable): VMM weight admission GPU test"
);
return;
};
let granule = 2usize << 20;
let governor = Arc::new(LedgerGovernor::new(LeaseLedger::new(
(granule * 2) as u64,
0,
0,
)));
let allocator = Arc::new(
crate::vmm_allocator::CudaVmmAllocator::new(
runtime.cuda_context(),
DeviceKey::device(0),
0,
64 << 20,
governor.as_ref(),
HolderId::new(736),
MemoryRole::Weights,
)
.expect("VMM allocator"),
);
let authority: Arc<dyn onnx_runtime_memory_governor::MemoryGovernor + Send + Sync> =
governor.clone();
let release_queue = CudaDeferredReleaseQueue::new(
Box::new(crate::deferred_release::CudaStreamFences::new(Arc::clone(
&runtime,
))),
crate::deferred_release::DEFAULT_DEFERRED_RELEASE_CAPACITY,
);
let residency = CudaWeightResidency::new(Arc::clone(&runtime), granule as u64)
.with_deferred_release_queue(Arc::clone(&release_queue))
.with_vmm_admission(Arc::clone(&allocator), authority)
.expect("install VMM admission");
residency
.adopt_governed_budget(governor.as_ref(), Tier::Device, HolderId::new(736))
.expect("reserve mapped weight allowance");
let before = runtime.allocation_counts();
let first_bytes = vec![0x31u8; granule];
let first = residency
.resident_vmm_with(
1,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |runtime, ptr| {
unsafe { runtime.htod(&first_bytes, ptr) }
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))
},
)
.map(VmmAdmit::expect_page)
.expect("first physical page");
let zone_error = residency
.resident_vmm_with(
2,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::StableResident,
false,
|_, _| -> Result<(), WeightHandleError> {
panic!("zone refusal must happen before copy")
},
)
.expect_err("global room cannot bypass a full mapped weight allowance");
assert!(zone_error.to_string().contains("weight-zone headroom"));
drop(first);
let second_bytes = vec![0x42u8; granule];
let second = residency
.resident_vmm_with(
2,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |runtime, ptr| {
unsafe { runtime.htod(&second_bytes, ptr) }
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))
},
)
.map(VmmAdmit::expect_page)
.expect("second page reuses the owned handle");
drop(second);
let stats = residency.stats();
assert_eq!(stats.resident_bytes, granule as u64);
assert_eq!(stats.mapped_physical_bytes, granule as u64);
assert_eq!(stats.physical_owned_bytes, granule as u64);
assert_eq!(stats.page_ins, 2);
assert_eq!(stats.evictions, 1);
assert_eq!(runtime.allocation_counts(), before);
assert_eq!(governor.used(Tier::Device), granule as u64);
let live_global = global_offload_stats();
assert!(live_global.content_resident_bytes >= granule as u64);
assert!(live_global.mapped_physical_bytes >= granule as u64);
assert!(live_global.budget_bytes >= granule as u64);
let second_authority: Arc<dyn onnx_runtime_memory_governor::MemoryGovernor + Send + Sync> =
governor.clone();
let other = CudaWeightResidency::new(Arc::clone(&runtime), granule as u64)
.with_vmm_admission(Arc::clone(&allocator), second_authority)
.expect("install second VMM admission");
other
.adopt_governed_budget(governor.as_ref(), Tier::Device, HolderId::new(737))
.expect("reserve second mapped allowance");
let _kv = governor
.reserve(
Tier::Device,
granule as u64,
MemoryRole::KvCache,
HolderId::new(9),
)
.expect("KV consumes remaining physical headroom");
let global_error = other
.resident_vmm_with(
3,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
|_, _| -> Result<(), WeightHandleError> {
panic!("global refusal must happen before copy")
},
)
.expect_err("zone room cannot bypass missing global physical headroom");
assert!(global_error.to_string().contains("physical headroom"));
drop(_kv);
drop(other);
drop(residency);
assert!(
release_queue.wait_until_idle(DEFERRED_RELEASE_WAIT_TIMEOUT),
"the final resident page must be released before checking the unloaded gauges"
);
let retained = allocator
.physical_pool_stats()
.expect("pool stats after unload")
.snapshot();
assert_eq!(retained.mapped_bytes, 0);
assert_eq!(retained.pooled_unmapped_bytes, granule as u64);
assert_eq!(retained.total_owned_bytes, granule as u64);
}
#[cfg_attr(
not(feature = "gpu-tests"),
ignore = "requires CUDA device; enable the gpu-tests feature on a CUDA runner"
)]
#[test]
fn vmm_retained_weight_key_keeps_a_stable_virtual_address_across_repage() {
use onnx_runtime_memory_governor::{
DeviceKey, HolderId, LeaseLedger, LedgerGovernor, MemoryRole,
};
let mut env = EnvVarGuard::acquire();
env.set(
crate::vmm_allocator::CUDA_PHYSICAL_HANDLE_POOL_BYTES_ENV,
&(64usize << 20).to_string(),
);
let Ok(runtime) = CudaRuntime::new(0).map(Arc::new) else {
eprintln!("SKIPPED (CUDA runtime dependencies unavailable): stable-VA repage GPU test");
return;
};
let granule = 2usize << 20;
let governor = Arc::new(LedgerGovernor::new(LeaseLedger::new(
(granule * 2) as u64,
0,
0,
)));
let allocator = Arc::new(
crate::vmm_allocator::CudaVmmAllocator::new(
runtime.cuda_context(),
DeviceKey::device(0),
0,
64 << 20,
governor.as_ref(),
HolderId::new(716),
MemoryRole::Weights,
)
.expect("VMM allocator"),
);
let authority: Arc<dyn onnx_runtime_memory_governor::MemoryGovernor + Send + Sync> =
governor.clone();
let residency = CudaWeightResidency::new(Arc::clone(&runtime), granule as u64)
.with_deferred_release_queue(CudaDeferredReleaseQueue::new(
Box::new(crate::deferred_release::CudaStreamFences::new(Arc::clone(
&runtime,
))),
crate::deferred_release::DEFAULT_DEFERRED_RELEASE_CAPACITY,
))
.with_vmm_admission(Arc::clone(&allocator), authority)
.expect("install VMM admission");
residency
.adopt_governed_budget(governor.as_ref(), Tier::Device, HolderId::new(716))
.expect("reserve mapped weight allowance");
assert!(
residency.stable_va_paging_active(),
"installing VMM admission must activate the stable-VA paging path"
);
let page_key_1 = |fill_byte: u8| {
let bytes = vec![fill_byte; granule];
residency
.resident_vmm_with(
1,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |runtime, ptr| {
unsafe { runtime.htod(&bytes, ptr) }
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))
},
)
.map(VmmAdmit::expect_page)
.expect("page key 1")
};
let first = page_key_1(0x11);
let stable_va = first.device_ptr();
drop(first);
let second_bytes = vec![0x22u8; granule];
let second = residency
.resident_vmm_with(
2,
DataType::Uint8,
vec![granule],
granule,
WeightEvictionPolicy::Lru,
false,
move |runtime, ptr| {
unsafe { runtime.htod(&second_bytes, ptr) }
.map_err(|error| WeightHandleError::DeviceBinding(error.to_string()))
},
)
.map(VmmAdmit::expect_page)
.expect("page key 2 evicts key 1");
assert_ne!(
second.device_ptr(),
stable_va,
"a different key must get a different VA slot"
);
drop(second);
let repaged = page_key_1(0x33);
assert_eq!(
repaged.device_ptr(),
stable_va,
"issue #716: a retained key must keep its stable VA across evict→repage"
);
let stats = residency.stats();
assert_eq!(stats.page_ins, 3);
assert_eq!(stats.evictions, 2);
assert_eq!(stats.physical_owned_bytes, granule as u64);
drop(repaged);
drop(residency);
}
#[test]
fn adopting_a_governed_budget_makes_the_claim_visible_to_other_holders() {
use onnx_runtime_memory_governor::{
HolderId, LeaseLedger, LedgerGovernor, MemoryGovernor, MemoryRole, Tier,
};
let Ok(runtime) = crate::runtime::CudaRuntime::new(0).map(std::sync::Arc::new) else {
eprintln!("SKIPPED (no CUDA runtime): the governed weight-budget check did NOT run.");
return;
};
let residency = CudaWeightResidency::new(runtime, 700);
assert_eq!(
residency.budget(),
(700, false),
"before adoption the budget answers to nobody"
);
let governor = LedgerGovernor::new(LeaseLedger::new(1000, 0, 0));
let granted = residency
.adopt_governed_budget(&governor, Tier::Device, HolderId::new(4))
.expect("700 of 1000 is affordable");
assert_eq!(granted, 700);
assert_eq!(residency.budget(), (700, true));
assert_eq!(governor.available(Tier::Device), 300);
let refused = governor
.reserve(Tier::Device, 700, MemoryRole::KvCache, HolderId::new(1))
.expect_err("the weights already hold 700 of the 1000");
assert!(matches!(
refused,
onnx_runtime_memory_governor::MemoryError::TierExhausted { .. }
));
}
#[test]
fn a_budget_already_governed_is_not_reserved_again() {
use onnx_runtime_memory_governor::{
HolderId, LeaseLedger, LedgerGovernor, MemoryGovernor, Tier,
};
let Ok(runtime) = crate::runtime::CudaRuntime::new(0).map(std::sync::Arc::new) else {
eprintln!("SKIPPED (no CUDA runtime): the double-adoption check did NOT run.");
return;
};
let residency = CudaWeightResidency::new(runtime, 400);
let governor = LedgerGovernor::new(LeaseLedger::new(1000, 0, 0));
residency
.adopt_governed_budget(&governor, Tier::Device, HolderId::new(4))
.expect("first adoption");
residency
.adopt_governed_budget(&governor, Tier::Device, HolderId::new(4))
.expect("second adoption is a no-op, not a second charge");
assert_eq!(
governor.available(Tier::Device),
600,
"the same budget was charged twice"
);
}
fn drive_repeated_scan(
policy: &mut WeightResidencyPolicy,
eviction: WeightEvictionPolicy,
working_set: u64,
cycles: u64,
) -> u64 {
let hits_before = policy.hits;
for _ in 0..cycles {
for key in 0..working_set {
let _ = policy.access(key, 1, eviction);
}
}
policy.hits - hits_before
}
fn measured_scan_hits(capacity: u64, eviction: WeightEvictionPolicy) -> u64 {
const WORKING_SET: u64 = 10;
const MEASURED_CYCLES: u64 = 5;
let mut policy = WeightResidencyPolicy::new(capacity);
let _ = drive_repeated_scan(&mut policy, eviction, WORKING_SET, 1);
drive_repeated_scan(&mut policy, eviction, WORKING_SET, MEASURED_CYCLES)
}
#[test]
fn lru_has_zero_steady_state_hits_across_cyclic_scan_capacity_sweep() {
const WORKING_SET: u64 = 10;
const MEASURED_CYCLES: u64 = 5;
let accesses = WORKING_SET * MEASURED_CYCLES;
for capacity in [2, 4, 6, 8] {
let hits = measured_scan_hits(capacity, WeightEvictionPolicy::Lru);
assert_eq!(
hits,
0,
"LRU should be pessimal for a clean cycle larger than capacity; \
B={capacity}/{WORKING_SET}, hit_rate={:.1}%",
(hits as f64 / accesses as f64) * 100.0
);
}
}
#[test]
fn stable_subset_recovers_capacity_fraction_across_cyclic_scan_sweep() {
const WORKING_SET: u64 = 10;
const MEASURED_CYCLES: u64 = 5;
let accesses = WORKING_SET * MEASURED_CYCLES;
for capacity in [2, 4, 6, 8] {
let hits = measured_scan_hits(capacity, WeightEvictionPolicy::StableResident);
assert_eq!(hits, capacity * MEASURED_CYCLES);
assert_eq!(
(hits as f64) / (accesses as f64),
(capacity as f64) / (WORKING_SET as f64),
"stable-subset residency should recover B/W for B={capacity}/{WORKING_SET}"
);
}
}
#[test]
fn byte_aware_beats_stable_subset_when_smalls_crowd_out_larges() {
let order: [(u64, u64); 8] = [
(0, 1),
(1, 1),
(2, 1),
(3, 1),
(4, 1),
(5, 1),
(6, 10),
(7, 10),
];
const BUDGET: u64 = 22;
const WARMUP: u64 = 4;
const MEASURED: u64 = 6;
let byte_hits = |byte_aware: bool| -> (u64, u64) {
let mut policy = WeightResidencyPolicy::new(BUDGET);
let mut hit_bytes = 0u64;
let mut streamed_bytes = 0u64;
for cycle in 0..(WARMUP + MEASURED) {
let measuring = cycle >= WARMUP;
for (key, bytes) in order {
let access = if byte_aware {
policy.access_byte_aware(key, bytes)
} else {
policy.access(key, bytes, WeightEvictionPolicy::StableResident)
};
if measuring {
if access.hit {
hit_bytes += bytes;
} else {
streamed_bytes += bytes;
}
}
}
}
(hit_bytes, streamed_bytes)
};
let (blind_hit, blind_stream) = byte_hits(false);
let (aware_hit, aware_stream) = byte_hits(true);
let blind_rate = blind_hit as f64 / (blind_hit + blind_stream) as f64;
let aware_rate = aware_hit as f64 / (aware_hit + aware_stream) as f64;
assert!(
aware_rate > blind_rate + 0.15,
"byte-aware residency must materially raise the byte-weighted hit rate when \
smalls crowd out larges: blind={blind_rate:.3} aware={aware_rate:.3}"
);
assert!(
aware_stream < blind_stream,
"byte-aware residency must stream fewer bytes: \
blind={blind_stream} aware={aware_stream}"
);
}
#[test]
fn byte_aware_does_not_thrash_when_larges_exceed_budget() {
let order: [(u64, u64); 5] = [(0, 10), (1, 10), (2, 10), (3, 10), (4, 10)];
const BUDGET: u64 = 22;
const WARMUP: u64 = 4;
const MEASURED: u64 = 6;
let steady_stream = |byte_aware: bool| -> u64 {
let mut policy = WeightResidencyPolicy::new(BUDGET);
let mut streamed = 0u64;
for cycle in 0..(WARMUP + MEASURED) {
for (key, bytes) in order {
let access = if byte_aware {
policy.access_byte_aware(key, bytes)
} else {
policy.access(key, bytes, WeightEvictionPolicy::StableResident)
};
if cycle >= WARMUP && !access.hit {
streamed += bytes;
}
}
}
streamed
};
assert!(
steady_stream(true) <= steady_stream(false),
"byte-aware must not stream more than size-blind when larges exceed budget: \
aware={} blind={}",
steady_stream(true),
steady_stream(false)
);
}
#[test]
fn stable_subset_matches_lru_when_whole_scan_fits() {
const WORKING_SET: u64 = 10;
const MEASURED_CYCLES: u64 = 5;
for capacity in [10, 12] {
let lru_hits = measured_scan_hits(capacity, WeightEvictionPolicy::Lru);
let stable_hits = measured_scan_hits(capacity, WeightEvictionPolicy::StableResident);
assert_eq!(lru_hits, WORKING_SET * MEASURED_CYCLES);
assert_eq!(stable_hits, lru_hits);
}
}
#[test]
fn scan_resistant_mode_leaves_moe_skew_on_lru() {
let selected = eviction_for_boundary(true, LazyWeightBoundary::QMoe);
assert_eq!(selected, WeightEvictionPolicy::Lru);
const CYCLES: u64 = 20;
let skewed = [0, 0, 0, 0, 1, 0, 1, 2, 0, 3, 0, 1];
let mut baseline = WeightResidencyPolicy::new(3);
let mut moe = WeightResidencyPolicy::new(3);
for _ in 0..CYCLES {
for &key in &skewed {
let _ = baseline.access(key, 1, WeightEvictionPolicy::Lru);
let _ = moe.access(key, 1, selected);
}
}
assert_eq!(moe.hits, baseline.hits);
assert_eq!(moe.page_ins, baseline.page_ins);
assert_eq!(moe.evictions, baseline.evictions);
assert_eq!(moe.hits, 178);
assert_eq!(moe.page_ins, 62);
assert_eq!(moe.evictions, 59);
}
#[test]
fn hot_set_policy_eviction_class_agrees_with_old_eviction_for_boundary() {
use onnx_runtime_ep_api::{EvictionClass, ResidencyPolicy as _};
let boundaries = [
LazyWeightBoundary::MatMul,
LazyWeightBoundary::MatMulNBits,
LazyWeightBoundary::QMoe,
];
for scan_resistant_dense in [false, true] {
for boundary in boundaries {
let old = eviction_for_boundary(scan_resistant_dense, boundary);
let new = CudaHotSetResidencyPolicy::from_env(scan_resistant_dense)
.eviction_class(boundary);
let new_as_old: WeightEvictionPolicy = new.into();
assert_eq!(
old, new_as_old,
"policy disagrees with legacy eviction_for_boundary for \
scan_resistant_dense={scan_resistant_dense} boundary={boundary:?}"
);
match boundary {
LazyWeightBoundary::MatMul | LazyWeightBoundary::MatMulNBits
if scan_resistant_dense =>
{
assert_eq!(new, EvictionClass::StableResident);
}
_ => assert_eq!(new, EvictionClass::Lru),
}
}
}
}
#[test]
fn hot_set_policy_defaults_never_pin_matching_shipped_env() {
use onnx_runtime_ep_api::{AdmissionPolicyInput, ResidencyPolicy as _};
let policy = CudaHotSetResidencyPolicy {
scan_resistant_dense: true,
static_pin_keys: None,
static_pin_config: None,
};
assert!(!policy.should_pin(&AdmissionPolicyInput {
key: 7,
len_bytes: u64::MAX,
already_pinned: false,
pinned_bytes_used: 0,
}));
}
#[test]
fn hot_set_policy_pin_keys_take_priority_over_threshold() {
use onnx_runtime_ep_api::{AdmissionPolicyInput, ResidencyPolicy as _};
let mut keys = HashSet::new();
keys.insert(42u64);
let keys: &'static HashSet<u64> = Box::leak(Box::new(keys));
let policy = CudaHotSetResidencyPolicy {
scan_resistant_dense: false,
static_pin_keys: Some(keys),
static_pin_config: Some((u64::MAX, u64::MAX)),
};
assert!(policy.should_pin(&AdmissionPolicyInput {
key: 42,
len_bytes: 1,
already_pinned: false,
pinned_bytes_used: 0,
}));
assert!(!policy.should_pin(&AdmissionPolicyInput {
key: 43,
len_bytes: 1,
already_pinned: false,
pinned_bytes_used: 0,
}));
assert!(!policy.should_pin(&AdmissionPolicyInput {
key: 42,
len_bytes: 1,
already_pinned: true,
pinned_bytes_used: 0,
}));
}
#[test]
fn hot_set_policy_threshold_path_respects_budget() {
use onnx_runtime_ep_api::{AdmissionPolicyInput, ResidencyPolicy as _};
let policy = CudaHotSetResidencyPolicy {
scan_resistant_dense: false,
static_pin_keys: None,
static_pin_config: Some((100, 250)),
};
assert!(!policy.should_pin(&AdmissionPolicyInput {
key: 1,
len_bytes: 99,
already_pinned: false,
pinned_bytes_used: 0,
}));
assert!(policy.should_pin(&AdmissionPolicyInput {
key: 1,
len_bytes: 200,
already_pinned: false,
pinned_bytes_used: 0,
}));
assert!(!policy.should_pin(&AdmissionPolicyInput {
key: 1,
len_bytes: 200,
already_pinned: false,
pinned_bytes_used: 100,
}));
}
#[test]
fn hot_set_policy_decide_delegates_to_whole_bank_default() {
use onnx_runtime_ep_api::{ResidencyDecision, ResidencyPolicy as _};
let policy = CudaHotSetResidencyPolicy {
scan_resistant_dense: true,
static_pin_keys: None,
static_pin_config: None,
};
let layout = onnx_runtime_loader::ExpertTensorLayout {
version: 1,
experts: 3,
rows_per_expert: 2,
storage_elements_per_row: 4,
order: onnx_runtime_loader::ExpertStorageOrder::ExpertMajor,
quantization: Some(onnx_runtime_loader::ExpertQuantization {
bits: 4,
block_size: 16,
blocks_per_row: 1,
}),
};
let weight = onnx_runtime_ir::WeightRef::External {
path: std::path::PathBuf::from("/nonexistent/weights.bin"),
offset: 16,
length: layout.experts * layout.rows_per_expert * layout.storage_elements_per_row,
dtype: DataType::Uint8,
dims: vec![
layout.experts,
layout.rows_per_expert,
layout.storage_elements_per_row,
],
};
let catalog = onnx_runtime_loader::WeightRegionCatalog::classify(&weight, layout);
let input = onnx_runtime_ep_api::ResidencyPolicyInput {
value_id: onnx_runtime_ir::ValueId(0),
boundary: LazyWeightBoundary::QMoe,
catalog: &catalog,
budget_bytes: None,
};
assert_eq!(
policy.decide(&input),
ResidencyDecision::WholeBankResident { reason: None }
);
}
fn grow_request(bytes: u64) -> onnx_runtime_ep_api::ResidencyResizeRequest {
onnx_runtime_ep_api::ResidencyResizeRequest {
direction: onnx_runtime_ep_api::ResizeDirection::Grow,
target_bytes: bytes,
priority: 0,
}
}
fn shrink_request(bytes: u64) -> onnx_runtime_ep_api::ResidencyResizeRequest {
onnx_runtime_ep_api::ResidencyResizeRequest {
direction: onnx_runtime_ep_api::ResizeDirection::Shrink,
target_bytes: bytes,
priority: 0,
}
}
#[test]
fn ungoverned_cache_refuses_grow_and_shrink_leaving_budget_untouched() {
let Ok(runtime) = crate::runtime::CudaRuntime::new(0).map(std::sync::Arc::new) else {
eprintln!(
"SKIPPED (no CUDA runtime): the ungoverned resize refusal check did NOT run."
);
return;
};
let residency = CudaWeightResidency::new(runtime, 500);
assert_eq!(residency.budget(), (500, false));
let plan =
onnx_runtime_ep_api::plan_resize(grow_request(100), residency.resize_safe_point(1));
let outcome = residency.execute_resize(plan, 1);
assert!(!outcome.is_success());
assert_eq!(residency.budget(), (500, false));
let plan =
onnx_runtime_ep_api::plan_resize(shrink_request(100), residency.resize_safe_point(1));
let outcome = residency.execute_resize(plan, 1);
assert!(!outcome.is_success());
assert_eq!(residency.budget(), (500, false));
}
#[test]
fn governed_grow_moves_exactly_the_requested_bytes() {
use onnx_runtime_memory_governor::{
HolderId, LeaseLedger, LedgerGovernor, MemoryGovernor, Tier,
};
let Ok(runtime) = crate::runtime::CudaRuntime::new(0).map(std::sync::Arc::new) else {
eprintln!("SKIPPED (no CUDA runtime): the governed grow check did NOT run.");
return;
};
let residency = CudaWeightResidency::new(runtime, 400);
let governor = LedgerGovernor::new(LeaseLedger::new(1000, 0, 0));
residency
.adopt_governed_budget(&governor, Tier::Device, HolderId::new(9))
.expect("400 of 1000 is affordable");
assert_eq!(residency.budget(), (400, true));
let plan =
onnx_runtime_ep_api::plan_resize(grow_request(200), residency.resize_safe_point(1));
let outcome = residency.execute_resize(plan, 1);
assert!(outcome.is_success(), "{outcome:?}");
assert_eq!(outcome.before_bytes, 400);
assert_eq!(outcome.after_bytes, 600);
assert_eq!(outcome.accepted_bytes, 200);
assert_eq!(residency.budget(), (600, true));
assert_eq!(governor.available(Tier::Device), 400);
}
#[test]
fn governed_grow_beyond_capacity_fails_and_leaves_budget_unchanged() {
use onnx_runtime_memory_governor::{HolderId, LeaseLedger, LedgerGovernor, Tier};
let Ok(runtime) = crate::runtime::CudaRuntime::new(0).map(std::sync::Arc::new) else {
eprintln!("SKIPPED (no CUDA runtime): the over-capacity grow check did NOT run.");
return;
};
let residency = CudaWeightResidency::new(runtime, 400);
let governor = LedgerGovernor::new(LeaseLedger::new(1000, 0, 0));
residency
.adopt_governed_budget(&governor, Tier::Device, HolderId::new(9))
.expect("400 of 1000 is affordable");
let plan = onnx_runtime_ep_api::plan_resize(
grow_request(1_000_000),
residency.resize_safe_point(1),
);
let outcome = residency.execute_resize(plan, 1);
assert!(!outcome.is_success());
assert_eq!(outcome.before_bytes, 400);
assert_eq!(outcome.after_bytes, 400);
assert_eq!(outcome.accepted_bytes, 0);
assert_eq!(residency.budget(), (400, true));
}
#[test]
fn zero_byte_resize_is_a_rejected_noop() {
let Ok(runtime) = crate::runtime::CudaRuntime::new(0).map(std::sync::Arc::new) else {
eprintln!("SKIPPED (no CUDA runtime): the no-op resize check did NOT run.");
return;
};
let residency = CudaWeightResidency::new(runtime, 400);
let plan =
onnx_runtime_ep_api::plan_resize(grow_request(0), residency.resize_safe_point(1));
assert!(matches!(
plan,
onnx_runtime_ep_api::ResidencyResizePlan::Rejected {
reason: onnx_runtime_ep_api::ResizeRejection::NoOp,
..
}
));
let outcome = residency.execute_resize(plan, 1);
assert!(!outcome.is_success());
assert_eq!(residency.budget(), (400, false));
}
#[test]
fn rejected_plan_outcome_preserves_the_original_shrink_direction() {
let Ok(runtime) = crate::runtime::CudaRuntime::new(0).map(std::sync::Arc::new) else {
eprintln!(
"SKIPPED (no CUDA runtime): the rejected-shrink-direction check did NOT run."
);
return;
};
let residency = CudaWeightResidency::new(runtime, 400);
let unsafe_point = onnx_runtime_ep_api::ResizeSafePoint {
multi_device: true,
..Default::default()
};
let plan = onnx_runtime_ep_api::plan_resize(shrink_request(64), unsafe_point);
let outcome = residency.execute_resize(plan, 2);
assert!(!outcome.is_success());
assert_eq!(
outcome.direction,
onnx_runtime_ep_api::ResizeDirection::Shrink
);
assert_eq!(outcome.requested_bytes, 64);
assert_eq!(outcome.accepted_bytes, 0);
}
#[test]
fn resize_safe_point_fails_closed_for_multi_device() {
let Ok(runtime) = crate::runtime::CudaRuntime::new(0).map(std::sync::Arc::new) else {
eprintln!("SKIPPED (no CUDA runtime): the multi-device fail-closed check did NOT run.");
return;
};
let residency = CudaWeightResidency::new(runtime, 400);
let point = residency.resize_safe_point(2);
assert!(point.multi_device);
assert!(!point.is_safe());
let plan = onnx_runtime_ep_api::plan_resize(grow_request(1), point);
assert!(matches!(
plan,
onnx_runtime_ep_api::ResidencyResizePlan::Rejected {
reason: onnx_runtime_ep_api::ResizeRejection::NotSafePoint(_),
..
}
));
}
#[test]
fn repeated_grow_shrink_oscillation_returns_to_the_starting_budget() {
use onnx_runtime_memory_governor::{
HolderId, LeaseLedger, LedgerGovernor, MemoryGovernor, Tier,
};
let Ok(runtime) = crate::runtime::CudaRuntime::new(0).map(std::sync::Arc::new) else {
eprintln!("SKIPPED (no CUDA runtime): the oscillation check did NOT run.");
return;
};
let residency = CudaWeightResidency::new(runtime, 400);
let governor = LedgerGovernor::new(LeaseLedger::new(1000, 0, 0));
residency
.adopt_governed_budget(&governor, Tier::Device, HolderId::new(9))
.expect("400 of 1000 is affordable");
for _ in 0..5 {
let grow_plan =
onnx_runtime_ep_api::plan_resize(grow_request(100), residency.resize_safe_point(1));
let grow_outcome = residency.execute_resize(grow_plan, 1);
assert!(grow_outcome.is_success(), "{grow_outcome:?}");
let shrink_plan = onnx_runtime_ep_api::plan_resize(
shrink_request(100),
residency.resize_safe_point(1),
);
let shrink_outcome = residency.execute_resize(shrink_plan, 1);
assert!(shrink_outcome.is_success(), "{shrink_outcome:?}");
}
assert_eq!(
residency.budget(),
(900, true),
"5 grows of 100 with nothing evictable to shrink net +500 over the 400 byte start"
);
assert_eq!(governor.available(Tier::Device), 100);
}
fn expert_catalog(pageable: bool) -> onnx_runtime_loader::WeightRegionCatalog {
let mut layout = onnx_runtime_loader::ExpertTensorLayout {
version: 1,
experts: 3,
rows_per_expert: 2,
storage_elements_per_row: 4,
order: onnx_runtime_loader::ExpertStorageOrder::ExpertMajor,
quantization: Some(onnx_runtime_loader::ExpertQuantization {
bits: 4,
block_size: 16,
blocks_per_row: 1,
}),
};
if !pageable {
layout.order = onnx_runtime_loader::ExpertStorageOrder::Interleaved;
}
let weight = onnx_runtime_ir::WeightRef::External {
path: std::path::PathBuf::from("/nonexistent/weights.bin"),
offset: 16,
length: layout.experts * layout.rows_per_expert * layout.storage_elements_per_row,
dtype: DataType::Uint8,
dims: vec![
layout.experts,
layout.rows_per_expert,
layout.storage_elements_per_row,
],
};
onnx_runtime_loader::WeightRegionCatalog::classify(&weight, layout)
}
#[test]
fn acquired_guard_reports_whole_bank_and_blocks_resize_until_dropped() {
let Ok(runtime) = crate::runtime::CudaRuntime::new(0).map(std::sync::Arc::new) else {
eprintln!("SKIPPED (no CUDA runtime): the routed-guard lifecycle check did NOT run.");
return;
};
let residency = std::sync::Arc::new(CudaWeightResidency::new(runtime, 400));
let catalog = expert_catalog(true);
assert!(residency.resize_safe_point(1).is_safe());
let guard = residency.acquire_routed_residency(
onnx_runtime_ep_api::RoutedResidencyRequirement::FusedRoutingUnknown,
&catalog,
);
assert_eq!(
guard.proof().coverage(),
&onnx_runtime_ep_api::RoutedResidencyCoverage::WholeBank {
reason: onnx_runtime_ep_api::WholeBankReason::FusedRoutingHasNoHostVisibility
}
);
let point = residency.resize_safe_point(1);
assert_eq!(point.routed_guards_active, 1);
assert!(!point.is_safe());
let plan =
onnx_runtime_ep_api::plan_resize(grow_request(1), residency.resize_safe_point(1));
assert!(matches!(
plan,
onnx_runtime_ep_api::ResidencyResizePlan::Rejected {
reason: onnx_runtime_ep_api::ResizeRejection::NotSafePoint(_),
..
}
));
drop(guard);
assert!(residency.resize_safe_point(1).is_safe());
}
#[test]
fn concurrent_guards_are_counted_independently_and_released_in_either_order() {
let Ok(runtime) = crate::runtime::CudaRuntime::new(0).map(std::sync::Arc::new) else {
eprintln!("SKIPPED (no CUDA runtime): the concurrent-guard count check did NOT run.");
return;
};
let residency = std::sync::Arc::new(CudaWeightResidency::new(runtime, 400));
let catalog = expert_catalog(true);
let guard_a = residency.acquire_routed_residency(
onnx_runtime_ep_api::RoutedResidencyRequirement::FusedRoutingUnknown,
&catalog,
);
let guard_b = residency.acquire_routed_residency(
onnx_runtime_ep_api::RoutedResidencyRequirement::FusedRoutingUnknown,
&catalog,
);
assert_eq!(residency.resize_safe_point(1).routed_guards_active, 2);
drop(guard_a);
let mid_point = residency.resize_safe_point(1);
assert_eq!(mid_point.routed_guards_active, 1);
assert!(!mid_point.is_safe());
drop(guard_b);
assert!(residency.resize_safe_point(1).is_safe());
}
#[test]
fn guard_over_non_pageable_catalog_degrades_to_whole_bank_and_still_blocks_resize() {
let Ok(runtime) = crate::runtime::CudaRuntime::new(0).map(std::sync::Arc::new) else {
eprintln!(
"SKIPPED (no CUDA runtime): the non-pageable guard degradation check did NOT run."
);
return;
};
let residency = std::sync::Arc::new(CudaWeightResidency::new(runtime, 400));
let catalog = expert_catalog(false);
let guard = residency.acquire_routed_residency(
onnx_runtime_ep_api::RoutedResidencyRequirement::HostKnownExperts { experts: vec![0] },
&catalog,
);
assert!(matches!(
guard.proof().coverage(),
onnx_runtime_ep_api::RoutedResidencyCoverage::WholeBank {
reason: onnx_runtime_ep_api::WholeBankReason::InvalidExactSet(_)
}
));
assert_eq!(residency.resize_safe_point(1).routed_guards_active, 1);
}
}