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::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: RwLock<Option<HashMap<ID, Arc<Tuner<AK>>>>>,
19    name: &'static str,
20    sets: Sets,
21    device_sets: RwLock<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: RwLock::new(None),
46            name,
47            sets: RwLock::new(None),
48            device_sets: RwLock::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 existing = self.device_sets.read().as_ref().and_then(|devices| devices.get(id)).cloned();
73        let sets = existing.unwrap_or_else(|| {
74            let mut devices = self.device_sets.write();
75            let devices = devices.get_or_insert_with(HashMap::new);
76            match devices.get(id) {
77                Some(sets) => sets.clone(),
78                None => devices.entry(id.clone())
79                    .or_insert_with(|| Arc::new(RwLock::new(None)))
80                    .clone(),
81            }
82        });
83        Self::init_set(&sets, init_set)
84    }
85
86    fn init_set<I, Out, F>(
87        sets: &Sets,
88        init_set: F,
89    ) -> Arc<TunableSet<AK, I, Out>>
90    where
91        F: Fn() -> TunableSet<AK, I, Out> + 'static + Send + Sync,
92        I: TuneInputs,
93        Out: AutotuneOutput,
94    {
95        let key = TypeId::of::<F>();
96        let read = sets.read();
97
98        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.";
99
100        if let Some(sets) = read.as_ref()
101            && let Some(set) = sets.get(&key)
102        {
103            return set.clone().downcast().expect(DOWNCAST_ERROR);
104        };
105
106        core::mem::drop(read);
107
108        let mut sets = sets.write();
109
110        if let Some(sets) = sets.as_ref()
111            && let Some(set) = sets.get(&key)
112        {
113            return set.clone().downcast().expect(DOWNCAST_ERROR);
114        };
115
116        let content = Arc::new(init_set());
117
118        if let Some(sets) = sets.as_mut() {
119            sets.insert(key, content.clone());
120        } else {
121            let mut map = HashMap::<TypeId, Arc<dyn Any + Send + Sync>>::new();
122            map.insert(key, content.clone());
123            *sets = Some(map);
124        };
125
126        content
127    }
128
129    /// Clear the autotune state.
130    pub fn clear(&self) {
131        if let Some(s) = self.state.write().as_mut() {
132            s.clear()
133        }
134    }
135
136    #[cfg(feature = "runtime-autotune-checks")]
137    fn checks<'a, I: TuneInputs, Out: AutotuneOutput>(
138        &self,
139        operations: &TunableSet<AK, I, Out>,
140        inputs: &<I as TuneInputs>::At<'a>,
141    ) where
142        <I as TuneInputs>::At<'a>: Clone + Send,
143    {
144        use alloc::vec::Vec;
145
146        let mut checks_outputs = Vec::new();
147        for i in 0..operations.len() {
148            let op = operations.fastest(i);
149            let result = op.execute(inputs.clone());
150            checks_outputs.push(result);
151        }
152        super::check_autotune_outputs(checks_outputs);
153    }
154
155    /// Execute the fastest operation in a [`TunableSet`], triggering a tuning pass on
156    /// the first call for a given key.
157    pub fn execute<'a, R: Runtime, I: TuneInputs, Out>(
158        &self,
159        id: &ID,
160        client: &ComputeClient<R>,
161        operations: Arc<TunableSet<AK, I, Out>>,
162        inputs: <I as TuneInputs>::At<'a>,
163    ) -> Out
164    where
165        <I as TuneInputs>::At<'a>: Clone + Send,
166        Out: AutotuneOutput,
167    {
168        #[cfg(std_io)]
169        if super::stack::stack_autotuner().is_some() && operations.stack_reference().is_some() {
170            return super::stack::try_execute_stack(
171                self.name, &id.to_string(), client, operations, inputs,
172            ).expect("Full-stack autotune failed; request was not replayed");
173        }
174        let key = operations.generate_key(&inputs);
175
176        let existing = self.state.read().as_ref().and_then(|state| state.get(id)).cloned();
177        let tuner = existing.unwrap_or_else(|| {
178            let mut state = self.state.write();
179            let state = state.get_or_insert_with(HashMap::new);
180            match state.get(id) {
181                Some(tuner) => tuner.clone(),
182                None => state.entry(id.clone())
183                    .or_insert_with(|| {
184                        let name = self.name.replace("::", "-");
185                        Arc::new(Tuner::new(&name, &id.to_string()))
186                    })
187                    .clone(),
188            }
189        });
190
191        // First, check for a cache hit under a read lock.
192        if let TuneCacheResult::Hit { fastest_index } = tuner.fastest(&key) {
193            #[cfg(feature = "runtime-autotune-checks")]
194            self.checks::<I, Out>(&operations, &inputs);
195            return operations
196                .fastest(fastest_index)
197                .execute(inputs)
198                .expect("Should run when selected by autotune.");
199        }
200
201        let fastest = tuner.check_tune::<R, I, Out>(
202            &key,
203            &inputs,
204            &operations,
205            || operations.compute_checksum(),
206            client,
207        );
208
209        // Run the execution depending on the cache state.
210        match fastest {
211            TuneCacheResult::Hit { fastest_index } => {
212                #[cfg(feature = "runtime-autotune-checks")]
213                self.checks::<I, Out>(&operations, &inputs);
214
215                operations
216                    .fastest(fastest_index)
217                    .execute(inputs)
218                    .expect("Should run when selected by autotune.")
219            }
220            TuneCacheResult::Unchecked | TuneCacheResult::Miss => {
221                panic!(
222                    "Somehow we STILL didn't check a tuning checksum or start tuning, something has gone wrong."
223                )
224            }
225            TuneCacheResult::Pending => {
226                // Still waiting (e.g. on wasm). Try all operations as a fallback.
227                for i in 0..operations.len() {
228                    if let Ok(output) = operations.fastest(i).execute(inputs.clone()) {
229                        return output;
230                    }
231                }
232                panic!("All autotune operations failed, no viable operation found.");
233            }
234        }
235    }
236}