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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
55pub enum CompilerInfo {
56 Vulkan { params_transfer: ParamsTransfer },
57 Metal,
58 WGSL,
59 None,
60}
61
62#[derive(Debug)]
64pub struct WgpuServer<C: WgpuCompiler> {
65 pub(crate) device: wgpu::Device,
66 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 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 #[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 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 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 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 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 let mut shared_inputs = self.shared_bindings_pool.acquire();
396 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 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 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 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 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}