Skip to main content

ruda_runtime/runtime/tune/
operation.rs

1use alloc::boxed::Box;
2use alloc::string::String;
3use alloc::sync::Arc;
4use alloc::vec::Vec;
5use core::fmt::{Debug, Display, Write};
6use core::hash::Hash;
7use spin::Once;
8
9use alloc::format;
10
11use super::{
12    AutotuneError, input_generator::InputGenerator, key_generator::KeyGenerator,
13    tune_inputs::TuneInputs,
14};
15use super::{Tunable, TunePlan};
16
17/// A type-erased delegate for a tunable function.
18///
19/// The lifetime `'inp` is the lifetime of the input data, the function must be defined such that
20/// it can be called for any lifetime `inp` and produce a `Result<Out, AutotuneError>`.
21type TuneDelegate<I, Out> =
22    dyn for<'inp> Fn(<I as TuneInputs>::At<'inp>) -> Result<Out, AutotuneError> + Send + Sync;
23
24/// A named, type-erased tunable function stored in a [`TunableSet`]. Constructed via
25/// [`Tunable::new`](super::Tunable::new); callers don't name this type directly.
26#[derive(new)]
27pub struct TuneFn<I: TuneInputs, Out> {
28    pub(crate) name: String,
29    func: Box<TuneDelegate<I, Out>>,
30}
31
32impl<I: TuneInputs, Out: 'static> TuneFn<I, Out> {
33    /// Run the wrapped function on the given inputs.
34    pub fn execute<'a>(&self, inputs: <I as TuneInputs>::At<'a>) -> Result<Out, AutotuneError> {
35        (self.func)(inputs)
36    }
37}
38
39/// A set of candidate tunable functions for autotune, sharing a key generator and an
40/// input generator. See [`TuneInputs`] for the `F` parameter.
41pub struct TunableSet<K: AutotuneKey, F: TuneInputs, Output: 'static> {
42    tunables: Vec<Tunable<K, F, Output>>,
43    key_gen: Arc<dyn KeyGenerator<K, F> + Send + Sync>,
44    input_gen: Arc<dyn InputGenerator<K, F> + Send + Sync>,
45    stack_reference: Option<usize>,
46    stack_revision: String,
47    stack_workload: Option<Arc<dyn for<'a> Fn(&F::At<'a>) -> String + Send + Sync>>,
48    checksum: Once<String>,
49    stack_manifest: Once<String>,
50}
51
52impl<K: AutotuneKey, F: TuneInputs, Output: 'static> TunableSet<K, F, Output> {
53    /// The number of tunables in the set.
54    pub fn len(&self) -> usize {
55        self.tunables.len()
56    }
57
58    /// Whether this set contains no tunables.
59    pub fn is_empty(&self) -> bool {
60        self.tunables.is_empty()
61    }
62
63    /// Create a tunable set from a key generator and an input generator.
64    pub fn new(key_gen: impl KeyGenerator<K, F>, input_gen: impl InputGenerator<K, F>) -> Self {
65        Self {
66            tunables: Default::default(),
67            input_gen: Arc::new(input_gen),
68            key_gen: Arc::new(key_gen),
69            stack_reference: None,
70            stack_revision: String::new(),
71            stack_workload: None,
72            checksum: Once::new(),
73            stack_manifest: Once::new(),
74        }
75    }
76
77    /// Register an exact workload signature and a known-safe reference for full-stack tuning.
78    /// The input generator MUST isolate all mutable state for each benchmark invocation.
79    /// Include dtype, exact shape/strides, precision, optional inputs and semantic parameters.
80    pub fn with_stack_tuning(
81        mut self, reference: usize, revision: &str,
82        workload: impl for<'a> Fn(&F::At<'a>) -> String + Send + Sync + 'static,
83    ) -> Self {
84        self.stack_reference = Some(reference);
85        self.stack_revision = revision.into();
86        self.stack_workload = Some(Arc::new(workload));
87        self.stack_manifest = Once::new();
88        self
89    }
90    pub fn stack_reference(&self) -> Option<usize> { self.stack_reference }
91    pub fn stack_workload<'a>(&self, inputs: &F::At<'a>) -> Option<String> {
92        self.stack_workload.as_ref().map(|f| f(inputs))
93    }
94
95    /// Shorthand for [`new`](Self::new) with a [`CloneInputGenerator`]: benchmarks run
96    /// on clones of the real call inputs.
97    pub fn new_cloning_inputs(key_gen: impl KeyGenerator<K, F>) -> Self {
98        Self::new(key_gen, super::CloneInputGenerator)
99    }
100
101    /// Register a tunable with this tunable set.
102    pub fn with(mut self, tunable: Tunable<K, F, Output>) -> Self {
103        self.tunables.push(tunable);
104        self.checksum = Once::new();
105        self.stack_manifest = Once::new();
106        self
107    }
108
109    /// All candidate operations in this set, in registration order.
110    pub fn autotunables(&self) -> impl Iterator<Item = &TuneFn<F, Output>> {
111        self.tunables.iter().map(|tunable| &tunable.function)
112    }
113
114    /// Returns the [autotune plan](TunePlan) for the given set.
115    pub(crate) fn plan(&self, key: &K) -> TunePlan {
116        TunePlan::new(key, &self.tunables)
117    }
118
119    /// Returns the operation for the given index, matching the order returned by
120    /// `autotunables`. Tunables are tried in order, so index 0 should be a good default.
121    pub fn fastest(&self, fastest_index: usize) -> &TuneFn<F, Output> {
122        &self.tunables[fastest_index].function
123    }
124
125    /// Compute a checksum that invalidates outdated cached auto-tune results when the
126    /// set of tunable names changes.
127    pub fn compute_checksum(&self) -> String {
128        // Preserve legacy cache identity when the new controller is not enabled.
129        self.checksum.call_once(|| {
130            let mut checksum = String::new();
131            for tune in &self.tunables { checksum += &tune.function.name; }
132            format!("{:x}", md5::compute(checksum))
133        }).clone()
134    }
135
136    /// Separate, length-delimited manifest for the opt-in full-stack cache.
137    pub fn stack_checksum(&self) -> String {
138        self.stack_manifest.call_once(|| {
139            let mut checksum = format!("stack-v1:{}:{};reference={:?};", self.stack_revision.len(), self.stack_revision, self.stack_reference);
140            for tune in &self.tunables {
141                let _ = write!(checksum, "{}:{}", tune.function.name.len(), tune.function.name);
142            }
143            format!("{:x}", md5::compute(checksum))
144        }).clone()
145    }
146
147    /// Generate a key from a set of inputs
148    pub fn generate_key<'a>(&self, inputs: &F::At<'a>) -> K {
149        self.key_gen.generate(inputs)
150    }
151
152    /// Generate a set of test inputs from a key and reference inputs.
153    pub fn generate_inputs<'a>(&self, key: &K, inputs: &F::At<'a>) -> F::At<'a> {
154        self.input_gen.generate(key, inputs)
155    }
156}
157
158#[cfg(std_io)]
159/// Trait alias with support for persistent caching
160pub trait AutotuneKey:
161    Clone
162    + Debug
163    + PartialEq
164    + Eq
165    + Hash
166    + Display
167    + serde::Serialize
168    + serde::de::DeserializeOwned
169    + Send
170    + Sync
171    + 'static
172{
173}
174#[cfg(not(std_io))]
175/// Trait alias
176pub trait AutotuneKey:
177    Clone + Debug + PartialEq + Eq + Hash + Display + Send + Sync + 'static
178{
179}
180
181impl AutotuneKey for String {}