ruda_runtime/runtime/
kernel.rs1use 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
21pub trait KernelMetadata: Send + Sync + 'static {
23 fn name(&self) -> &'static str {
25 core::any::type_name::<Self>()
26 }
27
28 fn id(&self) -> KernelId;
30
31 fn address_type(&self) -> StorageType;
33}
34
35pub use ruda_core::kernel::{KernelArg, KernelDefinition, KernelOptions, ScalarKernelArg, Visibility};
36
37pub struct CompiledKernel<C: Compiler> {
39 pub entrypoint_name: String,
49
50 pub debug_name: Option<&'static str>,
68
69 pub source: String,
71 pub repr: Option<C::Representation>,
73 pub ruda_dim: RudaDim,
75 pub debug_info: Option<DebugInformation>,
77}
78
79#[derive(new)]
81pub struct DebugInformation {
82 pub lang_tag: &'static str,
84 pub id: KernelId,
86}
87
88pub trait RudaKernel: KernelMetadata {
90 fn define(&self) -> KernelDefinition;
92}
93
94pub struct KernelTask<C: Compiler, K: RudaKernel> {
96 kernel_definition: K,
97 _compiler: PhantomData<C>,
98}
99
100pub struct RudaTaskKernel<C: Compiler> {
102 pub task: Box<dyn RudaTask<C>>,
104}
105
106impl<C: Compiler, K: RudaKernel> KernelTask<C, K> {
107 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 fn id(&self) -> KernelId {
147 self.kernel_definition.id()
148 }
149
150 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 fn id(&self) -> KernelId {
163 self.as_ref().id()
164 }
165
166 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}