use crate::ffi::DLDeviceType;
use snafu::Snafu;
mod consumer;
mod producer;
pub use consumer::BorrowedCudaSlice;
pub use producer::from_cuda_slice;
#[derive(Debug, Snafu)]
pub enum Error {
#[snafu(transparent)]
Metadata {
source: crate::metadata::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,
},
}
#[cfg(test)]
use crate::ffi::DLDevice;
#[cfg(test)]
use consumer::validated_cuda_parts;
#[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 tensor = unsafe { crate::tensor::TensorRef::from_raw(&tensor) }.unwrap();
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,
};
let tensor = unsafe { crate::tensor::TensorRef::from_raw(&tensor) }.unwrap();
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,
};
let tensor = unsafe { crate::tensor::TensorRef::from_raw(&tensor) }.unwrap();
assert!(matches!(
validated_cuda_parts::<i32>(&tensor),
Err(Error::InvalidDeviceId { device_id: -1 })
));
}
#[test]
#[ignore = "requires a CUDA device; run with --ignored"]
fn cuda_slice_roundtrips_with_stream_sync() {
use crate::{Managed, TryFromDlpack, ffi::DLManagedTensorVersioned};
use cudarc::driver::{CudaContext, CudaSlice};
use std::sync::Arc;
let ctx = CudaContext::new(0).expect("CUDA context");
let producer_stream = ctx.new_stream().expect("producer stream");
let data = vec![1i32, 2, 3, 4];
let slice: CudaSlice<i32> = producer_stream.clone_htod(&data).expect("htod copy");
let (initialized, stream) =
from_cuda_slice::<i32, DLManagedTensorVersioned>(Box::new(slice), &[4], &[1])
.expect("producer");
assert!(
Arc::ptr_eq(&stream, &producer_stream),
"producer returns the slice's stream"
);
let managed: Managed<DLManagedTensorVersioned> = unsafe { initialized.finish() };
let borrowed: BorrowedCudaSlice<DLManagedTensorVersioned, i32> =
unsafe { TryFromDlpack::try_from_dlpack(managed, producer_stream.clone()) }
.expect("consumer join");
let consumer_stream = borrowed.stream().clone();
let host: Vec<i32> = consumer_stream.clone_dtoh(&*borrowed).expect("dtoh copy");
assert_eq!(host, data);
}
#[test]
#[ignore = "requires a CUDA device; run with --ignored"]
fn cuda_slice_roundtrips_without_sync() {
use crate::{Managed, TryFromDlpack, ffi::DLManagedTensorVersioned};
use cudarc::driver::{CudaContext, CudaSlice};
let ctx = CudaContext::new(0).expect("CUDA context");
let producer_stream = ctx.new_stream().expect("producer stream");
let data = vec![5i32, 6, 7];
let slice: CudaSlice<i32> = producer_stream.clone_htod(&data).expect("htod copy");
let (initialized, _stream) =
from_cuda_slice::<i32, DLManagedTensorVersioned>(Box::new(slice), &[3], &[1])
.expect("producer");
let managed: Managed<DLManagedTensorVersioned> = unsafe { initialized.finish() };
let borrowed: BorrowedCudaSlice<DLManagedTensorVersioned, i32> =
unsafe { TryFromDlpack::try_from_dlpack(managed, ()) }.expect("consumer no-sync");
producer_stream.synchronize().expect("producer sync");
let consumer_stream = borrowed.stream().clone();
let host: Vec<i32> = consumer_stream.clone_dtoh(&*borrowed).expect("dtoh copy");
assert_eq!(host, data);
}
}