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 mut results: Vec<AutotuneResult> = tunables
174 .autotunables()
175 .map(|a| {
176 AutotuneResult::error(AutotuneError::Skip {
177 name: a.name.to_string(),
178 })
179 })
180 .collect();
181
182 #[cfg(std_io)]
183 let checksum = tunables.compute_checksum();
184
185 let test_inputs = tunables.generate_inputs(key, inputs);
186 client.flush().expect("Autotune input generation failed");
187 let mut plan = tunables.plan(key);
188 let mut context_logs = match self.logger.lock().log_level_autotune() {
189 AutotuneLogLevel::Full => Some(String::new()),
190 _ => None,
191 };
192
193 let mut pending = Vec::<PendingBench>::new();
198 loop {
199 let tunable_indices = plan.next(context_logs.as_mut());
200
201 if tunable_indices.is_empty() {
202 panic!(
203 "Can't execute the autotune plan for key: {key:?}\n - plan: {plan:?}\n - results: {results:?}"
204 );
205 }
206
207 for index in tunable_indices {
208 let op = tunables.fastest(index);
209
210 match tune_benchmark(op, test_inputs.clone(), client.clone()) {
211 Ok(profiles) => pending.push(PendingBench {
212 index,
213 name: op.name.clone(),
214 profiles,
215 }),
216 Err(err) => {
217 results[index] = AutotuneResult::error(err);
218 }
219 }
220 }
221
222 if !pending.is_empty() {
223 break;
224 }
225 }
226
227 let request = TuneRequest {
228 key: key.clone(),
229 results,
230 #[cfg(std_io)]
231 checksum,
232 context_logs,
233 pending,
234 };
235
236 #[cfg(target_family = "wasm")]
239 {
240 let cache = self.cache.clone();
241 let logger = self.logger.clone();
242 wasm_bindgen_futures::spawn_local(async move {
243 process_request(request, &cache, &logger).await;
244 });
245
246 return TuneCacheResult::Pending;
247 }
248
249 #[cfg(not(target_family = "wasm"))]
250 ruda_core::future::block_on(process_request(request, &self.cache, &self.logger))
251 }
252}
253
254async fn process_request<K: AutotuneKey>(
256 request: TuneRequest<K>,
257 cache: &spin::RwLock<TuneCache<K>>,
258 logger: &spin::Mutex<Logger>,
259) -> TuneCacheResult {
260 let TuneRequest {
261 key,
262 mut results,
263 #[cfg(std_io)]
264 checksum,
265 context_logs,
266 pending,
267 } = request;
268
269 for bench in pending {
270 let PendingBench {
271 index,
272 name,
273 profiles,
274 } = bench;
275
276 if profiles.is_empty() {
277 results[index] = AutotuneResult::error(AutotuneError::Unknown {
278 name,
279 err: "No profiling available".to_string(),
280 });
281 continue;
282 }
283
284 let timing_method = profiles.first().unwrap().timing_method();
285 let mut durations = Vec::with_capacity(profiles.len());
286 for profile in profiles {
287 durations.push(profile.resolve().await.duration());
288 }
289
290 results[index] = AutotuneResult::success(AutotuneOutcome::new(
291 name,
292 index,
293 BenchmarkComputations::new(&BenchmarkDurations::from_durations(
294 timing_method,
295 durations,
296 )),
297 ));
298 }
299
300 results.sort_by_cached_key(|result| {
301 result
302 .outcome
303 .as_ref()
304 .map(|r| r.computation.score())
305 .unwrap_or(u64::MAX)
306 });
307
308 let fastest_index = results
309 .first()
310 .expect("At least one kernel needed.")
311 .outcome
312 .as_ref()
313 .expect("At least one kernel has to succeed.")
314 .index;
315
316 {
317 log_result(&mut logger.lock(), &key, &results, context_logs.as_deref());
318 cache.write().cache_insert(key.clone(), fastest_index);
319 #[cfg(std_io)]
320 cache
321 .write()
322 .persistent_cache_insert(key, checksum, fastest_index, results);
323 }
324
325 TuneCacheResult::Hit { fastest_index }
326}
327
328fn log_result<K: AutotuneKey>(
330 logger: &mut Logger,
331 key: &K,
332 results: &[AutotuneResult],
333 context_logs: Option<&str>,
334) {
335 match logger.log_level_autotune() {
336 AutotuneLogLevel::Minimal => {
337 let top_times = results
338 .iter()
339 .map(|r| {
340 let time = r
341 .outcome
342 .as_ref()
343 .map(|r| r.computation.median)
344 .unwrap_or(Duration::MAX);
345
346 let index = r.outcome.as_ref().map(|r| r.index).unwrap_or_default();
347 (index, time)
348 })
349 .take(3)
350 .collect::<Vec<_>>();
351
352 let result = results
353 .first()
354 .expect("At least one kernel needed.")
355 .outcome
356 .as_ref()
357 .expect("At least one kernel has to succeed.");
358
359 let context = context_logs.unwrap_or("");
360 logger.log_autotune(&format!(
361 "Fastest result {}-{key}. \n Top 3 times: {top_times:?}, context: {context}",
362 result.name,
363 ));
364 }
365 AutotuneLogLevel::Full => {
366 let result = results
367 .first()
368 .expect("At least one kernel needed.")
369 .outcome
370 .as_ref()
371 .expect("At least one kernel has to succeed.");
372
373 let context = context_logs.unwrap_or("");
374 logger.log_autotune(&format!(
375 "Fastest result {}-{key}. Context: {context}",
376 result.name,
377 ));
378
379 for result in results.iter() {
380 match &result.outcome {
381 Ok(val) => {
382 logger.log_autotune(&format!("{val}"));
383 }
384 Err(err) => logger.log_autotune(&format!("{err:?}")),
385 }
386 }
387 }
388 AutotuneLogLevel::Disabled => {}
389 }
390}
391
392#[cfg(feature = "runtime-autotune-checks")]
393pub(crate) fn check_autotune_outputs<O: AutotuneOutput>(
394 mut checks_outputs: Vec<Result<O, AutotuneError>>,
395) {
396 let reference = checks_outputs.remove(checks_outputs.len() - 1);
397
398 if let Ok(reference) = reference {
399 for other in checks_outputs.into_iter().flatten() {
400 reference.check_equivalence(other);
401 }
402 }
403}