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