use crate::error::{DriverError, IntoResult};
use crate::simt::context::CudaContext;
use crate::simt::event::CudaEvent;
use std::ffi::c_void;
use std::mem::MaybeUninit;
use std::sync::atomic::Ordering;
use std::sync::Arc;
#[derive(Debug, PartialEq, Eq)]
pub struct CudaStream {
pub(crate) cu_stream: cuda_bindings::CUstream,
pub(crate) ctx: Arc<CudaContext>,
}
unsafe impl Send for CudaStream {}
unsafe impl Sync for CudaStream {}
impl Drop for CudaStream {
fn drop(&mut self) {
self.ctx.record_err(self.ctx.bind_to_thread());
if !self.cu_stream.is_null() {
self.ctx.num_streams.fetch_sub(1, Ordering::Relaxed);
self.ctx
.record_err(unsafe { cuda_bindings::cuStreamDestroy_v2(self.cu_stream).result() });
}
}
}
impl CudaStream {
pub fn cu_stream(&self) -> cuda_bindings::CUstream {
self.cu_stream
}
pub fn context(&self) -> &Arc<CudaContext> {
&self.ctx
}
pub fn priority(&self) -> Result<i32, DriverError> {
self.ctx.bind_to_thread()?;
let mut priority = MaybeUninit::uninit();
unsafe {
cuda_bindings::cuStreamGetPriority(self.cu_stream, priority.as_mut_ptr()).result()?;
Ok(priority.assume_init())
}
}
pub fn synchronize(&self) -> Result<(), DriverError> {
self.ctx.bind_to_thread()?;
unsafe { cuda_bindings::cuStreamSynchronize(self.cu_stream) }.result()
}
pub fn query(&self) -> Result<bool, DriverError> {
self.ctx.bind_to_thread()?;
match unsafe { cuda_bindings::cuStreamQuery(self.cu_stream) } {
cuda_bindings::cudaError_enum_CUDA_SUCCESS => Ok(true),
cuda_bindings::cudaError_enum_CUDA_ERROR_NOT_READY => Ok(false),
err => Err(DriverError(err)),
}
}
pub fn fork(&self) -> Result<Arc<Self>, DriverError> {
self.ctx.bind_to_thread()?;
self.ctx.num_streams.fetch_add(1, Ordering::Relaxed);
let mut cu_stream = MaybeUninit::uninit();
let cu_stream = unsafe {
cuda_bindings::cuStreamCreate(
cu_stream.as_mut_ptr(),
cuda_bindings::CUstream_flags_enum_CU_STREAM_NON_BLOCKING,
)
.result()?;
cu_stream.assume_init()
};
let stream = Arc::new(CudaStream {
cu_stream,
ctx: self.ctx.clone(),
});
stream.join(self)?;
Ok(stream)
}
pub fn join(&self, other: &CudaStream) -> Result<(), DriverError> {
self.wait(&other.record_event(None)?)
}
pub fn record_event(
&self,
flags: Option<cuda_bindings::CUevent_flags>,
) -> Result<CudaEvent, DriverError> {
let event = self.ctx.new_event(flags)?;
event.record(self)?;
Ok(event)
}
pub fn wait(&self, event: &CudaEvent) -> Result<(), DriverError> {
self.ctx.bind_to_thread()?;
unsafe {
cuda_bindings::cuStreamWaitEvent(
self.cu_stream,
event.cu_event(),
cuda_bindings::CUevent_wait_flags_enum_CU_EVENT_WAIT_DEFAULT,
)
.result()
}
}
pub fn launch_host_function<F: FnOnce() + Send>(
&self,
host_func: F,
) -> Result<(), DriverError> {
let boxed = Box::new(host_func);
unsafe {
cuda_bindings::cuLaunchHostFunc(
self.cu_stream,
Some(Self::callback_wrapper::<F>),
Box::into_raw(boxed) as *mut c_void,
)
.result()
}
}
unsafe extern "C" fn callback_wrapper<F: FnOnce() + Send>(callback: *mut c_void) {
let _ = std::panic::catch_unwind(|| {
let callback: Box<F> = unsafe { Box::from_raw(callback as *mut F) };
callback();
});
}
}