use std::any::Any;
use std::collections::VecDeque;
use std::fmt::Debug;
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, Condvar, Mutex, Weak};
use std::time::{Duration, Instant};
use onnx_runtime_memory_governor::{
AllocationReleaseOutcome, AllocationReleaseState, DeferredEnqueueError,
DeferredEnqueueRejection, DeferredReleaseQueue, PreparedAllocationRelease, QuarantineReason,
};
pub const DEFAULT_DEFERRED_RELEASE_CAPACITY: usize = usize::MAX;
pub const DEFAULT_POLL_INTERVAL: Duration = Duration::from_micros(250);
pub trait ReleaseFence: Send + Sync + Debug {
fn is_complete(&self) -> bool;
fn retain_after_device_loss(self: Box<Self>) {
std::mem::forget(self);
}
}
pub trait ReleaseFenceSource: Send + Sync + Debug {
fn record(&self) -> Result<Vec<Box<dyn ReleaseFence>>, String>;
}
pub struct RetainedOwnership {
pub bytes: u64,
pub detail: String,
pub keep_alive: Box<dyn Any + Send>,
}
impl Debug for RetainedOwnership {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("RetainedOwnership")
.field("bytes", &self.bytes)
.field("detail", &self.detail)
.finish_non_exhaustive()
}
}
#[derive(Debug)]
pub struct DeferredActionOutcome {
pub state: AllocationReleaseState,
pub unmapped_bytes: u64,
pub detail: Option<String>,
pub retained: Option<RetainedOwnership>,
}
impl DeferredActionOutcome {
pub fn released(unmapped_bytes: u64) -> Self {
Self {
state: AllocationReleaseState::Released,
unmapped_bytes,
detail: None,
retained: None,
}
}
pub fn quarantined(
state: AllocationReleaseState,
unmapped_bytes: u64,
detail: impl Into<String>,
retained: Option<RetainedOwnership>,
) -> Self {
Self {
state,
unmapped_bytes,
detail: Some(detail.into()),
retained,
}
}
pub fn is_complete(&self) -> bool {
self.state == AllocationReleaseState::Released
}
}
pub trait DeferredReleaseAction: Send + Debug + 'static {
fn execute(self: Box<Self>) -> DeferredActionOutcome;
fn settle_device_lost(self: Box<Self>, detail: &str) -> Option<RetainedOwnership> {
let bytes = self.bytes();
Some(RetainedOwnership {
bytes,
detail: detail.to_owned(),
keep_alive: Box::new(self),
})
}
fn label(&self) -> &'static str;
fn bytes(&self) -> u64 {
0
}
}
#[derive(Debug)]
pub struct RefusedRelease<A> {
pub rejection: DeferredEnqueueRejection,
pub action: A,
}
impl<A> RefusedRelease<A> {
fn new(rejection: DeferredEnqueueRejection, action: A) -> Self {
Self { rejection, action }
}
}
#[derive(Debug)]
pub struct RetainedRelease {
pub label: &'static str,
pub state: AllocationReleaseState,
pub bytes: u64,
pub detail: String,
#[allow(dead_code)]
ownership: Option<RetainedOwnership>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RetainedReleaseInfo {
pub label: &'static str,
pub state: AllocationReleaseState,
pub bytes: u64,
pub detail: String,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct DeferredReleaseStats {
pub pending: usize,
pub accepted: u64,
pub completed: u64,
pub quarantined: u64,
pub enqueue_failures: u64,
pub mapped_refunded_bytes: u64,
pub closed: bool,
pub draining: bool,
pub device_lost: bool,
pub retained: usize,
}
#[derive(Debug)]
struct PendingRelease {
fences: Vec<Box<dyn ReleaseFence>>,
action: Box<dyn DeferredReleaseAction>,
}
#[derive(Debug, Default)]
struct QueueState {
pending: VecDeque<PendingRelease>,
retained: Vec<RetainedRelease>,
closed: bool,
draining: bool,
device_lost: bool,
worker_started: bool,
}
#[derive(Debug, Default)]
struct ExecutionGateState {
owner: Option<std::thread::ThreadId>,
depth: usize,
}
#[derive(Debug, Default)]
struct ExecutionGate {
state: Mutex<ExecutionGateState>,
wake: Condvar,
}
impl ExecutionGate {
fn lock(&self) -> ExecutionGateGuard<'_> {
let current = std::thread::current().id();
let mut state = self
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
while state.owner.as_ref().is_some_and(|owner| *owner != current) {
state = self
.wake
.wait(state)
.unwrap_or_else(|poisoned| poisoned.into_inner());
}
state.owner = Some(current);
state.depth = state
.depth
.checked_add(1)
.expect("execution-gate recursion depth overflow");
ExecutionGateGuard { gate: self }
}
}
struct ExecutionGateGuard<'a> {
gate: &'a ExecutionGate,
}
impl Drop for ExecutionGateGuard<'_> {
fn drop(&mut self) {
let current = std::thread::current().id();
let mut state = self
.gate
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
debug_assert_eq!(state.owner.as_ref(), Some(¤t));
state.depth = state.depth.saturating_sub(1);
if state.depth == 0 {
state.owner = None;
self.gate.wake.notify_all();
}
}
}
#[derive(Debug, Default)]
struct Counters {
accepted: AtomicU64,
completed: AtomicU64,
quarantined: AtomicU64,
enqueue_failures: AtomicU64,
mapped_refunded_bytes: AtomicU64,
}
pub struct CudaDeferredReleaseQueue {
me: Weak<Self>,
fences: Box<dyn ReleaseFenceSource>,
capacity: usize,
poll_interval: Duration,
autonomous: bool,
execution_gate: ExecutionGate,
state: Mutex<QueueState>,
wake: Condvar,
outstanding: AtomicUsize,
device_lost: AtomicBool,
closed: AtomicBool,
draining: AtomicBool,
drain_callback: Mutex<Option<Box<dyn FnMut() -> bool + Send>>>,
counters: Counters,
}
impl Debug for CudaDeferredReleaseQueue {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let stats = self.stats();
formatter
.debug_struct("CudaDeferredReleaseQueue")
.field("capacity", &self.capacity)
.field("pending", &stats.pending)
.field("completed", &stats.completed)
.field("quarantined", &stats.quarantined)
.field("closed", &stats.closed)
.field("device_lost", &stats.device_lost)
.finish_non_exhaustive()
}
}
impl CudaDeferredReleaseQueue {
pub fn new(fences: Box<dyn ReleaseFenceSource>, capacity: usize) -> Arc<Self> {
Self::build(fences, capacity, DEFAULT_POLL_INTERVAL, true)
}
pub fn manual(fences: Box<dyn ReleaseFenceSource>, capacity: usize) -> Arc<Self> {
Self::build(fences, capacity, DEFAULT_POLL_INTERVAL, false)
}
pub fn with_poll_interval(
fences: Box<dyn ReleaseFenceSource>,
capacity: usize,
poll_interval: Duration,
) -> Arc<Self> {
Self::build(fences, capacity, poll_interval, true)
}
fn build(
fences: Box<dyn ReleaseFenceSource>,
capacity: usize,
poll_interval: Duration,
autonomous: bool,
) -> Arc<Self> {
Arc::new_cyclic(|me| Self {
me: me.clone(),
fences,
capacity: capacity.max(1),
poll_interval,
autonomous,
execution_gate: ExecutionGate::default(),
state: Mutex::new(QueueState::default()),
wake: Condvar::new(),
outstanding: AtomicUsize::new(0),
device_lost: AtomicBool::new(false),
closed: AtomicBool::new(false),
draining: AtomicBool::new(false),
drain_callback: Mutex::new(None),
counters: Counters::default(),
})
}
fn lock(&self) -> std::sync::MutexGuard<'_, QueueState> {
self.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
pub fn capacity(&self) -> usize {
self.capacity
}
pub fn pending(&self) -> usize {
self.outstanding.load(Ordering::Acquire)
}
pub fn is_closed(&self) -> bool {
self.closed.load(Ordering::Acquire)
}
pub fn is_device_lost(&self) -> bool {
self.device_lost.load(Ordering::Acquire)
}
pub fn is_draining(&self) -> bool {
self.draining.load(Ordering::Acquire)
}
pub fn stats(&self) -> DeferredReleaseStats {
let retained = self.lock().retained.len();
DeferredReleaseStats {
pending: self.pending(),
accepted: self.counters.accepted.load(Ordering::Relaxed),
completed: self.counters.completed.load(Ordering::Relaxed),
quarantined: self.counters.quarantined.load(Ordering::Relaxed),
enqueue_failures: self.counters.enqueue_failures.load(Ordering::Relaxed),
mapped_refunded_bytes: self.counters.mapped_refunded_bytes.load(Ordering::Relaxed),
closed: self.is_closed(),
draining: self.is_draining(),
device_lost: self.is_device_lost(),
retained,
}
}
pub fn retained(&self) -> Vec<RetainedReleaseInfo> {
self.lock()
.retained
.iter()
.map(|record| RetainedReleaseInfo {
label: record.label,
state: record.state,
bytes: record.bytes,
detail: record.detail.clone(),
})
.collect()
}
pub fn close(&self) {
self.closed.store(true, Ordering::Release);
self.lock().closed = true;
self.wake.notify_all();
}
pub fn close_after_drain(&self) {
self.draining.store(true, Ordering::Release);
self.lock().draining = true;
self.run_drain_callback_if_ready();
self.wake.notify_all();
}
pub fn set_drain_callback(&self, callback: impl FnMut() -> bool + Send + 'static) {
*self
.drain_callback
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = Some(Box::new(callback));
self.run_drain_callback_if_ready();
self.wake.notify_all();
}
fn run_drain_callback_if_ready(&self) {
if !self.is_draining() || self.is_device_lost() || self.pending() != 0 {
return;
}
let callback = self
.drain_callback
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.take();
let Some(mut callback) = callback else {
return;
};
if !callback() {
let mut slot = self
.drain_callback
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if slot.is_none() {
*slot = Some(callback);
}
}
}
pub fn mark_device_lost(&self, reason: impl Into<String>) {
let reason = reason.into();
let gate = self.execution_gate.lock();
if self.device_lost.swap(true, Ordering::AcqRel) {
return;
}
drop(gate);
{
let mut state = self.lock();
state.device_lost = true;
}
self.wake.notify_all();
self.retain_all_pending(&format!("device lost: {reason}"));
self.closed.store(true, Ordering::Release);
self.lock().closed = true;
self.wake.notify_all();
}
pub fn poll(&self) -> usize {
if self.is_device_lost() {
self.retain_all_pending("device lost before the release could be ordered");
return 0;
}
let drained: Vec<PendingRelease> = {
let mut state = self.lock();
state.pending.drain(..).collect()
};
if drained.is_empty() {
self.run_drain_callback_if_ready();
return 0;
}
let mut carry = VecDeque::with_capacity(drained.len());
let mut executed = 0usize;
let mut retained = Vec::new();
for entry in drained {
let execution = self.execution_gate.lock();
if self.is_device_lost() {
drop(execution);
retained.push(self.settle_lost_entry(
entry,
"device lost while a poller owned the deferred release",
));
continue;
}
if !entry.fences.iter().all(|fence| fence.is_complete()) {
drop(execution);
carry.push_back(entry);
continue;
}
let PendingRelease { fences, action } = entry;
let label = action.label();
drop(fences);
let outcome = action.execute();
drop(execution);
executed += 1;
self.record_outcome(label, outcome, &mut retained);
self.outstanding.fetch_sub(1, Ordering::AcqRel);
}
let carry_gate = self.execution_gate.lock();
if self.is_device_lost() {
while let Some(entry) = carry.pop_front() {
retained.push(self.settle_lost_entry(
entry,
"device lost after a poller observed an incomplete release fence",
));
}
}
if !carry.is_empty() || !retained.is_empty() {
let mut state = self.lock();
state.retained.append(&mut retained);
while let Some(entry) = carry.pop_back() {
state.pending.push_front(entry);
}
}
drop(carry_gate);
if executed > 0 {
self.wake.notify_all();
}
self.run_drain_callback_if_ready();
executed
}
fn record_outcome(
&self,
label: &'static str,
outcome: DeferredActionOutcome,
retained: &mut Vec<RetainedRelease>,
) {
if outcome.unmapped_bytes > 0 {
self.counters
.mapped_refunded_bytes
.fetch_add(outcome.unmapped_bytes, Ordering::Relaxed);
}
if outcome.is_complete() {
self.counters.completed.fetch_add(1, Ordering::Relaxed);
return;
}
self.counters.quarantined.fetch_add(1, Ordering::Relaxed);
let detail = outcome
.detail
.unwrap_or_else(|| String::from("the release did not complete"));
let bytes = outcome
.retained
.as_ref()
.map_or(0, |ownership| ownership.bytes);
eprintln!(
"cuda_ep: WARNING: deferred {label} release did not complete ({}): {detail}; \
{bytes} byte(s) of ownership are retained and will not be reused",
outcome.state
);
retained.push(RetainedRelease {
label,
state: outcome.state,
bytes,
detail,
ownership: outcome.retained,
});
}
fn settle_lost_entry(&self, entry: PendingRelease, detail: &str) -> RetainedRelease {
let PendingRelease { fences, action } = entry;
for fence in fences {
fence.retain_after_device_loss();
}
let label = action.label();
let bytes = action.bytes();
let ownership = action.settle_device_lost(detail);
self.counters.quarantined.fetch_add(1, Ordering::Relaxed);
self.outstanding.fetch_sub(1, Ordering::AcqRel);
RetainedRelease {
label,
state: AllocationReleaseState::DeviceLost,
bytes,
detail: detail.to_owned(),
ownership,
}
}
fn retain_all_pending(&self, detail: &str) {
let drained: Vec<PendingRelease> = {
let mut state = self.lock();
state.pending.drain(..).collect()
};
if drained.is_empty() {
return;
}
let mut records = Vec::with_capacity(drained.len());
for entry in drained {
records.push(self.settle_lost_entry(entry, detail));
}
let mut state = self.lock();
state.retained.append(&mut records);
}
pub fn enqueue<A: DeferredReleaseAction + 'static>(
&self,
action: A,
) -> Result<(), RefusedRelease<A>> {
if let Some(rejection) = self.refusal_reason() {
self.counters
.enqueue_failures
.fetch_add(1, Ordering::Relaxed);
return Err(RefusedRelease::new(rejection, action));
}
if let Err(rejection) = self.reserve_slot() {
self.counters
.enqueue_failures
.fetch_add(1, Ordering::Relaxed);
return Err(RefusedRelease::new(rejection, action));
}
let execution = self.execution_gate.lock();
if let Some(rejection) = self.refusal_reason() {
drop(execution);
self.outstanding.fetch_sub(1, Ordering::AcqRel);
self.counters
.enqueue_failures
.fetch_add(1, Ordering::Relaxed);
return Err(RefusedRelease::new(rejection, action));
}
let fences = match self.fences.record() {
Ok(fences) => fences,
Err(error) => {
drop(execution);
self.outstanding.fetch_sub(1, Ordering::AcqRel);
self.counters
.enqueue_failures
.fetch_add(1, Ordering::Relaxed);
eprintln!(
"cuda_ep: WARNING: could not record deferred-release ordering fences for a \
{} release: {error}",
action.label()
);
return Err(RefusedRelease::new(
DeferredEnqueueRejection::Refused,
action,
));
}
};
{
let mut state = self.lock();
if state.closed || state.device_lost {
let rejection = if state.device_lost {
DeferredEnqueueRejection::DeviceLost
} else {
DeferredEnqueueRejection::Closed
};
drop(state);
self.outstanding.fetch_sub(1, Ordering::AcqRel);
self.counters
.enqueue_failures
.fetch_add(1, Ordering::Relaxed);
for fence in fences {
if rejection == DeferredEnqueueRejection::DeviceLost {
fence.retain_after_device_loss();
}
}
return Err(RefusedRelease::new(rejection, action));
}
state.pending.push_back(PendingRelease {
fences,
action: Box::new(action),
});
}
drop(execution);
self.counters.accepted.fetch_add(1, Ordering::Relaxed);
self.wake.notify_all();
self.ensure_worker();
Ok(())
}
pub(crate) fn retain_refused<A: DeferredReleaseAction + 'static>(
&self,
refused: RefusedRelease<A>,
detail: String,
) {
let label = refused.action.label();
let bytes = refused.action.bytes();
let ownership = RetainedOwnership {
bytes,
detail: detail.clone(),
keep_alive: Box::new(refused.action),
};
self.lock().retained.push(RetainedRelease {
label,
state: AllocationReleaseState::Quarantined,
bytes,
detail,
ownership: Some(ownership),
});
self.counters.quarantined.fetch_add(1, Ordering::AcqRel);
}
pub fn enqueue_prepared(
&self,
request: PreparedAllocationRelease,
observer: Option<Arc<dyn ReleaseObserver>>,
) -> Result<(), DeferredEnqueueError> {
let action = PreparedReleaseAction {
request: Some(request),
observer,
};
match self.enqueue(action) {
Ok(()) => Ok(()),
Err(refused) => {
let request = refused
.action
.into_request()
.expect("a refused prepared release still holds its request");
Err(DeferredEnqueueError::new(refused.rejection, request))
}
}
}
fn refusal_reason(&self) -> Option<DeferredEnqueueRejection> {
if self.is_device_lost() {
return Some(DeferredEnqueueRejection::DeviceLost);
}
if self.is_closed() {
return Some(DeferredEnqueueRejection::Closed);
}
None
}
fn reserve_slot(&self) -> Result<(), DeferredEnqueueRejection> {
let mut current = self.outstanding.load(Ordering::Acquire);
loop {
if current >= self.capacity {
return Err(DeferredEnqueueRejection::Full);
}
match self.outstanding.compare_exchange_weak(
current,
current + 1,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return Ok(()),
Err(observed) => current = observed,
}
}
}
fn ensure_worker(&self) {
if !self.autonomous {
return;
}
{
let mut state = self.lock();
if state.worker_started {
return;
}
state.worker_started = true;
}
let Some(queue) = self.me.upgrade() else {
return;
};
let spawned = std::thread::Builder::new()
.name("cuda-deferred-release".into())
.spawn(move || queue.run_worker());
if let Err(error) = spawned {
self.lock().worker_started = false;
eprintln!(
"cuda_ep: WARNING: could not start the deferred-release worker ({error}); \
releases will be drained by later provider calls instead"
);
}
}
fn run_worker(self: Arc<Self>) {
loop {
if !self.is_device_lost() {
self.poll();
}
let mut state = self.lock();
let mut idle =
state.pending.is_empty() && self.outstanding.load(Ordering::Acquire) == 0;
if idle && state.closed {
return;
}
if idle && state.draining && !state.device_lost {
drop(state);
self.run_drain_callback_if_ready();
state = self.lock();
idle = state.pending.is_empty() && self.outstanding.load(Ordering::Acquire) == 0;
}
if idle && state.draining && Arc::strong_count(&self) == 1 {
drop(state);
self.closed.store(true, Ordering::Release);
self.lock().closed = true;
return;
}
let interval = if state.device_lost {
self.poll_interval.max(Duration::from_millis(50))
} else {
self.poll_interval
};
let (guard, _) = self
.wake
.wait_timeout(state, interval)
.unwrap_or_else(|poisoned| poisoned.into_inner());
drop(guard);
}
}
pub fn wait_until_idle(&self, timeout: Duration) -> bool {
let deadline = Instant::now() + timeout;
loop {
self.poll();
if self.pending() == 0 {
return true;
}
if Instant::now() >= deadline {
return false;
}
std::thread::sleep(self.poll_interval.min(Duration::from_millis(1)));
}
}
}
impl DeferredReleaseQueue for CudaDeferredReleaseQueue {
fn enqueue(&self, request: PreparedAllocationRelease) -> Result<(), DeferredEnqueueError> {
self.enqueue_prepared(request, None)
}
fn pending(&self) -> usize {
self.outstanding.load(Ordering::Acquire)
}
}
impl onnx_runtime_memory_governor::DeviceLossListener for CudaDeferredReleaseQueue {
fn mark_device_lost(&self, reason: &str) {
CudaDeferredReleaseQueue::mark_device_lost(self, reason);
}
}
pub trait ReleaseObserver: Send + Sync + Debug {
fn released(&self, outcome: &AllocationReleaseOutcome);
}
#[derive(Debug)]
pub struct PreparedReleaseAction {
request: Option<PreparedAllocationRelease>,
observer: Option<Arc<dyn ReleaseObserver>>,
}
impl PreparedReleaseAction {
pub fn new(
request: PreparedAllocationRelease,
observer: Option<Arc<dyn ReleaseObserver>>,
) -> Self {
Self {
request: Some(request),
observer,
}
}
fn into_request(mut self) -> Option<PreparedAllocationRelease> {
self.request.take()
}
}
impl DeferredReleaseAction for PreparedReleaseAction {
fn execute(mut self: Box<Self>) -> DeferredActionOutcome {
let Some(request) = self.request.take() else {
return DeferredActionOutcome::quarantined(
AllocationReleaseState::Quarantined,
0,
"a prepared release action was executed without its request",
None,
);
};
let allocation_bytes = request.len() as u64;
let allocator = Arc::clone(request.allocator());
let outcome = request.execute();
if let Some(observer) = self.observer.as_ref() {
observer.released(&outcome);
}
match outcome {
AllocationReleaseOutcome::Complete { accounting } => {
DeferredActionOutcome::released(accounting.unmapped_bytes)
}
AllocationReleaseOutcome::Quarantined {
accounting,
residual,
} => 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("binding-prepared provider allocation"),
keep_alive: Box::new(allocator),
}),
),
AllocationReleaseOutcome::Failed { failure } => DeferredActionOutcome::quarantined(
AllocationReleaseState::Quarantined,
0,
failure.to_string(),
Some(RetainedOwnership {
bytes: allocation_bytes,
detail: String::from("binding-prepared provider allocation"),
keep_alive: Box::new(allocator),
}),
),
}
}
fn settle_device_lost(mut self: Box<Self>, detail: &str) -> Option<RetainedOwnership> {
let request = self.request.take()?;
let bytes = request.len() as u64;
let allocator = Arc::clone(request.allocator());
let outcome = request.quarantine_device_lost();
if let Some(observer) = self.observer.as_ref() {
observer.released(&outcome);
}
Some(RetainedOwnership {
bytes,
detail: detail.to_owned(),
keep_alive: Box::new(allocator),
})
}
fn label(&self) -> &'static str {
"provider allocation"
}
fn bytes(&self) -> u64 {
self.request
.as_ref()
.map_or(0, |request| request.len() as u64)
}
}
impl Drop for PreparedReleaseAction {
fn drop(&mut self) {
if let Some(request) = self.request.take() {
let outcome = request.quarantine(QuarantineReason::AbandonedRequest);
if let Some(observer) = self.observer.as_ref() {
observer.released(&outcome);
}
}
}
}
#[derive(Debug)]
pub struct CudaStreamFences {
runtime: Arc<crate::runtime::CudaRuntime>,
}
impl CudaStreamFences {
pub fn new(runtime: Arc<crate::runtime::CudaRuntime>) -> Self {
Self { runtime }
}
}
impl ReleaseFenceSource for CudaStreamFences {
fn record(&self) -> Result<Vec<Box<dyn ReleaseFence>>, String> {
self.runtime
.bind()
.map_err(|error| format!("could not bind the CUDA context: {error}"))?;
let context = self.runtime.cuda_context();
let mut fences: Vec<Box<dyn ReleaseFence>> = Vec::with_capacity(2);
for (stream, name) in [
(self.runtime.stream(), "compute"),
(self.runtime.copy_stream(), "copy"),
] {
let event = context
.new_event(None)
.map_err(|error| format!("cuEventCreate for the {name} stream failed: {error}"))?;
event
.record(stream)
.map_err(|error| format!("cuEventRecord on the {name} stream failed: {error}"))?;
fences.push(Box::new(CudaEventFence { event }));
}
Ok(fences)
}
}
#[derive(Debug)]
struct CudaEventFence {
event: cudarc::driver::CudaEvent,
}
impl ReleaseFence for CudaEventFence {
fn is_complete(&self) -> bool {
self.event.is_complete()
}
}
#[derive(Debug)]
pub struct ReservationTeardownAction {
ticket: Option<crate::virtual_memory::ReservationTeardownTicket>,
}
impl DeferredReleaseAction for ReservationTeardownAction {
fn execute(mut self: Box<Self>) -> DeferredActionOutcome {
let Some(ticket) = self.ticket.take() else {
return DeferredActionOutcome::quarantined(
AllocationReleaseState::Quarantined,
0,
"a reservation teardown action was executed without its ticket",
None,
);
};
let bytes = ticket.len() as u64;
let outcome = ticket.execute_outcome();
let report = outcome.report;
if report.is_complete() {
return DeferredActionOutcome::released(report.unmapped_bytes);
}
DeferredActionOutcome::quarantined(
AllocationReleaseState::PartiallyUnmapped,
report.unmapped_bytes,
format!(
"{} reservation block(s) could not be released, so the {bytes} byte address \
range was not returned to the driver",
report.retained_blocks
),
outcome.retained.map(|ticket| RetainedOwnership {
bytes,
detail: String::from("CUDA reservation teardown"),
keep_alive: Box::new(ticket),
}),
)
}
fn label(&self) -> &'static str {
"reservation teardown"
}
fn bytes(&self) -> u64 {
self.ticket.as_ref().map_or(0, |ticket| ticket.len() as u64)
}
}
impl crate::virtual_memory::DeferredReservationQueue for CudaDeferredReleaseQueue {
fn enqueue_reservation(
&self,
ticket: crate::virtual_memory::ReservationTeardownTicket,
) -> Result<(), crate::virtual_memory::ReservationEnqueueError> {
match self.enqueue(ReservationTeardownAction {
ticket: Some(ticket),
}) {
Ok(()) => Ok(()),
Err(refused) => {
let rejection = match refused.rejection {
DeferredEnqueueRejection::Closed => {
crate::virtual_memory::ReservationEnqueueRejection::Closed
}
DeferredEnqueueRejection::Full => {
crate::virtual_memory::ReservationEnqueueRejection::Full
}
DeferredEnqueueRejection::DeviceLost => {
crate::virtual_memory::ReservationEnqueueRejection::DeviceLost
}
DeferredEnqueueRejection::Refused => {
crate::virtual_memory::ReservationEnqueueRejection::Refused
}
};
Err(crate::virtual_memory::ReservationEnqueueError {
rejection,
ticket: Box::new(
refused
.action
.ticket
.expect("a refused reservation teardown still holds its ticket"),
),
})
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicUsize;
#[derive(Debug, Default)]
struct FakeFence {
complete: Arc<AtomicBool>,
destroyed: Option<Arc<AtomicUsize>>,
}
impl ReleaseFence for FakeFence {
fn is_complete(&self) -> bool {
self.complete.load(Ordering::Acquire)
}
}
impl Drop for FakeFence {
fn drop(&mut self) {
if let Some(counter) = &self.destroyed {
counter.fetch_add(1, Ordering::Relaxed);
}
}
}
#[derive(Debug)]
struct FakeFenceSource {
compute: Arc<AtomicBool>,
copy: Arc<AtomicBool>,
destroyed: Arc<AtomicUsize>,
}
impl ReleaseFenceSource for FakeFenceSource {
fn record(&self) -> Result<Vec<Box<dyn ReleaseFence>>, String> {
Ok(vec![
Box::new(FakeFence {
complete: Arc::clone(&self.compute),
destroyed: Some(Arc::clone(&self.destroyed)),
}),
Box::new(FakeFence {
complete: Arc::clone(&self.copy),
destroyed: Some(Arc::clone(&self.destroyed)),
}),
])
}
}
#[derive(Debug)]
struct CountingAction {
executed: Arc<AtomicUsize>,
}
impl DeferredReleaseAction for CountingAction {
fn execute(self: Box<Self>) -> DeferredActionOutcome {
self.executed.fetch_add(1, Ordering::AcqRel);
DeferredActionOutcome::released(0)
}
fn label(&self) -> &'static str {
"test"
}
}
fn manual_queue(
capacity: usize,
) -> (
Arc<CudaDeferredReleaseQueue>,
Arc<AtomicBool>,
Arc<AtomicBool>,
) {
let compute = Arc::new(AtomicBool::new(false));
let copy = Arc::new(AtomicBool::new(false));
let queue = CudaDeferredReleaseQueue::manual(
Box::new(FakeFenceSource {
compute: Arc::clone(&compute),
copy: Arc::clone(©),
destroyed: Arc::new(AtomicUsize::new(0)),
}),
capacity,
);
(queue, compute, copy)
}
#[test]
fn drain_callback_runs_outside_queue_locks_and_retries_until_cleanup_succeeds() {
let (queue, _, _) = manual_queue(1);
let calls = Arc::new(AtomicUsize::new(0));
let callback_calls = Arc::clone(&calls);
let reentrant = Arc::clone(&queue);
queue.set_drain_callback(move || {
let call = callback_calls.fetch_add(1, Ordering::AcqRel);
let _ = reentrant.stats();
call != 0
});
queue.close_after_drain();
assert_eq!(calls.load(Ordering::Acquire), 1);
queue.poll();
assert_eq!(calls.load(Ordering::Acquire), 2);
queue.poll();
assert_eq!(calls.load(Ordering::Acquire), 2, "callback settles once");
}
#[test]
fn both_stream_fences_must_complete_before_release() {
let (queue, compute, copy) = manual_queue(8);
let executed = Arc::new(AtomicUsize::new(0));
queue
.enqueue(CountingAction {
executed: Arc::clone(&executed),
})
.expect("accepted");
assert_eq!(queue.poll(), 0, "neither stream has completed");
compute.store(true, Ordering::Release);
assert_eq!(queue.poll(), 0, "the copy stream is still in flight");
copy.store(true, Ordering::Release);
assert_eq!(queue.poll(), 1);
assert_eq!(executed.load(Ordering::Acquire), 1);
assert_eq!(queue.pending(), 0);
}
#[test]
fn a_bounded_queue_refuses_and_returns_the_exact_action() {
let (queue, _compute, _copy) = manual_queue(1);
let executed = Arc::new(AtomicUsize::new(0));
queue
.enqueue(CountingAction {
executed: Arc::clone(&executed),
})
.expect("first accepted");
let refused = queue
.enqueue(CountingAction {
executed: Arc::clone(&executed),
})
.expect_err("the bound is enforced");
assert_eq!(refused.rejection, DeferredEnqueueRejection::Full);
assert_eq!(queue.stats().enqueue_failures, 1);
assert_eq!(executed.load(Ordering::Acquire), 0);
drop(refused);
}
#[test]
fn refused_retirement_ownership_is_quarantined_and_accounted() {
let (queue, _compute, _copy) = manual_queue(1);
let executed = Arc::new(AtomicUsize::new(0));
queue
.enqueue(CountingAction {
executed: Arc::clone(&executed),
})
.expect("first accepted");
let refused = queue
.enqueue(CountingAction {
executed: Arc::clone(&executed),
})
.expect_err("the bound is enforced");
queue.retain_refused(refused, "retirement ownership retained".to_string());
assert_eq!(executed.load(Ordering::Acquire), 0);
assert_eq!(queue.stats().quarantined, 1);
let retained = queue.retained();
assert_eq!(retained.len(), 1);
assert_eq!(retained[0].label, "test");
assert_eq!(retained[0].state, AllocationReleaseState::Quarantined);
assert!(retained[0].detail.contains("retirement ownership"));
}
}