use crate::allocator::Pitch;
use crate::id::KernelId;
use crate::memory_management::drop_queue::{Fence, PendingDropQueue};
use crate::memory_management::{ManagedMemoryBinding, MemoryManagement};
use crate::metadata_cache::MetadataInfoCache;
use crate::server::{Handle, IoError, LaunchError};
use crate::storage::ComputeStorage;
use crate::stream::{EventStreamBackend, StreamCapture};
use cubecl_common::bytes::Bytes;
use cubecl_environment::backtrace::BackTrace;
use cubecl_zspace::{Shape, Strides, striding::has_pitched_row_major_strides};
pub trait DeviceStream {
type Fence: Fence + Send + 'static;
type DeviceStorage: ComputeStorage;
type HostStorage: ComputeStorage;
fn device_memory(&mut self) -> &mut MemoryManagement<Self::DeviceStorage>;
fn host_memory(&mut self) -> &mut MemoryManagement<Self::HostStorage>;
fn drop_queue(&mut self) -> &mut PendingDropQueue<Self::Fence>;
fn capturing(&mut self) -> &mut StreamCapture;
fn info_cache(&mut self) -> &mut MetadataInfoCache<Handle>;
type Signal: Copy;
fn signal(&self) -> Self::Signal;
fn fence(signal: Self::Signal) -> Self::Fence;
}
#[non_exhaustive]
pub struct CopyLayout<'a> {
pub shape: &'a Shape,
pub strides: &'a Strides,
pub elem_size: usize,
pub pitch: Option<Pitch>,
}
impl<'a> CopyLayout<'a> {
pub fn of(shape: &'a Shape, strides: &'a Strides, elem_size: usize) -> Result<Self, IoError> {
if !has_pitched_row_major_strides(shape, strides) {
return Err(IoError::UnsupportedStrides {
backtrace: BackTrace::capture(),
});
}
Ok(Self {
shape,
strides,
elem_size,
pitch: Pitch::of(shape, strides, elem_size),
})
}
}
pub trait Driver: Sized {
type Backend: EventStreamBackend<Stream = Self::Stream>;
type Stream: DeviceStream;
type Context;
type LaunchArgs: ?Sized;
unsafe fn pinned_bytes(
binding: ManagedMemoryBinding,
resource: HostResource<Self>,
size: usize,
) -> Bytes;
unsafe fn copy_to_host(
resource: &DeviceResource<Self>,
layout: &CopyLayout<'_>,
bytes: &mut Bytes,
stream: &Self::Stream,
) -> Result<(), IoError>;
unsafe fn copy_to_device(
resource: &DeviceResource<Self>,
layout: &CopyLayout<'_>,
data: &[u8],
stream: &Self::Stream,
) -> Result<(), IoError>;
fn launch(
ctx: &mut Self::Context,
stream: &mut Self::Stream,
kernel: KernelId,
count: (u32, u32, u32),
args: &mut Self::LaunchArgs,
) -> Result<(), LaunchError>;
}
pub type DeviceResource<D> =
<<<D as Driver>::Stream as DeviceStream>::DeviceStorage as ComputeStorage>::Resource;
pub type HostResource<D> =
<<<D as Driver>::Stream as DeviceStream>::HostStorage as ComputeStorage>::Resource;