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