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