use super::allocator::{TensorAllocator, TensorAllocatorError};
use arrow_buffer::{ArrowNativeType, Buffer};
use std::marker::PhantomData;
use std::sync::Arc;
use std::{alloc::Layout, ptr::NonNull};
pub struct TensorStorage<T: ArrowNativeType, A: TensorAllocator> {
pub data: Buffer,
alloc: A,
marker: PhantomData<T>,
}
impl<T, A: TensorAllocator> TensorStorage<T, A>
where
T: ArrowNativeType + std::panic::RefUnwindSafe,
{
pub fn new(len: usize, alloc: A) -> Result<Self, TensorAllocatorError> {
let ptr =
alloc.alloc(Layout::array::<T>(len).map_err(TensorAllocatorError::LayoutError)?)?;
let buffer = unsafe {
Buffer::from_custom_allocation(
NonNull::new_unchecked(ptr),
len * std::mem::size_of::<T>(),
Arc::new(Vec::<T>::with_capacity(len)),
)
};
Ok(Self {
data: buffer,
alloc,
marker: PhantomData,
})
}
pub fn from_vec(vec: Vec<T>, alloc: A) -> Result<Self, TensorAllocatorError> {
let buffer = unsafe {
Buffer::from_custom_allocation(
NonNull::new_unchecked(vec.as_ptr() as *mut u8),
vec.len() * std::mem::size_of::<T>(),
Arc::new(vec),
)
};
let storage = Self {
data: buffer,
alloc,
marker: PhantomData,
};
Ok(storage)
}
pub fn alloc(&self) -> &A {
&self.alloc
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tensor::allocator::CpuAllocator;
#[test]
fn test_tensor_storage() -> Result<(), TensorAllocatorError> {
let allocator = CpuAllocator;
let storage = TensorStorage::<u8, _>::new(1024, allocator)?;
assert_eq!(storage.data.len(), 1024);
Ok(())
}
#[test]
fn test_tensor_storage_ptr() -> Result<(), TensorAllocatorError> {
let allocator = CpuAllocator;
let storage = TensorStorage::<u8, _>::new(1024, allocator)?;
let ptr = storage.data.as_ptr();
assert!(!ptr.is_null());
Ok(())
}
#[test]
fn test_tensor_storage_from_vec() -> Result<(), TensorAllocatorError> {
type CpuStorage = TensorStorage<u8, CpuAllocator>;
let allocator = CpuAllocator;
let vec = vec![0, 1, 2, 3, 4, 5];
let storage = CpuStorage::from_vec(vec, allocator)?;
assert_eq!(storage.data.len(), 6);
Ok(())
}
}