use crate::{
ManagedBox, ManagedTensorBase,
builder::Builder,
ffi::{DLDataType, DLDataTypeCode, DLDevice},
metadata,
};
use candle_core::{
DType, Device, Storage, Tensor, backend::BackendStorage, cpu_backend::CpuStorage,
};
use snafu::Snafu;
use std::mem::size_of;
use std::os::raw::c_void;
#[derive(Debug, Snafu)]
pub enum 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 },
}
fn dl_dtype_from_candle(dtype: DType) -> Option<DLDataType> {
let (code, bits) = match dtype {
DType::U8 => (DLDataTypeCode::UINT, u8::BITS as u8),
DType::U32 => (DLDataTypeCode::UINT, u32::BITS as u8),
DType::I16 => (DLDataTypeCode::INT, i16::BITS as u8),
DType::I32 => (DLDataTypeCode::INT, i32::BITS as u8),
DType::I64 => (DLDataTypeCode::INT, i64::BITS as u8),
#[cfg(feature = "half")]
DType::BF16 => (DLDataTypeCode::BFLOAT, (size_of::<half::bf16>() * 8) as u8),
#[cfg(feature = "half")]
DType::F16 => (DLDataTypeCode::FLOAT, (size_of::<half::f16>() * 8) as u8),
DType::F32 => (DLDataTypeCode::FLOAT, (size_of::<f32>() * 8) as u8),
DType::F64 => (DLDataTypeCode::FLOAT, (size_of::<f64>() * 8) as u8),
DType::F8E4M3 => (DLDataTypeCode::FLOAT8_E4M3, 8),
DType::F8E8M0 => (DLDataTypeCode::FLOAT8_E8M0FNU, 8),
DType::F6E2M3 => (DLDataTypeCode::FLOAT6_E2M3FN, 6),
DType::F6E3M2 => (DLDataTypeCode::FLOAT6_E3M2FN, 6),
DType::F4 => (DLDataTypeCode::FLOAT4_E2M1FN, 4),
_ => return None,
};
Some(DLDataType {
code,
bits,
lanes: 1,
})
}
struct CandleLayout {
data_ptr: *mut c_void,
dtype: DLDataType,
dims: Vec<i64>,
strides: Vec<i64>,
}
fn dlpack_layout_from_candle(tensor: &Tensor) -> Result<CandleLayout, Error> {
let (storage, layout) = tensor.storage_and_layout();
let Storage::Cpu(cpu_storage) = &*storage else {
return Err(Error::UnsupportedDevice);
};
let dtype = dl_dtype_from_candle(cpu_storage.dtype()).ok_or(Error::UnsupportedCandleDType {
dtype: cpu_storage.dtype(),
})?;
let base_ptr = match cpu_storage {
CpuStorage::U8(v) => v.as_ptr() as *mut c_void,
CpuStorage::U32(v) => v.as_ptr() as *mut c_void,
CpuStorage::I16(v) => v.as_ptr() as *mut c_void,
CpuStorage::I32(v) => v.as_ptr() as *mut c_void,
CpuStorage::I64(v) => v.as_ptr() as *mut c_void,
#[cfg(feature = "half")]
CpuStorage::BF16(v) => v.as_ptr() as *mut c_void,
#[cfg(feature = "half")]
CpuStorage::F16(v) => v.as_ptr() as *mut c_void,
CpuStorage::F32(v) => v.as_ptr() as *mut c_void,
CpuStorage::F64(v) => v.as_ptr() as *mut c_void,
CpuStorage::F8E4M3(v) => v.as_ptr() as *mut c_void,
CpuStorage::F8E8M0(v) => v.as_ptr() as *mut c_void,
CpuStorage::F6E2M3(v) => v.as_ptr() as *mut c_void,
CpuStorage::F6E3M2(v) => v.as_ptr() as *mut c_void,
CpuStorage::F4(v) => v.as_ptr() as *mut c_void,
#[cfg(not(feature = "half"))]
other => {
return Err(Error::UnsupportedCandleDType {
dtype: other.dtype(),
});
}
};
if dtype.bits < 8 && layout.start_offset() != 0 {
return Err(Error::SubByteOffsetUnsupported { dtype });
}
let data_ptr =
unsafe { (base_ptr as *mut u8).add(layout.start_offset() * dtype.element_size()) }
as *mut c_void;
let dims = layout
.dims()
.iter()
.map(|&d| i64::try_from(d).map_err(|_| Error::DimensionOverflow))
.collect::<Result<Vec<_>, _>>()?;
let strides = layout
.stride()
.iter()
.map(|&s| i64::try_from(s).map_err(|_| Error::DimensionOverflow))
.collect::<Result<Vec<_>, _>>()?;
Ok(CandleLayout {
data_ptr,
dtype,
dims,
strides,
})
}
impl TryFrom<Box<Tensor>> for Builder<Box<Tensor>, metadata::CopiedSlice<Vec<i64>, Vec<i64>>> {
type Error = Error;
fn try_from(tensor: Box<Tensor>) -> Result<Self, Self::Error> {
if !tensor.device().is_cpu() {
return Err(Error::UnsupportedDevice);
}
let CandleLayout {
data_ptr,
dtype,
dims,
strides,
} = dlpack_layout_from_candle(&tensor)?;
let builder = Builder::new(tensor, metadata::CopiedSlice::new(dims, strides));
Ok(unsafe { builder.data(data_ptr) }
.dtype(dtype)
.device(DLDevice::CPU))
}
}
fn candle_dtype_from_dl(dtype: DLDataType) -> Option<DType> {
match (dtype.code, dtype.bits) {
(DLDataTypeCode::UINT, 8) => Some(DType::U8),
(DLDataTypeCode::UINT, 32) => Some(DType::U32),
(DLDataTypeCode::INT, 16) => Some(DType::I16),
(DLDataTypeCode::INT, 32) => Some(DType::I32),
(DLDataTypeCode::INT, 64) => Some(DType::I64),
#[cfg(feature = "half")]
(DLDataTypeCode::BFLOAT, 16) => Some(DType::BF16),
#[cfg(feature = "half")]
(DLDataTypeCode::FLOAT, 16) => Some(DType::F16),
(DLDataTypeCode::FLOAT, 32) => Some(DType::F32),
(DLDataTypeCode::FLOAT, 64) => Some(DType::F64),
(DLDataTypeCode::FLOAT8_E4M3, 8) => Some(DType::F8E4M3),
(DLDataTypeCode::FLOAT8_E8M0FNU, 8) => Some(DType::F8E8M0),
(DLDataTypeCode::FLOAT6_E2M3FN, 6) => Some(DType::F6E2M3),
(DLDataTypeCode::FLOAT6_E3M2FN, 6) => Some(DType::F6E3M2),
(DLDataTypeCode::FLOAT4_E2M1FN, 4) => Some(DType::F4),
_ => None,
}
}
fn validate_strided_span(
shape: &[i64],
strides: &[i64],
elem_size: usize,
num_bytes: usize,
) -> Result<usize, Error> {
let mut total: usize = 1;
let mut min_elem: i128 = 0;
let mut max_elem: i128 = 0;
for (&d, &s) in shape.iter().zip(strides) {
if d < 0 {
return Err(Error::Tensor {
source: crate::tensor::Error::NegativeDimension { axis: 0, value: d },
});
}
total = total.checked_mul(d as usize).ok_or(Error::Tensor {
source: crate::tensor::Error::NumElementsOverflow,
})?;
if d == 0 {
continue;
}
let last = (d - 1) as i128;
let s = s as i128;
let (lo, hi) = if s >= 0 { (0, last * s) } else { (last * s, 0) };
min_elem += lo;
max_elem += hi;
}
let min_byte = min_elem
.checked_mul(elem_size as i128)
.ok_or(Error::StridedSpanOverflow)?;
let max_byte = max_elem
.checked_add(1)
.and_then(|m| m.checked_mul(elem_size as i128))
.ok_or(Error::StridedSpanOverflow)?;
if min_byte < 0 || max_byte > num_bytes as i128 {
return Err(Error::StridedSpanOverflow);
}
Ok(total)
}
fn gather_strided_bytes(
ptr: *const u8,
elem_size: usize,
shape: &[i64],
strides: &[i64],
total: usize,
) -> Vec<u8> {
let ndim = shape.len();
let mut out = Vec::with_capacity(total * elem_size);
let mut idx = vec![0i64; ndim];
for _ in 0..total {
let elem_offset: i64 = idx.iter().zip(strides).map(|(&i, &s)| i * s).sum();
let byte_offset = elem_offset as isize * elem_size as isize;
unsafe {
let src = ptr.offset(byte_offset);
out.extend_from_slice(std::slice::from_raw_parts(src, elem_size));
}
for axis in (0..ndim).rev() {
idx[axis] += 1;
if idx[axis] < shape[axis] {
break;
}
idx[axis] = 0;
}
}
out
}
pub fn candle_tensor_from_dlpack<M: ManagedTensorBase>(
dlpack: &ManagedBox<M>,
) -> Result<Tensor, Error> {
let tensor = dlpack.tensor();
let dl_dtype = tensor.dtype;
let dtype =
candle_dtype_from_dl(dl_dtype).ok_or(Error::UnsupportedDlDataType { dtype: dl_dtype })?;
let shape = unsafe { tensor.shape()? };
let strides = unsafe { tensor.strides()? };
let ptr = unsafe { tensor.cpu_data_ptr_bytes()? };
let compact = match strides {
None => true,
Some(s) => crate::tensor::is_compact_strides(shape, Some(s))?,
};
let bytes: Vec<u8> = if compact {
unsafe { std::slice::from_raw_parts(ptr, tensor.num_bytes()?) }.to_vec()
} else if dl_dtype.bits < 8 {
return Err(Error::SubByteStridesUnsupported { dtype: dl_dtype });
} else {
let s = strides.unwrap();
let num_bytes = unsafe { tensor.num_bytes()? };
let total = validate_strided_span(shape, s, dl_dtype.element_size(), num_bytes)?;
gather_strided_bytes(ptr, dl_dtype.element_size(), shape, s, total)
};
let dims = shape
.iter()
.map(|&d| usize::try_from(d).map_err(|_| Error::DimensionOverflow))
.collect::<Result<Vec<_>, _>>()?;
Ok(Tensor::from_raw_buffer(&bytes, dtype, &dims, &Device::Cpu)?)
}
impl<'a, M> TryFrom<&'a ManagedBox<M>> for Tensor
where
M: ManagedTensorBase,
{
type Error = Error;
fn try_from(dlpack: &'a ManagedBox<M>) -> Result<Self, Self::Error> {
candle_tensor_from_dlpack(dlpack)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ffi::{DLDataTypeCode, DLManagedTensor, DLManagedTensorVersioned};
use crate::{DlpackElement, DlpackFlags, legacy, versioned};
#[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: legacy::Dlpack = Builder::try_from(Box::new(tensor))
.unwrap()
.try_build()
.unwrap();
assert_eq!(dlpack.shape().unwrap(), &[2, 3]);
assert_eq!(dlpack.strides().unwrap().unwrap(), &[3, 1]);
assert_eq!(dlpack.cpu_data_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: versioned::Dlpack = Builder::try_from(Box::new(tensor))
.unwrap()
.try_build()
.unwrap();
assert_eq!(dlpack.flags(), DlpackFlags::empty());
assert_eq!(dlpack.cpu_data_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: ManagedBox<DLManagedTensorVersioned> = Builder::try_from(Box::new(tensor))
.unwrap()
.insert_flags(DlpackFlags::READ_ONLY)
.unwrap()
.try_build()
.unwrap();
assert_eq!(dlpack.flags(), DlpackFlags::READ_ONLY);
assert_eq!(dlpack.cpu_data_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: legacy::Dlpack = Builder::try_from(Box::new(tensor))
.unwrap()
.try_build()
.unwrap();
assert_eq!(dlpack.shape().unwrap(), &[2, 3]);
let dtype = dlpack.tensor().dtype;
assert_eq!(dtype.code, DLDataTypeCode::FLOAT8_E4M3);
assert_eq!(dtype.bits, 8);
assert_eq!(dlpack.num_bytes().unwrap(), 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: legacy::Dlpack = Builder::try_from(Box::new(tensor))
.unwrap()
.try_build()
.unwrap();
assert_eq!(dlpack.shape().unwrap(), &[6]);
let dtype = dlpack.tensor().dtype;
assert_eq!(dtype.code, DLDataTypeCode::FLOAT4_E2M1FN);
assert_eq!(dtype.bits, 4);
assert_eq!(dlpack.num_bytes().unwrap(), 3);
}
#[test]
fn dlpack_f4_converts_to_candle_tensor_with_matching_dtype() {
let data = Box::new(vec![0xABu8, 0xCD, 0xEF]);
let data_ptr = data.as_ptr() as *mut c_void;
let dlpack = unsafe {
Builder::new(data, metadata::CopiedArray::new([6i64], [1i64])).data(data_ptr)
}
.dtype(DLDataType {
code: DLDataTypeCode::FLOAT4_E2M1FN,
bits: 4,
lanes: 1,
})
.build::<DLManagedTensor>();
let tensor = Tensor::try_from(&dlpack).unwrap();
assert_eq!(tensor.dims(), &[6]);
assert_eq!(tensor.dtype(), DType::F4);
}
#[test]
fn non_compact_sub_byte_packed_dlpack_is_rejected() {
let data = Box::new(vec![0xABu8, 0xCD, 0xEF]);
let data_ptr = data.as_ptr() as *mut c_void;
let dlpack = unsafe {
Builder::new(data, metadata::CopiedArray::new([3, 2], [1, 3])).data(data_ptr)
}
.dtype(DLDataType {
code: DLDataTypeCode::FLOAT4_E2M1FN,
bits: 4,
lanes: 1,
})
.build::<DLManagedTensor>();
let err = Tensor::try_from(&dlpack).unwrap_err();
assert!(matches!(err, Error::SubByteStridesUnsupported { .. }));
}
#[test]
fn compact_dlpack_to_candle_tensor_copies_values() {
let data = Box::new(vec![1i32, 2, 3, 4, 5, 6]);
let data_ptr = data.as_ptr() as *mut c_void;
let dlpack = unsafe {
Builder::new(data, metadata::CopiedArray::new([2, 3], [3, 1])).data(data_ptr)
}
.dtype(<i32 as DlpackElement>::DTYPE)
.build::<DLManagedTensor>();
let tensor = Tensor::try_from(&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 data = Box::new(vec![1i32, 2, 3, 4, 5, 6]);
let data_ptr = data.as_ptr() as *mut c_void;
let dlpack = unsafe {
Builder::new(data, metadata::CopiedArray::new([3, 2], [1, 3])).data(data_ptr)
}
.dtype(<i32 as DlpackElement>::DTYPE)
.build::<DLManagedTensor>();
let tensor = Tensor::try_from(&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 data = Box::new(vec![1i32, 2, 3, 4, 5, 6]);
let data_ptr = data.as_ptr() as *mut c_void;
let dlpack = unsafe {
Builder::new(data, metadata::CopiedArray::new([2, 2], [10, 1])).data(data_ptr)
}
.dtype(<i32 as DlpackElement>::DTYPE)
.build::<DLManagedTensor>();
let err = Tensor::try_from(&dlpack).unwrap_err();
assert!(matches!(err, Error::StridedSpanOverflow));
}
#[test]
fn non_compact_dlpack_with_negative_stride_is_rejected_as_out_of_bounds() {
let data = Box::new(vec![1i32, 2, 3, 4, 5, 6]);
let data_ptr = data.as_ptr() as *mut c_void;
let dlpack = unsafe {
Builder::new(data, metadata::CopiedArray::new([3, 2], [-1, -3])).data(data_ptr)
}
.dtype(<i32 as DlpackElement>::DTYPE)
.build::<DLManagedTensor>();
let err = Tensor::try_from(&dlpack).unwrap_err();
assert!(matches!(err, Error::StridedSpanOverflow));
}
#[test]
fn dlpack_f8e4m3_converts_to_candle_tensor_with_matching_dtype() {
let data = Box::new(vec![0u8; 6]);
let data_ptr = data.as_ptr() as *mut c_void;
let dlpack = unsafe {
Builder::new(data, metadata::CopiedArray::new([2, 3], [3, 1])).data(data_ptr)
}
.dtype(DLDataType {
code: DLDataTypeCode::FLOAT8_E4M3,
bits: 8,
lanes: 1,
})
.build::<DLManagedTensor>();
let tensor = Tensor::try_from(&dlpack).unwrap();
assert_eq!(tensor.dims(), &[2, 3]);
assert_eq!(tensor.dtype(), DType::F8E4M3);
}
#[test]
fn dlpack_with_unmatched_dtype_is_rejected() {
let data = Box::new(vec![0u8; 3]);
let data_ptr = data.as_ptr() as *mut c_void;
let dlpack = unsafe {
Builder::new(data, metadata::CopiedArray::new([3i64], [1i64])).data(data_ptr)
}
.dtype(crate::ffi::DLDataType {
code: DLDataTypeCode(99),
bits: 1,
lanes: 1,
})
.build::<DLManagedTensor>();
let err = Tensor::try_from(&dlpack).unwrap_err();
assert!(matches!(err, Error::UnsupportedDlDataType { .. }));
}
}