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("CUDA slice length {len} does not fit in i64"))]
LengthOverflow {
len: usize,
source: std::num::TryFromIntError,
},
#[snafu(display("CUDA device ordinal {ordinal} does not fit in i32"))]
DeviceIdOverflow {
ordinal: usize,
source: std::num::TryFromIntError,
},
#[snafu(display("CUDA device ID must be non-negative, got {device_id}"))]
InvalidDeviceId { device_id: i32 },
#[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 type CudaBuilder<T> = Builder<Box<CudaSlice<T>>, metadata::CopiedArray<[i64; 1], [i64; 1], 1>>;
impl<T: DlpackElement> TryFrom<CudaSlice<T>> for CudaBuilder<T> {
type Error = Error;
fn try_from(slice: CudaSlice<T>) -> Result<Self, Self::Error> {
let len = i64::try_from(slice.len()).map_err(|source| Error::LengthOverflow {
len: slice.len(),
source,
})?;
let device_id =
i32::try_from(slice.ordinal()).map_err(|source| Error::DeviceIdOverflow {
ordinal: slice.ordinal(),
source,
})?;
let data_ptr = device_ptr_of(&slice);
Ok(
Builder::new(Box::new(slice), metadata::CopiedArray::new([len], [1]))
.device(DLDevice::cuda(device_id))
.data(data_ptr)
.dtype(T::DTYPE),
)
}
}
pub fn from_cuda_slice<T: DlpackElement>(
slice: CudaSlice<T>,
shape: &[i64],
strides: &[i64],
) -> Result<legacy::Dlpack, Error> {
Ok(CudaBuilder::try_from(slice)?
.metadata(metadata::CopiedSlice::new(shape, strides))
.try_build::<DLManagedTensor>()?)
}
pub fn from_cuda_slice_versioned<T: DlpackElement>(
slice: CudaSlice<T>,
shape: &[i64],
strides: &[i64],
) -> Result<versioned::Dlpack, Error> {
Ok(CudaBuilder::try_from(slice)?
.metadata(metadata::CopiedSlice::new(shape, strides))
.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();
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();
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 validated_cuda_parts<T: DlpackElement>(
tensor: &crate::ffi::DLTensor,
) -> 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.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 = tensor.offset_data_ptr::<T>()? as usize as u64;
Ok((cu_device_ptr, len, device_id))
}
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
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ffi::{DLDataType, DLTensor};
#[test]
fn validated_cuda_parts_applies_byte_offset() {
let data = [0i32; 3];
let shape = [2i64];
let strides = [1i64];
let tensor = DLTensor {
data: data.as_ptr().cast_mut().cast(),
device: DLDevice::cuda(0),
ndim: 1,
dtype: DLDataType::of::<i32>(),
shape: shape.as_ptr().cast_mut(),
strides: strides.as_ptr().cast_mut(),
byte_offset: std::mem::size_of::<i32>() as u64,
};
let (ptr, len, device_id) = validated_cuda_parts::<i32>(&tensor).unwrap();
assert_eq!(ptr, unsafe { data.as_ptr().add(1) } as usize as u64);
assert_eq!(len, 2);
assert_eq!(device_id, 0);
}
#[test]
fn validated_cuda_parts_rejects_non_compact_strides() {
let data = [0i32; 5];
let shape = [2i64, 2];
let strides = [3i64, 1];
let tensor = DLTensor {
data: data.as_ptr().cast_mut().cast(),
device: DLDevice::cuda(0),
ndim: 2,
dtype: DLDataType::of::<i32>(),
shape: shape.as_ptr().cast_mut(),
strides: strides.as_ptr().cast_mut(),
byte_offset: 0,
};
assert!(matches!(
validated_cuda_parts::<i32>(&tensor),
Err(Error::Tensor {
source: crate::tensor::Error::NonCompactStrides
})
));
}
#[test]
fn validated_cuda_parts_rejects_negative_device_id() {
let data = [0i32; 1];
let shape = [1i64];
let strides = [1i64];
let tensor = DLTensor {
data: data.as_ptr().cast_mut().cast(),
device: DLDevice::cuda(-1),
ndim: 1,
dtype: DLDataType::of::<i32>(),
shape: shape.as_ptr().cast_mut(),
strides: strides.as_ptr().cast_mut(),
byte_offset: 0,
};
assert!(matches!(
validated_cuda_parts::<i32>(&tensor),
Err(Error::InvalidDeviceId { device_id: -1 })
));
}
}