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 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 pub fn read(&self, handles: Vec<Handle>) -> Vec<Bytes> {
163 ruda_core::reader::read_sync(self.read_async(handles)).expect("TODO")
164 }
165
166 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 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 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 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 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 pub fn read_one_unchecked_tensor(&self, descriptor: CopyDescriptor) -> Bytes {
227 self.read_tensor(vec![descriptor]).remove(0)
228 }
229
230 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 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 pub fn exclusive<'a, Re: Send + 'static, F: FnOnce() -> Re + Send + 'a>(
339 &'a self,
340 task: F,
341 ) -> Result<Re, ServerError> {
342 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 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 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 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 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 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 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 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 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 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 pub fn empty_tensors(&self, descriptors: Vec<MemoryLayoutDescriptor>) -> Vec<MemoryLayout> {
522 self.do_empty(descriptors).unwrap()
523 }
524
525 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 #[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 #[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 self.device.flush_queue();
618 }
619 }
620
621 #[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 self.device.flush_queue();
636 }
637
638 #[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 #[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 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 #[track_caller]
779 pub fn launch(
780 &self,
781 kernel: <R::Server as ComputeServer>::Kernel,
782 count: RudaCount,
783 bindings: KernelArguments,
784 ) {
785 unsafe {
787 self.launch_inner(
788 kernel,
789 count,
790 bindings,
791 ExecutionMode::Checked,
792 self.stream_id(),
793 )
794 }
795 }
796
797 #[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 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 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 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 pub fn properties(&self) -> &DeviceProperties {
858 &self.utilities.properties
859 }
860
861 pub fn features(&self) -> &Features {
863 &self.utilities.properties.features
864 }
865
866 pub fn properties_mut(&mut self) -> Option<&mut DeviceProperties> {
870 Arc::get_mut(&mut self.utilities).map(|state| &mut state.properties)
871 }
872
873 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 pub fn enumerate_devices(&self, type_id: u16) -> Vec<DeviceId> {
885 R::enumerate_devices(type_id, self.info())
886 }
887
888 pub fn enumerate_all_devices(&self) -> Vec<DeviceId> {
890 R::enumerate_all_devices(self.info())
891 }
892
893 pub fn device_count(&self, type_id: u16) -> usize {
895 self.enumerate_devices(type_id).len()
896 }
897
898 pub fn device_count_total(&self) -> usize {
900 self.enumerate_all_devices().len()
901 }
902
903 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 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 #[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 #[cfg(feature = "runtime-profile-tracy")]
934 let location = std::panic::Location::caller();
935
936 #[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 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 let out = func();
985
986 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 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 #[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 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 let max = usize::min(self.properties().hardware.max_vector_size, max).max(1);
1093
1094 let num_candidates = max.trailing_zeros() + 1;
1096
1097 (0..num_candidates).map(|i| 2usize.pow(i)).rev()
1098 }
1099}