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
31pub 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 pub fn info(&self) -> &<R::Server as ComputeServer>::Info {
52 &self.utilities.info
53 }
54
55 pub fn device_id(&self) -> DeviceId { self.device.device_id() }
57
58 pub fn properties_fingerprint(&self) -> u64 { self.utilities.properties_hash }
60
61 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 pub fn load<D: Device>(device: &D) -> Self {
76 let context = DeviceHandle::<R::Server>::new(device.to_id());
77
78 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 pub fn execution_stream(&self) -> StreamId { self.stream_id() }
102
103 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 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 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 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 pub fn read(&self, handles: Vec<Handle>) -> Vec<Bytes> {
158 ruda_core::reader::read_sync(self.read_async(handles)).expect("TODO")
159 }
160
161 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 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 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 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 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 pub fn read_one_unchecked_tensor(&self, descriptor: CopyDescriptor) -> Bytes {
219 self.read_tensor(vec![descriptor]).remove(0)
220 }
221
222 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 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 pub fn exclusive<'a, Re: Send + 'static, F: FnOnce() -> Re + Send + 'a>(
331 &'a self,
332 task: F,
333 ) -> Result<Re, ServerError> {
334 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 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 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 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 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 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 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 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 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 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 pub fn empty_tensors(&self, descriptors: Vec<MemoryLayoutDescriptor>) -> Vec<MemoryLayout> {
514 self.do_empty(descriptors).unwrap()
515 }
516
517 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 #[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 #[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 self.device.flush_queue();
610 }
611 }
612
613 #[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 self.device.flush_queue();
628 }
629
630 #[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 #[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 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 #[track_caller]
771 pub fn launch(
772 &self,
773 kernel: <R::Server as ComputeServer>::Kernel,
774 count: RudaCount,
775 bindings: KernelArguments,
776 ) {
777 unsafe {
779 self.launch_inner(
780 kernel,
781 count,
782 bindings,
783 ExecutionMode::Checked,
784 self.stream_id(),
785 )
786 }
787 }
788
789 #[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 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 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 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 pub fn properties(&self) -> &DeviceProperties {
846 &self.utilities.properties
847 }
848
849 pub fn features(&self) -> &Features {
851 &self.utilities.properties.features
852 }
853
854 pub fn properties_mut(&mut self) -> Option<&mut DeviceProperties> {
858 Arc::get_mut(&mut self.utilities).map(|state| &mut state.properties)
859 }
860
861 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 pub fn enumerate_devices(&self, type_id: u16) -> Vec<DeviceId> {
871 R::enumerate_devices(type_id, self.info())
872 }
873
874 pub fn enumerate_all_devices(&self) -> Vec<DeviceId> {
876 R::enumerate_all_devices(self.info())
877 }
878
879 pub fn device_count(&self, type_id: u16) -> usize {
881 self.enumerate_devices(type_id).len()
882 }
883
884 pub fn device_count_total(&self) -> usize {
886 self.enumerate_all_devices().len()
887 }
888
889 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 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 #[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 #[cfg(feature = "runtime-profile-tracy")]
920 let location = std::panic::Location::caller();
921
922 #[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 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 let out = func();
971
972 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 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 #[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 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 let max = usize::min(self.properties().hardware.max_vector_size, max).max(1);
1079
1080 let num_candidates = max.trailing_zeros() + 1;
1082
1083 (0..num_candidates).map(|i| 2usize.pow(i)).rev()
1084 }
1085}