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 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        // Walk the plan batch by batch, launching each benchmark synchronously. A
195        // successful launch queues a `PendingBench` for the async resolver below;
196        // launch errors go straight into `results`. Retry the next batch if a whole
197        // batch failed to queue anything.
198        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        // Resolve samples and commit the result. On wasm this runs on the browser
238        // event loop; elsewhere it blocks inline.
239        #[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
255/// Await every profile sample, pick the fastest tunable, commit to the cache.
256async 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
329/// Emit the autotune result through the logger at the currently configured level.
330fn 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}