Skip to main content

ruda_runtime/runtime/
client.rs

1use crate::runtime::{
2    config::{TypeNameFormatLevel, type_name_format},
3    kernel::KernelMetadata,
4    logging::ProfileLevel,
5    memory_management::{MemoryAllocationMode, MemoryUsage},
6    backend::Runtime,
7    server::{
8        CommunicationId, ComputeServer, CopyDescriptor, RudaCount, ExecutionMode, Handle, IoError,
9        KernelArguments, MemoryLayout, MemoryLayoutDescriptor, MemoryLayoutPolicy,
10        MemoryLayoutStrategy, ProfileError, ReduceOperation, ServerCommunication, ServerError,
11        ServerUtilities,
12    },
13    storage::{ComputeStorage, ManagedResource},
14};
15use alloc::{format, sync::Arc, vec, vec::Vec};
16use ruda_core::{
17    backtrace::BackTrace,
18    bytes::{AllocationProperty, Bytes},
19    device::{Device, DeviceId},
20    device_handle::DeviceHandle,
21    future::DynFut,
22    profile::ProfileDuration,
23};
24use ruda_core::ir::{DeviceProperties, ElemType, VectorSize, features::Features};
25use ruda_core::tensor::Shape;
26
27#[allow(unused)]
28use ruda_core::profile::TimingMethod;
29use ruda_core::stream_id::StreamId;
30
31/// The `ComputeClient` is the entry point to require tasks from the `ComputeServer`.
32/// It should be obtained for a specific device via the Compute struct.
33pub struct ComputeClient<R: Runtime> {
34    device: DeviceHandle<R::Server>,
35    utilities: Arc<ServerUtilities<R::Server>>,
36    stream_id: Option<StreamId>,
37}
38
39impl<R: Runtime> Clone for ComputeClient<R> {
40    fn clone(&self) -> Self {
41        Self {
42            device: self.device.clone(),
43            utilities: self.utilities.clone(),
44            stream_id: self.stream_id,
45        }
46    }
47}
48
49impl<R: Runtime> ComputeClient<R> {
50    /// Get the info of the current backend.
51    pub fn info(&self) -> &<R::Server as ComputeServer>::Info {
52        &self.utilities.info
53    }
54
55    /// Physical runtime device ordinal/type used by tuning identity probes.
56    pub fn device_id(&self) -> DeviceId { self.device.device_id() }
57
58    /// Cached capability/hardware checksum. Does not probe the driver on every operator call.
59    pub fn properties_fingerprint(&self) -> u64 { self.utilities.properties_hash }
60
61    /// Create a new client with a new server.
62    pub fn init<D: Device>(device: &D, server: R::Server) -> Self {
63        let utilities = server.utilities();
64        let context = DeviceHandle::<R::Server>::insert(device.to_id(), server)
65            .expect("Can't create a new client on an already registered server");
66
67        Self {
68            device: context,
69            utilities,
70            stream_id: None,
71        }
72    }
73
74    /// Load the client for the given device.
75    pub fn load<D: Device>(device: &D) -> Self {
76        let context = DeviceHandle::<R::Server>::new(device.to_id());
77
78        // This is safe because we now know the return type of [`DeviceHandle::utilities()`].
79        let utilities = context
80            .utilities()
81            .downcast::<ServerUtilities<R::Server>>()
82            .expect("Can downcast to `ServerUtilities`");
83
84        Self {
85            device: context,
86            utilities,
87            stream_id: None,
88        }
89    }
90
91    fn stream_id(&self) -> StreamId {
92        match self.stream_id {
93            Some(val) => val,
94            None => StreamId::current(),
95        }
96    }
97
98    /// Resolve the logical execution stream on the calling thread.
99    /// A reusable execution plan must also hold `fixed_execution_queue()`;
100    /// an implicit thread-local stream must not change between plan operations.
101    pub fn execution_stream(&self) -> StreamId { self.stream_id() }
102
103    /// Whether both clients currently submit to the same server and stream.
104    /// Implicit streams are resolved on the calling thread at the time of this check.
105    pub fn same_execution_queue(&self, other: &Self) -> bool {
106        self.device.device_id() == other.device.device_id() && self.stream_id() == other.stream_id()
107    }
108
109    /// Clone this client with its currently resolved execution queue fixed.
110    /// Useful for reusable device plans whose scratch must not silently move
111    /// between thread-local streams. Does not create a stream or synchronize.
112    pub fn fixed_execution_queue(&self) -> Self {
113        let mut client = self.clone();
114        client.stream_id = Some(self.stream_id());
115        client
116    }
117
118    /// Set the stream in which the current client is operating on.
119    ///
120    /// # Safety
121    ///
122    /// This is highly unsafe and should probably only be used by the Ruda/Ruda projects for now.
123    pub unsafe fn set_stream(&mut self, stream_id: StreamId) {
124        self.stream_id = Some(stream_id);
125    }
126
127    fn do_read(&self, descriptors: Vec<CopyDescriptor>) -> DynFut<Result<Vec<Bytes>, ServerError>> {
128        let stream_id = self.stream_id();
129        self.device
130            .submit_blocking(move |server| server.read(descriptors, stream_id))
131            .unwrap()
132    }
133
134    /// Given bindings, returns owned resources as bytes.
135    pub fn read_async(
136        &self,
137        handles: Vec<Handle>,
138    ) -> impl Future<Output = Result<Vec<Bytes>, ServerError>> + Send {
139        let shapes = handles
140            .iter()
141            .map(|it| [it.size_in_used() as usize].into())
142            .collect::<Vec<Shape>>();
143        let descriptors = handles
144            .into_iter()
145            .zip(shapes)
146            .map(|(handle, shape)| CopyDescriptor::new(handle.binding(), shape, [1].into(), 1))
147            .collect();
148
149        self.do_read(descriptors)
150    }
151
152    /// Given bindings, returns owned resources as bytes.
153    ///
154    /// # Remarks
155    ///
156    /// Panics if the read operation fails.
157    pub fn read(&self, handles: Vec<Handle>) -> Vec<Bytes> {
158        ruda_core::reader::read_sync(self.read_async(handles)).expect("TODO")
159    }
160
161    /// Given a binding, returns owned resource as bytes.
162    pub fn read_one(&self, handle: Handle) -> Result<Bytes, ServerError> {
163        Ok(ruda_core::reader::read_sync(self.read_async(vec![handle]))?.remove(0))
164    }
165
166    /// Given a binding, returns owned resource as bytes.
167    ///
168    /// # Remarks
169    ///
170    /// Panics if the read operation fails. Useful for tests.
171    pub fn read_one_unchecked(&self, handle: Handle) -> Bytes {
172        ruda_core::reader::read_sync(self.read_async(vec![handle]))
173            .unwrap()
174            .remove(0)
175    }
176
177    /// Given bindings, returns owned resources as bytes.
178    pub fn read_tensor_async(
179        &self,
180        descriptors: Vec<CopyDescriptor>,
181    ) -> impl Future<Output = Result<Vec<Bytes>, ServerError>> + Send {
182        self.do_read(descriptors)
183    }
184
185    /// Given bindings, returns owned resources as bytes.
186    ///
187    /// # Remarks
188    ///
189    /// Panics if the read operation fails.
190    ///
191    /// The tensor must be in the same layout as created by the runtime, or more strict.
192    /// Contiguous tensors are always fine, strided tensors are only ok if the stride is similar to
193    /// the one created by the runtime (i.e. padded on only the last dimension). A way to check
194    /// stride compatibility on the runtime will be added in the future.
195    ///
196    /// Also see [`ComputeClient::create_tensor`].
197    pub fn read_tensor(&self, descriptors: Vec<CopyDescriptor>) -> Vec<Bytes> {
198        ruda_core::reader::read_sync(self.read_tensor_async(descriptors)).expect("TODO")
199    }
200
201    /// Given a binding, returns owned resource as bytes.
202    /// See [`ComputeClient::read_tensor`]
203    pub fn read_one_tensor_async(
204        &self,
205        descriptor: CopyDescriptor,
206    ) -> impl Future<Output = Result<Bytes, ServerError>> + Send {
207        let fut = self.read_tensor_async(vec![descriptor]);
208
209        async { Ok(fut.await?.remove(0)) }
210    }
211
212    /// Given a binding, returns owned resource as bytes.
213    ///
214    /// # Remarks
215    ///
216    /// Panics if the read operation fails.
217    /// See [`ComputeClient::read_tensor`]
218    pub fn read_one_unchecked_tensor(&self, descriptor: CopyDescriptor) -> Bytes {
219        self.read_tensor(vec![descriptor]).remove(0)
220    }
221
222    /// Given a resource handle, returns the storage resource.
223    pub fn get_resource(
224        &self,
225        handle: Handle,
226    ) -> Result<
227        ManagedResource<<<R::Server as ComputeServer>::Storage as ComputeStorage>::Resource>,
228        ServerError,
229    > {
230        let stream_id = self.stream_id();
231        let binding = handle.binding();
232
233        self.device
234            .submit_blocking(move |state| state.get_resource(binding, stream_id))
235            .unwrap()
236    }
237
238    fn do_create_from_slices(
239        &self,
240        descriptors: Vec<MemoryLayoutDescriptor>,
241        slices: Vec<Vec<u8>>,
242    ) -> Result<Vec<MemoryLayout>, IoError> {
243        let stream_id = self.stream_id();
244        let (handle_base, layouts) = self.utilities.layout_policy.apply(stream_id, &descriptors);
245
246        let descriptors = descriptors
247            .into_iter()
248            .zip(layouts.iter())
249            .zip(slices)
250            .map(|((desc, alloc), data)| {
251                (
252                    CopyDescriptor::new(
253                        alloc.memory.clone().binding(),
254                        desc.shape,
255                        alloc.strides.clone(),
256                        desc.elem_size,
257                    ),
258                    Bytes::from_bytes_vec(data),
259                )
260            })
261            .collect::<Vec<_>>();
262
263        let (size, memory) = (handle_base.size(), handle_base.memory);
264        self.device.submit(move |server| {
265            server.initialize_memory(memory, size, stream_id);
266            server.write(descriptors, stream_id);
267        });
268
269        Ok(layouts)
270    }
271
272    fn do_create(
273        &self,
274        descriptors: Vec<MemoryLayoutDescriptor>,
275        mut data: Vec<Bytes>,
276    ) -> Result<Vec<MemoryLayout>, IoError> {
277        self.staging(data.iter_mut(), true);
278
279        let stream_id = self.stream_id();
280        let (handle_base, layouts) = self.utilities.layout_policy.apply(stream_id, &descriptors);
281
282        let descriptors = descriptors
283            .into_iter()
284            .zip(layouts.iter())
285            .zip(data)
286            .map(|((desc, layout), data)| {
287                (
288                    CopyDescriptor::new(
289                        layout.memory.clone().binding(),
290                        desc.shape,
291                        layout.strides.clone(),
292                        desc.elem_size,
293                    ),
294                    data,
295                )
296            })
297            .collect::<Vec<_>>();
298
299        let (size, memory) = (handle_base.size(), handle_base.memory);
300        self.device.submit(move |server| {
301            server.initialize_memory(memory, size, stream_id);
302            server.write(descriptors, stream_id);
303        });
304
305        Ok(layouts)
306    }
307
308    /// Returns a resource handle containing the given data.
309    ///
310    /// # Notes
311    ///
312    /// Prefer using the more efficient [`Self::create`] function.
313    pub fn create_from_slice(&self, slice: &[u8]) -> Handle {
314        let shape: Shape = [slice.len()].into();
315
316        self.do_create_from_slices(
317            vec![MemoryLayoutDescriptor::new(
318                MemoryLayoutStrategy::Contiguous,
319                shape,
320                1,
321            )],
322            vec![slice.to_vec()],
323        )
324        .unwrap()
325        .remove(0)
326        .memory
327    }
328
329    /// todo: docs
330    pub fn exclusive<'a, Re: Send + 'static, F: FnOnce() -> Re + Send + 'a>(
331        &'a self,
332        task: F,
333    ) -> Result<Re, ServerError> {
334        // We then launch the task.
335        self.device
336            .exclusive(task)
337            .map_err(|err| ServerError::Generic {
338                reason: format!("Communication channel with the server is down: {err:?}"),
339                backtrace: BackTrace::capture(),
340            })
341    }
342
343    /// dodo: Docs
344    pub fn memory_persistent_allocation<
345        'a,
346        Re: Send,
347        Input: Send,
348        F: FnOnce(Input) -> Re + Send + 'a,
349    >(
350        &'a self,
351        input: Input,
352        task: F,
353    ) -> Result<Re, ServerError> {
354        let stream_id = StreamId::current();
355
356        self.device.submit(move |server| {
357            server.allocation_mode(MemoryAllocationMode::Persistent, stream_id);
358        });
359
360        // All tasks created on the same stream will have persistent memory.
361        let output = task(input);
362
363        self.device.submit(move |server| {
364            server.allocation_mode(MemoryAllocationMode::Auto, stream_id);
365        });
366
367        Ok(output)
368    }
369
370    /// Returns a resource handle containing the given [Bytes].
371    pub fn create(&self, data: Bytes) -> Handle {
372        let shape = [data.len()].into();
373
374        self.do_create(
375            vec![MemoryLayoutDescriptor::new(
376                MemoryLayoutStrategy::Contiguous,
377                shape,
378                1,
379            )],
380            vec![data],
381        )
382        .unwrap()
383        .remove(0)
384        .memory
385    }
386
387    /// Given a resource and shape, stores it and returns the tensor handle and strides.
388    /// This may or may not return contiguous strides. The layout is up to the runtime, and care
389    /// should be taken when indexing.
390    ///
391    /// Currently the tensor may either be contiguous (most runtimes), or "pitched", to use the CUDA
392    /// terminology. This means the last (contiguous) dimension is padded to fit a certain alignment,
393    /// and the strides are adjusted accordingly. This can make memory accesses significantly faster
394    /// since all rows are aligned to at least 16 bytes (the maximum load width), meaning the GPU
395    /// can load as much data as possible in a single instruction. It may be aligned even more to
396    /// also take cache lines into account.
397    ///
398    /// However, the stride must be taken into account when indexing and reading the tensor
399    /// (also see [`ComputeClient::read_tensor`]).
400    ///
401    /// # Notes
402    ///
403    /// Prefer using [`Self::create_tensor`] for better performance.
404    pub fn create_tensor_from_slice(
405        &self,
406        slice: &[u8],
407        shape: Shape,
408        elem_size: usize,
409    ) -> MemoryLayout {
410        self.do_create_from_slices(
411            vec![MemoryLayoutDescriptor::new(
412                MemoryLayoutStrategy::Optimized,
413                shape,
414                elem_size,
415            )],
416            vec![slice.to_vec()],
417        )
418        .unwrap()
419        .remove(0)
420    }
421
422    /// Given a resource and shape, stores it and returns the tensor handle and strides.
423    /// This may or may not return contiguous strides. The layout is up to the runtime, and care
424    /// should be taken when indexing.
425    ///
426    /// Currently the tensor may either be contiguous (most runtimes), or "pitched", to use the CUDA
427    /// terminology. This means the last (contiguous) dimension is padded to fit a certain alignment,
428    /// and the strides are adjusted accordingly. This can make memory accesses significantly faster
429    /// since all rows are aligned to at least 16 bytes (the maximum load width), meaning the GPU
430    /// can load as much data as possible in a single instruction. It may be aligned even more to
431    /// also take cache lines into account.
432    ///
433    /// However, the stride must be taken into account when indexing and reading the tensor
434    /// (also see [`ComputeClient::read_tensor`]).
435    pub fn create_tensor(&self, bytes: Bytes, shape: Shape, elem_size: usize) -> MemoryLayout {
436        self.do_create(
437            vec![MemoryLayoutDescriptor::new(
438                MemoryLayoutStrategy::Optimized,
439                shape,
440                elem_size,
441            )],
442            vec![bytes],
443        )
444        .unwrap()
445        .remove(0)
446    }
447
448    /// Reserves all `shapes` in a single storage buffer, copies the corresponding `data` into each
449    /// handle, and returns the handles for them.
450    /// See [`ComputeClient::create_tensor`]
451    ///
452    /// # Notes
453    ///
454    /// Prefer using [`Self::create_tensors`] for better performance.
455    pub fn create_tensors_from_slices(
456        &self,
457        descriptors: Vec<(MemoryLayoutDescriptor, &[u8])>,
458    ) -> Vec<MemoryLayout> {
459        let mut data = Vec::with_capacity(descriptors.len());
460        let mut descriptors_ = Vec::with_capacity(descriptors.len());
461        for (a, b) in descriptors {
462            data.push(b.to_vec());
463            descriptors_.push(a);
464        }
465
466        self.do_create_from_slices(descriptors_, data).unwrap()
467    }
468
469    /// Reserves all `shapes` in a single storage buffer, copies the corresponding `data` into each
470    /// handle, and returns the handles for them.
471    /// See [`ComputeClient::create_tensor`]
472    pub fn create_tensors(
473        &self,
474        descriptors: Vec<(MemoryLayoutDescriptor, Bytes)>,
475    ) -> Vec<MemoryLayout> {
476        let (descriptors, data) = descriptors.into_iter().unzip();
477
478        self.do_create(descriptors, data).unwrap()
479    }
480
481    fn do_empty(
482        &self,
483        descriptors: Vec<MemoryLayoutDescriptor>,
484    ) -> Result<Vec<MemoryLayout>, IoError> {
485        let stream_id = self.stream_id();
486        let (handle_base, layouts) = self.utilities.layout_policy.apply(stream_id, &descriptors);
487
488        let (size, memory) = (handle_base.size(), handle_base.memory);
489        self.device.submit(move |server| {
490            server.initialize_memory(memory, size, stream_id);
491        });
492
493        Ok(layouts)
494    }
495
496    /// Reserves `size` bytes in the storage, and returns a handle over them.
497    pub fn empty(&self, size: usize) -> Handle {
498        let shape: Shape = [size].into();
499        let descriptor = MemoryLayoutDescriptor::new(MemoryLayoutStrategy::Contiguous, shape, 1);
500        self.do_empty(vec![descriptor]).unwrap().remove(0).memory
501    }
502
503    /// Reserves `shape` in the storage, and returns a tensor handle for it.
504    /// See [`ComputeClient::create_tensor`]
505    pub fn empty_tensor(&self, shape: Shape, elem_size: usize) -> MemoryLayout {
506        let descriptor =
507            MemoryLayoutDescriptor::new(MemoryLayoutStrategy::Optimized, shape, elem_size);
508        self.do_empty(vec![descriptor]).unwrap().remove(0)
509    }
510
511    /// Reserves all `shapes` in a single storage buffer, and returns the handles for them.
512    /// See [`ComputeClient::create_tensor`]
513    pub fn empty_tensors(&self, descriptors: Vec<MemoryLayoutDescriptor>) -> Vec<MemoryLayout> {
514        self.do_empty(descriptors).unwrap()
515    }
516
517    /// Marks the given [Bytes] as being a staging buffer, maybe transferring it to pinned memory
518    /// for faster data transfer with compute device.
519    ///
520    /// TODO: This blocks the compute queue, so it will drop the compute utilization.
521    pub fn staging<'a, I>(&self, bytes: I, file_only: bool)
522    where
523        I: Iterator<Item = &'a mut Bytes>,
524    {
525        let has_staging = |b: &Bytes| match b.property() {
526            AllocationProperty::Pinned => false,
527            AllocationProperty::File => true,
528            AllocationProperty::Native | AllocationProperty::Other => !file_only,
529        };
530
531        let mut to_be_updated = Vec::new();
532        let sizes = bytes
533            .filter_map(|b| match has_staging(b) {
534                true => {
535                    let len = b.len();
536                    to_be_updated.push(b);
537                    Some(len)
538                }
539                false => None,
540            })
541            .collect::<Vec<usize>>();
542
543        if sizes.is_empty() {
544            return;
545        }
546
547        let stream_id = self.stream_id();
548        let sizes = sizes.to_vec();
549        let stagings = self
550            .device
551            .submit_blocking(move |server| server.staging(&sizes, stream_id))
552            .unwrap();
553
554        let stagings = match stagings {
555            Ok(val) => val,
556            Err(_) => return,
557        };
558
559        to_be_updated
560            .into_iter()
561            .zip(stagings)
562            .for_each(|(b, mut staging)| {
563                b.copy_into(&mut staging);
564                core::mem::swap(b, &mut staging);
565            });
566    }
567
568    /// Transfer data from one client to another
569    #[cfg_attr(
570        feature = "runtime-tracing",
571        tracing::instrument(level = "trace", skip(self, src, dst_server))
572    )]
573    pub fn to_client(&mut self, src: Handle, dst_server: &Self, dtype: ElemType) -> Handle {
574        let shape = [src.size_in_used() as usize];
575        let src_descriptor = src.copy_descriptor(shape.into(), [1].into(), 1);
576
577        if R::Server::SERVER_COMM_ENABLED {
578            self.to_client_tensor(src_descriptor, dst_server, dtype)
579        } else {
580            let alloc_desc = MemoryLayoutDescriptor::new(
581                MemoryLayoutStrategy::Contiguous,
582                src_descriptor.shape.clone(),
583                src_descriptor.elem_size,
584            );
585            self.change_client_sync(src_descriptor, alloc_desc, dst_server)
586                .memory
587        }
588    }
589
590    /// Perform an `all_reduce` operation on the given devices.
591    #[cfg_attr(
592        feature = "runtime-tracing",
593        tracing::instrument(level = "trace", skip(self, device_ids))
594    )]
595    pub fn ensure_init_collective(&mut self, device_ids: Vec<DeviceId>) {
596        let comm_id = CommunicationId::from(device_ids.clone());
597        let is_comms_init = self
598            .utilities
599            .initialized_comms
600            .read()
601            .unwrap()
602            .contains(&comm_id);
603        if !is_comms_init {
604            self.device
605                .submit(move |server| server.comm_init(device_ids).unwrap());
606            let mut initialized_comms = self.utilities.initialized_comms.write().unwrap();
607            initialized_comms.insert(comm_id);
608            // Flush immediately so other devices aren't blocked waiting on this initialization.
609            self.device.flush_queue();
610        }
611    }
612
613    /// Wait on the communication stream.
614    #[cfg_attr(feature = "runtime-tracing", tracing::instrument(level = "trace", skip(self)))]
615    pub fn sync_collective(&self) {
616        if DeviceHandle::<R::Server>::is_blocking() {
617            panic!("Can't use `sync_collective` with a blocking device handle");
618        }
619        let stream_id = self.stream_id();
620
621        self.device.submit(move |server| {
622            server.sync_collective(stream_id).unwrap();
623        });
624
625        // We don't actually need or want to sync the server here, but we need to make sure any
626        // task enqueued on the communication channel is done.
627        self.device.flush_queue();
628    }
629
630    /// Perform an `all_reduce` operation on the given devices.
631    #[cfg_attr(
632        feature = "runtime-tracing",
633        tracing::instrument(level = "trace", skip(self, src, dst, dtype, device_ids, op))
634    )]
635    pub fn all_reduce(
636        &mut self,
637        src: Handle,
638        dst: Handle,
639        dtype: ElemType,
640        device_ids: Vec<DeviceId>,
641        op: ReduceOperation,
642    ) {
643        if DeviceHandle::<R::Server>::is_blocking() {
644            panic!("Can't use `all_reduce` with a blocking device handle");
645        }
646
647        let stream_id = self.stream_id();
648        let src = src.binding();
649        let dst = dst.binding();
650
651        self.ensure_init_collective(device_ids.clone());
652
653        self.device.submit(move |server| {
654            server
655                .all_reduce(src, dst, dtype, stream_id, op, device_ids)
656                .unwrap();
657        });
658    }
659
660    /// Transfer data from one client to another
661    ///
662    /// Make sure the source description can be read in a contiguous manner.
663    #[cfg_attr(
664        feature = "runtime-tracing",
665        tracing::instrument(level = "trace", skip(self, src_descriptor, dst_server))
666    )]
667    pub fn to_client_tensor(
668        &mut self,
669        src_descriptor: CopyDescriptor,
670        dst_server: &Self,
671        dtype: ElemType,
672    ) -> Handle {
673        let stream_id_src = self.stream_id();
674        let stream_id_dst = dst_server.stream_id();
675
676        let device_id_src = self.device.device_id();
677        let device_id_dst = dst_server.device.device_id();
678
679        let mut dst_server = dst_server.clone();
680        let handle = Handle::new(stream_id_dst, src_descriptor.handle.size_in_used());
681        let handle_cloned = handle.clone();
682
683        let device_ids = vec![device_id_src, device_id_dst];
684        self.ensure_init_collective(device_ids.clone());
685        dst_server.ensure_init_collective(device_ids);
686
687        self.device.submit(move |server_src| {
688            server_src
689                .send(src_descriptor, dtype, stream_id_src, device_id_dst)
690                .unwrap()
691        });
692
693        dst_server.device.submit(move |server_dst| {
694            server_dst
695                .recv(handle_cloned, dtype, stream_id_dst, device_id_src)
696                .unwrap();
697            server_dst.sync_collective(stream_id_dst).unwrap();
698        });
699
700        // `ServerCommunication::send` and`ServerCommunication::recv` are blocking: they each wait for the corresponding recv/send
701        // call to be made. We flush the operations right away so that the neither server ends up in a deadlock.
702        // The actual data transfer is still executed asynchronously on the communication stream.
703        self.device.flush_queue();
704        dst_server.device.flush_queue();
705
706        handle
707    }
708
709    #[track_caller]
710    #[cfg_attr(feature = "runtime-tracing", tracing::instrument(level="trace",
711        skip(self, kernel, bindings),
712        fields(
713            kernel.name = %kernel.name(),
714            kernel.id = %kernel.id(),
715        )
716    ))]
717    unsafe fn launch_inner(
718        &self,
719        kernel: <R::Server as ComputeServer>::Kernel,
720        count: RudaCount,
721        bindings: KernelArguments,
722        mode: ExecutionMode,
723        stream_id: StreamId,
724    ) {
725        let level = self.utilities.logger.profile_level();
726
727        match level {
728            None | Some(ProfileLevel::ExecutionOnly) => {
729                let utilities = self.utilities.clone();
730                self.device.submit(move |state| {
731                    let name = kernel.name();
732                    unsafe { state.launch(kernel, count, bindings, mode, stream_id) };
733
734                    if matches!(level, Some(ProfileLevel::ExecutionOnly)) {
735                        let info = type_name_format(name, TypeNameFormatLevel::Balanced);
736                        utilities.logger.register_execution(info);
737                    }
738                });
739            }
740            Some(level) => {
741                let name = kernel.name();
742                let kernel_id = kernel.id();
743                let context = self.device.clone();
744                let count_moved = count.clone();
745                let (result, profile) = self
746                    .profile(
747                        move || {
748                            context
749                                .submit_blocking(move |state| unsafe {
750                                    state.launch(kernel, count_moved, bindings, mode, stream_id)
751                                })
752                                .unwrap()
753                        },
754                        name,
755                    )
756                    .unwrap();
757                let info = match level {
758                    ProfileLevel::Full => {
759                        format!("{name}: {kernel_id} RudaCount {count:?}")
760                    }
761                    _ => type_name_format(name, TypeNameFormatLevel::Balanced),
762                };
763                self.utilities.logger.register_profiled(info, profile);
764                result
765            }
766        }
767    }
768
769    /// Launches the `kernel` with the given `bindings`.
770    #[track_caller]
771    pub fn launch(
772        &self,
773        kernel: <R::Server as ComputeServer>::Kernel,
774        count: RudaCount,
775        bindings: KernelArguments,
776    ) {
777        // SAFETY: Using checked execution mode.
778        unsafe {
779            self.launch_inner(
780                kernel,
781                count,
782                bindings,
783                ExecutionMode::Checked,
784                self.stream_id(),
785            )
786        }
787    }
788
789    /// Launches the `kernel` with the given `bindings` without performing any bound checks.
790    ///
791    /// # Safety
792    ///
793    /// To ensure this is safe, you must verify your kernel:
794    /// - Has no out-of-bound reads and writes that can happen.
795    /// - Has no infinite loops that might never terminate.
796    #[track_caller]
797    pub unsafe fn launch_unchecked(
798        &self,
799        kernel: <R::Server as ComputeServer>::Kernel,
800        count: RudaCount,
801        bindings: KernelArguments,
802    ) {
803        // SAFETY: Caller has to uphold kernel being safe.
804        unsafe {
805            self.launch_inner(
806                kernel,
807                count,
808                bindings,
809                match self.utilities.check_mode {
810                    crate::runtime::config::compilation::BoundsCheckMode::Enforce => ExecutionMode::Checked,
811                    crate::runtime::config::compilation::BoundsCheckMode::Validate => {
812                        ExecutionMode::Validate
813                    }
814                    crate::runtime::config::compilation::BoundsCheckMode::Auto => ExecutionMode::Unchecked,
815                },
816                self.stream_id(),
817            )
818        }
819    }
820
821    /// Flush all outstanding commands.
822    pub fn flush(&self) -> Result<(), ServerError> {
823        let stream_id = self.stream_id();
824
825        self.device
826            .submit_blocking(move |server| server.flush(stream_id))
827            .unwrap()
828    }
829
830    /// Wait for the completion of every task in the server.
831    pub fn sync(&self) -> DynFut<Result<(), ServerError>> {
832        let stream_id = self.stream_id();
833
834        let fut = self
835            .device
836            .submit_blocking(move |server| server.sync(stream_id))
837            .unwrap();
838
839        self.utilities.logger.profile_summary();
840
841        fut
842    }
843
844    /// Get the features supported by the compute server.
845    pub fn properties(&self) -> &DeviceProperties {
846        &self.utilities.properties
847    }
848
849    /// Get the features supported by the compute server.
850    pub fn features(&self) -> &Features {
851        &self.utilities.properties.features
852    }
853
854    /// # Warning
855    ///
856    /// For private use only.
857    pub fn properties_mut(&mut self) -> Option<&mut DeviceProperties> {
858        Arc::get_mut(&mut self.utilities).map(|state| &mut state.properties)
859    }
860
861    /// Get the current memory usage of this client.
862    pub fn memory_usage(&self) -> Result<MemoryUsage, ServerError> {
863        let stream_id = self.stream_id();
864        self.device
865            .submit_blocking(move |server| server.memory_usage(stream_id))
866            .unwrap()
867    }
868
869    /// Get all devices of a specific type available to this runtime
870    pub fn enumerate_devices(&self, type_id: u16) -> Vec<DeviceId> {
871        R::enumerate_devices(type_id, self.info())
872    }
873
874    /// Get all devices available to this runtime
875    pub fn enumerate_all_devices(&self) -> Vec<DeviceId> {
876        R::enumerate_all_devices(self.info())
877    }
878
879    /// Get the number of devices of a specific type available to this runtime
880    pub fn device_count(&self, type_id: u16) -> usize {
881        self.enumerate_devices(type_id).len()
882    }
883
884    /// Get the number of devices of a specific type available to this runtime
885    pub fn device_count_total(&self) -> usize {
886        self.enumerate_all_devices().len()
887    }
888
889    /// Change the memory allocation mode.
890    ///
891    /// # Safety
892    ///
893    /// This function isn't thread safe and might create memory leaks.
894    pub unsafe fn allocation_mode(&self, mode: MemoryAllocationMode) {
895        let stream_id = self.stream_id();
896        self.device
897            .submit(move |server| server.allocation_mode(mode, stream_id));
898    }
899
900    /// Ask the client to release memory that it can release.
901    ///
902    /// Nb: Results will vary on what the memory allocator deems beneficial,
903    /// so it's not guaranteed any memory is freed.
904    pub fn memory_cleanup(&self) {
905        let stream_id = self.stream_id();
906        self.device
907            .submit(move |server| server.memory_cleanup(stream_id));
908    }
909
910    /// Measure the execution time of some inner operations.
911    #[track_caller]
912    pub fn profile<O: Send + 'static>(
913        &self,
914        func: impl FnOnce() -> O + Send,
915        #[allow(unused)] func_name: &str,
916    ) -> Result<(O, ProfileDuration), ProfileError> {
917        // Get the outer caller. For execute() this points straight to the
918        // ruda kernel. For general profiling it points to whoever calls profile.
919        #[cfg(feature = "runtime-profile-tracy")]
920        let location = std::panic::Location::caller();
921
922        // Make a CPU span. If the server has system profiling this is all you need.
923        #[cfg(feature = "runtime-profile-tracy")]
924        let _span = tracy_client::Client::running().unwrap().span_alloc(
925            None,
926            func_name,
927            location.file(),
928            location.line(),
929            0,
930        );
931
932        let stream_id = self.stream_id();
933
934        #[cfg(feature = "runtime-profile-tracy")]
935        let gpu_span = if self.utilities.properties.timing_method == TimingMethod::Device {
936            let gpu_span = self
937                .utilities
938                .gpu_client
939                .span_alloc(func_name, "profile", location.file(), location.line())
940                .unwrap();
941            Some(gpu_span)
942        } else {
943            None
944        };
945
946        let device = self.device.clone();
947        #[allow(unused_mut, reason = "Used in profile-tracy")]
948        let mut result = self
949            .device
950            .exclusive(move || {
951                // We first get mut access to the server to create a token.
952                // Then we free to server, since it's going to be accessed in `func()`.
953                let token =
954                    match device.submit_blocking(move |server| server.start_profile(stream_id)) {
955                        Ok(token) => match token {
956                            Ok(token) => token,
957                            Err(err) => return Err(err),
958                        },
959                        Err(err) => {
960                            return Err(ServerError::Generic {
961                                reason: alloc::format!(
962                                    "Can't start profiling because of a call error: {err:?}"
963                                ),
964                                backtrace: BackTrace::capture(),
965                            });
966                        }
967                    };
968
969                // We execute `func()` which will recursibly access the server.
970                let out = func();
971
972                // Finally we get the result from the token.
973                let result = device
974                    .submit_blocking(move |server| {
975                        let mut result = server.end_profile(stream_id, token);
976
977                        match result {
978                            Ok(result) => Ok((out, result)),
979                            Err(err) => Err(err),
980                        }
981                    })
982                    .unwrap();
983
984                Ok(result)
985            })
986            .unwrap()
987            .map_err(|err| ProfileError::Unknown {
988                reason: alloc::format!("{err:?}"),
989                backtrace: BackTrace::capture(),
990            })?;
991
992        #[cfg(feature = "runtime-profile-tracy")]
993        if let Some(mut gpu_span) = gpu_span {
994            gpu_span.end_zone();
995            let epoch = self.utilities.epoch_time;
996            // Add in the work to upload the timestamp data.
997            result = result.map(|(o, result)| {
998                (
999                    o,
1000                    ProfileDuration::new(
1001                        alloc::boxed::Box::pin(async move {
1002                            let ticks = result.resolve().await;
1003                            let start_duration =
1004                                ticks.start_duration_since(epoch).as_nanos() as i64;
1005                            let end_duration = ticks.end_duration_since(epoch).as_nanos() as i64;
1006                            gpu_span.upload_timestamp_start(start_duration);
1007                            gpu_span.upload_timestamp_end(end_duration);
1008                            ticks
1009                        }),
1010                        TimingMethod::Device,
1011                    ),
1012                )
1013            });
1014        }
1015
1016        result
1017    }
1018
1019    /// Transfer data from one client to another
1020    #[cfg_attr(
1021        feature = "runtime-tracing",
1022        tracing::instrument(
1023            level = "trace",
1024            skip(self, src_descriptor, alloc_descriptor, dst_server)
1025        )
1026    )]
1027    fn change_client_sync(
1028        &self,
1029        src_descriptor: CopyDescriptor,
1030        alloc_descriptor: MemoryLayoutDescriptor,
1031        dst_server: &Self,
1032    ) -> MemoryLayout {
1033        let shape = src_descriptor.shape.clone();
1034        let elem_size = src_descriptor.elem_size;
1035        let stream_id = self.stream_id();
1036
1037        let read = self
1038            .device
1039            .submit_blocking(move |server| server.read(vec![src_descriptor], stream_id))
1040            .unwrap();
1041
1042        let mut data = ruda_core::future::block_on(read).unwrap();
1043
1044        let (handle_base, mut layouts) = self
1045            .utilities
1046            .layout_policy
1047            .apply(stream_id, &[alloc_descriptor]);
1048        let alloc = layouts.remove(0);
1049
1050        let desc_descriptor = CopyDescriptor {
1051            handle: handle_base.clone().binding(),
1052            shape,
1053            strides: alloc.strides.clone(),
1054            elem_size,
1055        };
1056
1057        let (size, memory) = (handle_base.size(), handle_base.memory);
1058        dst_server.device.submit(move |server| {
1059            server.initialize_memory(memory, size, stream_id);
1060            server.write(vec![(desc_descriptor, data.remove(0))], stream_id)
1061        });
1062
1063        alloc
1064    }
1065
1066    /// Returns all vector sizes that are useful to perform optimal IO operation on the given element.
1067    pub fn io_optimized_vector_sizes(
1068        &self,
1069        size: usize,
1070    ) -> impl Iterator<Item = VectorSize> + Clone {
1071        let load_width = self.properties().hardware.load_width as usize;
1072        let size_bits = size * 8;
1073        let max = load_width / size_bits;
1074        // Scalar IO is still the only valid choice when one element is wider
1075        // than the native load width. Leaving `max` at zero makes
1076        // `trailing_zeros() + 1` enumerate through the machine word size and
1077        // eventually overflow while constructing `2^i`.
1078        let max = usize::min(self.properties().hardware.max_vector_size, max).max(1);
1079
1080        // If the max is 8, we want to test 1, 2, 4, 8 which is log2(8) + 1.
1081        let num_candidates = max.trailing_zeros() + 1;
1082
1083        (0..num_candidates).map(|i| 2usize.pow(i)).rev()
1084    }
1085}