Skip to main content

tenferro_gpu/cubecl/
workspace_retirement.rs

1//! Deferred retirement of vendor device workspaces.
2//!
3//! A retired workspace returns its CubeCL handle to the shared pool, so the
4//! handle may only be released once the vendor work that used it has completed
5//! on its stream. `Workspace::drop` used to synchronize that stream, which
6//! drains the pipeline: `gpu/tensornetwork` trace execution spent 54% of each
7//! call with the device idle between kernels because plan-cache eviction
8//! synchronized once per evicted workspace.
9//!
10//! Retirement instead records a CUDA event on the workspace's stream and
11//! releases the handle once the event reports completion. The event is the
12//! completion witness that the stream barrier provided before, so a block is
13//! still never returned to the pool while vendor work may reference it.
14
15use std::collections::VecDeque;
16
17use cudarc::driver::result as cuda_result;
18use cudarc::driver::sys::{CUevent, CUevent_flags, CUstream};
19
20use super::runtime::CudaRuntimeState;
21
22/// In-flight retirements allowed before retirement falls back to a barrier.
23///
24/// The queue exists to avoid draining the stream on eviction, not to buffer an
25/// unbounded number of workspaces. Retirements normally resolve within the next
26/// contraction, so the depth stays in the single digits.
27pub(crate) const DEFAULT_WORKSPACE_RETIREMENT_CAPACITY: usize = 16;
28
29/// One workspace waiting for its stream to reach a recorded event.
30#[derive(Debug)]
31struct RetiredWorkspace {
32    event: CUevent,
33    handle: cubecl_runtime::server::Handle,
34    stream: u64,
35}
36
37/// Counters for deferred workspace retirement.
38#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
39pub struct WorkspaceRetirementStats {
40    /// Workspaces handed to the queue instead of a stream barrier.
41    pub deferred: u64,
42    /// Queued workspaces whose handle has been released.
43    pub released: u64,
44    /// Retirements that fell back to synchronizing the stream.
45    pub barrier_fallbacks: u64,
46    /// Handles leaked because no completion could be proven.
47    pub leaked: u64,
48    /// Retirements currently waiting for completion.
49    pub in_flight: usize,
50    /// High-water mark of `in_flight`.
51    pub max_in_flight: usize,
52}
53
54impl WorkspaceRetirementStats {
55    fn record_deferred(&mut self) {
56        self.deferred += 1;
57        self.in_flight += 1;
58        self.max_in_flight = self.max_in_flight.max(self.in_flight);
59    }
60}
61
62/// Bounded queue of workspaces awaiting stream completion.
63#[derive(Debug)]
64pub(crate) struct WorkspaceRetirementQueue {
65    entries: VecDeque<RetiredWorkspace>,
66    capacity: usize,
67    stats: WorkspaceRetirementStats,
68}
69
70impl Default for WorkspaceRetirementQueue {
71    fn default() -> Self {
72        Self::new(DEFAULT_WORKSPACE_RETIREMENT_CAPACITY)
73    }
74}
75
76impl WorkspaceRetirementQueue {
77    pub(crate) fn new(capacity: usize) -> Self {
78        Self {
79            entries: VecDeque::new(),
80            capacity,
81            stats: WorkspaceRetirementStats::default(),
82        }
83    }
84
85    pub(crate) fn stats(&self) -> WorkspaceRetirementStats {
86        self.stats
87    }
88
89    /// Defer `handle` until the work already enqueued on `stream` completes.
90    ///
91    /// Falls back to a stream barrier when no event can be recorded, when the
92    /// queue is at capacity, or when a previous retirement could not be
93    /// resolved. A handle is never released without a completion witness: if
94    /// even the barrier fails, the handle is leaked.
95    pub(crate) fn retire(
96        &mut self,
97        runtime: &CudaRuntimeState,
98        stream: u64,
99        handle: cubecl_runtime::server::Handle,
100    ) {
101        self.drain(runtime);
102        if self.entries.len() >= self.capacity {
103            self.stats.barrier_fallbacks += 1;
104            self.barrier_oldest(runtime);
105        }
106        match self.record_event(runtime, stream) {
107            Some(event) => {
108                self.entries.push_back(RetiredWorkspace {
109                    event,
110                    handle,
111                    stream,
112                });
113                self.stats.record_deferred();
114            }
115            None => {
116                self.stats.barrier_fallbacks += 1;
117                self.barrier(runtime, stream, handle);
118            }
119        }
120    }
121
122    /// Release every queued workspace whose event has completed.
123    pub(crate) fn drain(&mut self, runtime: &CudaRuntimeState) {
124        let mut index = 0;
125        while index < self.entries.len() {
126            let completed = {
127                let entry = &self.entries[index];
128                unsafe { cuda_result::event::query(entry.event) }.is_ok()
129            };
130            if completed {
131                let entry = self
132                    .entries
133                    .remove(index)
134                    .expect("index is within the retirement queue");
135                self.release(runtime, entry);
136            } else {
137                index += 1;
138            }
139        }
140    }
141
142    /// Resolve every queued retirement before returning, for explicit barriers
143    /// and teardown. Waits on each recorded event rather than on the stream, so
144    /// work enqueued after the retirement point stays asynchronous.
145    pub(crate) fn drain_blocking(&mut self, runtime: &CudaRuntimeState) {
146        while let Some(entry) = self.entries.pop_front() {
147            self.wait_for(runtime, entry);
148        }
149    }
150
151    fn record_event(&self, runtime: &CudaRuntimeState, stream: u64) -> Option<CUevent> {
152        if runtime
153            .set_current_cuda_context("cutensor_workspace_retire")
154            .is_err()
155        {
156            return None;
157        }
158        let event = cuda_result::event::create(CUevent_flags::CU_EVENT_DISABLE_TIMING).ok()?;
159        // SAFETY: `event` was just created and `stream` is a live CUDA stream
160        // owned by this runtime; recording binds the event to that stream.
161        if unsafe { cuda_result::event::record(event, stream as CUstream) }.is_err() {
162            // SAFETY: the event is not recorded anywhere and is destroyed once.
163            let _ = unsafe { cuda_result::event::destroy(event) };
164            return None;
165        }
166        Some(event)
167    }
168
169    fn barrier_oldest(&mut self, runtime: &CudaRuntimeState) {
170        if let Some(entry) = self.entries.pop_front() {
171            self.wait_for(runtime, entry);
172        }
173    }
174
175    /// Block until this retirement's recorded point is reached, then release.
176    fn wait_for(&mut self, runtime: &CudaRuntimeState, entry: RetiredWorkspace) {
177        let waited = runtime
178            .set_current_cuda_context("cutensor_workspace_wait")
179            .is_ok()
180            // SAFETY: `entry.event` was recorded on `entry.stream` and is
181            // destroyed exactly once below.
182            && unsafe { cuda_result::event::synchronize(entry.event) }.is_ok();
183        if !waited
184            && runtime
185                .synchronize_raw_stream(entry.stream, "cutensor_workspace_wait")
186                .is_err()
187        {
188            self.stats.leaked += 1;
189            // SAFETY: a leaked event is never recorded or queried again.
190            let _ = unsafe { cuda_result::event::destroy(entry.event) };
191            std::mem::forget(entry.handle);
192            self.stats.in_flight = self.stats.in_flight.saturating_sub(1);
193            return;
194        }
195        self.release(runtime, entry);
196    }
197
198    /// Release a retirement whose completion has already been observed.
199    fn release(&mut self, _runtime: &CudaRuntimeState, entry: RetiredWorkspace) {
200        // SAFETY: the caller observed completion (event query, event wait, or a
201        // stream barrier) and the event is destroyed exactly once here.
202        let _ = unsafe { cuda_result::event::destroy(entry.event) };
203        drop(entry.handle);
204        self.stats.released += 1;
205        self.stats.in_flight = self.stats.in_flight.saturating_sub(1);
206    }
207
208    fn barrier(
209        &mut self,
210        runtime: &CudaRuntimeState,
211        stream: u64,
212        handle: cubecl_runtime::server::Handle,
213    ) {
214        if runtime
215            .synchronize_raw_stream(stream, "cutensor_workspace_drop")
216            .is_err()
217        {
218            self.stats.leaked += 1;
219            std::mem::forget(handle);
220            return;
221        }
222        drop(handle);
223        self.stats.released += 1;
224    }
225}