use cubecl_common::{
backtrace::BackTrace,
bytes::{AllocationController, AllocationProperty},
};
use cubecl_core::server::{Binding, IoError};
use cubecl_runtime::{
memory_management::MemoryManagement,
storage::{BytesResource, BytesStorage},
};
pub struct CpuAllocController {
resource: BytesResource,
_binding: Binding,
}
impl AllocationController for CpuAllocController {
fn alloc_align(&self) -> usize {
align_of::<u8>()
}
fn property(&self) -> AllocationProperty {
AllocationProperty::Other
}
unsafe fn memory_mut(&mut self) -> &mut [std::mem::MaybeUninit<u8>] {
let slice = self.resource.write();
unsafe {
std::slice::from_raw_parts_mut(
slice.as_mut_ptr() as *mut std::mem::MaybeUninit<u8>,
slice.len(),
)
}
}
fn memory(&self) -> &[std::mem::MaybeUninit<u8>] {
let slice = self.resource.read();
unsafe {
std::slice::from_raw_parts(
slice.as_ptr() as *const std::mem::MaybeUninit<u8>,
slice.len(),
)
}
}
}
impl CpuAllocController {
pub fn init(
binding: Binding,
memory_management: &mut MemoryManagement<BytesStorage>,
) -> Result<Self, IoError> {
let resource = memory_management
.get_resource(
binding.memory.clone(),
binding.offset_start,
binding.offset_end,
)
.ok_or(IoError::InvalidHandle {
backtrace: BackTrace::capture(),
})?;
Ok(Self {
_binding: binding,
resource,
})
}
}