use crate::simt::device_context::with_deallocator_stream;
use crate::simt::error::DeviceError;
use crate::simt::launch::{AsyncKernelLaunchBuilder, KernelArgument};
use cuda_bindings::CUdeviceptr;
use cuda_core::simt::memory::free_async;
use std::io::{self, Write};
use std::marker::PhantomData;
#[derive(Debug, Copy, Clone)]
pub struct DevicePointer<T> {
_marker: PhantomData<T>,
pub dptr: CUdeviceptr,
}
unsafe impl<T> Send for DevicePointer<T> {}
impl<T> DevicePointer<T> {
pub fn cu_deviceptr(&self) -> CUdeviceptr {
self.dptr
}
}
impl<T: Send + Sized> KernelArgument for DevicePointer<T> {
fn push_arg(self, launcher: &mut AsyncKernelLaunchBuilder<'_>) {
launcher.push_arg(self.cu_deviceptr());
}
}
#[derive(Debug)]
pub struct DeviceBox<T: Send + ?Sized> {
device_id: usize,
cudptr: CUdeviceptr,
len: usize,
_marker: PhantomData<T>,
}
unsafe impl<T: Send + ?Sized> Send for DeviceBox<T> {}
unsafe impl<T: Send + ?Sized> Sync for DeviceBox<T> {}
impl<T: Send + ?Sized> Drop for DeviceBox<T> {
fn drop(&mut self) {
let result = unsafe {
with_deallocator_stream(self.device_id, |stream| {
free_async(self.cudptr, stream.cu_stream()).map_err(DeviceError::Driver)
})
};
match result {
Ok(Ok(())) => {}
Ok(Err(err)) | Err(err) => {
let mut stderr = io::stderr().lock();
let _ = writeln!(
stderr,
"cuda-async: failed to enqueue async free for device pointer on device_id={}: {}",
self.device_id, err
);
}
}
}
}
impl<DType: Send + Sized> DeviceBox<[DType]> {
pub unsafe fn from_raw_parts(dptr: CUdeviceptr, len: usize, device_id: usize) -> Self {
Self {
_marker: PhantomData,
cudptr: dptr,
len,
device_id,
}
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
pub fn len(&self) -> usize {
self.len
}
pub fn cu_deviceptr(&self) -> CUdeviceptr {
self.cudptr
}
pub fn device_id(&self) -> usize {
self.device_id
}
pub fn device_pointer(&self) -> DevicePointer<DType> {
DevicePointer {
_marker: PhantomData,
dptr: self.cudptr,
}
}
}