use cubecl_llvm::PlironOptions;
use crate::{
CpuCompiler,
compute::{
cpu_kernel::CpuKernel,
schedule::{BindingsResource, ScheduleTask, ScheduledCpuBackend},
},
};
use cubecl_common::{bytes::Bytes, profile::ProfileDuration};
use cubecl_core::server::ServerStorage;
use cubecl_core::{
CompilationError, CubeCount, MemoryConfiguration, MemoryUsage,
ir::MemoryDeviceProperties,
server::{
BufferBinding, CopyDescriptor, IoError, KernelArguments, KernelResource, LaunchError,
ProfileError, ProfilingToken, Server, ServerCommunication, ServerError, ServerUtilities,
},
zspace::{Shape, Strides, strides},
};
use cubecl_environment::backtrace::BackTrace;
use cubecl_environment::future::DynFut;
use cubecl_environment::stream::StreamId;
use cubecl_server::{
config::{CubeClRuntimeConfig, RuntimeConfig, compilation::F16Evaluation},
dry_run::LaunchMode,
id::KernelId,
kernel::{CompiledKernel, CubeKernel},
logging::ServerLogger,
memory_management::{ManagedMemoryHandle, MemoryAllocationMode},
storage::{BytesStorage, ComputeStorage, ManagedResource},
stream::scheduler::{SchedulerMultiStream, SchedulerMultiStreamOptions, SchedulerStrategy},
stream::{ExecuteScope, FailureStore, WriteScoped, failed_writing},
};
use std::{collections::HashMap, sync::Arc};
#[derive(Debug)]
pub struct CpuServer {
scheduler: SchedulerMultiStream<ScheduledCpuBackend>,
utilities: Arc<ServerUtilities>,
compilation_cache: HashMap<(KernelId, u32), CpuKernel>,
compilation_options: PlironOptions,
streams_pool: Vec<StreamId>,
}
impl WriteScoped for CpuServer {
type Streams = SchedulerMultiStream<ScheduledCpuBackend>;
fn write_streams(&mut self) -> &mut Self::Streams {
&mut self.scheduler
}
fn on_failure(&mut self, stream: StreamId, error: &ServerError) {
self.scheduler.stream(&stream).profile_failure(error);
}
}
impl CpuServer {
pub fn new(
memory_properties: MemoryDeviceProperties,
memory_config: MemoryConfiguration,
f16_evaluation: F16Evaluation,
utilities: Arc<ServerUtilities>,
) -> Self {
let backend =
ScheduledCpuBackend::new(memory_properties, memory_config, utilities.logger.clone());
let config = CubeClRuntimeConfig::get();
let max_streams = config.streaming.max_streams;
let scheduler = SchedulerMultiStream::new(
utilities.logger.clone(),
backend,
SchedulerMultiStreamOptions {
max_streams,
max_tasks: 8,
strategy: SchedulerStrategy::Interleave,
},
);
Self {
scheduler,
utilities,
compilation_cache: HashMap::new(),
compilation_options: PlironOptions {
f16_evaluation,
..Default::default()
},
streams_pool: Vec::new(),
}
}
fn prepare_bindings(&mut self, bindings: KernelArguments) -> BindingsResource {
let resources = bindings
.resources
.into_iter()
.filter_map(|binding| {
let KernelResource::Buffer(binding) = binding else {
return None;
};
let stream = self.scheduler.stream(&binding.stream);
let memory = binding.memory.clone();
let resource = stream
.memory_management
.get_resource(binding.memory, binding.offset_start, binding.offset_end)
.unwrap();
Some(ManagedResource::new(memory, resource))
})
.collect::<Vec<_>>();
BindingsResource {
resources,
info: bindings.info,
}
}
fn prepare_task(
&mut self,
kernel_id: (KernelId, u32),
count: CubeCount,
bindings: BindingsResource,
stream_id: StreamId,
) -> Result<ScheduleTask, CompilationError> {
let cube_count = match count {
CubeCount::Static(x, y, z) => [x, y, z],
CubeCount::Dynamic(binding) => {
let stream = self.scheduler.stream(&binding.stream);
let resource = stream
.memory_management
.get_resource(binding.memory, binding.offset_start, binding.offset_end)
.unwrap();
stream.submit();
let bytes = resource.read();
let x = u32::from_ne_bytes(bytes[0..4].try_into().unwrap());
let y = u32::from_ne_bytes(bytes[4..8].try_into().unwrap());
let z = u32::from_ne_bytes(bytes[8..12].try_into().unwrap());
[x, y, z]
}
};
self.prepare_task_inner(kernel_id, cube_count, bindings, stream_id)
}
fn compile_only(
&mut self,
kernel: &dyn CubeKernel,
alignment: u32,
) -> Result<(), CompilationError> {
let kernel_id = (kernel.id(), alignment);
if self.compilation_cache.contains_key(&kernel_id) {
return Ok(());
}
let definition = kernel.define();
let options = PlironOptions {
cpu_buffer_alignment: Some(alignment),
..self.compilation_options.clone()
};
let compiled =
CompiledKernel::compile(kernel, definition, &mut CpuCompiler::default(), &options)?;
if compiled.repr.is_none() {
return Err(CompilationError::Generic {
reason: format!(
"the CPU runtime cannot load the precompiled kernel `{}`: it runs compiled IR, not source text",
kernel.name()
),
backtrace: BackTrace::capture(),
});
}
self.compilation_cache
.insert(kernel_id, CpuKernel::new(compiled));
Ok(())
}
fn prepare_task_inner(
&mut self,
kernel_id: (KernelId, u32),
cube_count: [u32; 3],
bindings: BindingsResource,
stream_id: StreamId,
) -> Result<ScheduleTask, CompilationError> {
let kernel = self
.compilation_cache
.get_mut(&kernel_id)
.expect("compiled before the write scope was entered");
let cube_dim = kernel.mlir.cube_dim;
let mlir_engine = kernel
.mlir
.repr
.clone()
.expect("compile_only refuses a kernel without a representation")
.expect_jit();
let task = ScheduleTask::Execute {
stream_id,
pliron_engine: mlir_engine,
bindings,
cube_dim,
cube_count,
};
Ok(task)
}
pub(crate) fn utilities(&self) -> Arc<ServerUtilities> {
self.utilities.clone()
}
}
impl Server for CpuServer {
fn logger(&self) -> Arc<ServerLogger> {
self.scheduler.logger.clone()
}
fn staging(
&mut self,
_sizes: &[usize],
_stream_id: StreamId,
) -> Result<Vec<Bytes>, ServerError> {
Err(IoError::UnsupportedIoOperation {
backtrace: BackTrace::capture(),
}
.into())
}
fn utilities(&self) -> Arc<ServerUtilities> {
self.utilities.clone()
}
fn initialize_memory(&mut self, memory: ManagedMemoryHandle, size: u64, stream_id: StreamId) {
let (stream, failures) = self.scheduler.stream_and_failures(&stream_id);
let reserved = stream
.empty(size, failures)
.unwrap_or_else(|err| panic!("failed to reserve {size} bytes of host memory: {err}"));
stream.bind(reserved, memory, failures);
}
fn read(
&mut self,
descriptors: Vec<CopyDescriptor>,
stream_id: StreamId,
) -> DynFut<Result<Vec<Bytes>, ServerError>> {
if let Err(err) = self
.scheduler
.ensure_written(descriptors.iter().map(|d| &d.handle))
{
return Box::pin(async move { Err(err) });
}
let mut streams = vec![stream_id];
let mut results = Vec::with_capacity(descriptors.len());
let mut resources = Vec::with_capacity(descriptors.len());
for desc in descriptors {
if !streams.contains(&desc.handle.stream) {
streams.push(desc.handle.stream);
}
let stream = self.scheduler.stream(&desc.handle.stream);
let result = stream.read_async(desc);
results.push(result);
}
self.scheduler.execute_streams(streams);
if let Err(err) = self.scheduler.stream(&stream_id).flush(stream_id) {
return Box::pin(async move { Err(err) });
}
Box::pin(async move {
for result in results {
match result.await {
Ok(val) => resources.push(val),
Err(err) => return Err(err.into()),
}
}
Ok(resources)
})
}
fn write(&mut self, descriptors: Vec<(CopyDescriptor, Bytes)>, stream_id: StreamId) {
for (desc, data) in descriptors {
let mut written = self.write_set();
written.push(desc.handle.clone());
ExecuteScope::over(self, stream_id, written).execute(|server| {
if contiguous_strides(&desc.shape) != desc.strides {
return Err(ServerError::Io(IoError::UnsupportedStrides {
backtrace: BackTrace::capture(),
}));
}
let owner = desc.handle.stream;
let memory = desc.handle.memory.clone();
let stream = server.scheduler.stream(&owner);
let resource = stream.get_resource(desc.handle).map_err(ServerError::Io)?;
let task = ScheduleTask::Write {
data,
buffer: ManagedResource::new(memory, resource),
};
server.scheduler.register(stream_id, task, &[owner]);
Ok(())
});
}
}
fn memory_usage(&mut self, stream_id: StreamId) -> MemoryUsage {
self.scheduler
.stream(&stream_id)
.memory_management
.memory_usage()
}
fn memory_report(
&mut self,
stream_id: StreamId,
) -> cubecl_server::memory_management::MemoryReport {
self.scheduler
.stream(&stream_id)
.memory_management
.memory_report()
}
fn stream_ids(&self) -> Vec<StreamId> {
self.scheduler.stream_ids().collect()
}
fn memory_cleanup(&mut self, stream_id: StreamId) {
let (stream, failures) = self.scheduler.stream_and_failures(&stream_id);
stream.memory_management.cleanup(true, failures)
}
unsafe fn launch(
&mut self,
kernel: Box<dyn CubeKernel>,
count: CubeCount,
bindings: KernelArguments,
stream_id: StreamId,
launch_mode: LaunchMode,
) {
let alignment =
bindings
.resources
.iter()
.fold(
BytesStorage::ALIGNMENT as u32,
|align, resource| match resource {
KernelResource::Buffer(binding) => {
let bits = binding.offset_start.unwrap_or(0).trailing_zeros();
1u32 << bits.min(align.trailing_zeros())
}
_ => align,
},
);
let kernel_id = kernel.id();
let cache_key = (kernel_id.clone(), alignment);
if let Err(err) = self.compile_only(kernel.as_ref(), alignment) {
let error = ServerError::Launch(LaunchError::CompilationError(err));
self.scheduler.stream(&stream_id).profile_failure(&error);
if !launch_mode.is_skipped() {
let mut written = self.write_set();
written.extend(bindings.buffers_written(None).cloned());
failed_writing(self, stream_id, written, error);
}
return;
}
if launch_mode.is_skipped() {
return;
}
let io = self
.compilation_cache
.get(&cache_key)
.and_then(|kernel| kernel.mlir.io.clone());
let mut written = self.write_set();
written.extend(bindings.buffers_written(io.as_deref()).cloned());
let count_read = match &count {
CubeCount::Dynamic(binding) => Some(binding),
CubeCount::Static(..) => None,
};
ExecuteScope::launching(
self,
kernel_id.clone(),
stream_id,
bindings.buffers_read(io.as_deref()).chain(count_read),
written,
)
.execute(|server| {
server.streams_pool.clear();
bindings
.resources
.iter()
.filter_map(|b| {
let KernelResource::Buffer(b) = b else {
return None;
};
Some(b)
})
.for_each(|b| server.streams_pool.push(b.stream));
let bindings = server.prepare_bindings(bindings);
let task = server
.prepare_task(cache_key, count, bindings, stream_id)
.map_err(|err| ServerError::Launch(LaunchError::CompilationError(err)))?;
server
.scheduler
.register(stream_id, task, &server.streams_pool);
Ok(())
});
}
fn check(
&mut self,
handles: Vec<BufferBinding>,
_stream_id: StreamId,
) -> Result<(), ServerError> {
self.scheduler.ensure_written(handles.iter())
}
fn flush(&mut self, stream_id: StreamId) -> Result<(), ServerError> {
self.scheduler.execute_streams(vec![stream_id]);
let stream = self.scheduler.stream(&stream_id);
stream.flush(stream_id)
}
fn sync(
&mut self,
handles: Vec<BufferBinding>,
stream_id: StreamId,
) -> DynFut<Result<(), ServerError>> {
self.scheduler.execute_streams(vec![stream_id]);
let stream = self.scheduler.stream(&stream_id);
let mut result = stream.flush(stream_id);
if result.is_ok() {
result = self.scheduler.ensure_written(handles.iter());
}
Box::pin(async move { result })
}
fn start_profile(&mut self, stream_id: StreamId) -> Result<ProfilingToken, ServerError> {
self.scheduler.execute_streams(vec![stream_id]);
let stream = self.scheduler.stream(&stream_id);
stream.start_profile(stream_id)
}
fn end_profile(
&mut self,
stream_id: StreamId,
token: ProfilingToken,
) -> Result<ProfileDuration, ProfileError> {
self.scheduler.execute_streams(vec![stream_id]);
let stream = self.scheduler.stream(&stream_id);
stream.end_profile(token, stream_id)
}
fn abandon_profile(&mut self, stream_id: StreamId, token: ProfilingToken) {
self.scheduler.stream(&stream_id).abandon_profile(token);
}
fn allocation_mode(&mut self, mode: MemoryAllocationMode, stream_id: StreamId) {
let stream = self.scheduler.stream(&stream_id);
stream.allocation_mode(mode);
}
}
impl ServerCommunication for CpuServer {}
pub(crate) fn contiguous_strides(shape: &Shape) -> Strides {
let rank = shape.len();
let mut strides = strides![1; rank];
for i in (0..rank - 1).rev() {
strides[i] = strides[i + 1] * shape[i + 1];
}
strides
}
impl ServerStorage for CpuServer {
type Storage = BytesStorage;
fn get_resource(
&mut self,
binding: BufferBinding,
stream_id: StreamId,
) -> Result<ManagedResource<<Self::Storage as ComputeStorage>::Resource>, ServerError> {
self.scheduler.ensure_written([&binding].into_iter())?;
let mut streams = vec![stream_id];
if binding.stream != stream_id {
streams.push(binding.stream);
}
self.scheduler.execute_streams(streams);
let stream = self.scheduler.stream(&binding.stream);
let memory = binding.memory.clone();
let resource = stream.get_resource(binding)?;
Ok(ManagedResource::new(memory, resource))
}
}