1use super::Handle;
2use crate::runtime::{
3 client::ComputeClient,
4 compiler::CompilationError,
5 config::{RudaRuntimeConfig, RuntimeConfig, compilation::BoundsCheckMode},
6 kernel::KernelMetadata,
7 logging::ServerLogger,
8 memory_management::{ManagedMemoryHandle, MemoryAllocationMode, MemoryUsage},
9 backend::Runtime,
10 server::Binding,
11 storage::{ComputeStorage, ManagedResource},
12 tma::{OobFill, TensorMapFormat, TensorMapInterleave, TensorMapPrefetch, TensorMapSwizzle},
13};
14use alloc::boxed::Box;
15#[cfg(feature = "runtime-profile-tracy")]
16use alloc::format;
17use alloc::string::String;
18use alloc::sync::Arc;
19use alloc::vec::Vec;
20use core::fmt::Debug;
21use ruda_core::{
22 backtrace::BackTrace,
23 bytes::Bytes,
24 device::{self, DeviceId},
25 future::DynFut,
26 profile::ProfileDuration,
27 stream_id::StreamId,
28 stub::RwLock,
29};
30use ruda_core::ir::{DeviceProperties, ElemType, StorageType};
31use ruda_core::tensor::{Shape, Strides, metadata::Metadata};
32use hashbrown::HashSet;
33use thiserror::Error;
34
35#[derive(Error, Clone)]
36#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
37pub enum ProfileError {
39 #[error(
41 "An unknown error happened during profiling\nCaused by:\n {reason}\nBacktrace:\n{backtrace}"
42 )]
43 Unknown {
44 reason: String,
46 #[cfg_attr(std_io, serde(skip))]
48 backtrace: BackTrace,
49 },
50
51 #[error("No profiling registered\nBacktrace:\n{backtrace}")]
53 NotRegistered {
54 #[cfg_attr(std_io, serde(skip))]
56 backtrace: BackTrace,
57 },
58
59 #[error("A launch error happened during profiling\nCaused by:\n {0}")]
61 Launch(#[from] LaunchError),
62
63 #[error("An execution error happened during profiling\nCaused by:\n {0}")]
65 Server(#[from] Box<ServerError>),
66}
67
68impl core::fmt::Debug for ProfileError {
69 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
70 f.write_fmt(format_args!("{self}"))
71 }
72}
73
74pub struct ServerUtilities<Server: ComputeServer> {
76 #[cfg(feature = "runtime-profile-tracy")]
78 pub epoch_time: web_time::Instant,
79 #[cfg(feature = "runtime-profile-tracy")]
81 pub gpu_client: tracy_client::GpuContext,
82 pub properties: DeviceProperties,
84 pub properties_hash: u64,
86 pub info: Server::Info,
88 pub logger: Arc<ServerLogger>,
90 pub layout_policy: Server::MemoryLayoutPolicy,
92 pub check_mode: BoundsCheckMode,
94 pub initialized_comms: RwLock<HashSet<CommunicationId>>,
96}
97
98pub trait MemoryLayoutPolicy: Send + Sync + 'static {
100 fn apply(
105 &self,
106 stream_id: StreamId,
107 descriptors: &[MemoryLayoutDescriptor],
108 ) -> (Handle, Vec<MemoryLayout>);
109}
110
111impl<Server: core::fmt::Debug> core::fmt::Debug for ServerUtilities<Server>
112where
113 Server: ComputeServer,
114 Server::Info: core::fmt::Debug,
115{
116 fn fmt(&self, f: &mut core::fmt::Formatter) -> core::fmt::Result {
117 f.debug_struct("ServerUtilities")
118 .field("properties", &self.properties)
119 .field("info", &self.info)
120 .field("logger", &self.logger)
121 .finish()
122 }
123}
124
125impl<S: ComputeServer> ServerUtilities<S> {
126 pub fn new(
128 properties: DeviceProperties,
129 logger: Arc<ServerLogger>,
130 info: S::Info,
131 allocator: S::MemoryLayoutPolicy,
132 ) -> Self {
133 #[cfg(feature = "runtime-profile-tracy")]
135 let client = tracy_client::Client::start();
136
137 Self {
138 properties_hash: properties.checksum(),
139 properties,
140 logger,
141 #[cfg(feature = "runtime-profile-tracy")]
143 gpu_client: client
144 .clone()
145 .new_gpu_context(
146 Some(&format!("{info:?}")),
147 tracy_client::GpuContextType::Invalid,
149 0, 1.0, )
152 .unwrap(),
153 #[cfg(feature = "runtime-profile-tracy")]
154 epoch_time: web_time::Instant::now(),
155 info,
156 layout_policy: allocator,
157 check_mode: RudaRuntimeConfig::get().compilation.check_mode,
158 initialized_comms: RwLock::new(HashSet::default()),
159 }
160 }
161}
162
163#[derive(Error, Clone)]
165#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
166pub enum LaunchError {
167 #[error("A compilation error happened during launch\nCaused by:\n {0}")]
169 CompilationError(#[from] CompilationError),
170
171 #[error(
173 "An out-of-memory error happened during launch\nCaused by:\n {reason}\nBacktrace\n{backtrace}"
174 )]
175 OutOfMemory {
176 reason: String,
178 #[cfg_attr(std_io, serde(skip))]
180 backtrace: BackTrace,
181 },
182
183 #[error("Too many resources were requested during launch\n{0}")]
185 TooManyResources(#[from] ResourceLimitError),
186
187 #[error(
189 "An unknown error happened during launch\nCaused by:\n {reason}\nBacktrace\n{backtrace}"
190 )]
191 Unknown {
192 reason: String,
194 #[cfg_attr(std_io, serde(skip))]
196 backtrace: BackTrace,
197 },
198
199 #[error("An io error happened during launch\nCaused by:\n {0}")]
201 IoError(#[from] IoError),
202}
203
204#[derive(Error, Clone)]
206#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
207pub enum ResourceLimitError {
208 #[error(
210 "Too much shared memory requested.\nRequested {requested} bytes, maximum {max} bytes available.\nBacktrace\n{backtrace}"
211 )]
212 SharedMemory {
213 requested: usize,
215 max: usize,
217 #[cfg_attr(std_io, serde(skip))]
219 backtrace: BackTrace,
220 },
221 #[error(
223 "Total unit count exceeds maximum.\nRequested {requested} units, max units is {max}.\nBacktrace\n{backtrace}"
224 )]
225 Units {
226 requested: u32,
228 max: u32,
230 #[cfg_attr(std_io, serde(skip))]
232 backtrace: BackTrace,
233 },
234 #[error(
236 "Ruda dim exceeds maximum bounds.\nRequested {requested:?}, max is {max:?}.\nBacktrace\n{backtrace}"
237 )]
238 RudaDim {
239 requested: (u32, u32, u32),
241 max: (u32, u32, u32),
243 #[cfg_attr(std_io, serde(skip))]
245 backtrace: BackTrace,
246 },
247}
248
249impl core::fmt::Debug for LaunchError {
250 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
251 f.write_fmt(format_args!("{self}"))
252 }
253}
254
255impl core::fmt::Debug for ResourceLimitError {
256 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
257 f.write_fmt(format_args!("{self}"))
258 }
259}
260
261#[derive(Error, Debug, Clone)]
263#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
264pub enum ServerError {
265 #[error("An error happened during execution\nCaused by:\n {reason}\nBacktrace:\n{backtrace}")]
267 Generic {
268 reason: String,
270 #[cfg_attr(std_io, serde(skip))]
272 backtrace: BackTrace,
273 },
274
275 #[error("A launch error happened during profiling\nCaused by:\n {0}")]
277 Launch(#[from] LaunchError),
278
279 #[error("An execution error happened during profiling\nCaused by:\n {0}")]
281 Profile(#[from] ProfileError),
282
283 #[error("An execution error happened during profiling\nCaused by:\n {0}")]
285 Io(#[from] IoError),
286
287 #[error("The server is in an invalid state\nCaused by:\n {errors:?}")]
289 ServerUnhealthy {
290 errors: Vec<Self>,
292 #[cfg_attr(std_io, serde(skip))]
294 backtrace: BackTrace,
295 },
296}
297
298#[derive(Clone, Copy)]
300pub struct StreamErrorMode {
301 pub ignore: bool,
303 pub flush: bool,
305}
306
307pub trait ComputeServer:
312 Send + core::fmt::Debug + ServerCommunication + device::DeviceService + 'static
313where
314 Self: Sized,
315{
316 type Kernel: KernelMetadata;
318 type Info: Debug + Send + Sync;
320 type MemoryLayoutPolicy: MemoryLayoutPolicy;
322 type Storage: ComputeStorage;
324
325 fn initialize_memory(&mut self, memory: ManagedMemoryHandle, size: u64, stream_id: StreamId);
327
328 fn staging(
330 &mut self,
331 _sizes: &[usize],
332 _stream_id: StreamId,
333 ) -> Result<Vec<Bytes>, ServerError> {
334 Err(IoError::UnsupportedIoOperation {
335 backtrace: BackTrace::capture(),
336 }
337 .into())
338 }
339
340 fn logger(&self) -> Arc<ServerLogger>;
342
343 fn utilities(&self) -> Arc<ServerUtilities<Self>>;
345
346 fn read(
348 &mut self,
349 descriptors: Vec<CopyDescriptor>,
350 stream_id: StreamId,
351 ) -> DynFut<Result<Vec<Bytes>, ServerError>>;
352
353 fn write(&mut self, descriptors: Vec<(CopyDescriptor, Bytes)>, stream_id: StreamId);
355
356 fn sync(&mut self, stream_id: StreamId) -> DynFut<Result<(), ServerError>>;
358
359 fn get_resource(
361 &mut self,
362 binding: Binding,
363 stream_id: StreamId,
364 ) -> Result<ManagedResource<<Self::Storage as ComputeStorage>::Resource>, ServerError>;
365
366 unsafe fn launch(
375 &mut self,
376 kernel: Self::Kernel,
377 count: RudaCount,
378 bindings: KernelArguments,
379 kind: ExecutionMode,
380 stream_id: StreamId,
381 );
382
383 fn flush(&mut self, stream_id: StreamId) -> Result<(), ServerError>;
385
386 fn memory_usage(&mut self, stream_id: StreamId) -> Result<MemoryUsage, ServerError>;
388
389 fn memory_cleanup(&mut self, stream_id: StreamId);
391
392 fn start_profile(&mut self, stream_id: StreamId) -> Result<ProfilingToken, ServerError>;
394
395 fn end_profile(
397 &mut self,
398 stream_id: StreamId,
399 token: ProfilingToken,
400 ) -> Result<ProfileDuration, ProfileError>;
401
402 fn allocation_mode(&mut self, mode: MemoryAllocationMode, stream_id: StreamId);
404}
405
406pub use ruda_core::device::CommunicationId;
407
408pub enum ReduceOperation {
410 Sum,
412 Mean,
414}
415
416pub trait ServerCommunication {
419 const SERVER_COMM_ENABLED: bool;
421
422 #[allow(unused_variables)]
432 fn sync_collective(&mut self, stream_id: StreamId) -> Result<(), ServerError> {
433 todo!() }
435
436 #[allow(unused_variables)]
446 fn comm_init(&mut self, device_ids: Vec<DeviceId>) -> Result<(), ServerError> {
447 unimplemented!()
448 }
449
450 #[allow(unused_variables)]
466 fn all_reduce(
467 &mut self,
468 src: Binding,
469 dst: Binding,
470 dtype: ElemType,
471 stream_id: StreamId,
472 op: ReduceOperation,
473 device_ids: Vec<DeviceId>,
474 ) -> Result<(), ServerError> {
475 unimplemented!()
476 }
477
478 #[allow(unused_variables)]
491 fn send(
492 &mut self,
493 desc: CopyDescriptor,
494 dtype: ElemType,
495 stream_id: StreamId,
496 device_id_dst: DeviceId,
497 ) -> Result<(), ServerError> {
498 unimplemented!()
499 }
500
501 #[allow(unused_variables)]
514 fn recv(
515 &mut self,
516 handle: Handle,
517 dtype: ElemType,
518 stream_id: StreamId,
519 device_id_src: DeviceId,
520 ) -> Result<(), ServerError> {
521 unimplemented!()
522 }
523}
524
525#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq)]
526pub struct ProfilingToken {
528 pub id: u64,
530}
531
532#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq)]
534pub enum MemoryLayoutStrategy {
535 Contiguous,
537 Optimized,
540}
541
542#[derive(new, Debug, Clone)]
544pub struct MemoryLayoutDescriptor {
545 pub strategy: MemoryLayoutStrategy,
547 pub shape: Shape,
549 pub elem_size: usize,
551}
552
553impl MemoryLayoutDescriptor {
554 pub fn optimized(shape: Shape, elem_size: usize) -> Self {
556 MemoryLayoutDescriptor::new(MemoryLayoutStrategy::Optimized, shape, elem_size)
557 }
558
559 pub fn contiguous(shape: Shape, elem_size: usize) -> Self {
561 MemoryLayoutDescriptor::new(MemoryLayoutStrategy::Contiguous, shape, elem_size)
562 }
563}
564
565#[derive(Debug, Clone)]
567pub struct MemoryLayout {
568 pub memory: Handle,
570 pub strides: Strides,
574}
575
576impl MemoryLayout {
577 pub fn new(handle: Handle, strides: impl Into<Strides>) -> Self {
579 MemoryLayout {
580 memory: handle,
581 strides: strides.into(),
582 }
583 }
584}
585
586#[derive(Default, Clone)]
588pub struct Reason {
589 inner: ReasonInner,
590}
591
592#[cfg(std_io)]
593mod _reason_serde {
594 use super::*;
595
596 use alloc::string::ToString;
597 use serde::{Deserialize, Deserializer, Serialize, Serializer};
598
599 impl Serialize for Reason {
600 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
601 where
602 S: Serializer,
603 {
604 serializer.serialize_str(&self.to_string())
606 }
607 }
608
609 impl<'de> Deserialize<'de> for Reason {
610 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
611 where
612 D: Deserializer<'de>,
613 {
614 let s = String::deserialize(deserializer)?;
616
617 Ok(Reason {
620 inner: ReasonInner::Dynamic(Arc::new(s)),
621 })
622 }
623 }
624}
625
626#[derive(Default, Clone)]
627enum ReasonInner {
628 Static(&'static str),
629 Dynamic(Arc<String>),
630 #[default]
631 NotProvided,
632}
633
634impl core::fmt::Display for Reason {
635 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
636 match &self.inner {
637 ReasonInner::Static(content) => f.write_str(content),
638 ReasonInner::Dynamic(content) => f.write_str(content),
639 ReasonInner::NotProvided => f.write_str("No reason provided for the error"),
640 }
641 }
642}
643
644impl core::fmt::Debug for Reason {
645 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
646 core::fmt::Display::fmt(&self, f)
647 }
648}
649
650impl From<&'static str> for Reason {
651 fn from(value: &'static str) -> Self {
652 Self {
653 inner: ReasonInner::Static(value),
654 }
655 }
656}
657
658impl From<String> for Reason {
659 fn from(value: String) -> Self {
660 Self {
661 inner: ReasonInner::Dynamic(Arc::new(value)),
662 }
663 }
664}
665
666#[derive(Error, Clone)]
669#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
670pub enum IoError {
671 #[error("can't allocate buffer of size: {size}\n{backtrace}")]
673 BufferTooBig {
674 size: u64,
676 #[cfg_attr(std_io, serde(skip))]
678 backtrace: BackTrace,
679 },
680
681 #[error("the provided strides are not supported for this operation\n{backtrace}")]
683 UnsupportedStrides {
684 #[cfg_attr(std_io, serde(skip))]
686 backtrace: BackTrace,
687 },
688
689 #[error("couldn't find resource for that handle: {reason}\n{backtrace}")]
691 NotFound {
692 #[cfg_attr(std_io, serde(skip))]
694 backtrace: BackTrace,
695 reason: Reason,
697 },
698
699 #[error("couldn't free the handle, since it is currently in used. \n{backtrace}")]
701 FreeError {
702 #[cfg_attr(std_io, serde(skip))]
704 backtrace: BackTrace,
705 },
706
707 #[error("Unknown error happened during execution\n{backtrace}")]
709 Unknown {
710 description: String,
712 #[cfg_attr(std_io, serde(skip))]
714 backtrace: BackTrace,
715 },
716
717 #[error("The current IO operation is not supported\n{backtrace}")]
719 UnsupportedIoOperation {
720 #[cfg_attr(std_io, serde(skip))]
722 backtrace: BackTrace,
723 },
724
725 #[error("Can't perform the IO operation because of a runtime error: {0}")]
727 Execution(#[from] Box<ServerError>),
728}
729
730impl core::fmt::Debug for IoError {
731 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
732 f.write_fmt(format_args!("{self}"))
733 }
734}
735
736#[derive(Debug, Default)]
738pub struct KernelArguments {
739 pub buffers: Vec<Binding>,
741 pub info: MetadataBindingInfo,
744 pub tensor_maps: Vec<TensorMapBinding>,
746}
747
748impl core::fmt::Display for KernelArguments {
749 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
750 f.write_str("KernelArguments")?;
751 for b in self.buffers.iter() {
752 f.write_fmt(format_args!("\n - buffer: {b:?}\n"))?;
753 }
754
755 Ok(())
756 }
757}
758
759impl KernelArguments {
760 pub fn new() -> Self {
762 Self::default()
763 }
764
765 pub fn with_buffer(mut self, binding: Binding) -> Self {
767 self.buffers.push(binding);
768 self
769 }
770
771 pub fn with_buffers(mut self, bindings: Vec<Binding>) -> Self {
773 self.buffers.extend(bindings);
774 self
775 }
776
777 pub fn with_info(mut self, info: MetadataBindingInfo) -> Self {
779 self.info = info;
780 self
781 }
782
783 pub fn with_tensor_maps(mut self, bindings: Vec<TensorMapBinding>) -> Self {
785 self.tensor_maps.extend(bindings);
786 self
787 }
788}
789
790#[derive(new, Debug, Default)]
795pub struct MetadataBindingInfo {
796 pub data: Vec<u64>,
798 pub dynamic_metadata_offset: usize,
800}
801
802impl MetadataBindingInfo {
803 pub fn custom(data: Vec<u64>) -> Self {
805 Self::new(data, 0)
806 }
807}
808
809#[derive(new, Debug)]
811pub struct CopyDescriptor {
812 pub handle: Binding,
814 pub shape: Shape,
816 pub strides: Strides,
818 pub elem_size: usize,
820}
821
822#[derive(new, Debug)]
824pub struct TensorMapBinding {
825 pub binding: Binding,
827 pub map: TensorMapMeta,
829}
830
831#[derive(Debug, Clone)]
833pub struct TensorMapMeta {
834 pub format: TensorMapFormat,
836 pub metadata: Metadata,
838 pub elem_stride: Strides,
841 pub interleave: TensorMapInterleave,
843 pub swizzle: TensorMapSwizzle,
845 pub prefetch: TensorMapPrefetch,
847 pub oob_fill: OobFill,
849 pub storage_ty: StorageType,
851}
852
853#[allow(clippy::large_enum_variant)]
857pub enum RudaCount {
858 Static(u32, u32, u32),
860 Dynamic(Binding),
862}
863
864pub enum RudaCountSelection {
866 Exact(RudaCount),
868 Approx(RudaCount, u32),
872}
873
874impl RudaCountSelection {
875 pub fn new<R: Runtime>(client: &ComputeClient<R>, num_rudas: u32) -> Self {
877 let ruda_count = ruda_count_spread(&client.properties().hardware.max_ruda_count, num_rudas);
878
879 let num_rudas_actual = ruda_count[0] * ruda_count[1] * ruda_count[2];
880 let ruda_count = RudaCount::Static(ruda_count[0], ruda_count[1], ruda_count[2]);
881
882 match num_rudas_actual == num_rudas {
883 true => RudaCountSelection::Exact(ruda_count),
884 false => RudaCountSelection::Approx(ruda_count, num_rudas_actual),
885 }
886 }
887
888 pub fn has_idle(&self) -> bool {
890 matches!(self, Self::Approx(..))
891 }
892
893 pub fn ruda_count(self) -> RudaCount {
895 match self {
896 RudaCountSelection::Exact(ruda_count) => ruda_count,
897 RudaCountSelection::Approx(ruda_count, _) => ruda_count,
898 }
899 }
900}
901
902impl From<RudaCountSelection> for RudaCount {
903 fn from(value: RudaCountSelection) -> Self {
904 value.ruda_count()
905 }
906}
907
908impl RudaCount {
909 pub fn new_single() -> Self {
911 RudaCount::Static(1, 1, 1)
912 }
913
914 pub fn new_1d(x: u32) -> Self {
916 RudaCount::Static(x, 1, 1)
917 }
918
919 pub fn new_2d(x: u32, y: u32) -> Self {
921 RudaCount::Static(x, y, 1)
922 }
923
924 pub fn new_3d(x: u32, y: u32, z: u32) -> Self {
926 RudaCount::Static(x, y, z)
927 }
928
929 pub fn is_empty(&self) -> bool {
931 match self {
932 Self::Static(x, y, z) => *x == 0 || *y == 0 || *z == 0,
933 Self::Dynamic(_) => false,
934 }
935 }
936}
937
938impl Debug for RudaCount {
939 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
940 match self {
941 RudaCount::Static(x, y, z) => f.write_fmt(format_args!("({x}, {y}, {z})")),
942 RudaCount::Dynamic(_) => f.write_str("binding"),
943 }
944 }
945}
946
947impl Clone for RudaCount {
948 fn clone(&self) -> Self {
949 match self {
950 Self::Static(x, y, z) => Self::Static(*x, *y, *z),
951 Self::Dynamic(binding) => Self::Dynamic(binding.clone()),
952 }
953 }
954}
955
956pub use ruda_core::launch::{RudaDim, ExecutionMode};
957
958fn ruda_count_spread(max: &(u32, u32, u32), num_rudas: u32) -> [u32; 3] {
959 let max_ruda_counts = [max.0, max.1, max.2];
960 let mut num_rudas = [num_rudas, 1, 1];
961 let base = 2;
962
963 let mut reduce_count = |i: usize| {
964 if num_rudas[i] <= max_ruda_counts[i] {
965 return true;
966 }
967
968 loop {
969 num_rudas[i] = num_rudas[i].div_ceil(base);
970 num_rudas[i + 1] *= base;
971
972 if num_rudas[i] <= max_ruda_counts[i] {
973 return false;
974 }
975 }
976 };
977
978 for i in 0..2 {
979 if reduce_count(i) {
980 break;
981 }
982 }
983
984 num_rudas
985}
986
987#[cfg(test)]
988mod tests {
989 use super::*;
990
991 #[test_log::test]
992 fn safe_num_rudas_even() {
993 let max = (32, 32, 32);
994 let required = 2048;
995
996 let actual = ruda_count_spread(&max, required);
997 let expected = [32, 32, 2];
998 assert_eq!(actual, expected);
999 }
1000
1001 #[test_log::test]
1002 fn safe_num_rudas_odd() {
1003 let max = (48, 32, 16);
1004 let required = 3177;
1005
1006 let actual = ruda_count_spread(&max, required);
1007 let expected = [25, 32, 4];
1008 assert_eq!(actual, expected);
1009 }
1010}