baracuda_runtime/
stream.rs1use 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#[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 pub fn new() -> Result<Self> {
43 Self::with_flags(cudaStreamFlags::DEFAULT)
44 }
45
46 pub fn non_blocking() -> Result<Self> {
49 Self::with_flags(cudaStreamFlags::NON_BLOCKING)
50 }
51
52 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 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 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 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 #[inline]
111 pub fn is_idle(&self) -> Result<bool> {
112 self.is_complete()
113 }
114
115 #[inline]
117 pub fn device(&self) -> Device {
118 self.inner.device
119 }
120
121 #[inline]
123 pub fn as_raw(&self) -> cudaStream_t {
124 self.inner.handle
125 }
126
127 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 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 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 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
171pub 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 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 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 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 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 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 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 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 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 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}