1use crate::{
2 config::memory::MemoryPoolsConfig,
3 config::{TypeNameFormatLevel, type_name_format},
4 id::GraphId,
5 kernel::KernelMetadata,
6 logging::ProfileLevel,
7 memory_management::{
8 InstallMemoryPoolsError, MemoryAllocationMode, MemoryConfiguration, MemoryReport,
9 MemoryUsage,
10 },
11 runtime::Runtime,
12 server::{
13 CommunicationId, ComputeServer, CopyDescriptor, CubeCount, Handle, IoError,
14 KernelArguments, MemoryLayout, MemoryLayoutDescriptor, MemoryLayoutPolicy,
15 MemoryLayoutStrategy, ProfileError, ReduceOperation, ServerCommunication, ServerError,
16 ServerUtilities,
17 },
18 storage::{ComputeStorage, ManagedResource},
19 throughput::{
20 KernelConfig, ThroughputBenchmarker, ThroughputCache, ThroughputKey, ThroughputValue,
21 },
22};
23use alloc::{format, string::String, sync::Arc, vec, vec::Vec};
24
25#[cfg(not(target_family = "wasm"))]
26mod lazy;
27use cubecl_common::{
28 bytes::{AllocationProperty, Bytes},
29 device::{Device, DeviceId},
30 device_handle::{CallResultExt, DeviceHandle},
31 profile::ProfileDuration,
32};
33use cubecl_environment::backtrace::BackTrace;
34use cubecl_environment::future::DynFut;
35use cubecl_ir::{DeviceProperties, ElemType, VectorSize, features::Features};
36use cubecl_zspace::Shape;
37
38#[allow(unused)]
39use cubecl_common::profile::TimingMethod;
40use cubecl_environment::stream::StreamId;
41
42pub struct ComputeClient<R: Runtime> {
45 device: DeviceHandle<R::Server>,
46 utilities: Arc<ServerUtilities<R::Server>>,
47 stream_id: Option<StreamId>,
48}
49
50pub struct Graph<R: Runtime> {
72 inner: Arc<GraphHandle<R>>,
73}
74
75impl<R: Runtime> core::fmt::Debug for Graph<R> {
76 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
77 f.debug_struct("Graph")
78 .field("id", &self.inner.id)
79 .field("stream_id", &self.inner.stream_id)
80 .finish()
81 }
82}
83
84struct GraphHandle<R: Runtime> {
88 id: GraphId,
89 device: DeviceHandle<R::Server>,
90 stream_id: StreamId,
91}
92
93impl<R: Runtime> Graph<R> {
94 pub unsafe fn replay(&self) {
130 let id = self.inner.id;
131 let stream_id = self.inner.stream_id;
132 self.inner
133 .device
134 .submit(move |server| server.replay(id, stream_id));
135 }
136}
137
138impl<R: Runtime> Clone for Graph<R> {
139 fn clone(&self) -> Self {
140 Self {
141 inner: self.inner.clone(),
142 }
143 }
144}
145
146impl<R: Runtime> Drop for GraphHandle<R> {
147 fn drop(&mut self) {
148 let id = self.id;
149 let stream_id = self.stream_id;
150 self.device
155 .submit(move |server| server.graph_destroy(id, stream_id));
156 }
157}
158
159impl<R: Runtime> Clone for ComputeClient<R> {
160 fn clone(&self) -> Self {
161 Self {
162 device: self.device.clone(),
163 utilities: self.utilities.clone(),
164 stream_id: self.stream_id,
165 }
166 }
167}
168
169impl<R: Runtime> ComputeClient<R> {
170 pub fn info(&self) -> &<R::Server as ComputeServer>::Info {
172 &self.utilities.info
173 }
174
175 pub fn init<D: Device>(device: &D, server: R::Server) -> Self {
177 let utilities = server.utilities();
178 let context = DeviceHandle::<R::Server>::insert(device.to_id(), server)
179 .expect("Can't create a new client on an already registered server");
180
181 Self {
182 device: context,
183 utilities,
184 stream_id: None,
185 }
186 }
187
188 pub fn load<D: Device>(device: &D) -> Self {
190 let context = DeviceHandle::<R::Server>::new(device.to_id());
191
192 let utilities = context
194 .utilities()
195 .downcast::<ServerUtilities<R::Server>>()
196 .expect("Can downcast to `ServerUtilities`");
197
198 Self {
199 device: context,
200 utilities,
201 stream_id: None,
202 }
203 }
204
205 fn stream_id(&self) -> StreamId {
206 match self.stream_id {
207 Some(val) => val,
208 None => StreamId::current(),
209 }
210 }
211
212 pub unsafe fn set_stream(&mut self, stream_id: StreamId) {
218 self.stream_id = Some(stream_id);
219 }
220
221 fn do_read(&self, descriptors: Vec<CopyDescriptor>) -> DynFut<Result<Vec<Bytes>, ServerError>> {
222 let stream_id = self.stream_id();
223 self.device
224 .submit_blocking(move |server| server.read(descriptors, stream_id))
225 .unwrap_or_resume()
226 }
227
228 pub fn read_async(
230 &self,
231 handles: Vec<Handle>,
232 ) -> impl Future<Output = Result<Vec<Bytes>, ServerError>> + Send {
233 let shapes = handles
234 .iter()
235 .map(|it| [it.size_in_used() as usize].into())
236 .collect::<Vec<Shape>>();
237 let descriptors = handles
238 .into_iter()
239 .zip(shapes)
240 .map(|(handle, shape)| CopyDescriptor::new(handle.binding(), shape, [1].into(), 1))
241 .collect();
242
243 self.do_read(descriptors)
244 }
245
246 pub fn read(&self, handles: Vec<Handle>) -> Vec<Bytes> {
252 cubecl_environment::future::reader::read_sync(self.read_async(handles)).expect("TODO")
253 }
254
255 pub fn read_one(&self, handle: Handle) -> Result<Bytes, ServerError> {
257 Ok(cubecl_environment::future::reader::read_sync(self.read_async(vec![handle]))?.remove(0))
258 }
259
260 pub fn read_one_unchecked(&self, handle: Handle) -> Bytes {
266 cubecl_environment::future::reader::read_sync(self.read_async(vec![handle]))
267 .unwrap()
268 .remove(0)
269 }
270
271 pub fn read_tensor_async(
273 &self,
274 descriptors: Vec<CopyDescriptor>,
275 ) -> impl Future<Output = Result<Vec<Bytes>, ServerError>> + Send {
276 self.do_read(descriptors)
277 }
278
279 pub fn read_tensor(&self, descriptors: Vec<CopyDescriptor>) -> Vec<Bytes> {
292 cubecl_environment::future::reader::read_sync(self.read_tensor_async(descriptors))
293 .expect("TODO")
294 }
295
296 pub fn read_one_tensor_async(
299 &self,
300 descriptor: CopyDescriptor,
301 ) -> impl Future<Output = Result<Bytes, ServerError>> + Send {
302 let fut = self.read_tensor_async(vec![descriptor]);
303
304 async { Ok(fut.await?.remove(0)) }
305 }
306
307 pub fn read_one_unchecked_tensor(&self, descriptor: CopyDescriptor) -> Bytes {
314 self.read_tensor(vec![descriptor]).remove(0)
315 }
316
317 #[cfg(not(target_family = "wasm"))]
327 pub fn read_lazy(&self, descriptor: CopyDescriptor) -> Bytes {
328 let len = descriptor.shape.iter().product::<usize>() * descriptor.elem_size;
329 let controller = lazy::LazyDeviceController::new(self.clone(), Arc::new(descriptor));
330 unsafe { Bytes::from_controller(alloc::boxed::Box::new(controller), len) }
332 }
333
334 #[cfg(not(target_family = "wasm"))]
339 pub fn read_lazy_async(
340 &self,
341 descriptor: CopyDescriptor,
342 ) -> impl Future<Output = Result<Bytes, ServerError>> + Send {
343 let len = descriptor.shape.iter().product::<usize>() * descriptor.elem_size;
344 let controller = lazy::LazyDeviceController::new(self.clone(), Arc::new(descriptor));
345 let bytes = unsafe { Bytes::from_controller(alloc::boxed::Box::new(controller), len) };
347 core::future::ready(Ok(bytes))
348 }
349
350 #[cfg(target_family = "wasm")]
356 pub fn read_lazy_async(
357 &self,
358 descriptor: CopyDescriptor,
359 ) -> impl Future<Output = Result<Bytes, ServerError>> + Send {
360 self.read_one_tensor_async(descriptor)
361 }
362
363 pub fn get_resource(
365 &self,
366 handle: Handle,
367 ) -> Result<
368 ManagedResource<<<R::Server as ComputeServer>::Storage as ComputeStorage>::Resource>,
369 ServerError,
370 > {
371 let stream_id = self.stream_id();
372 let binding = handle.binding();
373
374 self.device
375 .submit_blocking(move |state| state.get_resource(binding, stream_id))
376 .unwrap_or_resume()
377 }
378
379 fn do_create_from_slices(
380 &self,
381 descriptors: Vec<MemoryLayoutDescriptor>,
382 slices: Vec<Vec<u8>>,
383 ) -> Result<Vec<MemoryLayout>, IoError> {
384 let stream_id = self.stream_id();
385 let (handle_base, layouts) = self.utilities.layout_policy.apply(stream_id, &descriptors);
386
387 let descriptors = descriptors
388 .into_iter()
389 .zip(layouts.iter())
390 .zip(slices)
391 .map(|((desc, alloc), data)| {
392 (
393 CopyDescriptor::new(
394 alloc.memory.clone().binding(),
395 desc.shape,
396 alloc.strides.clone(),
397 desc.elem_size,
398 ),
399 Bytes::from_bytes_vec(data.to_vec()),
400 )
401 })
402 .collect::<Vec<_>>();
403
404 let (size, memory) = (handle_base.size(), handle_base.memory);
405 self.device.submit(move |server| {
406 server.initialize_memory(memory, size, stream_id);
407 server.write(descriptors, stream_id);
408 });
409
410 Ok(layouts)
411 }
412
413 fn do_create(
414 &self,
415 descriptors: Vec<MemoryLayoutDescriptor>,
416 data: Vec<Bytes>,
417 ) -> Result<Vec<MemoryLayout>, IoError> {
418 let stream_id = self.stream_id();
419 let (handle_base, layouts) = self.utilities.layout_policy.apply(stream_id, &descriptors);
420
421 let descriptors = descriptors
422 .into_iter()
423 .zip(layouts.iter())
424 .zip(data)
425 .map(|((desc, layout), data)| {
426 (
427 CopyDescriptor::new(
428 layout.memory.clone().binding(),
429 desc.shape,
430 layout.strides.clone(),
431 desc.elem_size,
432 ),
433 data,
434 )
435 })
436 .collect::<Vec<_>>();
437
438 let (size, memory) = (handle_base.size(), handle_base.memory);
439 self.device.submit(move |server| {
440 server.initialize_memory(memory, size, stream_id);
441 server.write(descriptors, stream_id);
442 });
443
444 Ok(layouts)
445 }
446
447 pub fn create_from_slice(&self, slice: &[u8]) -> Handle {
453 let shape: Shape = [slice.len()].into();
454
455 self.do_create_from_slices(
456 vec![MemoryLayoutDescriptor::new(
457 MemoryLayoutStrategy::Contiguous,
458 shape,
459 1,
460 )],
461 vec![slice.to_vec()],
462 )
463 .unwrap()
464 .remove(0)
465 .memory
466 }
467
468 pub fn exclusive<'a, Re: Send + 'static, F: FnOnce() -> Re + Send + 'a>(
470 &'a self,
471 task: F,
472 ) -> Result<Re, ServerError> {
473 self.device
475 .exclusive(task)
476 .map_err(|err| ServerError::Generic {
477 reason: format!("{err:?}"),
478 backtrace: BackTrace::capture(),
479 })
480 }
481
482 pub fn memory_persistent_allocation<
484 'a,
485 Re: Send,
486 Input: Send,
487 F: FnOnce(Input) -> Re + Send + 'a,
488 >(
489 &'a self,
490 input: Input,
491 task: F,
492 ) -> Result<Re, ServerError> {
493 let stream_id = StreamId::current();
494
495 self.device.submit(move |server| {
496 server.allocation_mode(MemoryAllocationMode::Persistent, stream_id);
497 });
498
499 let output = task(input);
501
502 self.device.submit(move |server| {
503 server.allocation_mode(MemoryAllocationMode::Auto, stream_id);
504 });
505
506 Ok(output)
507 }
508
509 pub fn write(&self, handle: &Handle, data: Bytes) {
519 let stream_id = self.stream_id();
520 let descriptor =
521 CopyDescriptor::new(handle.clone().binding(), [data.len()].into(), [1].into(), 1);
522 self.device.submit(move |server| {
523 server.write(vec![(descriptor, data)], stream_id);
524 });
525 }
526
527 pub fn create(&self, data: Bytes) -> Handle {
529 let shape = [data.len()].into();
530
531 self.do_create(
532 vec![MemoryLayoutDescriptor::new(
533 MemoryLayoutStrategy::Contiguous,
534 shape,
535 1,
536 )],
537 vec![data],
538 )
539 .unwrap()
540 .remove(0)
541 .memory
542 }
543
544 pub fn create_tensor_from_slice(
562 &self,
563 slice: &[u8],
564 shape: Shape,
565 elem_size: usize,
566 ) -> MemoryLayout {
567 self.do_create_from_slices(
568 vec![MemoryLayoutDescriptor::new(
569 MemoryLayoutStrategy::Optimized,
570 shape,
571 elem_size,
572 )],
573 vec![slice.to_vec()],
574 )
575 .unwrap()
576 .remove(0)
577 }
578
579 pub fn create_tensor(&self, bytes: Bytes, shape: Shape, elem_size: usize) -> MemoryLayout {
593 self.do_create(
594 vec![MemoryLayoutDescriptor::new(
595 MemoryLayoutStrategy::Optimized,
596 shape,
597 elem_size,
598 )],
599 vec![bytes],
600 )
601 .unwrap()
602 .remove(0)
603 }
604
605 pub fn create_tensors_from_slices(
613 &self,
614 descriptors: Vec<(MemoryLayoutDescriptor, &[u8])>,
615 ) -> Vec<MemoryLayout> {
616 let mut data = Vec::with_capacity(descriptors.len());
617 let mut descriptors_ = Vec::with_capacity(descriptors.len());
618 for (a, b) in descriptors {
619 data.push(b.to_vec());
620 descriptors_.push(a);
621 }
622
623 self.do_create_from_slices(descriptors_, data).unwrap()
624 }
625
626 pub fn create_tensors(
630 &self,
631 descriptors: Vec<(MemoryLayoutDescriptor, Bytes)>,
632 ) -> Vec<MemoryLayout> {
633 let (descriptors, data) = descriptors.into_iter().unzip();
634
635 self.do_create(descriptors, data).unwrap()
636 }
637
638 fn do_empty(
639 &self,
640 descriptors: Vec<MemoryLayoutDescriptor>,
641 ) -> Result<Vec<MemoryLayout>, IoError> {
642 let stream_id = self.stream_id();
643 let (handle_base, layouts) = self.utilities.layout_policy.apply(stream_id, &descriptors);
644
645 let (size, memory) = (handle_base.size(), handle_base.memory);
646 self.device.submit(move |server| {
647 server.initialize_memory(memory, size, stream_id);
648 });
649
650 Ok(layouts)
651 }
652
653 pub fn empty(&self, size: usize) -> Handle {
655 let shape: Shape = [size].into();
656 let descriptor = MemoryLayoutDescriptor::new(MemoryLayoutStrategy::Contiguous, shape, 1);
657 self.do_empty(vec![descriptor]).unwrap().remove(0).memory
658 }
659
660 pub fn empty_tensor(&self, shape: Shape, elem_size: usize) -> MemoryLayout {
663 let descriptor =
664 MemoryLayoutDescriptor::new(MemoryLayoutStrategy::Optimized, shape, elem_size);
665 self.do_empty(vec![descriptor]).unwrap().remove(0)
666 }
667
668 pub fn empty_tensors(&self, descriptors: Vec<MemoryLayoutDescriptor>) -> Vec<MemoryLayout> {
671 self.do_empty(descriptors).unwrap()
672 }
673
674 pub fn staging<'a, I>(&self, bytes: I, file_only: bool)
679 where
680 I: Iterator<Item = &'a mut Bytes>,
681 {
682 let has_staging = |b: &Bytes| match b.property() {
683 AllocationProperty::Pinned => false,
684 AllocationProperty::File => true,
685 AllocationProperty::Device => false,
688 AllocationProperty::Native | AllocationProperty::Other => !file_only,
689 };
690
691 let mut to_be_updated = Vec::new();
692 let sizes = bytes
693 .filter_map(|b| match has_staging(b) {
694 true => {
695 let len = b.len();
696 to_be_updated.push(b);
697 Some(len)
698 }
699 false => None,
700 })
701 .collect::<Vec<usize>>();
702
703 if sizes.is_empty() {
704 return;
705 }
706
707 let stream_id = self.stream_id();
708 let sizes = sizes.to_vec();
709 let stagings = self
710 .device
711 .submit_blocking(move |server| server.staging(&sizes, stream_id))
712 .unwrap_or_resume();
713
714 let stagings = match stagings {
715 Ok(val) => val,
716 Err(_) => return,
717 };
718
719 to_be_updated
720 .into_iter()
721 .zip(stagings)
722 .for_each(|(b, mut staging)| {
723 b.copy_into(&mut staging);
724 core::mem::swap(b, &mut staging);
725 });
726 }
727
728 #[cfg_attr(
730 feature = "tracing",
731 tracing::instrument(level = "trace", skip(self, src, dst_server))
732 )]
733 pub fn to_client(&mut self, src: Handle, dst_server: &Self, dtype: ElemType) -> Handle {
734 let shape = [src.size_in_used() as usize];
735 let src_descriptor = src.copy_descriptor(shape.into(), [1].into(), 1);
736
737 if R::Server::SERVER_COMM_ENABLED {
738 self.to_client_tensor(src_descriptor, dst_server, dtype)
739 } else {
740 let alloc_desc = MemoryLayoutDescriptor::new(
741 MemoryLayoutStrategy::Contiguous,
742 src_descriptor.shape.clone(),
743 src_descriptor.elem_size,
744 );
745 self.change_client_sync(src_descriptor, alloc_desc, dst_server)
746 .memory
747 }
748 }
749
750 #[cfg_attr(
752 feature = "tracing",
753 tracing::instrument(level = "trace", skip(self, device_ids))
754 )]
755 pub fn ensure_init_collective(&mut self, device_ids: Vec<DeviceId>) {
756 let comm_id = CommunicationId::from(device_ids.clone());
757 let is_comms_init = self.utilities.initialized_comms.read().contains(&comm_id);
758 if !is_comms_init {
759 self.device
760 .submit(move |server| server.comm_init(device_ids).unwrap());
761 let mut initialized_comms = self.utilities.initialized_comms.write();
762 initialized_comms.insert(comm_id);
763 self.device.flush_queue();
765 }
766 }
767
768 #[cfg_attr(feature = "tracing", tracing::instrument(level = "trace", skip(self)))]
770 pub fn sync_collective(&self) {
771 if DeviceHandle::<R::Server>::is_blocking() {
772 panic!("Can't use `sync_collective` with a blocking device handle");
773 }
774 let stream_id = self.stream_id();
775
776 self.device.submit(move |server| {
777 server.sync_collective(stream_id).unwrap();
778 });
779
780 self.device.flush_queue();
783 }
784
785 #[cfg_attr(
787 feature = "tracing",
788 tracing::instrument(level = "trace", skip(self, src, dst, dtype, device_ids, op))
789 )]
790 pub fn all_reduce(
791 &mut self,
792 src: Handle,
793 dst: Handle,
794 dtype: ElemType,
795 device_ids: Vec<DeviceId>,
796 op: ReduceOperation,
797 ) {
798 if DeviceHandle::<R::Server>::is_blocking() {
799 panic!("Can't use `all_reduce` with a blocking device handle");
800 }
801
802 let stream_id = self.stream_id();
803 let src = src.binding();
804 let dst = dst.binding();
805
806 self.ensure_init_collective(device_ids.clone());
807
808 self.device.submit(move |server| {
809 server
810 .all_reduce(src, dst, dtype, stream_id, op, device_ids)
811 .unwrap();
812 });
813 }
814
815 #[cfg_attr(
819 feature = "tracing",
820 tracing::instrument(level = "trace", skip(self, src_descriptor, dst_server))
821 )]
822 pub fn to_client_tensor(
823 &mut self,
824 src_descriptor: CopyDescriptor,
825 dst_server: &Self,
826 dtype: ElemType,
827 ) -> Handle {
828 let stream_id_src = self.stream_id();
829 let stream_id_dst = dst_server.stream_id();
830
831 let device_id_src = self.device.device_id();
832 let device_id_dst = dst_server.device.device_id();
833
834 let mut dst_server = dst_server.clone();
835 let handle = Handle::new(stream_id_dst, src_descriptor.handle.size_in_used());
836 let handle_cloned = handle.clone();
837
838 let device_ids = vec![device_id_src, device_id_dst];
839 self.ensure_init_collective(device_ids.clone());
840 dst_server.ensure_init_collective(device_ids);
841
842 self.device.submit(move |server_src| {
843 server_src
844 .send(src_descriptor, dtype, stream_id_src, device_id_dst)
845 .unwrap()
846 });
847
848 dst_server.device.submit(move |server_dst| {
849 server_dst
850 .recv(handle_cloned, dtype, stream_id_dst, device_id_src)
851 .unwrap();
852 server_dst.sync_collective(stream_id_dst).unwrap();
853 });
854
855 self.device.flush_queue();
859 dst_server.device.flush_queue();
860
861 handle
862 }
863
864 #[track_caller]
865 #[cfg_attr(feature = "tracing", tracing::instrument(level="trace",
866 skip(self, kernel, bindings),
867 fields(
868 kernel.name = %kernel.name(),
869 kernel.id = %kernel.id(),
870 )
871 ))]
872 unsafe fn launch_inner(
873 &self,
874 kernel: <R::Server as ComputeServer>::Kernel,
875 count: CubeCount,
876 bindings: KernelArguments,
877 stream_id: StreamId,
878 ) {
879 if let CubeCount::Static(x, y, z) = &count
881 && (*x == 0 || *y == 0 || *z == 0)
882 {
883 return;
884 }
885
886 let launch_mode = crate::dry_run::launch_mode();
890
891 let level = self.utilities.logger.profile_level();
892
893 match level {
894 None | Some(ProfileLevel::ExecutionOnly) => {
895 let utilities = self.utilities.clone();
896 self.device.submit(move |state| {
897 let name = kernel.name();
898 unsafe { state.launch(kernel, count, bindings, stream_id, launch_mode) };
899
900 if matches!(level, Some(ProfileLevel::ExecutionOnly)) {
901 let info = type_name_format(name, TypeNameFormatLevel::Balanced);
902 utilities.logger.register_execution(info);
903 }
904 });
905 }
906 Some(level) => {
907 let name = kernel.name();
908 let kernel_id = kernel.id();
909 let context = self.device.clone();
910 let count_moved = count.clone();
911 let (result, profile) = self
912 .profile(
913 move || {
914 context
915 .submit_blocking(move |state| unsafe {
916 state.launch(
917 kernel,
918 count_moved,
919 bindings,
920 stream_id,
921 launch_mode,
922 )
923 })
924 .unwrap_or_resume()
925 },
926 name,
927 )
928 .unwrap();
929 let info = match level {
930 ProfileLevel::Full => {
931 format!("{name}: {kernel_id} CubeCount {count:?}")
932 }
933 _ => type_name_format(name, TypeNameFormatLevel::Balanced),
934 };
935 self.utilities.logger.register_profiled(info, profile);
936 result
937 }
938 }
939 }
940
941 #[track_caller]
943 pub fn launch(
944 &self,
945 kernel: <R::Server as ComputeServer>::Kernel,
946 count: CubeCount,
947 bindings: KernelArguments,
948 ) {
949 unsafe { self.launch_inner(kernel, count, bindings, self.stream_id()) }
950 }
951
952 pub fn flush(&self) -> Result<(), ServerError> {
954 let stream_id = self.stream_id();
955
956 self.device
957 .submit_blocking(move |server| server.flush(stream_id))
958 .unwrap_or_resume()
959 }
960
961 pub fn graph_prepare(&self) -> Result<(), ServerError> {
966 let stream_id = self.stream_id();
967 self.device
968 .submit_blocking(move |server| server.graph_prepare(stream_id))
969 .unwrap_or_resume()
970 }
971
972 pub fn start_capture(&self) -> Result<(), ServerError> {
988 let stream_id = self.stream_id();
989 self.device
990 .submit_blocking(move |server| server.begin_capture(stream_id))
991 .unwrap_or_resume()
992 }
993
994 pub fn stop_capture(&self) -> Result<Graph<R>, ServerError> {
997 let stream_id = self.stream_id();
998 let id = self
999 .device
1000 .submit_blocking(move |server| server.end_capture(stream_id))
1001 .unwrap_or_resume()?;
1002
1003 Ok(Graph {
1004 inner: Arc::new(GraphHandle {
1005 id,
1006 device: self.device.clone(),
1007 stream_id,
1008 }),
1009 })
1010 }
1011
1012 pub fn sync(&self) -> DynFut<Result<(), ServerError>> {
1014 let stream_id = self.stream_id();
1015
1016 let fut = self
1017 .device
1018 .submit_blocking(move |server| server.sync(stream_id))
1019 .unwrap_or_resume();
1020
1021 self.utilities.logger.profile_summary();
1022
1023 fut
1024 }
1025
1026 pub fn properties(&self) -> &DeviceProperties {
1028 &self.utilities.properties
1029 }
1030
1031 pub fn features(&self) -> &Features {
1033 &self.utilities.properties.features
1034 }
1035
1036 pub fn properties_mut(&mut self) -> Option<&mut DeviceProperties> {
1040 Arc::get_mut(&mut self.utilities).map(|state| &mut state.properties)
1041 }
1042
1043 pub fn memory_usage(&self) -> Result<MemoryUsage, ServerError> {
1049 self.device
1050 .submit_blocking(move |server| {
1051 server
1052 .stream_ids()
1053 .into_iter()
1054 .try_fold(MemoryUsage::default(), |acc, id| {
1055 Ok(acc.combine(server.memory_usage(id)?))
1056 })
1057 })
1058 .unwrap_or_resume()
1059 }
1060
1061 pub fn memory_report(&self) -> Result<MemoryReport, ServerError> {
1074 let stream_id = self.stream_id();
1075 self.device
1076 .submit_blocking(move |server| server.memory_report(stream_id))
1077 .unwrap_or_resume()
1078 }
1079
1080 pub fn enumerate_devices(&self, type_id: u16) -> Vec<DeviceId> {
1082 R::enumerate_devices(type_id, self.info())
1083 }
1084
1085 pub fn enumerate_all_devices(&self) -> Vec<DeviceId> {
1087 R::enumerate_all_devices(self.info())
1088 }
1089
1090 pub fn device_count(&self, type_id: u16) -> usize {
1092 self.enumerate_devices(type_id).len()
1093 }
1094
1095 pub fn device_count_total(&self) -> usize {
1097 self.enumerate_all_devices().len()
1098 }
1099
1100 pub unsafe fn allocation_mode(&self, mode: MemoryAllocationMode) {
1106 let stream_id = self.stream_id();
1107 self.device
1108 .submit(move |server| server.allocation_mode(mode, stream_id));
1109 }
1110
1111 pub fn memory_cleanup(&self) {
1116 self.device.submit(move |server| {
1117 for id in server.stream_ids() {
1118 server.memory_cleanup(id);
1119 }
1120 });
1121 }
1122
1123 pub fn install_memory_pools(
1168 &self,
1169 pools: &MemoryPoolsConfig,
1170 ) -> Result<(), InstallMemoryPoolsError> {
1171 let config =
1172 match MemoryConfiguration::default().resolve(Some(pools), &self.properties().memory) {
1173 Ok(config) => config,
1174 Err(err) => panic!("Invalid memory pools configuration: {err}"),
1175 };
1176 let stream_id = self.stream_id();
1177 self.device
1178 .submit_blocking(move |server| server.install_memory_pools(config, stream_id))
1179 .unwrap_or_resume()
1180 }
1181
1182 #[track_caller]
1184 pub fn profile<O: Send + 'static>(
1185 &self,
1186 func: impl FnOnce() -> O + Send,
1187 #[allow(unused)] func_name: &str,
1188 ) -> Result<(O, ProfileDuration), ProfileError> {
1189 #[cfg(feature = "profile-tracy")]
1192 let location = std::panic::Location::caller();
1193
1194 #[cfg(feature = "profile-tracy")]
1196 let _span = tracy_client::Client::running().unwrap().span_alloc(
1197 None,
1198 func_name,
1199 location.file(),
1200 location.line(),
1201 0,
1202 );
1203
1204 let stream_id = self.stream_id();
1205
1206 #[cfg(feature = "profile-tracy")]
1207 let gpu_span = if self.utilities.properties.timing_method == TimingMethod::Device {
1208 let gpu_span = self
1209 .utilities
1210 .gpu_client
1211 .span_alloc(func_name, "profile", location.file(), location.line())
1212 .unwrap();
1213 Some(gpu_span)
1214 } else {
1215 None
1216 };
1217
1218 let device = self.device.clone();
1219 #[allow(unused_mut, reason = "Used in profile-tracy")]
1220 let mut result = self
1221 .device
1222 .exclusive(move || {
1223 let token =
1226 match device.submit_blocking(move |server| server.start_profile(stream_id)) {
1227 Ok(token) => match token {
1228 Ok(token) => token,
1229 Err(err) => return Err(err),
1230 },
1231 Err(err) => {
1232 return Err(ServerError::Generic {
1233 reason: alloc::format!(
1234 "Can't start profiling because of a call error: {err:?}"
1235 ),
1236 backtrace: BackTrace::capture(),
1237 });
1238 }
1239 };
1240
1241 let out = func();
1243
1244 let result = device
1246 .submit_blocking(move |server| {
1247 let mut result = server.end_profile(stream_id, token);
1248
1249 match result {
1250 Ok(result) => Ok((out, result)),
1251 Err(err) => Err(err),
1252 }
1253 })
1254 .unwrap_or_resume();
1255
1256 Ok(result)
1257 })
1258 .unwrap_or_resume()
1259 .map_err(|err| ProfileError::Unknown {
1260 reason: alloc::format!("{err}"),
1261 backtrace: BackTrace::capture(),
1262 })?;
1263
1264 #[cfg(feature = "profile-tracy")]
1265 if let Some(mut gpu_span) = gpu_span {
1266 gpu_span.end_zone();
1267 let epoch = self.utilities.epoch_time;
1268 result = result.map(|(o, result)| {
1270 (
1271 o,
1272 ProfileDuration::new(
1273 alloc::boxed::Box::pin(async move {
1274 let ticks = result.resolve().await;
1275 let start_duration =
1276 ticks.start_duration_since(epoch).as_nanos() as i64;
1277 let end_duration = ticks.end_duration_since(epoch).as_nanos() as i64;
1278 gpu_span.upload_timestamp_start(start_duration);
1279 gpu_span.upload_timestamp_end(end_duration);
1280 ticks
1281 }),
1282 TimingMethod::Device,
1283 ),
1284 )
1285 });
1286 }
1287
1288 result
1289 }
1290
1291 #[cfg_attr(
1293 feature = "tracing",
1294 tracing::instrument(
1295 level = "trace",
1296 skip(self, src_descriptor, alloc_descriptor, dst_server)
1297 )
1298 )]
1299 fn change_client_sync(
1300 &self,
1301 src_descriptor: CopyDescriptor,
1302 alloc_descriptor: MemoryLayoutDescriptor,
1303 dst_server: &Self,
1304 ) -> MemoryLayout {
1305 let shape = src_descriptor.shape.clone();
1306 let elem_size = src_descriptor.elem_size;
1307 let stream_id = self.stream_id();
1308
1309 let read = self
1310 .device
1311 .submit_blocking(move |server| server.read(vec![src_descriptor], stream_id))
1312 .unwrap_or_resume();
1313
1314 let mut data = cubecl_environment::future::block_on(read).unwrap();
1315
1316 let (handle_base, mut layouts) = self
1317 .utilities
1318 .layout_policy
1319 .apply(stream_id, &[alloc_descriptor]);
1320 let alloc = layouts.remove(0);
1321
1322 let desc_descriptor = CopyDescriptor {
1323 handle: handle_base.clone().binding(),
1324 shape,
1325 strides: alloc.strides.clone(),
1326 elem_size,
1327 };
1328
1329 let (size, memory) = (handle_base.size(), handle_base.memory);
1330 dst_server.device.submit(move |server| {
1331 server.initialize_memory(memory, size, stream_id);
1332 server.write(vec![(desc_descriptor, data.remove(0))], stream_id)
1333 });
1334
1335 alloc
1336 }
1337
1338 pub fn io_optimized_vector_sizes(
1340 &self,
1341 size: usize,
1342 ) -> impl Iterator<Item = VectorSize> + Clone {
1343 let load_width = self.properties().hardware.load_width as usize;
1344 let size_bits = size * 8;
1345 let max = load_width / size_bits;
1346 let max = usize::min(self.properties().hardware.max_vector_size, max);
1347
1348 let num_candidates = max.trailing_zeros() + 1;
1350
1351 (0..num_candidates).map(|i| 2usize.pow(i)).rev()
1352 }
1353
1354 fn device_key(&self) -> String {
1356 format!("{}_dev{}", R::name(self), self.device.device_id().index_id)
1357 }
1358
1359 pub fn measure_throughput(
1361 &self,
1362 key: ThroughputKey,
1363 kernel_config: KernelConfig,
1364 ) -> ThroughputValue {
1365 let cache = ThroughputCache::get_for_device(&self.device_key());
1366 let mut throughputs = ThroughputBenchmarker::new(cache);
1367 throughputs.measure(key, kernel_config)
1368 }
1369}