use snafu::Snafu;
use std::{alloc::Layout, ptr::NonNull};
pub mod dynamic;
pub mod fixed;
#[derive(Debug, Snafu)]
pub enum Error {
#[snafu(display("dimension count ({ndim}) exceeds i32::MAX"))]
NdimOverflow {
ndim: usize,
},
#[snafu(display("managed tensor allocation layout overflows usize"))]
LayoutOverflow,
}
pub struct Initialized<M: crate::ManagedTensorBase, Storage> {
pub(super) managed: crate::Managed<M>,
pub(super) storage: Storage,
}
impl<M: crate::ManagedTensorBase, Storage> Initialized<M, Storage> {
pub fn tensor_mut(&mut self) -> &mut crate::ffi::DLTensor {
unsafe { (&mut *self.managed.as_ptr()).tensor_mut() }
}
pub fn set_data(&mut self, data: *mut std::ffi::c_void) -> &mut Self {
self.tensor_mut().data = data;
self
}
pub fn set_device(&mut self, device: crate::ffi::DLDevice) -> &mut Self {
self.tensor_mut().device = device;
self
}
pub fn set_dtype(&mut self, dtype: crate::ffi::DLDataType) -> &mut Self {
self.tensor_mut().dtype = dtype;
self
}
pub fn set_byte_offset(&mut self, byte_offset: u64) -> &mut Self {
self.tensor_mut().byte_offset = byte_offset;
self
}
pub fn set_flags(
&mut self,
flags: crate::DlpackFlags,
) -> Result<&mut Self, crate::tensor::Error> {
if flags.newly_asserts_is_copied(self.managed.flags()) {
return Err(crate::tensor::Error::CannotAssertIsCopied);
}
unsafe { (&mut *self.managed.as_ptr()).set_flags_unchecked(flags) };
Ok(self)
}
pub fn set_flags_unchecked(&mut self, flags: crate::DlpackFlags) -> &mut Self {
unsafe { (&mut *self.managed.as_ptr()).set_flags_unchecked(flags) };
self
}
pub unsafe fn finish(self) -> crate::Managed<M> {
self.managed
}
}
impl<Storage> Initialized<crate::ffi::DLManagedTensorVersioned, Storage> {
pub fn version(&self) -> crate::ffi::DLPackVersion {
unsafe { (*self.managed.as_ptr()).version }
}
pub fn set_version(
&mut self,
version: crate::ffi::DLPackVersion,
) -> Result<&mut Self, crate::VersionError> {
crate::version::validate_version(version)?;
unsafe { (*self.managed.as_ptr()).version = version };
Ok(self)
}
}
fn allocate<M>(layout: Layout) -> NonNull<M> {
let pointer = unsafe { std::alloc::alloc(layout) }.cast::<M>();
NonNull::new(pointer).unwrap_or_else(|| std::alloc::handle_alloc_error(layout))
}
fn empty_tensor(ndim: i32) -> crate::ffi::DLTensor {
crate::ffi::DLTensor::from_parts(std::ptr::null_mut(), std::ptr::null_mut(), ndim)
}