use super::ComputeClient;
use crate::runtime::Runtime;
use crate::server::CopyDescriptor;
use alloc::boxed::Box;
use alloc::format;
use alloc::sync::Arc;
use core::mem::MaybeUninit;
use cubecl_common::bytes::{
AccessError, AccessPolicy, AllocationController, AllocationProperty, Bytes,
};
use cubecl_zspace::striding::has_contiguous_row_major_strides;
use spin::Once;
pub struct LazyDeviceController<R: Runtime> {
client: ComputeClient<R>,
descriptor: Arc<CopyDescriptor>,
materialized: Once<Bytes>,
}
impl<R: Runtime> LazyDeviceController<R> {
pub(super) fn new(client: ComputeClient<R>, descriptor: Arc<CopyDescriptor>) -> Self {
Self {
client,
descriptor,
materialized: Once::new(),
}
}
fn ensure_init(&self, policy: AccessPolicy) -> Result<&Bytes, AccessError> {
if let Some(bytes) = self.materialized.get() {
return Ok(bytes);
}
if !policy.copy_allowed() {
return Err(AccessError::WouldCopy);
}
self.materialized
.try_call_once(|| -> Result<Bytes, AccessError> {
let desc = self.descriptor.as_ref();
let descriptor = CopyDescriptor::new(
desc.handle.clone(),
desc.shape.clone(),
desc.strides.clone(),
desc.elem_size,
);
cubecl_common::reader::read_sync(self.client.read_one_tensor_async(descriptor))
.map_err(|err| AccessError::Read(format!("{err:?}")))
})
}
fn byte_len(&self) -> usize {
let desc = self.descriptor.as_ref();
desc.shape.iter().product::<usize>() * desc.elem_size
}
}
impl<R: Runtime> AllocationController for LazyDeviceController<R> {
fn alloc_align(&self) -> usize {
match self.materialized.get() {
Some(bytes) => bytes.align(),
None => core::mem::align_of::<u128>(),
}
}
fn property(&self) -> AllocationProperty {
match self.materialized.get() {
Some(bytes) => bytes.property(),
None => AllocationProperty::Device,
}
}
fn capacity(&self) -> usize {
self.byte_len()
}
fn view(&self, start: usize, end: usize) -> Option<Box<dyn AllocationController>> {
if self.materialized.get().is_some() {
return None;
}
if start > end || end > self.byte_len() {
return None;
}
let desc = self.descriptor.as_ref();
if !has_contiguous_row_major_strides(&desc.shape, &desc.strides) {
return None;
}
let base = desc.handle.offset_start.unwrap_or(0);
let size = desc.handle.size;
let mut binding = desc.handle.clone();
binding.offset_start = Some(base + start as u64);
binding.offset_end = Some(size - (base + end as u64));
let descriptor = CopyDescriptor::new(binding, [end - start].into(), [1].into(), 1);
Some(Box::new(LazyDeviceController::new(
self.client.clone(),
Arc::new(descriptor),
)))
}
fn memory(&self, policy: AccessPolicy) -> Result<&[MaybeUninit<u8>], AccessError> {
let bytes = self.ensure_init(policy)?;
let slice: &[u8] = bytes;
Ok(unsafe { core::slice::from_raw_parts(slice.as_ptr().cast(), slice.len()) })
}
unsafe fn memory_mut(
&mut self,
policy: AccessPolicy,
) -> Result<&mut [MaybeUninit<u8>], AccessError> {
self.ensure_init(policy)?;
let bytes = self
.materialized
.get_mut()
.expect("materialized must be set after init");
let slice: &mut [u8] = bytes;
Ok(unsafe { core::slice::from_raw_parts_mut(slice.as_mut_ptr().cast(), slice.len()) })
}
unsafe fn copy_into(&self, buf: &mut [u8]) {
let bytes = self
.ensure_init(AccessPolicy::default())
.expect("device: host access failed");
let len = buf.len().min(bytes.len());
buf[..len].copy_from_slice(&bytes[..len]);
}
}