Skip to main content

ruda_runtime/runtime/server/
base.rs

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))]
37/// An error during profiling.
38pub enum ProfileError {
39    /// An unknown error happened during profiling
40    #[error(
41        "An unknown error happened during profiling\nCaused by:\n  {reason}\nBacktrace:\n{backtrace}"
42    )]
43    Unknown {
44        /// The caused of the error
45        reason: String,
46        /// The captured backtrace.
47        #[cfg_attr(std_io, serde(skip))]
48        backtrace: BackTrace,
49    },
50
51    /// No profiling was registered
52    #[error("No profiling registered\nBacktrace:\n{backtrace}")]
53    NotRegistered {
54        /// The captured backtrace.
55        #[cfg_attr(std_io, serde(skip))]
56        backtrace: BackTrace,
57    },
58
59    /// A launch error happened during profiling
60    #[error("A launch error happened during profiling\nCaused by:\n  {0}")]
61    Launch(#[from] LaunchError),
62
63    /// An execution error happened during profiling
64    #[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
74/// Contains many different types that are useful for server implementations and compute clients.
75pub struct ServerUtilities<Server: ComputeServer> {
76    /// The time when `profile-tracy` is activated.
77    #[cfg(feature = "runtime-profile-tracy")]
78    pub epoch_time: web_time::Instant,
79    /// The GPU client when `profile-tracy` is activated.
80    #[cfg(feature = "runtime-profile-tracy")]
81    pub gpu_client: tracy_client::GpuContext,
82    /// Information shared between all servers.
83    pub properties: DeviceProperties,
84    /// Stable hash of the device properties
85    pub properties_hash: u64,
86    /// Information specific to the current server.
87    pub info: Server::Info,
88    /// The logger based on global ruda configs.
89    pub logger: Arc<ServerLogger>,
90    /// How to create the allocation.
91    pub layout_policy: Server::MemoryLayoutPolicy,
92    /// How to enforce bounds checking on kernels.
93    pub check_mode: BoundsCheckMode,
94    /// A set containing the ids for which the inter-device communication has already been initialized.
95    pub initialized_comms: RwLock<HashSet<CommunicationId>>,
96}
97
98/// Defines how the memory layout is determined.
99pub trait MemoryLayoutPolicy: Send + Sync + 'static {
100    /// Applies the memory layout policy to a list of descriptors.
101    ///
102    /// Returns a vector of `MemoryLayout`, one per descriptor, with layouts that share a
103    /// single `Binding`.
104    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    /// Creates a new server utilities.
127    pub fn new(
128        properties: DeviceProperties,
129        logger: Arc<ServerLogger>,
130        info: S::Info,
131        allocator: S::MemoryLayoutPolicy,
132    ) -> Self {
133        // Start a tracy client if needed.
134        #[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            // Create the GPU client if needed.
142            #[cfg(feature = "runtime-profile-tracy")]
143            gpu_client: client
144                .clone()
145                .new_gpu_context(
146                    Some(&format!("{info:?}")),
147                    // In the future should ask the server what makes sense here. 'Invalid' atm is a generic stand-in (Tracy doesn't have CUDA/RocM atm anyway).
148                    tracy_client::GpuContextType::Invalid,
149                    0,   // Timestamps are manually aligned to this epoch so start at 0.
150                    1.0, // Timestamps are manually converted to be nanoseconds so period is 1.
151                )
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/// Kernel Launch Errors.
164#[derive(Error, Clone)]
165#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
166pub enum LaunchError {
167    /// The given kernel can't be compiled.
168    #[error("A compilation error happened during launch\nCaused by:\n  {0}")]
169    CompilationError(#[from] CompilationError),
170
171    /// The server is out of memory.
172    #[error(
173        "An out-of-memory error happened during launch\nCaused by:\n  {reason}\nBacktrace\n{backtrace}"
174    )]
175    OutOfMemory {
176        /// The caused of the memory error.
177        reason: String,
178        /// The backtrace for this error.
179        #[cfg_attr(std_io, serde(skip))]
180        backtrace: BackTrace,
181    },
182
183    /// Too many resources were requested
184    #[error("Too many resources were requested during launch\n{0}")]
185    TooManyResources(#[from] ResourceLimitError),
186
187    /// Unknown launch error.
188    #[error(
189        "An unknown error happened during launch\nCaused by:\n  {reason}\nBacktrace\n{backtrace}"
190    )]
191    Unknown {
192        /// The caused of the unknown error.
193        reason: String,
194        /// The backtrace for this error.
195        #[cfg_attr(std_io, serde(skip))]
196        backtrace: BackTrace,
197    },
198
199    /// Can't launch because of an IO Error.
200    #[error("An io error happened during launch\nCaused by:\n  {0}")]
201    IoError(#[from] IoError),
202}
203
204/// Resource limit errors.
205#[derive(Error, Clone)]
206#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
207pub enum ResourceLimitError {
208    /// Shared memory exceeds maximum
209    #[error(
210        "Too much shared memory requested.\nRequested {requested} bytes, maximum {max} bytes available.\nBacktrace\n{backtrace}"
211    )]
212    SharedMemory {
213        /// Value requested
214        requested: usize,
215        /// Maximum value
216        max: usize,
217        /// The backtrace for this error.
218        #[cfg_attr(std_io, serde(skip))]
219        backtrace: BackTrace,
220    },
221    /// Total units exceeds maximum
222    #[error(
223        "Total unit count exceeds maximum.\nRequested {requested} units, max units is {max}.\nBacktrace\n{backtrace}"
224    )]
225    Units {
226        /// Requested value
227        requested: u32,
228        /// Maximum value
229        max: u32,
230        /// The backtrace for this error.
231        #[cfg_attr(std_io, serde(skip))]
232        backtrace: BackTrace,
233    },
234    /// `RudaDim` exceeds maximum
235    #[error(
236        "Ruda dim exceeds maximum bounds.\nRequested {requested:?}, max is {max:?}.\nBacktrace\n{backtrace}"
237    )]
238    RudaDim {
239        /// Requested value
240        requested: (u32, u32, u32),
241        /// Maximum value
242        max: (u32, u32, u32),
243        /// The backtrace for this error.
244        #[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/// Error that can happen asynchronously while executing registered kernels.
262#[derive(Error, Debug, Clone)]
263#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
264pub enum ServerError {
265    /// A generic runtime error.
266    #[error("An error happened during execution\nCaused by:\n  {reason}\nBacktrace:\n{backtrace}")]
267    Generic {
268        /// The details of the generic error.
269        reason: String,
270        /// The backtrace for this error.
271        #[cfg_attr(std_io, serde(skip))]
272        backtrace: BackTrace,
273    },
274
275    /// A launch error happened during profiling
276    #[error("A launch error happened during profiling\nCaused by:\n  {0}")]
277    Launch(#[from] LaunchError),
278
279    /// An execution error happened during profiling
280    #[error("An execution error happened during profiling\nCaused by:\n  {0}")]
281    Profile(#[from] ProfileError),
282
283    /// An execution error happened during profiling
284    #[error("An execution error happened during profiling\nCaused by:\n  {0}")]
285    Io(#[from] IoError),
286
287    /// The server is an invalid state.
288    #[error("The server is in an invalid state\nCaused by:\n  {errors:?}")]
289    ServerUnhealthy {
290        /// The details of the generic error.
291        errors: Vec<Self>,
292        /// The backtrace for this error.
293        #[cfg_attr(std_io, serde(skip))]
294        backtrace: BackTrace,
295    },
296}
297
298/// How errors are handled in a stream when executing a task.
299#[derive(Clone, Copy)]
300pub struct StreamErrorMode {
301    /// Whether the task still executes even if the stream is in error.
302    pub ignore: bool,
303    /// Whether the errors are flushed by the current task.
304    pub flush: bool,
305}
306
307/// The compute server is responsible for handling resources and computations over resources.
308///
309/// Everything in the server is mutable, therefore it should be solely accessed through the
310/// [`ComputeClient`] for thread safety.
311pub trait ComputeServer:
312    Send + core::fmt::Debug + ServerCommunication + device::DeviceService + 'static
313where
314    Self: Sized,
315{
316    /// The kernel type defines the computation algorithms.
317    type Kernel: KernelMetadata;
318    /// Information that can be retrieved for the runtime.
319    type Info: Debug + Send + Sync;
320    /// Manages how allocations are performed for a server.
321    type MemoryLayoutPolicy: MemoryLayoutPolicy;
322    /// The [storage](ComputeStorage) type defines how data is stored and accessed.
323    type Storage: ComputeStorage;
324
325    /// Initializes [memory](ManagedMemoryHandle) on the given [stream](StreamId) with the given size.
326    fn initialize_memory(&mut self, memory: ManagedMemoryHandle, size: u64, stream_id: StreamId);
327
328    /// Reserves N [Bytes] of the provided sizes to be used as staging to load data.
329    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    /// Retrieve the server logger.
341    fn logger(&self) -> Arc<ServerLogger>;
342
343    /// Retrieve the server utilities.
344    fn utilities(&self) -> Arc<ServerUtilities<Self>>;
345
346    /// Given bindings, returns the owned resources as bytes.
347    fn read(
348        &mut self,
349        descriptors: Vec<CopyDescriptor>,
350        stream_id: StreamId,
351    ) -> DynFut<Result<Vec<Bytes>, ServerError>>;
352
353    /// Writes the specified bytes into the buffers given
354    fn write(&mut self, descriptors: Vec<(CopyDescriptor, Bytes)>, stream_id: StreamId);
355
356    /// Wait for the completion of every task in the server.
357    fn sync(&mut self, stream_id: StreamId) -> DynFut<Result<(), ServerError>>;
358
359    /// Given a resource handle, returns the storage resource.
360    fn get_resource(
361        &mut self,
362        binding: Binding,
363        stream_id: StreamId,
364    ) -> Result<ManagedResource<<Self::Storage as ComputeStorage>::Resource>, ServerError>;
365
366    /// Executes the `kernel` over the given memory `handles`.
367    ///
368    /// Kernels have mutable access to every resource they are given
369    /// and are responsible of determining which should be read or written.
370    ///
371    /// # Safety
372    ///
373    /// When executing with mode [`ExecutionMode::Unchecked`], out-of-bound reads and writes can happen.
374    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    /// Flush all outstanding tasks in the server.
384    fn flush(&mut self, stream_id: StreamId) -> Result<(), ServerError>;
385
386    /// The current memory usage of the server.
387    fn memory_usage(&mut self, stream_id: StreamId) -> Result<MemoryUsage, ServerError>;
388
389    /// Ask the server to release memory that it can release.
390    fn memory_cleanup(&mut self, stream_id: StreamId);
391
392    /// Enable collecting timestamps.
393    fn start_profile(&mut self, stream_id: StreamId) -> Result<ProfilingToken, ServerError>;
394
395    /// Disable collecting timestamps.
396    fn end_profile(
397        &mut self,
398        stream_id: StreamId,
399        token: ProfilingToken,
400    ) -> Result<ProfileDuration, ProfileError>;
401
402    /// Update the memory mode of allocation in the server.
403    fn allocation_mode(&mut self, mode: MemoryAllocationMode, stream_id: StreamId);
404}
405
406pub use ruda_core::device::CommunicationId;
407
408/// Different reduce operations.
409pub enum ReduceOperation {
410    /// Sum.
411    Sum,
412    /// Mean.
413    Mean,
414}
415
416/// Defines functions for optimized data transfer between servers, supporting custom communication
417/// mechanisms such as peer-to-peer communication or specialized implementations.
418pub trait ServerCommunication {
419    /// Indicates whether server-to-server communication is enabled for this implementation.
420    const SERVER_COMM_ENABLED: bool;
421
422    /// Ensure that all queued collective operations have been executed.
423    ///
424    /// # Arguments
425    ///
426    /// * `stream_id` - The [`StreamId`] of the stream waiting for the sync.
427    ///
428    /// # Returns
429    ///
430    /// Returns a `Result` containing an `ServerError` if the operation fails.
431    #[allow(unused_variables)]
432    fn sync_collective(&mut self, stream_id: StreamId) -> Result<(), ServerError> {
433        todo!() // For backends other than cuda.
434    }
435
436    /// Initialize the communication between the devices in `device_ids`.
437    ///
438    /// # Arguments
439    ///
440    /// * `device_ids` - The IDs of the devices that need communication.
441    ///
442    /// # Returns
443    ///
444    /// Returns a `Result` containing an `ServerError` if the operation fails.
445    #[allow(unused_variables)]
446    fn comm_init(&mut self, device_ids: Vec<DeviceId>) -> Result<(), ServerError> {
447        unimplemented!()
448    }
449
450    /// Performs an `all_reduce` operation on the input data and writes it to the output buffer.
451    /// see <https://docs.nvidia.com/deeplearning/nccl/user-guide/docs/usage/collectives.html#allreduce>
452    ///
453    /// # Arguments
454    ///
455    /// * `src` - The data to be reduced.
456    /// * `dst` - Where to write the result.
457    /// * `dtype` - The element type of the data being reduced
458    /// * `stream_id` - The data's stream id.
459    /// * `op` - The reduce's aggregation operation e.g. mean, sum, etc.
460    /// * `device_ids` - The list of device ids from which to `all_reduce`.
461    ///
462    /// # Returns
463    ///
464    /// Returns a `Result` containing an `ServerError` if the operation fails.
465    #[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    /// Sends data from this server to a destination server.
479    ///
480    /// # Arguments
481    ///
482    /// * `desc` - A descriptor specifying the data to be sent, including shape, strides, and binding.
483    /// * `dtype` - The element type of the data being sent.
484    /// * `stream_id` - The stream ID associated with the server's operation.
485    /// * `device_id_dst` - ID of the device receiving the data.
486    ///
487    /// # Returns
488    ///
489    /// Returns a `Result` containing an `ServerError` if the operation fails.
490    #[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    /// Receive data from another server.
502    ///
503    /// # Arguments
504    ///
505    /// * `handle` - The handle in which the received data is written.
506    /// * `dtype` - The element type of the data being sent.
507    /// * `stream_id` - The stream ID associated with the server's operation.
508    /// * `device_id_src` - ID of the device sending the data.
509    ///
510    /// # Returns
511    ///
512    /// Returns a `Result` containing an `ServerError` if the operation fails.
513    #[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)]
526/// Profiling identification so that the server can support recursive and overlapping profilings.
527pub struct ProfilingToken {
528    /// The token value.
529    pub id: u64,
530}
531
532/// Type of allocation, either contiguous or optimized (row-aligned when possible)
533#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq)]
534pub enum MemoryLayoutStrategy {
535    /// Contiguous layout, with no padding
536    Contiguous,
537    /// Optimized for access speed. In practice this means row-aligned with padding for runtimes
538    /// that support it.
539    Optimized,
540}
541
542/// Descriptor for a new tensor allocation
543#[derive(new, Debug, Clone)]
544pub struct MemoryLayoutDescriptor {
545    /// Strategy used to create the memory layout.
546    pub strategy: MemoryLayoutStrategy,
547    /// Shape of the tensor
548    pub shape: Shape,
549    /// Size of each element in the tensor (used for conversion of shape to bytes)
550    pub elem_size: usize,
551}
552
553impl MemoryLayoutDescriptor {
554    /// Create an optimized allocation descriptor
555    pub fn optimized(shape: Shape, elem_size: usize) -> Self {
556        MemoryLayoutDescriptor::new(MemoryLayoutStrategy::Optimized, shape, elem_size)
557    }
558
559    /// Create a contiguous allocation descriptor
560    pub fn contiguous(shape: Shape, elem_size: usize) -> Self {
561        MemoryLayoutDescriptor::new(MemoryLayoutStrategy::Contiguous, shape, elem_size)
562    }
563}
564
565/// An allocation with associated strides. Strides depend on tensor layout.
566#[derive(Debug, Clone)]
567pub struct MemoryLayout {
568    /// The handle for the memory resource
569    pub memory: Handle,
570    /// TODO: `Strides` should become `Layout`.
571    ///
572    /// The strides of the tensor
573    pub strides: Strides,
574}
575
576impl MemoryLayout {
577    /// Create a new memory layout.
578    pub fn new(handle: Handle, strides: impl Into<Strides>) -> Self {
579        MemoryLayout {
580            memory: handle,
581            strides: strides.into(),
582        }
583    }
584}
585
586/// A reason for an error.
587#[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            // Use the Display implementation (via to_string) to flatten the enum
605            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            // Deserialize into a standard String first
615            let s = String::deserialize(deserializer)?;
616
617            // Wrap it in the Dynamic variant since we can't safely
618            // reconstruct a 'static str from a runtime string.
619            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/// Error returned from `create`/`read`/`write` functions. Due to async execution not all errors
667/// are able to be caught, so some IO errors will still panic.
668#[derive(Error, Clone)]
669#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
670pub enum IoError {
671    /// Buffer size exceeds the max available
672    #[error("can't allocate buffer of size: {size}\n{backtrace}")]
673    BufferTooBig {
674        /// The size of the buffer in bytes.
675        size: u64,
676        /// The captured backtrace.
677        #[cfg_attr(std_io, serde(skip))]
678        backtrace: BackTrace,
679    },
680
681    /// Strides aren't supported for this copy operation on this runtime
682    #[error("the provided strides are not supported for this operation\n{backtrace}")]
683    UnsupportedStrides {
684        /// The backtrace.
685        #[cfg_attr(std_io, serde(skip))]
686        backtrace: BackTrace,
687    },
688
689    /// Memory wasn't found in the memory pool
690    #[error("couldn't find resource for that handle: {reason}\n{backtrace}")]
691    NotFound {
692        /// The backtrace.
693        #[cfg_attr(std_io, serde(skip))]
694        backtrace: BackTrace,
695        /// The reason the handle is invalid.
696        reason: Reason,
697    },
698
699    /// Handle wasn't found in the memory pool
700    #[error("couldn't free the handle, since it is currently in used. \n{backtrace}")]
701    FreeError {
702        /// The backtrace.
703        #[cfg_attr(std_io, serde(skip))]
704        backtrace: BackTrace,
705    },
706
707    /// Unknown error happened during execution
708    #[error("Unknown error happened during execution\n{backtrace}")]
709    Unknown {
710        /// Details of the error
711        description: String,
712        /// The backtrace.
713        #[cfg_attr(std_io, serde(skip))]
714        backtrace: BackTrace,
715    },
716
717    /// The current IO operation is not supported
718    #[error("The current IO operation is not supported\n{backtrace}")]
719    UnsupportedIoOperation {
720        /// The backtrace.
721        #[cfg_attr(std_io, serde(skip))]
722        backtrace: BackTrace,
723    },
724
725    /// Can't perform the IO operation because of a runtime error.
726    #[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/// Arguments to execute a kernel.
737#[derive(Debug, Default)]
738pub struct KernelArguments {
739    /// Buffer bindings
740    pub buffers: Vec<Binding>,
741    /// Packed scalars and metadata. First scalars sorted by type, then static metadata,
742    /// then dynamic metadata.
743    pub info: MetadataBindingInfo,
744    /// Tensor map bindings
745    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    /// Create a new bindings struct
761    pub fn new() -> Self {
762        Self::default()
763    }
764
765    /// Add a buffer binding
766    pub fn with_buffer(mut self, binding: Binding) -> Self {
767        self.buffers.push(binding);
768        self
769    }
770
771    /// Extend the buffers with `bindings`
772    pub fn with_buffers(mut self, bindings: Vec<Binding>) -> Self {
773        self.buffers.extend(bindings);
774        self
775    }
776
777    /// Set the info to `info`
778    pub fn with_info(mut self, info: MetadataBindingInfo) -> Self {
779        self.info = info;
780        self
781    }
782
783    /// Extend the tensor maps with `bindings`
784    pub fn with_tensor_maps(mut self, bindings: Vec<TensorMapBinding>) -> Self {
785        self.tensor_maps.extend(bindings);
786        self
787    }
788}
789
790/// Binding of a set of scalars of the same type to execute a kernel.
791///
792/// The [`ComputeServer`] is responsible to convert those info into actual [`Binding`] when launching
793/// kernels.
794#[derive(new, Debug, Default)]
795pub struct MetadataBindingInfo {
796    /// Scalar and metadata values
797    pub data: Vec<u64>,
798    /// Start of the dynamically sized portion of the metadata, relative to the entire info buffer
799    pub dynamic_metadata_offset: usize,
800}
801
802impl MetadataBindingInfo {
803    /// Create a new binding info for custom data, for externally compiled kernels.
804    pub fn custom(data: Vec<u64>) -> Self {
805        Self::new(data, 0)
806    }
807}
808
809/// A binding with shape and stride info for non-contiguous reading
810#[derive(new, Debug)]
811pub struct CopyDescriptor {
812    /// Binding for the memory resource
813    pub handle: Binding,
814    /// Shape of the resource
815    pub shape: Shape,
816    /// Strides of the resource
817    pub strides: Strides,
818    /// Size of each element in the resource
819    pub elem_size: usize,
820}
821
822/// A tensor map used with TMA ops
823#[derive(new, Debug)]
824pub struct TensorMapBinding {
825    /// The binding for the backing tensor
826    pub binding: Binding,
827    /// The tensormap metadata
828    pub map: TensorMapMeta,
829}
830
831/// `TensorMap` metadata for the opaque proxy used in TMA copies
832#[derive(Debug, Clone)]
833pub struct TensorMapMeta {
834    /// Tensormap format (tiled or im2col)
835    pub format: TensorMapFormat,
836    /// Metadata of the backing tensor
837    pub metadata: Metadata,
838    /// Element stride, usually 1 but may be 2 for complex tensors
839    /// For im2col, this is equivalent to the kernel stride
840    pub elem_stride: Strides,
841    /// Interleave mode
842    pub interleave: TensorMapInterleave,
843    /// Swizzle mode
844    pub swizzle: TensorMapSwizzle,
845    /// Prefetch settings
846    pub prefetch: TensorMapPrefetch,
847    /// OOB fill value
848    pub oob_fill: OobFill,
849    /// Storage type
850    pub storage_ty: StorageType,
851}
852
853/// Specifieds the number of rudas to be dispatched for a kernel.
854///
855/// This translates to eg. a grid for CUDA, or to `num_workgroups` for wgsl.
856#[allow(clippy::large_enum_variant)]
857pub enum RudaCount {
858    /// Dispatch a known count of x, y, z rudas.
859    Static(u32, u32, u32),
860    /// Dispatch an amount based on the values in this buffer. The buffer should contain a u32 array [x, y, z].
861    Dynamic(Binding),
862}
863
864/// Defines how to select ruda count based on the number of rudas required.
865pub enum RudaCountSelection {
866    /// If the number of rudas is the same as required.
867    Exact(RudaCount),
868    /// If the number of rudas isn't the same as required.
869    ///
870    /// This can happen based on the hardware limit, requiring the kernel to perform OOB checks.
871    Approx(RudaCount, u32),
872}
873
874impl RudaCountSelection {
875    /// Creates a [`RudaCount`] while respecting the hardware limits.
876    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    /// If some rudas will be idle.
889    pub fn has_idle(&self) -> bool {
890        matches!(self, Self::Approx(..))
891    }
892
893    /// Converts into [`RudaCount`].
894    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    /// Create a new static ruda count with the given x = y = z = 1.
910    pub fn new_single() -> Self {
911        RudaCount::Static(1, 1, 1)
912    }
913
914    /// Create a new static ruda count with the given x, and y = z = 1.
915    pub fn new_1d(x: u32) -> Self {
916        RudaCount::Static(x, 1, 1)
917    }
918
919    /// Create a new static ruda count with the given x and y, and z = 1.
920    pub fn new_2d(x: u32, y: u32) -> Self {
921        RudaCount::Static(x, y, 1)
922    }
923
924    /// Create a new static ruda count with the given x, y and z.
925    pub fn new_3d(x: u32, y: u32, z: u32) -> Self {
926        RudaCount::Static(x, y, z)
927    }
928
929    /// Checks whether the ruda count is definitely empty, i.e. has 0 dispatches.
930    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}