use super::Error;
use crate::{
ManagedTensorBase,
allocation::dynamic,
ffi::{DLDataType, DLDataTypeCode, DLDevice},
metadata::{Copied, Dynamic},
};
use candle_core::{DType, Storage, Tensor, backend::BackendStorage, cpu_backend::CpuStorage};
use std::{mem::size_of, os::raw::c_void};
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<M: ManagedTensorBase> TryFrom<Box<Tensor>> for dynamic::Initialized<M> {
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 prepared = Dynamic::new(Copied(dims), Copied(strides)).prepare::<M>()?;
let mut initialized = prepared
.initialize(tensor)
.map_err(crate::metadata::Error::from)?;
initialized.set_data(data_ptr);
initialized.set_dtype(dtype);
initialized.set_device(DLDevice::CPU);
Ok(initialized)
}
}