use std::fmt;
use std::marker::PhantomData;
use std::ops::Deref;
use tinyvec::TinyVec;
pub use ffi::{DLDataType, DLDataTypeCode, DLDevice, DLDeviceType, DLManagedTensor, DLTensor};
pub(crate) trait DeviceTypeExt {
fn is_device_compatible(&self) -> bool;
}
impl DeviceTypeExt for ffi::DLDeviceType {
fn is_device_compatible(&self) -> bool {
matches!(*self, Self::kDLCUDA | Self::kDLCUDAManaged)
}
}
const INLINE_DIMS: usize = 3;
pub(crate) type TensorDims = TinyVec<[i64; INLINE_DIMS]>;
pub trait AsDlTensor {
fn as_dl_tensor(&self) -> std::result::Result<DLTensorView<'_>, DLPackError>;
}
pub trait AsDlTensorMut {
fn as_dl_tensor_mut(&mut self) -> std::result::Result<DLTensorViewMut<'_>, DLPackError>;
}
pub trait DType {
fn dl_dtype() -> ffi::DLDataType;
}
macro_rules! impl_dtype {
($ty:ty, $code:expr, $bits:expr) => {
impl DType for $ty {
fn dl_dtype() -> ffi::DLDataType {
ffi::DLDataType { code: $code as u8, bits: $bits, lanes: 1 }
}
}
};
}
impl_dtype!(f32, ffi::DLDataTypeCode::kDLFloat, 32);
impl_dtype!(f64, ffi::DLDataTypeCode::kDLFloat, 64);
impl_dtype!(i32, ffi::DLDataTypeCode::kDLInt, 32);
impl_dtype!(i64, ffi::DLDataTypeCode::kDLInt, 64);
impl_dtype!(u32, ffi::DLDataTypeCode::kDLUInt, 32);
impl_dtype!(u64, ffi::DLDataTypeCode::kDLUInt, 64);
impl_dtype!(u8, ffi::DLDataTypeCode::kDLUInt, 8);
impl_dtype!(i8, ffi::DLDataTypeCode::kDLInt, 8);
impl_dtype!(u16, ffi::DLDataTypeCode::kDLUInt, 16);
impl_dtype!(i16, ffi::DLDataTypeCode::kDLInt, 16);
#[derive(Debug, Clone, thiserror::Error)]
#[non_exhaustive]
pub enum DLPackError {
#[error("unsupported tensor device: {0}")]
UnsupportedDevice(String),
#[error("unsupported tensor dtype: {0}")]
UnsupportedDType(String),
#[error("strides length {strides} does not match tensor rank {ndim}")]
StridesLenMismatch { ndim: usize, strides: usize },
#[error("invalid DLPack metadata: {0}")]
InvalidMetadata(&'static str),
}
pub(crate) struct ManagedTensorRef<'a> {
pub(crate) inner: ffi::DLManagedTensor,
_borrow: PhantomData<&'a ()>,
}
impl ManagedTensorRef<'_> {
pub(crate) fn as_mut_ptr(&mut self) -> *mut ffi::DLManagedTensor {
&mut self.inner
}
}
#[must_use]
pub struct DLTensorView<'a> {
data: *mut std::ffi::c_void,
device: ffi::DLDevice,
dtype: ffi::DLDataType,
shape: TensorDims,
strides: Option<TensorDims>,
_marker: PhantomData<&'a ()>,
}
impl<'a> DLTensorView<'a> {
pub unsafe fn from_raw_parts(
data: *mut std::ffi::c_void,
device: ffi::DLDevice,
shape: &[i64],
strides: Option<&[i64]>,
dtype: ffi::DLDataType,
) -> std::result::Result<Self, DLPackError> {
if let Some(s) = strides
&& s.len() != shape.len()
{
return Err(DLPackError::StridesLenMismatch { ndim: shape.len(), strides: s.len() });
}
Ok(Self {
data,
device,
dtype,
shape: shape.iter().copied().collect(),
strides: strides.map(|s| s.iter().copied().collect()),
_marker: PhantomData,
})
}
pub(crate) fn to_c(&self) -> ManagedTensorRef<'_> {
ManagedTensorRef {
inner: ffi::DLManagedTensor {
dl_tensor: ffi::DLTensor {
data: self.data,
device: self.device,
ndim: self.shape.len() as i32,
dtype: self.dtype,
shape: self.shape.as_ptr() as *mut _,
strides: self
.strides
.as_ref()
.map_or(std::ptr::null_mut(), |s| s.as_ptr() as *mut _),
byte_offset: 0,
},
manager_ctx: std::ptr::null_mut(),
deleter: None,
},
_borrow: PhantomData,
}
}
pub fn ndim(&self) -> usize {
self.shape.len()
}
pub fn shape(&self) -> &[i64] {
&self.shape
}
pub fn strides(&self) -> Option<&[i64]> {
self.strides.as_deref()
}
pub fn dtype(&self) -> ffi::DLDataType {
self.dtype
}
pub fn device(&self) -> ffi::DLDevice {
self.device
}
}
impl fmt::Debug for DLTensorView<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DLTensorView")
.field("shape", &self.shape.as_slice())
.field("strides", &self.strides.as_deref())
.finish()
}
}
#[must_use]
pub struct DLTensorViewMut<'a> {
base: DLTensorView<'a>,
_unique: PhantomData<&'a mut ()>,
}
impl<'a> DLTensorViewMut<'a> {
pub unsafe fn from_raw_parts(
data: *mut std::ffi::c_void,
device: ffi::DLDevice,
shape: &[i64],
strides: Option<&[i64]>,
dtype: ffi::DLDataType,
) -> std::result::Result<Self, DLPackError> {
Ok(Self {
base: unsafe { DLTensorView::from_raw_parts(data, device, shape, strides, dtype)? },
_unique: PhantomData,
})
}
}
impl<'a> Deref for DLTensorViewMut<'a> {
type Target = DLTensorView<'a>;
fn deref(&self) -> &Self::Target {
&self.base
}
}
impl fmt::Debug for DLTensorViewMut<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DLTensorViewMut")
.field("shape", &self.base.shape.as_slice())
.field("strides", &self.base.strides.as_deref())
.finish()
}
}
impl AsDlTensor for DLTensorView<'_> {
fn as_dl_tensor(&self) -> std::result::Result<DLTensorView<'_>, DLPackError> {
Ok(DLTensorView {
data: self.data,
device: self.device,
dtype: self.dtype,
shape: self.shape.clone(),
strides: self.strides.clone(),
_marker: PhantomData,
})
}
}
impl AsDlTensor for DLTensorViewMut<'_> {
fn as_dl_tensor(&self) -> std::result::Result<DLTensorView<'_>, DLPackError> {
self.base.as_dl_tensor()
}
}
impl AsDlTensorMut for DLTensorViewMut<'_> {
fn as_dl_tensor_mut(&mut self) -> std::result::Result<DLTensorViewMut<'_>, DLPackError> {
Ok(DLTensorViewMut { base: self.base.as_dl_tensor()?, _unique: PhantomData })
}
}
#[cfg(test)]
mod tests {
use super::*;
fn cpu() -> DLDevice {
DLDevice { device_type: DLDeviceType::kDLCPU, device_id: 0 }
}
#[test]
fn to_c_translates_contiguous_metadata() {
let data = [0.0f32; 6];
let view = unsafe {
DLTensorView::from_raw_parts(
data.as_ptr() as *mut _,
cpu(),
&[2, 3],
None,
f32::dl_dtype(),
)
}
.unwrap();
let managed = view.to_c();
let t = &managed.inner.dl_tensor;
assert_eq!(t.ndim, 2);
assert_eq!(t.data as *const f32, data.as_ptr());
assert_eq!(t.dtype.code, DLDataTypeCode::kDLFloat as u8);
assert_eq!(t.dtype.bits, 32);
assert_eq!(t.dtype.lanes, 1);
assert_eq!(t.byte_offset, 0);
assert!(managed.inner.manager_ctx.is_null());
assert!(managed.inner.deleter.is_none());
assert_eq!(unsafe { std::slice::from_raw_parts(t.shape, 2) }, &[2, 3]);
assert!(t.strides.is_null());
}
#[test]
fn to_c_preserves_explicit_strides() {
let data = [0.0f32; 6];
let view = unsafe {
DLTensorView::from_raw_parts(
data.as_ptr() as *mut _,
cpu(),
&[2, 3],
Some(&[1, 2]),
f32::dl_dtype(),
)
}
.unwrap();
let managed = view.to_c();
let t = &managed.inner.dl_tensor;
assert_eq!(unsafe { std::slice::from_raw_parts(t.shape, 2) }, &[2, 3]);
assert!(!t.strides.is_null());
assert_eq!(unsafe { std::slice::from_raw_parts(t.strides, 2) }, &[1, 2]);
}
#[test]
fn from_raw_parts_rejects_mismatched_strides_len() {
let data = [0.0f32; 6];
let err = unsafe {
DLTensorView::from_raw_parts(
data.as_ptr() as *mut _,
cpu(),
&[2, 3],
Some(&[1]),
f32::dl_dtype(),
)
}
.unwrap_err();
assert!(matches!(err, DLPackError::StridesLenMismatch { ndim: 2, strides: 1 }));
}
}