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::{check, Result};
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<()> { unsafe {
225        let r = runtime()?;
226        let cu = r.cuda_stream_write_value_32()?;
227        check(cu(self.inner.handle, addr, value, flags))
228    }}
229
230    /// # Safety
231    ///
232    /// Same as [`write_value_32`].
233    pub unsafe fn write_value_64(
234        &self,
235        addr: *mut core::ffi::c_void,
236        value: u64,
237        flags: u32,
238    ) -> Result<()> { unsafe {
239        let r = runtime()?;
240        let cu = r.cuda_stream_write_value_64()?;
241        check(cu(self.inner.handle, addr, value, flags))
242    }}
243
244    /// Block the stream until the 32-bit device memory at `addr` satisfies
245    /// the condition selected by `flags` (GEQ / EQ / AND / NOR, optionally
246    /// OR'd with FLUSH).
247    ///
248    /// # Safety
249    ///
250    /// `addr` must be a live device-addressable pointer.
251    pub unsafe fn wait_value_32(
252        &self,
253        addr: *mut core::ffi::c_void,
254        value: u32,
255        flags: u32,
256    ) -> Result<()> { unsafe {
257        let r = runtime()?;
258        let cu = r.cuda_stream_wait_value_32()?;
259        check(cu(self.inner.handle, addr, value, flags))
260    }}
261
262    /// # Safety
263    ///
264    /// Same as [`wait_value_32`].
265    pub unsafe fn wait_value_64(
266        &self,
267        addr: *mut core::ffi::c_void,
268        value: u64,
269        flags: u32,
270    ) -> Result<()> { unsafe {
271        let r = runtime()?;
272        let cu = r.cuda_stream_wait_value_64()?;
273        check(cu(self.inner.handle, addr, value, flags))
274    }}
275
276    /// Associate a managed-memory region with this stream
277    /// (`cudaStreamAttachMemAsync`). Pass `flags = 0` for the default.
278    ///
279    /// # Safety
280    ///
281    /// `dev_ptr` must be a managed-memory allocation.
282    pub unsafe fn attach_mem_async(
283        &self,
284        dev_ptr: *mut core::ffi::c_void,
285        length: usize,
286        flags: u32,
287    ) -> Result<()> { unsafe {
288        let r = runtime()?;
289        let cu = r.cuda_stream_attach_mem_async()?;
290        check(cu(self.inner.handle, dev_ptr, length, flags))
291    }}
292
293    /// Copy CUDA-managed attributes (access-policy window, sync policy)
294    /// from `src` onto `self`.
295    pub fn copy_attributes_from(&self, src: &Stream) -> Result<()> {
296        let r = runtime()?;
297        let cu = r.cuda_stream_copy_attributes()?;
298        check(unsafe { cu(self.inner.handle, src.inner.handle) })
299    }
300
301    /// Enqueue a batch of stream mem-ops (`WAIT_VALUE_32/64`,
302    /// `WRITE_VALUE_32/64`) atomically. Much cheaper than issuing the
303    /// ops one at a time.
304    ///
305    /// Build entries with [`baracuda_cuda_sys::types::CUstreamBatchMemOpParams::write_value_32`]
306    /// etc. Pass `flags = 0` for the default.
307    ///
308    /// # Safety
309    ///
310    /// Every entry's `address` must be a live device-addressable pointer.
311    pub unsafe fn batch_mem_op(
312        &self,
313        params: &mut [baracuda_cuda_sys::types::CUstreamBatchMemOpParams],
314        flags: u32,
315    ) -> Result<()> { unsafe {
316        let r = runtime()?;
317        let cu = r.cuda_stream_batch_mem_op()?;
318        check(cu(
319            self.inner.handle,
320            params.len() as core::ffi::c_uint,
321            params.as_mut_ptr(),
322            flags,
323        ))
324    }}
325}
326
327impl Drop for StreamInner {
328    fn drop(&mut self) {
329        if let Ok(r) = runtime() {
330            if let Ok(cu) = r.cuda_stream_destroy() {
331                let _ = unsafe { cu(self.handle) };
332            }
333        }
334    }
335}