use super::{CopyLayout, DeviceResource, DeviceStream, Driver, Staging};
use crate::id::KernelId;
use crate::memory_management::drop_queue::Fence;
use crate::memory_management::{
InstallMemoryPoolsError, ManagedMemoryHandle, MemoryAllocationMode, MemoryConfiguration,
MemoryHandle, MemoryReport, MemoryUsage,
};
use crate::server::{BufferBinding, CopyDescriptor, Handle, IoError, LaunchError, ServerError};
use crate::stream::ResolvedStreams;
use alloc::boxed::Box;
use alloc::vec;
use alloc::vec::Vec;
use cubecl_common::{bytes::Bytes, device::ServiceId};
use cubecl_environment::backtrace::BackTrace;
use cubecl_environment::future::DynFut;
use cubecl_environment::stream::StreamId;
use cubecl_ir::MemoryDeviceProperties;
pub struct Command<'a, D: Driver> {
ctx: &'a mut D::Context,
streams: ResolvedStreams<'a, D::Backend>,
service: ServiceId,
}
impl<'a, D: Driver> Command<'a, D> {
pub fn new(
ctx: &'a mut D::Context,
streams: ResolvedStreams<'a, D::Backend>,
service: ServiceId,
) -> Self {
Self {
ctx,
streams,
service,
}
}
pub fn stream(&mut self) -> &mut D::Stream {
self.streams.current()
}
pub fn resource(&mut self, binding: BufferBinding) -> Result<DeviceResource<D>, IoError> {
self.streams
.get(&binding.stream)
.device_memory()
.get_resource(binding.memory, binding.offset_start, binding.offset_end)
}
pub fn memory_usage(&mut self) -> MemoryUsage {
self.streams.current().device_memory().memory_usage()
}
pub fn memory_report(&mut self) -> MemoryReport {
self.streams.current().device_memory().memory_report()
}
pub fn memory_cleanup(&mut self) {
let stream = self.streams.current();
if !stream.capturing().is_recording() {
let signal = stream.signal();
stream.drop_queue().drain(|| D::Stream::fence(signal));
stream.info_cache().clear_unpinned();
}
let (stream, failures) = self.streams.current_and_failures();
stream.device_memory().cleanup(true, failures);
stream.host_memory().cleanup(true, failures);
}
pub fn flush_drops(&mut self) {
let stream = self.streams.current();
if stream.capturing().is_recording() {
return;
}
let signal = stream.signal();
stream.drop_queue().flush(|| D::Stream::fence(signal));
}
pub fn allocation_mode(&mut self, mode: MemoryAllocationMode) {
self.streams.current().device_memory().mode(mode)
}
pub fn install_memory_pools(
&mut self,
config: MemoryConfiguration,
props: &MemoryDeviceProperties,
) -> Result<(), InstallMemoryPoolsError> {
let (stream, failures) = self.streams.current_and_failures();
stream
.device_memory()
.install_pools(config, props, failures)
}
pub fn reserve(&mut self, size: u64) -> Result<ManagedMemoryHandle, IoError> {
let (stream, failures) = self.streams.current_and_failures();
match stream.device_memory().reserve(size, failures) {
Ok(handle) => Ok(handle),
Err(err) if !err.may_succeed_after_reclaim() => Err(err),
Err(err) => {
log::warn!("device allocation of {size} B failed ({err}); reclaiming and retrying");
self.memory_cleanup();
let (stream, failures) = self.streams.current_and_failures();
stream.device_memory().reserve(size, failures)
}
}
}
pub fn cursor(&self) -> u64 {
self.streams.cursor
}
pub fn empty(&mut self, size: u64) -> Result<Handle, IoError> {
let handle = Handle::new(self.service, self.streams.current, size);
let reserved = self.reserve(size)?;
self.bind(reserved, handle.memory.clone())?;
Ok(handle)
}
pub fn bind(
&mut self,
reserved: ManagedMemoryHandle,
new: ManagedMemoryHandle,
) -> Result<(), IoError> {
let cursor = self.cursor();
let (stream, failures) = self.streams.current_and_failures();
stream.device_memory().bind(reserved, new, cursor, failures)
}
pub fn reserve_cpu(&mut self, size: usize, origin: Option<StreamId>) -> Bytes {
self.reserve_pinned(size, origin)
.unwrap_or_else(|| Bytes::from_bytes_vec(vec![0; size]))
}
fn reserve_pinned(&mut self, size: usize, origin: Option<StreamId>) -> Option<Bytes> {
let (stream, failures) = match origin {
Some(id) => self.streams.get_and_failures(&id),
None => self.streams.current_and_failures(),
};
let handle = stream.host_memory().reserve(size as u64, failures).ok()?;
let binding = MemoryHandle::binding(handle);
let resource = stream
.host_memory()
.get_resource(binding.clone(), None, None)
.ok()?;
Some(unsafe { D::pinned_bytes(binding, resource, size) })
}
pub fn read_async(
&mut self,
descriptors: Vec<CopyDescriptor>,
) -> impl Future<Output = Result<Vec<Bytes>, ServerError>> + Send + use<D> {
let held = descriptors
.iter()
.map(|descriptor| descriptor.handle.clone())
.collect::<Vec<_>>();
let result = self.copies_to_bytes(descriptors);
let fence = D::Stream::fence(self.streams.current().signal());
async move {
let synced = fence.wait();
core::mem::drop(held);
synced?;
result.map_err(Into::into)
}
}
fn copies_to_bytes(&mut self, descriptors: Vec<CopyDescriptor>) -> Result<Vec<Bytes>, IoError> {
let mut result = Vec::with_capacity(descriptors.len());
for descriptor in descriptors {
match self.copy_to_bytes(descriptor, None) {
Ok(bytes) => result.push(bytes),
Err(err) => {
if !result.is_empty() {
D::Stream::fence(self.streams.current().signal()).sync();
}
return Err(err);
}
}
}
Ok(result)
}
fn copy_to_bytes(
&mut self,
descriptor: CopyDescriptor,
stream_id: Option<StreamId>,
) -> Result<Bytes, IoError> {
let num_bytes = descriptor.shape.iter().product::<usize>() * descriptor.elem_size;
let mut bytes = self.reserve_cpu(num_bytes, stream_id);
self.write_to_cpu(descriptor, &mut bytes, stream_id)?;
Ok(bytes)
}
pub fn write_to_cpu(
&mut self,
descriptor: CopyDescriptor,
bytes: &mut Bytes,
stream_id: Option<StreamId>,
) -> Result<(), IoError> {
let CopyDescriptor {
handle: binding,
shape,
strides,
elem_size,
} = descriptor;
if bytes.is_empty() {
return Ok(());
}
let layout = CopyLayout::of(&shape, &strides, elem_size)?;
let resource = self.resource(binding)?;
let stream = match stream_id {
Some(id) => self.streams.get(&id),
None => self.streams.current(),
};
unsafe { D::copy_to_host(&resource, &layout, bytes, stream) }
}
pub fn write_to_gpu(&mut self, descriptor: CopyDescriptor, data: Bytes) -> Result<(), IoError> {
let CopyDescriptor {
handle: binding,
shape,
strides,
elem_size,
} = descriptor;
let size = data.len();
if size == 0 {
return Ok(());
}
let layout = CopyLayout::of(&shape, &strides, elem_size)?;
let resource = self.resource(binding)?;
let staging = Staging::of(size, data.property());
let data = match staging.through_pinned {
true => {
let mut buffer = self
.reserve_pinned(size, None)
.unwrap_or_else(|| Bytes::from_bytes_vec(vec![0; size]));
data.copy_into(&mut buffer);
buffer
}
false => data,
};
let current = self.streams.current();
unsafe { D::copy_to_device(&resource, &layout, &data, current)? };
if current.capturing().is_recording() {
current.capturing().retain_host(data);
} else {
current.drop_queue().push(data);
if staging.flush_after || current.drop_queue().should_flush() {
let signal = current.signal();
current.drop_queue().flush(|| D::Stream::fence(signal));
}
}
Ok(())
}
pub fn create_with_data(&mut self, data: &[u8]) -> Result<Handle, IoError> {
let mut staging =
self.reserve_pinned(data.len(), None)
.ok_or_else(|| IoError::Unknown {
backtrace: BackTrace::capture(),
description: "Unable to reserve pinned memory".into(),
})?;
staging.copy_from_slice(data);
let handle = self.empty(staging.len() as u64)?;
self.write_to_gpu(
CopyDescriptor {
handle: handle.clone().binding(),
shape: [data.len()].into(),
strides: [1].into(),
elem_size: 1,
},
staging,
)?;
Ok(handle)
}
pub fn sync(&mut self) -> DynFut<Result<(), ServerError>> {
let fence = D::Stream::fence(self.streams.current().signal());
Box::pin(async move { fence.wait() })
}
pub fn kernel(
&mut self,
kernel: KernelId,
count: (u32, u32, u32),
args: &mut D::LaunchArgs,
) -> Result<(), LaunchError> {
let stream = self.streams.current();
let result = D::launch(self.ctx, stream, kernel, count, args);
if !stream.capturing().is_recording() && stream.drop_queue().should_flush() {
let signal = stream.signal();
stream.drop_queue().flush(|| D::Stream::fence(signal));
}
result
}
}