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    /// Explicitly compact eligible allocations and release empty pages.
393    /// Backends must submit pending address users and confirm relocation completion.
394    fn memory_compact(&mut self, _stream_id: StreamId) -> Result<(), ServerError> {
395        Err(IoError::UnsupportedIoOperation { backtrace: BackTrace::capture() }.into())
396    }
397
398    /// Enable collecting timestamps.
399    fn start_profile(&mut self, stream_id: StreamId) -> Result<ProfilingToken, ServerError>;
400
401    /// Disable collecting timestamps.
402    fn end_profile(
403        &mut self,
404        stream_id: StreamId,
405        token: ProfilingToken,
406    ) -> Result<ProfileDuration, ProfileError>;
407
408    /// Update the memory mode of allocation in the server.
409    fn allocation_mode(&mut self, mode: MemoryAllocationMode, stream_id: StreamId);
410}
411
412pub use ruda_core::device::CommunicationId;
413
414/// Different reduce operations.
415pub enum ReduceOperation {
416    /// Sum.
417    Sum,
418    /// Mean.
419    Mean,
420}
421
422/// Defines functions for optimized data transfer between servers, supporting custom communication
423/// mechanisms such as peer-to-peer communication or specialized implementations.
424pub trait ServerCommunication {
425    /// Indicates whether server-to-server communication is enabled for this implementation.
426    const SERVER_COMM_ENABLED: bool;
427
428    /// Ensure that all queued collective operations have been executed.
429    ///
430    /// # Arguments
431    ///
432    /// * `stream_id` - The [`StreamId`] of the stream waiting for the sync.
433    ///
434    /// # Returns
435    ///
436    /// Returns a `Result` containing an `ServerError` if the operation fails.
437    #[allow(unused_variables)]
438    fn sync_collective(&mut self, stream_id: StreamId) -> Result<(), ServerError> {
439        todo!() // For backends other than cuda.
440    }
441
442    /// Initialize the communication between the devices in `device_ids`.
443    ///
444    /// # Arguments
445    ///
446    /// * `device_ids` - The IDs of the devices that need communication.
447    ///
448    /// # Returns
449    ///
450    /// Returns a `Result` containing an `ServerError` if the operation fails.
451    #[allow(unused_variables)]
452    fn comm_init(&mut self, device_ids: Vec<DeviceId>) -> Result<(), ServerError> {
453        unimplemented!()
454    }
455
456    /// Performs an `all_reduce` operation on the input data and writes it to the output buffer.
457    /// see <https://docs.nvidia.com/deeplearning/nccl/user-guide/docs/usage/collectives.html#allreduce>
458    ///
459    /// # Arguments
460    ///
461    /// * `src` - The data to be reduced.
462    /// * `dst` - Where to write the result.
463    /// * `dtype` - The element type of the data being reduced
464    /// * `stream_id` - The data's stream id.
465    /// * `op` - The reduce's aggregation operation e.g. mean, sum, etc.
466    /// * `device_ids` - The list of device ids from which to `all_reduce`.
467    ///
468    /// # Returns
469    ///
470    /// Returns a `Result` containing an `ServerError` if the operation fails.
471    #[allow(unused_variables)]
472    fn all_reduce(
473        &mut self,
474        src: Binding,
475        dst: Binding,
476        dtype: ElemType,
477        stream_id: StreamId,
478        op: ReduceOperation,
479        device_ids: Vec<DeviceId>,
480    ) -> Result<(), ServerError> {
481        unimplemented!()
482    }
483
484    /// Sends data from this server to a destination server.
485    ///
486    /// # Arguments
487    ///
488    /// * `desc` - A descriptor specifying the data to be sent, including shape, strides, and binding.
489    /// * `dtype` - The element type of the data being sent.
490    /// * `stream_id` - The stream ID associated with the server's operation.
491    /// * `device_id_dst` - ID of the device receiving the data.
492    ///
493    /// # Returns
494    ///
495    /// Returns a `Result` containing an `ServerError` if the operation fails.
496    #[allow(unused_variables)]
497    fn send(
498        &mut self,
499        desc: CopyDescriptor,
500        dtype: ElemType,
501        stream_id: StreamId,
502        device_id_dst: DeviceId,
503    ) -> Result<(), ServerError> {
504        unimplemented!()
505    }
506
507    /// Receive data from another server.
508    ///
509    /// # Arguments
510    ///
511    /// * `handle` - The handle in which the received data is written.
512    /// * `dtype` - The element type of the data being sent.
513    /// * `stream_id` - The stream ID associated with the server's operation.
514    /// * `device_id_src` - ID of the device sending the data.
515    ///
516    /// # Returns
517    ///
518    /// Returns a `Result` containing an `ServerError` if the operation fails.
519    #[allow(unused_variables)]
520    fn recv(
521        &mut self,
522        handle: Handle,
523        dtype: ElemType,
524        stream_id: StreamId,
525        device_id_src: DeviceId,
526    ) -> Result<(), ServerError> {
527        unimplemented!()
528    }
529}
530
531#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq)]
532/// Profiling identification so that the server can support recursive and overlapping profilings.
533pub struct ProfilingToken {
534    /// The token value.
535    pub id: u64,
536}
537
538/// Type of allocation, either contiguous or optimized (row-aligned when possible)
539#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq)]
540pub enum MemoryLayoutStrategy {
541    /// Contiguous layout, with no padding
542    Contiguous,
543    /// Optimized for access speed. In practice this means row-aligned with padding for runtimes
544    /// that support it.
545    Optimized,
546}
547
548/// Descriptor for a new tensor allocation
549#[derive(new, Debug, Clone)]
550pub struct MemoryLayoutDescriptor {
551    /// Strategy used to create the memory layout.
552    pub strategy: MemoryLayoutStrategy,
553    /// Shape of the tensor
554    pub shape: Shape,
555    /// Size of each element in the tensor (used for conversion of shape to bytes)
556    pub elem_size: usize,
557}
558
559impl MemoryLayoutDescriptor {
560    /// Create an optimized allocation descriptor
561    pub fn optimized(shape: Shape, elem_size: usize) -> Self {
562        MemoryLayoutDescriptor::new(MemoryLayoutStrategy::Optimized, shape, elem_size)
563    }
564
565    /// Create a contiguous allocation descriptor
566    pub fn contiguous(shape: Shape, elem_size: usize) -> Self {
567        MemoryLayoutDescriptor::new(MemoryLayoutStrategy::Contiguous, shape, elem_size)
568    }
569}
570
571/// An allocation with associated strides. Strides depend on tensor layout.
572#[derive(Debug, Clone)]
573pub struct MemoryLayout {
574    /// The handle for the memory resource
575    pub memory: Handle,
576    /// TODO: `Strides` should become `Layout`.
577    ///
578    /// The strides of the tensor
579    pub strides: Strides,
580}
581
582impl MemoryLayout {
583    /// Create a new memory layout.
584    pub fn new(handle: Handle, strides: impl Into<Strides>) -> Self {
585        MemoryLayout {
586            memory: handle,
587            strides: strides.into(),
588        }
589    }
590}
591
592/// A reason for an error.
593#[derive(Default, Clone)]
594pub struct Reason {
595    inner: ReasonInner,
596}
597
598#[cfg(std_io)]
599mod _reason_serde {
600    use super::*;
601
602    use alloc::string::ToString;
603    use serde::{Deserialize, Deserializer, Serialize, Serializer};
604
605    impl Serialize for Reason {
606        fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
607        where
608            S: Serializer,
609        {
610            // Use the Display implementation (via to_string) to flatten the enum
611            serializer.serialize_str(&self.to_string())
612        }
613    }
614
615    impl<'de> Deserialize<'de> for Reason {
616        fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
617        where
618            D: Deserializer<'de>,
619        {
620            // Deserialize into a standard String first
621            let s = String::deserialize(deserializer)?;
622
623            // Wrap it in the Dynamic variant since we can't safely
624            // reconstruct a 'static str from a runtime string.
625            Ok(Reason {
626                inner: ReasonInner::Dynamic(Arc::new(s)),
627            })
628        }
629    }
630}
631
632#[derive(Default, Clone)]
633enum ReasonInner {
634    Static(&'static str),
635    Dynamic(Arc<String>),
636    #[default]
637    NotProvided,
638}
639
640impl core::fmt::Display for Reason {
641    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
642        match &self.inner {
643            ReasonInner::Static(content) => f.write_str(content),
644            ReasonInner::Dynamic(content) => f.write_str(content),
645            ReasonInner::NotProvided => f.write_str("No reason provided for the error"),
646        }
647    }
648}
649
650impl core::fmt::Debug for Reason {
651    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
652        core::fmt::Display::fmt(&self, f)
653    }
654}
655
656impl From<&'static str> for Reason {
657    fn from(value: &'static str) -> Self {
658        Self {
659            inner: ReasonInner::Static(value),
660        }
661    }
662}
663
664impl From<String> for Reason {
665    fn from(value: String) -> Self {
666        Self {
667            inner: ReasonInner::Dynamic(Arc::new(value)),
668        }
669    }
670}
671
672/// Error returned from `create`/`read`/`write` functions. Due to async execution not all errors
673/// are able to be caught, so some IO errors will still panic.
674#[derive(Error, Clone)]
675#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
676pub enum IoError {
677    /// Buffer size exceeds the max available
678    #[error("can't allocate buffer of size: {size}\n{backtrace}")]
679    BufferTooBig {
680        /// The size of the buffer in bytes.
681        size: u64,
682        /// The captured backtrace.
683        #[cfg_attr(std_io, serde(skip))]
684        backtrace: BackTrace,
685    },
686
687    /// Strides aren't supported for this copy operation on this runtime
688    #[error("the provided strides are not supported for this operation\n{backtrace}")]
689    UnsupportedStrides {
690        /// The backtrace.
691        #[cfg_attr(std_io, serde(skip))]
692        backtrace: BackTrace,
693    },
694
695    /// Memory wasn't found in the memory pool
696    #[error("couldn't find resource for that handle: {reason}\n{backtrace}")]
697    NotFound {
698        /// The backtrace.
699        #[cfg_attr(std_io, serde(skip))]
700        backtrace: BackTrace,
701        /// The reason the handle is invalid.
702        reason: Reason,
703    },
704
705    /// Handle wasn't found in the memory pool
706    #[error("couldn't free the handle, since it is currently in used. \n{backtrace}")]
707    FreeError {
708        /// The backtrace.
709        #[cfg_attr(std_io, serde(skip))]
710        backtrace: BackTrace,
711    },
712
713    /// Unknown error happened during execution
714    #[error("Unknown error happened during execution\n{backtrace}")]
715    Unknown {
716        /// Details of the error
717        description: String,
718        /// The backtrace.
719        #[cfg_attr(std_io, serde(skip))]
720        backtrace: BackTrace,
721    },
722
723    /// The current IO operation is not supported
724    #[error("The current IO operation is not supported\n{backtrace}")]
725    UnsupportedIoOperation {
726        /// The backtrace.
727        #[cfg_attr(std_io, serde(skip))]
728        backtrace: BackTrace,
729    },
730
731    /// Can't perform the IO operation because of a runtime error.
732    #[error("Can't perform the IO operation because of a runtime error: {0}")]
733    Execution(#[from] Box<ServerError>),
734}
735
736impl core::fmt::Debug for IoError {
737    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
738        f.write_fmt(format_args!("{self}"))
739    }
740}
741
742/// Arguments to execute a kernel.
743#[derive(Debug, Default)]
744pub struct KernelArguments {
745    /// Buffer bindings
746    pub buffers: Vec<Binding>,
747    /// Packed scalars and metadata. First scalars sorted by type, then static metadata,
748    /// then dynamic metadata.
749    pub info: MetadataBindingInfo,
750    /// Tensor map bindings
751    pub tensor_maps: Vec<TensorMapBinding>,
752}
753
754impl core::fmt::Display for KernelArguments {
755    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
756        f.write_str("KernelArguments")?;
757        for b in self.buffers.iter() {
758            f.write_fmt(format_args!("\n - buffer: {b:?}\n"))?;
759        }
760
761        Ok(())
762    }
763}
764
765impl KernelArguments {
766    /// Create a new bindings struct
767    pub fn new() -> Self {
768        Self::default()
769    }
770
771    /// Add a buffer binding
772    pub fn with_buffer(mut self, binding: Binding) -> Self {
773        self.buffers.push(binding);
774        self
775    }
776
777    /// Extend the buffers with `bindings`
778    pub fn with_buffers(mut self, bindings: Vec<Binding>) -> Self {
779        self.buffers.extend(bindings);
780        self
781    }
782
783    /// Set the info to `info`
784    pub fn with_info(mut self, info: MetadataBindingInfo) -> Self {
785        self.info = info;
786        self
787    }
788
789    /// Extend the tensor maps with `bindings`
790    pub fn with_tensor_maps(mut self, bindings: Vec<TensorMapBinding>) -> Self {
791        self.tensor_maps.extend(bindings);
792        self
793    }
794}
795
796/// Binding of a set of scalars of the same type to execute a kernel.
797///
798/// The [`ComputeServer`] is responsible to convert those info into actual [`Binding`] when launching
799/// kernels.
800#[derive(new, Debug, Default)]
801pub struct MetadataBindingInfo {
802    /// Scalar and metadata values
803    pub data: Vec<u64>,
804    /// Start of the dynamically sized portion of the metadata, relative to the entire info buffer
805    pub dynamic_metadata_offset: usize,
806}
807
808impl MetadataBindingInfo {
809    /// Create a new binding info for custom data, for externally compiled kernels.
810    pub fn custom(data: Vec<u64>) -> Self {
811        Self::new(data, 0)
812    }
813}
814
815/// A binding with shape and stride info for non-contiguous reading
816#[derive(new, Debug)]
817pub struct CopyDescriptor {
818    /// Binding for the memory resource
819    pub handle: Binding,
820    /// Shape of the resource
821    pub shape: Shape,
822    /// Strides of the resource
823    pub strides: Strides,
824    /// Size of each element in the resource
825    pub elem_size: usize,
826}
827
828/// A tensor map used with TMA ops
829#[derive(new, Debug)]
830pub struct TensorMapBinding {
831    /// The binding for the backing tensor
832    pub binding: Binding,
833    /// The tensormap metadata
834    pub map: TensorMapMeta,
835}
836
837/// `TensorMap` metadata for the opaque proxy used in TMA copies
838#[derive(Debug, Clone)]
839pub struct TensorMapMeta {
840    /// Tensormap format (tiled or im2col)
841    pub format: TensorMapFormat,
842    /// Metadata of the backing tensor
843    pub metadata: Metadata,
844    /// Element stride, usually 1 but may be 2 for complex tensors
845    /// For im2col, this is equivalent to the kernel stride
846    pub elem_stride: Strides,
847    /// Interleave mode
848    pub interleave: TensorMapInterleave,
849    /// Swizzle mode
850    pub swizzle: TensorMapSwizzle,
851    /// Prefetch settings
852    pub prefetch: TensorMapPrefetch,
853    /// OOB fill value
854    pub oob_fill: OobFill,
855    /// Storage type
856    pub storage_ty: StorageType,
857}
858
859/// Specifieds the number of rudas to be dispatched for a kernel.
860///
861/// This translates to eg. a grid for CUDA, or to `num_workgroups` for wgsl.
862#[allow(clippy::large_enum_variant)]
863pub enum RudaCount {
864    /// Dispatch a known count of x, y, z rudas.
865    Static(u32, u32, u32),
866    /// Dispatch an amount based on the values in this buffer. The buffer should contain a u32 array [x, y, z].
867    Dynamic(Binding),
868}
869
870/// Defines how to select ruda count based on the number of rudas required.
871pub enum RudaCountSelection {
872    /// If the number of rudas is the same as required.
873    Exact(RudaCount),
874    /// If the number of rudas isn't the same as required.
875    ///
876    /// This can happen based on the hardware limit, requiring the kernel to perform OOB checks.
877    Approx(RudaCount, u32),
878}
879
880impl RudaCountSelection {
881    /// Creates a [`RudaCount`] while respecting the hardware limits.
882    pub fn new<R: Runtime>(client: &ComputeClient<R>, num_rudas: u32) -> Self {
883        let ruda_count = ruda_count_spread(&client.properties().hardware.max_ruda_count, num_rudas);
884
885        let num_rudas_actual = ruda_count[0] * ruda_count[1] * ruda_count[2];
886        let ruda_count = RudaCount::Static(ruda_count[0], ruda_count[1], ruda_count[2]);
887
888        match num_rudas_actual == num_rudas {
889            true => RudaCountSelection::Exact(ruda_count),
890            false => RudaCountSelection::Approx(ruda_count, num_rudas_actual),
891        }
892    }
893
894    /// If some rudas will be idle.
895    pub fn has_idle(&self) -> bool {
896        matches!(self, Self::Approx(..))
897    }
898
899    /// Converts into [`RudaCount`].
900    pub fn ruda_count(self) -> RudaCount {
901        match self {
902            RudaCountSelection::Exact(ruda_count) => ruda_count,
903            RudaCountSelection::Approx(ruda_count, _) => ruda_count,
904        }
905    }
906}
907
908impl From<RudaCountSelection> for RudaCount {
909    fn from(value: RudaCountSelection) -> Self {
910        value.ruda_count()
911    }
912}
913
914impl RudaCount {
915    /// Create a new static ruda count with the given x = y = z = 1.
916    pub fn new_single() -> Self {
917        RudaCount::Static(1, 1, 1)
918    }
919
920    /// Create a new static ruda count with the given x, and y = z = 1.
921    pub fn new_1d(x: u32) -> Self {
922        RudaCount::Static(x, 1, 1)
923    }
924
925    /// Create a new static ruda count with the given x and y, and z = 1.
926    pub fn new_2d(x: u32, y: u32) -> Self {
927        RudaCount::Static(x, y, 1)
928    }
929
930    /// Create a new static ruda count with the given x, y and z.
931    pub fn new_3d(x: u32, y: u32, z: u32) -> Self {
932        RudaCount::Static(x, y, z)
933    }
934
935    /// Checks whether the ruda count is definitely empty, i.e. has 0 dispatches.
936    pub fn is_empty(&self) -> bool {
937        match self {
938            Self::Static(x, y, z) => *x == 0 || *y == 0 || *z == 0,
939            Self::Dynamic(_) => false,
940        }
941    }
942}
943
944impl Debug for RudaCount {
945    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
946        match self {
947            RudaCount::Static(x, y, z) => f.write_fmt(format_args!("({x}, {y}, {z})")),
948            RudaCount::Dynamic(_) => f.write_str("binding"),
949        }
950    }
951}
952
953impl Clone for RudaCount {
954    fn clone(&self) -> Self {
955        match self {
956            Self::Static(x, y, z) => Self::Static(*x, *y, *z),
957            Self::Dynamic(binding) => Self::Dynamic(binding.clone()),
958        }
959    }
960}
961
962pub use ruda_core::launch::{RudaDim, ExecutionMode};
963
964fn ruda_count_spread(max: &(u32, u32, u32), num_rudas: u32) -> [u32; 3] {
965    let max_ruda_counts = [max.0, max.1, max.2];
966    let mut num_rudas = [num_rudas, 1, 1];
967    let base = 2;
968
969    let mut reduce_count = |i: usize| {
970        if num_rudas[i] <= max_ruda_counts[i] {
971            return true;
972        }
973
974        loop {
975            num_rudas[i] = num_rudas[i].div_ceil(base);
976            num_rudas[i + 1] *= base;
977
978            if num_rudas[i] <= max_ruda_counts[i] {
979                return false;
980            }
981        }
982    };
983
984    for i in 0..2 {
985        if reduce_count(i) {
986            break;
987        }
988    }
989
990    num_rudas
991}
992
993#[cfg(test)]
994mod tests {
995    use super::*;
996
997    #[test_log::test]
998    fn safe_num_rudas_even() {
999        let max = (32, 32, 32);
1000        let required = 2048;
1001
1002        let actual = ruda_count_spread(&max, required);
1003        let expected = [32, 32, 2];
1004        assert_eq!(actual, expected);
1005    }
1006
1007    #[test_log::test]
1008    fn safe_num_rudas_odd() {
1009        let max = (48, 32, 16);
1010        let required = 3177;
1011
1012        let actual = ruda_count_spread(&max, required);
1013        let expected = [25, 32, 4];
1014        assert_eq!(actual, expected);
1015    }
1016}