use crate::compute::storage::cpu::{PINNED_MEMORY_ALIGNMENT, PinnedMemoryResource};
use cubecl_common::bytes::{AccessError, AccessPolicy, AllocationController, AllocationProperty};
use cubecl_runtime::memory_management::ManagedMemoryBinding;
pub struct PinnedMemoryManagedAllocController {
resource: PinnedMemoryResource,
_binding: ManagedMemoryBinding,
}
impl PinnedMemoryManagedAllocController {
pub fn init(binding: ManagedMemoryBinding, resource: PinnedMemoryResource) -> Self {
Self {
_binding: binding,
resource,
}
}
}
impl AllocationController for PinnedMemoryManagedAllocController {
fn alloc_align(&self) -> usize {
PINNED_MEMORY_ALIGNMENT
}
unsafe fn memory_mut(
&mut self,
_policy: AccessPolicy,
) -> Result<&mut [std::mem::MaybeUninit<u8>], AccessError> {
if self.resource.size == 0 {
return Ok(empty_pinned_slice_mut());
}
Ok(unsafe {
std::slice::from_raw_parts_mut(
self.resource.ptr as *mut std::mem::MaybeUninit<u8>,
self.resource.size,
)
})
}
fn memory(&self, _policy: AccessPolicy) -> Result<&[std::mem::MaybeUninit<u8>], AccessError> {
if self.resource.size == 0 {
return Ok(empty_pinned_slice_mut());
}
Ok(unsafe {
std::slice::from_raw_parts(
self.resource.ptr as *mut std::mem::MaybeUninit<u8>,
self.resource.size,
)
})
}
fn property(&self) -> AllocationProperty {
AllocationProperty::Pinned
}
}
fn empty_pinned_slice_mut<'a>() -> &'a mut [std::mem::MaybeUninit<u8>] {
unsafe {
std::slice::from_raw_parts_mut(std::ptr::without_provenance_mut(PINNED_MEMORY_ALIGNMENT), 0)
}
}