Skip to main content

ruda_runtime/runtime/tune/
local.rs

1use super::{AutotuneKey, AutotuneOutput, TunableSet, TuneInputs, Tuner};
2use crate::runtime::{client::ComputeClient, backend::Runtime, tune::TuneCacheResult};
3use alloc::string::ToString;
4use alloc::sync::Arc;
5use core::{
6    any::{Any, TypeId},
7    fmt::Display,
8    hash::Hash,
9};
10use hashbrown::HashMap;
11use spin::{Mutex, RwLock};
12
13type Sets = RwLock<Option<HashMap<TypeId, Arc<dyn Any + Send + Sync>>>>;
14
15/// A local tuner allows to create a tuner for a specific key that can be different from the server
16/// key.
17pub struct LocalTuner<AK: AutotuneKey, ID> {
18    state: Mutex<Option<HashMap<ID, Arc<Tuner<AK>>>>>,
19    name: &'static str,
20    sets: Sets,
21    device_sets: Mutex<Option<HashMap<ID, Arc<Sets>>>>,
22}
23
24/// Create a local tuner with the provided name.
25#[macro_export]
26macro_rules! local_tuner {
27    ($name:expr) => {
28        LocalTuner::new(concat!(module_path!(), "-", $name));
29    };
30    () => {
31        LocalTuner::new(module_path!());
32    };
33}
34
35pub use local_tuner;
36
37impl<AK, ID> LocalTuner<AK, ID>
38where
39    AK: AutotuneKey + 'static,
40    ID: Hash + PartialEq + Eq + Clone + Display,
41{
42    /// Create a new local tuner.
43    pub const fn new(name: &'static str) -> Self {
44        Self {
45            state: Mutex::new(None),
46            name,
47            sets: RwLock::new(None),
48            device_sets: Mutex::new(None),
49        }
50    }
51
52    /// Get or initialize the [`TunableSet`] for this tuner.
53    ///
54    /// Returns a cached `Arc<TunableSet>` keyed by the `TypeId` of `init_set`. The
55    /// initializer runs at most once per process.
56    pub fn init<I, Out, F>(&self, init_set: F) -> Arc<TunableSet<AK, I, Out>>
57    where
58        F: Fn() -> TunableSet<AK, I, Out> + 'static + Send + Sync,
59        I: TuneInputs,
60        Out: AutotuneOutput,
61    {
62        Self::init_set(&self.sets, init_set)
63    }
64
65    /// Cache a candidate set per device and initializer, including device-dependent choices.
66    pub fn init_for_device<I, Out, F>(&self, id: &ID, init_set: F) -> Arc<TunableSet<AK, I, Out>>
67    where
68        F: Fn() -> TunableSet<AK, I, Out> + 'static + Send + Sync,
69        I: TuneInputs,
70        Out: AutotuneOutput,
71    {
72        let sets = {
73            let mut devices = self.device_sets.lock();
74            let devices = devices.get_or_insert_with(HashMap::new);
75            match devices.get(id) {
76                Some(sets) => sets.clone(),
77                None => devices.entry(id.clone())
78                    .or_insert_with(|| Arc::new(RwLock::new(None)))
79                    .clone(),
80            }
81        };
82        Self::init_set(&sets, init_set)
83    }
84
85    fn init_set<I, Out, F>(
86        sets: &Sets,
87        init_set: F,
88    ) -> Arc<TunableSet<AK, I, Out>>
89    where
90        F: Fn() -> TunableSet<AK, I, Out> + 'static + Send + Sync,
91        I: TuneInputs,
92        Out: AutotuneOutput,
93    {
94        let key = TypeId::of::<F>();
95        let read = sets.read();
96
97        static DOWNCAST_ERROR: &str = "Local tuner only support one set of tunable that must work on the same input and output declared with the init function.";
98
99        if let Some(sets) = read.as_ref()
100            && let Some(set) = sets.get(&key)
101        {
102            return set.clone().downcast().expect(DOWNCAST_ERROR);
103        };
104
105        core::mem::drop(read);
106
107        let mut sets = sets.write();
108
109        if let Some(sets) = sets.as_ref()
110            && let Some(set) = sets.get(&key)
111        {
112            return set.clone().downcast().expect(DOWNCAST_ERROR);
113        };
114
115        let content = Arc::new(init_set());
116
117        if let Some(sets) = sets.as_mut() {
118            sets.insert(key, content.clone());
119        } else {
120            let mut map = HashMap::<TypeId, Arc<dyn Any + Send + Sync>>::new();
121            map.insert(key, content.clone());
122            *sets = Some(map);
123        };
124
125        content
126    }
127
128    /// Clear the autotune state.
129    pub fn clear(&self) {
130        if let Some(s) = self.state.lock().as_mut() {
131            s.clear()
132        }
133    }
134
135    #[cfg(feature = "runtime-autotune-checks")]
136    fn checks<'a, I: TuneInputs, Out: AutotuneOutput>(
137        &self,
138        operations: &TunableSet<AK, I, Out>,
139        inputs: &<I as TuneInputs>::At<'a>,
140    ) where
141        <I as TuneInputs>::At<'a>: Clone + Send,
142    {
143        use alloc::vec::Vec;
144
145        let mut checks_outputs = Vec::new();
146        for i in 0..operations.len() {
147            let op = operations.fastest(i);
148            let result = op.execute(inputs.clone());
149            checks_outputs.push(result);
150        }
151        super::check_autotune_outputs(checks_outputs);
152    }
153
154    /// Execute the fastest operation in a [`TunableSet`], triggering a tuning pass on
155    /// the first call for a given key.
156    pub fn execute<'a, R: Runtime, I: TuneInputs, Out>(
157        &self,
158        id: &ID,
159        client: &ComputeClient<R>,
160        operations: Arc<TunableSet<AK, I, Out>>,
161        inputs: <I as TuneInputs>::At<'a>,
162    ) -> Out
163    where
164        <I as TuneInputs>::At<'a>: Clone + Send,
165        Out: AutotuneOutput,
166    {
167        #[cfg(std_io)]
168        if super::stack::stack_autotuner().is_some() && operations.stack_reference().is_some() {
169            return super::stack::try_execute_stack(
170                self.name, &id.to_string(), client, operations, inputs,
171            ).expect("Full-stack autotune failed; request was not replayed");
172        }
173        let key = operations.generate_key(&inputs);
174
175        let tuner = {
176            let mut state = self.state.lock();
177            let state = state.get_or_insert_with(HashMap::new);
178            match state.get(id) {
179                Some(tuner) => tuner.clone(),
180                None => state.entry(id.clone())
181                    .or_insert_with(|| {
182                        let name = self.name.replace("::", "-");
183                        Arc::new(Tuner::new(&name, &id.to_string()))
184                    })
185                    .clone(),
186            }
187        };
188
189        // First, check for a cache hit under a read lock.
190        if let TuneCacheResult::Hit { fastest_index } = tuner.fastest(&key) {
191            #[cfg(feature = "runtime-autotune-checks")]
192            self.checks::<I, Out>(&operations, &inputs);
193            return operations
194                .fastest(fastest_index)
195                .execute(inputs)
196                .expect("Should run when selected by autotune.");
197        }
198
199        let fastest = tuner.check_tune::<R, I, Out>(
200            &key,
201            &inputs,
202            &operations,
203            || operations.compute_checksum(),
204            client,
205        );
206
207        // Run the execution depending on the cache state.
208        match fastest {
209            TuneCacheResult::Hit { fastest_index } => {
210                #[cfg(feature = "runtime-autotune-checks")]
211                self.checks::<I, Out>(&operations, &inputs);
212
213                operations
214                    .fastest(fastest_index)
215                    .execute(inputs)
216                    .expect("Should run when selected by autotune.")
217            }
218            TuneCacheResult::Unchecked | TuneCacheResult::Miss => {
219                panic!(
220                    "Somehow we STILL didn't check a tuning checksum or start tuning, something has gone wrong."
221                )
222            }
223            TuneCacheResult::Pending => {
224                // Still waiting (e.g. on wasm). Try all operations as a fallback.
225                for i in 0..operations.len() {
226                    if let Ok(output) = operations.fastest(i).execute(inputs.clone()) {
227                        return output;
228                    }
229                }
230                panic!("All autotune operations failed, no viable operation found.");
231            }
232        }
233    }
234}