use crate::{
Borrowed, DlpackElement,
builder::Builder,
dlpack::ManagedBox,
ffi::{DLDevice, DLDeviceType, DLManagedTensor, DLManagedTensorVersioned},
legacy,
managed_tensor::ManagedTensorBase,
metadata, versioned,
};
use cudarc::driver::{CudaContext, CudaSlice, CudaStream, DevicePtr};
use snafu::{Snafu, ensure};
use std::{mem::ManuallyDrop, ops::Deref, os::raw::c_void, sync::Arc};
#[derive(Debug, Snafu)]
pub enum Error {
#[snafu(display("tensor is not on a CUDA device, got {:?}", device_type))]
NotCuda { device_type: DLDeviceType },
#[snafu(display("tensor data pointer is null"))]
NullData,
#[snafu(display("dtype mismatch: expected {expected:?}, got {actual:?}"))]
DtypeMismatch {
expected: crate::ffi::DLDataType,
actual: crate::ffi::DLDataType,
},
#[snafu(display("cudarc driver error: {source}"))]
Driver { source: cudarc::driver::DriverError },
#[snafu(transparent)]
Tensor { source: crate::tensor::Error },
#[snafu(transparent)]
Builder { source: crate::builder::Error },
}
pub fn from_cuda_slice<T: DlpackElement>(
slice: CudaSlice<T>,
shape: &[i64],
strides: &[i64],
) -> Result<legacy::Dlpack, Error> {
let device_id = slice.ordinal() as i32;
let data_ptr = device_ptr_of(&slice);
Ok(
Builder::new(Box::new(slice), metadata::CopiedSlice::new(shape, strides))
.device(DLDevice::cuda(device_id))
.data(data_ptr)
.dtype(T::DTYPE)
.try_build::<DLManagedTensor>()?,
)
}
pub fn from_cuda_slice_versioned<T: DlpackElement>(
slice: CudaSlice<T>,
shape: &[i64],
strides: &[i64],
) -> Result<versioned::Dlpack, Error> {
let device_id = slice.ordinal() as i32;
let data_ptr = device_ptr_of(&slice);
Ok(
Builder::new(Box::new(slice), metadata::CopiedSlice::new(shape, strides))
.device(DLDevice::cuda(device_id))
.data(data_ptr)
.dtype(T::DTYPE)
.try_build::<DLManagedTensorVersioned>()?,
)
}
pub struct BorrowedCudaSlice<M: ManagedTensorBase, T> {
inner: Borrowed<ManagedBox<M>, CudaSliceView<T>>,
}
struct CudaSliceView<T>(ManuallyDrop<CudaSlice<T>>);
impl<T> Drop for CudaSliceView<T> {
fn drop(&mut self) {
let slice = unsafe { ManuallyDrop::take(&mut self.0) };
slice.leak();
}
}
impl<T> Deref for CudaSliceView<T> {
type Target = CudaSlice<T>;
fn deref(&self) -> &CudaSlice<T> {
&self.0
}
}
impl<M: ManagedTensorBase, T> BorrowedCudaSlice<M, T> {
pub fn dlpack(&self) -> &ManagedBox<M> {
self.inner.owner()
}
}
impl<M: ManagedTensorBase, T> Deref for BorrowedCudaSlice<M, T> {
type Target = CudaSlice<T>;
fn deref(&self) -> &CudaSlice<T> {
&self.inner
}
}
impl<T, M> TryFrom<ManagedBox<M>> for BorrowedCudaSlice<M, T>
where
T: DlpackElement,
M: ManagedTensorBase,
{
type Error = Error;
fn try_from(dlpack: ManagedBox<M>) -> Result<Self, Self::Error> {
let tensor = dlpack.tensor();
ensure!(
tensor.device.device_type == DLDeviceType::CUDA,
NotCudaSnafu {
device_type: tensor.device.device_type
}
);
ensure!(
tensor.dtype.is::<T>(),
DtypeMismatchSnafu {
expected: T::DTYPE,
actual: tensor.dtype,
}
);
ensure!(!tensor.data.is_null(), NullDataSnafu);
let len = tensor.num_elements()?;
let cu_device_ptr = tensor.data as usize as u64;
let ctx = CudaContext::new(tensor.device.device_id as usize)
.map_err(|source| Error::Driver { source })?;
let stream: Arc<CudaStream> = ctx.default_stream();
let slice = unsafe { stream.upgrade_device_ptr::<T>(cu_device_ptr, len) };
let view = CudaSliceView(ManuallyDrop::new(slice));
let inner = unsafe { Borrowed::new_unchecked(dlpack, view) };
Ok(BorrowedCudaSlice { inner })
}
}
fn device_ptr_of<T>(slice: &CudaSlice<T>) -> *mut c_void {
let stream = slice.stream().clone();
let (cu_ptr, sync) = slice.device_ptr(&stream);
drop(sync);
cu_ptr as usize as *mut c_void
}