use super::{DtypeMismatchSnafu, Error, NotCudaSnafu, NullDataSnafu};
use crate::{
Borrowed, DlpackElement, Managed, ManagedTensorBase, TryFromDlpack, ffi::DLDeviceType,
};
use cudarc::driver::{CudaContext, CudaSlice, CudaStream};
use snafu::ensure;
use std::{mem::ManuallyDrop, ops::Deref, sync::Arc};
pub struct BorrowedCudaSlice<M: ManagedTensorBase, T> {
inner: Borrowed<Managed<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) -> &Managed<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> TryFromDlpack<Managed<M>, ()> for BorrowedCudaSlice<M, T>
where
T: DlpackElement,
M: ManagedTensorBase,
{
type Error = Error;
unsafe fn try_from_dlpack(dlpack: Managed<M>, _stream: ()) -> Result<Self, Self::Error> {
build(dlpack, None)
}
}
impl<T, M> TryFromDlpack<Managed<M>, Arc<CudaStream>> for BorrowedCudaSlice<M, T>
where
T: DlpackElement,
M: ManagedTensorBase,
{
type Error = Error;
unsafe fn try_from_dlpack(
dlpack: Managed<M>,
producer_stream: Arc<CudaStream>,
) -> Result<Self, Self::Error> {
build(dlpack, Some(&producer_stream))
}
}
fn build<T, M>(
dlpack: Managed<M>,
producer_stream: Option<&CudaStream>,
) -> Result<BorrowedCudaSlice<M, T>, Error>
where
T: DlpackElement,
M: ManagedTensorBase,
{
let tensor = dlpack.validate()?;
let (cu_device_ptr, len, device_id) = validated_cuda_parts::<T>(&tensor)?;
let ctx = CudaContext::new(device_id).map_err(|source| Error::Driver { source })?;
let stream: Arc<CudaStream> = ctx.default_stream();
if let Some(producer) = producer_stream {
stream
.join(producer)
.map_err(|source| Error::Driver { source })?;
}
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 })
}
pub(super) fn validated_cuda_parts<T: DlpackElement>(
tensor: &crate::tensor::TensorRef<'_>,
) -> Result<(u64, usize, usize), Error> {
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_ptr().is_null(), NullDataSnafu);
let device_id =
usize::try_from(tensor.device().device_id).map_err(|_| Error::InvalidDeviceId {
device_id: tensor.device().device_id,
})?;
let len = tensor.num_elements();
if !tensor.is_compact()? {
return Err(Error::Tensor {
source: crate::tensor::Error::NonCompactStrides,
});
}
let cu_device_ptr = unsafe { tensor.offset_data_ptr::<T>()? } as usize as u64;
Ok((cu_device_ptr, len, device_id))
}