ruda_runtime/runtime/tune/
tuner.rs1use alloc::format;
2use alloc::sync::Arc;
3use alloc::vec::Vec;
4use ruda_core::profile::ProfileDuration;
5
6use core::time::Duration;
7
8use alloc::string::{String, ToString};
9use ruda_core::benchmark::{BenchmarkComputations, BenchmarkDurations};
10
11use crate::runtime::config::{Logger, autotune::AutotuneLogLevel};
12use crate::runtime::server::LaunchError;
13use crate::runtime::tune::{AutotuneResult, TuneCache, tune_benchmark};
14use crate::runtime::{client::ComputeClient, backend::Runtime};
15
16use super::{AutotuneKey, AutotuneOutput, TunableSet, TuneCacheResult, TuneInputs};
17
18#[derive(Debug)]
19pub struct Tuner<K: AutotuneKey> {
25 cache: Arc<spin::RwLock<TuneCache<K>>>,
26 logger: Arc<spin::Mutex<Logger>>,
27}
28
29#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
31#[derive(new, Debug, Clone, PartialEq, Eq)]
32pub struct AutotuneOutcome {
33 name: String,
34 index: usize,
35 computation: BenchmarkComputations,
36}
37
38impl core::fmt::Display for AutotuneOutcome {
39 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
40 write!(
41 f,
42 "Autotune[{}] name {} => {:?}",
43 self.index, self.name, self.computation
44 )
45 }
46}
47
48#[derive(Debug, Clone)]
50#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
51pub enum AutotuneError {
52 Unknown {
54 name: String,
56 err: String,
58 },
59 InvalidSamples {
61 name: String,
63 },
64 NoValidKernelFound {
70 context: String,
72 },
73 Skip {
75 name: String,
77 },
78
79 Launch(LaunchError),
81}
82
83impl From<LaunchError> for AutotuneError {
84 fn from(value: LaunchError) -> Self {
85 Self::Launch(value)
86 }
87}
88
89struct PendingBench {
91 index: usize,
92 name: String,
93 profiles: Vec<ProfileDuration>,
94}
95
96struct TuneRequest<K: AutotuneKey> {
99 key: K,
100 results: Vec<AutotuneResult>,
101 #[cfg(std_io)]
102 checksum: String,
103 context_logs: Option<String>,
104 pending: Vec<PendingBench>,
105}
106
107#[allow(clippy::new_without_default)]
108impl<K: AutotuneKey> Tuner<K> {
109 pub fn new(name: &str, device_id: &str) -> Self {
112 Self {
113 cache: Arc::new(spin::RwLock::new(TuneCache::new(name, device_id))),
114 logger: Arc::new(spin::Mutex::new(Logger::new())),
115 }
116 }
117
118 pub fn fastest(&self, key: &K) -> TuneCacheResult {
120 self.cache.read().fastest(key)
121 }
122
123 pub fn check_tune<'a, R: Runtime, F: TuneInputs, Out: AutotuneOutput>(
126 &self,
127 key: &K,
128 inputs: &F::At<'a>,
129 tunables: &TunableSet<K, F, Out>,
130 #[cfg_attr(not(std_io), allow(unused))] checksum: impl FnOnce() -> String + Send + Sync,
131 client: &ComputeClient<R>,
132 ) -> TuneCacheResult
133 where
134 <F as TuneInputs>::At<'a>: Clone + Send,
135 {
136 client.flush().expect("Cannot autotune with pre-existing stream errors");
139 {
140 let mut cache = self.cache.write();
141 let cur = cache.fastest(key);
142
143 #[cfg(std_io)]
144 let cur = if matches!(cur, TuneCacheResult::Unchecked) {
145 let mut log = self.logger.lock();
146 let checksum = checksum();
147 if let AutotuneLogLevel::Full = log.log_level_autotune() {
148 log.log_autotune(&format!("validate checksum key={key}, checksum={checksum}"));
149 }
150 cache.validate_checksum(key, &checksum)
151 } else {
152 cur
153 };
154
155 match cur {
156 TuneCacheResult::Hit { .. } | TuneCacheResult::Pending => return cur,
157 TuneCacheResult::Miss | TuneCacheResult::Unchecked => {
158 cache.mark_pending(key.clone())
159 }
160 }
161 }
164
165 log::info!("Tuning {key}");
166
167 if tunables.len() == 1 {
169 self.cache.write().cache_insert(key.clone(), 0);
170 return TuneCacheResult::Hit { fastest_index: 0 };
171 }
172
173 let autotunables = tunables.autotunables().collect::<Vec<_>>();
174 let mut results: Vec<AutotuneResult> = autotunables
175 .iter()
176 .map(|a| {
177 AutotuneResult::error(AutotuneError::Skip {
178 name: a.name.to_string(),
179 })
180 })
181 .collect();
182
183 #[cfg(std_io)]
184 let checksum = tunables.compute_checksum();
185
186 let test_inputs = tunables.generate_inputs(key, inputs);
187 client.flush().expect("Autotune input generation failed");
188 let mut plan = tunables.plan(key);
189 let mut context_logs = match self.logger.lock().log_level_autotune() {
190 AutotuneLogLevel::Full => Some(String::new()),
191 _ => None,
192 };
193
194 let mut pending = Vec::<PendingBench>::new();
199 loop {
200 let tunable_indices = plan.next(context_logs.as_mut());
201
202 if tunable_indices.is_empty() {
203 panic!(
204 "Can't execute the autotune plan for key: {key:?}\n - plan: {plan:?}\n - results: {results:?}"
205 );
206 }
207
208 for index in tunable_indices {
209 let op = autotunables[index];
210
211 match tune_benchmark(op, test_inputs.clone(), client.clone()) {
212 Ok(profiles) => pending.push(PendingBench {
213 index,
214 name: op.name.clone(),
215 profiles,
216 }),
217 Err(err) => {
218 results[index] = AutotuneResult::error(err);
219 }
220 }
221 }
222
223 if !pending.is_empty() {
224 break;
225 }
226 }
227
228 let request = TuneRequest {
229 key: key.clone(),
230 results,
231 #[cfg(std_io)]
232 checksum,
233 context_logs,
234 pending,
235 };
236
237 #[cfg(target_family = "wasm")]
240 {
241 let cache = self.cache.clone();
242 let logger = self.logger.clone();
243 wasm_bindgen_futures::spawn_local(async move {
244 process_request(request, &cache, &logger).await;
245 });
246
247 return TuneCacheResult::Pending;
248 }
249
250 #[cfg(not(target_family = "wasm"))]
251 ruda_core::future::block_on(process_request(request, &self.cache, &self.logger))
252 }
253}
254
255async fn process_request<K: AutotuneKey>(
257 request: TuneRequest<K>,
258 cache: &spin::RwLock<TuneCache<K>>,
259 logger: &spin::Mutex<Logger>,
260) -> TuneCacheResult {
261 let TuneRequest {
262 key,
263 mut results,
264 #[cfg(std_io)]
265 checksum,
266 context_logs,
267 pending,
268 } = request;
269
270 for bench in pending {
271 let PendingBench {
272 index,
273 name,
274 profiles,
275 } = bench;
276
277 if profiles.is_empty() {
278 results[index] = AutotuneResult::error(AutotuneError::Unknown {
279 name,
280 err: "No profiling available".to_string(),
281 });
282 continue;
283 }
284
285 let timing_method = profiles.first().unwrap().timing_method();
286 let mut durations = Vec::with_capacity(profiles.len());
287 for profile in profiles {
288 durations.push(profile.resolve().await.duration());
289 }
290
291 results[index] = AutotuneResult::success(AutotuneOutcome::new(
292 name,
293 index,
294 BenchmarkComputations::new(&BenchmarkDurations::from_durations(
295 timing_method,
296 durations,
297 )),
298 ));
299 }
300
301 results.sort_by_cached_key(|result| {
302 result
303 .outcome
304 .as_ref()
305 .map(|r| r.computation.score())
306 .unwrap_or(u64::MAX)
307 });
308
309 let fastest_index = results
310 .first()
311 .expect("At least one kernel needed.")
312 .outcome
313 .as_ref()
314 .expect("At least one kernel has to succeed.")
315 .index;
316
317 {
318 log_result(&mut logger.lock(), &key, &results, context_logs.as_deref());
319 cache.write().cache_insert(key.clone(), fastest_index);
320 #[cfg(std_io)]
321 cache
322 .write()
323 .persistent_cache_insert(key, checksum, fastest_index, results);
324 }
325
326 TuneCacheResult::Hit { fastest_index }
327}
328
329fn log_result<K: AutotuneKey>(
331 logger: &mut Logger,
332 key: &K,
333 results: &[AutotuneResult],
334 context_logs: Option<&str>,
335) {
336 match logger.log_level_autotune() {
337 AutotuneLogLevel::Minimal => {
338 let top_times = results
339 .iter()
340 .map(|r| {
341 let time = r
342 .outcome
343 .as_ref()
344 .map(|r| r.computation.median)
345 .unwrap_or(Duration::MAX);
346
347 let index = r.outcome.as_ref().map(|r| r.index).unwrap_or_default();
348 (index, time)
349 })
350 .take(3)
351 .collect::<Vec<_>>();
352
353 let result = results
354 .first()
355 .expect("At least one kernel needed.")
356 .outcome
357 .as_ref()
358 .expect("At least one kernel has to succeed.");
359
360 let context = context_logs.unwrap_or("");
361 logger.log_autotune(&format!(
362 "Fastest result {}-{key}. \n Top 3 times: {top_times:?}, context: {context}",
363 result.name,
364 ));
365 }
366 AutotuneLogLevel::Full => {
367 let result = results
368 .first()
369 .expect("At least one kernel needed.")
370 .outcome
371 .as_ref()
372 .expect("At least one kernel has to succeed.");
373
374 let context = context_logs.unwrap_or("");
375 logger.log_autotune(&format!(
376 "Fastest result {}-{key}. Context: {context}",
377 result.name,
378 ));
379
380 for result in results.iter() {
381 match &result.outcome {
382 Ok(val) => {
383 logger.log_autotune(&format!("{val}"));
384 }
385 Err(err) => logger.log_autotune(&format!("{err:?}")),
386 }
387 }
388 }
389 AutotuneLogLevel::Disabled => {}
390 }
391}
392
393#[cfg(feature = "runtime-autotune-checks")]
394pub(crate) fn check_autotune_outputs<O: AutotuneOutput>(
395 mut checks_outputs: Vec<Result<O, AutotuneError>>,
396) {
397 let reference = checks_outputs.remove(checks_outputs.len() - 1);
398
399 if let Ok(reference) = reference {
400 for other in checks_outputs.into_iter().flatten() {
401 reference.check_equivalence(other);
402 }
403 }
404}