use std::sync::Arc;
use cudarc::driver::sys as cu;
use cudarc::driver::sys::CUdeviceptr;
use onnx_runtime_cuda_memory::capture_gate;
use onnx_runtime_cuda_memory::release::{DriverFault, MappedBlock};
use onnx_runtime_cuda_memory::virtual_memory::{
CudaReservation, CudaVirtualBacking, PhysicalHandlePool, PhysicalLocation,
};
use onnx_runtime_ep_api::ResizeSafePoint;
use onnx_runtime_virtual_memory::VirtualBacking;
use crate::runtime::CudaRuntime;
#[derive(Debug)]
pub enum TransitionOutcome {
Committed {
granules: usize,
new_owned_bytes: u64,
old_released_bytes: u64,
},
Rejected { reason: &'static str },
RolledBack { fault: DriverFault },
Fatal {
transition_fault: DriverFault,
rollback_fault: Option<DriverFault>,
committed_count: usize,
quarantined: Vec<MappedBlock>,
poisoned_range: Option<(usize, usize)>,
},
}
impl TransitionOutcome {
pub fn is_committed(&self) -> bool {
matches!(self, Self::Committed { .. })
}
pub fn stable_va_intact(&self) -> bool {
matches!(self, Self::RolledBack { .. } | Self::Rejected { .. })
|| matches!(self, Self::Committed { granules: 0, .. })
}
pub fn has_poisoned_range(&self) -> bool {
matches!(
self,
Self::Fatal {
poisoned_range: Some(_),
..
}
)
}
}
pub struct VerifiedSafePoint(#[allow(dead_code)] ResizeSafePoint);
pub fn verify_safe_point(point: ResizeSafePoint) -> Result<VerifiedSafePoint, &'static str> {
match point.blocking_reason() {
Some(reason) => Err(reason),
None => Ok(VerifiedSafePoint(point)),
}
}
#[allow(clippy::too_many_arguments)]
pub fn transition_granule_range(
runtime: &CudaRuntime,
reservation: &mut CudaReservation,
backing: &CudaVirtualBacking,
offset: usize,
len: usize,
new_location: PhysicalLocation,
old_pool: &Arc<PhysicalHandlePool>,
new_pool: &Arc<PhysicalHandlePool>,
safe_point: &VerifiedSafePoint,
recheck_safe_point: impl Fn() -> ResizeSafePoint,
) -> TransitionOutcome {
transition_granule_range_inner(
runtime,
reservation,
backing,
offset,
len,
new_location,
old_pool,
new_pool,
safe_point,
recheck_safe_point,
#[cfg(any(test, feature = "gpu-tests"))]
None,
)
}
#[cfg(any(test, feature = "gpu-tests"))]
#[allow(clippy::too_many_arguments)]
pub fn transition_granule_range_with_phase8_faults(
runtime: &CudaRuntime,
reservation: &mut CudaReservation,
backing: &CudaVirtualBacking,
offset: usize,
len: usize,
new_location: PhysicalLocation,
old_pool: &Arc<PhysicalHandlePool>,
new_pool: &Arc<PhysicalHandlePool>,
safe_point: &VerifiedSafePoint,
recheck_safe_point: impl Fn() -> ResizeSafePoint,
phase8_faults: Arc<onnx_runtime_cuda_memory::release::DriverFaultPlan>,
) -> TransitionOutcome {
transition_granule_range_inner(
runtime,
reservation,
backing,
offset,
len,
new_location,
old_pool,
new_pool,
safe_point,
recheck_safe_point,
Some(phase8_faults),
)
}
#[allow(clippy::too_many_arguments)]
fn transition_granule_range_inner(
runtime: &CudaRuntime,
reservation: &mut CudaReservation,
backing: &CudaVirtualBacking,
offset: usize,
len: usize,
new_location: PhysicalLocation,
old_pool: &Arc<PhysicalHandlePool>,
new_pool: &Arc<PhysicalHandlePool>,
_safe_point: &VerifiedSafePoint,
recheck_safe_point: impl Fn() -> ResizeSafePoint,
#[cfg(any(test, feature = "gpu-tests"))] phase8_faults: Option<
Arc<onnx_runtime_cuda_memory::release::DriverFaultPlan>,
>,
) -> TransitionOutcome {
if len == 0 {
return TransitionOutcome::Committed {
granules: 0,
new_owned_bytes: 0,
old_released_bytes: 0,
};
}
if old_pool.location() == new_location {
return TransitionOutcome::Committed {
granules: 0,
new_owned_bytes: 0,
old_released_bytes: 0,
};
}
let granularity = old_pool.granularity();
if granularity == 0 || !offset.is_multiple_of(granularity) || !len.is_multiple_of(granularity) {
return TransitionOutcome::RolledBack {
fault: DriverFault::new(
"transition_granule_range",
format!("offset/len not aligned to granularity {granularity}"),
),
};
}
let granule_count = len / granularity;
let old_blocks = backing.blocks_in_range_pub(reservation, offset, len);
if old_blocks.len() != granule_count {
return TransitionOutcome::RolledBack {
fault: DriverFault::new(
"transition_granule_range precondition",
format!(
"expected {granule_count} committed blocks in [{offset}, {}), found {}",
offset + len,
old_blocks.len()
),
),
};
}
let mut new_handles: Vec<(cu::CUmemGenericAllocationHandle, u64)> =
Vec::with_capacity(granule_count);
let mut acquire_fault: Option<DriverFault> = None;
for _ in 0..granule_count {
match new_pool.acquire_handle_raw() {
Ok((h, c)) => new_handles.push((h, c)),
Err(e) => {
acquire_fault = Some(DriverFault::new(
"acquire_handle_raw (new-location handle)",
e.to_string(),
));
break;
}
}
}
if let Some(fault) = acquire_fault {
for (h, _) in new_handles {
new_pool.return_handle_unmapped(h);
}
return TransitionOutcome::RolledBack { fault };
}
let staging_result = <CudaVirtualBacking as VirtualBacking>::reserve(backing, len);
let staging_reservation = match staging_result {
Ok(r) => r,
Err(e) => {
for (h, _) in new_handles {
new_pool.return_handle_unmapped(h);
}
return TransitionOutcome::RolledBack {
fault: DriverFault::new("cuMemAddressReserve (staging VMM)", e.to_string()),
};
}
};
let staging_base: CUdeviceptr = staging_reservation.base_ptr();
let stable_base: CUdeviceptr = reservation.base_ptr();
let device_ordinal = new_pool.device_ordinal_pub();
let mut staging_mapped: usize = 0;
let mut phase4_fault: Option<DriverFault> = None;
{
let _section = capture_gate::synchronizing_section();
for (i, &(handle, _)) in new_handles.iter().enumerate() {
let addr = staging_base + (i * granularity) as u64;
if unsafe { cu::cuMemMap(addr, granularity, 0, handle, 0) }
!= cu::CUresult::CUDA_SUCCESS
{
phase4_fault = Some(DriverFault::new(
"cuMemMap (staging VMM, new handle)",
format!("granule {i}"),
));
break;
}
staging_mapped += 1;
let mut access: cu::CUmemAccessDesc = unsafe { std::mem::zeroed() };
access.location.type_ = cu::CUmemLocationType::CU_MEM_LOCATION_TYPE_DEVICE;
access.location.id = device_ordinal;
access.flags = cu::CUmemAccess_flags::CU_MEM_ACCESS_FLAGS_PROT_READWRITE;
if unsafe { cu::cuMemSetAccess(addr, granularity, &access, 1) }
!= cu::CUresult::CUDA_SUCCESS
{
unsafe {
let _ = cu::cuMemUnmap(addr, granularity);
}
staging_mapped -= 1;
phase4_fault = Some(DriverFault::new(
"cuMemSetAccess (staging VMM, new handle)",
format!("granule {i}"),
));
break;
}
}
}
if let Some(fault) = phase4_fault {
{
let _section = capture_gate::synchronizing_section();
for i in 0..staging_mapped {
unsafe {
let _ = cu::cuMemUnmap(staging_base + (i * granularity) as u64, granularity);
}
}
}
drop(staging_reservation);
for (h, _) in new_handles {
new_pool.return_handle_unmapped(h);
}
return TransitionOutcome::RolledBack { fault };
}
if let Err(e) = unsafe { runtime.dtod_async(stable_base + offset as u64, staging_base, len) } {
{
let _section = capture_gate::synchronizing_section();
for i in 0..granule_count {
unsafe {
let _ = cu::cuMemUnmap(staging_base + (i * granularity) as u64, granularity);
}
}
}
drop(staging_reservation);
for (h, _) in new_handles {
new_pool.return_handle_unmapped(h);
}
return TransitionOutcome::RolledBack {
fault: DriverFault::new("cuMemcpyDtoDAsync (stable VA → staging VMM)", e.to_string()),
};
}
if let Err(e) = runtime.drain_for_unmap() {
{
let _section = capture_gate::synchronizing_section();
for i in 0..granule_count {
unsafe {
let _ = cu::cuMemUnmap(staging_base + (i * granularity) as u64, granularity);
}
}
}
drop(staging_reservation);
for (h, _) in new_handles {
new_pool.return_handle_unmapped(h);
}
return TransitionOutcome::RolledBack {
fault: DriverFault::new("drain_for_unmap (reader drain + copy sync)", e.to_string()),
};
}
let recheck = recheck_safe_point();
if let Some(reason) = recheck.blocking_reason() {
{
let _section = capture_gate::synchronizing_section();
for i in 0..granule_count {
unsafe {
let _ = cu::cuMemUnmap(staging_base + (i * granularity) as u64, granularity);
}
}
}
drop(staging_reservation);
for (h, _) in new_handles {
new_pool.return_handle_unmapped(h);
}
return TransitionOutcome::Rejected { reason };
}
let mut committed_count = 0usize;
let mut total_new_owned: u64 = 0;
let mut total_old_released: u64 = 0;
struct FatalState {
transition_fault: DriverFault,
rollback_fault: Option<DriverFault>,
quarantined: Vec<MappedBlock>,
poisoned_range: Option<(usize, usize)>,
}
let mut fatal: Option<FatalState> = None;
let set_access = |addr: CUdeviceptr| -> bool {
let mut access: cu::CUmemAccessDesc = unsafe { std::mem::zeroed() };
access.location.type_ = cu::CUmemLocationType::CU_MEM_LOCATION_TYPE_DEVICE;
access.location.id = device_ordinal;
access.flags = cu::CUmemAccess_flags::CU_MEM_ACCESS_FLAGS_PROT_READWRITE;
unsafe { cu::cuMemSetAccess(addr, granularity, &access, 1) == cu::CUresult::CUDA_SUCCESS }
};
#[cfg(any(test, feature = "gpu-tests"))]
let should_inject = |op: onnx_runtime_cuda_memory::release::DriverOperation| -> bool {
phase8_faults
.as_ref()
.is_some_and(|plan| plan.should_fail(op))
};
let unmap_checked = |addr: CUdeviceptr| -> bool {
#[cfg(any(test, feature = "gpu-tests"))]
if should_inject(onnx_runtime_cuda_memory::release::DriverOperation::Unmap) {
return false;
}
unsafe { cu::cuMemUnmap(addr, granularity) == cu::CUresult::CUDA_SUCCESS }
};
let map_checked = |addr: CUdeviceptr, handle: cu::CUmemGenericAllocationHandle| -> bool {
#[cfg(any(test, feature = "gpu-tests"))]
if should_inject(onnx_runtime_cuda_memory::release::DriverOperation::Remap) {
return false;
}
unsafe { cu::cuMemMap(addr, granularity, 0, handle, 0) == cu::CUresult::CUDA_SUCCESS }
};
let set_access_checked = |addr: CUdeviceptr| -> bool {
#[cfg(any(test, feature = "gpu-tests"))]
if should_inject(onnx_runtime_cuda_memory::release::DriverOperation::SetAccess) {
return false;
}
set_access(addr)
};
{
let _section = capture_gate::synchronizing_section();
'granules: for (i, &(new_handle, charged)) in new_handles.iter().enumerate() {
let old_block = old_blocks[i];
let old_handle = old_block.handle;
let stable_offset = old_block.offset;
let stable_addr = stable_base + stable_offset as u64;
let staging_addr = staging_base + (i * granularity) as u64;
let cleanup_remaining_staging_and_new = || {
for j in i..granule_count {
unsafe {
let _ =
cu::cuMemUnmap(staging_base + (j * granularity) as u64, granularity);
}
}
for &(handle, _) in &new_handles[i..] {
new_pool.return_handle_unmapped(handle);
}
};
if !unmap_checked(stable_addr) {
cleanup_remaining_staging_and_new();
let fault = DriverFault::new("cuMemUnmap (old, stable VA)", format!("granule {i}"));
if committed_count == 0 {
drop(_section);
drop(staging_reservation);
return TransitionOutcome::RolledBack { fault };
}
fatal = Some(FatalState {
transition_fault: fault,
rollback_fault: None,
quarantined: Vec::new(),
poisoned_range: None,
});
break 'granules;
}
if !map_checked(stable_addr, new_handle) {
let transition_fault =
DriverFault::new("cuMemMap (new, stable VA)", format!("granule {i}"));
let restore_map_ok = map_checked(stable_addr, old_handle);
let restore_ok = restore_map_ok && set_access_checked(stable_addr);
cleanup_remaining_staging_and_new();
if !restore_ok {
if restore_map_ok {
let _ = unsafe { cu::cuMemUnmap(stable_addr, granularity) };
}
reservation.push_quarantined_block(old_block);
let poisoned = MappedBlock::new(stable_offset, granularity, old_handle);
fatal = Some(FatalState {
transition_fault,
rollback_fault: Some(DriverFault::new(
"cuMemMap (restore old, stable VA)",
format!("granule {i}: restore also failed"),
)),
quarantined: vec![poisoned],
poisoned_range: Some((stable_offset, granularity)),
});
break 'granules;
}
if committed_count == 0 {
drop(_section);
drop(staging_reservation);
return TransitionOutcome::RolledBack {
fault: transition_fault,
};
}
fatal = Some(FatalState {
transition_fault,
rollback_fault: None,
quarantined: Vec::new(),
poisoned_range: None,
});
break 'granules;
}
if !set_access_checked(stable_addr) {
let _ = unmap_checked(stable_addr);
let transition_fault =
DriverFault::new("cuMemSetAccess (new, stable VA)", format!("granule {i}"));
let restore_map_ok = map_checked(stable_addr, old_handle);
let restore_ok = restore_map_ok && set_access_checked(stable_addr);
cleanup_remaining_staging_and_new();
if !restore_ok {
if restore_map_ok {
let _ = unmap_checked(stable_addr);
}
reservation.push_quarantined_block(old_block);
let poisoned = MappedBlock::new(stable_offset, granularity, old_handle);
fatal = Some(FatalState {
transition_fault,
rollback_fault: Some(DriverFault::new(
"cuMemMap/cuMemSetAccess (restore old after set-access fail)",
format!("granule {i}"),
)),
quarantined: vec![poisoned],
poisoned_range: Some((stable_offset, granularity)),
});
break 'granules;
}
if committed_count == 0 {
drop(_section);
drop(staging_reservation);
return TransitionOutcome::RolledBack {
fault: transition_fault,
};
}
fatal = Some(FatalState {
transition_fault,
rollback_fault: None,
quarantined: Vec::new(),
poisoned_range: None,
});
break 'granules;
}
unsafe {
let _ = cu::cuMemUnmap(staging_addr, granularity);
}
let _ = old_pool.return_after_unmap_pub(old_handle);
reservation.swap_block(
old_block,
MappedBlock::new(stable_offset, granularity, new_handle),
Arc::clone(new_pool),
);
total_new_owned = total_new_owned.saturating_add(charged);
total_old_released = total_old_released.saturating_add(granularity as u64);
committed_count += 1;
}
}
drop(staging_reservation);
if let Some(state) = fatal {
return TransitionOutcome::Fatal {
transition_fault: state.transition_fault,
rollback_fault: state.rollback_fault,
committed_count,
quarantined: state.quarantined,
poisoned_range: state.poisoned_range,
};
}
TransitionOutcome::Committed {
granules: committed_count,
new_owned_bytes: total_new_owned,
old_released_bytes: total_old_released,
}
}
#[derive(Clone, Debug, Default)]
pub struct TransitionTimings {
pub drain_us: f64,
pub total_us: f64,
}
#[allow(clippy::too_many_arguments)]
pub fn transition_granule_range_timed(
runtime: &CudaRuntime,
reservation: &mut CudaReservation,
backing: &CudaVirtualBacking,
offset: usize,
len: usize,
new_location: PhysicalLocation,
old_pool: &Arc<PhysicalHandlePool>,
new_pool: &Arc<PhysicalHandlePool>,
safe_point: &VerifiedSafePoint,
recheck_safe_point: impl Fn() -> ResizeSafePoint,
timings: &mut TransitionTimings,
) -> TransitionOutcome {
let t0 = std::time::Instant::now();
let result = transition_granule_range(
runtime,
reservation,
backing,
offset,
len,
new_location,
old_pool,
new_pool,
safe_point,
recheck_safe_point,
);
timings.total_us = t0.elapsed().as_secs_f64() * 1e6;
result
}