use crate::device_context::with_deallocator_stream;
use cuda_core::free_async;
use cuda_core::sys::CUdeviceptr;
use std::marker::PhantomData;
use std::sync::Arc;
pub unsafe trait DeviceAllocation: Send + Sync + 'static {
fn device_ptr(&self) -> CUdeviceptr;
fn len_bytes(&self) -> usize;
fn device_id(&self) -> usize;
}
enum Owner {
Owned,
Borrowed,
Foreign(#[allow(dead_code)] Arc<dyn DeviceAllocation>),
}
#[derive(Debug, Copy, Clone)]
pub struct DevicePointer<T> {
dtype: PhantomData<T>,
pub dptr: CUdeviceptr,
}
unsafe impl<T> Send for DevicePointer<T> {}
impl<T> DevicePointer<T> {
pub fn cu_deviceptr(&self) -> CUdeviceptr {
self.dptr
}
pub unsafe fn from_cu_deviceptr(dptr: CUdeviceptr) -> Self {
Self {
dtype: PhantomData,
dptr,
}
}
}
pub struct DeviceBuffer {
device_id: usize,
cudptr: CUdeviceptr,
len: usize,
owner: Owner,
}
impl std::fmt::Debug for DeviceBuffer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let kind = match self.owner {
Owner::Owned => "owned",
Owner::Borrowed => "borrowed",
Owner::Foreign(_) => "foreign",
};
f.debug_struct("DeviceBuffer")
.field("device_id", &self.device_id)
.field("cudptr", &self.cudptr)
.field("len", &self.len)
.field("owner", &kind)
.finish()
}
}
unsafe impl Send for DeviceBuffer {}
unsafe impl Sync for DeviceBuffer {}
impl Drop for DeviceBuffer {
fn drop(&mut self) {
if !matches!(self.owner, Owner::Owned) {
return;
}
unsafe {
with_deallocator_stream(self.device_id, |stream| {
free_async(self.cudptr, stream);
})
.unwrap_or_else(|_| {
panic!(
"Failed to free device pointer on device_id={}",
self.device_id
)
})
}
}
}
impl DeviceBuffer {
pub unsafe fn from_raw_parts(dptr: CUdeviceptr, len_bytes: usize, device_id: usize) -> Self {
Self {
cudptr: dptr,
len: len_bytes,
device_id,
owner: Owner::Owned,
}
}
pub fn foreign(owner: Arc<dyn DeviceAllocation>, len_bytes: usize) -> Self {
Self {
cudptr: owner.device_ptr(),
len: len_bytes,
device_id: owner.device_id(),
owner: Owner::Foreign(owner),
}
}
pub unsafe fn borrowed_from_raw_parts(
dptr: CUdeviceptr,
len_bytes: usize,
device_id: usize,
) -> Self {
Self {
cudptr: dptr,
len: len_bytes,
device_id,
owner: Owner::Borrowed,
}
}
pub fn is_empty(&self) -> bool {
self.len_bytes() == 0
}
pub fn len_bytes(&self) -> usize {
self.len
}
pub fn cu_deviceptr(&self) -> CUdeviceptr {
self.cudptr
}
pub fn device_id(&self) -> usize {
self.device_id
}
}