Skip to main content

baracuda_runtime/
stream.rs

1//! Runtime-API streams.
2
3use std::sync::Arc;
4
5use baracuda_cuda_sys::runtime::{cudaStream_t, runtime, types::cudaStreamFlags};
6
7use crate::device::Device;
8use crate::error::{Result, check};
9
10/// An asynchronous work queue on the current CUDA device.
11#[derive(Clone)]
12pub struct Stream {
13    inner: Arc<StreamInner>,
14}
15
16struct StreamInner {
17    handle: cudaStream_t,
18    device: Device,
19}
20
21unsafe impl Send for StreamInner {}
22unsafe impl Sync for StreamInner {}
23
24impl core::fmt::Debug for StreamInner {
25    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
26        f.debug_struct("Stream")
27            .field("handle", &self.handle)
28            .field("device", &self.device)
29            .finish()
30    }
31}
32
33impl core::fmt::Debug for Stream {
34    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
35        self.inner.fmt(f)
36    }
37}
38
39impl Stream {
40    /// Create a stream with default (legacy-default-stream-synchronizing) flags
41    /// on the current device.
42    pub fn new() -> Result<Self> {
43        Self::with_flags(cudaStreamFlags::DEFAULT)
44    }
45
46    /// Create a non-blocking stream — does not synchronize with the legacy
47    /// default stream.
48    pub fn non_blocking() -> Result<Self> {
49        Self::with_flags(cudaStreamFlags::NON_BLOCKING)
50    }
51
52    /// Adopt a raw `cudaStream_t` handle. The wrapper will call
53    /// `cudaStreamDestroy` on drop.
54    ///
55    /// # Safety
56    ///
57    /// `handle` must be a live stream on the current device. Do not
58    /// destroy it externally.
59    pub unsafe fn from_raw(handle: cudaStream_t) -> Self {
60        let device = Device::current().unwrap_or(Device::from_ordinal(0));
61        Self {
62            inner: Arc::new(StreamInner { handle, device }),
63        }
64    }
65
66    /// Create a stream with raw flags (see [`cudaStreamFlags`]).
67    pub fn with_flags(flags: u32) -> Result<Self> {
68        let r = runtime()?;
69        let cu = r.cuda_stream_create_with_flags()?;
70        let mut stream: cudaStream_t = core::ptr::null_mut();
71        check(unsafe { cu(&mut stream, flags) })?;
72        let device = Device::current()?;
73        Ok(Self {
74            inner: Arc::new(StreamInner {
75                handle: stream,
76                device,
77            }),
78        })
79    }
80
81    /// Block the calling thread until all prior work on this stream is complete.
82    pub fn synchronize(&self) -> Result<()> {
83        let r = runtime()?;
84        let cu = r.cuda_stream_synchronize()?;
85        check(unsafe { cu(self.inner.handle) })
86    }
87
88    /// `Ok(true)` if all queued work has finished, `Ok(false)` if work remains.
89    pub fn is_complete(&self) -> Result<bool> {
90        use baracuda_cuda_sys::runtime::cudaError_t;
91        let r = runtime()?;
92        let cu = r.cuda_stream_query()?;
93        match unsafe { cu(self.inner.handle) } {
94            cudaError_t::Success => Ok(true),
95            cudaError_t::NotReady => Ok(false),
96            other => Err(crate::error::Error::Status { status: other }),
97        }
98    }
99
100    /// `Ok(true)` if this stream is idle (every piece of work *this process*
101    /// submitted to it has drained), `Ok(false)` while work is still pending.
102    ///
103    /// A read-only alias of [`is_complete`](Self::is_complete) phrased for
104    /// load probes: it wraps `cudaStreamQuery`, so it never synchronizes the
105    /// stream or perturbs scheduling. The signal reflects only the calling
106    /// process's submissions to *this* stream — for cross-process device
107    /// load use NVML utilization (`baracuda_nvml::Device::gpu_utilization_percent`).
108    /// A scheduler that wants a plain `bool` can treat "can't tell" as busy
109    /// with `.unwrap_or(false)`.
110    #[inline]
111    pub fn is_idle(&self) -> Result<bool> {
112        self.is_complete()
113    }
114
115    /// The device this stream belongs to.
116    #[inline]
117    pub fn device(&self) -> Device {
118        self.inner.device
119    }
120
121    /// Raw `cudaStream_t` handle. Use with care.
122    #[inline]
123    pub fn as_raw(&self) -> cudaStream_t {
124        self.inner.handle
125    }
126
127    /// Create a stream with a specific scheduling priority (lower = higher
128    /// priority). Use [`stream_priority_range`] to discover the legal
129    /// range on the current device.
130    pub fn with_priority(flags: u32, priority: i32) -> Result<Self> {
131        let r = runtime()?;
132        let cu = r.cuda_stream_create_with_priority()?;
133        let mut stream: cudaStream_t = core::ptr::null_mut();
134        check(unsafe { cu(&mut stream, flags, priority) })?;
135        let device = Device::current()?;
136        Ok(Self {
137            inner: Arc::new(StreamInner {
138                handle: stream,
139                device,
140            }),
141        })
142    }
143
144    /// This stream's scheduling priority.
145    pub fn priority(&self) -> Result<i32> {
146        let r = runtime()?;
147        let cu = r.cuda_stream_get_priority()?;
148        let mut p: core::ffi::c_int = 0;
149        check(unsafe { cu(self.inner.handle, &mut p) })?;
150        Ok(p)
151    }
152
153    /// This stream's flags bitmask.
154    pub fn flags(&self) -> Result<u32> {
155        let r = runtime()?;
156        let cu = r.cuda_stream_get_flags()?;
157        let mut f: core::ffi::c_uint = 0;
158        check(unsafe { cu(self.inner.handle, &mut f) })?;
159        Ok(f)
160    }
161
162    /// Wait for `event` on this stream — blocks future work on `self`
163    /// until the event has completed.
164    pub fn wait_event(&self, event: &crate::Event, flags: u32) -> Result<()> {
165        let r = runtime()?;
166        let cu = r.cuda_stream_wait_event()?;
167        check(unsafe { cu(self.inner.handle, event.as_raw(), flags) })
168    }
169}
170
171/// Return `(least_priority, greatest_priority)` supported on the current
172/// device. Lower numbers are higher priority.
173pub fn stream_priority_range() -> Result<(i32, i32)> {
174    let r = runtime()?;
175    let cu = r.cuda_device_get_stream_priority_range()?;
176    let mut low: core::ffi::c_int = 0;
177    let mut high: core::ffi::c_int = 0;
178    check(unsafe { cu(&mut low, &mut high) })?;
179    Ok((low, high))
180}
181
182impl Stream {
183    /// Enqueue a host-side callback on this stream. Runs on a
184    /// driver-owned thread after prior stream work completes.
185    ///
186    /// The closure is boxed and freed after it runs; a panic inside
187    /// aborts the process.
188    pub fn launch_host_func<F>(&self, f: F) -> Result<()>
189    where
190        F: FnOnce() + Send + 'static,
191    {
192        use core::ffi::c_void;
193
194        let boxed: Box<Box<dyn FnOnce() + Send>> = Box::new(Box::new(f));
195        let raw = Box::into_raw(boxed) as *mut c_void;
196
197        unsafe extern "C" fn trampoline(user_data: *mut c_void) {
198            let f: Box<Box<dyn FnOnce() + Send>> =
199                unsafe { Box::from_raw(user_data as *mut Box<dyn FnOnce() + Send>) };
200            (*f)();
201        }
202
203        let r = runtime()?;
204        let cu = r.cuda_launch_host_func()?;
205        let rc = unsafe { cu(self.inner.handle, Some(trampoline), raw) };
206        if rc != baracuda_cuda_sys::runtime::cudaError_t::Success {
207            // Reclaim the box — cudaLaunchHostFunc didn't take ownership on error.
208            drop(unsafe { Box::from_raw(raw as *mut Box<dyn FnOnce() + Send>) });
209            return Err(crate::error::Error::Status { status: rc });
210        }
211        Ok(())
212    }
213
214    /// Enqueue a 32-bit write of `value` to device memory `addr`.
215    ///
216    /// # Safety
217    ///
218    /// `addr` must be a live device-addressable pointer.
219    pub unsafe fn write_value_32(
220        &self,
221        addr: *mut core::ffi::c_void,
222        value: u32,
223        flags: u32,
224    ) -> Result<()> {
225        unsafe {
226            let r = runtime()?;
227            let cu = r.cuda_stream_write_value_32()?;
228            check(cu(self.inner.handle, addr, value, flags))
229        }
230    }
231
232    /// # Safety
233    ///
234    /// Same as [`write_value_32`].
235    pub unsafe fn write_value_64(
236        &self,
237        addr: *mut core::ffi::c_void,
238        value: u64,
239        flags: u32,
240    ) -> Result<()> {
241        unsafe {
242            let r = runtime()?;
243            let cu = r.cuda_stream_write_value_64()?;
244            check(cu(self.inner.handle, addr, value, flags))
245        }
246    }
247
248    /// Block the stream until the 32-bit device memory at `addr` satisfies
249    /// the condition selected by `flags` (GEQ / EQ / AND / NOR, optionally
250    /// OR'd with FLUSH).
251    ///
252    /// # Safety
253    ///
254    /// `addr` must be a live device-addressable pointer.
255    pub unsafe fn wait_value_32(
256        &self,
257        addr: *mut core::ffi::c_void,
258        value: u32,
259        flags: u32,
260    ) -> Result<()> {
261        unsafe {
262            let r = runtime()?;
263            let cu = r.cuda_stream_wait_value_32()?;
264            check(cu(self.inner.handle, addr, value, flags))
265        }
266    }
267
268    /// # Safety
269    ///
270    /// Same as [`wait_value_32`].
271    pub unsafe fn wait_value_64(
272        &self,
273        addr: *mut core::ffi::c_void,
274        value: u64,
275        flags: u32,
276    ) -> Result<()> {
277        unsafe {
278            let r = runtime()?;
279            let cu = r.cuda_stream_wait_value_64()?;
280            check(cu(self.inner.handle, addr, value, flags))
281        }
282    }
283
284    /// Associate a managed-memory region with this stream
285    /// (`cudaStreamAttachMemAsync`). Pass `flags = 0` for the default.
286    ///
287    /// # Safety
288    ///
289    /// `dev_ptr` must be a managed-memory allocation.
290    pub unsafe fn attach_mem_async(
291        &self,
292        dev_ptr: *mut core::ffi::c_void,
293        length: usize,
294        flags: u32,
295    ) -> Result<()> {
296        unsafe {
297            let r = runtime()?;
298            let cu = r.cuda_stream_attach_mem_async()?;
299            check(cu(self.inner.handle, dev_ptr, length, flags))
300        }
301    }
302
303    /// Copy CUDA-managed attributes (access-policy window, sync policy)
304    /// from `src` onto `self`.
305    pub fn copy_attributes_from(&self, src: &Stream) -> Result<()> {
306        let r = runtime()?;
307        let cu = r.cuda_stream_copy_attributes()?;
308        check(unsafe { cu(self.inner.handle, src.inner.handle) })
309    }
310
311    /// Enqueue a batch of stream mem-ops (`WAIT_VALUE_32/64`,
312    /// `WRITE_VALUE_32/64`) atomically. Much cheaper than issuing the
313    /// ops one at a time.
314    ///
315    /// Build entries with [`baracuda_cuda_sys::types::CUstreamBatchMemOpParams::write_value_32`]
316    /// etc. Pass `flags = 0` for the default.
317    ///
318    /// # Safety
319    ///
320    /// Every entry's `address` must be a live device-addressable pointer.
321    pub unsafe fn batch_mem_op(
322        &self,
323        params: &mut [baracuda_cuda_sys::types::CUstreamBatchMemOpParams],
324        flags: u32,
325    ) -> Result<()> {
326        unsafe {
327            let r = runtime()?;
328            let cu = r.cuda_stream_batch_mem_op()?;
329            check(cu(
330                self.inner.handle,
331                params.len() as core::ffi::c_uint,
332                params.as_mut_ptr(),
333                flags,
334            ))
335        }
336    }
337}
338
339impl Drop for StreamInner {
340    fn drop(&mut self) {
341        if let Ok(r) = runtime() {
342            if let Ok(cu) = r.cuda_stream_destroy() {
343                let _ = unsafe { cu(self.handle) };
344            }
345        }
346    }
347}