use crate::compute::{
storage::{
cpu::{PINNED_MEMORY_ALIGNMENT, PinnedMemoryStorage},
gpu::GpuStorage,
},
sync::Fence,
};
use cubecl_core::{
MemoryConfiguration,
ir::MemoryDeviceProperties,
server::{Binding, Handle, ServerError},
};
use cubecl_runtime::{
config::streaming::StreamPriority,
logging::ServerLogger,
memory_management::{
MemoryAllocationMode, MemoryManagement, MemoryManagementOptions, drop_queue,
},
metadata_cache::{CacheMode, MetadataCachePolicy, MetadataInfoCache},
stream::EventStreamBackend,
};
use std::{mem::MaybeUninit, sync::Arc};
#[derive(Debug)]
pub struct Stream {
pub sys: cudarc::driver::sys::CUstream,
pub memory_management_gpu: MemoryManagement<GpuStorage>,
pub memory_management_cpu: MemoryManagement<PinnedMemoryStorage>,
pub errors: Vec<ServerError>,
pub drop_queue: drop_queue::PendingDropQueue<Fence>,
pub capturing: StreamCaptureState,
pub info_cache: MetadataInfoCache<Handle>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StreamCaptureState {
NoCapture,
Prepare,
Capture,
}
impl StreamCaptureState {
pub fn is_recording(&self) -> bool {
matches!(self, StreamCaptureState::Capture)
}
pub fn cache_mode(&self) -> CacheMode {
match self {
StreamCaptureState::NoCapture => CacheMode::Normal,
StreamCaptureState::Prepare | StreamCaptureState::Capture => CacheMode::Capture,
}
}
}
impl drop_queue::Fence for Fence {
fn sync(self) {
let _ = self.wait_sync().ok();
}
}
#[derive(new, Debug)]
pub struct CudaStreamBackend {
mem_props: MemoryDeviceProperties,
mem_config: MemoryConfiguration,
mem_alignment: usize,
logger: Arc<ServerLogger>,
priority: StreamPriority,
#[new(default)]
gpu_pools_override: Option<MemoryConfiguration>,
}
impl CudaStreamBackend {
pub(crate) fn gpu_pools(&self) -> (MemoryConfiguration, MemoryDeviceProperties) {
let config = self
.gpu_pools_override
.clone()
.unwrap_or_else(|| self.mem_config.clone());
(config, self.mem_props.clone())
}
pub(crate) fn set_gpu_pools(&mut self, config: MemoryConfiguration) {
self.gpu_pools_override = Some(config);
}
}
pub(crate) fn create_cuda_stream(priority: StreamPriority) -> cudarc::driver::sys::CUstream {
use cudarc::driver::sys::{self, CUstream_flags};
let use_greatest = match priority {
StreamPriority::Default => {
return cudarc::driver::result::stream::create(
cudarc::driver::result::stream::StreamKind::NonBlocking,
)
.expect("Can create a new stream.");
}
StreamPriority::High => true,
StreamPriority::Low => false,
};
let value = unsafe {
let mut least: i32 = 0;
let mut greatest: i32 = 0;
sys::cuCtxGetStreamPriorityRange(&mut least, &mut greatest)
.result()
.expect("Can query CUDA stream priority range.");
if use_greatest { greatest } else { least }
};
unsafe {
let mut stream = MaybeUninit::uninit();
sys::cuStreamCreateWithPriority(
stream.as_mut_ptr(),
CUstream_flags::CU_STREAM_NON_BLOCKING as u32,
value,
)
.result()
.expect("Can create a new CUDA stream with priority.");
stream.assume_init()
}
}
impl EventStreamBackend for CudaStreamBackend {
type Stream = Stream;
type Event = Fence;
fn create_stream(&self) -> Self::Stream {
let stream = create_cuda_stream(self.priority);
let storage = GpuStorage::new(self.mem_alignment, stream);
let (gpu_config, gpu_props) = self.gpu_pools();
let memory_management_gpu = MemoryManagement::from_configuration(
storage,
&gpu_props,
gpu_config,
self.logger.clone(),
MemoryManagementOptions::new("Main GPU Memory"),
);
let memory_management_cpu = MemoryManagement::from_configuration(
PinnedMemoryStorage::new(),
&MemoryDeviceProperties {
max_page_size: self.mem_props.max_page_size,
alignment: PINNED_MEMORY_ALIGNMENT as u64,
},
self.mem_config.clone(),
self.logger.clone(),
MemoryManagementOptions::new("Pinned CPU Memory").mode(MemoryAllocationMode::Auto),
);
Stream {
sys: stream,
memory_management_gpu,
memory_management_cpu,
errors: Vec::new(),
drop_queue: Default::default(),
capturing: StreamCaptureState::NoCapture,
info_cache: MetadataInfoCache::new(MetadataCachePolicy::default()),
}
}
fn flush(stream: &mut Self::Stream) -> Self::Event {
Fence::new(stream.sys)
}
fn wait_event(stream: &mut Self::Stream, event: Self::Event) {
event.wait_async(stream.sys);
}
fn wait_event_sync(event: Self::Event) -> Result<(), ServerError> {
event.wait_sync()
}
fn handle_cursor(stream: &Self::Stream, binding: &Binding) -> u64 {
stream
.memory_management_gpu
.get_cursor(binding.memory.clone())
.unwrap_or(u64::MAX)
}
fn is_healthy(stream: &Self::Stream) -> bool {
stream.errors.is_empty()
}
}