Skip to main content

ruda_runtime/runtime/
kernel.rs

1use alloc::{
2    boxed::Box,
3    string::{String, ToString},
4};
5use core::{
6    fmt::Display,
7    marker::PhantomData,
8    sync::atomic::{AtomicI8, Ordering},
9};
10
11use ruda_core::format::format_str;
12use ruda_core::ir::StorageType;
13
14use crate::runtime::{
15    compiler::{CompilationError, Compiler, RudaTask},
16    config::{RudaRuntimeConfig, RuntimeConfig, compilation::CompilationLogLevel},
17    id::KernelId,
18    server::{RudaDim, ExecutionMode},
19};
20
21/// Implement this trait to create a [kernel definition](KernelDefinition).
22pub trait KernelMetadata: Send + Sync + 'static {
23    /// Name of the kernel for debugging.
24    fn name(&self) -> &'static str {
25        core::any::type_name::<Self>()
26    }
27
28    /// Identifier for the kernel, used for caching kernel compilation.
29    fn id(&self) -> KernelId;
30
31    /// Type of addresses in this kernel
32    fn address_type(&self) -> StorageType;
33}
34
35pub use ruda_core::kernel::{KernelArg, KernelDefinition, KernelOptions, ScalarKernelArg, Visibility};
36
37/// A kernel, compiled in the target language
38pub struct CompiledKernel<C: Compiler> {
39    /// The name of the kernel entrypoint.
40    /// For example
41    ///
42    /// ```text
43    /// #[ruda(launch)]
44    /// fn gelu_array<F: Float, R: Runtime>() {}
45    /// ```
46    ///
47    /// would have the entrypoint name "`gelu_array`".
48    pub entrypoint_name: String,
49
50    /// A fully qualified debug name of the kernel.
51    ///
52    /// For example
53    ///
54    /// ```text
55    /// #[ruda(launch)]
56    /// fn gelu_array<F: Float, R: Runtime>() {}
57    /// ```
58    ///
59    /// would have a debug name such as
60    ///
61    /// ```text
62    /// gelu::gelu_array::GeluArray<
63    ///    ruda_kernel::dsl::frontend::element::float::F32,
64    ///    ruda_driver_cuda::runtime::CudaRuntime,
65    /// >
66    /// ```
67    pub debug_name: Option<&'static str>,
68
69    /// Source code of the kernel
70    pub source: String,
71    /// In-memory representation of the kernel
72    pub repr: Option<C::Representation>,
73    /// Size of a ruda for the compiled kernel
74    pub ruda_dim: RudaDim,
75    /// Extra debugging information about the compiled kernel.
76    pub debug_info: Option<DebugInformation>,
77}
78
79/// Extra debugging information about the compiled kernel.
80#[derive(new)]
81pub struct DebugInformation {
82    /// The language tag of the source..
83    pub lang_tag: &'static str,
84    /// The compilation id.
85    pub id: KernelId,
86}
87
88/// Kernel that can be defined
89pub trait RudaKernel: KernelMetadata {
90    /// Define the kernel for compilation
91    fn define(&self) -> KernelDefinition;
92}
93
94/// Wraps a [`RudaKernel`] to allow it be compiled.
95pub struct KernelTask<C: Compiler, K: RudaKernel> {
96    kernel_definition: K,
97    _compiler: PhantomData<C>,
98}
99
100/// Generic [`RudaTask`] for compiling kernels
101pub struct RudaTaskKernel<C: Compiler> {
102    /// The inner compilation task being wrapped
103    pub task: Box<dyn RudaTask<C>>,
104}
105
106impl<C: Compiler, K: RudaKernel> KernelTask<C, K> {
107    /// Create a new kernel task
108    pub fn new(kernel_definition: K) -> Self {
109        Self {
110            kernel_definition,
111            _compiler: PhantomData,
112        }
113    }
114}
115
116impl<C: Compiler, K: RudaKernel> RudaTask<C> for KernelTask<C, K> {
117    fn kernel_definition(&self) -> Option<KernelDefinition> {
118        Some(self.kernel_definition.define())
119    }
120
121    fn compile(
122        &self,
123        compiler: &mut C,
124        compilation_options: &C::CompilationOptions,
125        mode: ExecutionMode,
126        addr_type: StorageType,
127    ) -> Result<CompiledKernel<C>, CompilationError> {
128        let gpu_ir = self.kernel_definition.define();
129        let entrypoint_name = gpu_ir.options.kernel_name.clone();
130        let ruda_dim = gpu_ir.ruda_dim;
131        let lower_level_ir = compiler.compile(gpu_ir, compilation_options, mode, addr_type)?;
132
133        Ok(CompiledKernel {
134            entrypoint_name,
135            debug_name: Some(core::any::type_name::<K>()),
136            source: lower_level_ir.to_string(),
137            repr: Some(lower_level_ir),
138            ruda_dim,
139            debug_info: None,
140        })
141    }
142}
143
144impl<C: Compiler, K: RudaKernel> KernelMetadata for KernelTask<C, K> {
145    // Forward ID to underlying kernel definition.
146    fn id(&self) -> KernelId {
147        self.kernel_definition.id()
148    }
149
150    // Forward name to underlying kernel definition.
151    fn name(&self) -> &'static str {
152        self.kernel_definition.name()
153    }
154
155    fn address_type(&self) -> StorageType {
156        self.kernel_definition.address_type()
157    }
158}
159
160impl<C: Compiler> KernelMetadata for Box<dyn RudaTask<C>> {
161    // Deref and use existing ID.
162    fn id(&self) -> KernelId {
163        self.as_ref().id()
164    }
165
166    // Deref and use existing name.
167    fn name(&self) -> &'static str {
168        self.as_ref().name()
169    }
170
171    fn address_type(&self) -> StorageType {
172        self.as_ref().address_type()
173    }
174}
175
176static COMPILATION_LEVEL: AtomicI8 = AtomicI8::new(-1);
177
178fn compilation_level() -> u8 {
179    let compilation_level = COMPILATION_LEVEL.load(Ordering::Relaxed);
180    if compilation_level == -1 {
181        let val = match RudaRuntimeConfig::get().compilation.logger.level {
182            CompilationLogLevel::Full => 2,
183            CompilationLogLevel::Disabled => 0,
184            CompilationLogLevel::Basic => 1,
185        };
186
187        COMPILATION_LEVEL.store(val, Ordering::Relaxed);
188        val as u8
189    } else {
190        compilation_level as u8
191    }
192}
193
194impl<C: Compiler> Display for CompiledKernel<C> {
195    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
196        match compilation_level() {
197            2 => self.format_full(f),
198            _ => self.format_basic(f),
199        }
200    }
201}
202
203impl<C: Compiler> CompiledKernel<C> {
204    fn format_basic(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
205        f.write_str("[Compiling kernel]")?;
206        if let Some(name) = self.debug_name {
207            if name.len() <= 32 {
208                f.write_fmt(format_args!(" {name}"))?;
209            } else {
210                f.write_fmt(format_args!(" {}", name.split('<').next().unwrap_or("")))?;
211            }
212        }
213
214        Ok(())
215    }
216
217    fn format_full(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
218        f.write_str("[START_KERNEL_COMPILATION]")?;
219
220        if let Some(name) = self.debug_name {
221            if name.len() <= 32 {
222                f.write_fmt(format_args!("\nname: {name}"))?;
223            } else {
224                let name = format_str(name, &[('<', '>')], false);
225                f.write_fmt(format_args!("\nname: {name}"))?;
226            }
227        }
228
229        if let Some(info) = &self.debug_info {
230            f.write_fmt(format_args!("\nid: {:#?}", info.id))?;
231        }
232
233        f.write_fmt(format_args!(
234            "
235source:
236```{}
237{}
238```
239[END_KERNEL_COMPILATION]
240",
241            self.debug_info
242                .as_ref()
243                .map(|info| info.lang_tag)
244                .unwrap_or(""),
245            self.source
246        ))
247    }
248}