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 {
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 pub fn load<D: Device>(device: &D) -> Self {
81 let context = DeviceHandle::<R::Server>::new(device.to_id());
82
83 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 pub fn execution_stream(&self) -> StreamId { self.stream_id() }
107
108 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 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 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 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 pub fn read(&self, handles: Vec<Handle>) -> Vec<Bytes> {
161 ruda_core::reader::read_sync(self.read_async(handles)).expect("TODO")
162 }
163
164 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 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 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 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 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 pub fn read_one_unchecked_tensor(&self, descriptor: CopyDescriptor) -> Bytes {
225 self.read_tensor(vec![descriptor]).remove(0)
226 }
227
228 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 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 pub fn exclusive<'a, Re: Send + 'static, F: FnOnce() -> Re + Send + 'a>(
337 &'a self,
338 task: F,
339 ) -> Result<Re, ServerError> {
340 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 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 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 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 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 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 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 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 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 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 pub fn empty_tensors(&self, descriptors: Vec<MemoryLayoutDescriptor>) -> Vec<MemoryLayout> {
520 self.do_empty(descriptors).unwrap()
521 }
522
523 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 #[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 #[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 self.device.flush_queue();
616 }
617 }
618
619 #[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 self.device.flush_queue();
634 }
635
636 #[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 #[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 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 #[track_caller]
777 pub fn launch(
778 &self,
779 kernel: <R::Server as ComputeServer>::Kernel,
780 count: RudaCount,
781 bindings: KernelArguments,
782 ) {
783 unsafe {
785 self.launch_inner(
786 kernel,
787 count,
788 bindings,
789 ExecutionMode::Checked,
790 self.stream_id(),
791 )
792 }
793 }
794
795 #[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 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 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 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 pub fn properties(&self) -> &DeviceProperties {
856 &self.utilities.properties
857 }
858
859 pub fn features(&self) -> &Features {
861 &self.utilities.properties.features
862 }
863
864 pub fn properties_mut(&mut self) -> Option<&mut DeviceProperties> {
868 Arc::get_mut(&mut self.utilities).map(|state| &mut state.properties)
869 }
870
871 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 pub fn enumerate_devices(&self, type_id: u16) -> Vec<DeviceId> {
883 R::enumerate_devices(type_id, self.info())
884 }
885
886 pub fn enumerate_all_devices(&self) -> Vec<DeviceId> {
888 R::enumerate_all_devices(self.info())
889 }
890
891 pub fn device_count(&self, type_id: u16) -> usize {
893 self.enumerate_devices(type_id).len()
894 }
895
896 pub fn device_count_total(&self) -> usize {
898 self.enumerate_all_devices().len()
899 }
900
901 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 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 #[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 #[cfg(feature = "runtime-profile-tracy")]
932 let location = std::panic::Location::caller();
933
934 #[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 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 let out = func();
983
984 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 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 #[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 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 let max = usize::min(self.properties().hardware.max_vector_size, max).max(1);
1091
1092 let num_candidates = max.trailing_zeros() + 1;
1094
1095 (0..num_candidates).map(|i| 2usize.pow(i)).rev()
1096 }
1097}