Skip to main content

cubecl_wgpu/compute/
server.rs

1use std::marker::PhantomData;
2
3use super::storage::{WgpuResource, WgpuStorage};
4use crate::WgpuCompiler;
5use crate::schedule::{BindingsResource, ScheduleTask, ScheduledWgpuBackend};
6use alloc::sync::Arc;
7use cubecl_common::pool::LeasePool;
8use cubecl_common::{
9    backtrace::BackTrace,
10    bytes::Bytes,
11    profile::{ProfileDuration, TimingMethod},
12    stream_id::StreamId,
13};
14use cubecl_core::server::{Binding, StreamErrorMode};
15use cubecl_core::zspace::Shape;
16use cubecl_core::{
17    MemoryConfiguration, WgpuCompilationOptions,
18    future::DynFut,
19    prelude::*,
20    server::{
21        CopyDescriptor, IoError, KernelArguments, LaunchError, ProfileError, ProfilingToken,
22        ServerCommunication, ServerError, ServerUtilities,
23    },
24    zspace::{Strides, strides},
25};
26#[cfg(feature = "spirv")]
27use cubecl_core::{cache::CacheOption, compilation_cache::CompilationCache, hash::StableHash};
28use cubecl_ir::MemoryDeviceProperties;
29use cubecl_runtime::allocator::ContiguousMemoryLayoutPolicy;
30use cubecl_runtime::memory_management::{ManagedMemoryHandle, MemoryUsage, SharedMemoryBindings};
31use cubecl_runtime::{
32    compiler::CubeTask,
33    config::{CubeClRuntimeConfig, RuntimeConfig},
34    logging::ServerLogger,
35    memory_management::MemoryAllocationMode,
36    server::ComputeServer,
37    storage::ManagedResource,
38    stream::scheduler::{
39        SchedulerMultiStream, SchedulerMultiStreamOptions, SchedulerStrategy,
40        SchedulerStreamBackend,
41    },
42    validation::{validate_cube_dim, validate_units},
43};
44use hashbrown::HashMap;
45use wgpu::ComputePipeline;
46
47#[derive(Debug, Clone, Copy, PartialEq, Eq)]
48pub enum ParamsTransfer {
49    Immediate,
50    Uniform,
51}
52
53/// Compiler kind and info used when compiling a specific kernel. Used to determine parameter passing strategies.
54#[derive(Debug, Clone, Copy, PartialEq, Eq)]
55pub enum CompilerInfo {
56    Vulkan { params_transfer: ParamsTransfer },
57    Metal,
58    WGSL,
59    None,
60}
61
62/// Wgpu compute server.
63#[derive(Debug)]
64pub struct WgpuServer<C: WgpuCompiler> {
65    pub(crate) device: wgpu::Device,
66    // A buffer that can be used to store stream id without extra allocations.
67    streams_pool: Vec<StreamId>,
68    pipelines: HashMap<KernelId, (Arc<ComputePipeline>, CompilerInfo)>,
69    scheduler: SchedulerMultiStream<ScheduledWgpuBackend>,
70    #[cfg(feature = "spirv")]
71    pub(crate) spirv_cache:
72        Option<CompilationCache<(u64, StableHash), cubecl_spirv::SpirvCacheEntry>>,
73    pub compilation_options: WgpuCompilationOptions,
74    pub(crate) backend: wgpu::Backend,
75    pub(crate) utilities: Arc<ServerUtilities<Self>>,
76    /// Reusable buffers for the cross-stream input bindings of each launch.
77    shared_bindings_pool: LeasePool<SharedMemoryBindings>,
78    _compiler: PhantomData<C>,
79}
80
81impl<C: WgpuCompiler> ServerCommunication for WgpuServer<C> {
82    const SERVER_COMM_ENABLED: bool = false;
83}
84
85impl<C: WgpuCompiler> WgpuServer<C> {
86    /// Create a new server.
87    #[allow(clippy::too_many_arguments)]
88    pub fn new(
89        memory_properties: MemoryDeviceProperties,
90        memory_config: MemoryConfiguration,
91        compilation_options: WgpuCompilationOptions,
92        device: wgpu::Device,
93        queue: wgpu::Queue,
94        tasks_max: usize,
95        backend: wgpu::Backend,
96        timing_method: TimingMethod,
97        utilities: ServerUtilities<Self>,
98    ) -> Self {
99        #[cfg(feature = "spirv")]
100        let adapter_info = device.adapter_info();
101        let backend_scheduler = ScheduledWgpuBackend::new(
102            device.clone(),
103            queue.clone(),
104            memory_properties,
105            memory_config,
106            timing_method,
107            backend,
108            tasks_max,
109            utilities.logger.clone(),
110            compilation_options.supports_vulkan_compiler,
111        );
112
113        let config = CubeClRuntimeConfig::get();
114        let max_streams = config.streaming.max_streams;
115
116        Self {
117            compilation_options,
118            streams_pool: Vec::new(),
119            device,
120            pipelines: HashMap::new(),
121            scheduler: SchedulerMultiStream::new(
122                utilities.logger.clone(),
123                backend_scheduler,
124                SchedulerMultiStreamOptions {
125                    max_streams,
126                    max_tasks: tasks_max,
127                    strategy: SchedulerStrategy::Interleave,
128                },
129            ),
130            #[cfg(feature = "spirv")]
131            spirv_cache: {
132                let config = cubecl_runtime::config::CubeClRuntimeConfig::get();
133                if let Some(cache) = &config.compilation.cache {
134                    let root = cache.root();
135                    Some(CompilationCache::new(
136                        format!("spirv_{}_{}", adapter_info.vendor, adapter_info.device),
137                        CacheOption::default().name("vulkan").root(root),
138                    ))
139                } else {
140                    None
141                }
142            },
143            backend,
144            utilities: Arc::new(utilities),
145            shared_bindings_pool: LeasePool::with_capacity(tasks_max * max_streams as usize),
146            _compiler: PhantomData,
147        }
148    }
149
150    fn prepare_bindings(
151        &mut self,
152        bindings: KernelArguments,
153        compiler_info: CompilerInfo,
154    ) -> Result<BindingsResource, IoError> {
155        // Store all the resources we'll be using. This could be eliminated if
156        // there was a way to tie the lifetime of the resource to the memory handle.
157        let mut resources = Vec::with_capacity(bindings.buffers.len());
158
159        for b in bindings.buffers.into_iter() {
160            let stream = self.scheduler.stream(&b.stream);
161            let resource = stream.mem_manage.get_resource(b)?;
162            resources.push(resource);
163        }
164
165        Ok(BindingsResource {
166            resources,
167            info: bindings.info,
168            compiler_info,
169        })
170    }
171
172    fn pipeline(
173        &mut self,
174        kernel: <Self as ComputeServer>::Kernel,
175        bindings: &KernelArguments,
176        mode: ExecutionMode,
177    ) -> Result<(Arc<ComputePipeline>, CompilerInfo), LaunchError> {
178        let mut kernel_id = kernel.id();
179        kernel_id.mode(mode);
180
181        if let Some(pipeline) = self.pipelines.get(&kernel_id) {
182            return Ok(pipeline.clone());
183        }
184
185        let cached = self.load_cached_pipeline(&kernel_id, bindings, mode)?;
186
187        if let Some(Ok(pipeline)) = cached {
188            self.pipelines.insert(kernel_id, pipeline.clone());
189            return Ok(pipeline);
190        }
191
192        validate_cube_dim(&self.utilities.properties, &kernel_id)?;
193        validate_units(&self.utilities.properties, &kernel_id)?;
194
195        let mut compiler = C::init(self.backend, &self.compilation_options);
196        let mut compiled = compiler.compile_kernel(self, kernel, mode)?;
197
198        if self.scheduler.logger.compilation_source_activated() {
199            compiled.debug_info = Some(DebugInformation::new(
200                compiler.lang_tag(),
201                kernel_id.clone(),
202            ));
203        }
204        self.scheduler.logger.log_compilation(&compiled);
205
206        compiler.validate_ir(&compiled.repr, &self.utilities.properties)?;
207        let (compiler_info, auto_repr) = compiler.normalize_repr(compiled.repr);
208        let repr = auto_repr.as_ref().map(|r| r.as_ref());
209
210        // /!\ Do not delete the following commented code.
211        // This is useful while working on the metal compiler.
212        // Also the errors are printed nicely which is not the case when this is the runtime
213        // that does it.
214        // {
215        //     // Write shader in metal file then compile it for error
216        //     std::fs::write("shader.metal", &compiled.source).expect("should write to file");
217        //     let status = std::process::Command::new("xcrun")
218        //         .args(vec![
219        //             "-sdk",
220        //             "macosx",
221        //             "metal",
222        //             "-o",
223        //             "shader.ir",
224        //             "-c",
225        //             "shader.metal",
226        //             "-w",
227        //         ])
228        //         .status()
229        //         .expect("should launch the command");
230        //     if !status.success() {
231        //         println!("SOURCE:\n{}", compiled.source);
232        //         std::process::exit(status.code().unwrap());
233        //     }
234        // }
235
236        let module = self.create_module(
237            &compiled.entrypoint_name,
238            kernel_id.cube_dim,
239            repr,
240            &compiled.source,
241            mode,
242        )?;
243        let pipeline = self.create_pipeline(&compiled.entrypoint_name, repr, module, bindings);
244        self.pipelines
245            .insert(kernel_id.clone(), (pipeline.clone(), compiler_info));
246
247        #[cfg(feature = "spirv")]
248        if let Some(Err(key)) = cached
249            && let Some(crate::AutoRepresentation::SpirV(kernel)) = auto_repr
250        {
251            let cache = self.spirv_cache.as_mut().unwrap();
252            let result = cache.insert(
253                key,
254                cubecl_spirv::SpirvCacheEntry::new(compiled.entrypoint_name, kernel),
255            );
256            if let Err(err) = result {
257                log::warn!("Unable to save the SPIR-V {err:?}");
258            }
259        }
260
261        Ok((pipeline, compiler_info))
262    }
263}
264
265impl<C: WgpuCompiler> ComputeServer for WgpuServer<C> {
266    type Kernel = Box<dyn CubeTask<C>>;
267    type Storage = WgpuStorage;
268    type MemoryLayoutPolicy = ContiguousMemoryLayoutPolicy;
269    type Info = wgpu::Backend;
270
271    fn logger(&self) -> Arc<ServerLogger> {
272        self.scheduler.logger.clone()
273    }
274
275    fn utilities(&self) -> Arc<ServerUtilities<Self>> {
276        self.utilities.clone()
277    }
278
279    fn staging(
280        &mut self,
281        _sizes: &[usize],
282        _stream_id: StreamId,
283    ) -> Result<Vec<Bytes>, ServerError> {
284        // TODO: Check if using a staging buffer is useful here.
285        Err(IoError::UnsupportedIoOperation {
286            backtrace: BackTrace::capture(),
287        }
288        .into())
289    }
290
291    fn initialize_memory(&mut self, memory: ManagedMemoryHandle, size: u64, stream_id: StreamId) {
292        let stream = self.scheduler.stream(&stream_id);
293        let reserved = stream.empty(size).unwrap();
294        stream.mem_manage.bind(reserved, memory);
295    }
296
297    fn read(
298        &mut self,
299        descriptors: Vec<CopyDescriptor>,
300        stream_id: StreamId,
301    ) -> DynFut<Result<Vec<Bytes>, ServerError>> {
302        let mut streams = vec![stream_id];
303        let mut resources = Vec::with_capacity(descriptors.len());
304        for desc in descriptors {
305            if contiguous_strides(&desc.shape) != desc.strides {
306                return Box::pin(async {
307                    Err(IoError::UnsupportedStrides {
308                        backtrace: BackTrace::capture(),
309                    }
310                    .into())
311                });
312            }
313            if !streams.contains(&desc.handle.stream) {
314                streams.push(desc.handle.stream);
315            }
316            let stream = self.scheduler.stream(&desc.handle.stream);
317            let resource = match stream.mem_manage.get_resource(desc.handle) {
318                Ok(val) => val,
319                Err(err) => return Box::pin(async move { Err(err.into()) }),
320            };
321            resources.push((resource, desc.shape, desc.elem_size));
322        }
323
324        self.scheduler.execute_streams(streams);
325
326        let stream = self.scheduler.stream(&stream_id);
327        stream.read_resources(resources)
328    }
329
330    fn write(&mut self, descriptors: Vec<(CopyDescriptor, Bytes)>, stream_id: StreamId) {
331        for (desc, data) in descriptors {
332            let stream = self.scheduler.stream(&desc.handle.stream);
333
334            if contiguous_strides(&desc.shape) != desc.strides {
335                stream.error(ServerError::Io(IoError::UnsupportedStrides {
336                    backtrace: BackTrace::capture(),
337                }));
338                return;
339            }
340
341            let resource = match stream.mem_manage.get_resource(desc.handle) {
342                Ok(r) => r,
343                Err(err) => {
344                    stream.error(ServerError::Io(err));
345                    return;
346                }
347            };
348            let task = ScheduleTask::Write {
349                data,
350                buffer: resource,
351            };
352
353            self.scheduler.register(stream_id, task, &[]);
354        }
355    }
356
357    fn get_resource(
358        &mut self,
359        binding: Binding,
360        stream_id: StreamId,
361    ) -> Result<ManagedResource<WgpuResource>, ServerError> {
362        let mut streams = vec![stream_id];
363        if binding.stream != stream_id {
364            streams.push(binding.stream);
365        }
366        self.scheduler.execute_streams(streams);
367        let stream = self.scheduler.stream(&binding.stream);
368        let memory = binding.memory.clone();
369        let resource = stream.mem_manage.get_resource(binding)?;
370
371        Ok(ManagedResource::new(memory, resource))
372    }
373
374    unsafe fn launch(
375        &mut self,
376        kernel: Self::Kernel,
377        count: CubeCount,
378        args: KernelArguments,
379        mode: ExecutionMode,
380        stream_id: StreamId,
381    ) {
382        let (pipeline, compiler_info) = match self.pipeline(kernel, &args, mode) {
383            Ok(val) => val,
384            Err(err) => {
385                // We make the stream that would execute the kernel in error.
386                let stream = self.scheduler.stream(&stream_id);
387                stream.errors.push(ServerError::Launch(err));
388                return;
389            }
390        };
391
392        self.streams_pool.clear();
393        // Reuse a pooled buffer to avoid allocating on every launch; it returns to the pool
394        // automatically when the guard drops.
395        let mut shared_inputs = self.shared_bindings_pool.acquire();
396        // Pin the memory of every input that lives on another stream (released in `WgpuStream::flush`).
397        args.buffers.iter().for_each(|b| {
398            self.streams_pool.push(b.stream);
399            if b.stream != stream_id {
400                shared_inputs.push(b.memory.clone());
401            }
402        });
403
404        let resources = match self.prepare_bindings(args, compiler_info) {
405            Ok(val) => val,
406            Err(err) => {
407                // We make the stream that would execute the kernel in error.
408                let stream = self.scheduler.stream(&stream_id);
409                stream.errors.push(ServerError::Io(err));
410                return;
411            }
412        };
413        let task = ScheduleTask::Execute {
414            pipeline,
415            count,
416            resources,
417            shared_inputs,
418        };
419
420        self.scheduler.register(stream_id, task, &self.streams_pool);
421    }
422
423    fn flush(&mut self, stream_id: StreamId) -> Result<(), ServerError> {
424        self.scheduler.execute_streams(vec![stream_id]);
425
426        let stream = self.scheduler.stream(&stream_id);
427
428        stream.flush(StreamErrorMode {
429            ignore: false,
430            flush: true,
431        })
432    }
433
434    /// Returns the total time of GPU work this sync completes.
435    fn sync(&mut self, stream_id: StreamId) -> DynFut<Result<(), ServerError>> {
436        self.scheduler.execute_streams(vec![stream_id]);
437        let stream = self.scheduler.stream(&stream_id);
438
439        stream.sync()
440    }
441
442    fn start_profile(&mut self, stream_id: StreamId) -> Result<ProfilingToken, ServerError> {
443        self.scheduler.execute_streams(vec![stream_id]);
444        let stream = self.scheduler.stream(&stream_id);
445        stream.start_profile()
446    }
447
448    fn end_profile(
449        &mut self,
450        stream_id: StreamId,
451        token: ProfilingToken,
452    ) -> Result<ProfileDuration, ProfileError> {
453        self.scheduler.execute_streams(vec![stream_id]);
454        let stream = self.scheduler.stream(&stream_id);
455
456        stream.end_profile(token)
457    }
458
459    fn memory_usage(&mut self, stream_id: StreamId) -> Result<MemoryUsage, ServerError> {
460        self.scheduler.execute_streams(vec![stream_id]);
461        let stream = self.scheduler.stream(&stream_id);
462        Ok(stream.mem_manage.memory_usage())
463    }
464
465    fn stream_ids(&self) -> Vec<StreamId> {
466        self.scheduler.stream_ids().collect()
467    }
468
469    fn memory_cleanup(&mut self, stream_id: StreamId) {
470        self.scheduler.execute_streams(vec![stream_id]);
471        let stream = self.scheduler.stream(&stream_id);
472        stream.mem_manage.memory_cleanup(true);
473    }
474
475    fn allocation_mode(&mut self, mode: MemoryAllocationMode, stream_id: StreamId) {
476        self.scheduler.execute_streams(vec![stream_id]);
477        let stream = self.scheduler.stream(&stream_id);
478        stream.mem_manage.mode(mode);
479    }
480
481    fn configure_memory_pools(&mut self, config: MemoryConfiguration, stream_id: StreamId) -> bool {
482        // Streams created from now on build their main pool with the new
483        // layout; memory is per stream, so already-created streams keep theirs.
484        self.scheduler
485            .backend_mut()
486            .factory()
487            .set_gpu_pools(config.clone());
488        let (_, props) = self.scheduler.backend_mut().factory().gpu_pools();
489
490        // The calling stream's pools are rebuilt in place (kept, with a log,
491        // when something is still live in them).
492        self.scheduler.execute_streams(vec![stream_id]);
493        let stream = self.scheduler.stream(&stream_id);
494        stream.mem_manage.configure_memory_pools(config, &props)
495    }
496}
497
498pub(crate) fn contiguous_strides(shape: &Shape) -> Strides {
499    let rank = shape.len();
500    let mut strides = strides![1; rank];
501    for i in (0..rank - 1).rev() {
502        strides[i] = strides[i + 1] * shape[i + 1];
503    }
504    strides
505}