Skip to main content

ruda_runtime/runtime/tune/
tuner.rs

1use 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)]
19/// Runs autotune benchmarks for a single device and caches the results.
20///
21/// On wasm, [`tune`](Self::tune) spawns its work on the browser event loop; elsewhere
22/// it blocks inline. Either way the benchmarking itself is synchronous; only the
23/// per-sample profile resolution is awaited.
24pub struct Tuner<K: AutotuneKey> {
25    cache: Arc<spin::RwLock<TuneCache<K>>>,
26    logger: Arc<spin::Mutex<Logger>>,
27}
28
29/// The measured outcome for a given autotune invocation.
30#[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/// Error from running autotune.
49#[derive(Debug, Clone)]
50#[cfg_attr(std_io, derive(serde::Serialize, serde::Deserialize))]
51pub enum AutotuneError {
52    /// An unknown error happened.
53    Unknown {
54        /// The name of the tunable.
55        name: String,
56        /// The unknown error,
57        err: String,
58    },
59    /// All samples are invalid.
60    InvalidSamples {
61        /// The name of the tunable.
62        name: String,
63    },
64    /// No autotune was flagged as valid for the problem.
65    ///
66    /// # Warning
67    ///
68    /// This is an unrecoverable error and will cause a panic.
69    NoValidKernelFound {
70        /// The formatted context on why no valid kernel was found.
71        context: String,
72    },
73    /// The autotune is skipped manually.
74    Skip {
75        /// The name of the skipped kernel.
76        name: String,
77    },
78
79    /// An error happened when launching a kernel.
80    Launch(LaunchError),
81}
82
83impl From<LaunchError> for AutotuneError {
84    fn from(value: LaunchError) -> Self {
85        Self::Launch(value)
86    }
87}
88
89/// A successfully-queued benchmark: the profile futures for each sample, plus its metadata.
90struct PendingBench {
91    index: usize,
92    name: String,
93    profiles: Vec<ProfileDuration>,
94}
95
96/// A queued tuning job: all data needed to resolve samples and commit the result.
97/// Holds no references so it's trivially `Send + 'static` for the wasm spawn path.
98struct 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    /// Create a tuner. Its cache is seeded from the persistent on-disk cache when
110    /// `std_io` is enabled.
111    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    /// Fetch the fastest autotune operation index for an autotune key.
119    pub fn fastest(&self, key: &K) -> TuneCacheResult {
120        self.cache.read().fastest(key)
121    }
122
123    /// Check the cache, validate checksums if needed, and kick off a tuning job if the
124    /// key is a miss. Returns the resolved cache state.
125    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        // Failed input kernels must not be consumed as failed tuning candidates.
137        // Submit outside the cache lock to preserve the device/cache lock order.
138        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            // Scope the guard: the rest of this function re-locks `self.cache` (fast
162            // path insert, `process_request`), and the write lock is non-reentrant.
163        }
164
165        log::info!("Tuning {key}");
166
167        // Fast path: single tunable, no benchmarking needed.
168        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        // Walk the plan batch by batch, launching each benchmark synchronously. A
194        // successful launch queues a `PendingBench` for the async resolver below;
195        // launch errors go straight into `results`. Retry the next batch if a whole
196        // batch failed to queue anything.
197        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        // Resolve samples and commit the result. On wasm this runs on the browser
237        // event loop; elsewhere it blocks inline.
238        #[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
254/// Await every profile sample, pick the fastest tunable, commit to the cache.
255async 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
328/// Emit the autotune result through the logger at the currently configured level.
329fn 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}