ruda_runtime/runtime/tune/
operation.rs1use 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
17type TuneDelegate<I, Out> =
22 dyn for<'inp> Fn(<I as TuneInputs>::At<'inp>) -> Result<Out, AutotuneError> + Send + Sync;
23
24#[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 pub fn execute<'a>(&self, inputs: <I as TuneInputs>::At<'a>) -> Result<Out, AutotuneError> {
35 (self.func)(inputs)
36 }
37}
38
39pub 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 pub fn len(&self) -> usize {
55 self.tunables.len()
56 }
57
58 pub fn is_empty(&self) -> bool {
60 self.tunables.is_empty()
61 }
62
63 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 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 pub fn new_cloning_inputs(key_gen: impl KeyGenerator<K, F>) -> Self {
98 Self::new(key_gen, super::CloneInputGenerator)
99 }
100
101 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 pub fn autotunables(&self) -> impl Iterator<Item = &TuneFn<F, Output>> {
111 self.tunables.iter().map(|tunable| &tunable.function)
112 }
113
114 pub(crate) fn plan(&self, key: &K) -> TunePlan {
116 TunePlan::new(key, &self.tunables)
117 }
118
119 pub fn fastest(&self, fastest_index: usize) -> &TuneFn<F, Output> {
122 &self.tunables[fastest_index].function
123 }
124
125 pub fn compute_checksum(&self) -> String {
128 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 pub fn stack_checksum(&self) -> String {
138 String::from(self.stack_checksum_ref())
139 }
140
141 pub(crate) fn stack_checksum_ref(&self) -> &str {
142 self.stack_manifest.call_once(|| {
143 let mut checksum = format!("stack-v1:{}:{};reference={:?};", self.stack_revision.len(), self.stack_revision, self.stack_reference);
144 for tune in &self.tunables {
145 let _ = write!(checksum, "{}:{}", tune.function.name.len(), tune.function.name);
146 }
147 format!("{:x}", md5::compute(checksum))
148 }).as_str()
149 }
150
151 pub fn generate_key<'a>(&self, inputs: &F::At<'a>) -> K {
153 self.key_gen.generate(inputs)
154 }
155
156 pub fn generate_inputs<'a>(&self, key: &K, inputs: &F::At<'a>) -> F::At<'a> {
158 self.input_gen.generate(key, inputs)
159 }
160}
161
162#[cfg(std_io)]
163pub trait AutotuneKey:
165 Clone
166 + Debug
167 + PartialEq
168 + Eq
169 + Hash
170 + Display
171 + serde::Serialize
172 + serde::de::DeserializeOwned
173 + Send
174 + Sync
175 + 'static
176{
177}
178#[cfg(not(std_io))]
179pub trait AutotuneKey:
181 Clone + Debug + PartialEq + Eq + Hash + Display + Send + Sync + 'static
182{
183}
184
185impl AutotuneKey for String {}