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::{Result, check};
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<()> {
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 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 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 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 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 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 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}