use crate::ffi::DLDataType;
use candle_core::DType;
use snafu::Snafu;
mod consumer;
mod producer;
pub use consumer::candle_tensor_from_dlpack;
#[derive(Debug, Snafu)]
pub enum Error {
#[snafu(transparent)]
Metadata {
source: crate::metadata::Error,
},
#[snafu(display("candle tensor must be on CPU"))]
UnsupportedDevice,
#[snafu(display("unsupported candle dtype: {dtype:?}"))]
UnsupportedCandleDType {
dtype: DType,
},
#[snafu(display("no candle dtype matches DLPack dtype: {dtype:?}"))]
UnsupportedDlDataType {
dtype: crate::ffi::DLDataType,
},
#[snafu(display("shape/stride value does not fit the target integer type"))]
DimensionOverflow,
#[snafu(display(
"sub-byte packed dtype {dtype:?} with a nonzero element offset is not supported"
))]
SubByteOffsetUnsupported {
dtype: DLDataType,
},
#[snafu(display(
"sub-byte packed dtype {dtype:?} only supports compact (non-strided) tensors"
))]
SubByteStridesUnsupported {
dtype: DLDataType,
},
#[snafu(display("strided access spans outside the tensor data buffer"))]
StridedSpanOverflow,
#[snafu(transparent)]
Tensor {
source: crate::tensor::Error,
},
#[snafu(transparent)]
Candle {
source: candle_core::Error,
},
}
#[cfg(test)]
use crate::{ManagedTensorBase, TryFromDlpack, allocation::dynamic, ffi::DLDevice};
#[cfg(test)]
use candle_core::{Device, Tensor};
#[cfg(test)]
mod tests {
use super::*;
use crate::ffi::{DLDataTypeCode, DLManagedTensor, DLManagedTensorVersioned};
use crate::{DlpackElement, DlpackFlags, Managed, allocation::fixed::make_test_tensor};
type LegacyDlpack = Managed<DLManagedTensor>;
type VersionedDlpack = Managed<DLManagedTensorVersioned>;
fn managed_candle<M: ManagedTensorBase>(tensor: Tensor) -> Managed<M> {
let initialized: dynamic::Initialized<M> = Box::new(tensor).try_into().unwrap();
unsafe { initialized.finish() }
}
fn managed_candle_with_flags<M: ManagedTensorBase>(
tensor: Tensor,
flags: DlpackFlags,
) -> Managed<M> {
let mut initialized: dynamic::Initialized<M> = Box::new(tensor).try_into().unwrap();
initialized.set_flags(flags).unwrap();
unsafe { initialized.finish() }
}
fn raw_tensor<T, const N: usize>(
data: Vec<T>,
dtype: DLDataType,
shape: [i64; N],
strides: [i64; N],
) -> Managed<DLManagedTensor>
where
T: Send + 'static,
{
let data = Box::new(data);
let data_ptr = data.as_ptr().cast_mut().cast();
make_test_tensor::<_, DLManagedTensor, N>(
data,
data_ptr,
dtype,
DLDevice::CPU,
shape,
strides,
DlpackFlags::empty(),
)
}
#[test]
fn candle_tensor_to_legacy_dlpack_keeps_layout_and_data() {
let tensor = Tensor::from_vec(vec![1i32, 2, 3, 4, 5, 6], (2, 3), &Device::Cpu).unwrap();
let dlpack: LegacyDlpack = managed_candle(tensor);
assert_eq!(dlpack.validate().unwrap().shape(), &[2, 3]);
assert_eq!(dlpack.validate().unwrap().strides().unwrap(), &[3, 1]);
assert_eq!(
unsafe { dlpack.tensor().cpu_slice::<i32>() }.unwrap(),
&[1, 2, 3, 4, 5, 6]
);
}
#[test]
fn candle_builder_defaults_to_empty_flags() {
let tensor = Tensor::from_vec(vec![1f32, 2., 3., 4.], (2, 2), &Device::Cpu).unwrap();
let dlpack: VersionedDlpack = managed_candle(tensor);
assert_eq!(dlpack.flags(), DlpackFlags::empty());
assert_eq!(
unsafe { dlpack.tensor().cpu_slice::<f32>() }.unwrap(),
&[1., 2., 3., 4.]
);
}
#[test]
fn candle_tensor_to_versioned_builder_allows_flags_before_build() {
let tensor = Tensor::from_vec(vec![1f32, 2., 3., 4.], (2, 2), &Device::Cpu).unwrap();
let dlpack: Managed<DLManagedTensorVersioned> =
managed_candle_with_flags(tensor, DlpackFlags::READ_ONLY);
assert_eq!(dlpack.flags(), DlpackFlags::READ_ONLY);
assert_eq!(
unsafe { dlpack.tensor().cpu_slice::<f32>() }.unwrap(),
&[1., 2., 3., 4.]
);
}
#[test]
fn candle_f8e4m3_tensor_roundtrips_dtype_and_shape() {
let tensor = Tensor::zeros((2, 3), DType::F8E4M3, &Device::Cpu).unwrap();
let dlpack: LegacyDlpack = managed_candle(tensor);
assert_eq!(dlpack.validate().unwrap().shape(), &[2, 3]);
let dtype = dlpack.validate().unwrap().dtype();
assert_eq!(dtype.code, DLDataTypeCode::FLOAT8_E4M3);
assert_eq!(dtype.bits, 8);
assert_eq!(dlpack.validate().unwrap().num_bytes(), 6);
}
#[test]
fn candle_f4_tensor_roundtrips_via_from_raw_buffer() {
let tensor =
Tensor::from_raw_buffer(&[0xAB, 0xCD, 0xEF], DType::F4, &[6], &Device::Cpu).unwrap();
let dlpack: LegacyDlpack = managed_candle(tensor);
assert_eq!(dlpack.validate().unwrap().shape(), &[6]);
let dtype = dlpack.validate().unwrap().dtype();
assert_eq!(dtype.code, DLDataTypeCode::FLOAT4_E2M1FN);
assert_eq!(dtype.bits, 4);
assert_eq!(dlpack.validate().unwrap().num_bytes(), 3);
}
#[test]
fn dlpack_f4_converts_to_candle_tensor_with_matching_dtype() {
let dlpack = raw_tensor(
vec![0xABu8, 0xCD, 0xEF],
DLDataType {
code: DLDataTypeCode::FLOAT4_E2M1FN,
bits: 4,
lanes: 1,
},
[6],
[1],
);
let tensor = unsafe { Tensor::try_from_dlpack(&dlpack, ()) }.unwrap();
assert_eq!(tensor.dims(), &[6]);
assert_eq!(tensor.dtype(), DType::F4);
}
#[test]
fn non_compact_sub_byte_packed_dlpack_is_rejected() {
let dlpack = raw_tensor(
vec![0xABu8, 0xCD, 0xEF],
DLDataType {
code: DLDataTypeCode::FLOAT4_E2M1FN,
bits: 4,
lanes: 1,
},
[3, 2],
[1, 3],
);
let err = unsafe { Tensor::try_from_dlpack(&dlpack, ()) }.unwrap_err();
assert!(matches!(err, Error::SubByteStridesUnsupported { .. }));
}
#[test]
fn compact_dlpack_to_candle_tensor_copies_values() {
let dlpack = raw_tensor(
vec![1i32, 2, 3, 4, 5, 6],
<i32 as DlpackElement>::DTYPE,
[2, 3],
[3, 1],
);
let tensor = unsafe { Tensor::try_from_dlpack(&dlpack, ()) }.unwrap();
assert_eq!(tensor.dims(), &[2, 3]);
assert_eq!(
tensor.flatten_all().unwrap().to_vec1::<i32>().unwrap(),
vec![1, 2, 3, 4, 5, 6]
);
}
#[test]
fn non_compact_strided_dlpack_to_candle_tensor_gathers_in_row_major_order() {
let dlpack = raw_tensor(
vec![1i32, 2, 3, 4, 5, 6],
<i32 as DlpackElement>::DTYPE,
[3, 2],
[1, 3],
);
let tensor = unsafe { Tensor::try_from_dlpack(&dlpack, ()) }.unwrap();
assert_eq!(tensor.dims(), &[3, 2]);
assert_eq!(
tensor.flatten_all().unwrap().to_vec1::<i32>().unwrap(),
vec![1, 4, 2, 5, 3, 6]
);
}
#[test]
fn non_compact_dlpack_with_out_of_bounds_stride_is_rejected() {
let dlpack = raw_tensor(
vec![1i32, 2, 3, 4, 5, 6],
<i32 as DlpackElement>::DTYPE,
[2, 2],
[10, 1],
);
let err = unsafe { Tensor::try_from_dlpack(&dlpack, ()) }.unwrap_err();
assert!(matches!(err, Error::StridedSpanOverflow));
}
#[test]
fn non_compact_dlpack_with_negative_stride_is_rejected_as_out_of_bounds() {
let dlpack = raw_tensor(
vec![1i32, 2, 3, 4, 5, 6],
<i32 as DlpackElement>::DTYPE,
[3, 2],
[-1, -3],
);
let err = unsafe { Tensor::try_from_dlpack(&dlpack, ()) }.unwrap_err();
assert!(matches!(err, Error::StridedSpanOverflow));
}
#[test]
fn dlpack_f8e4m3_converts_to_candle_tensor_with_matching_dtype() {
let dlpack = raw_tensor(
vec![0u8; 6],
DLDataType {
code: DLDataTypeCode::FLOAT8_E4M3,
bits: 8,
lanes: 1,
},
[2, 3],
[3, 1],
);
let tensor = unsafe { Tensor::try_from_dlpack(&dlpack, ()) }.unwrap();
assert_eq!(tensor.dims(), &[2, 3]);
assert_eq!(tensor.dtype(), DType::F8E4M3);
}
#[test]
fn dlpack_with_unmatched_dtype_is_rejected() {
let dlpack = raw_tensor(
vec![0u8; 3],
DLDataType {
code: DLDataTypeCode(99),
bits: 1,
lanes: 1,
},
[3],
[1],
);
let err = unsafe { Tensor::try_from_dlpack(&dlpack, ()) }.unwrap_err();
assert!(matches!(err, Error::UnsupportedDlDataType { .. }));
}
}