tenferro_gpu/cubecl/
workspace_retirement.rs1use 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
22pub(crate) const DEFAULT_WORKSPACE_RETIREMENT_CAPACITY: usize = 16;
28
29#[derive(Debug)]
31struct RetiredWorkspace {
32 event: CUevent,
33 handle: cubecl_runtime::server::Handle,
34 stream: u64,
35}
36
37#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
39pub struct WorkspaceRetirementStats {
40 pub deferred: u64,
42 pub released: u64,
44 pub barrier_fallbacks: u64,
46 pub leaked: u64,
48 pub in_flight: usize,
50 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#[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 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 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 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 if unsafe { cuda_result::event::record(event, stream as CUstream) }.is_err() {
162 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 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 && 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 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 fn release(&mut self, _runtime: &CudaRuntimeState, entry: RetiredWorkspace) {
200 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}