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